mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -9,6 +9,7 @@
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.services.billing.models import (
|
||||
|
||||
371
src/services/billing/dimension_collector_service.py
Normal file
371
src/services/billing/dimension_collector_service.py
Normal file
@@ -0,0 +1,371 @@
|
||||
"""
|
||||
DimensionCollector 运行时维度采集
|
||||
|
||||
特性(与 .plans/humming-seeking-marble.md 对齐):
|
||||
- (api_format, task_type) 作用域
|
||||
- 同一维度支持多条 collector(priority 回退)
|
||||
- 支持 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,
|
||||
),
|
||||
)
|
||||
368
src/services/billing/formula_engine.py
Normal file
368
src/services/billing/formula_engine.py
Normal file
@@ -0,0 +1,368 @@
|
||||
"""
|
||||
FormulaEngine - 配置驱动的安全计费表达式引擎
|
||||
|
||||
目标:
|
||||
- 支持 billing_rules.expression 的安全求值(AST 白名单)
|
||||
- 支持 dimension_mappings(dimension/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
|
||||
@@ -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
|
||||
|
||||
101
src/services/billing/rule_service.py
Normal file
101
src/services/billing/rule_service.py
Normal file
@@ -0,0 +1,101 @@
|
||||
"""
|
||||
BillingRule 查找逻辑
|
||||
|
||||
查找顺序(与 .plans/humming-seeking-marble.md 一致):
|
||||
1) Model(Provider 级)→ 2) GlobalModel(默认)
|
||||
|
||||
注意:
|
||||
- CLI 在计费域等同于 chat:billing_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
|
||||
@@ -9,7 +9,6 @@
|
||||
- PER_REQUEST: 按次计费
|
||||
"""
|
||||
|
||||
|
||||
from src.services.billing.models import BillingDimension, BillingUnit
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""调度器是否在运行"""
|
||||
|
||||
@@ -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 {}
|
||||
# 情况1:status='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)
|
||||
|
||||
@@ -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:
|
||||
# 存储多视频 URL(Gemini 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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user