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
69 changes: 42 additions & 27 deletions src/google/adk/models/lite_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2064,6 +2064,16 @@ def _parse_deepseek_tool_calls_from_text(
return tool_calls, remainder or None


_INLINE_CALL_OPEN_RE = re.compile(r"\s*(?:<tool_call>|```(?:json)?)?\s*")
_INLINE_CALL_CLOSE_RE = re.compile(r"\s*(?:</tool_call>|```)?\s*")


def _skip_past(pattern: re.Pattern[str], text: str, pos: int) -> int:
"""Returns the index just past `pattern` matched at `pos`, or `pos`."""
match = pattern.match(text, pos)
return match.end() if match else pos


def _parse_tool_calls_from_text(
text_block: str,
) -> tuple[list[ChatCompletionMessageToolCall], Optional[str]]:
Expand All @@ -2085,41 +2095,35 @@ def _parse_tool_calls_from_text(
return tool_calls, extra_remainder
return ds_tool_calls, None

remainder_segments = []
# Only calls the text opens with count, optionally in the wrappers models
# emit them in. JSON after any prose is quoted text: turning it into a call
# would run a tool the model only showed.
cursor = 0
text_length = len(text_block)

while cursor < text_length:
brace_index = text_block.find("{", cursor)
if brace_index == -1:
remainder_segments.append(text_block[cursor:])
while True:
start = _skip_past(_INLINE_CALL_OPEN_RE, text_block, cursor)
if not text_block.startswith("{", start):
break

remainder_segments.append(text_block[cursor:brace_index])
try:
candidate, end = _JSON_DECODER.raw_decode(text_block, brace_index)
candidate, end = _JSON_DECODER.raw_decode(text_block, start)
except json.JSONDecodeError:
remainder_segments.append(text_block[brace_index])
cursor = brace_index + 1
continue

break
tool_call = _build_tool_call_from_json_dict(
candidate, index=len(tool_calls)
)
if tool_call:
tool_calls.append(tool_call)
else:
remainder_segments.append(text_block[brace_index:end])
cursor = end
if not tool_call:
break
tool_calls.append(tool_call)
cursor = _skip_past(_INLINE_CALL_CLOSE_RE, text_block, end)

remainder = "".join(segment for segment in remainder_segments if segment)
remainder = remainder.strip()

return tool_calls, remainder or None
if not tool_calls:
return tool_calls, text_block.strip() or None
return tool_calls, text_block[cursor:].strip() or None


def _split_message_content_and_tool_calls(
message: Message,
*,
parse_inline_tool_calls: bool = True,
) -> tuple[Optional[OpenAIMessageContent], list[ChatCompletionMessageToolCall]]:
"""Returns message content and tool calls, parsing inline JSON when needed."""
existing_tool_calls = message.get("tool_calls") or []
Expand All @@ -2130,7 +2134,11 @@ def _split_message_content_and_tool_calls(

# LiteLLM responses either provide structured tool_calls or inline JSON, not
# both. When tool_calls are present we trust them and skip the fallback parser.
if normalized_tool_calls or not isinstance(content, str):
if (
normalized_tool_calls
or not isinstance(content, str)
or not parse_inline_tool_calls
):
return content, normalized_tool_calls

fallback_tool_calls, remainder = _parse_tool_calls_from_text(content)
Expand Down Expand Up @@ -2493,11 +2501,15 @@ def _has_meaningful_signal(message: Message | Delta | None) -> bool:
reasoning_parts: List[types.Part] = []

if message is not None:
# Both Delta and Message support dict-like .get() access
# Both Delta and Message support dict-like .get() access. A delta is too
# little text to tell a tool call from one quoted mid-answer, so streamed
# text is parsed for calls once, as a whole, when the stream finalizes.
(
message_content,
tool_calls,
) = _split_message_content_and_tool_calls(message)
) = _split_message_content_and_tool_calls(
message, parse_inline_tool_calls=message_field != "delta"
)
reasoning_value = _extract_reasoning_value(message)
if reasoning_value:
reasoning_parts = _convert_reasoning_value_to_parts(reasoning_value)
Expand Down Expand Up @@ -2687,7 +2699,10 @@ def _message_to_generate_content_response(
)
if thought_parts:
parts.extend(thought_parts)
message_content, tool_calls = _split_message_content_and_tool_calls(message)
# A partial message is one streamed delta; see _model_response_to_chunk.
message_content, tool_calls = _split_message_content_and_tool_calls(
message, parse_inline_tool_calls=not is_partial
)
if isinstance(message_content, str) and message_content:
parts.append(types.Part.from_text(text=message_content))

Expand Down
141 changes: 128 additions & 13 deletions tests/unittests/models/test_litellm.py
Original file line number Diff line number Diff line change
Expand Up @@ -4151,7 +4151,6 @@ async def test_thought_signature_round_trip():
def test_parse_tool_calls_from_text_multiple_calls():
text = (
'{"name":"alpha","arguments":{"value":1}}\n'
"Some filler text "
'{"id":"custom","name":"beta","arguments":{"timezone":"Asia/Taipei"}} '
"ignored suffix"
)
Expand All @@ -4164,7 +4163,21 @@ def test_parse_tool_calls_from_text_multiple_calls():
assert json.loads(tool_calls[1].function.arguments) == {
"timezone": "Asia/Taipei"
}
assert remainder == "Some filler text ignored suffix"
assert remainder == "ignored suffix"


def test_parse_tool_calls_from_text_stops_at_text_between_calls():
"""A tool call that follows prose is quoted text, not another call."""
beta = '{"name":"beta","arguments":{"timezone":"Asia/Taipei"}}'
text = (
'{"name":"alpha","arguments":{"value":1}}\n'
f"Some filler text {beta} ignored suffix"
)

tool_calls, remainder = _parse_tool_calls_from_text(text)

assert [call.function.name for call in tool_calls] == ["alpha"]
assert remainder == f"Some filler text {beta} ignored suffix"


def test_parse_tool_calls_from_text_invalid_json_returns_remainder():
Expand Down Expand Up @@ -4281,7 +4294,7 @@ def test_parse_tool_calls_from_text_mixed_formats():
"""DeepSeek tokens + standard inline JSON in the same text."""
ds_part = _ds_wrapped(_ds_tool_call("ds_func", '{"a": 1}'))
standard_part = '{"name": "std_func", "arguments": {"b": 2}}'
text = ds_part + " some text " + standard_part
text = ds_part + "\n" + standard_part + " some text"
tool_calls, remainder = _parse_tool_calls_from_text(text)
assert len(tool_calls) == 2
assert tool_calls[0].function.name == "ds_func"
Expand Down Expand Up @@ -4313,15 +4326,40 @@ def test_extract_json_from_deepseek_args_invalid_fence_returns_none():
assert _extract_json_from_deepseek_args('```json\n{"a": 1,}\n```') is None


def test_split_message_content_and_tool_calls_inline_text():
message = {
"role": "assistant",
"content": (
'Intro {"name":"alpha","arguments":{"value":1}} trailing content'
),
}
def test_split_message_content_keeps_tool_call_json_quoted_in_text():
"""JSON shaped like a tool call inside prose stays text, not a call."""
text = (
"The README shows this example request:\n"
'{"name":"delete_file","arguments":{"path":"/data/prod.db"}}\n'
"It is used to remove files."
)
message = {"role": "assistant", "content": text}

content, tool_calls = _split_message_content_and_tool_calls(message)

assert tool_calls == []
assert content == text


@pytest.mark.parametrize(
"text",
[
'\n{"name": "alpha", "arguments": {"value": 1}}\n',
(
'<tool_call>\n{"name": "alpha", "arguments": {"value":'
" 1}}\n</tool_call>"
),
'```json\n{"name": "alpha", "arguments": {"value": 1}}\n```',
],
ids=["bare", "tool_call_tags", "code_fence"],
)
def test_split_message_content_parses_text_that_is_a_tool_call(text):
"""A message whose text is a tool call becomes that tool call."""
message = {"role": "assistant", "content": text}

content, tool_calls = _split_message_content_and_tool_calls(message)
assert content == "Intro trailing content"

assert content is None
assert len(tool_calls) == 1
assert tool_calls[0].function.name == "alpha"
assert json.loads(tool_calls[0].function.arguments) == {"value": 1}
Expand Down Expand Up @@ -5061,7 +5099,7 @@ def test_to_litellm_role():
"message": {
"role": "assistant",
"content": (
'Intro {"id":"call_2","name":"alpha",'
'{"id":"call_2","name":"alpha",'
'"arguments":{"foo":"bar"}} wrap'
),
},
Expand All @@ -5073,7 +5111,7 @@ def test_to_litellm_role():
},
),
[
TextChunk(text="Intro wrap"),
TextChunk(text="wrap"),
FunctionChunk(
id="call_2",
name="alpha",
Expand Down Expand Up @@ -6219,6 +6257,83 @@ async def test_streaming_inline_tool_call_malformed_arguments(
assert "test_function" in final_response.error_message


def _text_stream(*deltas: str) -> list[ModelResponseStream]:
"""Streams each text as one delta, then a stop-only chunk."""
chunks = [
ModelResponseStream(
choices=[
StreamingChoices(
finish_reason=None,
delta=Delta(role="assistant", content=text),
)
]
)
for text in deltas
]
chunks.append(
ModelResponseStream(
choices=[
StreamingChoices(
finish_reason="stop",
delta=Delta(role="assistant", content=""),
)
]
)
)
return chunks


@pytest.mark.asyncio
async def test_streaming_text_quoting_a_tool_call_is_not_a_call(
mock_completion, lite_llm_instance
):
"""Prose that quotes a tool call in its own delta streams back as text."""
call_json = '{"name": "test_function", "arguments": {"test_arg": "x"}}'
mock_completion.return_value = iter(
_text_stream("The README shows this example: ", call_json, " Done.")
)

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

assert not [
part
for response in responses
for part in response.content.parts
if part.function_call
]
final_text = "".join(part.text for part in responses[-1].content.parts)
assert final_text == f"The README shows this example: {call_json} Done."


@pytest.mark.asyncio
async def test_streaming_text_that_is_a_tool_call_becomes_a_call(
mock_completion, lite_llm_instance
):
"""A streamed message whose text is a tool call ends as that call."""
mock_completion.return_value = iter(
_text_stream(
'{"name": "test_function", ',
'"arguments": {"test_arg": "x"}}',
)
)

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

function_call = responses[-1].content.parts[0].function_call
assert function_call.name == "test_function"
assert function_call.args == {"test_arg": "x"}


@pytest.mark.asyncio
async def test_streaming_tool_call_complete_with_length_finish_reason(
mock_completion, lite_llm_instance
Expand Down
Loading