diff --git a/src/google/adk/models/lite_llm.py b/src/google/adk/models/lite_llm.py index 384d76da..f6705c1d 100644 --- a/src/google/adk/models/lite_llm.py +++ b/src/google/adk/models/lite_llm.py @@ -110,6 +110,18 @@ _MISSING_TOOL_RESULT_MESSAGE = ( ) +def _map_finish_reason( + finish_reason: Any, +) -> types.FinishReason | None: + """Maps a LiteLLM finish_reason value to a google-genai FinishReason enum.""" + if not finish_reason: + return None + if isinstance(finish_reason, types.FinishReason): + return finish_reason + finish_reason_str = str(finish_reason).lower() + return _FINISH_REASON_MAPPING.get(finish_reason_str, types.FinishReason.OTHER) + + def _get_provider_from_model(model: str) -> str: """Extracts the provider name from a LiteLLM model string. @@ -1840,6 +1852,9 @@ class LiteLlm(BaseLlm): else None, ) ) + aggregated_llm_response_with_tool_call.finish_reason = ( + _map_finish_reason(finish_reason) + ) text = "" reasoning_parts = [] function_calls.clear() @@ -1854,6 +1869,9 @@ class LiteLlm(BaseLlm): if reasoning_parts else None, ) + aggregated_llm_response.finish_reason = _map_finish_reason( + finish_reason + ) text = "" reasoning_parts = [] diff --git a/tests/unittests/models/test_litellm.py b/tests/unittests/models/test_litellm.py index c687ceb0..f6428087 100644 --- a/tests/unittests/models/test_litellm.py +++ b/tests/unittests/models/test_litellm.py @@ -2880,6 +2880,7 @@ async def test_generate_content_async_stream( "test_arg": "test_value" } assert responses[3].content.parts[-1].function_call.id == "test_tool_call_id" + assert responses[3].finish_reason == types.FinishReason.STOP assert responses[3].model_version == "test_model" mock_completion.assert_called_once() @@ -2900,6 +2901,55 @@ async def test_generate_content_async_stream( ) +@pytest.mark.asyncio +async def test_generate_content_async_stream_sets_finish_reason( + mock_completion, lite_llm_instance +): + mock_completion.return_value = iter([ + ModelResponse( + model="test_model", + choices=[ + StreamingChoices( + finish_reason=None, + delta=Delta(role="assistant", content="Hello "), + ) + ], + ), + ModelResponse( + model="test_model", + choices=[ + StreamingChoices( + finish_reason=None, + delta=Delta(role="assistant", content="world"), + ) + ], + ), + ModelResponse( + model="test_model", + choices=[StreamingChoices(finish_reason="stop", delta=Delta())], + ), + ]) + + llm_request = LlmRequest( + contents=[ + types.Content( + role="user", parts=[types.Part.from_text(text="Test prompt")] + ) + ], + ) + + responses = [ + response + async for response in lite_llm_instance.generate_content_async( + llm_request, stream=True + ) + ] + + assert responses[-1].partial is False + assert responses[-1].finish_reason == types.FinishReason.STOP + assert responses[-1].content.parts[0].text == "Hello world" + + @pytest.mark.asyncio async def test_generate_content_async_stream_with_usage_metadata( mock_completion, lite_llm_instance @@ -2944,6 +2994,7 @@ async def test_generate_content_async_stream_with_usage_metadata( "test_arg": "test_value" } assert responses[3].content.parts[-1].function_call.id == "test_tool_call_id" + assert responses[3].finish_reason == types.FinishReason.STOP assert responses[3].usage_metadata.prompt_token_count == 10 assert responses[3].usage_metadata.candidates_token_count == 5