mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
fix(auth_v2): authenticate OIDC login sessions and adapt SCIM to AuthSecurity
The OIDC callback now mints a server-side session and sets the shared session cookie (httponly, secure, samesite=lax) before redirecting, so login yields an authenticated session that the shared SessionAuthenticator resolves; previously it returned the user record as JSON and left the caller unauthenticated. The login flow owns its CSRF protection without Starlette SessionMiddleware: it stores state, nonce and the PKCE (S256) verifier in the short-lived oauth_txn_store keyed by a temporary cookie, then on callback consumes the transaction one-time, checks the returned state, exchanges the code with the verifier and validates the nonce against the id_token. A replayed or expired state finds no transaction and returns 400. SCIM moves onto the AuthSecurity DI surface: build_scim_router(auth) reads auth.resolver and guards writes with Security(auth.principal, scopes= ["scim:write"]) instead of app.state and get_current_principal.
This commit is contained in:
parent
4ec48302a5
commit
5f02c88369
3 changed files with 220 additions and 143 deletions
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue