feat: service token introduction
This commit is contained in:
17
core/auth/fetch_sub.py
Normal file
17
core/auth/fetch_sub.py
Normal file
@@ -0,0 +1,17 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from db.models import User
|
||||
from repositories.users import UserRepository
|
||||
from schemas.jwt import ServiceJWTPayload
|
||||
|
||||
|
||||
async def fetch_subject_from_service(
|
||||
payload: ServiceJWTPayload, session: AsyncSession
|
||||
) -> User | None:
|
||||
repo = UserRepository(session)
|
||||
|
||||
if payload.acting_as.startswith("telegram:"):
|
||||
telegram_id = int(payload.acting_as.split("telegram:")[1])
|
||||
return await repo.get_user_by_telegram_id(telegram_id)
|
||||
|
||||
return
|
||||
@@ -4,17 +4,47 @@ from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from core.secrets import decode_jwt
|
||||
from core.auth.fetch_sub import fetch_subject_from_service
|
||||
from core.secrets import decode_jwt, decode_user_jwt, get_kid_from_token
|
||||
from repositories.service_signatures import get_active_signature_by_kid
|
||||
from repositories.users import UserRepository
|
||||
from schemas.dto import AuthContext
|
||||
from schemas.jwt import JWTPayload
|
||||
from schemas.jwt import ServiceJWTPayload, UserJWTPayload
|
||||
|
||||
|
||||
async def authorize_bot(kid: str, token: str, session: AsyncSession) -> AuthContext:
|
||||
signature = await get_active_signature_by_kid(session, kid)
|
||||
|
||||
if not signature:
|
||||
raise HTTPException(401, detail="Invalid service signature")
|
||||
|
||||
content = decode_jwt(token, signature.public_key, algo="EdDSA")
|
||||
|
||||
try:
|
||||
payload = ServiceJWTPayload.model_validate(content)
|
||||
except ValidationError:
|
||||
raise HTTPException(status_code=401, detail="Invalid credentials") from None
|
||||
|
||||
if payload.exp < datetime.now(UTC).timestamp():
|
||||
raise HTTPException(status_code=401, detail="Access token expired")
|
||||
|
||||
subject = await fetch_subject_from_service(payload, session)
|
||||
if not subject:
|
||||
raise HTTPException(status_code=401, detail="User not found")
|
||||
|
||||
return AuthContext(subject, auth_method="service", service=kid)
|
||||
|
||||
|
||||
async def authorize(token: str, session: AsyncSession, service: str | None = None) -> AuthContext:
|
||||
content = decode_jwt(token)
|
||||
kid = get_kid_from_token(token)
|
||||
|
||||
if kid:
|
||||
return await authorize_bot(kid, token, session)
|
||||
|
||||
content = decode_user_jwt(token)
|
||||
|
||||
try:
|
||||
payload = JWTPayload.model_validate(content)
|
||||
payload = UserJWTPayload.model_validate(content)
|
||||
except ValidationError:
|
||||
raise HTTPException(status_code=401, detail="Invalid credentials") from None
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import hashlib
|
||||
import logging
|
||||
import secrets
|
||||
from typing import Any
|
||||
from typing import Any, Literal
|
||||
|
||||
import jwt
|
||||
from argon2 import PasswordHasher
|
||||
@@ -10,7 +10,7 @@ from zxcvbn import zxcvbn
|
||||
|
||||
from config import cfg
|
||||
from schemas.dto import KeyPair
|
||||
from schemas.jwt import JWTPayload
|
||||
from schemas.jwt import UserJWTPayload
|
||||
from schemas.providers import ProvidersType
|
||||
|
||||
ctx = PasswordHasher()
|
||||
@@ -36,15 +36,25 @@ def generate_jwt(payload: dict[str, Any]) -> str:
|
||||
return jwt.encode(payload, cfg.private_key, "EdDSA")
|
||||
|
||||
|
||||
def decode_jwt(token: str) -> dict[str, Any] | None:
|
||||
def get_kid_from_token(token: str) -> str | None:
|
||||
return jwt.get_unverified_header(token).get("kid")
|
||||
|
||||
|
||||
def decode_jwt(
|
||||
token: str, public_key: str, algo: Literal["EdDSA"] = "EdDSA"
|
||||
) -> dict[str, Any] | None:
|
||||
try:
|
||||
return jwt.decode(token, cfg.public_key, "EdDSA")
|
||||
return jwt.decode(token, public_key, algo)
|
||||
except jwt.ExpiredSignatureError:
|
||||
return
|
||||
|
||||
|
||||
def decode_user_jwt(token: str) -> dict[str, Any] | None:
|
||||
return decode_jwt(token, cfg.public_key, algo="EdDSA")
|
||||
|
||||
|
||||
def generate_pair(user_id: int, iss: ProvidersType) -> KeyPair:
|
||||
payload = JWTPayload(sub=str(user_id), iss=iss)
|
||||
payload = UserJWTPayload(sub=str(user_id), iss=iss)
|
||||
access_token = generate_jwt(payload.model_dump())
|
||||
refresh_token = secrets.token_urlsafe(32)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user