aboutsummaryrefslogtreecommitdiff
path: root/jb/api/auth.py
blob: 7d6c3523ccefc275b65cf3273a87c07cfd045f1f (plain)
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
import logging
from datetime import datetime, timedelta, timezone
from typing import Annotated
from uuid import uuid4

import jwt
from fastapi import Depends, HTTPException, Request, Response, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer

from jb.config import settings
from jb.models.auth import AuthenticatedUser

bearer = HTTPBearer(auto_error=False)
logger = logging.getLogger(__name__)

AUTH_CACHE_TTL_SECONDS = 60

SESSION_COOKIE_NAME = "jb_session"
JWT_ISSUER = "jamesbillings67"
JWT_AUDIENCE = "jamesbillings67"


def get_authenticated_user(
    request: Request,
    credentials: Annotated[HTTPAuthorizationCredentials | None, Depends(bearer)],
) -> AuthenticatedUser:
    """FastAPI dependency for endpoints requiring a valid session."""
    if settings.session_jwt_secret is None:
        raise HTTPException(
            status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
            detail="Account session signing key is not configured",
        )

    token = request.cookies.get(SESSION_COOKIE_NAME)
    if credentials is not None and credentials.scheme.lower() == "bearer":
        token = credentials.credentials

    if token is None:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="Missing session token",
            headers={"WWW-Authenticate": "Bearer"},
        )

    try:
        claims = jwt.decode(
            token,
            settings.session_jwt_secret.get_secret_value(),
            algorithms=["HS256"],
            issuer=JWT_ISSUER,
            audience=JWT_AUDIENCE,
            options={"require": ["sub", "iat", "exp", "jti", "iss", "aud", "type"]},
        )
    except jwt.PyJWTError:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="Invalid or expired session token",
            headers={"WWW-Authenticate": "Bearer"},
        )

    if claims["type"] != "session" or not isinstance(claims["sub"], str):
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="Invalid or expired session token",
            headers={"WWW-Authenticate": "Bearer"},
        )
    product_user_id = claims.get("sub")

    # todo: in here, hit THL by the user's bpuid (email hash),
    #  in order to 1) be sure user exists & 2) pull the display_name
    user_email = ...  # lookup in thl by product_user_id

    return AuthenticatedUser(email=user_email)


def create_session(product_user_id: str) -> str:
    now = datetime.now(timezone.utc)
    return jwt.encode(
        {
            "sub": product_user_id,
            "iat": now,
            "exp": now + timedelta(seconds=settings.session_token_ttl_seconds),
            "jti": uuid4().hex,
            "iss": JWT_ISSUER,
            "aud": JWT_AUDIENCE,
            "type": "session",
        },
        settings.session_jwt_secret.get_secret_value(),
        algorithm="HS256",
    )