mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(auth_v2): replace install_auth with AuthSecurity DI
Per user direction, drop install_auth/AuthContext/app.state entirely; the enforcement layer is now an AuthSecurity object whose bound methods are the FastAPI Security() dependencies. The app constructs AuthSecurity(config, resolver) once and passes auth.principal / auth.require_roles / auth.require_permission to Security(); routers take the instance explicitly via build_*_router(auth) and read auth.resolver/auth.config rather than request.app.state. Browser sessions are unified behind a shared SessionStore + SessionAuthenticator (session.py): one cookie, keyed on identity["method"], so SAML and the upcoming OIDC login flow share one store. AuthSecurity owns the post-login session_store and a short-TTL oauth_txn_store for OIDC state/nonce/PKCE; SessionConfig moves the cookie/TTL/redirect settings off SAMLConfig. SAML keeps its protocol-specific outstanding/replay state local. Fold in the standing judge findings: collapse the triplicated http/oauth2/oidc bearer paths into one _authenticate_bearer_jwt helper, inject the introspection async-client via a factory instead of importing litellm inline, drop the dead scheme attribute from the authenticators, and rename to PEP 8 acronym casing (JWTVerifier, OIDCAuthenticator, OIDCProviderConfig, SAMLConfig, RBACEngine, APIKeyAuthenticator, MutualTLSAuthenticator). __all__ now exports AuthSecurity, Role and the resolver protocols. scim.py and oidc.py move to build_*_router(auth) separately.
This commit is contained in:
parent
c512a49fa5
commit
8302f55995
7 changed files with 337 additions and 316 deletions
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
88
litellm/proxy/auth_v2/session.py
Normal file
88
litellm/proxy/auth_v2/session.py
Normal file
|
|
@ -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 ""
|
||||
Loading…
Add table
Reference in a new issue