auth_middleware.py 2.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081
  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/growth-category-tree"),
  19. ("GET", "/api/demand-grade"),
  20. ("GET", "/api/video-discovery/demands"),
  21. ("GET", "/api/video-discovery/runs"),
  22. ("GET", "/api/video-discovery/feedback"),
  23. ("POST", "/api/video-discovery/feedback"),
  24. }
  25. _NORMAL_USER_PATTERNS = (
  26. re.compile(r"^/api/video-discovery/demands/\d+$"),
  27. re.compile(r"^/api/video-discovery/runs/[^/]+$"),
  28. re.compile(r"^/api/video-discovery/runs/[^/]+/searches$"),
  29. re.compile(r"^/api/video-discovery/runs/[^/]+/candidates$"),
  30. re.compile(r"^/api/demand-grade/\d+/videos$"),
  31. )
  32. def normal_user_can_access(method: str, path: str) -> bool:
  33. if (method, path) in _AUTHENTICATED_USER_PATHS:
  34. return True
  35. if (method, path) in _NORMAL_USER_PATHS:
  36. return True
  37. return method == "GET" and any(pattern.fullmatch(path) for pattern in _NORMAL_USER_PATTERNS)
  38. def _is_protected_path(path: str) -> bool:
  39. return (
  40. path == "/api"
  41. or path.startswith("/api/")
  42. or path in {"/docs", "/redoc", "/openapi.json"}
  43. )
  44. class AuthenticationMiddleware(BaseHTTPMiddleware):
  45. """Authenticate API requests and enforce the two fixed application roles."""
  46. async def dispatch(
  47. self,
  48. request: Request,
  49. call_next: Callable[[Request], Awaitable[Response]],
  50. ) -> Response:
  51. path = request.url.path
  52. method = request.method.upper()
  53. if method == "OPTIONS" or not _is_protected_path(path) or path in _PUBLIC_PATHS:
  54. return await call_next(request)
  55. token = request.cookies.get(SESSION_COOKIE_NAME)
  56. user = await run_in_threadpool(resolve_session, token) if token else None
  57. if user is None:
  58. return JSONResponse(
  59. status_code=401,
  60. content={"detail": "authentication required"},
  61. )
  62. request.state.current_user = user
  63. if user["role"] == ADMIN_ROLE or normal_user_can_access(method, path):
  64. return await call_next(request)
  65. return JSONResponse(
  66. status_code=403,
  67. content={"detail": "permission denied"},
  68. )