diff --git a/.changeset/durable-batches-evaluate.md b/.changeset/durable-batches-evaluate.md new file mode 100644 index 000000000..a0486161b --- /dev/null +++ b/.changeset/durable-batches-evaluate.md @@ -0,0 +1,5 @@ +--- +"braintrust": minor +--- + +feat: Add batch/durable evals api diff --git a/.gitignore b/.gitignore index bf65941bd..d6531a16e 100644 --- a/.gitignore +++ b/.gitignore @@ -12,6 +12,7 @@ dist !.aiderignore .pnpm-store **/.bt-tmp +**/.braintrust/evals docker-compose.override.yml Dockerfile.local diff --git a/e2e/scenarios/ai-sdk-harness-instrumentation/scenario.test.ts b/e2e/scenarios/ai-sdk-harness-instrumentation/scenario.test.ts index 87054af1f..c934f6dc9 100644 --- a/e2e/scenarios/ai-sdk-harness-instrumentation/scenario.test.ts +++ b/e2e/scenarios/ai-sdk-harness-instrumentation/scenario.test.ts @@ -245,7 +245,16 @@ describe.sequential("HarnessAgent instrumentation variants", () => { ); const harnessSpans = findAllSpans(events, "harness"); expect(harnessSpans).toHaveLength(4); - const bashSpans = findAllSpans(events, "bash"); + // The harness may issue additional bash calls while coordinating a + // suspended turn. Assert only the two commands requested from the + // agent; coordination calls are not part of this contract. + const bashSpans = findAllSpans(events, "bash").filter((span) => { + const input = String(span.input); + return ( + input.includes("printf GENERATE_OK") || + input.includes("printf STREAM_OK") + ); + }); expect(bashSpans).toHaveLength(2); for (const bashSpan of bashSpans) { expect(bashSpan.span.type).toBe("tool"); diff --git a/e2e/scenarios/durable-eval-webhook/scenario.test.ts b/e2e/scenarios/durable-eval-webhook/scenario.test.ts new file mode 100644 index 000000000..ebe696f92 --- /dev/null +++ b/e2e/scenarios/durable-eval-webhook/scenario.test.ts @@ -0,0 +1,67 @@ +import { expect, test } from "vitest"; +import { + prepareScenarioDir, + resolveScenarioDir, + withScenarioHarness, +} from "../../helpers/scenario-harness"; +import { findAllSpans } from "../../helpers/trace-selectors"; + +const scenarioDir = await prepareScenarioDir({ + scenarioDir: resolveScenarioDir(import.meta.url), +}); + +test("durable eval collects webhook sub-batches and logs completed rows", async () => { + await withScenarioHarness( + async ({ events, runScenarioDir, testRunEvents }) => { + await runScenarioDir({ scenarioDir }); + + const evalSpans = findAllSpans(testRunEvents(), "eval"); + const webhookSpans = evalSpans.filter( + (event) => event.metadata?.kind === "webhook", + ); + expect(webhookSpans).toHaveLength(3); + expect(webhookSpans.map((event) => event.output).sort()).toEqual([ + 2, 4, 6, + ]); + expect( + webhookSpans + .map((event) => event.scores) + .sort((left, right) => + JSON.stringify(left).localeCompare(JSON.stringify(right)), + ), + ).toEqual([{ exact: 1 }, { exact: 1 }, { exact: 1 }]); + expect(webhookSpans.map((event) => event.metadata?.durable_eval)).toEqual( + [ + expect.objectContaining({ run_id: expect.any(String) }), + expect.objectContaining({ run_id: expect.any(String) }), + expect.objectContaining({ run_id: expect.any(String) }), + ], + ); + + const taskSpans = findAllSpans(events(), "task"); + expect(taskSpans).toHaveLength(3); + expect(taskSpans.map((event) => event.output).sort()).toEqual([2, 4, 6]); + + const scoreSpans = findAllSpans(events(), "exact"); + expect(scoreSpans).toHaveLength(3); + expect(scoreSpans.map((event) => event.scores)).toEqual([ + { exact: 1 }, + { exact: 1 }, + { exact: 1 }, + ]); + expect(scoreSpans.map((event) => event.metadata?.method)).toEqual([ + "shared-eval-runtime", + "shared-eval-runtime", + "shared-eval-runtime", + ]); + + const classifierSpans = findAllSpans(events(), "quality"); + expect(classifierSpans).toHaveLength(3); + expect(webhookSpans.map((event) => event.row.classifications)).toEqual([ + { quality: [{ id: "pass", label: "Pass" }] }, + { quality: [{ id: "pass", label: "Pass" }] }, + { quality: [{ id: "pass", label: "Pass" }] }, + ]); + }, + ); +}); diff --git a/e2e/scenarios/durable-eval-webhook/scenario.ts b/e2e/scenarios/durable-eval-webhook/scenario.ts new file mode 100644 index 000000000..0020dc8f6 --- /dev/null +++ b/e2e/scenarios/durable-eval-webhook/scenario.ts @@ -0,0 +1,127 @@ +import { BatchTask, DurableEval, type DurableEvalStore } from "braintrust"; +import { + getTestRunId, + runMain, + scopedName, +} from "../../helpers/scenario-runtime"; + +class MemoryStore implements DurableEvalStore { + private readonly values = new Map(); + + async read(key: string) { + return this.values.get(key)?.slice(); + } + + async write(key: string, value: Uint8Array) { + this.values.set(key, value.slice()); + } +} + +async function main() { + const testRunId = getTestRunId(); + const store = new MemoryStore(); + const jobs = new Map>(); + const task = BatchTask< + number, + number, + number, + { testRunId: string; kind: string }, + Record + >({ + workflow(workflow) { + const generated = workflow.batch("generate", { + batchSize: 2, + input: (item) => item.input, + async submit(items) { + const id = `generate-${jobs.size + 1}`; + jobs.set(id, items); + return { id }; + }, + completion: { + mode: "webhook", + externalId: (handle) => handle.id, + }, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: item.input * 2, + })); + }, + }); + return workflow.batch("finalize", { + needs: { generated }, + input: (_item, { generated }) => generated, + batchSize: 2, + async submit(items) { + const id = `finalize-${jobs.size + 1}`; + jobs.set(id, items); + return { id }; + }, + completion: { + mode: "webhook", + externalId: (handle) => handle.id, + }, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: item.input, + })); + }, + }); + }, + }); + const definition = DurableEval( + scopedName("e2e-durable-eval-webhook-project", testRunId), + { + store, + experimentName: scopedName( + "e2e-durable-eval-webhook-experiment", + testRunId, + ), + data: [1, 2, 3].map((input) => ({ + id: `case-${input}`, + input, + expected: input * 2, + metadata: { testRunId, kind: "webhook" }, + })), + task, + scores: [ + function exact({ output, expected }) { + return { + name: "exact", + score: output === expected ? 1 : 0, + metadata: { method: "shared-eval-runtime" }, + }; + }, + ], + classifiers: [ + function quality({ output, expected }) { + return { + name: "quality", + id: output === expected ? "pass" : "fail", + label: output === expected ? "Pass" : "Fail", + }; + }, + ], + }, + ); + + const waiting = await definition.start(); + if (waiting.status !== "waiting" || jobs.size !== 2) { + throw new Error("Durable eval did not pause with two webhook batches"); + } + + let processed = waiting; + const completedJobs = new Set(); + while (completedJobs.size < jobs.size || processed.status !== "completed") { + const externalId = [...jobs.keys()].find((id) => !completedJobs.has(id)); + if (!externalId) throw new Error("Durable eval stopped before completion"); + completedJobs.add(externalId); + processed = await definition.processBatchResult({ + runId: waiting.runId, + externalId, + }); + } +} + +runMain(main); diff --git a/js/README.md b/js/README.md index da6da81a6..13bf76b73 100644 --- a/js/README.md +++ b/js/README.md @@ -44,6 +44,195 @@ async function main() { main().catch(console.error); ``` +## Durable evaluations + +`DurableEval` runs tasks and scorers through asynchronous provider batch APIs. +`batchSize` splits a dataset into provider-sized sub-batches. A small external +store connects submitted jobs with later webhook callbacks; it is required on +the eval definition so every invocation uses the same persistence authority. +No Braintrust backend changes are required. + +Every case needs a stable `id` (or a `caseId` function). + +```typescript +import { BatchTask, DurableEval, type DurableEvalStore } from "braintrust"; + +const store: DurableEvalStore = checkpointStore; +const supportEval = DurableEval("Support bot", { + store, + data: [ + { + id: "password-reset", + input: "How do I reset my password?", + expected: "Open account settings...", + }, + ], + task: BatchTask({ + // Each provider job contains at most 500 eval cases. + batchSize: 500, + + // Submit one sub-batch and return a JSON-serializable provider handle. + async submit(items, context) { + const batch = await provider.submit({ + idempotencyKey: context.batchId, + metadata: { + durableRunId: context.runId, + durableBatchId: context.batchId, + }, + items, + }); + return { id: batch.id }; + }, + + completion: { + // "webhook" waits for processBatchResult(). Use "poll" with a poll() + // callback when the provider does not send completion events. + mode: "webhook", + externalId: (handle) => handle.id, + }, + + async collect(handle) { + return (await provider.results(handle.id)).map((item) => ({ + id: item.id, + output: item.output, + })); + }, + }), + scores: [ + function exact({ output, expected }) { + return output === expected ? 1 : 0; + }, + ], +}); + +const result = await supportEval.start(); +const { runId } = result; +``` + +`start()` initializes the run, submits every ready task sub-batch, and returns. +It never waits in a polling loop. When all task results are available, scoring +begins. `BatchScorer` uses the same `batchSize`, `submit`, `completion`, and +array-returning `collect` contract. + +### Polling + +Polling adapters report the provider's current status through `completion`: + +```typescript +completion: { + mode: "poll", + async poll(handle) { + const batch = await provider.getBatch(handle.id); + if (batch.status === "completed") return { status: "complete" }; + if (batch.status === "failed") { + return { status: "failed", error: batch.error }; + } + return { status: "pending" }; + }, +}, +``` + +Call `poll()` from a cron, queue worker, or another short-lived invocation. It +checks every previously submitted polling batch once, collects completed +results, submits newly ready work, and returns without sleeping: + +```typescript +const result = await supportEval.poll({ + runId, +}); + +if (result.status === "waiting" && result.pending.poll > 0) { + scheduleAnotherPoll(); +} +``` + +`start()`, `poll()`, and `processBatchResult()` return the current eval status. +A waiting result includes the number of submitted batches using each completion +mode: + +```typescript +{ + status: "waiting", + runId, + pending: { poll: 2, webhook: 1 }, +} +``` + +Use `status()` to read the same information without polling providers, +collecting results, or advancing the workflow: + +```typescript +const status = await supportEval.status({ + runId, +}); +``` + +Completed statuses have zero pending batches and include the saved experiment +summary. They can be read repeatedly without logging the eval again. + +### Multi-stage workflows + +The direct `BatchTask({ submit, completion, collect })` form remains the +one-batch shorthand. Use `workflow` when a task or scorer requires multiple +provider batch operations: + +```typescript +task: BatchTask({ + workflow(workflow) { + const draft = workflow.batch("draft", { + input: ({ input }) => ({ prompt: input }), + batchSize: 500, + submit: submitDraftBatch, + completion: draftCompletion, + collect: collectDraftBatch, + }); + + return workflow.batch("revise", { + needs: { draft }, + input: ({ input }, { draft }) => ({ original: input, draft }), + batchSize: 500, + submit: submitRevisionBatch, + completion: revisionCompletion, + collect: collectRevisionBatch, + }); + }, +}), +``` + +Every named batch is a persisted workflow node. `needs` can express sequential +operations, parallel branches, and joins. The returned node supplies the final +task output or scorer result. + +### Webhook processing + +When the provider reports that any task or scorer batch completed, fetch and +store its results through `processBatchResult()`: + +```typescript +app.post("/webhooks/provider", async (request, response) => { + const event = request.body; + const batch = await provider.getBatch(event.batchId); + const runId = batch.metadata.durableRunId; + + const result = await supportEval.processBatchResult({ + // Returned by start() and saved alongside the provider job. + runId, + // The provider's batch ID. DurableEval saved it from submit()'s handle. + externalId: batch.id, + // The SDK-generated ID passed to submit(); include it in provider metadata + // when the webhook cannot provide the external ID used by the handle. + batchId: batch.metadata?.durableBatchId, + }); + + response.status(result.status === "waiting" ? 202 : 200).end(); +}); +``` + +The method accepts either `externalId` or `batchId`. The stored batch locator +identifies the task or scorer workflow node, whose `collect()` results are +stored before the eval advances. Webhook idempotency and provider failure +handling remain application responsibilities for now. + ## Auto-Instrumentation Braintrust can automatically instrument popular AI SDKs (OpenAI, Anthropic, Vercel AI SDK, and others) to log calls without manual wrapper code. diff --git a/js/src/durable-eval.test.ts b/js/src/durable-eval.test.ts new file mode 100644 index 000000000..387fe5850 --- /dev/null +++ b/js/src/durable-eval.test.ts @@ -0,0 +1,448 @@ +import { describe, expect, test, vi } from "vitest"; +import { configureNode } from "./node/config"; +import { + BatchScorer, + BatchTask, + DurableEval, + type DurableBatchScorerItem, + type DurableBatchTaskItem, + type DurableEvalStore, +} from "./durable-eval"; + +configureNode(); + +class MemoryStore implements DurableEvalStore { + private readonly values = new Map(); + + async read(key: string) { + return this.values.get(key)?.slice(); + } + + async write(key: string, value: Uint8Array) { + this.values.set(key, value.slice()); + } +} + +describe("DurableEval", () => { + test("runs ordinary tasks and scorers", async () => { + const task = vi.fn((input: number) => input * 2); + const result = await DurableEval("local", { + store: new MemoryStore(), + data: [ + { id: "one", input: 1, expected: 2 }, + { id: "two", input: 2, expected: 4 }, + ], + task, + scores: [ + function exact({ output, expected }) { + return output === expected ? 1 : 0; + }, + ], + }).start({ noSendLogs: true }); + + expect(result).toMatchObject({ + status: "completed", + summary: { scores: { exact: { score: 1 } } }, + }); + expect(task).toHaveBeenCalledTimes(2); + }); + + test("generates a new run id for every start", async () => { + const durable = DurableEval("generated-runs", { + store: new MemoryStore(), + data: [{ input: 1 }], + task: (input) => input, + scores: [() => 1], + }); + + const first = await durable.start({ noSendLogs: true }); + const second = await durable.start({ noSendLogs: true }); + + expect(first.runId).not.toBe(second.runId); + }); + + test("polls each existing task and scorer sub-batch once", async () => { + const taskJobs = new Map< + string, + DurableBatchTaskItem>[] + >(); + const scoreJobs = new Map< + string, + DurableBatchScorerItem[] + >(); + const taskPoll = vi.fn(async () => ({ status: "complete" as const })); + const scorePoll = vi.fn(async () => ({ status: "complete" as const })); + + const task = BatchTask< + number, + number, + number, + void, + Record, + { id: string } + >({ + batchSize: 2, + async submit(items) { + const id = `task-${taskJobs.size + 1}`; + taskJobs.set(id, items); + return { id }; + }, + completion: { + mode: "poll", + poll: taskPoll, + }, + async collect(handle) { + return (taskJobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: item.input * 2, + })); + }, + }); + const scorer = BatchScorer({ + name: "exact", + batchSize: 2, + async submit(items) { + const id = `score-${scoreJobs.size + 1}`; + scoreJobs.set(id, items); + return { id }; + }, + completion: { + mode: "poll", + poll: scorePoll, + }, + async collect(handle) { + return (scoreJobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + score: item.output === item.expected ? 1 : 0, + })); + }, + }); + + const store = new MemoryStore(); + const durable = DurableEval("polling-batches", { + store, + data: [1, 2, 3].map((input) => ({ + id: `case-${input}`, + input, + expected: input * 2, + })), + task, + scores: [scorer], + }); + const waiting = await durable.start({ noSendLogs: true }); + expect(waiting).toMatchObject({ + status: "waiting", + runId: expect.any(String), + pending: { poll: 2, webhook: 0 }, + }); + const options = { runId: waiting.runId }; + expect(taskJobs.size).toBe(2); + expect(scoreJobs.size).toBe(0); + await expect(durable.status(options)).resolves.toEqual({ + status: "waiting", + runId: waiting.runId, + pending: { poll: 2, webhook: 0 }, + }); + expect(taskJobs.size).toBe(2); + expect(taskPoll).not.toHaveBeenCalled(); + + await expect(durable.poll(options)).resolves.toEqual({ + status: "waiting", + runId: waiting.runId, + pending: { poll: 2, webhook: 0 }, + }); + expect(scoreJobs.size).toBe(2); + expect(taskPoll).toHaveBeenCalledTimes(2); + expect(scorePoll).not.toHaveBeenCalled(); + + const result = await durable.poll(options); + + expect(result).toMatchObject({ + status: "completed", + pending: { poll: 0, webhook: 0 }, + summary: { scores: { exact: { score: 1 } } }, + }); + await expect(durable.status(options)).resolves.toMatchObject({ + status: "completed", + pending: { poll: 0, webhook: 0 }, + summary: { scores: { exact: { score: 1 } } }, + }); + expect(scorePoll).toHaveBeenCalledTimes(2); + expect([...taskJobs.values()].map((items) => items.length)).toEqual([2, 1]); + expect([...scoreJobs.values()].map((items) => items.length)).toEqual([ + 2, 1, + ]); + }); + + test("processes task and scorer webhook batches through one method", async () => { + const store = new MemoryStore(); + const taskJobs = new Map< + string, + DurableBatchTaskItem>[] + >(); + const scoreJobs = new Map< + string, + DurableBatchScorerItem[] + >(); + const task = BatchTask< + number, + number, + number, + void, + Record, + { id: string } + >({ + batchSize: 2, + async submit(items) { + const id = `task-provider-${taskJobs.size + 1}`; + taskJobs.set(id, items); + return { id }; + }, + completion: { + mode: "webhook", + externalId: (handle) => handle.id, + }, + async collect(handle) { + return (taskJobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: item.input * 2, + })); + }, + }); + const scorer = BatchScorer({ + name: "exact", + batchSize: 2, + async submit(items) { + const id = `score-provider-${scoreJobs.size + 1}`; + scoreJobs.set(id, items); + return { id }; + }, + completion: { + mode: "webhook", + externalId: (handle) => handle.id, + }, + async collect(handle) { + return (scoreJobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + score: item.output === item.expected ? 1 : 0, + })); + }, + }); + const durable = DurableEval("webhook-batches", { + store, + data: [1, 2, 3].map((input) => ({ + id: `case-${input}`, + input, + expected: input * 2, + })), + task, + scores: [scorer], + }); + + const waiting = await durable.start({ noSendLogs: true }); + expect(waiting).toMatchObject({ + status: "waiting", + runId: expect.any(String), + pending: { poll: 0, webhook: 2 }, + }); + expect(taskJobs.size).toBe(2); + const runId = waiting.runId; + + const taskIds = [...taskJobs.keys()]; + await expect( + durable.processBatchResult({ + runId: "missing-run", + externalId: taskIds[0], + }), + ).rejects.toThrow("DurableEval run missing-run is missing"); + for (const [index, externalId] of taskIds.entries()) { + const result = await durable.processBatchResult({ runId, externalId }); + expect(result).toMatchObject({ + status: "waiting", + pending: { poll: 0, webhook: index === 0 ? 1 : 2 }, + }); + } + expect(scoreJobs.size).toBe(2); + + const scoreIds = [...scoreJobs.keys()]; + let result; + for (const externalId of scoreIds) { + result = await durable.processBatchResult({ runId, externalId }); + } + expect(result).toMatchObject({ + status: "completed", + pending: { poll: 0, webhook: 0 }, + summary: { scores: { exact: { score: 1 } } }, + }); + }); + + test("runs multi-stage task and scorer workflows", async () => { + const jobs = new Map>(); + const submit = async ( + prefix: string, + items: Array<{ id: string; input: unknown }>, + ) => { + const id = `${prefix}-${jobs.size + 1}`; + jobs.set(id, items); + return { id }; + }; + const completion = { + mode: "poll" as const, + async poll() { + return { status: "complete" as const }; + }, + }; + + const task = BatchTask>( + { + workflow(w) { + const doubled = w.batch("double", { + batchSize: 2, + input: (item) => item.input, + submit: (items) => submit("double", items), + completion, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: (item.input as number) * 2, + })); + }, + }); + const incremented = w.batch("increment", { + needs: { doubled }, + input: (_item, outputs) => outputs.doubled, + submit: (items) => submit("increment", items), + completion, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: (item.input as number) + 1, + })); + }, + }); + const decremented = w.batch("decrement", { + needs: { doubled }, + input: (_item, outputs) => outputs.doubled, + submit: (items) => submit("decrement", items), + completion, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: (item.input as number) - 1, + })); + }, + }); + return w.batch("combine", { + needs: { incremented, decremented }, + input: (_item, outputs) => outputs, + submit: (items) => submit("combine", items), + completion, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => { + const value = item.input as { + incremented: number; + decremented: number; + }; + return { + id: item.id, + output: (value.incremented + value.decremented) / 2, + }; + }); + }, + }); + }, + }, + ); + const scorer = BatchScorer({ + name: "exact", + workflow(w) { + const comparison = w.batch("compare", { + input: (item) => ({ + output: item.output, + expected: item.expected, + }), + submit: (items) => submit("compare", items), + completion, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => { + const value = item.input as { + output: number; + expected: number; + }; + return { + id: item.id, + output: value.output === value.expected, + }; + }); + }, + }); + return w.batch("score", { + needs: { comparison }, + input: (_item, outputs) => outputs.comparison, + submit: (items) => submit("score", items), + completion, + async collect(handle) { + return (jobs.get(handle.id) ?? []).map((item) => ({ + id: item.id, + output: item.input ? 1 : 0, + })); + }, + }); + }, + }); + const store = new MemoryStore(); + const durable = DurableEval("workflow", { + store, + data: [1, 2, 3].map((input) => ({ + id: `case-${input}`, + input, + expected: input * 2, + })), + task, + scores: [scorer], + }); + const waiting = await durable.start({ noSendLogs: true }); + expect(waiting).toMatchObject({ + status: "waiting", + }); + const options = { runId: waiting.runId }; + await expect(durable.poll(options)).resolves.toMatchObject({ + status: "waiting", + }); + await expect(durable.poll(options)).resolves.toMatchObject({ + status: "waiting", + }); + await expect(durable.poll(options)).resolves.toMatchObject({ + status: "waiting", + }); + await expect(durable.poll(options)).resolves.toMatchObject({ + status: "waiting", + }); + await expect(durable.poll(options)).resolves.toMatchObject({ + status: "completed", + summary: { scores: { exact: { score: 1 } } }, + }); + }); + + test("requires stable case ids", async () => { + await expect( + DurableEval("missing-ids", { + store: new MemoryStore(), + data: [{ input: "hello" }], + task: BatchTask({ + async submit() { + return { id: "unused" }; + }, + completion: { + mode: "webhook", + externalId: (handle) => handle.id, + }, + async collect() { + return []; + }, + }), + scores: [], + }).start({ noSendLogs: true }), + ).rejects.toThrow("requires id, upsert_id, or caseId"); + }); +}); diff --git a/js/src/durable-eval.ts b/js/src/durable-eval.ts new file mode 100644 index 000000000..489698658 --- /dev/null +++ b/js/src/durable-eval.ts @@ -0,0 +1,1752 @@ +import { makeScorerPropagatedEvent, SpanTypeAttribute } from "../util/index"; +import { + type EvalParameters, + type InferParameters, + validateParameters, +} from "./eval-parameters"; +import { + _internalInitEvaluatorExperiment, + _internalPrepareEvaluatorClassification, + _internalPrepareEvaluatorScore, + _internalResolveEvaluatorData, + _internalRunEvaluatorTask, + buildLocalSummary as buildEvaluatorLocalSummary, + callEvaluatorData, + classifierName, + type EvalClassifier, + type Evaluator, + type EvaluatorDef, + type EvalResult, + type EvalScorer, + type EvalScorerArgs, + type EvalTask, + type OneOrMoreScores, + runEvaluator, +} from "./framework"; +import iso from "./isomorph"; +import { + type BaseMetadata, + type DefaultMetadataType, + type EvalCase, + type Experiment, + type ExperimentSummary, + NOOP_SPAN, + type Span, + _internalResumeSpan, + _internalStartSpanWithInitialMerge, + logError as logSpanError, + newId, +} from "./logger"; + +const encoder = new TextEncoder(); +const decoder = new TextDecoder(); +const BATCH_TASK_KIND = "braintrust.durable.batch-task"; +const BATCH_SCORER_KIND = "braintrust.durable.batch-scorer"; +const CHECKPOINT_VERSION = 2; +const DEFAULT_BATCH_SIZE = 1_000; + +type JsonPrimitive = string | number | boolean | null; +export type JsonValue = + | JsonPrimitive + | JsonValue[] + | { [key: string]: JsonValue }; + +/** + * Minimal persistence used to reconnect provider webhooks with submitted + * batches. DurableEval does not require any Braintrust backend changes. + */ +export interface DurableEvalStore { + read(key: string): Promise; + write(key: string, value: Uint8Array): Promise; +} + +export interface DurableBatchContext { + runId: string; + batchId: string; +} + +export type DurableBatchPoll = + | { status: "pending" } + | { status: "complete" } + | { status: "failed"; error: unknown }; + +export type DurableBatchCompletion = + | { + mode: "poll"; + poll( + handle: Handle, + context: DurableBatchContext, + ): Promise; + } + | { + mode: "webhook"; + externalId(handle: Handle, context: DurableBatchContext): string; + }; + +export interface DurableBatchTaskItem< + Input, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +> { + id: string; + input: Input; + expected: Expected; + metadata: Metadata; + tags: string[] | undefined; + parameters: InferParameters; + trialIndex: number; +} + +export type DurableBatchTaskResult = + | { + id: string; + output: Output; + metadata?: Metadata; + tags?: string[]; + } + | { id: string; error: unknown }; + +export type DurableBatchScorerItem< + Input, + Output, + Expected, + Metadata extends BaseMetadata, +> = EvalScorerArgs & { + id: string; + trialIndex: number; +}; + +export type DurableBatchScorerResult = + | { id: string; score: OneOrMoreScores } + | { id: string; error: unknown }; + +export interface DurableBatchProcessor { + batchSize?: number; + submit(items: Item[], context: DurableBatchContext): Promise; + completion: DurableBatchCompletion; + collect(handle: Handle, context: DurableBatchContext): Promise; +} + +export interface DurableWorkflowBatchItem { + id: string; + input: Input; +} + +export type DurableWorkflowBatchResult = + | { + id: string; + output: Output; + metadata?: Metadata; + tags?: string[]; + } + | { id: string; error: unknown }; + +const WORKFLOW_NODE_OUTPUT: unique symbol = Symbol("DurableWorkflowNodeOutput"); + +export interface DurableWorkflowNode { + readonly [WORKFLOW_NODE_OUTPUT]: Output; +} + +type DurableWorkflowNodeMap = Record>; + +type DurableWorkflowNodeOutputs = { + [Name in keyof Nodes]: Nodes[Name] extends DurableWorkflowNode + ? Output + : never; +}; + +export interface DurableWorkflowBuilder< + RootItem, + Metadata extends BaseMetadata, +> { + batch< + Output, + Needs extends DurableWorkflowNodeMap = Record, + Input = RootItem, + Handle extends JsonValue = JsonValue, + >( + name: string, + processor: DurableBatchProcessor< + DurableWorkflowBatchItem, + DurableWorkflowBatchResult, + Handle + > & { + needs?: Needs; + input?: ( + item: RootItem, + outputs: DurableWorkflowNodeOutputs, + ) => Input; + }, + ): DurableWorkflowNode; +} + +type DurableWorkflowNodeDefinition = { + name: string; + needs: Record; + item: ( + rootItem: unknown, + outputs: Record, + id: string, + ) => unknown; + processor: DurableBatchProcessor; + result: (result: any) => unknown; +}; + +type DurableWorkflowDefinition = { + nodes: DurableWorkflowNodeDefinition[]; + outputNode: string; +}; + +export interface DurableBatchTask< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, + Handle extends JsonValue, +> { + readonly kind: typeof BATCH_TASK_KIND; + readonly processor?: DurableBatchProcessor< + DurableBatchTaskItem, + DurableBatchTaskResult, + Handle + >; + readonly workflow?: DurableWorkflowDefinition; +} + +export interface DurableBatchScorer< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Handle extends JsonValue, +> { + readonly kind: typeof BATCH_SCORER_KIND; + name: string; + readonly processor?: DurableBatchProcessor< + DurableBatchScorerItem, + DurableBatchScorerResult, + Handle + >; + readonly workflow?: DurableWorkflowDefinition; +} + +export function BatchTask< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Parameters extends EvalParameters = EvalParameters, + Handle extends JsonValue = JsonValue, +>( + processor: DurableBatchProcessor< + DurableBatchTaskItem, + DurableBatchTaskResult, + Handle + >, +): DurableBatchTask; +export function BatchTask< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Parameters extends EvalParameters = EvalParameters, +>(config: { + workflow( + builder: DurableWorkflowBuilder< + DurableBatchTaskItem, + Metadata + >, + ): DurableWorkflowNode; +}): DurableBatchTask; +export function BatchTask( + config: + | DurableBatchProcessor + | { + workflow( + builder: DurableWorkflowBuilder, + ): DurableWorkflowNode; + }, +): DurableBatchTask< + unknown, + unknown, + unknown, + BaseMetadata, + EvalParameters, + JsonValue +> { + return { + kind: BATCH_TASK_KIND, + ...("workflow" in config + ? { workflow: buildWorkflow(config.workflow) } + : { processor: config }), + } as DurableBatchTask< + unknown, + unknown, + unknown, + BaseMetadata, + EvalParameters, + JsonValue + >; +} + +export function BatchScorer< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Handle extends JsonValue = JsonValue, +>( + processor: DurableBatchProcessor< + DurableBatchScorerItem, + DurableBatchScorerResult, + Handle + > & { name: string }, +): DurableBatchScorer; +export function BatchScorer< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, +>(config: { + name: string; + workflow( + builder: DurableWorkflowBuilder< + DurableBatchScorerItem, + Metadata + >, + ): DurableWorkflowNode; +}): DurableBatchScorer; +export function BatchScorer( + config: + | (DurableBatchProcessor & { name: string }) + | { + name: string; + workflow( + builder: DurableWorkflowBuilder, + ): DurableWorkflowNode; + }, +): DurableBatchScorer { + return { + kind: BATCH_SCORER_KIND, + name: config.name, + ...("workflow" in config + ? { workflow: buildWorkflow(config.workflow) } + : { processor: config }), + } as DurableBatchScorer; +} + +function buildWorkflow( + define: ( + builder: DurableWorkflowBuilder, + ) => DurableWorkflowNode, +): DurableWorkflowDefinition { + const nodes: DurableWorkflowNodeDefinition[] = []; + const nodeNames = new Map, string>(); + const builder: DurableWorkflowBuilder = { + batch(name, config) { + if (!name.trim()) throw new Error("Workflow batch names cannot be empty"); + if (nodes.some((node) => node.name === name)) { + throw new Error(`Duplicate workflow batch name: ${name}`); + } + const needs = Object.fromEntries( + Object.entries(config.needs ?? {}).map(([alias, dependency]) => { + const dependencyName = nodeNames.get(dependency); + if (!dependencyName) { + throw new Error( + `Workflow batch ${name} depends on an unknown or later batch`, + ); + } + return [alias, dependencyName]; + }), + ); + const { input, needs: _needs, ...processor } = config; + const handle = {} as DurableWorkflowNode; + nodeNames.set(handle, name); + nodes.push({ + name, + needs, + item(rootItem, outputs, id) { + return { + id, + input: input + ? input(rootItem as RootItem, outputs as never) + : rootItem, + }; + }, + processor: processor as DurableBatchProcessor, + result: (result) => result.output, + }); + return handle as never; + }, + }; + const output = define(builder); + const outputNode = nodeNames.get(output); + if (!outputNode) { + throw new Error("A batch workflow must return one of its batch nodes"); + } + const reachable = new Set(); + const visit = (name: string) => { + if (reachable.has(name)) return; + reachable.add(name); + const node = nodes.find((candidate) => candidate.name === name)!; + Object.values(node.needs).forEach(visit); + }; + visit(outputNode); + const unused = nodes.filter((node) => !reachable.has(node.name)); + if (unused.length > 0) { + throw new Error( + `Batch workflow contains nodes that do not contribute to its output: ${unused + .map((node) => node.name) + .join(", ")}`, + ); + } + return { nodes, outputNode }; +} + +export type DurableEvaluator< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Parameters extends EvalParameters = EvalParameters, +> = Omit< + Evaluator, + "task" | "scores" | "timeout" | "signal" | "maxConcurrency" | "update" +> & { + store: DurableEvalStore; + caseId?: ( + datum: EvalCase, + ) => string | Promise; + task: + | EvalTask + | DurableBatchTask< + Input, + Output, + Expected, + Metadata, + Parameters, + JsonValue + >; + scores?: Array< + | EvalScorer + | DurableBatchScorer + >; +}; + +export interface DurableEvalStartOptions< + Parameters extends EvalParameters = EvalParameters, +> { + parameters?: InferParameters; + noSendLogs?: boolean; +} + +export type DurableBatchResult = { + runId: string; + batchId?: string; + externalId?: string; +}; + +export type DurableEvalResult = + | { + status: "waiting"; + runId: string; + pending: { + poll: number; + webhook: number; + }; + } + | { + status: "completed"; + runId: string; + pending: { + poll: 0; + webhook: 0; + }; + summary: ExperimentSummary; + }; + +export interface DurableEvalDefinition< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Parameters extends EvalParameters = EvalParameters, +> { + readonly projectName: string; + readonly evalName: string; + readonly evaluator: DurableEvaluator< + Input, + Output, + Expected, + Metadata, + Parameters + >; + start( + options?: DurableEvalStartOptions, + ): Promise; + status(options: { runId: string }): Promise; + poll(options: { runId: string }): Promise; + processBatchResult(result: DurableBatchResult): Promise; +} + +type DurableCaseRecord = { + id: string; + caseId: string; + trialIndex: number; + datum: JsonValue; + metadata: JsonValue; + tags?: string[]; + taskComplete: boolean; + taskLogged: boolean; + output?: JsonValue; + rootSpan?: string; + taskNodeOutputs: Record; + scores: Record; + loggedScores: Record; + scoreNodeOutputs: Record>; + classifications: Record; + loggedClassifications: Record; +}; + +type DurableBatchRecord = { + id: string; + kind: "task" | "score"; + scorerName?: string; + nodeName: string; + itemIds: string[]; + handle: JsonValue; + externalId?: string; + status: "submitted" | "complete"; +}; + +type DurableRunState = { + schemaVersion: number; + runId: string; + projectName: string; + evalName: string; + experimentName: string; + noSendLogs: boolean; + parameters: JsonValue; + status: "running" | "completed"; + summary?: ExperimentSummary; + cases: DurableCaseRecord[]; + batches: DurableBatchRecord[]; +}; + +class DurableEvalDefinitionImpl< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Parameters extends EvalParameters = EvalParameters, +> implements DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters +> { + readonly evalName: string; + + constructor( + readonly projectName: string, + readonly evaluator: DurableEvaluator< + Input, + Output, + Expected, + Metadata, + Parameters + >, + ) { + this.evalName = evaluator.experimentName ?? projectName; + } + + start( + options: DurableEvalStartOptions = {}, + ): Promise { + return startDurableEval(this, options); + } + + status(options: { runId: string }): Promise { + return getDurableEvalStatus(this, options); + } + + poll(options: { runId: string }): Promise { + return pollDurableEval(this, options); + } + + processBatchResult(result: DurableBatchResult): Promise { + return processDurableBatchResult(this, result); + } +} + +export function DurableEval< + Input, + Output, + Expected = void, + Metadata extends BaseMetadata = DefaultMetadataType, + Parameters extends EvalParameters = EvalParameters, +>( + projectName: string, + evaluator: DurableEvaluator, +): DurableEvalDefinition { + return new DurableEvalDefinitionImpl(projectName, evaluator); +} + +async function startDurableEval< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + options: DurableEvalStartOptions, +): Promise { + const store = definition.evaluator.store; + const runId = newId(); + const key = runKey(definition.projectName, definition.evalName, runId); + const { data } = callEvaluatorData(definition.evaluator.data); + const parameters = await validateParameters( + options.parameters ?? {}, + definition.evaluator.parameters, + ); + const experimentName = + definition.evaluator.experimentName ?? `${definition.evalName}-${runId}`; + const experiment = await _internalInitEvaluatorExperiment( + definition.projectName, + { ...definition.evaluator, data } as unknown as Evaluator< + Input, + Output, + Expected, + Metadata, + Parameters + >, + data, + { + disabled: options.noSendLogs ?? false, + experimentName, + update: true, + }, + ); + if ( + !isBatchTask(definition.evaluator.task) && + !(definition.evaluator.scores ?? []).some(isBatchScorer) + ) { + const result = await runEvaluator( + experiment, + { + ...definition.evaluator, + projectName: definition.projectName, + evalName: definition.evalName, + data, + } as unknown as EvaluatorDef< + Input, + Output, + Expected, + Metadata, + Parameters + >, + { + start: () => undefined, + stop: () => undefined, + increment: () => undefined, + }, + [], + undefined, + parameters, + true, + true, + ); + const state: DurableRunState = { + schemaVersion: CHECKPOINT_VERSION, + runId, + projectName: definition.projectName, + evalName: definition.evalName, + experimentName, + noSendLogs: options.noSendLogs ?? false, + parameters: assertJsonValue(parameters, "eval parameters"), + status: "completed", + summary: result.summary, + cases: [], + batches: [], + }; + await experiment?.flush(); + await writeJson(store, key, state); + return currentStatus(definition, state); + } + const state: DurableRunState = { + schemaVersion: CHECKPOINT_VERSION, + runId, + projectName: definition.projectName, + evalName: definition.evalName, + experimentName, + noSendLogs: options.noSendLogs ?? false, + parameters: assertJsonValue(parameters, "eval parameters"), + status: "running", + cases: await materializeCases(definition, data, experiment), + batches: [], + }; + await writeJson(store, key, state); + return advanceDurableEval(definition, state, store, key, experiment); +} + +async function getDurableEvalStatus< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + options: { runId: string }, +): Promise { + const store = definition.evaluator.store; + const state = await readJson( + store, + runKey(definition.projectName, definition.evalName, options.runId), + ); + if (!state) throw new Error(`DurableEval run ${options.runId} is missing`); + return currentStatus(definition, state); +} + +async function processDurableBatchResult< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + result: DurableBatchResult, +): Promise { + if (!result.batchId && !result.externalId) { + throw new Error("Batch results require batchId or externalId"); + } + const store = definition.evaluator.store; + const key = runKey(definition.projectName, definition.evalName, result.runId); + const state = await readJson(store, key); + if (!state) throw new Error(`DurableEval run ${result.runId} is missing`); + const byBatch = result.batchId + ? state.batches.find((candidate) => candidate.id === result.batchId) + : undefined; + const byExternal = result.externalId + ? state.batches.find( + (candidate) => candidate.externalId === result.externalId, + ) + : undefined; + if (byBatch && byExternal && byBatch.id !== byExternal.id) { + throw new Error("batchId and externalId identify different batches"); + } + const batch = byBatch ?? byExternal; + if (!batch) throw new Error("No submitted batch matches this result"); + if (batch.status !== "complete") { + await collectBatch(definition, state, batch); + batch.status = "complete"; + await writeJson(store, key, state); + } + return advanceDurableEval(definition, state, store, key); +} + +async function pollDurableEval< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + options: { runId: string }, +): Promise { + const store = definition.evaluator.store; + const key = runKey( + definition.projectName, + definition.evalName, + options.runId, + ); + const state = await readJson(store, key); + if (!state) throw new Error(`DurableEval run ${options.runId} is missing`); + + const batches = state.batches.filter((batch) => { + if (batch.status === "complete") return false; + return processorForBatch(definition, batch).completion.mode === "poll"; + }); + const results = await Promise.all( + batches.map(async (batch) => ({ + batch, + result: await ( + processorForBatch(definition, batch).completion as Extract< + DurableBatchCompletion, + { mode: "poll" } + > + ).poll(batch.handle, { + runId: state.runId, + batchId: batch.id, + }), + })), + ); + for (const { batch, result } of results) { + if (result.status === "failed") throw asError(result.error); + if (result.status !== "complete") continue; + await collectBatch(definition, state, batch); + batch.status = "complete"; + await writeJson(store, key, state); + } + return advanceDurableEval(definition, state, store, key); +} + +async function openDurableExperiment( + definition: DurableEvalDefinition, + state: DurableRunState, +) { + const data: EvalCase[] = []; + return await _internalInitEvaluatorExperiment( + definition.projectName, + { ...definition.evaluator, data } as unknown as Evaluator< + unknown, + unknown, + unknown, + BaseMetadata, + EvalParameters + >, + data, + { + disabled: state.noSendLogs, + experimentName: state.experimentName, + update: true, + }, + ); +} + +async function advanceDurableEval< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + state: DurableRunState, + store: DurableEvalStore, + key: string, + existingExperiment?: Experiment | null, +): Promise { + if (state.status === "completed") return currentStatus(definition, state); + const experiment = + existingExperiment === undefined + ? await openDurableExperiment(definition, state) + : existingExperiment; + await runTaskStage(definition, state, store, key, experiment); + await logCompletedTasks(definition, state, store, key, experiment); + if (state.cases.some((record) => !record.taskComplete)) { + return currentStatus(definition, state); + } + + await runScoreStages(definition, state, store, key, experiment); + const scorerNames = resolveScorers(definition.evaluator.scores ?? []).map( + ({ name }) => name, + ); + if ( + state.cases.some((record) => + scorerNames.some((name) => !(name in record.scores)), + ) + ) { + return currentStatus(definition, state); + } + + state.summary = await finishExperiment(definition, state, experiment); + state.status = "completed"; + await writeJson(store, key, state); + return currentStatus(definition, state); +} + +function currentStatus( + definition: DurableEvalDefinition, + state: DurableRunState, +): DurableEvalResult { + if (state.status === "completed") { + if (!state.summary) { + throw new Error(`DurableEval run ${state.runId} has no saved summary`); + } + return { + status: "completed", + runId: state.runId, + pending: { poll: 0, webhook: 0 }, + summary: state.summary, + }; + } + const pending = { poll: 0, webhook: 0 }; + for (const batch of state.batches) { + if (batch.status === "complete") continue; + pending[processorForBatch(definition, batch).completion.mode]++; + } + return { status: "waiting", runId: state.runId, pending }; +} + +async function startCaseRoot( + definition: DurableEvalDefinition, + state: DurableRunState, + record: DurableCaseRecord, + experiment: Experiment | null, +): Promise { + if (!experiment) return NOOP_SPAN; + const datum = record.datum as EvalCase; + return _internalStartSpanWithInitialMerge({ + ...(definition.evaluator.state + ? { state: definition.evaluator.state } + : {}), + parent: await experiment.export(), + name: "eval", + spanId: deterministicId(`${state.runId}:${record.id}:span`), + spanAttributes: { type: SpanTypeAttribute.EVAL }, + event: { + id: deterministicId(`${state.runId}:${record.id}:row`), + input: datum.input, + expected: "expected" in datum ? datum.expected : undefined, + tags: datum.tags, + origin: datum.origin, + }, + }); +} + +async function logTaskResult( + definition: DurableEvalDefinition, + state: DurableRunState, + record: DurableCaseRecord, + experiment: Experiment | null, + task?: EvalTask, +) { + const datum = record.datum as EvalCase; + const root = await startCaseRoot(definition, state, record, experiment); + try { + if (task) { + const result = await root.traced( + (span) => + _internalRunEvaluatorTask( + task, + datum, + record.trialIndex, + state.parameters as Record, + span, + ), + { + name: "task", + spanId: deterministicId(`${state.runId}:${record.id}:task`), + spanAttributes: { type: SpanTypeAttribute.TASK }, + event: { input: datum.input }, + }, + ); + record.output = assertJsonValue( + result.output, + `task output for ${record.caseId}`, + ); + record.metadata = assertJsonValue(result.metadata, "task metadata"); + record.tags = result.tags; + record.taskComplete = true; + } else { + await root.traced((span) => span.log({ output: record.output }), { + name: "task", + spanId: deterministicId(`${state.runId}:${record.id}:task`), + spanAttributes: { type: SpanTypeAttribute.TASK }, + event: { input: datum.input }, + }); + } + root.log({ + output: record.output, + expected: "expected" in datum ? datum.expected : undefined, + metadata: { + ...(record.metadata as Record), + durable_eval: { + run_id: state.runId, + case_id: record.caseId, + trial_index: record.trialIndex, + }, + }, + tags: record.tags, + }); + record.rootSpan = await root.export(); + record.taskLogged = true; + } catch (error) { + logSpanError(root, error); + throw error; + } finally { + root.end(); + await experiment?.flush(); + } +} + +async function logCompletedTasks( + definition: DurableEvalDefinition, + state: DurableRunState, + store: DurableEvalStore, + key: string, + experiment: Experiment | null, +) { + for (const record of state.cases) { + if (!record.taskComplete || record.taskLogged) continue; + await logTaskResult(definition, state, record, experiment); + await writeJson(store, key, state); + } +} + +async function runTaskStage< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + state: DurableRunState, + store: DurableEvalStore, + key: string, + experiment: Experiment | null, +) { + if (isBatchTask(definition.evaluator.task)) { + await ensureWorkflowBatches(definition, state, store, key, "task"); + return; + } + + const task = definition.evaluator.task as EvalTask< + Input, + Output, + Expected, + Metadata, + Parameters + >; + for (const record of state.cases) { + if (record.taskComplete) continue; + await logTaskResult(definition, state, record, experiment, task); + await writeJson(store, key, state); + } +} + +async function runScoreStages< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + state: DurableRunState, + store: DurableEvalStore, + key: string, + experiment: Experiment | null, +) { + const scorers = resolveScorers(definition.evaluator.scores ?? []); + for (const { name, scorer } of scorers) { + if (isBatchScorer(scorer)) { + for (const record of state.cases) { + if (name in record.scores && !record.loggedScores[name]) { + await evaluateAndLogScore( + definition, + state, + record, + name, + experiment, + ); + await writeJson(store, key, state); + } + } + await ensureWorkflowBatches(definition, state, store, key, "score", name); + continue; + } + for (const record of state.cases) { + if (record.loggedScores[name]) continue; + await evaluateAndLogScore( + definition, + state, + record, + name, + experiment, + scorer, + ); + await writeJson(store, key, state); + } + } + + for (const [index, classifier] of ( + definition.evaluator.classifiers ?? [] + ).entries()) { + const name = classifierName(classifier, index); + for (const record of state.cases) { + if (record.loggedClassifications[name]) continue; + await evaluateAndLogClassification( + definition, + state, + record, + name, + classifier, + experiment, + ); + await writeJson(store, key, state); + } + } +} + +function scorerArgs(record: DurableCaseRecord) { + const datum = record.datum as EvalCase; + return { + ...datum, + metadata: record.metadata, + output: record.output, + } as EvalScorerArgs; +} + +function resumeCaseRoot( + definition: DurableEvalDefinition, + record: DurableCaseRecord, + experiment: Experiment | null, +) { + if (!experiment) return NOOP_SPAN; + if (!record.rootSpan) { + throw new Error(`Durable eval case ${record.caseId} has no root span`); + } + return _internalResumeSpan({ + exported: record.rootSpan, + state: definition.evaluator.state, + }); +} + +async function evaluateAndLogScore( + definition: DurableEvalDefinition, + state: DurableRunState, + record: DurableCaseRecord, + name: string, + experiment: Experiment | null, + scorer?: EvalScorer, +) { + const root = resumeCaseRoot(definition, record, experiment); + try { + const rootExport = await root.export(); + const prepared = await root.traced( + async (span) => { + const value = scorer + ? await scorer(scorerArgs(record)) + : (record.scores[name] as OneOrMoreScores); + if (scorer) { + record.scores[name] = assertJsonValue(value, `scorer ${name} output`); + } + const result = _internalPrepareEvaluatorScore(value, name); + if (result.results !== null) { + span.log({ + output: result.output, + metadata: result.metadata, + scores: result.scores, + }); + } + return result; + }, + { + name, + spanId: deterministicId(`${state.runId}:${record.id}:score:${name}`), + spanAttributes: { + type: SpanTypeAttribute.SCORE, + purpose: "scorer", + }, + propagatedEvent: makeScorerPropagatedEvent(rootExport || undefined), + event: { input: scorerArgs(record) }, + }, + ); + if (prepared.scores) root.log({ scores: prepared.scores }); + record.loggedScores[name] = true; + } catch (error) { + logSpanError(root, error); + throw error; + } finally { + root.end(); + await experiment?.flush(); + } +} + +async function evaluateAndLogClassification( + definition: DurableEvalDefinition, + state: DurableRunState, + record: DurableCaseRecord, + name: string, + classifier: EvalClassifier, + experiment: Experiment | null, +) { + const root = resumeCaseRoot(definition, record, experiment); + try { + const rootExport = await root.export(); + const prepared = await root.traced( + async (span) => { + const value = await classifier(scorerArgs(record)); + record.classifications[name] = assertJsonValue( + value, + `classifier ${name} output`, + ); + const result = _internalPrepareEvaluatorClassification(value, name); + if (result.results !== null) { + span.log({ output: result.output, metadata: result.metadata }); + } + return result; + }, + { + name, + spanId: deterministicId( + `${state.runId}:${record.id}:classification:${name}`, + ), + spanAttributes: { + type: SpanTypeAttribute.CLASSIFIER, + purpose: "scorer", + }, + propagatedEvent: makeScorerPropagatedEvent(rootExport || undefined), + event: { input: scorerArgs(record) }, + }, + ); + if (prepared.classifications) { + root.log({ classifications: prepared.classifications }); + } + record.loggedClassifications[name] = true; + } catch (error) { + logSpanError(root, error); + throw error; + } finally { + root.end(); + await experiment?.flush(); + } +} + +async function ensureWorkflowBatches( + definition: DurableEvalDefinition, + state: DurableRunState, + store: DurableEvalStore, + key: string, + kind: "task" | "score", + scorerName?: string, +) { + const workflow = workflowForStage(definition, kind, scorerName); + for (const node of workflow.nodes) { + const batchSize = node.processor.batchSize ?? DEFAULT_BATCH_SIZE; + if (!Number.isInteger(batchSize) || batchSize < 1) { + throw new Error( + `Invalid batchSize for ${scorerName ? `${scorerName}.${node.name}` : node.name}`, + ); + } + const assigned = new Set( + state.batches + .filter( + (batch) => + batch.kind === kind && + batch.scorerName === scorerName && + batch.nodeName === node.name, + ) + .flatMap((batch) => batch.itemIds), + ); + const eligible = state.cases.filter((record) => { + const outputs = nodeOutputsFor(record, kind, scorerName); + return ( + !assigned.has(record.id) && + !(node.name in outputs) && + Object.values(node.needs).every((dependency) => dependency in outputs) + ); + }); + for (let offset = 0; offset < eligible.length; offset += batchSize) { + const records = eligible.slice(offset, offset + batchSize); + const batchId = newId(); + const context = { runId: state.runId, batchId }; + const items = records.map((record) => + itemForNode(state.parameters, record, kind, scorerName, node), + ); + const handle = assertJsonValue( + await node.processor.submit(items, context), + `handle for batch ${batchId}`, + ); + const externalId = + node.processor.completion.mode === "webhook" + ? node.processor.completion.externalId(handle, context) + : undefined; + if (externalId !== undefined && !externalId.trim()) { + throw new Error(`Batch ${batchId} produced an empty externalId`); + } + const batch: DurableBatchRecord = { + id: batchId, + kind, + scorerName, + nodeName: node.name, + itemIds: records.map((record) => record.id), + handle, + externalId, + status: "submitted", + }; + state.batches.push(batch); + await writeJson(store, key, state); + } + } +} + +async function collectBatch( + definition: DurableEvalDefinition, + state: DurableRunState, + batch: DurableBatchRecord, +) { + const workflow = workflowForStage(definition, batch.kind, batch.scorerName); + const node = workflow.nodes.find( + (candidate) => candidate.name === batch.nodeName, + ); + if (!node) + throw new Error( + `Definition no longer contains batch node ${batch.nodeName}`, + ); + const processor = node.processor; + const context = { runId: state.runId, batchId: batch.id }; + const results = await processor.collect(batch.handle, context); + if (!Array.isArray(results)) { + throw new Error(`collect for batch ${batch.id} must return an array`); + } + const expectedIds = new Set(batch.itemIds); + const seen = new Set(); + for (const result of results) { + const id = resultItemId(result); + if (!expectedIds.has(id)) { + throw new Error(`Batch ${batch.id} returned unknown item ${id}`); + } + if (seen.has(id)) { + throw new Error(`Batch ${batch.id} returned item ${id} more than once`); + } + seen.add(id); + if ("error" in result) throw asError(result.error); + const record = state.cases.find((candidate) => candidate.id === id)!; + const output = assertJsonValue( + node.result(result), + `output for ${batch.nodeName} item ${id}`, + ); + nodeOutputsFor(record, batch.kind, batch.scorerName)[batch.nodeName] = + output; + if (batch.nodeName !== workflow.outputNode) continue; + if (batch.kind === "task") { + record.output = output; + if ("metadata" in result && result.metadata !== undefined) { + record.metadata = assertJsonValue( + result.metadata, + `metadata for ${id}`, + ); + } + if ("tags" in result && result.tags !== undefined) + record.tags = result.tags; + record.taskComplete = true; + } else { + record.scores[batch.scorerName!] = assertJsonValue( + output, + `score for ${id}`, + ); + } + } + const missing = batch.itemIds.filter((id) => !seen.has(id)); + if (missing.length > 0) { + throw new Error( + `Batch ${batch.id} did not return results for: ${missing.join(", ")}`, + ); + } +} + +function processorForBatch( + definition: DurableEvalDefinition, + batch: DurableBatchRecord, +): DurableBatchProcessor { + const node = workflowForStage( + definition, + batch.kind, + batch.scorerName, + ).nodes.find((candidate) => candidate.name === batch.nodeName); + if (!node) + throw new Error( + `Definition no longer contains batch node ${batch.nodeName}`, + ); + return node.processor; +} + +function workflowForStage( + definition: DurableEvalDefinition, + kind: "task" | "score", + scorerName?: string, +): DurableWorkflowDefinition { + let stage: + | DurableBatchTask + | DurableBatchScorer; + if (kind === "task") { + if (!isBatchTask(definition.evaluator.task)) { + throw new Error("Definition no longer contains the batch task"); + } + stage = definition.evaluator.task; + } else { + const scorer = resolveScorers(definition.evaluator.scores ?? []).find( + ({ name }) => name === scorerName, + )?.scorer; + if (!isBatchScorer(scorer)) { + throw new Error(`Definition no longer contains scorer ${scorerName}`); + } + stage = scorer; + } + if (stage.workflow) return stage.workflow; + if (!stage.processor) { + throw new Error("Batch definition has neither a processor nor a workflow"); + } + return { + outputNode: "$batch", + nodes: [ + { + name: "$batch", + needs: {}, + item: (rootItem) => rootItem, + processor: stage.processor as DurableBatchProcessor< + any, + any, + JsonValue + >, + result: + kind === "task" + ? (result) => result.output + : (result) => result.score, + }, + ], + }; +} + +function itemForNode( + parameters: JsonValue, + record: DurableCaseRecord, + kind: "task" | "score", + scorerName: string | undefined, + node: DurableWorkflowNodeDefinition, +) { + const rootItem = + kind === "task" + ? taskBatchItem(record, parameters) + : scorerBatchItem(record); + const outputs = nodeOutputsFor(record, kind, scorerName); + return node.item( + rootItem, + Object.fromEntries( + Object.entries(node.needs).map(([alias, dependency]) => [ + alias, + outputs[dependency], + ]), + ), + record.id, + ); +} + +function nodeOutputsFor( + record: DurableCaseRecord, + kind: "task" | "score", + scorerName?: string, +) { + if (kind === "task") return record.taskNodeOutputs; + return (record.scoreNodeOutputs[scorerName!] ??= {}); +} + +async function materializeCases< + Input, + Output, + Expected, + Metadata extends BaseMetadata, + Parameters extends EvalParameters, +>( + definition: DurableEvalDefinition< + Input, + Output, + Expected, + Metadata, + Parameters + >, + data: Evaluator["data"], + experiment: Experiment | null, +): Promise { + const evaluator = definition.evaluator; + const iterable = await _internalResolveEvaluatorData( + { + data, + projectName: definition.projectName, + projectId: evaluator.projectId, + state: evaluator.state, + }, + experiment, + ); + const records: DurableCaseRecord[] = []; + const seen = new Set(); + for await (const datum of iterable) { + const caseId = + datum.id ?? + datum.upsert_id ?? + (evaluator.caseId + ? await evaluator.caseId(datum as EvalCase) + : undefined); + if (!caseId) { + throw new Error( + "Every DurableEval case requires id, upsert_id, or caseId", + ); + } + if (seen.has(caseId)) + throw new Error(`Duplicate DurableEval case id: ${caseId}`); + seen.add(caseId); + const trialCount = datum.trialCount ?? evaluator.trialCount ?? 1; + if (!Number.isInteger(trialCount) || trialCount < 1) { + throw new Error(`Invalid trialCount for DurableEval case ${caseId}`); + } + for (let trialIndex = 0; trialIndex < trialCount; trialIndex++) { + records.push({ + id: `${caseId}:trial:${trialIndex}`, + caseId, + trialIndex, + datum: assertJsonValue(datum, `case ${caseId}`), + metadata: assertJsonValue( + "metadata" in datum ? datum.metadata : {}, + `metadata for ${caseId}`, + ), + tags: datum.tags, + taskComplete: false, + taskLogged: false, + taskNodeOutputs: {}, + scores: {}, + loggedScores: {}, + scoreNodeOutputs: {}, + classifications: {}, + loggedClassifications: {}, + }); + } + } + return records; +} + +function taskBatchItem(record: DurableCaseRecord, parameters: JsonValue) { + const datum = record.datum as EvalCase; + return { + id: record.id, + input: datum.input, + expected: "expected" in datum ? datum.expected : undefined, + metadata: record.metadata, + tags: record.tags, + parameters, + trialIndex: record.trialIndex, + }; +} + +function scorerBatchItem(record: DurableCaseRecord) { + const datum = record.datum as EvalCase; + return { + id: record.id, + input: datum.input, + output: record.output, + expected: "expected" in datum ? datum.expected : undefined, + metadata: record.metadata, + tags: record.tags, + trialIndex: record.trialIndex, + }; +} + +async function finishExperiment( + definition: DurableEvalDefinition, + state: DurableRunState, + experiment: Experiment | null, +) { + const scorerNames = resolveScorers(definition.evaluator.scores ?? []).map( + ({ name }) => name, + ); + const results = state.cases.map((record) => { + const datum = record.datum as EvalCase; + const scores = Object.assign( + {}, + ...scorerNames.map( + (name) => + _internalPrepareEvaluatorScore( + record.scores[name] as OneOrMoreScores, + name, + ).scores ?? {}, + ), + ); + const classifications = Object.assign( + {}, + ...Object.entries(record.classifications).map( + ([name, value]) => + _internalPrepareEvaluatorClassification(value as never, name) + .classifications ?? {}, + ), + ); + return { + ...datum, + output: record.output, + metadata: record.metadata, + tags: record.tags, + scores, + error: undefined, + ...(Object.keys(classifications).length > 0 ? { classifications } : {}), + } as EvalResult; + }); + if (!experiment) { + return buildEvaluatorLocalSummary( + { + ...definition.evaluator, + projectName: definition.projectName, + evalName: state.experimentName, + } as unknown as EvaluatorDef, + results, + ); + } + await experiment.flush(); + let comparisonExperimentId = definition.evaluator.baseExperimentId; + if (!comparisonExperimentId) { + try { + comparisonExperimentId = await experiment._getBaseExperimentId(); + } catch { + comparisonExperimentId = undefined; + } + } + return await experiment.summarize({ + summarizeScores: definition.evaluator.summarizeScores, + ...(comparisonExperimentId ? { comparisonExperimentId } : {}), + }); +} + +function resolveScorers( + scorers: Array< + | EvalScorer + | DurableBatchScorer + >, +) { + return scorers.map((scorer, index) => ({ + name: isBatchScorer(scorer) + ? scorer.name + : scorer.name || `scorer_${index}`, + scorer, + })); +} + +function runKey(projectName: string, evalName: string, runId: string) { + return `durable-eval/v2/runs/${contentVersion(encoder.encode(`${projectName}\0${evalName}\0${runId}`))}`; +} + +function deterministicId(value: string) { + const hex = contentVersion(encoder.encode(value)) + .padEnd(32, "0") + .slice(0, 32); + return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-5${hex.slice(13, 16)}-a${hex.slice(17, 20)}-${hex.slice(20, 32)}`; +} + +function contentVersion(value: Uint8Array) { + if (iso.hash) return iso.hash(decoder.decode(value)); + let hash = 2166136261; + for (const byte of value) { + hash ^= byte; + hash = Math.imul(hash, 16777619); + } + return (hash >>> 0).toString(16).padStart(8, "0"); +} + +async function readJson(store: DurableEvalStore, key: string) { + const value = await store.read(key); + return value ? (JSON.parse(decoder.decode(value)) as T) : undefined; +} + +async function writeJson(store: DurableEvalStore, key: string, value: unknown) { + await store.write(key, encoder.encode(stableStringify(value))); +} + +function stableStringify(value: unknown) { + return JSON.stringify(value, (_key, nested) => { + if (nested && typeof nested === "object" && !Array.isArray(nested)) { + return Object.fromEntries( + Object.entries(nested).sort(([left], [right]) => + left.localeCompare(right), + ), + ); + } + return nested; + }); +} + +function assertJsonValue(value: unknown, label: string): JsonValue { + try { + const serialized = JSON.stringify(value); + if (serialized === undefined) + throw new Error("value serializes to undefined"); + return JSON.parse(serialized) as JsonValue; + } catch (error) { + throw new Error(`${label} must be JSON serializable`, { cause: error }); + } +} + +function resultItemId(value: unknown) { + if ( + typeof value !== "object" || + value === null || + !("id" in value) || + typeof value.id !== "string" + ) { + throw new Error("Batch results must contain a string id"); + } + return value.id; +} + +function isBatchTask( + value: unknown, +): value is DurableBatchTask { + return ( + typeof value === "object" && + value !== null && + "kind" in value && + value.kind === BATCH_TASK_KIND + ); +} + +function isBatchScorer( + value: unknown, +): value is DurableBatchScorer { + return ( + typeof value === "object" && + value !== null && + "kind" in value && + value.kind === BATCH_SCORER_KIND + ); +} + +function asError(error: unknown) { + return error instanceof Error ? error : new Error(String(error)); +} diff --git a/js/src/exports.ts b/js/src/exports.ts index abd2bf611..e2098a7f9 100644 --- a/js/src/exports.ts +++ b/js/src/exports.ts @@ -266,6 +266,10 @@ export { defaultErrorScoreHandler, } from "./framework"; +export type { DurableEvalStore } from "./durable-eval"; + +export { BatchScorer, BatchTask, DurableEval } from "./durable-eval"; + export { agentAssertionScorer } from "./agent-assertions"; export { DatasetPipeline } from "./dataset-pipeline"; diff --git a/js/src/framework.ts b/js/src/framework.ts index 8628eb0e9..8dff8bafd 100644 --- a/js/src/framework.ts +++ b/js/src/framework.ts @@ -498,6 +498,38 @@ async function getExperimentParametersRef( }; } +export async function _internalInitEvaluatorExperiment( + projectName: string, + evaluator: Evaluator, + data: EvalData, + options: { + disabled?: boolean; + experimentName?: string; + update?: boolean; + } = {}, +): Promise { + if (options.disabled) return null; + const { baseExperiment } = callEvaluatorData(data); + const parameters = await getExperimentParametersRef(evaluator.parameters); + return initExperiment(evaluator.state, { + ...(evaluator.projectId + ? { projectId: evaluator.projectId } + : { project: projectName }), + experiment: options.experimentName ?? evaluator.experimentName, + description: evaluator.description, + metadata: evaluator.metadata, + tags: evaluator.tags, + isPublic: evaluator.isPublic, + update: options.update ?? evaluator.update, + baseExperiment: evaluator.baseExperimentName ?? baseExperiment, + baseExperimentId: evaluator.baseExperimentId, + gitMetadataSettings: evaluator.gitMetadataSettings, + repoInfo: evaluator.repoInfo, + dataset: Dataset.isDataset(data) ? data : undefined, + parameters, + }); +} + export function callEvaluatorData< Input, Expected, @@ -546,6 +578,65 @@ function isIterable(value: unknown): value is Iterable { ); } +export async function _internalResolveEvaluatorData( + evaluator: Pick< + EvaluatorDef, + "data" | "projectName" | "projectId" | "state" + >, + experiment: Experiment | null, +): Promise>> { + if (typeof evaluator.data === "string") { + throw new Error("Unimplemented: string data paths"); + } + let dataResult = + typeof evaluator.data === "function" ? evaluator.data() : evaluator.data; + + if ("_type" in dataResult) { + if (dataResult._type !== "BaseExperiment") { + throw new Error("Invalid _type"); + } + if (!experiment) { + throw new Error( + "Cannot use BaseExperiment() without connecting to Braintrust (you most likely set --no-send-logs)", + ); + } + let name = dataResult.name; + if (isEmpty(name)) { + const baseExperiment = await experiment.fetchBaseExperiment(); + if (!baseExperiment) { + throw new Error("BaseExperiment() failed to fetch base experiment"); + } + name = baseExperiment.name; + } + + dataResult = initExperiment(evaluator.state, { + ...(evaluator.projectId + ? { projectId: evaluator.projectId } + : { project: evaluator.projectName }), + experiment: name, + open: true, + }).asDataset(); + } + + const resolvedDataResult = + dataResult instanceof Promise ? await dataResult : dataResult; + if (isAsyncIterable>(resolvedDataResult)) { + return resolvedDataResult; + } + if ( + Array.isArray(resolvedDataResult) || + isIterable>(resolvedDataResult) + ) { + const iterable = resolvedDataResult as Iterable>; + return (async function* () { + for (const datum of iterable) yield datum; + })(); + } + throw new Error( + "Evaluator data must be an array, iterable, or async iterable", + ); +} + declare global { var _evals: EvaluatorFile; @@ -716,33 +807,13 @@ export async function Eval< const resolvedReporter = options.reporter || defaultReporter; try { - const { data, baseExperiment: defaultBaseExperiment } = callEvaluatorData( - evaluator.data, + const { data } = callEvaluatorData(evaluator.data); + const experiment = await _internalInitEvaluatorExperiment( + name, + evaluator, + data, + { disabled: Boolean(options.parent || options.noSendLogs) }, ); - const parameters = await getExperimentParametersRef(evaluator.parameters); - // NOTE: This code is duplicated with initExperiment in js/src/cli.ts. Make sure - // to update that if you change this. - const experiment = - options.parent || options.noSendLogs - ? null - : initExperiment(evaluator.state, { - ...(evaluator.projectId - ? { projectId: evaluator.projectId } - : { project: name }), - experiment: evaluator.experimentName, - description: evaluator.description, - metadata: evaluator.metadata, - tags: evaluator.tags, - isPublic: evaluator.isPublic, - update: evaluator.update, - baseExperiment: - evaluator.baseExperimentName ?? defaultBaseExperiment, - baseExperimentId: evaluator.baseExperimentId, - gitMetadataSettings: evaluator.gitMetadataSettings, - repoInfo: evaluator.repoInfo, - dataset: Dataset.isDataset(data) ? data : undefined, - parameters, - }); // Ensure experiment ID is resolved before tasks start for OTEL parent attribute support // The Experiment constructor starts resolution (fire-and-forget), but we await here to ensure completion @@ -905,6 +976,42 @@ export function classifierName( return classifier.name || `classifier_${classifier_idx}`; } +export async function _internalRunEvaluatorTask( + task: EvalTask, + datum: EvalCase, + trialIndex: number, + parameters: Record, + span: Span, + reportProgress: (event: TaskProgressEvent) => void = () => undefined, +): Promise<{ + output: unknown; + metadata: Record; + tags: string[]; +}> { + const metadata: Record = { + ...("metadata" in datum ? datum.metadata : {}), + }; + const hooks: EvalHooks, EvalParameters> = { + meta(value) { + Object.assign(metadata, value); + }, + metadata, + expected: "expected" in datum ? datum.expected : undefined, + span, + parameters, + reportProgress, + trialIndex, + tags: [...(datum.tags ?? [])], + }; + const output = await task(datum.input, hooks); + span.log({ output }); + return { + output, + metadata: hooks.metadata, + tags: hooks.tags ?? [], + }; +} + function buildSpanMetadata( results: Array<{ name: string; metadata?: Record }>, ) { @@ -930,6 +1037,50 @@ function buildSpanScores( return { resultMetadata: buildSpanMetadata(results), scoresRecord }; } +export function _internalPrepareEvaluatorScore( + scoreValue: OneOrMoreScores, + name: string, +): { + results: Score[] | null; + output?: unknown; + metadata?: Record; + scores?: Record; +} { + if (scoreValue === null) return { results: null }; + if (Array.isArray(scoreValue)) { + for (const score of scoreValue) { + if (!(typeof score === "object" && !isEmpty(score))) { + throw new Error( + `When returning an array of scores, each score must be a non-empty object. Got: ${JSON.stringify(score)}`, + ); + } + } + } + const results: Score[] = Array.isArray(scoreValue) + ? scoreValue + : typeof scoreValue === "object" && !isEmpty(scoreValue) + ? [scoreValue] + : [{ name, score: scoreValue }]; + const { resultMetadata, scoresRecord } = buildSpanScores(results); + const fields = (score: Score) => { + const { metadata: _metadata, name: _name, ...rest } = score; + return rest; + }; + return { + results, + output: + results.length === 1 + ? fields(results[0]) + : results.reduce( + (previous, score) => + mergeDicts(previous, { [score.name ?? name]: fields(score) }), + {}, + ), + metadata: resultMetadata, + scores: scoresRecord, + }; +} + async function runInScorerSpan( rootSpan: Span, spanName: string, @@ -996,6 +1147,40 @@ function toClassificationItem(c: Classification): ClassificationItem { }; } +export function _internalPrepareEvaluatorClassification( + value: OneOrMoreClassifications, + name: string, +): { + results: Classification[] | null; + output?: unknown; + metadata?: Record; + classifications?: Record; +} { + if (value === null) return { results: null }; + const results = (Array.isArray(value) ? value : [value]).map((result) => + validateClassificationResult(result, name), + ); + const classifications: Record = {}; + for (const result of results) { + (classifications[result.name] ??= []).push(toClassificationItem(result)); + } + return { + results, + output: + results.length === 1 + ? toClassificationItem(results[0]) + : results.reduce( + (previous, result) => + mergeDicts(previous, { + [result.name]: toClassificationItem(result), + }), + {}, + ), + metadata: buildSpanMetadata(results), + classifications, + }; +} + function logScoringFailures( kind: string, failures: { name: string; error: unknown }[], @@ -1076,69 +1261,14 @@ async function runEvaluatorInternal( (evaluator.state ?? _internalGetGlobalState())?.spanCache?.start(); } try { - if (typeof evaluator.data === "string") { - throw new Error("Unimplemented: string data paths"); - } - let dataResult = - typeof evaluator.data === "function" ? evaluator.data() : evaluator.data; - parameters = await validateParameters( parameters ?? {}, evaluator.parameters, ); - - if ("_type" in dataResult) { - if (dataResult._type !== "BaseExperiment") { - // For some reason, the typesystem won't let me check if dataResult._type === "BaseExperiment" - throw new Error("Invalid _type"); - } - if (!experiment) { - throw new Error( - "Cannot use BaseExperiment() without connecting to Braintrust (you most likely set --no-send-logs)", - ); - } - let name = dataResult.name; - if (isEmpty(name)) { - const baseExperiment = await experiment.fetchBaseExperiment(); - if (!baseExperiment) { - throw new Error("BaseExperiment() failed to fetch base experiment"); - } - name = baseExperiment.name; - } - - dataResult = initExperiment(evaluator.state, { - ...(evaluator.projectId - ? { projectId: evaluator.projectId } - : { project: evaluator.projectName }), - experiment: name, - open: true, - }).asDataset(); - } - - const resolvedDataResult = - dataResult instanceof Promise ? await dataResult : dataResult; - - const dataIterable: AsyncIterable> = (() => { - if (isAsyncIterable>(resolvedDataResult)) { - return resolvedDataResult; - } - if ( - Array.isArray(resolvedDataResult) || - isIterable>(resolvedDataResult) - ) { - const iterable = resolvedDataResult as Iterable< - EvalCase - >; - return (async function* () { - for (const datum of iterable) { - yield datum; - } - })(); - } - throw new Error( - "Evaluator data must be an array, iterable, or async iterable", - ); - })(); + const dataIterable = await _internalResolveEvaluatorData( + evaluator, + experiment, + ); progressReporter.start(evaluator.evalName, 0); @@ -1252,13 +1382,11 @@ async function runEvaluatorInternal( }) : undefined; - let metadata: Record = { - ...("metadata" in datum ? datum.metadata : {}), - }; + let metadata: Record = {}; const expected = "expected" in datum ? datum.expected : undefined; let output: unknown = undefined; let error: unknown | undefined = undefined; - let tags: string[] = [...(datum.tags ?? [])]; + let tags: string[] = []; const scores: Record = {}; const classifications: Record = {}; const scorerNames = (evaluator.scores ?? []).map(scorerName); @@ -1267,22 +1395,15 @@ async function runEvaluatorInternal( ); let unhandledScores: string[] | null = scorerNames; try { - const meta = (o: Record) => - (metadata = { ...metadata, ...o }); - - await rootSpan.traced( - async (span: Span) => { - const hooksForTask: EvalHooks< - unknown, - Record, - EvalParameters - > = { - meta, - metadata, - expected, + const taskResult = await rootSpan.traced( + (span: Span) => + _internalRunEvaluatorTask( + evaluator.task, + datum, + trialIndex, + parameters ?? {}, span, - parameters: parameters ?? {}, - reportProgress: (event: TaskProgressEvent) => { + (event) => { stream?.({ ...event, id: rootSpan.id, @@ -1291,27 +1412,16 @@ async function runEvaluatorInternal( object_type: "task", }); }, - trialIndex, - tags, - }; - - const outputResult = evaluator.task(datum.input, hooksForTask); - if (outputResult instanceof Promise) { - output = await outputResult; - } else { - output = outputResult; - } - - tags = hooksForTask.tags ?? []; - - span.log({ output }); - }, + ), { name: "task", spanAttributes: { type: SpanTypeAttribute.TASK }, event: { input: datum.input }, }, ); + output = taskResult.output; + metadata = taskResult.metadata; + tags = taskResult.tags; if (tags.length) { rootSpan.log({ output, metadata, expected, tags }); } else { @@ -1334,11 +1444,6 @@ async function runEvaluatorInternal( await rootSpan.export(), ); - const getOtherFields = (s: Score) => { - const { metadata: _metadata, name: _name, ...rest } = s; - return rest; - }; - const [scoreResults, classificationResults] = await Promise.all([ Promise.all( (evaluator.scores ?? []).map((score, score_idx) => @@ -1352,44 +1457,17 @@ async function runEvaluatorInternal( const scoreValue = await Promise.resolve( score(scoringArgs), ); - if (scoreValue === null) return null; - if (Array.isArray(scoreValue)) { - for (const s of scoreValue) { - if (!(typeof s === "object" && !isEmpty(s))) { - throw new Error( - `When returning an array of scores, each score must be a non-empty object. Got: ${JSON.stringify(s)}`, - ); - } - } - } - const results: Score[] = Array.isArray(scoreValue) - ? scoreValue - : typeof scoreValue === "object" && !isEmpty(scoreValue) - ? [scoreValue] - : [ - { - name: scorerNames[score_idx], - score: scoreValue, - }, - ]; - const { resultMetadata, scoresRecord } = - buildSpanScores(results); - const resultOutput = - results.length === 1 - ? getOtherFields(results[0]) - : results.reduce( - (prev, s) => - mergeDicts(prev, { - [s.name]: getOtherFields(s), - }), - {}, - ); + const prepared = _internalPrepareEvaluatorScore( + scoreValue, + scorerNames[score_idx], + ); + if (prepared.results === null) return null; span.log({ - output: resultOutput, - metadata: resultMetadata, - scores: scoresRecord, + output: prepared.output, + metadata: prepared.metadata, + scores: prepared.scores, }); - return results; + return prepared.results; }, ), ), @@ -1406,32 +1484,16 @@ async function runEvaluatorInternal( const classifierValue = await Promise.resolve( classifier(scoringArgs), ); - if (classifierValue === null) return null; - const rawResults = ( - Array.isArray(classifierValue) - ? classifierValue - : [classifierValue] - ).map((result) => - validateClassificationResult( - result, - classifierNames[idx], - ), + const prepared = _internalPrepareEvaluatorClassification( + classifierValue, + classifierNames[idx], ); - const resultOutput = - rawResults.length === 1 - ? toClassificationItem(rawResults[0]) - : rawResults.reduce( - (prev, r) => - mergeDicts(prev, { - [r.name]: toClassificationItem(r), - }), - {}, - ); + if (prepared.results === null) return null; span.log({ - output: resultOutput, - metadata: buildSpanMetadata(rawResults), + output: prepared.output, + metadata: prepared.metadata, }); - return rawResults; + return prepared.results; }, ), ), diff --git a/js/src/isomorph.ts b/js/src/isomorph.ts index 231d7e6f2..a4076bb6a 100644 --- a/js/src/isomorph.ts +++ b/js/src/isomorph.ts @@ -76,9 +76,10 @@ interface Common { path: string, opts?: { recursive?: boolean }, ) => Promise; - writeFile?: (filename: string, data: string) => Promise; + writeFile?: (filename: string, data: string | Uint8Array) => Promise; readFile?: (filename: string) => Promise; readdir?: (path: string) => Promise; + rename?: (oldPath: string, newPath: string) => Promise; utimes?: (path: string, atime: Date, mtime: Date) => Promise; unlink?: (path: string) => Promise; // eslint-disable-next-line @typescript-eslint/no-explicit-any diff --git a/js/src/logger.ts b/js/src/logger.ts index 8bdb8a67b..10662e68d 100644 --- a/js/src/logger.ts +++ b/js/src/logger.ts @@ -272,10 +272,14 @@ type StartSpanEventArgs = ExperimentLogPartialArgs & Partial; const INITIAL_SPAN_WRITE_AS_MERGE = Symbol( "braintrust.initial-span-write-as-merge", ); +const RESUME_SPAN_WITHOUT_INITIAL_WRITE = Symbol( + "braintrust.resume-span-without-initial-write", +); const INTERNAL_SPAN_CONTEXT = Symbol("braintrust.internal-span-context"); type InitialSpanWriteAsMergeArg = { readonly [INITIAL_SPAN_WRITE_AS_MERGE]?: true; + readonly [RESUME_SPAN_WITHOUT_INITIAL_WRITE]?: true; }; type InternalSpanContextArg = { @@ -2209,6 +2213,37 @@ export function updateSpan({ }); } +/** @internal Rehydrate an exported root span so work can continue in another process. */ +export function _internalResumeSpan({ + exported, + state, +}: { + exported: string; + state?: BraintrustState; +}): Span { + const resolvedState = state ?? _globalState; + const components = SpanComponentsV4.fromStr(exported); + const { row_id, root_span_id, span_id } = components.data; + if (!row_id || !root_span_id || !span_id) { + throw new Error("Only exported root spans can be resumed"); + } + return new SpanImpl({ + state: resolvedState, + parentObjectType: components.data.object_type, + parentObjectId: new LazyValue( + spanComponentsToObjectIdLambda(resolvedState, components), + ), + parentComputeObjectMetadataArgs: undefined, + parentSpanIds: { parentSpanIds: [], rootSpanId: root_span_id }, + spanId: span_id, + event: { id: row_id }, + propagatedEvent: (components.data.propagated_event ?? undefined) as + | StartSpanEventArgs + | undefined, + [RESUME_SPAN_WITHOUT_INITIAL_WRITE]: true, + }); +} + /** * An opaque W3C trace-context, as returned by * {@link extractTraceContextFromHeaders}. @@ -7716,7 +7751,9 @@ export class SpanImpl implements Span { // Deterministic spans can be initialized concurrently by separate // workflow executions, so their first write must not replace later merges. this.isMerge = args[INITIAL_SPAN_WRITE_AS_MERGE] === true; - this.logInternal({ event, internalData }); + if (!args[RESUME_SPAN_WITHOUT_INITIAL_WRITE]) { + this.logInternal({ event, internalData }); + } this.isMerge = true; } diff --git a/js/src/node/config.ts b/js/src/node/config.ts index 6b3f37802..5cfa5076b 100644 --- a/js/src/node/config.ts +++ b/js/src/node/config.ts @@ -120,6 +120,7 @@ export function configureNode() { iso.writeFile = fs.writeFile; iso.readFile = fs.readFile; iso.readdir = fs.readdir; + iso.rename = fs.rename; iso.stat = fs.stat; iso.statSync = fsSync.statSync; iso.utimes = fs.utimes; diff --git a/knip.jsonc b/knip.jsonc index 9752cb9a2..7337fa278 100644 --- a/knip.jsonc +++ b/knip.jsonc @@ -10,6 +10,9 @@ ], "ignoreIssues": { "**/generated_types.ts": ["exports", "types"], + // These support the inferred signatures of the intentionally small public + // DurableEval API and must remain exported for declaration bundling. + "js/src/durable-eval.ts": ["types"], }, "workspaces": { "dev-packages/seinfeld": {