From 894d8c6c2652492324c428e8dae68a8646b17485 Mon Sep 17 00:00:00 2001 From: George Weale Date: Thu, 11 Dec 2025 12:19:35 -0800 Subject: [PATCH] fix: Refactor LiteLLM response schema formatting for different models The `_to_litellm_response_format` function now adapts the output format based on the provided model. Gemini models continue to use the "response_schema" key, while OpenAI-compatible models (including Azure OpenAI and Anthropic) now use the "json_schema" key as per LiteLLM's documentation for JSON mode. The schema name is also included in the "json_schema" format. Close #3713 Close #3890 Co-authored-by: George Weale PiperOrigin-RevId: 843326850 --- src/google/adk/models/lite_llm.py | 69 +++++++-- tests/unittests/models/test_litellm.py | 202 +++++++++++++++++++++++-- 2 files changed, 247 insertions(+), 24 deletions(-) diff --git a/src/google/adk/models/lite_llm.py b/src/google/adk/models/lite_llm.py index 5098425e..dfbf15a7 100644 --- a/src/google/adk/models/lite_llm.py +++ b/src/google/adk/models/lite_llm.py @@ -982,8 +982,20 @@ def _message_to_generate_content_response( def _to_litellm_response_format( response_schema: types.SchemaUnion, -) -> Optional[Dict[str, Any]]: - """Converts ADK response schema objects into LiteLLM-compatible payloads.""" + model: str, +) -> dict[str, Any] | None: + """Converts ADK response schema objects into LiteLLM-compatible payloads. + + Args: + response_schema: The response schema to convert. + model: The model string to determine the appropriate format. Gemini models + use 'response_schema' key, while OpenAI-compatible models use + 'json_schema' key. + + Returns: + A dictionary with the appropriate response format for LiteLLM. + """ + schema_name = "response" if isinstance(response_schema, dict): schema_type = response_schema.get("type") @@ -993,18 +1005,25 @@ def _to_litellm_response_format( ): return response_schema schema_dict = dict(response_schema) + if "title" in schema_dict: + schema_name = str(schema_dict["title"]) elif isinstance(response_schema, type) and issubclass( response_schema, BaseModel ): schema_dict = response_schema.model_json_schema() + schema_name = response_schema.__name__ elif isinstance(response_schema, BaseModel): if isinstance(response_schema, types.Schema): # GenAI Schema instances already represent JSON schema definitions. schema_dict = response_schema.model_dump(exclude_none=True, mode="json") + if "title" in schema_dict: + schema_name = str(schema_dict["title"]) else: schema_dict = response_schema.__class__.model_json_schema() + schema_name = response_schema.__class__.__name__ elif hasattr(response_schema, "model_dump"): schema_dict = response_schema.model_dump(exclude_none=True, mode="json") + schema_name = response_schema.__class__.__name__ else: logger.warning( "Unsupported response_schema type %s for LiteLLM structured outputs.", @@ -1012,14 +1031,37 @@ def _to_litellm_response_format( ) return None + # Gemini models use a special response format with 'response_schema' key + if _is_litellm_gemini_model(model): + return { + "type": "json_object", + "response_schema": schema_dict, + } + + # OpenAI-compatible format (default) per LiteLLM docs: + # https://docs.litellm.ai/docs/completion/json_mode + if ( + isinstance(schema_dict, dict) + and schema_dict.get("type") == "object" + and "additionalProperties" not in schema_dict + ): + # OpenAI structured outputs require explicit additionalProperties: false. + schema_dict = dict(schema_dict) + schema_dict["additionalProperties"] = False + return { - "type": "json_object", - "response_schema": schema_dict, + "type": "json_schema", + "json_schema": { + "name": schema_name, + "strict": True, + "schema": schema_dict, + }, } async def _get_completion_inputs( llm_request: LlmRequest, + model: str, ) -> Tuple[ List[Message], Optional[List[Dict]], @@ -1030,13 +1072,14 @@ async def _get_completion_inputs( Args: llm_request: The LlmRequest to convert. + model: The model string to use for determining provider-specific behavior. Returns: The litellm inputs (message list, tool dictionary, response format and generation params). """ # Determine provider for file handling - provider = _get_provider_from_model(llm_request.model or "") + provider = _get_provider_from_model(model) # 1. Construct messages messages: List[Message] = [] @@ -1071,14 +1114,15 @@ async def _get_completion_inputs( ] # 3. Handle response format - response_format: Optional[Dict[str, Any]] = None + response_format: dict[str, Any] | None = None if llm_request.config and llm_request.config.response_schema: response_format = _to_litellm_response_format( - llm_request.config.response_schema + llm_request.config.response_schema, + model=model, ) # 4. Extract generation parameters - generation_params: Optional[Dict] = None + generation_params: dict | None = None if llm_request.config: config_dict = llm_request.config.model_dump(exclude_none=True) # Generate LiteLlm parameters here, @@ -1190,9 +1234,7 @@ def _is_litellm_gemini_model(model_string: str) -> bool: Returns: True if it's a Gemini model accessed via LiteLLM, False otherwise """ - # Matches "gemini/gemini-*" (Google AI Studio) or "vertex_ai/gemini-*" (Vertex AI). - pattern = r"^(gemini|vertex_ai)/gemini-" - return bool(re.match(pattern, model_string)) + return model_string.startswith(("gemini/gemini-", "vertex_ai/gemini-")) def _extract_gemini_model_from_litellm(litellm_model: str) -> str: @@ -1308,8 +1350,9 @@ class LiteLlm(BaseLlm): _append_fallback_user_content_if_missing(llm_request) logger.debug(_build_request_log(llm_request)) + model = llm_request.model or self.model messages, tools, response_format, generation_params = ( - await _get_completion_inputs(llm_request) + await _get_completion_inputs(llm_request, model) ) if "functions" in self._additional_args: @@ -1317,7 +1360,7 @@ class LiteLlm(BaseLlm): tools = None completion_args = { - "model": llm_request.model or self.model, + "model": model, "messages": messages, "tools": tools, "response_format": response_format, diff --git a/tests/unittests/models/test_litellm.py b/tests/unittests/models/test_litellm.py index 094fde77..d3d58d44 100644 --- a/tests/unittests/models/test_litellm.py +++ b/tests/unittests/models/test_litellm.py @@ -236,7 +236,9 @@ async def test_get_completion_inputs_formats_pydantic_schema_for_litellm(): config=types.GenerateContentConfig(response_schema=_StructuredOutput) ) - _, _, response_format, _ = await _get_completion_inputs(llm_request) + _, _, response_format, _ = await _get_completion_inputs( + llm_request, model="gemini/gemini-2.0-flash" + ) assert response_format == { "type": "json_object", @@ -253,7 +255,12 @@ def test_to_litellm_response_format_passes_preformatted_dict(): }, } - assert _to_litellm_response_format(response_format) == response_format + assert ( + _to_litellm_response_format( + response_format, model="gemini/gemini-2.0-flash" + ) + == response_format + ) def test_to_litellm_response_format_wraps_json_schema_dict(): @@ -262,7 +269,9 @@ def test_to_litellm_response_format_wraps_json_schema_dict(): "properties": {"foo": {"type": "string"}}, } - formatted = _to_litellm_response_format(schema) + formatted = _to_litellm_response_format( + schema, model="gemini/gemini-2.0-flash" + ) assert formatted["type"] == "json_object" assert formatted["response_schema"] == schema @@ -270,7 +279,9 @@ def test_to_litellm_response_format_wraps_json_schema_dict(): def test_to_litellm_response_format_handles_model_dump_object(): schema_obj = _ModelDumpOnly() - formatted = _to_litellm_response_format(schema_obj) + formatted = _to_litellm_response_format( + schema_obj, model="gemini/gemini-2.0-flash" + ) assert formatted["type"] == "json_object" assert formatted["response_schema"] == schema_obj.model_dump() @@ -283,13 +294,174 @@ def test_to_litellm_response_format_handles_genai_schema_instance(): required=["foo"], ) - formatted = _to_litellm_response_format(schema_instance) + formatted = _to_litellm_response_format( + schema_instance, model="gemini/gemini-2.0-flash" + ) assert formatted["type"] == "json_object" assert formatted["response_schema"] == schema_instance.model_dump( exclude_none=True, mode="json" ) +def test_to_litellm_response_format_uses_json_schema_for_openai_model(): + """Test that OpenAI models use json_schema format instead of response_schema.""" + formatted = _to_litellm_response_format( + _StructuredOutput, model="gpt-4o-mini" + ) + + assert formatted["type"] == "json_schema" + assert "json_schema" in formatted + assert formatted["json_schema"]["name"] == "_StructuredOutput" + assert formatted["json_schema"]["strict"] is True + assert formatted["json_schema"]["schema"]["additionalProperties"] is False + assert "additionalProperties" in formatted["json_schema"]["schema"] + + +def test_to_litellm_response_format_uses_response_schema_for_gemini_model(): + """Test that Gemini models continue to use response_schema format.""" + formatted = _to_litellm_response_format( + _StructuredOutput, model="gemini/gemini-2.0-flash" + ) + + assert formatted["type"] == "json_object" + assert "response_schema" in formatted + assert formatted["response_schema"] == _StructuredOutput.model_json_schema() + + +def test_to_litellm_response_format_uses_response_schema_for_vertex_gemini(): + """Test that Vertex AI Gemini models use response_schema format.""" + formatted = _to_litellm_response_format( + _StructuredOutput, model="vertex_ai/gemini-2.0-flash" + ) + + assert formatted["type"] == "json_object" + assert "response_schema" in formatted + assert formatted["response_schema"] == _StructuredOutput.model_json_schema() + + +def test_to_litellm_response_format_uses_json_schema_for_azure_openai(): + """Test that Azure OpenAI models use json_schema format.""" + formatted = _to_litellm_response_format( + _StructuredOutput, model="azure/gpt-4o" + ) + + assert formatted["type"] == "json_schema" + assert "json_schema" in formatted + assert formatted["json_schema"]["name"] == "_StructuredOutput" + assert formatted["json_schema"]["strict"] is True + assert formatted["json_schema"]["schema"]["additionalProperties"] is False + assert "additionalProperties" in formatted["json_schema"]["schema"] + + +def test_to_litellm_response_format_uses_json_schema_for_anthropic(): + """Test that Anthropic models use json_schema format.""" + formatted = _to_litellm_response_format( + _StructuredOutput, model="anthropic/claude-3-5-sonnet" + ) + + assert formatted["type"] == "json_schema" + assert "json_schema" in formatted + assert formatted["json_schema"]["name"] == "_StructuredOutput" + assert formatted["json_schema"]["strict"] is True + assert formatted["json_schema"]["schema"]["additionalProperties"] is False + assert "additionalProperties" in formatted["json_schema"]["schema"] + + +def test_to_litellm_response_format_with_dict_schema_for_openai(): + """Test dict schema with OpenAI model uses json_schema format.""" + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + } + + formatted = _to_litellm_response_format(schema, model="gpt-4o") + + assert formatted["type"] == "json_schema" + assert formatted["json_schema"]["name"] == "response" + assert formatted["json_schema"]["strict"] is True + assert formatted["json_schema"]["schema"]["additionalProperties"] is False + + +async def test_get_completion_inputs_uses_openai_format_for_openai_model(): + """Test that _get_completion_inputs produces OpenAI-compatible format.""" + llm_request = LlmRequest( + model="gpt-4o-mini", + config=types.GenerateContentConfig(response_schema=_StructuredOutput), + ) + + _, _, response_format, _ = await _get_completion_inputs( + llm_request, model="gpt-4o-mini" + ) + + assert response_format["type"] == "json_schema" + assert "json_schema" in response_format + assert response_format["json_schema"]["name"] == "_StructuredOutput" + assert response_format["json_schema"]["strict"] is True + assert ( + response_format["json_schema"]["schema"]["additionalProperties"] is False + ) + + +async def test_get_completion_inputs_uses_gemini_format_for_gemini_model(): + """Test that _get_completion_inputs produces Gemini-compatible format.""" + llm_request = LlmRequest( + model="gemini/gemini-2.0-flash", + config=types.GenerateContentConfig(response_schema=_StructuredOutput), + ) + + _, _, response_format, _ = await _get_completion_inputs( + llm_request, model="gemini/gemini-2.0-flash" + ) + + assert response_format["type"] == "json_object" + assert "response_schema" in response_format + + +async def test_get_completion_inputs_uses_passed_model_for_response_format(): + """Test that _get_completion_inputs uses the passed model parameter for response format. + + This verifies that when llm_request.model is None, the explicit model parameter + is used to determine the correct response format (Gemini vs OpenAI). + """ + llm_request = LlmRequest( + model=None, # No model in request + config=types.GenerateContentConfig(response_schema=_StructuredOutput), + ) + + # Pass OpenAI model explicitly - should use json_schema format + _, _, response_format, _ = await _get_completion_inputs( + llm_request, model="gpt-4o-mini" + ) + + assert response_format["type"] == "json_schema" + assert "json_schema" in response_format + assert response_format["json_schema"]["name"] == "_StructuredOutput" + assert response_format["json_schema"]["strict"] is True + assert ( + response_format["json_schema"]["schema"]["additionalProperties"] is False + ) + + +async def test_get_completion_inputs_uses_passed_model_for_gemini_format(): + """Test that _get_completion_inputs uses passed model for Gemini response format. + + This verifies that when self.model is a Gemini model and passed explicitly, + the response format uses the Gemini-specific format. + """ + llm_request = LlmRequest( + model=None, # No model in request + config=types.GenerateContentConfig(response_schema=_StructuredOutput), + ) + + # Pass Gemini model explicitly - should use response_schema format + _, _, response_format, _ = await _get_completion_inputs( + llm_request, model="gemini/gemini-2.0-flash" + ) + + assert response_format["type"] == "json_object" + assert "response_schema" in response_format + + def test_schema_to_dict_filters_none_enum_values(): # Use model_construct to bypass strict enum validation. top_level_schema = types.Schema.model_construct( @@ -2421,7 +2593,9 @@ async def test_get_completion_inputs_generation_params(): ), ) - _, _, _, generation_params = await _get_completion_inputs(req) + _, _, _, generation_params = await _get_completion_inputs( + req, model="gpt-4o-mini" + ) assert generation_params["temperature"] == 0.33 assert generation_params["max_completion_tokens"] == 123 assert generation_params["top_p"] == 0.88 @@ -2444,7 +2618,9 @@ async def test_get_completion_inputs_empty_generation_params(): config=types.GenerateContentConfig(), ) - _, _, _, generation_params = await _get_completion_inputs(req) + _, _, _, generation_params = await _get_completion_inputs( + req, model="gpt-4o-mini" + ) assert generation_params is None @@ -2460,7 +2636,9 @@ async def test_get_completion_inputs_minimal_config(): ), ) - _, _, _, generation_params = await _get_completion_inputs(req) + _, _, _, generation_params = await _get_completion_inputs( + req, model="gpt-4o-mini" + ) assert generation_params is None @@ -2477,7 +2655,9 @@ async def test_get_completion_inputs_partial_generation_params(): ), ) - _, _, _, generation_params = await _get_completion_inputs(req) + _, _, _, generation_params = await _get_completion_inputs( + req, model="gpt-4o-mini" + ) assert generation_params is not None assert generation_params["temperature"] == 0.7 # Should only contain the temperature parameter @@ -2830,7 +3010,7 @@ async def test_get_completion_inputs_openai_file_upload(mocker): ) messages, tools, response_format, generation_params = ( - await _get_completion_inputs(llm_request) + await _get_completion_inputs(llm_request, model="openai/gpt-4o") ) assert len(messages) == 1 @@ -2869,7 +3049,7 @@ async def test_get_completion_inputs_non_openai_no_file_upload(mocker): ) messages, tools, response_format, generation_params = ( - await _get_completion_inputs(llm_request) + await _get_completion_inputs(llm_request, model="anthropic/claude-3-opus") ) assert len(messages) == 1