auth_middleware.py 2.6 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980
  1. from __future__ import annotations
  2. import re
  3. from collections.abc import Awaitable, Callable
  4. from fastapi import Request
  5. from starlette.concurrency import run_in_threadpool
  6. from starlette.middleware.base import BaseHTTPMiddleware
  7. from starlette.responses import JSONResponse, Response
  8. from api.services.auth import ADMIN_ROLE, SESSION_COOKIE_NAME, resolve_session
  9. _PUBLIC_PATHS = {
  10. "/api/auth/login",
  11. }
  12. _AUTHENTICATED_USER_PATHS = {
  13. ("GET", "/api/auth/me"),
  14. ("POST", "/api/auth/logout"),
  15. }
  16. _NORMAL_USER_PATHS = {
  17. ("GET", "/api/category-tree"),
  18. ("GET", "/api/demand-grade"),
  19. ("GET", "/api/video-discovery/demands"),
  20. ("GET", "/api/video-discovery/runs"),
  21. ("GET", "/api/video-discovery/feedback"),
  22. ("POST", "/api/video-discovery/feedback"),
  23. }
  24. _NORMAL_USER_PATTERNS = (
  25. re.compile(r"^/api/video-discovery/demands/\d+$"),
  26. re.compile(r"^/api/video-discovery/runs/[^/]+$"),
  27. re.compile(r"^/api/video-discovery/runs/[^/]+/searches$"),
  28. re.compile(r"^/api/video-discovery/runs/[^/]+/candidates$"),
  29. re.compile(r"^/api/demand-grade/\d+/videos$"),
  30. )
  31. def normal_user_can_access(method: str, path: str) -> bool:
  32. if (method, path) in _AUTHENTICATED_USER_PATHS:
  33. return True
  34. if (method, path) in _NORMAL_USER_PATHS:
  35. return True
  36. return method == "GET" and any(pattern.fullmatch(path) for pattern in _NORMAL_USER_PATTERNS)
  37. def _is_protected_path(path: str) -> bool:
  38. return (
  39. path == "/api"
  40. or path.startswith("/api/")
  41. or path in {"/docs", "/redoc", "/openapi.json"}
  42. )
  43. class AuthenticationMiddleware(BaseHTTPMiddleware):
  44. """Authenticate API requests and enforce the two fixed application roles."""
  45. async def dispatch(
  46. self,
  47. request: Request,
  48. call_next: Callable[[Request], Awaitable[Response]],
  49. ) -> Response:
  50. path = request.url.path
  51. method = request.method.upper()
  52. if method == "OPTIONS" or not _is_protected_path(path) or path in _PUBLIC_PATHS:
  53. return await call_next(request)
  54. token = request.cookies.get(SESSION_COOKIE_NAME)
  55. user = await run_in_threadpool(resolve_session, token) if token else None
  56. if user is None:
  57. return JSONResponse(
  58. status_code=401,
  59. content={"detail": "authentication required"},
  60. )
  61. request.state.current_user = user
  62. if user["role"] == ADMIN_ROLE or normal_user_can_access(method, path):
  63. return await call_next(request)
  64. return JSONResponse(
  65. status_code=403,
  66. content={"detail": "permission denied"},
  67. )