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
50 changes: 43 additions & 7 deletions src/google/adk/models/lite_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -709,6 +709,16 @@ def _is_gemma4_model(model: str) -> bool:
return bool(_GEMMA4_MODEL_PATTERN.search(model.lower()))


def _resolve_tool_result_role(
model: str,
tool_result_role: Literal["tool", "tool_responses"] | None = None,
) -> Literal["tool", "tool_responses"]:
"""Returns the override, or Gemma-4 auto-detect when unset."""
if tool_result_role is not None:
return tool_result_role
return "tool_responses" if _is_gemma4_model(model) else "tool"


class ChatCompletionFileUrlObject(TypedDict, total=False):
file_data: str
file_id: str
Expand Down Expand Up @@ -1269,6 +1279,7 @@ async def _content_to_message_param(
*,
provider: str = "",
model: str = "",
tool_result_role: Literal["tool", "tool_responses"] | None = None,
) -> Union[Message, list[Message]] | None:
"""Converts a types.Content to a litellm Message or list of Messages.

Expand All @@ -1279,6 +1290,8 @@ async def _content_to_message_param(
content: The content to convert.
provider: The LLM provider name (e.g., "openai", "azure").
model: The LiteLLM model string, used for provider-specific behavior.
tool_result_role: Override for tool-result message role. None keeps
Gemma-4 auto-detect.

Returns:
A litellm Message, a list of litellm Messages, or None if skipped.
Expand All @@ -1305,8 +1318,8 @@ async def _content_to_message_param(
# from the tool call, instead of OpenAI-compatible 'tool' role used by other models.
# Earlier Gemma versions before version 4 do not support tool use,
# so this check is intentionally scoped to only look for "gemma4" in the model name.
tool_role: Literal["tool", "tool_responses"] = (
"tool_responses" if _is_gemma4_model(model) else "tool"
tool_role: Literal["tool", "tool_responses"] = _resolve_tool_result_role(
model, tool_result_role
)
tool_messages.append(
_tool_message(
Expand All @@ -1330,6 +1343,7 @@ async def _content_to_message_param(
types.Content(role=content.role, parts=non_tool_parts),
provider=provider,
model=model,
tool_result_role=tool_result_role,
)
follow_up_messages = (
follow_up if isinstance(follow_up, list) else [follow_up]
Expand Down Expand Up @@ -1463,7 +1477,12 @@ async def _content_to_message_param(
)


def _ensure_tool_results(messages: List[Message], model: str) -> List[Message]:
def _ensure_tool_results(
messages: List[Message],
model: str,
*,
tool_result_role: Literal["tool", "tool_responses"] | None = None,
) -> List[Message]:
"""Insert placeholder tool messages for missing tool results.

LiteLLM-backed providers like OpenAI and Anthropic reject histories where an
Expand All @@ -1482,7 +1501,7 @@ def _ensure_tool_results(messages: List[Message], model: str) -> List[Message]:
healed_messages: List[Message] = []
pending_tool_call_ids: List[str] = []
expected_tool_role: Literal["tool", "tool_responses"] = (
"tool_responses" if _is_gemma4_model(model) else "tool"
_resolve_tool_result_role(model, tool_result_role)
)

for message in messages:
Expand Down Expand Up @@ -2711,6 +2730,8 @@ def _to_litellm_response_format(
async def _get_completion_inputs(
llm_request: LlmRequest,
model: str,
*,
tool_result_role: Literal["tool", "tool_responses"] | None = None,
) -> Tuple[
List[Message],
Optional[List[Dict[str, Any]]],
Expand All @@ -2737,7 +2758,10 @@ async def _get_completion_inputs(
messages: List[Message] = []
for content in llm_request.contents or []:
message_param_or_list = await _content_to_message_param(
content, provider=provider, model=model
content,
provider=provider,
model=model,
tool_result_role=tool_result_role,
)
if isinstance(message_param_or_list, list):
messages.extend(message_param_or_list)
Expand All @@ -2753,7 +2777,9 @@ async def _get_completion_inputs(
content=system_instruction,
),
)
messages = _ensure_tool_results(messages, model)
messages = _ensure_tool_results(
messages, model, tool_result_role=tool_result_role
)

# 2. Convert tool declarations
tools: Optional[List[Dict[str, Any]]] = None
Expand Down Expand Up @@ -3098,13 +3124,18 @@ class LiteLlm(BaseLlm):
Attributes:
model: The name of the LiteLlm model.
llm_client: The LLM client to use for the model.
tool_result_role: Override for tool-result message role. None keeps
Gemma-4 auto-detect.
"""

# LiteLLMClient has no JSON serializer, so it is excluded from dumps to keep
# model_dump(mode="json") from raising.
llm_client: LiteLLMClient = Field(default_factory=LiteLLMClient, exclude=True)
"""The LLM client to use for the model."""

tool_result_role: Optional[Literal["tool", "tool_responses"]] = None
"""Override for tool-result message role. None keeps Gemma-4 auto-detect."""

_additional_args: Dict[str, Any] = PrivateAttr(default_factory=dict)

def __init__(self, model: str, **kwargs: Any) -> None:
Expand All @@ -3122,6 +3153,7 @@ def __init__(self, model: str, **kwargs: Any) -> None:
# preventing generation call with llm_client
# and overriding messages, tools and stream which are managed internally
self._additional_args.pop("llm_client", None)
self._additional_args.pop("tool_result_role", None)
self._additional_args.pop("messages", None)
self._additional_args.pop("tools", None)
# public api called from runner determines to stream or not
Expand Down Expand Up @@ -3158,7 +3190,11 @@ async def generate_content_async(

effective_model = llm_request.model or self.model
messages, tools, response_format, generation_params, tool_choice = (
await _get_completion_inputs(llm_request, effective_model)
await _get_completion_inputs(
llm_request,
effective_model,
tool_result_role=self.tool_result_role,
)
)
normalized_messages = _normalize_ollama_chat_messages(
messages,
Expand Down
120 changes: 120 additions & 0 deletions tests/unittests/models/test_lite_llm_gemma_tool_role.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,15 @@
from typing import Any

from google.adk.models.lite_llm import _content_to_message_param
from google.adk.models.lite_llm import _get_completion_inputs
from google.adk.models.lite_llm import LiteLlm
from google.adk.models.lite_llm import LiteLLMClient
from google.adk.models.llm_request import LlmRequest
from google.genai import types
from litellm import ChatCompletionAssistantMessage
from litellm.types.utils import Choices
from litellm.types.utils import ModelResponse
from pydantic import ValidationError
import pytest


Expand Down Expand Up @@ -189,3 +197,115 @@ async def test_non_gemma_multi_response_uses_tool_role(self):
assert isinstance(result, list)
for msg in result:
assert _extract_role(msg) == "tool"


@pytest.mark.asyncio
async def test_tool_result_role_override_forces_tool_on_gemma4():
content = _make_function_response_content()

result = await _content_to_message_param(
content,
model="google/gemma-4-e4b",
tool_result_role="tool",
)

assert _extract_role(result) == "tool"


@pytest.mark.asyncio
async def test_tool_result_role_override_forces_tool_responses_on_openai():
content = _make_function_response_content()

result = await _content_to_message_param(
content,
model="openai/gpt-4o",
tool_result_role="tool_responses",
)

assert _extract_role(result) == "tool_responses"


def test_litellm_stores_tool_result_role_and_does_not_forward_it():
lite_llm = LiteLlm(model="google/gemma-4-e4b", tool_result_role="tool")

assert lite_llm.tool_result_role == "tool"
assert "tool_result_role" not in lite_llm._additional_args
restored = LiteLlm(**lite_llm.model_dump())
assert restored.tool_result_role == "tool"


def test_litellm_rejects_invalid_tool_result_role():
with pytest.raises(ValidationError, match="tool_result_role"):
LiteLlm(model="google/gemma-4-e4b", tool_result_role="assistant")


@pytest.mark.asyncio
async def test_get_completion_inputs_override_heals_with_tool_role():
assistant_content = types.Content(
role="model",
parts=[
types.Part.from_function_call(
name="get_weather", args={"location": "Seoul"}
),
],
)
assistant_content.parts[0].function_call.id = "tool_call_1"
llm_request = LlmRequest(
contents=[
types.Content(role="user", parts=[types.Part.from_text(text="Hi")]),
assistant_content,
types.Content(role="user", parts=[types.Part.from_text(text="Next")]),
]
)

messages, _, _, _, _ = await _get_completion_inputs(
llm_request,
model="google/gemma-4-e4b",
tool_result_role="tool",
)

roles = [_extract_role(m) for m in messages]
assert "tool" in roles
assert "tool_responses" not in roles


@pytest.mark.asyncio
async def test_generate_content_uses_tool_result_role_override():
captured = {}

class _Client(LiteLLMClient):

async def acompletion(self, **kwargs):
captured.update(kwargs)
return ModelResponse(
model="google/gemma-4-e4b",
choices=[
Choices(
message=ChatCompletionAssistantMessage(
role="assistant", content="ok"
)
)
],
)

lite_llm = LiteLlm(
model="google/gemma-4-e4b",
llm_client=_Client(),
tool_result_role="tool",
)
llm_request = LlmRequest(
contents=[
types.Content(role="user", parts=[types.Part.from_text(text="Hi")]),
_make_function_response_content(),
]
)

_ = [
r
async for r in lite_llm.generate_content_async(llm_request, stream=False)
]

roles = [_extract_role(m) for m in captured["messages"]]
assert "tool" in roles
assert "tool_responses" not in roles
assert "tool_result_role" not in captured