feat(auth_v2): implement full SAML 2.0 SP via pysaml2

Replace the deferred SAML thin-adapter stub with a working Service Provider
built on pysaml2: an SP metadata endpoint, an SP-initiated /login that
redirects to the IdP, and an ACS handling the HTTP-POST binding that verifies
the signed assertion, maps NameID and attribute statements into a
scim2_models.User, and upserts it through the same ProvisioningStore seam SCIM
and OIDC use. A SamlAuthenticator reads the post-ACS session cookie and resolves
to the one normalized Principal like every other scheme; AuthMethod gains a SAML
member. IdP metadata loads from a file path or inline XML via SamlConfig, and
install_auth mounts the router and authenticator when SAML is enabled.

pysaml2 pulls pyOpenSSL transitively without pinning it, and older pyOpenSSL
caps cryptography below 46 and breaks at import against the version this proxy
already requires; pin pyOpenSSL>=26 so the resolver stays on a
cryptography-46-compatible release. pysaml2 also needs the system xmlsec1
binary at runtime (brew install libxmlsec1 on macOS, apt-get install xmlsec1
libxmlsec1-dev on Debian); SamlConfig.xmlsec_binary can point at it when it is
not on PATH.
This commit is contained in:
Yassin Kortam 2026-06-10 17:33:02 -07:00
parent a0a59a2197
commit 0b74ffa9c6
6 changed files with 283 additions and 19 deletions

View file

@ -1,11 +1,22 @@
from __future__ import annotations
from typing import List, Optional
from typing import Dict, List, Optional
from pydantic import AnyHttpUrl, BaseModel, Field, SecretStr
from pydantic import AnyHttpUrl, BaseModel, Field, SecretStr, model_validator
from .models import SecuritySchemeType
DEFAULT_SAML_ATTRIBUTE_MAP = {
"email": "email",
"mail": "email",
"emailAddress": "email",
"displayName": "display_name",
"cn": "display_name",
"userName": "user_name",
"uid": "user_name",
"sAMAccountName": "user_name",
}
class ApiKeySchemeConfig(BaseModel):
header_name: str = "x-litellm-api-key"
@ -46,6 +57,31 @@ class TrustedProxyConfig(BaseModel):
trusted_proxy_cidrs: List[str] = Field(default_factory=list)
class SamlConfig(BaseModel):
enabled: bool = False
sp_entity_id: str
acs_url: str
idp_metadata_path: Optional[str] = None
idp_metadata_inline: Optional[str] = None
sp_key_file: Optional[str] = None
sp_cert_file: Optional[str] = None
want_assertions_signed: bool = True
allow_unsolicited: bool = True
session_cookie: str = "saml_session"
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":
if self.enabled and not (self.idp_metadata_path or self.idp_metadata_inline):
raise ValueError(
"SAML enabled but no IdP metadata: set idp_metadata_path or idp_metadata_inline"
)
return self
class AuthConfig(BaseModel):
scheme_order: List[SecuritySchemeType] = Field(
default_factory=lambda: [
@ -62,3 +98,4 @@ class AuthConfig(BaseModel):
oauth2_introspection: Optional[OAuth2IntrospectionConfig] = None
mutual_tls: MutualTlsConfig = Field(default_factory=MutualTlsConfig)
network: TrustedProxyConfig = Field(default_factory=TrustedProxyConfig)
saml: Optional[SamlConfig] = None

View file

@ -22,6 +22,7 @@ class AuthMethod(str, Enum):
BEARER_JWT = "bearer_jwt"
OAUTH2_INTROSPECTION = "oauth2_introspection"
OIDC = "oidc"
SAML = "saml"
MUTUAL_TLS = "mutual_tls"

View file

@ -1,35 +1,184 @@
from __future__ import annotations
from typing import Optional
import secrets
from typing import Any, Dict, Optional
from fastapi import APIRouter, Request
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import JSONResponse, RedirectResponse, Response
from saml2 import BINDING_HTTP_POST
from saml2.client import Saml2Client
from saml2.config import SPConfig
from saml2.metadata import entity_descriptor
from scim2_models import User as ScimUser
from .models import Credential, SecuritySchemeType
from .config import SamlConfig
from .models import AuthMethod, Credential, CredentialRef, SecuritySchemeType
from .resolver import ProvisioningStore
_SINGLE_VALUE_CLAIM = {
"email": "email",
"user_name": "preferred_username",
"display_name": "name",
}
_MULTI_VALUE_ATTRS = ("groups", "roles")
def _normalize_attributes(
ava: Dict[str, Any], attribute_map: Dict[str, str]
) -> Dict[str, Any]:
claims: Dict[str, Any] = {}
for saml_attr, target in attribute_map.items():
if saml_attr not in ava or target not in _SINGLE_VALUE_CLAIM:
continue
value = ava[saml_attr]
scalar = value[0] if isinstance(value, list) and value else value
claims.setdefault(_SINGLE_VALUE_CLAIM[target], scalar)
for attr in _MULTI_VALUE_ATTRS:
value = ava.get(attr)
if isinstance(value, list):
claims[attr] = value
return claims
def _user_from_claims(name_id: str, claims: Dict[str, Any]) -> ScimUser:
return ScimUser(
external_id=name_id,
user_name=claims.get("preferred_username") or claims.get("email") or name_id,
display_name=claims.get("name"),
)
def _sp_config_dict(config: SamlConfig) -> Dict[str, Any]:
cfg: Dict[str, Any] = {
"entityid": config.sp_entity_id,
"service": {
"sp": {
"endpoints": {
"assertion_consumer_service": [(config.acs_url, BINDING_HTTP_POST)]
},
"allow_unsolicited": config.allow_unsolicited,
"authn_requests_signed": False,
"want_assertions_signed": config.want_assertions_signed,
"want_response_signed": False,
}
},
"allow_unknown_attributes": True,
}
if config.idp_metadata_path:
cfg["metadata"] = {"local": [config.idp_metadata_path]}
elif config.idp_metadata_inline:
cfg["metadata"] = {"inline": [config.idp_metadata_inline]}
if config.sp_key_file:
cfg["key_file"] = config.sp_key_file
if config.sp_cert_file:
cfg["cert_file"] = config.sp_cert_file
if config.xmlsec_binary:
cfg["xmlsec_binary"] = config.xmlsec_binary
return cfg
def build_sp_client(config: SamlConfig) -> Saml2Client:
conf = SPConfig()
conf.load(_sp_config_dict(config))
return Saml2Client(config=conf)
class SamlSessionStore:
def __init__(self) -> None:
self._sessions: Dict[str, Dict[str, Any]] = {}
self.outstanding: Dict[str, str] = {}
def remember_request(self, request_id: str, relay_state: str = "/") -> None:
self.outstanding[request_id] = relay_state
def create_session(self, identity: Dict[str, Any]) -> str:
session_id = secrets.token_urlsafe(32)
self._sessions[session_id] = identity
return session_id
def get(self, session_id: str) -> Optional[Dict[str, Any]]:
return self._sessions.get(session_id)
class SamlAuthenticator:
"""Thin SAML SP seam. Full pysaml2 wiring (system libxmlsec1, pinned
xmlsec/lxml, multi-IdP metadata) is deferred; the ACS maps a SAML assertion's
NameID and attribute statements into the same scim2_models.User upsert as
OIDC and SCIM."""
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]:
raise NotImplementedError("SAML SP deferred; see 03-design.md cut list")
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() -> APIRouter:
def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> APIRouter:
client = build_sp_client(config)
router = APIRouter(prefix="/auth/saml", tags=["saml"])
@router.get("/metadata")
async def metadata() -> None:
raise NotImplementedError("requires pysaml2 + system libxmlsec1")
@router.post("/acs")
async def assertion_consumer_service(request: Request) -> None:
raise NotImplementedError(
"parse assertion -> scim2_models.User -> store.upsert_user"
async def metadata() -> Response:
return Response(
content=str(entity_descriptor(client.config)),
media_type="application/samlmetadata+xml",
)
@router.get("/login")
async def login() -> RedirectResponse:
request_id, info = client.prepare_for_authenticate()
session_store.remember_request(request_id)
location = dict(info["headers"]).get("Location")
if not location:
raise HTTPException(status_code=500, detail="no SAML redirect produced")
return RedirectResponse(location, status_code=303)
@router.post("/acs")
async def assertion_consumer_service(request: Request) -> Response:
form = await request.form()
saml_response = form.get("SAMLResponse")
if not isinstance(saml_response, str):
raise HTTPException(status_code=400, detail="missing SAMLResponse")
authn_response = client.parse_authn_request_response(
saml_response,
BINDING_HTTP_POST,
outstanding=session_store.outstanding or None,
)
if authn_response is None:
raise HTTPException(status_code=401, detail="invalid SAML response")
name_id = authn_response.get_subject().text
ava = authn_response.get_identity() or {}
claims = _normalize_attributes(ava, config.attribute_map)
user = _user_from_claims(name_id, claims)
store: ProvisioningStore = request.app.state.auth_v2.resolver
await store.upsert_user(user)
in_response_to = getattr(authn_response, "in_response_to", None)
if in_response_to:
session_store.outstanding.pop(in_response_to, None)
session_id = session_store.create_session(
{"name_id": name_id, "issuer": authn_response.issuer(), "claims": claims}
)
response = JSONResponse(content=user.model_dump())
response.set_cookie(
config.session_cookie, session_id, httponly=True, samesite="lax"
)
return response
return router

View file

@ -29,6 +29,7 @@ def install_auth(
*,
mount_scim: bool = True,
mount_oidc: bool = True,
mount_saml: bool = True,
) -> AuthContext:
ctx = AuthContext(config, build_authenticators(config), resolver)
app.state.auth_v2 = ctx
@ -40,6 +41,12 @@ def install_auth(
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()
ctx.authenticators.append(SamlAuthenticator(config.saml, session_store))
app.include_router(build_saml_router(config.saml, session_store))
return ctx

View file

@ -56,6 +56,11 @@ proxy = [
"PyJWT[crypto]>=2.13.0,<3.0",
"Authlib>=1.6.0,<2.0",
"scim2-models>=0.6.0,<1.0",
"pysaml2>=7.5.0,<8.0",
# pysaml2 pulls pyOpenSSL transitively without pinning it; force a floor that
# supports cryptography 46 (older pyOpenSSL caps cryptography below 46 and
# breaks at import against the version this proxy already requires).
"pyOpenSSL>=26.0.0,<27.0",
"python-multipart>=0.0.27,<1.0",
"cryptography>=46.0.7,<47.0",
"pynacl>=1.6.2,<2.0",

65
uv.lock generated
View file

@ -1297,6 +1297,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/26/8c/e1a7043e562b5b29fb5d0930630a18078fecb1c30ca6776221ce0dab6f95/ddtrace-2.19.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:d36a16e8746cb38a143faa6e1cd10927bf4a482c29f4010afde4bd0f4bb89db4", size = 7390107, upload-time = "2025-01-16T17:17:54.826Z" },
]
[[package]]
name = "defusedxml"
version = "0.7.1"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/0f/d5/c66da9b79e5bdb124974bfe172b4daf3c984ebd9c2a06e2b8a4dc7331c72/defusedxml-0.7.1.tar.gz", hash = "sha256:1bb3032db185915b62d7c6209c5a8792be6a32ab2fedacc84e01b52c51aa3e69", size = 75520, upload-time = "2021-03-08T10:59:26.269Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/07/6c/aa3f2f849e01cb6a001cd8554a88d4c77c5c1a31c95bdf1cf9301e6d9ef4/defusedxml-0.7.1-py2.py3-none-any.whl", hash = "sha256:a352e7e428770286cc899e2542b6cdaedb2b4953ff269a210103ec58f6198a61", size = 25604, upload-time = "2021-03-08T10:59:24.45Z" },
]
[[package]]
name = "deprecated"
version = "1.3.1"
@ -1422,6 +1431,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/02/10/5da547df7a391dcde17f59520a231527b8571e6f46fc8efb02ccb370ab12/docutils-0.22.4-py3-none-any.whl", hash = "sha256:d0013f540772d1420576855455d050a2180186c91c15779301ac2ccb3eeb68de", size = 633196, upload-time = "2025-12-18T19:00:18.077Z" },
]
[[package]]
name = "elementpath"
version = "4.8.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/ac/41/afdd82534c80e9675d1c51dc21d0889b72d023bfe395a2f5a44d751d3a73/elementpath-4.8.0.tar.gz", hash = "sha256:5822a2560d99e2633d95f78694c7ff9646adaa187db520da200a8e9479dc46ae", size = 358528, upload-time = "2025-03-03T20:51:08.397Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/45/95/615af832e7f507fe5ce4562b4be1bd2fec080c4ff6da88dcd0c2dbfca582/elementpath-4.8.0-py3-none-any.whl", hash = "sha256:5393191f84969bcf8033b05ec4593ef940e58622ea13cefe60ecefbbf09d58d9", size = 243271, upload-time = "2025-03-03T20:51:03.027Z" },
]
[[package]]
name = "email-validator"
version = "2.3.0"
@ -3366,7 +3384,9 @@ proxy = [
{ name = "pydantic-settings" },
{ name = "pyjwt", extra = ["crypto"] },
{ name = "pynacl" },
{ name = "pyopenssl" },
{ name = "pyroscope-io", marker = "sys_platform != 'win32'" },
{ name = "pysaml2" },
{ name = "python-multipart" },
{ name = "pyyaml" },
{ name = "restrictedpython" },
@ -3557,8 +3577,10 @@ requires-dist = [
{ name = "pydantic-settings", marker = "extra == 'proxy'", specifier = ">=2.14.1,<3.0" },
{ name = "pyjwt", extras = ["crypto"], marker = "extra == 'proxy'", specifier = ">=2.13.0,<3.0" },
{ name = "pynacl", marker = "extra == 'proxy'", specifier = ">=1.6.2,<2.0" },
{ name = "pyopenssl", marker = "extra == 'proxy'", specifier = ">=26.0.0,<27.0" },
{ name = "pypdf", marker = "python_full_version < '3.14' and extra == 'proxy-runtime'", specifier = ">=6.10.2,<7.0" },
{ name = "pyroscope-io", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.8.16,<1.0" },
{ name = "pysaml2", marker = "extra == 'proxy'", specifier = ">=7.5.0,<8.0" },
{ name = "python-dotenv", specifier = ">=1.0.0,<2.0" },
{ name = "python-multipart", marker = "extra == 'proxy'", specifier = ">=0.0.27,<1.0" },
{ name = "pyyaml", marker = "extra == 'cli'", specifier = ">=6.0.3,<7.0" },
@ -6077,6 +6099,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/29/7d/5945b5af29534641820d3bd7b00962abbbdfee84ec7e19f0d5b3175f9a31/pynacl-1.6.2-cp38-abi3-win_arm64.whl", hash = "sha256:834a43af110f743a754448463e8fd61259cd4ab5bbedcf70f9dabad1d28a394c", size = 184801, upload-time = "2026-01-01T17:32:36.309Z" },
]
[[package]]
name = "pyopenssl"
version = "26.2.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cryptography" },
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/1a/51/27a5ad5f939d08f690a326ef9582cda7140555180db71695f6fb747d6a36/pyopenssl-26.2.0.tar.gz", hash = "sha256:8c6fcecd1183a7fc897548dfe388b0cdb7f37e018200d8409cf33959dbe35387", size = 182195, upload-time = "2026-05-04T23:06:09.72Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/73/b8/a0e2790ae249d6f38c9f66de7a211621a7ab2650217bcd04e1262f578a56/pyopenssl-26.2.0-py3-none-any.whl", hash = "sha256:4f9d971bc5298b8bc1fab282803da04bf000c755d4ad9d99b52de2569ca19a70", size = 55823, upload-time = "2026-05-04T23:06:08.395Z" },
]
[[package]]
name = "pyparsing"
version = "3.3.2"
@ -6134,6 +6169,24 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/eb/8f/88d792e9cacd6ff3bd9a50100586ddc665e02a917662c17d30931f778542/pyroscope_io-0.8.16-py2.py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:6b91ce5b240f8de756c16a17022ca8e25ef8a4eed461c7d074b8a0841cf7b445", size = 3485288, upload-time = "2026-01-22T06:23:32Z" },
]
[[package]]
name = "pysaml2"
version = "7.5.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cryptography" },
{ name = "defusedxml" },
{ name = "pyopenssl" },
{ name = "python-dateutil" },
{ name = "pytz" },
{ name = "requests" },
{ name = "xmlschema" },
]
sdist = { url = "https://files.pythonhosted.org/packages/76/02/e8ecb5d1574a2add1431c8ec16dff137610f30580a7c1d6205929b3db3ee/pysaml2-7.5.0.tar.gz", hash = "sha256:f36871d4e5ee857c6b85532e942550d2cf90ea4ee943d75eb681044bbc4f54f7", size = 340338, upload-time = "2024-01-30T11:49:08.589Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/fe/d1/92d84ae0e80e829e84785c6e4e425ff6d447116289f0ecf2af068f771a73/pysaml2-7.5.0-py3-none-any.whl", hash = "sha256:bc6627cc344476a83c757f440a73fda1369f13b6fda1b4e16bca63ffbabb5318", size = 419304, upload-time = "2024-01-30T11:49:04.5Z" },
]
[[package]]
name = "pytest"
version = "9.0.3"
@ -8153,6 +8206,18 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/a4/f5/10b68b7b1544245097b2a1b8238f66f2fc6dcaeb24ba5d917f52bd2eed4f/wsproto-1.3.2-py3-none-any.whl", hash = "sha256:61eea322cdf56e8cc904bd3ad7573359a242ba65688716b0710a5eb12beab584", size = 24405, upload-time = "2025-11-20T18:18:00.454Z" },
]
[[package]]
name = "xmlschema"
version = "2.5.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "elementpath" },
]
sdist = { url = "https://files.pythonhosted.org/packages/59/af/42e9e773eaa6bc8e8c322f93c75454b0d370979048b250bfef7786ff26ec/xmlschema-2.5.1.tar.gz", hash = "sha256:4f7497de6c8b6dc2c28ad7b9ed6e21d186f4afe248a5bea4f54eedab4da44083", size = 539267, upload-time = "2023-12-19T15:51:57.663Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/5e/2a/c2bc97fd20efe65cfcfc21666d1b0213969133d37ea093761d264d9ed9f8/xmlschema-2.5.1-py3-none-any.whl", hash = "sha256:ec2b2a15c8896c1fcd14dcee34ca30032b99456c3c43ce793fdb9dca2fb4b869", size = 395065, upload-time = "2023-12-19T15:51:53.136Z" },
]
[[package]]
name = "xmltodict"
version = "1.0.4"