Files
xiaoxia-saas/packages/middleware/points_gate.py
T
CI Bot 0c135ca206
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 4s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 4s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m32s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m0s
AI Code Review / AI Code Review (pull_request) Successful in 6m55s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 16s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 16s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 9m57s
CI/CD Pipeline / Integration Tests (pull_request) Failing after 9m35s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 11m16s
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been cancelled
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
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 / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (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 / CI Gate (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 16h43m2s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 16h42m46s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 16h42m46s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 16h42m46s
CI/CD Pipeline / PR Build Web Image (pull_request) Failing after 16h42m18s
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 16h42m19s
CI/CD Pipeline / Check push changed paths (pull_request) Failing after 17h1m48s
style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
2026-09-15 14:07:45 +00:00

218 lines
7.3 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
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 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。
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)
@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
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))