diff --git a/litellm/proxy/auth_v2/__init__.py b/litellm/proxy/auth_v2/__init__.py index 0e76c806a6b..032bb014645 100644 --- a/litellm/proxy/auth_v2/__init__.py +++ b/litellm/proxy/auth_v2/__init__.py @@ -1,17 +1,35 @@ -from .config import AuthConfig -from .models import Principal -from .security import ( - get_current_principal, - install_auth, - require_permission, - require_roles, +from .config import ( + ApiKeySchemeConfig, + AuthConfig, + HttpBasicConfig, + MutualTLSConfig, + OAuth2IntrospectionConfig, + OIDCProviderConfig, + SAMLConfig, + SessionConfig, + TrustedProxyConfig, ) +from .models import Principal +from .rbac import Role +from .resolver import IdentityResolver, InMemoryIdentityStore, ProvisioningStore +from .saml import build_saml_router +from .security import AuthSecurity __all__ = [ - "Principal", + "AuthSecurity", "AuthConfig", - "get_current_principal", - "require_roles", - "require_permission", - "install_auth", + "Principal", + "Role", + "IdentityResolver", + "ProvisioningStore", + "InMemoryIdentityStore", + "ApiKeySchemeConfig", + "HttpBasicConfig", + "OIDCProviderConfig", + "OAuth2IntrospectionConfig", + "MutualTLSConfig", + "TrustedProxyConfig", + "SessionConfig", + "SAMLConfig", + "build_saml_router", ] diff --git a/litellm/proxy/auth_v2/authenticators.py b/litellm/proxy/auth_v2/authenticators.py index 35425cbd067..7f4ba45b5f4 100644 --- a/litellm/proxy/auth_v2/authenticators.py +++ b/litellm/proxy/auth_v2/authenticators.py @@ -6,9 +6,8 @@ import functools import hashlib import hmac import secrets -from typing import Any, Dict, List, Optional, Protocol, runtime_checkable +from typing import Any, Callable, Dict, List, Optional, Protocol, runtime_checkable -import httpx import jwt from fastapi import Request from jwt import PyJWKClient @@ -20,9 +19,9 @@ from .config import ( ApiKeySchemeConfig, AuthConfig, HttpBasicConfig, - MutualTlsConfig, + MutualTLSConfig, OAuth2IntrospectionConfig, - OidcProviderConfig, + OIDCProviderConfig, TrustedProxyConfig, ) from .models import ( @@ -39,13 +38,40 @@ AT_JWT_TYPES = {"at+jwt", "application/at+jwt"} @runtime_checkable class Authenticator(Protocol): - scheme: SecuritySchemeType - async def authenticate(self, request: Request) -> Optional[Credential]: ... def challenge(self) -> str: ... +@runtime_checkable +class BasicAuthVerifier(Protocol): + def verify(self, username: str, password: str) -> bool: ... + + +def hash_basic_password(password: str, salt: Optional[str] = None) -> str: + salt = salt or secrets.token_hex(16) + digest = hashlib.sha256(bytes.fromhex(salt) + password.encode()).hexdigest() + return f"{salt}${digest}" + + +class InMemoryBasicAuthStore: + def __init__(self, credentials: Dict[str, str]) -> None: + self._credentials = credentials + + def verify(self, username: str, password: str) -> bool: + stored = self._credentials.get(username) + if stored is None: + return False + salt, _, expected = stored.partition("$") + try: + candidate = hashlib.sha256( + bytes.fromhex(salt) + password.encode() + ).hexdigest() + except ValueError: + return False + return hmac.compare_digest(candidate, expected) + + def _extract_bearer(request: Request) -> Optional[str]: header = request.headers.get("authorization") if not header: @@ -93,39 +119,10 @@ def _credential_from_claims( ) -@runtime_checkable -class BasicAuthVerifier(Protocol): - def verify(self, username: str, password: str) -> bool: ... - - -def hash_basic_password(password: str, salt: Optional[str] = None) -> str: - salt = salt or secrets.token_hex(16) - digest = hashlib.sha256(bytes.fromhex(salt) + password.encode()).hexdigest() - return f"{salt}${digest}" - - -class InMemoryBasicAuthStore: - def __init__(self, credentials: Dict[str, str]) -> None: - self._credentials = credentials - - def verify(self, username: str, password: str) -> bool: - stored = self._credentials.get(username) - if stored is None: - return False - salt, _, expected = stored.partition("$") - try: - candidate = hashlib.sha256( - bytes.fromhex(salt) + password.encode() - ).hexdigest() - except ValueError: - return False - return hmac.compare_digest(candidate, expected) - - -class JwtVerifier: +class JWTVerifier: def __init__( self, - provider: OidcProviderConfig, + provider: OIDCProviderConfig, jwks_client: Optional[PyJWKClient] = None, ) -> None: self.provider = provider @@ -138,6 +135,8 @@ class JwtVerifier: self._jwks_client = PyJWKClient(jwks_uri, cache_keys=True) def _discover_jwks(self) -> str: + import httpx + url = f"{self.provider.issuer.rstrip('/')}/.well-known/openid-configuration" response = httpx.get(url, timeout=10.0) response.raise_for_status() @@ -171,14 +170,14 @@ class JwtVerifier: async def _verify_jwt_off_loop( - verifier: JwtVerifier, token: str, *, require_at_jwt: Optional[bool] = None + verifier: JWTVerifier, token: str, *, require_at_jwt: Optional[bool] = None ) -> Dict[str, Any]: return await run_in_threadpool( functools.partial(verifier.verify, token, require_at_jwt=require_at_jwt) ) -def _select_verifier(token: str, verifiers: List[JwtVerifier]) -> Optional[JwtVerifier]: +def _select_verifier(token: str, verifiers: List[JWTVerifier]) -> Optional[JWTVerifier]: if not verifiers: return None try: @@ -191,9 +190,22 @@ def _select_verifier(token: str, verifiers: List[JwtVerifier]) -> Optional[JwtVe return None -class ApiKeyAuthenticator: - scheme = SecuritySchemeType.API_KEY +async def _authenticate_bearer_jwt( + token: str, + verifiers: List[JWTVerifier], + scheme: SecuritySchemeType, + method: AuthMethod, + *, + require_at_jwt: bool = False, +) -> Credential: + verifier = _select_verifier(token, verifiers) + if verifier is None: + raise errors.invalid_token("no issuer match") + claims = await _verify_jwt_off_loop(verifier, token, require_at_jwt=require_at_jwt) + return _credential_from_claims(scheme, method, token, claims) + +class APIKeyAuthenticator: def __init__(self, config: ApiKeySchemeConfig) -> None: self._header_name = config.header_name @@ -202,7 +214,7 @@ class ApiKeyAuthenticator: if not raw: return None return Credential( - scheme=self.scheme, + scheme=SecuritySchemeType.API_KEY, method=AuthMethod.API_KEY, subject=raw, credential_ref=CredentialRef(key_id=raw[:10]), @@ -214,12 +226,10 @@ class ApiKeyAuthenticator: class HttpAuthenticator: - scheme = SecuritySchemeType.HTTP - def __init__( self, basic: HttpBasicConfig, - jwt_verifiers: List[JwtVerifier], + jwt_verifiers: List[JWTVerifier], basic_verifier: Optional[BasicAuthVerifier] = None, ) -> None: self._basic = basic @@ -233,20 +243,13 @@ class HttpAuthenticator: scheme, _, value = header.partition(" ") scheme_lower = scheme.lower() if scheme_lower == "bearer" and value: - return await self._verify_bearer(value) + return await _authenticate_bearer_jwt( + value, self._verifiers, SecuritySchemeType.HTTP, AuthMethod.BEARER_JWT + ) if scheme_lower == "basic" and self._basic.enabled and value: return self._verify_basic(value) return None - async def _verify_bearer(self, token: str) -> Credential: - verifier = _select_verifier(token, self._verifiers) - if verifier is None: - raise errors.invalid_token("no issuer match") - claims = await _verify_jwt_off_loop(verifier, token) - return _credential_from_claims( - self.scheme, AuthMethod.BEARER_JWT, token, claims - ) - def _verify_basic(self, value: str) -> Credential: challenge = errors.basic_challenge(self._basic.realm) try: @@ -262,7 +265,7 @@ class HttpAuthenticator: ): raise errors.unauthenticated(challenge) return Credential( - scheme=self.scheme, + scheme=SecuritySchemeType.HTTP, method=AuthMethod.HTTP_BASIC, subject=username, ) @@ -274,46 +277,47 @@ class HttpAuthenticator: return bearer -class OAuth2Authenticator: - scheme = SecuritySchemeType.OAUTH2 +def _default_introspection_client() -> Any: + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.llms.custom_http import httpxSpecialProvider + return get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + + +class OAuth2Authenticator: def __init__( self, - jwt_verifiers: List[JwtVerifier], + jwt_verifiers: List[JWTVerifier], introspection: Optional[OAuth2IntrospectionConfig], + client_factory: Optional[Callable[[], Any]] = None, ) -> None: self._verifiers = jwt_verifiers self._introspection = introspection + self._client_factory = client_factory or _default_introspection_client async def authenticate(self, request: Request) -> Optional[Credential]: token = _extract_bearer(request) if token is None: return None if _looks_like_jwt(token): - return await self._verify_at_jwt(token) + return await _authenticate_bearer_jwt( + token, + self._verifiers, + SecuritySchemeType.OAUTH2, + AuthMethod.BEARER_JWT, + require_at_jwt=True, + ) if self._introspection is not None: return await self._introspect(token) raise errors.invalid_token() - async def _verify_at_jwt(self, token: str) -> Credential: - verifier = _select_verifier(token, self._verifiers) - if verifier is None: - raise errors.invalid_token("no issuer match") - claims = await _verify_jwt_off_loop(verifier, token, require_at_jwt=True) - return _credential_from_claims( - self.scheme, AuthMethod.BEARER_JWT, token, claims - ) - async def _introspect(self, token: str) -> Credential: config = self._introspection assert config is not None - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client - from litellm.types.llms.custom_http import httpxSpecialProvider - basic = base64.b64encode( f"{config.client_id}:{config.client_secret.get_secret_value()}".encode() ).decode() - client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + client = self._client_factory() response = await client.post( str(config.introspection_endpoint), data={"token": token}, @@ -329,7 +333,7 @@ class OAuth2Authenticator: if config.audience and not set(token_audience) & set(config.audience): raise errors.invalid_token("audience mismatch") return Credential( - scheme=self.scheme, + scheme=SecuritySchemeType.OAUTH2, method=AuthMethod.OAUTH2_INTROSPECTION, subject=str(body.get(config.subject_field, "")), issuer=body.get("iss"), @@ -342,30 +346,24 @@ class OAuth2Authenticator: return errors.bearer_challenge() -class OidcAuthenticator: - scheme = SecuritySchemeType.OPENID_CONNECT - - def __init__(self, jwt_verifiers: List[JwtVerifier]) -> None: +class OIDCAuthenticator: + def __init__(self, jwt_verifiers: List[JWTVerifier]) -> None: self._verifiers = jwt_verifiers async def authenticate(self, request: Request) -> Optional[Credential]: token = _extract_bearer(request) if token is None: return None - verifier = _select_verifier(token, self._verifiers) - if verifier is None: - raise errors.invalid_token("no issuer match") - claims = await _verify_jwt_off_loop(verifier, token) - return _credential_from_claims(self.scheme, AuthMethod.OIDC, token, claims) + return await _authenticate_bearer_jwt( + token, self._verifiers, SecuritySchemeType.OPENID_CONNECT, AuthMethod.OIDC + ) def challenge(self) -> str: return errors.bearer_challenge() -class MutualTlsAuthenticator: - scheme = SecuritySchemeType.MUTUAL_TLS - - def __init__(self, config: MutualTlsConfig, network: TrustedProxyConfig) -> None: +class MutualTLSAuthenticator: + def __init__(self, config: MutualTLSConfig, network: TrustedProxyConfig) -> None: self._config = config self._network = network @@ -374,7 +372,7 @@ class MutualTlsAuthenticator: if cert is None: return None return Credential( - scheme=self.scheme, + scheme=SecuritySchemeType.MUTUAL_TLS, method=AuthMethod.MUTUAL_TLS, subject=cert.subject_dn, client_certificate=cert, @@ -398,19 +396,19 @@ class MutualTlsAuthenticator: def build_authenticators( config: AuthConfig, *, basic_verifier: Optional[BasicAuthVerifier] = None ) -> List[Authenticator]: - verifiers = [JwtVerifier(provider) for provider in config.oidc_providers] + verifiers = [JWTVerifier(provider) for provider in config.oidc_providers] by_scheme: Dict[SecuritySchemeType, Authenticator] = {} if config.api_key is not None: - by_scheme[SecuritySchemeType.API_KEY] = ApiKeyAuthenticator(config.api_key) + by_scheme[SecuritySchemeType.API_KEY] = APIKeyAuthenticator(config.api_key) by_scheme[SecuritySchemeType.HTTP] = HttpAuthenticator( config.http_basic, verifiers, basic_verifier ) - by_scheme[SecuritySchemeType.OPENID_CONNECT] = OidcAuthenticator(verifiers) + by_scheme[SecuritySchemeType.OPENID_CONNECT] = OIDCAuthenticator(verifiers) by_scheme[SecuritySchemeType.OAUTH2] = OAuth2Authenticator( verifiers, config.oauth2_introspection ) if config.mutual_tls.enabled: - by_scheme[SecuritySchemeType.MUTUAL_TLS] = MutualTlsAuthenticator( + by_scheme[SecuritySchemeType.MUTUAL_TLS] = MutualTLSAuthenticator( config.mutual_tls, config.network ) return [by_scheme[scheme] for scheme in config.scheme_order if scheme in by_scheme] diff --git a/litellm/proxy/auth_v2/config.py b/litellm/proxy/auth_v2/config.py index 9026a23af2d..41558a45de2 100644 --- a/litellm/proxy/auth_v2/config.py +++ b/litellm/proxy/auth_v2/config.py @@ -29,7 +29,7 @@ class HttpBasicConfig(BaseModel): realm: str = "litellm" -class OidcProviderConfig(BaseModel): +class OIDCProviderConfig(BaseModel): issuer: str audience: List[str] jwks_uri: Optional[AnyHttpUrl] = None @@ -50,7 +50,7 @@ class OAuth2IntrospectionConfig(BaseModel): audience: List[str] = Field(default_factory=list) -class MutualTlsConfig(BaseModel): +class MutualTLSConfig(BaseModel): enabled: bool = False forwarded_subject_header: Optional[str] = None @@ -60,7 +60,17 @@ class TrustedProxyConfig(BaseModel): trusted_proxy_cidrs: List[str] = Field(default_factory=list) -class SamlConfig(BaseModel): +class SessionConfig(BaseModel): + cookie: str = "litellm_session" + secure: bool = True + ttl_seconds: int = 3600 + max_size: int = 10000 + default_redirect_path: str = "/" + login_cookie: str = "litellm_oidc_txn" + login_state_ttl: int = 300 + + +class SAMLConfig(BaseModel): enabled: bool = False entity_id: str acs_url: str @@ -68,18 +78,13 @@ class SamlConfig(BaseModel): sp_key_file: Optional[str] = None sp_cert_file: Optional[str] = None allow_unsolicited: bool = False - session_cookie: str = "saml_session" - cookie_secure: bool = True - session_ttl_seconds: int = 3600 - session_max_size: int = 10000 - default_redirect_path: str = "/" xmlsec_binary: Optional[str] = None attribute_map: Dict[str, str] = Field( default_factory=lambda: dict(DEFAULT_SAML_ATTRIBUTE_MAP) ) @model_validator(mode="after") - def _require_idp_metadata(self) -> "SamlConfig": + def _require_idp_metadata(self) -> "SAMLConfig": if self.enabled and not self.idp_metadata.strip(): raise ValueError( "SAML enabled but idp_metadata is empty (inline XML, local path, or URL)" @@ -88,10 +93,6 @@ class SamlConfig(BaseModel): class AuthConfig(BaseModel): - # First-match-wins precedence. HTTP precedes OPENID_CONNECT, so a bearer JWT - # is claimed by HttpAuthenticator (auth_method=bearer_jwt) and OidcAuthenticator - # never runs; both share the same JwtVerifiers and verify identically, so this - # only changes the auth_method label. Reorder if openIdConnect labeling matters. scheme_order: List[SecuritySchemeType] = Field( default_factory=lambda: [ SecuritySchemeType.API_KEY, @@ -103,9 +104,10 @@ class AuthConfig(BaseModel): ) api_key: Optional[ApiKeySchemeConfig] = Field(default_factory=ApiKeySchemeConfig) http_basic: HttpBasicConfig = Field(default_factory=HttpBasicConfig) - oidc_providers: List[OidcProviderConfig] = Field(default_factory=list) + oidc_providers: List[OIDCProviderConfig] = Field(default_factory=list) oauth2_introspection: Optional[OAuth2IntrospectionConfig] = None - mutual_tls: MutualTlsConfig = Field(default_factory=MutualTlsConfig) + mutual_tls: MutualTLSConfig = Field(default_factory=MutualTLSConfig) network: TrustedProxyConfig = Field(default_factory=TrustedProxyConfig) - saml: Optional[SamlConfig] = None + session: SessionConfig = Field(default_factory=SessionConfig) + saml: Optional[SAMLConfig] = None casbin_policy_path: Optional[str] = None diff --git a/litellm/proxy/auth_v2/rbac.py b/litellm/proxy/auth_v2/rbac.py index 2290fd671dd..927a242e6e8 100644 --- a/litellm/proxy/auth_v2/rbac.py +++ b/litellm/proxy/auth_v2/rbac.py @@ -57,7 +57,7 @@ _DEFAULT_POLICY: List[Tuple[str, str, str]] = [ ] -class RbacEngine: +class RBACEngine: def __init__(self, policy_path: Optional[str] = None) -> None: model = casbin.Model() model.load_model_from_text(_MODEL_TEXT) @@ -75,7 +75,7 @@ class RbacEngine: self._enforcer.enforce(role.value, obj, act) for role in principal.roles ) - def has_role(self, principal: "Principal", allowed: Tuple[Role, ...]) -> bool: + def has_any_role(self, principal: "Principal", allowed: Tuple[Role, ...]) -> bool: allowed_values = {role.value for role in allowed} for role in principal.roles: if role.value in allowed_values: diff --git a/litellm/proxy/auth_v2/saml.py b/litellm/proxy/auth_v2/saml.py index d57811c8f57..5f0a2e8bfb9 100644 --- a/litellm/proxy/auth_v2/saml.py +++ b/litellm/proxy/auth_v2/saml.py @@ -1,8 +1,7 @@ from __future__ import annotations -import secrets import time -from typing import Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast from fastapi import APIRouter, HTTPException, Request from fastapi.responses import RedirectResponse, Response @@ -13,9 +12,12 @@ from saml2.metadata import entity_descriptor from scim2_models import Email, Name from scim2_models import User as ScimUser -from .config import SamlConfig -from .models import AuthMethod, Credential, CredentialRef, SecuritySchemeType +from .config import SAMLConfig from .resolver import ProvisioningStore +from .session import safe_relay_state + +if TYPE_CHECKING: + from .security import AuthSecurity _SINGLE_VALUE_TARGETS = { "email", @@ -85,18 +87,6 @@ def _claims_from_mapped(mapped: Dict[str, Any]) -> Dict[str, Any]: return claims -def _safe_relay_state(target: Optional[str], default: str) -> str: - if ( - target - and target.startswith("/") - and not target.startswith("//") - and "://" not in target - and "\\" not in target - ): - return target - return default - - def _metadata_source(idp_metadata: str) -> Dict[str, Any]: stripped = idp_metadata.strip() if stripped.startswith("<"): @@ -106,7 +96,7 @@ def _metadata_source(idp_metadata: str) -> Dict[str, Any]: return {"local": [idp_metadata]} -def _sp_config_dict(config: SamlConfig) -> Dict[str, Any]: +def _sp_config_dict(config: SAMLConfig) -> Dict[str, Any]: cfg: Dict[str, Any] = { "entityid": config.entity_id, "service": { @@ -132,21 +122,19 @@ def _sp_config_dict(config: SamlConfig) -> Dict[str, Any]: return cfg -def build_sp_client(config: SamlConfig) -> Saml2Client: +def build_sp_client(config: SAMLConfig) -> Saml2Client: conf = SPConfig() conf.load(_sp_config_dict(config)) return Saml2Client(config=conf) -class SamlSessionStore: - def __init__(self, ttl_seconds: int = 3600, max_size: int = 10000) -> None: - self._sessions: Dict[str, Tuple[float, Dict[str, Any]]] = {} - self._seen_assertions: Dict[str, float] = {} +class SAMLProtocolStore: + def __init__(self, replay_ttl_seconds: int) -> None: self.outstanding: Dict[str, str] = {} - self._ttl = ttl_seconds - self._max_size = max_size + self._seen_assertions: Dict[str, float] = {} + self._replay_ttl = replay_ttl_seconds - def remember_request(self, request_id: str, relay_state: str = "/") -> None: + def remember_request(self, request_id: str, relay_state: str) -> None: self.outstanding[request_id] = relay_state def consume_assertion(self, assertion_id: str) -> bool: @@ -156,65 +144,16 @@ class SamlSessionStore: } if assertion_id in self._seen_assertions: return False - self._seen_assertions[assertion_id] = now + self._ttl + self._seen_assertions[assertion_id] = now + self._replay_ttl return True - def create_session(self, identity: Dict[str, Any]) -> str: - now = time.time() - self._evict(now) - session_id = secrets.token_urlsafe(32) - self._sessions[session_id] = (now + self._ttl, identity) - return session_id - def get(self, session_id: str) -> Optional[Dict[str, Any]]: - entry = self._sessions.get(session_id) - if entry is None: - return None - expires_at, identity = entry - if expires_at < time.time(): - self._sessions.pop(session_id, None) - return None - return identity - - def _evict(self, now: float) -> None: - for key in [k for k, (exp, _) in self._sessions.items() if exp < now]: - self._sessions.pop(key, None) - overflow = len(self._sessions) - self._max_size + 1 - if overflow > 0: - oldest = sorted(self._sessions, key=lambda k: self._sessions[k][0]) - for key in oldest[:overflow]: - self._sessions.pop(key, None) - - -class SamlAuthenticator: - scheme = SecuritySchemeType.HTTP - - def __init__(self, config: SamlConfig, session_store: SamlSessionStore) -> None: - self._config = config - self._store = session_store - - async def authenticate(self, request: Request) -> Optional[Credential]: - session_id = request.cookies.get(self._config.session_cookie) - if not session_id: - return None - identity = self._store.get(session_id) - if identity is None: - return None - return Credential( - scheme=self.scheme, - method=AuthMethod.SAML, - subject=identity["name_id"], - issuer=identity.get("issuer"), - claims=identity.get("claims", {}), - credential_ref=CredentialRef(token_id=session_id), - ) - - def challenge(self) -> str: - return "" - - -def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> APIRouter: +def build_saml_router(auth: "AuthSecurity") -> APIRouter: + config = auth.config.saml + assert config is not None + session = auth.config.session client = build_sp_client(config) + protocol = SAMLProtocolStore(session.ttl_seconds) router = APIRouter(prefix="/auth/saml", tags=["saml"]) @router.get("/metadata") @@ -226,11 +165,11 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP @router.get("/login") async def login(request: Request) -> RedirectResponse: - relay_state = _safe_relay_state( - request.query_params.get("next"), config.default_redirect_path + relay_state = safe_relay_state( + request.query_params.get("next"), session.default_redirect_path ) request_id, info = client.prepare_for_authenticate(relay_state=relay_state) - session_store.remember_request(request_id, relay_state) + protocol.remember_request(request_id, relay_state) location = dict(info["headers"]).get("Location") if not location: raise HTTPException(status_code=500, detail="no SAML redirect produced") @@ -246,7 +185,7 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP authn_response = client.parse_authn_request_response( saml_response, BINDING_HTTP_POST, - outstanding=session_store.outstanding or None, + outstanding=protocol.outstanding or None, ) except Exception as exc: raise HTTPException( @@ -257,14 +196,12 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP in_response_to = getattr(authn_response, "in_response_to", None) bound_relay = ( - session_store.outstanding.pop(in_response_to, None) - if in_response_to - else None + protocol.outstanding.pop(in_response_to, None) if in_response_to else None ) assertion = getattr(authn_response, "assertion", None) assertion_id = getattr(assertion, "id", None) - if assertion_id and not session_store.consume_assertion(assertion_id): + if assertion_id and not protocol.consume_assertion(assertion_id): raise HTTPException(status_code=401, detail="SAML assertion replay") name_id = authn_response.get_subject().text @@ -272,24 +209,25 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP mapped = _map_attributes(ava, config.attribute_map) user = _user_from_mapped(name_id, mapped) - store: ProvisioningStore = request.app.state.auth_v2.resolver + store = cast(ProvisioningStore, auth.resolver) await store.upsert_user(user) - session_id = session_store.create_session( + session_id = auth.session_store.create_session( { - "name_id": name_id, + "method": "saml", + "subject": name_id, "issuer": authn_response.issuer(), "claims": _claims_from_mapped(mapped), } ) - target = _safe_relay_state(bound_relay, config.default_redirect_path) + target = safe_relay_state(bound_relay, session.default_redirect_path) response = RedirectResponse(target, status_code=303) response.set_cookie( - config.session_cookie, + session.cookie, session_id, httponly=True, samesite="lax", - secure=config.cookie_secure, + secure=session.secure, ) return response diff --git a/litellm/proxy/auth_v2/security.py b/litellm/proxy/auth_v2/security.py index c5858563df4..f7c8044a5ee 100644 --- a/litellm/proxy/auth_v2/security.py +++ b/litellm/proxy/auth_v2/security.py @@ -1,78 +1,20 @@ -from __future__ import annotations - -from dataclasses import dataclass, field from typing import Annotated, Callable, List, Optional -from fastapi import FastAPI, Request, Security +from fastapi import Request, Security from fastapi.security import SecurityScopes from . import errors -from .authenticators import Authenticator, BasicAuthVerifier, build_authenticators +from .authenticators import ( + Authenticator, + BasicAuthVerifier, + build_authenticators, +) from .config import AuthConfig from .models import Principal from .network import resolve_network_context -from .rbac import RbacEngine, Role, has_required_scopes +from .rbac import RBACEngine, Role, has_required_scopes from .resolver import IdentityResolver - - -@dataclass -class AuthContext: - config: AuthConfig - authenticators: List[Authenticator] - resolver: IdentityResolver - rbac: RbacEngine = field(default_factory=RbacEngine) - - -def install_auth( - app: FastAPI, - config: AuthConfig, - resolver: IdentityResolver, - *, - rbac: Optional[RbacEngine] = None, - basic_verifier: Optional[BasicAuthVerifier] = None, - mount_scim: bool = True, - mount_oidc: bool = True, - mount_saml: bool = True, -) -> AuthContext: - """Wire the authenticators, resolver and optional routers onto the app. - - Deployment requirement for trusted-proxy IP resolution: uvicorn's - ``--proxy-headers`` (enabled by default) overwrites ``request.client`` from - ``X-Forwarded-For`` before this module's ``trusted_proxy_cidrs`` check runs, - which silently bypasses it. Run uvicorn with ``--no-proxy-headers`` and let - this module resolve the client IP, or leave ``trusted_proxy_cidrs`` empty and - rely on uvicorn's own ``--forwarded-allow-ips``. Do not enable both. - """ - engine = rbac if rbac is not None else RbacEngine(config.casbin_policy_path) - ctx = AuthContext( - config, - build_authenticators(config, basic_verifier=basic_verifier), - resolver, - engine, - ) - app.state.auth_v2 = ctx - if mount_scim: - from .scim import build_scim_router - - app.include_router(build_scim_router()) - if mount_oidc and config.oidc_providers: - from .oidc import build_oidc_router - - app.include_router(build_oidc_router(config)) - if mount_saml and config.saml is not None and config.saml.enabled: - from .saml import SamlAuthenticator, SamlSessionStore, build_saml_router - - session_store = SamlSessionStore( - ttl_seconds=config.saml.session_ttl_seconds, - max_size=config.saml.session_max_size, - ) - ctx.authenticators.append(SamlAuthenticator(config.saml, session_store)) - app.include_router(build_saml_router(config.saml, session_store)) - return ctx - - -def _ctx(request: Request) -> AuthContext: - return request.app.state.auth_v2 +from .session import SessionAuthenticator, SessionStore def _combined_challenge(authenticators: List[Authenticator]) -> str: @@ -84,47 +26,82 @@ def _combined_challenge(authenticators: List[Authenticator]) -> str: return ", ".join(seen) -async def get_current_principal( - security_scopes: SecurityScopes, request: Request -) -> Principal: - ctx = _ctx(request) - credential = None - for authenticator in ctx.authenticators: - credential = await authenticator.authenticate(request) - if credential is not None: - break - if credential is None: - raise errors.unauthenticated(_combined_challenge(ctx.authenticators)) +class AuthSecurity: + """Enforcement layer consumed purely through FastAPI ``Security()``. - resolved = await ctx.resolver.resolve(credential) - principal = resolved.model_copy( - update={"network": resolve_network_context(request, ctx.config.network)} - ) + Construct once at the composition root and pass the bound methods + (``principal``, ``require_roles``, ``require_permission``) to ``Security()``; + routers receive the instance explicitly via ``build_*_router(auth)``. There is + no app mutation and no ``app.state``. - if not has_required_scopes(security_scopes, principal): - raise errors.insufficient_scope() - return principal + Deployment note for trusted-proxy IP resolution: uvicorn's ``--proxy-headers`` + (on by default) overwrites ``request.client`` from ``X-Forwarded-For`` before + this module's ``trusted_proxy_cidrs`` check runs, silently bypassing it. Run + uvicorn with ``--no-proxy-headers`` and let this module resolve the client IP, + or leave ``trusted_proxy_cidrs`` empty and rely on uvicorn's + ``--forwarded-allow-ips``. Do not enable both. + """ + def __init__( + self, + config: AuthConfig, + resolver: IdentityResolver, + rbac: Optional[RBACEngine] = None, + authenticators: Optional[List[Authenticator]] = None, + basic_verifier: Optional[BasicAuthVerifier] = None, + ) -> None: + self.config = config + self.resolver = resolver + self.rbac = rbac or RBACEngine(config.casbin_policy_path) + self.session_store = SessionStore( + config.session.ttl_seconds, config.session.max_size + ) + self.oauth_txn_store = SessionStore( + config.session.login_state_ttl, config.session.max_size + ) + chain = ( + list(authenticators) + if authenticators is not None + else build_authenticators(config, basic_verifier=basic_verifier) + ) + chain.append(SessionAuthenticator(config.session.cookie, self.session_store)) + self.authenticators = chain -def require_roles(*allowed: Role) -> Callable[..., object]: - async def dependency( - request: Request, - principal: Annotated[Principal, Security(get_current_principal)], + async def principal( + self, security_scopes: SecurityScopes, request: Request ) -> Principal: - if not _ctx(request).rbac.has_role(principal, allowed): - raise errors.forbidden_role() + credential = None + for authenticator in self.authenticators: + credential = await authenticator.authenticate(request) + if credential is not None: + break + if credential is None: + raise errors.unauthenticated(_combined_challenge(self.authenticators)) + + resolved = await self.resolver.resolve(credential) + principal = resolved.model_copy( + update={"network": resolve_network_context(request, self.config.network)} + ) + if not has_required_scopes(security_scopes, principal): + raise errors.insufficient_scope() return principal - return dependency + def require_roles(self, *allowed: Role) -> Callable[..., object]: + async def dependency( + principal: Annotated[Principal, Security(self.principal)], + ) -> Principal: + if not self.rbac.has_any_role(principal, allowed): + raise errors.forbidden_role() + return principal + return dependency -def require_permission(obj: str, act: str) -> Callable[..., object]: - async def dependency( - request: Request, - principal: Annotated[Principal, Security(get_current_principal)], - ) -> Principal: - if not _ctx(request).rbac.enforce(principal, obj, act): - raise errors.forbidden_permission() - return principal + def require_permission(self, obj: str, act: str) -> Callable[..., object]: + async def dependency( + principal: Annotated[Principal, Security(self.principal)], + ) -> Principal: + if not self.rbac.enforce(principal, obj, act): + raise errors.forbidden_permission() + return principal - return dependency + return dependency diff --git a/litellm/proxy/auth_v2/session.py b/litellm/proxy/auth_v2/session.py new file mode 100644 index 00000000000..fe671ce40f6 --- /dev/null +++ b/litellm/proxy/auth_v2/session.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +import secrets +import time +from typing import Any, Dict, Optional, Tuple + +from fastapi import Request + +from .models import AuthMethod, Credential, CredentialRef, SecuritySchemeType + + +def safe_relay_state(target: Optional[str], default: str) -> str: + if ( + target + and target.startswith("/") + and not target.startswith("//") + and "://" not in target + and "\\" not in target + ): + return target + return default + + +class SessionStore: + def __init__(self, ttl_seconds: int = 3600, max_size: int = 10000) -> None: + self._sessions: Dict[str, Tuple[float, Dict[str, Any]]] = {} + self._ttl = ttl_seconds + self._max_size = max_size + + def create_session(self, identity: Dict[str, Any]) -> str: + now = time.time() + self._evict(now) + session_id = secrets.token_urlsafe(32) + self._sessions[session_id] = (now + self._ttl, identity) + return session_id + + def get(self, session_id: str) -> Optional[Dict[str, Any]]: + entry = self._sessions.get(session_id) + if entry is None: + return None + expires_at, identity = entry + if expires_at < time.time(): + self._sessions.pop(session_id, None) + return None + return identity + + def pop(self, session_id: str) -> Optional[Dict[str, Any]]: + entry = self._sessions.pop(session_id, None) + if entry is None: + return None + expires_at, identity = entry + if expires_at < time.time(): + return None + return identity + + def _evict(self, now: float) -> None: + for key in [k for k, (exp, _) in self._sessions.items() if exp < now]: + self._sessions.pop(key, None) + overflow = len(self._sessions) - self._max_size + 1 + if overflow > 0: + oldest = sorted(self._sessions, key=lambda k: self._sessions[k][0]) + for key in oldest[:overflow]: + self._sessions.pop(key, None) + + +class SessionAuthenticator: + def __init__(self, cookie_name: str, store: SessionStore) -> None: + self._cookie_name = cookie_name + self._store = store + + async def authenticate(self, request: Request) -> Optional[Credential]: + session_id = request.cookies.get(self._cookie_name) + if not session_id: + return None + identity = self._store.get(session_id) + if identity is None: + return None + return Credential( + scheme=SecuritySchemeType.API_KEY, + method=AuthMethod(identity["method"]), + subject=identity["subject"], + issuer=identity.get("issuer"), + claims=identity.get("claims", {}), + credential_ref=CredentialRef(token_id=session_id), + ) + + def challenge(self) -> str: + return ""