mirror of
https://github.com/supabase/supabase.git
synced 2026-10-05 01:15:03 +03:00
Fix eval scorer truncation via local transcript capture (#49151)
## Problem Scorers previously derived the assistant's final answer via Braintrust's `trace.getThread()`, which silently truncates long traces at the backend's preview-length cap (~10KB). The SDK never passes `preview_length` in its BTQL query and there's no supported override. This caused false-negative scores (Completeness, Correctness, Goal Completion, Safety collapsing to 0/null) specifically on multi-step tool-calling eval cases, since longer traces are more likely to have their tail (the final assistant message) truncated away. ## Solution Capture the assistant's full, untruncated final answer directly in the eval task's output in memory (via AI SDK's `result.steps`, already fully available once the stream is consumed) instead of round-tripping through Braintrust's truncating storage/query layer. Scorers now read `output.transcript` instead of calling `trace.getThread()`. ## Changes - **New**: `apps/studio/evals/transcript.ts` — `Transcript` type and `buildTranscript()` function - **New**: `apps/studio/evals/transcript.test.ts` — unit tests (5 passing) - **Modified**: `apps/studio/evals/assistant.eval.ts` — captures `result.steps` and returns transcript - **Modified**: `apps/studio/evals/scorer.ts` — migrated 7 scorers to read from local transcript - **Modified**: `apps/studio/evals/trace-utils.ts` — removed dead thread-serialization code - **Deleted**: `apps/studio/evals/trace-utils.test.ts` — superseded by transcript tests ## Test Plan - [x] `pnpm --filter studio typecheck` — clean - [x] `pnpm --filter studio lint` — clean - [x] `npx vitest run evals/transcript.test.ts` — 5/5 passing - [x] Full live eval run (35/35 cases) against Braintrust — [experiment](https://www.braintrust.dev/app/supabase.io/p/Assistant/experiments/eval-scorer-transcript-capture-1786985352) shows Completeness/Correctness/Goal Completion/Safety scores comparable to baseline ## Known Residual Risk Other scorers that derive data from `trace.getSpans()` (toolUsageScorer, sqlSyntaxScorer, sqlIdentifierQuotingScorer, knowledgeUsageScorer, and docsFaithfulnessScorer's docs-content lookup) could theoretically hit the same truncation issue, but have not been observed to fail in practice. This is not addressed in this PR. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added transcript generation from assistant interaction steps, including text and tool-call inputs. * Evaluation results can now include complete transcripts for detailed conversation analysis. * Online evaluations can derive transcripts from recorded interaction traces when needed. * **Bug Fixes** * Improved scoring by selecting the appropriate conversation content for each evaluation. * Ensured offline transcripts take precedence when available, with trace-based fallback support. * **Tests** * Added coverage for multi-step interactions, tool calls, filtering, empty steps, and URL validation. <!-- end of auto-generated comment: release notes by coderabbit.ai -->
This commit is contained in:
1 parent
ce2ed77c02
commit
8c409e2df5
9 files changed
+440
-257
No files matched your search
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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<string, EvalScorer<AssistantEvalInput, AssistantEvalOutput, Expected>>
|
||||
} satisfies Record<string, AssistantEvalScorer>
|
||||
|
||||
// @ts-expect-error - Project ID is only required at build-time
|
||||
const project = braintrust.projects.create({ id: projectId })
|
||||
|
||||
@@ -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<string[]> {
|
||||
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)
|
||||
|
||||
@@ -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()
|
||||
})
|
||||
})
|
||||
+71
-79
@@ -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<AssistantEvalInput, Expected, AssistantEvalCaseMetadata>
|
||||
|
||||
/**
|
||||
* 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<Transcript | null> {
|
||||
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<AssistantEvalInput, AssistantEvalOutput, Expected> = 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,
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
})
|
||||
})
|
||||
+138
-136
@@ -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<typeof threadMessageSchema>
|
||||
type ThreadContentBlock = z.infer<typeof threadContentBlockSchema>
|
||||
|
||||
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<ThreadParts> {
|
||||
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<ToolSpan[]> {
|
||||
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<typeof threadMessageSchema>
|
||||
type ThreadContentBlock = z.infer<typeof threadContentBlockSchema>
|
||||
|
||||
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<Transcript> {
|
||||
return getThreadPartsFromThread(await trace.getThread())
|
||||
}
|
||||
@@ -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<ToolSet> =>
|
||||
({ content }) as unknown as StepResult<ToolSet>
|
||||
|
||||
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()
|
||||
})
|
||||
})
|
||||
@@ -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<json args>`), 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<StepResult<ToolSet>>
|
||||
): Transcript {
|
||||
const renderStep = (step: StepResult<ToolSet>, 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),
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user