mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
feat: 添加多维度计费系统和视频任务管理功能
计费系统: - 新增 BillingRule 和 DimensionCollector 数据模型 - 实现 FormulaEngine 安全表达式求值引擎 (AST 白名单) - 支持 dimension/matrix/tiered/constant 多种维度映射 - BillingRuleService 支持 Provider Model -> GlobalModel 规则回退 - CLI task_type 在计费域自动映射为 chat 视频任务增强: - 添加 request_metadata 字段记录候选 key 和计费规则快照 - 后台轮询支持并发控制 (Semaphore + 独立 session) - 任务终态自动写入 Usage 记录并计算成本 - 新增视频任务管理 API 和前端界面 其他改进: - UsageService 新增 record_usage_with_custom_cost 方法 - StandardizedUsage 支持 dimensions 字段 (兼容 extra) - 配置新增 BILLING_REQUIRE_RULE 和 BILLING_STRICT_MODE
This commit is contained in:
@@ -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 维度缺失时是否拒绝请求/标记任务失败(默认 false:cost=0 + 标记 incomplete)
|
||||
# BILLING_STRICT_MODE=false
|
||||
|
||||
1
.gitignore
vendored
1
.gitignore
vendored
@@ -220,6 +220,7 @@ frontend/public/*-firework.svg
|
||||
# Debug and experimental files
|
||||
debug_*.html
|
||||
extracted_*.ts
|
||||
test.py
|
||||
|
||||
# Deploy script cache
|
||||
.deps-hash
|
||||
|
||||
@@ -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")
|
||||
158
frontend/src/api/video-tasks.ts
Normal file
158
frontend/src/api/video-tasks.ts
Normal 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
|
||||
@@ -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 },
|
||||
]
|
||||
},
|
||||
|
||||
@@ -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'))
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
629
frontend/src/views/admin/VideoTasks.vue
Normal file
629
frontend/src/views/admin/VideoTasks.vue
Normal 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>
|
||||
@@ -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 时才会注册路由
|
||||
|
||||
5
src/api/admin/billing/__init__.py
Normal file
5
src/api/admin/billing/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
"""Billing 配置管理 API 模块(billing_rules / dimension_collectors)。"""
|
||||
|
||||
from .routes import router
|
||||
|
||||
__all__ = ["router"]
|
||||
527
src/api/admin/billing/routes.py
Normal file
527
src/api/admin/billing/routes.py
Normal 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 在计费域等同于 chat,billing_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)"
|
||||
)
|
||||
5
src/api/admin/video_tasks/__init__.py
Normal file
5
src/api/admin/video_tasks/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
"""视频任务管理 API 模块。"""
|
||||
|
||||
from .routes import router
|
||||
|
||||
__all__ = ["router"]
|
||||
424
src/api/admin/video_tasks/routes.py
Normal file
424
src/api/admin/video_tasks/routes.py
Normal 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",
|
||||
}
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.services.billing.models import (
|
||||
|
||||
371
src/services/billing/dimension_collector_service.py
Normal file
371
src/services/billing/dimension_collector_service.py
Normal file
@@ -0,0 +1,371 @@
|
||||
"""
|
||||
DimensionCollector 运行时维度采集
|
||||
|
||||
特性(与 .plans/humming-seeking-marble.md 对齐):
|
||||
- (api_format, task_type) 作用域
|
||||
- 同一维度支持多条 collector(priority 回退)
|
||||
- 支持 transform_expression(与 billing expression 共用 AST 安全规范)
|
||||
- computed 维度支持依赖拓扑排序,并对环依赖做保护性降级
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.database import DimensionCollector
|
||||
from src.services.billing.formula_engine import (
|
||||
ExpressionEvaluationError,
|
||||
SafeExpressionEvaluator,
|
||||
UnsafeExpressionError,
|
||||
extract_variable_names,
|
||||
)
|
||||
|
||||
ValueType = Literal["float", "int", "string"]
|
||||
|
||||
|
||||
def _normalize_api_format(api_format: str | None) -> str:
|
||||
return (api_format or "").upper()
|
||||
|
||||
|
||||
def _normalize_task_type(task_type: str | None) -> str:
|
||||
return (task_type or "").lower()
|
||||
|
||||
|
||||
def _get_nested_value(data: Any, path: str) -> Any:
|
||||
"""
|
||||
简单 JSON path:
|
||||
- a.b.c
|
||||
- 列表索引用数字:items.0.id
|
||||
"""
|
||||
if data is None or path is None or path == "":
|
||||
return None
|
||||
|
||||
value: Any = data
|
||||
for key in path.split("."):
|
||||
if isinstance(value, dict):
|
||||
value = value.get(key)
|
||||
elif isinstance(value, list):
|
||||
if not key.isdigit():
|
||||
return None
|
||||
idx = int(key)
|
||||
if idx < 0 or idx >= len(value):
|
||||
return None
|
||||
value = value[idx]
|
||||
else:
|
||||
return None
|
||||
if value is None:
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
def _cast_value(value: Any, value_type: ValueType) -> Any:
|
||||
if value_type == "string":
|
||||
return "" if value is None else str(value)
|
||||
if value_type == "int":
|
||||
if value is None:
|
||||
return 0
|
||||
if isinstance(value, bool):
|
||||
raise ValueError("bool is not a valid int dimension value")
|
||||
return int(float(value))
|
||||
# float
|
||||
if value is None:
|
||||
return 0.0
|
||||
if isinstance(value, bool):
|
||||
raise ValueError("bool is not a valid float dimension value")
|
||||
return float(value)
|
||||
|
||||
|
||||
def _type_default(value_type: ValueType) -> Any:
|
||||
return "" if value_type == "string" else (0 if value_type == "int" else 0.0)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DimensionCollectInput:
|
||||
request: dict[str, Any] | None = None
|
||||
response: dict[str, Any] | None = None
|
||||
metadata: dict[str, Any] | None = None
|
||||
base_dimensions: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class DimensionCollectorRuntime:
|
||||
"""纯运行时逻辑(不依赖 DB),便于测试与复用。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._evaluator = SafeExpressionEvaluator()
|
||||
|
||||
def collect(
|
||||
self,
|
||||
*,
|
||||
collectors: list[DimensionCollector],
|
||||
inp: DimensionCollectInput,
|
||||
) -> dict[str, Any]:
|
||||
dims: dict[str, Any] = dict(inp.base_dimensions or {})
|
||||
|
||||
# dimension_name -> collectors (priority desc)
|
||||
grouped: dict[str, list[DimensionCollector]] = {}
|
||||
for c in collectors:
|
||||
grouped.setdefault(c.dimension_name, []).append(c)
|
||||
for name in grouped:
|
||||
grouped[name].sort(key=lambda x: (x.priority or 0), reverse=True)
|
||||
|
||||
# 1) 先收集非 computed
|
||||
computed_only: set[str] = set()
|
||||
for dim_name, cs in grouped.items():
|
||||
non_computed = [c for c in cs if (c.source_type or "").lower() != "computed"]
|
||||
if not non_computed:
|
||||
computed_only.add(dim_name)
|
||||
continue
|
||||
value = self._resolve_dimension(dim_name, non_computed, dims, inp)
|
||||
dims[dim_name] = value
|
||||
|
||||
# 2) computed 维度拓扑排序
|
||||
ordered = self._toposort_computed(grouped, computed_only)
|
||||
for dim_name in ordered:
|
||||
cs = [
|
||||
c for c in grouped.get(dim_name, []) if (c.source_type or "").lower() == "computed"
|
||||
]
|
||||
if not cs:
|
||||
continue
|
||||
cs.sort(key=lambda x: (x.priority or 0), reverse=True)
|
||||
value = self._resolve_computed_dimension(dim_name, cs, dims)
|
||||
dims[dim_name] = value
|
||||
|
||||
return dims
|
||||
|
||||
def _resolve_dimension(
|
||||
self,
|
||||
dim_name: str,
|
||||
collectors: list[DimensionCollector],
|
||||
dims: dict[str, Any],
|
||||
inp: DimensionCollectInput,
|
||||
) -> Any:
|
||||
fallback_default: str | None = None
|
||||
fallback_value_type: ValueType | None = None
|
||||
value_type: ValueType = (
|
||||
(collectors[0].value_type or "float").lower() # type: ignore[assignment]
|
||||
if collectors
|
||||
else "float"
|
||||
)
|
||||
|
||||
for c in collectors:
|
||||
value_type = (c.value_type or "float").lower() # type: ignore[assignment]
|
||||
if c.default_value is not None and fallback_default is None:
|
||||
fallback_default = c.default_value
|
||||
fallback_value_type = value_type
|
||||
|
||||
src = (c.source_type or "").lower()
|
||||
path = c.source_path or ""
|
||||
|
||||
if src == "request":
|
||||
raw = _get_nested_value(inp.request or {}, path)
|
||||
elif src == "response":
|
||||
raw = _get_nested_value(inp.response or {}, path)
|
||||
elif src == "metadata":
|
||||
raw = _get_nested_value(inp.metadata or {}, path)
|
||||
else:
|
||||
# 未知 source:跳过尝试
|
||||
continue
|
||||
|
||||
if raw is None:
|
||||
continue
|
||||
|
||||
try:
|
||||
value: Any = raw
|
||||
if c.transform_expression:
|
||||
# transform_expression 仅允许使用 value
|
||||
value = self._evaluator.eval_number(c.transform_expression, {"value": value})
|
||||
casted = _cast_value(value, value_type)
|
||||
return casted
|
||||
except (ValueError, UnsafeExpressionError, ExpressionEvaluationError, Exception) as exc:
|
||||
# 注意:这里选择“不中断,尝试下一优先级”
|
||||
logger.debug(
|
||||
"Dimension collector failed (dim=%s, id=%s): %s",
|
||||
dim_name,
|
||||
getattr(c, "id", None),
|
||||
str(exc),
|
||||
)
|
||||
continue
|
||||
|
||||
# 兜底:default_value(仅允许配置一条,但这里不依赖 DB 校验)
|
||||
if fallback_default is not None:
|
||||
try:
|
||||
return _cast_value(fallback_default, fallback_value_type or value_type)
|
||||
except Exception:
|
||||
return _type_default(fallback_value_type or value_type)
|
||||
|
||||
return _type_default(value_type)
|
||||
|
||||
def _resolve_computed_dimension(
|
||||
self,
|
||||
dim_name: str,
|
||||
collectors: list[DimensionCollector],
|
||||
dims: dict[str, Any],
|
||||
) -> Any:
|
||||
fallback_default: str | None = None
|
||||
fallback_value_type: ValueType | None = None
|
||||
value_type: ValueType = (
|
||||
(collectors[0].value_type or "float").lower() # type: ignore[assignment]
|
||||
if collectors
|
||||
else "float"
|
||||
)
|
||||
|
||||
for c in collectors:
|
||||
value_type = (c.value_type or "float").lower() # type: ignore[assignment]
|
||||
if c.default_value is not None and fallback_default is None:
|
||||
fallback_default = c.default_value
|
||||
fallback_value_type = value_type
|
||||
|
||||
expr = c.transform_expression
|
||||
if not expr:
|
||||
continue
|
||||
|
||||
try:
|
||||
value = self._evaluator.eval_number(expr, dims)
|
||||
casted = _cast_value(value, value_type)
|
||||
return casted
|
||||
except (ValueError, ExpressionEvaluationError, UnsafeExpressionError, Exception):
|
||||
continue
|
||||
|
||||
if fallback_default is not None:
|
||||
try:
|
||||
return _cast_value(fallback_default, fallback_value_type or value_type)
|
||||
except Exception:
|
||||
return _type_default(fallback_value_type or value_type)
|
||||
|
||||
return _type_default(value_type)
|
||||
|
||||
def _toposort_computed(
|
||||
self,
|
||||
grouped: dict[str, list[DimensionCollector]],
|
||||
computed_only: set[str],
|
||||
) -> list[str]:
|
||||
# 建图:dependency -> dim
|
||||
allowed_func_names = set(self._evaluator.ALLOWED_FUNCS.keys())
|
||||
|
||||
deps: dict[str, set[str]] = {d: set() for d in computed_only}
|
||||
for dim_name in computed_only:
|
||||
for c in grouped.get(dim_name, []):
|
||||
if (c.source_type or "").lower() != "computed" or not c.transform_expression:
|
||||
continue
|
||||
try:
|
||||
names = extract_variable_names(c.transform_expression)
|
||||
except UnsafeExpressionError:
|
||||
# 配置错误:按无依赖处理,避免阻塞
|
||||
logger.error(
|
||||
"Invalid computed transform_expression (dim=%s, id=%s)",
|
||||
dim_name,
|
||||
getattr(c, "id", None),
|
||||
)
|
||||
names = set()
|
||||
names.discard("value")
|
||||
names -= allowed_func_names
|
||||
# 仅关心依赖的 computed 维度(非 computed 会在前一步收集)
|
||||
deps[dim_name] |= {n for n in names if n in computed_only and n != dim_name}
|
||||
|
||||
# Kahn
|
||||
in_degree: dict[str, int] = {d: 0 for d in computed_only}
|
||||
forward: dict[str, set[str]] = {d: set() for d in computed_only}
|
||||
for dim_name, dim_deps in deps.items():
|
||||
for dep in dim_deps:
|
||||
forward[dep].add(dim_name)
|
||||
in_degree[dim_name] += 1
|
||||
|
||||
queue = deque(sorted(d for d, deg in in_degree.items() if deg == 0))
|
||||
ordered: list[str] = []
|
||||
|
||||
while queue:
|
||||
node = queue.popleft()
|
||||
ordered.append(node)
|
||||
for nxt in sorted(forward.get(node, set())):
|
||||
in_degree[nxt] -= 1
|
||||
if in_degree[nxt] == 0:
|
||||
queue.append(nxt)
|
||||
|
||||
if len(ordered) != len(computed_only):
|
||||
# 有环依赖:保护性降级(按名称补齐),避免阻塞整条计费链路
|
||||
remaining = sorted(list(computed_only - set(ordered)))
|
||||
logger.error("Computed dimension cycle detected: %s", remaining)
|
||||
ordered.extend(remaining)
|
||||
|
||||
return ordered
|
||||
|
||||
|
||||
class DimensionCollectorService:
|
||||
"""DB + runtime 的封装:读取 collectors 并执行采集。"""
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
self._runtime = DimensionCollectorRuntime()
|
||||
|
||||
def list_enabled_collectors(
|
||||
self,
|
||||
*,
|
||||
api_format: str | None,
|
||||
task_type: str | None,
|
||||
) -> list[DimensionCollector]:
|
||||
api = _normalize_api_format(api_format)
|
||||
task = _normalize_task_type(task_type)
|
||||
api_variants = list({api, api.lower()})
|
||||
|
||||
if task == "cli":
|
||||
# CLI → chat:按维度回退(维度存在 cli collector 则用 cli,否则用 chat)
|
||||
cli_collectors = (
|
||||
self.db.query(DimensionCollector)
|
||||
.filter(
|
||||
DimensionCollector.api_format.in_(api_variants),
|
||||
DimensionCollector.task_type == "cli",
|
||||
DimensionCollector.is_enabled == True, # noqa: E712
|
||||
)
|
||||
.all()
|
||||
)
|
||||
chat_collectors = (
|
||||
self.db.query(DimensionCollector)
|
||||
.filter(
|
||||
DimensionCollector.api_format.in_(api_variants),
|
||||
DimensionCollector.task_type == "chat",
|
||||
DimensionCollector.is_enabled == True, # noqa: E712
|
||||
)
|
||||
.all()
|
||||
)
|
||||
cli_dims: set[str] = {c.dimension_name for c in cli_collectors}
|
||||
result: list[DimensionCollector] = list(cli_collectors)
|
||||
for c in chat_collectors:
|
||||
if c.dimension_name not in cli_dims:
|
||||
result.append(c)
|
||||
return result
|
||||
|
||||
return (
|
||||
self.db.query(DimensionCollector)
|
||||
.filter(
|
||||
DimensionCollector.api_format.in_(api_variants),
|
||||
DimensionCollector.task_type == task,
|
||||
DimensionCollector.is_enabled == True, # noqa: E712
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
def collect_dimensions(
|
||||
self,
|
||||
*,
|
||||
api_format: str | None,
|
||||
task_type: str | None,
|
||||
request: dict[str, Any] | None = None,
|
||||
response: dict[str, Any] | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
base_dimensions: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
collectors = self.list_enabled_collectors(api_format=api_format, task_type=task_type)
|
||||
return self._runtime.collect(
|
||||
collectors=collectors,
|
||||
inp=DimensionCollectInput(
|
||||
request=request,
|
||||
response=response,
|
||||
metadata=metadata,
|
||||
base_dimensions=base_dimensions,
|
||||
),
|
||||
)
|
||||
368
src/services/billing/formula_engine.py
Normal file
368
src/services/billing/formula_engine.py
Normal file
@@ -0,0 +1,368 @@
|
||||
"""
|
||||
FormulaEngine - 配置驱动的安全计费表达式引擎
|
||||
|
||||
目标:
|
||||
- 支持 billing_rules.expression 的安全求值(AST 白名单)
|
||||
- 支持 dimension_mappings(dimension/matrix/tiered/constant)
|
||||
- 支持 required/allow_zero 机制,避免维度缺失导致静默少收
|
||||
|
||||
注意:该模块不直接依赖数据库;规则查找、维度采集在上层服务完成。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable, Literal
|
||||
|
||||
|
||||
class UnsafeExpressionError(ValueError):
|
||||
"""表达式包含不安全/不支持的 AST 结构。"""
|
||||
|
||||
|
||||
class ExpressionEvaluationError(RuntimeError):
|
||||
"""表达式在安全求值阶段失败(如 NameError/ZeroDivision)。"""
|
||||
|
||||
|
||||
class BillingIncompleteError(RuntimeError):
|
||||
"""required 维度缺失且 strict_mode=true 时抛出,用于上层拒绝请求/标记任务失败。"""
|
||||
|
||||
def __init__(self, message: str, *, missing_required: list[str]):
|
||||
super().__init__(message)
|
||||
self.missing_required = missing_required
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FormulaEvaluationResult:
|
||||
status: Literal["complete", "incomplete"]
|
||||
cost: float
|
||||
resolved_values: dict[str, Any]
|
||||
missing_required: list[str]
|
||||
error: str | None = None
|
||||
|
||||
|
||||
_ALLOWED_BINOPS = (
|
||||
ast.Add,
|
||||
ast.Sub,
|
||||
ast.Mult,
|
||||
ast.Div,
|
||||
ast.Pow,
|
||||
ast.FloorDiv,
|
||||
ast.Mod,
|
||||
)
|
||||
_ALLOWED_UNARYOPS = (ast.UAdd, ast.USub)
|
||||
_ALLOWED_OP_NODES = _ALLOWED_BINOPS + _ALLOWED_UNARYOPS
|
||||
|
||||
|
||||
def _iter_ast_nodes(node: ast.AST) -> Iterable[ast.AST]:
|
||||
yield node
|
||||
for child in ast.iter_child_nodes(node):
|
||||
yield from _iter_ast_nodes(child)
|
||||
|
||||
|
||||
def extract_variable_names(expression: str) -> set[str]:
|
||||
"""提取表达式中出现的变量名(不含函数名)。"""
|
||||
try:
|
||||
tree = ast.parse(expression, mode="eval")
|
||||
except SyntaxError as exc:
|
||||
raise UnsafeExpressionError(f"Invalid expression syntax: {exc}") from exc
|
||||
|
||||
names: set[str] = set()
|
||||
for node in _iter_ast_nodes(tree):
|
||||
if isinstance(node, ast.Name):
|
||||
names.add(node.id)
|
||||
if isinstance(node, ast.Call):
|
||||
# Call 的函数名会以 ast.Name 出现,需要从结果中过滤掉
|
||||
if isinstance(node.func, ast.Name):
|
||||
names.discard(node.func.id)
|
||||
return names
|
||||
|
||||
|
||||
class SafeExpressionEvaluator:
|
||||
"""AST 白名单 + 无 builtins 的安全求值器。"""
|
||||
|
||||
ALLOWED_FUNCS: dict[str, Any] = {
|
||||
"min": min,
|
||||
"max": max,
|
||||
"abs": abs,
|
||||
"round": round,
|
||||
"int": int,
|
||||
"float": float,
|
||||
}
|
||||
|
||||
def validate(self, expression: str) -> ast.Expression:
|
||||
try:
|
||||
tree = ast.parse(expression, mode="eval")
|
||||
except SyntaxError as exc:
|
||||
raise UnsafeExpressionError(f"Invalid expression syntax: {exc}") from exc
|
||||
|
||||
for node in _iter_ast_nodes(tree):
|
||||
if isinstance(node, ast.Expression):
|
||||
continue
|
||||
# 运算符节点本身也会出现在 iter_child_nodes 中
|
||||
if isinstance(node, _ALLOWED_OP_NODES):
|
||||
continue
|
||||
if isinstance(node, ast.Constant):
|
||||
# 仅允许数字常量(bool 是 int 子类,需要显式排除)
|
||||
if isinstance(node.value, bool) or not isinstance(node.value, (int, float)):
|
||||
raise UnsafeExpressionError("Only int/float constants are allowed")
|
||||
continue
|
||||
if isinstance(node, ast.BinOp):
|
||||
if not isinstance(node.op, _ALLOWED_BINOPS):
|
||||
raise UnsafeExpressionError(f"Operator not allowed: {type(node.op).__name__}")
|
||||
continue
|
||||
if isinstance(node, ast.UnaryOp):
|
||||
if not isinstance(node.op, _ALLOWED_UNARYOPS):
|
||||
raise UnsafeExpressionError(
|
||||
f"Unary operator not allowed: {type(node.op).__name__}"
|
||||
)
|
||||
continue
|
||||
if isinstance(node, ast.Name):
|
||||
# 防御:拒绝双下划线变量名
|
||||
if node.id.startswith("__"):
|
||||
raise UnsafeExpressionError("Dunder names are not allowed")
|
||||
continue
|
||||
if isinstance(node, ast.Load):
|
||||
continue
|
||||
if isinstance(node, ast.keyword):
|
||||
continue
|
||||
if isinstance(node, ast.Call):
|
||||
if not isinstance(node.func, ast.Name):
|
||||
raise UnsafeExpressionError("Only direct function calls are allowed")
|
||||
func_name = node.func.id
|
||||
if func_name not in self.ALLOWED_FUNCS:
|
||||
raise UnsafeExpressionError(f"Function not allowed: {func_name}")
|
||||
if any(k.arg is None for k in node.keywords):
|
||||
raise UnsafeExpressionError("**kwargs is not allowed")
|
||||
continue
|
||||
|
||||
# 明确禁止的/不需要的节点类型(属性访问、下标、推导式、比较等)
|
||||
if isinstance(
|
||||
node,
|
||||
(
|
||||
ast.Attribute,
|
||||
ast.Subscript,
|
||||
ast.Compare,
|
||||
ast.BoolOp,
|
||||
ast.IfExp,
|
||||
ast.Lambda,
|
||||
ast.Dict,
|
||||
ast.List,
|
||||
ast.Tuple,
|
||||
ast.Set,
|
||||
ast.ListComp,
|
||||
ast.SetComp,
|
||||
ast.DictComp,
|
||||
ast.GeneratorExp,
|
||||
ast.Await,
|
||||
ast.Yield,
|
||||
ast.YieldFrom,
|
||||
),
|
||||
):
|
||||
raise UnsafeExpressionError(f"AST node not allowed: {type(node).__name__}")
|
||||
|
||||
raise UnsafeExpressionError(f"AST node not allowed: {type(node).__name__}")
|
||||
|
||||
assert isinstance(tree, ast.Expression)
|
||||
return tree
|
||||
|
||||
def eval_number(self, expression: str, variables: dict[str, Any]) -> float:
|
||||
tree = self.validate(expression)
|
||||
|
||||
safe_globals = {"__builtins__": {}}
|
||||
safe_locals = dict(self.ALLOWED_FUNCS)
|
||||
safe_locals.update(variables or {})
|
||||
|
||||
try:
|
||||
compiled = compile(tree, "<billing_expr>", "eval")
|
||||
value = eval(compiled, safe_globals, safe_locals) # noqa: S307 - validated AST
|
||||
except Exception as exc:
|
||||
raise ExpressionEvaluationError(str(exc)) from exc
|
||||
|
||||
try:
|
||||
return float(value)
|
||||
except Exception as exc:
|
||||
raise ExpressionEvaluationError(f"Expression result is not numeric: {value!r}") from exc
|
||||
|
||||
|
||||
class FormulaEngine:
|
||||
"""计费表达式引擎:解析 dimension_mappings 并进行安全求值。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._evaluator = SafeExpressionEvaluator()
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
*,
|
||||
expression: str,
|
||||
variables: dict[str, Any] | None,
|
||||
dimensions: dict[str, Any] | None,
|
||||
dimension_mappings: dict[str, dict[str, Any]] | None,
|
||||
strict_mode: bool = False,
|
||||
) -> FormulaEvaluationResult:
|
||||
dims = dimensions or {}
|
||||
mappings = dimension_mappings or {}
|
||||
resolved: dict[str, Any] = dict(variables or {})
|
||||
|
||||
missing_required: list[str] = []
|
||||
|
||||
# 先解析 dimension_mappings,产出 expression 变量表
|
||||
for var_name, mapping in mappings.items():
|
||||
source = (mapping.get("source") or "constant").lower()
|
||||
# 显式 constant 映射属于“兜底行为”:如果 variables 已经提供该变量,则不覆盖。
|
||||
if source == "constant" and var_name in resolved:
|
||||
continue
|
||||
value, is_missing = self._resolve_mapping(var_name, mapping, dims)
|
||||
if is_missing:
|
||||
missing_required.append(var_name)
|
||||
continue
|
||||
resolved[var_name] = value
|
||||
|
||||
# required 维度缺失:直接标记 incomplete(并由 strict_mode 决定是否抛错)
|
||||
if missing_required:
|
||||
if strict_mode:
|
||||
raise BillingIncompleteError(
|
||||
f"Missing required dimensions: {missing_required}",
|
||||
missing_required=missing_required,
|
||||
)
|
||||
return FormulaEvaluationResult(
|
||||
status="incomplete",
|
||||
cost=0.0,
|
||||
resolved_values=resolved,
|
||||
missing_required=missing_required,
|
||||
)
|
||||
|
||||
try:
|
||||
cost = self._evaluator.eval_number(expression, resolved)
|
||||
if cost < 0:
|
||||
# 防御:不允许负数成本(通常表示配置错误)
|
||||
return FormulaEvaluationResult(
|
||||
status="incomplete",
|
||||
cost=0.0,
|
||||
resolved_values=resolved,
|
||||
missing_required=[],
|
||||
error="negative_cost",
|
||||
)
|
||||
return FormulaEvaluationResult(
|
||||
status="complete",
|
||||
cost=cost,
|
||||
resolved_values=resolved,
|
||||
missing_required=[],
|
||||
)
|
||||
except (UnsafeExpressionError, ExpressionEvaluationError) as exc:
|
||||
if strict_mode:
|
||||
raise
|
||||
return FormulaEvaluationResult(
|
||||
status="incomplete",
|
||||
cost=0.0,
|
||||
resolved_values=resolved,
|
||||
missing_required=[],
|
||||
error=str(exc),
|
||||
)
|
||||
|
||||
def _resolve_mapping(
|
||||
self,
|
||||
var_name: str,
|
||||
mapping: dict[str, Any],
|
||||
dims: dict[str, Any],
|
||||
) -> tuple[Any, bool]:
|
||||
"""
|
||||
Returns:
|
||||
(value, is_missing_required)
|
||||
|
||||
说明:
|
||||
- is_missing_required 仅在 required=true 且缺失时为 True
|
||||
- required=false 的缺失会使用 default 或 0 兜底,并返回 is_missing_required=False
|
||||
"""
|
||||
source = (mapping.get("source") or "constant").lower()
|
||||
required = bool(mapping.get("required", False))
|
||||
allow_zero = bool(mapping.get("allow_zero", False))
|
||||
|
||||
default = mapping.get("default", 0)
|
||||
|
||||
def _missing() -> tuple[Any, bool]:
|
||||
if required:
|
||||
return None, True
|
||||
return default, False
|
||||
|
||||
if source == "constant":
|
||||
# constant 默认行为:由 variables 提供;dimension_mappings 显式 constant 时仅做兜底
|
||||
return default, False
|
||||
|
||||
if source == "dimension":
|
||||
key = mapping.get("key") or var_name
|
||||
raw = dims.get(key)
|
||||
if raw is None:
|
||||
return _missing()
|
||||
if isinstance(raw, str):
|
||||
if raw == "":
|
||||
return _missing()
|
||||
# 尝试将字符串解析为数字,否则按字符串返回(供上层自行决定)
|
||||
try:
|
||||
num = float(raw)
|
||||
if num == 0 and not allow_zero:
|
||||
return _missing()
|
||||
return num, False
|
||||
except Exception:
|
||||
return raw, False
|
||||
if isinstance(raw, (int, float)):
|
||||
if float(raw) == 0 and not allow_zero:
|
||||
return _missing()
|
||||
return raw, False
|
||||
# 其他类型:尽量转为 float,否则视为缺失
|
||||
try:
|
||||
num = float(raw)
|
||||
if num == 0 and not allow_zero:
|
||||
return _missing()
|
||||
return num, False
|
||||
except Exception:
|
||||
return _missing()
|
||||
|
||||
if source == "matrix":
|
||||
key = mapping.get("key") or var_name
|
||||
raw = dims.get(key)
|
||||
if raw is None or raw == "":
|
||||
return _missing()
|
||||
raw_key = str(raw)
|
||||
matrix = mapping.get("map") or {}
|
||||
if raw_key in matrix:
|
||||
return matrix[raw_key], False
|
||||
# matrix 未命中:若 required=true 则仍视为缺失;否则使用 default
|
||||
if required:
|
||||
return None, True
|
||||
return default, False
|
||||
|
||||
if source == "tiered":
|
||||
tier_key = mapping.get("tier_key")
|
||||
if not tier_key:
|
||||
return _missing()
|
||||
raw_tier_value = dims.get(tier_key)
|
||||
if raw_tier_value is None:
|
||||
return _missing()
|
||||
try:
|
||||
tier_value = float(raw_tier_value)
|
||||
except Exception:
|
||||
return _missing()
|
||||
|
||||
if tier_value == 0 and not allow_zero:
|
||||
return _missing()
|
||||
|
||||
tiers = mapping.get("tiers") or []
|
||||
# tiers: [{up_to: 128000, value: 2.5}, {up_to: null, value: 1.25}]
|
||||
for tier in tiers:
|
||||
up_to = tier.get("up_to")
|
||||
if up_to is None:
|
||||
return tier.get("value", default), False
|
||||
try:
|
||||
if tier_value <= float(up_to):
|
||||
return tier.get("value", default), False
|
||||
except Exception:
|
||||
# up_to 配置异常:忽略并继续
|
||||
continue
|
||||
# 无匹配:使用最后一个或 default
|
||||
if tiers:
|
||||
return tiers[-1].get("value", default), False
|
||||
return default, False
|
||||
|
||||
# 未知 source:视为配置错误,但不直接中断计费(返回 default)
|
||||
return default, False
|
||||
@@ -9,6 +9,7 @@
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
@@ -89,7 +90,7 @@ class BillingDimension:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@dataclass(init=False)
|
||||
class StandardizedUsage:
|
||||
"""
|
||||
标准化的 Usage 数据
|
||||
@@ -114,8 +115,49 @@ class StandardizedUsage:
|
||||
# 请求计数(用于按次计费)
|
||||
request_count: int = 1
|
||||
|
||||
# 扩展字段(未来可能需要的额外维度)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
# 任意维度存储(用于多维度计费;数值/字符串均可)
|
||||
# 兼容旧字段名:extra 作为 dimensions 的别名
|
||||
dimensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
reasoning_tokens: int = 0,
|
||||
cache_storage_token_hours: float = 0.0,
|
||||
request_count: int = 1,
|
||||
dimensions: dict[str, Any] | None = None,
|
||||
extra: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
# 基础字段
|
||||
self.input_tokens = input_tokens
|
||||
self.output_tokens = output_tokens
|
||||
self.cache_creation_tokens = cache_creation_tokens
|
||||
self.cache_read_tokens = cache_read_tokens
|
||||
self.reasoning_tokens = reasoning_tokens
|
||||
self.cache_storage_token_hours = cache_storage_token_hours
|
||||
self.request_count = request_count
|
||||
|
||||
# 兼容:支持 extra 与 dimensions 同时传入(dimensions 优先级更高)
|
||||
merged: dict[str, Any] = {}
|
||||
if isinstance(extra, dict):
|
||||
merged.update(extra)
|
||||
if isinstance(dimensions, dict):
|
||||
merged.update(dimensions)
|
||||
self.dimensions = merged
|
||||
|
||||
@property
|
||||
def extra(self) -> dict[str, Any]:
|
||||
"""向后兼容:旧代码使用 usage.extra 访问扩展维度。"""
|
||||
return self.dimensions
|
||||
|
||||
@extra.setter
|
||||
def extra(self, value: dict[str, Any]) -> None:
|
||||
"""向后兼容:允许旧代码写入 usage.extra。"""
|
||||
self.dimensions = value or {}
|
||||
|
||||
def get(self, field_name: str, default: Any = 0) -> Any:
|
||||
"""
|
||||
@@ -130,12 +172,14 @@ class StandardizedUsage:
|
||||
Returns:
|
||||
字段值
|
||||
"""
|
||||
if hasattr(self, field_name):
|
||||
value = getattr(self, field_name)
|
||||
# 对于 extra 字段,不直接返回
|
||||
if field_name != "extra":
|
||||
return value
|
||||
return self.extra.get(field_name, default)
|
||||
# 兼容旧字段名
|
||||
if field_name == "extra":
|
||||
return self.dimensions
|
||||
|
||||
if hasattr(self, field_name) and field_name not in {"dimensions"}:
|
||||
return getattr(self, field_name)
|
||||
|
||||
return self.dimensions.get(field_name, default)
|
||||
|
||||
def set(self, field_name: str, value: Any) -> None:
|
||||
"""
|
||||
@@ -145,10 +189,16 @@ class StandardizedUsage:
|
||||
field_name: 字段名
|
||||
value: 字段值
|
||||
"""
|
||||
if hasattr(self, field_name) and field_name != "extra":
|
||||
# 兼容旧字段名
|
||||
if field_name == "extra":
|
||||
self.dimensions = value or {}
|
||||
return
|
||||
|
||||
if hasattr(self, field_name) and field_name not in {"dimensions"}:
|
||||
setattr(self, field_name, value)
|
||||
else:
|
||||
self.extra[field_name] = value
|
||||
return
|
||||
|
||||
self.dimensions[field_name] = value
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""转换为字典"""
|
||||
@@ -161,14 +211,24 @@ class StandardizedUsage:
|
||||
"cache_storage_token_hours": self.cache_storage_token_hours,
|
||||
"request_count": self.request_count,
|
||||
}
|
||||
if self.extra:
|
||||
result["extra"] = self.extra
|
||||
if self.dimensions:
|
||||
# 新字段名
|
||||
result["dimensions"] = self.dimensions
|
||||
# 旧字段名(兼容)
|
||||
result["extra"] = self.dimensions
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> StandardizedUsage:
|
||||
"""从字典创建实例"""
|
||||
# 兼容:支持 extra / dimensions 两种键名
|
||||
extra = data.pop("extra", {}) if "extra" in data else {}
|
||||
dimensions = data.pop("dimensions", {}) if "dimensions" in data else {}
|
||||
merged_dimensions: dict[str, Any] = {}
|
||||
if isinstance(extra, dict):
|
||||
merged_dimensions.update(extra)
|
||||
if isinstance(dimensions, dict):
|
||||
merged_dimensions.update(dimensions)
|
||||
# 只取已知字段
|
||||
known_fields = {
|
||||
"input_tokens",
|
||||
@@ -180,7 +240,7 @@ class StandardizedUsage:
|
||||
"request_count",
|
||||
}
|
||||
filtered = {k: v for k, v in data.items() if k in known_fields}
|
||||
return cls(**filtered, extra=extra)
|
||||
return cls(**filtered, dimensions=merged_dimensions)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
101
src/services/billing/rule_service.py
Normal file
101
src/services/billing/rule_service.py
Normal file
@@ -0,0 +1,101 @@
|
||||
"""
|
||||
BillingRule 查找逻辑
|
||||
|
||||
查找顺序(与 .plans/humming-seeking-marble.md 一致):
|
||||
1) Model(Provider 级)→ 2) GlobalModel(默认)
|
||||
|
||||
注意:
|
||||
- CLI 在计费域等同于 chat:billing_rules.task_type 不含 "cli"
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import BillingRule, GlobalModel, Model
|
||||
|
||||
TaskType = Literal["chat", "cli", "video", "image", "audio"]
|
||||
|
||||
|
||||
def effective_rule_task_type(task_type: str) -> str:
|
||||
"""CLI 在计费规则域里恒等于 chat。"""
|
||||
t = (task_type or "").lower()
|
||||
return "chat" if t == "cli" else t
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BillingRuleLookupResult:
|
||||
rule: BillingRule
|
||||
scope: Literal["model", "global"]
|
||||
effective_task_type: str
|
||||
|
||||
|
||||
class BillingRuleService:
|
||||
@staticmethod
|
||||
def find_rule(
|
||||
db: Session,
|
||||
*,
|
||||
provider_id: str | None,
|
||||
model_name: str,
|
||||
task_type: str,
|
||||
) -> BillingRuleLookupResult | None:
|
||||
effective_task = effective_rule_task_type(task_type)
|
||||
|
||||
global_model = (
|
||||
db.query(GlobalModel)
|
||||
.filter(
|
||||
GlobalModel.name == model_name,
|
||||
GlobalModel.is_active == True, # noqa: E712
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not global_model:
|
||||
return None
|
||||
|
||||
# 1) Provider Model 覆盖
|
||||
if provider_id:
|
||||
model_obj = (
|
||||
db.query(Model)
|
||||
.filter(
|
||||
Model.provider_id == provider_id,
|
||||
Model.global_model_id == global_model.id,
|
||||
Model.is_active == True, # noqa: E712
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if model_obj:
|
||||
rule = (
|
||||
db.query(BillingRule)
|
||||
.filter(
|
||||
BillingRule.model_id == model_obj.id,
|
||||
BillingRule.task_type == effective_task,
|
||||
BillingRule.is_enabled == True, # noqa: E712
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if rule:
|
||||
return BillingRuleLookupResult(
|
||||
rule=rule,
|
||||
scope="model",
|
||||
effective_task_type=effective_task,
|
||||
)
|
||||
|
||||
# 2) GlobalModel 默认规则
|
||||
rule = (
|
||||
db.query(BillingRule)
|
||||
.filter(
|
||||
BillingRule.global_model_id == global_model.id,
|
||||
BillingRule.task_type == effective_task,
|
||||
BillingRule.is_enabled == True, # noqa: E712
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if rule:
|
||||
return BillingRuleLookupResult(
|
||||
rule=rule, scope="global", effective_task_type=effective_task
|
||||
)
|
||||
|
||||
return None
|
||||
@@ -9,7 +9,6 @@
|
||||
- PER_REQUEST: 按次计费
|
||||
"""
|
||||
|
||||
|
||||
from src.services.billing.models import BillingDimension, BillingUnit
|
||||
|
||||
|
||||
|
||||
@@ -8,9 +8,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Callable
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import Any, Callable
|
||||
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from apscheduler.triggers.cron import CronTrigger
|
||||
@@ -177,6 +177,19 @@ class TaskScheduler:
|
||||
f"({hours}小时{minutes}分钟后)"
|
||||
)
|
||||
|
||||
def remove_job(self, job_id: str) -> None:
|
||||
"""
|
||||
移除指定的定时任务
|
||||
|
||||
Args:
|
||||
job_id: 任务ID
|
||||
"""
|
||||
try:
|
||||
self.scheduler.remove_job(job_id)
|
||||
logger.info(f"已移除定时任务: {job_id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"移除定时任务失败 {job_id}: {e}")
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
"""调度器是否在运行"""
|
||||
|
||||
@@ -30,6 +30,7 @@ from src.services.system.config import SystemConfigService
|
||||
@dataclass
|
||||
class UsageRecordParams:
|
||||
"""用量记录参数数据类,用于在内部方法间传递数据"""
|
||||
|
||||
db: Session
|
||||
user: User | None
|
||||
api_key: ApiKey | None
|
||||
@@ -76,9 +77,7 @@ class UsageRecordParams:
|
||||
f"cache_creation_input_tokens 不能为负数: {self.cache_creation_input_tokens}"
|
||||
)
|
||||
if self.cache_read_input_tokens < 0:
|
||||
raise ValueError(
|
||||
f"cache_read_input_tokens 不能为负数: {self.cache_read_input_tokens}"
|
||||
)
|
||||
raise ValueError(f"cache_read_input_tokens 不能为负数: {self.cache_read_input_tokens}")
|
||||
|
||||
# 响应时间不能为负数
|
||||
if self.response_time_ms is not None and self.response_time_ms < 0:
|
||||
@@ -170,9 +169,10 @@ class UsageService:
|
||||
Returns:
|
||||
热力图数据字典
|
||||
"""
|
||||
import json
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.config.constants import CacheTTL
|
||||
import json
|
||||
|
||||
cache_key = cls._get_heatmap_cache_key(user_id, include_actual_cost)
|
||||
|
||||
@@ -261,8 +261,8 @@ class UsageService:
|
||||
request_cost: float,
|
||||
total_cost: float,
|
||||
# 价格信息
|
||||
input_price: float,
|
||||
output_price: float,
|
||||
input_price: float | None,
|
||||
output_price: float | None,
|
||||
cache_creation_price: float | None,
|
||||
cache_read_price: float | None,
|
||||
request_price: float | None,
|
||||
@@ -415,8 +415,21 @@ class UsageService:
|
||||
cache_ttl_minutes: int | None,
|
||||
use_tiered_pricing: bool,
|
||||
is_failed_request: bool,
|
||||
) -> tuple[float, float, float, float, float, float, float, float, float,
|
||||
float | None, float | None, float | None, int | None]:
|
||||
) -> tuple[
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float,
|
||||
float | None,
|
||||
float | None,
|
||||
float | None,
|
||||
int | None,
|
||||
]:
|
||||
"""计算所有成本相关数据
|
||||
|
||||
Returns:
|
||||
@@ -538,9 +551,19 @@ class UsageService:
|
||||
)
|
||||
|
||||
return (
|
||||
input_price, output_price, cache_creation_price, cache_read_price, request_price,
|
||||
input_cost, output_cost, cache_creation_cost, cache_read_cost, cache_cost,
|
||||
request_cost, total_cost, tier_index
|
||||
input_price,
|
||||
output_price,
|
||||
cache_creation_price,
|
||||
cache_read_price,
|
||||
request_price,
|
||||
input_cost,
|
||||
output_cost,
|
||||
cache_creation_cost,
|
||||
cache_read_cost,
|
||||
cache_cost,
|
||||
request_cost,
|
||||
total_cost,
|
||||
tier_index,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -584,7 +607,9 @@ class UsageService:
|
||||
existing_usage.total_cost_usd = usage_params["total_cost_usd"]
|
||||
existing_usage.actual_input_cost_usd = usage_params["actual_input_cost_usd"]
|
||||
existing_usage.actual_output_cost_usd = usage_params["actual_output_cost_usd"]
|
||||
existing_usage.actual_cache_creation_cost_usd = usage_params["actual_cache_creation_cost_usd"]
|
||||
existing_usage.actual_cache_creation_cost_usd = usage_params[
|
||||
"actual_cache_creation_cost_usd"
|
||||
]
|
||||
existing_usage.actual_cache_read_cost_usd = usage_params["actual_cache_read_cost_usd"]
|
||||
existing_usage.actual_request_cost_usd = usage_params["actual_request_cost_usd"]
|
||||
existing_usage.actual_total_cost_usd = usage_params["actual_total_cost_usd"]
|
||||
@@ -646,9 +671,7 @@ class UsageService:
|
||||
return service.get_cache_prices(provider, model, input_price)
|
||||
|
||||
@classmethod
|
||||
async def get_request_price_async(
|
||||
cls, db: Session, provider: str, model: str
|
||||
) -> float | None:
|
||||
async def get_request_price_async(cls, db: Session, provider: str, model: str) -> float | None:
|
||||
"""异步获取模型按次计费价格"""
|
||||
service = ModelCostService(db)
|
||||
return await service.get_request_price_async(provider, model)
|
||||
@@ -748,9 +771,19 @@ class UsageService:
|
||||
# 计算成本
|
||||
is_failed_request = params.status_code >= 400 or params.error_message is not None
|
||||
(
|
||||
input_price, output_price, cache_creation_price, cache_read_price, request_price,
|
||||
input_cost, output_cost, cache_creation_cost, cache_read_cost, cache_cost,
|
||||
request_cost, total_cost, _tier_index
|
||||
input_price,
|
||||
output_price,
|
||||
cache_creation_price,
|
||||
cache_read_price,
|
||||
request_price,
|
||||
input_cost,
|
||||
output_cost,
|
||||
cache_creation_cost,
|
||||
cache_read_cost,
|
||||
cache_cost,
|
||||
request_cost,
|
||||
total_cost,
|
||||
_tier_index,
|
||||
) = await cls._calculate_costs(
|
||||
db=params.db,
|
||||
provider=params.provider,
|
||||
@@ -834,7 +867,9 @@ class UsageService:
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
async def prepare_single(params: UsageRecordParams) -> tuple[dict[str, Any], float, Exception | None]:
|
||||
async def prepare_single(
|
||||
params: UsageRecordParams,
|
||||
) -> tuple[dict[str, Any], float, Exception | None]:
|
||||
try:
|
||||
usage_params, total_cost = await cls._prepare_usage_record(params)
|
||||
return (usage_params, total_cost, None)
|
||||
@@ -904,23 +939,38 @@ class UsageService:
|
||||
|
||||
# 使用共享逻辑准备记录参数
|
||||
params = UsageRecordParams(
|
||||
db=db, user=user, api_key=api_key, provider=provider, model=model,
|
||||
input_tokens=input_tokens, output_tokens=output_tokens,
|
||||
db=db,
|
||||
user=user,
|
||||
api_key=api_key,
|
||||
provider=provider,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_input_tokens,
|
||||
cache_read_input_tokens=cache_read_input_tokens,
|
||||
request_type=request_type, api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format, has_format_conversion=has_format_conversion,
|
||||
request_type=request_type,
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms, first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code, error_message=error_message, metadata=metadata,
|
||||
request_headers=request_headers, request_body=request_body,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=error_message,
|
||||
metadata=metadata,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers,
|
||||
response_headers=response_headers, client_response_headers=client_response_headers,
|
||||
response_headers=response_headers,
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body,
|
||||
request_id=request_id, provider_id=provider_id,
|
||||
request_id=request_id,
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=provider_endpoint_id,
|
||||
provider_api_key_id=provider_api_key_id, status=status,
|
||||
cache_ttl_minutes=cache_ttl_minutes, use_tiered_pricing=use_tiered_pricing,
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
status=status,
|
||||
cache_ttl_minutes=cache_ttl_minutes,
|
||||
use_tiered_pricing=use_tiered_pricing,
|
||||
target_model=target_model,
|
||||
)
|
||||
usage_params, _ = await cls._prepare_usage_record(params)
|
||||
@@ -931,6 +981,7 @@ class UsageService:
|
||||
|
||||
# 更新 GlobalModel 使用计数(原子操作)
|
||||
from sqlalchemy import update
|
||||
|
||||
from src.models.database import GlobalModel
|
||||
|
||||
db.execute(
|
||||
@@ -1003,23 +1054,38 @@ class UsageService:
|
||||
|
||||
# 使用共享逻辑准备记录参数
|
||||
params = UsageRecordParams(
|
||||
db=db, user=user, api_key=api_key, provider=provider, model=model,
|
||||
input_tokens=input_tokens, output_tokens=output_tokens,
|
||||
db=db,
|
||||
user=user,
|
||||
api_key=api_key,
|
||||
provider=provider,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_input_tokens,
|
||||
cache_read_input_tokens=cache_read_input_tokens,
|
||||
request_type=request_type, api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format, has_format_conversion=has_format_conversion,
|
||||
request_type=request_type,
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms, first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code, error_message=error_message, metadata=metadata,
|
||||
request_headers=request_headers, request_body=request_body,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=error_message,
|
||||
metadata=metadata,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers,
|
||||
response_headers=response_headers, client_response_headers=client_response_headers,
|
||||
response_headers=response_headers,
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body,
|
||||
request_id=request_id, provider_id=provider_id,
|
||||
request_id=request_id,
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=provider_endpoint_id,
|
||||
provider_api_key_id=provider_api_key_id, status=status,
|
||||
cache_ttl_minutes=cache_ttl_minutes, use_tiered_pricing=use_tiered_pricing,
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
status=status,
|
||||
cache_ttl_minutes=cache_ttl_minutes,
|
||||
use_tiered_pricing=use_tiered_pricing,
|
||||
target_model=target_model,
|
||||
)
|
||||
usage_params, total_cost = await cls._prepare_usage_record(params)
|
||||
@@ -1044,8 +1110,12 @@ class UsageService:
|
||||
api_key = db.merge(api_key)
|
||||
|
||||
# 使用原子更新避免并发竞态条件
|
||||
from sqlalchemy import func as sql_func, update
|
||||
from src.models.database import ApiKey as ApiKeyModel, User as UserModel, GlobalModel
|
||||
from sqlalchemy import func as sql_func
|
||||
from sqlalchemy import update
|
||||
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
from src.models.database import GlobalModel
|
||||
from src.models.database import User as UserModel
|
||||
|
||||
# 更新用户使用量(独立 Key 不计入创建者的使用记录)
|
||||
if user and not (api_key and api_key.is_standalone):
|
||||
@@ -1111,6 +1181,227 @@ class UsageService:
|
||||
|
||||
return usage
|
||||
|
||||
@classmethod
|
||||
async def record_usage_with_custom_cost(
|
||||
cls,
|
||||
*,
|
||||
db: Session,
|
||||
user: User | None,
|
||||
api_key: ApiKey | None,
|
||||
provider: str,
|
||||
model: str,
|
||||
request_type: str,
|
||||
total_cost_usd: float,
|
||||
request_cost_usd: float | None = None,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_input_tokens: int = 0,
|
||||
cache_read_input_tokens: int = 0,
|
||||
api_format: str | None = None,
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
is_stream: bool = False,
|
||||
response_time_ms: int | None = None,
|
||||
first_byte_time_ms: int | None = None,
|
||||
status_code: int = 200,
|
||||
error_message: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
request_headers: dict[str, Any] | None = None,
|
||||
request_body: Any | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
response_body: Any | None = None,
|
||||
request_id: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
provider_endpoint_id: str | None = None,
|
||||
provider_api_key_id: str | None = None,
|
||||
status: str = "completed",
|
||||
target_model: str | None = None,
|
||||
) -> Usage:
|
||||
"""
|
||||
记录“已计算好的”成本(用于 Video/Image/Audio 等异步任务的 FormulaEngine 计费结果)。
|
||||
|
||||
说明:
|
||||
- 仍然会应用 ProviderAPIKey.rate_multipliers 计算 actual_* 成本
|
||||
- 会更新 User/APIKey/GlobalModel/Provider 的统计(与 record_usage 行为一致)
|
||||
- 若 request_id 已存在则更新记录(避免重复写入)
|
||||
"""
|
||||
# 生成 request_id
|
||||
if request_id is None:
|
||||
request_id = str(uuid.uuid4())[:8]
|
||||
|
||||
# 获取费率倍数与免费套餐
|
||||
actual_rate_multiplier, is_free_tier = await cls._get_rate_multiplier_and_free_tier(
|
||||
db, provider_api_key_id, provider_id, api_format
|
||||
)
|
||||
|
||||
# 成本拆分:非 token 计费默认计入 request_cost
|
||||
input_cost = 0.0
|
||||
output_cost = 0.0
|
||||
cache_creation_cost = 0.0
|
||||
cache_read_cost = 0.0
|
||||
cache_cost = 0.0
|
||||
request_cost = (
|
||||
float(request_cost_usd) if request_cost_usd is not None else float(total_cost_usd)
|
||||
)
|
||||
total_cost = float(total_cost_usd)
|
||||
|
||||
usage_params = cls._build_usage_params(
|
||||
db=db,
|
||||
user=user,
|
||||
api_key=api_key,
|
||||
provider=provider,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_input_tokens,
|
||||
cache_read_input_tokens=cache_read_input_tokens,
|
||||
request_type=request_type,
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=error_message,
|
||||
metadata=metadata,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers,
|
||||
response_headers=response_headers,
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body,
|
||||
request_id=request_id,
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=provider_endpoint_id,
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
status=status,
|
||||
target_model=target_model,
|
||||
input_cost=input_cost,
|
||||
output_cost=output_cost,
|
||||
cache_creation_cost=cache_creation_cost,
|
||||
cache_read_cost=cache_read_cost,
|
||||
cache_cost=cache_cost,
|
||||
request_cost=request_cost,
|
||||
total_cost=total_cost,
|
||||
# token 价格对异步任务不适用,保持 None
|
||||
input_price=None,
|
||||
output_price=None,
|
||||
cache_creation_price=None,
|
||||
cache_read_price=None,
|
||||
request_price=None,
|
||||
actual_rate_multiplier=actual_rate_multiplier,
|
||||
is_free_tier=is_free_tier,
|
||||
)
|
||||
|
||||
# Upsert(与 record_usage 保持一致)
|
||||
existing_usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if existing_usage:
|
||||
# 避免重复记账:若已是终态记录,直接返回(批量接口也采用该策略)
|
||||
if existing_usage.status not in ("pending", "streaming"):
|
||||
logger.debug(
|
||||
"record_usage_with_custom_cost: request_id=%s already finalized (status=%s), skip",
|
||||
request_id,
|
||||
existing_usage.status,
|
||||
)
|
||||
return existing_usage
|
||||
cls._update_existing_usage(existing_usage, usage_params, target_model)
|
||||
usage = existing_usage
|
||||
else:
|
||||
usage = Usage(**usage_params)
|
||||
db.add(usage)
|
||||
|
||||
# 确保 user 和 api_key 在会话中(与 record_usage 保持一致)
|
||||
if user and not db.object_session(user):
|
||||
user = db.merge(user)
|
||||
if api_key and not db.object_session(api_key):
|
||||
api_key = db.merge(api_key)
|
||||
|
||||
# 原子更新统计
|
||||
from sqlalchemy import func as sql_func
|
||||
from sqlalchemy import update
|
||||
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
from src.models.database import GlobalModel
|
||||
from src.models.database import User as UserModel
|
||||
|
||||
# 更新用户使用量(独立 Key 不计入创建者)
|
||||
if user and not (api_key and api_key.is_standalone):
|
||||
db.execute(
|
||||
update(UserModel)
|
||||
.where(UserModel.id == user.id)
|
||||
.values(
|
||||
used_usd=UserModel.used_usd + total_cost,
|
||||
total_usd=UserModel.total_usd + total_cost,
|
||||
updated_at=sql_func.now(),
|
||||
)
|
||||
)
|
||||
|
||||
# 更新 API 密钥使用量
|
||||
if api_key:
|
||||
if api_key.is_standalone:
|
||||
db.execute(
|
||||
update(ApiKeyModel)
|
||||
.where(ApiKeyModel.id == api_key.id)
|
||||
.values(
|
||||
total_requests=ApiKeyModel.total_requests + 1,
|
||||
total_cost_usd=ApiKeyModel.total_cost_usd + total_cost,
|
||||
balance_used_usd=ApiKeyModel.balance_used_usd + total_cost,
|
||||
last_used_at=sql_func.now(),
|
||||
updated_at=sql_func.now(),
|
||||
)
|
||||
)
|
||||
else:
|
||||
db.execute(
|
||||
update(ApiKeyModel)
|
||||
.where(ApiKeyModel.id == api_key.id)
|
||||
.values(
|
||||
total_requests=ApiKeyModel.total_requests + 1,
|
||||
total_cost_usd=ApiKeyModel.total_cost_usd + total_cost,
|
||||
last_used_at=sql_func.now(),
|
||||
updated_at=sql_func.now(),
|
||||
)
|
||||
)
|
||||
|
||||
# 更新 GlobalModel 使用计数
|
||||
db.execute(
|
||||
update(GlobalModel)
|
||||
.where(GlobalModel.name == model)
|
||||
.values(usage_count=GlobalModel.usage_count + 1)
|
||||
)
|
||||
|
||||
# 更新 Provider 月度使用量(使用 actual_total_cost)
|
||||
if provider_id:
|
||||
actual_total_cost = usage_params["actual_total_cost_usd"]
|
||||
db.execute(
|
||||
update(Provider)
|
||||
.where(Provider.id == provider_id)
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||
)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
# 并发场景可能触发唯一约束冲突:降级为读取已存在记录
|
||||
try:
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
if isinstance(e, IntegrityError):
|
||||
db.rollback()
|
||||
existing = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if existing:
|
||||
return existing
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.error(f"提交使用记录时出错: {e}")
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
return usage
|
||||
|
||||
@classmethod
|
||||
async def record_usage_batch(
|
||||
cls,
|
||||
@@ -1136,8 +1427,12 @@ class UsageService:
|
||||
return []
|
||||
|
||||
from collections import defaultdict
|
||||
|
||||
from sqlalchemy import update
|
||||
from src.models.database import ApiKey as ApiKeyModel, User as UserModel, GlobalModel
|
||||
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
from src.models.database import GlobalModel
|
||||
from src.models.database import User as UserModel
|
||||
|
||||
# 分离需要更新和需要新建的记录
|
||||
request_ids = [r.get("request_id") for r in records if r.get("request_id")]
|
||||
@@ -1147,11 +1442,7 @@ class UsageService:
|
||||
|
||||
if request_ids:
|
||||
# 查询已存在的 Usage 记录(包括 pending/streaming 状态)
|
||||
existing_records = (
|
||||
db.query(Usage)
|
||||
.filter(Usage.request_id.in_(request_ids))
|
||||
.all()
|
||||
)
|
||||
existing_records = db.query(Usage).filter(Usage.request_id.in_(request_ids)).all()
|
||||
existing_usages = {u.request_id: u for u in existing_records}
|
||||
|
||||
for record in records:
|
||||
@@ -1280,8 +1571,8 @@ class UsageService:
|
||||
prepared_results = []
|
||||
|
||||
# 分配准备结果
|
||||
update_results = prepared_results[:len(update_params_list)]
|
||||
insert_results = prepared_results[len(update_params_list):]
|
||||
update_results = prepared_results[: len(update_params_list)]
|
||||
insert_results = prepared_results[len(update_params_list) :]
|
||||
|
||||
# 1. 处理需要更新的记录
|
||||
for i, (record, request_id, params) in enumerate(update_params_list):
|
||||
@@ -1369,12 +1660,12 @@ class UsageService:
|
||||
if skip_ratio > 0.1:
|
||||
logger.error(
|
||||
"批量记录失败率过高: %d/%d (%.1f%%) 条记录被跳过",
|
||||
skipped_count, total_count, skip_ratio * 100
|
||||
skipped_count,
|
||||
total_count,
|
||||
skip_ratio * 100,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"批量记录部分失败: %d/%d 条记录被跳过", skipped_count, total_count
|
||||
)
|
||||
logger.warning("批量记录部分失败: %d/%d 条记录被跳过", skipped_count, total_count)
|
||||
|
||||
# 批量更新 GlobalModel 使用计数
|
||||
for model_name, count in model_counts.items():
|
||||
@@ -1395,6 +1686,7 @@ class UsageService:
|
||||
|
||||
# 批量更新用户使用量
|
||||
from sqlalchemy import func as sql_func
|
||||
|
||||
for user_id, cost in user_costs.items():
|
||||
if cost > 0:
|
||||
db.execute(
|
||||
@@ -1438,9 +1730,7 @@ class UsageService:
|
||||
db.commit()
|
||||
inserted_count = len(usages) - updated_count
|
||||
if updated_count > 0:
|
||||
logger.debug(
|
||||
f"批量记录成功: 更新 {updated_count} 条, 新建 {inserted_count} 条"
|
||||
)
|
||||
logger.debug(f"批量记录成功: 更新 {updated_count} 条, 新建 {inserted_count} 条")
|
||||
else:
|
||||
logger.debug(f"批量记录 {len(usages)} 条使用记录成功")
|
||||
except Exception as e:
|
||||
@@ -1832,10 +2122,7 @@ class UsageService:
|
||||
while True:
|
||||
# 查询待删除的 ID(使用新索引 idx_usage_user_created)
|
||||
batch_ids = (
|
||||
db.query(Usage.id)
|
||||
.filter(Usage.created_at < cutoff_date)
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
db.query(Usage.id).filter(Usage.created_at < cutoff_date).limit(batch_size).all()
|
||||
)
|
||||
|
||||
if not batch_ids:
|
||||
@@ -2112,7 +2399,9 @@ class UsageService:
|
||||
|
||||
if count > 0:
|
||||
db.commit()
|
||||
logger.info(f"清理超时请求: 将 {count} 条超过 {timeout_minutes} 分钟的 pending/streaming 请求标记为 failed")
|
||||
logger.info(
|
||||
f"清理超时请求: 将 {count} 条超过 {timeout_minutes} 分钟的 pending/streaming 请求标记为 failed"
|
||||
)
|
||||
|
||||
return count
|
||||
|
||||
@@ -2240,9 +2529,7 @@ class UsageService:
|
||||
# 如果流已经成功完成(stream_completed: true),不应该标记为超时
|
||||
# 先获取这些 Usage 的 request_id
|
||||
usage_request_ids = (
|
||||
db.query(Usage.id, Usage.request_id)
|
||||
.filter(Usage.id.in_(timeout_candidates))
|
||||
.all()
|
||||
db.query(Usage.id, Usage.request_id).filter(Usage.id.in_(timeout_candidates)).all()
|
||||
)
|
||||
usage_id_to_request_id = {u.id: u.request_id for u in usage_request_ids}
|
||||
request_id_to_usage_id = {u.request_id: u.id for u in usage_request_ids}
|
||||
@@ -2278,9 +2565,7 @@ class UsageService:
|
||||
for candidate in candidates:
|
||||
extra_data = candidate.extra_data or {}
|
||||
# 情况1:status='success' 且 stream_completed=True
|
||||
if candidate.status == "success" and extra_data.get(
|
||||
"stream_completed", False
|
||||
):
|
||||
if candidate.status == "success" and extra_data.get("stream_completed", False):
|
||||
usage_id = request_id_to_usage_id.get(candidate.request_id)
|
||||
if usage_id:
|
||||
completed_usage_ids.add(usage_id)
|
||||
@@ -2321,11 +2606,7 @@ class UsageService:
|
||||
has_format_conversion = getattr(r, "has_format_conversion", None)
|
||||
|
||||
# 兼容历史数据:当 streaming 状态已拿到两个格式但 has_format_conversion 为空时,回填推断结果
|
||||
if (
|
||||
has_format_conversion is None
|
||||
and api_format
|
||||
and endpoint_api_format
|
||||
):
|
||||
if has_format_conversion is None and api_format and endpoint_api_format:
|
||||
has_format_conversion = not can_passthrough(api_format, endpoint_api_format)
|
||||
|
||||
item: dict[str, Any] = {
|
||||
@@ -2336,8 +2617,12 @@ class UsageService:
|
||||
"cache_creation_input_tokens": r.cache_creation_input_tokens,
|
||||
"cache_read_input_tokens": r.cache_read_input_tokens,
|
||||
"cost": float(r.total_cost_usd) if r.total_cost_usd else 0,
|
||||
"actual_cost": float(r.actual_total_cost_usd) if r.actual_total_cost_usd is not None else None,
|
||||
"rate_multiplier": float(r.rate_multiplier) if r.rate_multiplier is not None else None,
|
||||
"actual_cost": (
|
||||
float(r.actual_total_cost_usd) if r.actual_total_cost_usd is not None else None
|
||||
),
|
||||
"rate_multiplier": (
|
||||
float(r.rate_multiplier) if r.rate_multiplier is not None else None
|
||||
),
|
||||
"response_time_ms": r.response_time_ms,
|
||||
"first_byte_time_ms": r.first_byte_time_ms, # 首字时间 (TTFB)
|
||||
}
|
||||
@@ -2484,47 +2769,47 @@ class UsageService:
|
||||
) = row
|
||||
|
||||
# 计算推荐 TTL
|
||||
recommended_ttl = UsageService._calculate_recommended_ttl(
|
||||
p75_interval, p90_interval
|
||||
)
|
||||
recommended_ttl = UsageService._calculate_recommended_ttl(p75_interval, p90_interval)
|
||||
|
||||
# 获取用户信息
|
||||
user_info = user_info_map.get(str(group_id), {})
|
||||
|
||||
# 计算各区间占比
|
||||
total_intervals = request_count
|
||||
users_analysis.append({
|
||||
"group_id": group_id,
|
||||
"username": user_info.get("username"),
|
||||
"email": user_info.get("email"),
|
||||
"request_count": request_count,
|
||||
"interval_distribution": {
|
||||
"within_5min": within_5min,
|
||||
"within_15min": within_15min,
|
||||
"within_30min": within_30min,
|
||||
"within_60min": within_60min,
|
||||
"over_60min": over_60min,
|
||||
},
|
||||
"interval_percentages": {
|
||||
"within_5min": round(within_5min / total_intervals * 100, 1),
|
||||
"within_15min": round(within_15min / total_intervals * 100, 1),
|
||||
"within_30min": round(within_30min / total_intervals * 100, 1),
|
||||
"within_60min": round(within_60min / total_intervals * 100, 1),
|
||||
"over_60min": round(over_60min / total_intervals * 100, 1),
|
||||
},
|
||||
"percentiles": {
|
||||
"p50": round(float(median_interval), 2) if median_interval else None,
|
||||
"p75": round(float(p75_interval), 2) if p75_interval else None,
|
||||
"p90": round(float(p90_interval), 2) if p90_interval else None,
|
||||
},
|
||||
"avg_interval_minutes": round(float(avg_interval), 2) if avg_interval else None,
|
||||
"min_interval_minutes": round(float(min_interval), 2) if min_interval else None,
|
||||
"max_interval_minutes": round(float(max_interval), 2) if max_interval else None,
|
||||
"recommended_ttl_minutes": recommended_ttl,
|
||||
"recommendation_reason": UsageService._get_ttl_recommendation_reason(
|
||||
recommended_ttl, p75_interval, p90_interval
|
||||
),
|
||||
})
|
||||
users_analysis.append(
|
||||
{
|
||||
"group_id": group_id,
|
||||
"username": user_info.get("username"),
|
||||
"email": user_info.get("email"),
|
||||
"request_count": request_count,
|
||||
"interval_distribution": {
|
||||
"within_5min": within_5min,
|
||||
"within_15min": within_15min,
|
||||
"within_30min": within_30min,
|
||||
"within_60min": within_60min,
|
||||
"over_60min": over_60min,
|
||||
},
|
||||
"interval_percentages": {
|
||||
"within_5min": round(within_5min / total_intervals * 100, 1),
|
||||
"within_15min": round(within_15min / total_intervals * 100, 1),
|
||||
"within_30min": round(within_30min / total_intervals * 100, 1),
|
||||
"within_60min": round(within_60min / total_intervals * 100, 1),
|
||||
"over_60min": round(over_60min / total_intervals * 100, 1),
|
||||
},
|
||||
"percentiles": {
|
||||
"p50": round(float(median_interval), 2) if median_interval else None,
|
||||
"p75": round(float(p75_interval), 2) if p75_interval else None,
|
||||
"p90": round(float(p90_interval), 2) if p90_interval else None,
|
||||
},
|
||||
"avg_interval_minutes": round(float(avg_interval), 2) if avg_interval else None,
|
||||
"min_interval_minutes": round(float(min_interval), 2) if min_interval else None,
|
||||
"max_interval_minutes": round(float(max_interval), 2) if max_interval else None,
|
||||
"recommended_ttl_minutes": recommended_ttl,
|
||||
"recommendation_reason": UsageService._get_ttl_recommendation_reason(
|
||||
recommended_ttl, p75_interval, p90_interval
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
# 汇总统计
|
||||
ttl_distribution = {"5min": 0, "15min": 0, "30min": 0, "60min": 0}
|
||||
@@ -2684,7 +2969,11 @@ class UsageService:
|
||||
"analysis_period_hours": hours,
|
||||
"total_requests": total_requests,
|
||||
"requests_with_cache_hit": requests_with_cache_hit_count,
|
||||
"request_cache_hit_rate": round(requests_with_cache_hit_count / total_requests * 100, 2) if total_requests > 0 else 0,
|
||||
"request_cache_hit_rate": (
|
||||
round(requests_with_cache_hit_count / total_requests * 100, 2)
|
||||
if total_requests > 0
|
||||
else 0
|
||||
),
|
||||
"total_input_tokens": total_input_tokens,
|
||||
"total_cache_read_tokens": total_cache_read_tokens,
|
||||
"total_cache_creation_tokens": total_cache_creation_tokens,
|
||||
@@ -2839,10 +3128,7 @@ class UsageService:
|
||||
else:
|
||||
for row in rows:
|
||||
created_at, model, interval_minutes = row
|
||||
point_data = {
|
||||
"x": created_at.isoformat(),
|
||||
"y": round(float(interval_minutes), 2)
|
||||
}
|
||||
point_data = {"x": created_at.isoformat(), "y": round(float(interval_minutes), 2)}
|
||||
if model:
|
||||
point_data["model"] = model
|
||||
models_set.add(model)
|
||||
|
||||
@@ -6,14 +6,19 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import sanitize_error_message
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
|
||||
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
@@ -21,8 +26,12 @@ from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.database import create_session
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
|
||||
from src.models.database import ApiKey, Provider, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.dimension_collector_service import DimensionCollectorService
|
||||
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
|
||||
from src.services.billing.rule_service import BillingRuleService
|
||||
from src.services.system.scheduler import get_scheduler
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
# 永久性错误指示词(用于降级判断,不应重试)
|
||||
_PERMANENT_ERROR_INDICATORS = frozenset(
|
||||
@@ -53,7 +62,6 @@ class VideoTaskPollerService:
|
||||
|
||||
LOCK_KEY = "video_task_poller:lock"
|
||||
LOCK_TTL = 60
|
||||
BATCH_SIZE = 50
|
||||
MAX_BACKOFF_SECONDS = 300
|
||||
# 连续失败告警阈值
|
||||
CONSECUTIVE_FAILURE_ALERT_THRESHOLD = 5
|
||||
@@ -63,17 +71,26 @@ class VideoTaskPollerService:
|
||||
self.redis = None
|
||||
self._openai_normalizer = OpenAINormalizer()
|
||||
self._gemini_normalizer = GeminiNormalizer()
|
||||
self._formula_engine = FormulaEngine()
|
||||
# 追踪连续失败次数(用于告警)
|
||||
self._consecutive_failures = 0
|
||||
# 从配置读取参数
|
||||
self._batch_size = config.video_poll_batch_size
|
||||
self._concurrency = config.video_poll_concurrency
|
||||
# Semaphore 延迟初始化,避免在事件循环外创建
|
||||
self._semaphore: asyncio.Semaphore | None = None
|
||||
|
||||
async def start(self) -> None:
|
||||
# 在事件循环内初始化 Semaphore
|
||||
if self._semaphore is None:
|
||||
self._semaphore = asyncio.Semaphore(self._concurrency)
|
||||
if self.redis is None:
|
||||
self.redis = await get_redis_client(require_redis=False)
|
||||
|
||||
scheduler = get_scheduler()
|
||||
scheduler.add_interval_job(
|
||||
self.poll_pending_tasks,
|
||||
seconds=10,
|
||||
seconds=config.video_poll_interval_seconds,
|
||||
job_id="video_task_poller",
|
||||
name="视频任务轮询",
|
||||
)
|
||||
@@ -106,7 +123,7 @@ class VideoTaskPollerService:
|
||||
VideoTask.poll_count < VideoTask.max_poll_count,
|
||||
)
|
||||
.order_by(VideoTask.next_poll_at.asc())
|
||||
.limit(self.BATCH_SIZE)
|
||||
.limit(self._batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
@@ -115,32 +132,55 @@ class VideoTaskPollerService:
|
||||
self._consecutive_failures = 0
|
||||
return
|
||||
|
||||
batch_failures = 0
|
||||
for task in tasks:
|
||||
# 提取任务 ID 列表,释放查询 session 后逐个轮询
|
||||
task_ids = [t.id for t in tasks]
|
||||
|
||||
# 并发轮询:每个任务使用独立 session,避免共享 session 的并发风险
|
||||
poll_results: list[bool] = []
|
||||
|
||||
# 确保 semaphore 已初始化(在 start 中初始化,此处防御性检查)
|
||||
if self._semaphore is None:
|
||||
self._semaphore = asyncio.Semaphore(self._concurrency)
|
||||
semaphore = self._semaphore
|
||||
|
||||
async def poll_with_semaphore(task_id: str) -> None:
|
||||
"""带信号量的轮询,结果写入 poll_results"""
|
||||
async with semaphore:
|
||||
try:
|
||||
await self._poll_single_task(db, task)
|
||||
with create_session() as task_db:
|
||||
task_obj = task_db.query(VideoTask).get(task_id)
|
||||
if not task_obj:
|
||||
logger.warning("Task %s disappeared during poll", task_id)
|
||||
poll_results.append(True)
|
||||
return
|
||||
await self._poll_single_task(task_db, task_obj)
|
||||
task_db.commit()
|
||||
poll_results.append(True)
|
||||
except Exception as exc:
|
||||
batch_failures += 1
|
||||
# 单个任务失败不影响其他任务处理
|
||||
logger.exception(
|
||||
"Unexpected error polling task %s: %s",
|
||||
task.id,
|
||||
task_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
poll_results.append(False)
|
||||
|
||||
# 更新连续失败计数并检查告警阈值
|
||||
if batch_failures == len(tasks):
|
||||
self._consecutive_failures += 1
|
||||
if self._consecutive_failures >= self.CONSECUTIVE_FAILURE_ALERT_THRESHOLD:
|
||||
logger.error(
|
||||
"[ALERT] Video task poller: %d consecutive batches failed. "
|
||||
"Provider connectivity or configuration issue suspected.",
|
||||
self._consecutive_failures,
|
||||
)
|
||||
else:
|
||||
self._consecutive_failures = 0
|
||||
async with asyncio.TaskGroup() as tg:
|
||||
for tid in task_ids:
|
||||
tg.create_task(poll_with_semaphore(tid))
|
||||
|
||||
db.commit()
|
||||
batch_failures = sum(1 for r in poll_results if r is False)
|
||||
|
||||
# 更新连续失败计数并检查告警阈值
|
||||
if batch_failures == len(task_ids):
|
||||
self._consecutive_failures += 1
|
||||
if self._consecutive_failures >= self.CONSECUTIVE_FAILURE_ALERT_THRESHOLD:
|
||||
logger.error(
|
||||
"[ALERT] Video task poller: %d consecutive batches failed. "
|
||||
"Provider connectivity or configuration issue suspected.",
|
||||
self._consecutive_failures,
|
||||
)
|
||||
else:
|
||||
self._consecutive_failures = 0
|
||||
finally:
|
||||
await self._release_redis_lock(token)
|
||||
|
||||
@@ -156,11 +196,14 @@ class VideoTaskPollerService:
|
||||
# 存储多视频 URL(Gemini sampleCount > 1 时)
|
||||
if result.video_urls:
|
||||
task.video_urls = result.video_urls
|
||||
# 保存上游原始响应(用于审计/重算)
|
||||
self._attach_poll_raw_response(task, result)
|
||||
elif result.status == VideoStatus.FAILED:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = result.error_code
|
||||
task.error_message = result.error_message
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
self._attach_poll_raw_response(task, result)
|
||||
else:
|
||||
task.poll_count += 1
|
||||
task.progress_percent = result.progress_percent
|
||||
@@ -202,6 +245,308 @@ class VideoTaskPollerService:
|
||||
task.error_message = f"Task timed out after {task.poll_count} polls"
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
|
||||
# 终态写入 Usage(复用外层 per-task session)
|
||||
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||
try:
|
||||
await self._record_terminal_usage(db, task)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Failed to record video usage for task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
|
||||
if not result.raw_response:
|
||||
return
|
||||
if task.request_metadata is None:
|
||||
task.request_metadata = {}
|
||||
# 仅在终态写一次,避免污染 request_metadata
|
||||
task.request_metadata["poll_raw_response"] = result.raw_response
|
||||
|
||||
async def _record_terminal_usage(self, db: Session, task: VideoTask) -> None:
|
||||
"""
|
||||
为视频任务终态写入 Usage:
|
||||
- COMPLETED: 使用 FormulaEngine 计算 cost(或 no_rule / incomplete -> cost=0)
|
||||
- FAILED: cost=0
|
||||
"""
|
||||
request_id = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
request_id = task.request_metadata.get("request_id")
|
||||
request_id = request_id or task.id
|
||||
|
||||
# 计算异步任务总耗时(ms)
|
||||
response_time_ms = None
|
||||
if task.submitted_at and task.completed_at:
|
||||
delta = task.completed_at - task.submitted_at
|
||||
response_time_ms = int(delta.total_seconds() * 1000)
|
||||
|
||||
# 基础维度(无需 collectors 也可计费)
|
||||
base_dimensions: dict[str, Any] = {
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size or "",
|
||||
"retry_count": task.retry_count,
|
||||
}
|
||||
|
||||
# collectors 可用的 metadata(结构稳定,便于配置 path)
|
||||
collector_metadata: dict[str, Any] = {
|
||||
"task": {
|
||||
"id": task.id,
|
||||
"external_task_id": task.external_task_id,
|
||||
"model": task.model,
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size,
|
||||
"retry_count": task.retry_count,
|
||||
"video_size_bytes": task.video_size_bytes,
|
||||
},
|
||||
"result": {
|
||||
"video_url": task.video_url,
|
||||
"video_urls": task.video_urls or [],
|
||||
},
|
||||
}
|
||||
|
||||
# 维度采集:base + collectors 覆盖/补全
|
||||
dims = DimensionCollectorService(db).collect_dimensions(
|
||||
api_format=task.provider_api_format,
|
||||
task_type="video",
|
||||
request=task.original_request_body or {},
|
||||
response=(
|
||||
(task.request_metadata or {}).get("poll_raw_response")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
metadata=collector_metadata,
|
||||
base_dimensions=base_dimensions,
|
||||
)
|
||||
|
||||
# 取冻结的 rule_snapshot;若缺失则回退 DB 查找(兼容旧任务)
|
||||
rule_snapshot = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
|
||||
|
||||
# 构建 billing_snapshot(写入 Usage.request_metadata)
|
||||
billing_snapshot: dict[str, Any] = {
|
||||
"status": "complete",
|
||||
"missing_required": [],
|
||||
"strict_mode": config.billing_strict_mode,
|
||||
}
|
||||
cost = 0.0
|
||||
|
||||
if task.status == VideoStatus.FAILED.value:
|
||||
billing_snapshot["billed_reason"] = "task_failed"
|
||||
else:
|
||||
# COMPLETED:计算成本
|
||||
expression = None
|
||||
variables = None
|
||||
dimension_mappings = None
|
||||
rule_id = None
|
||||
rule_name = None
|
||||
rule_scope = None
|
||||
|
||||
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
|
||||
rule_id = rule_snapshot.get("rule_id")
|
||||
rule_name = rule_snapshot.get("rule_name")
|
||||
rule_scope = rule_snapshot.get("scope")
|
||||
expression = rule_snapshot.get("expression")
|
||||
variables = rule_snapshot.get("variables")
|
||||
dimension_mappings = rule_snapshot.get("dimension_mappings")
|
||||
else:
|
||||
lookup = BillingRuleService.find_rule(
|
||||
db,
|
||||
provider_id=task.provider_id,
|
||||
model_name=task.model,
|
||||
task_type="video",
|
||||
)
|
||||
if lookup:
|
||||
rule = lookup.rule
|
||||
rule_id = rule.id
|
||||
rule_name = rule.name
|
||||
rule_scope = lookup.scope
|
||||
expression = rule.expression
|
||||
variables = rule.variables
|
||||
dimension_mappings = rule.dimension_mappings
|
||||
|
||||
if not expression:
|
||||
billing_snapshot["status"] = "no_rule"
|
||||
billing_snapshot["cost_breakdown"] = {"total": 0.0}
|
||||
logger.warning(
|
||||
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
task.provider_id,
|
||||
)
|
||||
else:
|
||||
billing_snapshot.update(
|
||||
{
|
||||
"rule_id": rule_id,
|
||||
"rule_name": rule_name,
|
||||
"rule_scope": rule_scope,
|
||||
"expression": expression,
|
||||
"variables": variables or {},
|
||||
}
|
||||
)
|
||||
try:
|
||||
result = self._formula_engine.evaluate(
|
||||
expression=expression,
|
||||
variables=variables or {},
|
||||
dimensions=dims,
|
||||
dimension_mappings=dimension_mappings or {},
|
||||
strict_mode=config.billing_strict_mode,
|
||||
)
|
||||
billing_snapshot["status"] = result.status
|
||||
billing_snapshot["missing_required"] = result.missing_required
|
||||
billing_snapshot["resolved_values"] = result.resolved_values
|
||||
if result.status == "complete":
|
||||
cost = result.cost
|
||||
else:
|
||||
logger.error(
|
||||
"Billing incomplete due to missing required dimensions "
|
||||
"(request_id=%s, model=%s, missing=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
result.missing_required,
|
||||
)
|
||||
cost = 0.0
|
||||
await self._maybe_alert_missing_required(
|
||||
model=task.model,
|
||||
missing_required=result.missing_required,
|
||||
)
|
||||
if result.error:
|
||||
billing_snapshot["error"] = result.error
|
||||
except BillingIncompleteError as exc:
|
||||
logger.error(
|
||||
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
exc.missing_required,
|
||||
)
|
||||
billing_snapshot["status"] = "incomplete"
|
||||
billing_snapshot["missing_required"] = exc.missing_required
|
||||
billing_snapshot["resolved_values"] = {}
|
||||
billing_snapshot["error"] = "strict_mode_missing_required"
|
||||
cost = 0.0
|
||||
|
||||
# strict_mode=true:标记任务失败并隐藏产物,避免"免费放行"
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "billing_incomplete"
|
||||
task.error_message = f"Missing required dimensions: {exc.missing_required}"
|
||||
task.video_url = None
|
||||
task.video_urls = None
|
||||
|
||||
await self._maybe_alert_missing_required(
|
||||
model=task.model,
|
||||
missing_required=exc.missing_required,
|
||||
)
|
||||
|
||||
billing_snapshot["cost_breakdown"] = {"total": cost}
|
||||
|
||||
# Usage 元数据(包含 snapshot + dimensions + raw_response_ref)
|
||||
usage_metadata: dict[str, Any] = {
|
||||
"billing_snapshot": billing_snapshot,
|
||||
"dimensions": dims,
|
||||
"raw_response_ref": {
|
||||
"video_task_id": task.id,
|
||||
"field": "video_tasks.request_metadata.poll_raw_response",
|
||||
},
|
||||
}
|
||||
|
||||
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
|
||||
user_obj = db.query(User).filter(User.id == task.user_id).first()
|
||||
api_key_obj = (
|
||||
db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
|
||||
if task.api_key_id
|
||||
else None
|
||||
)
|
||||
provider_obj = (
|
||||
db.query(Provider).filter(Provider.id == task.provider_id).first()
|
||||
if task.provider_id
|
||||
else None
|
||||
)
|
||||
provider_name = provider_obj.name if provider_obj else "unknown"
|
||||
|
||||
await UsageService.record_usage_with_custom_cost(
|
||||
db=db,
|
||||
user=user_obj,
|
||||
api_key=api_key_obj,
|
||||
provider=provider_name,
|
||||
model=task.model,
|
||||
request_type="video",
|
||||
total_cost_usd=cost,
|
||||
request_cost_usd=cost,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
api_format=task.client_api_format,
|
||||
endpoint_api_format=task.provider_api_format,
|
||||
has_format_conversion=bool(task.format_converted),
|
||||
is_stream=False,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=None,
|
||||
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
|
||||
error_message=(
|
||||
None
|
||||
if task.status == VideoStatus.COMPLETED.value
|
||||
else (task.error_message or task.error_code or "video_task_failed")
|
||||
),
|
||||
metadata=usage_metadata,
|
||||
request_headers=(
|
||||
(task.request_metadata or {}).get("request_headers")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
request_body=task.original_request_body,
|
||||
provider_request_headers=None,
|
||||
response_headers=None,
|
||||
client_response_headers=None,
|
||||
response_body=None,
|
||||
request_id=request_id,
|
||||
provider_id=task.provider_id,
|
||||
provider_endpoint_id=task.endpoint_id,
|
||||
provider_api_key_id=task.key_id,
|
||||
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
|
||||
target_model=None,
|
||||
)
|
||||
|
||||
async def _maybe_alert_missing_required(
|
||||
self, *, model: str, missing_required: list[str]
|
||||
) -> None:
|
||||
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
|
||||
if not missing_required:
|
||||
return
|
||||
if not self.redis:
|
||||
# Redis 不可用:降级为日志
|
||||
logger.error(
|
||||
"Missing required billing dimensions (model=%s): %s", model, missing_required
|
||||
)
|
||||
return
|
||||
|
||||
# 按小时 bucket 聚合
|
||||
now = datetime.now(timezone.utc)
|
||||
hour_bucket = now.strftime("%Y%m%d%H")
|
||||
for dim in missing_required:
|
||||
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
|
||||
try:
|
||||
count = await self.redis.incr(key)
|
||||
# TTL 略大于 1h,避免边界抖动
|
||||
if count == 1:
|
||||
await self.redis.expire(key, 3700)
|
||||
if count >= 10:
|
||||
logger.warning(
|
||||
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
|
||||
model,
|
||||
dim,
|
||||
count,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
|
||||
)
|
||||
|
||||
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
|
||||
"""判断是否为永久性错误(不应重试)"""
|
||||
# 优先使用 HTTP 状态码判断
|
||||
@@ -282,9 +627,7 @@ class VideoTaskPollerService:
|
||||
error_code="missing_external_task_id",
|
||||
error_message="Task missing external_task_id",
|
||||
)
|
||||
operation_name = task.external_task_id
|
||||
if not operation_name.startswith("operations/"):
|
||||
operation_name = f"operations/{operation_name}"
|
||||
operation_name = normalize_gemini_operation_id(task.external_task_id)
|
||||
url = self._build_gemini_url(endpoint.base_url, operation_name)
|
||||
headers = self._build_headers(APIFormat.GEMINI, upstream_key, endpoint, auth_info)
|
||||
|
||||
|
||||
105
tests/services/billing/test_dimension_collector_service.py
Normal file
105
tests/services/billing/test_dimension_collector_service.py
Normal 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
|
||||
169
tests/services/billing/test_formula_engine.py
Normal file
169
tests/services/billing/test_formula_engine.py
Normal 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=true:0 合法
|
||||
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
|
||||
Reference in New Issue
Block a user