account-kit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- account_kit/__init__.py +9 -0
- account_kit/admin_router.py +181 -0
- account_kit/admin_schemas.py +58 -0
- account_kit/catalog.py +167 -0
- account_kit/config.py +81 -0
- account_kit/crypto.py +11 -0
- account_kit/db.py +14 -0
- account_kit/deps.py +74 -0
- account_kit/emailer.py +71 -0
- account_kit/models.py +167 -0
- account_kit/otp.py +99 -0
- account_kit/profile.py +54 -0
- account_kit/router.py +174 -0
- account_kit/schema_setup.py +11 -0
- account_kit/schemas.py +115 -0
- account_kit/security.py +47 -0
- account_kit/seed.py +59 -0
- account_kit/service.py +329 -0
- account_kit/templates/emails/verification_en.html +10 -0
- account_kit/templates/emails/verification_zh.html +10 -0
- account_kit/two_factor/__init__.py +1 -0
- account_kit/two_factor/challenges.py +69 -0
- account_kit/two_factor/router.py +165 -0
- account_kit/two_factor/service.py +308 -0
- account_kit/two_factor/totp.py +65 -0
- account_kit-0.1.0.dist-info/METADATA +108 -0
- account_kit-0.1.0.dist-info/RECORD +30 -0
- account_kit-0.1.0.dist-info/WHEEL +5 -0
- account_kit-0.1.0.dist-info/licenses/LICENSE +21 -0
- account_kit-0.1.0.dist-info/top_level.txt +1 -0
account_kit/__init__.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
"""Shared account kit: one async user model for every app."""
|
|
2
|
+
|
|
3
|
+
from account_kit.router import mount_account
|
|
4
|
+
from account_kit.schema_setup import init_db
|
|
5
|
+
from account_kit.seed import seed_defaults
|
|
6
|
+
|
|
7
|
+
__version__ = "0.1.0"
|
|
8
|
+
|
|
9
|
+
__all__ = ["init_db", "mount_account", "seed_defaults", "__version__"]
|
|
@@ -0,0 +1,181 @@
|
|
|
1
|
+
import uuid
|
|
2
|
+
|
|
3
|
+
from fastapi import APIRouter, Depends, HTTPException
|
|
4
|
+
from sqlalchemy import select
|
|
5
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
6
|
+
|
|
7
|
+
from account_kit.admin_schemas import AdminUserPatch, RoleBody, RoleChangeCreate, RoleChangeReview, RolePatch, TierBody, TierPatch
|
|
8
|
+
from account_kit.catalog import apply_admin_user_patch, create_role, create_tier, delete_role, delete_tier, patch_role, patch_tier
|
|
9
|
+
from account_kit.config import get_config
|
|
10
|
+
from account_kit.deps import get_current_user, get_db, require_admin
|
|
11
|
+
from account_kit.models import Role, RoleChangeRequest, User, UserTier
|
|
12
|
+
from account_kit.service import assert_role_code, get_user_by_id, to_response
|
|
13
|
+
|
|
14
|
+
admin_router = APIRouter(tags=["account-admin"])
|
|
15
|
+
public_extra = APIRouter(tags=["auth"])
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _role_dict(row: Role) -> dict:
|
|
19
|
+
return {
|
|
20
|
+
"code": row.code,
|
|
21
|
+
"name": row.name,
|
|
22
|
+
"sort_order": row.sort_order,
|
|
23
|
+
"description": row.description,
|
|
24
|
+
"is_default": row.is_default,
|
|
25
|
+
"allow_register": row.allow_register,
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _tier_dict(row: UserTier) -> dict:
|
|
30
|
+
return {
|
|
31
|
+
"code": row.code,
|
|
32
|
+
"name": row.name,
|
|
33
|
+
"sort_order": row.sort_order,
|
|
34
|
+
"badge_color": row.badge_color,
|
|
35
|
+
"description": row.description,
|
|
36
|
+
"is_default": row.is_default,
|
|
37
|
+
}
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
@public_extra.get("/roles")
|
|
41
|
+
async def public_roles(db: AsyncSession = Depends(get_db)):
|
|
42
|
+
result = await db.execute(select(Role).where(Role.allow_register.is_(True)).order_by(Role.sort_order, Role.code))
|
|
43
|
+
return [_role_dict(row) for row in result.scalars().all()]
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@public_extra.get("/tiers")
|
|
47
|
+
async def public_tiers(db: AsyncSession = Depends(get_db)):
|
|
48
|
+
result = await db.execute(select(UserTier).order_by(UserTier.sort_order, UserTier.code))
|
|
49
|
+
return [_tier_dict(row) for row in result.scalars().all()]
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@public_extra.post("/role-change-requests")
|
|
53
|
+
async def request_role_change(
|
|
54
|
+
body: RoleChangeCreate,
|
|
55
|
+
current_user=Depends(get_current_user),
|
|
56
|
+
db: AsyncSession = Depends(get_db),
|
|
57
|
+
):
|
|
58
|
+
if not get_config().role_change_enabled:
|
|
59
|
+
raise HTTPException(status_code=404, detail="Not found")
|
|
60
|
+
code = assert_role_code(body.requested_role)
|
|
61
|
+
role = await db.get(Role, code)
|
|
62
|
+
if role is None:
|
|
63
|
+
raise HTTPException(status_code=400, detail="角色不存在")
|
|
64
|
+
if code == current_user.role:
|
|
65
|
+
raise HTTPException(status_code=400, detail="已经是该角色")
|
|
66
|
+
existing = await db.execute(
|
|
67
|
+
select(RoleChangeRequest).where(
|
|
68
|
+
RoleChangeRequest.user_id == current_user.id,
|
|
69
|
+
RoleChangeRequest.status == "pending",
|
|
70
|
+
)
|
|
71
|
+
)
|
|
72
|
+
if existing.scalars().first():
|
|
73
|
+
raise HTTPException(status_code=409, detail="已有待审核的角色申请")
|
|
74
|
+
row = RoleChangeRequest(user_id=current_user.id, from_role=current_user.role, to_role=code, status="pending")
|
|
75
|
+
db.add(row)
|
|
76
|
+
await db.commit()
|
|
77
|
+
await db.refresh(row)
|
|
78
|
+
return {"id": row.id, "from_role": row.from_role, "to_role": row.to_role, "status": row.status}
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
@admin_router.get("/roles")
|
|
82
|
+
async def list_roles(_admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
83
|
+
result = await db.execute(select(Role).order_by(Role.sort_order, Role.code))
|
|
84
|
+
return [_role_dict(row) for row in result.scalars().all()]
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
@admin_router.post("/roles", status_code=201)
|
|
88
|
+
async def add_role(body: RoleBody, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
89
|
+
return _role_dict(await create_role(db, body))
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@admin_router.patch("/roles/{code}")
|
|
93
|
+
async def edit_role(code: str, body: RolePatch, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
94
|
+
return _role_dict(await patch_role(db, code, body))
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
@admin_router.delete("/roles/{code}", status_code=204)
|
|
98
|
+
async def remove_role(code: str, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
99
|
+
await delete_role(db, code)
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
@admin_router.get("/tiers")
|
|
103
|
+
async def list_tiers(_admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
104
|
+
result = await db.execute(select(UserTier).order_by(UserTier.sort_order, UserTier.code))
|
|
105
|
+
return [_tier_dict(row) for row in result.scalars().all()]
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
@admin_router.post("/tiers", status_code=201)
|
|
109
|
+
async def add_tier(body: TierBody, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
110
|
+
return _tier_dict(await create_tier(db, body))
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
@admin_router.patch("/tiers/{code}")
|
|
114
|
+
async def edit_tier(code: str, body: TierPatch, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
115
|
+
return _tier_dict(await patch_tier(db, code, body))
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
@admin_router.delete("/tiers/{code}", status_code=204)
|
|
119
|
+
async def remove_tier(code: str, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
120
|
+
await delete_tier(db, code)
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
@admin_router.get("/users")
|
|
124
|
+
async def list_users(_admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
125
|
+
result = await db.execute(select(User).order_by(User.created_at.desc()))
|
|
126
|
+
return [await to_response(db, user) for user in result.scalars().all()]
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
@admin_router.patch("/users/{user_id}")
|
|
130
|
+
async def patch_user(
|
|
131
|
+
user_id: str,
|
|
132
|
+
body: AdminUserPatch,
|
|
133
|
+
_admin=Depends(require_admin),
|
|
134
|
+
db: AsyncSession = Depends(get_db),
|
|
135
|
+
):
|
|
136
|
+
user = await get_user_by_id(db, user_id)
|
|
137
|
+
if user is None:
|
|
138
|
+
raise HTTPException(status_code=404, detail="用户不存在")
|
|
139
|
+
user = await apply_admin_user_patch(db, user, body)
|
|
140
|
+
return await to_response(db, user)
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
@admin_router.delete("/users/{user_id}", status_code=204)
|
|
144
|
+
async def delete_user(user_id: str, _admin=Depends(require_admin), db: AsyncSession = Depends(get_db)):
|
|
145
|
+
config = get_config()
|
|
146
|
+
user = await get_user_by_id(db, user_id)
|
|
147
|
+
if user is None:
|
|
148
|
+
raise HTTPException(status_code=404, detail="用户不存在")
|
|
149
|
+
if config.on_deleted is not None:
|
|
150
|
+
await config.on_deleted(db, user)
|
|
151
|
+
await db.delete(user)
|
|
152
|
+
await db.commit()
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
@admin_router.post("/role-change-requests/{request_id}/review")
|
|
156
|
+
async def review_role_change(
|
|
157
|
+
request_id: str,
|
|
158
|
+
body: RoleChangeReview,
|
|
159
|
+
admin=Depends(require_admin),
|
|
160
|
+
db: AsyncSession = Depends(get_db),
|
|
161
|
+
):
|
|
162
|
+
if body.status not in ("approved", "rejected"):
|
|
163
|
+
raise HTTPException(status_code=400, detail="状态无效")
|
|
164
|
+
try:
|
|
165
|
+
request_uuid = uuid.UUID(request_id)
|
|
166
|
+
except ValueError:
|
|
167
|
+
raise HTTPException(status_code=404, detail="申请不存在")
|
|
168
|
+
row = await db.get(RoleChangeRequest, request_uuid)
|
|
169
|
+
if row is None or row.status != "pending":
|
|
170
|
+
raise HTTPException(status_code=404, detail="申请不存在")
|
|
171
|
+
row.status = body.status
|
|
172
|
+
from datetime import datetime, timezone
|
|
173
|
+
|
|
174
|
+
row.reviewed_at = datetime.now(timezone.utc)
|
|
175
|
+
row.reviewed_by = admin.id
|
|
176
|
+
if body.status == "approved":
|
|
177
|
+
user = await get_user_by_id(db, row.user_id)
|
|
178
|
+
if user is not None:
|
|
179
|
+
user.role = row.to_role
|
|
180
|
+
await db.commit()
|
|
181
|
+
return {"id": row.id, "status": row.status}
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import uuid
|
|
4
|
+
from typing import Optional
|
|
5
|
+
|
|
6
|
+
from pydantic import BaseModel, Field
|
|
7
|
+
|
|
8
|
+
from account_kit.schemas import ApprovalStatus
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class RoleBody(BaseModel):
|
|
12
|
+
code: str = Field(min_length=1, max_length=50)
|
|
13
|
+
name: str = Field(min_length=1, max_length=50)
|
|
14
|
+
sort_order: int = 0
|
|
15
|
+
description: Optional[str] = None
|
|
16
|
+
is_default: bool = False
|
|
17
|
+
allow_register: bool = False
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class RolePatch(BaseModel):
|
|
21
|
+
name: Optional[str] = Field(default=None, min_length=1, max_length=50)
|
|
22
|
+
sort_order: Optional[int] = None
|
|
23
|
+
description: Optional[str] = None
|
|
24
|
+
is_default: Optional[bool] = None
|
|
25
|
+
allow_register: Optional[bool] = None
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class TierBody(BaseModel):
|
|
29
|
+
code: str = Field(min_length=1, max_length=20)
|
|
30
|
+
name: str = Field(min_length=1, max_length=50)
|
|
31
|
+
sort_order: int = 0
|
|
32
|
+
badge_color: Optional[str] = None
|
|
33
|
+
description: Optional[str] = None
|
|
34
|
+
is_default: bool = False
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class TierPatch(BaseModel):
|
|
38
|
+
name: Optional[str] = Field(default=None, min_length=1, max_length=50)
|
|
39
|
+
sort_order: Optional[int] = None
|
|
40
|
+
badge_color: Optional[str] = None
|
|
41
|
+
description: Optional[str] = None
|
|
42
|
+
is_default: Optional[bool] = None
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class AdminUserPatch(BaseModel):
|
|
46
|
+
role: Optional[str] = None
|
|
47
|
+
is_admin: Optional[bool] = None
|
|
48
|
+
is_active: Optional[bool] = None
|
|
49
|
+
approval_status: Optional[ApprovalStatus] = None
|
|
50
|
+
tier_code: Optional[str] = None
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class RoleChangeCreate(BaseModel):
|
|
54
|
+
requested_role: str = Field(min_length=1, max_length=50)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class RoleChangeReview(BaseModel):
|
|
58
|
+
status: ApprovalStatus
|
account_kit/catalog.py
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
1
|
+
"""Role and tier catalog. Limits stay in the host; this only stores names."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from datetime import datetime, timezone
|
|
6
|
+
|
|
7
|
+
from fastapi import HTTPException
|
|
8
|
+
from sqlalchemy import func, select, update
|
|
9
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
10
|
+
|
|
11
|
+
from account_kit.models import Role, User, UserTier, UserTierAssignment
|
|
12
|
+
from account_kit.service import assert_role_code
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
async def _clear_default_role(db: AsyncSession) -> None:
|
|
16
|
+
await db.execute(update(Role).values(is_default=False))
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
async def _clear_default_tier(db: AsyncSession) -> None:
|
|
20
|
+
await db.execute(update(UserTier).values(is_default=False))
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
async def create_role(db: AsyncSession, body) -> Role:
|
|
24
|
+
code = assert_role_code(body.code)
|
|
25
|
+
if await db.get(Role, code):
|
|
26
|
+
raise HTTPException(status_code=409, detail="角色已存在")
|
|
27
|
+
if body.is_default:
|
|
28
|
+
await _clear_default_role(db)
|
|
29
|
+
row = Role(
|
|
30
|
+
code=code,
|
|
31
|
+
name=body.name.strip(),
|
|
32
|
+
sort_order=body.sort_order,
|
|
33
|
+
description=body.description,
|
|
34
|
+
is_default=body.is_default,
|
|
35
|
+
allow_register=body.allow_register,
|
|
36
|
+
)
|
|
37
|
+
db.add(row)
|
|
38
|
+
await db.commit()
|
|
39
|
+
await db.refresh(row)
|
|
40
|
+
return row
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
async def patch_role(db: AsyncSession, code: str, body) -> Role:
|
|
44
|
+
row = await db.get(Role, code)
|
|
45
|
+
if row is None:
|
|
46
|
+
raise HTTPException(status_code=404, detail="角色不存在")
|
|
47
|
+
if body.name is not None:
|
|
48
|
+
row.name = body.name.strip()
|
|
49
|
+
if body.sort_order is not None:
|
|
50
|
+
row.sort_order = body.sort_order
|
|
51
|
+
if body.description is not None:
|
|
52
|
+
row.description = body.description
|
|
53
|
+
if body.allow_register is not None:
|
|
54
|
+
row.allow_register = body.allow_register
|
|
55
|
+
if body.is_default is True:
|
|
56
|
+
await _clear_default_role(db)
|
|
57
|
+
row.is_default = True
|
|
58
|
+
elif body.is_default is False and row.is_default:
|
|
59
|
+
raise HTTPException(status_code=400, detail="请先把另一个角色设为默认")
|
|
60
|
+
await db.commit()
|
|
61
|
+
await db.refresh(row)
|
|
62
|
+
return row
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
async def delete_role(db: AsyncSession, code: str) -> None:
|
|
66
|
+
row = await db.get(Role, code)
|
|
67
|
+
if row is None:
|
|
68
|
+
raise HTTPException(status_code=404, detail="角色不存在")
|
|
69
|
+
if row.is_default:
|
|
70
|
+
raise HTTPException(status_code=400, detail="不能删除默认角色")
|
|
71
|
+
used = await db.scalar(select(func.count()).select_from(User).where(User.role == code))
|
|
72
|
+
if used:
|
|
73
|
+
raise HTTPException(status_code=409, detail="仍有用户使用该角色")
|
|
74
|
+
await db.delete(row)
|
|
75
|
+
await db.commit()
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
async def create_tier(db: AsyncSession, body) -> UserTier:
|
|
79
|
+
code = (body.code or "").strip()
|
|
80
|
+
if not code:
|
|
81
|
+
raise HTTPException(status_code=400, detail="等级代码无效")
|
|
82
|
+
if await db.get(UserTier, code):
|
|
83
|
+
raise HTTPException(status_code=409, detail="等级已存在")
|
|
84
|
+
if body.is_default:
|
|
85
|
+
await _clear_default_tier(db)
|
|
86
|
+
row = UserTier(
|
|
87
|
+
code=code,
|
|
88
|
+
name=body.name.strip(),
|
|
89
|
+
sort_order=body.sort_order,
|
|
90
|
+
badge_color=body.badge_color,
|
|
91
|
+
description=body.description,
|
|
92
|
+
is_default=body.is_default,
|
|
93
|
+
)
|
|
94
|
+
db.add(row)
|
|
95
|
+
await db.commit()
|
|
96
|
+
await db.refresh(row)
|
|
97
|
+
return row
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
async def patch_tier(db: AsyncSession, code: str, body) -> UserTier:
|
|
101
|
+
row = await db.get(UserTier, code)
|
|
102
|
+
if row is None:
|
|
103
|
+
raise HTTPException(status_code=404, detail="等级不存在")
|
|
104
|
+
if body.name is not None:
|
|
105
|
+
row.name = body.name.strip()
|
|
106
|
+
if body.sort_order is not None:
|
|
107
|
+
row.sort_order = body.sort_order
|
|
108
|
+
if body.badge_color is not None:
|
|
109
|
+
row.badge_color = body.badge_color
|
|
110
|
+
if body.description is not None:
|
|
111
|
+
row.description = body.description
|
|
112
|
+
if body.is_default is True:
|
|
113
|
+
await _clear_default_tier(db)
|
|
114
|
+
row.is_default = True
|
|
115
|
+
elif body.is_default is False and row.is_default:
|
|
116
|
+
raise HTTPException(status_code=400, detail="请先把另一个等级设为默认")
|
|
117
|
+
await db.commit()
|
|
118
|
+
await db.refresh(row)
|
|
119
|
+
return row
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
async def delete_tier(db: AsyncSession, code: str) -> None:
|
|
123
|
+
row = await db.get(UserTier, code)
|
|
124
|
+
if row is None:
|
|
125
|
+
raise HTTPException(status_code=404, detail="等级不存在")
|
|
126
|
+
if row.is_default:
|
|
127
|
+
raise HTTPException(status_code=400, detail="不能删除默认等级")
|
|
128
|
+
used = await db.scalar(
|
|
129
|
+
select(func.count()).select_from(UserTierAssignment).where(UserTierAssignment.tier_code == code)
|
|
130
|
+
)
|
|
131
|
+
if used:
|
|
132
|
+
raise HTTPException(status_code=409, detail="仍有用户属于该等级")
|
|
133
|
+
await db.delete(row)
|
|
134
|
+
await db.commit()
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
async def set_user_tier(db: AsyncSession, user: User, tier_code: str) -> None:
|
|
138
|
+
tier = await db.get(UserTier, tier_code)
|
|
139
|
+
if tier is None:
|
|
140
|
+
raise HTTPException(status_code=400, detail="等级不存在")
|
|
141
|
+
row = await db.get(UserTierAssignment, user.id)
|
|
142
|
+
if row is None:
|
|
143
|
+
db.add(UserTierAssignment(user_id=user.id, tier_code=tier.code))
|
|
144
|
+
else:
|
|
145
|
+
row.tier_code = tier.code
|
|
146
|
+
row.updated_at = datetime.now(timezone.utc)
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
async def apply_admin_user_patch(db, user: User, body) -> User:
|
|
150
|
+
if body.role is not None:
|
|
151
|
+
code = assert_role_code(body.role)
|
|
152
|
+
if await db.get(Role, code) is None:
|
|
153
|
+
raise HTTPException(status_code=400, detail="角色不存在")
|
|
154
|
+
user.role = code
|
|
155
|
+
if body.is_admin is not None:
|
|
156
|
+
user.is_admin = body.is_admin
|
|
157
|
+
if body.is_active is not None:
|
|
158
|
+
user.is_active = body.is_active
|
|
159
|
+
if body.approval_status is not None:
|
|
160
|
+
user.approval_status = body.approval_status
|
|
161
|
+
if body.approval_status == "approved" and user.approved_at is None:
|
|
162
|
+
user.approved_at = datetime.now(timezone.utc)
|
|
163
|
+
if body.tier_code is not None:
|
|
164
|
+
await set_user_tier(db, user, body.tier_code)
|
|
165
|
+
await db.commit()
|
|
166
|
+
await db.refresh(user)
|
|
167
|
+
return user
|
account_kit/config.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
from typing import Awaitable, Callable, Optional, Tuple
|
|
5
|
+
|
|
6
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class SmtpConfig:
|
|
11
|
+
host: str = ""
|
|
12
|
+
port: int = 587
|
|
13
|
+
user: str = ""
|
|
14
|
+
password: str = ""
|
|
15
|
+
from_email: str = ""
|
|
16
|
+
tls: bool = True
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
# to, purpose, code, language
|
|
20
|
+
Mailer = Callable[[str, str, str, str], Awaitable[None]]
|
|
21
|
+
UserHook = Callable[..., Awaitable[None]]
|
|
22
|
+
# request, username
|
|
23
|
+
BeforeLogin = Callable[[object, str], Awaitable[None]]
|
|
24
|
+
# request
|
|
25
|
+
BeforeRegister = Callable[[object], Awaitable[None]]
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass
|
|
29
|
+
class AccountKitConfig:
|
|
30
|
+
"""Runtime settings supplied by the host app. Nothing here is a quota."""
|
|
31
|
+
|
|
32
|
+
jwt_secret: str
|
|
33
|
+
jwt_algorithm: str = "HS256"
|
|
34
|
+
access_token_expire_minutes: int = 1440
|
|
35
|
+
# stateless | single_device
|
|
36
|
+
session_mode: str = "stateless"
|
|
37
|
+
session_active_hours: int = 12
|
|
38
|
+
require_approval: bool = False
|
|
39
|
+
email_domain_restriction: bool = False
|
|
40
|
+
allowed_email_domains: Tuple[str, ...] = ()
|
|
41
|
+
brand_name: str = "Account"
|
|
42
|
+
password_min_length: int = 6
|
|
43
|
+
smtp: SmtpConfig = field(default_factory=SmtpConfig)
|
|
44
|
+
# HMAC material for email codes and recovery codes. Falls back to jwt_secret.
|
|
45
|
+
secret_key: str = ""
|
|
46
|
+
# Fernet material for TOTP secrets. Falls back to secret_key, then jwt_secret.
|
|
47
|
+
file_encryption_key: str = ""
|
|
48
|
+
two_factor_enabled: bool = True
|
|
49
|
+
role_change_enabled: bool = False
|
|
50
|
+
trusted_device_days: int = 30
|
|
51
|
+
forbid_admin_like_usernames: bool = False
|
|
52
|
+
reserved_usernames: Tuple[str, ...] = ()
|
|
53
|
+
api_prefix: str = "/api/auth"
|
|
54
|
+
admin_prefix: str = "/api/admin/account"
|
|
55
|
+
mailer: Optional[Mailer] = None
|
|
56
|
+
before_login: Optional[BeforeLogin] = None
|
|
57
|
+
before_register: Optional[BeforeRegister] = None
|
|
58
|
+
on_registered: Optional[UserHook] = None
|
|
59
|
+
on_login: Optional[UserHook] = None
|
|
60
|
+
on_password_changed: Optional[UserHook] = None
|
|
61
|
+
on_deleted: Optional[UserHook] = None
|
|
62
|
+
|
|
63
|
+
def code_secret(self) -> str:
|
|
64
|
+
return self.secret_key or self.jwt_secret
|
|
65
|
+
|
|
66
|
+
def encryption_material(self) -> str:
|
|
67
|
+
return self.file_encryption_key or self.secret_key or self.jwt_secret
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
_config: Optional[AccountKitConfig] = None
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def set_config(config: AccountKitConfig) -> None:
|
|
74
|
+
global _config
|
|
75
|
+
_config = config
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def get_config() -> AccountKitConfig:
|
|
79
|
+
if _config is None:
|
|
80
|
+
raise RuntimeError("account-kit is not configured")
|
|
81
|
+
return _config
|
account_kit/crypto.py
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import base64
|
|
4
|
+
import hashlib
|
|
5
|
+
|
|
6
|
+
from cryptography.fernet import Fernet
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def fernet_for(material: str) -> Fernet:
|
|
10
|
+
key = base64.urlsafe_b64encode(hashlib.sha256(material.encode("utf-8")).digest())
|
|
11
|
+
return Fernet(key)
|
account_kit/db.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
from sqlalchemy import MetaData
|
|
2
|
+
from sqlalchemy.orm import DeclarativeBase
|
|
3
|
+
|
|
4
|
+
NAMING = {
|
|
5
|
+
"ix": "ix_%(column_0_label)s",
|
|
6
|
+
"uq": "uq_%(table_name)s_%(column_0_name)s",
|
|
7
|
+
"ck": "ck_%(table_name)s_%(constraint_name)s",
|
|
8
|
+
"fk": "fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s",
|
|
9
|
+
"pk": "pk_%(table_name)s",
|
|
10
|
+
}
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class Base(DeclarativeBase):
|
|
14
|
+
metadata = MetaData(naming_convention=NAMING)
|
account_kit/deps.py
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import uuid
|
|
4
|
+
from datetime import datetime, timedelta, timezone
|
|
5
|
+
from typing import Optional
|
|
6
|
+
|
|
7
|
+
from fastapi import Depends, HTTPException, Query, Request, status
|
|
8
|
+
from fastapi.security import OAuth2PasswordBearer
|
|
9
|
+
from jose import JWTError, jwt
|
|
10
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
11
|
+
|
|
12
|
+
from account_kit.config import get_config
|
|
13
|
+
from account_kit.service import get_user_by_id
|
|
14
|
+
|
|
15
|
+
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
async def get_db(request: Request):
|
|
19
|
+
agen = request.app.state.account_get_db()
|
|
20
|
+
session = await agen.__anext__()
|
|
21
|
+
try:
|
|
22
|
+
yield session
|
|
23
|
+
finally:
|
|
24
|
+
await agen.aclose()
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
async def user_from_access_token(db: AsyncSession, raw: Optional[str]):
|
|
28
|
+
credentials = HTTPException(
|
|
29
|
+
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
30
|
+
detail="Could not validate credentials",
|
|
31
|
+
headers={"WWW-Authenticate": "Bearer"},
|
|
32
|
+
)
|
|
33
|
+
if not raw:
|
|
34
|
+
raise credentials
|
|
35
|
+
config = get_config()
|
|
36
|
+
try:
|
|
37
|
+
payload = jwt.decode(raw, config.jwt_secret, algorithms=[config.jwt_algorithm])
|
|
38
|
+
sub = payload.get("sub")
|
|
39
|
+
user_id = uuid.UUID(str(sub))
|
|
40
|
+
except (JWTError, ValueError, TypeError):
|
|
41
|
+
raise credentials
|
|
42
|
+
|
|
43
|
+
user = await get_user_by_id(db, user_id)
|
|
44
|
+
if user is None:
|
|
45
|
+
raise credentials
|
|
46
|
+
if not user.is_active:
|
|
47
|
+
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="User is disabled")
|
|
48
|
+
|
|
49
|
+
if config.session_mode == "single_device":
|
|
50
|
+
sid = payload.get("sid")
|
|
51
|
+
if not sid or sid != user.current_session_id:
|
|
52
|
+
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="SESSION_REPLACED")
|
|
53
|
+
now = datetime.now(timezone.utc)
|
|
54
|
+
last = user.session_last_seen_at
|
|
55
|
+
if last is not None and last.tzinfo is None:
|
|
56
|
+
last = last.replace(tzinfo=timezone.utc)
|
|
57
|
+
if last is None or now - last > timedelta(seconds=60):
|
|
58
|
+
user.session_last_seen_at = now
|
|
59
|
+
await db.commit()
|
|
60
|
+
return user
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
async def get_current_user(
|
|
64
|
+
token: Optional[str] = Depends(oauth2_scheme),
|
|
65
|
+
query_token: Optional[str] = Query(None, alias="token"),
|
|
66
|
+
db: AsyncSession = Depends(get_db),
|
|
67
|
+
):
|
|
68
|
+
return await user_from_access_token(db, token or query_token)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
async def require_admin(current_user=Depends(get_current_user)):
|
|
72
|
+
if not current_user.is_admin:
|
|
73
|
+
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin required")
|
|
74
|
+
return current_user
|
account_kit/emailer.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import smtplib
|
|
5
|
+
from email.mime.multipart import MIMEMultipart
|
|
6
|
+
from email.mime.text import MIMEText
|
|
7
|
+
|
|
8
|
+
from jinja2 import Environment, FileSystemLoader
|
|
9
|
+
|
|
10
|
+
from account_kit.config import AccountKitConfig
|
|
11
|
+
|
|
12
|
+
_template_dir = os.path.join(os.path.dirname(__file__), "templates", "emails")
|
|
13
|
+
_jinja = Environment(loader=FileSystemLoader(_template_dir))
|
|
14
|
+
|
|
15
|
+
_ACTIONS = {
|
|
16
|
+
"zh": {
|
|
17
|
+
"register": "注册账户",
|
|
18
|
+
"reset_password": "重置密码",
|
|
19
|
+
"change_password": "修改密码",
|
|
20
|
+
"login_2fa": "登录验证",
|
|
21
|
+
"disable_2fa": "关闭两步验证",
|
|
22
|
+
},
|
|
23
|
+
"en": {
|
|
24
|
+
"register": "register",
|
|
25
|
+
"reset_password": "reset your password",
|
|
26
|
+
"change_password": "change your password",
|
|
27
|
+
"login_2fa": "sign in",
|
|
28
|
+
"disable_2fa": "turn off two-factor authentication",
|
|
29
|
+
},
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def render_code_email(config: AccountKitConfig, purpose: str, code: str, language: str) -> tuple[str, str]:
|
|
34
|
+
lang = language if language in ("zh", "en") else "zh"
|
|
35
|
+
action = _ACTIONS[lang].get(purpose, purpose)
|
|
36
|
+
brand = config.brand_name
|
|
37
|
+
if lang == "zh":
|
|
38
|
+
subject = f"【{brand}】验证码"
|
|
39
|
+
else:
|
|
40
|
+
subject = f"[{brand}] Verification code"
|
|
41
|
+
body = _jinja.get_template(f"verification_{lang}.html").render(action=action, code=code, brand=brand)
|
|
42
|
+
return subject, body
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def deliver_smtp(config: AccountKitConfig, to_email: str, subject: str, body: str) -> None:
|
|
46
|
+
cfg = config.smtp
|
|
47
|
+
if not cfg.user or not cfg.password:
|
|
48
|
+
return
|
|
49
|
+
msg = MIMEMultipart()
|
|
50
|
+
sender = cfg.from_email or cfg.user
|
|
51
|
+
msg["From"] = f"{config.brand_name} <{sender}>"
|
|
52
|
+
msg["To"] = to_email
|
|
53
|
+
msg["Subject"] = subject
|
|
54
|
+
msg.attach(MIMEText(body, "html", "utf-8"))
|
|
55
|
+
if cfg.port == 465:
|
|
56
|
+
server = smtplib.SMTP_SSL(cfg.host, cfg.port)
|
|
57
|
+
else:
|
|
58
|
+
server = smtplib.SMTP(cfg.host, cfg.port)
|
|
59
|
+
if cfg.tls:
|
|
60
|
+
server.starttls()
|
|
61
|
+
server.login(cfg.user, cfg.password)
|
|
62
|
+
server.send_message(msg)
|
|
63
|
+
server.quit()
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
async def send_code_email(config: AccountKitConfig, to_email: str, purpose: str, code: str, language: str = "zh") -> None:
|
|
67
|
+
if config.mailer is not None:
|
|
68
|
+
await config.mailer(to_email, purpose, code, language)
|
|
69
|
+
return
|
|
70
|
+
subject, body = render_code_email(config, purpose, code, language)
|
|
71
|
+
deliver_smtp(config, to_email, subject, body)
|