小程序后端接入微信登录与家长头像能力
This commit is contained in:
@@ -11,12 +11,19 @@ DB_USER=root
|
|||||||
DB_PASSWORD=change_me
|
DB_PASSWORD=change_me
|
||||||
DB_NAME=mini_program
|
DB_NAME=mini_program
|
||||||
DB_PATH=./data.db
|
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_ID=change_me
|
||||||
COS_SECRET_KEY=change_me
|
COS_SECRET_KEY=change_me
|
||||||
COS_REGION=ap-guangzhou
|
COS_REGION=ap-guangzhou
|
||||||
COS_BUCKET=voicemessage-1320289366
|
COS_BUCKET_MESSAGE=message-1320289366
|
||||||
COS_PUBLIC_BASE_URL=https://voicemessage-1320289366.cos.ap-guangzhou.myqcloud.com
|
COS_BUCKET_AVA=ava-1320289366
|
||||||
COS_AUDIO_PREFIX=voiceMessage/
|
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_SECRET=change_me_to_a_long_random_string
|
||||||
JWT_ALGORITHM=HS256
|
JWT_ALGORITHM=HS256
|
||||||
JWT_ACCESS_TOKEN_EXPIRE_MINUTES=60
|
JWT_ACCESS_TOKEN_EXPIRE_MINUTES=60
|
||||||
|
|||||||
@@ -2,8 +2,10 @@ from collections.abc import Mapping
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from sqlalchemy import text
|
from sqlalchemy import text
|
||||||
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
|
||||||
from app.dao import BaseDAO
|
from app.dao import BaseDAO
|
||||||
|
from app.db_compat import inserted_primary_key
|
||||||
|
|
||||||
|
|
||||||
class ParentDAO(BaseDAO):
|
class ParentDAO(BaseDAO):
|
||||||
@@ -23,8 +25,9 @@ class ParentDAO(BaseDAO):
|
|||||||
),
|
),
|
||||||
{"openid": openid, "unionid": unionid, "nickname": nickname, "avatar_url": avatar_url},
|
{"openid": openid, "unionid": unionid, "nickname": nickname, "avatar_url": avatar_url},
|
||||||
)
|
)
|
||||||
|
user_id = inserted_primary_key(result)
|
||||||
self.commit()
|
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]:
|
def get_by_id(self, user_id: int) -> Optional[Mapping]:
|
||||||
return (
|
return (
|
||||||
@@ -67,6 +70,49 @@ class ParentDAO(BaseDAO):
|
|||||||
)
|
)
|
||||||
self.commit()
|
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(
|
def upsert(
|
||||||
self,
|
self,
|
||||||
openid: str,
|
openid: str,
|
||||||
@@ -74,27 +120,18 @@ class ParentDAO(BaseDAO):
|
|||||||
nickname: Optional[str] = None,
|
nickname: Optional[str] = None,
|
||||||
avatar_url: Optional[str] = None,
|
avatar_url: Optional[str] = None,
|
||||||
) -> int:
|
) -> int:
|
||||||
from app.settings import settings
|
|
||||||
|
|
||||||
existing = self.get_by_openid(openid)
|
existing = self.get_by_openid(openid)
|
||||||
if existing:
|
if existing:
|
||||||
self.update(existing["user_id"], nickname, avatar_url)
|
self.update_from_wechat_login(existing["user_id"], unionid, nickname, avatar_url)
|
||||||
return existing["user_id"]
|
return int(existing["user_id"])
|
||||||
|
|
||||||
if settings.db_type == "sqlite":
|
try:
|
||||||
return self.create(openid, unionid, nickname, avatar_url)
|
return self.create(openid, unionid, nickname, avatar_url)
|
||||||
|
except IntegrityError:
|
||||||
self.db.execute(
|
self.db.rollback()
|
||||||
text(
|
existing = self.get_by_openid(openid)
|
||||||
"""
|
if not existing:
|
||||||
INSERT INTO parents (openid, unionid, nickname, avatar_url, status)
|
raise
|
||||||
VALUES (:openid, :unionid, :nickname, :avatar_url, 1)
|
if unionid or nickname or avatar_url:
|
||||||
ON DUPLICATE KEY UPDATE
|
self.update_from_wechat_login(existing["user_id"], unionid, nickname, avatar_url)
|
||||||
user_id = LAST_INSERT_ID(user_id),
|
return int(existing["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())
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ class Parent(Base):
|
|||||||
unionid: Mapped[Optional[str]] = mapped_column(String(64))
|
unionid: Mapped[Optional[str]] = mapped_column(String(64))
|
||||||
nickname: 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_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))
|
phone: Mapped[Optional[str]] = mapped_column(String(20))
|
||||||
status: Mapped[int] = mapped_column(Integer, server_default="1")
|
status: Mapped[int] = mapped_column(Integer, server_default="1")
|
||||||
created_at: Mapped[Optional[datetime]] = mapped_column(DateTime, server_default=text("CURRENT_TIMESTAMP"))
|
created_at: Mapped[Optional[datetime]] = mapped_column(DateTime, server_default=text("CURRENT_TIMESTAMP"))
|
||||||
|
|||||||
@@ -1,14 +1,18 @@
|
|||||||
import logging
|
import logging
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from app.service.parent import ParentService
|
from app.service.parent import ParentService
|
||||||
from app.service import get_db_session
|
from app.service import get_db_session
|
||||||
|
from app.service.avatar_storage import AvatarStorageError
|
||||||
|
from app.security import get_current_user_id
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
from service.parent import ParentService
|
from service.parent import ParentService
|
||||||
from service import get_db_session
|
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"])
|
router = APIRouter(prefix="/parents", tags=["parents"])
|
||||||
@@ -38,6 +42,11 @@ class ParentUpdateRequest(BaseModel):
|
|||||||
phone: str | None = None
|
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)
|
@router.post("", response_model=ParentResponse, status_code=status.HTTP_201_CREATED)
|
||||||
def create_parent(payload: ParentCreateRequest, request: Request, db=Depends(get_db_session)) -> ParentResponse:
|
def create_parent(payload: ParentCreateRequest, request: Request, db=Depends(get_db_session)) -> ParentResponse:
|
||||||
service = ParentService(db)
|
service = ParentService(db)
|
||||||
@@ -45,6 +54,43 @@ def create_parent(payload: ParentCreateRequest, request: Request, db=Depends(get
|
|||||||
return ParentResponse(**parent)
|
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)
|
@router.get("/{user_id}", response_model=ParentResponse)
|
||||||
def get_parent(user_id: int, request: Request, db=Depends(get_db_session)) -> ParentResponse:
|
def get_parent(user_id: int, request: Request, db=Depends(get_db_session)) -> ParentResponse:
|
||||||
service = ParentService(db)
|
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)
|
parent = service.update(user_id, payload.nickname, payload.avatar_url, payload.phone)
|
||||||
if not parent:
|
if not parent:
|
||||||
raise HTTPException(status_code=404, detail="parent not found")
|
raise HTTPException(status_code=404, detail="parent not found")
|
||||||
return ParentResponse(**parent)
|
return ParentResponse(**parent)
|
||||||
|
|||||||
@@ -1,17 +1,23 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel, Field
|
||||||
from sqlalchemy import text
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from app.db import get_db
|
from app.db import get_db
|
||||||
from app.security import create_access_token
|
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:
|
except ModuleNotFoundError:
|
||||||
from db import get_db
|
from db import get_db
|
||||||
from security import create_access_token
|
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"])
|
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||||
@@ -19,15 +25,14 @@ logger = logging.getLogger("app.auth")
|
|||||||
|
|
||||||
|
|
||||||
class LoginRequest(BaseModel):
|
class LoginRequest(BaseModel):
|
||||||
username: Optional[str] = None
|
code: str = Field(min_length=1, max_length=191)
|
||||||
password: Optional[str] = None
|
nickname: str | None = Field(default=None, max_length=64)
|
||||||
code: Optional[str] = None
|
avatar_url: str | None = Field(default=None, max_length=255)
|
||||||
nickname: Optional[str] = None
|
|
||||||
avatar_url: Optional[str] = None
|
|
||||||
|
|
||||||
|
|
||||||
class LoginResponse(BaseModel):
|
class LoginResponse(BaseModel):
|
||||||
access_token: str
|
access_token: str
|
||||||
|
token_type: str = "bearer"
|
||||||
expires_in: int
|
expires_in: int
|
||||||
user_id: int
|
user_id: int
|
||||||
|
|
||||||
@@ -37,51 +42,38 @@ def login(
|
|||||||
payload: LoginRequest,
|
payload: LoginRequest,
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
|
wechat_auth_service: WechatAuthService = Depends(get_wechat_auth_service),
|
||||||
) -> LoginResponse:
|
) -> LoginResponse:
|
||||||
logger.info(f"/auth/login called with code={payload.code[:20] if payload.code else None}...")
|
del request
|
||||||
identifier = payload.code or payload.username
|
logger.info("/auth/login called")
|
||||||
if not identifier:
|
|
||||||
raise HTTPException(status_code=400, detail="username or code is required")
|
|
||||||
|
|
||||||
with db.begin():
|
try:
|
||||||
row = (
|
wechat_session = wechat_auth_service.exchange_code(payload.code)
|
||||||
db.execute(
|
except WechatAuthError as exc:
|
||||||
text("SELECT user_id FROM parents WHERE openid = :openid"),
|
logger.warning(
|
||||||
{"openid": identifier},
|
"wechat login failed",
|
||||||
).mappings().first()
|
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:
|
parent_service = ParentService(db)
|
||||||
user_id = int(row["user_id"])
|
parent = parent_service.create(
|
||||||
logger.info(f"Existing user found: user_id={user_id}")
|
openid=wechat_session.openid,
|
||||||
if payload.nickname or payload.avatar_url:
|
unionid=wechat_session.unionid,
|
||||||
db.execute(
|
nickname=payload.nickname,
|
||||||
text(
|
avatar_url=payload.avatar_url,
|
||||||
"""
|
)
|
||||||
UPDATE parents
|
user_id = int(parent["user_id"])
|
||||||
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}")
|
|
||||||
|
|
||||||
access_token, expires_in = create_access_token(user_id=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)
|
return LoginResponse(access_token=access_token, expires_in=expires_in, user_id=user_id)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/logout")
|
@router.post("/logout")
|
||||||
def logout(request: Request):
|
def logout(request: Request):
|
||||||
return {"message": "logged out"}
|
return {"message": "logged out"}
|
||||||
|
|||||||
134
mini-program/app/service/avatar_storage.py
Normal file
134
mini-program/app/service/avatar_storage.py
Normal 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}"
|
||||||
|
)
|
||||||
@@ -2,12 +2,14 @@ from collections.abc import Mapping
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from app.dao.parent import ParentDAO
|
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:
|
class ParentService:
|
||||||
def __init__(self, db):
|
def __init__(self, db, avatar_storage: AvatarStorageService | None = None):
|
||||||
self.dao = ParentDAO(db)
|
self.dao = ParentDAO(db)
|
||||||
|
self.avatar_storage = avatar_storage or AvatarStorageService()
|
||||||
|
|
||||||
def create(
|
def create(
|
||||||
self,
|
self,
|
||||||
@@ -17,10 +19,10 @@ class ParentService:
|
|||||||
avatar_url: Optional[str] = None,
|
avatar_url: Optional[str] = None,
|
||||||
) -> Mapping:
|
) -> Mapping:
|
||||||
user_id = self.dao.upsert(openid, unionid, nickname, avatar_url)
|
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]:
|
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(
|
def update(
|
||||||
self,
|
self,
|
||||||
@@ -30,4 +32,71 @@ class ParentService:
|
|||||||
phone: Optional[str] = None,
|
phone: Optional[str] = None,
|
||||||
) -> Mapping:
|
) -> Mapping:
|
||||||
self.dao.update(user_id, nickname, avatar_url, phone)
|
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
|
||||||
|
|||||||
125
mini-program/app/service/wechat_login.py
Normal file
125
mini-program/app/service/wechat_login.py
Normal 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()
|
||||||
@@ -29,6 +29,31 @@ class Settings(BaseSettings):
|
|||||||
db_password: str = Field(default="", validation_alias="DB_PASSWORD")
|
db_password: str = Field(default="", validation_alias="DB_PASSWORD")
|
||||||
db_name: str = Field(default="mini_program", validation_alias="DB_NAME")
|
db_name: str = Field(default="mini_program", validation_alias="DB_NAME")
|
||||||
db_path: str = Field(default="./data.db", validation_alias="DB_PATH")
|
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(
|
jwt_secret: str = Field(
|
||||||
default="dev_only_change_jwt_secret",
|
default="dev_only_change_jwt_secret",
|
||||||
validation_alias="JWT_SECRET",
|
validation_alias="JWT_SECRET",
|
||||||
|
|||||||
@@ -2,9 +2,11 @@ fastapi==0.135.3
|
|||||||
uvicorn[standard]==0.44.0
|
uvicorn[standard]==0.44.0
|
||||||
pydantic-settings==2.8.1
|
pydantic-settings==2.8.1
|
||||||
python-dotenv==1.0.1
|
python-dotenv==1.0.1
|
||||||
|
python-multipart==0.0.26
|
||||||
SQLAlchemy==2.0.38
|
SQLAlchemy==2.0.38
|
||||||
PyMySQL==1.1.1
|
PyMySQL==1.1.1
|
||||||
PyJWT==2.10.1
|
PyJWT==2.10.1
|
||||||
httpx==0.28.1
|
httpx==0.28.1
|
||||||
|
cos-python-sdk-v5==1.9.41
|
||||||
pytest==8.3.4
|
pytest==8.3.4
|
||||||
pytest-asyncio==0.24.0
|
pytest-asyncio==0.24.0
|
||||||
|
|||||||
Reference in New Issue
Block a user