diff --git a/py/src/braintrust/framework2.py b/py/src/braintrust/framework2.py index 4bda4dbb5..51fb23630 100644 --- a/py/src/braintrust/framework2.py +++ b/py/src/braintrust/framework2.py @@ -1,7 +1,7 @@ import dataclasses import json from collections.abc import Callable, Mapping, Sequence -from typing import Any, overload +from typing import Any, cast, overload import slugify from braintrust.logger import _internal_get_global_state, api_conn, login @@ -26,16 +26,19 @@ def __init__(self): self._cache: dict[Project, str] = {} self._name_cache: dict[str, str] = {} - def get_by_name(self, project_name: str) -> str: + def get_by_name(self, project_name: str, project_group_name: str | None = None) -> str: if project_name not in self._name_cache: state = _internal_get_global_state() - project = state.api_client().projects.post_project(body={"name": project_name, "org_name": state.org_name}) + body: dict[str, Any] = {"name": project_name, "org_name": state.org_name} + if project_group_name is not None: + body["project_group_name"] = project_group_name + project = state.api_client().projects.post_project(body=cast(Any, body)) self._name_cache[project_name] = project["id"] return self._name_cache[project_name] def get(self, project: "Project") -> str: if project not in self._cache: - self._cache[project] = self.get_by_name(project.name) + self._cache[project] = self.get_by_name(project.name, project.project_group_name) return self._cache[project] @@ -607,8 +610,9 @@ def create( class Project: """A handle to a Braintrust project.""" - def __init__(self, name: str): + def __init__(self, name: str, project_group_name: str | None = None): self.name = name + self.project_group_name = project_group_name self.tools = ToolBuilder(self) self.prompts = PromptBuilder(self) self.parameters = ParametersBuilder(self) @@ -659,8 +663,13 @@ def publish(self): class ProjectBuilder: """Creates handles to Braintrust projects.""" - def create(self, name: str) -> Project: - return Project(name) + def create(self, name: str, project_group_name: str | None = None) -> Project: + """Create a handle to a Braintrust project. + + :param name: The name of the project. + :param project_group_name: (Optional) If specified, creates the project inside the project group with this name when the project does not already exist. Requires permission to create projects in that group. + """ + return Project(name, project_group_name=project_group_name) projects = ProjectBuilder() diff --git a/py/src/braintrust/test_framework2.py b/py/src/braintrust/test_framework2.py index c49b6d796..aa7079da9 100644 --- a/py/src/braintrust/test_framework2.py +++ b/py/src/braintrust/test_framework2.py @@ -30,6 +30,38 @@ def test_project_id_cache_uses_generated_project_registration(): mock_state.app_conn.assert_not_called() +def test_project_id_cache_creates_the_project_in_its_project_group(): + mock_state = MagicMock() + mock_state.org_name = "test-org" + mock_state.api_client.return_value.projects.post_project.return_value = { + "id": "generated-project-id", + "name": "test-project", + } + project = projects.create("test-project", project_group_name="my-group") + with patch("braintrust.logger._state", mock_state): + project_id = ProjectIdCache().get(project) + + assert project_id == "generated-project-id" + mock_state.api_client.return_value.projects.post_project.assert_called_once_with( + body={"name": "test-project", "org_name": "test-org", "project_group_name": "my-group"} + ) + + +def test_project_id_cache_omits_project_group_name_when_unspecified(): + mock_state = MagicMock() + mock_state.org_name = "test-org" + mock_state.api_client.return_value.projects.post_project.return_value = { + "id": "generated-project-id", + "name": "test-project", + } + with patch("braintrust.logger._state", mock_state): + ProjectIdCache().get(projects.create("test-project")) + + mock_state.api_client.return_value.projects.post_project.assert_called_once_with( + body={"name": "test-project", "org_name": "test-org"} + ) + + class TestCodeFunctionMetadata: """Tests for CodeFunction metadata support."""