File size: 3,313 Bytes
0dddf11
0d29f85
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
"""Supabase JWT authentication for the Lingo World API."""

from __future__ import annotations

import os
from dataclasses import dataclass
from functools import lru_cache

import jwt
from fastapi import HTTPException, Security, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer


_bearer = HTTPBearer(auto_error=False)


@dataclass(frozen=True)
class AuthenticatedUser:
    id: str
    email: str | None = None


class SupabaseJWTVerifier:
    def __init__(self) -> None:
        self.supabase_url = os.environ.get("SUPABASE_URL", "").rstrip("/")
        self.jwt_secret = os.environ.get("SUPABASE_JWT_SECRET", "")
        self.audience = os.environ.get("SUPABASE_JWT_AUDIENCE", "authenticated")
        self.issuer = (
            os.environ.get("SUPABASE_JWT_ISSUER")
            or (f"{self.supabase_url}/auth/v1" if self.supabase_url else "")
        )
        self._jwks_client = (
            jwt.PyJWKClient(f"{self.supabase_url}/auth/v1/.well-known/jwks.json")
            if self.supabase_url and not self.jwt_secret
            else None
        )

    @property
    def configured(self) -> bool:
        return bool(self.supabase_url and (self.jwt_secret or self._jwks_client))

    def verify(self, token: str) -> AuthenticatedUser:
        if not self.configured:
            raise HTTPException(
                status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
                detail="Authentication is not configured",
            )

        try:
            if self.jwt_secret:
                payload = jwt.decode(
                    token,
                    self.jwt_secret,
                    algorithms=["HS256"],
                    audience=self.audience,
                    issuer=self.issuer or None,
                )
            else:
                signing_key = self._jwks_client.get_signing_key_from_jwt(token)
                payload = jwt.decode(
                    token,
                    signing_key.key,
                    algorithms=["RS256", "ES256"],
                    audience=self.audience,
                    issuer=self.issuer or None,
                )
        except jwt.PyJWTError as exc:
            raise HTTPException(
                status_code=status.HTTP_401_UNAUTHORIZED,
                detail="Invalid or expired session",
                headers={"WWW-Authenticate": "Bearer"},
            ) from exc

        user_id = str(payload.get("sub") or "").strip()
        if not user_id:
            raise HTTPException(
                status_code=status.HTTP_401_UNAUTHORIZED,
                detail="Session is missing a user identifier",
                headers={"WWW-Authenticate": "Bearer"},
            )
        return AuthenticatedUser(id=user_id, email=payload.get("email"))


@lru_cache(maxsize=1)
def get_verifier() -> SupabaseJWTVerifier:
    return SupabaseJWTVerifier()


def require_user(
    credentials: HTTPAuthorizationCredentials | None = Security(_bearer),
) -> AuthenticatedUser:
    if credentials is None or credentials.scheme.lower() != "bearer":
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="Authentication required",
            headers={"WWW-Authenticate": "Bearer"},
        )
    return get_verifier().verify(credentials.credentials)