diff --git a/litellm/auth_v2/config.py b/litellm/auth_v2/config.py index c4a51afbabb..8b2fdc98ec5 100644 --- a/litellm/auth_v2/config.py +++ b/litellm/auth_v2/config.py @@ -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 diff --git a/litellm/auth_v2/models.py b/litellm/auth_v2/models.py index d23306abf83..3c6ba79de13 100644 --- a/litellm/auth_v2/models.py +++ b/litellm/auth_v2/models.py @@ -22,6 +22,7 @@ class AuthMethod(str, Enum): BEARER_JWT = "bearer_jwt" OAUTH2_INTROSPECTION = "oauth2_introspection" OIDC = "oidc" + SAML = "saml" MUTUAL_TLS = "mutual_tls" diff --git a/litellm/auth_v2/saml.py b/litellm/auth_v2/saml.py index 706880b7d2b..a36d6cc4103 100644 --- a/litellm/auth_v2/saml.py +++ b/litellm/auth_v2/saml.py @@ -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 diff --git a/litellm/auth_v2/security.py b/litellm/auth_v2/security.py index 0e016f3a5c2..5330509f617 100644 --- a/litellm/auth_v2/security.py +++ b/litellm/auth_v2/security.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 610305e1905..d388c29ab33 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/uv.lock b/uv.lock index e1b84f4e629..4f3e18d573b 100644 --- a/uv.lock +++ b/uv.lock @@ -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"