Files
dodo/backend/auth.py
T

86 lines
2.9 KiB
Python

import hashlib
import secrets
from datetime import UTC, datetime, timedelta
from argon2 import PasswordHasher
from argon2.exceptions import InvalidHashError, VerificationError
from fastapi import Cookie, Depends, HTTPException, Request, Response, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from .config import get_settings
from .db import get_db
from .models import Session, User
password_hasher = PasswordHasher()
COOKIE_NAME = "dodo_session"
CSRF_COOKIE_NAME = "dodo_csrf"
def hash_token(token: str) -> str:
return hashlib.sha256(token.encode()).hexdigest()
def hash_password(password: str) -> str:
return password_hasher.hash(password)
def verify_password(password_hash: str, password: str) -> bool:
try:
return password_hasher.verify(password_hash, password)
except (InvalidHashError, VerificationError):
return False
async def issue_session(db: AsyncSession, response: Response, user: User, request: Request | None = None) -> None:
token = secrets.token_urlsafe(32)
csrf = secrets.token_urlsafe(24)
expires = datetime.now(UTC) + timedelta(days=get_settings().session_days)
db.add(Session(
token_hash=hash_token(token), user_id=user.id, expires_at=expires,
ip_address=request.client.host if request and request.client else None,
user_agent=request.headers.get("user-agent", "")[:500] if request else None,
))
await db.commit()
response.set_cookie(
COOKIE_NAME, token, max_age=get_settings().session_days * 86400,
httponly=True, secure=get_settings().cookie_secure, samesite="lax", path="/",
)
response.set_cookie(
CSRF_COOKIE_NAME, csrf, max_age=get_settings().session_days * 86400,
httponly=False, secure=get_settings().cookie_secure, samesite="lax", path="/",
)
async def session_token(
token: str | None = Cookie(default=None, alias=COOKIE_NAME),
) -> str:
if not token:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录")
return token
async def current_user(
token: str = Depends(session_token),
db: AsyncSession = Depends(get_db),
) -> User:
result = await db.execute(
select(User).join(Session).where(
Session.token_hash == hash_token(token), Session.expires_at > datetime.now(UTC)
)
)
user = result.scalar_one_or_none()
if user is None:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="会话已失效,请重新登录")
return user
async def current_session(
token: str = Depends(session_token),
db: AsyncSession = Depends(get_db),
) -> Session:
row = await db.scalar(select(Session).where(Session.token_hash == hash_token(token), Session.expires_at > datetime.now(UTC)))
if row is None:
raise HTTPException(status_code=401, detail="会话已失效,请重新登录")
return row