fix(gemini): close async handler and harden scope fallback detection

This commit is contained in:
balazss 2026-03-17 20:41:21 -07:00
parent f7374ad800
commit 3c143336f7
5 changed files with 212 additions and 91 deletions

View file

@ -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

View file

@ -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)

View file

@ -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 = {

View file

@ -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

View file

@ -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()