From 045771a6c3ddab1bf5f254d5ce9578055f513fab Mon Sep 17 00:00:00 2001 From: Ankur Goyal Date: Thu, 11 Jun 2026 22:52:41 -0700 Subject: [PATCH 1/3] add agent type --- py/src/braintrust/logger.py | 155 +++++++++++++++++++++++++++---- py/src/braintrust/test_logger.py | 36 +++++++ 2 files changed, 173 insertions(+), 18 deletions(-) diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index 9a65dee30..ee0a0b77e 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -1522,6 +1522,10 @@ class ProjectDatasetMetadata: class OrgProjectMetadata: org_id: str project: ObjectMetadata + agent: ObjectMetadata | None = None + + +DefaultSpanAttributes = dict[str, Any] | LazyValue[dict[str, Any] | None] # Pyright produces an error for overlapping overloads @@ -1834,7 +1838,40 @@ def compute_metadata(): ) -def _compute_logger_metadata(project_name: str | None = None, project_id: str | None = None): +def _agent_register_endpoint_not_found(error: AugmentedHTTPError) -> bool: + cause = error.__cause__ + response = getattr(cause, "response", None) + return getattr(response, "status_code", None) == 404 + + +def _register_agent_metadata(project_id: str, agent_name: str) -> ObjectMetadata: + try: + response = _state.app_conn().post_json( + "api/agent/register", + { + "project_id": project_id, + "agent_name": agent_name, + }, + ) + except AugmentedHTTPError as e: + if _agent_register_endpoint_not_found(e): + _logger.debug("Agent registration endpoint not found; using agent name as span attribute") + return ObjectMetadata(id=agent_name, name=agent_name, full_info={}) + raise + + resp_agent = response["agent"] + return ObjectMetadata( + id=resp_agent["id"], + name=resp_agent["name"], + full_info=resp_agent, + ) + + +def _compute_logger_metadata( + project_name: str | None = None, + project_id: str | None = None, + agent_name: str | None = None, +): login() org_id = _state.org_id if project_id is None: @@ -1846,24 +1883,42 @@ def _compute_logger_metadata(project_name: str | None = None, project_id: str | }, ) resp_project = response["project"] - return OrgProjectMetadata( + metadata = OrgProjectMetadata( org_id=org_id, - project=ObjectMetadata(id=resp_project["id"], name=resp_project["name"], full_info=resp_project), + project=ObjectMetadata( + id=resp_project["id"], + name=resp_project["name"], + full_info=resp_project, + ), ) elif project_name is None: response = _state.app_conn().get_json("api/project", {"id": project_id}) - return OrgProjectMetadata( - org_id=org_id, project=ObjectMetadata(id=project_id, name=response["name"], full_info=response) + metadata = OrgProjectMetadata( + org_id=org_id, + project=ObjectMetadata( + id=project_id, + name=response["name"], + full_info=response, + ), ) else: - return OrgProjectMetadata( - org_id=org_id, project=ObjectMetadata(id=project_id, name=project_name, full_info=dict()) + metadata = OrgProjectMetadata( + org_id=org_id, + project=ObjectMetadata( + id=project_id, + name=project_name, + full_info=dict(), + ), ) + if agent_name is not None: + metadata.agent = _register_agent_metadata(metadata.project.id, agent_name) + return metadata def init_logger( project: str | None = None, project_id: str | None = None, + agent: str | None = None, async_flush: bool = True, app_url: str | None = None, api_key: str | None = None, @@ -1877,6 +1932,7 @@ def init_logger( :param project: The name of the project to log into. If unspecified, will default to the Global project. :param project_id: The id of the project to log into. This takes precedence over project if specified. + :param agent: The name of the agent to register and associate with logged spans. The registered agent id is stored as `span_attributes.agent_id`. :param async_flush: If true (the default), log events will be batched and sent asynchronously in a background thread. If false, log events will be sent synchronously. Set to false in serverless environments. :param app_url: The URL of the Braintrust API. 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 @@ -1888,7 +1944,11 @@ def init_logger( """ state = state or _state - compute_metadata_args = dict(project_name=project, project_id=project_id) + compute_metadata_args = dict( + project_name=project, + project_id=project_id, + agent_name=agent, + ) link_args = { "app_url": app_url, @@ -1898,17 +1958,38 @@ def init_logger( } def compute_metadata(): - state.login(org_name=org_name, api_key=api_key, app_url=app_url, force_login=force_login) + state.login( + org_name=org_name, + api_key=api_key, + app_url=app_url, + force_login=force_login, + ) return _compute_logger_metadata(**compute_metadata_args) # For loggers, enable queue size limit enforcement (bounded queue) state.enforce_queue_size_limit(True) + lazy_metadata = LazyValue(compute_metadata, use_mutex=True) + + def compute_default_span_attributes() -> dict[str, Any]: + metadata = lazy_metadata.get() + if metadata.agent: + return {"agent_id": metadata.agent.id} + return {"agent_id": agent} + ret = Logger( - lazy_metadata=LazyValue(compute_metadata, use_mutex=True), + lazy_metadata=lazy_metadata, async_flush=async_flush, compute_metadata_args=compute_metadata_args, link_args=link_args, + default_span_attributes=( + None + if agent is None + else LazyValue( + compute_default_span_attributes, + use_mutex=False, + ) + ), state=state, ) if set_current: @@ -4163,6 +4244,7 @@ def __init__( name: str | None = None, type: SpanTypeAttribute | None = None, default_root_type: SpanTypeAttribute | None = None, + default_span_attributes: DefaultSpanAttributes | None = None, span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, @@ -4175,6 +4257,11 @@ def __init__( ): if span_attributes is None: span_attributes = SpanAttributes() + lazy_default_span_attributes = ( + default_span_attributes if isinstance(default_span_attributes, LazyValue) else None + ) + static_default_span_attributes = None if lazy_default_span_attributes else default_span_attributes + span_attributes = {**(static_default_span_attributes or {}), **span_attributes} if event is None: event = {} if type is None and not parent_span_ids: @@ -4192,6 +4279,7 @@ def __init__( self.parent_object_type = parent_object_type self.parent_object_id = parent_object_id self.parent_compute_object_metadata_args = parent_compute_object_metadata_args + self.default_span_attributes = default_span_attributes # Merge propagated_event into event. The propagated_event data will get # propagated-and-merged into every subspan. @@ -4222,16 +4310,35 @@ def __init__( _EXEC_COUNTER += 1 exec_counter = _EXEC_COUNTER + base_span_attributes = dict( + **{"type": type, "name": name, **span_attributes}, + exec_counter=exec_counter, + ) internal_data: dict[str, Any] = dict( metrics=dict( start=start_time or time.time(), ), # Set type first, in case they override it in `span_attributes`. - span_attributes=dict(**{"type": type, "name": name, **span_attributes}, exec_counter=exec_counter), + span_attributes=base_span_attributes, created=datetime.datetime.now(datetime.timezone.utc).isoformat(), ) if caller_location: internal_data["context"] = caller_location + lazy_internal_data: dict[str, LazyValue[Any]] | None = None + if lazy_default_span_attributes: + + def compute_span_attributes() -> dict[str, Any]: + return { + **(lazy_default_span_attributes.get() or {}), + **base_span_attributes, + } + + lazy_internal_data = { + "span_attributes": LazyValue( + compute_span_attributes, + use_mutex=False, + ) + } # TODO: can be simplified after `event` is typed. id = event.pop("id", None) @@ -4255,7 +4362,7 @@ def __init__( # The first log is a replacement, but subsequent logs to the same span # object will be merges. self._is_merge = False - self.log_internal(event=event, internal_data=internal_data) + self.log_internal(event=event, internal_data=internal_data, lazy_internal_data=lazy_internal_data) self._is_merge = True @property @@ -4290,8 +4397,15 @@ def set_attributes( def log(self, **event: Any) -> None: return self.log_internal(event=event, internal_data=None) - def log_internal(self, event: dict[str, Any] | None = None, internal_data: dict[str, Any] | None = None) -> None: + def log_internal( + self, + event: dict[str, Any] | None = None, + internal_data: dict[str, Any] | None = None, + lazy_internal_data: dict[str, LazyValue[Any]] | None = None, + ) -> None: serializable_partial_record, lazy_partial_record = split_logging_data(event, internal_data) + if lazy_internal_data: + lazy_partial_record = {**lazy_partial_record, **lazy_internal_data} # We both check for serializability and round-trip `partial_record` # through JSON in order to create a "deep copy". This has the benefit of @@ -4328,14 +4442,15 @@ def log_internal(self, event: dict[str, Any] | None = None, internal_data: dict[ def compute_record() -> dict[str, Any]: exporter = _get_exporter() - return dict( - **serializable_partial_record, - **{k: v.get() for k, v in lazy_partial_record.items()}, - **exporter( + record = dict(serializable_partial_record) + record.update({k: v.get() for k, v in lazy_partial_record.items()}) + record.update( + exporter( object_type=self.parent_object_type, object_id=self.parent_object_id.get(), - ).object_id_fields(), + ).object_id_fields() ) + return record self.state.global_bg_logger().log(LazyValue(compute_record, use_mutex=False)) @@ -4379,6 +4494,7 @@ def start_span( ), name=name, type=type, + default_span_attributes=self.default_span_attributes, span_attributes=span_attributes, start_time=start_time, set_current=set_current, @@ -5213,11 +5329,13 @@ def __init__( async_flush: bool = True, compute_metadata_args: dict | None = None, link_args: dict | None = None, + default_span_attributes: DefaultSpanAttributes | None = None, state: BraintrustState | None = None, ): self._lazy_metadata = lazy_metadata self.async_flush = async_flush self._compute_metadata_args = compute_metadata_args + self._default_span_attributes = default_span_attributes self.last_start_time = time.time() self._lazy_id = LazyValue(lambda: self.id, use_mutex=False) self._called_start_span = False @@ -5413,6 +5531,7 @@ def _start_span_impl( name=name, type=type, default_root_type=SpanTypeAttribute.TASK, + default_span_attributes=self._default_span_attributes, span_attributes=span_attributes, start_time=start_time, set_current=set_current, diff --git a/py/src/braintrust/test_logger.py b/py/src/braintrust/test_logger.py index 3fc5c833c..dbf14ed33 100644 --- a/py/src/braintrust/test_logger.py +++ b/py/src/braintrust/test_logger.py @@ -862,6 +862,42 @@ class MetadataModel(BaseModel): assert logs[0]["metadata"] == {"foo": "bar"} +def test_init_logger_agent_sets_span_attribute(with_memory_logger, with_simulate_login, monkeypatch): + def post_json(object_type, args=None): + assert object_type == "api/agent/register" + assert args == { + "project_id": "test-project-id", + "agent_name": "support-agent", + } + return { + "agent": { + "id": "agent-id", + "name": "support-agent", + }, + "found_existing": False, + } + + monkeypatch.setattr(logger._state.app_conn(), "post_json", post_json) + + bt_logger = init_logger( + project="test-project", + project_id="test-project-id", + agent="support-agent", + ) + + parent = bt_logger.start_span(name="parent") + parent.start_span(name="child").end() + parent.start_span(name="override", span_attributes={"agent_id": "override-agent"}).end() + parent.end() + + logs = with_memory_logger.pop() + by_name = {log["span_attributes"]["name"]: log for log in logs} + + assert by_name["parent"]["span_attributes"]["agent_id"] == "agent-id" + assert by_name["child"]["span_attributes"]["agent_id"] == "agent-id" + assert by_name["override"]["span_attributes"]["agent_id"] == "override-agent" + + class _ModelDumpMetadata: def __init__(self, **values): self.values = values From 381463ff3e5cc24acb945dc6b768cc17a1ab9bbd Mon Sep 17 00:00:00 2001 From: Ankur Goyal Date: Thu, 11 Jun 2026 23:06:27 -0700 Subject: [PATCH 2/3] fix --- py/src/braintrust/logger.py | 50 +++++++++++++++++++++++++------- py/src/braintrust/test_logger.py | 14 +++++++-- 2 files changed, 51 insertions(+), 13 deletions(-) diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index ee0a0b77e..f60ed43b0 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -1838,10 +1838,27 @@ def compute_metadata(): ) -def _agent_register_endpoint_not_found(error: AugmentedHTTPError) -> bool: +def _http_error_status(error: AugmentedHTTPError) -> int | None: cause = error.__cause__ response = getattr(cause, "response", None) - return getattr(response, "status_code", None) == 404 + return getattr(response, "status_code", None) + + +def _agent_register_endpoint_not_found(error: AugmentedHTTPError) -> bool: + return _http_error_status(error) == 404 + + +def _agent_metadata_from_register_response(response: Mapping[str, Any]) -> ObjectMetadata | None: + resp_agent = response.get("agent") + if not isinstance(resp_agent, Mapping): + return None + + agent_id = resp_agent.get("id") + agent_name = resp_agent.get("name") + if not isinstance(agent_id, str) or not isinstance(agent_name, str): + return None + + return ObjectMetadata(id=agent_id, name=agent_name, full_info=dict(resp_agent)) def _register_agent_metadata(project_id: str, agent_name: str) -> ObjectMetadata: @@ -1875,13 +1892,25 @@ def _compute_logger_metadata( login() org_id = _state.org_id if project_id is None: - response = _state.app_conn().post_json( - "api/project/register", - { - "project_name": project_name or GLOBAL_PROJECT, - "org_id": org_id, - }, - ) + register_args: dict[str, Any] = { + "project_name": project_name or GLOBAL_PROJECT, + "org_id": org_id, + } + if agent_name is not None: + register_args["agent_name"] = agent_name + try: + response = _state.app_conn().post_json("api/project/register", register_args) + except AugmentedHTTPError as e: + if agent_name is None or _http_error_status(e) != 400: + raise + _logger.debug("Project registration did not accept agent_name; retrying without it") + response = _state.app_conn().post_json( + "api/project/register", + { + "project_name": project_name or GLOBAL_PROJECT, + "org_id": org_id, + }, + ) resp_project = response["project"] metadata = OrgProjectMetadata( org_id=org_id, @@ -1890,6 +1919,7 @@ def _compute_logger_metadata( name=resp_project["name"], full_info=resp_project, ), + agent=_agent_metadata_from_register_response(response), ) elif project_name is None: response = _state.app_conn().get_json("api/project", {"id": project_id}) @@ -1910,7 +1940,7 @@ def _compute_logger_metadata( full_info=dict(), ), ) - if agent_name is not None: + if agent_name is not None and metadata.agent is None: metadata.agent = _register_agent_metadata(metadata.project.id, agent_name) return metadata diff --git a/py/src/braintrust/test_logger.py b/py/src/braintrust/test_logger.py index dbf14ed33..8e0424166 100644 --- a/py/src/braintrust/test_logger.py +++ b/py/src/braintrust/test_logger.py @@ -863,13 +863,21 @@ class MetadataModel(BaseModel): def test_init_logger_agent_sets_span_attribute(with_memory_logger, with_simulate_login, monkeypatch): + post_json_calls = [] + def post_json(object_type, args=None): - assert object_type == "api/agent/register" + post_json_calls.append((object_type, args)) + assert object_type == "api/project/register" assert args == { - "project_id": "test-project-id", + "project_name": "test-project", + "org_id": "test-org-id", "agent_name": "support-agent", } return { + "project": { + "id": "test-project-id", + "name": "test-project", + }, "agent": { "id": "agent-id", "name": "support-agent", @@ -881,7 +889,6 @@ def post_json(object_type, args=None): bt_logger = init_logger( project="test-project", - project_id="test-project-id", agent="support-agent", ) @@ -896,6 +903,7 @@ def post_json(object_type, args=None): assert by_name["parent"]["span_attributes"]["agent_id"] == "agent-id" assert by_name["child"]["span_attributes"]["agent_id"] == "agent-id" assert by_name["override"]["span_attributes"]["agent_id"] == "override-agent" + assert len(post_json_calls) == 1 class _ModelDumpMetadata: From b19fbe5855ddd5e624fd19cd79db108ca7551e42 Mon Sep 17 00:00:00 2001 From: Ankur Goyal Date: Fri, 12 Jun 2026 23:54:50 -0700 Subject: [PATCH 3/3] bump --- py/src/braintrust/logger.py | 3 ++- py/src/braintrust/test_logger.py | 12 ++++++++++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index ce24f912d..383a05652 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -1947,8 +1947,9 @@ def init_logger( compute_metadata_args = dict( project_name=project, project_id=project_id, - agent_name=agent, ) + if agent is not None: + compute_metadata_args["agent_name"] = agent link_args = { "app_url": app_url, diff --git a/py/src/braintrust/test_logger.py b/py/src/braintrust/test_logger.py index cf6862c7d..753157160 100644 --- a/py/src/braintrust/test_logger.py +++ b/py/src/braintrust/test_logger.py @@ -979,6 +979,18 @@ def post_json(object_type, args=None): assert len(post_json_calls) == 1 +def test_init_logger_without_agent_omits_agent_from_exported_metadata(with_memory_logger, with_simulate_login): + from braintrust.span_identifier_v3 import SpanComponentsV3 + + bt_logger = init_logger(project="test-project") + + with bt_logger.start_span(name="parent") as span: + components = SpanComponentsV3.from_str(span.export()) + + assert components.compute_object_metadata_args is not None + assert "agent_name" not in components.compute_object_metadata_args + + class _ModelDumpMetadata: def __init__(self, **values): self.values = values