Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
85 changes: 68 additions & 17 deletions py/src/braintrust/framework.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,10 +20,12 @@
Protocol,
TypedDict,
TypeVar,
cast,
)

from tqdm.asyncio import tqdm as async_tqdm
from tqdm.auto import tqdm as std_tqdm
from typing_extensions import NotRequired

from .generated_types import FunctionFormat, FunctionOutputType, ObjectReference
from .git_fields import GitMetadataSettings, RepoInfo
Expand Down Expand Up @@ -90,6 +92,7 @@ class EvalCase(SerializableDataClass, Generic[Input, Expected]):
id: str | None = None
_xact_id: str | None = None
created: str | None = None
origin: ObjectReference | None = None


# Inheritance doesn't quite work for dataclasses, so we redefine the fields
Expand All @@ -107,6 +110,7 @@ class EvalResult(SerializableDataClass, Generic[Input, Output, Expected]):
tags: list[str] | None = None
error: Exception | None = None
exc_info: str | None = None
origin: ObjectReference | None = None


class TaskProgressEvent(TypedDict):
Expand Down Expand Up @@ -135,7 +139,7 @@ class SSEProgressEvent(TaskProgressEvent):

id: str
object_type: str
origin: ObjectReference
origin: NotRequired[ObjectReference]
name: str


Expand Down Expand Up @@ -1372,6 +1376,40 @@ def _get_persisted_base_experiment_id(experiment: Experiment) -> str | None:
return base_experiment_id if isinstance(base_experiment_id, str) and base_experiment_id else None


_OBJECT_REFERENCE_TYPES = frozenset({"project_logs", "experiment", "dataset", "prompt", "function", "prompt_session"})


def _validated_object_reference(value: Any) -> ObjectReference | None:
if not isinstance(value, Mapping):
return None

object_type = value.get("object_type")
object_id = value.get("object_id")
row_id = value.get("id")
if (
not isinstance(object_type, str)
or object_type not in _OBJECT_REFERENCE_TYPES
or not isinstance(object_id, str)
or not isinstance(row_id, str)
):
return None

validated: dict[str, Any] = {
"object_type": object_type,
"object_id": object_id,
"id": row_id,
}
for optional_field in ("_xact_id", "created"):
if optional_field not in value:
continue
optional_value = value[optional_field]
if optional_value is not None and not isinstance(optional_value, str):
return None
validated[optional_field] = optional_value

return cast(ObjectReference, validated)


async def run_evaluator(
experiment: Experiment | None,
evaluator: Evaluator[Input, Output, Expected],
Expand Down Expand Up @@ -1561,24 +1599,32 @@ async def run_evaluator_task(datum, trial_index=0):
experiment.dataset if experiment else evaluator.data if isinstance(evaluator.data, Dataset) else None
)

origin = (
{
"object_type": "dataset",
"object_id": event_dataset.id,
"id": datum.id,
"created": datum.created,
"_xact_id": datum._xact_id,
}
if event_dataset and datum.id and datum._xact_id
else None
)
if (
event_dataset
and isinstance(datum.id, str)
and datum.id
and isinstance(datum._xact_id, str)
and datum._xact_id
):
origin = cast(
ObjectReference,
{
"object_type": "dataset",
"object_id": event_dataset.id,
"id": datum.id,
"_xact_id": datum._xact_id,
**({"created": datum.created} if isinstance(datum.created, str) else {}),
},
)
else:
origin = _validated_object_reference(datum.origin)
base_event = dict(
name="eval",
span_attributes={"type": SpanTypeAttribute.EVAL},
input=datum.input,
expected=datum.expected,
tags=tags,
origin=origin,
**({"origin": origin} if origin is not None else {}),
)

if experiment:
Expand All @@ -1593,11 +1639,15 @@ async def run_evaluator_task(datum, trial_index=0):
def report_progress(event: TaskProgressEvent):
if not stream:
return
stream(
SSEProgressEvent(
id=root_span.id, origin=origin, name=evaluator.eval_name, object_type="task", **event
)
progress = SSEProgressEvent(
id=root_span.id,
name=evaluator.eval_name,
object_type="task",
**event,
)
if origin is not None:
progress["origin"] = origin
stream(progress)

hooks = DictEvalHooks(
metadata,
Expand Down Expand Up @@ -1786,6 +1836,7 @@ async def ensure_spans_flushed():
classifications=classifications or None,
error=error,
exc_info=exc_info,
origin=origin,
)

data_iterator = evaluator.data
Expand Down
150 changes: 149 additions & 1 deletion py/src/braintrust/test_framework.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@
from unittest.mock import MagicMock, patch

import pytest
from braintrust.logger import BraintrustState
from braintrust.logger import BraintrustState, Dataset, ObjectMetadata, ProjectDatasetMetadata
from braintrust.util import LazyValue

from .framework import (
Eval,
Expand All @@ -25,6 +26,153 @@
HAS_PYDANTIC = importlib.util.find_spec("pydantic") is not None


def make_dataset(dataset_id, row):
project_metadata = ObjectMetadata(id="test-project", name="test-project", full_info={})
dataset_metadata = ObjectMetadata(id=dataset_id, name="test-dataset", full_info={})
dataset = Dataset(
lazy_metadata=LazyValue(
lambda: ProjectDatasetMetadata(project=project_metadata, dataset=dataset_metadata),
use_mutex=False,
),
state=BraintrustState(),
)
return dataset, patch.object(dataset, "_refetch", return_value=[row])


SOURCE_ORIGIN = {
"object_type": "project_logs",
"object_id": "source-project",
"id": "source-row",
"_xact_id": "source-xact",
"created": "2026-06-01T00:00:00.000Z",
}
INVALID_ORIGIN = {
"object_type": "dataset",
"object_id": "source-dataset",
"id": 123,
}


def reporting_task(input_value, hooks):
hooks.report_progress(
{
"format": "code",
"output_type": "completion",
"event": "progress",
"data": "0.5",
}
)
return input_value * 2


def test_eval_case_from_dict_preserves_valid_origin():
assert EvalCase.from_dict({"input": 1, "origin": SOURCE_ORIGIN}).origin == SOURCE_ORIGIN


@pytest.mark.parametrize(
("inline_origin", "expected_origin"),
[(SOURCE_ORIGIN, SOURCE_ORIGIN), (INVALID_ORIGIN, None)],
ids=["valid", "invalid"],
)
@pytest.mark.asyncio
async def test_run_evaluator_propagates_only_valid_inline_origins(
inline_origin, expected_origin, with_memory_logger, with_simulate_login
):
streamed_events = []
evaluator = Evaluator(
project_name="test-project",
eval_name="test-evaluator",
data=[{"input": 1, "origin": inline_origin}],
task=reporting_task,
scores=[],
experiment_name=None,
metadata=None,
summarize_scores=False,
)

result = await run_evaluator(
experiment=init_test_exp("test-evaluator", "test-project"),
evaluator=evaluator,
position=None,
filters=[],
stream=streamed_events.append,
)

assert result.results[0].origin == expected_origin
root_spans = [log for log in with_memory_logger.pop() if not log["span_parents"]]
for event in (streamed_events[0], root_spans[0]):
assert event.get("origin") == expected_origin
assert ("origin" in event) is (expected_origin is not None)


@pytest.mark.asyncio
async def test_dataset_row_origin_precedes_preserved_source_origin():
row = {
"input": 1,
"id": "dataset-row",
"_xact_id": "dataset-xact",
"created": "2026-06-02T00:00:00.000Z",
"origin": SOURCE_ORIGIN,
}
dataset, patched_refetch = make_dataset("active-dataset", row)
streamed_events = []

evaluator = Evaluator(
project_name="test-project",
eval_name="test-evaluator",
data=dataset,
task=reporting_task,
scores=[],
experiment_name=None,
metadata=None,
summarize_scores=False,
)

with patched_refetch:
result = await run_evaluator(
experiment=None,
evaluator=evaluator,
position=None,
filters=[],
stream=streamed_events.append,
)

expected_origin = {
"object_type": "dataset",
"object_id": "active-dataset",
"id": "dataset-row",
"_xact_id": "dataset-xact",
"created": "2026-06-02T00:00:00.000Z",
}
assert result.results[0].origin == expected_origin
assert streamed_events[0]["origin"] == expected_origin


@pytest.mark.asyncio
async def test_incomplete_dataset_row_origin_falls_back_to_preserved_source_origin():
row = {
"input": 1,
"id": "dataset-row",
"created": "2026-06-02T00:00:00.000Z",
"origin": SOURCE_ORIGIN,
}
dataset, patched_refetch = make_dataset("active-dataset", row)
evaluator = Evaluator(
project_name="test-project",
eval_name="test-evaluator",
data=dataset,
task=lambda input_value: input_value * 2,
scores=[],
experiment_name=None,
metadata=None,
)

with patched_refetch:
result = await run_evaluator(experiment=None, evaluator=evaluator, position=None, filters=[])

assert result.results[0].origin == SOURCE_ORIGIN


@pytest.mark.asyncio
async def test_run_evaluator_basic():
"""Test that run_evaluator correctly processes a simple evaluation."""
Expand Down
23 changes: 23 additions & 0 deletions py/src/braintrust/type_tests/test_eval_generics.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

import pytest
from braintrust.framework import EvalAsync, EvalCase, EvalResultWithSummary
from braintrust.generated_types import ObjectReference
from braintrust.score import Score


Expand Down Expand Up @@ -118,3 +119,25 @@ async def test_eval_divergent_output_and_expected():
assert result.results[0].output == ModelOutput(answer="4", confidence=0.99)
assert isinstance(result.results[0].expected, frozenset)
assert result.results[0].scores.get("match") == 1.0


@pytest.mark.asyncio
async def test_eval_origin_types():
"""EvalCase, dictionary inputs, and EvalResult expose row origins."""
origin: ObjectReference = {
"object_type": "project_logs",
"object_id": "source-project",
"id": "source-row",
}

eval_case = EvalCase(input="case", origin=origin)
assert eval_case.origin == origin

result = await EvalAsync(
"test-origin-types",
data=[{"input": "dictionary", "origin": origin}],
task=lambda input_value: input_value,
scores=[],
no_send_logs=True,
)
assert result.results[0].origin == origin
6 changes: 6 additions & 0 deletions py/src/braintrust/types/_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@

from typing_extensions import NotRequired, TypedDict

from ..generated_types import ObjectReference


Input = TypeVar("Input")
Expected = TypeVar("Expected")
Expand All @@ -30,6 +32,8 @@ class EvalCaseDictNoOutput(Generic[Input], TypedDict):

id: NotRequired[str | None]
_xact_id: NotRequired[str | None]
created: NotRequired[str | None]
origin: NotRequired[ObjectReference | None]


class EvalCaseDict(Generic[Input, Expected], EvalCaseDictNoOutput[Input]):
Expand All @@ -51,3 +55,5 @@ class ExperimentDatasetEvent(TypedDict):
input: Any | None
expected: Any | None
tags: Sequence[str] | None
created: NotRequired[str | None]
origin: NotRequired[ObjectReference | None]