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
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -1036,13 +1040,15 @@ 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",
prompt="",
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. "
Expand All @@ -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()
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down
26 changes: 21 additions & 5 deletions py/src/braintrust/integrations/claude_agent_sdk/tracing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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")
Expand Down
Loading