SPB Git forge

spb/polyllm

Public
15commits 1branches 0releases
2.2 MBsize
maindefault branch
14 days agolast push
TypeScript 97.4% SQL 1% JavaScript 0.9% CSS 0.6%
19.8 KB · 324 lines typescript
Raw Blame History
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