feat: simplify auth flow - replace JWT with session cookies
Replace complex JWT + refresh token authentication with simple session-based auth using signed cookies. **Removed:** - JWT token service (jwt_service.py) - Refresh token store (refresh_store.py) - Refresh token model and database table - JWKS fetching and OIDC token verification - python-jose dependency **Added:** - Session service (session.py) with HMAC-SHA256 signed cookies - Auth dependencies module for shared auth logic - Session-based auth endpoints **Updated:** - All API endpoints to use session-based auth - Config: removed JWT settings, added SESSION_SECRET/SESSION_TTL_HOURS - Tests: rewritten for session-based flow - Frontend: no changes needed (already uses cookies) Quality gates: ruff ✓, mypy ✓, typecheck ✓, lint ✓
This commit is contained in:
@@ -1,12 +1,10 @@
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import build_login_redirect_url
|
||||
from src.auth.refresh_store import hash_refresh_token
|
||||
from src.auth.session import create_session_cookie, decode_session_cookie
|
||||
|
||||
__all__ = [
|
||||
"build_cookie_options",
|
||||
"build_login_redirect_url",
|
||||
"decode_access_token",
|
||||
"hash_refresh_token",
|
||||
"mint_access_token",
|
||||
"create_session_cookie",
|
||||
"decode_session_cookie",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import Cookie, Depends, HTTPException, status
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.session import decode_session_cookie
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
session_cookie: Annotated[str | None, Cookie()] = None,
|
||||
) -> uuid.UUID:
|
||||
if not session_cookie:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
return uuid.UUID(str(payload["user_id"]))
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid session")
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
session_cookie: Annotated[str | None, Cookie()] = None,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> User:
|
||||
if not session_cookie:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
user_id = uuid.UUID(str(payload["user_id"]))
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid session")
|
||||
|
||||
user = await db_session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
@@ -1,27 +0,0 @@
|
||||
from datetime import datetime
|
||||
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def mint_access_token(
|
||||
*,
|
||||
settings: Settings,
|
||||
subject: str,
|
||||
email: str,
|
||||
name: str,
|
||||
expires_at: datetime,
|
||||
) -> str:
|
||||
payload = {
|
||||
"sub": subject,
|
||||
"email": email,
|
||||
"name": name,
|
||||
"exp": expires_at,
|
||||
}
|
||||
return jwt.encode(payload, settings.jwt_secret, algorithm=settings.jwt_algorithm)
|
||||
|
||||
|
||||
def decode_access_token(*, settings: Settings, token: str) -> dict[str, str | int]:
|
||||
claims = jwt.decode(token, settings.jwt_secret, algorithms=[settings.jwt_algorithm])
|
||||
return dict(claims)
|
||||
+11
-24
@@ -1,7 +1,7 @@
|
||||
from typing import Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
@@ -11,7 +11,6 @@ def build_login_redirect_url(
|
||||
settings: Settings,
|
||||
redirect_uri: str,
|
||||
state: str,
|
||||
nonce: str,
|
||||
) -> str:
|
||||
query = urlencode(
|
||||
{
|
||||
@@ -20,7 +19,6 @@ def build_login_redirect_url(
|
||||
"redirect_uri": redirect_uri,
|
||||
"scope": "openid profile email",
|
||||
"state": state,
|
||||
"nonce": nonce,
|
||||
}
|
||||
)
|
||||
return f"{settings.resolved_authentik_authorize_url}?{query}"
|
||||
@@ -51,27 +49,16 @@ async def exchange_code_for_tokens(
|
||||
}
|
||||
|
||||
|
||||
async def fetch_jwks(*, settings: Settings, client: httpx.AsyncClient) -> dict[str, list[dict[str, str]]]:
|
||||
response = await client.get(settings.resolved_authentik_jwks_url)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
return {"keys": payload["keys"]}
|
||||
|
||||
|
||||
def verify_provider_access_token(
|
||||
async def fetch_user_info(
|
||||
*,
|
||||
settings: Settings,
|
||||
token: str,
|
||||
jwks: dict[str, list[dict[str, str]]],
|
||||
) -> dict[str, str | int]:
|
||||
unverified_header = jwt.get_unverified_header(token)
|
||||
key_id = unverified_header["kid"]
|
||||
jwk_key = next(key for key in jwks["keys"] if key.get("kid") == key_id)
|
||||
claims = jwt.decode(
|
||||
token,
|
||||
jwk_key,
|
||||
algorithms=[jwk_key.get("alg", "HS256")],
|
||||
audience=settings.authentik_audience,
|
||||
issuer=settings.resolved_authentik_issuer,
|
||||
access_token: str,
|
||||
client: httpx.AsyncClient,
|
||||
) -> dict[str, Any]:
|
||||
"""Fetch user info from Authentik userinfo endpoint."""
|
||||
response = await client.get(
|
||||
f"{settings.authentik_base_url}/application/o/userinfo/",
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
)
|
||||
return dict(claims)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
@@ -1,79 +0,0 @@
|
||||
from datetime import UTC, datetime
|
||||
from hashlib import sha256
|
||||
from secrets import token_urlsafe
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.models.refresh_token import RefreshToken
|
||||
|
||||
|
||||
def hash_refresh_token(raw_token: str) -> str:
|
||||
return sha256(raw_token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
async def create_refresh_token(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
user_id: object,
|
||||
expires_at: datetime,
|
||||
user_agent: str | None,
|
||||
ip_address: str | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
raw_token = token_urlsafe(48)
|
||||
record = RefreshToken(
|
||||
user_id=user_id,
|
||||
token_hash=hash_refresh_token(raw_token),
|
||||
expires_at=expires_at,
|
||||
created_at=datetime.now(UTC),
|
||||
user_agent=user_agent,
|
||||
ip_address=ip_address,
|
||||
)
|
||||
session.add(record)
|
||||
await session.commit()
|
||||
await session.refresh(record)
|
||||
return raw_token, record
|
||||
|
||||
|
||||
async def rotate_refresh_token(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
raw_token: str,
|
||||
user_agent: str | None,
|
||||
ip_address: str | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
existing_hash = hash_refresh_token(raw_token)
|
||||
existing = await session.scalar(
|
||||
select(RefreshToken).where(
|
||||
RefreshToken.token_hash == existing_hash,
|
||||
RefreshToken.revoked_at.is_(None),
|
||||
)
|
||||
)
|
||||
if existing is None:
|
||||
raise ValueError("refresh token not found")
|
||||
if existing.expires_at <= datetime.now(UTC):
|
||||
raise ValueError("refresh token expired")
|
||||
|
||||
existing.revoked_at = datetime.now(UTC)
|
||||
await session.flush()
|
||||
|
||||
return await create_refresh_token(
|
||||
session=session,
|
||||
user_id=existing.user_id,
|
||||
expires_at=existing.expires_at,
|
||||
user_agent=user_agent,
|
||||
ip_address=ip_address,
|
||||
)
|
||||
|
||||
|
||||
async def revoke_refresh_token(*, session: AsyncSession, raw_token: str) -> bool:
|
||||
token_hash = hash_refresh_token(raw_token)
|
||||
existing = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
if existing is None:
|
||||
return False
|
||||
if existing.revoked_at is not None:
|
||||
return True
|
||||
|
||||
existing.revoked_at = datetime.now(UTC)
|
||||
await session.commit()
|
||||
return True
|
||||
@@ -0,0 +1,71 @@
|
||||
import hmac
|
||||
import hashlib
|
||||
import json
|
||||
import base64
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def _base64url_encode(data: bytes) -> str:
|
||||
return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
||||
def _base64url_decode(data: str) -> bytes:
|
||||
padding = 4 - len(data) % 4
|
||||
if padding != 4:
|
||||
data += "=" * padding
|
||||
return base64.urlsafe_b64decode(data)
|
||||
|
||||
|
||||
def create_session_cookie(*, settings: Settings, user_id: str) -> str:
|
||||
"""Create a signed session cookie value."""
|
||||
payload = {
|
||||
"user_id": user_id,
|
||||
"exp": int((datetime.now(UTC) + timedelta(hours=settings.session_ttl_hours)).timestamp()),
|
||||
}
|
||||
|
||||
header = _base64url_encode(json.dumps({"alg": "HS256", "typ": "session"}).encode())
|
||||
payload_encoded = _base64url_encode(json.dumps(payload).encode())
|
||||
message = f"{header}.{payload_encoded}"
|
||||
|
||||
signature = hmac.new(
|
||||
settings.session_secret.encode(),
|
||||
message.encode(),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
signature_encoded = _base64url_encode(signature)
|
||||
|
||||
return f"{message}.{signature_encoded}"
|
||||
|
||||
|
||||
def decode_session_cookie(*, settings: Settings, cookie_value: str) -> dict[str, Any]:
|
||||
"""Decode and verify a session cookie. Returns payload or raises ValueError."""
|
||||
parts = cookie_value.split(".")
|
||||
if len(parts) != 3:
|
||||
raise ValueError("invalid session format")
|
||||
|
||||
header, payload_encoded, signature_encoded = parts
|
||||
message = f"{header}.{payload_encoded}"
|
||||
|
||||
# Verify signature
|
||||
expected_signature = hmac.new(
|
||||
settings.session_secret.encode(),
|
||||
message.encode(),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
expected_signature_encoded = _base64url_encode(expected_signature)
|
||||
|
||||
if not hmac.compare_digest(signature_encoded, expected_signature_encoded):
|
||||
raise ValueError("invalid session signature")
|
||||
|
||||
# Decode payload
|
||||
payload_bytes = _base64url_decode(payload_encoded)
|
||||
payload = json.loads(payload_bytes)
|
||||
|
||||
# Check expiry
|
||||
if payload.get("exp", 0) < int(datetime.now(UTC).timestamp()):
|
||||
raise ValueError("session expired")
|
||||
|
||||
return payload
|
||||
Reference in New Issue
Block a user