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
23 changes: 16 additions & 7 deletions py/src/braintrust/framework2.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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]


Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()
32 changes: 32 additions & 0 deletions py/src/braintrust/test_framework2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down