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..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: @@ -1036,6 +1040,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 +1048,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 +1057,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 +1086,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 +1104,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..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 @@ -901,11 +904,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: @@ -1008,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( @@ -1047,6 +1057,12 @@ def _log_assistant_usage( _, metadata = extract_anthropic_usage(raw_message_usage, include_output=False) 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")