diff --git a/py/src/braintrust/framework.py b/py/src/braintrust/framework.py index 26780076..471a2e37 100644 --- a/py/src/braintrust/framework.py +++ b/py/src/braintrust/framework.py @@ -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 @@ -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 @@ -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) diff --git a/py/src/braintrust/test_framework.py b/py/src/braintrust/test_framework.py index 9a59c987..9dac6952 100644 --- a/py/src/braintrust/test_framework.py +++ b/py/src/braintrust/test_framework.py @@ -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 ( @@ -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)], diff --git a/py/src/braintrust/type_tests/test_eval_generics.py b/py/src/braintrust/type_tests/test_eval_generics.py index 05335547..9d655e91 100644 --- a/py/src/braintrust/type_tests/test_eval_generics.py +++ b/py/src/braintrust/type_tests/test_eval_generics.py @@ -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 --- @@ -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" diff --git a/py/src/braintrust/types/_eval.py b/py/src/braintrust/types/_eval.py index caf0ae16..31d2cb8d 100644 --- a/py/src/braintrust/types/_eval.py +++ b/py/src/braintrust/types/_eval.py @@ -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]): @@ -57,3 +58,4 @@ class ExperimentDatasetEvent(TypedDict): tags: Sequence[str] | None created: NotRequired[str | None] origin: NotRequired[ObjectReference | None] + upsert_id: NotRequired[str | None]