diff --git a/litellm/proxy/auth_v2/__init__.py b/litellm/proxy/auth_v2/__init__.py index 032bb014645..2bd9527afc7 100644 --- a/litellm/proxy/auth_v2/__init__.py +++ b/litellm/proxy/auth_v2/__init__.py @@ -10,9 +10,11 @@ from .config import ( TrustedProxyConfig, ) from .models import Principal +from .oidc import build_oidc_router from .rbac import Role from .resolver import IdentityResolver, InMemoryIdentityStore, ProvisioningStore from .saml import build_saml_router +from .scim import build_scim_router from .security import AuthSecurity __all__ = [ @@ -32,4 +34,6 @@ __all__ = [ "SessionConfig", "SAMLConfig", "build_saml_router", + "build_scim_router", + "build_oidc_router", ] diff --git a/litellm/proxy/auth_v2/oidc.py b/litellm/proxy/auth_v2/oidc.py index e463981f6e6..28712597846 100644 --- a/litellm/proxy/auth_v2/oidc.py +++ b/litellm/proxy/auth_v2/oidc.py @@ -1,18 +1,24 @@ from __future__ import annotations import re -from typing import Any, Dict +from typing import TYPE_CHECKING, Any, Dict, cast from authlib.integrations.starlette_client import OAuth from fastapi import APIRouter, HTTPException, Request -from fastapi.responses import JSONResponse +from fastapi.responses import RedirectResponse from scim2_models import User as ScimUser -from .config import AuthConfig, OidcProviderConfig +from .config import OIDCProviderConfig from .resolver import ProvisioningStore +from .session import safe_relay_state + +if TYPE_CHECKING: + from .security import AuthSecurity + +_CLAIM_KEYS = ("email", "preferred_username", "name", "groups", "roles") -def _provider_key(provider: OidcProviderConfig) -> str: +def _provider_key(provider: OIDCProviderConfig) -> str: return re.sub(r"[^a-z0-9]+", "-", provider.issuer.lower()).strip("-") @@ -24,9 +30,15 @@ def _user_from_userinfo(userinfo: Dict[str, Any]) -> ScimUser: ) -def build_oidc_router(config: AuthConfig) -> APIRouter: +def _mapped_claims(userinfo: Dict[str, Any]) -> Dict[str, Any]: + return {key: userinfo[key] for key in _CLAIM_KEYS if userinfo.get(key) is not None} + + +def build_oidc_router(auth: AuthSecurity) -> APIRouter: + session = auth.config.session + issuers = {_provider_key(p): p.issuer for p in auth.config.oidc_providers} oauth = OAuth() - for provider in config.oidc_providers: + for provider in auth.config.oidc_providers: oauth.register( name=_provider_key(provider), server_metadata_url=f"{provider.issuer.rstrip('/')}/.well-known/openid-configuration", @@ -36,30 +48,98 @@ def build_oidc_router(config: AuthConfig) -> APIRouter: if provider.client_secret else None ), - client_kwargs={"scope": " ".join(provider.login_scopes)}, + client_kwargs={ + "scope": " ".join(provider.login_scopes), + "code_challenge_method": "S256", + }, ) router = APIRouter(prefix="/auth/oidc", tags=["oidc"]) @router.get("/{provider}/login") - async def login(provider: str, request: Request) -> Any: + async def login(provider: str, request: Request) -> RedirectResponse: client = oauth.create_client(provider) if client is None: raise HTTPException(status_code=404, detail="unknown provider") - redirect_uri = request.url_for("oidc_callback", provider=provider) - return await client.authorize_redirect(request, str(redirect_uri)) + redirect_uri = str(request.url_for("oidc_callback", provider=provider)) + relay = safe_relay_state( + request.query_params.get("next"), session.default_redirect_path + ) + authorization = await client.create_authorization_url(redirect_uri) + txn_id = auth.oauth_txn_store.create_session( + { + "provider": provider, + "state": authorization["state"], + "nonce": authorization.get("nonce"), + "code_verifier": authorization.get("code_verifier"), + "redirect_uri": redirect_uri, + "relay": relay, + } + ) + response = RedirectResponse(authorization["url"], status_code=303) + response.set_cookie( + session.login_cookie, + txn_id, + httponly=True, + samesite="lax", + secure=session.secure, + max_age=session.login_state_ttl, + ) + return response @router.get("/{provider}/callback", name="oidc_callback") - async def callback(provider: str, request: Request) -> JSONResponse: + async def callback(provider: str, request: Request) -> RedirectResponse: client = oauth.create_client(provider) if client is None: raise HTTPException(status_code=404, detail="unknown provider") - token = await client.authorize_access_token(request) - userinfo = token.get("userinfo") - if userinfo is None: + txn_id = request.cookies.get(session.login_cookie) + txn = auth.oauth_txn_store.pop(txn_id) if txn_id else None + if txn is None or txn.get("provider") != provider: + raise HTTPException( + status_code=400, detail="invalid or expired login state" + ) + returned_state = request.query_params.get("state") + if not returned_state or returned_state != txn["state"]: + raise HTTPException(status_code=400, detail="state mismatch") + error = request.query_params.get("error") + if error: + raise HTTPException(status_code=400, detail=error) + code = request.query_params.get("code") + if not code: + raise HTTPException(status_code=400, detail="missing authorization code") + + token = await client.fetch_access_token( + redirect_uri=txn["redirect_uri"], + code=code, + code_verifier=txn.get("code_verifier"), + state=txn["state"], + ) + if token.get("id_token"): + userinfo = await client.parse_id_token(token, nonce=txn.get("nonce")) + else: userinfo = await client.userinfo(token=token) - store: ProvisioningStore = request.app.state.auth_v2.resolver - stored = await store.upsert_user(_user_from_userinfo(dict(userinfo))) - return JSONResponse(content=stored.model_dump()) + info = dict(userinfo) + + store = cast(ProvisioningStore, auth.resolver) + await store.upsert_user(_user_from_userinfo(info)) + + session_id = auth.session_store.create_session( + { + "method": "oidc", + "subject": info.get("sub"), + "issuer": info.get("iss") or issuers.get(provider), + "claims": _mapped_claims(info), + } + ) + target = safe_relay_state(txn.get("relay"), session.default_redirect_path) + response = RedirectResponse(target, status_code=303) + response.set_cookie( + session.cookie, + session_id, + httponly=True, + samesite="lax", + secure=session.secure, + ) + return response return router diff --git a/litellm/proxy/auth_v2/scim.py b/litellm/proxy/auth_v2/scim.py index eafe83f8878..64255db266a 100644 --- a/litellm/proxy/auth_v2/scim.py +++ b/litellm/proxy/auth_v2/scim.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Dict, Optional, Type, TypeVar +from typing import TYPE_CHECKING, Any, Dict, Optional, Type, TypeVar, cast from fastapi import APIRouter, Query, Request, Response, Security, status from fastapi.responses import JSONResponse @@ -24,15 +24,13 @@ from scim2_models import ( ) from .resolver import ProvisioningStore -from .security import get_current_principal + +if TYPE_CHECKING: + from .security import AuthSecurity R = TypeVar("R", bound=Resource) -def _store(request: Request) -> ProvisioningStore: - return request.app.state.auth_v2.resolver - - def _error(status_code: int, detail: str) -> JSONResponse: return JSONResponse( status_code=status_code, @@ -89,6 +87,119 @@ def _dump(resource: Resource, ctx: Context) -> Dict[str, Any]: return resource.model_dump(scim_ctx=ctx) +def _build_protected_router(auth: AuthSecurity) -> APIRouter: + store = cast(ProvisioningStore, auth.resolver) + protected = APIRouter( + dependencies=[Security(auth.principal, scopes=["scim:write"])], + ) + + @protected.post("/Users", status_code=status.HTTP_201_CREATED) + async def create_user(request: Request) -> Response: + try: + user = await _parse(request, User) + except ValidationError as exc: + return _error(status.HTTP_400_BAD_REQUEST, str(exc)) + stored = await store.upsert_user(user) + return JSONResponse( + status_code=status.HTTP_201_CREATED, + content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE), + ) + + @protected.get("/Users/{resource_id}") + async def get_user(resource_id: str) -> Response: + user = await store.get_user(resource_id) + if user is None: + return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found") + return JSONResponse(content=_dump(user, Context.RESOURCE_QUERY_RESPONSE)) + + @protected.patch("/Users/{resource_id}") + async def patch_user(resource_id: str, request: Request) -> Response: + user = await store.get_user(resource_id) + if user is None: + return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found") + try: + patch = PatchOp[User].model_validate(await request.json()) + patched = _apply_patch(user, patch) + except (ValidationError, ValueError) as exc: + return _error(status.HTTP_400_BAD_REQUEST, str(exc)) + updated = await store.upsert_user(patched) + return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE)) + + @protected.delete("/Users/{resource_id}", status_code=status.HTTP_204_NO_CONTENT) + async def deactivate_user(resource_id: str) -> Response: + if await store.get_user(resource_id) is None: + return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found") + await store.deactivate_user(resource_id) + return Response(status_code=status.HTTP_204_NO_CONTENT) + + @protected.get("/Users") + async def list_users( + filter_expr: Optional[str] = Query(default=None, alias="filter"), + ) -> Response: + users = await store.list_users(filter_expr) + listing: ListResponse[User] = ListResponse[User]( + total_results=len(users), + start_index=1, + items_per_page=len(users), + resources=users or None, + ) + return JSONResponse(content=_dump(listing, Context.RESOURCE_QUERY_RESPONSE)) + + @protected.post("/Groups", status_code=status.HTTP_201_CREATED) + async def create_group(request: Request) -> Response: + try: + group = await _parse(request, Group) + except ValidationError as exc: + return _error(status.HTTP_400_BAD_REQUEST, str(exc)) + stored = await store.upsert_group(group) + return JSONResponse( + status_code=status.HTTP_201_CREATED, + content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE), + ) + + @protected.get("/Groups/{resource_id}") + async def get_group(resource_id: str) -> Response: + group = await store.get_group(resource_id) + if group is None: + return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found") + return JSONResponse(content=_dump(group, Context.RESOURCE_QUERY_RESPONSE)) + + @protected.patch("/Groups/{resource_id}") + async def patch_group(resource_id: str, request: Request) -> Response: + group = await store.get_group(resource_id) + if group is None: + return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found") + try: + patch = PatchOp[Group].model_validate(await request.json()) + patched = _apply_patch(group, patch) + except (ValidationError, ValueError) as exc: + return _error(status.HTTP_400_BAD_REQUEST, str(exc)) + updated = await store.upsert_group(patched) + return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE)) + + @protected.delete("/Groups/{resource_id}", status_code=status.HTTP_204_NO_CONTENT) + async def delete_group(resource_id: str) -> Response: + if await store.get_group(resource_id) is None: + return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found") + await store.delete_group(resource_id) + return Response(status_code=status.HTTP_204_NO_CONTENT) + + @protected.get("/Groups") + async def list_groups( + filter_expr: Optional[str] = Query(default=None, alias="filter"), + ) -> Response: + groups = await store.list_groups(filter_expr) + listing: ListResponse[Group] = ListResponse[Group]( + total_results=len(groups), + start_index=1, + items_per_page=len(groups), + resources=groups or None, + ) + return JSONResponse(content=_dump(listing, Context.RESOURCE_QUERY_RESPONSE)) + + return protected + + def _build_discovery_router() -> APIRouter: router = APIRouter() @@ -143,126 +254,8 @@ def _build_discovery_router() -> APIRouter: return router -def _build_protected_router() -> APIRouter: - protected = APIRouter( - dependencies=[Security(get_current_principal, scopes=["scim:write"])], - ) - - @protected.post("/Users", status_code=status.HTTP_201_CREATED) - async def create_user(request: Request) -> Response: - try: - user = await _parse(request, User) - except ValidationError as exc: - return _error(status.HTTP_400_BAD_REQUEST, str(exc)) - stored = await _store(request).upsert_user(user) - return JSONResponse( - status_code=status.HTTP_201_CREATED, - content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE), - ) - - @protected.get("/Users/{resource_id}") - async def get_user(resource_id: str, request: Request) -> Response: - user = await _store(request).get_user(resource_id) - if user is None: - return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found") - return JSONResponse(content=_dump(user, Context.RESOURCE_QUERY_RESPONSE)) - - @protected.patch("/Users/{resource_id}") - async def patch_user(resource_id: str, request: Request) -> Response: - store = _store(request) - user = await store.get_user(resource_id) - if user is None: - return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found") - try: - patch = PatchOp[User].model_validate(await request.json()) - patched = _apply_patch(user, patch) - except (ValidationError, ValueError) as exc: - return _error(status.HTTP_400_BAD_REQUEST, str(exc)) - updated = await store.upsert_user(patched) - return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE)) - - @protected.delete("/Users/{resource_id}", status_code=status.HTTP_204_NO_CONTENT) - async def deactivate_user(resource_id: str, request: Request) -> Response: - store = _store(request) - if await store.get_user(resource_id) is None: - return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found") - await store.deactivate_user(resource_id) - return Response(status_code=status.HTTP_204_NO_CONTENT) - - @protected.get("/Users") - async def list_users( - request: Request, - filter_expr: Optional[str] = Query(default=None, alias="filter"), - ) -> Response: - users = await _store(request).list_users(filter_expr) - listing: ListResponse[User] = ListResponse[User]( - total_results=len(users), - start_index=1, - items_per_page=len(users), - resources=users or None, - ) - return JSONResponse(content=_dump(listing, Context.RESOURCE_QUERY_RESPONSE)) - - @protected.post("/Groups", status_code=status.HTTP_201_CREATED) - async def create_group(request: Request) -> Response: - try: - group = await _parse(request, Group) - except ValidationError as exc: - return _error(status.HTTP_400_BAD_REQUEST, str(exc)) - stored = await _store(request).upsert_group(group) - return JSONResponse( - status_code=status.HTTP_201_CREATED, - content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE), - ) - - @protected.get("/Groups/{resource_id}") - async def get_group(resource_id: str, request: Request) -> Response: - group = await _store(request).get_group(resource_id) - if group is None: - return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found") - return JSONResponse(content=_dump(group, Context.RESOURCE_QUERY_RESPONSE)) - - @protected.patch("/Groups/{resource_id}") - async def patch_group(resource_id: str, request: Request) -> Response: - store = _store(request) - group = await store.get_group(resource_id) - if group is None: - return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found") - try: - patch = PatchOp[Group].model_validate(await request.json()) - patched = _apply_patch(group, patch) - except (ValidationError, ValueError) as exc: - return _error(status.HTTP_400_BAD_REQUEST, str(exc)) - updated = await store.upsert_group(patched) - return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE)) - - @protected.delete("/Groups/{resource_id}", status_code=status.HTTP_204_NO_CONTENT) - async def delete_group(resource_id: str, request: Request) -> Response: - store = _store(request) - if await store.get_group(resource_id) is None: - return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found") - await store.delete_group(resource_id) - return Response(status_code=status.HTTP_204_NO_CONTENT) - - @protected.get("/Groups") - async def list_groups( - request: Request, - filter_expr: Optional[str] = Query(default=None, alias="filter"), - ) -> Response: - groups = await _store(request).list_groups(filter_expr) - listing: ListResponse[Group] = ListResponse[Group]( - total_results=len(groups), - start_index=1, - items_per_page=len(groups), - resources=groups or None, - ) - return JSONResponse(content=_dump(listing, Context.RESOURCE_QUERY_RESPONSE)) - - return protected - - -def build_scim_router() -> APIRouter: +def build_scim_router(auth: AuthSecurity) -> APIRouter: router = APIRouter(prefix="/scim/v2", tags=["scim"]) - router.include_router(_build_protected_router()) + router.include_router(_build_protected_router(auth)) router.include_router(_build_discovery_router()) return router