小程序后端接入微信登录与家长头像能力

This commit is contained in:
stu2not
2026-04-16 18:19:19 +08:00
parent db888df50f
commit 2526365af1
10 changed files with 515 additions and 77 deletions

View File

@@ -11,12 +11,19 @@ DB_USER=root
DB_PASSWORD=change_me
DB_NAME=mini_program
DB_PATH=./data.db
WECHAT_APP_ID=wx_change_me
WECHAT_APP_SECRET=change_me
WECHAT_API_BASE_URL=https://api.weixin.qq.com
WECHAT_HTTP_TIMEOUT_SECONDS=5
COS_SECRET_ID=change_me
COS_SECRET_KEY=change_me
COS_REGION=ap-guangzhou
COS_BUCKET=voicemessage-1320289366
COS_PUBLIC_BASE_URL=https://voicemessage-1320289366.cos.ap-guangzhou.myqcloud.com
COS_AUDIO_PREFIX=voiceMessage/
COS_BUCKET_MESSAGE=message-1320289366
COS_BUCKET_AVA=ava-1320289366
COS_PUBLIC_BASE_URL=
COS_AVATAR_PREFIX=avatars/
COS_AVATAR_URL_EXPIRE_SECONDS=86400
COS_AVATAR_MAX_BYTES=2097152
JWT_SECRET=change_me_to_a_long_random_string
JWT_ALGORITHM=HS256
JWT_ACCESS_TOKEN_EXPIRE_MINUTES=60

View File

@@ -2,8 +2,10 @@ from collections.abc import Mapping
from typing import Optional
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from app.dao import BaseDAO
from app.db_compat import inserted_primary_key
class ParentDAO(BaseDAO):
@@ -23,8 +25,9 @@ class ParentDAO(BaseDAO):
),
{"openid": openid, "unionid": unionid, "nickname": nickname, "avatar_url": avatar_url},
)
user_id = inserted_primary_key(result)
self.commit()
return int(self.db.execute(text("SELECT last_insert_rowid()")).scalar_one())
return user_id
def get_by_id(self, user_id: int) -> Optional[Mapping]:
return (
@@ -67,6 +70,49 @@ class ParentDAO(BaseDAO):
)
self.commit()
def set_avatar_file_key(self, user_id: int, avatar_file_key: str) -> None:
self.db.execute(
text(
"""
UPDATE parents
SET avatar_file_key = :avatar_file_key,
avatar_url = NULL
WHERE user_id = :user_id
"""
),
{"user_id": user_id, "avatar_file_key": avatar_file_key},
)
self.commit()
def update_from_wechat_login(
self,
user_id: int,
unionid: Optional[str] = None,
nickname: Optional[str] = None,
avatar_url: Optional[str] = None,
) -> None:
self.db.execute(
text(
"""
UPDATE parents
SET unionid = CASE
WHEN unionid IS NULL AND :unionid IS NOT NULL THEN :unionid
ELSE unionid
END,
nickname = COALESCE(:nickname, nickname),
avatar_url = COALESCE(:avatar_url, avatar_url)
WHERE user_id = :user_id
"""
),
{
"user_id": user_id,
"unionid": unionid,
"nickname": nickname,
"avatar_url": avatar_url,
},
)
self.commit()
def upsert(
self,
openid: str,
@@ -74,27 +120,18 @@ class ParentDAO(BaseDAO):
nickname: Optional[str] = None,
avatar_url: Optional[str] = None,
) -> int:
from app.settings import settings
existing = self.get_by_openid(openid)
if existing:
self.update(existing["user_id"], nickname, avatar_url)
return existing["user_id"]
self.update_from_wechat_login(existing["user_id"], unionid, nickname, avatar_url)
return int(existing["user_id"])
if settings.db_type == "sqlite":
try:
return self.create(openid, unionid, nickname, avatar_url)
self.db.execute(
text(
"""
INSERT INTO parents (openid, unionid, nickname, avatar_url, status)
VALUES (:openid, :unionid, :nickname, :avatar_url, 1)
ON DUPLICATE KEY UPDATE
user_id = LAST_INSERT_ID(user_id),
updated_at = CURRENT_TIMESTAMP
"""
),
{"openid": openid, "unionid": unionid, "nickname": nickname, "avatar_url": avatar_url},
)
self.commit()
return int(self.db.execute(text("SELECT LAST_INSERT_ID()")).scalar_one())
except IntegrityError:
self.db.rollback()
existing = self.get_by_openid(openid)
if not existing:
raise
if unionid or nickname or avatar_url:
self.update_from_wechat_login(existing["user_id"], unionid, nickname, avatar_url)
return int(existing["user_id"])

View File

@@ -30,6 +30,7 @@ class Parent(Base):
unionid: Mapped[Optional[str]] = mapped_column(String(64))
nickname: Mapped[Optional[str]] = mapped_column(String(64))
avatar_url: Mapped[Optional[str]] = mapped_column(String(255))
avatar_file_key: Mapped[Optional[str]] = mapped_column(String(255))
phone: Mapped[Optional[str]] = mapped_column(String(20))
status: Mapped[int] = mapped_column(Integer, server_default="1")
created_at: Mapped[Optional[datetime]] = mapped_column(DateTime, server_default=text("CURRENT_TIMESTAMP"))

View File

@@ -1,14 +1,18 @@
import logging
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status
from pydantic import BaseModel
try:
from app.service.parent import ParentService
from app.service import get_db_session
from app.service.avatar_storage import AvatarStorageError
from app.security import get_current_user_id
except ModuleNotFoundError:
from service.parent import ParentService
from service import get_db_session
from service.avatar_storage import AvatarStorageError
from security import get_current_user_id
router = APIRouter(prefix="/parents", tags=["parents"])
@@ -38,6 +42,11 @@ class ParentUpdateRequest(BaseModel):
phone: str | None = None
class AvatarDownloadResponse(BaseModel):
avatar_url: str
expires_in: int | None = None
@router.post("", response_model=ParentResponse, status_code=status.HTTP_201_CREATED)
def create_parent(payload: ParentCreateRequest, request: Request, db=Depends(get_db_session)) -> ParentResponse:
service = ParentService(db)
@@ -45,6 +54,43 @@ def create_parent(payload: ParentCreateRequest, request: Request, db=Depends(get
return ParentResponse(**parent)
@router.post("/me/avatar", response_model=ParentResponse)
def upload_my_avatar(
request: Request,
file: UploadFile = File(...),
current_user_id: int = Depends(get_current_user_id),
db=Depends(get_db_session),
) -> ParentResponse:
del request
service = ParentService(db)
try:
content = file.file.read()
parent = service.upload_avatar(
user_id=current_user_id,
filename=file.filename,
content_type=file.content_type,
content=content,
)
except AvatarStorageError as exc:
raise HTTPException(status_code=exc.status_code, detail=str(exc)) from exc
finally:
file.file.close()
if not parent:
raise HTTPException(status_code=404, detail="parent not found")
return ParentResponse(**parent)
@router.get("/{user_id}/avatar", response_model=AvatarDownloadResponse)
def get_parent_avatar(user_id: int, request: Request, db=Depends(get_db_session)) -> AvatarDownloadResponse:
del request
service = ParentService(db)
avatar = service.get_avatar_download(user_id)
if not avatar:
raise HTTPException(status_code=404, detail="avatar not found")
return AvatarDownloadResponse(**avatar)
@router.get("/{user_id}", response_model=ParentResponse)
def get_parent(user_id: int, request: Request, db=Depends(get_db_session)) -> ParentResponse:
service = ParentService(db)
@@ -60,4 +106,4 @@ def update_parent(user_id: int, payload: ParentUpdateRequest, request: Request,
parent = service.update(user_id, payload.nickname, payload.avatar_url, payload.phone)
if not parent:
raise HTTPException(status_code=404, detail="parent not found")
return ParentResponse(**parent)
return ParentResponse(**parent)

View File

@@ -1,17 +1,23 @@
import logging
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel
from sqlalchemy import text
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
try:
from app.db import get_db
from app.security import create_access_token
from app.service.parent import ParentService
from app.service.wechat_login import (
WechatAuthError,
WechatAuthService,
get_wechat_auth_service,
)
except ModuleNotFoundError:
from db import get_db
from security import create_access_token
from service.parent import ParentService
from service.wechat_login import WechatAuthError, WechatAuthService, get_wechat_auth_service
router = APIRouter(prefix="/auth", tags=["auth"])
@@ -19,15 +25,14 @@ logger = logging.getLogger("app.auth")
class LoginRequest(BaseModel):
username: Optional[str] = None
password: Optional[str] = None
code: Optional[str] = None
nickname: Optional[str] = None
avatar_url: Optional[str] = None
code: str = Field(min_length=1, max_length=191)
nickname: str | None = Field(default=None, max_length=64)
avatar_url: str | None = Field(default=None, max_length=255)
class LoginResponse(BaseModel):
access_token: str
token_type: str = "bearer"
expires_in: int
user_id: int
@@ -37,51 +42,38 @@ def login(
payload: LoginRequest,
request: Request,
db: Session = Depends(get_db),
wechat_auth_service: WechatAuthService = Depends(get_wechat_auth_service),
) -> LoginResponse:
logger.info(f"/auth/login called with code={payload.code[:20] if payload.code else None}...")
identifier = payload.code or payload.username
if not identifier:
raise HTTPException(status_code=400, detail="username or code is required")
del request
logger.info("/auth/login called")
with db.begin():
row = (
db.execute(
text("SELECT user_id FROM parents WHERE openid = :openid"),
{"openid": identifier},
).mappings().first()
try:
wechat_session = wechat_auth_service.exchange_code(payload.code)
except WechatAuthError as exc:
logger.warning(
"wechat login failed",
extra={
"event": "wechat_login_failed",
"status_code": exc.status_code,
"errcode": exc.errcode,
},
)
raise HTTPException(status_code=exc.status_code, detail=str(exc)) from exc
if row:
user_id = int(row["user_id"])
logger.info(f"Existing user found: user_id={user_id}")
if payload.nickname or payload.avatar_url:
db.execute(
text(
"""
UPDATE parents
SET nickname = COALESCE(:nickname, nickname),
avatar_url = COALESCE(:avatar_url, avatar_url)
WHERE user_id = :user_id
"""
),
{"user_id": user_id, "nickname": payload.nickname, "avatar_url": payload.avatar_url},
)
else:
logger.info("Creating new user...")
db.execute(
text(
"INSERT INTO parents (openid, nickname, avatar_url, status) VALUES (:openid, :nickname, :avatar_url, 1)"
),
{"openid": identifier, "nickname": payload.nickname, "avatar_url": payload.avatar_url},
)
user_id = int(db.execute(text("SELECT last_insert_rowid()")).scalar_one())
logger.info(f"New user created: user_id={user_id}")
parent_service = ParentService(db)
parent = parent_service.create(
openid=wechat_session.openid,
unionid=wechat_session.unionid,
nickname=payload.nickname,
avatar_url=payload.avatar_url,
)
user_id = int(parent["user_id"])
access_token, expires_in = create_access_token(user_id=user_id)
logger.info(f"Login succeeded: user_id={user_id}, expires_in={expires_in}")
logger.info("wechat login succeeded", extra={"event": "wechat_login_succeeded", "user_id": user_id})
return LoginResponse(access_token=access_token, expires_in=expires_in, user_id=user_id)
@router.post("/logout")
def logout(request: Request):
return {"message": "logged out"}
return {"message": "logged out"}

View File

@@ -0,0 +1,134 @@
from dataclasses import dataclass
from datetime import UTC, datetime
from pathlib import Path
from uuid import uuid4
try:
from qcloud_cos import CosConfig, CosS3Client
except ModuleNotFoundError: # pragma: no cover - exercised in runtime env
CosConfig = None
CosS3Client = None
from app.settings import settings
_CONTENT_TYPE_TO_EXT = {
"image/jpeg": "jpg",
"image/png": "png",
"image/webp": "webp",
}
_EXTENSION_ALIASES = {
".jpg": "jpg",
".jpeg": "jpg",
".png": "png",
".webp": "webp",
}
_EXT_TO_CONTENT_TYPE = {
"jpg": "image/jpeg",
"png": "image/png",
"webp": "image/webp",
}
class AvatarStorageError(Exception):
def __init__(self, message: str, status_code: int = 400) -> None:
super().__init__(message)
self.status_code = status_code
@dataclass(frozen=True)
class StoredAvatar:
file_key: str
class AvatarStorageService:
def __init__(self) -> None:
self._client = None
def upload_avatar(
self,
*,
user_id: int,
filename: str | None,
content_type: str | None,
content: bytes,
) -> StoredAvatar:
self._assert_ready()
normalized_ext = self._normalize_extension(filename=filename, content_type=content_type)
if not content:
raise AvatarStorageError("avatar file is empty")
if len(content) > settings.cos_avatar_max_bytes:
raise AvatarStorageError("avatar file too large", status_code=413)
key = self._build_key(user_id=user_id, extension=normalized_ext)
self._get_client().put_object(
Bucket=settings.cos_bucket_ava,
Body=content,
Key=key,
ContentType=_EXT_TO_CONTENT_TYPE[normalized_ext],
EnableMD5=False,
)
return StoredAvatar(file_key=key)
def get_avatar_url(self, file_key: str) -> str:
self._assert_ready()
if not file_key:
raise AvatarStorageError("avatar file key is required", status_code=500)
return self._get_client().get_presigned_url(
Bucket=settings.cos_bucket_ava,
Key=file_key,
Method="GET",
Expired=settings.cos_avatar_url_expire_seconds,
)
def delete_avatar(self, file_key: str) -> None:
self._assert_ready()
if not file_key:
return
self._get_client().delete_object(Bucket=settings.cos_bucket_ava, Key=file_key)
def _assert_ready(self) -> None:
if CosConfig is None or CosS3Client is None:
raise AvatarStorageError("COS SDK is not installed", status_code=500)
required_pairs = {
"COS_SECRET_ID": settings.cos_secret_id,
"COS_SECRET_KEY": settings.cos_secret_key,
"COS_REGION": settings.cos_region,
"COS_BUCKET_AVA": settings.cos_bucket_ava,
}
missing = [key for key, value in required_pairs.items() if not value]
if missing:
raise AvatarStorageError(
f"missing COS avatar config: {', '.join(missing)}",
status_code=500,
)
def _get_client(self):
if self._client is None:
config = CosConfig(
Region=settings.cos_region,
SecretId=settings.cos_secret_id,
SecretKey=settings.cos_secret_key,
Scheme="https",
)
self._client = CosS3Client(config)
return self._client
def _normalize_extension(self, *, filename: str | None, content_type: str | None) -> str:
if content_type in _CONTENT_TYPE_TO_EXT:
return _CONTENT_TYPE_TO_EXT[content_type]
suffix = Path(filename or "").suffix.lower()
if suffix in _EXTENSION_ALIASES:
return _EXTENSION_ALIASES[suffix]
raise AvatarStorageError("unsupported avatar file type", status_code=415)
def _build_key(self, *, user_id: int, extension: str) -> str:
prefix = settings.cos_avatar_prefix.strip("/") or "avatars"
now = datetime.now(UTC)
return (
f"{prefix}/{user_id}/{now.strftime('%Y/%m/%d')}/"
f"{uuid4().hex}.{extension}"
)

View File

@@ -2,12 +2,14 @@ from collections.abc import Mapping
from typing import Optional
from app.dao.parent import ParentDAO
from app.service import get_db_session
from app.service.avatar_storage import AvatarStorageService
from app.settings import settings
class ParentService:
def __init__(self, db):
def __init__(self, db, avatar_storage: AvatarStorageService | None = None):
self.dao = ParentDAO(db)
self.avatar_storage = avatar_storage or AvatarStorageService()
def create(
self,
@@ -17,10 +19,10 @@ class ParentService:
avatar_url: Optional[str] = None,
) -> Mapping:
user_id = self.dao.upsert(openid, unionid, nickname, avatar_url)
return self.dao.get_by_id(user_id)
return self.get(user_id)
def get(self, user_id: int) -> Optional[Mapping]:
return self.dao.get_by_id(user_id)
return self._present_parent(self.dao.get_by_id(user_id))
def update(
self,
@@ -30,4 +32,71 @@ class ParentService:
phone: Optional[str] = None,
) -> Mapping:
self.dao.update(user_id, nickname, avatar_url, phone)
return self.dao.get_by_id(user_id)
return self.get(user_id)
def upload_avatar(
self,
*,
user_id: int,
filename: str | None,
content_type: str | None,
content: bytes,
) -> Optional[Mapping]:
existing = self.dao.get_by_id(user_id)
if not existing:
return None
stored = self.avatar_storage.upload_avatar(
user_id=user_id,
filename=filename,
content_type=content_type,
content=content,
)
old_avatar_file_key = existing.get("avatar_file_key")
try:
self.dao.set_avatar_file_key(user_id, stored.file_key)
except Exception:
try:
self.avatar_storage.delete_avatar(stored.file_key)
except Exception:
pass
raise
if old_avatar_file_key and old_avatar_file_key != stored.file_key:
try:
self.avatar_storage.delete_avatar(old_avatar_file_key)
except Exception:
pass
return self.get(user_id)
def get_avatar_download(self, user_id: int) -> Optional[dict]:
parent = self.dao.get_by_id(user_id)
if not parent:
return None
avatar_file_key = parent.get("avatar_file_key")
if avatar_file_key:
return {
"avatar_url": self.avatar_storage.get_avatar_url(avatar_file_key),
"expires_in": settings.cos_avatar_url_expire_seconds,
}
avatar_url = parent.get("avatar_url")
if avatar_url:
return {
"avatar_url": avatar_url,
"expires_in": None,
}
return None
def _present_parent(self, parent: Optional[Mapping]) -> Optional[dict]:
if not parent:
return None
data = dict(parent)
avatar_file_key = data.get("avatar_file_key")
if avatar_file_key:
data["avatar_url"] = self.avatar_storage.get_avatar_url(avatar_file_key)
return data

View File

@@ -0,0 +1,125 @@
import logging
from dataclasses import dataclass
import httpx
try:
from app.settings import settings
except ModuleNotFoundError:
from settings import settings
logger = logging.getLogger("app.wechat_login")
INVALID_CODE_ERRCODES = {40029, 40163}
MISCONFIGURED_APP_ERRCODES = {40013, 40125}
@dataclass(frozen=True)
class WechatCodeSession:
openid: str
session_key: str
unionid: str | None = None
class WechatAuthError(Exception):
def __init__(self, detail: str, *, status_code: int, errcode: int | None = None):
super().__init__(detail)
self.status_code = status_code
self.errcode = errcode
class WechatAuthService:
def __init__(self, client: httpx.Client | None = None):
self._client = client
def exchange_code(self, code: str) -> WechatCodeSession:
if not settings.wechat_app_id or not settings.wechat_app_secret:
raise WechatAuthError("wechat login is not configured", status_code=503)
response = self._request_code2session(code)
data = self._parse_response_json(response)
errcode = self._parse_errcode(data.get("errcode"))
if errcode not in (None, 0):
errmsg = data.get("errmsg")
logger.warning(
"wechat code2session rejected login code",
extra={
"event": "wechat_code2session_rejected",
"errcode": errcode,
"errmsg": errmsg,
},
)
raise self._map_exchange_error(errcode)
openid = data.get("openid")
session_key = data.get("session_key")
unionid = data.get("unionid")
if not isinstance(openid, str) or not openid:
raise WechatAuthError("wechat login response missing openid", status_code=502)
if not isinstance(session_key, str) or not session_key:
raise WechatAuthError("wechat login response missing session_key", status_code=502)
if not isinstance(unionid, str) or not unionid:
unionid = None
return WechatCodeSession(openid=openid, session_key=session_key, unionid=unionid)
def _request_code2session(self, code: str) -> httpx.Response:
params = {
"appid": settings.wechat_app_id,
"secret": settings.wechat_app_secret,
"js_code": code,
"grant_type": "authorization_code",
}
client = self._client
owns_client = client is None
if client is None:
client = httpx.Client(
base_url=settings.wechat_api_base_url.rstrip("/"),
timeout=settings.wechat_http_timeout_seconds,
)
try:
response = client.get("/sns/jscode2session", params=params)
response.raise_for_status()
return response
except httpx.HTTPError as exc:
logger.warning(
"wechat code2session request failed",
extra={"event": "wechat_code2session_request_failed"},
)
raise WechatAuthError("wechat login service unavailable", status_code=502) from exc
finally:
if owns_client:
client.close()
def _parse_response_json(self, response: httpx.Response) -> dict:
try:
data = response.json()
except ValueError as exc:
raise WechatAuthError("invalid response from wechat login service", status_code=502) from exc
if not isinstance(data, dict):
raise WechatAuthError("invalid response from wechat login service", status_code=502)
return data
def _parse_errcode(self, value: object) -> int | None:
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def _map_exchange_error(self, errcode: int) -> WechatAuthError:
if errcode in INVALID_CODE_ERRCODES:
return WechatAuthError("invalid or expired wechat login code", status_code=401, errcode=errcode)
if errcode in MISCONFIGURED_APP_ERRCODES:
return WechatAuthError("wechat login is not configured correctly", status_code=503, errcode=errcode)
return WechatAuthError("wechat login service unavailable", status_code=502, errcode=errcode)
def get_wechat_auth_service() -> WechatAuthService:
return WechatAuthService()

View File

@@ -29,6 +29,31 @@ class Settings(BaseSettings):
db_password: str = Field(default="", validation_alias="DB_PASSWORD")
db_name: str = Field(default="mini_program", validation_alias="DB_NAME")
db_path: str = Field(default="./data.db", validation_alias="DB_PATH")
wechat_app_id: str = Field(default="", validation_alias="WECHAT_APP_ID")
wechat_app_secret: str = Field(default="", validation_alias="WECHAT_APP_SECRET")
wechat_api_base_url: str = Field(
default="https://api.weixin.qq.com",
validation_alias="WECHAT_API_BASE_URL",
)
wechat_http_timeout_seconds: float = Field(
default=5.0,
validation_alias="WECHAT_HTTP_TIMEOUT_SECONDS",
)
cos_secret_id: str = Field(default="", validation_alias="COS_SECRET_ID")
cos_secret_key: str = Field(default="", validation_alias="COS_SECRET_KEY")
cos_region: str = Field(default="", validation_alias="COS_REGION")
cos_bucket_message: str = Field(default="", validation_alias="COS_BUCKET_MESSAGE")
cos_bucket_ava: str = Field(default="", validation_alias="COS_BUCKET_AVA")
cos_public_base_url: str = Field(default="", validation_alias="COS_PUBLIC_BASE_URL")
cos_avatar_prefix: str = Field(default="avatars/", validation_alias="COS_AVATAR_PREFIX")
cos_avatar_url_expire_seconds: int = Field(
default=86400,
validation_alias="COS_AVATAR_URL_EXPIRE_SECONDS",
)
cos_avatar_max_bytes: int = Field(
default=2 * 1024 * 1024,
validation_alias="COS_AVATAR_MAX_BYTES",
)
jwt_secret: str = Field(
default="dev_only_change_jwt_secret",
validation_alias="JWT_SECRET",

View File

@@ -2,9 +2,11 @@ fastapi==0.135.3
uvicorn[standard]==0.44.0
pydantic-settings==2.8.1
python-dotenv==1.0.1
python-multipart==0.0.26
SQLAlchemy==2.0.38
PyMySQL==1.1.1
PyJWT==2.10.1
httpx==0.28.1
cos-python-sdk-v5==1.9.41
pytest==8.3.4
pytest-asyncio==0.24.0