diff --git a/packages/middleware/points_gate.py b/packages/middleware/points_gate.py index c24240e51..54f8b7187 100644 --- a/packages/middleware/points_gate.py +++ b/packages/middleware/points_gate.py @@ -9,7 +9,6 @@ import asyncio import functools import inspect import logging -import types from collections.abc import Callable from typing import Any @@ -27,65 +26,10 @@ def _points_gate_enabled() -> bool: from app.config import settings as _settings return bool(_settings.points_enabled) - except Exception: # pragma: no cover + 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, @@ -96,17 +40,51 @@ def points_gate( 当 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) - return _make_wrapper(func, is_async, scene_key, per_unit, unit_field, quantity_field) + + @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 @@ -117,7 +95,6 @@ def _filter_kwargs(func: Callable, kwargs: dict) -> dict: except (ValueError, TypeError): return kwargs - def _extract_kwargs(func: Callable, args: tuple, kwargs: dict) -> dict: """将位置参数映射到函数签名中的参数名,便于统一按 kwargs 提取。""" sig = inspect.signature(func) @@ -143,6 +120,7 @@ def _execute_with_gate( # 提取 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 @@ -195,6 +173,7 @@ def _execute_with_gate( member_type=member_type, ) + # 零消耗场景(如免费的声音克隆训练)直接放行 if total_points == 0: kwargs["_points_deducted"] = 0 if is_async: