187 lines
6.4 KiB
Python
187 lines
6.4 KiB
Python
import asyncio
|
|
import json
|
|
import logging
|
|
import random
|
|
import time
|
|
from datetime import datetime
|
|
from typing import Any
|
|
|
|
from app.config import settings
|
|
from app.utils.redis import get_redis
|
|
|
|
logger = logging.getLogger("videogen")
|
|
|
|
# In-memory fallback for local development when Redis is disabled.
|
|
_sms_code_store: dict[str, tuple[str, float]] = {}
|
|
_sms_send_interval_store: dict[str, float] = {}
|
|
_sms_daily_count_store: dict[str, tuple[str, int]] = {}
|
|
|
|
|
|
def _generate_code(length: int | None = None) -> str:
|
|
code_length = int(length or settings.SMS_CODE_LENGTH or 4)
|
|
code_length = max(4, min(code_length, 8))
|
|
return "".join(random.choices("0123456789", k=code_length))
|
|
|
|
|
|
def _normalize_scene(scene: str | None) -> str:
|
|
value = (scene or "").strip().lower()
|
|
if value not in {"register", "login", "set_password"}:
|
|
value = "login"
|
|
return value
|
|
|
|
|
|
def _code_key(phone: str, scene: str | None) -> str:
|
|
return f"sms_code:{_normalize_scene(scene)}:{phone}"
|
|
|
|
|
|
def _interval_key(phone: str, scene: str | None) -> str:
|
|
return f"sms_interval:{_normalize_scene(scene)}:{phone}"
|
|
|
|
|
|
def _daily_key(phone: str, scene: str | None) -> str:
|
|
today = datetime.now().strftime("%Y%m%d")
|
|
return f"sms_daily:{_normalize_scene(scene)}:{phone}:{today}"
|
|
|
|
|
|
def _today() -> str:
|
|
return datetime.now().strftime("%Y%m%d")
|
|
|
|
|
|
async def _check_send_limit(phone: str, scene: str) -> None:
|
|
"""发送频控:同场景同手机号间隔限制 + 每日次数限制。"""
|
|
redis = get_redis()
|
|
interval_seconds = int(settings.SMS_SEND_INTERVAL_SECONDS or 60)
|
|
daily_limit = int(settings.SMS_DAILY_LIMIT or 20)
|
|
|
|
if redis:
|
|
interval_key = _interval_key(phone, scene)
|
|
if await redis.get(interval_key):
|
|
raise ValueError(f"短信发送过于频繁,请{interval_seconds}秒后再试")
|
|
|
|
daily_key = _daily_key(phone, scene)
|
|
count = await redis.incr(daily_key)
|
|
if count == 1:
|
|
await redis.expire(daily_key, 24 * 60 * 60)
|
|
if count > daily_limit:
|
|
raise ValueError("今日短信发送次数已达上限,请明天再试")
|
|
|
|
await redis.setex(interval_key, interval_seconds, "1")
|
|
return
|
|
|
|
now_ts = time.time()
|
|
interval_key = _interval_key(phone, scene)
|
|
last_send_at = _sms_send_interval_store.get(interval_key)
|
|
if last_send_at and now_ts - last_send_at < interval_seconds:
|
|
raise ValueError(f"短信发送过于频繁,请{interval_seconds}秒后再试")
|
|
_sms_send_interval_store[interval_key] = now_ts
|
|
|
|
daily_key = f"{_normalize_scene(scene)}:{phone}"
|
|
current_day = _today()
|
|
stored_day, count = _sms_daily_count_store.get(daily_key, (current_day, 0))
|
|
if stored_day != current_day:
|
|
stored_day, count = current_day, 0
|
|
count += 1
|
|
_sms_daily_count_store[daily_key] = (stored_day, count)
|
|
if count > daily_limit:
|
|
raise ValueError("今日短信发送次数已达上限,请明天再试")
|
|
|
|
|
|
def _send_volc_sms_sync(phone: str, code: str) -> dict[str, Any]:
|
|
from volcengine.sms.SmsService import SmsService
|
|
|
|
if not settings.VOLC_SMS_ACCESS_KEY_ID or not settings.VOLC_SMS_SECRET_ACCESS_KEY:
|
|
raise RuntimeError("火山短信 AK/SK 未配置")
|
|
if not settings.VOLC_SMS_ACCOUNT:
|
|
raise RuntimeError("火山短信消息组ID VOLC_SMS_ACCOUNT 未配置")
|
|
if not settings.VOLC_SMS_TEMPLATE_ID:
|
|
raise RuntimeError("火山短信模板ID VOLC_SMS_TEMPLATE_ID 未配置")
|
|
if not settings.VOLC_SMS_SIGN:
|
|
raise RuntimeError("火山短信签名 VOLC_SMS_SIGN 未配置")
|
|
|
|
sms_service = SmsService()
|
|
sms_service.set_ak(settings.VOLC_SMS_ACCESS_KEY_ID)
|
|
sms_service.set_sk(settings.VOLC_SMS_SECRET_ACCESS_KEY)
|
|
|
|
body = {
|
|
"SmsAccount": settings.VOLC_SMS_ACCOUNT,
|
|
"Sign": settings.VOLC_SMS_SIGN,
|
|
"TemplateID": settings.VOLC_SMS_TEMPLATE_ID,
|
|
"TemplateParam": json.dumps({"xxxx": str(code)}, ensure_ascii=False, separators=(",", ":")),
|
|
"Tag": f"{phone}:{int(time.time())}",
|
|
"PhoneNumbers": phone,
|
|
}
|
|
raw_resp = sms_service.send_sms(json.dumps(body, ensure_ascii=False, separators=(",", ":")))
|
|
|
|
if isinstance(raw_resp, str):
|
|
try:
|
|
resp: dict[str, Any] = json.loads(raw_resp)
|
|
except json.JSONDecodeError:
|
|
resp = {"raw": raw_resp}
|
|
elif isinstance(raw_resp, dict):
|
|
resp = raw_resp
|
|
else:
|
|
resp = {"raw": raw_resp}
|
|
|
|
error = (resp.get("ResponseMetadata") or {}).get("Error") if isinstance(resp, dict) else None
|
|
if error:
|
|
raise RuntimeError(f"火山短信发送失败:{error.get('Code')} {error.get('Message')}")
|
|
|
|
return resp
|
|
|
|
|
|
async def send_sms(phone: str, code: str, scene: str = "login") -> bool:
|
|
"""发送短信验证码。
|
|
|
|
SMS_MOCK=true 时只写日志,方便本地调试;否则使用火山引擎短信 SDK。
|
|
火山 SDK 为同步调用,这里放到线程中执行,避免阻塞 FastAPI event loop。
|
|
"""
|
|
scene = _normalize_scene(scene)
|
|
|
|
try:
|
|
logger.info("[SMS] scene=%s, To=%s, Code=%s", scene, phone, code)
|
|
await asyncio.to_thread(_send_volc_sms_sync, phone, code)
|
|
return True
|
|
except Exception:
|
|
logger.exception("SMS send failed. scene=%s phone=%s", scene, phone)
|
|
return False
|
|
|
|
|
|
async def store_sms_code(phone: str, code: str, scene: str = "login", ttl: int | None = None) -> None:
|
|
ttl_seconds = int(ttl or settings.SMS_CODE_TTL_SECONDS or 300)
|
|
key = _code_key(phone, scene)
|
|
redis = get_redis()
|
|
if redis:
|
|
await redis.setex(key, ttl_seconds, code)
|
|
else:
|
|
_sms_code_store[key] = (code, time.time() + ttl_seconds)
|
|
|
|
|
|
async def generate_and_send_sms(phone: str, scene: str = "login") -> bool:
|
|
scene = _normalize_scene(scene)
|
|
await _check_send_limit(phone, scene)
|
|
code = _generate_code()
|
|
ok = await send_sms(phone, code, scene)
|
|
if ok:
|
|
await store_sms_code(phone, code, scene)
|
|
return ok
|
|
|
|
|
|
async def verify_sms_code(phone: str, code: str, scene: str = "login") -> bool:
|
|
key = _code_key(phone, scene)
|
|
redis = get_redis()
|
|
if redis:
|
|
stored = await redis.get(key)
|
|
if isinstance(stored, bytes):
|
|
stored = stored.decode()
|
|
if stored and str(stored) == str(code):
|
|
await redis.delete(key)
|
|
return True
|
|
return False
|
|
|
|
entry = _sms_code_store.pop(key, None)
|
|
if entry:
|
|
stored_code, expires = entry
|
|
if time.time() < expires and stored_code == code:
|
|
return True
|
|
return False
|