Files
xiaobaifupan/next/backend/features/accounts/routes.py
T

331 lines
11 KiB
Python

from __future__ import annotations
from datetime import date, time
from typing import Annotated
from fastapi import APIRouter, Path, Request, Response
from backend.features.accounts.auth import (
CSRF_COOKIE,
SESSION_COOKIE,
AdminPrincipal,
AdminWritePrincipal,
AuthenticatedPrincipal,
CsrfPrincipal,
)
from backend.features.accounts.models import (
BirthProfile,
MembershipAccountView,
MembershipView,
Principal,
SessionIssue,
)
from backend.features.accounts.schemas import (
AccountIdentityResponse,
AuthResponse,
BirthProfileInput,
BirthProfileResponse,
CredentialInput,
CredentialsInput,
CredentialStatusResponse,
MembershipAdminResponse,
MembershipStatusResponse,
MembershipUpdateInput,
MessageResponse,
ModelInput,
ModelPoolItemResponse,
ModelSelectionInput,
ModelTestResponse,
ModelUpdateInput,
PasswordChangeInput,
)
from backend.http.errors import AppError
from backend.llm.gateway import LLMGatewayError
router = APIRouter()
def _membership_response(view: MembershipView, smart_access: bool) -> MembershipStatusResponse:
return MembershipStatusResponse(
status=view.status,
active=view.active,
is_permanent=view.is_permanent,
expires_at=view.expires_at,
remaining_days=view.remaining_days,
daily_limit=view.daily_limit,
used_today=view.used_today,
remaining_today=view.remaining_today,
quota_exempt=view.quota_exempt,
smart_access=smart_access,
)
def _identity(request: Request, principal: Principal) -> AccountIdentityResponse:
memberships = request.app.state.container.memberships
view = memberships.view_for(principal)
badges = (["admin"] if principal.user.is_admin else []) + (["member"] if view.active else [])
return AccountIdentityResponse(
id=principal.user.id,
username=principal.user.username,
is_admin=principal.user.is_admin,
membership_status=view.status,
membership_active=view.active,
smart_access=memberships.can_use_smart_features(principal),
badges=badges,
)
def _set_session_cookies(request: Request, response: Response, issue: SessionIssue) -> None:
secure = request.app.state.settings.environment == "production"
max_age = 30 * 24 * 60 * 60
response.set_cookie(
SESSION_COOKIE,
issue.token,
max_age=max_age,
httponly=True,
secure=secure,
samesite="lax",
path="/",
)
response.set_cookie(
CSRF_COOKIE,
issue.csrf_token,
max_age=max_age,
httponly=False,
secure=secure,
samesite="lax",
path="/",
)
def _clear_session_cookies(response: Response) -> None:
response.delete_cookie(SESSION_COOKIE, path="/")
response.delete_cookie(CSRF_COOKIE, path="/")
def _profile_response(profile: BirthProfile | None) -> BirthProfileResponse:
if profile is None:
return BirthProfileResponse(configured=False)
return BirthProfileResponse(
configured=True,
birth_date=date.fromisoformat(profile.birth_date),
birth_time=time.fromisoformat(profile.birth_time),
gender=profile.gender,
updated_at=profile.updated_at,
)
def _admin_membership_response(
account: MembershipAccountView,
) -> MembershipAdminResponse:
return MembershipAdminResponse(
user_id=account.user.id,
username=account.user.username,
is_admin=account.user.is_admin,
membership=_membership_response(
account.membership,
account.user.is_admin or account.membership.active,
),
)
@router.post("/auth/register", response_model=AuthResponse, status_code=201)
def register(payload: CredentialsInput, request: Request, response: Response) -> AuthResponse:
issue = request.app.state.container.accounts.register(payload.username, payload.password)
_set_session_cookies(request, response, issue)
return AuthResponse(account=_identity(request, issue.principal), csrf_token=issue.csrf_token)
@router.post("/auth/login", response_model=AuthResponse)
def login(payload: CredentialsInput, request: Request, response: Response) -> AuthResponse:
issue = request.app.state.container.accounts.login(payload.username, payload.password)
_set_session_cookies(request, response, issue)
return AuthResponse(account=_identity(request, issue.principal), csrf_token=issue.csrf_token)
@router.get("/auth/session", response_model=AccountIdentityResponse)
def session(request: Request, principal: AuthenticatedPrincipal):
return _identity(request, principal)
@router.post("/auth/logout", response_model=MessageResponse)
@router.post("/auth/switch-account", response_model=MessageResponse)
def logout(
request: Request,
response: Response,
principal: CsrfPrincipal,
) -> MessageResponse:
request.app.state.container.accounts.logout(principal)
_clear_session_cookies(response)
return MessageResponse(message="已退出当前账号。")
@router.patch("/account/password", response_model=MessageResponse)
def change_password(
payload: PasswordChangeInput,
request: Request,
principal: CsrfPrincipal,
) -> MessageResponse:
request.app.state.container.accounts.change_password(
principal,
payload.current_password,
payload.new_password,
payload.confirmation,
)
return MessageResponse(message="密码已修改。")
@router.get("/account/profile", response_model=BirthProfileResponse)
def get_profile(request: Request, principal: AuthenticatedPrincipal) -> BirthProfileResponse:
return _profile_response(request.app.state.container.accounts.get_profile(principal.user.id))
@router.put("/account/profile", response_model=BirthProfileResponse)
def save_profile(
payload: BirthProfileInput,
request: Request,
principal: CsrfPrincipal,
) -> BirthProfileResponse:
profile = request.app.state.container.accounts.save_profile(
principal.user.id,
payload.birth_date,
payload.birth_time.isoformat(timespec="minutes"),
payload.gender,
)
return _profile_response(profile)
@router.delete("/account/profile", response_model=MessageResponse)
def delete_profile(request: Request, principal: CsrfPrincipal) -> MessageResponse:
request.app.state.container.accounts.delete_profile(principal.user.id)
return MessageResponse(message="个人出生资料已删除。")
@router.get("/account/membership", response_model=MembershipStatusResponse)
def get_membership(request: Request, principal: AuthenticatedPrincipal) -> MembershipStatusResponse:
memberships = request.app.state.container.memberships
view = memberships.view_for(principal)
return _membership_response(view, memberships.can_use_smart_features(principal))
@router.get(
"/admin/memberships",
response_model=list[MembershipAdminResponse],
)
def list_memberships(request: Request, _principal: AdminPrincipal) -> list[MembershipAdminResponse]:
return [
_admin_membership_response(account)
for account in request.app.state.container.memberships.list_accounts()
]
@router.patch(
"/admin/memberships/{user_id}",
response_model=MembershipAdminResponse,
)
def update_membership(
payload: MembershipUpdateInput,
request: Request,
user_id: Annotated[int, Path(ge=1)],
principal: AdminWritePrincipal,
) -> MembershipAdminResponse:
request.app.state.container.memberships.update(
principal, user_id, payload.action, payload.duration, payload.daily_limit
)
return _admin_membership_response(request.app.state.container.memberships.get_account(user_id))
@router.get(
"/admin/system/credentials",
response_model=list[CredentialStatusResponse],
)
def list_credentials(
request: Request, _principal: AdminPrincipal
) -> tuple[dict[str, str | bool | None], ...]:
return request.app.state.container.system_credentials.list_status()
@router.put("/admin/system/credentials/{name}", response_model=MessageResponse)
def save_credential(
payload: CredentialInput,
request: Request,
name: Annotated[str, Path(min_length=1, max_length=64)],
principal: AdminWritePrincipal,
) -> MessageResponse:
request.app.state.container.system_credentials.save(principal, name, payload.value)
return MessageResponse(message="系统凭据已保存。")
@router.get("/admin/models", response_model=list[ModelPoolItemResponse])
def list_models(
request: Request, _principal: AdminPrincipal
) -> tuple[object, ...]:
return request.app.state.container.model_pool.list()
@router.post("/admin/models", response_model=ModelPoolItemResponse, status_code=201)
def create_model(
payload: ModelInput,
request: Request,
principal: AdminWritePrincipal,
) -> ModelPoolItemResponse:
return request.app.state.container.model_pool.create(
principal,
payload.display_name,
payload.base_url,
payload.model_identifier,
payload.api_key,
)
@router.put("/admin/models/selection", response_model=MessageResponse)
def select_models(
payload: ModelSelectionInput,
request: Request,
principal: AdminWritePrincipal,
) -> MessageResponse:
request.app.state.container.model_pool.select(
principal, payload.primary_model_id, payload.fallback_model_id
)
return MessageResponse(message="主模型和辅助模型已更新。")
@router.put("/admin/models/{model_id}", response_model=ModelPoolItemResponse)
def update_model(
payload: ModelUpdateInput,
request: Request,
model_id: Annotated[int, Path(ge=1)],
principal: AdminWritePrincipal,
) -> ModelPoolItemResponse:
return request.app.state.container.model_pool.update(
principal,
model_id,
payload.display_name,
payload.base_url,
payload.model_identifier,
payload.api_key,
)
@router.delete("/admin/models/{model_id}", response_model=MessageResponse)
def delete_model(
request: Request,
model_id: Annotated[int, Path(ge=1)],
_principal: AdminWritePrincipal,
) -> MessageResponse:
request.app.state.container.model_pool.delete(model_id)
return MessageResponse(message="模型已删除。")
@router.post("/admin/models/{model_id}/test", response_model=ModelTestResponse)
def test_model(
request: Request,
model_id: Annotated[int, Path(ge=1)],
principal: AdminWritePrincipal,
) -> dict[str, object]:
try:
return request.app.state.container.llm.test_model(principal, model_id)
except LLMGatewayError as exc:
status = 404 if exc.code == "model_not_found" else 503
raise AppError(exc.code, str(exc), status) from exc