mirror of
https://github.com/supabase/supabase.git
synced 2026-10-11 20:35:07 +03:00
feat: split AssertionScorer into separate scorers
This commit is contained in:
1 parent
4e9d7d9c38
commit
5f601bc9d3
2 files changed
+91
-115
No files matched your search
@@ -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],
|
||||
})
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user