mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge b461a67d97 into 22a349ee70
This commit is contained in:
commit
a5cd51482f
2 changed files with 43 additions and 4 deletions
|
|
@ -38,6 +38,15 @@ else:
|
|||
GoogleCredentialsObject = Any
|
||||
|
||||
|
||||
def _import_token_state() -> "type[TokenState]":
|
||||
try:
|
||||
from google.auth.credentials import TokenState
|
||||
except ImportError:
|
||||
raise ImportError(GOOGLE_IMPORT_ERROR_MESSAGE)
|
||||
|
||||
return TokenState
|
||||
|
||||
|
||||
class VertexBase:
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
|
@ -399,7 +408,7 @@ class VertexBase:
|
|||
Look up cached credentials and return (token, project_id) if the token
|
||||
is FRESH. Returns None if not cached or not fresh.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
TokenState: Final = _import_token_state()
|
||||
|
||||
creds, cached_project_id = self._unpack_cached_credentials(credential_cache_key)
|
||||
if (
|
||||
|
|
@ -425,7 +434,7 @@ class VertexBase:
|
|||
credentials object so the caller can schedule a background refresh
|
||||
without holding the per-key async lock.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
TokenState: Final = _import_token_state()
|
||||
|
||||
creds, cached_project_id = self._unpack_cached_credentials(credential_cache_key)
|
||||
if creds is None:
|
||||
|
|
@ -459,7 +468,7 @@ class VertexBase:
|
|||
Falls back to expired/valid checks if token_state is unavailable
|
||||
(e.g. older google-auth versions or mock objects in tests).
|
||||
"""
|
||||
from google.auth.credentials import TokenState as _TokenState
|
||||
_TokenState: Final = _import_token_state()
|
||||
|
||||
token_state: Final = getattr(credentials, "token_state", None)
|
||||
if isinstance(token_state, _TokenState):
|
||||
|
|
@ -980,7 +989,7 @@ class VertexBase:
|
|||
only one coroutine refreshes while others wait on the lock. Uses native
|
||||
async refresh for service_account and authorized_user credentials.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
TokenState: Final = _import_token_state()
|
||||
|
||||
cache_credentials: Final = json.dumps(credentials) if isinstance(credentials, dict) else credentials
|
||||
credential_cache_key: Final = (cache_credentials, project_id)
|
||||
|
|
|
|||
|
|
@ -2083,3 +2083,33 @@ class TestVertexBase:
|
|||
|
||||
assert token == "cached-token"
|
||||
assert not mock_get_lock.called, "Fast path should not acquire lock"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_google_auth_raises_actionable_import_error(self):
|
||||
"""Test that a missing google-auth surfaces the install hint, not a raw ModuleNotFoundError"""
|
||||
vertex_base = VertexBase()
|
||||
|
||||
with patch.dict("sys.modules", {"google.auth.credentials": None}):
|
||||
with pytest.raises(ImportError, match="Google Cloud SDK not found"):
|
||||
await vertex_base._ensure_access_token_async(
|
||||
credentials={"type": "service_account", "project_id": "project-1"},
|
||||
project_id="project-1",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call_helper",
|
||||
[
|
||||
lambda vb: vb._try_get_cached_token(("key", None), "project-1"),
|
||||
lambda vb: vb._try_get_usable_cached_token(("key", None), "project-1"),
|
||||
lambda vb: vb._get_token_state(MagicMock()),
|
||||
],
|
||||
ids=["try_get_cached_token", "try_get_usable_cached_token", "get_token_state"],
|
||||
)
|
||||
def test_token_state_helpers_raise_actionable_import_error(self, call_helper):
|
||||
"""Test that every TokenState call site surfaces the install hint"""
|
||||
vertex_base = VertexBase()
|
||||
|
||||
with patch.dict("sys.modules", {"google.auth.credentials": None}):
|
||||
with pytest.raises(ImportError, match="Google Cloud SDK not found"):
|
||||
call_helper(vertex_base)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue