feat: split AssertionScorer into separate scorers

This commit is contained in:
Matt Rossman committed 2025-12-12 11:44:31 -05:00
1 parent 4e9d7d9c38
commit 5f601bc9d3
2 files changed
+91 -115

No files matched your search

+18 -21
View File
@@ -1,7 +1,13 @@
import { openai } from '@ai-sdk/openai'
import { Eval } from 'braintrust'
import { stripIndent } from 'common-tags'
import { Assertion, createAssertionScorer } from 'lib/ai/evals/scorer'
import {
AssistantEvalCase,
criteriaMetScorer,
sqlSimilarityScorer,
textIncludesScorer,
toolUsageScorer,
} from 'lib/ai/evals/scorer'
import { generateAssistantResponse } from 'lib/ai/generate-assistant-response'
import { getMockTools } from 'lib/ai/tools/mock-tools'
import assert from 'node:assert'
@@ -9,41 +15,32 @@ import assert from 'node:assert'
assert(process.env.BRAINTRUST_PROJECT_ID, 'BRAINTRUST_PROJECT_ID is not set')
assert(process.env.OPENAI_API_KEY, 'OPENAI_API_KEY is not set')
type MockToolName = keyof Awaited<ReturnType<typeof getMockTools>>
type EvalAssertion = Assertion<MockToolName>
const AssertionScorer = createAssertionScorer<MockToolName>()
Eval('Assistant', {
projectId: process.env.BRAINTRUST_PROJECT_ID,
data: () => {
return [
{
input: 'Hello!',
expected: [{ type: 'text_includes', substring: 'Hi' }],
expected: { textIncludes: 'Hi' },
},
{
input: 'How do I implement IP address rate limiting?',
expected: [{ type: 'tools_include', tool: 'search_docs' }],
expected: { requiredTools: ['search_docs'] },
},
{
input: 'Check if my project is having issues right now and tell me what to fix first.',
expected: [
{ type: 'tools_include', tool: 'get_advisors' },
{ type: 'tools_include', tool: 'get_logs' },
{ type: 'llm_criteria_met', criteria: 'Response reflects there are RLS issues to fix.' },
],
expected: {
requiredTools: ['get_advisors', 'get_logs'],
criteria: 'Response reflects there are RLS issues to fix.',
},
},
{
input: 'Create a new table "foods" with columns for "name" and "color"',
expected: [
{
type: 'sql_similar',
sql: stripIndent`CREATE TABLE IF NOT EXISTS public.foods ( id bigserial PRIMARY KEY, name text NOT NULL, color text );`,
},
],
expected: {
sqlQuery: stripIndent`CREATE TABLE IF NOT EXISTS public.foods ( id bigserial PRIMARY KEY, name text NOT NULL, color text );`,
},
},
] satisfies Array<{ input: string; expected: EvalAssertion[] }>
] satisfies AssistantEvalCase[]
},
task: async (input) => {
const result = await generateAssistantResponse({
@@ -76,5 +73,5 @@ Eval('Assistant', {
sqlQueries,
}
},
scores: [AssertionScorer],
scores: [toolUsageScorer, sqlSimilarityScorer, criteriaMetScorer, textIncludesScorer],
})
+73 -94
View File
@@ -1,4 +1,4 @@
import { EvalScorer } from 'braintrust'
import { EvalCase, EvalScorer } from 'braintrust'
import { ClosedQA, Sql } from 'autoevals'
type Input = string
@@ -9,106 +9,85 @@ type Output = {
sqlQueries?: string[]
}
export type Assertion<ToolName extends string = string> =
| {
type: 'tools_include'
tool: ToolName
}
| {
type: 'text_includes'
substring: string
}
| {
type: 'llm_criteria_met'
criteria: string
}
| {
type: 'sql_similar'
sql: string
}
export type Expected = {
requiredTools?: string[]
sqlQuery?: string
criteria?: string
textIncludes?: string
}
export function createAssertionScorer<ToolName extends string>(): EvalScorer<
Input,
Output,
Assertion<ToolName>[]
> {
return async ({ input, output, expected: assertions }) => {
const assertionResults: {
status: string
statusDetail?: string
assertion: Assertion<ToolName>
}[] = []
export type AssistantEvalCase = EvalCase<Input, Expected, void>
for (const assertion of assertions) {
let passedTest = false
let statusDetail: string | undefined
export const toolUsageScorer: EvalScorer<Input, Output, Expected> = async ({
output,
expected,
}) => {
if (!expected.requiredTools) return null
try {
switch (assertion.type) {
case 'tools_include': {
passedTest = output.tools.includes(assertion.tool)
break
}
case 'text_includes': {
passedTest = output.text.includes(assertion.substring)
break
}
case 'llm_criteria_met': {
const closedQA = await ClosedQA({
input: 'According to the provided criterion is the submission correct?',
criteria: assertion.criteria,
output: output.text,
})
passedTest = closedQA.score !== null && closedQA.score > 0.5
break
}
case 'sql_similar': {
const sqlQueries = output.sqlQueries || []
if (sqlQueries.length === 0) {
passedTest = false
break
}
// Check if any of the generated SQL queries are similar to the expected SQL
const similarityScores = await Promise.all(
sqlQueries.map(async (generatedSql) => {
const sqlScore = await Sql({
input,
output: generatedSql,
expected: assertion.sql,
})
return sqlScore.score ?? 0
})
)
passedTest = similarityScores.some((score) => score > 0.5)
statusDetail = similarityScores
.map((score, index) => `SQL ${index + 1}: ${score}`)
.join('\n')
break
}
default:
throw new Error(`Unknown assertion type`)
}
} catch (e) {
passedTest = false
}
const presentCount = expected.requiredTools.filter((tool) => output.tools.includes(tool)).length
const totalCount = expected.requiredTools.length
const ratio = totalCount === 0 ? 1 : presentCount / totalCount
assertionResults.push({
status: passedTest ? 'passed' : 'failed',
statusDetail,
assertion,
})
}
return {
name: 'tool_usage',
score: ratio,
}
}
const passedCount = assertionResults.filter((r) => r.status === 'passed').length
const totalCount = assertionResults.length
const ratioPassed = totalCount === 0 ? 1 : passedCount / totalCount
export const sqlSimilarityScorer: EvalScorer<Input, Output, Expected> = async ({
input,
output,
expected,
}) => {
if (!expected.sqlQuery) return null
const sqlQuery = output.sqlQueries?.[0]
if (!sqlQuery) {
return {
name: 'Assertions Score',
score: ratioPassed,
metadata: {
assertionResults,
},
name: 'sql_similarity',
score: 0,
}
}
const sqlResult = await Sql({
input,
output: sqlQuery,
expected: expected.sqlQuery,
})
return {
name: 'sql_similarity',
score: sqlResult.score ?? 0,
}
}
export const criteriaMetScorer: EvalScorer<Input, Output, Expected> = async ({
output,
expected,
}) => {
if (!expected.criteria) return null
const qaResult = await ClosedQA({
input: 'According to the provided criterion is the submission correct?',
output: output.text,
criteria: expected.criteria,
})
return {
name: 'criteria_met',
score: qaResult.score ?? 0,
}
}
export const textIncludesScorer: EvalScorer<Input, Output, Expected> = async ({
output,
expected,
}) => {
if (!expected.textIncludes) return null
const includes = output.text.includes(expected.textIncludes)
return {
name: 'text_includes',
score: includes ? 1 : 0,
}
}