From a3dea6046c6fe5e3972adc2d3d40bc815afa6eb3 Mon Sep 17 00:00:00 2001 From: Curtis Galione Date: Tue, 28 Jul 2026 22:47:08 -0400 Subject: [PATCH 1/3] fix(devserver): preserve remote dataset version selection --- py/src/braintrust/devserver/dataset.py | 4 ++ py/src/braintrust/devserver/test_dataset.py | 73 +++++++++++++++++++++ py/src/braintrust/logger.py | 12 ++++ py/src/braintrust/test_logger.py | 58 ++++++++++++++++ 4 files changed, 147 insertions(+) create mode 100644 py/src/braintrust/devserver/test_dataset.py diff --git a/py/src/braintrust/devserver/dataset.py b/py/src/braintrust/devserver/dataset.py index de222efb..8ace8a20 100644 --- a/py/src/braintrust/devserver/dataset.py +++ b/py/src/braintrust/devserver/dataset.py @@ -40,6 +40,8 @@ async def get_dataset(state: BraintrustState, data: RunEvalData | RunEvalData1 | state=state, project=data["project_name"], name=data["dataset_name"], + **({"version": data["dataset_version"]} if "dataset_version" in data else {}), + **({"environment": data["dataset_environment"]} if "dataset_environment" in data else {}), # _internal_btql is optional **({"_internal_btql": data["_internal_btql"]} if "_internal_btql" in data else {}), ) @@ -50,6 +52,8 @@ async def get_dataset(state: BraintrustState, data: RunEvalData | RunEvalData1 | state=state, project_id=dataset_info["project_id"], name=dataset_info["dataset"], + **({"version": data["dataset_version"]} if "dataset_version" in data else {}), + **({"environment": data["dataset_environment"]} if "dataset_environment" in data else {}), # _internal_btql is optional **({"_internal_btql": data["_internal_btql"]} if "_internal_btql" in data else {}), ) diff --git a/py/src/braintrust/devserver/test_dataset.py b/py/src/braintrust/devserver/test_dataset.py new file mode 100644 index 00000000..a96f2d67 --- /dev/null +++ b/py/src/braintrust/devserver/test_dataset.py @@ -0,0 +1,73 @@ +from unittest.mock import AsyncMock, Mock + +import pytest + +from braintrust.devserver.dataset import get_dataset + + +@pytest.mark.asyncio +async def test_get_dataset_named_reference_forwards_version_and_environment(monkeypatch): + state = Mock() + dataset = Mock() + init_dataset = Mock(return_value=dataset) + monkeypatch.setattr("braintrust.devserver.dataset.init_dataset", init_dataset) + + result = await get_dataset( + state, + { + "project_name": "project-name", + "dataset_name": "dataset-name", + "dataset_version": "version-1", + "dataset_environment": "production", + "_internal_btql": {"limit": 10}, + }, + ) + + assert result is dataset + init_dataset.assert_called_once_with( + state=state, + project="project-name", + name="dataset-name", + version="version-1", + environment="production", + _internal_btql={"limit": 10}, + ) + + +@pytest.mark.asyncio +async def test_get_dataset_id_reference_forwards_version_and_environment(monkeypatch): + state = Mock() + dataset = Mock() + get_dataset_by_id = AsyncMock(return_value={"project_id": "project-id", "dataset": "dataset-name"}) + init_dataset = Mock(return_value=dataset) + monkeypatch.setattr("braintrust.devserver.dataset.get_dataset_by_id", get_dataset_by_id) + monkeypatch.setattr("braintrust.devserver.dataset.init_dataset", init_dataset) + + result = await get_dataset( + state, + { + "dataset_id": "dataset-id", + "dataset_version": "version-1", + "dataset_environment": "production", + "_internal_btql": {"limit": 10}, + }, + ) + + assert result is dataset + get_dataset_by_id.assert_awaited_once_with(state, "dataset-id") + init_dataset.assert_called_once_with( + state=state, + project_id="project-id", + name="dataset-name", + version="version-1", + environment="production", + _internal_btql={"limit": 10}, + ) + + +@pytest.mark.asyncio +async def test_get_dataset_returns_inline_data_unchanged(): + state = Mock() + data = [{"input": "hello", "expected": "world"}] + + assert await get_dataset(state, {"data": data}) is data diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index 7c89d51a..a5e277c2 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -16,6 +16,7 @@ import textwrap import threading import time +from urllib.parse import quote import traceback import types import uuid @@ -1818,6 +1819,7 @@ def init_dataset( use_output: bool = DEFAULT_IS_LEGACY_DATASET, _internal_btql: dict[str, Any] | None = None, state: BraintrustState | None = None, + environment: str | None = None, ) -> "Dataset": """ Create a new dataset in a specified project. If the project does not exist, it will be created. @@ -1826,6 +1828,7 @@ def init_dataset( :param name: The name of the dataset to create. If not specified, a name will be generated automatically. :param description: An optional description of the dataset. :param version: An optional version of the dataset (to read). If not specified, the latest version will be used. + :param environment: The environment to load the dataset from. If both `version` and `environment` are provided, `version` takes precedence. :param app_url: The URL of the Braintrust App. Defaults to https://www.braintrust.dev. :param api_key: The API key to use. If the parameter is not specified, will try to use the `BRAINTRUST_API_KEY` environment variable. If no API key is specified, will prompt the user to login. @@ -1846,6 +1849,7 @@ def init_dataset( _internal_btql = dict(cli_internal_btql) else: _internal_btql = {**cli_internal_btql, **_internal_btql} + effective_environment = None if version is not None else environment def compute_metadata(): state.login(org_name=org_name, api_key=api_key, app_url=app_url) @@ -1866,6 +1870,7 @@ def compute_metadata(): return Dataset( lazy_metadata=LazyValue(compute_metadata, use_mutex=True), version=version, + environment=effective_environment, legacy=use_output, _internal_btql=_internal_btql, state=state, @@ -5029,6 +5034,7 @@ def __init__( legacy: bool = DEFAULT_IS_LEGACY_DATASET, _internal_btql: dict[str, Any] | None = None, state: BraintrustState | None = None, + environment: str | None = None, ): if legacy: eprint( @@ -5040,6 +5046,7 @@ def mutate_record(r: DatasetEvent) -> DatasetEvent: return ensure_dataset_record(r, legacy) self._lazy_metadata = lazy_metadata + self._environment = environment self.new_records = 0 ObjectFetcher.__init__( @@ -5079,6 +5086,11 @@ def __getattr__(self, name: str) -> Any: def _get_state(self) -> BraintrustState: # Ensure the login state is populated by fetching the lazy_metadata. self._lazy_metadata.get() + if self._pinned_version is None and self._environment is not None: + response = self.state.api_conn().get_json( + f"environment-object/dataset/{self.id}/{quote(self._environment, safe='')}" + ) + self._pinned_version = response["object_version"] return self.state def _validate_event( diff --git a/py/src/braintrust/test_logger.py b/py/src/braintrust/test_logger.py index f7b47a19..84e80733 100644 --- a/py/src/braintrust/test_logger.py +++ b/py/src/braintrust/test_logger.py @@ -115,6 +115,64 @@ def test_init_with_dataset_id_and_version(self): assert dataset_dict["id"] == "dataset-id-123" assert dataset_dict["version"] == "v2" + def test_init_dataset_prefers_version_over_environment(self): + state = MagicMock() + app_conn = state.app_conn.return_value + app_conn.post_json.return_value = { + "project": {"id": "project-id", "name": "project-name"}, + "dataset": {"id": "dataset-id", "name": "dataset-name"}, + } + + dataset = braintrust.init_dataset( + state=state, + project="project-name", + name="dataset-name", + version="1", + environment="production", + ) + + assert dataset.id == "dataset-id" + app_conn.post_json.assert_called_once_with( + "api/dataset/register", + { + "project_name": "project-name", + "project_id": None, + "org_id": state.org_id, + "dataset_name": "dataset-name", + }, + ) + + def test_init_dataset_uses_environment_without_version(self): + state = MagicMock() + app_conn = state.app_conn.return_value + app_conn.post_json.return_value = { + "project": {"id": "project-id", "name": "project-name"}, + "dataset": {"id": "dataset-id", "name": "dataset-name"}, + } + state.api_conn.return_value.get_json.return_value = {"object_version": "2"} + + dataset = braintrust.init_dataset( + state=state, + project="project-name", + name="dataset-name", + environment="production", + ) + + assert dataset._get_state() is state + assert dataset._pinned_version == "2" + app_conn.post_json.assert_called_once_with( + "api/dataset/register", + { + "project_name": "project-name", + "project_id": None, + "org_id": state.org_id, + "dataset_name": "dataset-name", + }, + ) + state.api_conn.return_value.get_json.assert_called_once_with( + "environment-object/dataset/dataset-id/production" + ) + def test_init_with_repo_info_does_not_raise(self): """Test that passing repo_info to init() doesn't cause an UnboundLocalError. From c3d8e05ae4bd0f7557f504ad479ae6e17f7c262a Mon Sep 17 00:00:00 2001 From: Abhijeet Prasad Date: Wed, 29 Jul 2026 12:16:49 -0400 Subject: [PATCH 2/3] cleanup impl --- py/src/braintrust/devserver/test_dataset.py | 176 ++++++++++++++------ py/src/braintrust/logger.py | 6 +- py/src/braintrust/test_logger.py | 58 ------- 3 files changed, 130 insertions(+), 110 deletions(-) diff --git a/py/src/braintrust/devserver/test_dataset.py b/py/src/braintrust/devserver/test_dataset.py index a96f2d67..1e820814 100644 --- a/py/src/braintrust/devserver/test_dataset.py +++ b/py/src/braintrust/devserver/test_dataset.py @@ -1,73 +1,151 @@ -from unittest.mock import AsyncMock, Mock +import http.server +import json +import threading +from collections.abc import Iterator +from typing import Any +from urllib.parse import parse_qs, urlsplit import pytest - from braintrust.devserver.dataset import get_dataset +from braintrust.logger import BraintrustState + + +class _DatasetAPIHandler(http.server.BaseHTTPRequestHandler): + requests: list[tuple[str, str, dict[str, list[str]], Any]] = [] + + def log_message(self, format: str, *args: Any) -> None: + pass + + def _send_json(self, value: Any, status: int = 200) -> None: + body = json.dumps(value).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self) -> None: + parsed_url = urlsplit(self.path) + self.requests.append(("GET", parsed_url.path, parse_qs(parsed_url.query), None)) + + if parsed_url.path == "/v1/dataset/dataset-reference": + self._send_json({"project_id": "project-id", "name": "dataset-name"}) + elif parsed_url.path == "/environment-object/dataset/dataset-id/prod%2Fstable": + self._send_json({"object_version": "2"}) + else: + self._send_json({"error": "not found"}, status=404) + + def do_POST(self) -> None: + parsed_url = urlsplit(self.path) + content_length = int(self.headers.get("Content-Length", "0")) + body = json.loads(self.rfile.read(content_length)) + self.requests.append(("POST", parsed_url.path, parse_qs(parsed_url.query), body)) + + if parsed_url.path == "/api/dataset/register": + self._send_json( + { + "project": {"id": "project-id", "name": "project-name"}, + "dataset": {"id": "dataset-id", "name": "dataset-name"}, + } + ) + elif parsed_url.path == "/btql": + self._send_json({"data": []}) + else: + self._send_json({"error": "not found"}, status=404) + + +@pytest.fixture +def dataset_api_server() -> Iterator[str]: + _DatasetAPIHandler.requests = [] + server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), _DatasetAPIHandler) + thread = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.01}, daemon=True) + thread.start() + + try: + host, port = server.server_address + yield f"http://{host}:{port}" + finally: + server.shutdown() + server.server_close() + thread.join() + + +def _logged_in_state(base_url: str) -> BraintrustState: + state = BraintrustState() + state.logged_in = True + state.app_url = base_url + state.api_url = base_url + state.org_id = "org-id" + state.org_name = "test org" + return state + + +def _btql_request_body() -> dict[str, Any]: + return next( + body for method, path, _query, body in _DatasetAPIHandler.requests if method == "POST" and path == "/btql" + ) @pytest.mark.asyncio -async def test_get_dataset_named_reference_forwards_version_and_environment(monkeypatch): - state = Mock() - dataset = Mock() - init_dataset = Mock(return_value=dataset) - monkeypatch.setattr("braintrust.devserver.dataset.init_dataset", init_dataset) - - result = await get_dataset( - state, +@pytest.mark.parametrize( + "reference", + [ + pytest.param({"project_name": "project-name", "dataset_name": "dataset-name"}, id="name"), + pytest.param({"dataset_id": "dataset-reference"}, id="id"), + ], +) +async def test_get_dataset_resolves_environment_to_pinned_version( + dataset_api_server: str, reference: dict[str, str] +) -> None: + dataset = await get_dataset( + _logged_in_state(dataset_api_server), { - "project_name": "project-name", - "dataset_name": "dataset-name", - "dataset_version": "version-1", - "dataset_environment": "production", + **reference, + "dataset_environment": "prod/stable", "_internal_btql": {"limit": 10}, }, ) - assert result is dataset - init_dataset.assert_called_once_with( - state=state, - project="project-name", - name="dataset-name", - version="version-1", - environment="production", - _internal_btql={"limit": 10}, - ) + assert list(dataset) == [] + assert ( + "GET", + "/environment-object/dataset/dataset-id/prod%2Fstable", + {"org_name": ["test org"]}, + None, + ) in _DatasetAPIHandler.requests + assert _btql_request_body()["version"] == "2" + assert _btql_request_body()["query"]["limit"] == 10 @pytest.mark.asyncio -async def test_get_dataset_id_reference_forwards_version_and_environment(monkeypatch): - state = Mock() - dataset = Mock() - get_dataset_by_id = AsyncMock(return_value={"project_id": "project-id", "dataset": "dataset-name"}) - init_dataset = Mock(return_value=dataset) - monkeypatch.setattr("braintrust.devserver.dataset.get_dataset_by_id", get_dataset_by_id) - monkeypatch.setattr("braintrust.devserver.dataset.init_dataset", init_dataset) - - result = await get_dataset( - state, +@pytest.mark.parametrize( + "reference", + [ + pytest.param({"project_name": "project-name", "dataset_name": "dataset-name"}, id="name"), + pytest.param({"dataset_id": "dataset-reference"}, id="id"), + ], +) +async def test_get_dataset_prefers_explicit_version_over_environment( + dataset_api_server: str, reference: dict[str, str] +) -> None: + dataset = await get_dataset( + _logged_in_state(dataset_api_server), { - "dataset_id": "dataset-id", - "dataset_version": "version-1", - "dataset_environment": "production", - "_internal_btql": {"limit": 10}, + **reference, + "dataset_version": "1", + "dataset_environment": "prod/stable", }, ) - assert result is dataset - get_dataset_by_id.assert_awaited_once_with(state, "dataset-id") - init_dataset.assert_called_once_with( - state=state, - project_id="project-id", - name="dataset-name", - version="version-1", - environment="production", - _internal_btql={"limit": 10}, + assert list(dataset) == [] + assert not any( + path.startswith("/environment-object/") for _method, path, _query, _body in _DatasetAPIHandler.requests ) + assert _btql_request_body()["version"] == "1" @pytest.mark.asyncio -async def test_get_dataset_returns_inline_data_unchanged(): - state = Mock() +async def test_get_dataset_returns_inline_data_unchanged() -> None: data = [{"input": "hello", "expected": "world"}] - assert await get_dataset(state, {"data": data}) is data + assert await get_dataset(BraintrustState(), {"data": data}) is data diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index a5e277c2..550b0d1d 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -16,7 +16,6 @@ import textwrap import threading import time -from urllib.parse import quote import traceback import types import uuid @@ -34,7 +33,7 @@ cast, overload, ) -from urllib.parse import urlencode +from urllib.parse import quote, urlencode import chevron import exceptiongroup @@ -5087,8 +5086,9 @@ def _get_state(self) -> BraintrustState: # Ensure the login state is populated by fetching the lazy_metadata. self._lazy_metadata.get() if self._pinned_version is None and self._environment is not None: + environment_path = f"environment-object/dataset/{self.id}/{quote(self._environment, safe='')}" response = self.state.api_conn().get_json( - f"environment-object/dataset/{self.id}/{quote(self._environment, safe='')}" + environment_path, _populate_args({}, org_name=self.state.org_name) ) self._pinned_version = response["object_version"] return self.state diff --git a/py/src/braintrust/test_logger.py b/py/src/braintrust/test_logger.py index 84e80733..f7b47a19 100644 --- a/py/src/braintrust/test_logger.py +++ b/py/src/braintrust/test_logger.py @@ -115,64 +115,6 @@ def test_init_with_dataset_id_and_version(self): assert dataset_dict["id"] == "dataset-id-123" assert dataset_dict["version"] == "v2" - def test_init_dataset_prefers_version_over_environment(self): - state = MagicMock() - app_conn = state.app_conn.return_value - app_conn.post_json.return_value = { - "project": {"id": "project-id", "name": "project-name"}, - "dataset": {"id": "dataset-id", "name": "dataset-name"}, - } - - dataset = braintrust.init_dataset( - state=state, - project="project-name", - name="dataset-name", - version="1", - environment="production", - ) - - assert dataset.id == "dataset-id" - app_conn.post_json.assert_called_once_with( - "api/dataset/register", - { - "project_name": "project-name", - "project_id": None, - "org_id": state.org_id, - "dataset_name": "dataset-name", - }, - ) - - def test_init_dataset_uses_environment_without_version(self): - state = MagicMock() - app_conn = state.app_conn.return_value - app_conn.post_json.return_value = { - "project": {"id": "project-id", "name": "project-name"}, - "dataset": {"id": "dataset-id", "name": "dataset-name"}, - } - state.api_conn.return_value.get_json.return_value = {"object_version": "2"} - - dataset = braintrust.init_dataset( - state=state, - project="project-name", - name="dataset-name", - environment="production", - ) - - assert dataset._get_state() is state - assert dataset._pinned_version == "2" - app_conn.post_json.assert_called_once_with( - "api/dataset/register", - { - "project_name": "project-name", - "project_id": None, - "org_id": state.org_id, - "dataset_name": "dataset-name", - }, - ) - state.api_conn.return_value.get_json.assert_called_once_with( - "environment-object/dataset/dataset-id/production" - ) - def test_init_with_repo_info_does_not_raise(self): """Test that passing repo_info to init() doesn't cause an UnboundLocalError. From caeb7b49f5c538875cba37f5945ef4da7218f113 Mon Sep 17 00:00:00 2001 From: Abhijeet Prasad Date: Wed, 29 Jul 2026 12:31:09 -0400 Subject: [PATCH 3/3] more robust tests --- py/src/braintrust/devserver/test_dataset.py | 20 +++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/py/src/braintrust/devserver/test_dataset.py b/py/src/braintrust/devserver/test_dataset.py index 1e820814..5da4d0f4 100644 --- a/py/src/braintrust/devserver/test_dataset.py +++ b/py/src/braintrust/devserver/test_dataset.py @@ -70,13 +70,13 @@ def dataset_api_server() -> Iterator[str]: thread.join() -def _logged_in_state(base_url: str) -> BraintrustState: +def _logged_in_state(base_url: str, *, org_name: str | None = "test org") -> BraintrustState: state = BraintrustState() state.logged_in = True state.app_url = base_url state.api_url = base_url state.org_id = "org-id" - state.org_name = "test org" + state.org_name = org_name return state @@ -94,11 +94,21 @@ def _btql_request_body() -> dict[str, Any]: pytest.param({"dataset_id": "dataset-reference"}, id="id"), ], ) +@pytest.mark.parametrize( + ("org_name", "expected_query"), + [ + pytest.param("test org", {"org_name": ["test org"]}, id="with-org"), + pytest.param(None, {}, id="without-org"), + ], +) async def test_get_dataset_resolves_environment_to_pinned_version( - dataset_api_server: str, reference: dict[str, str] + dataset_api_server: str, + reference: dict[str, str], + org_name: str | None, + expected_query: dict[str, list[str]], ) -> None: dataset = await get_dataset( - _logged_in_state(dataset_api_server), + _logged_in_state(dataset_api_server, org_name=org_name), { **reference, "dataset_environment": "prod/stable", @@ -110,7 +120,7 @@ async def test_get_dataset_resolves_environment_to_pinned_version( assert ( "GET", "/environment-object/dataset/dataset-id/prod%2Fstable", - {"org_name": ["test org"]}, + expected_query, None, ) in _DatasetAPIHandler.requests assert _btql_request_body()["version"] == "2"