diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index f93a105f09f..a1cf0564f68 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -1,5 +1,6 @@ import base64 import datetime +import json from typing import Any, Dict, List, Optional, Union import httpx @@ -159,7 +160,78 @@ def should_fallback_to_google_code_assist(error: Exception) -> bool: """ Returns True if the error indicates missing OAuth scope for Gemini calls. """ - return "ACCESS_TOKEN_SCOPE_INSUFFICIENT" in str(error) + + def _iter_exception_chain(exc: BaseException): + seen = set() + current: Optional[BaseException] = exc + while current is not None and id(current) not in seen: + seen.add(id(current)) + yield current + current = current.__cause__ or current.__context__ + + def _contains_scope_code(value: Any) -> bool: + if isinstance(value, str): + return "ACCESS_TOKEN_SCOPE_INSUFFICIENT" in value + if isinstance(value, dict): + for v in value.values(): + if _contains_scope_code(v): + return True + return False + if isinstance(value, list): + for item in value: + if _contains_scope_code(item): + return True + return False + return False + + def _extract_json_payload_from_response(response: Any) -> Optional[dict]: + if response is None: + return None + try: + payload = response.json() + if isinstance(payload, dict): + return payload + except Exception: + pass + try: + text = getattr(response, "text", None) + if isinstance(text, str): + payload = json.loads(text) + if isinstance(payload, dict): + return payload + except Exception: + pass + return None + + for exc in _iter_exception_chain(error): + if isinstance(exc, httpx.HTTPStatusError): + response = getattr(exc, "response", None) + if response is None or getattr(response, "status_code", None) != 403: + continue + + payload = _extract_json_payload_from_response(response) + if payload and _contains_scope_code(payload): + return True + continue + + status_code = getattr(exc, "status_code", None) + if status_code != 403: + continue + + body = getattr(exc, "body", None) + if body and _contains_scope_code(body): + return True + + response = getattr(exc, "response", None) + payload = _extract_json_payload_from_response(response) + if payload and _contains_scope_code(payload): + return True + + message = getattr(exc, "message", None) + if message and _contains_scope_code(message): + return True + + return False def get_gemini_oauth_token() -> Optional[dict]: # noqa: PLR0915 diff --git a/litellm/llms/google_code_assist/chat.py b/litellm/llms/google_code_assist/chat.py index af6fa7abc01..7bb74cb4a54 100644 --- a/litellm/llms/google_code_assist/chat.py +++ b/litellm/llms/google_code_assist/chat.py @@ -113,32 +113,39 @@ class GoogleCodeAssistChat: initial_project_id = gemini_auth_data.get("project_id") async_handler = AsyncHTTPHandler() + try: + final_project_id = await self._ahandle_handshake( + async_handler, token, initial_project_id + ) + litellm_params["google_code_assist_project"] = final_project_id - final_project_id = await self._ahandle_handshake( - async_handler, token, initial_project_id - ) - litellm_params["google_code_assist_project"] = final_project_id + data = self.config.transform_request( + model, messages, optional_params, litellm_params + ) + url = "https://cloudcode-pa.googleapis.com/v1internal:generateContent" + headers = self._get_headers(token) - data = self.config.transform_request( - model, messages, optional_params, litellm_params - ) - url = "https://cloudcode-pa.googleapis.com/v1internal:generateContent" - headers = self._get_headers(token) + response = await async_handler.post(url=url, headers=headers, json=data) + response.raise_for_status() - response = await async_handler.post(url=url, headers=headers, json=data) - response.raise_for_status() - - return self.config.transform_response( - model=model, - raw_response=response, - model_response=model_response, - logging_obj=logging_obj, - request_data=data, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - encoding=None, - ) + return self.config.transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + request_data=data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=None, + ) + finally: + try: + await async_handler.close() + except Exception as close_error: + verbose_logger.debug( + f"Failed to close Google Code Assist async HTTP handler: {close_error}" + ) except Exception as e: raise self._handle_error(e) diff --git a/litellm/llms/google_code_assist/transformation.py b/litellm/llms/google_code_assist/transformation.py index 5de7dd725ed..b58c7ed780d 100644 --- a/litellm/llms/google_code_assist/transformation.py +++ b/litellm/llms/google_code_assist/transformation.py @@ -148,7 +148,7 @@ class GoogleCodeAssistConfig(VertexGeminiConfig): vertex_request["generationConfig"] = generation_config # 3. Wrap in Code Assist envelope (matches verified gemini-cli structure) - user_prompt_id = f"litellm-{uuid.uuid4()}"[:13] + user_prompt_id = f"litellm-{uuid.uuid4()}" model_name = model.split("/")[-1] ca_request = { diff --git a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py index c31ff308c61..92dde78c2f3 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py +++ b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py @@ -1,8 +1,13 @@ from unittest.mock import AsyncMock, patch +import httpx import pytest -from litellm.llms.gemini.common_utils import GeminiModelInfo, GoogleAIStudioTokenCounter +from litellm.llms.gemini.common_utils import ( + GeminiModelInfo, + GoogleAIStudioTokenCounter, + should_fallback_to_google_code_assist, +) class TestGeminiModelInfo: @@ -94,17 +99,23 @@ class TestGoogleAIStudioTokenCounter: def test_should_use_token_counting_api(self): """Test should_use_token_counting_api method with different provider values""" from litellm.types.utils import LlmProviders - + token_counter = GoogleAIStudioTokenCounter() - + # Test with gemini provider - should return True - assert token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value) is True - + assert ( + token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value) + is True + ) + # Test with other providers - should return False - assert token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value) is False + assert ( + token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value) + is False + ) assert token_counter.should_use_token_counting_api("anthropic") is False assert token_counter.should_use_token_counting_api("vertex_ai") is False - + # Test with None - should return False assert token_counter.should_use_token_counting_api(None) is False @@ -112,39 +123,36 @@ class TestGoogleAIStudioTokenCounter: async def test_count_tokens(self): """Test count_tokens method with mocked API response""" from litellm.types.utils import TokenCountResponse - + token_counter = GoogleAIStudioTokenCounter() - + # Mock the GoogleAIStudioTokenCounter from handler module mock_response = { "totalTokens": 31, "totalBillableCharacters": 96, - "promptTokensDetails": [ - { - "modality": "TEXT", - "tokenCount": 31 - } - ] + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 31}], } - - with patch('litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens', - new_callable=AsyncMock) as mock_acount_tokens: + + with patch( + "litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens", + new_callable=AsyncMock, + ) as mock_acount_tokens: mock_acount_tokens.return_value = mock_response - + # Test data model_to_use = "gemini-1.5-flash" contents = [{"parts": [{"text": "Hello world"}]}] request_model = "gemini/gemini-1.5-flash" - + # Call the method result = await token_counter.count_tokens( model_to_use=model_to_use, messages=None, contents=contents, deployment=None, - request_model=request_model + request_model=request_model, ) - + # Verify the result assert result is not None assert isinstance(result, TokenCountResponse) @@ -152,29 +160,21 @@ class TestGoogleAIStudioTokenCounter: assert result.request_model == request_model assert result.model_used == model_to_use assert result.original_response == mock_response - + # Verify the mock was called correctly mock_acount_tokens.assert_called_once_with( - model=model_to_use, - contents=contents + model=model_to_use, contents=contents ) def test_clean_contents_for_gemini_api_removes_id_field(self): """Test that _clean_contents_for_gemini_api removes unsupported 'id' field from function responses""" from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter - + token_counter = GoogleAIStudioTokenCounter() - + # Test contents with function response containing 'id' field (camelCase) contents_with_id = [ - { - "parts": [ - { - "text": "Hello world" - } - ], - "role": "user" - }, + {"parts": [{"text": "Hello world"}], "role": "user"}, { "parts": [ { @@ -183,56 +183,91 @@ class TestGoogleAIStudioTokenCounter: "name": "read_many_files", "response": { "output": "No files matching the criteria were found or all were skipped." - } + }, } } ], - "role": "user" - } + "role": "user", + }, ] - + # Clean the contents - cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_with_id) - + cleaned_contents = token_counter._clean_contents_for_gemini_api( + contents_with_id + ) + # Verify the 'id' field was removed function_response = cleaned_contents[1]["parts"][0]["functionResponse"] assert "id" not in function_response assert "name" in function_response assert "response" in function_response assert function_response["name"] == "read_many_files" - assert function_response["response"]["output"] == "No files matching the criteria were found or all were skipped." - + assert ( + function_response["response"]["output"] + == "No files matching the criteria were found or all were skipped." + ) def test_clean_contents_for_gemini_api_preserves_other_fields(self): """Test that _clean_contents_for_gemini_api preserves other fields and structure""" from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter - + token_counter = GoogleAIStudioTokenCounter() - + # Test contents without function responses contents_without_function_response = [ - { - "parts": [ - { - "text": "This is a regular message" - } - ], - "role": "user" - }, - { - "parts": [ - { - "text": "This is a model response" - } - ], - "role": "model" - } + {"parts": [{"text": "This is a regular message"}], "role": "user"}, + {"parts": [{"text": "This is a model response"}], "role": "model"}, ] - + # Clean the contents - cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_without_function_response) - + cleaned_contents = token_counter._clean_contents_for_gemini_api( + contents_without_function_response + ) + # Verify the contents are unchanged assert cleaned_contents == contents_without_function_response +class TestGeminiFallbackDetection: + def test_should_fallback_to_google_code_assist_for_structured_403_scope_error(self): + request = httpx.Request( + "POST", "https://generativelanguage.googleapis.com/test" + ) + response = httpx.Response( + status_code=403, + json={ + "error": { + "code": 403, + "message": "Request had insufficient authentication scopes.", + "status": "PERMISSION_DENIED", + "details": [{"reason": "ACCESS_TOKEN_SCOPE_INSUFFICIENT"}], + } + }, + request=request, + ) + err = httpx.HTTPStatusError( + "403 Client Error: Forbidden for url", request=request, response=response + ) + + assert should_fallback_to_google_code_assist(err) is True + + def test_should_not_fallback_to_google_code_assist_for_non_403_scope_string(self): + request = httpx.Request( + "POST", "https://generativelanguage.googleapis.com/test" + ) + response = httpx.Response( + status_code=500, + json={ + "error": { + "message": "ACCESS_TOKEN_SCOPE_INSUFFICIENT", + } + }, + request=request, + ) + err = httpx.HTTPStatusError( + "500 Server Error: Internal Server Error for url", + request=request, + response=response, + ) + + assert should_fallback_to_google_code_assist(err) is False diff --git a/tests/test_litellm/llms/google_code_assist/test_google_code_assist.py b/tests/test_litellm/llms/google_code_assist/test_google_code_assist.py index 61bda83eb21..5340443775e 100644 --- a/tests/test_litellm/llms/google_code_assist/test_google_code_assist.py +++ b/tests/test_litellm/llms/google_code_assist/test_google_code_assist.py @@ -1,5 +1,5 @@ import pytest -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import httpx import json from litellm.llms.google_code_assist.chat import GoogleCodeAssistChat @@ -72,8 +72,14 @@ class TestGoogleCodeAssist: @pytest.mark.asyncio @patch("litellm.llms.google_code_assist.chat.AsyncHTTPHandler.post") + @patch( + "litellm.llms.google_code_assist.chat.AsyncHTTPHandler.close", + new_callable=AsyncMock, + ) @patch("litellm.llms.gemini.common_utils.get_gemini_oauth_token") - async def test_acompletion_basic(self, mock_get_token, mock_async_post): + async def test_acompletion_basic( + self, mock_get_token, mock_async_close, mock_async_post + ): """ Test async completion. """ @@ -125,3 +131,4 @@ class TestGoogleCodeAssist: ) assert response.choices[0].message.content == "Async success" + mock_async_close.assert_awaited_once()