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
4 changes: 4 additions & 0 deletions py/src/braintrust/devserver/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 {}),
)
Expand All @@ -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 {}),
)
Expand Down
161 changes: 161 additions & 0 deletions py/src/braintrust/devserver/test_dataset.py
Original file line number Diff line number Diff line change
@@ -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
14 changes: 13 additions & 1 deletion py/src/braintrust/logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@
cast,
overload,
)
from urllib.parse import urlencode
from urllib.parse import quote, urlencode

import chevron
import exceptiongroup
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand All @@ -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)
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand All @@ -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__(
Expand Down Expand Up @@ -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(
Expand Down