Skip to content

Commit 9b19e77

Browse files
authored
Update accounts.py
1 parent 09bce55 commit 9b19e77

1 file changed

Lines changed: 129 additions & 30 deletions

File tree

‎app/controllers/accounts.py‎

Lines changed: 129 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -1,36 +1,135 @@
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.
95
"""
106

117
from __future__ import annotations
128

13-
from typing import TypeVar
14-
9+
import jwt as pyjwt
10+
from sqlalchemy import func, select
1511
from sqlalchemy.orm import Session
1612

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

Comments
 (0)