409 lines
15 KiB
Python
409 lines
15 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.security import create_access_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 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(request.headers.get("x-ctms-client-type"))
|
|
|
|
|
|
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)
|
|
|
|
|
|
def issue_user_token(db_user, request: Request) -> Token:
|
|
session_start = datetime.now(timezone.utc)
|
|
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,
|
|
)
|
|
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 issue_user_token(db_user, request)
|
|
|
|
|
|
@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 issue_user_token(db_user, request)
|
|
|
|
|
|
@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")),
|
|
)
|
|
expires_at = policy_expires_at(issued_at, session_start, policy)
|
|
return ExtendResponse(accessToken=new_token, expiresAt=expires_at)
|
|
|
|
|
|
@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)
|