217 lines
7.3 KiB
Python
217 lines
7.3 KiB
Python
"""AI 功能入口的积分扣费装饰器 (#1895)
|
||
|
||
支持 sync 和 async 函数。业务失败时自动退还积分。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import functools
|
||
import inspect
|
||
import logging
|
||
from collections.abc import Callable
|
||
from typing import Any
|
||
|
||
from fastapi import HTTPException
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def _points_gate_enabled() -> bool:
|
||
"""读取配置开关:总开关关闭时装饰器完全放行(零开销/零副作用)。
|
||
|
||
放在模块顶层便于测试时 monkeypatch。
|
||
"""
|
||
try:
|
||
from app.config import settings as _settings
|
||
|
||
return bool(_settings.points_enabled)
|
||
except Exception: # pragma: no cover - 配置加载失败默认关闭
|
||
return False
|
||
|
||
|
||
def points_gate(
|
||
scene_key: str,
|
||
per_unit: int | None = None,
|
||
unit_field: str | None = None,
|
||
quantity_field: str | None = None,
|
||
) -> Callable:
|
||
"""AI 功能入口积分扣费装饰器。
|
||
|
||
当 POINTS_ENABLED=false(默认)时,装饰器完全透传原函数,零副作用。
|
||
开启后才会执行扣费逻辑:业务异常自动退费,积分不足返回 402。
|
||
|
||
Args:
|
||
scene_key: 消耗场景标识(对应 points_rules.POINTS_SCENES 的 key)
|
||
per_unit: 固定消耗积分(直接指定,不走规则计算)
|
||
unit_field: 从 request body 取时长字段名(按时长计费场景)
|
||
quantity_field: 从 request body 取数量字段名(按次计费场景)
|
||
|
||
使用示例::
|
||
|
||
@router.post("/ai/voice")
|
||
@points_gate("ai_voice", unit_field="duration_minutes")
|
||
async def create_ai_voice(body: VoiceRequest, current_user=Depends(get_current_user), db=Depends(get_db_session)):
|
||
...
|
||
"""
|
||
|
||
def decorator(func: Callable) -> Callable:
|
||
is_async = asyncio.iscoroutinefunction(func)
|
||
|
||
@functools.wraps(func)
|
||
async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
|
||
if not _points_gate_enabled():
|
||
return await func(*args, **kwargs)
|
||
return await _execute_with_gate(
|
||
func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, is_async=True
|
||
)
|
||
|
||
@functools.wraps(func)
|
||
def sync_wrapper(*args: Any, **kwargs: Any) -> Any:
|
||
if not _points_gate_enabled():
|
||
return func(*args, **_filter_kwargs(func, kwargs))
|
||
return _execute_with_gate(
|
||
func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, is_async=False
|
||
)
|
||
|
||
if is_async:
|
||
return async_wrapper
|
||
return sync_wrapper
|
||
|
||
return decorator
|
||
|
||
|
||
|
||
def _filter_kwargs(func: Callable, kwargs: dict) -> dict:
|
||
"""过滤掉目标函数签名不接受的 kwargs(避免 TypeError)。"""
|
||
import inspect
|
||
try:
|
||
sig = inspect.signature(func)
|
||
params = sig.parameters
|
||
has_var_keyword = any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values())
|
||
if has_var_keyword:
|
||
return kwargs
|
||
return {k: v for k, v in kwargs.items() if k in params}
|
||
except (ValueError, TypeError):
|
||
return kwargs
|
||
|
||
def _extract_kwargs(func: Callable, args: tuple, kwargs: dict) -> dict:
|
||
"""将位置参数映射到函数签名中的参数名,便于统一按 kwargs 提取。"""
|
||
sig = inspect.signature(func)
|
||
bound = sig.bind_partial(*args, **kwargs)
|
||
merged = dict(bound.arguments)
|
||
merged.update(kwargs)
|
||
return merged
|
||
|
||
|
||
def _execute_with_gate(
|
||
func: Callable,
|
||
args: tuple,
|
||
kwargs: dict,
|
||
scene_key: str,
|
||
per_unit: int | None,
|
||
unit_field: str | None,
|
||
quantity_field: str | None,
|
||
is_async: bool,
|
||
) -> Any:
|
||
"""积分扣费核心逻辑。"""
|
||
merged = _extract_kwargs(func, args, kwargs)
|
||
|
||
# 提取 current_user(兼容 authenticated_user 命名)
|
||
current_user = merged.get("current_user") or merged.get("authenticated_user")
|
||
if current_user is None:
|
||
# 尝试从位置参数中找
|
||
for arg in args:
|
||
if hasattr(arg, "user"):
|
||
current_user = arg
|
||
break
|
||
if not current_user:
|
||
raise HTTPException(status_code=401, detail="未登录")
|
||
|
||
# 提取 db session
|
||
db = merged.get("db")
|
||
if db is None:
|
||
raise HTTPException(status_code=500, detail="缺少数据库 session")
|
||
|
||
user = current_user.user
|
||
is_member = getattr(user, "is_member", False)
|
||
member_type = getattr(user, "member_type", None)
|
||
|
||
# ── 混剪场景:先检查免费额度 ──
|
||
if scene_key == "ai_video":
|
||
from packages.domain.points_service import PointsService
|
||
|
||
svc = PointsService()
|
||
if not is_member:
|
||
if svc.check_daily_free_clip(user.id, db):
|
||
svc.record_daily_free_clip(user.id, db)
|
||
kwargs["_points_deducted"] = 0
|
||
kwargs["_is_free_quota"] = True
|
||
if is_async:
|
||
return _run_async(func, args, _filter_kwargs(func, kwargs))
|
||
return func(*args, **_filter_kwargs(func, kwargs))
|
||
|
||
# ── 计算积分消耗 ──
|
||
if per_unit is not None:
|
||
total_points = per_unit
|
||
else:
|
||
from packages.domain.points_rules import calculate_points_cost
|
||
|
||
quantity = 1
|
||
duration = 0.0
|
||
request_body = merged.get("body") or merged.get("request") or merged.get("payload")
|
||
if request_body and unit_field:
|
||
duration = float(getattr(request_body, unit_field, 0) or 0)
|
||
if request_body and quantity_field:
|
||
quantity = int(getattr(request_body, quantity_field, 1) or 1)
|
||
|
||
total_points = calculate_points_cost(
|
||
scene_key,
|
||
is_member,
|
||
quantity=quantity,
|
||
duration_minutes=duration,
|
||
member_type=member_type,
|
||
)
|
||
|
||
# 零消耗场景(如免费的声音克隆训练)直接放行
|
||
if total_points == 0:
|
||
kwargs["_points_deducted"] = 0
|
||
if is_async:
|
||
return _run_async(func, args, _filter_kwargs(func, kwargs))
|
||
return func(*args, **_filter_kwargs(func, kwargs))
|
||
|
||
# ── 扣减积分 ──
|
||
from packages.domain.points_service import PointsService
|
||
|
||
svc = PointsService()
|
||
job_id = merged.get("job_id", "") or ""
|
||
result = svc.deduct_points(user.id, total_points, scene_key, db, ref_id=str(job_id))
|
||
|
||
if not result["success"]:
|
||
raise HTTPException(
|
||
status_code=402,
|
||
detail={
|
||
"code": "INSUFFICIENT_POINTS",
|
||
"message": f"积分不足,需要 {total_points} 积分,当前余额 {result['balance']}",
|
||
"required": total_points,
|
||
"balance": result["balance"],
|
||
},
|
||
)
|
||
|
||
kwargs["_points_deducted"] = total_points
|
||
kwargs["_points_transaction_id"] = result["transaction_id"]
|
||
|
||
# ── 执行业务函数,失败则退还积分 ──
|
||
try:
|
||
if is_async:
|
||
return _run_async(func, args, _filter_kwargs(func, kwargs))
|
||
return func(*args, **_filter_kwargs(func, kwargs))
|
||
except Exception:
|
||
svc.refund_points(user.id, total_points, scene_key, db, ref_id=str(job_id))
|
||
raise
|
||
|
||
|
||
def _run_async(func: Callable, args: tuple, kwargs: dict):
|
||
"""在 async wrapper 中 await 原始 async 函数。"""
|
||
return func(*args, **_filter_kwargs(func, kwargs))
|