fix(points): @points_gate decorator filter internal kwargs to avoid TypeError on routes without **kwargs

This commit is contained in:
xiaoxia-agent
2026-09-15 20:11:34 +08:00
parent a082090edf
commit 42c122690d
+38 -59
View File
@@ -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: