TypeScript 97.4%
SQL 1%
JavaScript 0.9%
CSS 0.6%
1import "server-only";2import { and, desc, eq, inArray, isNull, sql } from "drizzle-orm";3import { getDb, arenaSessions, arenaResponses, arenaVotes, sharedArenaSessions, usageRecords, messageAttachments } from "@/db";4import { ids, newId } from "@/lib/ids";5import { ApiError } from "@/lib/api";6import { getModel, touchRecent } from "@/lib/ai/registry";7import { getAdapter } from "@/lib/ai/providers";8import { getDecryptedKey, recordProviderOutcome } from "@/lib/providers/keys";9import { isCustomModelKey, resolveCustomEndpoint } from "@/lib/endpoints/service";10import { estimateCost } from "@/lib/ai/core/pricing";11import { filterSettings } from "@/lib/ai/core/normalize";12import { StreamAccumulator } from "@/lib/ai/core/stream-utils";13import type { ContentPart, PolyProviderErrorShape, UnifiedGenerationSettings, UnifiedStreamEvent, Usage } from "@/lib/ai/core/types";14import type { GenerationSettingsInput } from "@/lib/chat/schemas";15import { analyzePrompt } from "@/lib/client/router";16import { categoryFromTask, computeScoreboard, isTaskCategory, ratingKeyFor, shuffledOrder, type ScoreboardFilter, type TaskCategory } from "./scoring";17import { arenaExportFilename, buildArenaJson, buildArenaMarkdown, buildArenaShareSnapshot, type ArenaExportInput, type ExportModelInfo } from "./export";1819/** Keys stored in `arena_sessions.settings` that are not generation parameters. */20const ARENA_SETTING_KEYS = ["attachmentIds", "blind", "blindOrder", "tools"] as const;2122export async function createArenaSession(userId: string, input: { prompt: string; systemPrompt?: string; modelKeys: string[]; settings?: GenerationSettingsInput; attachmentIds?: string[]; blind?: boolean }) {23 const db = getDb();24 // validate models + keys up front so the UI can show a clear error before streaming25 const missing: string[] = [];26 for (const key of input.modelKeys) {27 const custom = isCustomModelKey(key) ? await resolveCustomEndpoint(userId, key) : null;28 const m = custom?.model ?? (await getModel(key));29 if (!m) throw new ApiError(404, `Unknown model ${key}`, "MODEL_NOT_FOUND");30 if (!custom && !(await getDecryptedKey(userId, m.provider))) missing.push(m.provider);31 }32 if (missing.length) throw new ApiError(400, `No API key for ${[...new Set(missing)].join(", ")}.`, "NO_PROVIDER_KEY", { providers: [...new Set(missing)] });33 const settings: Record<string, unknown> = { ...(input.settings ?? {}), attachmentIds: input.attachmentIds ?? [] };34 if (input.blind) {35 settings.blind = true;36 settings.blindOrder = shuffledOrder(input.modelKeys.length);37 }38 const [session] = await db.insert(arenaSessions).values({ id: ids.arena(), userId, prompt: input.prompt, systemPrompt: input.systemPrompt ?? null, modelKeys: input.modelKeys, settings }).returning();39 return session;40}4142function serializeResponse<T extends { createdAt: Date }>(r: T) {43 return { ...r, createdAt: r.createdAt.toISOString() };44}4546export async function listArenaSessions(userId: string, limit = 30) {47 const db = getDb();48 const sessions = await db.select().from(arenaSessions).where(eq(arenaSessions.userId, userId)).orderBy(desc(arenaSessions.createdAt)).limit(limit);49 if (!sessions.length) return [];50 const sessionIds = sessions.map((s) => s.id);51 const [responses, votes] = await Promise.all([db.select().from(arenaResponses).where(inArray(arenaResponses.sessionId, sessionIds)), db.select().from(arenaVotes).where(inArray(arenaVotes.sessionId, sessionIds))]);52 return sessions.map((s) => ({53 ...s,54 createdAt: s.createdAt.toISOString(),55 responses: responses.filter((r) => r.sessionId === s.id).map(serializeResponse),56 votes: votes.filter((v) => v.sessionId === s.id).map(serializeResponse),57 }));58}5960/** One session with its responses and votes (owner only). */61export async function getArenaSession(userId: string, id: string) {62 const db = getDb();63 const [s] = await db64 .select()65 .from(arenaSessions)66 .where(and(eq(arenaSessions.id, id), eq(arenaSessions.userId, userId)))67 .limit(1);68 if (!s) throw new ApiError(404, "Arena session not found", "NOT_FOUND");69 const [responses, votes] = await Promise.all([db.select().from(arenaResponses).where(eq(arenaResponses.sessionId, id)), db.select().from(arenaVotes).where(eq(arenaVotes.sessionId, id))]);70 return { ...s, createdAt: s.createdAt.toISOString(), responses: responses.map(serializeResponse), votes: votes.map(serializeResponse) };71}7273export async function deleteArenaSession(userId: string, id: string) {74 const db = getDb();75 const res = await db76 .delete(arenaSessions)77 .where(and(eq(arenaSessions.id, id), eq(arenaSessions.userId, userId)))78 .returning({ id: arenaSessions.id });79 if (!res.length) throw new ApiError(404, "Arena session not found", "NOT_FOUND");80 // responses cascade; votes and share links have no FK (kept in the workspace schema)81 await db.delete(arenaVotes).where(and(eq(arenaVotes.sessionId, id), eq(arenaVotes.userId, userId)));82 await db.delete(sharedArenaSessions).where(and(eq(sharedArenaSessions.sessionId, id), eq(sharedArenaSessions.userId, userId)));83}8485export async function rateArenaResponse(userId: string, responseId: string, ratings: Record<string, boolean>) {86 const db = getDb();87 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);88 if (!r) throw new ApiError(404, "Response not found", "NOT_FOUND");89 // "best" style ratings are exclusive within a session90 for (const [k, v] of Object.entries(ratings)) {91 if (v && k.startsWith("best")) {92 const others = await db.select({ id: arenaResponses.id, ratings: arenaResponses.ratings }).from(arenaResponses).where(eq(arenaResponses.sessionId, r.sessionId));93 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));94 }95 }96 const [cur] = await db.select({ ratings: arenaResponses.ratings }).from(arenaResponses).where(eq(arenaResponses.id, responseId));97 const [updated] = await db98 .update(arenaResponses)99 .set({ ratings: { ...cur.ratings, ...ratings } })100 .where(eq(arenaResponses.id, responseId))101 .returning();102 return updated;103}104105// ---------------------------------------------------------------------------106// Votes (one row per session + criterion; legacy `ratings` kept in sync)107// ---------------------------------------------------------------------------108export async function castArenaVote(userId: string, input: { sessionId: string; responseId: string; criterion: string; category?: TaskCategory }) {109 const db = getDb();110 const [r] = await db111 .select({ id: arenaResponses.id, sessionId: arenaResponses.sessionId, modelKey: arenaResponses.modelKey, status: arenaResponses.status })112 .from(arenaResponses)113 .innerJoin(arenaSessions, eq(arenaResponses.sessionId, arenaSessions.id))114 .where(and(eq(arenaResponses.id, input.responseId), eq(arenaResponses.sessionId, input.sessionId), eq(arenaSessions.userId, userId)))115 .limit(1);116 if (!r) throw new ApiError(404, "Response not found", "NOT_FOUND");117 if (r.status === "streaming") throw new ApiError(409, "Wait for the response to finish before voting", "NOT_FINAL");118 await db119 .insert(arenaVotes)120 .values({ id: newId("arv"), userId, sessionId: input.sessionId, responseId: input.responseId, modelKey: r.modelKey, criterion: input.criterion, category: input.category ?? null })121 .onConflictDoUpdate({ target: [arenaVotes.sessionId, arenaVotes.criterion], set: { responseId: input.responseId, modelKey: r.modelKey, category: input.category ?? null, createdAt: new Date() } });122 // Backward compatibility: mirror into arena_responses.ratings (exclusive per session).123 await rateArenaResponse(userId, input.responseId, { [ratingKeyFor(input.criterion)]: true }).catch(() => {});124 return getSessionVotesAndRatings(input.sessionId);125}126127export async function retractArenaVote(userId: string, input: { sessionId: string; criterion: string }) {128 const db = getDb();129 const removed = await db130 .delete(arenaVotes)131 .where(and(eq(arenaVotes.sessionId, input.sessionId), eq(arenaVotes.criterion, input.criterion), eq(arenaVotes.userId, userId)))132 .returning({ responseId: arenaVotes.responseId });133 for (const v of removed) await rateArenaResponse(userId, v.responseId, { [ratingKeyFor(input.criterion)]: false }).catch(() => {});134 return getSessionVotesAndRatings(input.sessionId);135}136137async function getSessionVotesAndRatings(sessionId: string) {138 const db = getDb();139 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))]);140 return { votes: votes.map(serializeResponse), ratings: Object.fromEntries(ratings.map((r) => [r.id, r.ratings])) as Record<string, Record<string, boolean>> };141}142143// ---------------------------------------------------------------------------144// Scoreboard145// ---------------------------------------------------------------------------146export async function getArenaScoreboard(userId: string, filter: ScoreboardFilter, limit = 1000) {147 const db = getDb();148 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);149 if (!sessions.length) return { rows: [], sessions: 0, votes: 0 };150 const sessionIds = sessions.map((s) => s.id);151 const [responses, votes] = await Promise.all([152 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)),153 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))),154 ]);155 // Session category: the voters' classification when present, else the router's analysis of the prompt.156 const scored = sessions.map((s) => {157 const voted = votes.find((v) => v.sessionId === s.id && isTaskCategory(v.category));158 const category = voted && isTaskCategory(voted.category) ? voted.category : categoryFromTask(analyzePrompt(s.prompt).task);159 return { id: s.id, modelKeys: s.modelKeys, category };160 });161 const rows = computeScoreboard({ sessions: scored, responses: responses.map((r) => ({ ...r, usage: r.usage as { inputTokens?: number; outputTokens?: number } | null })), votes }, filter);162 return { rows, sessions: scored.filter((s) => !filter.category || s.category === filter.category).length, votes: votes.filter((v) => !filter.criterion || v.criterion === filter.criterion).length };163}164165// ---------------------------------------------------------------------------166// Export & share167// ---------------------------------------------------------------------------168async function exportInput(userId: string, id: string): Promise<ArenaExportInput> {169 const s = await getArenaSession(userId, id);170 const models: Record<string, ExportModelInfo> = {};171 for (const key of s.modelKeys) {172 const m = await getModel(key).catch(() => null);173 models[key] = { displayName: m?.displayName ?? key.split("/").slice(1).join("/"), provider: m?.provider ?? key.split("/")[0] };174 }175 return {176 session: { id: s.id, prompt: s.prompt, systemPrompt: s.systemPrompt, modelKeys: s.modelKeys, settings: s.settings ?? {}, createdAt: s.createdAt },177 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 })),178 votes: s.votes.map((v) => ({ criterion: v.criterion, responseId: v.responseId, modelKey: v.modelKey, category: v.category, createdAt: v.createdAt })),179 models,180 };181}182183export async function exportArenaSession(userId: string, id: string, format: "json" | "markdown") {184 const input = await exportInput(userId, id);185 if (format === "json") return { filename: arenaExportFilename(input.session.prompt, "json"), contentType: "application/json", body: buildArenaJson(input) };186 return { filename: arenaExportFilename(input.session.prompt, "md"), contentType: "text/markdown; charset=utf-8", body: buildArenaMarkdown(input) };187}188189export async function shareArenaSession(userId: string, id: string) {190 const db = getDb();191 const input = await exportInput(userId, id);192 const snapshot = buildArenaShareSnapshot(input) as unknown as Record<string, unknown>;193 const [existing] = await db194 .select({ id: sharedArenaSessions.id })195 .from(sharedArenaSessions)196 .where(and(eq(sharedArenaSessions.sessionId, id), eq(sharedArenaSessions.userId, userId), isNull(sharedArenaSessions.revokedAt)))197 .limit(1);198 if (existing) {199 await db.update(sharedArenaSessions).set({ snapshot }).where(eq(sharedArenaSessions.id, existing.id));200 return getArenaShare(userId, id);201 }202 const shareId = ids.share();203 await db.insert(sharedArenaSessions).values({ id: shareId, sessionId: id, userId, snapshot });204 return getArenaShare(userId, id);205}206207export async function revokeArenaShare(userId: string, sessionId: string) {208 await getDb()209 .update(sharedArenaSessions)210 .set({ revokedAt: new Date() })211 .where(and(eq(sharedArenaSessions.sessionId, sessionId), eq(sharedArenaSessions.userId, userId), isNull(sharedArenaSessions.revokedAt)));212}213214export async function getArenaShare(userId: string, sessionId: string) {215 const [row] = await getDb()216 .select({ id: sharedArenaSessions.id, createdAt: sharedArenaSessions.createdAt, viewCount: sharedArenaSessions.viewCount })217 .from(sharedArenaSessions)218 .where(and(eq(sharedArenaSessions.sessionId, sessionId), eq(sharedArenaSessions.userId, userId), isNull(sharedArenaSessions.revokedAt)))219 .limit(1);220 return row ? { ...row, createdAt: row.createdAt.toISOString() } : null;221}222223/** Public read (no auth): increments the view counter. */224export async function getPublicArenaShare(shareId: string) {225 const db = getDb();226 const [row] = await db227 .select()228 .from(sharedArenaSessions)229 .where(and(eq(sharedArenaSessions.id, shareId), isNull(sharedArenaSessions.revokedAt), eq(sharedArenaSessions.isPublic, true)))230 .limit(1);231 if (!row) return null;232 await db.update(sharedArenaSessions).set({ viewCount: sql`${sharedArenaSessions.viewCount} + 1` }).where(eq(sharedArenaSessions.id, shareId));233 return row;234}235236// ---------------------------------------------------------------------------237// Streaming238// ---------------------------------------------------------------------------239/** Streams one model's answer for an arena session and persists it. */240export async function runArenaModel(ctx: { userId: string; requestId: string }, sessionId: string, modelKey: string, emit: (ev: unknown) => void, signal: AbortSignal) {241 const db = getDb();242 const [session] = await db243 .select()244 .from(arenaSessions)245 .where(and(eq(arenaSessions.id, sessionId), eq(arenaSessions.userId, ctx.userId)))246 .limit(1);247 if (!session) throw new ApiError(404, "Arena session not found", "NOT_FOUND");248 if (!session.modelKeys.includes(modelKey)) throw new ApiError(400, "Model is not part of this session", "BAD_REQUEST");249 const custom = isCustomModelKey(modelKey) ? await resolveCustomEndpoint(ctx.userId, modelKey) : null;250 const model = custom?.model ?? (await getModel(modelKey));251 if (!model) throw new ApiError(404, "Unknown model", "MODEL_NOT_FOUND");252 const apiKey = custom?.apiKey ?? (await getDecryptedKey(ctx.userId, model.provider));253 if (!apiKey) throw new ApiError(400, `No API key for ${model.provider}`, "NO_PROVIDER_KEY");254255 const stored = (session.settings ?? {}) as Record<string, unknown>;256 const attachmentIds = Array.isArray(stored.attachmentIds) ? (stored.attachmentIds as string[]) : [];257 const rawSettings = Object.fromEntries(Object.entries(stored).filter(([k]) => !(ARENA_SETTING_KEYS as readonly string[]).includes(k)));258 const { settings } = filterSettings(rawSettings as UnifiedGenerationSettings, model);259 const parts: ContentPart[] = [{ type: "text", text: session.prompt }];260 if (attachmentIds.length) {261 const atts = await db262 .select()263 .from(messageAttachments)264 .where(and(eq(messageAttachments.userId, ctx.userId), inArray(messageAttachments.id, attachmentIds)));265 for (const a of atts) {266 if (a.kind === "image" && model.capabilities.vision) parts.push({ type: "image", mimeType: a.mimeType, data: a.dataBase64, name: a.name });267 else if (a.kind !== "image") parts.push({ type: "file", mimeType: a.mimeType, data: a.dataBase64, name: a.name });268 }269 }270271 const [resp] = await db.insert(arenaResponses).values({ id: ids.arenaResponse(), sessionId, modelKey, provider: model.provider, status: "streaming" }).returning();272 emit({ type: "meta", responseId: resp.id, modelKey });273274 const t0 = Date.now();275 const acc = new StreamAccumulator(t0);276 let error: PolyProviderErrorShape | undefined;277 let usage: Usage | undefined;278 let exactCost: number | null = null;279 const adapter = custom?.adapter ?? getAdapter(model.provider);280 try {281 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 });282 for await (const ev of stream as AsyncIterable<UnifiedStreamEvent>) {283 if (signal.aborted) break;284 acc.push(ev);285 if (ev.type === "text-delta" || ev.type === "reasoning-delta" || ev.type === "citation" || ev.type === "server-tool") emit(ev);286 else if (ev.type === "usage") usage = ev.usage;287 else if (ev.type === "error") error = ev.error;288 else if (ev.type === "provider-data" && typeof ev.data.exactCostUsd === "number") exactCost = ev.data.exactCostUsd as number;289 }290 } catch (e) {291 error = adapter.normalizeError(e).toJSON();292 }293 const latencyMs = Date.now() - t0;294 const cost = estimateCost(usage, model.pricing);295 const costUsd = exactCost ?? (cost.known ? cost.totalUsd : null);296 const status = error ? "error" : signal.aborted ? "stopped" : "complete";297 const [saved] = await db298 .update(arenaResponses)299 .set({ content: acc.text, reasoning: acc.reasoning || null, status, error: error ? { code: error.code, message: error.message } : null, usage: usage as unknown as Record<string, number> | null, latencyMs, ttftMs: acc.ttftMs ?? null, costUsd })300 .where(eq(arenaResponses.id, resp.id))301 .returning();302 await db.insert(usageRecords).values({303 id: ids.usage(),304 userId: ctx.userId,305 arenaSessionId: sessionId,306 provider: model.provider,307 modelKey,308 kind: "arena",309 status: error ? "error" : "ok",310 errorCode: error?.code ?? null,311 inputTokens: usage?.inputTokens ?? 0,312 outputTokens: usage?.outputTokens ?? 0,313 cachedTokens: usage?.cachedInputTokens ?? 0,314 reasoningTokens: usage?.reasoningTokens ?? 0,315 costUsd,316 latencyMs,317 ttftMs: acc.ttftMs ?? null,318 });319 await touchRecent(ctx.userId, modelKey).catch(() => {});320 await recordProviderOutcome(ctx.userId, model.provider, error ? { ok: false, code: error.code } : { ok: true });321 if (error) emit({ type: "error", error });322 emit({ type: "done", response: { ...saved, createdAt: saved.createdAt.toISOString() }, usage, costUsd, latencyMs, ttftMs: acc.ttftMs, status });323}324