mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(auth_v2): harden HTTP basic, SAML sessions, and JWKS fetch
Address Greptile security findings in the authenticator, SAML and config layers: - HTTP Basic accepted any password and copied the cleartext password into Principal.claims. Verify the password against an injected BasicAuthVerifier (InMemoryBasicAuthStore holds username -> salted sha256, constant-time compared with hmac.compare_digest) and stop putting the password in the credential; basic with no configured verifier now rejects rather than trusting the caller. - SAML session cookie gains the Secure flag (httponly and samesite=lax already set), gated by SamlConfig.cookie_secure. - SAML session store gains TTL expiry and max-size eviction (SamlConfig.session_ttl_seconds / session_max_size) so it can no longer grow unbounded or hand out stale sessions. - JWKS signing-key lookup ran synchronously inside the async request path and blocked the event loop on a cache miss; run the JWT verify off-loop via starlette run_in_threadpool on the http-bearer, oauth2 at+jwt and oidc paths.
This commit is contained in:
parent
2dd336a795
commit
71a189bf65
5 changed files with 117 additions and 27 deletions
|
|
@ -2,6 +2,10 @@ from __future__ import annotations
|
||||||
|
|
||||||
import base64
|
import base64
|
||||||
import binascii
|
import binascii
|
||||||
|
import functools
|
||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import secrets
|
||||||
from typing import Any, Dict, List, Optional, Protocol, runtime_checkable
|
from typing import Any, Dict, List, Optional, Protocol, runtime_checkable
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
@ -9,6 +13,7 @@ import jwt
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
from jwt import PyJWKClient
|
from jwt import PyJWKClient
|
||||||
from jwt import decode as jwt_decode
|
from jwt import decode as jwt_decode
|
||||||
|
from starlette.concurrency import run_in_threadpool
|
||||||
|
|
||||||
from . import errors
|
from . import errors
|
||||||
from .config import (
|
from .config import (
|
||||||
|
|
@ -86,6 +91,35 @@ 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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|
@ -134,6 +168,14 @@ class JwtVerifier:
|
||||||
raise errors.invalid_token(str(exc)) from exc
|
raise errors.invalid_token(str(exc)) from exc
|
||||||
|
|
||||||
|
|
||||||
|
async def _verify_jwt_off_loop(
|
||||||
|
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:
|
if not verifiers:
|
||||||
return None
|
return None
|
||||||
|
|
@ -173,10 +215,14 @@ class HttpAuthenticator:
|
||||||
scheme = SecuritySchemeType.HTTP
|
scheme = SecuritySchemeType.HTTP
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, basic: HttpBasicConfig, jwt_verifiers: List[JwtVerifier]
|
self,
|
||||||
|
basic: HttpBasicConfig,
|
||||||
|
jwt_verifiers: List[JwtVerifier],
|
||||||
|
basic_verifier: Optional[BasicAuthVerifier] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._basic = basic
|
self._basic = basic
|
||||||
self._verifiers = jwt_verifiers
|
self._verifiers = jwt_verifiers
|
||||||
|
self._basic_verifier = basic_verifier
|
||||||
|
|
||||||
async def authenticate(self, request: Request) -> Optional[Credential]:
|
async def authenticate(self, request: Request) -> Optional[Credential]:
|
||||||
header = request.headers.get("authorization")
|
header = request.headers.get("authorization")
|
||||||
|
|
@ -185,35 +231,38 @@ class HttpAuthenticator:
|
||||||
scheme, _, value = header.partition(" ")
|
scheme, _, value = header.partition(" ")
|
||||||
scheme_lower = scheme.lower()
|
scheme_lower = scheme.lower()
|
||||||
if scheme_lower == "bearer" and value:
|
if scheme_lower == "bearer" and value:
|
||||||
return self._verify_bearer(value)
|
return await self._verify_bearer(value)
|
||||||
if scheme_lower == "basic" and self._basic.enabled and value:
|
if scheme_lower == "basic" and self._basic.enabled and value:
|
||||||
return self._verify_basic(value)
|
return self._verify_basic(value)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _verify_bearer(self, token: str) -> Credential:
|
async def _verify_bearer(self, token: str) -> Credential:
|
||||||
verifier = _select_verifier(token, self._verifiers)
|
verifier = _select_verifier(token, self._verifiers)
|
||||||
if verifier is None:
|
if verifier is None:
|
||||||
raise errors.invalid_token("no issuer match")
|
raise errors.invalid_token("no issuer match")
|
||||||
claims = verifier.verify(token)
|
claims = await _verify_jwt_off_loop(verifier, token)
|
||||||
return _credential_from_claims(
|
return _credential_from_claims(
|
||||||
self.scheme, AuthMethod.BEARER_JWT, token, claims
|
self.scheme, AuthMethod.BEARER_JWT, token, claims
|
||||||
)
|
)
|
||||||
|
|
||||||
def _verify_basic(self, value: str) -> Credential:
|
def _verify_basic(self, value: str) -> Credential:
|
||||||
|
challenge = errors.basic_challenge(self._basic.realm)
|
||||||
try:
|
try:
|
||||||
decoded = base64.b64decode(value).decode("utf-8")
|
decoded = base64.b64decode(value).decode("utf-8")
|
||||||
except (binascii.Error, UnicodeDecodeError) as exc:
|
except (binascii.Error, UnicodeDecodeError) as exc:
|
||||||
raise errors.unauthenticated(
|
raise errors.unauthenticated(challenge) from exc
|
||||||
errors.basic_challenge(self._basic.realm)
|
username, separator, password = decoded.partition(":")
|
||||||
) from exc
|
if (
|
||||||
username, _, password = decoded.partition(":")
|
not username
|
||||||
if not username:
|
or separator != ":"
|
||||||
raise errors.unauthenticated(errors.basic_challenge(self._basic.realm))
|
or self._basic_verifier is None
|
||||||
|
or not self._basic_verifier.verify(username, password)
|
||||||
|
):
|
||||||
|
raise errors.unauthenticated(challenge)
|
||||||
return Credential(
|
return Credential(
|
||||||
scheme=self.scheme,
|
scheme=self.scheme,
|
||||||
method=AuthMethod.HTTP_BASIC,
|
method=AuthMethod.HTTP_BASIC,
|
||||||
subject=username,
|
subject=username,
|
||||||
claims={"_basic_password": password},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def challenge(self) -> str:
|
def challenge(self) -> str:
|
||||||
|
|
@ -239,16 +288,16 @@ class OAuth2Authenticator:
|
||||||
if token is None:
|
if token is None:
|
||||||
return None
|
return None
|
||||||
if _looks_like_jwt(token):
|
if _looks_like_jwt(token):
|
||||||
return self._verify_at_jwt(token)
|
return await self._verify_at_jwt(token)
|
||||||
if self._introspection is not None:
|
if self._introspection is not None:
|
||||||
return await self._introspect(token)
|
return await self._introspect(token)
|
||||||
raise errors.invalid_token()
|
raise errors.invalid_token()
|
||||||
|
|
||||||
def _verify_at_jwt(self, token: str) -> Credential:
|
async def _verify_at_jwt(self, token: str) -> Credential:
|
||||||
verifier = _select_verifier(token, self._verifiers)
|
verifier = _select_verifier(token, self._verifiers)
|
||||||
if verifier is None:
|
if verifier is None:
|
||||||
raise errors.invalid_token("no issuer match")
|
raise errors.invalid_token("no issuer match")
|
||||||
claims = verifier.verify(token, require_at_jwt=True)
|
claims = await _verify_jwt_off_loop(verifier, token, require_at_jwt=True)
|
||||||
return _credential_from_claims(
|
return _credential_from_claims(
|
||||||
self.scheme, AuthMethod.BEARER_JWT, token, claims
|
self.scheme, AuthMethod.BEARER_JWT, token, claims
|
||||||
)
|
)
|
||||||
|
|
@ -301,7 +350,7 @@ class OidcAuthenticator:
|
||||||
verifier = _select_verifier(token, self._verifiers)
|
verifier = _select_verifier(token, self._verifiers)
|
||||||
if verifier is None:
|
if verifier is None:
|
||||||
raise errors.invalid_token("no issuer match")
|
raise errors.invalid_token("no issuer match")
|
||||||
claims = verifier.verify(token)
|
claims = await _verify_jwt_off_loop(verifier, token)
|
||||||
return _credential_from_claims(self.scheme, AuthMethod.OIDC, token, claims)
|
return _credential_from_claims(self.scheme, AuthMethod.OIDC, token, claims)
|
||||||
|
|
||||||
def challenge(self) -> str:
|
def challenge(self) -> str:
|
||||||
|
|
@ -337,12 +386,16 @@ class MutualTlsAuthenticator:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
|
|
||||||
def build_authenticators(config: AuthConfig) -> List[Authenticator]:
|
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] = {}
|
by_scheme: Dict[SecuritySchemeType, Authenticator] = {}
|
||||||
if config.api_key is not None:
|
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)
|
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(
|
by_scheme[SecuritySchemeType.OAUTH2] = OAuth2Authenticator(
|
||||||
verifiers, config.oauth2_introspection
|
verifiers, config.oauth2_introspection
|
||||||
|
|
|
||||||
|
|
@ -68,6 +68,9 @@ class SamlConfig(BaseModel):
|
||||||
sp_cert_file: Optional[str] = None
|
sp_cert_file: Optional[str] = None
|
||||||
allow_unsolicited: bool = True
|
allow_unsolicited: bool = True
|
||||||
session_cookie: str = "saml_session"
|
session_cookie: str = "saml_session"
|
||||||
|
cookie_secure: bool = True
|
||||||
|
session_ttl_seconds: int = 3600
|
||||||
|
session_max_size: int = 10000
|
||||||
default_redirect_path: str = "/"
|
default_redirect_path: str = "/"
|
||||||
xmlsec_binary: Optional[str] = None
|
xmlsec_binary: Optional[str] = None
|
||||||
attribute_map: Dict[str, str] = Field(
|
attribute_map: Dict[str, str] = Field(
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,8 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import secrets
|
import secrets
|
||||||
from typing import Any, Dict, List, Optional
|
import time
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
from fastapi import APIRouter, HTTPException, Request
|
from fastapi import APIRouter, HTTPException, Request
|
||||||
from fastapi.responses import RedirectResponse, Response
|
from fastapi.responses import RedirectResponse, Response
|
||||||
|
|
@ -138,20 +139,40 @@ def build_sp_client(config: SamlConfig) -> Saml2Client:
|
||||||
|
|
||||||
|
|
||||||
class SamlSessionStore:
|
class SamlSessionStore:
|
||||||
def __init__(self) -> None:
|
def __init__(self, ttl_seconds: int = 3600, max_size: int = 10000) -> None:
|
||||||
self._sessions: Dict[str, Dict[str, Any]] = {}
|
self._sessions: Dict[str, Tuple[float, Dict[str, Any]]] = {}
|
||||||
self.outstanding: Dict[str, str] = {}
|
self.outstanding: Dict[str, str] = {}
|
||||||
|
self._ttl = ttl_seconds
|
||||||
|
self._max_size = max_size
|
||||||
|
|
||||||
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
|
self.outstanding[request_id] = relay_state
|
||||||
|
|
||||||
def create_session(self, identity: Dict[str, Any]) -> str:
|
def create_session(self, identity: Dict[str, Any]) -> str:
|
||||||
|
now = time.time()
|
||||||
|
self._evict(now)
|
||||||
session_id = secrets.token_urlsafe(32)
|
session_id = secrets.token_urlsafe(32)
|
||||||
self._sessions[session_id] = identity
|
self._sessions[session_id] = (now + self._ttl, identity)
|
||||||
return session_id
|
return session_id
|
||||||
|
|
||||||
def get(self, session_id: str) -> Optional[Dict[str, Any]]:
|
def get(self, session_id: str) -> Optional[Dict[str, Any]]:
|
||||||
return self._sessions.get(session_id)
|
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:
|
class SamlAuthenticator:
|
||||||
|
|
@ -248,7 +269,11 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP
|
||||||
)
|
)
|
||||||
response = RedirectResponse(target, status_code=303)
|
response = RedirectResponse(target, status_code=303)
|
||||||
response.set_cookie(
|
response.set_cookie(
|
||||||
config.session_cookie, session_id, httponly=True, samesite="lax"
|
config.session_cookie,
|
||||||
|
session_id,
|
||||||
|
httponly=True,
|
||||||
|
samesite="lax",
|
||||||
|
secure=config.cookie_secure,
|
||||||
)
|
)
|
||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,7 @@ from fastapi import FastAPI, Request, Security
|
||||||
from fastapi.security import SecurityScopes
|
from fastapi.security import SecurityScopes
|
||||||
|
|
||||||
from . import errors
|
from . import errors
|
||||||
from .authenticators import Authenticator, build_authenticators
|
from .authenticators import Authenticator, BasicAuthVerifier, build_authenticators
|
||||||
from .config import AuthConfig
|
from .config import AuthConfig
|
||||||
from .models import Principal
|
from .models import Principal
|
||||||
from .network import resolve_network_context
|
from .network import resolve_network_context
|
||||||
|
|
@ -29,6 +29,7 @@ def install_auth(
|
||||||
resolver: IdentityResolver,
|
resolver: IdentityResolver,
|
||||||
*,
|
*,
|
||||||
rbac: Optional[RbacEngine] = None,
|
rbac: Optional[RbacEngine] = None,
|
||||||
|
basic_verifier: Optional[BasicAuthVerifier] = None,
|
||||||
mount_scim: bool = True,
|
mount_scim: bool = True,
|
||||||
mount_oidc: bool = True,
|
mount_oidc: bool = True,
|
||||||
mount_saml: bool = True,
|
mount_saml: bool = True,
|
||||||
|
|
@ -43,7 +44,12 @@ def install_auth(
|
||||||
rely on uvicorn's own ``--forwarded-allow-ips``. Do not enable both.
|
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)
|
engine = rbac if rbac is not None else RbacEngine(config.casbin_policy_path)
|
||||||
ctx = AuthContext(config, build_authenticators(config), resolver, engine)
|
ctx = AuthContext(
|
||||||
|
config,
|
||||||
|
build_authenticators(config, basic_verifier=basic_verifier),
|
||||||
|
resolver,
|
||||||
|
engine,
|
||||||
|
)
|
||||||
app.state.auth_v2 = ctx
|
app.state.auth_v2 = ctx
|
||||||
if mount_scim:
|
if mount_scim:
|
||||||
from .scim import build_scim_router
|
from .scim import build_scim_router
|
||||||
|
|
@ -56,7 +62,10 @@ def install_auth(
|
||||||
if mount_saml and config.saml is not None and config.saml.enabled:
|
if mount_saml and config.saml is not None and config.saml.enabled:
|
||||||
from .saml import SamlAuthenticator, SamlSessionStore, build_saml_router
|
from .saml import SamlAuthenticator, SamlSessionStore, build_saml_router
|
||||||
|
|
||||||
session_store = SamlSessionStore()
|
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))
|
ctx.authenticators.append(SamlAuthenticator(config.saml, session_store))
|
||||||
app.include_router(build_saml_router(config.saml, session_store))
|
app.include_router(build_saml_router(config.saml, session_store))
|
||||||
return ctx
|
return ctx
|
||||||
|
|
|
||||||
|
|
@ -287,7 +287,7 @@ plugins = "pydantic.mypy"
|
||||||
|
|
||||||
# scim2-models ships py.typed, but its generic, alias-driven SCIM models report
|
# scim2-models ships py.typed, but its generic, alias-driven SCIM models report
|
||||||
# phantom call-arg errors under mypy though they work at runtime. Treat the
|
# phantom call-arg errors under mypy though they work at runtime. Treat the
|
||||||
# library as untyped at the boundary; litellm/auth_v2 is its only consumer.
|
# library as untyped at the boundary; litellm/proxy/auth_v2 is its only consumer.
|
||||||
# CI runs mypy from litellm/ against litellm/mypy.ini, which carries the same
|
# CI runs mypy from litellm/ against litellm/mypy.ini, which carries the same
|
||||||
# override; this block keeps root-level mypy runs consistent.
|
# override; this block keeps root-level mypy runs consistent.
|
||||||
[[tool.mypy.overrides]]
|
[[tool.mypy.overrides]]
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue