mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix: identity module
This commit is contained in:
parent
f5b11b72a6
commit
cfc57d709f
21 changed files with 1130 additions and 0 deletions
39
litellm/identity/__init__.py
Normal file
39
litellm/identity/__init__.py
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
"""Caller-identity module.
|
||||
|
||||
Phase 1: domain types + extractors + bidirectional adapter to
|
||||
``UserAPIKeyAuth``. The proxy still drives identity through
|
||||
``litellm/proxy/auth/`` today; this module is the new home those flows
|
||||
will migrate to.
|
||||
|
||||
The public surface is small on purpose; downstream code should depend on
|
||||
``IdentityContext`` and the ``Principal`` union, not on individual
|
||||
extractor internals.
|
||||
"""
|
||||
|
||||
from litellm.identity.context import (
|
||||
AuditInfo,
|
||||
ClientInfo,
|
||||
IdentityContext,
|
||||
RequestIds,
|
||||
)
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
ApiKeyPrincipal,
|
||||
JWTPrincipal,
|
||||
Principal,
|
||||
SSOPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AnonymousPrincipal",
|
||||
"ApiKeyPrincipal",
|
||||
"AuditInfo",
|
||||
"ClientInfo",
|
||||
"IdentityContext",
|
||||
"JWTPrincipal",
|
||||
"Principal",
|
||||
"RequestIds",
|
||||
"SSOPrincipal",
|
||||
"ServiceAccountPrincipal",
|
||||
]
|
||||
151
litellm/identity/adapter.py
Normal file
151
litellm/identity/adapter.py
Normal file
|
|
@ -0,0 +1,151 @@
|
|||
"""Bidirectional bridge between ``IdentityContext`` and ``UserAPIKeyAuth``.
|
||||
|
||||
Phase 1 keeps the legacy Pydantic model as the universal carrier. These
|
||||
two pure functions let new code work in terms of ``IdentityContext``
|
||||
without forcing call sites to migrate today.
|
||||
|
||||
Invariants:
|
||||
- ``identity_context_to_user_api_key_auth(uak.to_identity_context())``
|
||||
preserves every identity-relevant field on ``uak``.
|
||||
- ``ApiKeyPrincipal.token_hash`` is treated as already-hashed; the
|
||||
Pydantic ``check_api_key`` validator does not re-hash it.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
)
|
||||
from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext, RequestIds
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
ApiKeyPrincipal,
|
||||
JWTPrincipal,
|
||||
Principal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_SERVICE_ACCOUNT_NAMES = frozenset(
|
||||
{
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _principal_from_uak(uak: "UserAPIKeyAuth") -> Principal:
|
||||
if uak.api_key in _SERVICE_ACCOUNT_NAMES or uak.key_alias in _SERVICE_ACCOUNT_NAMES:
|
||||
return ServiceAccountPrincipal(
|
||||
name=uak.api_key if uak.api_key in _SERVICE_ACCOUNT_NAMES else uak.key_alias # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
if uak.jwt_claims:
|
||||
claims = uak.jwt_claims
|
||||
scope_claim = claims.get("scope") or claims.get("scp") or ""
|
||||
if isinstance(scope_claim, list):
|
||||
scopes = [str(s) for s in scope_claim if s]
|
||||
elif isinstance(scope_claim, str):
|
||||
scopes = [s for s in scope_claim.split(" ") if s]
|
||||
else:
|
||||
scopes = []
|
||||
return JWTPrincipal(
|
||||
sub=claims.get("sub"),
|
||||
iss=claims.get("iss"),
|
||||
aud=claims.get("aud"),
|
||||
scopes=scopes,
|
||||
claims=dict(claims),
|
||||
mapped_user_id=uak.user_id,
|
||||
mapped_team_id=uak.team_id,
|
||||
mapped_org_id=uak.org_id,
|
||||
)
|
||||
|
||||
if uak.token:
|
||||
return ApiKeyPrincipal(
|
||||
token_hash=uak.token,
|
||||
key_alias=uak.key_alias,
|
||||
user_id=uak.user_id,
|
||||
team_id=uak.team_id,
|
||||
org_id=uak.org_id,
|
||||
project_id=uak.project_id,
|
||||
agent_id=uak.agent_id,
|
||||
)
|
||||
|
||||
return AnonymousPrincipal()
|
||||
|
||||
|
||||
def user_api_key_auth_to_identity_context(
|
||||
uak: "UserAPIKeyAuth",
|
||||
) -> IdentityContext:
|
||||
principal = _principal_from_uak(uak)
|
||||
return IdentityContext(
|
||||
principal=principal,
|
||||
end_user_id=uak.end_user_id,
|
||||
tags=[],
|
||||
access_group_ids=list(uak.access_group_ids or []),
|
||||
request=RequestIds(),
|
||||
client=ClientInfo(),
|
||||
audit=AuditInfo(),
|
||||
)
|
||||
|
||||
|
||||
def identity_context_to_user_api_key_auth(
|
||||
ctx: IdentityContext,
|
||||
) -> "UserAPIKeyAuth":
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
principal = ctx.principal
|
||||
kwargs: dict = {
|
||||
"end_user_id": ctx.end_user_id,
|
||||
"access_group_ids": list(ctx.access_group_ids) if ctx.access_group_ids else None,
|
||||
}
|
||||
|
||||
if isinstance(principal, ApiKeyPrincipal):
|
||||
kwargs.update(
|
||||
{
|
||||
"token": principal.token_hash,
|
||||
"key_alias": principal.key_alias,
|
||||
"user_id": principal.user_id,
|
||||
"team_id": principal.team_id,
|
||||
"org_id": principal.org_id,
|
||||
"project_id": principal.project_id,
|
||||
"agent_id": principal.agent_id,
|
||||
}
|
||||
)
|
||||
elif isinstance(principal, JWTPrincipal):
|
||||
kwargs.update(
|
||||
{
|
||||
"jwt_claims": dict(principal.claims),
|
||||
"user_id": principal.mapped_user_id,
|
||||
"team_id": principal.mapped_team_id,
|
||||
"org_id": principal.mapped_org_id,
|
||||
}
|
||||
)
|
||||
elif isinstance(principal, ServiceAccountPrincipal):
|
||||
if principal.name == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME:
|
||||
kwargs.update(
|
||||
{
|
||||
"api_key": principal.name,
|
||||
"team_id": "system",
|
||||
"key_alias": principal.name,
|
||||
"team_alias": "system",
|
||||
"user_id": "system",
|
||||
"user_role": LitellmUserRoles.PROXY_ADMIN,
|
||||
}
|
||||
)
|
||||
else:
|
||||
kwargs.update(
|
||||
{
|
||||
"api_key": principal.name,
|
||||
"team_id": principal.name,
|
||||
"key_alias": principal.name,
|
||||
"team_alias": principal.name,
|
||||
}
|
||||
)
|
||||
|
||||
return UserAPIKeyAuth(**{k: v for k, v in kwargs.items() if v is not None})
|
||||
47
litellm/identity/context.py
Normal file
47
litellm/identity/context.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
"""The per-request identity bundle.
|
||||
|
||||
``IdentityContext`` is what downstream consumers (auth, spend, guardrails,
|
||||
logging, audit) should read identity from once Phase 2 migration is done.
|
||||
In Phase 1 it travels alongside the legacy ``UserAPIKeyAuth`` via the
|
||||
adapter functions in ``litellm.identity.adapter``.
|
||||
|
||||
The bundle is mutable on purpose: identity fields like ``end_user_id`` are
|
||||
sometimes resolved or overridden after initial extraction, and the
|
||||
existing ``UserAPIKeyAuth`` mutation patterns must keep working.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
|
||||
from litellm.identity.principal import AnonymousPrincipal, Principal
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestIds:
|
||||
request_id: Optional[str] = None
|
||||
trace_id: Optional[str] = None
|
||||
session_id: Optional[str] = None
|
||||
mcp_session_id: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClientInfo:
|
||||
ip: Optional[str] = None
|
||||
user_agent: Optional[str] = None
|
||||
forwarded_chain: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AuditInfo:
|
||||
changed_by: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class IdentityContext:
|
||||
principal: Principal = field(default_factory=AnonymousPrincipal)
|
||||
end_user_id: Optional[str] = None
|
||||
tags: List[str] = field(default_factory=list)
|
||||
access_group_ids: List[str] = field(default_factory=list)
|
||||
request: RequestIds = field(default_factory=RequestIds)
|
||||
client: ClientInfo = field(default_factory=ClientInfo)
|
||||
audit: AuditInfo = field(default_factory=AuditInfo)
|
||||
7
litellm/identity/extractors/__init__.py
Normal file
7
litellm/identity/extractors/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""Identity extractors.
|
||||
|
||||
Each extractor wraps an existing helper in ``litellm/proxy/auth/`` and
|
||||
returns a piece of an ``IdentityContext``. Extractors must not introduce
|
||||
new behavior. If you need to change *how* a field is resolved, change the
|
||||
underlying helper and update the extractor's tests.
|
||||
"""
|
||||
36
litellm/identity/extractors/api_key.py
Normal file
36
litellm/identity/extractors/api_key.py
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
"""API-key principal extraction.
|
||||
|
||||
Wraps the existing ``get_api_key`` (header extraction) and
|
||||
``UserAPIKeyAuth._safe_hash_litellm_api_key`` (hashing, including the
|
||||
"hashed-jwt-..." prefix for JWT-shaped values).
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from litellm.identity.principal import ApiKeyPrincipal
|
||||
|
||||
|
||||
def hash_principal_token(api_key: str) -> str:
|
||||
"""Hash an API key the same way the legacy auth path does.
|
||||
|
||||
Centralized so future principal types (e.g. SSO-issued ephemeral
|
||||
keys) can reuse the same hashing without re-importing the Pydantic
|
||||
model.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
return UserAPIKeyAuth._safe_hash_litellm_api_key(api_key)
|
||||
|
||||
|
||||
def extract_api_key_principal(api_key: Optional[str]) -> Optional[ApiKeyPrincipal]:
|
||||
"""Build an ``ApiKeyPrincipal`` from a raw API key string.
|
||||
|
||||
Returns ``None`` when no key is supplied. Callers that need to pull
|
||||
the key out of a FastAPI request should call ``get_api_key`` from
|
||||
``litellm.proxy.auth.user_api_key_auth`` first; this extractor stays
|
||||
framework-agnostic on purpose so it can be reused by non-FastAPI
|
||||
entrypoints (CLI, MCP, background jobs).
|
||||
"""
|
||||
if not api_key:
|
||||
return None
|
||||
return ApiKeyPrincipal(token_hash=hash_principal_token(api_key))
|
||||
71
litellm/identity/extractors/client.py
Normal file
71
litellm/identity/extractors/client.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
"""Client/network identity extraction.
|
||||
|
||||
Builds a ``ClientInfo`` from a FastAPI request. ``X-Forwarded-For`` is
|
||||
only honored when the direct peer is in a configured trusted-proxy CIDR.
|
||||
The trust logic is delegated to ``IPAddressUtils.is_request_from_trusted_proxy``
|
||||
so we stay in sync with the rest of the proxy.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from litellm.identity.context import ClientInfo
|
||||
|
||||
|
||||
def _split_forwarded_chain(raw: Optional[str]) -> List[str]:
|
||||
if not raw or not isinstance(raw, str):
|
||||
return []
|
||||
return [hop.strip() for hop in raw.split(",") if hop.strip()]
|
||||
|
||||
|
||||
def _direct_client_host(request: Any) -> Optional[str]:
|
||||
client = getattr(request, "client", None)
|
||||
host = getattr(client, "host", None)
|
||||
if isinstance(host, str) and host:
|
||||
return host
|
||||
return None
|
||||
|
||||
|
||||
def extract_client_info(
|
||||
request: Any,
|
||||
general_settings: Optional[Dict[str, Any]] = None,
|
||||
) -> ClientInfo:
|
||||
headers = getattr(request, "headers", {}) or {}
|
||||
forwarded_chain: List[str] = []
|
||||
ip: Optional[str] = None
|
||||
|
||||
# Headers in FastAPI are case-insensitive; normalize defensively for dicts.
|
||||
def _get_header(name: str) -> Optional[str]:
|
||||
try:
|
||||
value = headers.get(name)
|
||||
except AttributeError:
|
||||
return None
|
||||
if value is not None:
|
||||
return value
|
||||
try:
|
||||
for key, val in headers.items():
|
||||
if isinstance(key, str) and key.lower() == name:
|
||||
return val
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
xff = _get_header("x-forwarded-for")
|
||||
if xff:
|
||||
forwarded_chain = _split_forwarded_chain(xff)
|
||||
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
||||
if forwarded_chain and IPAddressUtils.is_request_from_trusted_proxy(
|
||||
request=request, general_settings=general_settings
|
||||
):
|
||||
ip = forwarded_chain[0]
|
||||
else:
|
||||
ip = _direct_client_host(request)
|
||||
|
||||
user_agent = _get_header("user-agent")
|
||||
|
||||
return ClientInfo(
|
||||
ip=ip,
|
||||
user_agent=user_agent if isinstance(user_agent, str) else None,
|
||||
forwarded_chain=forwarded_chain,
|
||||
)
|
||||
24
litellm/identity/extractors/end_user.py
Normal file
24
litellm/identity/extractors/end_user.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
"""End-user extraction.
|
||||
|
||||
Thin wrapper over the existing six-check chain in
|
||||
``litellm.proxy.auth.auth_utils.get_end_user_id_from_request_body``.
|
||||
Validation against the DB stays in ``resolve_and_validate_end_user_id``
|
||||
and runs from the legacy auth path; this extractor returns the raw
|
||||
identifier only.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def extract_end_user_id(
|
||||
body: Optional[dict],
|
||||
headers: Optional[dict] = None,
|
||||
) -> Optional[str]:
|
||||
if body is None:
|
||||
body = {}
|
||||
|
||||
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
||||
|
||||
return get_end_user_id_from_request_body(
|
||||
request_body=body, request_headers=headers
|
||||
)
|
||||
22
litellm/identity/extractors/header.py
Normal file
22
litellm/identity/extractors/header.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
"""Header-driven identity extractors.
|
||||
|
||||
These pull non-credential identity fields out of request headers. They
|
||||
do not perform authorization decisions; that stays in the auth chain.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
AUDIT_CHANGED_BY_HEADER = "litellm-changed-by"
|
||||
|
||||
|
||||
def extract_audit_changed_by(headers: Optional[dict]) -> Optional[str]:
|
||||
"""Read the ``litellm-changed-by`` header used for management-API audit."""
|
||||
if not headers:
|
||||
return None
|
||||
|
||||
for key, value in headers.items():
|
||||
if isinstance(key, str) and key.lower() == AUDIT_CHANGED_BY_HEADER:
|
||||
if isinstance(value, str) and value:
|
||||
return value
|
||||
return None
|
||||
46
litellm/identity/extractors/jwt.py
Normal file
46
litellm/identity/extractors/jwt.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
"""JWT principal extraction.
|
||||
|
||||
Uses ``JWTHandler.is_jwt`` for shape detection and
|
||||
``JWTHandler.get_unverified_claims`` for claim peek. Signature
|
||||
verification and DB-backed claim mapping live in ``JWTHandler.auth_builder``;
|
||||
that path runs from the resolver when DB access is available.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from litellm.identity.principal import JWTPrincipal
|
||||
|
||||
|
||||
def extract_jwt_principal(token: Optional[str]) -> Optional[JWTPrincipal]:
|
||||
"""Decode JWT claims without verification and build a ``JWTPrincipal``.
|
||||
|
||||
Returns ``None`` when the token is missing or not JWT-shaped. The
|
||||
caller is responsible for invoking ``auth_builder`` (or another
|
||||
verifier) before trusting the principal for authorization.
|
||||
"""
|
||||
if not token:
|
||||
return None
|
||||
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
|
||||
if not JWTHandler.is_jwt(token=token):
|
||||
return None
|
||||
|
||||
claims = JWTHandler.get_unverified_claims(token=token) or {}
|
||||
|
||||
aud = claims.get("aud")
|
||||
scope_claim = claims.get("scope") or claims.get("scp") or ""
|
||||
if isinstance(scope_claim, list):
|
||||
scopes = [str(s) for s in scope_claim if s]
|
||||
elif isinstance(scope_claim, str):
|
||||
scopes = [s for s in scope_claim.split(" ") if s]
|
||||
else:
|
||||
scopes = []
|
||||
|
||||
return JWTPrincipal(
|
||||
sub=claims.get("sub"),
|
||||
iss=claims.get("iss"),
|
||||
aud=aud,
|
||||
scopes=scopes,
|
||||
claims=claims,
|
||||
)
|
||||
66
litellm/identity/principal.py
Normal file
66
litellm/identity/principal.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
"""Caller-identity primitives.
|
||||
|
||||
A ``Principal`` answers "who is making this request" using only the fields
|
||||
that uniquely identify the caller. Per-row enrichment (budgets, team rows,
|
||||
object permissions) is intentionally not modeled here; that data continues
|
||||
to ride on ``UserAPIKeyAuth`` in Phase 1.
|
||||
|
||||
Each subtype is a frozen dataclass with a ``kind`` discriminator suitable
|
||||
for ``match``-style dispatch.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Literal, Optional, Union
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ApiKeyPrincipal:
|
||||
kind: Literal["api_key"] = field(default="api_key", init=False)
|
||||
token_hash: str
|
||||
key_alias: Optional[str] = None
|
||||
user_id: Optional[str] = None
|
||||
team_id: Optional[str] = None
|
||||
org_id: Optional[str] = None
|
||||
project_id: Optional[str] = None
|
||||
agent_id: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class JWTPrincipal:
|
||||
kind: Literal["jwt"] = field(default="jwt", init=False)
|
||||
sub: Optional[str] = None
|
||||
iss: Optional[str] = None
|
||||
aud: Optional[Union[str, List[str]]] = None
|
||||
scopes: List[str] = field(default_factory=list)
|
||||
claims: Dict[str, Any] = field(default_factory=dict)
|
||||
mapped_user_id: Optional[str] = None
|
||||
mapped_team_id: Optional[str] = None
|
||||
mapped_org_id: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SSOPrincipal:
|
||||
kind: Literal["sso"] = field(default="sso", init=False)
|
||||
sso_user_id: str
|
||||
email: Optional[str] = None
|
||||
provider: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ServiceAccountPrincipal:
|
||||
kind: Literal["service_account"] = field(default="service_account", init=False)
|
||||
name: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AnonymousPrincipal:
|
||||
kind: Literal["anonymous"] = field(default="anonymous", init=False)
|
||||
|
||||
|
||||
Principal = Union[
|
||||
ApiKeyPrincipal,
|
||||
JWTPrincipal,
|
||||
SSOPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
AnonymousPrincipal,
|
||||
]
|
||||
75
litellm/identity/resolver.py
Normal file
75
litellm/identity/resolver.py
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
"""Compose extractors 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.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
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.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
|
||||
from litellm.identity.extractors.jwt import extract_jwt_principal
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
Principal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
|
||||
_SERVICE_ACCOUNT_API_KEYS = frozenset(
|
||||
{
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _resolve_principal(api_key: Optional[str]) -> Principal:
|
||||
if api_key and api_key in _SERVICE_ACCOUNT_API_KEYS:
|
||||
return ServiceAccountPrincipal(name=api_key)
|
||||
|
||||
jwt_principal = extract_jwt_principal(api_key)
|
||||
if jwt_principal is not None:
|
||||
return jwt_principal
|
||||
|
||||
api_key_principal = extract_api_key_principal(api_key)
|
||||
if api_key_principal is not None:
|
||||
return api_key_principal
|
||||
|
||||
return AnonymousPrincipal()
|
||||
|
||||
|
||||
def resolve_identity(
|
||||
*,
|
||||
api_key: Optional[str] = None,
|
||||
request: Any = None,
|
||||
body: Optional[Dict[str, Any]] = None,
|
||||
headers: Optional[Dict[str, Any]] = None,
|
||||
general_settings: Optional[Dict[str, Any]] = None,
|
||||
) -> IdentityContext:
|
||||
principal = _resolve_principal(api_key)
|
||||
end_user_id = extract_end_user_id(body=body, headers=headers)
|
||||
audit = AuditInfo(changed_by=extract_audit_changed_by(headers))
|
||||
client: ClientInfo
|
||||
if request is not None:
|
||||
client = extract_client_info(
|
||||
request=request, general_settings=general_settings
|
||||
)
|
||||
else:
|
||||
client = ClientInfo()
|
||||
|
||||
return IdentityContext(
|
||||
principal=principal,
|
||||
end_user_id=end_user_id,
|
||||
audit=audit,
|
||||
client=client,
|
||||
)
|
||||
|
|
@ -2539,6 +2539,17 @@ class UserAPIKeyAuth(
|
|||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
def to_identity_context(self):
|
||||
from litellm.identity.adapter import user_api_key_auth_to_identity_context
|
||||
|
||||
return user_api_key_auth_to_identity_context(self)
|
||||
|
||||
@classmethod
|
||||
def from_identity_context(cls, ctx) -> "UserAPIKeyAuth":
|
||||
from litellm.identity.adapter import identity_context_to_user_api_key_auth
|
||||
|
||||
return identity_context_to_user_api_key_auth(ctx)
|
||||
|
||||
|
||||
def user_api_key_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
"""Return True if the caller's role grants unscoped read access to all
|
||||
|
|
|
|||
52
tests/test_litellm/identity/extractors/test_api_key.py
Normal file
52
tests/test_litellm/identity/extractors/test_api_key.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.identity.extractors.api_key import (
|
||||
extract_api_key_principal,
|
||||
hash_principal_token,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
def test_returns_none_for_empty_input():
|
||||
assert extract_api_key_principal(None) is None
|
||||
assert extract_api_key_principal("") is None
|
||||
|
||||
|
||||
def test_sk_key_is_hashed_with_legacy_helper():
|
||||
raw = "sk-abc123"
|
||||
principal = extract_api_key_principal(raw)
|
||||
assert principal is not None
|
||||
assert principal.token_hash == UserAPIKeyAuth._safe_hash_litellm_api_key(raw)
|
||||
assert principal.token_hash != raw
|
||||
|
||||
|
||||
def test_bearer_prefix_normalized_then_hashed():
|
||||
raw = "Bearer sk-abc123"
|
||||
principal = extract_api_key_principal(raw)
|
||||
assert principal is not None
|
||||
assert principal.token_hash == UserAPIKeyAuth._safe_hash_litellm_api_key(raw)
|
||||
sk_only_hash = UserAPIKeyAuth._safe_hash_litellm_api_key("sk-abc123")
|
||||
assert principal.token_hash == sk_only_hash
|
||||
|
||||
|
||||
def test_jwt_shaped_token_gets_hashed_jwt_prefix():
|
||||
fake_jwt = "aaaa.bbbb.cccc"
|
||||
principal = extract_api_key_principal(fake_jwt)
|
||||
assert principal is not None
|
||||
assert principal.token_hash.startswith("hashed-jwt-")
|
||||
|
||||
|
||||
def test_non_sk_non_jwt_returned_unhashed():
|
||||
raw = "custom-key-without-prefix"
|
||||
principal = extract_api_key_principal(raw)
|
||||
assert principal is not None
|
||||
assert principal.token_hash == raw
|
||||
|
||||
|
||||
def test_hash_helper_delegates_to_legacy_path():
|
||||
assert hash_principal_token("sk-x") == UserAPIKeyAuth._safe_hash_litellm_api_key(
|
||||
"sk-x"
|
||||
)
|
||||
47
tests/test_litellm/identity/extractors/test_client.py
Normal file
47
tests/test_litellm/identity/extractors/test_client.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
import os
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.identity.extractors.client import extract_client_info
|
||||
|
||||
|
||||
def _fake_request(headers, client_host=None):
|
||||
return SimpleNamespace(
|
||||
headers=headers,
|
||||
client=SimpleNamespace(host=client_host) if client_host else None,
|
||||
)
|
||||
|
||||
|
||||
def test_falls_back_to_direct_peer_when_not_trusted():
|
||||
req = _fake_request({"x-forwarded-for": "9.9.9.9"}, client_host="10.0.0.1")
|
||||
info = extract_client_info(req, general_settings={})
|
||||
assert info.ip == "10.0.0.1"
|
||||
assert info.forwarded_chain == ["9.9.9.9"]
|
||||
|
||||
|
||||
def test_uses_xff_first_hop_when_proxy_trusted():
|
||||
req = _fake_request(
|
||||
{"x-forwarded-for": "1.2.3.4, 10.0.0.1"}, client_host="10.0.0.1"
|
||||
)
|
||||
settings = {
|
||||
"use_x_forwarded_for": True,
|
||||
"mcp_trusted_proxy_ranges": ["10.0.0.0/8"],
|
||||
}
|
||||
info = extract_client_info(req, general_settings=settings)
|
||||
assert info.ip == "1.2.3.4"
|
||||
assert info.forwarded_chain == ["1.2.3.4", "10.0.0.1"]
|
||||
|
||||
|
||||
def test_no_xff_returns_direct_peer():
|
||||
req = _fake_request({}, client_host="127.0.0.1")
|
||||
info = extract_client_info(req, general_settings={})
|
||||
assert info.ip == "127.0.0.1"
|
||||
assert info.forwarded_chain == []
|
||||
|
||||
|
||||
def test_user_agent_passthrough():
|
||||
req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1")
|
||||
info = extract_client_info(req, general_settings={})
|
||||
assert info.user_agent == "curl/8"
|
||||
45
tests/test_litellm/identity/extractors/test_end_user.py
Normal file
45
tests/test_litellm/identity/extractors/test_end_user.py
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.identity.extractors.end_user import extract_end_user_id
|
||||
|
||||
|
||||
def test_user_body_field():
|
||||
assert extract_end_user_id({"user": "eu-1"}, {}) == "eu-1"
|
||||
|
||||
|
||||
def test_litellm_metadata_user():
|
||||
assert (
|
||||
extract_end_user_id({"litellm_metadata": {"user": "eu-meta"}}, {}) == "eu-meta"
|
||||
)
|
||||
|
||||
|
||||
def test_metadata_user_id():
|
||||
assert extract_end_user_id({"metadata": {"user_id": "eu-md"}}, {}) == "eu-md"
|
||||
|
||||
|
||||
def test_safety_identifier_fallback():
|
||||
assert extract_end_user_id({"safety_identifier": "eu-safety"}, {}) == "eu-safety"
|
||||
|
||||
|
||||
def test_returns_none_when_empty():
|
||||
assert extract_end_user_id(None, None) is None
|
||||
assert extract_end_user_id({}, {}) is None
|
||||
|
||||
|
||||
def test_user_field_wins_over_metadata():
|
||||
body = {"user": "eu-primary", "metadata": {"user_id": "eu-secondary"}}
|
||||
assert extract_end_user_id(body, {}) == "eu-primary"
|
||||
|
||||
|
||||
def test_anthropic_standard_customer_id_header(monkeypatch):
|
||||
from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS
|
||||
|
||||
if not STANDARD_CUSTOMER_ID_HEADERS:
|
||||
pytest.skip("no standard customer headers configured")
|
||||
header_name = STANDARD_CUSTOMER_ID_HEADERS[0]
|
||||
assert extract_end_user_id({}, {header_name: "eu-hdr"}) == "eu-hdr"
|
||||
23
tests/test_litellm/identity/extractors/test_header.py
Normal file
23
tests/test_litellm/identity/extractors/test_header.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.identity.extractors.header import extract_audit_changed_by
|
||||
|
||||
|
||||
def test_returns_value_when_present():
|
||||
assert extract_audit_changed_by({"litellm-changed-by": "alice"}) == "alice"
|
||||
|
||||
|
||||
def test_case_insensitive():
|
||||
assert extract_audit_changed_by({"Litellm-Changed-By": "bob"}) == "bob"
|
||||
|
||||
|
||||
def test_returns_none_when_missing():
|
||||
assert extract_audit_changed_by({}) is None
|
||||
assert extract_audit_changed_by(None) is None
|
||||
|
||||
|
||||
def test_empty_string_value_is_none():
|
||||
assert extract_audit_changed_by({"litellm-changed-by": ""}) is None
|
||||
65
tests/test_litellm/identity/extractors/test_jwt.py
Normal file
65
tests/test_litellm/identity/extractors/test_jwt.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.identity.extractors.jwt import extract_jwt_principal
|
||||
|
||||
|
||||
def _b64url(payload: dict) -> str:
|
||||
raw = json.dumps(payload).encode()
|
||||
return base64.urlsafe_b64encode(raw).rstrip(b"=").decode()
|
||||
|
||||
|
||||
def _build_unverified_jwt(claims: dict) -> str:
|
||||
header = _b64url({"alg": "HS256", "typ": "JWT"})
|
||||
body = _b64url(claims)
|
||||
return f"{header}.{body}.sig"
|
||||
|
||||
|
||||
def test_returns_none_for_non_jwt():
|
||||
assert extract_jwt_principal("sk-abc") is None
|
||||
assert extract_jwt_principal(None) is None
|
||||
assert extract_jwt_principal("") is None
|
||||
|
||||
|
||||
def test_extracts_sub_iss_aud():
|
||||
token = _build_unverified_jwt(
|
||||
{"sub": "user-1", "iss": "https://idp.example", "aud": "litellm"}
|
||||
)
|
||||
p = extract_jwt_principal(token)
|
||||
assert p is not None
|
||||
assert p.sub == "user-1"
|
||||
assert p.iss == "https://idp.example"
|
||||
assert p.aud == "litellm"
|
||||
|
||||
|
||||
def test_string_scope_is_split():
|
||||
token = _build_unverified_jwt({"sub": "u", "scope": "read write admin"})
|
||||
p = extract_jwt_principal(token)
|
||||
assert p is not None
|
||||
assert p.scopes == ["read", "write", "admin"]
|
||||
|
||||
|
||||
def test_list_scope_is_preserved():
|
||||
token = _build_unverified_jwt({"sub": "u", "scope": ["a", "b"]})
|
||||
p = extract_jwt_principal(token)
|
||||
assert p is not None
|
||||
assert p.scopes == ["a", "b"]
|
||||
|
||||
|
||||
def test_scp_claim_supported():
|
||||
token = _build_unverified_jwt({"sub": "u", "scp": "read"})
|
||||
p = extract_jwt_principal(token)
|
||||
assert p is not None
|
||||
assert p.scopes == ["read"]
|
||||
|
||||
|
||||
def test_raw_claims_preserved():
|
||||
claims = {"sub": "u", "custom": {"groups": ["g1"]}}
|
||||
token = _build_unverified_jwt(claims)
|
||||
p = extract_jwt_principal(token)
|
||||
assert p is not None
|
||||
assert p.claims["custom"] == {"groups": ["g1"]}
|
||||
148
tests/test_litellm/identity/test_adapter.py
Normal file
148
tests/test_litellm/identity/test_adapter.py
Normal file
|
|
@ -0,0 +1,148 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
)
|
||||
from litellm.identity import (
|
||||
AnonymousPrincipal,
|
||||
ApiKeyPrincipal,
|
||||
IdentityContext,
|
||||
JWTPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
from litellm.identity.adapter import (
|
||||
identity_context_to_user_api_key_auth,
|
||||
user_api_key_auth_to_identity_context,
|
||||
)
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
|
||||
def test_roundtrip_preserves_api_key_identity_fields():
|
||||
uak = UserAPIKeyAuth(
|
||||
api_key="sk-abc",
|
||||
user_id="u1",
|
||||
team_id="t1",
|
||||
org_id="o1",
|
||||
project_id="p1",
|
||||
agent_id="a1",
|
||||
key_alias="alias-1",
|
||||
end_user_id="eu-1",
|
||||
access_group_ids=["g1", "g2"],
|
||||
)
|
||||
ctx = user_api_key_auth_to_identity_context(uak)
|
||||
assert isinstance(ctx.principal, ApiKeyPrincipal)
|
||||
assert ctx.principal.user_id == "u1"
|
||||
assert ctx.principal.team_id == "t1"
|
||||
assert ctx.principal.org_id == "o1"
|
||||
assert ctx.principal.project_id == "p1"
|
||||
assert ctx.principal.agent_id == "a1"
|
||||
assert ctx.principal.key_alias == "alias-1"
|
||||
assert ctx.end_user_id == "eu-1"
|
||||
assert ctx.access_group_ids == ["g1", "g2"]
|
||||
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.user_id == "u1"
|
||||
assert back.team_id == "t1"
|
||||
assert back.org_id == "o1"
|
||||
assert back.project_id == "p1"
|
||||
assert back.agent_id == "a1"
|
||||
assert back.key_alias == "alias-1"
|
||||
assert back.end_user_id == "eu-1"
|
||||
assert back.access_group_ids == ["g1", "g2"]
|
||||
assert back.token == uak.token
|
||||
|
||||
|
||||
def test_token_hash_not_double_hashed():
|
||||
ctx = IdentityContext(principal=ApiKeyPrincipal(token_hash="abc123"))
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.token == "abc123"
|
||||
|
||||
|
||||
def test_jwt_principal_roundtrip():
|
||||
uak = UserAPIKeyAuth(
|
||||
api_key="aaaa.bbbb.cccc",
|
||||
user_id="jwt-user",
|
||||
team_id="jwt-team",
|
||||
org_id="jwt-org",
|
||||
jwt_claims={"sub": "jwt-user", "iss": "idp", "scope": "read write"},
|
||||
)
|
||||
ctx = user_api_key_auth_to_identity_context(uak)
|
||||
assert isinstance(ctx.principal, JWTPrincipal)
|
||||
assert ctx.principal.sub == "jwt-user"
|
||||
assert ctx.principal.iss == "idp"
|
||||
assert ctx.principal.scopes == ["read", "write"]
|
||||
assert ctx.principal.mapped_user_id == "jwt-user"
|
||||
assert ctx.principal.mapped_team_id == "jwt-team"
|
||||
assert ctx.principal.mapped_org_id == "jwt-org"
|
||||
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.user_id == "jwt-user"
|
||||
assert back.team_id == "jwt-team"
|
||||
assert back.org_id == "jwt-org"
|
||||
assert back.jwt_claims is not None
|
||||
assert back.jwt_claims.get("sub") == "jwt-user"
|
||||
|
||||
|
||||
def test_service_account_jobs_principal_roundtrip():
|
||||
uak = UserAPIKeyAuth.get_litellm_internal_jobs_user_api_key_auth()
|
||||
ctx = user_api_key_auth_to_identity_context(uak)
|
||||
assert isinstance(ctx.principal, ServiceAccountPrincipal)
|
||||
assert ctx.principal.name == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
|
||||
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.user_id == "system"
|
||||
assert back.team_id == "system"
|
||||
assert back.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
assert back.key_alias == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
|
||||
|
||||
|
||||
def test_service_account_health_check_roundtrip():
|
||||
uak = UserAPIKeyAuth.get_litellm_internal_health_check_user_api_key_auth()
|
||||
ctx = user_api_key_auth_to_identity_context(uak)
|
||||
assert isinstance(ctx.principal, ServiceAccountPrincipal)
|
||||
assert ctx.principal.name == LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
|
||||
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.team_id == LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
|
||||
assert back.team_alias == LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
|
||||
|
||||
|
||||
def test_service_account_cli_roundtrip():
|
||||
uak = UserAPIKeyAuth.get_litellm_cli_user_api_key_auth()
|
||||
ctx = user_api_key_auth_to_identity_context(uak)
|
||||
assert isinstance(ctx.principal, ServiceAccountPrincipal)
|
||||
assert ctx.principal.name == LITTELM_CLI_SERVICE_ACCOUNT_NAME
|
||||
|
||||
|
||||
def test_anonymous_principal_when_no_token():
|
||||
uak = UserAPIKeyAuth()
|
||||
ctx = user_api_key_auth_to_identity_context(uak)
|
||||
assert isinstance(ctx.principal, AnonymousPrincipal)
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.token is None
|
||||
assert back.user_id is None
|
||||
|
||||
|
||||
def test_end_user_does_not_leak_into_principal():
|
||||
ctx = IdentityContext(
|
||||
principal=ApiKeyPrincipal(token_hash="t", user_id="u"),
|
||||
end_user_id="customer-99",
|
||||
)
|
||||
back = identity_context_to_user_api_key_auth(ctx)
|
||||
assert back.user_id == "u"
|
||||
assert back.end_user_id == "customer-99"
|
||||
assert ctx.principal.user_id == "u"
|
||||
|
||||
|
||||
def test_uak_methods_delegate_to_adapter():
|
||||
uak = UserAPIKeyAuth(api_key="sk-test", user_id="u", team_id="t")
|
||||
ctx = uak.to_identity_context()
|
||||
assert isinstance(ctx, IdentityContext)
|
||||
back = UserAPIKeyAuth.from_identity_context(ctx)
|
||||
assert back.user_id == "u"
|
||||
assert back.team_id == "t"
|
||||
46
tests/test_litellm/identity/test_context.py
Normal file
46
tests/test_litellm/identity/test_context.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from litellm.identity.context import (
|
||||
AuditInfo,
|
||||
ClientInfo,
|
||||
IdentityContext,
|
||||
RequestIds,
|
||||
)
|
||||
from litellm.identity.principal import AnonymousPrincipal, ApiKeyPrincipal
|
||||
|
||||
|
||||
def test_default_principal_is_anonymous():
|
||||
ctx = IdentityContext()
|
||||
assert isinstance(ctx.principal, AnonymousPrincipal)
|
||||
assert ctx.end_user_id is None
|
||||
assert ctx.tags == []
|
||||
assert ctx.access_group_ids == []
|
||||
assert ctx.request == RequestIds()
|
||||
assert ctx.client == ClientInfo()
|
||||
assert ctx.audit == AuditInfo()
|
||||
|
||||
|
||||
def test_context_is_mutable():
|
||||
ctx = IdentityContext()
|
||||
ctx.end_user_id = "eu-1"
|
||||
ctx.tags.append("env:prod")
|
||||
assert ctx.end_user_id == "eu-1"
|
||||
assert "env:prod" in ctx.tags
|
||||
|
||||
|
||||
def test_context_carries_principal():
|
||||
p = ApiKeyPrincipal(token_hash="abc", user_id="u1", team_id="t1")
|
||||
ctx = IdentityContext(principal=p)
|
||||
assert ctx.principal is p
|
||||
|
||||
|
||||
def test_subobjects_are_independent_between_instances():
|
||||
a = IdentityContext()
|
||||
b = IdentityContext()
|
||||
a.tags.append("x")
|
||||
assert b.tags == []
|
||||
a.client.forwarded_chain.append("1.2.3.4")
|
||||
assert b.client.forwarded_chain == []
|
||||
39
tests/test_litellm/identity/test_principal.py
Normal file
39
tests/test_litellm/identity/test_principal.py
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
ApiKeyPrincipal,
|
||||
JWTPrincipal,
|
||||
SSOPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
|
||||
|
||||
def test_principal_kind_discriminators_are_fixed():
|
||||
assert ApiKeyPrincipal(token_hash="x").kind == "api_key"
|
||||
assert JWTPrincipal().kind == "jwt"
|
||||
assert SSOPrincipal(sso_user_id="s").kind == "sso"
|
||||
assert ServiceAccountPrincipal(name="n").kind == "service_account"
|
||||
assert AnonymousPrincipal().kind == "anonymous"
|
||||
|
||||
|
||||
def test_principals_are_frozen():
|
||||
p = ApiKeyPrincipal(token_hash="x", user_id="u1")
|
||||
with pytest.raises(Exception):
|
||||
p.user_id = "u2" # type: ignore[misc]
|
||||
|
||||
|
||||
def test_principals_are_hashable():
|
||||
a = ApiKeyPrincipal(token_hash="x", user_id="u1")
|
||||
b = ApiKeyPrincipal(token_hash="x", user_id="u1")
|
||||
assert {a, b} == {a}
|
||||
|
||||
|
||||
def test_kind_is_not_constructor_arg():
|
||||
with pytest.raises(TypeError):
|
||||
ApiKeyPrincipal(kind="api_key", token_hash="x") # type: ignore[call-arg]
|
||||
70
tests/test_litellm/identity/test_resolver.py
Normal file
70
tests/test_litellm/identity/test_resolver.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from litellm.constants import LITTELM_CLI_SERVICE_ACCOUNT_NAME
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
ApiKeyPrincipal,
|
||||
JWTPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
from litellm.identity.resolver import resolve_identity
|
||||
|
||||
|
||||
def _fake_request(headers=None, client_host=None):
|
||||
return SimpleNamespace(
|
||||
headers=headers or {},
|
||||
client=SimpleNamespace(host=client_host) if client_host else None,
|
||||
)
|
||||
|
||||
|
||||
def _jwt(claims):
|
||||
def b(d):
|
||||
return (
|
||||
base64.urlsafe_b64encode(json.dumps(d).encode()).rstrip(b"=").decode()
|
||||
)
|
||||
|
||||
return f"{b({'alg':'HS256','typ':'JWT'})}.{b(claims)}.sig"
|
||||
|
||||
|
||||
def test_anonymous_when_no_credentials():
|
||||
ctx = resolve_identity()
|
||||
assert isinstance(ctx.principal, AnonymousPrincipal)
|
||||
|
||||
|
||||
def test_api_key_principal_for_sk_key():
|
||||
ctx = 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"}))
|
||||
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)
|
||||
assert isinstance(ctx.principal, ServiceAccountPrincipal)
|
||||
assert ctx.principal.name == LITTELM_CLI_SERVICE_ACCOUNT_NAME
|
||||
|
||||
|
||||
def test_end_user_and_audit_propagate():
|
||||
ctx = resolve_identity(
|
||||
body={"user": "eu-42"},
|
||||
headers={"litellm-changed-by": "admin"},
|
||||
)
|
||||
assert ctx.end_user_id == "eu-42"
|
||||
assert ctx.audit.changed_by == "admin"
|
||||
|
||||
|
||||
def test_client_info_from_request():
|
||||
req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1")
|
||||
ctx = resolve_identity(request=req)
|
||||
assert ctx.client.ip == "127.0.0.1"
|
||||
assert ctx.client.user_agent == "curl/8"
|
||||
Loading…
Add table
Reference in a new issue