diff --git a/apps/api/app/api/routes/auth.py b/apps/api/app/api/routes/auth.py old mode 100644 new mode 100755 index 4855f3a07..45286aa19 --- a/apps/api/app/api/routes/auth.py +++ b/apps/api/app/api/routes/auth.py @@ -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) diff --git a/packages/application/auth/wechat_oauth_service.py b/packages/application/auth/wechat_oauth_service.py index 23e7d8b28..16cdf5d47 100755 --- a/packages/application/auth/wechat_oauth_service.py +++ b/packages/application/auth/wechat_oauth_service.py @@ -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 用户信息") diff --git a/tests/unit/test_wechat_login_and_verification.py b/tests/unit/test_wechat_login_and_verification.py index 7a85ad681..9cd55481c 100755 --- a/tests/unit/test_wechat_login_and_verification.py +++ b/tests/unit/test_wechat_login_and_verification.py @@ -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