Skip to content
Open
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
8 changes: 8 additions & 0 deletions py/src/braintrust/framework.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
import re
import sys
import traceback
import uuid
import warnings
from collections import defaultdict
from collections.abc import Awaitable, Callable, Coroutine, Iterable, Iterator, Mapping, Sequence
Expand Down Expand Up @@ -104,6 +105,7 @@ class EvalCase(SerializableDataClass, Generic[Input, Expected]):
_xact_id: str | None = None
created: str | None = None
origin: ObjectReference | None = None
upsert_id: str | None = None


# Inheritance doesn't quite work for dataclasses, so we redefine the fields
Expand Down Expand Up @@ -1648,6 +1650,12 @@ async def run_evaluator_task(datum, trial_index=0):
tags=tags,
**({"origin": origin} if origin is not None else {}),
)
if datum.upsert_id:
base_event["id"] = (
datum.upsert_id
if trial_index == 0
else str(uuid.uuid5(uuid.NAMESPACE_URL, f"braintrust:eval:{datum.upsert_id}:trial:{trial_index}"))
)

if experiment:
root_span = experiment.start_span(**base_event)
Expand Down
55 changes: 54 additions & 1 deletion py/src/braintrust/test_framework.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from unittest.mock import MagicMock, patch

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

from .framework import (
Expand Down Expand Up @@ -69,6 +69,59 @@ def test_eval_case_from_dict_preserves_valid_origin():
assert EvalCase.from_dict({"input": 1, "origin": SOURCE_ORIGIN}).origin == SOURCE_ORIGIN


@pytest.mark.parametrize("upsert_id", ["eval-row", None, ""])
@pytest.mark.parametrize("as_dict", [True, False], ids=["dict", "dataclass"])
@pytest.mark.parametrize("use_experiment", [True, False], ids=["experiment", "parent-context"])
@pytest.mark.parametrize(("trial_count", "row_trial_count"), [(1, None), (3, None), (1, 3), (3, 1)])
@pytest.mark.asyncio
async def test_run_evaluator_upsert_id(
upsert_id, as_dict, use_experiment, trial_count, row_trial_count, with_memory_logger, with_simulate_login
):
experiment = init_test_exp("test-upsert", "test-project")
expected_trials = row_trial_count if row_trial_count is not None else trial_count
root_ids = []
for input_value in [1, 2]:
data = {"input": input_value, "id": "dataset-row", "origin": SOURCE_ORIGIN, "trial_count": row_trial_count}
if upsert_id is not None:
data["upsert_id"] = upsert_id
evaluator = Evaluator(
project_name="test-project",
eval_name="test-upsert",
data=[data if as_dict else EvalCase(**data)],
task=lambda input_value, hooks: [input_value * 2, hooks.trial_index],
scores=[],
experiment_name=None,
metadata=None,
summarize_scores=False,
trial_count=trial_count,
)
with parent_context(experiment.export()):
await run_evaluator(
experiment=experiment if use_experiment else None,
evaluator=evaluator,
position=None,
filters=[],
)
logs = with_memory_logger.pop()
roots = [log for log in logs if not log["span_parents"]]
children = [log for log in logs if log["span_parents"]]
assert len(roots) == len(children) == expected_trials
assert sorted(root["output"] for root in roots) == [
[input_value * 2, trial_index] for trial_index in range(expected_trials)
]
assert all(root["origin"] == SOURCE_ORIGIN for root in roots)
root_outputs = {root["root_span_id"]: root["output"] for root in roots}
assert all(child["output"] == root_outputs[child["root_span_id"]] for child in children)
root_ids.append({root["output"][1]: root["id"] for root in roots})

if upsert_id:
assert root_ids[0] == root_ids[1]
assert root_ids[0][0] == upsert_id
else:
assert set(root_ids[0].values()).isdisjoint(root_ids[1].values())
assert all("dataset-row" not in ids.values() for ids in root_ids)


@pytest.mark.parametrize(
("inline_origin", "expected_origin"),
[(SOURCE_ORIGIN, SOURCE_ORIGIN), (INVALID_ORIGIN, None)],
Expand Down
15 changes: 15 additions & 0 deletions py/src/braintrust/type_tests/test_eval_generics.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from braintrust.framework import EvalAsync, EvalCase, EvalResultWithSummary
from braintrust.generated_types import ObjectReference
from braintrust.score import Score
from braintrust.types._eval import EvalCaseDict, EvalCaseDictNoOutput


# --- Domain types for testing ---
Expand Down Expand Up @@ -141,3 +142,17 @@ async def test_eval_origin_types():
no_send_logs=True,
)
assert result.results[0].origin == origin


def test_eval_upsert_id_types() -> None:
eval_case: EvalCase[str, str] = EvalCase(input="case", upsert_id="case-root")
assert eval_case.upsert_id == "case-root"
dict_case: EvalCaseDictNoOutput[str] = {"input": "dictionary", "upsert_id": "dict-root"}
expected_case: EvalCaseDict[str, str] = {
"input": "expected",
"expected": "expected",
"upsert_id": "expected-root",
}

assert EvalCase.from_dict(dict(dict_case)).upsert_id == "dict-root"
assert EvalCase.from_dict(dict(expected_case)).upsert_id == "expected-root"
2 changes: 2 additions & 0 deletions py/src/braintrust/types/_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ class EvalCaseDictNoOutput(Generic[Input], TypedDict):
_xact_id: NotRequired[str | None]
created: NotRequired[str | None]
origin: NotRequired[ObjectReference | None]
upsert_id: NotRequired[str | None]


class EvalCaseDict(Generic[Input, Expected], EvalCaseDictNoOutput[Input]):
Expand All @@ -57,3 +58,4 @@ class ExperimentDatasetEvent(TypedDict):
tags: Sequence[str] | None
created: NotRequired[str | None]
origin: NotRequired[ObjectReference | None]
upsert_id: NotRequired[str | None]