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..5da4d0f4 --- /dev/null +++ b/py/src/braintrust/devserver/test_dataset.py @@ -0,0 +1,161 @@ +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, *, 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 = org_name + 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 +@pytest.mark.parametrize( + "reference", + [ + pytest.param({"project_name": "project-name", "dataset_name": "dataset-name"}, id="name"), + 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], + org_name: str | None, + expected_query: dict[str, list[str]], +) -> None: + dataset = await get_dataset( + _logged_in_state(dataset_api_server, org_name=org_name), + { + **reference, + "dataset_environment": "prod/stable", + "_internal_btql": {"limit": 10}, + }, + ) + + assert list(dataset) == [] + assert ( + "GET", + "/environment-object/dataset/dataset-id/prod%2Fstable", + expected_query, + None, + ) in _DatasetAPIHandler.requests + assert _btql_request_body()["version"] == "2" + assert _btql_request_body()["query"]["limit"] == 10 + + +@pytest.mark.asyncio +@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), + { + **reference, + "dataset_version": "1", + "dataset_environment": "prod/stable", + }, + ) + + 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() -> None: + data = [{"input": "hello", "expected": "world"}] + + assert await get_dataset(BraintrustState(), {"data": data}) is data diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index 7c89d51a..550b0d1d 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -33,7 +33,7 @@ cast, overload, ) -from urllib.parse import urlencode +from urllib.parse import quote, urlencode import chevron import exceptiongroup @@ -1818,6 +1818,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 +1827,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 +1848,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 +1869,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 +5033,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 +5045,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 +5085,12 @@ 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: + environment_path = f"environment-object/dataset/{self.id}/{quote(self._environment, safe='')}" + response = self.state.api_conn().get_json( + environment_path, _populate_args({}, org_name=self.state.org_name) + ) + self._pinned_version = response["object_version"] return self.state def _validate_event(