mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(gemini): harden code-assist params and atomic oauth creds writes
This commit is contained in:
parent
90a7177db8
commit
8eaa1582eb
4 changed files with 54 additions and 28 deletions
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
import webbrowser
|
||||
import http.server
|
||||
|
|
@ -143,17 +144,22 @@ class GeminiAuthenticator:
|
|||
"""
|
||||
Write oauth credentials with user-only permissions.
|
||||
"""
|
||||
fd = os.open(
|
||||
self.oauth_creds_file,
|
||||
os.O_WRONLY | os.O_CREAT | os.O_TRUNC,
|
||||
0o600,
|
||||
)
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||
json.dump(creds, f)
|
||||
fd, tmp_path = tempfile.mkstemp(dir=self.token_dir, prefix=".tmp_creds_")
|
||||
try:
|
||||
os.chmod(self.oauth_creds_file, 0o600)
|
||||
except OSError:
|
||||
pass
|
||||
os.chmod(tmp_path, 0o600)
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||
json.dump(creds, f)
|
||||
os.replace(tmp_path, self.oauth_creds_file)
|
||||
try:
|
||||
os.chmod(self.oauth_creds_file, 0o600)
|
||||
except OSError:
|
||||
pass
|
||||
except Exception:
|
||||
try:
|
||||
os.unlink(tmp_path)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
|
||||
def _login(self) -> Dict[str, Any]:
|
||||
"""Perform loopback flow login."""
|
||||
|
|
|
|||
|
|
@ -99,13 +99,14 @@ class GoogleCodeAssistConfig(VertexGeminiConfig):
|
|||
"includeThoughts": base_params.pop("include_thoughts")
|
||||
}
|
||||
|
||||
if (
|
||||
"thinkingConfig" in optional_params
|
||||
and "thinkingConfig" not in generation_config
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"google_code_assist: `thinkingConfig` was provided but not mapped into generationConfig."
|
||||
)
|
||||
# Forward any remaining mapped params into generation_config so they are not silently dropped.
|
||||
for key, value in base_params.items():
|
||||
if key not in generation_config:
|
||||
verbose_logger.debug(
|
||||
"google_code_assist: forwarding mapped param '%s' into generationConfig",
|
||||
key,
|
||||
)
|
||||
generation_config[key] = value
|
||||
|
||||
vertex_request = {
|
||||
"contents": contents,
|
||||
|
|
|
|||
|
|
@ -3692,13 +3692,14 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
response = model_response
|
||||
|
||||
elif custom_llm_provider == "google_code_assist":
|
||||
_ca_optional_params = optional_params or {}
|
||||
if acompletion is True:
|
||||
response = _google_code_assist_chat.acompletion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
optional_params=optional_params,
|
||||
optional_params=_ca_optional_params,
|
||||
litellm_params=litellm_params, # type: ignore
|
||||
logging_obj=logging,
|
||||
logger_fn=logger_fn,
|
||||
|
|
@ -3709,7 +3710,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
messages=messages,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
optional_params=optional_params,
|
||||
optional_params=_ca_optional_params,
|
||||
litellm_params=litellm_params, # type: ignore
|
||||
logging_obj=logging,
|
||||
logger_fn=logger_fn,
|
||||
|
|
|
|||
|
|
@ -12,15 +12,9 @@ def test_run_gemini_completion_with_code_assist_fallback_disabled():
|
|||
def _raise_scope_error():
|
||||
raise Exception("ACCESS_TOKEN_SCOPE_INSUFFICIENT")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.llms.gemini.fallback_handler.should_fallback_to_google_code_assist",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"litellm.llms.gemini.fallback_handler._google_code_assist_chat.completion"
|
||||
) as mock_completion,
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.gemini.fallback_handler._google_code_assist_chat.completion"
|
||||
) as mock_completion:
|
||||
with pytest.raises(Exception, match="ACCESS_TOKEN_SCOPE_INSUFFICIENT"):
|
||||
run_gemini_completion_with_code_assist_fallback(
|
||||
primary_call=_raise_scope_error,
|
||||
|
|
@ -31,6 +25,30 @@ def test_run_gemini_completion_with_code_assist_fallback_disabled():
|
|||
mock_completion.assert_not_called()
|
||||
|
||||
|
||||
def test_run_gemini_completion_with_code_assist_fallback_enabled_but_not_match():
|
||||
def _raise_other_error():
|
||||
raise Exception("SOME_OTHER_ERROR")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.llms.gemini.fallback_handler.should_fallback_to_google_code_assist",
|
||||
return_value=False,
|
||||
) as mock_should_fallback,
|
||||
patch(
|
||||
"litellm.llms.gemini.fallback_handler._google_code_assist_chat.completion"
|
||||
) as mock_completion,
|
||||
):
|
||||
with pytest.raises(Exception, match="SOME_OTHER_ERROR"):
|
||||
run_gemini_completion_with_code_assist_fallback(
|
||||
primary_call=_raise_other_error,
|
||||
fallback_kwargs={},
|
||||
auto_fallback_to_google_code_assist=True,
|
||||
)
|
||||
|
||||
mock_should_fallback.assert_called_once()
|
||||
mock_completion.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_gemini_acompletion_with_code_assist_fallback_enabled():
|
||||
async def _raise_scope_error():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue