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
6 changes: 6 additions & 0 deletions src/google/adk/models/anthropic_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -711,6 +711,7 @@ def message_to_generate_content_response(
),
usage_metadata=usage_metadata,
finish_reason=to_google_genai_finish_reason(message.stop_reason),
model_version=message.model,
)


Expand Down Expand Up @@ -1012,6 +1013,7 @@ 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: Optional[str] = None

async for event in raw_stream:
if event.type == "message_start":
Expand All @@ -1022,6 +1024,7 @@ async def _generate_content_streaming(
cache_creation_tokens = _extract_cache_creation_token_count(
event.message.usage
)
model_version = event.message.model

elif event.type == "content_block_start":
block = event.content_block
Expand Down Expand Up @@ -1055,6 +1058,7 @@ async def _generate_content_streaming(
role="model",
parts=[types.Part(text=delta.thinking, thought=True)],
),
model_version=model_version,
partial=True,
)
elif isinstance(delta, anthropic_types.SignatureDelta):
Expand All @@ -1079,6 +1083,7 @@ async def _generate_content_streaming(
role="model",
parts=[types.Part.from_text(text=delta.text)],
),
model_version=model_version,
partial=True,
)
elif isinstance(delta, anthropic_types.InputJSONDelta):
Expand Down Expand Up @@ -1151,6 +1156,7 @@ async def _generate_content_streaming(
content=types.Content(role="model", parts=all_parts),
usage_metadata=usage_metadata,
finish_reason=to_google_genai_finish_reason(stop_reason),
model_version=model_version,
partial=False,
)

Expand Down
150 changes: 139 additions & 11 deletions tests/unittests/models/test_anthropic_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -1335,7 +1335,10 @@ async def test_streaming_text_yields_partial_and_final():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=10, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=10, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -1400,7 +1403,10 @@ async def test_streaming_tool_use_yields_function_call():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=20, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=20, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -1483,7 +1489,10 @@ async def test_streaming_passes_stream_true_to_create():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=5, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=5, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -1821,6 +1830,40 @@ def test_message_to_generate_content_response_reports_cache_read_tokens():
assert response.usage_metadata.cached_content_token_count == 75


def test_message_to_generate_content_response_sets_model_version():
"""LlmResponse.model_version reflects the resolved snapshot Anthropic served.

The requested model can be an alias (e.g. "claude-sonnet-4-5"); the
response's `model` field is the concrete snapshot that actually served it,
and that is what should end up in model_version.
"""
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="hi", 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=100,
output_tokens=20,
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"


def test_message_to_generate_content_response_no_cache_read_tokens():
"""Absent cache_read_input_tokens yields cached_content_token_count=None."""
from google.adk.models.anthropic_llm import message_to_generate_content_response
Expand Down Expand Up @@ -2052,12 +2095,13 @@ async def test_streaming_reports_cache_creation_tokens():
MagicMock(
type="message_start",
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(
input_tokens=100,
output_tokens=0,
cache_creation_input_tokens=50,
cache_read_input_tokens=0,
)
),
),
),
MagicMock(
Expand Down Expand Up @@ -2101,6 +2145,71 @@ async def test_streaming_reports_cache_creation_tokens():
assert "usage_metadata" in dumped


async def test_streaming_sets_model_version():
"""model_version is set on every streamed response, partials included.

message_start carries the resolved snapshot before any content arrives, so
every LlmResponse yielded afterwards -- partial deltas and the final
aggregated response alike -- should carry it, matching the parity
lite_llm.py already has for its own partial yields.
"""
llm = AnthropicLlm(model="claude-sonnet-4-20250514")

events = [
MagicMock(
type="message_start",
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(
input_tokens=100,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_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=20),
),
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)
]

assert len(responses) == 2
partial_response, final_response = responses
assert partial_response.partial is True
assert partial_response.model_version == "claude-sonnet-4-20250514"
assert final_response.partial is False
assert final_response.model_version == "claude-sonnet-4-20250514"


def test_part_to_message_block_thinking_roundtrip():
"""Part with thought=True and signature creates ThinkingBlockParam."""
part = Part(
Expand Down Expand Up @@ -2229,7 +2338,10 @@ async def test_streaming_thinking_yields_partial_and_final():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=15, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=15, output_tokens=0),
),
),
# Thinking block start
MagicMock(
Expand Down Expand Up @@ -2337,12 +2449,13 @@ async def test_streaming_reports_thinking_tokens_disjoint_from_candidates():
MagicMock(
type="message_start",
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=anthropic_types.Usage(
input_tokens=15,
output_tokens=0,
cache_read_input_tokens=5,
cache_creation_input_tokens=0,
)
),
),
),
MagicMock(
Expand Down Expand Up @@ -2423,7 +2536,10 @@ async def test_streaming_thinking_captures_signature_delta():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=15, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=15, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -2495,7 +2611,10 @@ async def test_streaming_passes_thinking_param():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=5, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=5, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -2549,7 +2668,10 @@ async def test_streaming_redacted_thinking_block_preserved_in_final():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=8, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=8, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -2814,7 +2936,10 @@ async def test_streaming_no_system_instruction_passes_not_given():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=1, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=1, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down Expand Up @@ -3212,7 +3337,10 @@ async def test_streaming_sets_finish_reason():
events = [
MagicMock(
type="message_start",
message=MagicMock(usage=MagicMock(input_tokens=5, output_tokens=0)),
message=MagicMock(
model="claude-sonnet-4-20250514",
usage=MagicMock(input_tokens=5, output_tokens=0),
),
),
MagicMock(
type="content_block_start",
Expand Down