diff --git a/tests/v1/test_e2e.py b/tests/v1/test_e2e.py index 0c77546ae4..10f838b601 100644 --- a/tests/v1/test_e2e.py +++ b/tests/v1/test_e2e.py @@ -322,7 +322,7 @@ async def test_rubric_judge(run_v1, tmp_path): assert trace.ok assert trace.rewards["rubric"].score > 0 # the judge's verdict landed in the reward assert trace.metrics["rubric/always_yes"] == 1.0 - assert trace.info["judge"] # the call was recorded onto the trace + assert trace.judge_calls # the call was recorded onto the trace @pytest.mark.e2e diff --git a/tests/v1/test_judges.py b/tests/v1/test_judges.py index cf5e2f0b6f..d0a66bed74 100644 --- a/tests/v1/test_judges.py +++ b/tests/v1/test_judges.py @@ -1,15 +1,15 @@ -"""Pluggable judges: plugin resolution, base-`TaskConfig.judges` narrowing, the built-in -`reference` / `rubric` judges, and `Task.score` running plugged judges after the decorated -rewards. Judge model calls are faked at `Judge.complete` — no network.""" +"""Pluggable judges, scoring integration, and the typed provider-call lifecycle.""" import json import re import pytest +from openai.resources.chat.completions import AsyncCompletions +from openai.types.chat import ChatCompletion import verifiers.v1 as vf from verifiers.v1.graph import MessageNode -from verifiers.v1.judge import Judge, JudgeResponse +from verifiers.v1.judge import JudgeResponse from verifiers.v1.loaders import judge_class, judge_config_type, load_judge from verifiers.v1.types import AssistantMessage, UserMessage @@ -52,6 +52,33 @@ def make_trace( ) +def model_response(text: str) -> ChatCompletion: + return ChatCompletion.model_validate( + { + "id": "fake", + "object": "chat.completion", + "created": 0, + "model": "fake", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": text, + "refusal": None, + }, + } + ], + "usage": { + "prompt_tokens": 2, + "completion_tokens": 1, + "total_tokens": 3, + }, + } + ) + + @pytest.fixture def fake_judge_model(monkeypatch): """Fake the judge's model call, recording each prompt for assertions. Rubric calls (a JSON @@ -59,13 +86,12 @@ def fake_judge_model(monkeypatch): "yes" iff it mentions Paris; other judges reply plain "yes"/"no" by the response block.""" prompts: list[str] = [] - async def fake_complete( - self, messages, *, trace=None, schema=None, parse=None, **sampling - ): + async def fake_response(self, **kwargs): + messages = kwargs["messages"][0]["content"] prompts.append(messages) # Rubric calls carry criteria lines + a JSON `verdicts` instruction; reply one verdict per # criterion (yes iff its text mentions Paris) with a reason. Other judges get plain yes/no. - if schema is not None or '"verdicts"' in messages: + if "response_format" in kwargs or '"verdicts"' in messages: verdicts = [ { "name": name, @@ -74,23 +100,12 @@ async def fake_complete( } for name, text in re.findall(r"^- ([^:]+): (.+)$", messages, re.M) ] - response = JudgeResponse( - text=json.dumps({"verdicts": verdicts}), - parsed=schema.model_validate({"verdicts": verdicts}) - if schema - else None, - ) + text = json.dumps({"verdicts": verdicts}) else: - response = JudgeResponse( - text="yes" if "Paris" in messages.split("Response:")[-1] else "no" - ) - if parse is not None: - response.parsed = parse(response) - if trace is not None: - trace.record_judge(response) - return response - - monkeypatch.setattr(Judge, "complete", fake_complete) + text = "yes" if "Paris" in messages.split("Response:")[-1] else "no" + return model_response(text) + + monkeypatch.setattr(AsyncCompletions, "create", fake_response) return prompts @@ -215,10 +230,12 @@ async def test_reference_score(fake_judge_model): assert ( "Capital of France?" in fake_judge_model[0] ) # the task prompt is in the judge prompt - assert len(trace.info["judge"]) == 1 # the call is recorded onto the trace + assert len(trace.judge_calls) == 1 # the call is recorded onto the trace + assert trace.judge_calls[0].error is None trace = make_trace(reply="It is Rome.") assert await vf.ReferenceJudge().score(trace.task.data, trace) == 0.0 + assert trace.judge_calls[0].error is None judge = vf.ReferenceJudge(vf.ReferenceJudgeConfig(answer_field="gold")) with pytest.raises(ValueError, match="no 'gold' field"): # misconfig raises, not 0 @@ -432,21 +449,12 @@ async def test_reference_choices(fake_judge_model): async def test_error_attribution(monkeypatch, tmp_path): # The policy: a MODEL failure scores 0.0; a JUDGE failure errors the rollout (raises - # out of Task.score as a TaskError) so training skips the sample instead of + # out of Task.score) so training skips the sample instead of # punishing the model for a broken judge. - async def gibberish_judge( - self, messages, *, trace=None, schema=None, parse=None, **s - ): - response = JudgeResponse(text="as an AI language model I cannot grade this") - try: - if parse is not None: - response.parsed = parse(response) - return response - finally: - if trace is not None: - trace.record_judge(response) - - monkeypatch.setattr(Judge, "complete", gibberish_judge) + async def gibberish_judge(self, **kwargs): + return model_response("as an AI language model I cannot grade this") + + monkeypatch.setattr(AsyncCompletions, "create", gibberish_judge) taskset = JudgedTaskset( JudgedConfig.model_validate({"task": {"judges": [{"id": "reference"}]}}) ) @@ -461,7 +469,7 @@ async def gibberish_judge( trace, runtime=None ) assert "reference" not in trace.rewards - assert len(trace.info["judge"]) == 1 # the billed call is still recorded + assert trace.judge_calls[0].error is not None # --- rubric -------------------------------------------------------------------------------- @@ -521,21 +529,22 @@ async def test_rubric_score(tmp_path, fake_judge_model): trace = make_trace() assert await judge.score(trace.task.data, trace) == 0.75 assert trace.metrics == {"rubric/mentions_paris": 1.0, "rubric/is_polite": 0.0} - assert len(trace.info["judge"]) == 1 # one call for the whole rubric + assert len(trace.judge_calls) == 1 # one call for the whole rubric async def test_rubric_verdict_mismatch_raises(tmp_path, monkeypatch): # A reply that doesn't verdict exactly the rubric's criteria is a judge failure: raise # (-> rollout error), don't guess or silently score 0. - async def wrong_names(self, messages, *, trace=None, schema=None, parse=None, **s): + async def wrong_names(self, **kwargs): verdicts = {"verdicts": [{"name": "typo", "reason": "x", "verdict": "yes"}]} - return JudgeResponse(text=json.dumps(verdicts), parsed=None) + return model_response(json.dumps(verdicts)) - monkeypatch.setattr(Judge, "complete", wrong_names) + monkeypatch.setattr(AsyncCompletions, "create", wrong_names) judge = rubric_judge(tmp_path) trace = make_trace() with pytest.raises(ValueError, match="expected the batch"): await judge.score(trace.task.data, trace) + assert trace.judge_calls[0].error is not None CHOICES_TOML = '[[criteria]]\nname = "depth"\ntext = "How thorough?"\nchoices = ["none", "partial", "good"]\n' @@ -543,11 +552,11 @@ async def wrong_names(self, messages, *, trace=None, schema=None, parse=None, ** async def test_rubric_choices_normalize(tmp_path, monkeypatch): # Ordered choices (worst→best) score by rank: "partial" of ["none","partial","good"] -> 0.5. - async def graded(self, messages, *, trace=None, schema=None, parse=None, **s): + async def graded(self, **kwargs): v = {"verdicts": [{"name": "depth", "reason": "r", "verdict": "partial"}]} - return JudgeResponse(text=json.dumps(v), parsed=None) + return model_response(json.dumps(v)) - monkeypatch.setattr(Judge, "complete", graded) + monkeypatch.setattr(AsyncCompletions, "create", graded) judge = rubric_judge(tmp_path, body=CHOICES_TOML, name="q") trace = make_trace() assert await judge.score(trace.task.data, trace) == 0.5 @@ -556,11 +565,11 @@ async def graded(self, messages, *, trace=None, schema=None, parse=None, **s): async def test_rubric_off_menu_answer_raises(tmp_path, monkeypatch): # A verdict that isn't one of the criterion's choices is a judge failure, not a 0. - async def off_menu(self, messages, *, trace=None, schema=None, parse=None, **s): + async def off_menu(self, **kwargs): v = {"verdicts": [{"name": "depth", "reason": "r", "verdict": "maybe"}]} - return JudgeResponse(text=json.dumps(v), parsed=None) + return model_response(json.dumps(v)) - monkeypatch.setattr(Judge, "complete", off_menu) + monkeypatch.setattr(AsyncCompletions, "create", off_menu) judge = rubric_judge(tmp_path, body=CHOICES_TOML) trace = make_trace() with pytest.raises(ValueError, match="expected one of"): @@ -630,9 +639,7 @@ async def test_task_score_runs_plugged_judges(tmp_path, fake_judge_model): assert ( trace.rewards["quality"].score == 0.75 ) # the rubric's aggregate, under its `name` - assert ( - len(trace.info["judge"]) == 2 - ) # every judge call recorded (rubric = one call) + assert len(trace.judge_calls) == 2 # every judge call recorded (rubric = one call) async def test_task_without_judges_scores_as_before(): diff --git a/verifiers/v1/__init__.py b/verifiers/v1/__init__.py index 55ab5998b7..788ddb09f5 100644 --- a/verifiers/v1/__init__.py +++ b/verifiers/v1/__init__.py @@ -106,6 +106,7 @@ Error, EvalRunInfo, GenerationSpan, + JudgeCall, ModelCall, Reward, RunInfo, @@ -179,6 +180,7 @@ "AgentInfo", "RunInfo", "EvalRunInfo", + "JudgeCall", "ModelCall", "TrainRunInfo", "VersionInfo", diff --git a/verifiers/v1/cli/dashboard/eval.py b/verifiers/v1/cli/dashboard/eval.py index 67f2ab7f02..d37e9f36bd 100644 --- a/verifiers/v1/cli/dashboard/eval.py +++ b/verifiers/v1/cli/dashboard/eval.py @@ -405,8 +405,10 @@ def _breakdown(scored: list[Trace], done: list[Trace]) -> Table | None: if trace.usage is not None and trace.usage.cost is not None: total_cost += trace.usage.cost have_cost = True - # Judge / auxiliary scoring calls (off the message graph) shown separately from the agent's. - judge = Usage.aggregate(trace.extra_usage) + # Judge calls (off the message graph) shown separately from the agent's. + judge = Usage.aggregate( + call.call.usage for call in trace.judge_calls if call.call.usage is not None + ) if judge is not None: total_judge_in += judge.input_tokens total_judge_out += judge.completion_tokens diff --git a/verifiers/v1/cli/replay.py b/verifiers/v1/cli/replay.py index a535430221..b450e43891 100644 --- a/verifiers/v1/cli/replay.py +++ b/verifiers/v1/cli/replay.py @@ -153,14 +153,13 @@ async def rescore( ) -> None: async with sem or contextlib.nullcontext(): st.start = time.time() - # Generation failures have no complete transcript to score. - if trace.stop_condition == "error": + # Failed traces may include scoring components replay cannot rerun offline. + if not trace.ok: st.state, st.detail, st.end = "skipped", "rollout errored", time.time() else: st.state = "running" # Clear the recorded scores; only offline-recomputable ones come back. - trace.info.pop("judge", None) - trace.rewards, trace.metrics, trace.extra_usage = {}, {}, [] + trace.judge_calls, trace.rewards, trace.metrics = [], {}, {} # The declared Task for a rebuilt row, the base Task for the # WireTaskData fallback. task = (task_cls if isinstance(row, data_cls) else Task)( diff --git a/verifiers/v1/configs/judge.py b/verifiers/v1/configs/judge.py index 2ac8cca380..468419beba 100644 --- a/verifiers/v1/configs/judge.py +++ b/verifiers/v1/configs/judge.py @@ -13,7 +13,10 @@ class JudgeSamplingConfig(SamplingConfig): - pass + """Keep provider token-limit aliases distinct on the wire.""" + + max_tokens: int | None = None + max_completion_tokens: int | None = None class JudgeConfig(BaseClientConfig): diff --git a/verifiers/v1/interception/server.py b/verifiers/v1/interception/server.py index 13a7e300f4..6495860864 100644 --- a/verifiers/v1/interception/server.py +++ b/verifiers/v1/interception/server.py @@ -24,7 +24,6 @@ import logging import secrets import time -import traceback from collections.abc import AsyncIterator from contextlib import asynccontextmanager from typing import Literal @@ -38,7 +37,6 @@ from verifiers.v1 import graph from verifiers.v1.errors import ( OverlongPromptError, - ProviderError, RolloutError, TaskError, ) @@ -271,19 +269,7 @@ def record_call( finish_reason=finish_reason, usage=usage, time=TimeSpan(start=started, end=time.time()), - error=None - if error is None - else Error( - type=type(error).__name__, - message=str(error), - status_code=getattr(error, "status_code", None), - # Provider errors already carry the actionable upstream diagnostic. - # Format from the exception object: the record is written in a - # `finally`, where the ambient exception state is already cleared. - traceback=None - if isinstance(error, ProviderError) - else "".join(traceback.format_exception(error)), - ), + error=None if error is None else Error.from_exception(error), ) ) diff --git a/verifiers/v1/judge.py b/verifiers/v1/judge.py index 53d224e408..6ce499586e 100644 --- a/verifiers/v1/judge.py +++ b/verifiers/v1/judge.py @@ -32,11 +32,9 @@ async def correct(self, trace) -> float: A judge is cheap to construct (the HTTP client is opened per call, inside `complete`, and closed when the call returns), so build it where you use it. -Passing `trace=` records the call onto it — a typed record appended to `trace.info["judge"]` and -the call's tokens + cost added to `trace.extra_usage` (kept separate from the agent's `trace.usage`), -so judge behaviour and spend are no longer invisible. The record lands even if the judge refuses, an -empty structured output comes back, or `parse` raises (the request was already billed). Omit `trace` -for a pure call (e.g. in tests). +Passing `trace=` appends a small typed `JudgeCall` before calling the provider. Its nested +`ModelCall` uses the same model, sampling, endpoint, timing, usage, and error fields as an +agentic judge trace without copying gold-bearing prompts into the trace. A judge can also be *plugged* rather than called from task code: a judge with an `id` and a `score` implementation is a plugin (like a taskset or harness — see `verifiers.v1.judges` for the @@ -49,10 +47,13 @@ async def correct(self, trace) -> float: from __future__ import annotations import re +import time from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Callable, Generic, Literal, cast +from openai import ContentFilterFinishReasonError, OpenAIError from pydantic import BaseModel +from openai.lib._parsing import type_to_response_format_param from typing_extensions import TypeVar from verifiers.v1.clients.config import build_async_openai @@ -60,8 +61,15 @@ async def correct(self, trace) -> float: JudgeConfig, judge_key, ) -from verifiers.v1.dialects.chat import message_to_wire +from verifiers.v1.dialects.chat import ChatDialect, message_to_wire, response_from_wire +from verifiers.v1.errors import model_error from verifiers.v1.scoring import parse_judge_choice +from verifiers.v1.trace import ( + Error, + JudgeCall, + ModelCall, + TimeSpan, +) from verifiers.v1.utils.generic import generic_type from verifiers.v1.types import Messages, StrictBaseModel, Usage @@ -170,50 +178,79 @@ async def complete( parse: Callable[[JudgeResponse[Any]], Any] | None = None, **sampling: Any, ) -> JudgeResponse[Any]: - """Call the judge and record billed usage even when parsing fails.""" + """Call the judge and preserve its provider and parsing failures.""" wire = ( [{"role": "user", "content": messages}] if isinstance(messages, str) else [message_to_wire(m) for m in messages] ) - kwargs: dict[str, Any] = {"model": self.config.model, "messages": wire} - kwargs.update(self.config.sampling.model_dump(exclude_none=True)) - kwargs.update(sampling) + kwargs: dict[str, Any] = { + **self.config.sampling.model_dump(exclude_none=True), + **sampling, + "model": self.config.model, + "messages": wire, + } + if sampling.get("max_tokens") is not None: + kwargs.pop("max_completion_tokens", None) + elif sampling.get("max_completion_tokens") is not None: + kwargs.pop("max_tokens", None) + response_format = ( + type_to_response_format_param(schema) if schema is not None else None + ) + dialect = ChatDialect() + provider_call = ModelCall( + model=self.config.model, + sampling=dialect.parse_sampling(kwargs), + endpoint=dialect.upstream_path, + time=TimeSpan(start=time.time()), + ) + call = JudgeCall( + judge=f"{type(self).__module__}.{type(self).__qualname__}", + call=provider_call, + ) + if trace is not None: + trace.judge_calls.append(call) + if response_format is not None: + kwargs["response_format"] = response_format - response: JudgeResponse[Any] | None = None try: async with build_async_openai(self.config) as client: - if schema is not None: - completion = await client.beta.chat.completions.parse( - response_format=schema, **kwargs - ) - choice = completion.choices[0] - response = JudgeResponse( - text=choice.message.content or "", - parsed=choice.message.parsed, - usage=Usage.from_openai(completion.usage), - ) - if choice.message.refusal is not None: - raise RuntimeError( - f"judge refused structured output: {choice.message.refusal}" - ) - if response.parsed is None: - raise RuntimeError( - f"judge returned no parseable structured output " - f"(finish_reason={choice.finish_reason})" - ) - else: - completion = await client.chat.completions.create(**kwargs) - response = JudgeResponse( - text=completion.choices[0].message.content or "", - usage=Usage.from_openai(completion.usage), - ) + completion = await client.chat.completions.create(**kwargs) + provider_response = response_from_wire(completion) + provider_call.finish_reason = provider_response.finish_reason + provider_call.usage = provider_response.usage + except OpenAIError as error: + provider_error = model_error(error) + provider_call.error = Error.from_exception(provider_error) + raise provider_error from error + except BaseException as error: + provider_call.error = Error.from_exception(error) + raise + finally: + provider_call.time.end = time.time() + + try: + response = JudgeResponse[Any]( + text=provider_response.message.content or "", + usage=provider_response.usage, + ) + call.text = response.text + refusal = completion.choices[0].message.refusal + if refusal is not None: + call.refusal = refusal + raise ValueError(f"judge refused output: {refusal}") + if completion.choices[0].finish_reason == "content_filter": + raise ContentFilterFinishReasonError() + if schema is not None and provider_response.finish_reason == "length": + raise ValueError("judge structured output was truncated") + if schema is not None: + response.parsed = schema.model_validate_json(response.text) if parse is not None: response.parsed = parse(response) return response - finally: - if trace is not None and response is not None: - trace.record_judge(response) + except BaseException as error: + call.error = Error.from_exception(error) + raise async def evaluate( self, *, trace: "Trace | None" = None, **fields: Any diff --git a/verifiers/v1/judges/rubric.py b/verifiers/v1/judges/rubric.py index c7a84856c0..5f66fd0aee 100644 --- a/verifiers/v1/judges/rubric.py +++ b/verifiers/v1/judges/rubric.py @@ -12,7 +12,13 @@ from pydantic import field_validator from verifiers.v1.configs.judge import JudgeConfig -from verifiers.v1.judge import Judge, JudgeView, judge_question, judge_response +from verifiers.v1.judge import ( + Judge, + JudgeResponse, + JudgeView, + judge_question, + judge_response, +) from verifiers.v1.task import TaskData from verifiers.v1.trace import Trace from verifiers.v1.types import ID, StrictBaseModel @@ -23,7 +29,7 @@ # Appended to the prompt when `structured_output` is off, so the reply is JSON we can parse. Each # criterion carries a one-sentence `reason` written *before* the verdict (chain-of-thought), so the -# verdict follows from it and the reasoning is auditable in `trace.info["judge"]`. +# verdict follows from it and the reasoning is auditable in `trace.judge_calls`. JSON_SUFFIX = ( "\n\nRespond with ONLY a JSON object and nothing else, in exactly this shape:\n" '{"verdicts": [{"name": "", "reason": " str: criteria="\n".join(render(c) for c in batch), reference=reference, ) - if self.config.structured_output: - result = await self.evaluate(trace=trace, **fields) - verdicts = cast(RubricVerdicts, result.parsed).verdicts - else: - messages = cast(str, self.build_messages(**fields)) + JSON_SUFFIX - result = await self.complete(messages, trace=trace) - obj = first_verdicts_object(result.text) - if obj is None: + by_criterion = {c.name: c for c in batch} + + def parse(response: JudgeResponse) -> dict[str, float]: + parsed = response.parsed or first_verdicts_object(response.text) + if parsed is None: raise ValueError( - f"judge returned no verdicts JSON object: {result.text!r}" + f"judge returned no verdicts JSON object: {response.text!r}" ) - verdicts = RubricVerdicts.model_validate(obj).verdicts - # Exactly one verdict per criterion in the batch, matched by name — anything else is a - # judge failure and must error the rollout, not score the model (see `judge_verdict`). - by_criterion = {c.name: c for c in batch} - if sorted(v.name for v in verdicts) != sorted(by_criterion): - raise ValueError( - f"judge verdicts name {sorted(v.name for v in verdicts)}, expected the " - f"batch's {sorted(by_criterion)}" - ) - scores: dict[str, float] = {} - for v in verdicts: - choices = by_criterion[v.name].choices - # An off-menu answer is a judge failure, not a zero score. - if v.verdict not in choices: + verdicts = RubricVerdicts.model_validate(parsed).verdicts + if sorted(v.name for v in verdicts) != sorted(by_criterion): raise ValueError( - f"judge answered {v.verdict!r} for '{v.name}', expected one of {choices}" + f"judge verdicts name {sorted(v.name for v in verdicts)}, expected the " + f"batch's {sorted(by_criterion)}" ) - scores[v.name] = normalize_choice(v.verdict, choices) - return scores + scores: dict[str, float] = {} + for verdict in verdicts: + choices = by_criterion[verdict.name].choices + if verdict.verdict not in choices: + raise ValueError( + f"judge answered {verdict.verdict!r} for '{verdict.name}', " + f"expected one of {choices}" + ) + scores[verdict.name] = normalize_choice(verdict.verdict, choices) + return scores + + messages = self.build_messages(**fields) + schema = self.schema if self.config.structured_output else None + if schema is None: + messages = cast(str, messages) + JSON_SUFFIX + result = await self.complete(messages, trace=trace, schema=schema, parse=parse) + return cast(dict[str, float], result.parsed) async def score(self, task: TaskData, trace: Trace) -> float: criteria = self.criteria diff --git a/verifiers/v1/push.py b/verifiers/v1/push.py index f16fde4745..e76d48833f 100644 --- a/verifiers/v1/push.py +++ b/verifiers/v1/push.py @@ -53,6 +53,12 @@ def dump(messages): task = trace.task.data.model_dump(mode="json", exclude_none=True) branches = trace.branches + info = dict(trace.info) + if trace.judge_calls: + info["judge_calls"] = [ + call.model_dump(mode="json", exclude_none=True) + for call in trace.judge_calls + ] sample = { "sample_id": trace.id, "example_id": trace.task.data.idx, @@ -88,7 +94,7 @@ def dump(messages): "token_usage": trace.usage.model_dump(mode="json", exclude_none=True) if trace.usage else None, - "info": dict(trace.info) or None, + "info": info or None, } # Flatten sub-rewards to top-level keys the way v0 does (raw scores, as v0's # per-function outputs were); env metrics stay nested. diff --git a/verifiers/v1/trace.py b/verifiers/v1/trace.py index ae90009e99..477b2ee2e3 100644 --- a/verifiers/v1/trace.py +++ b/verifiers/v1/trace.py @@ -5,7 +5,7 @@ import traceback import uuid from collections.abc import Mapping -from typing import TYPE_CHECKING, Annotated, Any, Generic, Literal +from typing import Annotated, Any, Generic, Literal from typing_extensions import TypeVar @@ -13,9 +13,6 @@ from pydantic import Field, PrivateAttr from renderers.base import MultiModalData -if TYPE_CHECKING: - from verifiers.v1.judge import JudgeResponse - from verifiers.v1 import graph from verifiers.v1.errors import ProviderError from verifiers.v1.graph import MessageNode @@ -86,6 +83,18 @@ class Error(StrictBaseModel): chosen for a transport fault); None when the failure carried no HTTP exchange.""" traceback: str | None = None + @classmethod + def from_exception(cls, error: BaseException) -> "Error": + """Capture one exception without depending on ambient ``except`` state.""" + return cls( + type=type(error).__name__, + message=str(error) or type(error).__name__, + status_code=getattr(error, "status_code", None), + traceback=None + if isinstance(error, ProviderError) + else "".join(traceback.format_exception(error)), + ) + class ModelCall(StrictBaseModel): """One provider exchange behind a sampled turn; its conversation is the linked @@ -94,7 +103,7 @@ class ModelCall(StrictBaseModel): node: int | None = None """Index into `Trace.nodes` of the assistant node this call committed — the link into the message graph (the call's conversation is that node's root-to-self path). None for - a call that committed no turn (see `error`).""" + a failed agent call or an off-graph call nested under `JudgeCall`.""" model: str | None = None """The model requested from the provider. The rollout's model override makes this `agent.config.model` on every call; recorded per call because it is cheap and provable.""" @@ -118,6 +127,17 @@ class ModelCall(StrictBaseModel): success. A failed call still records the settings it was sent with.""" +class JudgeCall(StrictBaseModel): + """One off-graph judge decision and its shared model-call record.""" + + judge: str + """Fully qualified Judge class.""" + call: ModelCall + text: str | None = None + refusal: str | None = None + error: Error | None = None + + class Branch(StrictBaseModel): """A root-to-leaf graph path; each branch becomes one training sample.""" @@ -266,7 +286,7 @@ def num_input_tokens(self) -> int: """Raw tensor fields kept on the msgpack wire but excluded from JSON records.""" -TRACE_VERSION = 4 +TRACE_VERSION = 5 """Version of the trace record schema (see `Trace.model_json_schema()`). Bumped on breaking shape changes; optional-with-default fields are additive and don't bump it.""" @@ -383,6 +403,8 @@ class Trace(StrictBaseModel, Generic[DataT, StateT, AgentConfigT]): calls: list[ModelCall] = Field(default_factory=list) """Every provider exchange behind the sampled turns, in order: raw wire request/response plus per-call timing and errors, linked into `nodes` via `ModelCall.node`.""" + judge_calls: list[JudgeCall] = Field(default_factory=list) + """Every LLM judge decision made while scoring this trace.""" rewards: dict[str, Reward] = Field(default_factory=dict) """Named rewards from tasks, judges, and the env's `score()` — each keeps its @@ -394,9 +416,6 @@ class Trace(StrictBaseModel, Generic[DataT, StateT, AgentConfigT]): state: StateT = Field(default_factory=State, exclude=True) """Transient state shared with servers and scoring; excluded from every dump.""" - extra_usage: list[Usage] = Field(default_factory=list) - """Usage from judges and other calls outside the agent's message graph.""" - is_completed: bool = False ok: bool = False """THE success sentinel, stamped by the engine when the rollout ran to @@ -569,11 +588,6 @@ def record_metrics(self, values: Mapping[str, float]) -> None: for name, value in values.items(): self.record_metric(name, value) - def record_judge(self, response: JudgeResponse) -> None: - self.info.setdefault("judge", []).append(response.model_dump()) - if response.usage is not None: - self.extra_usage.append(response.usage) - def record_reward(self, name: str, value: float, weight: float = 1.0) -> None: reward = Reward(score=float(value), weight=float(weight)) if name in self.rewards: @@ -607,18 +621,7 @@ def split_generation(self) -> None: gen.harness.duration = gen.duration - gen.model.duration def capture_error(self, error: Exception) -> None: - self.errors.append( - Error( - type=type(error).__name__, - message=str(error), - status_code=getattr(error, "status_code", None), - # Provider errors already carry the actionable upstream diagnostic. - # Keep full tracebacks for every other failure. - traceback=None - if isinstance(error, ProviderError) - else traceback.format_exc(), - ) - ) + self.errors.append(Error.from_exception(error)) self.ok = False self.stop("error")