import "server-only"; import { and, desc, eq, inArray, isNull, sql } from "drizzle-orm"; import { getDb, arenaSessions, arenaResponses, arenaVotes, sharedArenaSessions, usageRecords, messageAttachments } from "@/db"; import { ids, newId } from "@/lib/ids"; import { ApiError } from "@/lib/api"; import { getModel, touchRecent } from "@/lib/ai/registry"; import { getAdapter } from "@/lib/ai/providers"; import { getDecryptedKey, recordProviderOutcome } from "@/lib/providers/keys"; import { isCustomModelKey, resolveCustomEndpoint } from "@/lib/endpoints/service"; import { estimateCost } from "@/lib/ai/core/pricing"; import { filterSettings } from "@/lib/ai/core/normalize"; import { StreamAccumulator } from "@/lib/ai/core/stream-utils"; import type { ContentPart, PolyProviderErrorShape, UnifiedGenerationSettings, UnifiedStreamEvent, Usage } from "@/lib/ai/core/types"; import type { GenerationSettingsInput } from "@/lib/chat/schemas"; import { analyzePrompt } from "@/lib/client/router"; import { categoryFromTask, computeScoreboard, isTaskCategory, ratingKeyFor, shuffledOrder, type ScoreboardFilter, type TaskCategory } from "./scoring"; import { arenaExportFilename, buildArenaJson, buildArenaMarkdown, buildArenaShareSnapshot, type ArenaExportInput, type ExportModelInfo } from "./export"; /** Keys stored in `arena_sessions.settings` that are not generation parameters. */ const ARENA_SETTING_KEYS = ["attachmentIds", "blind", "blindOrder", "tools"] as const; export async function createArenaSession(userId: string, input: { prompt: string; systemPrompt?: string; modelKeys: string[]; settings?: GenerationSettingsInput; attachmentIds?: string[]; blind?: boolean }) { const db = getDb(); // validate models + keys up front so the UI can show a clear error before streaming const missing: string[] = []; for (const key of input.modelKeys) { const custom = isCustomModelKey(key) ? await resolveCustomEndpoint(userId, key) : null; const m = custom?.model ?? (await getModel(key)); if (!m) throw new ApiError(404, `Unknown model ${key}`, "MODEL_NOT_FOUND"); if (!custom && !(await getDecryptedKey(userId, m.provider))) missing.push(m.provider); } if (missing.length) throw new ApiError(400, `No API key for ${[...new Set(missing)].join(", ")}.`, "NO_PROVIDER_KEY", { providers: [...new Set(missing)] }); const settings: Record = { ...(input.settings ?? {}), attachmentIds: input.attachmentIds ?? [] }; if (input.blind) { settings.blind = true; settings.blindOrder = shuffledOrder(input.modelKeys.length); } const [session] = await db.insert(arenaSessions).values({ id: ids.arena(), userId, prompt: input.prompt, systemPrompt: input.systemPrompt ?? null, modelKeys: input.modelKeys, settings }).returning(); return session; } function serializeResponse(r: T) { return { ...r, createdAt: r.createdAt.toISOString() }; } export async function listArenaSessions(userId: string, limit = 30) { const db = getDb(); const sessions = await db.select().from(arenaSessions).where(eq(arenaSessions.userId, userId)).orderBy(desc(arenaSessions.createdAt)).limit(limit); if (!sessions.length) return []; const sessionIds = sessions.map((s) => s.id); const [responses, votes] = await Promise.all([db.select().from(arenaResponses).where(inArray(arenaResponses.sessionId, sessionIds)), db.select().from(arenaVotes).where(inArray(arenaVotes.sessionId, sessionIds))]); return sessions.map((s) => ({ ...s, createdAt: s.createdAt.toISOString(), responses: responses.filter((r) => r.sessionId === s.id).map(serializeResponse), votes: votes.filter((v) => v.sessionId === s.id).map(serializeResponse), })); } /** One session with its responses and votes (owner only). */ export async function getArenaSession(userId: string, id: string) { const db = getDb(); const [s] = await db .select() .from(arenaSessions) .where(and(eq(arenaSessions.id, id), eq(arenaSessions.userId, userId))) .limit(1); if (!s) throw new ApiError(404, "Arena session not found", "NOT_FOUND"); const [responses, votes] = await Promise.all([db.select().from(arenaResponses).where(eq(arenaResponses.sessionId, id)), db.select().from(arenaVotes).where(eq(arenaVotes.sessionId, id))]); return { ...s, createdAt: s.createdAt.toISOString(), responses: responses.map(serializeResponse), votes: votes.map(serializeResponse) }; } export async function deleteArenaSession(userId: string, id: string) { const db = getDb(); const res = await db .delete(arenaSessions) .where(and(eq(arenaSessions.id, id), eq(arenaSessions.userId, userId))) .returning({ id: arenaSessions.id }); if (!res.length) throw new ApiError(404, "Arena session not found", "NOT_FOUND"); // responses cascade; votes and share links have no FK (kept in the workspace schema) await db.delete(arenaVotes).where(and(eq(arenaVotes.sessionId, id), eq(arenaVotes.userId, userId))); await db.delete(sharedArenaSessions).where(and(eq(sharedArenaSessions.sessionId, id), eq(sharedArenaSessions.userId, userId))); } export async function rateArenaResponse(userId: string, responseId: string, ratings: Record) { const db = getDb(); const [r] = await db.select({ id: arenaResponses.id, sessionId: arenaResponses.sessionId }).from(arenaResponses).innerJoin(arenaSessions, eq(arenaResponses.sessionId, arenaSessions.id)).where(and(eq(arenaResponses.id, responseId), eq(arenaSessions.userId, userId))).limit(1); if (!r) throw new ApiError(404, "Response not found", "NOT_FOUND"); // "best" style ratings are exclusive within a session for (const [k, v] of Object.entries(ratings)) { if (v && k.startsWith("best")) { const others = await db.select({ id: arenaResponses.id, ratings: arenaResponses.ratings }).from(arenaResponses).where(eq(arenaResponses.sessionId, r.sessionId)); for (const o of others) if (o.id !== responseId && o.ratings[k]) await db.update(arenaResponses).set({ ratings: { ...o.ratings, [k]: false } }).where(eq(arenaResponses.id, o.id)); } } const [cur] = await db.select({ ratings: arenaResponses.ratings }).from(arenaResponses).where(eq(arenaResponses.id, responseId)); const [updated] = await db .update(arenaResponses) .set({ ratings: { ...cur.ratings, ...ratings } }) .where(eq(arenaResponses.id, responseId)) .returning(); return updated; } // --------------------------------------------------------------------------- // Votes (one row per session + criterion; legacy `ratings` kept in sync) // --------------------------------------------------------------------------- export async function castArenaVote(userId: string, input: { sessionId: string; responseId: string; criterion: string; category?: TaskCategory }) { const db = getDb(); const [r] = await db .select({ id: arenaResponses.id, sessionId: arenaResponses.sessionId, modelKey: arenaResponses.modelKey, status: arenaResponses.status }) .from(arenaResponses) .innerJoin(arenaSessions, eq(arenaResponses.sessionId, arenaSessions.id)) .where(and(eq(arenaResponses.id, input.responseId), eq(arenaResponses.sessionId, input.sessionId), eq(arenaSessions.userId, userId))) .limit(1); if (!r) throw new ApiError(404, "Response not found", "NOT_FOUND"); if (r.status === "streaming") throw new ApiError(409, "Wait for the response to finish before voting", "NOT_FINAL"); await db .insert(arenaVotes) .values({ id: newId("arv"), userId, sessionId: input.sessionId, responseId: input.responseId, modelKey: r.modelKey, criterion: input.criterion, category: input.category ?? null }) .onConflictDoUpdate({ target: [arenaVotes.sessionId, arenaVotes.criterion], set: { responseId: input.responseId, modelKey: r.modelKey, category: input.category ?? null, createdAt: new Date() } }); // Backward compatibility: mirror into arena_responses.ratings (exclusive per session). await rateArenaResponse(userId, input.responseId, { [ratingKeyFor(input.criterion)]: true }).catch(() => {}); return getSessionVotesAndRatings(input.sessionId); } export async function retractArenaVote(userId: string, input: { sessionId: string; criterion: string }) { const db = getDb(); const removed = await db .delete(arenaVotes) .where(and(eq(arenaVotes.sessionId, input.sessionId), eq(arenaVotes.criterion, input.criterion), eq(arenaVotes.userId, userId))) .returning({ responseId: arenaVotes.responseId }); for (const v of removed) await rateArenaResponse(userId, v.responseId, { [ratingKeyFor(input.criterion)]: false }).catch(() => {}); return getSessionVotesAndRatings(input.sessionId); } async function getSessionVotesAndRatings(sessionId: string) { const db = getDb(); const [votes, ratings] = await Promise.all([db.select().from(arenaVotes).where(eq(arenaVotes.sessionId, sessionId)), db.select({ id: arenaResponses.id, ratings: arenaResponses.ratings }).from(arenaResponses).where(eq(arenaResponses.sessionId, sessionId))]); return { votes: votes.map(serializeResponse), ratings: Object.fromEntries(ratings.map((r) => [r.id, r.ratings])) as Record> }; } // --------------------------------------------------------------------------- // Scoreboard // --------------------------------------------------------------------------- export async function getArenaScoreboard(userId: string, filter: ScoreboardFilter, limit = 1000) { const db = getDb(); const sessions = await db.select({ id: arenaSessions.id, modelKeys: arenaSessions.modelKeys, prompt: arenaSessions.prompt }).from(arenaSessions).where(eq(arenaSessions.userId, userId)).orderBy(desc(arenaSessions.createdAt)).limit(limit); if (!sessions.length) return { rows: [], sessions: 0, votes: 0 }; const sessionIds = sessions.map((s) => s.id); const [responses, votes] = await Promise.all([ db.select({ id: arenaResponses.id, sessionId: arenaResponses.sessionId, modelKey: arenaResponses.modelKey, status: arenaResponses.status, ttftMs: arenaResponses.ttftMs, latencyMs: arenaResponses.latencyMs, costUsd: arenaResponses.costUsd, usage: arenaResponses.usage }).from(arenaResponses).where(inArray(arenaResponses.sessionId, sessionIds)), db.select({ sessionId: arenaVotes.sessionId, responseId: arenaVotes.responseId, modelKey: arenaVotes.modelKey, criterion: arenaVotes.criterion, category: arenaVotes.category }).from(arenaVotes).where(and(eq(arenaVotes.userId, userId), inArray(arenaVotes.sessionId, sessionIds))), ]); // Session category: the voters' classification when present, else the router's analysis of the prompt. const scored = sessions.map((s) => { const voted = votes.find((v) => v.sessionId === s.id && isTaskCategory(v.category)); const category = voted && isTaskCategory(voted.category) ? voted.category : categoryFromTask(analyzePrompt(s.prompt).task); return { id: s.id, modelKeys: s.modelKeys, category }; }); const rows = computeScoreboard({ sessions: scored, responses: responses.map((r) => ({ ...r, usage: r.usage as { inputTokens?: number; outputTokens?: number } | null })), votes }, filter); return { rows, sessions: scored.filter((s) => !filter.category || s.category === filter.category).length, votes: votes.filter((v) => !filter.criterion || v.criterion === filter.criterion).length }; } // --------------------------------------------------------------------------- // Export & share // --------------------------------------------------------------------------- async function exportInput(userId: string, id: string): Promise { const s = await getArenaSession(userId, id); const models: Record = {}; for (const key of s.modelKeys) { const m = await getModel(key).catch(() => null); models[key] = { displayName: m?.displayName ?? key.split("/").slice(1).join("/"), provider: m?.provider ?? key.split("/")[0] }; } return { session: { id: s.id, prompt: s.prompt, systemPrompt: s.systemPrompt, modelKeys: s.modelKeys, settings: s.settings ?? {}, createdAt: s.createdAt }, responses: s.responses.map((r) => ({ id: r.id, modelKey: r.modelKey, provider: r.provider, status: r.status, content: r.content, reasoning: r.reasoning, error: r.error ?? null, ttftMs: r.ttftMs, latencyMs: r.latencyMs, costUsd: r.costUsd, usage: r.usage as { inputTokens?: number; outputTokens?: number } | null, ratings: r.ratings ?? {}, createdAt: r.createdAt })), votes: s.votes.map((v) => ({ criterion: v.criterion, responseId: v.responseId, modelKey: v.modelKey, category: v.category, createdAt: v.createdAt })), models, }; } export async function exportArenaSession(userId: string, id: string, format: "json" | "markdown") { const input = await exportInput(userId, id); if (format === "json") return { filename: arenaExportFilename(input.session.prompt, "json"), contentType: "application/json", body: buildArenaJson(input) }; return { filename: arenaExportFilename(input.session.prompt, "md"), contentType: "text/markdown; charset=utf-8", body: buildArenaMarkdown(input) }; } export async function shareArenaSession(userId: string, id: string) { const db = getDb(); const input = await exportInput(userId, id); const snapshot = buildArenaShareSnapshot(input) as unknown as Record; const [existing] = await db .select({ id: sharedArenaSessions.id }) .from(sharedArenaSessions) .where(and(eq(sharedArenaSessions.sessionId, id), eq(sharedArenaSessions.userId, userId), isNull(sharedArenaSessions.revokedAt))) .limit(1); if (existing) { await db.update(sharedArenaSessions).set({ snapshot }).where(eq(sharedArenaSessions.id, existing.id)); return getArenaShare(userId, id); } const shareId = ids.share(); await db.insert(sharedArenaSessions).values({ id: shareId, sessionId: id, userId, snapshot }); return getArenaShare(userId, id); } export async function revokeArenaShare(userId: string, sessionId: string) { await getDb() .update(sharedArenaSessions) .set({ revokedAt: new Date() }) .where(and(eq(sharedArenaSessions.sessionId, sessionId), eq(sharedArenaSessions.userId, userId), isNull(sharedArenaSessions.revokedAt))); } export async function getArenaShare(userId: string, sessionId: string) { const [row] = await getDb() .select({ id: sharedArenaSessions.id, createdAt: sharedArenaSessions.createdAt, viewCount: sharedArenaSessions.viewCount }) .from(sharedArenaSessions) .where(and(eq(sharedArenaSessions.sessionId, sessionId), eq(sharedArenaSessions.userId, userId), isNull(sharedArenaSessions.revokedAt))) .limit(1); return row ? { ...row, createdAt: row.createdAt.toISOString() } : null; } /** Public read (no auth): increments the view counter. */ export async function getPublicArenaShare(shareId: string) { const db = getDb(); const [row] = await db .select() .from(sharedArenaSessions) .where(and(eq(sharedArenaSessions.id, shareId), isNull(sharedArenaSessions.revokedAt), eq(sharedArenaSessions.isPublic, true))) .limit(1); if (!row) return null; await db.update(sharedArenaSessions).set({ viewCount: sql`${sharedArenaSessions.viewCount} + 1` }).where(eq(sharedArenaSessions.id, shareId)); return row; } // --------------------------------------------------------------------------- // Streaming // --------------------------------------------------------------------------- /** Streams one model's answer for an arena session and persists it. */ export async function runArenaModel(ctx: { userId: string; requestId: string }, sessionId: string, modelKey: string, emit: (ev: unknown) => void, signal: AbortSignal) { const db = getDb(); const [session] = await db .select() .from(arenaSessions) .where(and(eq(arenaSessions.id, sessionId), eq(arenaSessions.userId, ctx.userId))) .limit(1); if (!session) throw new ApiError(404, "Arena session not found", "NOT_FOUND"); if (!session.modelKeys.includes(modelKey)) throw new ApiError(400, "Model is not part of this session", "BAD_REQUEST"); const custom = isCustomModelKey(modelKey) ? await resolveCustomEndpoint(ctx.userId, modelKey) : null; const model = custom?.model ?? (await getModel(modelKey)); if (!model) throw new ApiError(404, "Unknown model", "MODEL_NOT_FOUND"); const apiKey = custom?.apiKey ?? (await getDecryptedKey(ctx.userId, model.provider)); if (!apiKey) throw new ApiError(400, `No API key for ${model.provider}`, "NO_PROVIDER_KEY"); const stored = (session.settings ?? {}) as Record; const attachmentIds = Array.isArray(stored.attachmentIds) ? (stored.attachmentIds as string[]) : []; const rawSettings = Object.fromEntries(Object.entries(stored).filter(([k]) => !(ARENA_SETTING_KEYS as readonly string[]).includes(k))); const { settings } = filterSettings(rawSettings as UnifiedGenerationSettings, model); const parts: ContentPart[] = [{ type: "text", text: session.prompt }]; if (attachmentIds.length) { const atts = await db .select() .from(messageAttachments) .where(and(eq(messageAttachments.userId, ctx.userId), inArray(messageAttachments.id, attachmentIds))); for (const a of atts) { if (a.kind === "image" && model.capabilities.vision) parts.push({ type: "image", mimeType: a.mimeType, data: a.dataBase64, name: a.name }); else if (a.kind !== "image") parts.push({ type: "file", mimeType: a.mimeType, data: a.dataBase64, name: a.name }); } } const [resp] = await db.insert(arenaResponses).values({ id: ids.arenaResponse(), sessionId, modelKey, provider: model.provider, status: "streaming" }).returning(); emit({ type: "meta", responseId: resp.id, modelKey }); const t0 = Date.now(); const acc = new StreamAccumulator(t0); let error: PolyProviderErrorShape | undefined; let usage: Usage | undefined; let exactCost: number | null = null; const adapter = custom?.adapter ?? getAdapter(model.provider); try { const stream = adapter.streamChat({ provider: model.provider, model: model.id, apiKey, system: session.systemPrompt ?? undefined, messages: [{ role: "user", content: parts }], settings, modelInfo: model, signal, requestId: ctx.requestId }); for await (const ev of stream as AsyncIterable) { if (signal.aborted) break; acc.push(ev); if (ev.type === "text-delta" || ev.type === "reasoning-delta" || ev.type === "citation" || ev.type === "server-tool") emit(ev); else if (ev.type === "usage") usage = ev.usage; else if (ev.type === "error") error = ev.error; else if (ev.type === "provider-data" && typeof ev.data.exactCostUsd === "number") exactCost = ev.data.exactCostUsd as number; } } catch (e) { error = adapter.normalizeError(e).toJSON(); } const latencyMs = Date.now() - t0; const cost = estimateCost(usage, model.pricing); const costUsd = exactCost ?? (cost.known ? cost.totalUsd : null); const status = error ? "error" : signal.aborted ? "stopped" : "complete"; const [saved] = await db .update(arenaResponses) .set({ content: acc.text, reasoning: acc.reasoning || null, status, error: error ? { code: error.code, message: error.message } : null, usage: usage as unknown as Record | null, latencyMs, ttftMs: acc.ttftMs ?? null, costUsd }) .where(eq(arenaResponses.id, resp.id)) .returning(); await db.insert(usageRecords).values({ id: ids.usage(), userId: ctx.userId, arenaSessionId: sessionId, provider: model.provider, modelKey, kind: "arena", status: error ? "error" : "ok", errorCode: error?.code ?? null, inputTokens: usage?.inputTokens ?? 0, outputTokens: usage?.outputTokens ?? 0, cachedTokens: usage?.cachedInputTokens ?? 0, reasoningTokens: usage?.reasoningTokens ?? 0, costUsd, latencyMs, ttftMs: acc.ttftMs ?? null, }); await touchRecent(ctx.userId, modelKey).catch(() => {}); await recordProviderOutcome(ctx.userId, model.provider, error ? { ok: false, code: error.code } : { ok: true }); if (error) emit({ type: "error", error }); emit({ type: "done", response: { ...saved, createdAt: saved.createdAt.toISOString() }, usage, costUsd, latencyMs, ttftMs: acc.ttftMs, status }); }