|
1 | | -"""Shared ownership lookup for controllers: fetch-or-404 with an owner |
2 | | -check. |
3 | | -
|
4 | | -Several controllers fetch a row by id and 404 unless it belongs to the |
5 | | -caller's parent (a conversation to its user, a run/file to its |
6 | | -conversation). This centralizes that "missing OR not yours -> NotFound" |
7 | | -branch; ownership failures surface as 404 rather than 403 so an id can't |
8 | | -be probed for existence. |
| 1 | +"""Account business logic: registration, credentials, token identity. |
| 2 | +
|
| 3 | +The routes keep only FastAPI's dependency plumbing; deciding whether a |
| 4 | +credential or token grants access is decided here. |
9 | 5 | """ |
10 | 6 |
|
11 | 7 | from __future__ import annotations |
12 | 8 |
|
13 | | -from typing import TypeVar |
14 | | - |
| 9 | +import jwt as pyjwt |
| 10 | +from sqlalchemy import func, select |
15 | 11 | from sqlalchemy.orm import Session |
16 | 12 |
|
17 | | -from .errors import NotFound |
18 | | - |
19 | | -M = TypeVar("M") |
20 | | - |
21 | | - |
22 | | -def get_owned( |
23 | | - db: Session, |
24 | | - model: type[M], |
25 | | - id_: str, |
26 | | - *, |
27 | | - owner_field: str, |
28 | | - owner_value: object, |
29 | | - detail: str, |
30 | | -) -> M: |
31 | | - """Return the row iff it exists and its ``owner_field`` matches |
32 | | - ``owner_value``; otherwise raise ``NotFound(detail)``.""" |
33 | | - obj = db.get(model, id_) |
34 | | - if obj is None or getattr(obj, owner_field) != owner_value: |
35 | | - raise NotFound(detail) |
36 | | - return obj |
| 13 | +from ..infra.db import commit_now |
| 14 | +from ..infra.security import ( |
| 15 | + create_access_token, |
| 16 | + create_refresh_token, |
| 17 | + decode_token, |
| 18 | + hash_password, |
| 19 | + verify_password, |
| 20 | +) |
| 21 | +from ..models import Conversation, User, new_id |
| 22 | +from .errors import Conflict, Forbidden, Unauthorized |
| 23 | + |
| 24 | + |
| 25 | +def token_pair(user: User) -> tuple[str, str]: |
| 26 | + """Fresh (access, refresh) pair for a user, stamped with the user's |
| 27 | + current token version so a later revocation invalidates them.""" |
| 28 | + return ( |
| 29 | + create_access_token(user.id, user.token_version), |
| 30 | + create_refresh_token(user.id, user.token_version), |
| 31 | + ) |
| 32 | + |
| 33 | + |
| 34 | +def register(db: Session, email: str, password: str) -> User: |
| 35 | + """Create an account; the email must be unused.""" |
| 36 | + if db.scalar(select(User).where(User.email == email)) is not None: |
| 37 | + raise Conflict("email already registered") |
| 38 | + user = User(id=new_id("usr"), email=email, password_hash=hash_password(password)) |
| 39 | + db.add(user) |
| 40 | + db.flush() |
| 41 | + commit_now(db) |
| 42 | + return user |
| 43 | + |
| 44 | + |
| 45 | +def login(db: Session, email: str, password: str) -> User: |
| 46 | + """Authenticate a password. |
| 47 | +
|
| 48 | + A wrong email and a wrong password give the same 401, so neither |
| 49 | + can be used to enumerate accounts; a disabled account is told |
| 50 | + apart (403) only once the password checked out. |
| 51 | + """ |
| 52 | + user = db.scalar(select(User).where(User.email == email)) |
| 53 | + if user is None or not verify_password(password, user.password_hash): |
| 54 | + raise Unauthorized("invalid email or password") |
| 55 | + if not user.is_active: |
| 56 | + raise Forbidden("account disabled") |
| 57 | + return user |
| 58 | + |
| 59 | + |
| 60 | +def _user_of_token(db: Session, token: str, *, expected_type: str) -> User: |
| 61 | + try: |
| 62 | + payload = decode_token(token, expected_type=expected_type) |
| 63 | + except pyjwt.PyJWTError as exc: |
| 64 | + raise Unauthorized(f"invalid token: {exc}") from exc |
| 65 | + user = db.get(User, payload["sub"]) |
| 66 | + if user is None or not user.is_active: |
| 67 | + raise Unauthorized("unknown or inactive user") |
| 68 | + # Tokens minted before the last revocation (logout / password change) |
| 69 | + # carry an older version and are rejected. |
| 70 | + if payload.get("ver", 0) != user.token_version: |
| 71 | + raise Unauthorized("token revoked") |
| 72 | + return user |
| 73 | + |
| 74 | + |
| 75 | +def user_for_access_token(db: Session, token: str | None) -> User: |
| 76 | + """Identify the caller from a bearer access token.""" |
| 77 | + if not token: |
| 78 | + raise Unauthorized("missing bearer token") |
| 79 | + return _user_of_token(db, token, expected_type="access") |
| 80 | + |
| 81 | + |
| 82 | +def user_for_refresh_token(db: Session, token: str) -> User: |
| 83 | + """Identify the holder of a refresh token.""" |
| 84 | + return _user_of_token(db, token, expected_type="refresh") |
| 85 | + |
| 86 | + |
| 87 | +def require_admin(user: User) -> User: |
| 88 | + if not user.is_admin: |
| 89 | + raise Forbidden("admin required") |
| 90 | + return user |
| 91 | + |
| 92 | + |
| 93 | +def revoke_all_tokens(db: Session, user: User) -> None: |
| 94 | + """Invalidate every outstanding token for the user (logout). |
| 95 | +
|
| 96 | + Bumps ``token_version``; tokens carrying the old version are then |
| 97 | + rejected on their next use. Note this logs out *all* of the user's |
| 98 | + sessions, not just one — fine for this app's single-client UI. |
| 99 | + """ |
| 100 | + user.token_version += 1 |
| 101 | + commit_now(db) |
| 102 | + |
| 103 | + |
| 104 | +def change_password(db: Session, user: User, old_password: str, new_password: str) -> None: |
| 105 | + """Set a new password after verifying the current one, and revoke |
| 106 | + all existing sessions so a leaked/old token can't outlive the change.""" |
| 107 | + if not verify_password(old_password, user.password_hash): |
| 108 | + raise Unauthorized("current password is incorrect") |
| 109 | + user.password_hash = hash_password(new_password) |
| 110 | + user.token_version += 1 |
| 111 | + commit_now(db) |
| 112 | + |
| 113 | + |
| 114 | +def profile(db: Session, user: User) -> dict: |
| 115 | + """Read model for the account dropdown and 个人中心 page. |
| 116 | +
|
| 117 | + One query joins the stored account fields with the user's live |
| 118 | + conversation count (the app's "创作"). ``points`` is derived from |
| 119 | + the two point buckets so the total can never drift from its parts. |
| 120 | + ``name`` is the email local-part until a display-name column exists. |
| 121 | + """ |
| 122 | + works = ( |
| 123 | + db.scalar(select(func.count(Conversation.id)).where(Conversation.user_id == user.id)) or 0 |
| 124 | + ) |
| 125 | + return { |
| 126 | + "name": user.email.split("@")[0], |
| 127 | + "email": user.email, |
| 128 | + "plan": user.plan, |
| 129 | + "phone": user.phone, |
| 130 | + "points": user.plan_points + user.pack_points, |
| 131 | + "plan_points": user.plan_points, |
| 132 | + "pack_points": user.pack_points, |
| 133 | + "points_expire_at": user.points_expire_at, |
| 134 | + "works": int(works), |
| 135 | + } |
0 commit comments