"""AI 功能入口的积分扣费装饰器 (#1895) 支持 sync 和 async 函数。业务失败时自动退还积分。 """ from __future__ import annotations import asyncio import functools import inspect import logging import types 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 _make_wrapper( func: Callable, is_async: bool, scene_key: str, per_unit, unit_field, quantity_field, ) -> Callable: """在被装饰函数所在模块的 globals 下创建 wrapper。 关键点:Python 闭包默认在「定义闭包的模块」globals 下查找自由变量。若直接 在 points_gate 模块内定义 wrapper,wrapper.__globals__ 将指向本模块,导致 FastAPI/Pydantic 在解析函数类型注解(ForwardRef)时找不到路由模块中 导入的 Pydantic Model,出现 PydanticUndefinedAnnotation。 这里通过 types.FunctionType 把 wrapper 的 code 绑定到「被装饰函数所在模块 的 globals(补充装饰器内部符号)」,使 wrapper 的 ForwardRef 解析行为 与原路由函数一致。 """ if is_async: async def 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 ) else: def 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 ) # 把闭包所需的内部符号注入到被装饰模块的 globals,避免闭包找不到名字 merged_globals: dict = dict(func.__globals__) for _k, _v in ( ("_points_gate_enabled", _points_gate_enabled), ("_execute_with_gate", _execute_with_gate), ("_run_async", _run_async), ("_filter_kwargs", _filter_kwargs), ): merged_globals.setdefault(_k, _v) new_wrapper = types.FunctionType( wrapper.__code__, merged_globals, wrapper.__name__, wrapper.__defaults__, wrapper.__closure__, ) new_wrapper = functools.wraps(func)(new_wrapper) return new_wrapper 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。 """ def decorator(func: Callable) -> Callable: is_async = asyncio.iscoroutinefunction(func) return _make_wrapper(func, is_async, scene_key, per_unit, unit_field, quantity_field) return decorator def _filter_kwargs(func: Callable, kwargs: dict) -> dict: """过滤掉目标函数签名不接受的 kwargs(避免 TypeError)。""" 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))