Skip to content
Open
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
7 changes: 7 additions & 0 deletions src/google/adk/models/anthropic_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -709,6 +709,7 @@ def message_to_generate_content_response(
role="model",
parts=parts,
),
model_version=message.model,
usage_metadata=usage_metadata,
finish_reason=to_google_genai_finish_reason(message.stop_reason),
)
Expand Down Expand Up @@ -1012,13 +1013,18 @@ async def _generate_content_streaming(
cached_input_tokens: int | None = None
cache_creation_tokens: int | None = None
stop_reason: Optional[anthropic_types.StopReason] = None
model_version: str | None = None

async for event in raw_stream:
if event.type == "message_start":
input_tokens = _extract_prompt_token_count(event.message.usage)
output_tokens = event.message.usage.output_tokens
thinking_tokens = _extract_thinking_token_count(event.message.usage)
cached_input_tokens = _extract_cached_token_count(event.message.usage)
# Guard on str: keeps partial/mock streams that omit the field from
# poisoning model_version with a non-string.
if isinstance(event.message.model, str):
model_version = event.message.model
cache_creation_tokens = _extract_cache_creation_token_count(
event.message.usage
)
Expand Down Expand Up @@ -1149,6 +1155,7 @@ async def _generate_content_streaming(

yield LlmResponse(
content=types.Content(role="model", parts=all_parts),
model_version=model_version,
usage_metadata=usage_metadata,
finish_reason=to_google_genai_finish_reason(stop_reason),
partial=False,
Expand Down
85 changes: 83 additions & 2 deletions tests/unittests/models/test_anthropic_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@
import json
import os
import re
import sys
from unittest import mock
from unittest.mock import AsyncMock
from unittest.mock import MagicMock
Expand All @@ -38,7 +37,6 @@
from google.adk.models.llm_request import LlmRequest
from google.adk.models.llm_response import LlmResponse
from google.genai import types
from google.genai import version as genai_version
from google.genai.types import Content
from google.genai.types import Part
import httpx
Expand Down Expand Up @@ -3457,3 +3455,86 @@ def test_anthropic_config_allows_thinking_budget_without_thinking_level():
assert config.effort == "high"
assert config.thinking_config.thinking_budget == 2048
assert config.thinking_config.thinking_level is None


def test_message_to_generate_content_response_sets_model_version():
"""The served snapshot id must land in LlmResponse.model_version (#6847)."""
from google.adk.models.anthropic_llm import (
message_to_generate_content_response,
)

message = anthropic_types.Message(
id="msg_model_version",
content=[
anthropic_types.TextBlock(text="ok", type="text", citations=None)
],
model="claude-sonnet-4-20250514",
role="assistant",
stop_reason="end_turn",
stop_sequence=None,
type="message",
usage=anthropic_types.Usage(
input_tokens=1,
output_tokens=2,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
server_tool_use=None,
service_tier=None,
),
)

response = message_to_generate_content_response(message)

assert response.model_version == "claude-sonnet-4-20250514"


@pytest.mark.asyncio
async def test_streaming_final_response_carries_model_version():
"""Streaming: message_start's resolved model id reaches the final yield."""
llm = AnthropicLlm(model="claude-sonnet-4-20250514")

events = [
MagicMock(
type="message_start",
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=10, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
index=0,
content_block=anthropic_types.TextBlock(text="", type="text"),
),
MagicMock(
type="content_block_delta",
index=0,
delta=anthropic_types.TextDelta(text="Hi!", type="text_delta"),
),
MagicMock(type="content_block_stop", index=0),
MagicMock(
type="message_delta",
delta=MagicMock(stop_reason="end_turn"),
usage=MagicMock(output_tokens=3),
),
MagicMock(type="message_stop"),
]

mock_client = MagicMock()
mock_client.messages.create = AsyncMock(
return_value=_make_mock_stream_events(events)
)

llm_request = LlmRequest(
model="claude-sonnet-4-20250514",
contents=[Content(role="user", parts=[Part.from_text(text="Hi")])],
)

with mock.patch.object(llm, "_anthropic_client", mock_client):
responses = [
r async for r in llm.generate_content_async(llm_request, stream=True)
]

final = responses[-1]
assert final.partial is False
assert final.model_version == "claude-sonnet-4-20250514"