fix(gemini): harden code-assist params and atomic oauth creds writes

This commit is contained in:
balazss 2026-03-18 18:32:49 -07:00
parent 90a7177db8
commit 8eaa1582eb
4 changed files with 54 additions and 28 deletions

View file

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

View file

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

View file

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

View file

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