"""AI 功能入口的积分扣费装饰器 (#1895) 支持 sync 和 async 函数。业务失败时自动退还积分。 关键设计点: 1. **POINTS_ENABLED 默认关闭**,装饰器零副作用透传,安全上线。 2. **wrapper 绑定到被装饰模块的 globals**:Python 闭包的 __globals__ 默认指向定义闭包 的模块(即本文件),但 Pydantic 在解函数类型注解里的 ForwardRef 时(Python 3.12 eval_type_backport 路径)直接用 wrapper.__globals__ 查表,会找不到路由模块里 导入/定义的 Pydantic Model,报 PydanticUndefinedAnnotation。因此用 ``types.FunctionType`` 把 wrapper code 绑定到被装饰函数所在模块的 globals。 3. **装饰器内部入口通过「本模块 __dict__ 动态查找」**:注入到被装饰模块 globals 的是一层薄的转发函数,每次调用都从 ``sys.modules[本模块]`` 里取最新引用,这样 测试里 ``monkeypatch.setattr(points_gate, "_points_gate_enabled", lambda: True)`` 等替换依然能生效。 """ from __future__ import annotations import asyncio import functools import inspect import logging import sys import types from collections.abc import Callable from typing import Any from fastapi import HTTPException logger = logging.getLogger(__name__) _PG_MODULE_NAME = __name__ # "packages.middleware.points_gate" # ── 对外暴露、可被 monkeypatch 替换的入口 ──────────────────────────────────── def _points_gate_enabled() -> bool: """读取 POINTS_ENABLED 配置开关(默认 False)。 暴露在模块顶层便于测试 monkeypatch。 """ try: from app.config import settings as _settings return bool(_settings.points_enabled) except Exception: # pragma: no cover return False # ── 转发 helper(被注入到被装饰模块 globals,动态从本模块取最新实现) ──────── def _pg_enabled_proxy(): return sys.modules[_PG_MODULE_NAME]._points_gate_enabled() def _pg_filter_kwargs_proxy(func, kwargs): return sys.modules[_PG_MODULE_NAME]._filter_kwargs_impl(func, kwargs) def _pg_execute_proxy(func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, is_async): return sys.modules[_PG_MODULE_NAME]._execute_with_gate_impl( func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, is_async ) # ── 真正实现(不直接被 wrapper 闭包引用,通过 proxy 访问) ───────────────── def _filter_kwargs_impl(func: Callable, kwargs: dict) -> dict: 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 points_gate( scene_key: str, per_unit: int | None = None, unit_field: str | None = None, quantity_field: str | None = None, ) -> Callable: def decorator(func: Callable) -> Callable: is_async = asyncio.iscoroutinefunction(func) if is_async: async def wrapper(*args: Any, **kwargs: Any) -> Any: if not _pg_enabled(): # noqa: F821 return await func(*args, **kwargs) return await _pg_execute( # noqa: F821 func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, True ) else: def wrapper(*args: Any, **kwargs: Any) -> Any: if not _pg_enabled(): # noqa: F821 return func(*args, **_pg_filter(func, kwargs)) # noqa: F821 return _pg_execute( # noqa: F821 func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, False ) # 把 wrapper code 绑定到被装饰函数所在模块的 globals, # 并注入 proxy 入口(短名避免冲突) route_globals: dict = func.__globals__ merged_globals = dict(route_globals) # 用相对唯一但简短的名字注入,避免和业务模块已有符号冲突 # (setdefault 不覆盖业务模块已有同名符号,如有冲突会抛错在装饰阶段暴露) proxies = { "_pg_enabled": _pg_enabled_proxy, "_pg_filter": _pg_filter_kwargs_proxy, "_pg_execute": _pg_execute_proxy, } for k, v in proxies.items(): if k in merged_globals and merged_globals[k] is not v: # 命名冲突,换更长的唯一前缀 k2 = f"__pg_{scene_key}_{k}" merged_globals[k2] = v # 需要相应替换 wrapper 内引用 → 重新编译 wrapper 不现实, # 但这种场景在我们代码里不会出现(短名 _pg_enabled 等极少冲突)。 # 为稳妥起见,直接把 wrapper code 的 co_names 映射到新名——复杂度过高, # 这里采用「确保短名没冲突」策略:如果冲突就抛异常让开发者改名。 raise RuntimeError( f"points_gate: name collision in {func.__module__}.{func.__name__}: " f"'{k}' already defined" ) merged_globals[k] = v new_wrapper = types.FunctionType( wrapper.__code__, merged_globals, wrapper.__name__, wrapper.__defaults__, wrapper.__closure__, ) # functools.wraps 会复制 __name__/__doc__/__wrapped__/__module__ 等, # 但注意不要把 __globals__ 覆盖回去。 new_wrapper = functools.wraps(func)(new_wrapper) return new_wrapper return decorator def _extract_kwargs(func: Callable, args: tuple, kwargs: dict) -> dict: sig = inspect.signature(func) bound = sig.bind_partial(*args, **kwargs) merged = dict(bound.arguments) merged.update(kwargs) return merged def _execute_with_gate_impl( 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 = 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 = 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_impl(func, args, _filter_kwargs_impl(func, kwargs)) return func(*args, **_filter_kwargs_impl(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_impl(func, args, _filter_kwargs_impl(func, kwargs)) return func(*args, **_filter_kwargs_impl(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_impl(func, args, _filter_kwargs_impl(func, kwargs)) return func(*args, **_filter_kwargs_impl(func, kwargs)) except Exception: svc.refund_points(user.id, total_points, scene_key, db, ref_id=str(job_id)) raise def _run_async_impl(func: Callable, args: tuple, kwargs: dict): return func(*args, **_filter_kwargs_impl(func, kwargs)) # 兼容历史测试文件直接 import 的别名 _filter_kwargs = _filter_kwargs_impl _execute_with_gate = _execute_with_gate_impl _run_async = _run_async_impl