feat: 添加多维度计费系统和视频任务管理功能

计费系统:
- 新增 BillingRule 和 DimensionCollector 数据模型
- 实现 FormulaEngine 安全表达式求值引擎 (AST 白名单)
- 支持 dimension/matrix/tiered/constant 多种维度映射
- BillingRuleService 支持 Provider Model -> GlobalModel 规则回退
- CLI task_type 在计费域自动映射为 chat

视频任务增强:
- 添加 request_metadata 字段记录候选 key 和计费规则快照
- 后台轮询支持并发控制 (Semaphore + 独立 session)
- 任务终态自动写入 Usage 记录并计算成本
- 新增视频任务管理 API 和前端界面

其他改进:
- UsageService 新增 record_usage_with_custom_cost 方法
- StandardizedUsage 支持 dimensions 字段 (兼容 extra)
- 配置新增 BILLING_REQUIRE_RULE 和 BILLING_STRICT_MODE
This commit is contained in:
fawney19
2026-01-31 19:11:25 +08:00
parent dc4bb25cc2
commit 97b15afe7c
30 changed files with 4356 additions and 237 deletions

View File

@@ -9,6 +9,7 @@
"""
from __future__ import annotations
from typing import Any
from src.services.billing.models import (

View File

@@ -0,0 +1,371 @@
"""
DimensionCollector 运行时维度采集
特性(与 .plans/humming-seeking-marble.md 对齐):
- (api_format, task_type) 作用域
- 同一维度支持多条 collectorpriority 回退)
- 支持 transform_expression与 billing expression 共用 AST 安全规范)
- computed 维度支持依赖拓扑排序,并对环依赖做保护性降级
"""
from __future__ import annotations
from collections import deque
from dataclasses import dataclass
from typing import Any, Literal
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import DimensionCollector
from src.services.billing.formula_engine import (
ExpressionEvaluationError,
SafeExpressionEvaluator,
UnsafeExpressionError,
extract_variable_names,
)
ValueType = Literal["float", "int", "string"]
def _normalize_api_format(api_format: str | None) -> str:
return (api_format or "").upper()
def _normalize_task_type(task_type: str | None) -> str:
return (task_type or "").lower()
def _get_nested_value(data: Any, path: str) -> Any:
"""
简单 JSON path
- a.b.c
- 列表索引用数字items.0.id
"""
if data is None or path is None or path == "":
return None
value: Any = data
for key in path.split("."):
if isinstance(value, dict):
value = value.get(key)
elif isinstance(value, list):
if not key.isdigit():
return None
idx = int(key)
if idx < 0 or idx >= len(value):
return None
value = value[idx]
else:
return None
if value is None:
return None
return value
def _cast_value(value: Any, value_type: ValueType) -> Any:
if value_type == "string":
return "" if value is None else str(value)
if value_type == "int":
if value is None:
return 0
if isinstance(value, bool):
raise ValueError("bool is not a valid int dimension value")
return int(float(value))
# float
if value is None:
return 0.0
if isinstance(value, bool):
raise ValueError("bool is not a valid float dimension value")
return float(value)
def _type_default(value_type: ValueType) -> Any:
return "" if value_type == "string" else (0 if value_type == "int" else 0.0)
@dataclass(frozen=True)
class DimensionCollectInput:
request: dict[str, Any] | None = None
response: dict[str, Any] | None = None
metadata: dict[str, Any] | None = None
base_dimensions: dict[str, Any] | None = None
class DimensionCollectorRuntime:
"""纯运行时逻辑(不依赖 DB便于测试与复用。"""
def __init__(self) -> None:
self._evaluator = SafeExpressionEvaluator()
def collect(
self,
*,
collectors: list[DimensionCollector],
inp: DimensionCollectInput,
) -> dict[str, Any]:
dims: dict[str, Any] = dict(inp.base_dimensions or {})
# dimension_name -> collectors (priority desc)
grouped: dict[str, list[DimensionCollector]] = {}
for c in collectors:
grouped.setdefault(c.dimension_name, []).append(c)
for name in grouped:
grouped[name].sort(key=lambda x: (x.priority or 0), reverse=True)
# 1) 先收集非 computed
computed_only: set[str] = set()
for dim_name, cs in grouped.items():
non_computed = [c for c in cs if (c.source_type or "").lower() != "computed"]
if not non_computed:
computed_only.add(dim_name)
continue
value = self._resolve_dimension(dim_name, non_computed, dims, inp)
dims[dim_name] = value
# 2) computed 维度拓扑排序
ordered = self._toposort_computed(grouped, computed_only)
for dim_name in ordered:
cs = [
c for c in grouped.get(dim_name, []) if (c.source_type or "").lower() == "computed"
]
if not cs:
continue
cs.sort(key=lambda x: (x.priority or 0), reverse=True)
value = self._resolve_computed_dimension(dim_name, cs, dims)
dims[dim_name] = value
return dims
def _resolve_dimension(
self,
dim_name: str,
collectors: list[DimensionCollector],
dims: dict[str, Any],
inp: DimensionCollectInput,
) -> Any:
fallback_default: str | None = None
fallback_value_type: ValueType | None = None
value_type: ValueType = (
(collectors[0].value_type or "float").lower() # type: ignore[assignment]
if collectors
else "float"
)
for c in collectors:
value_type = (c.value_type or "float").lower() # type: ignore[assignment]
if c.default_value is not None and fallback_default is None:
fallback_default = c.default_value
fallback_value_type = value_type
src = (c.source_type or "").lower()
path = c.source_path or ""
if src == "request":
raw = _get_nested_value(inp.request or {}, path)
elif src == "response":
raw = _get_nested_value(inp.response or {}, path)
elif src == "metadata":
raw = _get_nested_value(inp.metadata or {}, path)
else:
# 未知 source跳过尝试
continue
if raw is None:
continue
try:
value: Any = raw
if c.transform_expression:
# transform_expression 仅允许使用 value
value = self._evaluator.eval_number(c.transform_expression, {"value": value})
casted = _cast_value(value, value_type)
return casted
except (ValueError, UnsafeExpressionError, ExpressionEvaluationError, Exception) as exc:
# 注意:这里选择“不中断,尝试下一优先级”
logger.debug(
"Dimension collector failed (dim=%s, id=%s): %s",
dim_name,
getattr(c, "id", None),
str(exc),
)
continue
# 兜底default_value仅允许配置一条但这里不依赖 DB 校验)
if fallback_default is not None:
try:
return _cast_value(fallback_default, fallback_value_type or value_type)
except Exception:
return _type_default(fallback_value_type or value_type)
return _type_default(value_type)
def _resolve_computed_dimension(
self,
dim_name: str,
collectors: list[DimensionCollector],
dims: dict[str, Any],
) -> Any:
fallback_default: str | None = None
fallback_value_type: ValueType | None = None
value_type: ValueType = (
(collectors[0].value_type or "float").lower() # type: ignore[assignment]
if collectors
else "float"
)
for c in collectors:
value_type = (c.value_type or "float").lower() # type: ignore[assignment]
if c.default_value is not None and fallback_default is None:
fallback_default = c.default_value
fallback_value_type = value_type
expr = c.transform_expression
if not expr:
continue
try:
value = self._evaluator.eval_number(expr, dims)
casted = _cast_value(value, value_type)
return casted
except (ValueError, ExpressionEvaluationError, UnsafeExpressionError, Exception):
continue
if fallback_default is not None:
try:
return _cast_value(fallback_default, fallback_value_type or value_type)
except Exception:
return _type_default(fallback_value_type or value_type)
return _type_default(value_type)
def _toposort_computed(
self,
grouped: dict[str, list[DimensionCollector]],
computed_only: set[str],
) -> list[str]:
# 建图dependency -> dim
allowed_func_names = set(self._evaluator.ALLOWED_FUNCS.keys())
deps: dict[str, set[str]] = {d: set() for d in computed_only}
for dim_name in computed_only:
for c in grouped.get(dim_name, []):
if (c.source_type or "").lower() != "computed" or not c.transform_expression:
continue
try:
names = extract_variable_names(c.transform_expression)
except UnsafeExpressionError:
# 配置错误:按无依赖处理,避免阻塞
logger.error(
"Invalid computed transform_expression (dim=%s, id=%s)",
dim_name,
getattr(c, "id", None),
)
names = set()
names.discard("value")
names -= allowed_func_names
# 仅关心依赖的 computed 维度(非 computed 会在前一步收集)
deps[dim_name] |= {n for n in names if n in computed_only and n != dim_name}
# Kahn
in_degree: dict[str, int] = {d: 0 for d in computed_only}
forward: dict[str, set[str]] = {d: set() for d in computed_only}
for dim_name, dim_deps in deps.items():
for dep in dim_deps:
forward[dep].add(dim_name)
in_degree[dim_name] += 1
queue = deque(sorted(d for d, deg in in_degree.items() if deg == 0))
ordered: list[str] = []
while queue:
node = queue.popleft()
ordered.append(node)
for nxt in sorted(forward.get(node, set())):
in_degree[nxt] -= 1
if in_degree[nxt] == 0:
queue.append(nxt)
if len(ordered) != len(computed_only):
# 有环依赖:保护性降级(按名称补齐),避免阻塞整条计费链路
remaining = sorted(list(computed_only - set(ordered)))
logger.error("Computed dimension cycle detected: %s", remaining)
ordered.extend(remaining)
return ordered
class DimensionCollectorService:
"""DB + runtime 的封装:读取 collectors 并执行采集。"""
def __init__(self, db: Session):
self.db = db
self._runtime = DimensionCollectorRuntime()
def list_enabled_collectors(
self,
*,
api_format: str | None,
task_type: str | None,
) -> list[DimensionCollector]:
api = _normalize_api_format(api_format)
task = _normalize_task_type(task_type)
api_variants = list({api, api.lower()})
if task == "cli":
# CLI → chat按维度回退维度存在 cli collector 则用 cli否则用 chat
cli_collectors = (
self.db.query(DimensionCollector)
.filter(
DimensionCollector.api_format.in_(api_variants),
DimensionCollector.task_type == "cli",
DimensionCollector.is_enabled == True, # noqa: E712
)
.all()
)
chat_collectors = (
self.db.query(DimensionCollector)
.filter(
DimensionCollector.api_format.in_(api_variants),
DimensionCollector.task_type == "chat",
DimensionCollector.is_enabled == True, # noqa: E712
)
.all()
)
cli_dims: set[str] = {c.dimension_name for c in cli_collectors}
result: list[DimensionCollector] = list(cli_collectors)
for c in chat_collectors:
if c.dimension_name not in cli_dims:
result.append(c)
return result
return (
self.db.query(DimensionCollector)
.filter(
DimensionCollector.api_format.in_(api_variants),
DimensionCollector.task_type == task,
DimensionCollector.is_enabled == True, # noqa: E712
)
.all()
)
def collect_dimensions(
self,
*,
api_format: str | None,
task_type: str | None,
request: dict[str, Any] | None = None,
response: dict[str, Any] | None = None,
metadata: dict[str, Any] | None = None,
base_dimensions: dict[str, Any] | None = None,
) -> dict[str, Any]:
collectors = self.list_enabled_collectors(api_format=api_format, task_type=task_type)
return self._runtime.collect(
collectors=collectors,
inp=DimensionCollectInput(
request=request,
response=response,
metadata=metadata,
base_dimensions=base_dimensions,
),
)

View File

@@ -0,0 +1,368 @@
"""
FormulaEngine - 配置驱动的安全计费表达式引擎
目标:
- 支持 billing_rules.expression 的安全求值AST 白名单)
- 支持 dimension_mappingsdimension/matrix/tiered/constant
- 支持 required/allow_zero 机制,避免维度缺失导致静默少收
注意:该模块不直接依赖数据库;规则查找、维度采集在上层服务完成。
"""
from __future__ import annotations
import ast
from dataclasses import dataclass
from typing import Any, Iterable, Literal
class UnsafeExpressionError(ValueError):
"""表达式包含不安全/不支持的 AST 结构。"""
class ExpressionEvaluationError(RuntimeError):
"""表达式在安全求值阶段失败(如 NameError/ZeroDivision"""
class BillingIncompleteError(RuntimeError):
"""required 维度缺失且 strict_mode=true 时抛出,用于上层拒绝请求/标记任务失败。"""
def __init__(self, message: str, *, missing_required: list[str]):
super().__init__(message)
self.missing_required = missing_required
@dataclass(frozen=True)
class FormulaEvaluationResult:
status: Literal["complete", "incomplete"]
cost: float
resolved_values: dict[str, Any]
missing_required: list[str]
error: str | None = None
_ALLOWED_BINOPS = (
ast.Add,
ast.Sub,
ast.Mult,
ast.Div,
ast.Pow,
ast.FloorDiv,
ast.Mod,
)
_ALLOWED_UNARYOPS = (ast.UAdd, ast.USub)
_ALLOWED_OP_NODES = _ALLOWED_BINOPS + _ALLOWED_UNARYOPS
def _iter_ast_nodes(node: ast.AST) -> Iterable[ast.AST]:
yield node
for child in ast.iter_child_nodes(node):
yield from _iter_ast_nodes(child)
def extract_variable_names(expression: str) -> set[str]:
"""提取表达式中出现的变量名(不含函数名)。"""
try:
tree = ast.parse(expression, mode="eval")
except SyntaxError as exc:
raise UnsafeExpressionError(f"Invalid expression syntax: {exc}") from exc
names: set[str] = set()
for node in _iter_ast_nodes(tree):
if isinstance(node, ast.Name):
names.add(node.id)
if isinstance(node, ast.Call):
# Call 的函数名会以 ast.Name 出现,需要从结果中过滤掉
if isinstance(node.func, ast.Name):
names.discard(node.func.id)
return names
class SafeExpressionEvaluator:
"""AST 白名单 + 无 builtins 的安全求值器。"""
ALLOWED_FUNCS: dict[str, Any] = {
"min": min,
"max": max,
"abs": abs,
"round": round,
"int": int,
"float": float,
}
def validate(self, expression: str) -> ast.Expression:
try:
tree = ast.parse(expression, mode="eval")
except SyntaxError as exc:
raise UnsafeExpressionError(f"Invalid expression syntax: {exc}") from exc
for node in _iter_ast_nodes(tree):
if isinstance(node, ast.Expression):
continue
# 运算符节点本身也会出现在 iter_child_nodes 中
if isinstance(node, _ALLOWED_OP_NODES):
continue
if isinstance(node, ast.Constant):
# 仅允许数字常量bool 是 int 子类,需要显式排除)
if isinstance(node.value, bool) or not isinstance(node.value, (int, float)):
raise UnsafeExpressionError("Only int/float constants are allowed")
continue
if isinstance(node, ast.BinOp):
if not isinstance(node.op, _ALLOWED_BINOPS):
raise UnsafeExpressionError(f"Operator not allowed: {type(node.op).__name__}")
continue
if isinstance(node, ast.UnaryOp):
if not isinstance(node.op, _ALLOWED_UNARYOPS):
raise UnsafeExpressionError(
f"Unary operator not allowed: {type(node.op).__name__}"
)
continue
if isinstance(node, ast.Name):
# 防御:拒绝双下划线变量名
if node.id.startswith("__"):
raise UnsafeExpressionError("Dunder names are not allowed")
continue
if isinstance(node, ast.Load):
continue
if isinstance(node, ast.keyword):
continue
if isinstance(node, ast.Call):
if not isinstance(node.func, ast.Name):
raise UnsafeExpressionError("Only direct function calls are allowed")
func_name = node.func.id
if func_name not in self.ALLOWED_FUNCS:
raise UnsafeExpressionError(f"Function not allowed: {func_name}")
if any(k.arg is None for k in node.keywords):
raise UnsafeExpressionError("**kwargs is not allowed")
continue
# 明确禁止的/不需要的节点类型(属性访问、下标、推导式、比较等)
if isinstance(
node,
(
ast.Attribute,
ast.Subscript,
ast.Compare,
ast.BoolOp,
ast.IfExp,
ast.Lambda,
ast.Dict,
ast.List,
ast.Tuple,
ast.Set,
ast.ListComp,
ast.SetComp,
ast.DictComp,
ast.GeneratorExp,
ast.Await,
ast.Yield,
ast.YieldFrom,
),
):
raise UnsafeExpressionError(f"AST node not allowed: {type(node).__name__}")
raise UnsafeExpressionError(f"AST node not allowed: {type(node).__name__}")
assert isinstance(tree, ast.Expression)
return tree
def eval_number(self, expression: str, variables: dict[str, Any]) -> float:
tree = self.validate(expression)
safe_globals = {"__builtins__": {}}
safe_locals = dict(self.ALLOWED_FUNCS)
safe_locals.update(variables or {})
try:
compiled = compile(tree, "<billing_expr>", "eval")
value = eval(compiled, safe_globals, safe_locals) # noqa: S307 - validated AST
except Exception as exc:
raise ExpressionEvaluationError(str(exc)) from exc
try:
return float(value)
except Exception as exc:
raise ExpressionEvaluationError(f"Expression result is not numeric: {value!r}") from exc
class FormulaEngine:
"""计费表达式引擎:解析 dimension_mappings 并进行安全求值。"""
def __init__(self) -> None:
self._evaluator = SafeExpressionEvaluator()
def evaluate(
self,
*,
expression: str,
variables: dict[str, Any] | None,
dimensions: dict[str, Any] | None,
dimension_mappings: dict[str, dict[str, Any]] | None,
strict_mode: bool = False,
) -> FormulaEvaluationResult:
dims = dimensions or {}
mappings = dimension_mappings or {}
resolved: dict[str, Any] = dict(variables or {})
missing_required: list[str] = []
# 先解析 dimension_mappings产出 expression 变量表
for var_name, mapping in mappings.items():
source = (mapping.get("source") or "constant").lower()
# 显式 constant 映射属于“兜底行为”:如果 variables 已经提供该变量,则不覆盖。
if source == "constant" and var_name in resolved:
continue
value, is_missing = self._resolve_mapping(var_name, mapping, dims)
if is_missing:
missing_required.append(var_name)
continue
resolved[var_name] = value
# required 维度缺失:直接标记 incomplete并由 strict_mode 决定是否抛错)
if missing_required:
if strict_mode:
raise BillingIncompleteError(
f"Missing required dimensions: {missing_required}",
missing_required=missing_required,
)
return FormulaEvaluationResult(
status="incomplete",
cost=0.0,
resolved_values=resolved,
missing_required=missing_required,
)
try:
cost = self._evaluator.eval_number(expression, resolved)
if cost < 0:
# 防御:不允许负数成本(通常表示配置错误)
return FormulaEvaluationResult(
status="incomplete",
cost=0.0,
resolved_values=resolved,
missing_required=[],
error="negative_cost",
)
return FormulaEvaluationResult(
status="complete",
cost=cost,
resolved_values=resolved,
missing_required=[],
)
except (UnsafeExpressionError, ExpressionEvaluationError) as exc:
if strict_mode:
raise
return FormulaEvaluationResult(
status="incomplete",
cost=0.0,
resolved_values=resolved,
missing_required=[],
error=str(exc),
)
def _resolve_mapping(
self,
var_name: str,
mapping: dict[str, Any],
dims: dict[str, Any],
) -> tuple[Any, bool]:
"""
Returns:
(value, is_missing_required)
说明:
- is_missing_required 仅在 required=true 且缺失时为 True
- required=false 的缺失会使用 default 或 0 兜底,并返回 is_missing_required=False
"""
source = (mapping.get("source") or "constant").lower()
required = bool(mapping.get("required", False))
allow_zero = bool(mapping.get("allow_zero", False))
default = mapping.get("default", 0)
def _missing() -> tuple[Any, bool]:
if required:
return None, True
return default, False
if source == "constant":
# constant 默认行为:由 variables 提供dimension_mappings 显式 constant 时仅做兜底
return default, False
if source == "dimension":
key = mapping.get("key") or var_name
raw = dims.get(key)
if raw is None:
return _missing()
if isinstance(raw, str):
if raw == "":
return _missing()
# 尝试将字符串解析为数字,否则按字符串返回(供上层自行决定)
try:
num = float(raw)
if num == 0 and not allow_zero:
return _missing()
return num, False
except Exception:
return raw, False
if isinstance(raw, (int, float)):
if float(raw) == 0 and not allow_zero:
return _missing()
return raw, False
# 其他类型:尽量转为 float否则视为缺失
try:
num = float(raw)
if num == 0 and not allow_zero:
return _missing()
return num, False
except Exception:
return _missing()
if source == "matrix":
key = mapping.get("key") or var_name
raw = dims.get(key)
if raw is None or raw == "":
return _missing()
raw_key = str(raw)
matrix = mapping.get("map") or {}
if raw_key in matrix:
return matrix[raw_key], False
# matrix 未命中:若 required=true 则仍视为缺失;否则使用 default
if required:
return None, True
return default, False
if source == "tiered":
tier_key = mapping.get("tier_key")
if not tier_key:
return _missing()
raw_tier_value = dims.get(tier_key)
if raw_tier_value is None:
return _missing()
try:
tier_value = float(raw_tier_value)
except Exception:
return _missing()
if tier_value == 0 and not allow_zero:
return _missing()
tiers = mapping.get("tiers") or []
# tiers: [{up_to: 128000, value: 2.5}, {up_to: null, value: 1.25}]
for tier in tiers:
up_to = tier.get("up_to")
if up_to is None:
return tier.get("value", default), False
try:
if tier_value <= float(up_to):
return tier.get("value", default), False
except Exception:
# up_to 配置异常:忽略并继续
continue
# 无匹配:使用最后一个或 default
if tiers:
return tiers[-1].get("value", default), False
return default, False
# 未知 source视为配置错误但不直接中断计费返回 default
return default, False

View File

@@ -9,6 +9,7 @@
"""
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
@@ -89,7 +90,7 @@ class BillingDimension:
)
@dataclass
@dataclass(init=False)
class StandardizedUsage:
"""
标准化的 Usage 数据
@@ -114,8 +115,49 @@ class StandardizedUsage:
# 请求计数(用于按次计费)
request_count: int = 1
# 扩展字段(未来可能需要的额外维度
extra: dict[str, Any] = field(default_factory=dict)
# 任意维度存储(用于多维度计费;数值/字符串均可
# 兼容旧字段名extra 作为 dimensions 的别名
dimensions: dict[str, Any] = field(default_factory=dict)
def __init__(
self,
*,
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
reasoning_tokens: int = 0,
cache_storage_token_hours: float = 0.0,
request_count: int = 1,
dimensions: dict[str, Any] | None = None,
extra: dict[str, Any] | None = None,
) -> None:
# 基础字段
self.input_tokens = input_tokens
self.output_tokens = output_tokens
self.cache_creation_tokens = cache_creation_tokens
self.cache_read_tokens = cache_read_tokens
self.reasoning_tokens = reasoning_tokens
self.cache_storage_token_hours = cache_storage_token_hours
self.request_count = request_count
# 兼容:支持 extra 与 dimensions 同时传入dimensions 优先级更高)
merged: dict[str, Any] = {}
if isinstance(extra, dict):
merged.update(extra)
if isinstance(dimensions, dict):
merged.update(dimensions)
self.dimensions = merged
@property
def extra(self) -> dict[str, Any]:
"""向后兼容:旧代码使用 usage.extra 访问扩展维度。"""
return self.dimensions
@extra.setter
def extra(self, value: dict[str, Any]) -> None:
"""向后兼容:允许旧代码写入 usage.extra。"""
self.dimensions = value or {}
def get(self, field_name: str, default: Any = 0) -> Any:
"""
@@ -130,12 +172,14 @@ class StandardizedUsage:
Returns:
字段值
"""
if hasattr(self, field_name):
value = getattr(self, field_name)
# 对于 extra 字段,不直接返回
if field_name != "extra":
return value
return self.extra.get(field_name, default)
# 兼容旧字段名
if field_name == "extra":
return self.dimensions
if hasattr(self, field_name) and field_name not in {"dimensions"}:
return getattr(self, field_name)
return self.dimensions.get(field_name, default)
def set(self, field_name: str, value: Any) -> None:
"""
@@ -145,10 +189,16 @@ class StandardizedUsage:
field_name: 字段名
value: 字段值
"""
if hasattr(self, field_name) and field_name != "extra":
# 兼容旧字段名
if field_name == "extra":
self.dimensions = value or {}
return
if hasattr(self, field_name) and field_name not in {"dimensions"}:
setattr(self, field_name, value)
else:
self.extra[field_name] = value
return
self.dimensions[field_name] = value
def to_dict(self) -> dict[str, Any]:
"""转换为字典"""
@@ -161,14 +211,24 @@ class StandardizedUsage:
"cache_storage_token_hours": self.cache_storage_token_hours,
"request_count": self.request_count,
}
if self.extra:
result["extra"] = self.extra
if self.dimensions:
# 新字段名
result["dimensions"] = self.dimensions
# 旧字段名(兼容)
result["extra"] = self.dimensions
return result
@classmethod
def from_dict(cls, data: dict[str, Any]) -> StandardizedUsage:
"""从字典创建实例"""
# 兼容:支持 extra / dimensions 两种键名
extra = data.pop("extra", {}) if "extra" in data else {}
dimensions = data.pop("dimensions", {}) if "dimensions" in data else {}
merged_dimensions: dict[str, Any] = {}
if isinstance(extra, dict):
merged_dimensions.update(extra)
if isinstance(dimensions, dict):
merged_dimensions.update(dimensions)
# 只取已知字段
known_fields = {
"input_tokens",
@@ -180,7 +240,7 @@ class StandardizedUsage:
"request_count",
}
filtered = {k: v for k, v in data.items() if k in known_fields}
return cls(**filtered, extra=extra)
return cls(**filtered, dimensions=merged_dimensions)
@dataclass

View File

@@ -0,0 +1,101 @@
"""
BillingRule 查找逻辑
查找顺序(与 .plans/humming-seeking-marble.md 一致):
1) ModelProvider 级)→ 2) GlobalModel默认
注意:
- CLI 在计费域等同于 chatbilling_rules.task_type 不含 "cli"
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal
from sqlalchemy.orm import Session
from src.models.database import BillingRule, GlobalModel, Model
TaskType = Literal["chat", "cli", "video", "image", "audio"]
def effective_rule_task_type(task_type: str) -> str:
"""CLI 在计费规则域里恒等于 chat。"""
t = (task_type or "").lower()
return "chat" if t == "cli" else t
@dataclass(frozen=True)
class BillingRuleLookupResult:
rule: BillingRule
scope: Literal["model", "global"]
effective_task_type: str
class BillingRuleService:
@staticmethod
def find_rule(
db: Session,
*,
provider_id: str | None,
model_name: str,
task_type: str,
) -> BillingRuleLookupResult | None:
effective_task = effective_rule_task_type(task_type)
global_model = (
db.query(GlobalModel)
.filter(
GlobalModel.name == model_name,
GlobalModel.is_active == True, # noqa: E712
)
.first()
)
if not global_model:
return None
# 1) Provider Model 覆盖
if provider_id:
model_obj = (
db.query(Model)
.filter(
Model.provider_id == provider_id,
Model.global_model_id == global_model.id,
Model.is_active == True, # noqa: E712
)
.first()
)
if model_obj:
rule = (
db.query(BillingRule)
.filter(
BillingRule.model_id == model_obj.id,
BillingRule.task_type == effective_task,
BillingRule.is_enabled == True, # noqa: E712
)
.first()
)
if rule:
return BillingRuleLookupResult(
rule=rule,
scope="model",
effective_task_type=effective_task,
)
# 2) GlobalModel 默认规则
rule = (
db.query(BillingRule)
.filter(
BillingRule.global_model_id == global_model.id,
BillingRule.task_type == effective_task,
BillingRule.is_enabled == True, # noqa: E712
)
.first()
)
if rule:
return BillingRuleLookupResult(
rule=rule, scope="global", effective_task_type=effective_task
)
return None

View File

@@ -9,7 +9,6 @@
- PER_REQUEST: 按次计费
"""
from src.services.billing.models import BillingDimension, BillingUnit

View File

@@ -8,9 +8,9 @@
from __future__ import annotations
from typing import Any, Callable
import os
from datetime import datetime
from typing import Any, Callable
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.cron import CronTrigger
@@ -177,6 +177,19 @@ class TaskScheduler:
f"({hours}小时{minutes}分钟后)"
)
def remove_job(self, job_id: str) -> None:
"""
移除指定的定时任务
Args:
job_id: 任务ID
"""
try:
self.scheduler.remove_job(job_id)
logger.info(f"已移除定时任务: {job_id}")
except Exception as e:
logger.warning(f"移除定时任务失败 {job_id}: {e}")
@property
def is_running(self) -> bool:
"""调度器是否在运行"""

View File

@@ -30,6 +30,7 @@ from src.services.system.config import SystemConfigService
@dataclass
class UsageRecordParams:
"""用量记录参数数据类,用于在内部方法间传递数据"""
db: Session
user: User | None
api_key: ApiKey | None
@@ -76,9 +77,7 @@ class UsageRecordParams:
f"cache_creation_input_tokens 不能为负数: {self.cache_creation_input_tokens}"
)
if self.cache_read_input_tokens < 0:
raise ValueError(
f"cache_read_input_tokens 不能为负数: {self.cache_read_input_tokens}"
)
raise ValueError(f"cache_read_input_tokens 不能为负数: {self.cache_read_input_tokens}")
# 响应时间不能为负数
if self.response_time_ms is not None and self.response_time_ms < 0:
@@ -170,9 +169,10 @@ class UsageService:
Returns:
热力图数据字典
"""
import json
from src.clients.redis_client import get_redis_client
from src.config.constants import CacheTTL
import json
cache_key = cls._get_heatmap_cache_key(user_id, include_actual_cost)
@@ -261,8 +261,8 @@ class UsageService:
request_cost: float,
total_cost: float,
# 价格信息
input_price: float,
output_price: float,
input_price: float | None,
output_price: float | None,
cache_creation_price: float | None,
cache_read_price: float | None,
request_price: float | None,
@@ -415,8 +415,21 @@ class UsageService:
cache_ttl_minutes: int | None,
use_tiered_pricing: bool,
is_failed_request: bool,
) -> tuple[float, float, float, float, float, float, float, float, float,
float | None, float | None, float | None, int | None]:
) -> tuple[
float,
float,
float,
float,
float,
float,
float,
float,
float,
float | None,
float | None,
float | None,
int | None,
]:
"""计算所有成本相关数据
Returns:
@@ -538,9 +551,19 @@ class UsageService:
)
return (
input_price, output_price, cache_creation_price, cache_read_price, request_price,
input_cost, output_cost, cache_creation_cost, cache_read_cost, cache_cost,
request_cost, total_cost, tier_index
input_price,
output_price,
cache_creation_price,
cache_read_price,
request_price,
input_cost,
output_cost,
cache_creation_cost,
cache_read_cost,
cache_cost,
request_cost,
total_cost,
tier_index,
)
@staticmethod
@@ -584,7 +607,9 @@ class UsageService:
existing_usage.total_cost_usd = usage_params["total_cost_usd"]
existing_usage.actual_input_cost_usd = usage_params["actual_input_cost_usd"]
existing_usage.actual_output_cost_usd = usage_params["actual_output_cost_usd"]
existing_usage.actual_cache_creation_cost_usd = usage_params["actual_cache_creation_cost_usd"]
existing_usage.actual_cache_creation_cost_usd = usage_params[
"actual_cache_creation_cost_usd"
]
existing_usage.actual_cache_read_cost_usd = usage_params["actual_cache_read_cost_usd"]
existing_usage.actual_request_cost_usd = usage_params["actual_request_cost_usd"]
existing_usage.actual_total_cost_usd = usage_params["actual_total_cost_usd"]
@@ -646,9 +671,7 @@ class UsageService:
return service.get_cache_prices(provider, model, input_price)
@classmethod
async def get_request_price_async(
cls, db: Session, provider: str, model: str
) -> float | None:
async def get_request_price_async(cls, db: Session, provider: str, model: str) -> float | None:
"""异步获取模型按次计费价格"""
service = ModelCostService(db)
return await service.get_request_price_async(provider, model)
@@ -748,9 +771,19 @@ class UsageService:
# 计算成本
is_failed_request = params.status_code >= 400 or params.error_message is not None
(
input_price, output_price, cache_creation_price, cache_read_price, request_price,
input_cost, output_cost, cache_creation_cost, cache_read_cost, cache_cost,
request_cost, total_cost, _tier_index
input_price,
output_price,
cache_creation_price,
cache_read_price,
request_price,
input_cost,
output_cost,
cache_creation_cost,
cache_read_cost,
cache_cost,
request_cost,
total_cost,
_tier_index,
) = await cls._calculate_costs(
db=params.db,
provider=params.provider,
@@ -834,7 +867,9 @@ class UsageService:
"""
import asyncio
async def prepare_single(params: UsageRecordParams) -> tuple[dict[str, Any], float, Exception | None]:
async def prepare_single(
params: UsageRecordParams,
) -> tuple[dict[str, Any], float, Exception | None]:
try:
usage_params, total_cost = await cls._prepare_usage_record(params)
return (usage_params, total_cost, None)
@@ -904,23 +939,38 @@ class UsageService:
# 使用共享逻辑准备记录参数
params = UsageRecordParams(
db=db, user=user, api_key=api_key, provider=provider, model=model,
input_tokens=input_tokens, output_tokens=output_tokens,
db=db,
user=user,
api_key=api_key,
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
request_type=request_type, api_format=api_format,
endpoint_api_format=endpoint_api_format, has_format_conversion=has_format_conversion,
request_type=request_type,
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
is_stream=is_stream,
response_time_ms=response_time_ms, first_byte_time_ms=first_byte_time_ms,
status_code=status_code, error_message=error_message, metadata=metadata,
request_headers=request_headers, request_body=request_body,
response_time_ms=response_time_ms,
first_byte_time_ms=first_byte_time_ms,
status_code=status_code,
error_message=error_message,
metadata=metadata,
request_headers=request_headers,
request_body=request_body,
provider_request_headers=provider_request_headers,
response_headers=response_headers, client_response_headers=client_response_headers,
response_headers=response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
request_id=request_id, provider_id=provider_id,
request_id=request_id,
provider_id=provider_id,
provider_endpoint_id=provider_endpoint_id,
provider_api_key_id=provider_api_key_id, status=status,
cache_ttl_minutes=cache_ttl_minutes, use_tiered_pricing=use_tiered_pricing,
provider_api_key_id=provider_api_key_id,
status=status,
cache_ttl_minutes=cache_ttl_minutes,
use_tiered_pricing=use_tiered_pricing,
target_model=target_model,
)
usage_params, _ = await cls._prepare_usage_record(params)
@@ -931,6 +981,7 @@ class UsageService:
# 更新 GlobalModel 使用计数(原子操作)
from sqlalchemy import update
from src.models.database import GlobalModel
db.execute(
@@ -1003,23 +1054,38 @@ class UsageService:
# 使用共享逻辑准备记录参数
params = UsageRecordParams(
db=db, user=user, api_key=api_key, provider=provider, model=model,
input_tokens=input_tokens, output_tokens=output_tokens,
db=db,
user=user,
api_key=api_key,
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
request_type=request_type, api_format=api_format,
endpoint_api_format=endpoint_api_format, has_format_conversion=has_format_conversion,
request_type=request_type,
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
is_stream=is_stream,
response_time_ms=response_time_ms, first_byte_time_ms=first_byte_time_ms,
status_code=status_code, error_message=error_message, metadata=metadata,
request_headers=request_headers, request_body=request_body,
response_time_ms=response_time_ms,
first_byte_time_ms=first_byte_time_ms,
status_code=status_code,
error_message=error_message,
metadata=metadata,
request_headers=request_headers,
request_body=request_body,
provider_request_headers=provider_request_headers,
response_headers=response_headers, client_response_headers=client_response_headers,
response_headers=response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
request_id=request_id, provider_id=provider_id,
request_id=request_id,
provider_id=provider_id,
provider_endpoint_id=provider_endpoint_id,
provider_api_key_id=provider_api_key_id, status=status,
cache_ttl_minutes=cache_ttl_minutes, use_tiered_pricing=use_tiered_pricing,
provider_api_key_id=provider_api_key_id,
status=status,
cache_ttl_minutes=cache_ttl_minutes,
use_tiered_pricing=use_tiered_pricing,
target_model=target_model,
)
usage_params, total_cost = await cls._prepare_usage_record(params)
@@ -1044,8 +1110,12 @@ class UsageService:
api_key = db.merge(api_key)
# 使用原子更新避免并发竞态条件
from sqlalchemy import func as sql_func, update
from src.models.database import ApiKey as ApiKeyModel, User as UserModel, GlobalModel
from sqlalchemy import func as sql_func
from sqlalchemy import update
from src.models.database import ApiKey as ApiKeyModel
from src.models.database import GlobalModel
from src.models.database import User as UserModel
# 更新用户使用量(独立 Key 不计入创建者的使用记录)
if user and not (api_key and api_key.is_standalone):
@@ -1111,6 +1181,227 @@ class UsageService:
return usage
@classmethod
async def record_usage_with_custom_cost(
cls,
*,
db: Session,
user: User | None,
api_key: ApiKey | None,
provider: str,
model: str,
request_type: str,
total_cost_usd: float,
request_cost_usd: float | None = None,
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_input_tokens: int = 0,
cache_read_input_tokens: int = 0,
api_format: str | None = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
is_stream: bool = False,
response_time_ms: int | None = None,
first_byte_time_ms: int | None = None,
status_code: int = 200,
error_message: str | None = None,
metadata: dict[str, Any] | None = None,
request_headers: dict[str, Any] | None = None,
request_body: Any | None = None,
provider_request_headers: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
response_body: Any | None = None,
request_id: str | None = None,
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
status: str = "completed",
target_model: str | None = None,
) -> Usage:
"""
记录“已计算好的”成本(用于 Video/Image/Audio 等异步任务的 FormulaEngine 计费结果)。
说明:
- 仍然会应用 ProviderAPIKey.rate_multipliers 计算 actual_* 成本
- 会更新 User/APIKey/GlobalModel/Provider 的统计(与 record_usage 行为一致)
- 若 request_id 已存在则更新记录(避免重复写入)
"""
# 生成 request_id
if request_id is None:
request_id = str(uuid.uuid4())[:8]
# 获取费率倍数与免费套餐
actual_rate_multiplier, is_free_tier = await cls._get_rate_multiplier_and_free_tier(
db, provider_api_key_id, provider_id, api_format
)
# 成本拆分:非 token 计费默认计入 request_cost
input_cost = 0.0
output_cost = 0.0
cache_creation_cost = 0.0
cache_read_cost = 0.0
cache_cost = 0.0
request_cost = (
float(request_cost_usd) if request_cost_usd is not None else float(total_cost_usd)
)
total_cost = float(total_cost_usd)
usage_params = cls._build_usage_params(
db=db,
user=user,
api_key=api_key,
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
request_type=request_type,
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
is_stream=is_stream,
response_time_ms=response_time_ms,
first_byte_time_ms=first_byte_time_ms,
status_code=status_code,
error_message=error_message,
metadata=metadata,
request_headers=request_headers,
request_body=request_body,
provider_request_headers=provider_request_headers,
response_headers=response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
request_id=request_id,
provider_id=provider_id,
provider_endpoint_id=provider_endpoint_id,
provider_api_key_id=provider_api_key_id,
status=status,
target_model=target_model,
input_cost=input_cost,
output_cost=output_cost,
cache_creation_cost=cache_creation_cost,
cache_read_cost=cache_read_cost,
cache_cost=cache_cost,
request_cost=request_cost,
total_cost=total_cost,
# token 价格对异步任务不适用,保持 None
input_price=None,
output_price=None,
cache_creation_price=None,
cache_read_price=None,
request_price=None,
actual_rate_multiplier=actual_rate_multiplier,
is_free_tier=is_free_tier,
)
# Upsert与 record_usage 保持一致)
existing_usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if existing_usage:
# 避免重复记账:若已是终态记录,直接返回(批量接口也采用该策略)
if existing_usage.status not in ("pending", "streaming"):
logger.debug(
"record_usage_with_custom_cost: request_id=%s already finalized (status=%s), skip",
request_id,
existing_usage.status,
)
return existing_usage
cls._update_existing_usage(existing_usage, usage_params, target_model)
usage = existing_usage
else:
usage = Usage(**usage_params)
db.add(usage)
# 确保 user 和 api_key 在会话中(与 record_usage 保持一致)
if user and not db.object_session(user):
user = db.merge(user)
if api_key and not db.object_session(api_key):
api_key = db.merge(api_key)
# 原子更新统计
from sqlalchemy import func as sql_func
from sqlalchemy import update
from src.models.database import ApiKey as ApiKeyModel
from src.models.database import GlobalModel
from src.models.database import User as UserModel
# 更新用户使用量(独立 Key 不计入创建者)
if user and not (api_key and api_key.is_standalone):
db.execute(
update(UserModel)
.where(UserModel.id == user.id)
.values(
used_usd=UserModel.used_usd + total_cost,
total_usd=UserModel.total_usd + total_cost,
updated_at=sql_func.now(),
)
)
# 更新 API 密钥使用量
if api_key:
if api_key.is_standalone:
db.execute(
update(ApiKeyModel)
.where(ApiKeyModel.id == api_key.id)
.values(
total_requests=ApiKeyModel.total_requests + 1,
total_cost_usd=ApiKeyModel.total_cost_usd + total_cost,
balance_used_usd=ApiKeyModel.balance_used_usd + total_cost,
last_used_at=sql_func.now(),
updated_at=sql_func.now(),
)
)
else:
db.execute(
update(ApiKeyModel)
.where(ApiKeyModel.id == api_key.id)
.values(
total_requests=ApiKeyModel.total_requests + 1,
total_cost_usd=ApiKeyModel.total_cost_usd + total_cost,
last_used_at=sql_func.now(),
updated_at=sql_func.now(),
)
)
# 更新 GlobalModel 使用计数
db.execute(
update(GlobalModel)
.where(GlobalModel.name == model)
.values(usage_count=GlobalModel.usage_count + 1)
)
# 更新 Provider 月度使用量(使用 actual_total_cost
if provider_id:
actual_total_cost = usage_params["actual_total_cost_usd"]
db.execute(
update(Provider)
.where(Provider.id == provider_id)
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
)
try:
db.commit()
except Exception as e:
# 并发场景可能触发唯一约束冲突:降级为读取已存在记录
try:
from sqlalchemy.exc import IntegrityError
if isinstance(e, IntegrityError):
db.rollback()
existing = db.query(Usage).filter(Usage.request_id == request_id).first()
if existing:
return existing
except Exception:
pass
logger.error(f"提交使用记录时出错: {e}")
db.rollback()
raise
return usage
@classmethod
async def record_usage_batch(
cls,
@@ -1136,8 +1427,12 @@ class UsageService:
return []
from collections import defaultdict
from sqlalchemy import update
from src.models.database import ApiKey as ApiKeyModel, User as UserModel, GlobalModel
from src.models.database import ApiKey as ApiKeyModel
from src.models.database import GlobalModel
from src.models.database import User as UserModel
# 分离需要更新和需要新建的记录
request_ids = [r.get("request_id") for r in records if r.get("request_id")]
@@ -1147,11 +1442,7 @@ class UsageService:
if request_ids:
# 查询已存在的 Usage 记录(包括 pending/streaming 状态)
existing_records = (
db.query(Usage)
.filter(Usage.request_id.in_(request_ids))
.all()
)
existing_records = db.query(Usage).filter(Usage.request_id.in_(request_ids)).all()
existing_usages = {u.request_id: u for u in existing_records}
for record in records:
@@ -1280,8 +1571,8 @@ class UsageService:
prepared_results = []
# 分配准备结果
update_results = prepared_results[:len(update_params_list)]
insert_results = prepared_results[len(update_params_list):]
update_results = prepared_results[: len(update_params_list)]
insert_results = prepared_results[len(update_params_list) :]
# 1. 处理需要更新的记录
for i, (record, request_id, params) in enumerate(update_params_list):
@@ -1369,12 +1660,12 @@ class UsageService:
if skip_ratio > 0.1:
logger.error(
"批量记录失败率过高: %d/%d (%.1f%%) 条记录被跳过",
skipped_count, total_count, skip_ratio * 100
skipped_count,
total_count,
skip_ratio * 100,
)
else:
logger.warning(
"批量记录部分失败: %d/%d 条记录被跳过", skipped_count, total_count
)
logger.warning("批量记录部分失败: %d/%d 条记录被跳过", skipped_count, total_count)
# 批量更新 GlobalModel 使用计数
for model_name, count in model_counts.items():
@@ -1395,6 +1686,7 @@ class UsageService:
# 批量更新用户使用量
from sqlalchemy import func as sql_func
for user_id, cost in user_costs.items():
if cost > 0:
db.execute(
@@ -1438,9 +1730,7 @@ class UsageService:
db.commit()
inserted_count = len(usages) - updated_count
if updated_count > 0:
logger.debug(
f"批量记录成功: 更新 {updated_count} 条, 新建 {inserted_count}"
)
logger.debug(f"批量记录成功: 更新 {updated_count} 条, 新建 {inserted_count}")
else:
logger.debug(f"批量记录 {len(usages)} 条使用记录成功")
except Exception as e:
@@ -1832,10 +2122,7 @@ class UsageService:
while True:
# 查询待删除的 ID使用新索引 idx_usage_user_created
batch_ids = (
db.query(Usage.id)
.filter(Usage.created_at < cutoff_date)
.limit(batch_size)
.all()
db.query(Usage.id).filter(Usage.created_at < cutoff_date).limit(batch_size).all()
)
if not batch_ids:
@@ -2112,7 +2399,9 @@ class UsageService:
if count > 0:
db.commit()
logger.info(f"清理超时请求: 将 {count} 条超过 {timeout_minutes} 分钟的 pending/streaming 请求标记为 failed")
logger.info(
f"清理超时请求: 将 {count} 条超过 {timeout_minutes} 分钟的 pending/streaming 请求标记为 failed"
)
return count
@@ -2240,9 +2529,7 @@ class UsageService:
# 如果流已经成功完成stream_completed: true不应该标记为超时
# 先获取这些 Usage 的 request_id
usage_request_ids = (
db.query(Usage.id, Usage.request_id)
.filter(Usage.id.in_(timeout_candidates))
.all()
db.query(Usage.id, Usage.request_id).filter(Usage.id.in_(timeout_candidates)).all()
)
usage_id_to_request_id = {u.id: u.request_id for u in usage_request_ids}
request_id_to_usage_id = {u.request_id: u.id for u in usage_request_ids}
@@ -2278,9 +2565,7 @@ class UsageService:
for candidate in candidates:
extra_data = candidate.extra_data or {}
# 情况1status='success' 且 stream_completed=True
if candidate.status == "success" and extra_data.get(
"stream_completed", False
):
if candidate.status == "success" and extra_data.get("stream_completed", False):
usage_id = request_id_to_usage_id.get(candidate.request_id)
if usage_id:
completed_usage_ids.add(usage_id)
@@ -2321,11 +2606,7 @@ class UsageService:
has_format_conversion = getattr(r, "has_format_conversion", None)
# 兼容历史数据:当 streaming 状态已拿到两个格式但 has_format_conversion 为空时,回填推断结果
if (
has_format_conversion is None
and api_format
and endpoint_api_format
):
if has_format_conversion is None and api_format and endpoint_api_format:
has_format_conversion = not can_passthrough(api_format, endpoint_api_format)
item: dict[str, Any] = {
@@ -2336,8 +2617,12 @@ class UsageService:
"cache_creation_input_tokens": r.cache_creation_input_tokens,
"cache_read_input_tokens": r.cache_read_input_tokens,
"cost": float(r.total_cost_usd) if r.total_cost_usd else 0,
"actual_cost": float(r.actual_total_cost_usd) if r.actual_total_cost_usd is not None else None,
"rate_multiplier": float(r.rate_multiplier) if r.rate_multiplier is not None else None,
"actual_cost": (
float(r.actual_total_cost_usd) if r.actual_total_cost_usd is not None else None
),
"rate_multiplier": (
float(r.rate_multiplier) if r.rate_multiplier is not None else None
),
"response_time_ms": r.response_time_ms,
"first_byte_time_ms": r.first_byte_time_ms, # 首字时间 (TTFB)
}
@@ -2484,47 +2769,47 @@ class UsageService:
) = row
# 计算推荐 TTL
recommended_ttl = UsageService._calculate_recommended_ttl(
p75_interval, p90_interval
)
recommended_ttl = UsageService._calculate_recommended_ttl(p75_interval, p90_interval)
# 获取用户信息
user_info = user_info_map.get(str(group_id), {})
# 计算各区间占比
total_intervals = request_count
users_analysis.append({
"group_id": group_id,
"username": user_info.get("username"),
"email": user_info.get("email"),
"request_count": request_count,
"interval_distribution": {
"within_5min": within_5min,
"within_15min": within_15min,
"within_30min": within_30min,
"within_60min": within_60min,
"over_60min": over_60min,
},
"interval_percentages": {
"within_5min": round(within_5min / total_intervals * 100, 1),
"within_15min": round(within_15min / total_intervals * 100, 1),
"within_30min": round(within_30min / total_intervals * 100, 1),
"within_60min": round(within_60min / total_intervals * 100, 1),
"over_60min": round(over_60min / total_intervals * 100, 1),
},
"percentiles": {
"p50": round(float(median_interval), 2) if median_interval else None,
"p75": round(float(p75_interval), 2) if p75_interval else None,
"p90": round(float(p90_interval), 2) if p90_interval else None,
},
"avg_interval_minutes": round(float(avg_interval), 2) if avg_interval else None,
"min_interval_minutes": round(float(min_interval), 2) if min_interval else None,
"max_interval_minutes": round(float(max_interval), 2) if max_interval else None,
"recommended_ttl_minutes": recommended_ttl,
"recommendation_reason": UsageService._get_ttl_recommendation_reason(
recommended_ttl, p75_interval, p90_interval
),
})
users_analysis.append(
{
"group_id": group_id,
"username": user_info.get("username"),
"email": user_info.get("email"),
"request_count": request_count,
"interval_distribution": {
"within_5min": within_5min,
"within_15min": within_15min,
"within_30min": within_30min,
"within_60min": within_60min,
"over_60min": over_60min,
},
"interval_percentages": {
"within_5min": round(within_5min / total_intervals * 100, 1),
"within_15min": round(within_15min / total_intervals * 100, 1),
"within_30min": round(within_30min / total_intervals * 100, 1),
"within_60min": round(within_60min / total_intervals * 100, 1),
"over_60min": round(over_60min / total_intervals * 100, 1),
},
"percentiles": {
"p50": round(float(median_interval), 2) if median_interval else None,
"p75": round(float(p75_interval), 2) if p75_interval else None,
"p90": round(float(p90_interval), 2) if p90_interval else None,
},
"avg_interval_minutes": round(float(avg_interval), 2) if avg_interval else None,
"min_interval_minutes": round(float(min_interval), 2) if min_interval else None,
"max_interval_minutes": round(float(max_interval), 2) if max_interval else None,
"recommended_ttl_minutes": recommended_ttl,
"recommendation_reason": UsageService._get_ttl_recommendation_reason(
recommended_ttl, p75_interval, p90_interval
),
}
)
# 汇总统计
ttl_distribution = {"5min": 0, "15min": 0, "30min": 0, "60min": 0}
@@ -2684,7 +2969,11 @@ class UsageService:
"analysis_period_hours": hours,
"total_requests": total_requests,
"requests_with_cache_hit": requests_with_cache_hit_count,
"request_cache_hit_rate": round(requests_with_cache_hit_count / total_requests * 100, 2) if total_requests > 0 else 0,
"request_cache_hit_rate": (
round(requests_with_cache_hit_count / total_requests * 100, 2)
if total_requests > 0
else 0
),
"total_input_tokens": total_input_tokens,
"total_cache_read_tokens": total_cache_read_tokens,
"total_cache_creation_tokens": total_cache_creation_tokens,
@@ -2839,10 +3128,7 @@ class UsageService:
else:
for row in rows:
created_at, model, interval_minutes = row
point_data = {
"x": created_at.isoformat(),
"y": round(float(interval_minutes), 2)
}
point_data = {"x": created_at.isoformat(), "y": round(float(interval_minutes), 2)}
if model:
point_data["model"] = model
models_set.add(model)

View File

@@ -6,14 +6,19 @@ from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Any
from uuid import uuid4
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
from src.api.handlers.base.video_handler_base import sanitize_error_message
from src.api.handlers.base.video_handler_base import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.clients.http_client import HTTPClientPool
from src.clients.redis_client import get_redis_client
from src.config.settings import config
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
@@ -21,8 +26,12 @@ from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.models.database import ApiKey, Provider, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.system.scheduler import get_scheduler
from src.services.usage.service import UsageService
# 永久性错误指示词(用于降级判断,不应重试)
_PERMANENT_ERROR_INDICATORS = frozenset(
@@ -53,7 +62,6 @@ class VideoTaskPollerService:
LOCK_KEY = "video_task_poller:lock"
LOCK_TTL = 60
BATCH_SIZE = 50
MAX_BACKOFF_SECONDS = 300
# 连续失败告警阈值
CONSECUTIVE_FAILURE_ALERT_THRESHOLD = 5
@@ -63,17 +71,26 @@ class VideoTaskPollerService:
self.redis = None
self._openai_normalizer = OpenAINormalizer()
self._gemini_normalizer = GeminiNormalizer()
self._formula_engine = FormulaEngine()
# 追踪连续失败次数(用于告警)
self._consecutive_failures = 0
# 从配置读取参数
self._batch_size = config.video_poll_batch_size
self._concurrency = config.video_poll_concurrency
# Semaphore 延迟初始化,避免在事件循环外创建
self._semaphore: asyncio.Semaphore | None = None
async def start(self) -> None:
# 在事件循环内初始化 Semaphore
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self._concurrency)
if self.redis is None:
self.redis = await get_redis_client(require_redis=False)
scheduler = get_scheduler()
scheduler.add_interval_job(
self.poll_pending_tasks,
seconds=10,
seconds=config.video_poll_interval_seconds,
job_id="video_task_poller",
name="视频任务轮询",
)
@@ -106,7 +123,7 @@ class VideoTaskPollerService:
VideoTask.poll_count < VideoTask.max_poll_count,
)
.order_by(VideoTask.next_poll_at.asc())
.limit(self.BATCH_SIZE)
.limit(self._batch_size)
.all()
)
@@ -115,32 +132,55 @@ class VideoTaskPollerService:
self._consecutive_failures = 0
return
batch_failures = 0
for task in tasks:
# 提取任务 ID 列表,释放查询 session 后逐个轮询
task_ids = [t.id for t in tasks]
# 并发轮询:每个任务使用独立 session避免共享 session 的并发风险
poll_results: list[bool] = []
# 确保 semaphore 已初始化(在 start 中初始化,此处防御性检查)
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self._concurrency)
semaphore = self._semaphore
async def poll_with_semaphore(task_id: str) -> None:
"""带信号量的轮询,结果写入 poll_results"""
async with semaphore:
try:
await self._poll_single_task(db, task)
with create_session() as task_db:
task_obj = task_db.query(VideoTask).get(task_id)
if not task_obj:
logger.warning("Task %s disappeared during poll", task_id)
poll_results.append(True)
return
await self._poll_single_task(task_db, task_obj)
task_db.commit()
poll_results.append(True)
except Exception as exc:
batch_failures += 1
# 单个任务失败不影响其他任务处理
logger.exception(
"Unexpected error polling task %s: %s",
task.id,
task_id,
sanitize_error_message(str(exc)),
)
poll_results.append(False)
# 更新连续失败计数并检查告警阈值
if batch_failures == len(tasks):
self._consecutive_failures += 1
if self._consecutive_failures >= self.CONSECUTIVE_FAILURE_ALERT_THRESHOLD:
logger.error(
"[ALERT] Video task poller: %d consecutive batches failed. "
"Provider connectivity or configuration issue suspected.",
self._consecutive_failures,
)
else:
self._consecutive_failures = 0
async with asyncio.TaskGroup() as tg:
for tid in task_ids:
tg.create_task(poll_with_semaphore(tid))
db.commit()
batch_failures = sum(1 for r in poll_results if r is False)
# 更新连续失败计数并检查告警阈值
if batch_failures == len(task_ids):
self._consecutive_failures += 1
if self._consecutive_failures >= self.CONSECUTIVE_FAILURE_ALERT_THRESHOLD:
logger.error(
"[ALERT] Video task poller: %d consecutive batches failed. "
"Provider connectivity or configuration issue suspected.",
self._consecutive_failures,
)
else:
self._consecutive_failures = 0
finally:
await self._release_redis_lock(token)
@@ -156,11 +196,14 @@ class VideoTaskPollerService:
# 存储多视频 URLGemini sampleCount > 1 时)
if result.video_urls:
task.video_urls = result.video_urls
# 保存上游原始响应(用于审计/重算)
self._attach_poll_raw_response(task, result)
elif result.status == VideoStatus.FAILED:
task.status = VideoStatus.FAILED.value
task.error_code = result.error_code
task.error_message = result.error_message
task.completed_at = datetime.now(timezone.utc)
self._attach_poll_raw_response(task, result)
else:
task.poll_count += 1
task.progress_percent = result.progress_percent
@@ -202,6 +245,308 @@ class VideoTaskPollerService:
task.error_message = f"Task timed out after {task.poll_count} polls"
task.completed_at = datetime.now(timezone.utc)
# 终态写入 Usage复用外层 per-task session
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await self._record_terminal_usage(db, task)
except Exception as exc:
logger.exception(
"Failed to record video usage for task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
if not result.raw_response:
return
if task.request_metadata is None:
task.request_metadata = {}
# 仅在终态写一次,避免污染 request_metadata
task.request_metadata["poll_raw_response"] = result.raw_response
async def _record_terminal_usage(self, db: Session, task: VideoTask) -> None:
"""
为视频任务终态写入 Usage
- COMPLETED: 使用 FormulaEngine 计算 cost或 no_rule / incomplete -> cost=0
- FAILED: cost=0
"""
request_id = None
if isinstance(task.request_metadata, dict):
request_id = task.request_metadata.get("request_id")
request_id = request_id or task.id
# 计算异步任务总耗时ms
response_time_ms = None
if task.submitted_at and task.completed_at:
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
# 基础维度(无需 collectors 也可计费)
base_dimensions: dict[str, Any] = {
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size or "",
"retry_count": task.retry_count,
}
# collectors 可用的 metadata结构稳定便于配置 path
collector_metadata: dict[str, Any] = {
"task": {
"id": task.id,
"external_task_id": task.external_task_id,
"model": task.model,
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
"retry_count": task.retry_count,
"video_size_bytes": task.video_size_bytes,
},
"result": {
"video_url": task.video_url,
"video_urls": task.video_urls or [],
},
}
# 维度采集base + collectors 覆盖/补全
dims = DimensionCollectorService(db).collect_dimensions(
api_format=task.provider_api_format,
task_type="video",
request=task.original_request_body or {},
response=(
(task.request_metadata or {}).get("poll_raw_response")
if isinstance(task.request_metadata, dict)
else None
),
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# 取冻结的 rule_snapshot若缺失则回退 DB 查找(兼容旧任务)
rule_snapshot = None
if isinstance(task.request_metadata, dict):
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
# 构建 billing_snapshot写入 Usage.request_metadata
billing_snapshot: dict[str, Any] = {
"status": "complete",
"missing_required": [],
"strict_mode": config.billing_strict_mode,
}
cost = 0.0
if task.status == VideoStatus.FAILED.value:
billing_snapshot["billed_reason"] = "task_failed"
else:
# COMPLETED计算成本
expression = None
variables = None
dimension_mappings = None
rule_id = None
rule_name = None
rule_scope = None
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
rule_id = rule_snapshot.get("rule_id")
rule_name = rule_snapshot.get("rule_name")
rule_scope = rule_snapshot.get("scope")
expression = rule_snapshot.get("expression")
variables = rule_snapshot.get("variables")
dimension_mappings = rule_snapshot.get("dimension_mappings")
else:
lookup = BillingRuleService.find_rule(
db,
provider_id=task.provider_id,
model_name=task.model,
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = rule.id
rule_name = rule.name
rule_scope = lookup.scope
expression = rule.expression
variables = rule.variables
dimension_mappings = rule.dimension_mappings
if not expression:
billing_snapshot["status"] = "no_rule"
billing_snapshot["cost_breakdown"] = {"total": 0.0}
logger.warning(
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
request_id,
task.model,
task.provider_id,
)
else:
billing_snapshot.update(
{
"rule_id": rule_id,
"rule_name": rule_name,
"rule_scope": rule_scope,
"expression": expression,
"variables": variables or {},
}
)
try:
result = self._formula_engine.evaluate(
expression=expression,
variables=variables or {},
dimensions=dims,
dimension_mappings=dimension_mappings or {},
strict_mode=config.billing_strict_mode,
)
billing_snapshot["status"] = result.status
billing_snapshot["missing_required"] = result.missing_required
billing_snapshot["resolved_values"] = result.resolved_values
if result.status == "complete":
cost = result.cost
else:
logger.error(
"Billing incomplete due to missing required dimensions "
"(request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
result.missing_required,
)
cost = 0.0
await self._maybe_alert_missing_required(
model=task.model,
missing_required=result.missing_required,
)
if result.error:
billing_snapshot["error"] = result.error
except BillingIncompleteError as exc:
logger.error(
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
exc.missing_required,
)
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["resolved_values"] = {}
billing_snapshot["error"] = "strict_mode_missing_required"
cost = 0.0
# strict_mode=true标记任务失败并隐藏产物避免"免费放行"
task.status = VideoStatus.FAILED.value
task.error_code = "billing_incomplete"
task.error_message = f"Missing required dimensions: {exc.missing_required}"
task.video_url = None
task.video_urls = None
await self._maybe_alert_missing_required(
model=task.model,
missing_required=exc.missing_required,
)
billing_snapshot["cost_breakdown"] = {"total": cost}
# Usage 元数据(包含 snapshot + dimensions + raw_response_ref
usage_metadata: dict[str, Any] = {
"billing_snapshot": billing_snapshot,
"dimensions": dims,
"raw_response_ref": {
"video_task_id": task.id,
"field": "video_tasks.request_metadata.poll_raw_response",
},
}
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
user_obj = db.query(User).filter(User.id == task.user_id).first()
api_key_obj = (
db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
if task.api_key_id
else None
)
provider_obj = (
db.query(Provider).filter(Provider.id == task.provider_id).first()
if task.provider_id
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
await UsageService.record_usage_with_custom_cost(
db=db,
user=user_obj,
api_key=api_key_obj,
provider=provider_name,
model=task.model,
request_type="video",
total_cost_usd=cost,
request_cost_usd=cost,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=task.client_api_format,
endpoint_api_format=task.provider_api_format,
has_format_conversion=bool(task.format_converted),
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
error_message=(
None
if task.status == VideoStatus.COMPLETED.value
else (task.error_message or task.error_code or "video_task_failed")
),
metadata=usage_metadata,
request_headers=(
(task.request_metadata or {}).get("request_headers")
if isinstance(task.request_metadata, dict)
else None
),
request_body=task.original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=task.provider_id,
provider_endpoint_id=task.endpoint_id,
provider_api_key_id=task.key_id,
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
target_model=None,
)
async def _maybe_alert_missing_required(
self, *, model: str, missing_required: list[str]
) -> None:
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
if not missing_required:
return
if not self.redis:
# Redis 不可用:降级为日志
logger.error(
"Missing required billing dimensions (model=%s): %s", model, missing_required
)
return
# 按小时 bucket 聚合
now = datetime.now(timezone.utc)
hour_bucket = now.strftime("%Y%m%d%H")
for dim in missing_required:
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
try:
count = await self.redis.incr(key)
# TTL 略大于 1h避免边界抖动
if count == 1:
await self.redis.expire(key, 3700)
if count >= 10:
logger.warning(
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
model,
dim,
count,
)
except Exception as exc:
logger.warning(
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
)
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
"""判断是否为永久性错误(不应重试)"""
# 优先使用 HTTP 状态码判断
@@ -282,9 +627,7 @@ class VideoTaskPollerService:
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
operation_name = task.external_task_id
if not operation_name.startswith("operations/"):
operation_name = f"operations/{operation_name}"
operation_name = normalize_gemini_operation_id(task.external_task_id)
url = self._build_gemini_url(endpoint.base_url, operation_name)
headers = self._build_headers(APIFormat.GEMINI, upstream_key, endpoint, auth_info)