194 lines
6.2 KiB
Python
194 lines
6.2 KiB
Python
from datetime import timedelta
|
|
from typing import Annotated
|
|
|
|
from fastapi import Depends, HTTPException, Request, Response, status
|
|
from sqlalchemy import select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.config import settings
|
|
from app.db import get_db
|
|
from app.models import Setting, User, UserSession
|
|
from app.security import (
|
|
hash_token,
|
|
new_token,
|
|
tokens_equal,
|
|
utcnow,
|
|
)
|
|
|
|
SESSION_COOKIE = "ea_session"
|
|
CSRF_COOKIE = "ea_csrf"
|
|
CSRF_HEADER = "X-CSRF-Token"
|
|
|
|
DbSession = Annotated[Session, Depends(get_db)]
|
|
|
|
|
|
# --------------------------------------------------------------------------
|
|
# Laufzeit-Einstellungen
|
|
# --------------------------------------------------------------------------
|
|
|
|
def get_setting(db: Session, key: str, default: str = "") -> str:
|
|
row = db.get(Setting, key)
|
|
return row.value if row else default
|
|
|
|
|
|
def set_setting(db: Session, key: str, value: str) -> None:
|
|
row = db.get(Setting, key)
|
|
if row is None:
|
|
db.add(Setting(key=key, value=value))
|
|
else:
|
|
row.value = value
|
|
|
|
|
|
def registration_locked_by_env() -> bool:
|
|
"""True, wenn ALLOW_SELF_REGISTRATION hart auf true/false steht.
|
|
Dann darf die Admin-Oberflaeche den Wert nicht aendern."""
|
|
return settings.allow_self_registration in ("true", "false")
|
|
|
|
|
|
def self_registration_enabled(db: Session) -> bool:
|
|
mode = settings.allow_self_registration
|
|
if mode == "true":
|
|
return True
|
|
if mode == "false":
|
|
return False
|
|
# mode == "admin": die Laufzeiteinstellung entscheidet.
|
|
return get_setting(db, "allow_self_registration", "true") == "true"
|
|
|
|
|
|
# --------------------------------------------------------------------------
|
|
# Sessions
|
|
# --------------------------------------------------------------------------
|
|
|
|
def create_session(db: Session, user: User, response: Response) -> UserSession:
|
|
# Spät importiert: runtime_settings greift auf get_setting in dieser
|
|
# Datei zu, ein Import auf Modulebene wäre ein Kreis.
|
|
from app.runtime_settings import get_duration
|
|
|
|
days = get_duration(db, "session_days")
|
|
|
|
raw = new_token()
|
|
csrf = new_token()
|
|
sess = UserSession(
|
|
token_hash=hash_token(raw),
|
|
user_id=user.id,
|
|
csrf_token=csrf,
|
|
expires_at=utcnow() + timedelta(days=days),
|
|
)
|
|
db.add(sess)
|
|
db.flush()
|
|
|
|
max_age = days * 24 * 3600
|
|
# HttpOnly: fuer JavaScript unsichtbar, damit ein XSS-Fund das Token
|
|
# nicht abgreifen kann.
|
|
response.set_cookie(
|
|
SESSION_COOKIE, raw, max_age=max_age, httponly=True,
|
|
secure=settings.cookie_secure, samesite="lax", path="/",
|
|
)
|
|
# Bewusst NICHT HttpOnly: der Client muss den Wert lesen und als
|
|
# Header zurueckschicken koennen (Double-Submit-Verfahren).
|
|
response.set_cookie(
|
|
CSRF_COOKIE, csrf, max_age=max_age, httponly=False,
|
|
secure=settings.cookie_secure, samesite="lax", path="/",
|
|
)
|
|
return sess
|
|
|
|
|
|
def destroy_session(db: Session, request: Request, response: Response) -> None:
|
|
raw = request.cookies.get(SESSION_COOKIE)
|
|
if raw:
|
|
sess = db.scalar(
|
|
select(UserSession).where(UserSession.token_hash == hash_token(raw))
|
|
)
|
|
if sess:
|
|
db.delete(sess)
|
|
response.delete_cookie(SESSION_COOKIE, path="/")
|
|
response.delete_cookie(CSRF_COOKIE, path="/")
|
|
|
|
|
|
def _load_session(db: Session, request: Request) -> UserSession | None:
|
|
raw = request.cookies.get(SESSION_COOKIE)
|
|
if not raw:
|
|
return None
|
|
sess = db.scalar(
|
|
select(UserSession).where(UserSession.token_hash == hash_token(raw))
|
|
)
|
|
if sess is None:
|
|
return None
|
|
if sess.expires_at <= utcnow():
|
|
db.delete(sess)
|
|
db.commit()
|
|
return None
|
|
return sess
|
|
|
|
|
|
# --------------------------------------------------------------------------
|
|
# Abhaengigkeiten fuer Routen
|
|
# --------------------------------------------------------------------------
|
|
|
|
_UNSAFE = {"POST", "PUT", "PATCH", "DELETE"}
|
|
|
|
|
|
def current_user(request: Request, db: DbSession) -> User:
|
|
sess = _load_session(db, request)
|
|
if sess is None:
|
|
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "Nicht angemeldet")
|
|
|
|
if request.method in _UNSAFE:
|
|
supplied = request.headers.get(CSRF_HEADER, "")
|
|
if not supplied or not tokens_equal(supplied, sess.csrf_token):
|
|
raise HTTPException(status.HTTP_403_FORBIDDEN, "CSRF-Token fehlt oder ungültig")
|
|
|
|
user = db.get(User, sess.user_id)
|
|
if user is None or not user.is_active:
|
|
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "Konto nicht aktiv")
|
|
|
|
# Nur einmal pro Stunde schreiben, sonst erzeugt jeder Request ein
|
|
# UPDATE. Der Zeitstempel am Konto ist die Grundlage der
|
|
# automatischen Deaktivierung - er muss auch dann mitlaufen, wenn
|
|
# sich jemand monatelang nicht neu anmeldet, weil die Sitzung hält.
|
|
if utcnow() - sess.last_seen_at > timedelta(hours=1):
|
|
sess.last_seen_at = utcnow()
|
|
user.last_seen_at = utcnow()
|
|
db.commit()
|
|
|
|
return user
|
|
|
|
|
|
CurrentUser = Annotated[User, Depends(current_user)]
|
|
|
|
|
|
def verified_user(user: CurrentUser) -> User:
|
|
if user.verified_at is None:
|
|
raise HTTPException(
|
|
status.HTTP_403_FORBIDDEN, "E-Mail-Adresse noch nicht bestätigt"
|
|
)
|
|
# Alles ausser GET /api/auth/me und POST /api/auth/password/change
|
|
# haengt an dieser Abhaengigkeit - der Zwang wirkt also flaechendeckend,
|
|
# ohne dass jede Route ihn einzeln pruefen muesste.
|
|
if user.must_change_password:
|
|
raise HTTPException(
|
|
status.HTTP_403_FORBIDDEN,
|
|
"Das Startpasswort muss zuerst geändert werden: "
|
|
"POST /api/auth/password/change",
|
|
)
|
|
return user
|
|
|
|
|
|
VerifiedUser = Annotated[User, Depends(verified_user)]
|
|
|
|
|
|
def admin_user(user: VerifiedUser) -> User:
|
|
if not user.is_admin:
|
|
raise HTTPException(status.HTTP_403_FORBIDDEN, "Administratorrechte erforderlich")
|
|
return user
|
|
|
|
|
|
AdminUser = Annotated[User, Depends(admin_user)]
|
|
|
|
|
|
def client_ip(request: Request) -> str:
|
|
"""Fuer Rate Limiting. Hinter einem Reverse Proxy liefert
|
|
request.client.host die Proxy-IP - ab Phase 3 setzen wir dafuer
|
|
ProxyHeadersMiddleware mit einer Liste vertrauenswuerdiger Hosts."""
|
|
return request.client.host if request.client else "unknown"
|