mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(gemini): close async handler and harden scope fallback detection
This commit is contained in:
parent
f7374ad800
commit
3c143336f7
5 changed files with 212 additions and 91 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue