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

@@ -42,3 +42,10 @@ ADMIN_PASSWORD=admin123456
# 示例: http://localhost:3000,https://example.com
# 默认: * (允许所有源)
# CORS_ORIGINS=*
# ==================== 计费系统(可选) ====================
# Video/Image/Audio 缺失 billing_rule 时是否拒绝请求(默认 false允许请求但 cost=0 并告警)
# BILLING_REQUIRE_RULE=false
#
# required 维度缺失时是否拒绝请求/标记任务失败(默认 falsecost=0 + 标记 incomplete
# BILLING_STRICT_MODE=false

1
.gitignore vendored
View File

@@ -220,6 +220,7 @@ frontend/public/*-firework.svg
# Debug and experimental files
debug_*.html
extracted_*.ts
test.py
# Deploy script cache
.deps-hash

View File

@@ -0,0 +1,180 @@
"""Add billing system tables and video_tasks.request_metadata
Revision ID: c8d2e4f6a1b3
Revises: b6f1a2c5d8e9
Create Date: 2026-01-31 12:00:00.000000
"""
from __future__ import annotations
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from sqlalchemy.dialects.postgresql import JSONB
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "c8d2e4f6a1b3"
down_revision: Union[str, None] = "b6f1a2c5d8e9"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
try:
indexes = inspector.get_indexes(table_name)
except Exception:
return False
return any(idx.get("name") == index_name for idx in indexes)
def upgrade() -> None:
# ==================== video_tasks.request_metadata ====================
if not column_exists("video_tasks", "request_metadata"):
op.add_column(
"video_tasks",
sa.Column("request_metadata", sa.JSON(), nullable=True),
)
# ==================== billing_rules ====================
if not table_exists("billing_rules"):
op.create_table(
"billing_rules",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column(
"global_model_id",
sa.String(36),
sa.ForeignKey("global_models.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column(
"model_id",
sa.String(36),
sa.ForeignKey("models.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column("name", sa.String(100), nullable=False),
sa.Column("task_type", sa.String(20), nullable=False, server_default="chat"),
sa.Column("expression", sa.Text(), nullable=False),
sa.Column("variables", JSONB, nullable=False, server_default=sa.text("'{}'::jsonb")),
sa.Column(
"dimension_mappings", JSONB, nullable=False, server_default=sa.text("'{}'::jsonb")
),
sa.Column("is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("true")),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.CheckConstraint(
"(global_model_id IS NOT NULL AND model_id IS NULL) OR "
"(global_model_id IS NULL AND model_id IS NOT NULL)",
name="chk_billing_rules_model_ref",
),
)
# Partial unique indexes for enabled rules
if table_exists("billing_rules"):
if not index_exists("billing_rules", "uq_billing_rules_global_model_task"):
op.create_index(
"uq_billing_rules_global_model_task",
"billing_rules",
["global_model_id", "task_type"],
unique=True,
postgresql_where=sa.text("is_enabled = TRUE AND global_model_id IS NOT NULL"),
)
if not index_exists("billing_rules", "uq_billing_rules_model_task"):
op.create_index(
"uq_billing_rules_model_task",
"billing_rules",
["model_id", "task_type"],
unique=True,
postgresql_where=sa.text("is_enabled = TRUE AND model_id IS NOT NULL"),
)
# ==================== dimension_collectors ====================
if not table_exists("dimension_collectors"):
op.create_table(
"dimension_collectors",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("api_format", sa.String(50), nullable=False),
sa.Column("task_type", sa.String(20), nullable=False),
sa.Column("dimension_name", sa.String(100), nullable=False),
sa.Column("source_type", sa.String(20), nullable=False),
sa.Column("source_path", sa.String(200), nullable=True),
sa.Column("value_type", sa.String(20), nullable=False, server_default="float"),
sa.Column("transform_expression", sa.Text(), nullable=True),
sa.Column("default_value", sa.String(100), nullable=True),
sa.Column("priority", sa.Integer(), nullable=False, server_default="0"),
sa.Column("is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("true")),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.CheckConstraint(
"(source_type = 'computed' AND source_path IS NULL AND transform_expression IS NOT NULL) OR "
"(source_type != 'computed' AND source_path IS NOT NULL)",
name="chk_dimension_collectors_source_config",
),
)
if table_exists("dimension_collectors"):
if not index_exists("dimension_collectors", "uq_dimension_collectors_enabled"):
op.create_index(
"uq_dimension_collectors_enabled",
"dimension_collectors",
["api_format", "task_type", "dimension_name", "priority"],
unique=True,
postgresql_where=sa.text("is_enabled = TRUE"),
)
def downgrade() -> None:
# Drop in reverse order
if table_exists("dimension_collectors"):
if index_exists("dimension_collectors", "uq_dimension_collectors_enabled"):
op.drop_index("uq_dimension_collectors_enabled", table_name="dimension_collectors")
op.drop_table("dimension_collectors")
if table_exists("billing_rules"):
if index_exists("billing_rules", "uq_billing_rules_model_task"):
op.drop_index("uq_billing_rules_model_task", table_name="billing_rules")
if index_exists("billing_rules", "uq_billing_rules_global_model_task"):
op.drop_index("uq_billing_rules_global_model_task", table_name="billing_rules")
op.drop_table("billing_rules")
if column_exists("video_tasks", "request_metadata"):
op.drop_column("video_tasks", "request_metadata")

View File

@@ -0,0 +1,158 @@
import apiClient from './client'
// 视频任务状态
export type VideoTaskStatus = 'pending' | 'submitted' | 'queued' | 'processing' | 'completed' | 'failed' | 'cancelled'
// 视频任务列表项
export interface VideoTaskItem {
id: string
external_task_id: string
user_id: string
username: string
model: string
prompt: string
status: VideoTaskStatus
progress_percent: number
progress_message: string | null
provider_id: string
provider_name: string
duration_seconds: number
resolution: string
aspect_ratio: string
video_url: string | null
error_code: string | null
error_message: string | null
poll_count: number
max_poll_count: number
created_at: string
completed_at: string | null
submitted_at: string | null
}
// 候选 Key 信息
export interface CandidateKeyInfo {
index: number
provider_id: string
provider_name: string
endpoint_id: string
key_id: string
key_name: string | null
auth_type: string
has_billing_rule: boolean
priority: number
selected?: boolean
}
// 请求元数据
export interface VideoTaskRequestMetadata {
candidate_keys: CandidateKeyInfo[]
selected_key_id: string
selected_endpoint_id: string
client_ip: string
user_agent: string
request_id: string
request_headers?: Record<string, string>
}
// 视频任务详情
export interface VideoTaskDetail extends VideoTaskItem {
api_key_id: string
endpoint_id: string
key_id: string
client_api_format: string
provider_api_format: string
format_converted: boolean
original_request_body: any
converted_request_body: any
size: string | null
video_urls: string[] | null
thumbnail_url: string | null
video_size_bytes: number | null
video_expires_at: string | null
stored_video_path: string | null
storage_provider: string | null
retry_count: number
max_retries: number
poll_interval_seconds: number
next_poll_at: string | null
updated_at: string | null
endpoint: {
id: string
base_url: string
api_format: string
} | null
request_metadata: VideoTaskRequestMetadata | null
}
// 视频任务列表响应
export interface VideoTaskListResponse {
items: VideoTaskItem[]
total: number
page: number
page_size: number
pages: number
}
// 视频任务统计响应
export interface VideoTaskStatsResponse {
total: number
by_status: Record<VideoTaskStatus, number>
by_model: Record<string, number>
today_count: number
active_users?: number // 仅管理员
processing_count?: number // 仅管理员
}
// 视频任务查询参数
export interface VideoTaskQueryParams {
status?: VideoTaskStatus
user_id?: string
model?: string
page?: number
page_size?: number
}
export const videoTasksApi = {
/**
* 获取视频任务列表
*/
async list(params: VideoTaskQueryParams = {}): Promise<VideoTaskListResponse> {
const searchParams = new URLSearchParams()
if (params.status) searchParams.append('status', params.status)
if (params.user_id) searchParams.append('user_id', params.user_id)
if (params.model) searchParams.append('model', params.model)
if (params.page) searchParams.append('page', params.page.toString())
if (params.page_size) searchParams.append('page_size', params.page_size.toString())
const query = searchParams.toString()
const url = query ? `/api/admin/video-tasks?${query}` : '/api/admin/video-tasks'
const response = await apiClient.get(url)
return response.data
},
/**
* 获取视频任务统计
*/
async getStats(): Promise<VideoTaskStatsResponse> {
const response = await apiClient.get('/api/admin/video-tasks/stats')
return response.data
},
/**
* 获取视频任务详情
*/
async getDetail(taskId: string): Promise<VideoTaskDetail> {
const response = await apiClient.get(`/api/admin/video-tasks/${taskId}`)
return response.data
},
/**
* 取消视频任务
*/
async cancel(taskId: string): Promise<{ id: string; status: string; message: string }> {
const response = await apiClient.post(`/api/admin/video-tasks/${taskId}/cancel`)
return response.data
},
}
export default videoTasksApi

View File

@@ -365,6 +365,7 @@ import {
X,
Mail,
Puzzle,
Video,
type LucideIcon,
} from 'lucide-vue-next'
@@ -530,6 +531,7 @@ const navigation = computed(() => {
{ name: '模型管理', href: '/admin/models', icon: Layers },
{ name: '独立密钥', href: '/admin/keys', icon: Key },
{ name: '访问令牌', href: '/admin/management-tokens', icon: KeyRound },
{ name: '视频任务', href: '/admin/video-tasks', icon: Video },
{ name: '使用记录', href: '/admin/usage', icon: BarChart3 },
]
},

View File

@@ -207,6 +207,11 @@ const routes: RouteRecordRaw[] = [
path: 'announcements',
name: 'AnnouncementManagement',
component: () => importWithRetry(() => import('@/views/user/Announcements.vue'))
},
{
path: 'video-tasks',
name: 'VideoTasks',
component: () => importWithRetry(() => import('@/views/admin/VideoTasks.vue'))
}
]
}

View File

@@ -0,0 +1,629 @@
<template>
<div class="space-y-6 pb-8">
<!-- 统计卡片 -->
<div class="grid grid-cols-2 lg:grid-cols-4 gap-4">
<Card variant="default" class="p-4">
<div class="flex items-center gap-3">
<div class="w-10 h-10 rounded-lg bg-primary/10 flex items-center justify-center">
<Video class="w-5 h-5 text-primary" />
</div>
<div>
<p class="text-2xl font-bold">{{ stats?.total ?? '-' }}</p>
<p class="text-xs text-muted-foreground">总任务数</p>
</div>
</div>
</Card>
<Card variant="default" class="p-4">
<div class="flex items-center gap-3">
<div class="w-10 h-10 rounded-lg bg-blue-500/10 flex items-center justify-center">
<Loader2 class="w-5 h-5 text-blue-500" :class="{ 'animate-spin': (stats?.processing_count ?? 0) > 0 }" />
</div>
<div>
<p class="text-2xl font-bold">{{ stats?.processing_count ?? stats?.by_status?.processing ?? '-' }}</p>
<p class="text-xs text-muted-foreground">处理中</p>
</div>
</div>
</Card>
<Card variant="default" class="p-4">
<div class="flex items-center gap-3">
<div class="w-10 h-10 rounded-lg bg-green-500/10 flex items-center justify-center">
<CheckCircle class="w-5 h-5 text-green-500" />
</div>
<div>
<p class="text-2xl font-bold">{{ stats?.by_status?.completed ?? '-' }}</p>
<p class="text-xs text-muted-foreground">已完成</p>
</div>
</div>
</Card>
<Card variant="default" class="p-4">
<div class="flex items-center gap-3">
<div class="w-10 h-10 rounded-lg bg-amber-500/10 flex items-center justify-center">
<Calendar class="w-5 h-5 text-amber-500" />
</div>
<div>
<p class="text-2xl font-bold">{{ stats?.today_count ?? '-' }}</p>
<p class="text-xs text-muted-foreground">今日任务</p>
</div>
</div>
</Card>
</div>
<!-- 任务表格 -->
<Card variant="default" class="overflow-hidden">
<!-- 标题和筛选器 -->
<div class="px-4 sm:px-6 py-3.5 border-b border-border/60">
<div class="flex flex-col gap-3 sm:flex-row sm:items-center sm:justify-between">
<h3 class="text-base font-semibold">视频任务</h3>
<div class="flex items-center gap-2">
<!-- 状态筛选 -->
<Select v-model="filterStatus">
<SelectTrigger class="w-28 h-8 text-xs border-border/60">
<SelectValue placeholder="状态" />
</SelectTrigger>
<SelectContent>
<SelectItem value="all">全部状态</SelectItem>
<SelectItem value="submitted">已提交</SelectItem>
<SelectItem value="processing">处理中</SelectItem>
<SelectItem value="completed">已完成</SelectItem>
<SelectItem value="failed">失败</SelectItem>
<SelectItem value="cancelled">已取消</SelectItem>
</SelectContent>
</Select>
<!-- 模型筛选 -->
<Input
v-model="filterModel"
type="text"
placeholder="模型..."
class="w-32 h-8 text-xs"
/>
<!-- 刷新按钮 -->
<Button
variant="ghost"
size="icon"
class="h-8 w-8"
:disabled="loading"
@click="fetchTasks"
>
<RefreshCw class="w-3.5 h-3.5" :class="{ 'animate-spin': loading }" />
</Button>
</div>
</div>
</div>
<!-- 加载状态 -->
<div v-if="loading && !tasks.length" class="p-8 text-center">
<Loader2 class="w-8 h-8 animate-spin mx-auto text-muted-foreground" />
<p class="mt-2 text-sm text-muted-foreground">加载中...</p>
</div>
<!-- 空状态 -->
<div v-else-if="!tasks.length" class="p-8 text-center">
<Video class="w-12 h-12 mx-auto text-muted-foreground/50" />
<p class="mt-2 text-sm text-muted-foreground">暂无视频任务</p>
</div>
<!-- 任务列表 -->
<div v-else class="divide-y divide-border/60">
<div
v-for="task in tasks"
:key="task.id"
class="px-4 sm:px-6 py-4 hover:bg-muted/30 cursor-pointer transition-colors"
@click="openTaskDetail(task)"
>
<div class="flex items-start justify-between gap-4">
<div class="flex-1 min-w-0">
<!-- 模型和状态 -->
<div class="flex items-center gap-2 mb-1">
<span class="font-medium text-sm">{{ task.model }}</span>
<Badge :variant="getStatusVariant(task.status)">
{{ getStatusLabel(task.status) }}
</Badge>
<span v-if="task.progress_percent > 0 && task.status === 'processing'" class="text-xs text-muted-foreground">
{{ task.progress_percent }}%
</span>
</div>
<!-- Prompt 摘要 -->
<p class="text-sm text-muted-foreground truncate">{{ task.prompt }}</p>
<!-- 元信息 -->
<div class="flex items-center gap-4 mt-2 text-xs text-muted-foreground">
<span class="flex items-center gap-1">
<User class="w-3 h-3" />
{{ task.username }}
</span>
<span class="flex items-center gap-1">
<Server class="w-3 h-3" />
{{ task.provider_name }}
</span>
<span class="flex items-center gap-1">
<Clock class="w-3 h-3" />
{{ formatDate(task.created_at) }}
</span>
<span v-if="task.duration_seconds" class="flex items-center gap-1">
<Timer class="w-3 h-3" />
{{ task.duration_seconds }}s
</span>
</div>
</div>
<!-- 操作 -->
<div class="flex items-center gap-2">
<Button
v-if="canCancel(task.status)"
variant="ghost"
size="sm"
class="text-red-500 hover:text-red-600 hover:bg-red-50"
@click.stop="cancelTask(task)"
>
<XCircle class="w-4 h-4" />
</Button>
<ChevronRight class="w-4 h-4 text-muted-foreground" />
</div>
</div>
</div>
</div>
<!-- 分页 -->
<div v-if="totalPages > 1" class="px-4 sm:px-6 py-3 border-t border-border/60 flex items-center justify-between">
<p class="text-xs text-muted-foreground">
{{ total }} {{ currentPage }}/{{ totalPages }}
</p>
<div class="flex items-center gap-2">
<Button
variant="outline"
size="sm"
:disabled="currentPage <= 1"
@click="goToPage(currentPage - 1)"
>
上一页
</Button>
<Button
variant="outline"
size="sm"
:disabled="currentPage >= totalPages"
@click="goToPage(currentPage + 1)"
>
下一页
</Button>
</div>
</div>
</Card>
<!-- 任务详情抽屉 -->
<Teleport to="body">
<Transition name="drawer">
<div
v-if="showDetail && selectedTask"
class="fixed inset-0 z-50 flex justify-end"
@click.self="showDetail = false"
>
<!-- 背景遮罩 -->
<div
class="absolute inset-0 bg-black/30 backdrop-blur-sm"
@click="showDetail = false"
/>
<!-- 抽屉内容 -->
<Card class="relative h-full w-full sm:w-[600px] sm:max-w-[90vw] rounded-none shadow-2xl overflow-y-auto">
<!-- 标题栏 -->
<div class="sticky top-0 z-10 bg-background border-b p-4 sm:p-6">
<div class="flex items-center justify-between">
<h3 class="text-lg font-semibold">任务详情</h3>
<Button variant="ghost" size="icon" @click="showDetail = false">
<X class="w-4 h-4" />
</Button>
</div>
</div>
<!-- 内容 -->
<div class="p-4 sm:p-6 space-y-6">
<!-- 状态和进度 -->
<div class="space-y-2">
<div class="flex items-center justify-between">
<span class="text-sm font-medium">状态</span>
<Badge :variant="getStatusVariant(selectedTask.status)">
{{ getStatusLabel(selectedTask.status) }}
</Badge>
</div>
<div v-if="selectedTask.progress_percent > 0" class="space-y-1">
<div class="flex justify-between text-xs text-muted-foreground">
<span>进度</span>
<span>{{ selectedTask.progress_percent }}%</span>
</div>
<div class="h-2 bg-muted rounded-full overflow-hidden">
<div
class="h-full bg-primary transition-all"
:style="{ width: `${selectedTask.progress_percent}%` }"
/>
</div>
</div>
<p v-if="selectedTask.progress_message" class="text-xs text-muted-foreground">
{{ selectedTask.progress_message }}
</p>
</div>
<!-- 错误信息 -->
<div v-if="selectedTask.error_message" class="p-3 bg-red-50 dark:bg-red-900/20 rounded-lg">
<p class="text-sm text-red-600 dark:text-red-400">
<span v-if="selectedTask.error_code" class="font-medium">[{{ selectedTask.error_code }}]</span>
{{ selectedTask.error_message }}
</p>
</div>
<!-- 基本信息 -->
<div class="space-y-3">
<h4 class="text-sm font-medium">基本信息</h4>
<div class="grid grid-cols-2 gap-3 text-sm">
<div>
<span class="text-muted-foreground">模型</span>
<p class="font-medium">{{ selectedTask.model }}</p>
</div>
<div>
<span class="text-muted-foreground">时长</span>
<p class="font-medium">{{ selectedTask.duration_seconds }}s</p>
</div>
<div>
<span class="text-muted-foreground">分辨率</span>
<p class="font-medium">{{ selectedTask.resolution }}</p>
</div>
<div>
<span class="text-muted-foreground">宽高比</span>
<p class="font-medium">{{ selectedTask.aspect_ratio }}</p>
</div>
</div>
</div>
<!-- Prompt -->
<div class="space-y-2">
<h4 class="text-sm font-medium">Prompt</h4>
<div class="p-3 bg-muted/50 rounded-lg text-sm whitespace-pre-wrap">
{{ selectedTask.prompt }}
</div>
</div>
<!-- Provider 信息 -->
<div class="space-y-3">
<h4 class="text-sm font-medium">Provider 信息</h4>
<div class="grid grid-cols-2 gap-3 text-sm">
<div>
<span class="text-muted-foreground">用户</span>
<p class="font-medium">{{ selectedTask.username }}</p>
</div>
<div>
<span class="text-muted-foreground">Provider</span>
<p class="font-medium">{{ selectedTask.provider_name }}</p>
</div>
<div>
<span class="text-muted-foreground">客户端格式</span>
<p class="font-medium">{{ selectedTask.client_api_format }}</p>
</div>
<div>
<span class="text-muted-foreground">Provider 格式</span>
<p class="font-medium">{{ selectedTask.provider_api_format }}</p>
</div>
</div>
</div>
<!-- 候选 Key 追踪 -->
<div v-if="selectedTask.request_metadata?.candidate_keys?.length" class="space-y-3">
<h4 class="text-sm font-medium flex items-center gap-2">
<Key class="w-4 h-4" />
候选 Key 列表
<span class="text-xs text-muted-foreground">({{ selectedTask.request_metadata.candidate_keys.length }} )</span>
</h4>
<div class="space-y-2">
<div
v-for="candidateKey in selectedTask.request_metadata.candidate_keys"
:key="candidateKey.key_id"
class="p-2 rounded-lg text-xs border"
:class="candidateKey.selected ? 'bg-primary/10 border-primary/30' : 'bg-muted/30 border-border/60'"
>
<div class="flex items-center justify-between">
<div class="flex items-center gap-2">
<span class="font-medium">{{ candidateKey.provider_name }}</span>
<span v-if="candidateKey.key_name" class="text-muted-foreground">/ {{ candidateKey.key_name }}</span>
<Badge v-if="candidateKey.selected" variant="default" class="text-[10px] px-1.5 py-0">已选中</Badge>
<Badge v-if="candidateKey.has_billing_rule === false" variant="outline" class="text-[10px] px-1.5 py-0 text-amber-500 border-amber-500/50">无计费规则</Badge>
</div>
<span class="text-muted-foreground">优先级: {{ candidateKey.priority }}</span>
</div>
<div class="flex items-center gap-3 mt-1 text-muted-foreground">
<span>Auth: {{ candidateKey.auth_type }}</span>
<span class="font-mono truncate max-w-[120px]" :title="candidateKey.key_id">Key: {{ candidateKey.key_id.slice(0, 8) }}...</span>
</div>
</div>
</div>
</div>
<!-- 请求追踪 -->
<div v-if="selectedTask.request_metadata" class="space-y-3">
<h4 class="text-sm font-medium">请求追踪</h4>
<div class="grid grid-cols-2 gap-3 text-sm">
<div>
<span class="text-muted-foreground">Request ID</span>
<p class="font-mono text-xs truncate" :title="selectedTask.request_metadata.request_id">{{ selectedTask.request_metadata.request_id }}</p>
</div>
<div>
<span class="text-muted-foreground">Client IP</span>
<p class="font-medium">{{ selectedTask.request_metadata.client_ip }}</p>
</div>
<div class="col-span-2">
<span class="text-muted-foreground">User Agent</span>
<p class="text-xs truncate" :title="selectedTask.request_metadata.user_agent">{{ selectedTask.request_metadata.user_agent }}</p>
</div>
</div>
</div>
<!-- 轮询信息 -->
<div class="space-y-3">
<h4 class="text-sm font-medium">轮询信息</h4>
<div class="grid grid-cols-2 gap-3 text-sm">
<div>
<span class="text-muted-foreground">轮询次数</span>
<p class="font-medium">{{ selectedTask.poll_count }} / {{ selectedTask.max_poll_count }}</p>
</div>
<div>
<span class="text-muted-foreground">轮询间隔</span>
<p class="font-medium">{{ selectedTask.poll_interval_seconds }}s</p>
</div>
<div v-if="selectedTask.next_poll_at">
<span class="text-muted-foreground">下次轮询</span>
<p class="font-medium">{{ formatDate(selectedTask.next_poll_at) }}</p>
</div>
</div>
</div>
<!-- 时间信息 -->
<div class="space-y-3">
<h4 class="text-sm font-medium">时间信息</h4>
<div class="grid grid-cols-2 gap-3 text-sm">
<div>
<span class="text-muted-foreground">创建时间</span>
<p class="font-medium">{{ formatDate(selectedTask.created_at) }}</p>
</div>
<div v-if="selectedTask.submitted_at">
<span class="text-muted-foreground">提交时间</span>
<p class="font-medium">{{ formatDate(selectedTask.submitted_at) }}</p>
</div>
<div v-if="selectedTask.completed_at">
<span class="text-muted-foreground">完成时间</span>
<p class="font-medium">{{ formatDate(selectedTask.completed_at) }}</p>
</div>
</div>
</div>
<!-- 视频结果 -->
<div v-if="selectedTask.video_url" class="space-y-3">
<h4 class="text-sm font-medium">视频结果</h4>
<div class="space-y-2">
<video
:src="selectedTask.video_url"
controls
class="w-full rounded-lg"
/>
<p v-if="selectedTask.video_expires_at" class="text-xs text-muted-foreground">
过期时间: {{ formatDate(selectedTask.video_expires_at) }}
</p>
</div>
</div>
<!-- 操作按钮 -->
<div v-if="canCancel(selectedTask.status)" class="pt-4 border-t">
<Button
variant="destructive"
class="w-full"
@click="cancelTask(selectedTask)"
>
<XCircle class="w-4 h-4 mr-2" />
取消任务
</Button>
</div>
</div>
</Card>
</div>
</Transition>
</Teleport>
</div>
</template>
<style scoped>
.drawer-enter-active,
.drawer-leave-active {
transition: all 0.3s ease;
}
.drawer-enter-active > div:first-child,
.drawer-leave-active > div:first-child {
transition: opacity 0.3s ease;
}
.drawer-enter-active > div:last-child,
.drawer-leave-active > div:last-child {
transition: transform 0.3s ease;
}
.drawer-enter-from,
.drawer-leave-to {
opacity: 0;
}
.drawer-enter-from > div:last-child,
.drawer-leave-to > div:last-child {
transform: translateX(100%);
}
</style>
<script setup lang="ts">
import { ref, computed, onMounted, watch } from 'vue'
import { videoTasksApi, type VideoTaskItem, type VideoTaskDetail, type VideoTaskStatsResponse, type VideoTaskStatus } from '@/api/video-tasks'
import { useToast } from '@/composables/useToast'
import Card from '@/components/ui/card.vue'
import Button from '@/components/ui/button.vue'
import Input from '@/components/ui/input.vue'
import Badge from '@/components/ui/badge.vue'
import Select from '@/components/ui/select.vue'
import SelectTrigger from '@/components/ui/select-trigger.vue'
import SelectValue from '@/components/ui/select-value.vue'
import SelectContent from '@/components/ui/select-content.vue'
import SelectItem from '@/components/ui/select-item.vue'
import {
Video,
Loader2,
CheckCircle,
Calendar,
RefreshCw,
User,
Server,
Clock,
Timer,
XCircle,
ChevronRight,
X,
Key,
} from 'lucide-vue-next'
const { toast } = useToast()
// 状态
const loading = ref(false)
const tasks = ref<VideoTaskItem[]>([])
const stats = ref<VideoTaskStatsResponse | null>(null)
const total = ref(0)
const currentPage = ref(1)
const pageSize = ref(20)
const filterStatus = ref('all')
const filterModel = ref('')
const showDetail = ref(false)
const selectedTask = ref<VideoTaskDetail | null>(null)
const totalPages = computed(() => Math.ceil(total.value / pageSize.value))
// 获取任务列表
async function fetchTasks() {
loading.value = true
try {
const response = await videoTasksApi.list({
status: filterStatus.value !== 'all' ? filterStatus.value as VideoTaskStatus : undefined,
model: filterModel.value || undefined,
page: currentPage.value,
page_size: pageSize.value,
})
tasks.value = response.items
total.value = response.total
} catch (error: any) {
toast({
title: '获取任务列表失败',
description: error.message,
variant: 'destructive',
})
} finally {
loading.value = false
}
}
// 获取统计数据
async function fetchStats() {
try {
stats.value = await videoTasksApi.getStats()
} catch (error) {
console.error('Failed to fetch stats:', error)
}
}
// 打开任务详情
async function openTaskDetail(task: VideoTaskItem) {
try {
selectedTask.value = await videoTasksApi.getDetail(task.id)
showDetail.value = true
} catch (error: any) {
toast({
title: '获取任务详情失败',
description: error.message,
variant: 'destructive',
})
}
}
// 取消任务
async function cancelTask(task: VideoTaskItem | VideoTaskDetail) {
if (!confirm('确定要取消这个任务吗?')) return
try {
await videoTasksApi.cancel(task.id)
toast({
title: '任务已取消',
})
fetchTasks()
fetchStats()
if (showDetail.value) {
showDetail.value = false
selectedTask.value = null
}
} catch (error: any) {
toast({
title: '取消任务失败',
description: error.message,
variant: 'destructive',
})
}
}
// 状态相关
function getStatusVariant(status: string): 'default' | 'secondary' | 'destructive' | 'outline' {
switch (status) {
case 'completed':
return 'default'
case 'failed':
return 'destructive'
case 'cancelled':
return 'outline'
default:
return 'secondary'
}
}
function getStatusLabel(status: string): string {
const labels: Record<string, string> = {
pending: '待处理',
submitted: '已提交',
queued: '排队中',
processing: '处理中',
completed: '已完成',
failed: '失败',
cancelled: '已取消',
}
return labels[status] || status
}
function canCancel(status: string): boolean {
return ['pending', 'submitted', 'queued', 'processing'].includes(status)
}
// 格式化日期
function formatDate(dateStr: string | null): string {
if (!dateStr) return '-'
const date = new Date(dateStr)
return date.toLocaleString('zh-CN', {
month: '2-digit',
day: '2-digit',
hour: '2-digit',
minute: '2-digit',
})
}
// 分页
function goToPage(page: number) {
currentPage.value = page
fetchTasks()
}
// 监听筛选条件变化
let filterTimeout: number
watch(filterStatus, () => {
currentPage.value = 1
fetchTasks()
})
watch(filterModel, () => {
clearTimeout(filterTimeout)
filterTimeout = window.setTimeout(() => {
currentPage.value = 1
fetchTasks()
}, 400)
})
onMounted(() => {
fetchTasks()
fetchStats()
})
</script>

View File

@@ -4,10 +4,11 @@ from fastapi import APIRouter
from .adaptive import router as adaptive_router
from .api_keys import router as api_keys_router
from .billing import router as billing_router
from .endpoints import router as endpoints_router
from .management_tokens import router as management_tokens_router
from .modules import router as modules_router
from .models import router as models_router
from .modules import router as modules_router
from .monitoring import router as monitoring_router
from .provider_ops import router as provider_ops_router
from .provider_query import router as provider_query_router
@@ -17,12 +18,14 @@ from .security import router as security_router
from .system import router as system_router
from .usage import router as usage_router
from .users import router as users_router
from .video_tasks import router as video_tasks_router
router = APIRouter()
router.include_router(system_router)
router.include_router(users_router)
router.include_router(providers_router)
router.include_router(api_keys_router)
router.include_router(billing_router)
router.include_router(usage_router)
router.include_router(monitoring_router)
router.include_router(endpoints_router)
@@ -34,6 +37,7 @@ router.include_router(provider_query_router)
router.include_router(management_tokens_router)
router.include_router(modules_router)
router.include_router(provider_ops_router)
router.include_router(video_tasks_router)
# 注意ldap_router 已迁移到模块系统,由 ModuleRegistry 动态注册
# 当 LDAP_AVAILABLE=true 时才会注册路由

View File

@@ -0,0 +1,5 @@
"""Billing 配置管理 API 模块billing_rules / dimension_collectors"""
from .routes import router
__all__ = ["router"]

View File

@@ -0,0 +1,527 @@
"""Billing 配置管理 API 路由。
包含:
- billing_rules: 计费规则(公式/变量/维度映射)
- dimension_collectors: 维度采集器request/response/metadata/computed
"""
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Literal
from fastapi import APIRouter, Depends, Query, Request
from pydantic import BaseModel, Field
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.database import get_db
from src.models.database import BillingRule, DimensionCollector
from src.services.billing.formula_engine import SafeExpressionEvaluator, UnsafeExpressionError
router = APIRouter(prefix="/api/admin/billing", tags=["Admin - Billing"])
pipeline = ApiRequestPipeline()
_expr_validator = SafeExpressionEvaluator()
AllowedTaskType = Literal["chat", "video", "image", "audio"]
AllowedCollectorSourceType = Literal["request", "response", "metadata", "computed"]
AllowedValueType = Literal["float", "int", "string"]
class BillingRuleUpsertRequest(BaseModel):
name: str = Field(..., min_length=1, max_length=100)
task_type: AllowedTaskType = "chat"
global_model_id: str | None = None
model_id: str | None = None
expression: str = Field(..., min_length=1)
variables: dict[str, Any] = Field(default_factory=dict)
dimension_mappings: dict[str, Any] = Field(default_factory=dict)
is_enabled: bool = True
class BillingRuleResponse(BaseModel):
id: str
name: str
task_type: str
global_model_id: str | None
model_id: str | None
expression: str
variables: dict[str, Any]
dimension_mappings: dict[str, Any]
is_enabled: bool
created_at: datetime
updated_at: datetime
@classmethod
def from_orm_obj(cls, rule: BillingRule) -> "BillingRuleResponse":
return cls(
id=rule.id,
name=rule.name,
task_type=rule.task_type,
global_model_id=rule.global_model_id,
model_id=rule.model_id,
expression=rule.expression,
variables=rule.variables or {},
dimension_mappings=rule.dimension_mappings or {},
is_enabled=bool(rule.is_enabled),
created_at=rule.created_at,
updated_at=rule.updated_at,
)
class DimensionCollectorUpsertRequest(BaseModel):
api_format: str = Field(..., min_length=1, max_length=50)
task_type: str = Field(..., min_length=1, max_length=20)
dimension_name: str = Field(..., min_length=1, max_length=100)
source_type: AllowedCollectorSourceType
source_path: str | None = None
value_type: AllowedValueType = "float"
transform_expression: str | None = None
default_value: str | None = None
priority: int = 0
is_enabled: bool = True
class DimensionCollectorResponse(BaseModel):
id: str
api_format: str
task_type: str
dimension_name: str
source_type: str
source_path: str | None
value_type: str
transform_expression: str | None
default_value: str | None
priority: int
is_enabled: bool
created_at: datetime
updated_at: datetime
@classmethod
def from_orm_obj(cls, c: DimensionCollector) -> "DimensionCollectorResponse":
return cls(
id=c.id,
api_format=c.api_format,
task_type=c.task_type,
dimension_name=c.dimension_name,
source_type=c.source_type,
source_path=c.source_path,
value_type=c.value_type,
transform_expression=c.transform_expression,
default_value=c.default_value,
priority=int(c.priority or 0),
is_enabled=bool(c.is_enabled),
created_at=c.created_at,
updated_at=c.updated_at,
)
@router.get("/rules")
async def list_billing_rules(
request: Request,
task_type: str | None = Query(None),
is_enabled: bool | None = Query(None),
page: int = Query(1, ge=1),
page_size: int = Query(50, ge=1, le=200),
db: Session = Depends(get_db),
) -> Any:
adapter = BillingRuleListAdapter(
task_type=task_type,
is_enabled=is_enabled,
page=page,
page_size=page_size,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/rules/{rule_id}")
async def get_billing_rule(rule_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
adapter = BillingRuleDetailAdapter(rule_id=rule_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/rules")
async def create_billing_rule(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = BillingRuleCreateAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.put("/rules/{rule_id}")
async def update_billing_rule(rule_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
adapter = BillingRuleUpdateAdapter(rule_id=rule_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/collectors")
async def list_dimension_collectors(
request: Request,
api_format: str | None = Query(None),
task_type: str | None = Query(None),
dimension_name: str | None = Query(None),
is_enabled: bool | None = Query(None),
page: int = Query(1, ge=1),
page_size: int = Query(50, ge=1, le=200),
db: Session = Depends(get_db),
) -> Any:
adapter = DimensionCollectorListAdapter(
api_format=api_format,
task_type=task_type,
dimension_name=dimension_name,
is_enabled=is_enabled,
page=page,
page_size=page_size,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/collectors/{collector_id}")
async def get_dimension_collector(
collector_id: str, request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = DimensionCollectorDetailAdapter(collector_id=collector_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/collectors")
async def create_dimension_collector(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = DimensionCollectorCreateAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.put("/collectors/{collector_id}")
async def update_dimension_collector(
collector_id: str, request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = DimensionCollectorUpdateAdapter(collector_id=collector_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# ---------------------------------------------------------------------------
# Adapters
# ---------------------------------------------------------------------------
@dataclass
class BillingRuleListAdapter(AdminApiAdapter):
page: int
page_size: int
task_type: str | None = None
is_enabled: bool | None = None
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
q = context.db.query(BillingRule)
if self.task_type:
q = q.filter(BillingRule.task_type == self.task_type.lower())
if self.is_enabled is not None:
q = q.filter(BillingRule.is_enabled == self.is_enabled)
total = q.count()
items = (
q.order_by(BillingRule.updated_at.desc())
.offset((self.page - 1) * self.page_size)
.limit(self.page_size)
.all()
)
return {
"items": [BillingRuleResponse.from_orm_obj(r).model_dump() for r in items],
"total": total,
"page": self.page,
"page_size": self.page_size,
"pages": (total + self.page_size - 1) // self.page_size,
}
@dataclass
class BillingRuleDetailAdapter(AdminApiAdapter):
rule_id: str
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
rule = context.db.query(BillingRule).filter(BillingRule.id == self.rule_id).first()
if not rule:
raise NotFoundException("Billing rule not found")
return BillingRuleResponse.from_orm_obj(rule).model_dump()
class BillingRuleCreateAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
payload = context.ensure_json_body()
try:
req = BillingRuleUpsertRequest.model_validate(payload)
except Exception as exc:
raise InvalidRequestException(f"Invalid request body: {exc}")
_validate_billing_rule_request(req)
rule = BillingRule(
name=req.name,
task_type=req.task_type,
global_model_id=req.global_model_id,
model_id=req.model_id,
expression=req.expression,
variables=req.variables,
dimension_mappings=req.dimension_mappings,
is_enabled=req.is_enabled,
)
context.db.add(rule)
try:
context.db.commit()
except IntegrityError as exc:
context.db.rollback()
raise InvalidRequestException(f"Integrity error: {exc}")
context.db.refresh(rule)
return BillingRuleResponse.from_orm_obj(rule).model_dump()
@dataclass
class BillingRuleUpdateAdapter(AdminApiAdapter):
rule_id: str
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
rule = context.db.query(BillingRule).filter(BillingRule.id == self.rule_id).first()
if not rule:
raise NotFoundException("Billing rule not found")
payload = context.ensure_json_body()
try:
req = BillingRuleUpsertRequest.model_validate(payload)
except Exception as exc:
raise InvalidRequestException(f"Invalid request body: {exc}")
_validate_billing_rule_request(req)
rule.name = req.name
rule.task_type = req.task_type
rule.global_model_id = req.global_model_id
rule.model_id = req.model_id
rule.expression = req.expression
rule.variables = req.variables
rule.dimension_mappings = req.dimension_mappings
rule.is_enabled = req.is_enabled
try:
context.db.commit()
except IntegrityError as exc:
context.db.rollback()
raise InvalidRequestException(f"Integrity error: {exc}")
context.db.refresh(rule)
return BillingRuleResponse.from_orm_obj(rule).model_dump()
@dataclass
class DimensionCollectorListAdapter(AdminApiAdapter):
page: int
page_size: int
api_format: str | None = None
task_type: str | None = None
dimension_name: str | None = None
is_enabled: bool | None = None
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
q = context.db.query(DimensionCollector)
if self.api_format:
q = q.filter(DimensionCollector.api_format == self.api_format.upper())
if self.task_type:
q = q.filter(DimensionCollector.task_type == self.task_type.lower())
if self.dimension_name:
q = q.filter(DimensionCollector.dimension_name == self.dimension_name)
if self.is_enabled is not None:
q = q.filter(DimensionCollector.is_enabled == self.is_enabled)
total = q.count()
items = (
q.order_by(DimensionCollector.updated_at.desc())
.offset((self.page - 1) * self.page_size)
.limit(self.page_size)
.all()
)
return {
"items": [DimensionCollectorResponse.from_orm_obj(c).model_dump() for c in items],
"total": total,
"page": self.page,
"page_size": self.page_size,
"pages": (total + self.page_size - 1) // self.page_size,
}
@dataclass
class DimensionCollectorDetailAdapter(AdminApiAdapter):
collector_id: str
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
c = (
context.db.query(DimensionCollector)
.filter(DimensionCollector.id == self.collector_id)
.first()
)
if not c:
raise NotFoundException("Dimension collector not found")
return DimensionCollectorResponse.from_orm_obj(c).model_dump()
class DimensionCollectorCreateAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
payload = context.ensure_json_body()
try:
req = DimensionCollectorUpsertRequest.model_validate(payload)
except Exception as exc:
raise InvalidRequestException(f"Invalid request body: {exc}")
_validate_dimension_collector_request(context.db, req, existing_id=None)
c = DimensionCollector(
api_format=req.api_format.upper(),
task_type=req.task_type.lower(),
dimension_name=req.dimension_name,
source_type=req.source_type,
source_path=req.source_path,
value_type=req.value_type,
transform_expression=req.transform_expression,
default_value=req.default_value,
priority=req.priority,
is_enabled=req.is_enabled,
)
context.db.add(c)
try:
context.db.commit()
except IntegrityError as exc:
context.db.rollback()
raise InvalidRequestException(f"Integrity error: {exc}")
context.db.refresh(c)
return DimensionCollectorResponse.from_orm_obj(c).model_dump()
@dataclass
class DimensionCollectorUpdateAdapter(AdminApiAdapter):
collector_id: str
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
c = (
context.db.query(DimensionCollector)
.filter(DimensionCollector.id == self.collector_id)
.first()
)
if not c:
raise NotFoundException("Dimension collector not found")
payload = context.ensure_json_body()
try:
req = DimensionCollectorUpsertRequest.model_validate(payload)
except Exception as exc:
raise InvalidRequestException(f"Invalid request body: {exc}")
_validate_dimension_collector_request(context.db, req, existing_id=self.collector_id)
c.api_format = req.api_format.upper()
c.task_type = req.task_type.lower()
c.dimension_name = req.dimension_name
c.source_type = req.source_type
c.source_path = req.source_path
c.value_type = req.value_type
c.transform_expression = req.transform_expression
c.default_value = req.default_value
c.priority = req.priority
c.is_enabled = req.is_enabled
try:
context.db.commit()
except IntegrityError as exc:
context.db.rollback()
raise InvalidRequestException(f"Integrity error: {exc}")
context.db.refresh(c)
return DimensionCollectorResponse.from_orm_obj(c).model_dump()
# ---------------------------------------------------------------------------
# Validation helpers
# ---------------------------------------------------------------------------
def _validate_billing_rule_request(req: BillingRuleUpsertRequest) -> None:
# model/global_model 二选一
if bool(req.global_model_id) == bool(req.model_id):
raise InvalidRequestException("Exactly one of global_model_id or model_id must be provided")
# task_type 校验Pydantic Literal 已限制为 "chat", "video", "image", "audio"
# 注CLI 在计费域等同于 chatbilling_rules 不存储 "cli"
# expression 安全校验
try:
_expr_validator.validate(req.expression)
except UnsafeExpressionError as exc:
raise InvalidRequestException(f"Invalid expression: {exc}")
# variables 必须为数值JSON 可包含 int/float
if not isinstance(req.variables, dict):
raise InvalidRequestException("variables must be a JSON object")
for k, v in req.variables.items():
if not isinstance(k, str) or not k:
raise InvalidRequestException("variables keys must be non-empty strings")
if isinstance(v, bool) or not isinstance(v, (int, float)):
raise InvalidRequestException(f"variables['{k}'] must be a number")
# dimension_mappings 结构做轻量校验(详细 schema 由业务侧保障)
if not isinstance(req.dimension_mappings, dict):
raise InvalidRequestException("dimension_mappings must be a JSON object")
for var_name, mapping in req.dimension_mappings.items():
if not isinstance(var_name, str) or not var_name:
raise InvalidRequestException("dimension_mappings keys must be non-empty strings")
if not isinstance(mapping, dict):
raise InvalidRequestException(f"dimension_mappings['{var_name}'] must be an object")
if "source" not in mapping:
raise InvalidRequestException(f"dimension_mappings['{var_name}'].source is required")
def _validate_dimension_collector_request(
db: Session,
req: DimensionCollectorUpsertRequest,
*,
existing_id: str | None,
) -> None:
src = req.source_type
if src == "computed":
if req.source_path is not None:
raise InvalidRequestException("computed collector must have source_path=null")
if not req.transform_expression:
raise InvalidRequestException("computed collector must have transform_expression")
else:
if not req.source_path:
raise InvalidRequestException("non-computed collector must have source_path")
# transform_expression 安全校验(如配置)
if req.transform_expression:
try:
_expr_validator.validate(req.transform_expression)
except UnsafeExpressionError as exc:
raise InvalidRequestException(f"Invalid transform_expression: {exc}")
# default_value 仅允许同一维度一条enabled=true
if req.default_value is not None and req.is_enabled:
q = db.query(DimensionCollector).filter(
DimensionCollector.api_format == req.api_format.upper(),
DimensionCollector.task_type == req.task_type.lower(),
DimensionCollector.dimension_name == req.dimension_name,
DimensionCollector.is_enabled.is_(True),
DimensionCollector.default_value.isnot(None),
)
if existing_id:
q = q.filter(DimensionCollector.id != existing_id)
exists = db.query(q.exists()).scalar()
if exists:
raise InvalidRequestException(
"default_value already exists for this (api_format, task_type, dimension_name)"
)

View File

@@ -0,0 +1,5 @@
"""视频任务管理 API 模块。"""
from .routes import router
__all__ = ["router"]

View File

@@ -0,0 +1,424 @@
"""视频任务管理 API 路由。
管理员可以查看所有视频任务,用户只能查看自己的任务。
"""
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.api.dashboard.routes import DashboardAdapter
from src.core.enums import UserRole
from src.database import get_db
from src.models.database import Provider, ProviderEndpoint, User, VideoTask
router = APIRouter(prefix="/api/admin/video-tasks", tags=["Admin - Video Tasks"])
pipeline = ApiRequestPipeline()
@router.get("")
async def list_video_tasks(
request: Request,
status: str | None = Query(None, description="Filter by status"),
user_id: str | None = Query(None, description="Filter by user ID (admin only)"),
model: str | None = Query(None, description="Filter by model"),
page: int = Query(1, ge=1, description="Page number"),
page_size: int = Query(20, ge=1, le=100, description="Items per page"),
db: Session = Depends(get_db),
) -> Any:
"""
获取视频任务列表
管理员可以查看所有用户的任务,普通用户只能查看自己的任务。
**查询参数**:
- `status`: 按状态筛选pending/submitted/processing/completed/failed/cancelled
- `user_id`: 按用户 ID 筛选(仅管理员)
- `model`: 按模型筛选
- `page`: 页码,默认 1
- `page_size`: 每页数量,默认 20最大 100
**返回字段**:
- `items`: 任务列表
- `total`: 总数
- `page`: 当前页码
- `page_size`: 每页数量
- `pages`: 总页数
"""
adapter = VideoTaskListAdapter(
status=status,
user_id=user_id,
model=model,
page=page,
page_size=page_size,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/stats")
async def get_video_task_stats(
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
获取视频任务统计
**返回字段**:
- `total`: 总任务数
- `by_status`: 按状态分组的数量
- `by_model`: 按模型分组的数量(前 10
- `today_count`: 今日任务数
"""
adapter = VideoTaskStatsAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/{task_id}")
async def get_video_task_detail(
task_id: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
获取视频任务详情
**路径参数**:
- `task_id`: 任务 ID
**返回字段**:
- 任务的完整信息,包括请求体、响应、状态等
"""
adapter = VideoTaskDetailAdapter(task_id=task_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/{task_id}/cancel")
async def cancel_video_task(
task_id: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
取消视频任务
**路径参数**:
- `task_id`: 任务 ID
**返回**:
- 更新后的任务信息
"""
adapter = VideoTaskCancelAdapter(task_id=task_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# ==================== Adapters ====================
@dataclass
class VideoTaskListAdapter(DashboardAdapter):
"""视频任务列表适配器"""
status: str | None
user_id: str | None
model: str | None
page: int
page_size: int
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
user = context.user
is_admin = user.role == UserRole.ADMIN
query = db.query(VideoTask)
# 权限过滤:普通用户只能看自己的任务
if not is_admin:
query = query.filter(VideoTask.user_id == user.id)
elif self.user_id:
# 管理员可以按用户筛选
query = query.filter(VideoTask.user_id == self.user_id)
# 状态筛选
if self.status:
query = query.filter(VideoTask.status == self.status)
# 模型筛选
if self.model:
escaped = self.model.replace("%", "\\%").replace("_", "\\_")
query = query.filter(VideoTask.model.ilike(f"%{escaped}%"))
# 统计总数
total = query.count()
# 分页
offset = (self.page - 1) * self.page_size
tasks = (
query.order_by(VideoTask.created_at.desc()).offset(offset).limit(self.page_size).all()
)
# 获取用户信息映射
user_ids = list(set(t.user_id for t in tasks if t.user_id))
users_map = {}
if user_ids:
users = db.query(User).filter(User.id.in_(user_ids)).all()
users_map = {u.id: u.username for u in users}
# 获取 Provider 信息映射
provider_ids = list(set(t.provider_id for t in tasks if t.provider_id))
providers_map = {}
if provider_ids:
providers = db.query(Provider).filter(Provider.id.in_(provider_ids)).all()
providers_map = {p.id: p.name for p in providers}
items = []
for task in tasks:
items.append(
{
"id": task.id,
"external_task_id": task.external_task_id,
"user_id": task.user_id,
"username": users_map.get(task.user_id, "Unknown"),
"model": task.model,
"prompt": (
task.prompt[:100] + "..."
if task.prompt and len(task.prompt) > 100
else task.prompt
),
"status": task.status,
"progress_percent": task.progress_percent,
"progress_message": task.progress_message,
"provider_id": task.provider_id,
"provider_name": providers_map.get(task.provider_id, "Unknown"),
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"video_url": task.video_url,
"error_code": task.error_code,
"error_message": task.error_message,
"poll_count": task.poll_count,
"max_poll_count": task.max_poll_count,
"created_at": task.created_at.isoformat() if task.created_at else None,
"completed_at": task.completed_at.isoformat() if task.completed_at else None,
"submitted_at": task.submitted_at.isoformat() if task.submitted_at else None,
}
)
pages = (total + self.page_size - 1) // self.page_size
return {
"items": items,
"total": total,
"page": self.page,
"page_size": self.page_size,
"pages": pages,
}
@dataclass
class VideoTaskStatsAdapter(DashboardAdapter):
"""视频任务统计适配器"""
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
user = context.user
is_admin = user.role == UserRole.ADMIN
base_query = db.query(VideoTask)
if not is_admin:
base_query = base_query.filter(VideoTask.user_id == user.id)
# 总数
total = base_query.count()
# 按状态分组
status_stats = (
base_query.with_entities(
VideoTask.status,
func.count(VideoTask.id).label("count"),
)
.group_by(VideoTask.status)
.all()
)
by_status = {stat.status: stat.count for stat in status_stats}
# 按模型分组(前 10
model_stats = (
base_query.with_entities(
VideoTask.model,
func.count(VideoTask.id).label("count"),
)
.group_by(VideoTask.model)
.order_by(func.count(VideoTask.id).desc())
.limit(10)
.all()
)
by_model = {stat.model: stat.count for stat in model_stats}
# 今日任务数
today = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
today_count = base_query.filter(VideoTask.created_at >= today).count()
# 管理员额外统计
result = {
"total": total,
"by_status": by_status,
"by_model": by_model,
"today_count": today_count,
}
if is_admin:
# 活跃用户数(有视频任务的用户)
active_users = db.query(func.count(func.distinct(VideoTask.user_id))).scalar() or 0
result["active_users"] = active_users
# 处理中的任务数
processing_count = (
db.query(func.count(VideoTask.id))
.filter(VideoTask.status.in_(["submitted", "queued", "processing"]))
.scalar()
or 0
)
result["processing_count"] = processing_count
return result
@dataclass
class VideoTaskDetailAdapter(DashboardAdapter):
"""视频任务详情适配器"""
task_id: str
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
user = context.user
is_admin = user.role == UserRole.ADMIN
query = db.query(VideoTask).filter(VideoTask.id == self.task_id)
if not is_admin:
query = query.filter(VideoTask.user_id == user.id)
task = query.first()
if not task:
raise HTTPException(status_code=404, detail="Video task not found")
# 获取用户信息
task_user = db.query(User).filter(User.id == task.user_id).first()
username = task_user.username if task_user else "Unknown"
# 获取 Provider 信息
provider = db.query(Provider).filter(Provider.id == task.provider_id).first()
provider_name = provider.name if provider else "Unknown"
# 获取 Endpoint 信息
endpoint = (
db.query(ProviderEndpoint).filter(ProviderEndpoint.id == task.endpoint_id).first()
)
endpoint_info = None
if endpoint:
endpoint_info = {
"id": endpoint.id,
"base_url": endpoint.base_url,
"api_format": str(endpoint.api_format),
}
return {
"id": task.id,
"external_task_id": task.external_task_id,
"user_id": task.user_id,
"username": username,
"api_key_id": task.api_key_id,
"provider_id": task.provider_id,
"provider_name": provider_name,
"endpoint_id": task.endpoint_id,
"endpoint": endpoint_info,
"key_id": task.key_id,
"client_api_format": task.client_api_format,
"provider_api_format": task.provider_api_format,
"format_converted": task.format_converted,
"model": task.model,
"prompt": task.prompt,
"original_request_body": task.original_request_body,
"converted_request_body": task.converted_request_body,
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
"status": task.status,
"progress_percent": task.progress_percent,
"progress_message": task.progress_message,
"video_url": task.video_url,
"video_urls": task.video_urls,
"thumbnail_url": task.thumbnail_url,
"video_size_bytes": task.video_size_bytes,
"video_expires_at": (
task.video_expires_at.isoformat() if task.video_expires_at else None
),
"stored_video_path": task.stored_video_path,
"storage_provider": task.storage_provider,
"error_code": task.error_code,
"error_message": task.error_message,
"retry_count": task.retry_count,
"max_retries": task.max_retries,
"poll_interval_seconds": task.poll_interval_seconds,
"next_poll_at": task.next_poll_at.isoformat() if task.next_poll_at else None,
"poll_count": task.poll_count,
"max_poll_count": task.max_poll_count,
"created_at": task.created_at.isoformat() if task.created_at else None,
"updated_at": task.updated_at.isoformat() if task.updated_at else None,
"submitted_at": task.submitted_at.isoformat() if task.submitted_at else None,
"completed_at": task.completed_at.isoformat() if task.completed_at else None,
"request_metadata": task.request_metadata,
}
@dataclass
class VideoTaskCancelAdapter(DashboardAdapter):
"""视频任务取消适配器"""
task_id: str
async def handle(self, context: ApiRequestContext) -> Any:
from src.core.api_format.conversion.internal_video import VideoStatus
db = context.db
user = context.user
is_admin = user.role == UserRole.ADMIN
query = db.query(VideoTask).filter(VideoTask.id == self.task_id)
if not is_admin:
query = query.filter(VideoTask.user_id == user.id)
task = query.first()
if not task:
raise HTTPException(status_code=404, detail="Video task not found")
# 只能取消进行中的任务
if task.status in [
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
]:
raise HTTPException(
status_code=400, detail=f"Cannot cancel task with status: {task.status}"
)
# 更新状态
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
task.completed_at = datetime.now(timezone.utc)
db.commit()
return {
"id": task.id,
"status": task.status,
"message": "Task cancelled successfully",
}

View File

@@ -16,6 +16,7 @@ from sqlalchemy.orm import Session
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult
if TYPE_CHECKING:
import httpx
@@ -45,6 +46,24 @@ def sanitize_error_message(message: str, max_length: int = 200) -> str:
return sanitized[:max_length]
def normalize_gemini_operation_id(operation_id: str) -> str:
"""
规范化 Gemini operation ID确保以 "operations/" 开头
Gemini API 返回的任务 ID 格式可能是 "operations/xxx""xxx"
此函数统一规范化为 "operations/xxx" 格式。
Args:
operation_id: 原始 operation ID
Returns:
规范化后的 operation ID
"""
if not operation_id.startswith("operations/"):
return f"operations/{operation_id}"
return operation_id
class VideoHandlerBase(ABC):
"""视频处理器基类"""
@@ -204,5 +223,28 @@ class VideoHandlerBase(ABC):
extra={"model": task.model},
)
def _build_billing_rule_snapshot(
self, rule_lookup: BillingRuleLookupResult | None
) -> dict[str, Any]:
"""
构建 billing_rule 快照,用于冻结到视频任务的 request_metadata 中。
__all__ = ["VideoHandlerBase", "sanitize_error_message"]
快照确保异步任务完成时使用创建时刻的计费规则,避免规则变更导致成本计算不一致。
"""
if not rule_lookup:
return {"status": "no_rule"}
rule = rule_lookup.rule
return {
"status": "ok",
"scope": rule_lookup.scope,
"effective_task_type": rule_lookup.effective_task_type,
"rule_id": rule.id,
"rule_name": rule.name,
"expression": rule.expression,
"variables": rule.variables,
"dimension_mappings": rule.dimension_mappings,
}
__all__ = ["VideoHandlerBase", "normalize_gemini_operation_id", "sanitize_error_message"]

View File

@@ -14,8 +14,13 @@ from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import get_provider_auth
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
from src.api.handlers.base.video_handler_base import (
VideoHandlerBase,
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.clients.http_client import HTTPClientPool
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 (
InternalVideoRequest,
@@ -27,6 +32,7 @@ from src.core.api_format.headers import HOP_BY_HOP_HEADERS
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
@@ -34,8 +40,6 @@ class GeminiVeoHandler(VideoHandlerBase):
FORMAT_ID = "GEMINI"
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
POLL_INTERVAL_SECONDS = 10
MAX_POLL_COUNT = 360
def __init__(
self,
@@ -80,16 +84,39 @@ class GeminiVeoHandler(VideoHandlerBase):
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
candidate = await self._select_candidate(internal_request.model)
candidate, candidate_keys, rule_lookup = await self._select_candidate(
internal_request.model
)
if not candidate:
raise HTTPException(
status_code=503, detail="No available provider for video generation"
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"
raise HTTPException(status_code=503, detail=detail)
# 冻结 billing_rule 配置(用于异步任务的成本一致性)
# 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查
if rule_lookup is None:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=internal_request.model,
task_type="video",
)
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
upstream_key, endpoint, key, auth_info = await self._resolve_upstream_key(candidate)
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint, auth_info)
api_key_header = (
headers.get("x-goog-api-key", "")[:10] + "..."
if headers.get("x-goog-api-key")
else "MISSING"
)
logger.info(
f"[GeminiVeoHandler] Create task: endpoint_id={endpoint.id}, base_url={endpoint.base_url}, upstream_url={upstream_url}, api_key_prefix={api_key_header}"
)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
if response.status_code >= 400:
@@ -99,18 +126,25 @@ class GeminiVeoHandler(VideoHandlerBase):
external_task_id = str(payload.get("name") or "")
if not external_task_id:
raise HTTPException(status_code=502, detail="Upstream returned empty task id")
external_task_id = normalize_gemini_operation_id(external_task_id)
task = self._create_task_record(
external_task_id=external_task_id,
candidate=candidate,
original_request_body=original_request_body,
internal_request=internal_request,
candidate_keys=candidate_keys,
original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot,
)
try:
self.db.add(task)
self.db.flush() # 先 flush 检测冲突
self.db.commit()
self.db.refresh(task)
logger.info(
f"[GeminiVeoHandler] Task created: id={task.id}, external_task_id={task.external_task_id}, user_id={task.user_id}"
)
except IntegrityError:
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
@@ -136,6 +170,8 @@ class GeminiVeoHandler(VideoHandlerBase):
) -> JSONResponse:
# Gemini 使用 operations/{id} 格式,需要按 external_task_id 查找
task = self._get_task_by_external_id(task_id)
# 直接从数据库返回任务状态(后台轮询服务会持续更新状态)
internal_task = self._task_to_internal(task)
response_body = self._normalizer.video_task_from_internal(internal_task)
return JSONResponse(response_body)
@@ -272,7 +308,10 @@ class GeminiVeoHandler(VideoHandlerBase):
# Helpers
# ------------------------------------------------------------------
async def _select_candidate(self, model_name: str) -> ProviderCandidate | None:
async def _select_candidate(
self, model_name: str
) -> tuple[ProviderCandidate | None, list[dict[str, Any]], BillingRuleLookupResult | None]:
"""选择候选 key返回 (选中的候选, 所有候选列表, 选中候选的 billing rule lookup)"""
scheduler = CacheAwareScheduler()
candidates, _ = await scheduler.list_all_candidates(
db=self.db,
@@ -282,11 +321,47 @@ class GeminiVeoHandler(VideoHandlerBase):
user_api_key=self.api_key,
max_candidates=10,
)
for candidate in candidates:
# 记录所有候选 key 信息
candidate_keys = []
selected_candidate = None
selected_index = -1
selected_rule_lookup: BillingRuleLookupResult | None = None
for idx, candidate in enumerate(candidates):
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
if auth_type in {"api_key", "vertex_ai"}:
return candidate
return None
has_billing_rule = True
rule_lookup: BillingRuleLookupResult | None = None
if config.billing_require_rule:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=model_name,
task_type="video",
)
has_billing_rule = rule_lookup is not None
candidate_info = {
"index": idx,
"provider_id": candidate.provider.id,
"provider_name": candidate.provider.name,
"endpoint_id": candidate.endpoint.id,
"key_id": candidate.key.id,
"key_name": candidate.key.name,
"auth_type": auth_type,
"has_billing_rule": has_billing_rule,
"priority": getattr(candidate.key, "priority", 0) or 0,
}
candidate_keys.append(candidate_info)
if (
selected_candidate is None
and auth_type in {"api_key", "vertex_ai"}
and has_billing_rule
):
selected_candidate = candidate
selected_index = idx
selected_rule_lookup = rule_lookup
# 标记选中的候选
if selected_index >= 0:
candidate_keys[selected_index]["selected"] = True
return selected_candidate, candidate_keys, selected_rule_lookup
async def _resolve_upstream_key(
self, candidate: ProviderCandidate
@@ -351,8 +426,31 @@ class GeminiVeoHandler(VideoHandlerBase):
candidate: ProviderCandidate,
original_request_body: dict[str, Any],
internal_request: Any,
candidate_keys: list[dict[str, Any]] | None = None,
original_headers: dict[str, str] | None = None,
billing_rule_snapshot: dict[str, Any] | None = None,
) -> VideoTask:
now = datetime.now(timezone.utc)
# 构建请求元数据(使用追踪信息)
request_metadata = {
"candidate_keys": candidate_keys or [],
"selected_key_id": candidate.key.id,
"selected_endpoint_id": candidate.endpoint.id,
"client_ip": self.client_ip,
"user_agent": self.user_agent,
"request_id": self.request_id,
"billing_rule_snapshot": billing_rule_snapshot,
}
# 记录请求头(脱敏处理)
if original_headers:
safe_headers = {
k: v
for k, v in original_headers.items()
if k.lower() not in {"authorization", "x-api-key", "x-goog-api-key", "cookie"}
}
request_metadata["request_headers"] = safe_headers
return VideoTask(
id=str(uuid4()),
external_task_id=external_task_id,
@@ -373,18 +471,21 @@ class GeminiVeoHandler(VideoHandlerBase):
aspect_ratio=internal_request.aspect_ratio,
status=VideoStatus.SUBMITTED.value,
progress_percent=0,
poll_interval_seconds=self.POLL_INTERVAL_SECONDS,
next_poll_at=now + timedelta(seconds=self.POLL_INTERVAL_SECONDS),
poll_interval_seconds=config.video_poll_interval_seconds,
next_poll_at=now + timedelta(seconds=config.video_poll_interval_seconds),
poll_count=0,
max_poll_count=self.MAX_POLL_COUNT,
max_poll_count=config.video_max_poll_count,
submitted_at=now,
request_metadata=request_metadata,
)
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
"""按 external_task_id 查找任务Gemini 使用 operations/{id} 格式)"""
normalized_id = external_id
if not normalized_id.startswith("operations/"):
normalized_id = f"operations/{normalized_id}"
normalized_id = normalize_gemini_operation_id(external_id)
logger.info(
f"[GeminiVeoHandler] Looking for task: normalized_id={normalized_id}, user_id={self.user.id}"
)
task = (
self.db.query(VideoTask)
@@ -395,7 +496,13 @@ class GeminiVeoHandler(VideoHandlerBase):
.first()
)
if not task:
logger.warning(
f"[GeminiVeoHandler] Task not found: normalized_id={normalized_id}, user_id={self.user.id}"
)
raise HTTPException(status_code=404, detail="Video task not found")
logger.info(
f"[GeminiVeoHandler] Task found: id={task.id}, external_task_id={task.external_task_id}"
)
return task

View File

@@ -16,6 +16,7 @@ from sqlalchemy.orm import Session
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
from src.clients.http_client import HTTPClientPool
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 (
InternalVideoRequest,
@@ -27,6 +28,7 @@ from src.core.api_format.headers import HOP_BY_HOP_HEADERS
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
@@ -34,8 +36,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
FORMAT_ID = "OPENAI"
DEFAULT_BASE_URL = "https://api.openai.com"
POLL_INTERVAL_SECONDS = 10
MAX_POLL_COUNT = 360
def __init__(
self,
@@ -73,11 +73,25 @@ class OpenAIVideoHandler(VideoHandlerBase):
internal_request = self._normalizer.video_request_to_internal(original_request_body)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
candidate = await self._select_candidate(internal_request.model)
candidate, candidate_keys, rule_lookup = await self._select_candidate(
internal_request.model
)
if not candidate:
raise HTTPException(
status_code=503, detail="No available provider for video generation"
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"
raise HTTPException(status_code=503, detail=detail)
# 冻结 billing_rule 配置(用于异步任务的成本一致性)
# 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查
if rule_lookup is None:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=internal_request.model,
task_type="video",
)
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
upstream_key, endpoint, provider_key = await self._resolve_upstream_key(candidate)
upstream_url = self._build_upstream_url(endpoint.base_url)
@@ -98,6 +112,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
candidate=candidate,
original_request_body=original_request_body,
internal_request=internal_request,
candidate_keys=candidate_keys,
original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot,
)
try:
self.db.add(task)
@@ -282,7 +299,10 @@ class OpenAIVideoHandler(VideoHandlerBase):
# Helpers
# ------------------------------------------------------------------
async def _select_candidate(self, model_name: str) -> ProviderCandidate | None:
async def _select_candidate(
self, model_name: str
) -> tuple[ProviderCandidate | None, list[dict[str, Any]], BillingRuleLookupResult | None]:
"""选择候选 key返回 (选中的候选, 所有候选列表, 选中候选的 billing rule lookup)"""
scheduler = CacheAwareScheduler()
candidates, _ = await scheduler.list_all_candidates(
db=self.db,
@@ -292,11 +312,43 @@ class OpenAIVideoHandler(VideoHandlerBase):
user_api_key=self.api_key,
max_candidates=10,
)
for candidate in candidates:
# 记录所有候选 key 信息
candidate_keys = []
selected_candidate = None
selected_index = -1
selected_rule_lookup: BillingRuleLookupResult | None = None
for idx, candidate in enumerate(candidates):
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
if auth_type == "api_key":
return candidate
return None
has_billing_rule = True
rule_lookup: BillingRuleLookupResult | None = None
if config.billing_require_rule:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=model_name,
task_type="video",
)
has_billing_rule = rule_lookup is not None
candidate_info = {
"index": idx,
"provider_id": candidate.provider.id,
"provider_name": candidate.provider.name,
"endpoint_id": candidate.endpoint.id,
"key_id": candidate.key.id,
"key_name": candidate.key.name,
"auth_type": auth_type,
"has_billing_rule": has_billing_rule,
"priority": getattr(candidate.key, "priority", 0) or 0,
}
candidate_keys.append(candidate_info)
if selected_candidate is None and auth_type == "api_key" and has_billing_rule:
selected_candidate = candidate
selected_index = idx
selected_rule_lookup = rule_lookup
# 标记选中的候选
if selected_index >= 0:
candidate_keys[selected_index]["selected"] = True
return selected_candidate, candidate_keys, selected_rule_lookup
async def _resolve_upstream_key(
self, candidate: ProviderCandidate
@@ -342,9 +394,32 @@ class OpenAIVideoHandler(VideoHandlerBase):
candidate: ProviderCandidate,
original_request_body: dict[str, Any],
internal_request: InternalVideoRequest,
candidate_keys: list[dict[str, Any]] | None = None,
original_headers: dict[str, str] | None = None,
billing_rule_snapshot: dict[str, Any] | None = None,
) -> VideoTask:
now = datetime.now(timezone.utc)
size = internal_request.extra.get("original_size")
# 构建请求元数据(使用追踪信息)
request_metadata = {
"candidate_keys": candidate_keys or [],
"selected_key_id": candidate.key.id,
"selected_endpoint_id": candidate.endpoint.id,
"client_ip": self.client_ip,
"user_agent": self.user_agent,
"request_id": self.request_id,
"billing_rule_snapshot": billing_rule_snapshot,
}
# 记录请求头(脱敏处理)
if original_headers:
safe_headers = {
k: v
for k, v in original_headers.items()
if k.lower() not in {"authorization", "x-api-key", "cookie"}
}
request_metadata["request_headers"] = safe_headers
return VideoTask(
id=str(uuid4()),
external_task_id=external_task_id,
@@ -366,11 +441,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
size=size,
status=VideoStatus.SUBMITTED.value,
progress_percent=0,
poll_interval_seconds=self.POLL_INTERVAL_SECONDS,
next_poll_at=now + timedelta(seconds=self.POLL_INTERVAL_SECONDS),
poll_interval_seconds=config.video_poll_interval_seconds,
next_poll_at=now + timedelta(seconds=config.video_poll_interval_seconds),
poll_count=0,
max_poll_count=self.MAX_POLL_COUNT,
max_poll_count=config.video_max_poll_count,
submitted_at=now,
request_metadata=request_metadata,
)
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:

View File

@@ -106,7 +106,7 @@ async def create_video_veo(model: str, http_request: Request, db: Session = Depe
)
@router.get("/v1beta/operations/{operation_id}")
@router.get("/v1beta/operations/{operation_id:path}")
async def get_video_veo(
operation_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
@@ -148,7 +148,7 @@ async def cancel_video_veo(
)
@router.get("/v1beta/operations/{operation_id}/content")
@router.get("/v1beta/operations/{operation_id:path}/content")
async def download_video_content_veo(
operation_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:

View File

@@ -3,9 +3,9 @@
从环境变量或 .env 文件加载配置
"""
from typing import Any
import os
from pathlib import Path
from typing import Any
# 尝试加载 .env 文件
try:
@@ -110,9 +110,9 @@ class Config:
# 异常处理配置
# 设置为 True 时ProxyException 会传播到路由层以便记录 provider_request_headers
# 设置为 False 时,使用全局异常处理器统一处理
self.propagate_provider_exceptions = os.getenv(
"PROPAGATE_PROVIDER_EXCEPTIONS", "true"
).lower() == "true"
self.propagate_provider_exceptions = (
os.getenv("PROPAGATE_PROVIDER_EXCEPTIONS", "true").lower() == "true"
)
# 数据库连接池配置 - 智能自动调整
# 系统会根据 Worker 数量和 PostgreSQL 限制自动计算安全值
@@ -151,17 +151,17 @@ class Config:
# 格式转换配置
# FORMAT_CONVERSION_ENABLED: 全局格式转换总开关,默认开启
# 注意:即使开启,也需要端点配置 format_acceptance_config.enabled=true 才能生效
self.format_conversion_enabled = os.getenv(
"FORMAT_CONVERSION_ENABLED", "true"
).lower() == "true"
self.format_conversion_enabled = (
os.getenv("FORMAT_CONVERSION_ENABLED", "true").lower() == "true"
)
# KEEP_PRIORITY_ON_CONVERSION: 格式转换时是否保持提供商原优先级,默认关闭
# - false默认: 需要格式转换的候选整体降级到不需要转换的候选之后
# - true: 所有提供商保持原优先级,不因格式转换降级
# 注意:即使全局关闭,单个提供商也可以通过 keep_priority_on_conversion 字段保持自己的优先级
self.keep_priority_on_conversion = os.getenv(
"KEEP_PRIORITY_ON_CONVERSION", "false"
).lower() == "true"
self.keep_priority_on_conversion = (
os.getenv("KEEP_PRIORITY_ON_CONVERSION", "false").lower() == "true"
)
# HTTP 连接池配置
# HTTP_MAX_CONNECTIONS: 最大连接数,影响并发能力
@@ -200,24 +200,14 @@ class Config:
os.getenv("USAGE_QUEUE_INCLUDE_BODIES", "true").lower() == "true"
)
# 0 表示不截断由系统设置max_request/response_body_size统一控制
self.usage_queue_body_max_bytes = int(
os.getenv("USAGE_QUEUE_BODY_MAX_BYTES", "0")
)
self.usage_queue_body_max_bytes = int(os.getenv("USAGE_QUEUE_BODY_MAX_BYTES", "0"))
self.usage_queue_stream_key = os.getenv("USAGE_QUEUE_STREAM_KEY", "usage:events")
self.usage_queue_stream_group = os.getenv(
"USAGE_QUEUE_STREAM_GROUP", "usage_consumers"
)
self.usage_queue_stream_maxlen = int(
os.getenv("USAGE_QUEUE_STREAM_MAXLEN", "200000")
)
self.usage_queue_stream_group = os.getenv("USAGE_QUEUE_STREAM_GROUP", "usage_consumers")
self.usage_queue_stream_maxlen = int(os.getenv("USAGE_QUEUE_STREAM_MAXLEN", "200000"))
self.usage_queue_dlq_key = os.getenv("USAGE_QUEUE_DLQ_KEY", "usage:events:dlq")
self.usage_queue_dlq_maxlen = int(os.getenv("USAGE_QUEUE_DLQ_MAXLEN", "5000"))
self.usage_queue_consumer_batch = int(
os.getenv("USAGE_QUEUE_CONSUMER_BATCH", "200")
)
self.usage_queue_consumer_block_ms = int(
os.getenv("USAGE_QUEUE_CONSUMER_BLOCK_MS", "500")
)
self.usage_queue_consumer_batch = int(os.getenv("USAGE_QUEUE_CONSUMER_BATCH", "200"))
self.usage_queue_consumer_block_ms = int(os.getenv("USAGE_QUEUE_CONSUMER_BLOCK_MS", "500"))
self.usage_queue_claim_idle_ms = int(os.getenv("USAGE_QUEUE_CLAIM_IDLE_MS", "30000"))
self.usage_queue_claim_interval_seconds = float(
os.getenv("USAGE_QUEUE_CLAIM_INTERVAL_SECONDS", "5")
@@ -245,12 +235,8 @@ class Config:
self.internal_user_agent_claude_cli = os.getenv(
"CLAUDE_CLI_USER_AGENT", "claude-code/1.0.1"
)
self.internal_user_agent_openai_cli = os.getenv(
"OPENAI_CLI_USER_AGENT", "openai-codex/1.0"
)
self.internal_user_agent_gemini_cli = os.getenv(
"GEMINI_CLI_USER_AGENT", "gemini-cli/0.1.0"
)
self.internal_user_agent_openai_cli = os.getenv("OPENAI_CLI_USER_AGENT", "openai-codex/1.0")
self.internal_user_agent_gemini_cli = os.getenv("GEMINI_CLI_USER_AGENT", "gemini-cli/0.1.0")
# 邮箱验证配置
# VERIFICATION_CODE_EXPIRE_MINUTES: 验证码有效期(分钟)
@@ -258,19 +244,29 @@ class Config:
self.verification_code_expire_minutes = int(
os.getenv("VERIFICATION_CODE_EXPIRE_MINUTES", "5")
)
self.verification_send_cooldown = int(
os.getenv("VERIFICATION_SEND_COOLDOWN", "60")
)
self.verification_send_cooldown = int(os.getenv("VERIFICATION_SEND_COOLDOWN", "60"))
# 计费系统配置(多维度计费 / 异步任务)
# BILLING_REQUIRE_RULE: Video/Image/Audio 缺失 billing_rule 时是否拒绝请求(默认 false缺失则 cost=0 并告警)
# BILLING_STRICT_MODE: required 维度缺失时是否拒绝请求/标记任务失败(默认 false缺失则 cost=0 + 标记 incomplete
self.billing_require_rule = os.getenv("BILLING_REQUIRE_RULE", "false").lower() == "true"
self.billing_strict_mode = os.getenv("BILLING_STRICT_MODE", "false").lower() == "true"
# 视频任务轮询配置
# VIDEO_POLL_INTERVAL_SECONDS: 轮询间隔(秒),默认 10 秒
# VIDEO_MAX_POLL_COUNT: 最大轮询次数,默认 360 次(约 1 小时)
# VIDEO_POLL_BATCH_SIZE: 每批处理任务数,默认 50
# VIDEO_POLL_CONCURRENCY: 并发轮询数,默认 10
self.video_poll_interval_seconds = int(os.getenv("VIDEO_POLL_INTERVAL_SECONDS", "10"))
self.video_max_poll_count = int(os.getenv("VIDEO_MAX_POLL_COUNT", "360"))
self.video_poll_batch_size = int(os.getenv("VIDEO_POLL_BATCH_SIZE", "50"))
self.video_poll_concurrency = int(os.getenv("VIDEO_POLL_CONCURRENCY", "10"))
# Management Token 速率限制(每分钟每 IP
self.management_token_rate_limit = int(
os.getenv("MANAGEMENT_TOKEN_RATE_LIMIT", "30")
)
self.management_token_rate_limit = int(os.getenv("MANAGEMENT_TOKEN_RATE_LIMIT", "30"))
# 每个用户最多可创建的 Management Token 数量
self.management_token_max_per_user = int(
os.getenv("MANAGEMENT_TOKEN_MAX_PER_USER", "20")
)
self.management_token_max_per_user = int(os.getenv("MANAGEMENT_TOKEN_MAX_PER_USER", "20"))
# API 文档配置
# DOCS_ENABLED: 是否启用 API 文档(/docs, /redoc, /openapi.json
@@ -416,9 +412,7 @@ class Config:
# 加密密钥警告
if not self.encryption_key and self.environment != "production":
logger.warning(
"ENCRYPTION_KEY 未设置,使用开发环境默认密钥。生产环境必须设置。"
)
logger.warning("ENCRYPTION_KEY 未设置,使用开发环境默认密钥。生产环境必须设置。")
# CORS 配置警告(生产环境)
if self.environment == "production" and not self.cors_origins:

View File

@@ -843,6 +843,16 @@ class GeminiNormalizer(FormatNormalizer):
},
}
if internal.status == VideoStatus.FAILED:
return {
"name": operation_name,
"done": True,
"error": {
"code": internal.error_code or "UNKNOWN",
"message": internal.error_message or "Video generation failed",
},
}
return {
"name": operation_name,
"done": False,

View File

@@ -28,6 +28,7 @@ from sqlalchemy import (
String,
Text,
UniqueConstraint,
text,
)
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import declarative_base, relationship
@@ -1097,6 +1098,122 @@ class Model(Base):
return names
class BillingRule(Base):
"""计费规则表(单条 formula 规则,支持 Model 覆盖 GlobalModel"""
__tablename__ = "billing_rules"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 规则关联(两者必有其一)
global_model_id = Column(
String(36), ForeignKey("global_models.id", ondelete="CASCADE"), nullable=True, index=True
)
model_id = Column(
String(36), ForeignKey("models.id", ondelete="CASCADE"), nullable=True, index=True
)
name = Column(String(100), nullable=False)
# 注CLI 在计费域里恒等于 chat不单独存 "cli"
task_type = Column(String(20), nullable=False, default="chat")
# Formula 表达式及其配置
expression = Column(Text, nullable=False)
variables = Column(JSONB, nullable=False, default=dict)
dimension_mappings = Column(JSONB, nullable=False, default=dict)
is_enabled = Column(Boolean, nullable=False, default=True)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
global_model = relationship("GlobalModel", foreign_keys=[global_model_id])
model = relationship("Model", foreign_keys=[model_id])
__table_args__ = (
CheckConstraint(
"(global_model_id IS NOT NULL AND model_id IS NULL) OR "
"(global_model_id IS NULL AND model_id IS NOT NULL)",
name="chk_billing_rules_model_ref",
),
# 同级同 task_type 只允许一条启用规则partial unique index
Index(
"uq_billing_rules_global_model_task",
"global_model_id",
"task_type",
unique=True,
postgresql_where=text("is_enabled = TRUE AND global_model_id IS NOT NULL"),
),
Index(
"uq_billing_rules_model_task",
"model_id",
"task_type",
unique=True,
postgresql_where=text("is_enabled = TRUE AND model_id IS NOT NULL"),
),
)
class DimensionCollector(Base):
"""维度收集器配置表(从请求/响应/元数据/派生计算收集维度)。"""
__tablename__ = "dimension_collectors"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
api_format = Column(String(50), nullable=False)
task_type = Column(String(20), nullable=False)
dimension_name = Column(String(100), nullable=False)
# 来源配置
# - response / request / metadata / computed
source_type = Column(String(20), nullable=False)
source_path = Column(String(200), nullable=True) # computed 允许为空
# 值类型与转换
value_type = Column(String(20), nullable=False, default="float") # float/int/string
transform_expression = Column(Text, nullable=True) # computed 时为派生公式
default_value = Column(String(100), nullable=True)
priority = Column(Integer, nullable=False, default=0)
is_enabled = Column(Boolean, nullable=False, default=True)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (
CheckConstraint(
"(source_type = 'computed' AND source_path IS NULL AND transform_expression IS NOT NULL) OR "
"(source_type != 'computed' AND source_path IS NOT NULL)",
name="chk_dimension_collectors_source_config",
),
# 同维度 + 同优先级 + enabled 才唯一(允许禁用旧配置后重建)
Index(
"uq_dimension_collectors_enabled",
"api_format",
"task_type",
"dimension_name",
"priority",
unique=True,
postgresql_where=text("is_enabled = TRUE"),
),
)
class ProviderAPIKey(Base):
"""Provider API密钥表 - 直接归属于 Provider支持多种 API 格式"""
@@ -1290,6 +1407,16 @@ class VideoTask(Base):
String(36), ForeignKey("video_tasks.id", ondelete="SET NULL"), nullable=True
)
# 使用追踪(候选 key、请求头等
request_metadata = Column(JSON, nullable=True) # 存储候选 key 列表、请求头等追踪信息
# 示例: {
# "candidate_keys": [{"key_id": "xxx", "endpoint_id": "yyy", "priority": 1}, ...],
# "selected_key_index": 0,
# "client_ip": "1.2.3.4",
# "user_agent": "...",
# "request_headers": {...}
# }
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False

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)

View File

@@ -0,0 +1,105 @@
from src.models.database import DimensionCollector
from src.services.billing.dimension_collector_service import (
DimensionCollectInput,
DimensionCollectorRuntime,
)
class TestDimensionCollectorRuntime:
def test_priority_fallback(self) -> None:
runtime = DimensionCollectorRuntime()
collectors = [
DimensionCollector(
api_format="OPENAI",
task_type="chat",
dimension_name="input_tokens",
source_type="response",
source_path="usage.prompt_tokens",
value_type="int",
priority=10,
is_enabled=True,
),
DimensionCollector(
api_format="OPENAI",
task_type="chat",
dimension_name="input_tokens",
source_type="response",
source_path="usageMetadata.promptTokenCount",
value_type="int",
priority=5,
is_enabled=True,
),
]
dims = runtime.collect(
collectors=collectors,
inp=DimensionCollectInput(
response={"usageMetadata": {"promptTokenCount": 123}},
),
)
assert dims["input_tokens"] == 123
def test_transform_expression_value(self) -> None:
runtime = DimensionCollectorRuntime()
collectors = [
DimensionCollector(
api_format="GEMINI",
task_type="video",
dimension_name="file_size_mb",
source_type="metadata",
source_path="result.file_size_bytes",
transform_expression="value / 1024 / 1024",
value_type="float",
priority=0,
is_enabled=True,
)
]
dims = runtime.collect(
collectors=collectors,
inp=DimensionCollectInput(metadata={"result": {"file_size_bytes": 1048576}}),
)
assert abs(dims["file_size_mb"] - 1.0) < 1e-9
def test_computed_dimension(self) -> None:
runtime = DimensionCollectorRuntime()
collectors = [
DimensionCollector(
api_format="CLAUDE",
task_type="chat",
dimension_name="input_tokens",
source_type="request",
source_path="usage.input_tokens",
value_type="int",
priority=0,
is_enabled=True,
),
DimensionCollector(
api_format="CLAUDE",
task_type="chat",
dimension_name="cache_read_tokens",
source_type="request",
source_path="usage.cache_read_tokens",
value_type="int",
priority=0,
is_enabled=True,
),
DimensionCollector(
api_format="CLAUDE",
task_type="chat",
dimension_name="total_input_tokens",
source_type="computed",
source_path=None,
transform_expression="input_tokens + cache_read_tokens",
value_type="int",
priority=0,
is_enabled=True,
),
]
dims = runtime.collect(
collectors=collectors,
inp=DimensionCollectInput(
request={"usage": {"input_tokens": 100, "cache_read_tokens": 20}}
),
)
assert dims["input_tokens"] == 100
assert dims["cache_read_tokens"] == 20
assert dims["total_input_tokens"] == 120

View File

@@ -0,0 +1,169 @@
import pytest
from src.services.billing.formula_engine import (
BillingIncompleteError,
FormulaEngine,
UnsafeExpressionError,
)
class TestSafeExpression:
def test_reject_attribute_access(self) -> None:
engine = FormulaEngine()
with pytest.raises(UnsafeExpressionError):
engine.evaluate(
expression="(1).__class__",
variables={},
dimensions={},
dimension_mappings={},
strict_mode=True, # 确保抛出异常,便于断言类型
)
def test_reject_import(self) -> None:
engine = FormulaEngine()
with pytest.raises(UnsafeExpressionError):
engine.evaluate(
expression="__import__('os').system('echo hacked')",
variables={},
dimensions={},
dimension_mappings={},
strict_mode=True,
)
class TestFormulaEngine:
def test_basic_video_formula(self) -> None:
engine = FormulaEngine()
result = engine.evaluate(
expression="(base_price + duration_seconds * price_per_second) * resolution_multiplier",
variables={"base_price": 0.05, "price_per_second": 0.02},
dimensions={"duration_seconds": 10, "resolution": "720p"},
dimension_mappings={
"duration_seconds": {
"source": "dimension",
"key": "duration_seconds",
"required": True,
},
"resolution_multiplier": {
"source": "matrix",
"key": "resolution",
"map": {"720p": 1.0, "1080p": 1.5},
"default": 1.0,
},
},
)
assert result.status == "complete"
assert result.missing_required == []
assert abs(result.cost - 0.25) < 1e-9
def test_required_dimension_missing_non_strict(self) -> None:
engine = FormulaEngine()
result = engine.evaluate(
expression="base_price + duration_seconds * price_per_second",
variables={"base_price": 0.05, "price_per_second": 0.02},
dimensions={"duration_seconds": None},
dimension_mappings={
"duration_seconds": {
"source": "dimension",
"key": "duration_seconds",
"required": True,
},
},
strict_mode=False,
)
assert result.status == "incomplete"
assert result.cost == 0.0
assert result.missing_required == ["duration_seconds"]
def test_required_dimension_missing_strict_raises(self) -> None:
engine = FormulaEngine()
with pytest.raises(BillingIncompleteError) as exc:
engine.evaluate(
expression="base_price + duration_seconds * price_per_second",
variables={"base_price": 0.05, "price_per_second": 0.02},
dimensions={"duration_seconds": None},
dimension_mappings={
"duration_seconds": {
"source": "dimension",
"key": "duration_seconds",
"required": True,
},
},
strict_mode=True,
)
assert exc.value.missing_required == ["duration_seconds"]
def test_allow_zero(self) -> None:
engine = FormulaEngine()
# allow_zero=false默认0 视为缺失
result = engine.evaluate(
expression="duration_seconds * 1",
variables={},
dimensions={"duration_seconds": 0},
dimension_mappings={
"duration_seconds": {
"source": "dimension",
"key": "duration_seconds",
"required": True,
},
},
)
assert result.status == "incomplete"
assert result.missing_required == ["duration_seconds"]
# allow_zero=true0 合法
result = engine.evaluate(
expression="duration_seconds * 1",
variables={},
dimensions={"duration_seconds": 0},
dimension_mappings={
"duration_seconds": {
"source": "dimension",
"key": "duration_seconds",
"required": True,
"allow_zero": True,
},
},
)
assert result.status == "complete"
assert result.cost == 0.0
def test_tiered_mapping(self) -> None:
engine = FormulaEngine()
result = engine.evaluate(
expression="input_price",
variables={},
dimensions={"total_input_tokens": 100_000},
dimension_mappings={
"input_price": {
"source": "tiered",
"tier_key": "total_input_tokens",
"tiers": [
{"up_to": 128_000, "value": 3.0},
{"up_to": None, "value": 1.5},
],
}
},
)
assert result.status == "complete"
assert result.cost == 3.0
result = engine.evaluate(
expression="input_price",
variables={},
dimensions={"total_input_tokens": 200_000},
dimension_mappings={
"input_price": {
"source": "tiered",
"tier_key": "total_input_tokens",
"tiers": [
{"up_to": 128_000, "value": 3.0},
{"up_to": None, "value": 1.5},
],
}
},
)
assert result.status == "complete"
assert result.cost == 1.5