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
39 changes: 32 additions & 7 deletions src/google/adk/models/lite_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2847,19 +2847,22 @@ def _has_meaningful_signal(message: Message | Delta | None) -> bool:
func_args = function_obj.get("arguments")
func_index = tool_call.get("index", idx)
tool_call_id = tool_call.get("id")
thought_signature = _extract_thought_signature_from_tool_call(
tool_call
)

# Ignore empty chunks that don't carry any information.
if not func_name and not func_args:
# Ignore empty chunks that don't carry any information. A chunk
# with only a signature still counts: the signature belongs to the
# call another chunk names.
if not func_name and not func_args and not thought_signature:
continue

yield FunctionChunk(
id=tool_call_id,
name=func_name,
args=func_args,
index=func_index,
thought_signature=_extract_thought_signature_from_tool_call(
tool_call
),
thought_signature=thought_signature,
), finish_reason

if finish_reason and not (message_content or tool_calls or reasoning_parts):
Expand Down Expand Up @@ -3943,6 +3946,8 @@ async def generate_content_async(
function_calls: dict[int, dict[str, Any]] = (
{}
) # index -> {name, args_parts, id, thought_signature}
# Signatures that arrived before any chunk naming their call.
pending_signatures: dict[int, bytes] = {}
tool_call_trackers: Dict[int, _BraceDepthTracker] = {}
completion_args["stream"] = True
completion_args["stream_options"] = {"include_usage": True}
Expand Down Expand Up @@ -4080,6 +4085,7 @@ def _reset_stream_buffers() -> None:
text_parts.clear()
reasoning_parts = []
function_calls.clear()
pending_signatures.clear()
tool_call_trackers.clear()
# The reason belongs to the segment just finalized; carrying it into
# the next one would stamp the wrong reason on the next response.
Expand All @@ -4106,14 +4112,33 @@ def _reset_stream_buffers() -> None:
for chunk, finish_reason in _model_response_to_chunk(part):
if finish_reason:
last_finish_reason = finish_reason
if isinstance(chunk, FunctionChunk):
if (
isinstance(chunk, FunctionChunk)
and not chunk.name
and not chunk.args
and chunk.thought_signature
):
# Only a signature. Opening a call for it would send the model a
# call with no name, so attach it to the call at its own index
# (the fallback index may already point past that call), or hold
# it for the call that starts there next.
signed_call = function_calls.get(
chunk.index if chunk.index is not None else fallback_index
)
if signed_call is not None:
signed_call["thought_signature"] = chunk.thought_signature
else:
pending_signatures[chunk.index or fallback_index] = (
chunk.thought_signature
)
elif isinstance(chunk, FunctionChunk):
index = chunk.index or fallback_index
if index not in function_calls:
function_calls[index] = {
"name": "",
"args_parts": [],
"id": None,
"thought_signature": None,
"thought_signature": pending_signatures.pop(index, None),
}

if chunk.name:
Expand Down
83 changes: 83 additions & 0 deletions tests/unittests/models/test_litellm.py
Original file line number Diff line number Diff line change
Expand Up @@ -4444,6 +4444,89 @@ async def test_streaming_parallel_tool_calls_keep_signature_per_call(
assert parts[1].thought_signature is None


def test_model_response_to_chunk_keeps_a_signature_only_delta():
"""A delta with only a signature is kept so it can reach its call."""
chunks = list(
_model_response_to_chunk(
_streamed_tool_call_chunk(index=0, signature=b"late_sig")
)
)

function_chunk = chunks[0][0]
assert isinstance(function_chunk, FunctionChunk)
assert not function_chunk.name and not function_chunk.args
assert function_chunk.thought_signature == b"late_sig"


@pytest.mark.asyncio
async def test_streaming_signature_after_its_call_signs_that_call(
mock_completion, lite_llm_instance
):
"""A signature-only delta that trails a later call still signs its own."""
mock_completion.return_value = iter([
_streamed_tool_call_chunk(
index=0,
call_id="call_1",
name="get_weather",
arguments='{"city": "Oslo"}',
),
_streamed_tool_call_chunk(
index=1,
call_id="call_2",
name="get_weather",
arguments='{"city": "Bergen"}',
),
_streamed_tool_call_chunk(index=0, signature=b"late_sig"),
_streamed_finish_chunk(),
])

responses = [
response
async for response in lite_llm_instance.generate_content_async(
LLM_REQUEST_WITH_FUNCTION_DECLARATION, stream=True
)
]

parts = responses[-1].content.parts
assert [p.function_call.id for p in parts] == ["call_1", "call_2"]
assert [p.thought_signature for p in parts] == [b"late_sig", None]


@pytest.mark.asyncio
async def test_streaming_signature_before_its_call_signs_that_call(
mock_completion, lite_llm_instance
):
"""A signature-only delta ahead of its call opens no call of its own."""
mock_completion.return_value = iter([
_streamed_tool_call_chunk(index=0, signature=b"early_sig"),
_streamed_tool_call_chunk(
index=0,
call_id="call_1",
name="get_weather",
arguments='{"city": "Oslo"}',
),
_streamed_finish_chunk(),
])

responses = [
response
async for response in lite_llm_instance.generate_content_async(
LLM_REQUEST_WITH_FUNCTION_DECLARATION, stream=True
)
]

streamed_calls = [
part.function_call
for response in responses
for part in (response.content.parts if response.content else [])
if part.function_call
]
assert all(call.name == "get_weather" for call in streamed_calls)
parts = responses[-1].content.parts
assert [p.function_call.id for p in parts] == ["call_1"]
assert parts[0].thought_signature == b"early_sig"


def test_message_to_generate_content_response_no_thought_signature():
"""Parts without thought_signature have thought_signature=None."""
message = ChatCompletionAssistantMessage(
Expand Down
Loading