Files
xiaoxia-saas/packages/middleware/points_gate.py
T
xiaoxia 53fb25efcf
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Successful in 2s
CI/CD Pipeline / Check push changed paths (push) Successful in 19s
CI/CD Pipeline / Build Staging API Image (push) Successful in 41s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 48s
CI/CD Pipeline / Integration Tests (push) Successful in 3m10s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 3m17s
CI/CD Pipeline / Build Staging Web Image (push) Successful in 3m30s
CI/CD Pipeline / Validate - Style (push) Successful in 4m17s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 59s
CI/CD Pipeline / Frontend Unit Tests (push) Failing after 5m33s
CI/CD Pipeline / Validate - Security (push) Successful in 7m12s
CI/CD Pipeline / ACR Image Cleanup (push) Successful in 2m38s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m13s
CI/CD Pipeline / Unit Tests (push) Successful in 10m11s
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (push) Failing after 26h14m3s
CI/CD Pipeline / PR Build Worker Image (push) Failing after 26h24m21s
CI/CD Pipeline / Retag skipped Staging API Image (push) Failing after 26h19m47s
CI/CD Pipeline / PR Build Web Image (push) Failing after 26h23m44s
CI/CD Pipeline / PR Build API Image (push) Failing after 26h23m44s
CI/CD Pipeline / Deploy Production (push) Failing after 26h13m23s
CI/CD Pipeline / Build Production Web Image (push) Failing after 26h13m26s
CI/CD Pipeline / CI Gate (push) Failing after 26h13m25s
CI/CD Pipeline / Build Production API Image (push) Failing after 26h13m26s
CI/CD Pipeline / Canary Release to Production (push) Failing after 26h13m23s
CI/CD Pipeline / Retag skipped Staging Web Image (push) Failing after 26h19m46s
CI/CD Pipeline / Frontend Lint (push) Failing after 26h23m37s
CI/CD Pipeline / Check if frontend-only change (push) Failing after 26h23m45s
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Failing after 26h19m46s
fix(#1834): 批量修复 UP 系列静态分析警告(UP007/UP006/UP017/UP035) (#1928)
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
2026-09-15 12:59:17 +08:00

183 lines
5.8 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(
scene_key: str,
per_unit: int | None = None,
unit_field: str | None = None,
quantity_field: str | None = None,
) -> Callable:
"""AI 功能入口积分扣费装饰器。
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:
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:
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 _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
current_user = merged.get("current_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, kwargs)
return func(*args, **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, kwargs)
return func(*args, **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, kwargs)
return func(*args, **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, **kwargs)