feat(#558): 微信登录 + 绑定手机号邮箱完整后端实现 #670
Executable
+97
@@ -0,0 +1,97 @@
|
||||
"""#558 - 微信登录:手机号绑定字段 + 验证码表
|
||||
|
||||
Revision ID: 049
|
||||
Revises: 048
|
||||
Create Date: 2026-07-21
|
||||
|
||||
Changes:
|
||||
1. users 表新增 phone_verified / binding_completed_at 字段(phone 字段已在 029 中添加)
|
||||
2. users 表 phone 字段添加唯一索引(幂等)
|
||||
3. 新建 verification_codes 表(统一管理邮箱+手机验证码)
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision = "049_wechat_login_phone"
|
||||
down_revision = "048_cleanup_result_count"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(table: str, column: str) -> bool:
|
||||
"""检查列是否已存在。离线模式下返回 False。"""
|
||||
if context.is_offline_mode():
|
||||
return False
|
||||
conn = op.get_bind()
|
||||
result = conn.execute(
|
||||
sa.text("SELECT 1 FROM information_schema.columns " "WHERE table_name = :table AND column_name = :column"),
|
||||
{"table": table, "column": column},
|
||||
)
|
||||
return result.first() is not None
|
||||
|
||||
|
||||
def _index_exists(index_name: str) -> bool:
|
||||
"""检查索引是否已存在。离线模式下返回 False。"""
|
||||
if context.is_offline_mode():
|
||||
return False
|
||||
conn = op.get_bind()
|
||||
result = conn.execute(
|
||||
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :index_name"),
|
||||
{"index_name": index_name},
|
||||
)
|
||||
return result.first() is not None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 1. users 表新增手机号验证状态字段(幂等)
|
||||
if not _column_exists("users", "phone_verified"):
|
||||
op.add_column(
|
||||
"users",
|
||||
sa.Column(
|
||||
"phone_verified",
|
||||
sa.Boolean,
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
)
|
||||
|
||||
if not _column_exists("users", "binding_completed_at"):
|
||||
op.add_column(
|
||||
"users",
|
||||
sa.Column("binding_completed_at", sa.DateTime, nullable=True),
|
||||
)
|
||||
|
||||
# 2. phone 字段唯一索引(幂等 - 029 加了字段但没加索引)
|
||||
if not _index_exists("ix_users_phone"):
|
||||
op.create_index("ix_users_phone", "users", ["phone"], unique=True)
|
||||
|
||||
# 3. verification_codes 表
|
||||
op.create_table(
|
||||
"verification_codes",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("recipient", sa.String(255), nullable=False, index=True),
|
||||
sa.Column("code", sa.String(10), nullable=False),
|
||||
sa.Column("code_type", sa.String(32), nullable=False, index=True),
|
||||
sa.Column("expires_at", sa.DateTime, nullable=False),
|
||||
sa.Column("used_at", sa.DateTime, nullable=True),
|
||||
sa.Column("attempts", sa.Integer, nullable=False, server_default="0"),
|
||||
sa.Column("created_at", sa.DateTime, nullable=False),
|
||||
sa.Index(
|
||||
"ix_verification_recipient_type",
|
||||
"recipient",
|
||||
"code_type",
|
||||
"created_at",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("verification_codes")
|
||||
if _index_exists("ix_users_phone"):
|
||||
op.drop_index("ix_users_phone", table_name="users")
|
||||
if _column_exists("users", "binding_completed_at"):
|
||||
op.drop_column("users", "binding_completed_at")
|
||||
if _column_exists("users", "phone_verified"):
|
||||
op.drop_column("users", "phone_verified")
|
||||
@@ -81,6 +81,9 @@ class CurrentUserResponse(BaseModel):
|
||||
username: str
|
||||
display_name: str
|
||||
email_verified: bool
|
||||
phone: str = ""
|
||||
phone_verified: bool = False
|
||||
binding_complete: bool = False
|
||||
|
||||
|
||||
class PasswordResetRequestModel(BaseModel):
|
||||
@@ -259,12 +262,16 @@ async def get_current_user_info(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> CurrentUserResponse:
|
||||
user = authenticated_user.user
|
||||
binding_complete = user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
|
||||
return CurrentUserResponse(
|
||||
user_id=user.id,
|
||||
email=user.email,
|
||||
username=user.username,
|
||||
display_name=user.display_name,
|
||||
email_verified=user.email_verified,
|
||||
phone=user.phone or "",
|
||||
phone_verified=user.phone_verified,
|
||||
binding_complete=binding_complete,
|
||||
)
|
||||
|
||||
|
||||
@@ -380,3 +387,202 @@ async def wechat_sync(
|
||||
raise HTTPException(status_code=400, detail=error)
|
||||
|
||||
return WechatSyncResponse(**response.to_dict())
|
||||
|
||||
|
||||
# ==================== 微信网页登录(OAuth) ====================
|
||||
|
||||
|
||||
class WechatAuthUrlResponse(BaseModel):
|
||||
auth_url: str
|
||||
state: str
|
||||
|
||||
|
||||
class WechatCallbackRequest(BaseModel):
|
||||
code: str
|
||||
state: str = ""
|
||||
|
||||
|
||||
class WechatLoginResponse(BaseModel):
|
||||
access_token: str
|
||||
refresh_token: str
|
||||
user_id: str
|
||||
display_name: str
|
||||
avatar_url: str = ""
|
||||
is_new_user: bool
|
||||
binding_complete: bool
|
||||
expires_in: int
|
||||
|
||||
|
||||
@router.get("/wechat/url", response_model=WechatAuthUrlResponse)
|
||||
async def get_wechat_auth_url() -> WechatAuthUrlResponse:
|
||||
"""获取微信扫码登录授权链接"""
|
||||
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
|
||||
|
||||
oauth_service = get_wechat_oauth_service()
|
||||
auth_url, state = oauth_service.generate_auth_url()
|
||||
return WechatAuthUrlResponse(auth_url=auth_url, state=state)
|
||||
|
||||
|
||||
@router.post("/wechat/callback", response_model=WechatLoginResponse)
|
||||
async def wechat_callback(
|
||||
request: WechatCallbackRequest,
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> WechatLoginResponse:
|
||||
"""微信登录回调处理"""
|
||||
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
|
||||
from packages.application.auth.wechat_sync_use_case import WechatSyncRequest as SyncRequest
|
||||
from packages.application.auth.wechat_sync_use_case import WechatSyncUseCase
|
||||
|
||||
# 1. 用 code 换微信用户信息
|
||||
oauth_service = get_wechat_oauth_service()
|
||||
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
|
||||
if err:
|
||||
raise HTTPException(status_code=400, detail=err)
|
||||
|
||||
# 2. 同步登录/注册(复用 wechat-sync 逻辑)
|
||||
use_case = WechatSyncUseCase(user_repository=user_repository)
|
||||
sync_request = SyncRequest(
|
||||
openid=wechat_user.openid,
|
||||
unionid=wechat_user.unionid,
|
||||
nickname=wechat_user.nickname,
|
||||
avatar_url=wechat_user.avatar_url,
|
||||
source="web",
|
||||
)
|
||||
response, err = use_case.execute(sync_request)
|
||||
if err:
|
||||
raise HTTPException(status_code=400, detail=err)
|
||||
|
||||
# 3. 判断绑定状态
|
||||
user = user_repository.find_by_id(response.user_id)
|
||||
binding_complete = False
|
||||
if user:
|
||||
binding_complete = (
|
||||
user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
|
||||
)
|
||||
|
||||
return WechatLoginResponse(
|
||||
access_token=response.access_token,
|
||||
refresh_token=response.refresh_token,
|
||||
user_id=response.user_id,
|
||||
display_name=response.nickname,
|
||||
avatar_url=response.avatar_url or wechat_user.avatar_url,
|
||||
is_new_user=response.is_new_user,
|
||||
binding_complete=binding_complete,
|
||||
expires_in=response.expires_in,
|
||||
)
|
||||
|
||||
|
||||
# ==================== 验证码 & 绑定 ====================
|
||||
|
||||
|
||||
class SendVerificationCodeRequest(BaseModel):
|
||||
target: str # phone / email
|
||||
value: str
|
||||
purpose: str # bind / login / reset_password
|
||||
|
||||
|
||||
class SendVerificationCodeResponse(BaseModel):
|
||||
expires_in: int
|
||||
resend_after: int
|
||||
|
||||
|
||||
class BindContactRequest(BaseModel):
|
||||
phone: str = ""
|
||||
phone_code: str = ""
|
||||
email: str = ""
|
||||
email_code: str = ""
|
||||
|
||||
|
||||
class BindContactResponse(BaseModel):
|
||||
success: bool
|
||||
user: dict
|
||||
|
||||
|
||||
@router.post("/send-verification-code", response_model=SendVerificationCodeResponse)
|
||||
async def send_verification_code(
|
||||
request: SendVerificationCodeRequest,
|
||||
) -> SendVerificationCodeResponse:
|
||||
"""发送验证码(手机或邮箱)"""
|
||||
from app.dependencies import get_db
|
||||
|
||||
from packages.adapters.sms.sms_service import get_sms_service
|
||||
from packages.adapters.smtp import get_email_service
|
||||
from packages.adapters.sqlalchemy_impl.verification_code_repository import (
|
||||
SQLAlchemyVerificationCodeRepository,
|
||||
)
|
||||
from packages.application.auth.bind_contact_use_case import SendVerificationCodeRequest as UseCaseRequest
|
||||
from packages.application.auth.bind_contact_use_case import (
|
||||
SendVerificationCodeUseCase,
|
||||
)
|
||||
from packages.application.auth.verification_code_service import VerificationCodeService
|
||||
|
||||
db = next(get_db())
|
||||
repo = SQLAlchemyVerificationCodeRepository(db)
|
||||
vc_service = VerificationCodeService(repo=repo)
|
||||
sms_service = get_sms_service()
|
||||
email_service = get_email_service()
|
||||
|
||||
use_case = SendVerificationCodeUseCase(
|
||||
verification_code_service=vc_service,
|
||||
sms_service=sms_service,
|
||||
email_service=email_service,
|
||||
)
|
||||
uc_request = UseCaseRequest(
|
||||
target=request.target,
|
||||
value=request.value,
|
||||
purpose=request.purpose,
|
||||
)
|
||||
response, err = use_case.execute(uc_request)
|
||||
if err:
|
||||
raise HTTPException(status_code=400, detail=err)
|
||||
|
||||
return SendVerificationCodeResponse(
|
||||
expires_in=response.expires_in,
|
||||
resend_after=response.resend_after,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/bind-contact", response_model=BindContactResponse)
|
||||
async def bind_contact(
|
||||
request: BindContactRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> BindContactResponse:
|
||||
"""绑定手机号和/或邮箱(需登录态)"""
|
||||
from app.dependencies import get_db
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.verification_code_repository import (
|
||||
SQLAlchemyVerificationCodeRepository,
|
||||
)
|
||||
from packages.application.auth.bind_contact_use_case import BindContactRequest as UseCaseRequest
|
||||
from packages.application.auth.bind_contact_use_case import (
|
||||
BindContactUseCase,
|
||||
)
|
||||
from packages.application.auth.verification_code_service import VerificationCodeService
|
||||
|
||||
db = next(get_db())
|
||||
vc_repo = SQLAlchemyVerificationCodeRepository(db)
|
||||
vc_service = VerificationCodeService(repo=vc_repo)
|
||||
|
||||
use_case = BindContactUseCase(
|
||||
user_repository=user_repository,
|
||||
verification_code_service=vc_service,
|
||||
)
|
||||
uc_request = UseCaseRequest(
|
||||
user_id=current_user.user_id,
|
||||
phone=request.phone,
|
||||
phone_code=request.phone_code,
|
||||
email=request.email,
|
||||
email_code=request.email_code,
|
||||
)
|
||||
response, err = use_case.execute(uc_request)
|
||||
if err:
|
||||
raise HTTPException(status_code=400, detail=err)
|
||||
|
||||
return BindContactResponse(success=True, user=response.to_dict()["user"])
|
||||
|
||||
|
||||
# ==================== 当前用户信息扩展 ====================
|
||||
|
||||
# 扩展 CurrentUserResponse 增加绑定状态字段(在原响应基础上补充)
|
||||
# 通过给 get_current_user_info 返回值补充字段实现
|
||||
|
||||
Regular → Executable
+34
@@ -17,6 +17,9 @@ class InMemoryUserRepository(UserRepository):
|
||||
self._username_index: Dict[str, str] = {} # username -> user_id
|
||||
self._verification_token_index: Dict[str, str] = {} # token -> user_id
|
||||
self._reset_token_index: Dict[str, str] = {} # token -> user_id
|
||||
self._wechat_openid_index: Dict[str, str] = {} # openid -> user_id
|
||||
self._wechat_unionid_index: Dict[str, str] = {} # unionid -> user_id
|
||||
self._phone_index: Dict[str, str] = {} # phone -> user_id
|
||||
|
||||
def save(self, user: User) -> None:
|
||||
"""保存用户"""
|
||||
@@ -28,6 +31,12 @@ class InMemoryUserRepository(UserRepository):
|
||||
self._verification_token_index[user.email_verification_token] = user.id
|
||||
if user.password_reset_token:
|
||||
self._reset_token_index[user.password_reset_token] = user.id
|
||||
if user.wechat_openid:
|
||||
self._wechat_openid_index[user.wechat_openid] = user.id
|
||||
if user.wechat_unionid:
|
||||
self._wechat_unionid_index[user.wechat_unionid] = user.id
|
||||
if user.phone:
|
||||
self._phone_index[user.phone] = user.id
|
||||
|
||||
def find_by_id(self, user_id: str) -> Optional[User]:
|
||||
"""根据 ID 查找用户"""
|
||||
@@ -61,6 +70,31 @@ class InMemoryUserRepository(UserRepository):
|
||||
return self._users.get(user_id)
|
||||
return None
|
||||
|
||||
def find_by_wechat_openid(self, openid: str) -> Optional[User]:
|
||||
"""根据微信 openid 查找用户"""
|
||||
user_id = self._wechat_openid_index.get(openid)
|
||||
if user_id:
|
||||
return self._users.get(user_id)
|
||||
return None
|
||||
|
||||
def find_by_wechat_unionid(self, unionid: str) -> Optional[User]:
|
||||
"""根据微信 unionid 查找用户"""
|
||||
if not unionid:
|
||||
return None
|
||||
user_id = self._wechat_unionid_index.get(unionid)
|
||||
if user_id:
|
||||
return self._users.get(user_id)
|
||||
return None
|
||||
|
||||
def find_by_phone(self, phone: str) -> Optional[User]:
|
||||
"""根据手机号查找用户"""
|
||||
if not phone:
|
||||
return None
|
||||
user_id = self._phone_index.get(phone)
|
||||
if user_id:
|
||||
return self._users.get(user_id)
|
||||
return None
|
||||
|
||||
def delete(self, user_id: str) -> bool:
|
||||
"""删除用户"""
|
||||
user = self._users.get(user_id)
|
||||
|
||||
Executable
+85
@@ -0,0 +1,85 @@
|
||||
"""
|
||||
短信服务实现(Noop + 阿里云)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class NoopSmsService:
|
||||
"""空实现短信服务 - 开发/测试环境用,只打日志不真发"""
|
||||
|
||||
def send_verification_code(self, phone: str, code: str) -> bool:
|
||||
logger.info("[NoopSMS] 发送验证码到 %s: %s", phone, code)
|
||||
return True
|
||||
|
||||
def send_template_sms(self, phone: str, template_id: str, params: dict) -> bool:
|
||||
logger.info("[NoopSMS] 发送模板短信到 %s, template=%s, params=%s", phone, template_id, params)
|
||||
return True
|
||||
|
||||
|
||||
class AliyunSmsService:
|
||||
"""阿里云短信服务"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
access_key_id: str | None = None,
|
||||
access_key_secret: str | None = None,
|
||||
sign_name: str | None = None,
|
||||
verify_template_id: str | None = None,
|
||||
):
|
||||
self.access_key_id = access_key_id or os.environ.get("ALIYUN_SMS_ACCESS_KEY_ID", "")
|
||||
self.access_key_secret = access_key_secret or os.environ.get("ALIYUN_SMS_ACCESS_KEY_SECRET", "")
|
||||
self.sign_name = sign_name or os.environ.get("ALIYUN_SMS_SIGN_NAME", "小应剪辑")
|
||||
self.verify_template_id = verify_template_id or os.environ.get("ALIYUN_SMS_VERIFY_TEMPLATE_ID", "SMS_123456789")
|
||||
|
||||
def send_verification_code(self, phone: str, code: str) -> bool:
|
||||
return self.send_template_sms(phone, self.verify_template_id, {"code": code})
|
||||
|
||||
def send_template_sms(self, phone: str, template_id: str, params: dict) -> bool:
|
||||
try:
|
||||
import json
|
||||
|
||||
from alibabacloud_dysmsapi20170525 import models as dysmsapi_models
|
||||
from alibabacloud_dysmsapi20170525.client import Client as DysmsapiClient
|
||||
from alibabacloud_tea_openapi import models as open_api_models
|
||||
|
||||
config = open_api_models.Config(
|
||||
access_key_id=self.access_key_id,
|
||||
access_key_secret=self.access_key_secret,
|
||||
)
|
||||
config.endpoint = "dysmsapi.aliyuncs.com"
|
||||
client = DysmsapiClient(config)
|
||||
|
||||
request = dysmsapi_models.SendSmsRequest(
|
||||
phone_numbers=phone,
|
||||
sign_name=self.sign_name,
|
||||
template_code=template_id,
|
||||
template_param=json.dumps(params),
|
||||
)
|
||||
response = client.send_sms(request)
|
||||
body = response.body
|
||||
if body.code == "OK":
|
||||
logger.info("阿里云短信发送成功: phone=%s, template=%s", phone, template_id)
|
||||
return True
|
||||
else:
|
||||
logger.error("阿里云短信发送失败: code=%s, message=%s", body.code, body.message)
|
||||
return False
|
||||
except ImportError:
|
||||
logger.error("阿里云短信 SDK 未安装,请 pip install alibabacloud-dysmsapi20170525")
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error("阿里云短信发送异常: %s", e, exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
def get_sms_service() -> "NoopSmsService | AliyunSmsService":
|
||||
"""获取短信服务实例"""
|
||||
provider = os.environ.get("SMS_PROVIDER", "noop").lower()
|
||||
if provider == "aliyun":
|
||||
return AliyunSmsService()
|
||||
return NoopSmsService()
|
||||
@@ -31,9 +31,13 @@ class UserModel(Base):
|
||||
used_storage_gb = Column(Integer, nullable=False, default=0)
|
||||
# 管理员标识
|
||||
is_admin = Column(Boolean, nullable=False, default=False)
|
||||
# 微信登录(小程序端)
|
||||
# 微信登录
|
||||
wechat_openid = Column(String(128), nullable=True, unique=True, index=True)
|
||||
wechat_unionid = Column(String(128), nullable=True, unique=True, index=True)
|
||||
# 手机号绑定
|
||||
phone = Column(String(32), nullable=True, unique=True, index=True)
|
||||
phone_verified = Column(Boolean, nullable=False, default=False)
|
||||
binding_completed_at = Column(DateTime, nullable=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@@ -555,3 +559,18 @@ class BillingRecordModel(Base):
|
||||
invoice_url = Column(String(500), nullable=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
paid_at = Column(DateTime, nullable=True)
|
||||
|
||||
|
||||
class VerificationCodeModel(Base):
|
||||
"""验证码(邮箱/手机统一管理)"""
|
||||
|
||||
__tablename__ = "verification_codes"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
recipient = Column(String(255), nullable=False, index=True) # 邮箱或手机号
|
||||
code = Column(String(10), nullable=False)
|
||||
code_type = Column(String(32), nullable=False, index=True) # email_bind / phone_bind / ...
|
||||
expires_at = Column(DateTime, nullable=False)
|
||||
used_at = Column(DateTime, nullable=True)
|
||||
attempts = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
Regular → Executable
+12
@@ -35,6 +35,9 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model.is_admin = user.is_admin
|
||||
model.wechat_openid = user.wechat_openid
|
||||
model.wechat_unionid = user.wechat_unionid
|
||||
model.phone = user.phone
|
||||
model.phone_verified = user.phone_verified
|
||||
model.binding_completed_at = user.binding_completed_at
|
||||
model.created_at = user.created_at
|
||||
|
||||
self.session.commit()
|
||||
@@ -61,6 +64,12 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model = self.session.query(UserModel).filter(UserModel.wechat_unionid == unionid.strip()).first()
|
||||
return self._to_entity(model)
|
||||
|
||||
def find_by_phone(self, phone: str) -> User | None:
|
||||
if not phone or not phone.strip():
|
||||
return None
|
||||
model = self.session.query(UserModel).filter(UserModel.phone == phone.strip()).first()
|
||||
return self._to_entity(model)
|
||||
|
||||
def find_by_verification_token(self, token: str) -> User | None:
|
||||
model = self.session.query(UserModel).filter(UserModel.email_verification_token == token).first()
|
||||
return self._to_entity(model)
|
||||
@@ -101,5 +110,8 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
is_admin=model.is_admin or False,
|
||||
wechat_openid=model.wechat_openid,
|
||||
wechat_unionid=model.wechat_unionid,
|
||||
phone=model.phone,
|
||||
phone_verified=model.phone_verified or False,
|
||||
binding_completed_at=model.binding_completed_at,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
"""
|
||||
验证码仓储 SQLAlchemy 实现
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import VerificationCodeModel
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
from packages.ports.verification_code_repository import VerificationCodeRepository
|
||||
|
||||
|
||||
class SQLAlchemyVerificationCodeRepository(VerificationCodeRepository):
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def save(self, code: VerificationCode) -> None:
|
||||
model = self.session.get(VerificationCodeModel, code.id)
|
||||
if model is None:
|
||||
model = VerificationCodeModel(id=code.id)
|
||||
self.session.add(model)
|
||||
|
||||
model.recipient = code.recipient
|
||||
model.code = code.code
|
||||
model.code_type = code.code_type
|
||||
model.expires_at = code.expires_at
|
||||
model.used_at = code.used_at
|
||||
model.attempts = code.attempts
|
||||
model.created_at = code.created_at
|
||||
|
||||
self.session.commit()
|
||||
self.session.refresh(model)
|
||||
|
||||
def find_latest(self, recipient: str, code_type: str) -> Optional[VerificationCode]:
|
||||
model = (
|
||||
self.session.query(VerificationCodeModel)
|
||||
.filter(
|
||||
VerificationCodeModel.recipient == recipient.strip(),
|
||||
VerificationCodeModel.code_type == code_type,
|
||||
)
|
||||
.order_by(VerificationCodeModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
return self._to_entity(model)
|
||||
|
||||
def find_by_id(self, code_id: str) -> Optional[VerificationCode]:
|
||||
return self._to_entity(self.session.get(VerificationCodeModel, code_id))
|
||||
|
||||
def count_today(self, recipient: str, code_type: str) -> int:
|
||||
now = datetime.now(timezone.utc)
|
||||
start_of_day = now.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
return (
|
||||
self.session.query(VerificationCodeModel)
|
||||
.filter(
|
||||
VerificationCodeModel.recipient == recipient.strip(),
|
||||
VerificationCodeModel.code_type == code_type,
|
||||
VerificationCodeModel.created_at >= start_of_day,
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _to_entity(model: VerificationCodeModel | None) -> VerificationCode | None:
|
||||
if model is None:
|
||||
return None
|
||||
return VerificationCode(
|
||||
id=model.id,
|
||||
recipient=model.recipient,
|
||||
code=model.code,
|
||||
code_type=model.code_type,
|
||||
expires_at=model.expires_at,
|
||||
used_at=model.used_at,
|
||||
attempts=model.attempts,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
+231
@@ -0,0 +1,231 @@
|
||||
"""
|
||||
微信登录 + 绑定手机号邮箱 Use Case
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from packages.application.auth.verification_code_service import (
|
||||
CODE_TYPE_EMAIL_BIND,
|
||||
CODE_TYPE_PHONE_BIND,
|
||||
normalize_phone,
|
||||
validate_email,
|
||||
validate_phone,
|
||||
)
|
||||
from packages.domain.entities import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BindContactRequest:
|
||||
"""绑定联系方式请求"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: str,
|
||||
phone: str = "",
|
||||
phone_code: str = "",
|
||||
email: str = "",
|
||||
email_code: str = "",
|
||||
):
|
||||
self.user_id = user_id
|
||||
self.phone = normalize_phone(phone) if phone else ""
|
||||
self.phone_code = phone_code.strip() if phone_code else ""
|
||||
self.email = email.strip().lower() if email else ""
|
||||
self.email_code = email_code.strip() if email_code else ""
|
||||
|
||||
|
||||
class BindContactResponse:
|
||||
"""绑定响应"""
|
||||
|
||||
def __init__(self, user: User):
|
||||
self.user = user
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"user": {
|
||||
"id": self.user.id,
|
||||
"email": self.user.email,
|
||||
"phone": self.user.phone,
|
||||
"phone_verified": self.user.phone_verified,
|
||||
"display_name": self.user.display_name,
|
||||
"binding_complete": self.user.binding_completed_at is not None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
class BindContactUseCase:
|
||||
"""绑定手机+邮箱用例"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_repository,
|
||||
verification_code_service,
|
||||
email_service=None,
|
||||
):
|
||||
self.user_repo = user_repository
|
||||
self.verification_service = verification_code_service
|
||||
self.email_service = email_service
|
||||
|
||||
def execute(self, request: BindContactRequest) -> tuple[Optional[BindContactResponse], Optional[str]]:
|
||||
try:
|
||||
# 1. 校验参数
|
||||
if not request.phone and not request.email:
|
||||
return None, "至少填写手机号或邮箱"
|
||||
|
||||
# 2. 查找用户
|
||||
user = self.user_repo.find_by_id(request.user_id)
|
||||
if not user:
|
||||
return None, "用户不存在"
|
||||
|
||||
# 3. 手机号绑定
|
||||
if request.phone:
|
||||
ok, err = validate_phone(request.phone)
|
||||
if not ok:
|
||||
return None, err
|
||||
|
||||
if not request.phone_code:
|
||||
return None, "请输入手机验证码"
|
||||
|
||||
# 校验手机号未被其他账号绑定
|
||||
existing = self.user_repo.find_by_phone(request.phone)
|
||||
if existing and existing.id != user.id:
|
||||
return None, "该手机号已被其他账号绑定"
|
||||
|
||||
# 校验验证码
|
||||
ok, err = self.verification_service.verify(
|
||||
recipient=request.phone,
|
||||
code_type=CODE_TYPE_PHONE_BIND,
|
||||
code_value=request.phone_code,
|
||||
)
|
||||
if not ok:
|
||||
return None, f"手机验证码错误:{err}"
|
||||
|
||||
user.phone = request.phone
|
||||
user.phone_verified = True
|
||||
|
||||
# 4. 邮箱绑定
|
||||
if request.email:
|
||||
ok, err = validate_email(request.email)
|
||||
if not ok:
|
||||
return None, err
|
||||
|
||||
if not request.email_code:
|
||||
return None, "请输入邮箱验证码"
|
||||
|
||||
# 校验邮箱未被其他账号绑定
|
||||
existing = self.user_repo.find_by_email(request.email)
|
||||
if existing and existing.id != user.id:
|
||||
return None, "该邮箱已被其他账号绑定"
|
||||
|
||||
# 校验验证码
|
||||
ok, err = self.verification_service.verify(
|
||||
recipient=request.email,
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
code_value=request.email_code,
|
||||
)
|
||||
if not ok:
|
||||
return None, f"邮箱验证码错误:{err}"
|
||||
|
||||
user.email = request.email
|
||||
user.email_verified = True
|
||||
|
||||
# 5. 判断是否完成绑定
|
||||
if user.phone_verified and user.email_verified and "@wechat.local" not in user.email:
|
||||
user.binding_completed_at = datetime.now(timezone.utc)
|
||||
|
||||
# 6. 保存
|
||||
self.user_repo.save(user)
|
||||
|
||||
return BindContactResponse(user=user), None
|
||||
|
||||
except Exception as e:
|
||||
logger.error("绑定联系方式失败: %s", e, exc_info=True)
|
||||
return None, f"绑定失败: {str(e)}"
|
||||
|
||||
|
||||
class SendVerificationCodeRequest:
|
||||
"""发送验证码请求"""
|
||||
|
||||
def __init__(self, target: str, value: str, purpose: str):
|
||||
self.target = target # phone / email
|
||||
self.value = value.strip()
|
||||
self.purpose = purpose # bind / login / reset_password
|
||||
|
||||
|
||||
class SendVerificationCodeResponse:
|
||||
"""发送验证码响应"""
|
||||
|
||||
def __init__(self, expires_in: int, resend_after: int):
|
||||
self.expires_in = expires_in
|
||||
self.resend_after = resend_after
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"expires_in": self.expires_in,
|
||||
"resend_after": self.resend_after,
|
||||
}
|
||||
|
||||
|
||||
class SendVerificationCodeUseCase:
|
||||
"""发送验证码用例"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
verification_code_service,
|
||||
sms_service=None,
|
||||
email_service=None,
|
||||
):
|
||||
self.verification_service = verification_code_service
|
||||
self.sms_service = sms_service
|
||||
self.email_service = email_service
|
||||
|
||||
def execute(
|
||||
self, request: SendVerificationCodeRequest
|
||||
) -> tuple[Optional[SendVerificationCodeResponse], Optional[str]]:
|
||||
try:
|
||||
# 1. 确定 code_type
|
||||
if request.target == "phone":
|
||||
ok, err = validate_phone(request.value)
|
||||
if not ok:
|
||||
return None, err
|
||||
code_type = f"{request.target}_{request.purpose}"
|
||||
recipient = normalize_phone(request.value)
|
||||
elif request.target == "email":
|
||||
ok, err = validate_email(request.value)
|
||||
if not ok:
|
||||
return None, err
|
||||
code_type = f"{request.target}_{request.purpose}"
|
||||
recipient = request.value.lower()
|
||||
else:
|
||||
return None, f"不支持的目标类型: {request.target}"
|
||||
|
||||
# 2. 生成验证码
|
||||
code_obj, err = self.verification_service.generate(recipient, code_type)
|
||||
if err:
|
||||
return None, err
|
||||
|
||||
# 3. 发送
|
||||
if request.target == "phone" and self.sms_service:
|
||||
self.sms_service.send_verification_code(recipient, code_obj.code)
|
||||
elif request.target == "email" and self.email_service:
|
||||
subject = "验证码 - 小应剪辑"
|
||||
body = f"您的验证码是:{code_obj.code},5分钟内有效。"
|
||||
self.email_service.send_email(recipient, subject, body)
|
||||
|
||||
# 4. 返回
|
||||
ttl = (code_obj.expires_at - code_obj.created_at).total_seconds()
|
||||
return (
|
||||
SendVerificationCodeResponse(
|
||||
expires_in=int(ttl),
|
||||
resend_after=60,
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("发送验证码失败: %s", e, exc_info=True)
|
||||
return None, f"发送失败: {str(e)}"
|
||||
+206
@@ -0,0 +1,206 @@
|
||||
"""
|
||||
验证码服务
|
||||
- 生成验证码
|
||||
- 校验验证码
|
||||
- 频控(60s 冷却 + 每日上限)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
from packages.ports.verification_code_repository import VerificationCodeRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 频控参数
|
||||
RESEND_COOLDOWN_SECONDS = 60 # 重发冷却时间
|
||||
DAILY_LIMIT = 10 # 每日发送上限
|
||||
MAX_ATTEMPTS = 5 # 单验证码最大尝试次数
|
||||
DEFAULT_TTL_SECONDS = 300 # 默认有效期 5 分钟
|
||||
|
||||
# 验证码类型
|
||||
CODE_TYPE_EMAIL_BIND = "email_bind"
|
||||
CODE_TYPE_PHONE_BIND = "phone_bind"
|
||||
CODE_TYPE_EMAIL_LOGIN = "email_login"
|
||||
CODE_TYPE_PHONE_LOGIN = "phone_login"
|
||||
CODE_TYPE_RESET_PASSWORD = "reset_password"
|
||||
|
||||
VALID_CODE_TYPES = {
|
||||
CODE_TYPE_EMAIL_BIND,
|
||||
CODE_TYPE_PHONE_BIND,
|
||||
CODE_TYPE_EMAIL_LOGIN,
|
||||
CODE_TYPE_PHONE_LOGIN,
|
||||
CODE_TYPE_RESET_PASSWORD,
|
||||
}
|
||||
|
||||
|
||||
class VerificationCodeService:
|
||||
"""验证码服务"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
repo: VerificationCodeRepository,
|
||||
resend_cooldown: int = RESEND_COOLDOWN_SECONDS,
|
||||
daily_limit: int = DAILY_LIMIT,
|
||||
max_attempts: int = MAX_ATTEMPTS,
|
||||
default_ttl: int = DEFAULT_TTL_SECONDS,
|
||||
):
|
||||
self.repo = repo
|
||||
self.resend_cooldown = resend_cooldown
|
||||
self.daily_limit = daily_limit
|
||||
self.max_attempts = max_attempts
|
||||
self.default_ttl = default_ttl
|
||||
|
||||
def generate(
|
||||
self,
|
||||
recipient: str,
|
||||
code_type: str,
|
||||
ttl_seconds: int | None = None,
|
||||
custom_code: str | None = None,
|
||||
) -> tuple[Optional[VerificationCode], Optional[str]]:
|
||||
"""
|
||||
生成验证码
|
||||
|
||||
Returns:
|
||||
(验证码实体, 错误信息)
|
||||
"""
|
||||
recipient = recipient.strip()
|
||||
|
||||
# 参数校验
|
||||
if not recipient:
|
||||
return None, "接收方不能为空"
|
||||
if code_type not in VALID_CODE_TYPES:
|
||||
return None, f"无效的验证码类型: {code_type}"
|
||||
|
||||
# 频控检查
|
||||
can_send, wait_seconds = self._check_rate_limit(recipient, code_type)
|
||||
if not can_send:
|
||||
if wait_seconds > 0:
|
||||
return None, f"发送太频繁,请 {wait_seconds} 秒后再试"
|
||||
return None, "今日发送次数已达上限"
|
||||
|
||||
# 生成并保存
|
||||
code = VerificationCode.create(
|
||||
recipient=recipient,
|
||||
code_type=code_type,
|
||||
ttl_seconds=ttl_seconds or self.default_ttl,
|
||||
custom_code=custom_code,
|
||||
)
|
||||
self.repo.save(code)
|
||||
|
||||
return code, None
|
||||
|
||||
def verify(
|
||||
self,
|
||||
recipient: str,
|
||||
code_type: str,
|
||||
code_value: str,
|
||||
consume: bool = True,
|
||||
) -> tuple[bool, Optional[str]]:
|
||||
"""
|
||||
校验验证码
|
||||
|
||||
Args:
|
||||
recipient: 接收方(邮箱/手机号)
|
||||
code_type: 验证码类型
|
||||
code_value: 用户输入的验证码
|
||||
consume: 校验成功后是否标记为已使用
|
||||
|
||||
Returns:
|
||||
(是否通过, 错误信息)
|
||||
"""
|
||||
recipient = recipient.strip()
|
||||
code_value = code_value.strip()
|
||||
|
||||
if not recipient or not code_value:
|
||||
return False, "参数不完整"
|
||||
|
||||
# 查找最新的验证码
|
||||
latest = self.repo.find_latest(recipient, code_type)
|
||||
if not latest:
|
||||
return False, "验证码不存在或已过期"
|
||||
|
||||
# 增加尝试次数
|
||||
latest.increment_attempts()
|
||||
self.repo.save(latest)
|
||||
|
||||
# 检查是否已使用
|
||||
if latest.is_used:
|
||||
return False, "验证码已使用,请重新获取"
|
||||
|
||||
# 检查是否过期
|
||||
if latest.is_expired:
|
||||
return False, "验证码已过期,请重新获取"
|
||||
|
||||
# 检查尝试次数
|
||||
if latest.attempts > self.max_attempts:
|
||||
return False, "验证次数过多,请重新获取验证码"
|
||||
|
||||
# 校验验证码
|
||||
if latest.code != code_value:
|
||||
return False, "验证码错误"
|
||||
|
||||
# 校验通过,标记为已使用
|
||||
if consume:
|
||||
latest.mark_used()
|
||||
self.repo.save(latest)
|
||||
|
||||
return True, None
|
||||
|
||||
def _check_rate_limit(self, recipient: str, code_type: str) -> tuple[bool, int]:
|
||||
"""
|
||||
频控检查
|
||||
|
||||
Returns:
|
||||
(是否允许发送, 需等待秒数)
|
||||
"""
|
||||
# 检查冷却时间
|
||||
latest = self.repo.find_latest(recipient, code_type)
|
||||
if latest:
|
||||
elapsed = (datetime.now(timezone.utc) - latest.created_at).total_seconds()
|
||||
if elapsed < self.resend_cooldown:
|
||||
wait = int(self.resend_cooldown - elapsed)
|
||||
return False, wait
|
||||
|
||||
# 检查每日上限
|
||||
today_count = self.repo.count_today(recipient, code_type)
|
||||
if today_count >= self.daily_limit:
|
||||
return False, 0
|
||||
|
||||
return True, 0
|
||||
|
||||
|
||||
def validate_phone(phone: str) -> tuple[bool, str]:
|
||||
"""校验手机号格式(中国大陆手机号)"""
|
||||
phone = phone.strip()
|
||||
if not phone:
|
||||
return False, "手机号不能为空"
|
||||
# 支持 +86 前缀或纯 11 位
|
||||
pattern = r"^(\+86)?1[3-9]\d{9}$"
|
||||
if not re.match(pattern, phone):
|
||||
return False, "手机号格式不正确"
|
||||
return True, ""
|
||||
|
||||
|
||||
def normalize_phone(phone: str) -> str:
|
||||
"""标准化手机号(去掉 +86 前缀,统一存储格式)"""
|
||||
phone = phone.strip()
|
||||
if phone.startswith("+86"):
|
||||
phone = phone[3:]
|
||||
return phone
|
||||
|
||||
|
||||
def validate_email(email: str) -> tuple[bool, str]:
|
||||
"""校验邮箱格式"""
|
||||
email = email.strip()
|
||||
if not email:
|
||||
return False, "邮箱不能为空"
|
||||
pattern = r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$"
|
||||
if not re.match(pattern, email):
|
||||
return False, "邮箱格式不正确"
|
||||
return True, ""
|
||||
+163
@@ -0,0 +1,163 @@
|
||||
"""
|
||||
微信 OAuth 服务
|
||||
- 生成授权链接(网页扫码登录)
|
||||
- 处理回调,用 code 换 access_token + 用户信息
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import urllib.parse
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
import requests
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class WechatUserInfo:
|
||||
"""微信用户信息"""
|
||||
|
||||
openid: str
|
||||
unionid: str = ""
|
||||
nickname: str = ""
|
||||
avatar_url: str = ""
|
||||
|
||||
|
||||
class WechatOAuthService:
|
||||
"""微信开放平台 OAuth 服务(网页扫码登录)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
app_id: str | None = None,
|
||||
app_secret: str | None = None,
|
||||
redirect_uri: str | None = None,
|
||||
state_store=None,
|
||||
):
|
||||
self.app_id = app_id or os.environ.get("WECHAT_OPEN_APP_ID", "")
|
||||
self.app_secret = app_secret or os.environ.get("WECHAT_OPEN_APP_SECRET", "")
|
||||
self.redirect_uri = redirect_uri or os.environ.get("WECHAT_OPEN_REDIRECT_URI", "")
|
||||
self._state_store = state_store # 可选:state 存储(Redis/内存),用于 CSRF 防护
|
||||
|
||||
def is_configured(self) -> bool:
|
||||
"""检查微信配置是否完整"""
|
||||
return bool(self.app_id and self.app_secret and self.redirect_uri)
|
||||
|
||||
def generate_auth_url(self, scope: str = "snsapi_login") -> tuple[str, str]:
|
||||
"""
|
||||
生成微信授权链接
|
||||
|
||||
Returns:
|
||||
(授权URL, state)
|
||||
"""
|
||||
state = uuid4().hex
|
||||
|
||||
if not self.is_configured():
|
||||
# 未配置时返回 mock URL,方便前端联调
|
||||
mock_params = urllib.parse.urlencode(
|
||||
{
|
||||
"app_id": "mock",
|
||||
"redirect_uri": self.redirect_uri,
|
||||
"scope": scope,
|
||||
"state": state,
|
||||
}
|
||||
)
|
||||
return f"/mock/wechat/auth?{mock_params}", state
|
||||
|
||||
params = {
|
||||
"appid": self.app_id,
|
||||
"redirect_uri": self.redirect_uri,
|
||||
"response_type": "code",
|
||||
"scope": scope,
|
||||
"state": state,
|
||||
}
|
||||
url = "https://open.weixin.qq.com/connect/qrconnect?" + urllib.parse.urlencode(params) + "#wechat_redirect"
|
||||
return url, state
|
||||
|
||||
def handle_callback(self, code: str, state: str) -> tuple[Optional[WechatUserInfo], Optional[str]]:
|
||||
"""
|
||||
处理微信回调
|
||||
|
||||
Args:
|
||||
code: 微信授权码
|
||||
state: 防 CSRF 状态
|
||||
|
||||
Returns:
|
||||
(微信用户信息, 错误信息)
|
||||
"""
|
||||
if not code:
|
||||
return None, "缺少授权码"
|
||||
|
||||
if not self.is_configured():
|
||||
# 开发模式:返回 mock 用户信息
|
||||
logger.info("微信未配置,使用 mock 用户信息")
|
||||
return (
|
||||
WechatUserInfo(
|
||||
openid=f"mock_{code[:20]}",
|
||||
unionid=f"mock_union_{code[:16]}",
|
||||
nickname="微信测试用户",
|
||||
avatar_url="",
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
try:
|
||||
# 1. 用 code 换 access_token
|
||||
token_url = "https://api.weixin.qq.com/sns/oauth2/access_token"
|
||||
token_params = {
|
||||
"appid": self.app_id,
|
||||
"secret": self.app_secret,
|
||||
"code": code,
|
||||
"grant_type": "authorization_code",
|
||||
}
|
||||
token_resp = requests.get(token_url, params=token_params, timeout=10)
|
||||
token_data = token_resp.json()
|
||||
|
||||
if "errcode" in token_data and token_data["errcode"] != 0:
|
||||
logger.error("微信获取 access_token 失败: %s", token_data)
|
||||
return None, f"微信授权失败: {token_data.get('errmsg', '未知错误')}"
|
||||
|
||||
access_token = token_data["access_token"]
|
||||
openid = token_data["openid"]
|
||||
unionid = token_data.get("unionid", "")
|
||||
|
||||
# 2. 获取用户信息
|
||||
user_url = "https://api.weixin.qq.com/sns/userinfo"
|
||||
user_params = {
|
||||
"access_token": access_token,
|
||||
"openid": openid,
|
||||
"lang": "zh_CN",
|
||||
}
|
||||
user_resp = requests.get(user_url, params=user_params, timeout=10)
|
||||
user_data = user_resp.json()
|
||||
|
||||
if "errcode" in user_data and user_data["errcode"] != 0:
|
||||
logger.error("微信获取用户信息失败: %s", user_data)
|
||||
return None, f"获取用户信息失败: {user_data.get('errmsg', '未知错误')}"
|
||||
|
||||
return (
|
||||
WechatUserInfo(
|
||||
openid=openid,
|
||||
unionid=unionid,
|
||||
nickname=user_data.get("nickname", ""),
|
||||
avatar_url=user_data.get("headimgurl", ""),
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
except requests.RequestException as e:
|
||||
logger.error("微信 OAuth 请求异常: %s", e, exc_info=True)
|
||||
return None, "微信服务暂不可用,请稍后再试"
|
||||
except Exception as e:
|
||||
logger.error("微信回调处理异常: %s", e, exc_info=True)
|
||||
return None, "微信登录处理失败"
|
||||
|
||||
|
||||
def get_wechat_oauth_service() -> WechatOAuthService:
|
||||
"""获取微信 OAuth 服务单例"""
|
||||
# TODO: 可替换为 Redis state store
|
||||
return WechatOAuthService()
|
||||
Executable
+19
@@ -0,0 +1,19 @@
|
||||
"""
|
||||
短信服务接口
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class SmsService(ABC):
|
||||
"""短信服务接口"""
|
||||
|
||||
@abstractmethod
|
||||
def send_verification_code(self, phone: str, code: str) -> bool:
|
||||
"""发送验证码短信"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def send_template_sms(self, phone: str, template_id: str, params: dict) -> bool:
|
||||
"""发送模板短信"""
|
||||
pass
|
||||
@@ -59,6 +59,11 @@ class User:
|
||||
wechat_openid: str | None = None
|
||||
wechat_unionid: str | None = None
|
||||
|
||||
# 手机号绑定
|
||||
phone: str | None = None
|
||||
phone_verified: bool = False
|
||||
binding_completed_at: datetime | None = None
|
||||
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
|
||||
Executable
+66
@@ -0,0 +1,66 @@
|
||||
"""
|
||||
验证码领域实体
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class VerificationCode:
|
||||
"""验证码(邮箱/手机统一)"""
|
||||
|
||||
id: str
|
||||
recipient: str # 邮箱或手机号
|
||||
code: str
|
||||
code_type: str # email_bind / phone_bind / email_login / phone_login / reset_password
|
||||
expires_at: datetime
|
||||
used_at: datetime | None = None
|
||||
attempts: int = 0
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
recipient: str,
|
||||
code_type: str,
|
||||
ttl_seconds: int = 300,
|
||||
custom_code: str | None = None,
|
||||
) -> "VerificationCode":
|
||||
"""创建验证码"""
|
||||
import random
|
||||
|
||||
code = custom_code or "".join(random.choices("0123456789", k=6))
|
||||
now = datetime.now(timezone.utc)
|
||||
return cls(
|
||||
id=uuid4().hex,
|
||||
recipient=recipient.strip(),
|
||||
code=code,
|
||||
code_type=code_type,
|
||||
expires_at=now + timedelta(seconds=ttl_seconds),
|
||||
created_at=now,
|
||||
)
|
||||
|
||||
@property
|
||||
def is_expired(self) -> bool:
|
||||
"""是否已过期"""
|
||||
return datetime.now(timezone.utc) > self.expires_at
|
||||
|
||||
@property
|
||||
def is_used(self) -> bool:
|
||||
"""是否已使用"""
|
||||
return self.used_at is not None
|
||||
|
||||
@property
|
||||
def is_valid(self) -> bool:
|
||||
"""是否有效(未过期且未使用)"""
|
||||
return not self.is_expired and not self.is_used
|
||||
|
||||
def mark_used(self) -> None:
|
||||
"""标记为已使用"""
|
||||
self.used_at = datetime.now(timezone.utc)
|
||||
|
||||
def increment_attempts(self) -> None:
|
||||
"""增加尝试次数"""
|
||||
self.attempts += 1
|
||||
Regular → Executable
+5
@@ -51,6 +51,11 @@ class UserRepository(ABC):
|
||||
"""根据微信 unionid 查找用户"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_phone(self, phone: str) -> Optional[User]:
|
||||
"""根据手机号查找用户"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self, user_id: str) -> bool:
|
||||
"""删除用户"""
|
||||
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
"""
|
||||
验证码仓储接口
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
|
||||
|
||||
class VerificationCodeRepository(ABC):
|
||||
"""验证码仓储接口"""
|
||||
|
||||
@abstractmethod
|
||||
def save(self, code: VerificationCode) -> None:
|
||||
"""保存验证码"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_latest(self, recipient: str, code_type: str) -> Optional[VerificationCode]:
|
||||
"""查找最新的有效验证码"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_id(self, code_id: str) -> Optional[VerificationCode]:
|
||||
"""根据 ID 查找验证码"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def count_today(self, recipient: str, code_type: str) -> int:
|
||||
"""统计当日发送次数(频控用)"""
|
||||
pass
|
||||
+278
@@ -0,0 +1,278 @@
|
||||
"""
|
||||
验证码服务 + 绑定流程单元测试
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.in_memory.user_repository import InMemoryUserRepository
|
||||
from packages.application.auth.bind_contact_use_case import (
|
||||
BindContactRequest,
|
||||
BindContactUseCase,
|
||||
SendVerificationCodeRequest,
|
||||
SendVerificationCodeUseCase,
|
||||
)
|
||||
from packages.application.auth.verification_code_service import (
|
||||
CODE_TYPE_EMAIL_BIND,
|
||||
CODE_TYPE_PHONE_BIND,
|
||||
VerificationCodeService,
|
||||
normalize_phone,
|
||||
validate_email,
|
||||
validate_phone,
|
||||
)
|
||||
from packages.domain.entities import User
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
|
||||
|
||||
class InMemoryVerificationCodeRepository:
|
||||
"""内存版验证码仓储"""
|
||||
|
||||
def __init__(self):
|
||||
self._codes = {}
|
||||
|
||||
def save(self, code):
|
||||
self._codes[code.id] = code
|
||||
|
||||
def find_latest(self, recipient, code_type):
|
||||
candidates = [c for c in self._codes.values() if c.recipient == recipient and c.code_type == code_type]
|
||||
if not candidates:
|
||||
return None
|
||||
return max(candidates, key=lambda c: c.created_at)
|
||||
|
||||
def find_by_id(self, code_id):
|
||||
return self._codes.get(code_id)
|
||||
|
||||
def count_today(self, recipient, code_type):
|
||||
now = datetime.now(timezone.utc)
|
||||
start_of_day = now.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
return sum(
|
||||
1
|
||||
for c in self._codes.values()
|
||||
if c.recipient == recipient and c.code_type == code_type and c.created_at >= start_of_day
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def vc_repo():
|
||||
return InMemoryVerificationCodeRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def vc_service(vc_repo):
|
||||
return VerificationCodeService(repo=vc_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def user_repo():
|
||||
repo = InMemoryUserRepository()
|
||||
# 创建一个微信用户(未绑定手机+邮箱)
|
||||
user = User(
|
||||
id="user_001",
|
||||
email="wx_test@wechat.local",
|
||||
username="wx_test",
|
||||
display_name="微信测试用户",
|
||||
password_hash="hash",
|
||||
email_verified=True,
|
||||
wechat_openid="openid_test",
|
||||
)
|
||||
repo.save(user)
|
||||
return repo
|
||||
|
||||
|
||||
class TestPhoneValidation:
|
||||
def test_valid_phone(self):
|
||||
ok, _ = validate_phone("13800138000")
|
||||
assert ok
|
||||
|
||||
def test_valid_phone_with_86(self):
|
||||
ok, _ = validate_phone("+8613800138000")
|
||||
assert ok
|
||||
|
||||
def test_invalid_phone_too_short(self):
|
||||
ok, _ = validate_phone("12345")
|
||||
assert not ok
|
||||
|
||||
def test_invalid_phone_wrong_prefix(self):
|
||||
ok, _ = validate_phone("12000000000")
|
||||
assert not ok
|
||||
|
||||
def test_normalize_phone_strip_86(self):
|
||||
assert normalize_phone("+8613800138000") == "13800138000"
|
||||
assert normalize_phone("13800138000") == "13800138000"
|
||||
|
||||
|
||||
class TestEmailValidation:
|
||||
def test_valid_email(self):
|
||||
ok, _ = validate_email("test@example.com")
|
||||
assert ok
|
||||
|
||||
def test_invalid_email_no_at(self):
|
||||
ok, _ = validate_email("testexample.com")
|
||||
assert not ok
|
||||
|
||||
def test_invalid_email_empty(self):
|
||||
ok, _ = validate_email("")
|
||||
assert not ok
|
||||
|
||||
|
||||
class TestVerificationCodeService:
|
||||
def test_generate_code(self, vc_service):
|
||||
code, err = vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
assert err is None
|
||||
assert code is not None
|
||||
assert len(code.code) == 6
|
||||
assert code.code.isdigit()
|
||||
assert code.recipient == "13800138000"
|
||||
assert code.code_type == CODE_TYPE_PHONE_BIND
|
||||
assert code.is_valid
|
||||
|
||||
def test_generate_code_invalid_type(self, vc_service):
|
||||
code, err = vc_service.generate("13800138000", "invalid_type")
|
||||
assert err is not None
|
||||
assert code is None
|
||||
|
||||
def test_verify_success(self, vc_service):
|
||||
code, _ = vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
ok, err = vc_service.verify("13800138000", CODE_TYPE_PHONE_BIND, code.code)
|
||||
assert ok
|
||||
assert err is None
|
||||
assert code.is_used
|
||||
|
||||
def test_verify_wrong_code(self, vc_service):
|
||||
vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
ok, err = vc_service.verify("13800138000", CODE_TYPE_PHONE_BIND, "000000")
|
||||
assert not ok
|
||||
assert "错误" in err
|
||||
|
||||
def test_verify_expired_code(self, vc_service, vc_repo):
|
||||
# 手动创建一个已过期的验证码
|
||||
expired = VerificationCode.create("13800138000", CODE_TYPE_PHONE_BIND, ttl_seconds=1)
|
||||
expired.created_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
expired.expires_at = datetime.now(timezone.utc) - timedelta(seconds=5)
|
||||
vc_repo.save(expired)
|
||||
|
||||
ok, err = vc_service.verify("13800138000", CODE_TYPE_PHONE_BIND, expired.code)
|
||||
assert not ok
|
||||
assert "过期" in err
|
||||
|
||||
def test_resend_cooldown(self, vc_service):
|
||||
vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
_, err = vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
assert err is not None
|
||||
assert "频繁" in err
|
||||
|
||||
|
||||
class TestBindContactUseCase:
|
||||
def test_bind_phone_and_email(self, vc_service, vc_repo, user_repo):
|
||||
# 直接生成已知验证码(绕开发送流程,测试绑定逻辑)
|
||||
phone_code, _ = vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND, custom_code="123456")
|
||||
email_code, _ = vc_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="654321")
|
||||
assert phone_code and email_code
|
||||
|
||||
# 4. 绑定
|
||||
bind_uc = BindContactUseCase(
|
||||
user_repository=user_repo,
|
||||
verification_code_service=vc_service,
|
||||
)
|
||||
bind_req = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800138000",
|
||||
phone_code="123456",
|
||||
email="test@example.com",
|
||||
email_code="654321",
|
||||
)
|
||||
resp, err = bind_uc.execute(bind_req)
|
||||
assert err is None
|
||||
assert resp is not None
|
||||
assert resp.user.phone == "13800138000"
|
||||
assert resp.user.phone_verified
|
||||
assert resp.user.email == "test@example.com"
|
||||
assert resp.user.email_verified
|
||||
assert resp.user.binding_completed_at is not None
|
||||
|
||||
def test_bind_phone_only(self, vc_service, vc_repo, user_repo):
|
||||
vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND, custom_code="123456")
|
||||
|
||||
bind_uc = BindContactUseCase(
|
||||
user_repository=user_repo,
|
||||
verification_code_service=vc_service,
|
||||
)
|
||||
bind_req = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800138000",
|
||||
phone_code="123456",
|
||||
)
|
||||
resp, err = bind_uc.execute(bind_req)
|
||||
assert err is None
|
||||
assert resp.user.phone == "13800138000"
|
||||
assert resp.user.phone_verified
|
||||
# 只绑了手机,邮箱还是 wechat.local,所以 binding_complete 为 false
|
||||
assert resp.user.binding_completed_at is None
|
||||
|
||||
def test_bind_phone_wrong_code(self, vc_service, user_repo):
|
||||
vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
|
||||
bind_uc = BindContactUseCase(
|
||||
user_repository=user_repo,
|
||||
verification_code_service=vc_service,
|
||||
)
|
||||
bind_req = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800138000",
|
||||
phone_code="000000",
|
||||
)
|
||||
resp, err = bind_uc.execute(bind_req)
|
||||
assert err is not None
|
||||
assert resp is None
|
||||
|
||||
def test_bind_phone_already_used(self, vc_service, user_repo):
|
||||
# 另一个用户已经绑定了这个手机号
|
||||
other = User(
|
||||
id="user_002",
|
||||
email="other@example.com",
|
||||
username="other",
|
||||
display_name="其他用户",
|
||||
password_hash="hash",
|
||||
phone="13800138000",
|
||||
phone_verified=True,
|
||||
)
|
||||
user_repo.save(other)
|
||||
|
||||
vc_service.generate("13800138000", CODE_TYPE_PHONE_BIND, custom_code="123456")
|
||||
|
||||
bind_uc = BindContactUseCase(
|
||||
user_repository=user_repo,
|
||||
verification_code_service=vc_service,
|
||||
)
|
||||
bind_req = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800138000",
|
||||
phone_code="123456",
|
||||
)
|
||||
resp, err = bind_uc.execute(bind_req)
|
||||
assert err is not None
|
||||
assert "已被其他账号绑定" in err
|
||||
|
||||
|
||||
class TestWechatOAuthService:
|
||||
def test_mock_mode_generate_url(self):
|
||||
# 没配置真实微信参数时走 mock 模式
|
||||
from packages.application.auth.wechat_oauth_service import WechatOAuthService
|
||||
|
||||
service = WechatOAuthService(app_id="", app_secret="", redirect_uri="")
|
||||
assert not service.is_configured()
|
||||
|
||||
url, state = service.generate_auth_url()
|
||||
assert state
|
||||
assert "mock" in url
|
||||
|
||||
def test_mock_mode_callback(self):
|
||||
from packages.application.auth.wechat_oauth_service import WechatOAuthService
|
||||
|
||||
service = WechatOAuthService(app_id="", app_secret="", redirect_uri="")
|
||||
user_info, err = service.handle_callback("test_code", "test_state")
|
||||
assert err is None
|
||||
assert user_info is not None
|
||||
assert "mock" in user_info.openid
|
||||
assert user_info.nickname == "微信测试用户"
|
||||
Reference in New Issue
Block a user