315 lines
10 KiB
Python
315 lines
10 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,
|
|
ModelUpdateInput,
|
|
PasswordChangeInput,
|
|
)
|
|
|
|
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="模型已删除。")
|