Files
xiaoxia-saas/packages/middleware/points_gate.py
T
CI Bot ade0cf78e7
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m13s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m58s
AI Code Review / AI Code Review (pull_request) Successful in 6m48s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 10s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 10s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 8m57s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 7m50s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 7m36s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 8m43s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 17m11s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m53s
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 14h57m54s
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Failing after 15h3m28s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 15h27m29s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 15h27m49s
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Failing after 15h2m45s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 15h26m45s
CI/CD Pipeline / PR Build Web Image (pull_request) Failing after 15h26m55s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 14h57m11s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 14h57m11s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 14h55m35s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Failing after 15h2m45s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 15h26m45s
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 15h27m5s
CI/CD Pipeline / Check push changed paths (pull_request) Failing after 15h36m4s
style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
2026-09-15 15:18:24 +00:00

241 lines
7.9 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 函数。业务失败时自动退还积分。
"""
from __future__ import annotations
import asyncio
import functools
import inspect
import logging
import types
from collections.abc import Callable
from typing import Any
from fastapi import HTTPException
logger = logging.getLogger(__name__)
def _points_gate_enabled() -> bool:
"""读取配置开关:总开关关闭时装饰器完全放行(零开销/零副作用)。
放在模块顶层便于测试时 monkeypatch。
"""
try:
from app.config import settings as _settings
return bool(_settings.points_enabled)
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,
unit_field: str | None = None,
quantity_field: str | None = None,
) -> Callable:
"""AI 功能入口积分扣费装饰器。
当 POINTS_ENABLED=false(默认)时,装饰器完全透传原函数,零副作用。
开启后才会执行扣费逻辑:业务异常自动退费,积分不足返回 402。
"""
def decorator(func: Callable) -> Callable:
is_async = asyncio.iscoroutinefunction(func)
return _make_wrapper(func, is_async, scene_key, per_unit, unit_field, quantity_field)
return decorator
def _filter_kwargs(func: Callable, kwargs: dict) -> dict:
"""过滤掉目标函数签名不接受的 kwargs(避免 TypeError)。"""
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 _extract_kwargs(func: Callable, args: tuple, kwargs: dict) -> dict:
"""将位置参数映射到函数签名中的参数名,便于统一按 kwargs 提取。"""
sig = inspect.signature(func)
bound = sig.bind_partial(*args, **kwargs)
merged = dict(bound.arguments)
merged.update(kwargs)
return merged
def _execute_with_gate(
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(兼容 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
break
if not current_user:
raise HTTPException(status_code=401, detail="未登录")
# 提取 db session
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(func, args, _filter_kwargs(func, kwargs))
return func(*args, **_filter_kwargs(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(func, args, _filter_kwargs(func, kwargs))
return func(*args, **_filter_kwargs(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(func, args, _filter_kwargs(func, kwargs))
return func(*args, **_filter_kwargs(func, kwargs))
except Exception:
svc.refund_points(user.id, total_points, scene_key, db, ref_id=str(job_id))
raise
def _run_async(func: Callable, args: tuple, kwargs: dict):
"""在 async wrapper 中 await 原始 async 函数。"""
return func(*args, **_filter_kwargs(func, kwargs))