diff --git a/video-gen-api/app/api/v1/sms.py b/video-gen-api/app/api/v1/sms.py index c53b335c..e842ed4f 100644 --- a/video-gen-api/app/api/v1/sms.py +++ b/video-gen-api/app/api/v1/sms.py @@ -1,7 +1,10 @@ -from fastapi import APIRouter, HTTPException, status +from fastapi import APIRouter, Depends, HTTPException, status +from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings -from app.schemas.sms import SmsSendRequest, SmsVerifyRequest, SmsResponse +from app.dependencies import get_db +from app.schemas.sms import SmsResponse, SmsScene, SmsSendRequest, SmsVerifyRequest +from app.services.auth import get_user_by_phone from app.services.sms import generate_and_send_sms, verify_sms_code router = APIRouter(prefix="/sms", tags=["短信验证码"]) @@ -10,6 +13,7 @@ router = APIRouter(prefix="/sms", tags=["短信验证码"]) def _validate_captcha_token(captcha_token: str | None) -> None: if settings.SMS_MOCK: return + if not (req_captcha := captcha_token): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -26,15 +30,74 @@ def _validate_captcha_token(captcha_token: str | None) -> None: ) +async def _validate_sms_scene_phone( + *, + db: AsyncSession, + phone: str, + scene: SmsScene, +) -> None: + """ + 按短信来源场景做手机号注册状态前置校验。 + + 规则: + 1. login:必须是已注册手机号,否则拦截发送。 + 2. register:必须是未注册手机号,否则拦截发送。 + 3. common:通用场景,暂不校验手机号注册状态。 + 4. set_password:保留旧场景,按通用场景处理,暂不校验手机号注册状态。 + 5. 其他未确认场景:理论上会被 SmsScene 枚举拦截;这里额外兜底。 + """ + if scene == SmsScene.login: + user = await get_user_by_phone(db, phone) + if not user: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="手机号未注册,请先注册", + ) + return + + if scene == SmsScene.register: + user = await get_user_by_phone(db, phone) + if user: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="该手机号已注册,请直接登录", + ) + return + + if scene in {SmsScene.common, SmsScene.set_password}: + return + + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="未确认短信场景,禁止发送验证码", + ) + + @router.post( "/send", response_model=SmsResponse, summary="发送短信验证码", - description="客户端发送短信验证码。scene=register 用于注册,scene=login 用于短信登录,scene=set_password 用于设置密码。正式环境会按配置校验图形验证码。", + description=( + "客户端发送短信验证码。" + "scene=register 用于注册,发送前会校验手机号未注册;" + "scene=login 用于短信登录,发送前会校验手机号已注册;" + "scene=common 用于通用短信场景,暂不校验手机号注册状态;" + "scene=set_password 用于设置密码,保留旧场景,暂按通用场景处理。" + "正式环境会按配置校验图形验证码。" + ), ) -async def send_sms_code(req: SmsSendRequest): +async def send_sms_code( + req: SmsSendRequest, + db: AsyncSession = Depends(get_db), +): _validate_captcha_token(req.captcha_token) + await _validate_sms_scene_phone( + db=db, + phone=req.phone, + scene=req.scene, + ) + try: ok, code = await generate_and_send_sms(req.phone, req.scene.value) except ValueError as exc: @@ -66,4 +129,4 @@ async def verify_sms(req: SmsVerifyRequest): status_code=status.HTTP_400_BAD_REQUEST, detail="验证码错误或已过期", ) - return SmsResponse(message="验证成功", success=True) + return SmsResponse(message="验证成功", success=True) \ No newline at end of file diff --git a/video-gen-api/app/schemas/sms.py b/video-gen-api/app/schemas/sms.py index d070b1f9..1c84396a 100644 --- a/video-gen-api/app/schemas/sms.py +++ b/video-gen-api/app/schemas/sms.py @@ -6,21 +6,40 @@ from pydantic import BaseModel, Field class SmsScene(str, Enum): register = "register" login = "login" + common = "common" set_password = "set_password" class SmsSendRequest(BaseModel): phone: str = Field(..., pattern=r"^1[3-9]\d{9}$", description="手机号") - scene: SmsScene = Field(..., description="短信场景:register=注册,login=短信登录,set_password=设置密码") + scene: SmsScene = Field( + ..., + description=( + "短信场景:" + "register=注册," + "login=短信登录," + "common=通用场景," + "set_password=设置密码" + ), + ) captcha_token: str | None = Field(None, description="图形验证码令牌,正式环境按配置要求传入") class SmsVerifyRequest(BaseModel): phone: str = Field(..., pattern=r"^1[3-9]\d{9}$", description="手机号") - scene: SmsScene = Field(..., description="短信场景") + scene: SmsScene = Field( + ..., + description=( + "短信场景:" + "register=注册," + "login=短信登录," + "common=通用场景," + "set_password=设置密码" + ), + ) code: str = Field(..., min_length=4, max_length=8, description="验证码") class SmsResponse(BaseModel): message: str - success: bool + success: bool \ No newline at end of file