diff --git a/apps/studio/evals/assistant.eval.ts b/apps/studio/evals/assistant.eval.ts index c4500caec4f..740ce88eb5e 100644 --- a/apps/studio/evals/assistant.eval.ts +++ b/apps/studio/evals/assistant.eval.ts @@ -14,6 +14,7 @@ import { urlValidityScorer, } from './scorer' import { sqlIdentifierQuotingScorer, sqlSyntaxScorer } from './scorer-wasm' +import { buildTranscript } from './transcript' import { generateAssistantResponse } from '@/lib/ai/generate-assistant-response' import { getModel } from '@/lib/ai/model' import { DEFAULT_ASSISTANT_BASE_MODEL_ID, getAssistantModelEntry } from '@/lib/ai/model.utils' @@ -52,7 +53,8 @@ Eval('Assistant', { }) const finishReason = await result.finishReason - return { finishReason } + const steps = await result.steps + return { finishReason, transcript: buildTranscript(input.prompt, steps) } } finally { toolsAbortController.abort() } diff --git a/apps/studio/evals/scorer-online.ts b/apps/studio/evals/scorer-online.ts index f0d8386e656..a7b33f69486 100644 --- a/apps/studio/evals/scorer-online.ts +++ b/apps/studio/evals/scorer-online.ts @@ -8,7 +8,7 @@ * - correctnessScorer: requires ground truth (expected output), offline-eval-only. */ -import braintrust, { type EvalScorer } from 'braintrust' +import braintrust from 'braintrust' import { completenessScorer, @@ -17,9 +17,7 @@ import { goalCompletionScorer, safetyScorer, urlValidityScorer, - type AssistantEvalInput, - type AssistantEvalOutput, - type Expected, + type AssistantEvalScorer, } from './scorer' import manifest from './scorer-online-manifest.json' @@ -47,7 +45,7 @@ const handlers = { // Online traces have no expected, so we wrap it to always run. safety: (args) => safetyScorer({ ...args, expected: { ...args.expected, requiresSafetyCheck: true } }), -} satisfies Record> +} satisfies Record // @ts-expect-error - Project ID is only required at build-time const project = braintrust.projects.create({ id: projectId }) diff --git a/apps/studio/evals/scorer-wasm.ts b/apps/studio/evals/scorer-wasm.ts index 0856e8cd904..b7c91ea119f 100644 --- a/apps/studio/evals/scorer-wasm.ts +++ b/apps/studio/evals/scorer-wasm.ts @@ -1,7 +1,7 @@ -import { EvalScorer, Trace } from 'braintrust' +import { Trace } from 'braintrust' import { parse } from 'libpg-query' -import { AssistantEvalInput, AssistantEvalOutput, Expected } from './scorer' +import { AssistantEvalScorer } from './scorer' import { getParsedToolSpans } from './trace-utils' import { executeSqlInputSchema } from '@/lib/ai/tools/studio-tools' import { extractIdentifiers, isQuotedInSql, needsQuoting } from '@/lib/sql-identifier-quoting' @@ -14,11 +14,7 @@ async function getSqlQueries(trace: Trace): Promise { return spans.map((s) => s.input.sql) } -export const sqlSyntaxScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { +export const sqlSyntaxScorer: AssistantEvalScorer = async ({ trace }) => { if (!trace) return null const sqlQueries = await getSqlQueries(trace) @@ -44,11 +40,7 @@ export const sqlSyntaxScorer: EvalScorer< } } -export const sqlIdentifierQuotingScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { +export const sqlIdentifierQuotingScorer: AssistantEvalScorer = async ({ trace }) => { if (!trace) return null const sqlQueries = await getSqlQueries(trace) diff --git a/apps/studio/evals/scorer.test.ts b/apps/studio/evals/scorer.test.ts new file mode 100644 index 00000000000..a6150c369be --- /dev/null +++ b/apps/studio/evals/scorer.test.ts @@ -0,0 +1,69 @@ +import type { Trace } from 'braintrust' +import { describe, expect, it, vi } from 'vitest' + +import { urlValidityScorer, type AssistantEvalOutput } from './scorer' +import type { Transcript } from './transcript' + +const DOCS_URL = 'https://supabase.com/docs/guides/auth' + +/** + * Minimal stand-in for a live Trace. Only getThread is exercised by the online + * fallback path (see resolveTranscript in scorer.ts). + */ +const mockTrace = (thread: unknown[]) => { + const getThread = vi.fn().mockResolvedValue(thread) + return { trace: { getThread } as unknown as Trace, getThread } +} + +const THREAD_WITH_DOCS_URL = [ + { role: 'user', content: 'Where are the auth docs?' }, + { role: 'assistant', content: [{ type: 'text', text: `See ${DOCS_URL} for details.` }] }, +] + +const transcript = (lastAssistantTurn: string): Transcript => ({ + currentUserInput: 'Where are the auth docs?', + priorConversation: null, + lastAssistantTurn, + lastAssistantTurnWithToolInputs: lastAssistantTurn, +}) + +const runUrlValidityScorer = (output: AssistantEvalOutput | null, trace?: Trace) => + urlValidityScorer({ + input: { prompt: 'Where are the auth docs?' }, + expected: {}, + output, + trace, + }) + +describe('scorers with online (null) output', () => { + // Online scorers run against live production logs, which have no eval task + // and therefore no output at all — Braintrust passes null. + it('returns null instead of throwing when there is neither output nor trace', async () => { + await expect(runUrlValidityScorer(null)).resolves.toBeNull() + }) + + it('derives the transcript from the trace when output is null', async () => { + const fetchMock = vi.fn().mockResolvedValue({ ok: true, status: 200 }) + vi.stubGlobal('fetch', fetchMock) + + const { trace, getThread } = mockTrace(THREAD_WITH_DOCS_URL) + const result = await runUrlValidityScorer(null, trace) + + expect(getThread).toHaveBeenCalled() + expect(result).toMatchObject({ name: 'URL Validity', score: 1, metadata: { urls: [DOCS_URL] } }) + + vi.unstubAllGlobals() + }) + + it('prefers the offline task transcript over the trace when both are present', async () => { + const { trace, getThread } = mockTrace(THREAD_WITH_DOCS_URL) + const output = { + finishReason: 'stop' as const, + transcript: transcript('No links here, just prose.'), + } + + // No supabase URLs in the offline transcript, so the scorer opts out. + await expect(runUrlValidityScorer(output, trace)).resolves.toBeNull() + expect(getThread).not.toHaveBeenCalled() + }) +}) diff --git a/apps/studio/evals/scorer.ts b/apps/studio/evals/scorer.ts index 43b4080ce08..998e856a013 100644 --- a/apps/studio/evals/scorer.ts +++ b/apps/studio/evals/scorer.ts @@ -1,10 +1,11 @@ import { FinishReason } from 'ai' import { LLMClassifierFromTemplate } from 'autoevals' -import { EvalCase, EvalScorer } from 'braintrust' +import { EvalCase, EvalScorer, type Trace } from 'braintrust' import { stripIndent } from 'common-tags' import { z } from 'zod' import { getParsedToolSpans, getThreadParts, getToolSpans } from './trace-utils' +import type { Transcript } from './transcript' import { loadKnowledgeInputSchema } from '@/lib/ai/tools/studio-tools' import { extractUrls } from '@/lib/helpers' @@ -22,8 +23,10 @@ export type AssistantEvalInput = { > } +/** What the offline eval task in assistant.eval.ts returns. */ export type AssistantEvalOutput = { finishReason: FinishReason + transcript: Transcript } type ToolInputExactValue = string | number | boolean | null | string[] @@ -56,12 +59,38 @@ export type AssistantEvalCaseMetadata = { export type AssistantEvalCase = EvalCase +/** + * Note the nullable output: offline, scorers get the eval task's + * AssistantEvalOutput, but online scorers run against live production logs that + * have no eval task behind them, so Braintrust passes `null`. Those scorers + * derive an equivalent Transcript from `trace` instead — see resolveTranscript. + */ +export type AssistantEvalScorer = EvalScorer< + AssistantEvalInput, + AssistantEvalOutput | null, + Expected +> + // --- Trace helpers --- const mcpTextContentSpanOutputSchema = z.object({ content: z.array(z.object({ type: z.literal('text').optional(), text: z.string() })), }) +/** + * Prefers the offline eval task's in-memory transcript (untruncated); falls + * back to deriving one from the live trace for online scorers, which have no + * such task output to read. + */ +async function resolveTranscript( + output: AssistantEvalOutput | null, + trace: Trace | undefined +): Promise { + if (output?.transcript) return output.transcript + if (!trace) return null + return getThreadParts(trace) +} + // --- Scorers --- const matchesToolInputField = (actual: unknown, expected: ToolInputFieldExpectation) => { @@ -83,11 +112,7 @@ const matchesExpectedToolInput = ( }) } -export const toolUsageScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ expected, trace }) => { +export const toolUsageScorer: AssistantEvalScorer = async ({ expected, trace }) => { if (!expected.requiredTools || !trace) return null const toolSpans = await getToolSpans(trace) @@ -113,11 +138,7 @@ export const toolUsageScorer: EvalScorer< } } -export const knowledgeUsageScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ expected, trace }) => { +export const knowledgeUsageScorer: AssistantEvalScorer = async ({ expected, trace }) => { if (!expected.requiredKnowledge || !trace) return null const knowledgeSpans = await getParsedToolSpans(trace, 'load_knowledge', { @@ -157,17 +178,12 @@ const concisenessEvaluator = LLMClassifierFromTemplate<{ input: string }>({ model: LLM_AS_A_JUDGE_MODEL, }) -export const concisenessScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { - if (!trace) return null - const parts = await getThreadParts(trace) - if (!parts.currentUserInput || !parts.lastAssistantTurn) return null +export const concisenessScorer: AssistantEvalScorer = async ({ output, trace }) => { + const transcript = await resolveTranscript(output, trace) + if (!transcript?.lastAssistantTurn) return null return await concisenessEvaluator({ - input: parts.currentUserInput, - output: parts.lastAssistantTurn, + input: transcript.currentUserInput, + output: transcript.lastAssistantTurn, }) } @@ -188,17 +204,12 @@ const completenessEvaluator = LLMClassifierFromTemplate<{ input: string }>({ model: LLM_AS_A_JUDGE_MODEL, }) -export const completenessScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { - if (!trace) return null - const parts = await getThreadParts(trace, { includeToolCallInputs: true }) - if (!parts.currentUserInput || !parts.lastAssistantTurn) return null +export const completenessScorer: AssistantEvalScorer = async ({ output, trace }) => { + const transcript = await resolveTranscript(output, trace) + if (!transcript?.lastAssistantTurnWithToolInputs) return null return await completenessEvaluator({ - input: parts.currentUserInput, - output: parts.lastAssistantTurn, + input: transcript.currentUserInput, + output: transcript.lastAssistantTurnWithToolInputs, }) } @@ -229,18 +240,13 @@ const goalCompletionEvaluator = LLMClassifierFromTemplate<{ model: LLM_AS_A_JUDGE_MODEL, }) -export const goalCompletionScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { - if (!trace) return null - const parts = await getThreadParts(trace, { includeToolCallInputs: true }) - if (!parts.currentUserInput || !parts.lastAssistantTurn) return null +export const goalCompletionScorer: AssistantEvalScorer = async ({ output, trace }) => { + const transcript = await resolveTranscript(output, trace) + if (!transcript?.lastAssistantTurnWithToolInputs) return null return await goalCompletionEvaluator({ - input: parts.currentUserInput, - priorConversation: parts.priorConversation ?? 'None', - output: parts.lastAssistantTurn, + input: transcript.currentUserInput, + priorConversation: transcript.priorConversation ?? 'None', + output: transcript.lastAssistantTurnWithToolInputs, }) } @@ -265,11 +271,7 @@ const docsFaithfulnessEvaluator = LLMClassifierFromTemplate<{ docs: string }>({ model: LLM_AS_A_JUDGE_MODEL, }) -export const docsFaithfulnessScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { +export const docsFaithfulnessScorer: AssistantEvalScorer = async ({ output, trace }) => { if (!trace) return null const docsSpans = await getToolSpans(trace, 'search_docs') @@ -290,12 +292,12 @@ export const docsFaithfulnessScorer: EvalScorer< if (docs.length === 0) return null - const parts = await getThreadParts(trace, { includeToolCallInputs: true }) - if (!parts.lastAssistantTurn) return null + const transcript = await resolveTranscript(output, trace) + if (!transcript?.lastAssistantTurnWithToolInputs) return null return await docsFaithfulnessEvaluator({ docs: docs.join('\n\n'), - output: parts.lastAssistantTurn, + output: transcript.lastAssistantTurnWithToolInputs, }) } @@ -328,19 +330,17 @@ const correctnessEvaluator = LLMClassifierFromTemplate<{ input: string; expected model: LLM_AS_A_JUDGE_MODEL, }) -export const correctnessScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ expected, trace }) => { - if (!expected.correctAnswer || !trace) return null - const parts = await getThreadParts(trace, { includeToolCallInputs: true }) - if (!parts.currentUserInput || !parts.lastAssistantTurn) return null +export const correctnessScorer: AssistantEvalScorer = async ({ expected, output }) => { + if (!expected.correctAnswer) return null + // Correctness needs ground truth, so it only ever runs offline where the eval + // task's transcript is present — no trace fallback needed. + const transcript = output?.transcript + if (!transcript?.lastAssistantTurnWithToolInputs) return null return await correctnessEvaluator({ - input: parts.currentUserInput, + input: transcript.currentUserInput, expected: expected.correctAnswer, - output: parts.lastAssistantTurn, + output: transcript.lastAssistantTurnWithToolInputs, }) } @@ -371,32 +371,24 @@ const safetyEvaluator = LLMClassifierFromTemplate<{ input: string; priorConversa model: LLM_AS_A_JUDGE_MODEL, }) -export const safetyScorer: EvalScorer = async ({ - expected, - trace, -}) => { - if (!expected.requiresSafetyCheck || !trace) return null +export const safetyScorer: AssistantEvalScorer = async ({ expected, output, trace }) => { + if (!expected.requiresSafetyCheck) return null - const parts = await getThreadParts(trace, { includeToolCallInputs: true }) - if (!parts.currentUserInput || !parts.lastAssistantTurn) return null + const transcript = await resolveTranscript(output, trace) + if (!transcript?.lastAssistantTurnWithToolInputs) return null return await safetyEvaluator({ - input: parts.currentUserInput, - priorConversation: parts.priorConversation ?? 'None', - output: parts.lastAssistantTurn, + input: transcript.currentUserInput, + priorConversation: transcript.priorConversation ?? 'None', + output: transcript.lastAssistantTurnWithToolInputs, }) } -export const urlValidityScorer: EvalScorer< - AssistantEvalInput, - AssistantEvalOutput, - Expected -> = async ({ trace }) => { - if (!trace) return null - const parts = await getThreadParts(trace) - if (!parts.lastAssistantTurn) return null +export const urlValidityScorer: AssistantEvalScorer = async ({ output, trace }) => { + const transcript = await resolveTranscript(output, trace) + if (!transcript?.lastAssistantTurn) return null - const allUrls = extractUrls(parts.lastAssistantTurn, { + const allUrls = extractUrls(transcript.lastAssistantTurn, { excludeCodeBlocks: true, excludeTemplates: true, }) diff --git a/apps/studio/evals/trace-utils.test.ts b/apps/studio/evals/trace-utils.test.ts index 60c47409cae..e74aab1fb22 100644 --- a/apps/studio/evals/trace-utils.test.ts +++ b/apps/studio/evals/trace-utils.test.ts @@ -116,20 +116,14 @@ const MOCK_THREAD = [ ] describe('getThreadPartsFromThread', () => { - it('parses a sanitized Braintrust trace.getThread payload', () => { + it('parses a sanitized Braintrust trace.getThread payload, computing both tool-input variants', () => { expect(getThreadPartsFromThread(MOCK_THREAD)).toEqual({ - projectContext: "The user's current project is Acme Analytics.", + currentUserInput: 'Can you create that orders table now?', priorConversation: '[user]\nWhat did we decide earlier?\n\n[assistant]\nWe decided to add an orders table with RLS policies before generating sample data.', - currentUserInput: 'Can you create that orders table now?', lastAssistantTurn: '[assistant]\n[called rename_chat]\n\n[assistant]\n[called load_knowledge]\n[called execute_sql]\n\n[assistant]\nI created the public.orders table. You should add RLS policies before exposing it to users.', - }) - }) - - it('can include tool call inputs in serialized assistant turns', () => { - expect(getThreadPartsFromThread(MOCK_THREAD, { includeToolCallInputs: true })).toMatchObject({ - lastAssistantTurn: `\ + lastAssistantTurnWithToolInputs: `\ [assistant] [called rename_chat] { @@ -151,21 +145,24 @@ I created the public.orders table. You should add RLS policies before exposing i }) }) - it('uses the most recent project context message', () => { - expect( - getThreadPartsFromThread([ - { - role: 'assistant', - content: "The user's current project is Old Project.", - }, - ...MOCK_THREAD, - ]) - ).toMatchObject({ - projectContext: "The user's current project is Acme Analytics.", - }) + it('filters out project-context messages so they never leak into prior conversation', () => { + const threadWithExtraProjectContext = [ + { + role: 'assistant', + content: "The user's current project is Old Project.", + }, + ...MOCK_THREAD, + ] + + expect(getThreadPartsFromThread(threadWithExtraProjectContext)).toEqual( + getThreadPartsFromThread(MOCK_THREAD) + ) + expect(getThreadPartsFromThread(threadWithExtraProjectContext).priorConversation).not.toContain( + 'Old Project' + ) }) - it('returns prior conversation without current turn parts when there is no user message', () => { + it('treats all messages as prior conversation when there is no user message', () => { expect( getThreadPartsFromThread([ { @@ -174,10 +171,10 @@ I created the public.orders table. You should add RLS policies before exposing i }, ]) ).toEqual({ - projectContext: null, + currentUserInput: '', priorConversation: '[assistant]\nI can help with your Supabase project.', - currentUserInput: null, lastAssistantTurn: null, + lastAssistantTurnWithToolInputs: null, }) }) }) diff --git a/apps/studio/evals/trace-utils.ts b/apps/studio/evals/trace-utils.ts index 94114491382..3a7aef8ea7f 100644 --- a/apps/studio/evals/trace-utils.ts +++ b/apps/studio/evals/trace-utils.ts @@ -1,7 +1,7 @@ import type { SpanData, Trace } from 'braintrust' import { z } from 'zod' -const projectContextPrefix = "The user's current project is " +import type { Transcript } from './transcript' /** * Matches AI SDK tool spans as Braintrust records them: tool args first, @@ -17,35 +17,6 @@ const aiSdkToolSpanInputSchema = z.tuple([ .passthrough(), ]) -const threadTextBlockSchema = z.object({ type: z.literal('text'), text: z.string() }) -const threadToolCallArgumentsSchema = z.object({ type: z.literal('valid'), value: z.unknown() }) -const threadToolCallBlockSchema = z.object({ - type: z.literal('tool_call'), - tool_name: z.string(), - arguments: z.unknown().optional(), -}) -const threadContentBlockSchema = z.union([threadTextBlockSchema, threadToolCallBlockSchema]) -const threadContentSchema = z.union([ - z.string(), - z.array(z.unknown()).transform((blocks) => - blocks.flatMap((block) => { - const result = threadContentBlockSchema.safeParse(block) - return result.success ? [result.data] : [] - }) - ), -]) -const threadMessageSchema = z.object({ - role: z.enum(['system', 'user', 'assistant', 'tool']), - content: threadContentSchema, -}) - -type ThreadMessage = z.infer -type ThreadContentBlock = z.infer - -export type ThreadSerializationOptions = { - includeToolCallInputs?: boolean -} - /** Normalized Braintrust tool span with unwrapped tool input and raw output. */ export type ToolSpan = { span: SpanData @@ -53,13 +24,6 @@ export type ToolSpan = { output: unknown } -export type ThreadParts = { - projectContext: string | null - priorConversation: string | null - currentUserInput: string | null - lastAssistantTurn: string | null -} - /** Optional schemas used to validate and type a tool span's input and output. */ type ToolSpanSchemas< TInputSchema extends z.ZodType | undefined, @@ -85,105 +49,6 @@ function getToolSpanInput(span: SpanData): unknown { return result.success ? result.data[0] : span.input } -function unwrapToolCallArguments(args: unknown): unknown { - const result = threadToolCallArgumentsSchema.safeParse(args) - return result.success ? result.data.value : args -} - -function serializeContentBlock( - block: ThreadContentBlock, - options: ThreadSerializationOptions -): string { - if (block.type === 'text') return block.text - - const marker = `[called ${block.tool_name}]` - if (!options.includeToolCallInputs || typeof block.arguments === 'undefined') return marker - - return `${marker}\n${JSON.stringify(unwrapToolCallArguments(block.arguments), null, 2)}` -} - -function serializeMessageContent( - message: ThreadMessage | undefined, - options: ThreadSerializationOptions = {} -): string | null { - if (!message) return null - if (typeof message.content === 'string') return message.content || null - - const content = message.content.map((block) => serializeContentBlock(block, options)).join('\n') - - return content || null -} - -function serializeMessages( - messages: ThreadMessage[], - options: ThreadSerializationOptions = {} -): string | null { - const parts = messages.flatMap((message) => { - const content = serializeMessageContent(message, options) - return content ? [`[${message.role}]\n${content}`] : [] - }) - - return parts.length > 0 ? parts.join('\n\n') : null -} - -function isProjectContextMessage(message: ThreadMessage): boolean { - return ( - message.role === 'assistant' && - Boolean(serializeMessageContent(message)?.startsWith(projectContextPrefix)) - ) -} - -function findLastUserIndex(messages: ThreadMessage[]): number { - for (let i = messages.length - 1; i >= 0; i--) { - if (messages[i].role === 'user') return i - } - return -1 -} - -export function getThreadPartsFromThread( - thread: unknown[], - options: ThreadSerializationOptions = {} -): ThreadParts { - const messages = thread.flatMap((message) => { - const result = threadMessageSchema.safeParse(message) - if (!result.success || result.data.role === 'system' || result.data.role === 'tool') return [] - return [result.data] - }) - - const projectContextMessages = messages.filter(isProjectContextMessage) - const chatMessages = messages.filter((message) => !isProjectContextMessage(message)) - const lastUserIdx = findLastUserIndex(chatMessages) - const projectContext = serializeMessageContent( - projectContextMessages[projectContextMessages.length - 1] - ) - - if (lastUserIdx === -1) { - return { - projectContext, - priorConversation: serializeMessages(chatMessages, options), - currentUserInput: null, - lastAssistantTurn: null, - } - } - - return { - projectContext, - priorConversation: serializeMessages(chatMessages.slice(0, lastUserIdx), options), - currentUserInput: serializeMessageContent(chatMessages[lastUserIdx]), - lastAssistantTurn: serializeMessages( - chatMessages.slice(lastUserIdx + 1).filter((message) => message.role === 'assistant'), - options - ), - } -} - -export async function getThreadParts( - trace: Trace, - options: ThreadSerializationOptions = {} -): Promise { - return getThreadPartsFromThread(await trace.getThread(), options) -} - /** Returns normalized tool spans from the trace, optionally filtered to a specific tool name. */ export async function getToolSpans(trace: Trace, toolName?: string): Promise { const spans = await trace.getSpans({ spanType: ['tool'] }) @@ -223,3 +88,140 @@ export async function getParsedToolSpans< ] }) } + +// --- Thread parsing (fallback path for online scorers) --- +// +// Online scorers run against live production traces, which have no in-memory +// transcript to read (that only exists inside the offline eval task — see +// buildTranscript in transcript.ts). trace.getThread() is the only source of +// the full conversation available to them there, so we derive an equivalent +// Transcript from it. This is subject to Braintrust's getThread() +// preview-length truncation, unlike buildTranscript's in-memory path — that's +// a known, pre-existing limitation of the online-scoring path, not something +// this fallback introduces. + +const projectContextPrefix = "The user's current project is " + +const threadTextBlockSchema = z.object({ type: z.literal('text'), text: z.string() }) +const threadToolCallArgumentsSchema = z.object({ type: z.literal('valid'), value: z.unknown() }) +const threadToolCallBlockSchema = z.object({ + type: z.literal('tool_call'), + tool_name: z.string(), + arguments: z.unknown().optional(), +}) +const threadContentBlockSchema = z.union([threadTextBlockSchema, threadToolCallBlockSchema]) +const threadContentSchema = z.union([ + z.string(), + z.array(z.unknown()).transform((blocks) => + blocks.flatMap((block) => { + const result = threadContentBlockSchema.safeParse(block) + return result.success ? [result.data] : [] + }) + ), +]) +const threadMessageSchema = z.object({ + role: z.enum(['system', 'user', 'assistant', 'tool']), + content: threadContentSchema, +}) + +type ThreadMessage = z.infer +type ThreadContentBlock = z.infer + +function unwrapToolCallArguments(args: unknown): unknown { + const result = threadToolCallArgumentsSchema.safeParse(args) + return result.success ? result.data.value : args +} + +function serializeContentBlock(block: ThreadContentBlock, includeToolCallInputs: boolean): string { + if (block.type === 'text') return block.text + + const marker = `[called ${block.tool_name}]` + if (!includeToolCallInputs || typeof block.arguments === 'undefined') return marker + + return `${marker}\n${JSON.stringify(unwrapToolCallArguments(block.arguments), null, 2)}` +} + +function serializeMessageContent( + message: ThreadMessage | undefined, + includeToolCallInputs: boolean +): string | null { + if (!message) return null + if (typeof message.content === 'string') return message.content || null + + const content = message.content + .map((block) => serializeContentBlock(block, includeToolCallInputs)) + .join('\n') + + return content || null +} + +function serializeMessages( + messages: ThreadMessage[], + includeToolCallInputs: boolean +): string | null { + const parts = messages.flatMap((message) => { + const content = serializeMessageContent(message, includeToolCallInputs) + return content ? [`[${message.role}]\n${content}`] : [] + }) + + return parts.length > 0 ? parts.join('\n\n') : null +} + +function isProjectContextMessage(message: ThreadMessage): boolean { + return ( + message.role === 'assistant' && + Boolean(serializeMessageContent(message, false)?.startsWith(projectContextPrefix)) + ) +} + +function findLastUserIndex(messages: ThreadMessage[]): number { + for (let i = messages.length - 1; i >= 0; i--) { + if (messages[i].role === 'user') return i + } + return -1 +} + +/** + * Parses a raw trace.getThread() payload into a Transcript. Pulled out from + * getThreadParts so it can be unit tested without a live Trace. + */ +export function getThreadPartsFromThread(thread: unknown[]): Transcript { + const messages = thread.flatMap((message) => { + const result = threadMessageSchema.safeParse(message) + if (!result.success || result.data.role === 'system' || result.data.role === 'tool') return [] + return [result.data] + }) + + const chatMessages = messages.filter((message) => !isProjectContextMessage(message)) + const lastUserIdx = findLastUserIndex(chatMessages) + + if (lastUserIdx === -1) { + return { + currentUserInput: '', + priorConversation: serializeMessages(chatMessages, true), + lastAssistantTurn: null, + lastAssistantTurnWithToolInputs: null, + } + } + + const assistantMessages = chatMessages + .slice(lastUserIdx + 1) + .filter((message) => message.role === 'assistant') + + return { + currentUserInput: serializeMessageContent(chatMessages[lastUserIdx], false) ?? '', + priorConversation: serializeMessages(chatMessages.slice(0, lastUserIdx), true), + lastAssistantTurn: serializeMessages(assistantMessages, false), + lastAssistantTurnWithToolInputs: serializeMessages(assistantMessages, true), + } +} + +/** + * Derives a Transcript from a live trace's thread. Used as the fallback for + * online scorers, which have no in-memory transcript to read (see + * getThreadPartsFromThread above for why, and why this is truncation-prone + * in a way the offline path isn't). + */ +export async function getThreadParts(trace: Trace): Promise { + return getThreadPartsFromThread(await trace.getThread()) +} diff --git a/apps/studio/evals/transcript.test.ts b/apps/studio/evals/transcript.test.ts new file mode 100644 index 00000000000..a889d6507d7 --- /dev/null +++ b/apps/studio/evals/transcript.test.ts @@ -0,0 +1,79 @@ +import type { StepResult, ToolSet } from 'ai' +import { describe, expect, it } from 'vitest' + +import { buildTranscript } from './transcript' + +// Minimal fixtures matching only the `content` shape buildTranscript reads — +// no need to fully populate the rest of StepResult's fields. +const makeStep = (content: unknown[]): StepResult => + ({ content }) as unknown as StepResult + +describe('buildTranscript', () => { + it('serializes a single step with only a text part', () => { + const steps = [makeStep([{ type: 'text', text: 'Here is your answer.' }])] + + const transcript = buildTranscript('What is the answer?', steps) + + expect(transcript.currentUserInput).toBe('What is the answer?') + expect(transcript.priorConversation).toBeNull() + expect(transcript.lastAssistantTurn).toBe('Here is your answer.') + expect(transcript.lastAssistantTurnWithToolInputs).toBe('Here is your answer.') + }) + + it('serializes a single step with only a tool-call part', () => { + const steps = [ + makeStep([{ type: 'tool-call', toolName: 'execute_sql', input: { sql: 'select 1;' } }]), + ] + + const transcript = buildTranscript('Run a query.', steps) + + expect(transcript.lastAssistantTurn).toBe('[called execute_sql]') + expect(transcript.lastAssistantTurnWithToolInputs).toBe( + '[called execute_sql]\n' + JSON.stringify({ sql: 'select 1;' }, null, 2) + ) + }) + + it('joins multiple steps in order with a blank line between them', () => { + const steps = [ + makeStep([{ type: 'text', text: 'Let me check that.' }]), + makeStep([{ type: 'tool-call', toolName: 'execute_sql', input: { sql: 'select 1;' } }]), + makeStep([{ type: 'text', text: 'The answer is 1.' }]), + ] + + const transcript = buildTranscript('What is 1?', steps) + + expect(transcript.lastAssistantTurn).toBe( + 'Let me check that.\n\n[called execute_sql]\n\nThe answer is 1.' + ) + expect(transcript.lastAssistantTurnWithToolInputs).toBe( + 'Let me check that.\n\n[called execute_sql]\n' + + JSON.stringify({ sql: 'select 1;' }, null, 2) + + '\n\nThe answer is 1.' + ) + expect(transcript.lastAssistantTurn?.endsWith('The answer is 1.')).toBe(true) + }) + + it('skips reasoning and tool-result content parts', () => { + const steps = [ + makeStep([ + { type: 'reasoning', text: 'Thinking about the best approach...' }, + { type: 'tool-call', toolName: 'execute_sql', input: { sql: 'select 1;' } }, + { type: 'tool-result', toolName: 'execute_sql', output: { rows: [{ '1': 1 }] } }, + { type: 'text', text: 'The result is 1.' }, + ]), + ] + + const transcript = buildTranscript('What is 1?', steps) + + expect(transcript.lastAssistantTurn).toBe('[called execute_sql]\nThe result is 1.') + expect(transcript.lastAssistantTurn).not.toContain('Thinking about the best approach') + expect(transcript.lastAssistantTurn).not.toContain('rows') + }) + + it('returns null transcripts for an empty steps array', () => { + const transcript = buildTranscript('Hello', []) + + expect(transcript.lastAssistantTurn).toBeNull() + expect(transcript.lastAssistantTurnWithToolInputs).toBeNull() + }) +}) diff --git a/apps/studio/evals/transcript.ts b/apps/studio/evals/transcript.ts new file mode 100644 index 00000000000..13ed3a0e0b7 --- /dev/null +++ b/apps/studio/evals/transcript.ts @@ -0,0 +1,52 @@ +import type { StepResult, ToolSet } from 'ai' + +export type Transcript = { + /** The user's prompt for this eval case. Always available locally — no reconstruction needed. */ + currentUserInput: string + /** + * Serialized prior conversation before the current turn, or null. + */ + priorConversation: string | null + /** Text-only serialization of the assistant's turn: prose only, no tool call markers. */ + lastAssistantTurn: string | null + /** Same as lastAssistantTurn, but with `[called toolName]` markers (and JSON args) for each tool call, interleaved in order. */ + lastAssistantTurnWithToolInputs: string | null +} + +/** + * Builds a Transcript directly from the AI SDK's step history — no trace/network round-trip. + * Mirrors the marker format the old trace.getThread()-based serialization used + * (`[called toolName]\n`), so scorer prompts don't need to change: + * only text and tool-call content parts are represented; reasoning/source/file/tool-result/ + * tool-error parts are skipped, exactly as the old serializeContentBlock did. + */ +export function buildTranscript( + currentUserInput: string, + steps: ReadonlyArray> +): Transcript { + const renderStep = (step: StepResult, includeToolInputs: boolean): string => + step.content + .flatMap((part) => { + if (part.type === 'text') return part.text ? [part.text] : [] + if (part.type === 'tool-call') { + const marker = `[called ${part.toolName}]` + return [includeToolInputs ? `${marker}\n${JSON.stringify(part.input, null, 2)}` : marker] + } + return [] + }) + .join('\n') + + const joinSteps = (includeToolInputs: boolean): string | null => { + const rendered = steps + .map((step) => renderStep(step, includeToolInputs)) + .filter((s) => s.length > 0) + return rendered.length > 0 ? rendered.join('\n\n') : null + } + + return { + currentUserInput, + priorConversation: null, + lastAssistantTurn: joinSteps(false), + lastAssistantTurnWithToolInputs: joinSteps(true), + } +}