29cdd32203
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Failing after 26s
CI/CD Pipeline / Check push changed paths (push) Successful in 2m24s
CI/CD Pipeline / Build Staging API Image (push) Successful in 15s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 18s
CI/CD Pipeline / Build Staging Web Image (push) Successful in 1m26s
CI/CD Pipeline / Check if frontend-only change (push) Failing after 13m50s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 7m14s
CI/CD Pipeline / Validate - Style (push) Has been cancelled
CI/CD Pipeline / Validate - Security (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (push) Has been cancelled
CI/CD Pipeline / PR Build API Image (push) Has been cancelled
CI/CD Pipeline / PR Build Web Image (push) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (push) Has been cancelled
CI/CD Pipeline / Retag skipped Staging API Image (push) Has been cancelled
CI/CD Pipeline / Retag skipped Staging Web Image (push) Has been cancelled
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production API Image (push) Has been cancelled
CI/CD Pipeline / Build Production Web Image (push) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (push) Has been cancelled
CI/CD Pipeline / Canary Release to Production (push) Has been cancelled
CI/CD Pipeline / CI Gate (push) Has been cancelled
Fixes PydanticUndefinedAnnotation on CI (Python 3.12): the wrapper's __globals__ previously pointed at packages.middleware.points_gate, so ForwardRefs for request models (e.g. ExtractFromDouyinRequest) could not resolve. Rebind the wrapper code object to the decorated route's globals via types.FunctionType, and route internal helpers through thin sys.modules proxies so tests can monkeypatch them.
263 lines
9.6 KiB
Python
263 lines
9.6 KiB
Python
"""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
|