mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix: move remaining stuff
This commit is contained in:
parent
cfc57d709f
commit
043d2d65b1
22 changed files with 1629 additions and 166 deletions
|
|
@ -10,6 +10,10 @@ The public surface is small on purpose; downstream code should depend on
|
|||
extractor internals.
|
||||
"""
|
||||
|
||||
from litellm.identity.cache import IdentityCache
|
||||
from litellm.identity.jwt import build_user_api_key_auth_from_jwt_result
|
||||
from litellm.identity.oauth2 import build_user_api_key_auth_from_oauth2_response
|
||||
from litellm.identity.runtime import get_identity_cache
|
||||
from litellm.identity.context import (
|
||||
AuditInfo,
|
||||
ClientInfo,
|
||||
|
|
@ -24,16 +28,25 @@ from litellm.identity.principal import (
|
|||
SSOPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
from litellm.identity.resolver import resolve_identity, resolve_user_api_key_auth
|
||||
from litellm.identity.store import load_identity
|
||||
|
||||
__all__ = [
|
||||
"AnonymousPrincipal",
|
||||
"ApiKeyPrincipal",
|
||||
"AuditInfo",
|
||||
"ClientInfo",
|
||||
"IdentityCache",
|
||||
"IdentityContext",
|
||||
"JWTPrincipal",
|
||||
"Principal",
|
||||
"RequestIds",
|
||||
"SSOPrincipal",
|
||||
"ServiceAccountPrincipal",
|
||||
"build_user_api_key_auth_from_jwt_result",
|
||||
"build_user_api_key_auth_from_oauth2_response",
|
||||
"get_identity_cache",
|
||||
"load_identity",
|
||||
"resolve_identity",
|
||||
"resolve_user_api_key_auth",
|
||||
]
|
||||
|
|
|
|||
152
litellm/identity/cache.py
Normal file
152
litellm/identity/cache.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
"""Three-layer identity cache.
|
||||
|
||||
Layer 1 (per-process, ~5s TTL): bounded ``InMemoryCache`` inside the
|
||||
``DualCache`` we wrap. Bounds revocation staleness without round-tripping
|
||||
to Redis on every request.
|
||||
|
||||
Layer 2 (Redis, cross-replica): the ``redis_cache`` on the same
|
||||
``DualCache``. Writes go to both layers; reads fall through.
|
||||
|
||||
Layer 3 (Prisma): not owned here. ``store.load_identity`` calls the DB
|
||||
when both cache layers miss.
|
||||
|
||||
Cross-table fan-out is handled via *generation counters*: when a team or
|
||||
user changes, the counter for that team/user is bumped. Cached
|
||||
identities carry the team/user generation they were minted under, and a
|
||||
read that finds a stale generation treats the entry as a miss. This is
|
||||
cheaper than enumerating every key that references a given team.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
from litellm.integrations.otel.runtime import traced
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
IDENTITY_KEY_PREFIX = "identity:v1"
|
||||
IDENTITY_GENERATION_PREFIX = "identity:gen:v1"
|
||||
DEFAULT_IDENTITY_TTL_SECONDS = 5
|
||||
|
||||
|
||||
def identity_cache_key(token_hash: str) -> str:
|
||||
return f"{IDENTITY_KEY_PREFIX}:{token_hash}"
|
||||
|
||||
|
||||
def team_generation_key(team_id: str) -> str:
|
||||
return f"{IDENTITY_GENERATION_PREFIX}:team:{team_id}"
|
||||
|
||||
|
||||
def user_generation_key(user_id: str) -> str:
|
||||
return f"{IDENTITY_GENERATION_PREFIX}:user:{user_id}"
|
||||
|
||||
|
||||
def org_generation_key(org_id: str) -> str:
|
||||
return f"{IDENTITY_GENERATION_PREFIX}:org:{org_id}"
|
||||
|
||||
|
||||
def _generation_attr_key(scope: str) -> str:
|
||||
return f"identity_cache_generation_{scope}"
|
||||
|
||||
|
||||
def _attach_generations(
|
||||
uak: "UserAPIKeyAuth", generations: dict
|
||||
) -> None:
|
||||
"""Stash the generation counters this entry was minted under.
|
||||
|
||||
Stored on the model's ``metadata`` so the value survives Pydantic
|
||||
serialization round-trips through Redis (``CacheCodec``).
|
||||
"""
|
||||
if not isinstance(uak.metadata, dict):
|
||||
uak.metadata = {}
|
||||
uak.metadata["__identity_cache_generations__"] = generations
|
||||
|
||||
|
||||
def _read_generations(uak: "UserAPIKeyAuth") -> dict:
|
||||
metadata = uak.metadata if isinstance(uak.metadata, dict) else {}
|
||||
raw = metadata.get("__identity_cache_generations__")
|
||||
return raw if isinstance(raw, dict) else {}
|
||||
|
||||
|
||||
class IdentityCache:
|
||||
"""Memory + Redis facade keyed by the hashed virtual key.
|
||||
|
||||
Stores ``UserAPIKeyAuth`` because the legacy carrier is what the rest
|
||||
of the proxy consumes today, and ``UserApiKeyCache`` already knows
|
||||
how to round-trip it through ``CacheCodec``. Callers that want an
|
||||
``IdentityContext`` view should call ``uak.to_identity_context()``
|
||||
at the read site.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dual_cache: "DualCache",
|
||||
ttl_seconds: int = DEFAULT_IDENTITY_TTL_SECONDS,
|
||||
) -> None:
|
||||
self._cache = dual_cache
|
||||
self._ttl_seconds = ttl_seconds
|
||||
|
||||
@traced("identity.cache.get", role=SpanRole.DB_CALL,
|
||||
attrs=lambda result: {
|
||||
"identity.cache.layer": (
|
||||
"miss" if result is None else "memory_or_redis"
|
||||
),
|
||||
})
|
||||
async def get(self, token_hash: str) -> Optional["UserAPIKeyAuth"]:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
cache_key = identity_cache_key(token_hash)
|
||||
cached = await self._cache.async_get_cache(
|
||||
key=cache_key, model_type=UserAPIKeyAuth
|
||||
)
|
||||
if cached is None:
|
||||
return None
|
||||
if await self._is_stale(cached):
|
||||
await self._cache.async_delete_cache(cache_key)
|
||||
return None
|
||||
return cached
|
||||
|
||||
@traced("identity.cache.set", role=SpanRole.DB_CALL)
|
||||
async def set(
|
||||
self, token_hash: str, uak: "UserAPIKeyAuth"
|
||||
) -> None:
|
||||
generations = await self._snapshot_generations_for(uak)
|
||||
_attach_generations(uak, generations)
|
||||
await self._cache.async_set_cache(
|
||||
key=identity_cache_key(token_hash),
|
||||
value=uak,
|
||||
ttl=self._ttl_seconds,
|
||||
)
|
||||
|
||||
async def delete(self, token_hash: str) -> None:
|
||||
await self._cache.async_delete_cache(identity_cache_key(token_hash))
|
||||
|
||||
async def _is_stale(self, uak: "UserAPIKeyAuth") -> bool:
|
||||
stored = _read_generations(uak)
|
||||
if not stored:
|
||||
return False
|
||||
current = await self._snapshot_generations_for(uak)
|
||||
return any(stored.get(k) != current.get(k) for k in stored)
|
||||
|
||||
async def _snapshot_generations_for(
|
||||
self, uak: "UserAPIKeyAuth"
|
||||
) -> dict:
|
||||
scopes: list[tuple[str, str]] = []
|
||||
if uak.team_id:
|
||||
scopes.append(("team", team_generation_key(uak.team_id)))
|
||||
if uak.user_id:
|
||||
scopes.append(("user", user_generation_key(uak.user_id)))
|
||||
if uak.org_id:
|
||||
scopes.append(("org", org_generation_key(uak.org_id)))
|
||||
return {
|
||||
scope: await self._cache.async_get_cache(key=key) or 0
|
||||
for scope, key in scopes
|
||||
}
|
||||
|
||||
async def bump_generation(self, scope_key: str) -> None:
|
||||
await self._cache.async_increment_cache(key=scope_key, value=1)
|
||||
70
litellm/identity/invalidation.py
Normal file
70
litellm/identity/invalidation.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
"""Identity-cache invalidation hooks.
|
||||
|
||||
Two flavors:
|
||||
|
||||
- Per-token: a key was rotated, blocked, or deleted. We know the exact
|
||||
token hash, so we drop the entry from both memory and Redis.
|
||||
|
||||
- Per-scope (team / user / org): a row that fans out to many keys
|
||||
changed. Rather than enumerating every key that references the team,
|
||||
we bump a generation counter for that scope. Cached identities carry
|
||||
the scope generations they were minted under; reads compare and treat
|
||||
a mismatch as a miss.
|
||||
|
||||
The legacy ``_delete_cache_key_object`` and the per-table cache deletes
|
||||
in ``auth_checks.py`` are left in place by Phase 2. These hooks run
|
||||
side-by-side so we don't strand a partially-deployed fleet that's still
|
||||
reading from the legacy cache keys.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm.identity.cache import (
|
||||
IdentityCache,
|
||||
org_generation_key,
|
||||
team_generation_key,
|
||||
user_generation_key,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
||||
def _identity_cache_for(dual_cache: "DualCache") -> IdentityCache:
|
||||
return IdentityCache(dual_cache=dual_cache)
|
||||
|
||||
|
||||
async def invalidate_identity_for_token(
|
||||
*, token_hash: str, dual_cache: "DualCache"
|
||||
) -> None:
|
||||
"""Drop a single key's cached identity."""
|
||||
await _identity_cache_for(dual_cache).delete(token_hash)
|
||||
|
||||
|
||||
async def invalidate_identity_for_team(
|
||||
*, team_id: str, dual_cache: "DualCache"
|
||||
) -> None:
|
||||
"""Mark every identity that references this team as stale."""
|
||||
await _identity_cache_for(dual_cache).bump_generation(
|
||||
team_generation_key(team_id)
|
||||
)
|
||||
|
||||
|
||||
async def invalidate_identity_for_user(
|
||||
*, user_id: str, dual_cache: "DualCache"
|
||||
) -> None:
|
||||
"""Mark every identity that references this user as stale."""
|
||||
await _identity_cache_for(dual_cache).bump_generation(
|
||||
user_generation_key(user_id)
|
||||
)
|
||||
|
||||
|
||||
async def invalidate_identity_for_org(
|
||||
*, org_id: str, dual_cache: "DualCache"
|
||||
) -> None:
|
||||
"""Mark every identity that references this organization as stale."""
|
||||
await _identity_cache_for(dual_cache).bump_generation(
|
||||
org_generation_key(org_id)
|
||||
)
|
||||
121
litellm/identity/jwt.py
Normal file
121
litellm/identity/jwt.py
Normal file
|
|
@ -0,0 +1,121 @@
|
|||
"""JWT identity construction.
|
||||
|
||||
Owns the translation from a ``JWTAuthManager.auth_builder`` result (or
|
||||
any equivalent JWT-validated payload) into the proxy's carrier model
|
||||
``UserAPIKeyAuth``. The JWT branch of ``_user_api_key_auth_builder``
|
||||
used to inline this construction; centralizing it here means:
|
||||
|
||||
- Every JWT-derived ``UserAPIKeyAuth`` carries the same team / user /
|
||||
membership fields, so downstream auth checks see one shape.
|
||||
- The mapping from ``jwt_claims`` to a ``JWTPrincipal`` happens at the
|
||||
same boundary, so callers that want the typed principal can read it
|
||||
off ``UserAPIKeyAuth.to_identity_context()``.
|
||||
|
||||
This module does NOT perform JWT validation or policy checks. Signature
|
||||
verification, RBAC, scope, email-domain enforcement, and
|
||||
``custom_validate`` are all done by ``JWTAuthManager.auth_builder``
|
||||
upstream. Here we just build the carrier.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
|
||||
def build_user_api_key_auth_from_jwt_result(
|
||||
*,
|
||||
result: dict,
|
||||
parent_otel_span: Any = None,
|
||||
is_proxy_admin: bool,
|
||||
) -> "UserAPIKeyAuth":
|
||||
"""Build a ``UserAPIKeyAuth`` carrier from a JWT auth-builder result.
|
||||
|
||||
``result`` is the dict returned by
|
||||
``JWTAuthManager.auth_builder``; its shape (``team_id``,
|
||||
``team_object``, ``user_id``, ``user_object``, ``end_user_id``,
|
||||
``org_id``, ``team_membership``, ``jwt_claims``) is the public
|
||||
contract between the validation layer and this construction layer.
|
||||
|
||||
The ``is_proxy_admin`` flag is the caller's responsibility to pass
|
||||
because the auth_builder result reports it but the call site already
|
||||
branched on it; threading it through avoids re-checking.
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
team_id = result["team_id"]
|
||||
team_object = result["team_object"]
|
||||
user_id = result["user_id"]
|
||||
user_object = result["user_object"]
|
||||
end_user_id = result["end_user_id"]
|
||||
org_id = result["org_id"]
|
||||
team_membership = result.get("team_membership")
|
||||
jwt_claims = result.get("jwt_claims")
|
||||
|
||||
team_alias = team_object.team_alias if team_object is not None else None
|
||||
team_tpm_limit = team_object.tpm_limit if team_object is not None else None
|
||||
team_rpm_limit = team_object.rpm_limit if team_object is not None else None
|
||||
team_models = team_object.models if team_object is not None else []
|
||||
team_metadata = team_object.metadata if team_object is not None else None
|
||||
team_object_permission = (
|
||||
team_object.object_permission if team_object is not None else None
|
||||
)
|
||||
|
||||
if is_proxy_admin:
|
||||
return UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
team_alias=team_alias,
|
||||
team_tpm_limit=team_tpm_limit,
|
||||
team_rpm_limit=team_rpm_limit,
|
||||
team_models=team_models,
|
||||
team_metadata=team_metadata,
|
||||
org_id=org_id,
|
||||
end_user_id=end_user_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
jwt_claims=jwt_claims,
|
||||
)
|
||||
|
||||
user_role = (
|
||||
LitellmUserRoles(user_object.user_role)
|
||||
if user_object is not None and user_object.user_role is not None
|
||||
else LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
user_tpm_limit = user_object.tpm_limit if user_object is not None else None
|
||||
user_rpm_limit = user_object.rpm_limit if user_object is not None else None
|
||||
team_member_rpm_limit = (
|
||||
team_membership.safe_get_team_member_rpm_limit()
|
||||
if team_membership is not None
|
||||
else None
|
||||
)
|
||||
team_member_tpm_limit = (
|
||||
team_membership.safe_get_team_member_tpm_limit()
|
||||
if team_membership is not None
|
||||
else None
|
||||
)
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
team_id=team_id,
|
||||
team_alias=team_alias,
|
||||
team_tpm_limit=team_tpm_limit,
|
||||
team_rpm_limit=team_rpm_limit,
|
||||
team_models=team_models,
|
||||
user_role=user_role,
|
||||
user_id=user_id,
|
||||
org_id=org_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
end_user_id=end_user_id,
|
||||
user_tpm_limit=user_tpm_limit,
|
||||
user_rpm_limit=user_rpm_limit,
|
||||
team_member_rpm_limit=team_member_rpm_limit,
|
||||
team_member_tpm_limit=team_member_tpm_limit,
|
||||
team_metadata=team_metadata,
|
||||
jwt_claims=jwt_claims,
|
||||
)
|
||||
valid_token.team_object_permission = team_object_permission
|
||||
return valid_token
|
||||
49
litellm/identity/oauth2.py
Normal file
49
litellm/identity/oauth2.py
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
"""OAuth2 identity construction.
|
||||
|
||||
Owns the translation from an OAuth2 introspection / userinfo response
|
||||
into the proxy's carrier model ``UserAPIKeyAuth``. The HTTP plumbing —
|
||||
introspection endpoint detection, request signing, error handling —
|
||||
stays in ``litellm.proxy.auth.oauth2_check.Oauth2Handler``; this module
|
||||
just maps a validated response payload into the carrier so the JWT and
|
||||
OAuth2 paths converge on the same construction surface.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
|
||||
def build_user_api_key_auth_from_oauth2_response(
|
||||
*,
|
||||
token: str,
|
||||
response_data: dict,
|
||||
user_id_field_name: str = "sub",
|
||||
user_role_field_name: str = "role",
|
||||
user_team_id_field_name: str = "team_id",
|
||||
) -> "UserAPIKeyAuth":
|
||||
"""Build a ``UserAPIKeyAuth`` carrier from an OAuth2 response.
|
||||
|
||||
``response_data`` is the parsed JSON body of an OAuth2 introspection
|
||||
or userinfo response. The three ``*_field_name`` parameters let
|
||||
deployments map their IdP's claim names onto our canonical fields;
|
||||
defaults match the OAuth2 introspection spec (``sub``) plus
|
||||
``role`` / ``team_id`` for the role and team claims.
|
||||
|
||||
Active-token validation, scope checks, and token-not-active rejection
|
||||
happen upstream in ``Oauth2Handler``; this builder trusts its input.
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
user_id: Optional[str] = response_data.get(user_id_field_name)
|
||||
user_role: Optional[str] = response_data.get(user_role_field_name)
|
||||
user_team_id: Optional[str] = response_data.get(user_team_id_field_name)
|
||||
|
||||
return UserAPIKeyAuth(
|
||||
api_key=token,
|
||||
team_id=user_team_id,
|
||||
user_id=user_id,
|
||||
user_role=cast("LitellmUserRoles", user_role),
|
||||
)
|
||||
|
|
@ -1,11 +1,20 @@
|
|||
"""Compose extractors into a single ``IdentityContext`` per request.
|
||||
"""Compose extractors + DB load into a single ``IdentityContext`` per request.
|
||||
|
||||
This entrypoint is intentionally not yet wired into the proxy auth chain;
|
||||
it exists so Phase 2 can switch over and so unit tests can exercise the
|
||||
full path end-to-end.
|
||||
Two call shapes:
|
||||
|
||||
- ``resolve_identity_for_principal`` — given pre-extracted credentials, decide
|
||||
the principal kind and resolve it. Used by the proxy auth chain after it
|
||||
already pulled the api-key out of the request.
|
||||
|
||||
- ``resolve_identity`` — request-scoped composition for new entrypoints. Not
|
||||
yet wired into ``user_api_key_auth.py``; lives here so callers without a
|
||||
hashed-token-in-hand (CLI, MCP, background jobs) can still build a
|
||||
``IdentityContext``.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
|
|
@ -13,7 +22,10 @@ from litellm.constants import (
|
|||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
)
|
||||
from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext
|
||||
from litellm.identity.extractors.api_key import extract_api_key_principal
|
||||
from litellm.identity.extractors.api_key import (
|
||||
extract_api_key_principal,
|
||||
hash_principal_token,
|
||||
)
|
||||
from litellm.identity.extractors.client import extract_client_info
|
||||
from litellm.identity.extractors.end_user import extract_end_user_id
|
||||
from litellm.identity.extractors.header import extract_audit_changed_by
|
||||
|
|
@ -23,6 +35,15 @@ from litellm.identity.principal import (
|
|||
Principal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
from litellm.integrations.otel.runtime import traced
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
from litellm.identity.cache import IdentityCache
|
||||
|
||||
_SERVICE_ACCOUNT_API_KEYS = frozenset(
|
||||
{
|
||||
|
|
@ -33,7 +54,7 @@ _SERVICE_ACCOUNT_API_KEYS = frozenset(
|
|||
)
|
||||
|
||||
|
||||
def _resolve_principal(api_key: Optional[str]) -> Principal:
|
||||
def _principal_from_raw_key(api_key: Optional[str]) -> Principal:
|
||||
if api_key and api_key in _SERVICE_ACCOUNT_API_KEYS:
|
||||
return ServiceAccountPrincipal(name=api_key)
|
||||
|
||||
|
|
@ -48,7 +69,15 @@ def _resolve_principal(api_key: Optional[str]) -> Principal:
|
|||
return AnonymousPrincipal()
|
||||
|
||||
|
||||
def resolve_identity(
|
||||
@traced(
|
||||
"identity.resolve",
|
||||
role=SpanRole.SERVICE,
|
||||
attrs=lambda result: {
|
||||
"identity.principal.kind": result.principal.kind,
|
||||
"identity.end_user.present": result.end_user_id is not None,
|
||||
},
|
||||
)
|
||||
async def resolve_identity(
|
||||
*,
|
||||
api_key: Optional[str] = None,
|
||||
request: Any = None,
|
||||
|
|
@ -56,7 +85,13 @@ def resolve_identity(
|
|||
headers: Optional[Dict[str, Any]] = None,
|
||||
general_settings: Optional[Dict[str, Any]] = None,
|
||||
) -> IdentityContext:
|
||||
principal = _resolve_principal(api_key)
|
||||
"""Build an ``IdentityContext`` purely from request-side signals.
|
||||
|
||||
Does not touch the database. The hydrated-row variant lives in
|
||||
``store.load_identity`` and is composed by callers that have a
|
||||
``PrismaClient`` + ``IdentityCache`` in hand.
|
||||
"""
|
||||
principal = _principal_from_raw_key(api_key)
|
||||
end_user_id = extract_end_user_id(body=body, headers=headers)
|
||||
audit = AuditInfo(changed_by=extract_audit_changed_by(headers))
|
||||
client: ClientInfo
|
||||
|
|
@ -73,3 +108,32 @@ def resolve_identity(
|
|||
audit=audit,
|
||||
client=client,
|
||||
)
|
||||
|
||||
|
||||
async def resolve_user_api_key_auth(
|
||||
*,
|
||||
api_key: str,
|
||||
prisma_client: "PrismaClient",
|
||||
identity_cache: "IdentityCache",
|
||||
user_api_key_cache: "UserApiKeyCache",
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj: Optional["ProxyLogging"] = None,
|
||||
) -> "UserAPIKeyAuth":
|
||||
"""Cache-or-DB resolve a hashed virtual key into ``UserAPIKeyAuth``.
|
||||
|
||||
This is the surface ``_user_api_key_auth_builder`` calls in place of
|
||||
the legacy ``get_key_object``. The hashed-token computation is
|
||||
centralized via ``hash_principal_token`` so the auth chain doesn't
|
||||
re-implement the JWT-vs-key hashing rules.
|
||||
"""
|
||||
from litellm.identity.store import load_identity
|
||||
|
||||
hashed_token = hash_principal_token(api_key)
|
||||
return await load_identity(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
|
|||
50
litellm/identity/runtime.py
Normal file
50
litellm/identity/runtime.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
"""Process-wide accessor for the shared ``IdentityCache``.
|
||||
|
||||
The proxy already owns one ``DualCache`` for caller identity at
|
||||
``litellm.proxy.proxy_server.user_api_key_cache``. We layer an
|
||||
``IdentityCache`` on top of it so the new identity load path shares a
|
||||
single in-memory/Redis backend with the legacy caches. This avoids
|
||||
double-caching on a single deploy and keeps invalidation surfaces
|
||||
aligned.
|
||||
|
||||
Off-proxy callers (CLI, tests) can pass their own ``DualCache``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from litellm.identity.cache import IdentityCache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
||||
_identity_cache: Optional[IdentityCache] = None
|
||||
|
||||
|
||||
def get_identity_cache(
|
||||
dual_cache: Optional["DualCache"] = None,
|
||||
) -> IdentityCache:
|
||||
"""Return the shared ``IdentityCache``, building it on first call.
|
||||
|
||||
When ``dual_cache`` is omitted, the proxy's module-level cache is
|
||||
used. The first call wins; subsequent calls ignore the argument so
|
||||
that every consumer in a process sees the same instance.
|
||||
"""
|
||||
global _identity_cache
|
||||
if _identity_cache is not None:
|
||||
return _identity_cache
|
||||
|
||||
if dual_cache is None:
|
||||
from litellm.proxy.proxy_server import user_api_key_cache as _proxy_cache
|
||||
|
||||
dual_cache = _proxy_cache
|
||||
|
||||
_identity_cache = IdentityCache(dual_cache=dual_cache)
|
||||
return _identity_cache
|
||||
|
||||
|
||||
def reset_identity_cache_for_tests() -> None:
|
||||
global _identity_cache
|
||||
_identity_cache = None
|
||||
232
litellm/identity/store.py
Normal file
232
litellm/identity/store.py
Normal file
|
|
@ -0,0 +1,232 @@
|
|||
"""Cold-path identity loader.
|
||||
|
||||
One async function, one I/O round trip on cache miss. The actual SQL
|
||||
JOIN lives in ``litellm.proxy.utils.PrismaClient.get_data`` (the
|
||||
``combined_view`` query) and is reused as-is so we don't duplicate the
|
||||
schema-coupled SQL. Object-permission rows that are referenced but not
|
||||
joined are lazily filled exactly as the legacy ``get_key_object`` did.
|
||||
|
||||
The cached payload is ``UserAPIKeyAuth`` — the proxy's existing carrier.
|
||||
Callers that want an ``IdentityContext`` view should call
|
||||
``uak.to_identity_context()`` at the consumption site.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from fastapi import status
|
||||
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
from litellm.integrations.otel.runtime import traced
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
from litellm.identity.cache import IdentityCache
|
||||
|
||||
|
||||
@traced(
|
||||
"identity.db.combined_view",
|
||||
role=SpanRole.DB_CALL,
|
||||
attrs=lambda result: {
|
||||
"identity.load.outcome": "found" if result is not None else "missing",
|
||||
"db.system.name": "postgresql",
|
||||
},
|
||||
)
|
||||
async def _fetch_from_db(
|
||||
*,
|
||||
hashed_token: str,
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_cache: "UserApiKeyCache",
|
||||
parent_otel_span,
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
) -> Optional["UserAPIKeyAuth"]:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_fetch_key_object_from_db_with_reconnect,
|
||||
get_object_permission,
|
||||
get_user_object,
|
||||
)
|
||||
|
||||
row = await _fetch_key_object_from_db_with_reconnect(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
|
||||
uak = UserAPIKeyAuth(**row.model_dump(exclude_none=True))
|
||||
|
||||
if uak.object_permission_id and not uak.object_permission:
|
||||
try:
|
||||
uak.object_permission = await get_object_permission(
|
||||
object_permission_id=uak.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Bundle the user row into the identity cache so the auth chain's
|
||||
# follow-up `get_user_object` lookup is a free in-memory read on the
|
||||
# next request (the user object survives the same TTL as the rest of
|
||||
# the identity entry). The lookup itself is wrapped: failures fall
|
||||
# back to the legacy per-row cache that `get_user_object` populates.
|
||||
if uak.user_id:
|
||||
try:
|
||||
uak.user = await get_user_object(
|
||||
user_id=uak.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception:
|
||||
uak.user = None
|
||||
|
||||
return uak
|
||||
|
||||
|
||||
def _principal_kind(uak: "UserAPIKeyAuth") -> str:
|
||||
if uak.api_key and uak.api_key == uak.team_id:
|
||||
return "service_account"
|
||||
if uak.jwt_claims:
|
||||
return "jwt"
|
||||
if uak.token:
|
||||
return "api_key"
|
||||
return "anonymous"
|
||||
|
||||
|
||||
@traced(
|
||||
"identity.load",
|
||||
role=SpanRole.SERVICE,
|
||||
attrs=lambda result: {
|
||||
"identity.principal.kind": _principal_kind(result),
|
||||
"identity.principal.has_team": bool(result.team_id),
|
||||
"identity.principal.has_org": bool(result.org_id),
|
||||
"identity.principal.has_project": bool(result.project_id),
|
||||
"identity.principal.has_agent": bool(result.agent_id),
|
||||
},
|
||||
)
|
||||
async def load_identity(
|
||||
*,
|
||||
hashed_token: str,
|
||||
prisma_client: "PrismaClient",
|
||||
cache: "IdentityCache",
|
||||
user_api_key_cache: "UserApiKeyCache",
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj: Optional["ProxyLogging"] = None,
|
||||
) -> "UserAPIKeyAuth":
|
||||
"""Return the hydrated ``UserAPIKeyAuth`` for a hashed token.
|
||||
|
||||
Cache-hit returns immediately. Cache-miss issues exactly one combined
|
||||
view SQL query. A missing key raises ``ProxyException`` matching the
|
||||
legacy ``get_key_object`` contract so callers don't need to special
|
||||
case the cutover.
|
||||
"""
|
||||
cached = await cache.get(hashed_token)
|
||||
if cached is not None:
|
||||
_rehydrate_bundled_user(cached)
|
||||
return _hydrated_copy(cached)
|
||||
|
||||
if prisma_client is None:
|
||||
raise Exception(
|
||||
"No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys"
|
||||
)
|
||||
|
||||
uak = await _fetch_from_db(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if uak is None:
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"Authentication Error, Invalid proxy server token passed. "
|
||||
f"key={hashed_token}, not found in db. Create key via "
|
||||
"`/key/generate` call."
|
||||
),
|
||||
type=ProxyErrorTypes.token_not_found_in_db,
|
||||
param="key",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
|
||||
await cache.set(hashed_token, uak)
|
||||
await _populate_legacy_cache(
|
||||
hashed_token=hashed_token,
|
||||
uak=uak,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
return _hydrated_copy(uak)
|
||||
|
||||
|
||||
async def _populate_legacy_cache(
|
||||
*,
|
||||
hashed_token: str,
|
||||
uak: "UserAPIKeyAuth",
|
||||
user_api_key_cache: "UserApiKeyCache",
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
) -> None:
|
||||
"""Keep the legacy ``user_api_key_cache`` populated.
|
||||
|
||||
The pre-DB cache peek in ``_user_api_key_auth_builder`` (the
|
||||
``check_cache_only=True`` admin fast-path) reads from this cache, so
|
||||
populating it on every cold load preserves that fast path with no
|
||||
code change at the call site.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _cache_key_object
|
||||
|
||||
try:
|
||||
await _cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
user_api_key_obj=uak,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _rehydrate_bundled_user(uak: "UserAPIKeyAuth") -> None:
|
||||
"""Coerce the bundled user dict back into ``LiteLLM_UserTable``.
|
||||
|
||||
``UserAPIKeyAuth.user`` is typed ``Any`` so the codec round-trips it
|
||||
as a plain dict. Consumers read scalar attributes (`tpm_limit`,
|
||||
`metadata`, `user_role`, …) off the model, so we restore the typed
|
||||
form before handing the cached entry back.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
|
||||
raw = getattr(uak, "user", None)
|
||||
if isinstance(raw, dict):
|
||||
try:
|
||||
uak.user = LiteLLM_UserTable(**raw)
|
||||
except Exception:
|
||||
uak.user = None
|
||||
|
||||
|
||||
def _hydrated_copy(uak: "UserAPIKeyAuth") -> "UserAPIKeyAuth":
|
||||
"""Return a copy safe to mutate per-request.
|
||||
|
||||
The cached entry is shared across requests; consumers mutate fields
|
||||
like ``parent_otel_span``, ``request_route``, and ``end_user_id`` on
|
||||
the returned object, so we hand each caller its own model copy with
|
||||
those request-scoped fields cleared.
|
||||
"""
|
||||
copy = uak.model_copy()
|
||||
copy.parent_otel_span = None
|
||||
copy.request_route = None
|
||||
copy.budget_reservation = None
|
||||
return copy
|
||||
|
|
@ -1,14 +1,18 @@
|
|||
"""SDK-free entrypoints for proxy-core call sites (auth, …).
|
||||
"""SDK-free entrypoints for proxy-core call sites (auth, identity, …).
|
||||
|
||||
Proxy code may run without the OpenTelemetry SDK installed, so it must not import
|
||||
``litellm.integrations.otel.logger`` (which imports the SDK at module scope) at
|
||||
module load. These wrappers import it lazily and no-op when the SDK is absent or
|
||||
V2 is not the active logger — so a call site can wrap a request phase or seed
|
||||
identity unconditionally.
|
||||
V2 is not the active logger — so a call site can wrap a request phase, seed
|
||||
identity, or decorate an async function with ``@traced`` unconditionally.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import inspect
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Iterator
|
||||
from typing import Any, Callable, Iterator, Mapping, Optional
|
||||
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
|
||||
|
||||
@contextmanager
|
||||
|
|
@ -36,3 +40,94 @@ def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None:
|
|||
except Exception:
|
||||
return
|
||||
_seed_request_identity(user_api_key_dict, model=model)
|
||||
|
||||
|
||||
def _apply_span_attrs(span: Any, attrs: Optional[Mapping[str, Any]]) -> None:
|
||||
if span is None or not attrs:
|
||||
return
|
||||
for key, value in attrs.items():
|
||||
if value is None:
|
||||
continue
|
||||
try:
|
||||
span.set_attribute(key, value)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _resolve_attrs(
|
||||
builder: Optional[Callable[..., Mapping[str, Any]]],
|
||||
*,
|
||||
args: tuple,
|
||||
kwargs: dict,
|
||||
result: Any,
|
||||
) -> Mapping[str, Any]:
|
||||
if builder is None:
|
||||
return {}
|
||||
try:
|
||||
sig = inspect.signature(builder)
|
||||
accepts_result = "result" in sig.parameters or any(
|
||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||
)
|
||||
if accepts_result:
|
||||
payload = builder(*args, result=result, **kwargs)
|
||||
else:
|
||||
payload = builder(result)
|
||||
return payload or {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def traced(
|
||||
span_name: str,
|
||||
*,
|
||||
role: SpanRole = SpanRole.SERVICE,
|
||||
attrs: Optional[Callable[..., Mapping[str, Any]]] = None,
|
||||
) -> Callable:
|
||||
"""Wrap an async function so it runs inside an OTel v2 span.
|
||||
|
||||
``SpanRole.SERVICE`` opens a phase span under the active root proxy request
|
||||
span; nested DB/service calls become its children. ``SpanRole.DB_CALL`` (and
|
||||
any other non-SERVICE role) does not open a new span — it just attaches
|
||||
attributes to the current span, so we don't double-count for code paths
|
||||
that already produce a DB_CALL via service-logger plumbing.
|
||||
|
||||
``attrs`` is an optional callback that receives the function's return value
|
||||
(and the wrapped function's args/kwargs via ``**kwargs`` if it accepts
|
||||
them) and returns a mapping of span attributes to set on success. Any
|
||||
exception in ``attrs`` is swallowed — instrumentation must never break the
|
||||
request.
|
||||
"""
|
||||
|
||||
def _decorator(func: Callable) -> Callable:
|
||||
if not inspect.iscoroutinefunction(func):
|
||||
raise TypeError(
|
||||
f"@traced requires an async function; got {func!r}"
|
||||
)
|
||||
|
||||
@functools.wraps(func)
|
||||
async def _wrapper(*args, **kwargs):
|
||||
if role is SpanRole.SERVICE:
|
||||
with phase_span(span_name) as span:
|
||||
result = await func(*args, **kwargs)
|
||||
_apply_span_attrs(
|
||||
span,
|
||||
_resolve_attrs(attrs, args=args, kwargs=kwargs, result=result),
|
||||
)
|
||||
return result
|
||||
|
||||
result = await func(*args, **kwargs)
|
||||
try:
|
||||
from opentelemetry import trace as _otel_trace
|
||||
|
||||
current_span = _otel_trace.get_current_span()
|
||||
except Exception:
|
||||
current_span = None
|
||||
_apply_span_attrs(
|
||||
current_span,
|
||||
_resolve_attrs(attrs, args=args, kwargs=kwargs, result=result),
|
||||
)
|
||||
return result
|
||||
|
||||
return _wrapper
|
||||
|
||||
return _decorator
|
||||
|
|
|
|||
|
|
@ -1869,6 +1869,12 @@ async def _delete_cache_key_object(
|
|||
key=key
|
||||
)
|
||||
|
||||
from litellm.identity.invalidation import invalidate_identity_for_token
|
||||
|
||||
await invalidate_identity_for_token(
|
||||
token_hash=hashed_token, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def _get_team_db_check(
|
||||
|
|
|
|||
|
|
@ -1,15 +1,16 @@
|
|||
import base64
|
||||
import os
|
||||
from typing import Dict, Optional, Tuple, cast
|
||||
from typing import Dict, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.identity import build_user_api_key_auth_from_oauth2_response
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
|
||||
|
||||
class Oauth2Handler:
|
||||
|
|
@ -86,31 +87,6 @@ class Oauth2Handler:
|
|||
"""
|
||||
return {"Authorization": f"Bearer {token}", "Content-Type": "application/json"}
|
||||
|
||||
@staticmethod
|
||||
def _extract_user_info(
|
||||
response_data: Dict,
|
||||
user_id_field_name: str,
|
||||
user_role_field_name: str,
|
||||
user_team_id_field_name: str,
|
||||
) -> Tuple[Optional[str], Optional[str], Optional[str]]:
|
||||
"""
|
||||
Extract user information from OAuth2 response.
|
||||
|
||||
Args:
|
||||
response_data: The response data from OAuth2 endpoint
|
||||
user_id_field_name: Field name for user ID
|
||||
user_role_field_name: Field name for user role
|
||||
user_team_id_field_name: Field name for team ID
|
||||
|
||||
Returns:
|
||||
Tuple of (user_id, user_role, user_team_id)
|
||||
"""
|
||||
user_id = response_data.get(user_id_field_name)
|
||||
user_team_id = response_data.get(user_team_id_field_name)
|
||||
user_role = response_data.get(user_role_field_name)
|
||||
|
||||
return user_id, user_role, user_team_id
|
||||
|
||||
@staticmethod
|
||||
async def check_oauth2_token(token: str) -> UserAPIKeyAuth:
|
||||
"""
|
||||
|
|
@ -202,20 +178,13 @@ class Oauth2Handler:
|
|||
if is_introspection_endpoint and not data.get("active", True):
|
||||
raise ValueError("Token is not active")
|
||||
|
||||
# Extract user information from response
|
||||
user_id, user_role, user_team_id = Oauth2Handler._extract_user_info(
|
||||
return build_user_api_key_auth_from_oauth2_response(
|
||||
token=token,
|
||||
response_data=data,
|
||||
user_id_field_name=user_id_field_name,
|
||||
user_role_field_name=user_role_field_name,
|
||||
user_team_id_field_name=user_team_id_field_name,
|
||||
)
|
||||
|
||||
return UserAPIKeyAuth(
|
||||
api_key=token,
|
||||
team_id=user_team_id,
|
||||
user_id=user_id,
|
||||
user_role=cast(LitellmUserRoles, user_role),
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
# This will catch any 4xx or 5xx errors
|
||||
raise ValueError(f"Oauth 2.0 Token validation failed: {e}")
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ import litellm
|
|||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
|
||||
from litellm.identity import build_user_api_key_auth_from_jwt_result
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
from litellm.integrations.otel.runtime import phase_span, seed_request_identity
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
|
|
@ -1141,10 +1142,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
+ CommonProxyErrors.not_premium_user.value
|
||||
)
|
||||
|
||||
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
||||
with phase_span("identity.oauth2"):
|
||||
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
||||
|
||||
if general_settings.get("enable_oauth2_proxy_auth", False) is True:
|
||||
return await handle_oauth2_proxy_request(request=request)
|
||||
with phase_span("identity.oauth2_proxy"):
|
||||
return await handle_oauth2_proxy_request(request=request)
|
||||
|
||||
if general_settings.get("enable_jwt_auth", False) is True:
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
|
@ -1191,7 +1194,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
# standard JWT auth_builder below
|
||||
|
||||
if do_standard_jwt_auth:
|
||||
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
|
||||
with phase_span("identity.jwt"):
|
||||
result = await JWTAuthManager.auth_builder(
|
||||
request_data=request_data,
|
||||
general_settings=general_settings,
|
||||
|
|
@ -1210,14 +1213,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
is_proxy_admin = result["is_proxy_admin"]
|
||||
team_id = result["team_id"]
|
||||
team_object = result["team_object"]
|
||||
user_id = result["user_id"]
|
||||
user_object = result["user_object"]
|
||||
end_user_id = result["end_user_id"]
|
||||
org_id = result["org_id"]
|
||||
team_membership: Optional[LiteLLM_TeamMembership] = result.get(
|
||||
"team_membership", None
|
||||
)
|
||||
jwt_claims = result.get("jwt_claims", None)
|
||||
|
||||
if is_proxy_admin:
|
||||
|
|
@ -1234,90 +1232,16 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
value=_JWT_PROXY_ADMIN_SENTINEL,
|
||||
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
|
||||
)
|
||||
return UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
team_alias=(
|
||||
team_object.team_alias
|
||||
if team_object is not None
|
||||
else None
|
||||
),
|
||||
team_tpm_limit=(
|
||||
team_object.tpm_limit
|
||||
if team_object is not None
|
||||
else None
|
||||
),
|
||||
team_rpm_limit=(
|
||||
team_object.rpm_limit
|
||||
if team_object is not None
|
||||
else None
|
||||
),
|
||||
team_models=(
|
||||
team_object.models if team_object is not None else []
|
||||
),
|
||||
team_metadata=(
|
||||
team_object.metadata
|
||||
if team_object is not None
|
||||
else None
|
||||
),
|
||||
org_id=org_id,
|
||||
end_user_id=end_user_id,
|
||||
return build_user_api_key_auth_from_jwt_result(
|
||||
result=result,
|
||||
parent_otel_span=parent_otel_span,
|
||||
jwt_claims=jwt_claims,
|
||||
is_proxy_admin=True,
|
||||
)
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
team_id=team_id,
|
||||
team_alias=(
|
||||
team_object.team_alias if team_object is not None else None
|
||||
),
|
||||
team_tpm_limit=(
|
||||
team_object.tpm_limit if team_object is not None else None
|
||||
),
|
||||
team_rpm_limit=(
|
||||
team_object.rpm_limit if team_object is not None else None
|
||||
),
|
||||
team_models=(
|
||||
team_object.models if team_object is not None else []
|
||||
),
|
||||
user_role=(
|
||||
LitellmUserRoles(user_object.user_role)
|
||||
if user_object is not None
|
||||
and user_object.user_role is not None
|
||||
else LitellmUserRoles.INTERNAL_USER
|
||||
),
|
||||
user_id=user_id,
|
||||
org_id=org_id,
|
||||
valid_token = build_user_api_key_auth_from_jwt_result(
|
||||
result=result,
|
||||
parent_otel_span=parent_otel_span,
|
||||
end_user_id=end_user_id,
|
||||
user_tpm_limit=(
|
||||
user_object.tpm_limit if user_object is not None else None
|
||||
),
|
||||
user_rpm_limit=(
|
||||
user_object.rpm_limit if user_object is not None else None
|
||||
),
|
||||
team_member_rpm_limit=(
|
||||
team_membership.safe_get_team_member_rpm_limit()
|
||||
if team_membership is not None
|
||||
else None
|
||||
),
|
||||
team_member_tpm_limit=(
|
||||
team_membership.safe_get_team_member_tpm_limit()
|
||||
if team_membership is not None
|
||||
else None
|
||||
),
|
||||
team_metadata=(
|
||||
team_object.metadata if team_object is not None else None
|
||||
),
|
||||
jwt_claims=jwt_claims,
|
||||
)
|
||||
valid_token.team_object_permission = (
|
||||
team_object.object_permission
|
||||
if team_object is not None
|
||||
else None
|
||||
is_proxy_admin=False,
|
||||
)
|
||||
|
||||
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
|
||||
|
|
@ -1679,9 +1603,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
try:
|
||||
with tracer.trace("litellm.proxy.auth.get_key_object_from_db"):
|
||||
valid_token = await get_key_object(
|
||||
from litellm.identity import (
|
||||
get_identity_cache,
|
||||
load_identity,
|
||||
)
|
||||
|
||||
valid_token = await load_identity(
|
||||
hashed_token=api_key,
|
||||
prisma_client=prisma_client,
|
||||
cache=get_identity_cache(user_api_key_cache),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
|
|
@ -1737,23 +1667,27 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
# Check 2. If user_id for this token is in budget - done in common_checks()
|
||||
if valid_token.user_id is not None:
|
||||
try:
|
||||
with tracer.trace("litellm.proxy.auth.get_user_object"):
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
bundled_user = getattr(valid_token, "user", None)
|
||||
if isinstance(bundled_user, LiteLLM_UserTable):
|
||||
user_obj = bundled_user
|
||||
else:
|
||||
try:
|
||||
with tracer.trace("litellm.proxy.auth.get_user_object"):
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
user_obj = None
|
||||
user_obj = None
|
||||
|
||||
if (
|
||||
user_obj is not None
|
||||
|
|
|
|||
|
|
@ -1373,6 +1373,16 @@ async def _update_single_user_helper(
|
|||
detail={"error": "Failed to update user"},
|
||||
)
|
||||
_strip_password_from_response(response)
|
||||
|
||||
updated_user_id = response.get("user_id") if isinstance(response, dict) else None
|
||||
if updated_user_id:
|
||||
from litellm.identity.invalidation import invalidate_identity_for_user
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
await invalidate_identity_for_user(
|
||||
user_id=updated_user_id, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -2325,6 +2335,14 @@ async def delete_user(
|
|||
where={"user_id": {"in": data.user_ids}}
|
||||
)
|
||||
|
||||
from litellm.identity.invalidation import invalidate_identity_for_user
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
for user_id in data.user_ids:
|
||||
await invalidate_identity_for_user(
|
||||
user_id=user_id, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
return deleted_users
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -589,6 +589,13 @@ async def update_organization(
|
|||
include={"members": True, "teams": True, "litellm_budget_table": True},
|
||||
)
|
||||
|
||||
from litellm.identity.invalidation import invalidate_identity_for_org
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
await invalidate_identity_for_org(
|
||||
org_id=data.organization_id, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -676,6 +683,14 @@ async def delete_organization(
|
|||
)
|
||||
deleted_orgs.append(deleted_org)
|
||||
|
||||
from litellm.identity.invalidation import invalidate_identity_for_org
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
for organization_id in data.organization_ids:
|
||||
await invalidate_identity_for_org(
|
||||
org_id=organization_id, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
return deleted_orgs
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -171,6 +171,12 @@ async def _refresh_cached_team(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
from litellm.identity.invalidation import invalidate_identity_for_team
|
||||
|
||||
await invalidate_identity_for_team(
|
||||
team_id=team_row.team_id, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
|
||||
async def _verify_team_access(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
|
|
@ -3294,6 +3300,15 @@ async def delete_team(
|
|||
deleted_teams = await prisma_client.delete_data(
|
||||
team_id_list=data.team_ids, table_name="team"
|
||||
)
|
||||
|
||||
from litellm.identity.invalidation import invalidate_identity_for_team
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
for team_id in data.team_ids:
|
||||
await invalidate_identity_for_team(
|
||||
team_id=team_id, dual_cache=user_api_key_cache
|
||||
)
|
||||
|
||||
return deleted_teams
|
||||
|
||||
|
||||
|
|
|
|||
77
tests/test_litellm/identity/test_cache.py
Normal file
77
tests/test_litellm/identity/test_cache.py
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.identity.cache import (
|
||||
IdentityCache,
|
||||
identity_cache_key,
|
||||
team_generation_key,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
|
||||
def _user_api_key_cache() -> UserApiKeyCache:
|
||||
return UserApiKeyCache()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_then_get_returns_same_uak():
|
||||
cache = IdentityCache(dual_cache=_user_api_key_cache())
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t1")
|
||||
token_hash = uak.token
|
||||
assert token_hash is not None
|
||||
|
||||
await cache.set(token_hash, uak)
|
||||
got = await cache.get(token_hash)
|
||||
assert got is not None
|
||||
assert got.user_id == "u1"
|
||||
assert got.team_id == "t1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_miss_returns_none():
|
||||
cache = IdentityCache(dual_cache=_user_api_key_cache())
|
||||
assert await cache.get("missing-hash") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_clears_entry():
|
||||
dual = _user_api_key_cache()
|
||||
cache = IdentityCache(dual_cache=dual)
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1")
|
||||
await cache.set(uak.token, uak)
|
||||
assert await cache.get(uak.token) is not None
|
||||
await cache.delete(uak.token)
|
||||
assert await cache.get(uak.token) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generation_bump_invalidates_team_scoped_entry():
|
||||
dual = _user_api_key_cache()
|
||||
cache = IdentityCache(dual_cache=dual)
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t-rotate")
|
||||
await cache.set(uak.token, uak)
|
||||
assert await cache.get(uak.token) is not None
|
||||
|
||||
await cache.bump_generation(team_generation_key("t-rotate"))
|
||||
assert await cache.get(uak.token) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generation_bump_for_unrelated_team_keeps_entry():
|
||||
dual = _user_api_key_cache()
|
||||
cache = IdentityCache(dual_cache=dual)
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t-keep")
|
||||
await cache.set(uak.token, uak)
|
||||
|
||||
await cache.bump_generation(team_generation_key("t-other"))
|
||||
assert await cache.get(uak.token) is not None
|
||||
|
||||
|
||||
def test_key_format_is_versioned():
|
||||
assert identity_cache_key("abc").startswith("identity:v1:")
|
||||
73
tests/test_litellm/identity/test_invalidation.py
Normal file
73
tests/test_litellm/identity/test_invalidation.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.identity.cache import IdentityCache, team_generation_key
|
||||
from litellm.identity.invalidation import (
|
||||
invalidate_identity_for_org,
|
||||
invalidate_identity_for_team,
|
||||
invalidate_identity_for_token,
|
||||
invalidate_identity_for_user,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_invalidation_drops_entry():
|
||||
backend = UserApiKeyCache()
|
||||
cache = IdentityCache(dual_cache=backend)
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1")
|
||||
await cache.set(uak.token, uak)
|
||||
assert await cache.get(uak.token) is not None
|
||||
|
||||
await invalidate_identity_for_token(
|
||||
token_hash=uak.token, dual_cache=backend
|
||||
)
|
||||
assert await cache.get(uak.token) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_invalidation_bumps_generation():
|
||||
backend = UserApiKeyCache()
|
||||
cache = IdentityCache(dual_cache=backend)
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t-rotate")
|
||||
await cache.set(uak.token, uak)
|
||||
|
||||
await invalidate_identity_for_team(
|
||||
team_id="t-rotate", dual_cache=backend
|
||||
)
|
||||
assert await cache.get(uak.token) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_invalidation_is_scoped_to_user():
|
||||
backend = UserApiKeyCache()
|
||||
cache = IdentityCache(dual_cache=backend)
|
||||
uak_a = UserAPIKeyAuth(api_key="sk-a", user_id="u-rotate", team_id="t1")
|
||||
uak_b = UserAPIKeyAuth(api_key="sk-b", user_id="u-keep", team_id="t1")
|
||||
await cache.set(uak_a.token, uak_a)
|
||||
await cache.set(uak_b.token, uak_b)
|
||||
|
||||
await invalidate_identity_for_user(
|
||||
user_id="u-rotate", dual_cache=backend
|
||||
)
|
||||
|
||||
assert await cache.get(uak_a.token) is None
|
||||
assert await cache.get(uak_b.token) is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_org_invalidation_drops_org_scoped_entry():
|
||||
backend = UserApiKeyCache()
|
||||
cache = IdentityCache(dual_cache=backend)
|
||||
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", org_id="org-rotate")
|
||||
await cache.set(uak.token, uak)
|
||||
|
||||
await invalidate_identity_for_org(
|
||||
org_id="org-rotate", dual_cache=backend
|
||||
)
|
||||
assert await cache.get(uak.token) is None
|
||||
144
tests/test_litellm/identity/test_jwt_builder.py
Normal file
144
tests/test_litellm/identity/test_jwt_builder.py
Normal file
|
|
@ -0,0 +1,144 @@
|
|||
import os
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from litellm.identity import build_user_api_key_auth_from_jwt_result
|
||||
from litellm.identity.principal import JWTPrincipal
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
|
||||
|
||||
def _team(team_id="t-jwt"):
|
||||
return LiteLLM_TeamTableCachedObj(
|
||||
team_id=team_id,
|
||||
team_alias=f"alias-{team_id}",
|
||||
tpm_limit=100,
|
||||
rpm_limit=10,
|
||||
models=["gpt-4o"],
|
||||
metadata={"env": "prod"},
|
||||
)
|
||||
|
||||
|
||||
def _user(user_id="u-jwt", role="internal_user"):
|
||||
return LiteLLM_UserTable(
|
||||
user_id=user_id,
|
||||
user_email="jwt@litellm.io",
|
||||
user_role=role,
|
||||
tpm_limit=50,
|
||||
rpm_limit=5,
|
||||
)
|
||||
|
||||
|
||||
def _membership(tpm=42, rpm=7):
|
||||
return SimpleNamespace(
|
||||
safe_get_team_member_tpm_limit=lambda: tpm,
|
||||
safe_get_team_member_rpm_limit=lambda: rpm,
|
||||
)
|
||||
|
||||
|
||||
def _auth_builder_result(
|
||||
*,
|
||||
is_proxy_admin: bool = False,
|
||||
team=None,
|
||||
user=None,
|
||||
team_membership=None,
|
||||
jwt_claims=None,
|
||||
):
|
||||
return {
|
||||
"is_proxy_admin": is_proxy_admin,
|
||||
"team_id": team.team_id if team is not None else None,
|
||||
"team_object": team,
|
||||
"user_id": user.user_id if user is not None else None,
|
||||
"user_object": user,
|
||||
"end_user_id": "eu-jwt",
|
||||
"org_id": "org-jwt",
|
||||
"team_membership": team_membership,
|
||||
"jwt_claims": jwt_claims
|
||||
or {"sub": "u-jwt", "iss": "https://idp", "scope": "read write"},
|
||||
}
|
||||
|
||||
|
||||
def test_proxy_admin_path_marks_role_and_drops_membership_fields():
|
||||
result = _auth_builder_result(is_proxy_admin=True, team=_team(), user=_user())
|
||||
uak = build_user_api_key_auth_from_jwt_result(
|
||||
result=result, parent_otel_span=None, is_proxy_admin=True
|
||||
)
|
||||
assert isinstance(uak, UserAPIKeyAuth)
|
||||
assert uak.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
assert uak.team_id == "t-jwt"
|
||||
assert uak.team_alias == "alias-t-jwt"
|
||||
assert uak.team_tpm_limit == 100
|
||||
assert uak.team_metadata == {"env": "prod"}
|
||||
assert uak.team_member_rpm_limit is None
|
||||
assert uak.team_member_tpm_limit is None
|
||||
assert uak.user_tpm_limit is None
|
||||
assert uak.jwt_claims == result["jwt_claims"]
|
||||
|
||||
|
||||
def test_regular_path_layers_team_user_membership_limits():
|
||||
membership = _membership(tpm=42, rpm=7)
|
||||
result = _auth_builder_result(
|
||||
team=_team(), user=_user(role="internal_user"), team_membership=membership
|
||||
)
|
||||
uak = build_user_api_key_auth_from_jwt_result(
|
||||
result=result, parent_otel_span=None, is_proxy_admin=False
|
||||
)
|
||||
assert uak.user_role == LitellmUserRoles.INTERNAL_USER
|
||||
assert uak.user_tpm_limit == 50
|
||||
assert uak.user_rpm_limit == 5
|
||||
assert uak.team_member_tpm_limit == 42
|
||||
assert uak.team_member_rpm_limit == 7
|
||||
assert uak.org_id == "org-jwt"
|
||||
assert uak.end_user_id == "eu-jwt"
|
||||
|
||||
|
||||
def test_missing_team_object_defaults_team_fields_safely():
|
||||
result = _auth_builder_result(team=None, user=_user())
|
||||
uak = build_user_api_key_auth_from_jwt_result(
|
||||
result=result, parent_otel_span=None, is_proxy_admin=False
|
||||
)
|
||||
assert uak.team_alias is None
|
||||
assert uak.team_tpm_limit is None
|
||||
assert uak.team_models == []
|
||||
assert uak.team_metadata is None
|
||||
|
||||
|
||||
def test_missing_user_object_defaults_user_role_to_internal():
|
||||
result = _auth_builder_result(team=_team(), user=None)
|
||||
uak = build_user_api_key_auth_from_jwt_result(
|
||||
result=result, parent_otel_span=None, is_proxy_admin=False
|
||||
)
|
||||
assert uak.user_role == LitellmUserRoles.INTERNAL_USER
|
||||
assert uak.user_tpm_limit is None
|
||||
|
||||
|
||||
def test_adapter_roundtrip_produces_jwt_principal():
|
||||
result = _auth_builder_result(team=_team(), user=_user())
|
||||
uak = build_user_api_key_auth_from_jwt_result(
|
||||
result=result, parent_otel_span=None, is_proxy_admin=False
|
||||
)
|
||||
ctx = uak.to_identity_context()
|
||||
assert isinstance(ctx.principal, JWTPrincipal)
|
||||
assert ctx.principal.sub == "u-jwt"
|
||||
assert ctx.principal.iss == "https://idp"
|
||||
assert ctx.principal.scopes == ["read", "write"]
|
||||
assert ctx.principal.mapped_user_id == "u-jwt"
|
||||
assert ctx.principal.mapped_team_id == "t-jwt"
|
||||
assert ctx.principal.mapped_org_id == "org-jwt"
|
||||
|
||||
|
||||
def test_team_object_permission_propagates_when_present():
|
||||
team = _team()
|
||||
team.object_permission = SimpleNamespace(object_permission_id="perm-1")
|
||||
result = _auth_builder_result(team=team, user=_user())
|
||||
uak = build_user_api_key_auth_from_jwt_result(
|
||||
result=result, parent_otel_span=None, is_proxy_admin=False
|
||||
)
|
||||
assert uak.team_object_permission is not None
|
||||
assert uak.team_object_permission.object_permission_id == "perm-1"
|
||||
55
tests/test_litellm/identity/test_oauth2_builder.py
Normal file
55
tests/test_litellm/identity/test_oauth2_builder.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from litellm.identity import build_user_api_key_auth_from_oauth2_response
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
def test_default_field_names_extract_from_introspection_response():
|
||||
response = {"sub": "u-1", "role": "internal_user", "team_id": "t-1"}
|
||||
uak = build_user_api_key_auth_from_oauth2_response(
|
||||
token="opaque-token", response_data=response
|
||||
)
|
||||
assert isinstance(uak, UserAPIKeyAuth)
|
||||
assert uak.user_id == "u-1"
|
||||
assert uak.user_role == "internal_user"
|
||||
assert uak.team_id == "t-1"
|
||||
|
||||
|
||||
def test_custom_field_names_override_defaults():
|
||||
response = {
|
||||
"preferred_username": "alice",
|
||||
"groups": "proxy_admin",
|
||||
"tenant_id": "tenant-9",
|
||||
}
|
||||
uak = build_user_api_key_auth_from_oauth2_response(
|
||||
token="t",
|
||||
response_data=response,
|
||||
user_id_field_name="preferred_username",
|
||||
user_role_field_name="groups",
|
||||
user_team_id_field_name="tenant_id",
|
||||
)
|
||||
assert uak.user_id == "alice"
|
||||
assert uak.user_role == "proxy_admin"
|
||||
assert uak.team_id == "tenant-9"
|
||||
|
||||
|
||||
def test_missing_fields_default_to_none():
|
||||
uak = build_user_api_key_auth_from_oauth2_response(
|
||||
token="t", response_data={}
|
||||
)
|
||||
assert uak.user_id is None
|
||||
assert uak.user_role is None
|
||||
assert uak.team_id is None
|
||||
|
||||
|
||||
def test_token_is_hashed_into_token_field():
|
||||
"""The api_key is hashed by the UserAPIKeyAuth validator; the
|
||||
OAuth2 builder must not bypass that path."""
|
||||
uak = build_user_api_key_auth_from_oauth2_response(
|
||||
token="sk-oauth2-test", response_data={"sub": "u"}
|
||||
)
|
||||
assert uak.token is not None
|
||||
assert uak.token != "sk-oauth2-test"
|
||||
74
tests/test_litellm/identity/test_observability.py
Normal file
74
tests/test_litellm/identity/test_observability.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
from litellm.integrations.otel.runtime import traced
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decorator_is_passthrough_when_v2_disabled():
|
||||
calls = []
|
||||
|
||||
@traced("identity.unit-test", role=SpanRole.SERVICE)
|
||||
async def work(x):
|
||||
calls.append(x)
|
||||
return x * 2
|
||||
|
||||
result = await work(5)
|
||||
assert result == 10
|
||||
assert calls == [5]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attrs_callback_invoked_with_result():
|
||||
captured = {}
|
||||
|
||||
@traced(
|
||||
"identity.attrs-test",
|
||||
role=SpanRole.SERVICE,
|
||||
attrs=lambda result: captured.setdefault("result", result) or {},
|
||||
)
|
||||
async def work():
|
||||
return {"principal_kind": "api_key"}
|
||||
|
||||
await work()
|
||||
assert captured["result"] == {"principal_kind": "api_key"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attrs_exception_is_swallowed():
|
||||
@traced(
|
||||
"identity.attrs-error",
|
||||
role=SpanRole.SERVICE,
|
||||
attrs=lambda result: 1 / 0,
|
||||
)
|
||||
async def work():
|
||||
return "ok"
|
||||
|
||||
assert await work() == "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_call_role_does_not_open_a_new_span():
|
||||
"""DB_CALL spans piggyback on the current span instead of starting one.
|
||||
|
||||
With V2 disabled there is no current span; the decorator should still
|
||||
run the function and swallow any attribute-application errors.
|
||||
"""
|
||||
|
||||
@traced("identity.db-test", role=SpanRole.DB_CALL)
|
||||
async def work():
|
||||
return "value"
|
||||
|
||||
assert await work() == "value"
|
||||
|
||||
|
||||
def test_decorator_rejects_sync_functions():
|
||||
with pytest.raises(TypeError):
|
||||
@traced("identity.sync-not-allowed", role=SpanRole.SERVICE)
|
||||
def sync_fn():
|
||||
return None
|
||||
|
|
@ -6,6 +6,8 @@ from types import SimpleNamespace
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.constants import LITTELM_CLI_SERVICE_ACCOUNT_NAME
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
|
|
@ -32,30 +34,35 @@ def _jwt(claims):
|
|||
return f"{b({'alg':'HS256','typ':'JWT'})}.{b(claims)}.sig"
|
||||
|
||||
|
||||
def test_anonymous_when_no_credentials():
|
||||
ctx = resolve_identity()
|
||||
@pytest.mark.asyncio
|
||||
async def test_anonymous_when_no_credentials():
|
||||
ctx = await resolve_identity()
|
||||
assert isinstance(ctx.principal, AnonymousPrincipal)
|
||||
|
||||
|
||||
def test_api_key_principal_for_sk_key():
|
||||
ctx = resolve_identity(api_key="sk-test")
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_principal_for_sk_key():
|
||||
ctx = await resolve_identity(api_key="sk-test")
|
||||
assert isinstance(ctx.principal, ApiKeyPrincipal)
|
||||
|
||||
|
||||
def test_jwt_principal_for_jwt_shaped_key():
|
||||
ctx = resolve_identity(api_key=_jwt({"sub": "u1"}))
|
||||
@pytest.mark.asyncio
|
||||
async def test_jwt_principal_for_jwt_shaped_key():
|
||||
ctx = await resolve_identity(api_key=_jwt({"sub": "u1"}))
|
||||
assert isinstance(ctx.principal, JWTPrincipal)
|
||||
assert ctx.principal.sub == "u1"
|
||||
|
||||
|
||||
def test_service_account_for_known_sentinel():
|
||||
ctx = resolve_identity(api_key=LITTELM_CLI_SERVICE_ACCOUNT_NAME)
|
||||
@pytest.mark.asyncio
|
||||
async def test_service_account_for_known_sentinel():
|
||||
ctx = await resolve_identity(api_key=LITTELM_CLI_SERVICE_ACCOUNT_NAME)
|
||||
assert isinstance(ctx.principal, ServiceAccountPrincipal)
|
||||
assert ctx.principal.name == LITTELM_CLI_SERVICE_ACCOUNT_NAME
|
||||
|
||||
|
||||
def test_end_user_and_audit_propagate():
|
||||
ctx = resolve_identity(
|
||||
@pytest.mark.asyncio
|
||||
async def test_end_user_and_audit_propagate():
|
||||
ctx = await resolve_identity(
|
||||
body={"user": "eu-42"},
|
||||
headers={"litellm-changed-by": "admin"},
|
||||
)
|
||||
|
|
@ -63,8 +70,9 @@ def test_end_user_and_audit_propagate():
|
|||
assert ctx.audit.changed_by == "admin"
|
||||
|
||||
|
||||
def test_client_info_from_request():
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_info_from_request():
|
||||
req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1")
|
||||
ctx = resolve_identity(request=req)
|
||||
ctx = await resolve_identity(request=req)
|
||||
assert ctx.client.ip == "127.0.0.1"
|
||||
assert ctx.client.user_agent == "curl/8"
|
||||
|
|
|
|||
229
tests/test_litellm/identity/test_store.py
Normal file
229
tests/test_litellm/identity/test_store.py
Normal file
|
|
@ -0,0 +1,229 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
from fastapi import status
|
||||
|
||||
from litellm.identity.cache import IdentityCache
|
||||
from litellm.identity.store import load_identity
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_VerificationTokenView,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
|
||||
def _stub_prisma_client():
|
||||
"""Minimal prisma_client stub with the get_data hook the store uses."""
|
||||
|
||||
class _Stub:
|
||||
def __init__(self):
|
||||
self.get_data = AsyncMock()
|
||||
self.calls = 0
|
||||
|
||||
return _Stub()
|
||||
|
||||
|
||||
def _verification_token_view(**fields) -> LiteLLM_VerificationTokenView:
|
||||
return LiteLLM_VerificationTokenView(
|
||||
token=fields.pop("token", "hash-x"),
|
||||
user_id=fields.pop("user_id", "u1"),
|
||||
team_id=fields.pop("team_id", "t1"),
|
||||
org_id=fields.pop("org_id", None),
|
||||
**fields,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_hit_skips_db():
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
prisma = _stub_prisma_client()
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
seed = UserAPIKeyAuth(
|
||||
token="hash-cached", user_id="u-cache", team_id="t-cache"
|
||||
)
|
||||
await identity_cache.set("hash-cached", seed)
|
||||
|
||||
result = await load_identity(
|
||||
hashed_token="hash-cached",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
|
||||
assert prisma.get_data.await_count == 0
|
||||
assert result.user_id == "u-cache"
|
||||
assert result.team_id == "t-cache"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_miss_hits_db_once_and_caches():
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
prisma = _stub_prisma_client()
|
||||
prisma.get_data.return_value = _verification_token_view(
|
||||
token="hash-db", user_id="u-db", team_id="t-db"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
|
||||
):
|
||||
result = await load_identity(
|
||||
hashed_token="hash-db",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
|
||||
assert prisma.get_data.await_count == 1
|
||||
assert result.user_id == "u-db"
|
||||
cached = await identity_cache.get("hash-db")
|
||||
assert cached is not None
|
||||
assert cached.user_id == "u-db"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_warm_load_after_cold_does_not_hit_db():
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
prisma = _stub_prisma_client()
|
||||
prisma.get_data.return_value = _verification_token_view(
|
||||
token="hash-warm", user_id="u-warm", team_id="t-warm"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
|
||||
):
|
||||
await load_identity(
|
||||
hashed_token="hash-warm",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
await load_identity(
|
||||
hashed_token="hash-warm",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
|
||||
assert prisma.get_data.await_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_token_raises_proxy_exception_with_401():
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
prisma = _stub_prisma_client()
|
||||
prisma.get_data.return_value = None
|
||||
|
||||
with pytest.raises(ProxyException) as excinfo:
|
||||
await load_identity(
|
||||
hashed_token="missing-hash",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
|
||||
assert int(excinfo.value.code) == status.HTTP_401_UNAUTHORIZED
|
||||
assert excinfo.value.type == ProxyErrorTypes.token_not_found_in_db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bundled_user_survives_cache_roundtrip_as_typed_model():
|
||||
from litellm.identity.cache import IdentityCache
|
||||
from litellm.identity.store import _rehydrate_bundled_user
|
||||
from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth
|
||||
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
uak = UserAPIKeyAuth(
|
||||
api_key="sk-bundle",
|
||||
user_id="u-bundle",
|
||||
user=LiteLLM_UserTable(
|
||||
user_id="u-bundle", user_email="x@y.com", tpm_limit=10
|
||||
),
|
||||
)
|
||||
await identity_cache.set(uak.token, uak)
|
||||
got = await identity_cache.get(uak.token)
|
||||
_rehydrate_bundled_user(got)
|
||||
|
||||
assert isinstance(got.user, LiteLLM_UserTable)
|
||||
assert got.user.user_email == "x@y.com"
|
||||
assert got.user.tpm_limit == 10
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cold_path_bundles_user_into_cache():
|
||||
"""``load_identity`` must populate ``user`` on first DB hit so the
|
||||
next request reads the user object from the identity cache."""
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
prisma = _stub_prisma_client()
|
||||
prisma.get_data.return_value = _verification_token_view(
|
||||
token="hash-user", user_id="u-bundle", team_id="t-bundle"
|
||||
)
|
||||
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
|
||||
async def _fake_get_user_object(*, user_id, **kwargs):
|
||||
return LiteLLM_UserTable(
|
||||
user_id=user_id, user_email="bundle@litellm.io", tpm_limit=42
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks.get_user_object",
|
||||
side_effect=_fake_get_user_object,
|
||||
):
|
||||
result = await load_identity(
|
||||
hashed_token="hash-user",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
|
||||
assert result.user is not None
|
||||
assert result.user.user_email == "bundle@litellm.io"
|
||||
cached = await identity_cache.get("hash-user")
|
||||
assert cached is not None
|
||||
assert cached.user is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hydrated_copy_is_request_scoped():
|
||||
cache_backend = UserApiKeyCache()
|
||||
identity_cache = IdentityCache(dual_cache=cache_backend)
|
||||
prisma = _stub_prisma_client()
|
||||
prisma.get_data.return_value = _verification_token_view(
|
||||
token="hash-copy", user_id="u-copy", team_id="t-copy"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
|
||||
):
|
||||
a = await load_identity(
|
||||
hashed_token="hash-copy",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
a.parent_otel_span = "request-A-span"
|
||||
a.request_route = "/v1/chat/completions"
|
||||
|
||||
b = await load_identity(
|
||||
hashed_token="hash-copy",
|
||||
prisma_client=prisma,
|
||||
cache=identity_cache,
|
||||
user_api_key_cache=cache_backend,
|
||||
)
|
||||
|
||||
assert b.parent_otel_span is None
|
||||
assert b.request_route is None
|
||||
Loading…
Add table
Reference in a new issue