76 lines
2.7 KiB
Python
76 lines
2.7 KiB
Python
"""Server-side authentication and encryption utilities."""
|
|
|
|
import hashlib
|
|
import secrets
|
|
from datetime import timedelta
|
|
|
|
from argon2 import PasswordHasher
|
|
from argon2.exceptions import VerificationError
|
|
from cryptography.fernet import Fernet
|
|
from fastapi import Depends, HTTPException, Request
|
|
from fastapi.security import APIKeyCookie
|
|
from sqlalchemy import delete
|
|
|
|
from .models import Account, Admin, LoginSession, now
|
|
|
|
password_hasher = PasswordHasher()
|
|
session_cookie = APIKeyCookie(name="wq_session", auto_error=False)
|
|
|
|
|
|
def token_hash(token: str) -> str:
|
|
return hashlib.sha256(token.encode()).hexdigest()
|
|
|
|
|
|
def cipher(settings) -> Fernet:
|
|
return Fernet(settings.encryption_key.get_secret_value().encode())
|
|
|
|
|
|
async def bootstrap(db, settings):
|
|
"""Initialize singletons and apply environment credentials without resetting the admin password."""
|
|
if not await db.get(Admin, 1):
|
|
db.add(
|
|
Admin(
|
|
id=1,
|
|
username=settings.admin_username,
|
|
password_hash=password_hasher.hash(settings.admin_password.get_secret_value()),
|
|
)
|
|
)
|
|
account = await db.get(Account, 1)
|
|
if account is None:
|
|
account = Account(id=1)
|
|
db.add(account)
|
|
if settings.wq_email:
|
|
# The environment must not bypass the single-account data ownership boundary.
|
|
if account.wq_user_id and (account.email or "").casefold() != settings.wq_email.casefold():
|
|
raise ValueError("WQ_EMAIL conflicts with the bound WorldQuant account")
|
|
account.email = settings.wq_email
|
|
account.password_encrypted = cipher(settings).encrypt(
|
|
settings.wq_password.get_secret_value().encode()
|
|
).decode()
|
|
await db.execute(delete(LoginSession).where(LoginSession.expires_at < now()))
|
|
await db.commit()
|
|
|
|
|
|
def valid_password(encoded: str, value: str) -> bool:
|
|
try:
|
|
return password_hasher.verify(encoded, value)
|
|
except VerificationError:
|
|
return False
|
|
|
|
|
|
async def issue_session(db, hours: int) -> str:
|
|
token = secrets.token_urlsafe(32)
|
|
db.add(LoginSession(token_hash=token_hash(token), expires_at=now() + timedelta(hours=hours)))
|
|
await db.commit()
|
|
return token
|
|
|
|
|
|
async def require_auth(request: Request, token: str | None = Depends(session_cookie)):
|
|
if not token:
|
|
raise HTTPException(401, "请先登录系统")
|
|
async with request.app.state.sessions() as db:
|
|
row = await db.get(LoginSession, token_hash(token))
|
|
# SQLite fixtures lose tzinfo; production PostgreSQL preserves it.
|
|
if row is None or row.expires_at.replace(tzinfo=now().tzinfo) <= now():
|
|
raise HTTPException(401, "系统登录已过期,请重新登录")
|