Files
xiaoxia-saas/packages/middleware/points_gate.py
T
CI Bot 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
style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
2026-09-15 17:18:08 +00:00

265 lines
9.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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