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
11 changes: 10 additions & 1 deletion src/google/adk/models/lite_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2178,9 +2178,18 @@ def _message_to_generate_content_response(
for tool_call in tool_calls:
if tool_call.type == "function":
thought_signature = _extract_thought_signature_from_tool_call(tool_call)
try:
args = _parse_tool_call_arguments(tool_call.function.arguments)
except json.JSONDecodeError:
logger.warning(
"Skipping tool call %s (id=REDACTED) with malformed JSON"
" arguments.",
tool_call.function.name,
)
continue
part = types.Part.from_function_call(
name=tool_call.function.name,
args=_parse_tool_call_arguments(tool_call.function.arguments),
args=args,
)
part.function_call.id = tool_call.id
if thought_signature:
Expand Down
114 changes: 114 additions & 0 deletions tests/unittests/models/test_litellm.py
Original file line number Diff line number Diff line change
Expand Up @@ -6760,3 +6760,117 @@ async def test_generate_content_async_omits_tool_choice_when_functions_override(
_, kwargs = mock_acompletion.call_args
assert kwargs.get("tools") is None
assert "tool_choice" not in kwargs


def test_message_to_generate_content_response_malformed_tool_call_json():
"""Malformed tool call arguments should be skipped, not crash."""
message = ChatCompletionAssistantMessage(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
type="function",
id="call_bad",
function=Function(
name="broken_tool",
arguments='{"a":"unterminated',
),
),
ChatCompletionMessageToolCall(
type="function",
id="call_good",
function=Function(
name="valid_tool",
arguments='{"key": "value"}',
),
),
],
)

response = _message_to_generate_content_response(message)
assert response.content.role == "model"
# The malformed tool call should be skipped; only the valid one remains
function_parts = [
p for p in response.content.parts if p.function_call is not None
]
assert len(function_parts) == 1
assert function_parts[0].function_call.name == "valid_tool"
assert function_parts[0].function_call.id == "call_good"


def test_message_to_generate_content_response_all_malformed_tool_calls():
"""When all tool calls have malformed JSON, response should have no parts."""
message = ChatCompletionAssistantMessage(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
type="function",
id="call_1",
function=Function(
name="broken",
arguments="not json at all",
),
),
],
)

response = _message_to_generate_content_response(message)
assert response.content.role == "model"
assert len(response.content.parts) == 0


def test_message_to_generate_content_response_tool_call_none_arguments():
"""Tool call with arguments=None should produce a function call with empty args."""
message = ChatCompletionAssistantMessage(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
type="function",
id="call_none_args",
function=Function(
name="no_args_tool",
arguments=None,
),
),
],
)

response = _message_to_generate_content_response(message)
assert response.content.role == "model"
function_parts = [
p for p in response.content.parts if p.function_call is not None
]
assert len(function_parts) == 1
assert function_parts[0].function_call.name == "no_args_tool"
assert function_parts[0].function_call.id == "call_none_args"
assert dict(function_parts[0].function_call.args) == {}


def test_message_to_generate_content_response_tool_call_empty_string_arguments():
"""Tool call with arguments='' should produce a function call with empty args."""
message = ChatCompletionAssistantMessage(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
type="function",
id="call_empty_args",
function=Function(
name="empty_args_tool",
arguments="",
),
),
],
)

response = _message_to_generate_content_response(message)
assert response.content.role == "model"
function_parts = [
p for p in response.content.parts if p.function_call is not None
]
assert len(function_parts) == 1
assert function_parts[0].function_call.name == "empty_args_tool"
assert function_parts[0].function_call.id == "call_empty_args"
assert dict(function_parts[0].function_call.args) == {}