From ad0a42f24f3f2efb91915aa6f2f60219479d36d7 Mon Sep 17 00:00:00 2001 From: Abhijeet Prasad Date: Mon, 28 Sep 2026 12:17:10 -0400 Subject: [PATCH 1/2] fix(claude-agent-sdk): preserve partial-mode usage Keep complete ResultMessage usage on the root span metadata and mark per-call output usage unknown when no final stream event arrives. Refs SDK-421 --- .../claude_agent_sdk/test_claude_agent_sdk.py | 11 +++++++++++ .../integrations/claude_agent_sdk/tracing.py | 18 +++++++++++++----- 2 files changed, 24 insertions(+), 5 deletions(-) diff --git a/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py b/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py index 2be9e56c0..53f1643ba 100644 --- a/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py +++ b/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py @@ -1036,6 +1036,7 @@ async def test_bundled_subagent_creates_task_span(memory_logger): cwd=REPO_ROOT, permission_mode="bypassPermissions", max_turns=8, + include_partial_messages=True, ) transport = make_cassette_transport( cassette_name="test_bundled_subagent_creates_task_span", @@ -1043,6 +1044,7 @@ async def test_bundled_subagent_creates_task_span(memory_logger): options=options, ) + result_message = None async with claude_agent_sdk.ClaudeSDKClient(options=options, transport=transport) as client: await client.query( "You must delegate this task to the bundled general-purpose agent. " @@ -1051,6 +1053,7 @@ async def test_bundled_subagent_creates_task_span(memory_logger): ) async for message in client.receive_response(): if type(message).__name__ == "ResultMessage": + result_message = message break spans = memory_logger.pop() @@ -1079,6 +1082,9 @@ async def test_bundled_subagent_creates_task_span(memory_logger): assert root_task_span["span_id"] in parents assert root_task_span.get("metadata", {}).get("task_events"), "Expected task events on root task span" + assert result_message is not None + assert root_task_span.get("metadata", {}).get("model_usage") == _aggregate_model_usage(result_message.model_usage) + assert not {"prompt_tokens", "completion_tokens", "tokens"}.intersection(root_task_span.get("metrics", {})) llm_spans = [s for s in spans if s["span_attributes"]["type"] == SpanTypeAttribute.LLM] _assert_llm_spans_have_time_to_first_token(llm_spans) @@ -1094,6 +1100,11 @@ async def test_bundled_subagent_creates_task_span(memory_logger): if any(subagent_span["span_id"] in llm_span["span_parents"] for subagent_span in subagent_spans) ] assert delegated_llm_spans, "Expected at least one delegated LLM span nested under a subagent task span" + for llm_span in delegated_llm_spans: + assert llm_span["metrics"].get("prompt_tokens", 0) > 0 + assert "completion_tokens" not in llm_span["metrics"] + assert "tokens" not in llm_span["metrics"] + assert llm_span.get("metadata", {}).get("usage_output_tokens_unknown") is True assert any( any(llm_span["span_id"] in tool_span["span_parents"] for llm_span in delegated_llm_spans) diff --git a/py/src/braintrust/integrations/claude_agent_sdk/tracing.py b/py/src/braintrust/integrations/claude_agent_sdk/tracing.py index 086088643..4f7a20d6e 100644 --- a/py/src/braintrust/integrations/claude_agent_sdk/tracing.py +++ b/py/src/braintrust/integrations/claude_agent_sdk/tracing.py @@ -901,11 +901,17 @@ def _handle_result(self, message: Any) -> None: if v is not None } result_metrics: dict[str, float] = {} - if not self._include_partial_messages: - raw_usage = getattr(message, "usage", None) - _, usage_metadata = extract_anthropic_usage(raw_usage) - result_metadata.update(usage_metadata) - aggregate_usage = _aggregate_model_usage(getattr(message, "model_usage", None)) + raw_usage = getattr(message, "usage", None) + _, usage_metadata = extract_anthropic_usage(raw_usage) + result_metadata.update(usage_metadata) + aggregate_usage = _aggregate_model_usage(getattr(message, "model_usage", None)) + if self._include_partial_messages: + # Keep the complete turn-level total visible without adding it as + # another token metric alongside the per-call LLM span metrics. + complete_usage = aggregate_usage or _copy_usage(raw_usage) + if complete_usage: + result_metadata["model_usage"] = complete_usage + else: usage = aggregate_usage or _copy_usage(raw_usage) result_metrics, _ = extract_anthropic_usage(usage) if result_metadata or result_metrics: @@ -1045,6 +1051,8 @@ def _log_assistant_usage( has_final_output = message_id in self._final_output_usage_message_ids if message_id else False metrics, _ = extract_anthropic_usage(usage, include_output=has_final_output) _, metadata = extract_anthropic_usage(raw_message_usage, include_output=False) + if not has_final_output and "prompt_tokens" in metrics: + metadata["usage_output_tokens_unknown"] = True ctx.llm_span.log(metrics=metrics or None, metadata=metadata or None) def _process_task_event(self, message: Any, agent_span_export: str | None) -> None: From d1ca0419870fa5795a7848700842d4bb51a401e6 Mon Sep 17 00:00:00 2001 From: Abhijeet Prasad Date: Mon, 28 Sep 2026 14:55:17 -0400 Subject: [PATCH 2/2] better output marker logic --- .../claude_agent_sdk/test_claude_agent_sdk.py | 4 ++++ .../integrations/claude_agent_sdk/tracing.py | 12 ++++++++++-- 2 files changed, 14 insertions(+), 2 deletions(-) diff --git a/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py b/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py index 53f1643ba..1404c4df1 100644 --- a/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py +++ b/py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py @@ -263,6 +263,10 @@ async def calculator_handler(args): assert len(ordered_llm_spans) == len(expected_partial_usage) for llm_span, usage in zip(ordered_llm_spans, expected_partial_usage, strict=True): assert llm_span["metrics"] | _metrics_from_exact_anthropic_usage(usage) == llm_span["metrics"] + llm_spans_with_output = [span for span in llm_spans if "completion_tokens" in span["metrics"]] + assert llm_spans_with_output + for llm_span in llm_spans_with_output: + assert "usage_output_tokens_unknown" not in llm_span.get("metadata", {}) elif _sdk_version_at_least("0.1.11"): expected_usage_by_message_id: dict[str, dict[str, Any]] = {} for message in received_messages: diff --git a/py/src/braintrust/integrations/claude_agent_sdk/tracing.py b/py/src/braintrust/integrations/claude_agent_sdk/tracing.py index 4f7a20d6e..d3c026d04 100644 --- a/py/src/braintrust/integrations/claude_agent_sdk/tracing.py +++ b/py/src/braintrust/integrations/claude_agent_sdk/tracing.py @@ -744,6 +744,7 @@ def cleanup(self) -> None: """End all open LLM spans, TASK spans, and TOOL spans; clear thread-local.""" for ctx in self._contexts.values(): if ctx.llm_span: + self._mark_output_usage_unknown(ctx) ctx.llm_span.end() ctx.llm_span = None if ctx.task_span: @@ -884,6 +885,8 @@ def _handle_stream_event(self, message: Any) -> None: def _handle_result(self, message: Any) -> None: self._active_key = None + for ctx in self._contexts.values(): + self._mark_output_usage_unknown(ctx) result_value = getattr(message, "result", None) if result_value is not None: self._result_output = result_value @@ -1014,6 +1017,7 @@ def _start_or_merge_llm_span( first_token_time = time.time() if ctx.llm_span: + self._mark_output_usage_unknown(ctx) ctx.llm_span.end(end_time=resolved_start) final_content, span = _create_llm_span_for_messages( @@ -1051,10 +1055,14 @@ def _log_assistant_usage( has_final_output = message_id in self._final_output_usage_message_ids if message_id else False metrics, _ = extract_anthropic_usage(usage, include_output=has_final_output) _, metadata = extract_anthropic_usage(raw_message_usage, include_output=False) - if not has_final_output and "prompt_tokens" in metrics: - metadata["usage_output_tokens_unknown"] = True ctx.llm_span.log(metrics=metrics or None, metadata=metadata or None) + def _mark_output_usage_unknown(self, ctx: _AgentContext) -> None: + """Mark missing output only once an LLM span is finalized.""" + if ctx.llm_span is None or ctx.llm_message_id in self._final_output_usage_message_ids: + return + ctx.llm_span.log(metadata={"usage_output_tokens_unknown": True}) + def _process_task_event(self, message: Any, agent_span_export: str | None) -> None: """Handle TaskStarted / TaskProgress / TaskNotification system messages.""" task_id = _msg_field(message, "task_id")