7715b789a8
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Successful in 4s
CI/CD Pipeline / Check push changed paths (push) Successful in 19s
CI/CD Pipeline / Build Staging API Image (push) Successful in 32s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 29s
CI/CD Pipeline / Validate - Style (push) Has been cancelled
CI/CD Pipeline / Validate - Security (push) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (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 Unit Tests (push) Has been cancelled
CI/CD Pipeline / Build Staging Web 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
CI/CD Pipeline / PR Build Worker Image (push) Failing after 13h49m8s
CI/CD Pipeline / PR Build Web Image (push) Failing after 13h48m24s
CI/CD Pipeline / PR Build API Image (push) Failing after 13h48m24s
CI/CD Pipeline / Frontend Lint (push) Failing after 13h46m3s
CI/CD Pipeline / Check if frontend-only change (push) Failing after 13h56m22s
265 lines
9.6 KiB
Python
265 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
|