Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f67095ade6 |
Regular → Executable
+4
-4
@@ -503,7 +503,7 @@ async def send_verification_code(
|
||||
request: SendVerificationCodeRequest,
|
||||
) -> SendVerificationCodeResponse:
|
||||
"""发送验证码(手机或邮箱)"""
|
||||
from app.dependencies import get_db
|
||||
from app.dependencies import get_db_session
|
||||
|
||||
from packages.adapters.sms.sms_service import get_sms_service
|
||||
from packages.adapters.smtp import get_email_service
|
||||
@@ -516,7 +516,7 @@ async def send_verification_code(
|
||||
)
|
||||
from packages.application.auth.verification_code_service import VerificationCodeService
|
||||
|
||||
db = next(get_db())
|
||||
db = next(get_db_session())
|
||||
repo = SQLAlchemyVerificationCodeRepository(db)
|
||||
vc_service = VerificationCodeService(repo=repo)
|
||||
sms_service = get_sms_service()
|
||||
@@ -549,7 +549,7 @@ async def bind_contact(
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> BindContactResponse:
|
||||
"""绑定手机号和/或邮箱(需登录态)"""
|
||||
from app.dependencies import get_db
|
||||
from app.dependencies import get_db_session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.verification_code_repository import (
|
||||
SQLAlchemyVerificationCodeRepository,
|
||||
@@ -560,7 +560,7 @@ async def bind_contact(
|
||||
)
|
||||
from packages.application.auth.verification_code_service import VerificationCodeService
|
||||
|
||||
db = next(get_db())
|
||||
db = next(get_db_session())
|
||||
vc_repo = SQLAlchemyVerificationCodeRepository(db)
|
||||
vc_service = VerificationCodeService(repo=vc_repo)
|
||||
|
||||
|
||||
@@ -8,8 +8,10 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import urllib.parse
|
||||
from dataclasses import dataclass
|
||||
from threading import Lock
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -17,6 +19,39 @@ import requests
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
STATE_TTL_SECONDS = 600 # state 有效期 10 分钟
|
||||
|
||||
|
||||
class MemoryStateStore:
|
||||
"""内存 state 存储(简单实现,单节点可用)
|
||||
|
||||
多实例部署时建议替换为 Redis 实现。
|
||||
"""
|
||||
|
||||
def __init__(self, ttl_seconds: int = STATE_TTL_SECONDS):
|
||||
self._ttl = ttl_seconds
|
||||
self._states: dict[str, float] = {} # state -> expire_at
|
||||
self._lock = Lock()
|
||||
|
||||
def put(self, state: str) -> None:
|
||||
with self._lock:
|
||||
self._clean_expired()
|
||||
self._states[state] = time.time() + self._ttl
|
||||
|
||||
def verify_and_consume(self, state: str) -> bool:
|
||||
with self._lock:
|
||||
self._clean_expired()
|
||||
if state in self._states:
|
||||
del self._states[state]
|
||||
return True
|
||||
return False
|
||||
|
||||
def _clean_expired(self) -> None:
|
||||
now = time.time()
|
||||
expired = [s for s, exp in self._states.items() if exp < now]
|
||||
for s in expired:
|
||||
del self._states[s]
|
||||
|
||||
|
||||
@dataclass
|
||||
class WechatUserInfo:
|
||||
@@ -41,7 +76,8 @@ class WechatOAuthService:
|
||||
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 防护
|
||||
# state 存储(CSRF 防护),默认内存实现
|
||||
self._state_store = state_store or MemoryStateStore()
|
||||
|
||||
def is_configured(self) -> bool:
|
||||
"""检查微信配置是否完整"""
|
||||
@@ -55,6 +91,8 @@ class WechatOAuthService:
|
||||
(授权URL, state)
|
||||
"""
|
||||
state = uuid4().hex
|
||||
# 保存 state 用于回调校验(防 CSRF)
|
||||
self._state_store.put(state)
|
||||
|
||||
if not self.is_configured():
|
||||
# 未配置时返回 mock URL,方便前端联调
|
||||
@@ -92,6 +130,11 @@ class WechatOAuthService:
|
||||
if not code:
|
||||
return None, "缺少授权码"
|
||||
|
||||
# 校验 state(防 CSRF)—— 一次性使用
|
||||
if not state or not self._state_store.verify_and_consume(state):
|
||||
logger.warning("微信回调 state 校验失败: state=%s", state)
|
||||
return None, "无效的 state 参数,请求可能已过期或被篡改"
|
||||
|
||||
if not self.is_configured():
|
||||
# 开发模式:返回 mock 用户信息
|
||||
logger.info("微信未配置,使用 mock 用户信息")
|
||||
|
||||
@@ -271,8 +271,49 @@ class TestWechatOAuthService:
|
||||
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")
|
||||
# 先生成授权 URL 获得有效 state(state 会被存入 store)
|
||||
_, valid_state = service.generate_auth_url()
|
||||
user_info, err = service.handle_callback("test_code", valid_state)
|
||||
assert err is None
|
||||
assert user_info is not None
|
||||
assert "mock" in user_info.openid
|
||||
assert user_info.nickname == "微信测试用户"
|
||||
|
||||
def test_callback_invalid_state_rejected(self):
|
||||
"""无效 state 应被拒绝(CSRF 防护)"""
|
||||
from packages.application.auth.wechat_oauth_service import WechatOAuthService
|
||||
|
||||
service = WechatOAuthService(app_id="", app_secret="", redirect_uri="")
|
||||
# 直接用随机 state 调用,未经过 generate_auth_url
|
||||
user_info, err = service.handle_callback("test_code", "random_fake_state")
|
||||
assert err is not None
|
||||
assert "state" in err
|
||||
assert user_info is None
|
||||
|
||||
def test_callback_state_single_use(self):
|
||||
"""state 只能使用一次(防重放)"""
|
||||
from packages.application.auth.wechat_oauth_service import WechatOAuthService
|
||||
|
||||
service = WechatOAuthService(app_id="", app_secret="", redirect_uri="")
|
||||
_, valid_state = service.generate_auth_url()
|
||||
|
||||
# 第一次使用:成功
|
||||
user_info, err = service.handle_callback("test_code", valid_state)
|
||||
assert err is None
|
||||
assert user_info is not None
|
||||
|
||||
# 第二次使用相同 state:失败(已被消费)
|
||||
user_info2, err2 = service.handle_callback("test_code", valid_state)
|
||||
assert err2 is not None
|
||||
assert "state" in err2
|
||||
assert user_info2 is None
|
||||
|
||||
def test_callback_empty_state_rejected(self):
|
||||
"""空 state 应被拒绝"""
|
||||
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", "")
|
||||
assert err is not None
|
||||
assert "state" in err
|
||||
assert user_info is None
|
||||
|
||||
Reference in New Issue
Block a user