Skip to content

Commit 495a7af

Browse files
authored
fix(claude-agent-sdk): preserve partial-mode usage (#810)
resolves https://linear.app/braintrustdata/issue/SDK-421
1 parent 57af883 commit 495a7af

2 files changed

Lines changed: 36 additions & 5 deletions

File tree

‎py/src/braintrust/integrations/claude_agent_sdk/test_claude_agent_sdk.py‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -263,6 +263,10 @@ async def calculator_handler(args):
263263
assert len(ordered_llm_spans) == len(expected_partial_usage)
264264
for llm_span, usage in zip(ordered_llm_spans, expected_partial_usage, strict=True):
265265
assert llm_span["metrics"] | _metrics_from_exact_anthropic_usage(usage) == llm_span["metrics"]
266+
llm_spans_with_output = [span for span in llm_spans if "completion_tokens" in span["metrics"]]
267+
assert llm_spans_with_output
268+
for llm_span in llm_spans_with_output:
269+
assert "usage_output_tokens_unknown" not in llm_span.get("metadata", {})
266270
elif _sdk_version_at_least("0.1.11"):
267271
expected_usage_by_message_id: dict[str, dict[str, Any]] = {}
268272
for message in received_messages:
@@ -1036,13 +1040,15 @@ async def test_bundled_subagent_creates_task_span(memory_logger):
10361040
cwd=REPO_ROOT,
10371041
permission_mode="bypassPermissions",
10381042
max_turns=8,
1043+
include_partial_messages=True,
10391044
)
10401045
transport = make_cassette_transport(
10411046
cassette_name="test_bundled_subagent_creates_task_span",
10421047
prompt="",
10431048
options=options,
10441049
)
10451050

1051+
result_message = None
10461052
async with claude_agent_sdk.ClaudeSDKClient(options=options, transport=transport) as client:
10471053
await client.query(
10481054
"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):
10511057
)
10521058
async for message in client.receive_response():
10531059
if type(message).__name__ == "ResultMessage":
1060+
result_message = message
10541061
break
10551062

10561063
spans = memory_logger.pop()
@@ -1079,6 +1086,9 @@ async def test_bundled_subagent_creates_task_span(memory_logger):
10791086
assert root_task_span["span_id"] in parents
10801087

10811088
assert root_task_span.get("metadata", {}).get("task_events"), "Expected task events on root task span"
1089+
assert result_message is not None
1090+
assert root_task_span.get("metadata", {}).get("model_usage") == _aggregate_model_usage(result_message.model_usage)
1091+
assert not {"prompt_tokens", "completion_tokens", "tokens"}.intersection(root_task_span.get("metrics", {}))
10821092

10831093
llm_spans = [s for s in spans if s["span_attributes"]["type"] == SpanTypeAttribute.LLM]
10841094
_assert_llm_spans_have_time_to_first_token(llm_spans)
@@ -1094,6 +1104,11 @@ async def test_bundled_subagent_creates_task_span(memory_logger):
10941104
if any(subagent_span["span_id"] in llm_span["span_parents"] for subagent_span in subagent_spans)
10951105
]
10961106
assert delegated_llm_spans, "Expected at least one delegated LLM span nested under a subagent task span"
1107+
for llm_span in delegated_llm_spans:
1108+
assert llm_span["metrics"].get("prompt_tokens", 0) > 0
1109+
assert "completion_tokens" not in llm_span["metrics"]
1110+
assert "tokens" not in llm_span["metrics"]
1111+
assert llm_span.get("metadata", {}).get("usage_output_tokens_unknown") is True
10971112

10981113
assert any(
10991114
any(llm_span["span_id"] in tool_span["span_parents"] for llm_span in delegated_llm_spans)

‎py/src/braintrust/integrations/claude_agent_sdk/tracing.py‎

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -744,6 +744,7 @@ def cleanup(self) -> None:
744744
"""End all open LLM spans, TASK spans, and TOOL spans; clear thread-local."""
745745
for ctx in self._contexts.values():
746746
if ctx.llm_span:
747+
self._mark_output_usage_unknown(ctx)
747748
ctx.llm_span.end()
748749
ctx.llm_span = None
749750
if ctx.task_span:
@@ -884,6 +885,8 @@ def _handle_stream_event(self, message: Any) -> None:
884885

885886
def _handle_result(self, message: Any) -> None:
886887
self._active_key = None
888+
for ctx in self._contexts.values():
889+
self._mark_output_usage_unknown(ctx)
887890
result_value = getattr(message, "result", None)
888891
if result_value is not None:
889892
self._result_output = result_value
@@ -901,11 +904,17 @@ def _handle_result(self, message: Any) -> None:
901904
if v is not None
902905
}
903906
result_metrics: dict[str, float] = {}
904-
if not self._include_partial_messages:
905-
raw_usage = getattr(message, "usage", None)
906-
_, usage_metadata = extract_anthropic_usage(raw_usage)
907-
result_metadata.update(usage_metadata)
908-
aggregate_usage = _aggregate_model_usage(getattr(message, "model_usage", None))
907+
raw_usage = getattr(message, "usage", None)
908+
_, usage_metadata = extract_anthropic_usage(raw_usage)
909+
result_metadata.update(usage_metadata)
910+
aggregate_usage = _aggregate_model_usage(getattr(message, "model_usage", None))
911+
if self._include_partial_messages:
912+
# Keep the complete turn-level total visible without adding it as
913+
# another token metric alongside the per-call LLM span metrics.
914+
complete_usage = aggregate_usage or _copy_usage(raw_usage)
915+
if complete_usage:
916+
result_metadata["model_usage"] = complete_usage
917+
else:
909918
usage = aggregate_usage or _copy_usage(raw_usage)
910919
result_metrics, _ = extract_anthropic_usage(usage)
911920
if result_metadata or result_metrics:
@@ -1008,6 +1017,7 @@ def _start_or_merge_llm_span(
10081017
first_token_time = time.time()
10091018

10101019
if ctx.llm_span:
1020+
self._mark_output_usage_unknown(ctx)
10111021
ctx.llm_span.end(end_time=resolved_start)
10121022

10131023
final_content, span = _create_llm_span_for_messages(
@@ -1047,6 +1057,12 @@ def _log_assistant_usage(
10471057
_, metadata = extract_anthropic_usage(raw_message_usage, include_output=False)
10481058
ctx.llm_span.log(metrics=metrics or None, metadata=metadata or None)
10491059

1060+
def _mark_output_usage_unknown(self, ctx: _AgentContext) -> None:
1061+
"""Mark missing output only once an LLM span is finalized."""
1062+
if ctx.llm_span is None or ctx.llm_message_id in self._final_output_usage_message_ids:
1063+
return
1064+
ctx.llm_span.log(metadata={"usage_output_tokens_unknown": True})
1065+
10501066
def _process_task_event(self, message: Any, agent_span_export: str | None) -> None:
10511067
"""Handle TaskStarted / TaskProgress / TaskNotification system messages."""
10521068
task_id = _msg_field(message, "task_id")

0 commit comments

Comments
 (0)