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
43 changes: 34 additions & 9 deletions py/src/braintrust/integrations/anthropic/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,11 +53,18 @@ def _set_numeric_metric(metrics: dict[str, float], name: str, value: Any) -> Non
metrics[name] = float(value)


def extract_anthropic_usage(usage: Any) -> tuple[dict[str, float], dict[str, Any]]:
def extract_anthropic_usage(
usage: Any,
*,
include_output: bool = True,
include_legacy_cache_creation: bool = True,
) -> tuple[dict[str, float], dict[str, Any]]:
"""Extract normalized metrics and allowlisted metadata from Anthropic usage.

Numeric usage fields are converted into Braintrust metrics. Allowlisted
non-numeric fields are attached as span metadata with a ``usage_`` prefix.
Anthropic's per-TTL cache creation breakdown supersedes the legacy aggregate
metric, and totals are emitted only when completion usage is known.
"""
usage = _try_to_dict(usage)
if usage is None:
Expand All @@ -66,6 +73,10 @@ def extract_anthropic_usage(usage: Any) -> tuple[dict[str, float], dict[str, Any
metrics: dict[str, float] = {}
metadata: dict[str, Any] = {}
for source_name, metric_name in _ANTHROPIC_USAGE_METRIC_FIELDS:
if metric_name == "completion_tokens" and not include_output:
continue
if metric_name == "prompt_cache_creation_tokens" and not include_legacy_cache_creation:
continue
_set_numeric_metric(metrics, metric_name, usage.get(source_name))

cache_creation = _try_to_dict(usage.get("cache_creation"))
Expand All @@ -77,22 +88,36 @@ def extract_anthropic_usage(usage: Any) -> tuple[dict[str, float], dict[str, Any
metrics[metric_name] = float(value)
cache_creation_breakdown.append(float(value))

if cache_creation_breakdown:
metrics.pop("prompt_cache_creation_tokens", None)

server_tool_use = _try_to_dict(usage.get("server_tool_use"))
if server_tool_use is not None:
for source_name, value in server_tool_use.items():
_set_numeric_metric(metrics, f"server_tool_use_{source_name}", value)

if "prompt_cache_creation_tokens" not in metrics and cache_creation_breakdown:
metrics["prompt_cache_creation_tokens"] = sum(cache_creation_breakdown)

if metrics:
has_prompt_usage = any(
metric_name in metrics
for metric_name in (
"prompt_tokens",
"prompt_cached_tokens",
"prompt_cache_creation_tokens",
"prompt_cache_creation_5m_tokens",
"prompt_cache_creation_1h_tokens",
)
)
if has_prompt_usage:
effective_cache_creation_tokens = (
sum(cache_creation_breakdown)
if cache_creation_breakdown
else metrics.get("prompt_cache_creation_tokens", 0)
)
total_prompt_tokens = (
metrics.get("prompt_tokens", 0)
+ metrics.get("prompt_cached_tokens", 0)
+ metrics.get("prompt_cache_creation_tokens", 0)
metrics.get("prompt_tokens", 0) + metrics.get("prompt_cached_tokens", 0) + effective_cache_creation_tokens
)
metrics["prompt_tokens"] = total_prompt_tokens
metrics["tokens"] = total_prompt_tokens + metrics.get("completion_tokens", 0)
if "completion_tokens" in metrics:
metrics["tokens"] = total_prompt_tokens + metrics["completion_tokens"]

for name, value in usage.items():
if name in _ANTHROPIC_USAGE_METADATA_FIELDS and value is not None:
Expand Down
41 changes: 32 additions & 9 deletions py/src/braintrust/integrations/anthropic/test_anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -376,7 +376,6 @@ def to_dict(self):
"prompt_tokens": 21.0,
"completion_tokens": 7.0,
"prompt_cached_tokens": 3.0,
"prompt_cache_creation_tokens": 7.0,
"prompt_cache_creation_5m_tokens": 2.0,
"prompt_cache_creation_1h_tokens": 5.0,
"server_tool_use_web_search_requests": 2.0,
Expand Down Expand Up @@ -410,7 +409,7 @@ def test_anthropic_messages_create_prompt_cache_5m_metrics(memory_logger):

span = find_span_by_name(memory_logger.pop(), "anthropic.messages.create")
assert span["output"]["role"] == response.role
assert span["metrics"]["prompt_cache_creation_tokens"] == response.usage.cache_creation_input_tokens
assert "prompt_cache_creation_tokens" not in span["metrics"]
assert (
span["metrics"]["prompt_cache_creation_5m_tokens"] == response.usage.cache_creation.ephemeral_5m_input_tokens
)
Expand Down Expand Up @@ -442,7 +441,7 @@ def test_anthropic_messages_create_prompt_cache_1h_metrics(memory_logger):

span = find_span_by_name(memory_logger.pop(), "anthropic.messages.create")
assert span["output"]["role"] == response.role
assert span["metrics"]["prompt_cache_creation_tokens"] == response.usage.cache_creation_input_tokens
assert "prompt_cache_creation_tokens" not in span["metrics"]
assert (
span["metrics"]["prompt_cache_creation_5m_tokens"] == response.usage.cache_creation.ephemeral_5m_input_tokens
)
Expand Down Expand Up @@ -851,7 +850,7 @@ async def test_anthropic_messages_streaming_async(memory_logger):
assert metrics["completion_tokens"] == usage.output_tokens
assert metrics["tokens"] == usage.input_tokens + usage.output_tokens
assert metrics["prompt_cached_tokens"] == usage.cache_read_input_tokens
assert metrics["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
_assert_cache_creation_metrics(metrics, usage)
assert log["metadata"]["model"] == MODEL
assert log["metadata"]["max_tokens"] == 1024

Expand Down Expand Up @@ -933,7 +932,7 @@ def test_anthropic_messages_streaming_sync(memory_logger):
assert log["metrics"]["completion_tokens"] == usage.output_tokens
assert log["metrics"]["tokens"] == usage.input_tokens + usage.output_tokens
assert log["metrics"]["prompt_cached_tokens"] == usage.cache_read_input_tokens
assert log["metrics"]["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
_assert_cache_creation_metrics(log["metrics"], usage)


@pytest.mark.vcr
Expand Down Expand Up @@ -973,7 +972,7 @@ def test_anthropic_messages_streaming_sync_text_stream(memory_logger):
assert log["metrics"]["completion_tokens"] == usage.output_tokens
assert log["metrics"]["tokens"] == usage.input_tokens + usage.output_tokens
assert log["metrics"]["prompt_cached_tokens"] == usage.cache_read_input_tokens
assert log["metrics"]["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
_assert_cache_creation_metrics(log["metrics"], usage)


@pytest.mark.vcr
Expand Down Expand Up @@ -1014,7 +1013,7 @@ async def test_anthropic_messages_streaming_async_text_stream(memory_logger):
assert log["metrics"]["completion_tokens"] == usage.output_tokens
assert log["metrics"]["tokens"] == usage.input_tokens + usage.output_tokens
assert log["metrics"]["prompt_cached_tokens"] == usage.cache_read_input_tokens
assert log["metrics"]["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
_assert_cache_creation_metrics(log["metrics"], usage)


@pytest.mark.vcr
Expand Down Expand Up @@ -1119,6 +1118,30 @@ def test_anthropic_messages_sync_server_tool_spans(memory_logger):
assert tool_span["root_span_id"] == llm_span["root_span_id"]


def _assert_cache_creation_metrics(metrics, usage):
cache_creation = getattr(usage, "cache_creation", None)
if cache_creation is None:
assert metrics["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
return

if isinstance(cache_creation, dict):
ephemeral_5m = cache_creation.get("ephemeral_5m_input_tokens")
ephemeral_1h = cache_creation.get("ephemeral_1h_input_tokens")
else:
ephemeral_5m = getattr(cache_creation, "ephemeral_5m_input_tokens", None)
ephemeral_1h = getattr(cache_creation, "ephemeral_1h_input_tokens", None)

if ephemeral_5m is None and ephemeral_1h is None:
assert metrics["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
return

assert "prompt_cache_creation_tokens" not in metrics
if ephemeral_5m is not None:
assert metrics["prompt_cache_creation_5m_tokens"] == ephemeral_5m
if ephemeral_1h is not None:
assert metrics["prompt_cache_creation_1h_tokens"] == ephemeral_1h


def _assert_metrics_are_valid(metrics, start, end):
assert metrics["tokens"] > 0
assert metrics["prompt_tokens"] > 0
Expand Down Expand Up @@ -1467,7 +1490,7 @@ def test_setup_creates_spans(memory_logger):
usage.input_tokens + usage.cache_read_input_tokens + usage.cache_creation_input_tokens
)
assert metrics["completion_tokens"] == usage.output_tokens
assert metrics["prompt_cache_creation_tokens"] == usage.cache_creation_input_tokens
assert "prompt_cache_creation_tokens" not in metrics
assert metrics["prompt_cache_creation_5m_tokens"] == ephemeral_5m
assert metrics["prompt_cache_creation_1h_tokens"] == ephemeral_1h
assert "service_tier" not in metrics
Expand Down Expand Up @@ -1498,7 +1521,7 @@ def test_extract_anthropic_usage_preserves_nested_numeric_fields():
assert metrics["prompt_tokens"] == 15
assert metrics["completion_tokens"] == 12
assert metrics["tokens"] == 27
assert metrics["prompt_cache_creation_tokens"] == 7
assert "prompt_cache_creation_tokens" not in metrics
assert metrics["prompt_cache_creation_5m_tokens"] == 3
assert metrics["prompt_cache_creation_1h_tokens"] == 4
assert metrics["server_tool_use_web_search_requests"] == 2
Expand Down
13 changes: 10 additions & 3 deletions py/src/braintrust/integrations/claude_agent_sdk/_test_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,8 @@ def _normalize_for_match(value: Any) -> Any:
return [_normalize_for_match(item) for item in value]
if isinstance(value, dict):
normalized = {key: _normalize_for_match(item) for key, item in value.items()}
if normalized.get("type") == "user" and isinstance(normalized.get("session_id"), str):
normalized["session_id"] = "SESSION_ID"
if normalized.get("type") == "control_request" and isinstance(normalized.get("request_id"), str):
normalized["request_id"] = "CONTROL_REQUEST_ID"
return normalized
Expand Down Expand Up @@ -116,8 +118,13 @@ def _compact_initialize_message_for_storage(value: dict[str, Any]) -> dict[str,
return value

compact_result: dict[str, Any] = {}
if "account" in result:
compact_result["account"] = result["account"]
account = result.get("account")
if isinstance(account, dict):
compact_result["account"] = {
key: account[key]
for key in ("apiKeySource", "apiProvider", "subscriptionType", "tokenSource")
if key in account
}

for key in ("available_output_styles", "commands", "models", "agents"):
if key in result:
Expand Down Expand Up @@ -185,7 +192,7 @@ def _sanitize_url_string(value: str) -> str:
return urlunsplit((parts.scheme, parts.netloc, parts.path, urlencode(query, doseq=True), parts.fragment))


_PATH_RE = re.compile(r"(?:(?:/Users|/home)/[^\s\"']+|[A-Za-z]:\\\\[^\s\"']+)")
_PATH_RE = re.compile(r"(?:(?:/Users|/home|/private/(?:tmp|var)|/tmp)/[^\s\"']+|[A-Za-z]:\\\\[^\s\"']+)")
_AUTH_BEARER_RE = re.compile(r"Bearer\s+[A-Za-z0-9._-]+")
_API_KEY_RE = re.compile(r"\bsk-[A-Za-z0-9_-]+\b")
_SENSITIVE_FIELDS = {
Expand Down
Loading