Files
ctms/backend/app/api/v1/auth.py
T

462 lines
17 KiB
Python

from datetime import datetime, timedelta, timezone
from dataclasses import dataclass
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response, status
from fastapi import File, UploadFile
from pydantic import BaseModel, EmailStr, Field
from sqlalchemy.ext.asyncio import AsyncSession
from pathlib import Path
import uuid
from app.core.config import settings
from app.core.login_crypto import create_login_challenge, decrypt_login_payload, get_public_key_pem
from app.core.request_context import resolve_client_ip, resolve_ctms_client_type
from app.core.security import create_access_token, decode_token, decode_token_allow_expired, oauth2_scheme, verify_password
from app.core.deps import get_current_user, get_db_session
from app.crud import user as user_crud
from app.models.user import UserStatus
from app.schemas.email_settings import (
EmailCodeResponse,
EmailCodeVerifyResponse,
PasswordResetCodeVerifyRequest,
PasswordResetCodeVerifyResponse,
PasswordResetLinkSendRequest,
PasswordResetRequest,
PasswordResetTokenRequest,
RegisterEmailCodeSendRequest,
RegisterEmailCodeVerifyRequest,
)
from app.schemas.user import Token, UserRead, UserRegisterRequest, UserSelfUpdate, UserUpdate
from app.services import email_service
from app.services.user_login_sessions import (
create_login_session,
end_login_session,
session_id_from_payload,
touch_login_session,
)
from fastapi.responses import FileResponse
class LoginRequest(BaseModel):
key_id: str = Field(min_length=1)
challenge: str = Field(min_length=16)
ciphertext: str = Field(min_length=1)
class DevLoginRequest(BaseModel):
email: str = Field(min_length=1)
password: str = Field(min_length=1)
class LoginKeyResponse(BaseModel):
key_id: str
public_key: str
challenge: str
expires_at: datetime
class ExtendResponse(BaseModel):
accessToken: str
expiresAt: datetime
class EmailAvailabilityResponse(BaseModel):
available: bool
class EmailDomainsResponse(BaseModel):
items: list[str]
router = APIRouter()
AVATAR_ROOT = Path(__file__).resolve().parent.parent.parent / "uploads" / "avatars"
AVATAR_ROOT.mkdir(parents=True, exist_ok=True)
AVATAR_ALLOWED_CONTENT_TYPES = {
"image/png": ".png",
"image/jpeg": ".jpg",
"image/gif": ".gif",
"image/webp": ".webp",
}
@dataclass(frozen=True)
class SessionPolicy:
access_minutes: int
absolute_max_seconds: int
def normalize_session_client_type(value: str | None) -> str:
return "desktop" if (value or "").strip().lower() == "desktop" else "web"
def get_session_policy_for_client_type(client_type: str) -> SessionPolicy:
if client_type == "desktop":
max_seconds = settings.DESKTOP_SESSION_MAX_DAYS * 24 * 3600
return SessionPolicy(
access_minutes=settings.DESKTOP_SESSION_MAX_DAYS * 24 * 60,
absolute_max_seconds=max_seconds,
)
return SessionPolicy(
access_minutes=settings.JWT_EXPIRE_MINUTES,
absolute_max_seconds=settings.ABSOLUTE_SESSION_MAX_HOURS * 3600,
)
def get_request_session_client_type(request: Request) -> str:
return normalize_session_client_type(resolve_ctms_client_type(request.headers))
def policy_expires_at(issued_at: datetime, session_start: datetime, policy: SessionPolicy) -> datetime:
access_expires_at = issued_at + timedelta(minutes=policy.access_minutes)
session_expires_at = session_start + timedelta(seconds=policy.absolute_max_seconds)
return min(access_expires_at, session_expires_at)
async def issue_user_token(db_user, request: Request, db: AsyncSession) -> Token:
session_start = datetime.now(timezone.utc)
session_id = uuid.uuid4()
client_type = get_request_session_client_type(request)
policy = get_session_policy_for_client_type(client_type)
access_token = create_access_token(
user_id=str(db_user.id),
expires_minutes=policy.access_minutes,
session_start=session_start,
max_age_seconds=policy.absolute_max_seconds,
issued_at=session_start,
client_type=client_type,
session_id=str(session_id),
)
await create_login_session(
db,
session_id=session_id,
user_id=db_user.id,
request=request,
login_at=session_start,
)
return Token(access_token=access_token, token_type="bearer")
async def authenticate_encrypted_password(payload: LoginRequest, db: AsyncSession):
decrypted = decrypt_login_payload(
key_id=payload.key_id,
challenge=payload.challenge,
ciphertext=payload.ciphertext,
)
if not decrypted:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无法验证登录凭据",
)
db_user = await user_crud.get_by_email(db, decrypted.email)
if not db_user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="账号不存在",
)
if not verify_password(decrypted.password, db_user.password_hash):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="密码错误",
)
return db_user
async def authenticate_plain_password(payload: DevLoginRequest, db: AsyncSession):
db_user = await user_crud.get_by_email(db, payload.email)
if not db_user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="账号不存在",
)
if not verify_password(payload.password, db_user.password_hash):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="密码错误",
)
return db_user
def ensure_user_active(db_user) -> None:
if db_user.status != UserStatus.ACTIVE:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="账号未审核或不可用",
)
@router.get("/email-domains", response_model=EmailDomainsResponse)
async def read_email_domains(
response: Response,
db: AsyncSession = Depends(get_db_session),
) -> EmailDomainsResponse:
response.headers["Cache-Control"] = "no-store"
rows = await email_service.list_email_settings(db)
return EmailDomainsResponse(items=[row.register_domain for row in rows])
@router.get("/register/email-availability", response_model=EmailAvailabilityResponse)
async def check_register_email_availability(
email: EmailStr = Query(...),
db: AsyncSession = Depends(get_db_session),
) -> EmailAvailabilityResponse:
existing = await user_crud.get_by_email(db, str(email))
return EmailAvailabilityResponse(available=existing is None)
@router.post("/register/email-code/send", response_model=EmailCodeResponse)
async def send_register_email_code(
payload: RegisterEmailCodeSendRequest,
db: AsyncSession = Depends(get_db_session),
) -> EmailCodeResponse:
await email_service.send_register_code(db, str(payload.email))
return EmailCodeResponse(message="验证码已发送")
@router.post("/register/email-code/verify", response_model=EmailCodeVerifyResponse)
async def verify_register_email_code(
payload: RegisterEmailCodeVerifyRequest,
db: AsyncSession = Depends(get_db_session),
) -> EmailCodeVerifyResponse:
await email_service.verify_register_code(db, str(payload.email), payload.code)
return EmailCodeVerifyResponse(verified=True)
@router.post("/password-reset/email-code/send", response_model=EmailCodeResponse)
async def send_password_reset_email_code(
payload: PasswordResetLinkSendRequest,
db: AsyncSession = Depends(get_db_session),
) -> EmailCodeResponse:
await email_service.send_password_reset_code(db, str(payload.email))
return EmailCodeResponse(message="验证码发送成功,请查收邮箱")
@router.post("/password-reset/email-code/verify", response_model=PasswordResetCodeVerifyResponse)
async def verify_password_reset_email_code(
payload: PasswordResetCodeVerifyRequest,
db: AsyncSession = Depends(get_db_session),
) -> PasswordResetCodeVerifyResponse:
reset_token = await email_service.verify_password_reset_code(db, str(payload.email), payload.code)
return PasswordResetCodeVerifyResponse(verified=True, reset_token=reset_token)
@router.post("/password-reset-link/send", response_model=EmailCodeResponse)
async def send_password_reset_link(
payload: PasswordResetLinkSendRequest,
db: AsyncSession = Depends(get_db_session),
) -> EmailCodeResponse:
await email_service.send_password_reset_link(
db,
str(payload.email),
frontend_origin=settings.FRONTEND_PUBLIC_URL,
)
return EmailCodeResponse(message="如果账号存在,重置链接已发送")
@router.post("/password-reset", response_model=EmailCodeResponse)
async def reset_password(
payload: PasswordResetRequest,
db: AsyncSession = Depends(get_db_session),
) -> EmailCodeResponse:
await email_service.reset_password_with_code(db, str(payload.email), payload.code, payload.password)
return EmailCodeResponse(message="密码已重置,请返回登录")
@router.post("/password-reset-link", response_model=EmailCodeResponse)
async def reset_password_with_link(
payload: PasswordResetTokenRequest,
db: AsyncSession = Depends(get_db_session),
) -> EmailCodeResponse:
await email_service.reset_password_with_token(db, payload.token, payload.password)
return EmailCodeResponse(message="密码已重置,请返回登录")
@router.post("/register", status_code=status.HTTP_201_CREATED)
async def register(
payload: UserRegisterRequest,
db: AsyncSession = Depends(get_db_session),
):
existing = await user_crud.get_by_email(db, payload.email)
if existing:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="邮箱已注册")
await email_service.ensure_register_email_verified(db, str(payload.email))
await user_crud.create_registered_user(db, payload)
return {"message": "注册成功,请登录"}
@router.get("/login-key", response_model=LoginKeyResponse)
async def get_login_key() -> LoginKeyResponse:
challenge = create_login_challenge()
return LoginKeyResponse(
key_id=settings.LOGIN_RSA_KEY_ID,
public_key=get_public_key_pem(),
challenge=challenge.value,
expires_at=challenge.expires_at,
)
@router.post("/login", response_model=Token)
async def login_for_access_token(
payload: LoginRequest, request: Request, db: AsyncSession = Depends(get_db_session)
) -> Token:
db_user = await authenticate_encrypted_password(payload, db)
ensure_user_active(db_user)
return await issue_user_token(db_user, request, db)
@router.post("/dev-login", response_model=Token)
async def dev_login_for_access_token(
payload: DevLoginRequest, request: Request, db: AsyncSession = Depends(get_db_session)
) -> Token:
if settings.ENV != "development":
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not found")
db_user = await authenticate_plain_password(payload, db)
ensure_user_active(db_user)
return await issue_user_token(db_user, request, db)
@router.get("/me", response_model=UserRead)
async def read_me(current_user=Depends(get_current_user)) -> UserRead:
return current_user
@router.post("/extend", response_model=ExtendResponse)
async def extend_access_token(
token: str = Depends(oauth2_scheme),
db: AsyncSession = Depends(get_db_session),
) -> ExtendResponse:
payload = decode_token_allow_expired(token)
exp = payload.get("exp")
if not exp:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="登录已过期")
now_ts = int(datetime.now(timezone.utc).timestamp())
if now_ts > int(exp) + settings.JWT_EXTEND_GRACE_SECONDS:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="登录已过期")
user_id = payload.get("sub")
if not user_id:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无法验证登录凭据")
db_user = await user_crud.get_by_id(db, uuid.UUID(str(user_id)))
if not db_user:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="账号不存在或已停用")
if db_user.status != UserStatus.ACTIVE:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="账号已停用")
session_start_ts = payload.get("orig_iat") or payload.get("iat")
policy = get_session_policy_for_client_type(normalize_session_client_type(payload.get("client_type")))
if session_start_ts:
if now_ts - int(session_start_ts) > policy.absolute_max_seconds:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="会话已到期,请重新登录")
session_start = datetime.fromtimestamp(int(session_start_ts), tz=timezone.utc)
else:
session_start = datetime.now(timezone.utc)
issued_at = datetime.now(timezone.utc)
new_token = create_access_token(
user_id=str(db_user.id),
expires_minutes=policy.access_minutes,
session_start=session_start,
max_age_seconds=policy.absolute_max_seconds,
issued_at=issued_at,
client_type=normalize_session_client_type(payload.get("client_type")),
session_id=str(session_id_from_payload(payload)),
)
expires_at = policy_expires_at(issued_at, session_start, policy)
return ExtendResponse(accessToken=new_token, expiresAt=expires_at)
@router.post("/session/heartbeat")
async def heartbeat_login_session(
request: Request,
token: str = Depends(oauth2_scheme),
current_user=Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> dict:
session = await touch_login_session(
db,
user_id=current_user.id,
payload=decode_token(token),
request=request,
)
if session is None:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="登录会话已结束")
return {
"status": "online",
"last_seen_at": session.last_seen_at.isoformat(),
"client_ip": resolve_client_ip(request),
}
@router.post("/session/logout", status_code=status.HTTP_204_NO_CONTENT)
async def logout_login_session(
token: str = Depends(oauth2_scheme),
current_user=Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> Response:
await end_login_session(
db,
user_id=current_user.id,
payload=decode_token(token),
)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@router.patch("/me", response_model=UserRead)
async def update_me(
payload: UserSelfUpdate,
current_user=Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> UserRead:
if payload.password:
if not payload.current_password or not verify_password(payload.current_password, current_user.password_hash):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="当前密码不正确")
update_data = {
"full_name": payload.full_name if payload.full_name is not None else current_user.full_name,
"clinical_department": (
payload.clinical_department if payload.clinical_department is not None else current_user.clinical_department
),
"password": payload.password if payload.password else None,
"avatar_url": payload.avatar_url if payload.avatar_url is not None else current_user.avatar_url,
}
updated = await user_crud.update_user(db, current_user, UserUpdate(**update_data))
return updated
@router.post("/me/avatar", response_model=UserRead)
async def upload_avatar(
file: UploadFile = File(...),
current_user=Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> UserRead:
ext = AVATAR_ALLOWED_CONTENT_TYPES.get(file.content_type or "")
if not ext:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="头像仅支持图片格式",
)
AVATAR_ROOT.mkdir(parents=True, exist_ok=True)
user_dir = AVATAR_ROOT / str(current_user.id)
user_dir.mkdir(parents=True, exist_ok=True)
filename = f"{uuid.uuid4()}{ext}"
dest = user_dir / filename
content = await file.read()
dest.write_bytes(content)
# Optionally clean old files
for old in user_dir.iterdir():
if old.name != filename:
try:
old.unlink()
except OSError:
pass
rel_url = f"/api/v1/auth/avatar/{current_user.id}/{filename}"
updated = await user_crud.update_user(
db,
current_user,
UserUpdate(avatar_url=rel_url),
)
return updated
@router.get("/avatar/{user_id}/{filename}")
async def get_avatar(user_id: str, filename: str):
file_path = AVATAR_ROOT / user_id / filename
if not file_path.exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="头像不存在")
return FileResponse(file_path)