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.
@@ -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)