mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
refactor: 重构异步任务系统和计费服务架构
- 重构任务系统:新增 lifecycle (TaskStatus/BillingStatus)、context、application 模块 - 将 video tasks 泛化为 async tasks,支持更通用的异步任务管理 - 新增 Gemini Files 管理模块和管理界面 - 重构 billing 服务:拆分 schema.py 和 service.py - 新增 candidate 服务模块用于请求候选管理 - 数据库迁移:添加 billing_status、request_id、gemini_file_mappings 表和索引 - 移除废弃的 video_telemetry、task orchestrator 等模块
This commit is contained in:
@@ -0,0 +1,324 @@
|
|||||||
|
"""Add usage billing, video_tasks fields, gemini_file_mappings, provider format conversion, and indexes
|
||||||
|
|
||||||
|
Revision ID: a2f1b3c4d5e6
|
||||||
|
Revises: cf40e6a5c5b1
|
||||||
|
Create Date: 2026-02-01 12:00:00+00:00
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
1. usage 表:
|
||||||
|
- 添加 billing_status (pending/settled/void),用于表示结算状态
|
||||||
|
- 添加 finalized_at,用于记录结算完成时间
|
||||||
|
- 添加 (provider_name, created_at) 和 (model, created_at) 索引
|
||||||
|
|
||||||
|
2. video_tasks 表:
|
||||||
|
- 添加 request_id(全局唯一),用于与 Usage/RequestCandidate 建立稳定关联
|
||||||
|
- 添加 short_id (Gemini-style short ID)
|
||||||
|
|
||||||
|
3. gemini_file_mappings 表:
|
||||||
|
- 创建新表用于文件映射
|
||||||
|
- 添加 source_hash 字段用于关联相同源文件
|
||||||
|
|
||||||
|
4. providers 表:
|
||||||
|
- 添加 enable_format_conversion 开关字段
|
||||||
|
|
||||||
|
5. request_candidates 表:
|
||||||
|
- 添加 created_at 索引
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import secrets
|
||||||
|
import string
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from sqlalchemy import inspect, text
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "a2f1b3c4d5e6"
|
||||||
|
down_revision: Union[str, None] = "cf40e6a5c5b1"
|
||||||
|
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)
|
||||||
|
indexes = inspector.get_indexes(table_name)
|
||||||
|
return any(idx.get("name") == index_name for idx in indexes)
|
||||||
|
|
||||||
|
|
||||||
|
def unique_constraint_exists(table_name: str, constraint_name: str) -> bool:
|
||||||
|
bind = op.get_bind()
|
||||||
|
inspector = inspect(bind)
|
||||||
|
constraints = inspector.get_unique_constraints(table_name)
|
||||||
|
return any(c.get("name") == constraint_name for c in constraints)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_short_id(length: int = 12) -> str:
|
||||||
|
"""Generate a Gemini-style short ID (lowercase letters + digits)"""
|
||||||
|
alphabet = string.ascii_lowercase + string.digits
|
||||||
|
return "".join(secrets.choice(alphabet) for _ in range(length))
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
bind = op.get_bind()
|
||||||
|
dialect = bind.dialect.name
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 1. usage 表: billing_status + finalized_at + 索引
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("usage"):
|
||||||
|
if not column_exists("usage", "billing_status"):
|
||||||
|
op.add_column(
|
||||||
|
"usage",
|
||||||
|
sa.Column(
|
||||||
|
"billing_status",
|
||||||
|
sa.String(20),
|
||||||
|
nullable=False,
|
||||||
|
server_default="settled",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
if not column_exists("usage", "finalized_at"):
|
||||||
|
op.add_column(
|
||||||
|
"usage",
|
||||||
|
sa.Column("finalized_at", sa.DateTime(timezone=True), nullable=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
if not index_exists("usage", "idx_usage_billing_status"):
|
||||||
|
op.create_index("idx_usage_billing_status", "usage", ["billing_status"])
|
||||||
|
|
||||||
|
# (provider_name, created_at) — provider list / dashboard queries
|
||||||
|
if (
|
||||||
|
column_exists("usage", "provider_name")
|
||||||
|
and column_exists("usage", "created_at")
|
||||||
|
and not index_exists("usage", "idx_usage_provider_created")
|
||||||
|
):
|
||||||
|
op.create_index("idx_usage_provider_created", "usage", ["provider_name", "created_at"])
|
||||||
|
|
||||||
|
# (model, created_at) — model analytics / recent requests queries
|
||||||
|
if (
|
||||||
|
column_exists("usage", "model")
|
||||||
|
and column_exists("usage", "created_at")
|
||||||
|
and not index_exists("usage", "idx_usage_model_created")
|
||||||
|
):
|
||||||
|
op.create_index("idx_usage_model_created", "usage", ["model", "created_at"])
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 2. video_tasks 表: request_id + short_id
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("video_tasks"):
|
||||||
|
# --- request_id ---
|
||||||
|
if not column_exists("video_tasks", "request_id"):
|
||||||
|
op.add_column(
|
||||||
|
"video_tasks",
|
||||||
|
sa.Column("request_id", sa.String(100), nullable=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 回填 request_id
|
||||||
|
if dialect == "postgresql":
|
||||||
|
op.execute("""
|
||||||
|
UPDATE video_tasks
|
||||||
|
SET request_id = COALESCE(request_metadata->>'request_id', id)
|
||||||
|
WHERE request_id IS NULL
|
||||||
|
""")
|
||||||
|
elif dialect == "sqlite":
|
||||||
|
op.execute("""
|
||||||
|
UPDATE video_tasks
|
||||||
|
SET request_id = COALESCE(json_extract(request_metadata, '$.request_id'), id)
|
||||||
|
WHERE request_id IS NULL
|
||||||
|
""")
|
||||||
|
else:
|
||||||
|
op.execute("""
|
||||||
|
UPDATE video_tasks
|
||||||
|
SET request_id = id
|
||||||
|
WHERE request_id IS NULL
|
||||||
|
""")
|
||||||
|
|
||||||
|
if dialect == "postgresql":
|
||||||
|
op.alter_column("video_tasks", "request_id", nullable=False)
|
||||||
|
|
||||||
|
if not index_exists("video_tasks", "idx_video_tasks_request_id"):
|
||||||
|
op.create_index("idx_video_tasks_request_id", "video_tasks", ["request_id"])
|
||||||
|
|
||||||
|
if not unique_constraint_exists("video_tasks", "uq_video_tasks_request_id"):
|
||||||
|
op.create_unique_constraint(
|
||||||
|
"uq_video_tasks_request_id",
|
||||||
|
"video_tasks",
|
||||||
|
["request_id"],
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- short_id ---
|
||||||
|
if not column_exists("video_tasks", "short_id"):
|
||||||
|
op.add_column(
|
||||||
|
"video_tasks",
|
||||||
|
sa.Column("short_id", sa.String(16), nullable=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Populate existing rows with unique short_ids
|
||||||
|
result = bind.execute(text("SELECT id FROM video_tasks WHERE short_id IS NULL"))
|
||||||
|
for row in result:
|
||||||
|
short_id = generate_short_id()
|
||||||
|
bind.execute(
|
||||||
|
text("UPDATE video_tasks SET short_id = :short_id WHERE id = :id"),
|
||||||
|
{"short_id": short_id, "id": row[0]},
|
||||||
|
)
|
||||||
|
|
||||||
|
op.alter_column("video_tasks", "short_id", nullable=False)
|
||||||
|
op.create_index("ix_video_tasks_short_id", "video_tasks", ["short_id"], unique=True)
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 3. gemini_file_mappings 表
|
||||||
|
# =========================================================================
|
||||||
|
if not table_exists("gemini_file_mappings"):
|
||||||
|
op.create_table(
|
||||||
|
"gemini_file_mappings",
|
||||||
|
sa.Column("id", sa.String(36), primary_key=True),
|
||||||
|
sa.Column("file_name", sa.String(255), nullable=False, unique=True),
|
||||||
|
sa.Column(
|
||||||
|
"key_id",
|
||||||
|
sa.String(36),
|
||||||
|
sa.ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column(
|
||||||
|
"user_id",
|
||||||
|
sa.String(36),
|
||||||
|
sa.ForeignKey("users.id", ondelete="CASCADE"),
|
||||||
|
nullable=True,
|
||||||
|
),
|
||||||
|
sa.Column("display_name", sa.String(255), nullable=True),
|
||||||
|
sa.Column("mime_type", sa.String(100), nullable=True),
|
||||||
|
sa.Column("source_hash", sa.String(64), nullable=True),
|
||||||
|
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||||
|
)
|
||||||
|
|
||||||
|
op.create_index("ix_gemini_file_mappings_id", "gemini_file_mappings", ["id"])
|
||||||
|
op.create_index(
|
||||||
|
"ix_gemini_file_mappings_file_name", "gemini_file_mappings", ["file_name"], unique=True
|
||||||
|
)
|
||||||
|
op.create_index("ix_gemini_file_mappings_key_id", "gemini_file_mappings", ["key_id"])
|
||||||
|
op.create_index("ix_gemini_file_mappings_user_id", "gemini_file_mappings", ["user_id"])
|
||||||
|
op.create_index("idx_gemini_file_mappings_expires", "gemini_file_mappings", ["expires_at"])
|
||||||
|
op.create_index(
|
||||||
|
"idx_gemini_file_mappings_source_hash", "gemini_file_mappings", ["source_hash"]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# 表已存在,只添加 source_hash
|
||||||
|
if not column_exists("gemini_file_mappings", "source_hash"):
|
||||||
|
op.add_column(
|
||||||
|
"gemini_file_mappings",
|
||||||
|
sa.Column("source_hash", sa.String(64), nullable=True),
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_gemini_file_mappings_source_hash",
|
||||||
|
"gemini_file_mappings",
|
||||||
|
["source_hash"],
|
||||||
|
)
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 4. providers 表: enable_format_conversion
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("providers") and not column_exists("providers", "enable_format_conversion"):
|
||||||
|
op.add_column(
|
||||||
|
"providers",
|
||||||
|
sa.Column(
|
||||||
|
"enable_format_conversion",
|
||||||
|
sa.Boolean(),
|
||||||
|
nullable=False,
|
||||||
|
server_default=sa.text("false"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 5. request_candidates 表: created_at 索引
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("request_candidates"):
|
||||||
|
if not index_exists("request_candidates", "idx_request_candidates_created_at"):
|
||||||
|
op.create_index(
|
||||||
|
"idx_request_candidates_created_at",
|
||||||
|
"request_candidates",
|
||||||
|
["created_at"],
|
||||||
|
unique=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
bind = op.get_bind()
|
||||||
|
dialect = bind.dialect.name
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 5. request_candidates 表回滚
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("request_candidates"):
|
||||||
|
if index_exists("request_candidates", "idx_request_candidates_created_at"):
|
||||||
|
op.drop_index("idx_request_candidates_created_at", table_name="request_candidates")
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 4. providers 表回滚
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("providers") and column_exists("providers", "enable_format_conversion"):
|
||||||
|
op.drop_column("providers", "enable_format_conversion")
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 3. gemini_file_mappings 表回滚
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("gemini_file_mappings"):
|
||||||
|
op.drop_index("idx_gemini_file_mappings_source_hash", table_name="gemini_file_mappings")
|
||||||
|
op.drop_index("idx_gemini_file_mappings_expires", table_name="gemini_file_mappings")
|
||||||
|
op.drop_index("ix_gemini_file_mappings_user_id", table_name="gemini_file_mappings")
|
||||||
|
op.drop_index("ix_gemini_file_mappings_key_id", table_name="gemini_file_mappings")
|
||||||
|
op.drop_index("ix_gemini_file_mappings_file_name", table_name="gemini_file_mappings")
|
||||||
|
op.drop_index("ix_gemini_file_mappings_id", table_name="gemini_file_mappings")
|
||||||
|
op.drop_table("gemini_file_mappings")
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 2. video_tasks 表回滚
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("video_tasks"):
|
||||||
|
# short_id
|
||||||
|
if column_exists("video_tasks", "short_id"):
|
||||||
|
if index_exists("video_tasks", "ix_video_tasks_short_id"):
|
||||||
|
op.drop_index("ix_video_tasks_short_id", table_name="video_tasks")
|
||||||
|
op.drop_column("video_tasks", "short_id")
|
||||||
|
|
||||||
|
# request_id
|
||||||
|
if column_exists("video_tasks", "request_id"):
|
||||||
|
if dialect == "postgresql":
|
||||||
|
if unique_constraint_exists("video_tasks", "uq_video_tasks_request_id"):
|
||||||
|
op.drop_constraint("uq_video_tasks_request_id", "video_tasks", type_="unique")
|
||||||
|
if index_exists("video_tasks", "idx_video_tasks_request_id"):
|
||||||
|
op.drop_index("idx_video_tasks_request_id", table_name="video_tasks")
|
||||||
|
op.drop_column("video_tasks", "request_id")
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# 1. usage 表回滚
|
||||||
|
# =========================================================================
|
||||||
|
if table_exists("usage"):
|
||||||
|
if index_exists("usage", "idx_usage_model_created"):
|
||||||
|
op.drop_index("idx_usage_model_created", table_name="usage")
|
||||||
|
if index_exists("usage", "idx_usage_provider_created"):
|
||||||
|
op.drop_index("idx_usage_provider_created", table_name="usage")
|
||||||
|
if index_exists("usage", "idx_usage_billing_status"):
|
||||||
|
op.drop_index("idx_usage_billing_status", table_name="usage")
|
||||||
|
if column_exists("usage", "finalized_at"):
|
||||||
|
op.drop_column("usage", "finalized_at")
|
||||||
|
if column_exists("usage", "billing_status"):
|
||||||
|
op.drop_column("usage", "billing_status")
|
||||||
@@ -1,17 +1,21 @@
|
|||||||
import apiClient from './client'
|
import apiClient from './client'
|
||||||
|
|
||||||
// 视频任务状态
|
// 异步任务状态
|
||||||
export type VideoTaskStatus = 'pending' | 'submitted' | 'queued' | 'processing' | 'completed' | 'failed' | 'cancelled'
|
export type AsyncTaskStatus = 'pending' | 'submitted' | 'queued' | 'processing' | 'completed' | 'failed' | 'cancelled'
|
||||||
|
|
||||||
// 视频任务列表项
|
// 异步任务类型
|
||||||
export interface VideoTaskItem {
|
export type AsyncTaskType = 'video'
|
||||||
|
|
||||||
|
// 异步任务列表项
|
||||||
|
export interface AsyncTaskItem {
|
||||||
id: string
|
id: string
|
||||||
external_task_id: string
|
external_task_id: string
|
||||||
user_id: string
|
user_id: string
|
||||||
username: string
|
username: string
|
||||||
|
task_type: AsyncTaskType
|
||||||
model: string
|
model: string
|
||||||
prompt: string
|
prompt: string
|
||||||
status: VideoTaskStatus
|
status: AsyncTaskStatus
|
||||||
progress_percent: number
|
progress_percent: number
|
||||||
progress_message: string | null
|
progress_message: string | null
|
||||||
provider_id: string
|
provider_id: string
|
||||||
@@ -44,7 +48,7 @@ export interface CandidateKeyInfo {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 请求元数据
|
// 请求元数据
|
||||||
export interface VideoTaskRequestMetadata {
|
export interface AsyncTaskRequestMetadata {
|
||||||
candidate_keys: CandidateKeyInfo[]
|
candidate_keys: CandidateKeyInfo[]
|
||||||
selected_key_id: string
|
selected_key_id: string
|
||||||
selected_endpoint_id: string
|
selected_endpoint_id: string
|
||||||
@@ -52,10 +56,12 @@ export interface VideoTaskRequestMetadata {
|
|||||||
user_agent: string
|
user_agent: string
|
||||||
request_id: string
|
request_id: string
|
||||||
request_headers?: Record<string, string>
|
request_headers?: Record<string, string>
|
||||||
|
poll_raw_response?: any // 轮询完成时的原始响应
|
||||||
|
billing_snapshot?: any // 计费快照
|
||||||
}
|
}
|
||||||
|
|
||||||
// 视频任务详情
|
// 异步任务详情
|
||||||
export interface VideoTaskDetail extends VideoTaskItem {
|
export interface AsyncTaskDetail extends AsyncTaskItem {
|
||||||
api_key_id: string
|
api_key_id: string
|
||||||
endpoint_id: string
|
endpoint_id: string
|
||||||
key_id: string
|
key_id: string
|
||||||
@@ -81,73 +87,76 @@ export interface VideoTaskDetail extends VideoTaskItem {
|
|||||||
base_url: string
|
base_url: string
|
||||||
api_format: string
|
api_format: string
|
||||||
} | null
|
} | null
|
||||||
request_metadata: VideoTaskRequestMetadata | null
|
request_metadata: AsyncTaskRequestMetadata | null
|
||||||
}
|
}
|
||||||
|
|
||||||
// 视频任务列表响应
|
// 异步任务列表响应
|
||||||
export interface VideoTaskListResponse {
|
export interface AsyncTaskListResponse {
|
||||||
items: VideoTaskItem[]
|
items: AsyncTaskItem[]
|
||||||
total: number
|
total: number
|
||||||
page: number
|
page: number
|
||||||
page_size: number
|
page_size: number
|
||||||
pages: number
|
pages: number
|
||||||
}
|
}
|
||||||
|
|
||||||
// 视频任务统计响应
|
// 异步任务统计响应
|
||||||
export interface VideoTaskStatsResponse {
|
export interface AsyncTaskStatsResponse {
|
||||||
total: number
|
total: number
|
||||||
by_status: Record<VideoTaskStatus, number>
|
by_status: Record<AsyncTaskStatus, number>
|
||||||
by_model: Record<string, number>
|
by_model: Record<string, number>
|
||||||
today_count: number
|
today_count: number
|
||||||
active_users?: number // 仅管理员
|
active_users?: number // 仅管理员
|
||||||
processing_count?: number // 仅管理员
|
processing_count?: number // 仅管理员
|
||||||
}
|
}
|
||||||
|
|
||||||
// 视频任务查询参数
|
// 异步任务查询参数
|
||||||
export interface VideoTaskQueryParams {
|
export interface AsyncTaskQueryParams {
|
||||||
status?: VideoTaskStatus
|
status?: AsyncTaskStatus
|
||||||
|
task_type?: AsyncTaskType
|
||||||
user_id?: string
|
user_id?: string
|
||||||
model?: string
|
model?: string
|
||||||
page?: number
|
page?: number
|
||||||
page_size?: number
|
page_size?: number
|
||||||
}
|
}
|
||||||
|
|
||||||
export const videoTasksApi = {
|
export const asyncTasksApi = {
|
||||||
/**
|
/**
|
||||||
* 获取视频任务列表
|
* 获取异步任务列表
|
||||||
*/
|
*/
|
||||||
async list(params: VideoTaskQueryParams = {}): Promise<VideoTaskListResponse> {
|
async list(params: AsyncTaskQueryParams = {}): Promise<AsyncTaskListResponse> {
|
||||||
const searchParams = new URLSearchParams()
|
const searchParams = new URLSearchParams()
|
||||||
if (params.status) searchParams.append('status', params.status)
|
if (params.status) searchParams.append('status', params.status)
|
||||||
|
if (params.task_type) searchParams.append('task_type', params.task_type)
|
||||||
if (params.user_id) searchParams.append('user_id', params.user_id)
|
if (params.user_id) searchParams.append('user_id', params.user_id)
|
||||||
if (params.model) searchParams.append('model', params.model)
|
if (params.model) searchParams.append('model', params.model)
|
||||||
if (params.page) searchParams.append('page', params.page.toString())
|
if (params.page) searchParams.append('page', params.page.toString())
|
||||||
if (params.page_size) searchParams.append('page_size', params.page_size.toString())
|
if (params.page_size) searchParams.append('page_size', params.page_size.toString())
|
||||||
|
|
||||||
const query = searchParams.toString()
|
const query = searchParams.toString()
|
||||||
|
// 后端 API 路径保持不变,前端抽象为异步任务
|
||||||
const url = query ? `/api/admin/video-tasks?${query}` : '/api/admin/video-tasks'
|
const url = query ? `/api/admin/video-tasks?${query}` : '/api/admin/video-tasks'
|
||||||
const response = await apiClient.get(url)
|
const response = await apiClient.get(url)
|
||||||
return response.data
|
return response.data
|
||||||
},
|
},
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 获取视频任务统计
|
* 获取异步任务统计
|
||||||
*/
|
*/
|
||||||
async getStats(): Promise<VideoTaskStatsResponse> {
|
async getStats(): Promise<AsyncTaskStatsResponse> {
|
||||||
const response = await apiClient.get('/api/admin/video-tasks/stats')
|
const response = await apiClient.get('/api/admin/video-tasks/stats')
|
||||||
return response.data
|
return response.data
|
||||||
},
|
},
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 获取视频任务详情
|
* 获取异步任务详情
|
||||||
*/
|
*/
|
||||||
async getDetail(taskId: string): Promise<VideoTaskDetail> {
|
async getDetail(taskId: string): Promise<AsyncTaskDetail> {
|
||||||
const response = await apiClient.get(`/api/admin/video-tasks/${taskId}`)
|
const response = await apiClient.get(`/api/admin/video-tasks/${taskId}`)
|
||||||
return response.data
|
return response.data
|
||||||
},
|
},
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 取消视频任务
|
* 取消异步任务
|
||||||
*/
|
*/
|
||||||
async cancel(taskId: string): Promise<{ id: string; status: string; message: string }> {
|
async cancel(taskId: string): Promise<{ id: string; status: string; message: string }> {
|
||||||
const response = await apiClient.post(`/api/admin/video-tasks/${taskId}/cancel`)
|
const response = await apiClient.post(`/api/admin/video-tasks/${taskId}/cancel`)
|
||||||
@@ -155,4 +164,4 @@ export const videoTasksApi = {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
export default videoTasksApi
|
export default asyncTasksApi
|
||||||
@@ -340,6 +340,7 @@ export interface ProviderWithEndpointsSummary {
|
|||||||
website?: string
|
website?: string
|
||||||
provider_priority: number
|
provider_priority: number
|
||||||
keep_priority_on_conversion: boolean // 格式转换时是否保持优先级
|
keep_priority_on_conversion: boolean // 格式转换时是否保持优先级
|
||||||
|
enable_format_conversion: boolean // 是否允许格式转换(提供商级别开关)
|
||||||
billing_type?: 'monthly_quota' | 'pay_as_you_go' | 'free_tier'
|
billing_type?: 'monthly_quota' | 'pay_as_you_go' | 'free_tier'
|
||||||
monthly_quota_usd?: number
|
monthly_quota_usd?: number
|
||||||
monthly_used_usd?: number
|
monthly_used_usd?: number
|
||||||
|
|||||||
124
frontend/src/api/gemini-files.ts
Normal file
124
frontend/src/api/gemini-files.ts
Normal file
@@ -0,0 +1,124 @@
|
|||||||
|
/**
|
||||||
|
* Gemini Files 管理 API
|
||||||
|
*/
|
||||||
|
|
||||||
|
import apiClient from './client'
|
||||||
|
|
||||||
|
export interface FileMappingResponse {
|
||||||
|
id: string
|
||||||
|
file_name: string
|
||||||
|
key_id: string
|
||||||
|
key_name: string | null
|
||||||
|
user_id: string | null
|
||||||
|
username: string | null
|
||||||
|
display_name: string | null
|
||||||
|
mime_type: string | null
|
||||||
|
created_at: string
|
||||||
|
expires_at: string
|
||||||
|
is_expired: boolean
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface FileMappingListResponse {
|
||||||
|
items: FileMappingResponse[]
|
||||||
|
total: number
|
||||||
|
page: number
|
||||||
|
page_size: number
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface FileMappingStatsResponse {
|
||||||
|
total_mappings: number
|
||||||
|
active_mappings: number
|
||||||
|
expired_mappings: number
|
||||||
|
by_mime_type: Record<string, number>
|
||||||
|
capable_keys_count: number
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface ListMappingsParams {
|
||||||
|
page?: number
|
||||||
|
page_size?: number
|
||||||
|
include_expired?: boolean
|
||||||
|
search?: string
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface CapableKeyResponse {
|
||||||
|
id: string
|
||||||
|
name: string
|
||||||
|
provider_name: string | null
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface UploadResultItem {
|
||||||
|
key_id: string
|
||||||
|
key_name: string | null
|
||||||
|
success: boolean
|
||||||
|
file_name: string | null
|
||||||
|
error: string | null
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface UploadResponse {
|
||||||
|
display_name: string
|
||||||
|
mime_type: string
|
||||||
|
size_bytes: number
|
||||||
|
results: UploadResultItem[]
|
||||||
|
success_count: number
|
||||||
|
fail_count: number
|
||||||
|
}
|
||||||
|
|
||||||
|
export const geminiFilesApi = {
|
||||||
|
/**
|
||||||
|
* 获取文件映射统计
|
||||||
|
*/
|
||||||
|
async getStats(): Promise<FileMappingStatsResponse> {
|
||||||
|
const response = await apiClient.get('/api/admin/gemini-files/stats')
|
||||||
|
return response.data
|
||||||
|
},
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 列出文件映射
|
||||||
|
*/
|
||||||
|
async listMappings(params?: ListMappingsParams): Promise<FileMappingListResponse> {
|
||||||
|
const response = await apiClient.get('/api/admin/gemini-files/mappings', { params })
|
||||||
|
return response.data
|
||||||
|
},
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 删除指定映射
|
||||||
|
*/
|
||||||
|
async deleteMapping(mappingId: string): Promise<{ message: string; file_name: string }> {
|
||||||
|
const response = await apiClient.delete(`/api/admin/gemini-files/mappings/${mappingId}`)
|
||||||
|
return response.data
|
||||||
|
},
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 清理过期映射
|
||||||
|
*/
|
||||||
|
async cleanupExpired(): Promise<{ message: string; deleted_count: number }> {
|
||||||
|
const response = await apiClient.delete('/api/admin/gemini-files/mappings')
|
||||||
|
return response.data
|
||||||
|
},
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 获取可用的 Key 列表
|
||||||
|
*/
|
||||||
|
async getCapableKeys(): Promise<CapableKeyResponse[]> {
|
||||||
|
const response = await apiClient.get('/api/admin/gemini-files/capable-keys')
|
||||||
|
return response.data
|
||||||
|
},
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 上传文件到指定的 Keys
|
||||||
|
*/
|
||||||
|
async uploadFile(file: File, keyIds: string[]): Promise<UploadResponse> {
|
||||||
|
const formData = new FormData()
|
||||||
|
formData.append('file', file)
|
||||||
|
const response = await apiClient.post(
|
||||||
|
`/api/admin/gemini-files/upload?key_ids=${keyIds.join(',')}`,
|
||||||
|
formData,
|
||||||
|
{
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'multipart/form-data'
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return response.data
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -242,25 +242,41 @@
|
|||||||
</form>
|
</form>
|
||||||
|
|
||||||
<template #footer>
|
<template #footer>
|
||||||
<Button
|
<div class="flex w-full items-center justify-between">
|
||||||
variant="outline"
|
<!-- 左侧:清除按钮(仅在已有配置时显示) -->
|
||||||
@click="$emit('update:open', false)"
|
<div>
|
||||||
>
|
<Button
|
||||||
取消
|
v-if="hasExistingConfig"
|
||||||
</Button>
|
variant="destructive"
|
||||||
<Button
|
:disabled="isClearing"
|
||||||
:disabled="isSaving || !canSave"
|
@click="handleClear"
|
||||||
@click="handleSave"
|
>
|
||||||
>
|
{{ isClearing ? '清除中...' : '清除' }}
|
||||||
{{ isSaving ? '保存中...' : '保存' }}
|
</Button>
|
||||||
</Button>
|
</div>
|
||||||
<Button
|
<!-- 右侧:验证、保存、取消按钮 -->
|
||||||
variant="outline"
|
<div class="flex gap-2">
|
||||||
:disabled="isVerifying || !canVerify"
|
<Button
|
||||||
@click="handleVerify"
|
variant="outline"
|
||||||
>
|
:disabled="isVerifying || !canVerify"
|
||||||
{{ isVerifying ? '验证中...' : '验证' }}
|
@click="handleVerify"
|
||||||
</Button>
|
>
|
||||||
|
{{ isVerifying ? '验证中...' : '验证' }}
|
||||||
|
</Button>
|
||||||
|
<Button
|
||||||
|
:disabled="isSaving || !canSave"
|
||||||
|
@click="handleSave"
|
||||||
|
>
|
||||||
|
{{ isSaving ? '保存中...' : '保存' }}
|
||||||
|
</Button>
|
||||||
|
<Button
|
||||||
|
variant="outline"
|
||||||
|
@click="$emit('update:open', false)"
|
||||||
|
>
|
||||||
|
取消
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</template>
|
</template>
|
||||||
</Dialog>
|
</Dialog>
|
||||||
</template>
|
</template>
|
||||||
@@ -281,8 +297,9 @@ import {
|
|||||||
SelectValue,
|
SelectValue,
|
||||||
Switch,
|
Switch,
|
||||||
} from '@/components/ui'
|
} from '@/components/ui'
|
||||||
import { saveProviderOpsConfig, verifyProviderAuth, getProviderOpsConfig } from '@/api/providerOps'
|
import { saveProviderOpsConfig, verifyProviderAuth, getProviderOpsConfig, deleteProviderOpsConfig } from '@/api/providerOps'
|
||||||
import { useToast } from '@/composables/useToast'
|
import { useToast } from '@/composables/useToast'
|
||||||
|
import { useConfirm } from '@/composables/useConfirm'
|
||||||
import {
|
import {
|
||||||
authTemplateRegistry,
|
authTemplateRegistry,
|
||||||
type AuthTemplate,
|
type AuthTemplate,
|
||||||
@@ -305,11 +322,13 @@ const emit = defineEmits<{
|
|||||||
const SENSITIVE_FIELDS = ['api_key', 'password', 'session_token', 'session_cookie', 'token_cookie', 'auth_cookie', 'cookie_string', 'cookie', 'proxy_password'] as const
|
const SENSITIVE_FIELDS = ['api_key', 'password', 'session_token', 'session_cookie', 'token_cookie', 'auth_cookie', 'cookie_string', 'cookie', 'proxy_password'] as const
|
||||||
|
|
||||||
const { success: showSuccess, error: showError } = useToast()
|
const { success: showSuccess, error: showError } = useToast()
|
||||||
|
const { confirmDanger } = useConfirm()
|
||||||
|
|
||||||
// State
|
// State
|
||||||
const isSaving = ref(false)
|
const isSaving = ref(false)
|
||||||
const isVerifying = ref(false)
|
const isVerifying = ref(false)
|
||||||
const isLoadingConfig = ref(false)
|
const isLoadingConfig = ref(false)
|
||||||
|
const isClearing = ref(false)
|
||||||
const verifyStatus = ref<'success' | 'error' | null>(null)
|
const verifyStatus = ref<'success' | 'error' | null>(null)
|
||||||
const formChanged = ref(false)
|
const formChanged = ref(false)
|
||||||
|
|
||||||
@@ -543,6 +562,40 @@ async function handleSave() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async function handleClear() {
|
||||||
|
if (!props.providerId) return
|
||||||
|
|
||||||
|
const confirmed = await confirmDanger(
|
||||||
|
'确定要清除该提供商的认证配置吗?清除后将无法进行余额查询、签到等操作。',
|
||||||
|
'清除认证',
|
||||||
|
'清除'
|
||||||
|
)
|
||||||
|
if (!confirmed) return
|
||||||
|
|
||||||
|
isClearing.value = true
|
||||||
|
try {
|
||||||
|
const result = await deleteProviderOpsConfig(props.providerId)
|
||||||
|
if (result.success) {
|
||||||
|
showSuccess(result.message || '认证信息已清除', '清除成功')
|
||||||
|
// 重置状态
|
||||||
|
hasExistingConfig.value = false
|
||||||
|
sensitivePlaceholders.value = {}
|
||||||
|
verifyStatus.value = null
|
||||||
|
formChanged.value = false
|
||||||
|
selectedTemplateId.value = 'new_api'
|
||||||
|
resetFormData()
|
||||||
|
emit('saved')
|
||||||
|
emit('update:open', false)
|
||||||
|
} else {
|
||||||
|
showError(result.message || '清除失败')
|
||||||
|
}
|
||||||
|
} catch (error: any) {
|
||||||
|
showError(error.response?.data?.detail || error.message, '清除失败')
|
||||||
|
} finally {
|
||||||
|
isClearing.value = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
function loadFromConfig(config: any) {
|
function loadFromConfig(config: any) {
|
||||||
if (!config?.connector) return
|
if (!config?.connector) return
|
||||||
|
|
||||||
|
|||||||
@@ -55,6 +55,15 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="flex items-center gap-1 shrink-0">
|
<div class="flex items-center gap-1 shrink-0">
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
:title="provider.enable_format_conversion ? '已启用格式转换(点击关闭)' : '启用格式转换'"
|
||||||
|
:class="provider.enable_format_conversion ? 'text-primary' : ''"
|
||||||
|
@click="toggleFormatConversion"
|
||||||
|
>
|
||||||
|
<Shuffle class="w-4 h-4" />
|
||||||
|
</Button>
|
||||||
<Button
|
<Button
|
||||||
variant="ghost"
|
variant="ghost"
|
||||||
size="icon"
|
size="icon"
|
||||||
@@ -502,7 +511,8 @@ import {
|
|||||||
Power,
|
Power,
|
||||||
GripVertical,
|
GripVertical,
|
||||||
Copy,
|
Copy,
|
||||||
Shield
|
Shield,
|
||||||
|
Shuffle
|
||||||
} from 'lucide-vue-next'
|
} from 'lucide-vue-next'
|
||||||
import { useEscapeKey } from '@/composables/useEscapeKey'
|
import { useEscapeKey } from '@/composables/useEscapeKey'
|
||||||
import Button from '@/components/ui/button.vue'
|
import Button from '@/components/ui/button.vue'
|
||||||
@@ -511,7 +521,7 @@ import Card from '@/components/ui/card.vue'
|
|||||||
import { useToast } from '@/composables/useToast'
|
import { useToast } from '@/composables/useToast'
|
||||||
import { useClipboard } from '@/composables/useClipboard'
|
import { useClipboard } from '@/composables/useClipboard'
|
||||||
import { useCountdownTimer, formatCountdown } from '@/composables/useCountdownTimer'
|
import { useCountdownTimer, formatCountdown } from '@/composables/useCountdownTimer'
|
||||||
import { getProvider, getProviderEndpoints } from '@/api/endpoints'
|
import { getProvider, getProviderEndpoints, updateProvider } from '@/api/endpoints'
|
||||||
import {
|
import {
|
||||||
KeyFormDialog,
|
KeyFormDialog,
|
||||||
KeyAllowedModelsEditDialog,
|
KeyAllowedModelsEditDialog,
|
||||||
@@ -702,6 +712,20 @@ function handleClose() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 切换格式转换开关
|
||||||
|
async function toggleFormatConversion() {
|
||||||
|
if (!provider.value) return
|
||||||
|
const newValue = !provider.value.enable_format_conversion
|
||||||
|
try {
|
||||||
|
await updateProvider(provider.value.id, { enable_format_conversion: newValue })
|
||||||
|
provider.value.enable_format_conversion = newValue
|
||||||
|
showSuccess(newValue ? '已启用格式转换' : '已禁用格式转换')
|
||||||
|
emit('refresh')
|
||||||
|
} catch {
|
||||||
|
showError('切换格式转换失败')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 显示端点管理对话框
|
// 显示端点管理对话框
|
||||||
function showAddEndpointDialog() {
|
function showAddEndpointDialog() {
|
||||||
endpointDialogOpen.value = true
|
endpointDialogOpen.value = true
|
||||||
|
|||||||
@@ -366,6 +366,8 @@ import {
|
|||||||
Mail,
|
Mail,
|
||||||
Puzzle,
|
Puzzle,
|
||||||
Video,
|
Video,
|
||||||
|
Zap,
|
||||||
|
FileUp,
|
||||||
type LucideIcon,
|
type LucideIcon,
|
||||||
} from 'lucide-vue-next'
|
} from 'lucide-vue-next'
|
||||||
|
|
||||||
@@ -511,6 +513,27 @@ const navigation = computed(() => {
|
|||||||
{ name: '邮件配置', href: '/admin/email', icon: Mail },
|
{ name: '邮件配置', href: '/admin/email', icon: Mail },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
// 动态添加已激活模块的菜单项
|
||||||
|
// 图标映射
|
||||||
|
const iconMap: Record<string, LucideIcon> = {
|
||||||
|
'Key': Key,
|
||||||
|
'FileUp': FileUp,
|
||||||
|
'Shield': Shield,
|
||||||
|
'Puzzle': Puzzle,
|
||||||
|
}
|
||||||
|
|
||||||
|
// 添加模块菜单项(按 admin_menu_order 排序,只显示已激活的)
|
||||||
|
const moduleMenuItems = Object.values(moduleStore.modules)
|
||||||
|
.filter(m => m.active && m.admin_route && m.admin_menu_group === 'system')
|
||||||
|
.sort((a, b) => a.admin_menu_order - b.admin_menu_order)
|
||||||
|
.map(m => ({
|
||||||
|
name: m.display_name,
|
||||||
|
href: m.admin_route!,
|
||||||
|
icon: iconMap[m.admin_menu_icon || ''] || Puzzle
|
||||||
|
}))
|
||||||
|
|
||||||
|
systemItems.push(...moduleMenuItems)
|
||||||
|
|
||||||
// 模块管理和系统设置放在最后
|
// 模块管理和系统设置放在最后
|
||||||
systemItems.push({ name: '模块管理', href: '/admin/modules', icon: Puzzle })
|
systemItems.push({ name: '模块管理', href: '/admin/modules', icon: Puzzle })
|
||||||
systemItems.push({ name: '系统设置', href: '/admin/system', icon: Cog })
|
systemItems.push({ name: '系统设置', href: '/admin/system', icon: Cog })
|
||||||
@@ -531,7 +554,7 @@ const navigation = computed(() => {
|
|||||||
{ name: '模型管理', href: '/admin/models', icon: Layers },
|
{ name: '模型管理', href: '/admin/models', icon: Layers },
|
||||||
{ name: '独立密钥', href: '/admin/keys', icon: Key },
|
{ name: '独立密钥', href: '/admin/keys', icon: Key },
|
||||||
{ name: '访问令牌', href: '/admin/management-tokens', icon: KeyRound },
|
{ name: '访问令牌', href: '/admin/management-tokens', icon: KeyRound },
|
||||||
{ name: '视频任务', href: '/admin/video-tasks', icon: Video },
|
{ name: '异步任务', href: '/admin/async-tasks', icon: Zap },
|
||||||
{ name: '使用记录', href: '/admin/usage', icon: BarChart3 },
|
{ name: '使用记录', href: '/admin/usage', icon: BarChart3 },
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
@@ -582,6 +605,18 @@ const breadcrumbs = computed((): BreadcrumbItem[] => {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Special case: module pages not in navigation (module not active)
|
||||||
|
// Check if current path matches a module's admin_route
|
||||||
|
const currentModule = Object.values(moduleStore.modules).find(
|
||||||
|
m => m.admin_route && route.path === m.admin_route
|
||||||
|
)
|
||||||
|
if (currentModule) {
|
||||||
|
return [
|
||||||
|
{ label: '模块管理', href: '/admin/modules' },
|
||||||
|
{ label: currentModule.display_name }
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
return [{ label: '仪表盘' }]
|
return [{ label: '仪表盘' }]
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
@@ -208,10 +208,20 @@ const routes: RouteRecordRaw[] = [
|
|||||||
name: 'AnnouncementManagement',
|
name: 'AnnouncementManagement',
|
||||||
component: () => importWithRetry(() => import('@/views/user/Announcements.vue'))
|
component: () => importWithRetry(() => import('@/views/user/Announcements.vue'))
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
path: 'async-tasks',
|
||||||
|
name: 'AsyncTasks',
|
||||||
|
component: () => importWithRetry(() => import('@/views/admin/AsyncTasks.vue'))
|
||||||
|
},
|
||||||
|
{
|
||||||
|
path: 'gemini-files',
|
||||||
|
name: 'GeminiFilesManagement',
|
||||||
|
component: () => importWithRetry(() => import('@/views/admin/GeminiFilesManagement.vue'))
|
||||||
|
},
|
||||||
|
// 保留旧路由兼容性
|
||||||
{
|
{
|
||||||
path: 'video-tasks',
|
path: 'video-tasks',
|
||||||
name: 'VideoTasks',
|
redirect: '/admin/async-tasks'
|
||||||
component: () => importWithRetry(() => import('@/views/admin/VideoTasks.vue'))
|
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,7 +5,7 @@
|
|||||||
<Card variant="default" class="p-4">
|
<Card variant="default" class="p-4">
|
||||||
<div class="flex items-center gap-3">
|
<div class="flex items-center gap-3">
|
||||||
<div class="w-10 h-10 rounded-lg bg-primary/10 flex items-center justify-center">
|
<div class="w-10 h-10 rounded-lg bg-primary/10 flex items-center justify-center">
|
||||||
<Video class="w-5 h-5 text-primary" />
|
<Zap class="w-5 h-5 text-primary" />
|
||||||
</div>
|
</div>
|
||||||
<div>
|
<div>
|
||||||
<p class="text-2xl font-bold">{{ stats?.total ?? '-' }}</p>
|
<p class="text-2xl font-bold">{{ stats?.total ?? '-' }}</p>
|
||||||
@@ -53,7 +53,7 @@
|
|||||||
<!-- 标题和筛选器 -->
|
<!-- 标题和筛选器 -->
|
||||||
<div class="px-4 sm:px-6 py-3.5 border-b border-border/60">
|
<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">
|
<div class="flex flex-col gap-3 sm:flex-row sm:items-center sm:justify-between">
|
||||||
<h3 class="text-base font-semibold">视频任务</h3>
|
<h3 class="text-base font-semibold">异步任务</h3>
|
||||||
<div class="flex items-center gap-2">
|
<div class="flex items-center gap-2">
|
||||||
<!-- 状态筛选 -->
|
<!-- 状态筛选 -->
|
||||||
<Select v-model="filterStatus">
|
<Select v-model="filterStatus">
|
||||||
@@ -98,8 +98,8 @@
|
|||||||
|
|
||||||
<!-- 空状态 -->
|
<!-- 空状态 -->
|
||||||
<div v-else-if="!tasks.length" class="p-8 text-center">
|
<div v-else-if="!tasks.length" class="p-8 text-center">
|
||||||
<Video class="w-12 h-12 mx-auto text-muted-foreground/50" />
|
<Zap class="w-12 h-12 mx-auto text-muted-foreground/50" />
|
||||||
<p class="mt-2 text-sm text-muted-foreground">暂无视频任务</p>
|
<p class="mt-2 text-sm text-muted-foreground">暂无异步任务</p>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 任务列表 -->
|
<!-- 任务列表 -->
|
||||||
@@ -114,6 +114,7 @@
|
|||||||
<div class="flex-1 min-w-0">
|
<div class="flex-1 min-w-0">
|
||||||
<!-- 模型和状态 -->
|
<!-- 模型和状态 -->
|
||||||
<div class="flex items-center gap-2 mb-1">
|
<div class="flex items-center gap-2 mb-1">
|
||||||
|
<Video v-if="isVideoTask(task)" class="w-4 h-4 text-muted-foreground" />
|
||||||
<span class="font-medium text-sm">{{ task.model }}</span>
|
<span class="font-medium text-sm">{{ task.model }}</span>
|
||||||
<Badge :variant="getStatusVariant(task.status)">
|
<Badge :variant="getStatusVariant(task.status)">
|
||||||
{{ getStatusLabel(task.status) }}
|
{{ getStatusLabel(task.status) }}
|
||||||
@@ -389,18 +390,86 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 视频结果 -->
|
<!-- 视频结果 -->
|
||||||
<div v-if="selectedTask.video_url" class="space-y-3">
|
<div v-if="selectedTask.status === 'completed' || selectedTask.video_url || selectedTask.video_urls?.length" class="space-y-3">
|
||||||
<h4 class="text-sm font-medium">视频结果</h4>
|
<h4 class="text-sm font-medium flex items-center gap-2">
|
||||||
<div class="space-y-2">
|
<Video class="w-4 h-4" />
|
||||||
|
视频结果
|
||||||
|
</h4>
|
||||||
|
|
||||||
|
<!-- 主视频 -->
|
||||||
|
<div v-if="selectedTask.video_url" class="space-y-2">
|
||||||
<video
|
<video
|
||||||
:src="selectedTask.video_url"
|
:src="selectedTask.video_url"
|
||||||
controls
|
controls
|
||||||
class="w-full rounded-lg"
|
class="w-full rounded-lg"
|
||||||
/>
|
/>
|
||||||
|
<!-- 视频链接 -->
|
||||||
|
<div class="p-2 bg-muted/50 rounded text-xs">
|
||||||
|
<div class="flex items-center justify-between gap-2">
|
||||||
|
<span class="text-muted-foreground truncate flex-1" :title="selectedTask.video_url">
|
||||||
|
{{ selectedTask.video_url }}
|
||||||
|
</span>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="sm"
|
||||||
|
class="h-6 px-2 text-xs"
|
||||||
|
@click="copyToClipboard(selectedTask.video_url)"
|
||||||
|
>
|
||||||
|
复制链接
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
<p v-if="selectedTask.video_expires_at" class="text-xs text-muted-foreground">
|
<p v-if="selectedTask.video_expires_at" class="text-xs text-muted-foreground">
|
||||||
过期时间: {{ formatDate(selectedTask.video_expires_at) }}
|
过期时间: {{ formatDate(selectedTask.video_expires_at) }}
|
||||||
</p>
|
</p>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
<!-- 多个视频(如果有) -->
|
||||||
|
<div v-else-if="selectedTask.video_urls?.length" class="space-y-3">
|
||||||
|
<div v-for="(url, index) in selectedTask.video_urls" :key="index" class="space-y-2">
|
||||||
|
<p class="text-xs text-muted-foreground">视频 {{ index + 1 }}</p>
|
||||||
|
<video :src="url" controls class="w-full rounded-lg" />
|
||||||
|
<div class="p-2 bg-muted/50 rounded text-xs">
|
||||||
|
<div class="flex items-center justify-between gap-2">
|
||||||
|
<span class="text-muted-foreground truncate flex-1" :title="url">{{ url }}</span>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="sm"
|
||||||
|
class="h-6 px-2 text-xs"
|
||||||
|
@click="copyToClipboard(url)"
|
||||||
|
>
|
||||||
|
复制链接
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- 任务完成但无视频 -->
|
||||||
|
<div v-else-if="selectedTask.status === 'completed'" class="p-4 bg-amber-50 dark:bg-amber-900/20 rounded-lg text-center">
|
||||||
|
<p class="text-sm text-amber-600 dark:text-amber-400">任务已完成,但视频链接不可用或已过期</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- 任务完成响应体 -->
|
||||||
|
<div v-if="selectedTask.request_metadata?.poll_raw_response" class="space-y-3">
|
||||||
|
<div class="flex items-center justify-between">
|
||||||
|
<h4 class="text-sm font-medium flex items-center gap-2">
|
||||||
|
<FileJson class="w-4 h-4" />
|
||||||
|
任务响应
|
||||||
|
</h4>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="sm"
|
||||||
|
class="h-6 px-2 text-xs"
|
||||||
|
@click="copyToClipboard(JSON.stringify(selectedTask.request_metadata.poll_raw_response, null, 2))"
|
||||||
|
>
|
||||||
|
复制
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
<div class="p-3 bg-muted/50 rounded-lg overflow-x-auto">
|
||||||
|
<pre class="text-xs font-mono whitespace-pre-wrap break-all">{{ formatJson(selectedTask.request_metadata.poll_raw_response) }}</pre>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 操作按钮 -->
|
<!-- 操作按钮 -->
|
||||||
@@ -447,7 +516,7 @@
|
|||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, computed, onMounted, watch } from 'vue'
|
import { ref, computed, onMounted, watch } from 'vue'
|
||||||
import { videoTasksApi, type VideoTaskItem, type VideoTaskDetail, type VideoTaskStatsResponse, type VideoTaskStatus } from '@/api/video-tasks'
|
import { asyncTasksApi, type AsyncTaskItem, type AsyncTaskDetail, type AsyncTaskStatsResponse, type AsyncTaskStatus } from '@/api/async-tasks'
|
||||||
import { useToast } from '@/composables/useToast'
|
import { useToast } from '@/composables/useToast'
|
||||||
import Card from '@/components/ui/card.vue'
|
import Card from '@/components/ui/card.vue'
|
||||||
import Button from '@/components/ui/button.vue'
|
import Button from '@/components/ui/button.vue'
|
||||||
@@ -459,8 +528,10 @@ import SelectValue from '@/components/ui/select-value.vue'
|
|||||||
import SelectContent from '@/components/ui/select-content.vue'
|
import SelectContent from '@/components/ui/select-content.vue'
|
||||||
import SelectItem from '@/components/ui/select-item.vue'
|
import SelectItem from '@/components/ui/select-item.vue'
|
||||||
import {
|
import {
|
||||||
|
Zap,
|
||||||
Video,
|
Video,
|
||||||
Loader2,
|
Loader2,
|
||||||
|
FileJson,
|
||||||
CheckCircle,
|
CheckCircle,
|
||||||
Calendar,
|
Calendar,
|
||||||
RefreshCw,
|
RefreshCw,
|
||||||
@@ -478,24 +549,29 @@ const { toast } = useToast()
|
|||||||
|
|
||||||
// 状态
|
// 状态
|
||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
const tasks = ref<VideoTaskItem[]>([])
|
const tasks = ref<AsyncTaskItem[]>([])
|
||||||
const stats = ref<VideoTaskStatsResponse | null>(null)
|
const stats = ref<AsyncTaskStatsResponse | null>(null)
|
||||||
const total = ref(0)
|
const total = ref(0)
|
||||||
const currentPage = ref(1)
|
const currentPage = ref(1)
|
||||||
const pageSize = ref(20)
|
const pageSize = ref(20)
|
||||||
const filterStatus = ref('all')
|
const filterStatus = ref('all')
|
||||||
const filterModel = ref('')
|
const filterModel = ref('')
|
||||||
const showDetail = ref(false)
|
const showDetail = ref(false)
|
||||||
const selectedTask = ref<VideoTaskDetail | null>(null)
|
const selectedTask = ref<AsyncTaskDetail | null>(null)
|
||||||
|
|
||||||
const totalPages = computed(() => Math.ceil(total.value / pageSize.value))
|
const totalPages = computed(() => Math.ceil(total.value / pageSize.value))
|
||||||
|
|
||||||
|
// 判断是否为视频任务
|
||||||
|
function isVideoTask(task: AsyncTaskItem): boolean {
|
||||||
|
return task.task_type === 'video' || !!task.video_url || !!task.duration_seconds
|
||||||
|
}
|
||||||
|
|
||||||
// 获取任务列表
|
// 获取任务列表
|
||||||
async function fetchTasks() {
|
async function fetchTasks() {
|
||||||
loading.value = true
|
loading.value = true
|
||||||
try {
|
try {
|
||||||
const response = await videoTasksApi.list({
|
const response = await asyncTasksApi.list({
|
||||||
status: filterStatus.value !== 'all' ? filterStatus.value as VideoTaskStatus : undefined,
|
status: filterStatus.value !== 'all' ? filterStatus.value as AsyncTaskStatus : undefined,
|
||||||
model: filterModel.value || undefined,
|
model: filterModel.value || undefined,
|
||||||
page: currentPage.value,
|
page: currentPage.value,
|
||||||
page_size: pageSize.value,
|
page_size: pageSize.value,
|
||||||
@@ -516,16 +592,16 @@ async function fetchTasks() {
|
|||||||
// 获取统计数据
|
// 获取统计数据
|
||||||
async function fetchStats() {
|
async function fetchStats() {
|
||||||
try {
|
try {
|
||||||
stats.value = await videoTasksApi.getStats()
|
stats.value = await asyncTasksApi.getStats()
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Failed to fetch stats:', error)
|
console.error('Failed to fetch stats:', error)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 打开任务详情
|
// 打开任务详情
|
||||||
async function openTaskDetail(task: VideoTaskItem) {
|
async function openTaskDetail(task: AsyncTaskItem) {
|
||||||
try {
|
try {
|
||||||
selectedTask.value = await videoTasksApi.getDetail(task.id)
|
selectedTask.value = await asyncTasksApi.getDetail(task.id)
|
||||||
showDetail.value = true
|
showDetail.value = true
|
||||||
} catch (error: any) {
|
} catch (error: any) {
|
||||||
toast({
|
toast({
|
||||||
@@ -537,10 +613,10 @@ async function openTaskDetail(task: VideoTaskItem) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 取消任务
|
// 取消任务
|
||||||
async function cancelTask(task: VideoTaskItem | VideoTaskDetail) {
|
async function cancelTask(task: AsyncTaskItem | AsyncTaskDetail) {
|
||||||
if (!confirm('确定要取消这个任务吗?')) return
|
if (!confirm('确定要取消这个任务吗?')) return
|
||||||
try {
|
try {
|
||||||
await videoTasksApi.cancel(task.id)
|
await asyncTasksApi.cancel(task.id)
|
||||||
toast({
|
toast({
|
||||||
title: '任务已取消',
|
title: '任务已取消',
|
||||||
})
|
})
|
||||||
@@ -602,6 +678,30 @@ function formatDate(dateStr: string | null): string {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 复制到剪贴板
|
||||||
|
async function copyToClipboard(text: string) {
|
||||||
|
try {
|
||||||
|
await navigator.clipboard.writeText(text)
|
||||||
|
toast({
|
||||||
|
title: '已复制到剪贴板',
|
||||||
|
})
|
||||||
|
} catch (error) {
|
||||||
|
toast({
|
||||||
|
title: '复制失败',
|
||||||
|
variant: 'destructive',
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 格式化 JSON
|
||||||
|
function formatJson(obj: any): string {
|
||||||
|
try {
|
||||||
|
return JSON.stringify(obj, null, 2)
|
||||||
|
} catch {
|
||||||
|
return String(obj)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 分页
|
// 分页
|
||||||
function goToPage(page: number) {
|
function goToPage(page: number) {
|
||||||
currentPage.value = page
|
currentPage.value = page
|
||||||
570
frontend/src/views/admin/GeminiFilesManagement.vue
Normal file
570
frontend/src/views/admin/GeminiFilesManagement.vue
Normal file
@@ -0,0 +1,570 @@
|
|||||||
|
<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">
|
||||||
|
<FileUp class="w-5 h-5 text-primary" />
|
||||||
|
</div>
|
||||||
|
<div>
|
||||||
|
<p class="text-2xl font-bold">{{ stats?.total_mappings ?? '-' }}</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?.active_mappings ?? '-' }}</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">
|
||||||
|
<Clock class="w-5 h-5 text-amber-500" />
|
||||||
|
</div>
|
||||||
|
<div>
|
||||||
|
<p class="text-2xl font-bold">{{ stats?.expired_mappings ?? '-' }}</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">
|
||||||
|
<Key class="w-5 h-5 text-blue-500" />
|
||||||
|
</div>
|
||||||
|
<div>
|
||||||
|
<p class="text-2xl font-bold">{{ stats?.capable_keys_count ?? '-' }}</p>
|
||||||
|
<p class="text-xs text-muted-foreground">支持的 Key</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</Card>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- 上传区域 -->
|
||||||
|
<Card variant="default" class="p-4">
|
||||||
|
<div class="flex items-center justify-between mb-3">
|
||||||
|
<h3 class="text-sm font-medium">上传文件</h3>
|
||||||
|
<Button
|
||||||
|
v-if="capableKeys.length > 0"
|
||||||
|
variant="ghost"
|
||||||
|
size="sm"
|
||||||
|
class="h-7 text-xs"
|
||||||
|
@click="toggleSelectAll"
|
||||||
|
>
|
||||||
|
{{ selectedKeyIds.length === capableKeys.length ? '取消全选' : '全选' }}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- Key 选择器 -->
|
||||||
|
<div v-if="capableKeys.length > 0" class="mb-4">
|
||||||
|
<p class="text-xs text-muted-foreground mb-2">选择要上传到的 Key(可多选):</p>
|
||||||
|
<div class="flex flex-wrap gap-2">
|
||||||
|
<button
|
||||||
|
v-for="key in capableKeys"
|
||||||
|
:key="key.id"
|
||||||
|
class="px-3 py-1.5 text-xs rounded-lg border transition-colors"
|
||||||
|
:class="selectedKeyIds.includes(key.id)
|
||||||
|
? 'border-primary bg-primary/10 text-primary'
|
||||||
|
: 'border-border hover:border-primary/50'"
|
||||||
|
@click="toggleKeySelection(key.id)"
|
||||||
|
>
|
||||||
|
<span class="font-medium">{{ key.name }}</span>
|
||||||
|
<span v-if="key.provider_name" class="text-muted-foreground ml-1">({{ key.provider_name }})</span>
|
||||||
|
</button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div v-else class="mb-4 text-sm text-amber-600 bg-amber-50 dark:bg-amber-950/30 rounded-lg p-3">
|
||||||
|
暂无可用的 Key,请先配置具有「Gemini 文件 API」能力的 Key
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- 拖拽上传区 -->
|
||||||
|
<div
|
||||||
|
class="border-2 border-dashed border-border/60 rounded-lg p-6 text-center transition-colors"
|
||||||
|
:class="{
|
||||||
|
'border-primary bg-primary/5': isDragging,
|
||||||
|
'hover:border-primary/50': !isDragging && !uploading && selectedKeyIds.length > 0,
|
||||||
|
'opacity-50 cursor-not-allowed': selectedKeyIds.length === 0
|
||||||
|
}"
|
||||||
|
@dragover.prevent="isDragging = true"
|
||||||
|
@dragleave.prevent="isDragging = false"
|
||||||
|
@drop.prevent="handleDrop"
|
||||||
|
>
|
||||||
|
<input
|
||||||
|
ref="fileInputRef"
|
||||||
|
type="file"
|
||||||
|
class="hidden"
|
||||||
|
@change="handleFileSelect"
|
||||||
|
/>
|
||||||
|
<div v-if="uploading" class="flex flex-col items-center gap-2">
|
||||||
|
<Loader2 class="w-8 h-8 animate-spin text-primary" />
|
||||||
|
<p class="text-sm text-muted-foreground">正在上传到 {{ selectedKeyIds.length }} 个 Key...</p>
|
||||||
|
</div>
|
||||||
|
<div v-else class="flex flex-col items-center gap-2">
|
||||||
|
<Upload class="w-8 h-8 text-muted-foreground" />
|
||||||
|
<p class="text-sm text-muted-foreground">
|
||||||
|
<template v-if="selectedKeyIds.length > 0">
|
||||||
|
拖拽文件到此处,或
|
||||||
|
<button
|
||||||
|
class="text-primary hover:underline"
|
||||||
|
@click="fileInputRef?.click()"
|
||||||
|
>
|
||||||
|
点击选择
|
||||||
|
</button>
|
||||||
|
</template>
|
||||||
|
<template v-else>
|
||||||
|
请先选择至少一个 Key
|
||||||
|
</template>
|
||||||
|
</p>
|
||||||
|
<p class="text-xs text-muted-foreground">
|
||||||
|
支持视频、图片、音频、文档等,最大 2GB,有效期 48 小时
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</Card>
|
||||||
|
|
||||||
|
<!-- MIME 类型分布 -->
|
||||||
|
<Card v-if="stats?.by_mime_type && Object.keys(stats.by_mime_type).length > 0" variant="default" class="p-4">
|
||||||
|
<h3 class="text-sm font-medium mb-3">文件类型分布</h3>
|
||||||
|
<div class="flex flex-wrap gap-2">
|
||||||
|
<Badge
|
||||||
|
v-for="(count, mimeType) in stats.by_mime_type"
|
||||||
|
:key="mimeType"
|
||||||
|
variant="secondary"
|
||||||
|
class="text-xs"
|
||||||
|
>
|
||||||
|
{{ mimeType }}: {{ count }}
|
||||||
|
</Badge>
|
||||||
|
</div>
|
||||||
|
</Card>
|
||||||
|
|
||||||
|
<!-- 文件映射表格 -->
|
||||||
|
<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">
|
||||||
|
<!-- 搜索 -->
|
||||||
|
<Input
|
||||||
|
v-model="searchQuery"
|
||||||
|
type="text"
|
||||||
|
placeholder="搜索文件名..."
|
||||||
|
class="w-40 h-8 text-xs"
|
||||||
|
/>
|
||||||
|
<!-- 包含过期 -->
|
||||||
|
<label class="flex items-center gap-1.5 text-xs text-muted-foreground cursor-pointer">
|
||||||
|
<input
|
||||||
|
v-model="includeExpired"
|
||||||
|
type="checkbox"
|
||||||
|
class="rounded border-border"
|
||||||
|
/>
|
||||||
|
包含过期
|
||||||
|
</label>
|
||||||
|
<!-- 清理过期按钮 -->
|
||||||
|
<Button
|
||||||
|
variant="outline"
|
||||||
|
size="sm"
|
||||||
|
class="h-8 text-xs"
|
||||||
|
:disabled="loading || (stats?.expired_mappings ?? 0) === 0"
|
||||||
|
@click="cleanupExpired"
|
||||||
|
>
|
||||||
|
<Trash2 class="w-3 h-3 mr-1" />
|
||||||
|
清理过期
|
||||||
|
</Button>
|
||||||
|
<!-- 刷新按钮 -->
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8"
|
||||||
|
:disabled="loading"
|
||||||
|
@click="fetchData"
|
||||||
|
>
|
||||||
|
<RefreshCw class="w-3.5 h-3.5" :class="{ 'animate-spin': loading }" />
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- 加载状态 -->
|
||||||
|
<div v-if="loading && !mappings.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="!mappings.length" class="p-8 text-center">
|
||||||
|
<FileUp class="w-12 h-12 mx-auto text-muted-foreground/50" />
|
||||||
|
<p class="mt-2 text-sm text-muted-foreground">暂无文件映射</p>
|
||||||
|
<p class="mt-1 text-xs text-muted-foreground">
|
||||||
|
用户通过 Gemini Files API 上传文件后会在此显示
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<!-- 文件列表 -->
|
||||||
|
<div v-else class="divide-y divide-border/60">
|
||||||
|
<div
|
||||||
|
v-for="mapping in mappings"
|
||||||
|
:key="mapping.id"
|
||||||
|
class="px-4 sm:px-6 py-4 hover:bg-muted/30 transition-colors"
|
||||||
|
:class="{ 'opacity-50': mapping.is_expired }"
|
||||||
|
>
|
||||||
|
<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">
|
||||||
|
<component :is="getFileIcon(mapping.mime_type)" class="w-4 h-4 text-muted-foreground" />
|
||||||
|
<span class="font-mono text-sm font-medium">{{ mapping.file_name }}</span>
|
||||||
|
<Badge v-if="mapping.is_expired" variant="secondary" class="text-xs">
|
||||||
|
已过期
|
||||||
|
</Badge>
|
||||||
|
<Badge v-else variant="outline" class="text-xs text-green-600">
|
||||||
|
有效
|
||||||
|
</Badge>
|
||||||
|
</div>
|
||||||
|
<!-- 显示名 -->
|
||||||
|
<p v-if="mapping.display_name" class="text-sm text-muted-foreground truncate">
|
||||||
|
{{ mapping.display_name }}
|
||||||
|
</p>
|
||||||
|
<!-- 元信息 -->
|
||||||
|
<div class="flex items-center gap-4 mt-2 text-xs text-muted-foreground">
|
||||||
|
<span v-if="mapping.mime_type" class="flex items-center gap-1">
|
||||||
|
<File class="w-3 h-3" />
|
||||||
|
{{ mapping.mime_type }}
|
||||||
|
</span>
|
||||||
|
<span v-if="mapping.username" class="flex items-center gap-1">
|
||||||
|
<User class="w-3 h-3" />
|
||||||
|
{{ mapping.username }}
|
||||||
|
</span>
|
||||||
|
<span v-if="mapping.key_name" class="flex items-center gap-1">
|
||||||
|
<Key class="w-3 h-3" />
|
||||||
|
{{ mapping.key_name }}
|
||||||
|
</span>
|
||||||
|
<span class="flex items-center gap-1">
|
||||||
|
<Clock class="w-3 h-3" />
|
||||||
|
{{ formatDate(mapping.created_at) }}
|
||||||
|
</span>
|
||||||
|
<span class="flex items-center gap-1" :class="{ 'text-red-500': mapping.is_expired }">
|
||||||
|
<Timer class="w-3 h-3" />
|
||||||
|
过期: {{ formatDate(mapping.expires_at) }}
|
||||||
|
</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<!-- 操作 -->
|
||||||
|
<div class="flex items-center gap-2">
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8 text-muted-foreground hover:text-red-500"
|
||||||
|
title="删除映射"
|
||||||
|
@click.stop="deleteMapping(mapping)"
|
||||||
|
>
|
||||||
|
<Trash2 class="w-4 h-4" />
|
||||||
|
</Button>
|
||||||
|
</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 }} 条记录
|
||||||
|
</p>
|
||||||
|
<div class="flex items-center gap-2">
|
||||||
|
<Button
|
||||||
|
variant="outline"
|
||||||
|
size="sm"
|
||||||
|
:disabled="currentPage <= 1"
|
||||||
|
@click="currentPage--"
|
||||||
|
>
|
||||||
|
上一页
|
||||||
|
</Button>
|
||||||
|
<span class="text-sm text-muted-foreground">
|
||||||
|
{{ currentPage }} / {{ totalPages }}
|
||||||
|
</span>
|
||||||
|
<Button
|
||||||
|
variant="outline"
|
||||||
|
size="sm"
|
||||||
|
:disabled="currentPage >= totalPages"
|
||||||
|
@click="currentPage++"
|
||||||
|
>
|
||||||
|
下一页
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</Card>
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup lang="ts">
|
||||||
|
import { ref, computed, watch, onMounted } from 'vue'
|
||||||
|
import { useToast } from '@/composables/useToast'
|
||||||
|
import Card from '@/components/ui/card.vue'
|
||||||
|
import Badge from '@/components/ui/badge.vue'
|
||||||
|
import Button from '@/components/ui/button.vue'
|
||||||
|
import Input from '@/components/ui/input.vue'
|
||||||
|
import {
|
||||||
|
FileUp,
|
||||||
|
CheckCircle,
|
||||||
|
Clock,
|
||||||
|
Key,
|
||||||
|
RefreshCw,
|
||||||
|
Loader2,
|
||||||
|
Trash2,
|
||||||
|
File,
|
||||||
|
User,
|
||||||
|
Timer,
|
||||||
|
Video,
|
||||||
|
Image,
|
||||||
|
FileText,
|
||||||
|
Music,
|
||||||
|
Upload
|
||||||
|
} from 'lucide-vue-next'
|
||||||
|
import { geminiFilesApi } from '@/api/gemini-files'
|
||||||
|
|
||||||
|
const { toast } = useToast()
|
||||||
|
|
||||||
|
// 状态
|
||||||
|
const loading = ref(false)
|
||||||
|
const stats = ref<any>(null)
|
||||||
|
const mappings = ref<any[]>([])
|
||||||
|
const total = ref(0)
|
||||||
|
const currentPage = ref(1)
|
||||||
|
const pageSize = 20
|
||||||
|
const searchQuery = ref('')
|
||||||
|
const includeExpired = ref(false)
|
||||||
|
|
||||||
|
// 上传状态
|
||||||
|
const uploading = ref(false)
|
||||||
|
const isDragging = ref(false)
|
||||||
|
const fileInputRef = ref<HTMLInputElement | null>(null)
|
||||||
|
const capableKeys = ref<any[]>([])
|
||||||
|
const selectedKeyIds = ref<string[]>([])
|
||||||
|
|
||||||
|
// 计算属性
|
||||||
|
const totalPages = computed(() => Math.ceil(total.value / pageSize))
|
||||||
|
|
||||||
|
// 监听筛选条件变化
|
||||||
|
watch([searchQuery, includeExpired], () => {
|
||||||
|
currentPage.value = 1
|
||||||
|
fetchMappings()
|
||||||
|
})
|
||||||
|
|
||||||
|
watch(currentPage, () => {
|
||||||
|
fetchMappings()
|
||||||
|
})
|
||||||
|
|
||||||
|
// 获取数据
|
||||||
|
async function fetchData() {
|
||||||
|
await Promise.all([fetchStats(), fetchMappings(), fetchCapableKeys()])
|
||||||
|
}
|
||||||
|
|
||||||
|
async function fetchCapableKeys() {
|
||||||
|
try {
|
||||||
|
const keys = await geminiFilesApi.getCapableKeys()
|
||||||
|
capableKeys.value = keys
|
||||||
|
// 默认全选
|
||||||
|
if (selectedKeyIds.value.length === 0 && keys.length > 0) {
|
||||||
|
selectedKeyIds.value = keys.map(k => k.id)
|
||||||
|
}
|
||||||
|
} catch (error: any) {
|
||||||
|
console.error('Failed to fetch capable keys:', error)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function toggleKeySelection(keyId: string) {
|
||||||
|
const index = selectedKeyIds.value.indexOf(keyId)
|
||||||
|
if (index === -1) {
|
||||||
|
selectedKeyIds.value.push(keyId)
|
||||||
|
} else {
|
||||||
|
selectedKeyIds.value.splice(index, 1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function toggleSelectAll() {
|
||||||
|
if (selectedKeyIds.value.length === capableKeys.value.length) {
|
||||||
|
selectedKeyIds.value = []
|
||||||
|
} else {
|
||||||
|
selectedKeyIds.value = capableKeys.value.map(k => k.id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function fetchStats() {
|
||||||
|
try {
|
||||||
|
const data = await geminiFilesApi.getStats()
|
||||||
|
stats.value = data
|
||||||
|
} catch (error: any) {
|
||||||
|
toast({
|
||||||
|
title: '获取统计失败',
|
||||||
|
description: error.message,
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function fetchMappings() {
|
||||||
|
loading.value = true
|
||||||
|
try {
|
||||||
|
const data = await geminiFilesApi.listMappings({
|
||||||
|
page: currentPage.value,
|
||||||
|
page_size: pageSize,
|
||||||
|
include_expired: includeExpired.value,
|
||||||
|
search: searchQuery.value || undefined
|
||||||
|
})
|
||||||
|
mappings.value = data.items
|
||||||
|
total.value = data.total
|
||||||
|
} catch (error: any) {
|
||||||
|
toast({
|
||||||
|
title: '获取文件列表失败',
|
||||||
|
description: error.message,
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
} finally {
|
||||||
|
loading.value = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function deleteMapping(mapping: any) {
|
||||||
|
if (!confirm(`确定要删除映射 "${mapping.file_name}" 吗?\n\n注意:这只会删除映射记录,不会删除 Google 上的实际文件。`)) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
await geminiFilesApi.deleteMapping(mapping.id)
|
||||||
|
toast({
|
||||||
|
title: '删除成功',
|
||||||
|
description: `已删除映射 ${mapping.file_name}`
|
||||||
|
})
|
||||||
|
await fetchData()
|
||||||
|
} catch (error: any) {
|
||||||
|
toast({
|
||||||
|
title: '删除失败',
|
||||||
|
description: error.message,
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function cleanupExpired() {
|
||||||
|
if (!confirm('确定要清理所有过期的文件映射吗?')) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
const result = await geminiFilesApi.cleanupExpired()
|
||||||
|
toast({
|
||||||
|
title: '清理完成',
|
||||||
|
description: `已清理 ${result.deleted_count} 条过期映射`
|
||||||
|
})
|
||||||
|
await fetchData()
|
||||||
|
} catch (error: any) {
|
||||||
|
toast({
|
||||||
|
title: '清理失败',
|
||||||
|
description: error.message,
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 上传相关
|
||||||
|
async function uploadFile(file: globalThis.File) {
|
||||||
|
if (selectedKeyIds.value.length === 0) {
|
||||||
|
toast({
|
||||||
|
title: '请选择 Key',
|
||||||
|
description: '请至少选择一个 Key 来上传文件',
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
uploading.value = true
|
||||||
|
let hasSuccess = false
|
||||||
|
try {
|
||||||
|
const result = await geminiFilesApi.uploadFile(file, selectedKeyIds.value)
|
||||||
|
if (result.fail_count === 0) {
|
||||||
|
toast({
|
||||||
|
title: '上传成功',
|
||||||
|
description: `文件 ${result.display_name} 已上传到 ${result.success_count} 个 Key`
|
||||||
|
})
|
||||||
|
hasSuccess = true
|
||||||
|
} else if (result.success_count > 0) {
|
||||||
|
toast({
|
||||||
|
title: '部分成功',
|
||||||
|
description: `成功 ${result.success_count} 个,失败 ${result.fail_count} 个`
|
||||||
|
})
|
||||||
|
hasSuccess = true
|
||||||
|
} else {
|
||||||
|
const errors = result.results.map(r => r.error).filter(Boolean).join('; ')
|
||||||
|
toast({
|
||||||
|
title: '上传失败',
|
||||||
|
description: errors || '所有 Key 上传都失败了',
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
}
|
||||||
|
} catch (error: any) {
|
||||||
|
toast({
|
||||||
|
title: '上传失败',
|
||||||
|
description: error.response?.data?.detail || error.message,
|
||||||
|
variant: 'destructive'
|
||||||
|
})
|
||||||
|
} finally {
|
||||||
|
uploading.value = false
|
||||||
|
isDragging.value = false
|
||||||
|
// 有成功上传时刷新列表,并重置到第一页
|
||||||
|
if (hasSuccess) {
|
||||||
|
currentPage.value = 1
|
||||||
|
await fetchData()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleDrop(e: DragEvent) {
|
||||||
|
isDragging.value = false
|
||||||
|
const files = e.dataTransfer?.files
|
||||||
|
if (files && files.length > 0) {
|
||||||
|
uploadFile(files[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleFileSelect(e: Event) {
|
||||||
|
const input = e.target as HTMLInputElement
|
||||||
|
if (input.files && input.files.length > 0) {
|
||||||
|
uploadFile(input.files[0])
|
||||||
|
input.value = '' // 清空以便重复选择同一文件
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 工具函数
|
||||||
|
function formatDate(dateStr: 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 getFileIcon(mimeType: string | null) {
|
||||||
|
if (!mimeType) return File
|
||||||
|
if (mimeType.startsWith('video/')) return Video
|
||||||
|
if (mimeType.startsWith('image/')) return Image
|
||||||
|
if (mimeType.startsWith('audio/')) return Music
|
||||||
|
if (mimeType.startsWith('text/') || mimeType.includes('pdf')) return FileText
|
||||||
|
return File
|
||||||
|
}
|
||||||
|
|
||||||
|
// 初始化
|
||||||
|
onMounted(() => {
|
||||||
|
fetchData()
|
||||||
|
})
|
||||||
|
</script>
|
||||||
@@ -187,6 +187,26 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
<div class="flex items-center h-full">
|
||||||
|
<div class="flex items-center space-x-2">
|
||||||
|
<Checkbox
|
||||||
|
id="enable-format-conversion"
|
||||||
|
v-model:checked="systemConfig.enable_format_conversion"
|
||||||
|
/>
|
||||||
|
<div>
|
||||||
|
<Label
|
||||||
|
for="enable-format-conversion"
|
||||||
|
class="cursor-pointer"
|
||||||
|
>
|
||||||
|
全局格式转换
|
||||||
|
</Label>
|
||||||
|
<p class="text-xs text-muted-foreground">
|
||||||
|
开启后强制允许所有提供商接受跨格式请求
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</CardSection>
|
</CardSection>
|
||||||
|
|
||||||
@@ -887,6 +907,8 @@ interface SystemConfig {
|
|||||||
enable_registration: boolean
|
enable_registration: boolean
|
||||||
// 独立余额 Key 过期管理
|
// 独立余额 Key 过期管理
|
||||||
auto_delete_expired_keys: boolean
|
auto_delete_expired_keys: boolean
|
||||||
|
// 格式转换
|
||||||
|
enable_format_conversion: boolean
|
||||||
// 日志记录
|
// 日志记录
|
||||||
request_log_level: string
|
request_log_level: string
|
||||||
max_request_body_size: number
|
max_request_body_size: number
|
||||||
@@ -941,6 +963,8 @@ const systemConfig = ref<SystemConfig>({
|
|||||||
enable_registration: false,
|
enable_registration: false,
|
||||||
// 独立余额 Key 过期管理
|
// 独立余额 Key 过期管理
|
||||||
auto_delete_expired_keys: false,
|
auto_delete_expired_keys: false,
|
||||||
|
// 格式转换
|
||||||
|
enable_format_conversion: false,
|
||||||
// 日志记录
|
// 日志记录
|
||||||
request_log_level: 'basic',
|
request_log_level: 'basic',
|
||||||
max_request_body_size: 1048576,
|
max_request_body_size: 1048576,
|
||||||
@@ -968,7 +992,8 @@ const hasBasicConfigChanges = computed(() => {
|
|||||||
systemConfig.value.default_user_quota_usd !== originalConfig.value.default_user_quota_usd ||
|
systemConfig.value.default_user_quota_usd !== originalConfig.value.default_user_quota_usd ||
|
||||||
systemConfig.value.rate_limit_per_minute !== originalConfig.value.rate_limit_per_minute ||
|
systemConfig.value.rate_limit_per_minute !== originalConfig.value.rate_limit_per_minute ||
|
||||||
systemConfig.value.enable_registration !== originalConfig.value.enable_registration ||
|
systemConfig.value.enable_registration !== originalConfig.value.enable_registration ||
|
||||||
systemConfig.value.auto_delete_expired_keys !== originalConfig.value.auto_delete_expired_keys
|
systemConfig.value.auto_delete_expired_keys !== originalConfig.value.auto_delete_expired_keys ||
|
||||||
|
systemConfig.value.enable_format_conversion !== originalConfig.value.enable_format_conversion
|
||||||
)
|
)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -1045,6 +1070,8 @@ async function loadSystemConfig() {
|
|||||||
'enable_registration',
|
'enable_registration',
|
||||||
// 独立余额 Key 过期管理
|
// 独立余额 Key 过期管理
|
||||||
'auto_delete_expired_keys',
|
'auto_delete_expired_keys',
|
||||||
|
// 格式转换
|
||||||
|
'enable_format_conversion',
|
||||||
// 日志记录
|
// 日志记录
|
||||||
'request_log_level',
|
'request_log_level',
|
||||||
'max_request_body_size',
|
'max_request_body_size',
|
||||||
@@ -1104,6 +1131,11 @@ async function saveBasicConfig() {
|
|||||||
value: systemConfig.value.auto_delete_expired_keys,
|
value: systemConfig.value.auto_delete_expired_keys,
|
||||||
description: '是否自动删除过期的API Key'
|
description: '是否自动删除过期的API Key'
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
key: 'enable_format_conversion',
|
||||||
|
value: systemConfig.value.enable_format_conversion,
|
||||||
|
description: '全局格式转换开关:开启时强制允许所有提供商的格式转换'
|
||||||
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
await Promise.all(
|
await Promise.all(
|
||||||
@@ -1117,6 +1149,7 @@ async function saveBasicConfig() {
|
|||||||
originalConfig.value.rate_limit_per_minute = systemConfig.value.rate_limit_per_minute
|
originalConfig.value.rate_limit_per_minute = systemConfig.value.rate_limit_per_minute
|
||||||
originalConfig.value.enable_registration = systemConfig.value.enable_registration
|
originalConfig.value.enable_registration = systemConfig.value.enable_registration
|
||||||
originalConfig.value.auto_delete_expired_keys = systemConfig.value.auto_delete_expired_keys
|
originalConfig.value.auto_delete_expired_keys = systemConfig.value.auto_delete_expired_keys
|
||||||
|
originalConfig.value.enable_format_conversion = systemConfig.value.enable_format_conversion
|
||||||
}
|
}
|
||||||
success('基础配置已保存')
|
success('基础配置已保存')
|
||||||
} catch (err) {
|
} catch (err) {
|
||||||
|
|||||||
543
src/api/admin/gemini_files.py
Normal file
543
src/api/admin/gemini_files.py
Normal file
@@ -0,0 +1,543 @@
|
|||||||
|
"""
|
||||||
|
Gemini Files 管理 API
|
||||||
|
|
||||||
|
提供文件映射的管理功能:
|
||||||
|
- 列出所有文件映射
|
||||||
|
- 删除文件映射
|
||||||
|
- 查看文件映射统计
|
||||||
|
- 上传文件到 Gemini
|
||||||
|
|
||||||
|
优化:HTTP 上传期间不持有数据库连接,避免阻塞其他请求。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import hashlib
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
|
||||||
|
from pydantic import BaseModel
|
||||||
|
from sqlalchemy import delete, func
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.clients.http_client import HTTPClientPool
|
||||||
|
from src.core.crypto import crypto_service
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.database import create_session, get_db
|
||||||
|
from src.models.database import GeminiFileMapping, ProviderAPIKey, User
|
||||||
|
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class KeyInfo:
|
||||||
|
"""Key 信息(用于 HTTP 上传,不依赖数据库会话)"""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str | None
|
||||||
|
decrypted_api_key: str
|
||||||
|
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api/admin/gemini-files", tags=["Gemini Files Management"])
|
||||||
|
|
||||||
|
# Gemini Files API 基础 URL
|
||||||
|
GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
|
||||||
|
|
||||||
|
|
||||||
|
# ============ Schema ============
|
||||||
|
|
||||||
|
|
||||||
|
class FileMappingResponse(BaseModel):
|
||||||
|
"""文件映射响应"""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
file_name: str
|
||||||
|
key_id: str
|
||||||
|
key_name: str | None = None
|
||||||
|
user_id: str | None = None
|
||||||
|
username: str | None = None
|
||||||
|
display_name: str | None = None
|
||||||
|
mime_type: str | None = None
|
||||||
|
created_at: datetime
|
||||||
|
expires_at: datetime
|
||||||
|
is_expired: bool
|
||||||
|
|
||||||
|
|
||||||
|
class FileMappingListResponse(BaseModel):
|
||||||
|
"""文件映射列表响应"""
|
||||||
|
|
||||||
|
items: list[FileMappingResponse]
|
||||||
|
total: int
|
||||||
|
page: int
|
||||||
|
page_size: int
|
||||||
|
|
||||||
|
|
||||||
|
class FileMappingStatsResponse(BaseModel):
|
||||||
|
"""文件映射统计响应"""
|
||||||
|
|
||||||
|
total_mappings: int
|
||||||
|
active_mappings: int
|
||||||
|
expired_mappings: int
|
||||||
|
by_mime_type: dict[str, int]
|
||||||
|
capable_keys_count: int
|
||||||
|
|
||||||
|
|
||||||
|
# ============ Routes ============
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/mappings", response_model=FileMappingListResponse)
|
||||||
|
async def list_file_mappings(
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(20, ge=1, le=100),
|
||||||
|
include_expired: bool = Query(False),
|
||||||
|
search: str | None = Query(None),
|
||||||
|
) -> Any:
|
||||||
|
"""
|
||||||
|
列出所有文件映射
|
||||||
|
|
||||||
|
- **page**: 页码
|
||||||
|
- **page_size**: 每页数量
|
||||||
|
- **include_expired**: 是否包含已过期的映射
|
||||||
|
- **search**: 搜索文件名或显示名
|
||||||
|
"""
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
query = db.query(GeminiFileMapping)
|
||||||
|
|
||||||
|
# 过滤过期
|
||||||
|
if not include_expired:
|
||||||
|
query = query.filter(GeminiFileMapping.expires_at > now)
|
||||||
|
|
||||||
|
# 搜索
|
||||||
|
if search:
|
||||||
|
search_pattern = f"%{search}%"
|
||||||
|
query = query.filter(
|
||||||
|
(GeminiFileMapping.file_name.ilike(search_pattern))
|
||||||
|
| (GeminiFileMapping.display_name.ilike(search_pattern))
|
||||||
|
)
|
||||||
|
|
||||||
|
# 总数
|
||||||
|
total = query.count()
|
||||||
|
|
||||||
|
# 分页
|
||||||
|
offset = (page - 1) * page_size
|
||||||
|
mappings = (
|
||||||
|
query.order_by(GeminiFileMapping.created_at.desc()).offset(offset).limit(page_size).all()
|
||||||
|
)
|
||||||
|
|
||||||
|
# 获取关联的 Key 和 User 信息
|
||||||
|
key_ids = {m.key_id for m in mappings}
|
||||||
|
user_ids = {m.user_id for m in mappings if m.user_id}
|
||||||
|
|
||||||
|
keys_map = {}
|
||||||
|
if key_ids:
|
||||||
|
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.id.in_(key_ids)).all()
|
||||||
|
keys_map = {str(k.id): k.name for k in keys}
|
||||||
|
|
||||||
|
users_map = {}
|
||||||
|
if user_ids:
|
||||||
|
users = db.query(User).filter(User.id.in_(user_ids)).all()
|
||||||
|
users_map = {str(u.id): u.username for u in users}
|
||||||
|
|
||||||
|
items = []
|
||||||
|
for m in mappings:
|
||||||
|
items.append(
|
||||||
|
FileMappingResponse(
|
||||||
|
id=str(m.id),
|
||||||
|
file_name=m.file_name,
|
||||||
|
key_id=str(m.key_id),
|
||||||
|
key_name=keys_map.get(str(m.key_id)),
|
||||||
|
user_id=str(m.user_id) if m.user_id else None,
|
||||||
|
username=users_map.get(str(m.user_id)) if m.user_id else None,
|
||||||
|
display_name=m.display_name,
|
||||||
|
mime_type=m.mime_type,
|
||||||
|
created_at=m.created_at,
|
||||||
|
expires_at=m.expires_at,
|
||||||
|
is_expired=m.expires_at <= now,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return FileMappingListResponse(
|
||||||
|
items=items,
|
||||||
|
total=total,
|
||||||
|
page=page,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/stats", response_model=FileMappingStatsResponse)
|
||||||
|
async def get_file_mapping_stats(
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> Any:
|
||||||
|
"""获取文件映射统计信息"""
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# 总数
|
||||||
|
total_mappings = db.query(func.count(GeminiFileMapping.id)).scalar() or 0
|
||||||
|
|
||||||
|
# 活跃数(未过期)
|
||||||
|
active_mappings = (
|
||||||
|
db.query(func.count(GeminiFileMapping.id))
|
||||||
|
.filter(GeminiFileMapping.expires_at > now)
|
||||||
|
.scalar()
|
||||||
|
or 0
|
||||||
|
)
|
||||||
|
|
||||||
|
# 过期数
|
||||||
|
expired_mappings = total_mappings - active_mappings
|
||||||
|
|
||||||
|
# 按 MIME 类型统计
|
||||||
|
mime_stats = (
|
||||||
|
db.query(GeminiFileMapping.mime_type, func.count(GeminiFileMapping.id))
|
||||||
|
.filter(GeminiFileMapping.expires_at > now)
|
||||||
|
.group_by(GeminiFileMapping.mime_type)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
by_mime_type = {(mt or "unknown"): count for mt, count in mime_stats}
|
||||||
|
|
||||||
|
# 有 gemini_files 能力的 Key 数量
|
||||||
|
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
|
||||||
|
capable_keys_count = sum(
|
||||||
|
1 for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
|
||||||
|
)
|
||||||
|
|
||||||
|
return FileMappingStatsResponse(
|
||||||
|
total_mappings=total_mappings,
|
||||||
|
active_mappings=active_mappings,
|
||||||
|
expired_mappings=expired_mappings,
|
||||||
|
by_mime_type=by_mime_type,
|
||||||
|
capable_keys_count=capable_keys_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/mappings/{mapping_id}")
|
||||||
|
async def delete_mapping(
|
||||||
|
mapping_id: str,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> Any:
|
||||||
|
"""
|
||||||
|
删除指定的文件映射
|
||||||
|
|
||||||
|
注意:这只删除映射记录,不会删除 Gemini 上的实际文件
|
||||||
|
"""
|
||||||
|
mapping = db.query(GeminiFileMapping).filter(GeminiFileMapping.id == mapping_id).first()
|
||||||
|
|
||||||
|
if not mapping:
|
||||||
|
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||||
|
|
||||||
|
file_name = mapping.file_name
|
||||||
|
|
||||||
|
# 从数据库删除
|
||||||
|
db.delete(mapping)
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
# 同时从 Redis 删除
|
||||||
|
await delete_file_key_mapping(file_name)
|
||||||
|
|
||||||
|
return {"message": "Mapping deleted successfully", "file_name": file_name}
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/mappings")
|
||||||
|
async def cleanup_expired_mappings(
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> Any:
|
||||||
|
"""清理所有过期的文件映射"""
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
deleted_count = result.rowcount
|
||||||
|
|
||||||
|
return {
|
||||||
|
"message": f"Cleaned up {deleted_count} expired mappings",
|
||||||
|
"deleted_count": deleted_count,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class CapableKeyResponse(BaseModel):
|
||||||
|
"""可用 Key 响应"""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
provider_name: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class UploadResultItem(BaseModel):
|
||||||
|
"""单个 Key 的上传结果"""
|
||||||
|
|
||||||
|
key_id: str
|
||||||
|
key_name: str | None = None
|
||||||
|
success: bool
|
||||||
|
file_name: str | None = None
|
||||||
|
error: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class UploadResponse(BaseModel):
|
||||||
|
"""上传响应"""
|
||||||
|
|
||||||
|
display_name: str
|
||||||
|
mime_type: str
|
||||||
|
size_bytes: int
|
||||||
|
results: list[UploadResultItem]
|
||||||
|
success_count: int
|
||||||
|
fail_count: int
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/capable-keys", response_model=list[CapableKeyResponse])
|
||||||
|
async def list_capable_keys(
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> Any:
|
||||||
|
"""获取所有具有 gemini_files 能力的 Key 列表"""
|
||||||
|
from src.models.database import Provider
|
||||||
|
|
||||||
|
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
|
||||||
|
capable_keys = [
|
||||||
|
key for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
|
||||||
|
]
|
||||||
|
|
||||||
|
# 获取 Provider 名称
|
||||||
|
provider_ids = {key.provider_id for key in capable_keys}
|
||||||
|
providers = db.query(Provider).filter(Provider.id.in_(provider_ids)).all()
|
||||||
|
provider_map = {str(p.id): p.name for p in providers}
|
||||||
|
|
||||||
|
return [
|
||||||
|
CapableKeyResponse(
|
||||||
|
id=str(key.id),
|
||||||
|
name=key.name,
|
||||||
|
provider_name=provider_map.get(str(key.provider_id)),
|
||||||
|
)
|
||||||
|
for key in capable_keys
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
async def _upload_to_key(
|
||||||
|
key_info: KeyInfo,
|
||||||
|
content: bytes,
|
||||||
|
file_size: int,
|
||||||
|
mime_type: str,
|
||||||
|
display_name: str,
|
||||||
|
source_hash: str,
|
||||||
|
) -> UploadResultItem:
|
||||||
|
"""上传文件到指定的 Key(不依赖数据库会话)"""
|
||||||
|
api_key = key_info.decrypted_api_key
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 第一步:初始化可恢复上传
|
||||||
|
init_url = f"{GEMINI_FILES_BASE_URL}/upload/v1beta/files?key={api_key}"
|
||||||
|
init_headers = {
|
||||||
|
"X-Goog-Upload-Protocol": "resumable",
|
||||||
|
"X-Goog-Upload-Command": "start",
|
||||||
|
"X-Goog-Upload-Header-Content-Length": str(file_size),
|
||||||
|
"X-Goog-Upload-Header-Content-Type": mime_type,
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
}
|
||||||
|
init_body = {"file": {"display_name": display_name}}
|
||||||
|
|
||||||
|
init_response = await client.post(init_url, headers=init_headers, json=init_body)
|
||||||
|
|
||||||
|
if init_response.status_code != 200:
|
||||||
|
logger.error(
|
||||||
|
f"Gemini upload init failed for key {key_info.id}: {init_response.status_code}"
|
||||||
|
)
|
||||||
|
return UploadResultItem(
|
||||||
|
key_id=key_info.id,
|
||||||
|
key_name=key_info.name,
|
||||||
|
success=False,
|
||||||
|
error=f"初始化失败: {init_response.status_code}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 获取上传 URL
|
||||||
|
upload_url = init_response.headers.get("X-Goog-Upload-URL")
|
||||||
|
if not upload_url:
|
||||||
|
return UploadResultItem(
|
||||||
|
key_id=key_info.id,
|
||||||
|
key_name=key_info.name,
|
||||||
|
success=False,
|
||||||
|
error="未获取到上传 URL",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 第二步:上传文件内容
|
||||||
|
upload_headers = {
|
||||||
|
"X-Goog-Upload-Command": "upload, finalize",
|
||||||
|
"X-Goog-Upload-Offset": "0",
|
||||||
|
"Content-Length": str(file_size),
|
||||||
|
"Content-Type": mime_type,
|
||||||
|
}
|
||||||
|
|
||||||
|
upload_response = await client.post(upload_url, headers=upload_headers, content=content)
|
||||||
|
|
||||||
|
if upload_response.status_code != 200:
|
||||||
|
logger.error(
|
||||||
|
f"Gemini upload failed for key {key_info.id}: {upload_response.status_code}"
|
||||||
|
)
|
||||||
|
return UploadResultItem(
|
||||||
|
key_id=key_info.id,
|
||||||
|
key_name=key_info.name,
|
||||||
|
success=False,
|
||||||
|
error=f"上传失败: {upload_response.status_code}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 解析响应
|
||||||
|
result = upload_response.json()
|
||||||
|
file_info = result.get("file", {})
|
||||||
|
file_name = file_info.get("name", "")
|
||||||
|
response_display_name = file_info.get("displayName", display_name)
|
||||||
|
response_mime_type = file_info.get("mimeType", mime_type)
|
||||||
|
|
||||||
|
# 存储文件映射(包含源文件哈希,用于关联相同源文件的不同上传)
|
||||||
|
await store_file_key_mapping(
|
||||||
|
file_name=file_name,
|
||||||
|
key_id=key_info.id,
|
||||||
|
user_id=None,
|
||||||
|
display_name=response_display_name,
|
||||||
|
mime_type=response_mime_type,
|
||||||
|
source_hash=source_hash,
|
||||||
|
)
|
||||||
|
|
||||||
|
return UploadResultItem(
|
||||||
|
key_id=key_info.id,
|
||||||
|
key_name=key_info.name,
|
||||||
|
success=True,
|
||||||
|
file_name=file_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Gemini upload error for key {key_info.id}: {exc}")
|
||||||
|
return UploadResultItem(
|
||||||
|
key_id=key_info.id,
|
||||||
|
key_name=key_info.name,
|
||||||
|
success=False,
|
||||||
|
error=str(exc),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/upload", response_model=UploadResponse)
|
||||||
|
async def upload_file(
|
||||||
|
file: UploadFile = File(...),
|
||||||
|
key_ids: str = Query(..., description="逗号分隔的 Key ID 列表"),
|
||||||
|
) -> Any:
|
||||||
|
"""
|
||||||
|
上传文件到 Gemini Files API
|
||||||
|
|
||||||
|
- **file**: 要上传的文件
|
||||||
|
- **key_ids**: 逗号分隔的 Key ID 列表,文件将上传到所有指定的 Key
|
||||||
|
|
||||||
|
支持的文件类型:视频、图片、音频、文档等
|
||||||
|
文件大小限制:2GB
|
||||||
|
文件有效期:48小时
|
||||||
|
|
||||||
|
优化:HTTP 上传期间不持有数据库连接
|
||||||
|
"""
|
||||||
|
# 解析 Key IDs
|
||||||
|
key_id_list = [kid.strip() for kid in key_ids.split(",") if kid.strip()]
|
||||||
|
if not key_id_list:
|
||||||
|
raise HTTPException(status_code=400, detail="请至少选择一个 Key")
|
||||||
|
|
||||||
|
# ========== 阶段 1:读取文件内容并计算哈希 ==========
|
||||||
|
content = await file.read()
|
||||||
|
file_size = len(content)
|
||||||
|
mime_type = file.content_type or "application/octet-stream"
|
||||||
|
display_name = file.filename or "uploaded_file"
|
||||||
|
|
||||||
|
# 计算源文件哈希(用于重复检测和关联相同源文件的不同上传)
|
||||||
|
source_hash = hashlib.sha256(content).hexdigest()
|
||||||
|
|
||||||
|
# ========== 阶段 2:查询数据库(短暂持有连接)==========
|
||||||
|
key_infos: list[KeyInfo] = []
|
||||||
|
existing_mappings: dict[str, str] = {} # key_id -> existing file_name
|
||||||
|
|
||||||
|
with create_session() as db:
|
||||||
|
keys = (
|
||||||
|
db.query(ProviderAPIKey)
|
||||||
|
.filter(
|
||||||
|
ProviderAPIKey.id.in_(key_id_list),
|
||||||
|
ProviderAPIKey.is_active.is_(True),
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
|
||||||
|
# 过滤有 gemini_files 能力的 Key,并提取必要信息
|
||||||
|
capable_key_ids = []
|
||||||
|
for key in keys:
|
||||||
|
if key.capabilities and key.capabilities.get("gemini_files", False):
|
||||||
|
try:
|
||||||
|
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||||
|
key_infos.append(
|
||||||
|
KeyInfo(
|
||||||
|
id=str(key.id),
|
||||||
|
name=key.name,
|
||||||
|
decrypted_api_key=decrypted_key,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
capable_key_ids.append(str(key.id))
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Failed to decrypt provider key {key.id}: {exc}")
|
||||||
|
|
||||||
|
# 检查是否已存在相同 source_hash 的映射(重复检测)
|
||||||
|
if capable_key_ids:
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
existing = (
|
||||||
|
db.query(GeminiFileMapping)
|
||||||
|
.filter(
|
||||||
|
GeminiFileMapping.source_hash == source_hash,
|
||||||
|
GeminiFileMapping.key_id.in_(capable_key_ids),
|
||||||
|
GeminiFileMapping.expires_at > now, # 只查未过期的
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
for mapping in existing:
|
||||||
|
existing_mappings[str(mapping.key_id)] = mapping.file_name
|
||||||
|
|
||||||
|
if not key_infos:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="选中的 Key 都没有「Gemini 文件 API」能力或解密失败",
|
||||||
|
)
|
||||||
|
|
||||||
|
# ========== 阶段 3:并发上传(跳过已有相同文件的 Key)==========
|
||||||
|
results: list[UploadResultItem] = []
|
||||||
|
keys_to_upload: list[KeyInfo] = []
|
||||||
|
|
||||||
|
for key_info in key_infos:
|
||||||
|
if key_info.id in existing_mappings:
|
||||||
|
# 该 Key 已有相同文件,跳过上传
|
||||||
|
results.append(
|
||||||
|
UploadResultItem(
|
||||||
|
key_id=key_info.id,
|
||||||
|
key_name=key_info.name,
|
||||||
|
success=True,
|
||||||
|
file_name=existing_mappings[key_info.id],
|
||||||
|
error=None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"跳过重复上传: Key {key_info.id} 已有文件 {existing_mappings[key_info.id]}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
keys_to_upload.append(key_info)
|
||||||
|
|
||||||
|
# 只对需要上传的 Key 执行上传
|
||||||
|
if keys_to_upload:
|
||||||
|
tasks = [
|
||||||
|
_upload_to_key(key_info, content, file_size, mime_type, display_name, source_hash)
|
||||||
|
for key_info in keys_to_upload
|
||||||
|
]
|
||||||
|
upload_results = await asyncio.gather(*tasks)
|
||||||
|
results.extend(upload_results)
|
||||||
|
|
||||||
|
success_count = sum(1 for r in results if r.success)
|
||||||
|
fail_count = len(results) - success_count
|
||||||
|
|
||||||
|
return UploadResponse(
|
||||||
|
display_name=display_name,
|
||||||
|
mime_type=mime_type,
|
||||||
|
size_bytes=file_size,
|
||||||
|
results=results,
|
||||||
|
success_count=success_count,
|
||||||
|
fail_count=fail_count,
|
||||||
|
)
|
||||||
@@ -20,6 +20,7 @@ from src.core.logger import logger
|
|||||||
from src.database.database import get_db
|
from src.database.database import get_db
|
||||||
from src.models.database import Provider, ProviderEndpoint, User
|
from src.models.database import Provider, ProviderEndpoint, User
|
||||||
from src.services.model.fetch_scheduler import (
|
from src.services.model.fetch_scheduler import (
|
||||||
|
MODEL_FETCH_HTTP_TIMEOUT,
|
||||||
get_upstream_models_from_cache,
|
get_upstream_models_from_cache,
|
||||||
set_upstream_models_to_cache,
|
set_upstream_models_to_cache,
|
||||||
)
|
)
|
||||||
@@ -135,7 +136,9 @@ async def query_available_models(
|
|||||||
return [], f"Key {api_key.name or api_key.id}: decrypt failed", False
|
return [], f"Key {api_key.name or api_key.id}: decrypt failed", False
|
||||||
|
|
||||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
|
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
|
||||||
models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
|
models, errors, has_success = await fetch_models_from_endpoints(
|
||||||
|
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
|
||||||
|
)
|
||||||
|
|
||||||
# 写入缓存
|
# 写入缓存
|
||||||
if models:
|
if models:
|
||||||
@@ -274,7 +277,9 @@ async def _fetch_models_for_single_key(
|
|||||||
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
|
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
|
||||||
|
|
||||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
|
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
|
||||||
all_models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
|
all_models, errors, has_success = await fetch_models_from_endpoints(
|
||||||
|
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
|
||||||
|
)
|
||||||
|
|
||||||
# 按 model id 聚合,合并所有 api_format
|
# 按 model id 聚合,合并所有 api_format
|
||||||
unique_models = _aggregate_models_by_id(all_models)
|
unique_models = _aggregate_models_by_id(all_models)
|
||||||
|
|||||||
@@ -309,6 +309,7 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
|
|||||||
website=provider.website,
|
website=provider.website,
|
||||||
provider_priority=provider.provider_priority,
|
provider_priority=provider.provider_priority,
|
||||||
keep_priority_on_conversion=provider.keep_priority_on_conversion,
|
keep_priority_on_conversion=provider.keep_priority_on_conversion,
|
||||||
|
enable_format_conversion=provider.enable_format_conversion,
|
||||||
is_active=provider.is_active,
|
is_active=provider.is_active,
|
||||||
billing_type=provider.billing_type.value if provider.billing_type else None,
|
billing_type=provider.billing_type.value if provider.billing_type else None,
|
||||||
monthly_quota_usd=provider.monthly_quota_usd,
|
monthly_quota_usd=provider.monthly_quota_usd,
|
||||||
|
|||||||
@@ -271,6 +271,9 @@ class AdminUsageStatsAdapter(AdminApiAdapter):
|
|||||||
self.end_date = end_date
|
self.end_date = end_date
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||||
|
# Perf: use a single aggregate query (avoid 3 full scans).
|
||||||
|
from sqlalchemy import case
|
||||||
|
|
||||||
db = context.db
|
db = context.db
|
||||||
query = db.query(Usage)
|
query = db.query(Usage)
|
||||||
if self.start_date:
|
if self.start_date:
|
||||||
@@ -278,56 +281,53 @@ class AdminUsageStatsAdapter(AdminApiAdapter):
|
|||||||
if self.end_date:
|
if self.end_date:
|
||||||
query = query.filter(Usage.created_at <= self.end_date)
|
query = query.filter(Usage.created_at <= self.end_date)
|
||||||
|
|
||||||
total_stats = query.with_entities(
|
stats = query.with_entities(
|
||||||
func.count(Usage.id).label("total_requests"),
|
func.count(Usage.id).label("total_requests"),
|
||||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||||
func.sum(Usage.total_cost_usd).label("total_cost"),
|
func.sum(Usage.total_cost_usd).label("total_cost"),
|
||||||
func.sum(Usage.actual_total_cost_usd).label("total_actual_cost"),
|
func.sum(Usage.actual_total_cost_usd).label("total_actual_cost"),
|
||||||
func.avg(Usage.response_time_ms).label("avg_response_time_ms"),
|
func.avg(Usage.response_time_ms).label("avg_response_time_ms"),
|
||||||
).first()
|
|
||||||
|
|
||||||
# 缓存统计
|
|
||||||
cache_stats = query.with_entities(
|
|
||||||
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
|
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
|
||||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||||
func.sum(Usage.cache_creation_cost_usd).label("cache_creation_cost"),
|
func.sum(Usage.cache_creation_cost_usd).label("cache_creation_cost"),
|
||||||
func.sum(Usage.cache_read_cost_usd).label("cache_read_cost"),
|
func.sum(Usage.cache_read_cost_usd).label("cache_read_cost"),
|
||||||
|
func.sum(
|
||||||
|
case(
|
||||||
|
(
|
||||||
|
(Usage.status_code >= 400) | (Usage.error_message.isnot(None)),
|
||||||
|
1,
|
||||||
|
),
|
||||||
|
else_=0,
|
||||||
|
)
|
||||||
|
).label("error_count"),
|
||||||
).first()
|
).first()
|
||||||
|
|
||||||
# 错误统计
|
|
||||||
error_count = query.filter(
|
|
||||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
|
||||||
).count()
|
|
||||||
|
|
||||||
context.add_audit_metadata(
|
context.add_audit_metadata(
|
||||||
action="usage_stats",
|
action="usage_stats",
|
||||||
start_date=self.start_date.isoformat() if self.start_date else None,
|
start_date=self.start_date.isoformat() if self.start_date else None,
|
||||||
end_date=self.end_date.isoformat() if self.end_date else None,
|
end_date=self.end_date.isoformat() if self.end_date else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
total_requests = total_stats.total_requests if total_stats else 0
|
total_requests = int(stats.total_requests or 0) if stats else 0
|
||||||
avg_response_time_ms = float(total_stats.avg_response_time_ms or 0) if total_stats else 0
|
avg_response_time_ms = float(stats.avg_response_time_ms or 0) if stats else 0
|
||||||
avg_response_time = avg_response_time_ms / 1000.0
|
avg_response_time = avg_response_time_ms / 1000.0
|
||||||
|
error_count = int(stats.error_count or 0) if stats else 0
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"total_requests": total_requests,
|
"total_requests": total_requests,
|
||||||
"total_tokens": int(total_stats.total_tokens or 0),
|
"total_tokens": int(stats.total_tokens or 0) if stats else 0,
|
||||||
"total_cost": float(total_stats.total_cost or 0),
|
"total_cost": float(stats.total_cost or 0) if stats else 0,
|
||||||
"total_actual_cost": float(total_stats.total_actual_cost or 0),
|
"total_actual_cost": float(stats.total_actual_cost or 0) if stats else 0,
|
||||||
"avg_response_time": round(avg_response_time, 2),
|
"avg_response_time": round(avg_response_time, 2),
|
||||||
"error_count": error_count,
|
"error_count": error_count,
|
||||||
"error_rate": (
|
"error_rate": (
|
||||||
round((error_count / total_requests) * 100, 2) if total_requests > 0 else 0
|
round((error_count / total_requests) * 100, 2) if total_requests > 0 else 0
|
||||||
),
|
),
|
||||||
"cache_stats": {
|
"cache_stats": {
|
||||||
"cache_creation_tokens": (
|
"cache_creation_tokens": (int(stats.cache_creation_tokens or 0) if stats else 0),
|
||||||
int(cache_stats.cache_creation_tokens or 0) if cache_stats else 0
|
"cache_read_tokens": int(stats.cache_read_tokens or 0) if stats else 0,
|
||||||
),
|
"cache_creation_cost": (float(stats.cache_creation_cost or 0) if stats else 0),
|
||||||
"cache_read_tokens": int(cache_stats.cache_read_tokens or 0) if cache_stats else 0,
|
"cache_read_cost": float(stats.cache_read_cost or 0) if stats else 0,
|
||||||
"cache_creation_cost": (
|
|
||||||
float(cache_stats.cache_creation_cost or 0) if cache_stats else 0
|
|
||||||
),
|
|
||||||
"cache_read_cost": float(cache_stats.cache_read_cost or 0) if cache_stats else 0,
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -637,6 +637,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||||
from sqlalchemy import or_
|
from sqlalchemy import or_
|
||||||
|
from sqlalchemy.orm import load_only
|
||||||
|
|
||||||
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
|
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
|
||||||
|
|
||||||
@@ -677,13 +678,13 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
escaped = escape_like_pattern(self.username)
|
escaped = escape_like_pattern(self.username)
|
||||||
query = query.filter(User.username.ilike(f"%{escaped}%", escape="\\"))
|
query = query.filter(User.username.ilike(f"%{escaped}%", escape="\\"))
|
||||||
if self.model:
|
if self.model:
|
||||||
# 支持模型名模糊搜索
|
# 模型筛选:前端为下拉框精确值,使用精确匹配以启用索引
|
||||||
escaped = escape_like_pattern(self.model)
|
# 如需模糊搜索,请使用 search 参数。
|
||||||
query = query.filter(Usage.model.ilike(f"%{escaped}%", escape="\\"))
|
query = query.filter(Usage.model == self.model)
|
||||||
if self.provider:
|
if self.provider:
|
||||||
# 支持提供商名称搜索
|
# 提供商筛选:前端为下拉框精确值,使用精确匹配以启用索引
|
||||||
escaped = escape_like_pattern(self.provider)
|
# 如需模糊搜索,请使用 search 参数。
|
||||||
query = query.filter(Provider.name.ilike(f"%{escaped}%", escape="\\"))
|
query = query.filter(Provider.name == self.provider)
|
||||||
if self.status:
|
if self.status:
|
||||||
# 状态筛选
|
# 状态筛选
|
||||||
# 旧的筛选值(基于 is_stream 和 status_code):stream, standard, error
|
# 旧的筛选值(基于 is_stream 和 status_code):stream, standard, error
|
||||||
@@ -714,7 +715,51 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
if self.end_date:
|
if self.end_date:
|
||||||
query = query.filter(Usage.created_at <= self.end_date)
|
query = query.filter(Usage.created_at <= self.end_date)
|
||||||
|
|
||||||
total = query.count()
|
# Perf: avoid Query.count() building a subquery selecting many columns
|
||||||
|
total = int(query.with_entities(func.count(Usage.id)).scalar() or 0)
|
||||||
|
|
||||||
|
# Perf: do not load large request/response columns for list view
|
||||||
|
query = query.options(
|
||||||
|
load_only(
|
||||||
|
Usage.id,
|
||||||
|
Usage.request_id,
|
||||||
|
Usage.user_id,
|
||||||
|
Usage.api_key_id,
|
||||||
|
Usage.provider_name,
|
||||||
|
Usage.provider_id,
|
||||||
|
Usage.provider_endpoint_id,
|
||||||
|
Usage.provider_api_key_id,
|
||||||
|
Usage.model,
|
||||||
|
Usage.target_model,
|
||||||
|
Usage.input_tokens,
|
||||||
|
Usage.output_tokens,
|
||||||
|
Usage.cache_creation_input_tokens,
|
||||||
|
Usage.cache_read_input_tokens,
|
||||||
|
Usage.total_tokens,
|
||||||
|
Usage.total_cost_usd,
|
||||||
|
Usage.actual_total_cost_usd,
|
||||||
|
Usage.rate_multiplier,
|
||||||
|
Usage.response_time_ms,
|
||||||
|
Usage.first_byte_time_ms,
|
||||||
|
Usage.created_at,
|
||||||
|
Usage.is_stream,
|
||||||
|
Usage.status_code,
|
||||||
|
Usage.error_message,
|
||||||
|
Usage.status,
|
||||||
|
Usage.api_format,
|
||||||
|
Usage.endpoint_api_format,
|
||||||
|
Usage.has_format_conversion,
|
||||||
|
Usage.request_metadata,
|
||||||
|
Usage.input_price_per_1m,
|
||||||
|
Usage.output_price_per_1m,
|
||||||
|
Usage.cache_creation_price_per_1m,
|
||||||
|
Usage.cache_read_price_per_1m,
|
||||||
|
),
|
||||||
|
load_only(User.id, User.email, User.username),
|
||||||
|
load_only(ProviderEndpoint.id, ProviderEndpoint.api_format),
|
||||||
|
load_only(ProviderAPIKey.id, ProviderAPIKey.name),
|
||||||
|
load_only(ApiKey.id, ApiKey.name, ApiKey.key_encrypted),
|
||||||
|
)
|
||||||
records = (
|
records = (
|
||||||
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
|
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
|
||||||
)
|
)
|
||||||
@@ -779,7 +824,9 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 构建 provider_id -> Provider 名称的映射,避免 N+1 查询
|
# 构建 provider_id -> Provider 名称的映射,避免 N+1 查询
|
||||||
provider_ids = [usage.provider_id for usage, _, _, _, _ in records if usage.provider_id]
|
provider_ids = list(
|
||||||
|
{usage.provider_id for usage, _, _, _, _ in records if usage.provider_id}
|
||||||
|
)
|
||||||
provider_map = {}
|
provider_map = {}
|
||||||
if provider_ids:
|
if provider_ids:
|
||||||
providers_data = (
|
providers_data = (
|
||||||
@@ -788,6 +835,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
provider_map = {str(p.id): p.name for p in providers_data}
|
provider_map = {str(p.id): p.name for p in providers_data}
|
||||||
|
|
||||||
data = []
|
data = []
|
||||||
|
api_key_display_cache: dict[str, str] = {}
|
||||||
for usage, user, endpoint, provider_api_key, user_api_key in records:
|
for usage, user, endpoint, provider_api_key, user_api_key in records:
|
||||||
actual_cost = (
|
actual_cost = (
|
||||||
float(usage.actual_total_cost_usd)
|
float(usage.actual_total_cost_usd)
|
||||||
@@ -829,7 +877,9 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
{
|
{
|
||||||
"id": user_api_key.id,
|
"id": user_api_key.id,
|
||||||
"name": user_api_key.name,
|
"name": user_api_key.name,
|
||||||
"display": user_api_key.get_display_key(),
|
"display": api_key_display_cache.setdefault(
|
||||||
|
user_api_key.id, user_api_key.get_display_key()
|
||||||
|
),
|
||||||
}
|
}
|
||||||
if user_api_key
|
if user_api_key
|
||||||
else None
|
else None
|
||||||
|
|||||||
@@ -311,7 +311,7 @@ def get_compatible_provider_formats(
|
|||||||
endpoint_format,
|
endpoint_format,
|
||||||
format_acceptance_config,
|
format_acceptance_config,
|
||||||
is_stream=False,
|
is_stream=False,
|
||||||
global_conversion_enabled=global_conversion_enabled,
|
effective_conversion_enabled=global_conversion_enabled,
|
||||||
)
|
)
|
||||||
if not is_compatible:
|
if not is_compatible:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -813,23 +813,31 @@ class DashboardRecentRequestsAdapter(DashboardAdapter):
|
|||||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
user = context.user
|
user = context.user
|
||||||
query = db.query(Usage)
|
# Perf: select only required columns (avoid loading large JSON/BLOB fields).
|
||||||
|
query = db.query(
|
||||||
|
Usage.id,
|
||||||
|
Usage.user_id,
|
||||||
|
Usage.model,
|
||||||
|
Usage.total_tokens,
|
||||||
|
Usage.created_at,
|
||||||
|
Usage.is_stream,
|
||||||
|
DBUser.username,
|
||||||
|
).outerjoin(DBUser, DBUser.id == Usage.user_id)
|
||||||
if user.role != UserRole.ADMIN:
|
if user.role != UserRole.ADMIN:
|
||||||
query = query.filter(Usage.user_id == user.id)
|
query = query.filter(Usage.user_id == user.id)
|
||||||
|
|
||||||
recent_requests = query.order_by(Usage.created_at.desc()).limit(self.limit).all()
|
rows = query.order_by(Usage.created_at.desc()).limit(self.limit).all()
|
||||||
|
|
||||||
results = []
|
results = []
|
||||||
for req in recent_requests:
|
for req_id, _user_id, model, total_tokens, created_at, is_stream, username in rows:
|
||||||
owner = db.query(DBUser).filter(DBUser.id == req.user_id).first()
|
|
||||||
results.append(
|
results.append(
|
||||||
{
|
{
|
||||||
"id": req.id,
|
"id": req_id,
|
||||||
"user": owner.username if owner else "Unknown",
|
"user": username or "Unknown",
|
||||||
"model": req.model or "N/A",
|
"model": model or "N/A",
|
||||||
"tokens": req.total_tokens,
|
"tokens": int(total_tokens or 0),
|
||||||
"time": req.created_at.strftime("%H:%M") if req.created_at else None,
|
"time": created_at.strftime("%H:%M") if created_at else None,
|
||||||
"is_stream": req.is_stream,
|
"is_stream": bool(is_stream),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -848,18 +856,25 @@ class DashboardProviderStatusAdapter(DashboardAdapter):
|
|||||||
providers = db.query(Provider).filter(Provider.is_active.is_(True)).all()
|
providers = db.query(Provider).filter(Provider.is_active.is_(True)).all()
|
||||||
since = datetime.now(timezone.utc) - timedelta(days=1)
|
since = datetime.now(timezone.utc) - timedelta(days=1)
|
||||||
|
|
||||||
|
# Avoid N+1: compute 24h request counts for all providers in one GROUP BY query.
|
||||||
|
provider_names = [p.name for p in providers if p and p.name]
|
||||||
|
counts: dict[str, int] = {}
|
||||||
|
if provider_names:
|
||||||
|
rows = (
|
||||||
|
db.query(Usage.provider_name, func.count(Usage.id))
|
||||||
|
.filter(and_(Usage.created_at >= since, Usage.provider_name.in_(provider_names)))
|
||||||
|
.group_by(Usage.provider_name)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
counts = {str(name): int(cnt or 0) for name, cnt in rows if name}
|
||||||
|
|
||||||
entries = []
|
entries = []
|
||||||
for provider in providers:
|
for provider in providers:
|
||||||
count = (
|
|
||||||
db.query(func.count(Usage.id))
|
|
||||||
.filter(and_(Usage.provider_name == provider.name, Usage.created_at >= since))
|
|
||||||
.scalar()
|
|
||||||
)
|
|
||||||
entries.append(
|
entries.append(
|
||||||
{
|
{
|
||||||
"name": provider.name,
|
"name": provider.name,
|
||||||
"status": "active" if provider.is_active else "inactive",
|
"status": "active" if provider.is_active else "inactive",
|
||||||
"requests": count,
|
"requests": int(counts.get(provider.name, 0)),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -79,10 +79,8 @@ class VideoAdapterBase(ApiAdapter):
|
|||||||
path_params=path_params,
|
path_params=path_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Cancel task
|
# Cancel task (POST /videos/{id}/cancel or explicit action=cancel)
|
||||||
if method in {"DELETE", "POST"} and (
|
if (method == "POST" and path.endswith("/cancel")) or path_params.get("action") == "cancel":
|
||||||
path.endswith("/cancel") or path_params.get("action") == "cancel"
|
|
||||||
):
|
|
||||||
if not task_id:
|
if not task_id:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=400, detail="Task ID is required for cancel operation"
|
status_code=400, detail="Task ID is required for cancel operation"
|
||||||
@@ -95,6 +93,16 @@ class VideoAdapterBase(ApiAdapter):
|
|||||||
path_params=path_params,
|
path_params=path_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Delete task (DELETE /videos/{id})
|
||||||
|
if method == "DELETE" and task_id:
|
||||||
|
return await handler.handle_delete_task(
|
||||||
|
task_id=task_id,
|
||||||
|
http_request=http_request,
|
||||||
|
original_headers=context.original_headers,
|
||||||
|
query_params=context.query_params,
|
||||||
|
path_params=path_params,
|
||||||
|
)
|
||||||
|
|
||||||
# Remix task
|
# Remix task
|
||||||
if method == "POST" and path.endswith("/remix") and task_id:
|
if method == "POST" and path.endswith("/remix") and task_id:
|
||||||
return await handler.handle_remix_task(
|
return await handler.handle_remix_task(
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from src.services.cache.aware_scheduler import ProviderCandidate
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from src.services.task.orchestrator import SubmitOutcome
|
from src.services.candidate.submit import SubmitOutcome
|
||||||
|
|
||||||
# 敏感信息匹配正则(预编译提升性能)
|
# 敏感信息匹配正则(预编译提升性能)
|
||||||
_SENSITIVE_PATTERN = re.compile(
|
_SENSITIVE_PATTERN = re.compile(
|
||||||
@@ -53,21 +53,39 @@ def sanitize_error_message(message: str, max_length: int = 200) -> str:
|
|||||||
return sanitized[:max_length]
|
return sanitized[:max_length]
|
||||||
|
|
||||||
|
|
||||||
|
def extract_short_id_from_operation(operation_id: str) -> str:
|
||||||
|
"""
|
||||||
|
从 operation ID 中提取短 ID
|
||||||
|
|
||||||
|
我们对外暴露的 operation name 格式是:
|
||||||
|
- models/{model}/operations/{short_id}
|
||||||
|
|
||||||
|
此函数提取最后一部分作为 short_id,用于在数据库中查找任务。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
operation_id: 原始 operation ID(如 "models/veo-3.1/operations/abc123")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
short_id(如 "abc123")
|
||||||
|
"""
|
||||||
|
# 格式: models/{model}/operations/{short_id}
|
||||||
|
# 或者直接是 short_id
|
||||||
|
if "/" in operation_id:
|
||||||
|
# 提取最后一部分
|
||||||
|
return operation_id.rsplit("/", 1)[-1]
|
||||||
|
return operation_id
|
||||||
|
|
||||||
|
|
||||||
def normalize_gemini_operation_id(operation_id: str) -> str:
|
def normalize_gemini_operation_id(operation_id: str) -> str:
|
||||||
"""
|
"""
|
||||||
规范化 Gemini operation ID,确保以 "operations/" 开头
|
规范化 Gemini operation ID(保留用于向后兼容)
|
||||||
|
|
||||||
Gemini API 返回的任务 ID 格式可能是 "operations/xxx" 或 "xxx",
|
|
||||||
此函数统一规范化为 "operations/xxx" 格式。
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
operation_id: 原始 operation ID
|
operation_id: 原始 operation ID
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
规范化后的 operation ID
|
规范化后的 operation ID(原样返回)
|
||||||
"""
|
"""
|
||||||
if not operation_id.startswith("operations/"):
|
|
||||||
return f"operations/{operation_id}"
|
|
||||||
return operation_id
|
return operation_id
|
||||||
|
|
||||||
|
|
||||||
@@ -143,6 +161,18 @@ class VideoHandlerBase(ABC):
|
|||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
"""取消任务"""
|
"""取消任务"""
|
||||||
|
|
||||||
|
async def handle_delete_task(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
task_id: str,
|
||||||
|
http_request: Request,
|
||||||
|
original_headers: dict[str, str],
|
||||||
|
query_params: dict[str, str] | None = None,
|
||||||
|
path_params: dict[str, Any] | None = None,
|
||||||
|
) -> JSONResponse:
|
||||||
|
"""删除已完成或失败的视频任务 - 可选实现"""
|
||||||
|
raise HTTPException(status_code=501, detail="Delete not supported for this provider")
|
||||||
|
|
||||||
async def handle_remix_task(
|
async def handle_remix_task(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -205,6 +235,7 @@ class VideoHandlerBase(ABC):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def _get_task(self, task_id: str) -> VideoTask:
|
def _get_task(self, task_id: str) -> VideoTask:
|
||||||
|
"""通过 UUID 查找任务(OpenAI Sora 风格)"""
|
||||||
task = (
|
task = (
|
||||||
self.db.query(VideoTask)
|
self.db.query(VideoTask)
|
||||||
.filter(VideoTask.id == task_id, VideoTask.user_id == self.user.id)
|
.filter(VideoTask.id == task_id, VideoTask.user_id == self.user.id)
|
||||||
@@ -229,7 +260,7 @@ class VideoHandlerBase(ABC):
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
status = VideoStatus.PENDING
|
status = VideoStatus.PENDING
|
||||||
return InternalVideoTask(
|
return InternalVideoTask(
|
||||||
id=task.id,
|
id=task.id, # OpenAI Sora 使用 UUID
|
||||||
external_id=task.external_task_id,
|
external_id=task.external_task_id,
|
||||||
status=status,
|
status=status,
|
||||||
progress_percent=task.progress_percent or 0,
|
progress_percent=task.progress_percent or 0,
|
||||||
@@ -243,6 +274,52 @@ class VideoHandlerBase(ABC):
|
|||||||
extra={"model": task.model},
|
extra={"model": task.model},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _finalize_usage_on_submit_failure(
|
||||||
|
self,
|
||||||
|
candidate_keys: list[dict[str, Any]],
|
||||||
|
status_code: int | None,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
提交失败时结算 pending usage(避免遗留 pending 状态)。
|
||||||
|
|
||||||
|
从 candidate_keys 中提取最后尝试的 provider 信息,更新 Usage 记录。
|
||||||
|
"""
|
||||||
|
from src.services.usage.service import UsageService
|
||||||
|
|
||||||
|
# 提取 provider 信息:优先取最后一个有 attempt 的候选
|
||||||
|
provider_name = "unknown"
|
||||||
|
provider_id = None
|
||||||
|
endpoint_id = None
|
||||||
|
key_id = None
|
||||||
|
|
||||||
|
for ck in reversed(candidate_keys):
|
||||||
|
if ck.get("attempt_status") or ck.get("selected"):
|
||||||
|
provider_name = ck.get("provider_name") or "unknown"
|
||||||
|
provider_id = ck.get("provider_id")
|
||||||
|
endpoint_id = ck.get("endpoint_id")
|
||||||
|
key_id = ck.get("key_id")
|
||||||
|
break
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 更新 usage 状态并设置 provider 信息
|
||||||
|
UsageService.update_usage_status(
|
||||||
|
self.db,
|
||||||
|
request_id=self.request_id,
|
||||||
|
status="failed",
|
||||||
|
error_message=f"submit_failed (status_code={status_code or 'unknown'})",
|
||||||
|
provider=provider_name,
|
||||||
|
provider_id=provider_id,
|
||||||
|
provider_endpoint_id=endpoint_id,
|
||||||
|
provider_api_key_id=key_id,
|
||||||
|
status_code=status_code,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to finalize usage on submit failure: request_id=%s, error=%s",
|
||||||
|
self.request_id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
|
||||||
def _build_billing_rule_snapshot(
|
def _build_billing_rule_snapshot(
|
||||||
self, rule_lookup: BillingRuleLookupResult | None
|
self, rule_lookup: BillingRuleLookupResult | None
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
@@ -290,16 +367,16 @@ class VideoHandlerBase(ABC):
|
|||||||
- 无可用候选 / 全部失败:抛 HTTPException(503)
|
- 无可用候选 / 全部失败:抛 HTTPException(503)
|
||||||
"""
|
"""
|
||||||
# 延迟导入,避免 handler 基类层引入过多依赖导致循环
|
# 延迟导入,避免 handler 基类层引入过多依赖导致循环
|
||||||
from src.services.task.orchestrator import (
|
from src.services.candidate.service import CandidateService
|
||||||
|
from src.services.candidate.submit import (
|
||||||
AllCandidatesFailedError,
|
AllCandidatesFailedError,
|
||||||
AsyncTaskOrchestrator,
|
|
||||||
SubmitOutcome,
|
SubmitOutcome,
|
||||||
UpstreamClientRequestError,
|
UpstreamClientRequestError,
|
||||||
)
|
)
|
||||||
|
|
||||||
orchestrator = AsyncTaskOrchestrator(self.db)
|
candidate_service = CandidateService(self.db)
|
||||||
try:
|
try:
|
||||||
return await orchestrator.submit_with_failover(
|
return await candidate_service.submit_with_failover(
|
||||||
api_format=api_format,
|
api_format=api_format,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
affinity_key=str(self.api_key.id),
|
affinity_key=str(self.api_key.id),
|
||||||
@@ -314,8 +391,12 @@ class VideoHandlerBase(ABC):
|
|||||||
max_candidates=max_candidates,
|
max_candidates=max_candidates,
|
||||||
)
|
)
|
||||||
except UpstreamClientRequestError as exc:
|
except UpstreamClientRequestError as exc:
|
||||||
|
# 将 pending usage 结算为 failed,并记录 provider 信息
|
||||||
|
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.response.status_code)
|
||||||
return self._build_error_response(exc.response)
|
return self._build_error_response(exc.response)
|
||||||
except AllCandidatesFailedError as exc:
|
except AllCandidatesFailedError as exc:
|
||||||
|
# 将 pending usage 结算为 failed
|
||||||
|
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.last_status_code)
|
||||||
detail = "No available provider for video generation"
|
detail = "No available provider for video generation"
|
||||||
if config.billing_require_rule:
|
if config.billing_require_rule:
|
||||||
detail = "No available provider with billing rule for video generation"
|
detail = "No available provider with billing rule for video generation"
|
||||||
|
|||||||
@@ -18,7 +18,6 @@ from src.core.api_format import ApiFamily, get_auth_handler
|
|||||||
from src.core.api_format.enums import AuthMethod
|
from src.core.api_format.enums import AuthMethod
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.gemini import GeminiRequest
|
from src.models.gemini import GeminiRequest
|
||||||
from src.services.gemini_files_mapping import extract_file_names_from_request
|
|
||||||
from src.services.provider.transport import redact_url_for_log
|
from src.services.provider.transport import redact_url_for_log
|
||||||
|
|
||||||
|
|
||||||
@@ -63,11 +62,9 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
def detect_capability_requirements(
|
def detect_capability_requirements(
|
||||||
self,
|
self,
|
||||||
headers: dict[str, str], # noqa: ARG002 - 预留
|
headers: dict[str, str], # noqa: ARG002 - 预留
|
||||||
request_body: dict[str, Any] | None = None,
|
request_body: dict[str, Any] | None = None, # noqa: ARG002 - 预留
|
||||||
) -> dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""检测是否需要 Gemini Files API 能力"""
|
"""Gemini API 无特殊能力要求"""
|
||||||
if request_body and extract_file_names_from_request(request_body):
|
|
||||||
return {"gemini_files_api": True}
|
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
def _merge_path_params(
|
def _merge_path_params(
|
||||||
|
|||||||
@@ -40,28 +40,31 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
|
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
|
||||||
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
|
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
|
||||||
|
|
||||||
|
当同一源文件被上传到多个 Key 时,会返回所有可用的 Key ID,
|
||||||
|
让系统能够选择任意可用的 Key。
|
||||||
|
|
||||||
注意事项:
|
注意事项:
|
||||||
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
|
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
|
||||||
- 如果多个文件属于不同 Key,只能使用其中一个,其他文件可能无法访问
|
- 优先返回所有支持该文件的 Key,让调度器选择可用的
|
||||||
"""
|
"""
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.services.gemini_files_mapping import (
|
from src.services.gemini_files_mapping import (
|
||||||
extract_file_names_from_request,
|
extract_file_names_from_request,
|
||||||
get_file_key_mapping,
|
get_all_key_ids_for_file,
|
||||||
)
|
)
|
||||||
|
|
||||||
file_names = extract_file_names_from_request(request_body or {})
|
file_names = extract_file_names_from_request(request_body or {})
|
||||||
if not file_names:
|
if not file_names:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
preferred_key_ids: list[str] = []
|
all_key_ids: set[str] = set()
|
||||||
unmapped_files: list[str] = [] # 记录找不到映射的文件
|
unmapped_files: list[str] = []
|
||||||
|
|
||||||
for file_name in file_names:
|
for file_name in file_names:
|
||||||
key_id = await get_file_key_mapping(file_name)
|
# 获取所有支持该文件的 Key(包括通过 source_hash 关联的)
|
||||||
if key_id:
|
key_ids = await get_all_key_ids_for_file(file_name)
|
||||||
if key_id not in preferred_key_ids:
|
if key_ids:
|
||||||
preferred_key_ids.append(key_id)
|
all_key_ids.update(key_ids)
|
||||||
else:
|
else:
|
||||||
unmapped_files.append(file_name)
|
unmapped_files.append(file_name)
|
||||||
|
|
||||||
@@ -72,14 +75,10 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
"请求可能失败(文件属于其他 Key 或映射已过期)"
|
"请求可能失败(文件属于其他 Key 或映射已过期)"
|
||||||
)
|
)
|
||||||
|
|
||||||
# 警告:多个文件属于不同 Key
|
if all_key_ids:
|
||||||
if len(preferred_key_ids) > 1:
|
logger.debug(f"[{self.request_id}] 文件引用可用的 Key: {list(all_key_ids)}")
|
||||||
logger.warning(
|
|
||||||
f"[{self.request_id}] 请求使用了多个文件,但它们属于不同的 Key: "
|
|
||||||
f"{preferred_key_ids},只能使用第一个 Key,其他文件可能无法访问"
|
|
||||||
)
|
|
||||||
|
|
||||||
return preferred_key_ids or None
|
return list(all_key_ids) if all_key_ids else None
|
||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ Gemini Video Handler - Veo 视频生成实现
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Any, AsyncIterator
|
from typing import Any, AsyncIterator
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
@@ -34,12 +35,14 @@ from src.core.api_format.conversion.internal_video import (
|
|||||||
VideoStatus,
|
VideoStatus,
|
||||||
)
|
)
|
||||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||||
|
from src.core.api_format.conversion.registry import format_conversion_registry
|
||||||
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||||
|
from src.services.usage.service import UsageService
|
||||||
|
|
||||||
|
|
||||||
class GeminiVeoHandler(VideoHandlerBase):
|
class GeminiVeoHandler(VideoHandlerBase):
|
||||||
@@ -92,20 +95,93 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
raise HTTPException(status_code=400, detail=str(e))
|
raise HTTPException(status_code=400, detail=str(e))
|
||||||
|
|
||||||
|
# 异步任务:提前创建 pending usage,便于前端看到“处理中”
|
||||||
|
try:
|
||||||
|
UsageService.create_pending_usage(
|
||||||
|
db=self.db,
|
||||||
|
request_id=self.request_id,
|
||||||
|
user=self.user,
|
||||||
|
api_key=self.api_key,
|
||||||
|
model=internal_request.model,
|
||||||
|
is_stream=False,
|
||||||
|
request_type="video",
|
||||||
|
api_format=self.FORMAT_ID,
|
||||||
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to create pending usage for video request_id=%s: %s",
|
||||||
|
self.request_id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 用于跟踪是否发生了格式转换
|
||||||
|
format_conversion_info: dict[str, Any] = {
|
||||||
|
"converted": False,
|
||||||
|
"provider_format": None,
|
||||||
|
}
|
||||||
|
|
||||||
async def _submit(candidate: ProviderCandidate) -> Any:
|
async def _submit(candidate: ProviderCandidate) -> Any:
|
||||||
upstream_key, endpoint, _key, auth_info = await self._resolve_upstream_key(candidate)
|
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
|
provider_format = make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
)
|
)
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
|
||||||
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
format_conversion_info["provider_format"] = provider_format
|
||||||
|
format_conversion_info["converted"] = needs_conversion
|
||||||
|
|
||||||
|
if needs_conversion and provider_format.upper().startswith("OPENAI:"):
|
||||||
|
# Gemini -> OpenAI 格式转换
|
||||||
|
converted_body = format_conversion_registry.convert_video_request(
|
||||||
|
original_request_body,
|
||||||
|
self.FORMAT_ID,
|
||||||
|
provider_format,
|
||||||
|
)
|
||||||
|
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string)
|
||||||
|
if "seconds" in converted_body and converted_body["seconds"] is not None:
|
||||||
|
converted_body["seconds"] = str(converted_body["seconds"])
|
||||||
|
|
||||||
|
# 构建 OpenAI 风格的 URL
|
||||||
|
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
|
||||||
|
|
||||||
|
# 构建 OpenAI 风格的请求头
|
||||||
|
headers = self._build_openai_upstream_headers(
|
||||||
|
original_headers, upstream_key, endpoint
|
||||||
|
)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
return await client.post(upstream_url, headers=headers, json=converted_body)
|
||||||
|
else:
|
||||||
|
# 原始 Gemini 格式
|
||||||
|
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
|
||||||
|
headers = self._build_upstream_headers(
|
||||||
|
original_headers, upstream_key, endpoint, auth_info
|
||||||
|
)
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||||
|
|
||||||
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
||||||
value = payload.get("name")
|
# 根据响应格式提取 task ID
|
||||||
if not value:
|
# Gemini: {"name": "operations/..."}
|
||||||
return None
|
# OpenAI: {"id": "..."}
|
||||||
return normalize_gemini_operation_id(str(value))
|
if "name" in payload:
|
||||||
|
value = payload.get("name")
|
||||||
|
logger.debug(
|
||||||
|
"[GeminiVeoHandler] Upstream response name=%s, keys=%s",
|
||||||
|
value,
|
||||||
|
list(payload.keys()) if isinstance(payload, dict) else type(payload),
|
||||||
|
)
|
||||||
|
if not value:
|
||||||
|
return None
|
||||||
|
return normalize_gemini_operation_id(str(value))
|
||||||
|
if "id" in payload:
|
||||||
|
# OpenAI 格式
|
||||||
|
return str(payload["id"])
|
||||||
|
return None
|
||||||
|
|
||||||
outcome_or_response = await self._submit_with_failover(
|
outcome_or_response = await self._submit_with_failover(
|
||||||
api_format=self.FORMAT_ID,
|
api_format=self.FORMAT_ID,
|
||||||
@@ -114,7 +190,7 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
submit_func=_submit,
|
submit_func=_submit,
|
||||||
extract_external_task_id=_extract_task_id,
|
extract_external_task_id=_extract_task_id,
|
||||||
supported_auth_types={"api_key", "vertex_ai"},
|
supported_auth_types={"api_key", "vertex_ai"},
|
||||||
allow_format_conversion=False,
|
allow_format_conversion=True,
|
||||||
max_candidates=10,
|
max_candidates=10,
|
||||||
)
|
)
|
||||||
if isinstance(outcome_or_response, JSONResponse):
|
if isinstance(outcome_or_response, JSONResponse):
|
||||||
@@ -135,35 +211,92 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
|
|
||||||
external_task_id = outcome.external_task_id
|
external_task_id = outcome.external_task_id
|
||||||
|
|
||||||
|
# 如果发生了格式转换,记录转换后的请求体
|
||||||
|
converted_request_body = original_request_body
|
||||||
|
if format_conversion_info["converted"]:
|
||||||
|
try:
|
||||||
|
converted_request_body = format_conversion_registry.convert_video_request(
|
||||||
|
original_request_body,
|
||||||
|
self.FORMAT_ID,
|
||||||
|
format_conversion_info["provider_format"],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"[GeminiVeoHandler] Failed to record converted request: %s",
|
||||||
|
sanitize_error_message(str(e)),
|
||||||
|
)
|
||||||
|
|
||||||
task = self._create_task_record(
|
task = self._create_task_record(
|
||||||
external_task_id=external_task_id,
|
external_task_id=external_task_id,
|
||||||
candidate=outcome.candidate,
|
candidate=outcome.candidate,
|
||||||
original_request_body=original_request_body,
|
original_request_body=original_request_body,
|
||||||
|
converted_request_body=converted_request_body,
|
||||||
internal_request=internal_request,
|
internal_request=internal_request,
|
||||||
candidate_keys=outcome.candidate_keys,
|
candidate_keys=outcome.candidate_keys,
|
||||||
original_headers=original_headers,
|
original_headers=original_headers,
|
||||||
billing_rule_snapshot=billing_rule_snapshot,
|
billing_rule_snapshot=billing_rule_snapshot,
|
||||||
|
format_converted=format_conversion_info["converted"],
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
self.db.add(task)
|
self.db.add(task)
|
||||||
self.db.flush() # 先 flush 检测冲突
|
self.db.flush() # 先 flush 检测冲突
|
||||||
self.db.commit()
|
self.db.commit()
|
||||||
self.db.refresh(task)
|
self.db.refresh(task)
|
||||||
logger.info(
|
logger.debug(
|
||||||
f"[GeminiVeoHandler] Task created: id={task.id}, external_task_id={task.external_task_id}, user_id={task.user_id}"
|
"[GeminiVeoHandler] Task created: id=%s, external_task_id=%s",
|
||||||
|
task.id,
|
||||||
|
task.external_task_id,
|
||||||
)
|
)
|
||||||
except IntegrityError:
|
except IntegrityError:
|
||||||
self.db.rollback()
|
self.db.rollback()
|
||||||
raise HTTPException(status_code=409, detail="Task already exists")
|
raise HTTPException(status_code=409, detail="Task already exists")
|
||||||
|
|
||||||
|
# 先构建返回给客户端的响应(使用短 ID 对外暴露)
|
||||||
internal_task = InternalVideoTask(
|
internal_task = InternalVideoTask(
|
||||||
id=task.id,
|
id=task.short_id,
|
||||||
external_id=external_task_id,
|
external_id=external_task_id,
|
||||||
status=VideoStatus.SUBMITTED,
|
status=VideoStatus.SUBMITTED,
|
||||||
created_at=task.created_at,
|
created_at=task.created_at,
|
||||||
original_request=internal_request,
|
original_request=internal_request,
|
||||||
)
|
)
|
||||||
response_body = self._normalizer.video_task_from_internal(internal_task)
|
response_body = self._normalizer.video_task_from_internal(internal_task)
|
||||||
|
|
||||||
|
# 提交成功后立即结算 Usage(费用暂时为 0,轮询完成后更新)
|
||||||
|
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||||
|
try:
|
||||||
|
# 构建发送给上游的请求头(脱敏)
|
||||||
|
upstream_request_headers = self._build_upstream_headers(
|
||||||
|
original_headers,
|
||||||
|
"", # key 不重要,只是用于记录
|
||||||
|
outcome.candidate.endpoint,
|
||||||
|
None, # auth_info
|
||||||
|
)
|
||||||
|
|
||||||
|
UsageService.finalize_submitted(
|
||||||
|
self.db,
|
||||||
|
request_id=self.request_id,
|
||||||
|
provider_name=outcome.candidate.provider.name,
|
||||||
|
provider_id=outcome.candidate.provider.id,
|
||||||
|
provider_endpoint_id=outcome.candidate.endpoint.id,
|
||||||
|
provider_api_key_id=outcome.candidate.key.id,
|
||||||
|
response_time_ms=response_time_ms,
|
||||||
|
status_code=outcome.upstream_status_code or 200,
|
||||||
|
endpoint_api_format=make_signature_key(
|
||||||
|
str(getattr(outcome.candidate.endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(outcome.candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
),
|
||||||
|
provider_request_headers=upstream_request_headers,
|
||||||
|
response_headers=outcome.upstream_headers,
|
||||||
|
response_body=response_body, # 使用我们转换后的响应(包含我们的 ID)
|
||||||
|
)
|
||||||
|
self.db.commit()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to finalize submitted usage for video request_id=%s: %s",
|
||||||
|
self.request_id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
|
||||||
return JSONResponse(response_body)
|
return JSONResponse(response_body)
|
||||||
|
|
||||||
async def handle_get_task(
|
async def handle_get_task(
|
||||||
@@ -234,6 +367,28 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
|
|
||||||
task.status = VideoStatus.CANCELLED.value
|
task.status = VideoStatus.CANCELLED.value
|
||||||
task.updated_at = datetime.now(timezone.utc)
|
task.updated_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# 将 Usage 作废(不收费)
|
||||||
|
# 尝试 finalize_void(处理 pending)和 void_settled(处理已 settled)
|
||||||
|
try:
|
||||||
|
voided = UsageService.finalize_void(
|
||||||
|
self.db,
|
||||||
|
request_id=task.request_id,
|
||||||
|
reason="cancelled_by_user",
|
||||||
|
)
|
||||||
|
if not voided:
|
||||||
|
# pending 状态未找到,尝试处理已 settled 的记录
|
||||||
|
UsageService.void_settled(
|
||||||
|
self.db,
|
||||||
|
request_id=task.request_id,
|
||||||
|
reason="cancelled_by_user",
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to void usage for cancelled task=%s: %s",
|
||||||
|
task.id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
self.db.commit()
|
self.db.commit()
|
||||||
return JSONResponse({})
|
return JSONResponse({})
|
||||||
|
|
||||||
@@ -275,12 +430,38 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
if task.video_expires_at < now:
|
if task.video_expires_at < now:
|
||||||
raise HTTPException(status_code=410, detail="Video URL has expired")
|
raise HTTPException(status_code=410, detail="Video URL has expired")
|
||||||
|
|
||||||
|
# 获取 provider 的认证信息(Gemini 下载视频需要带 API Key)
|
||||||
|
endpoint, key = self._get_endpoint_and_key(task)
|
||||||
|
download_headers: dict[str, str] = {}
|
||||||
|
if key.api_key:
|
||||||
|
try:
|
||||||
|
upstream_key = crypto_service.decrypt(key.api_key)
|
||||||
|
# Gemini API 使用 x-goog-api-key 头进行认证
|
||||||
|
download_headers["x-goog-api-key"] = upstream_key
|
||||||
|
|
||||||
|
# 如果是 Vertex AI,需要使用 OAuth Bearer token
|
||||||
|
auth_info = await get_provider_auth(endpoint, key)
|
||||||
|
if auth_info:
|
||||||
|
download_headers.pop("x-goog-api-key", None)
|
||||||
|
download_headers[auth_info.auth_header] = auth_info.auth_value
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"[VideoDownload] Failed to get auth for download task=%s: %s",
|
||||||
|
task.id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
# 继续尝试无认证下载(某些 URL 可能是预签名的)
|
||||||
|
|
||||||
# 代理下载而非直接重定向,避免暴露上游存储 URL
|
# 代理下载而非直接重定向,避免暴露上游存储 URL
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
# 使用 httpx 支持重定向(Gemini 视频 URL 会重定向到实际存储位置)
|
||||||
|
import httpx
|
||||||
|
|
||||||
try:
|
try:
|
||||||
request = client.build_request("GET", task.video_url)
|
# 使用 follow_redirects=True 跟随重定向
|
||||||
# 视频下载可能较大,设置 5 分钟超时
|
async with httpx.AsyncClient(
|
||||||
response = await client.send(request, stream=True, timeout=300.0)
|
follow_redirects=True, timeout=httpx.Timeout(300.0)
|
||||||
|
) as client:
|
||||||
|
response = await client.get(task.video_url, headers=download_headers)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error(
|
logger.error(
|
||||||
"[VideoDownload] Upstream fetch failed user=%s task=%s: %s",
|
"[VideoDownload] Upstream fetch failed user=%s task=%s: %s",
|
||||||
@@ -291,21 +472,14 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
raise HTTPException(status_code=502, detail="Failed to fetch video")
|
raise HTTPException(status_code=502, detail="Failed to fetch video")
|
||||||
|
|
||||||
if response.status_code >= 400:
|
if response.status_code >= 400:
|
||||||
await response.aclose()
|
|
||||||
raise HTTPException(status_code=response.status_code, detail="Upstream error")
|
raise HTTPException(status_code=response.status_code, detail="Upstream error")
|
||||||
|
|
||||||
async def _iter_bytes() -> AsyncIterator[bytes]:
|
# 返回完整的视频内容(非 streaming,因为需要跟随重定向)
|
||||||
try:
|
|
||||||
async for chunk in response.aiter_bytes():
|
|
||||||
yield chunk
|
|
||||||
finally:
|
|
||||||
await response.aclose()
|
|
||||||
|
|
||||||
safe_headers = {
|
safe_headers = {
|
||||||
k: v for k, v in response.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
|
k: v for k, v in response.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
|
||||||
}
|
}
|
||||||
return StreamingResponse(
|
return Response(
|
||||||
_iter_bytes(),
|
content=response.content,
|
||||||
status_code=response.status_code,
|
status_code=response.status_code,
|
||||||
headers=safe_headers,
|
headers=safe_headers,
|
||||||
media_type=response.headers.get("content-type", "video/mp4"),
|
media_type=response.headers.get("content-type", "video/mp4"),
|
||||||
@@ -375,6 +549,36 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
"status": error.get("status", "BAD_GATEWAY"),
|
"status": error.get("status", "BAD_GATEWAY"),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# OpenAI format conversion helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _build_openai_upstream_url(self, base_url: str | None) -> str:
|
||||||
|
"""构建 OpenAI Sora API 的上游 URL"""
|
||||||
|
base = (base_url or "https://api.openai.com").rstrip("/")
|
||||||
|
if base.endswith("/v1"):
|
||||||
|
return f"{base}/videos"
|
||||||
|
return f"{base}/v1/videos"
|
||||||
|
|
||||||
|
def _build_openai_upstream_headers(
|
||||||
|
self,
|
||||||
|
original_headers: dict[str, str],
|
||||||
|
upstream_key: str,
|
||||||
|
endpoint: ProviderEndpoint,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
"""构建 OpenAI 格式的请求头"""
|
||||||
|
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||||
|
endpoint_sig = make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
return build_upstream_headers_for_endpoint(
|
||||||
|
original_headers,
|
||||||
|
endpoint_sig,
|
||||||
|
upstream_key,
|
||||||
|
endpoint_headers=extra_headers,
|
||||||
|
)
|
||||||
|
|
||||||
def _create_task_record(
|
def _create_task_record(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -385,6 +589,8 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
candidate_keys: list[dict[str, Any]] | None = None,
|
candidate_keys: list[dict[str, Any]] | None = None,
|
||||||
original_headers: dict[str, str] | None = None,
|
original_headers: dict[str, str] | None = None,
|
||||||
billing_rule_snapshot: dict[str, Any] | None = None,
|
billing_rule_snapshot: dict[str, Any] | None = None,
|
||||||
|
converted_request_body: dict[str, Any] | None = None,
|
||||||
|
format_converted: bool = False,
|
||||||
) -> VideoTask:
|
) -> VideoTask:
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
@@ -407,8 +613,14 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
}
|
}
|
||||||
request_metadata["request_headers"] = safe_headers
|
request_metadata["request_headers"] = safe_headers
|
||||||
|
|
||||||
|
provider_api_format = make_signature_key(
|
||||||
|
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
|
||||||
return VideoTask(
|
return VideoTask(
|
||||||
id=str(uuid4()),
|
id=str(uuid4()),
|
||||||
|
request_id=self.request_id,
|
||||||
external_task_id=external_task_id,
|
external_task_id=external_task_id,
|
||||||
user_id=self.user.id,
|
user_id=self.user.id,
|
||||||
api_key_id=self.api_key.id,
|
api_key_id=self.api_key.id,
|
||||||
@@ -416,15 +628,12 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
endpoint_id=candidate.endpoint.id,
|
endpoint_id=candidate.endpoint.id,
|
||||||
key_id=candidate.key.id,
|
key_id=candidate.key.id,
|
||||||
client_api_format=self.FORMAT_ID,
|
client_api_format=self.FORMAT_ID,
|
||||||
provider_api_format=make_signature_key(
|
provider_api_format=provider_api_format,
|
||||||
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
format_converted=format_converted,
|
||||||
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
|
||||||
),
|
|
||||||
format_converted=False,
|
|
||||||
model=internal_request.model,
|
model=internal_request.model,
|
||||||
prompt=internal_request.prompt,
|
prompt=internal_request.prompt,
|
||||||
original_request_body=original_request_body,
|
original_request_body=original_request_body,
|
||||||
converted_request_body=original_request_body,
|
converted_request_body=converted_request_body or original_request_body,
|
||||||
duration_seconds=internal_request.duration_seconds,
|
duration_seconds=internal_request.duration_seconds,
|
||||||
resolution=internal_request.resolution,
|
resolution=internal_request.resolution,
|
||||||
aspect_ratio=internal_request.aspect_ratio,
|
aspect_ratio=internal_request.aspect_ratio,
|
||||||
@@ -438,30 +647,45 @@ class GeminiVeoHandler(VideoHandlerBase):
|
|||||||
request_metadata=request_metadata,
|
request_metadata=request_metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
|
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
|
||||||
"""按 external_task_id 查找任务(Gemini 使用 operations/{id} 格式)"""
|
"""覆盖父类方法,Gemini 使用 short_id 作为对外暴露的 ID"""
|
||||||
normalized_id = normalize_gemini_operation_id(external_id)
|
try:
|
||||||
|
status = VideoStatus(task.status)
|
||||||
logger.info(
|
except ValueError:
|
||||||
f"[GeminiVeoHandler] Looking for task: normalized_id={normalized_id}, user_id={self.user.id}"
|
status = VideoStatus.PENDING
|
||||||
|
return InternalVideoTask(
|
||||||
|
id=task.short_id, # Gemini 使用短 ID
|
||||||
|
external_id=task.external_task_id,
|
||||||
|
status=status,
|
||||||
|
progress_percent=task.progress_percent or 0,
|
||||||
|
progress_message=task.progress_message,
|
||||||
|
video_url=task.video_url,
|
||||||
|
video_urls=task.video_urls or [],
|
||||||
|
created_at=task.created_at,
|
||||||
|
completed_at=task.completed_at,
|
||||||
|
error_code=task.error_code,
|
||||||
|
error_message=task.error_message,
|
||||||
|
extra={"model": task.model},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
|
||||||
|
"""按 short_id 查找任务(我们对外暴露的 operation 格式是 models/{model}/operations/{short_id})"""
|
||||||
|
from src.api.handlers.base.video_handler_base import extract_short_id_from_operation
|
||||||
|
|
||||||
|
short_id = extract_short_id_from_operation(external_id)
|
||||||
|
|
||||||
|
# 通过 short_id 查找任务
|
||||||
task = (
|
task = (
|
||||||
self.db.query(VideoTask)
|
self.db.query(VideoTask)
|
||||||
.filter(
|
.filter(
|
||||||
VideoTask.external_task_id == normalized_id,
|
VideoTask.short_id == short_id,
|
||||||
VideoTask.user_id == self.user.id,
|
VideoTask.user_id == self.user.id,
|
||||||
)
|
)
|
||||||
.first()
|
.first()
|
||||||
)
|
)
|
||||||
if not task:
|
if not task:
|
||||||
logger.warning(
|
logger.debug("[GeminiVeoHandler] Task not found: short_id=%s", short_id)
|
||||||
f"[GeminiVeoHandler] Task not found: normalized_id={normalized_id}, user_id={self.user.id}"
|
|
||||||
)
|
|
||||||
raise HTTPException(status_code=404, detail="Video task not found")
|
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
|
return task
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ OpenAI Video Handler - Sora 视频生成实现
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import time
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Any, AsyncIterator
|
from typing import Any, AsyncIterator
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
@@ -14,6 +15,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
|
|||||||
from sqlalchemy.exc import IntegrityError
|
from sqlalchemy.exc import IntegrityError
|
||||||
from sqlalchemy.orm import Session
|
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, sanitize_error_message
|
||||||
from src.clients.http_client import HTTPClientPool
|
from src.clients.http_client import HTTPClientPool
|
||||||
from src.config.settings import config
|
from src.config.settings import config
|
||||||
@@ -30,6 +32,7 @@ from src.core.api_format.conversion.internal_video import (
|
|||||||
VideoStatus,
|
VideoStatus,
|
||||||
)
|
)
|
||||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||||
|
from src.core.api_format.conversion.registry import format_conversion_registry
|
||||||
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
@@ -91,16 +94,91 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
)
|
)
|
||||||
raise HTTPException(status_code=400, detail=str(e))
|
raise HTTPException(status_code=400, detail=str(e))
|
||||||
|
|
||||||
|
# 异步任务:提前创建 pending usage,便于前端看到“处理中”
|
||||||
|
try:
|
||||||
|
UsageService.create_pending_usage(
|
||||||
|
db=self.db,
|
||||||
|
request_id=self.request_id,
|
||||||
|
user=self.user,
|
||||||
|
api_key=self.api_key,
|
||||||
|
model=internal_request.model,
|
||||||
|
is_stream=False,
|
||||||
|
request_type="video",
|
||||||
|
api_format=self.FORMAT_ID,
|
||||||
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to create pending usage for video request_id=%s: %s",
|
||||||
|
self.request_id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 用于跟踪是否发生了格式转换
|
||||||
|
format_conversion_info: dict[str, Any] = {
|
||||||
|
"converted": False,
|
||||||
|
"provider_format": None,
|
||||||
|
}
|
||||||
|
|
||||||
async def _submit(candidate: ProviderCandidate) -> Any:
|
async def _submit(candidate: ProviderCandidate) -> Any:
|
||||||
upstream_key, endpoint, _provider_key = await self._resolve_upstream_key(candidate)
|
upstream_key, endpoint, _provider_key = await self._resolve_upstream_key(candidate)
|
||||||
upstream_url = self._build_upstream_url(endpoint.base_url)
|
|
||||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
# 检测目标格式
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
provider_format = make_signature_key(
|
||||||
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
|
||||||
|
format_conversion_info["provider_format"] = provider_format
|
||||||
|
format_conversion_info["converted"] = needs_conversion
|
||||||
|
|
||||||
|
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string)
|
||||||
|
request_body = original_request_body.copy()
|
||||||
|
if "seconds" in request_body and request_body["seconds"] is not None:
|
||||||
|
request_body["seconds"] = str(request_body["seconds"])
|
||||||
|
|
||||||
|
if needs_conversion and provider_format.upper().startswith("GEMINI:"):
|
||||||
|
# OpenAI -> Gemini 格式转换
|
||||||
|
converted_body = format_conversion_registry.convert_video_request(
|
||||||
|
request_body,
|
||||||
|
self.FORMAT_ID,
|
||||||
|
provider_format,
|
||||||
|
)
|
||||||
|
# 如果 model 不在请求体中,从路径或内部请求中获取
|
||||||
|
if "model" not in converted_body:
|
||||||
|
converted_body["model"] = internal_request.model
|
||||||
|
|
||||||
|
# 构建 Gemini 风格的 URL
|
||||||
|
upstream_url = self._build_gemini_upstream_url(
|
||||||
|
endpoint.base_url, internal_request.model
|
||||||
|
)
|
||||||
|
|
||||||
|
# 构建 Gemini 风格的请求头
|
||||||
|
auth_info = await get_provider_auth(endpoint, _provider_key)
|
||||||
|
headers = self._build_gemini_upstream_headers(
|
||||||
|
original_headers, upstream_key, endpoint, auth_info
|
||||||
|
)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
return await client.post(upstream_url, headers=headers, json=converted_body)
|
||||||
|
else:
|
||||||
|
# 原始 OpenAI 格式
|
||||||
|
upstream_url = self._build_upstream_url(endpoint.base_url)
|
||||||
|
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
return await client.post(upstream_url, headers=headers, json=request_body)
|
||||||
|
|
||||||
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
||||||
value = payload.get("id")
|
# 根据响应格式提取 task ID
|
||||||
return str(value) if value else None
|
# OpenAI: {"id": "..."}
|
||||||
|
# Gemini: {"name": "operations/..."}
|
||||||
|
if "id" in payload:
|
||||||
|
return str(payload["id"])
|
||||||
|
if "name" in payload:
|
||||||
|
# Gemini 格式
|
||||||
|
return str(payload["name"])
|
||||||
|
return None
|
||||||
|
|
||||||
# 捕获提交阶段的所有错误,记录失败任务
|
# 捕获提交阶段的所有错误,记录失败任务
|
||||||
try:
|
try:
|
||||||
@@ -110,8 +188,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
task_type="video",
|
task_type="video",
|
||||||
submit_func=_submit,
|
submit_func=_submit,
|
||||||
extract_external_task_id=_extract_task_id,
|
extract_external_task_id=_extract_task_id,
|
||||||
supported_auth_types={"api_key"},
|
supported_auth_types={"api_key", "vertex_ai"},
|
||||||
allow_format_conversion=False,
|
allow_format_conversion=True,
|
||||||
max_candidates=10,
|
max_candidates=10,
|
||||||
)
|
)
|
||||||
except HTTPException as exc:
|
except HTTPException as exc:
|
||||||
@@ -156,14 +234,31 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
|
|
||||||
external_task_id = outcome.external_task_id
|
external_task_id = outcome.external_task_id
|
||||||
|
|
||||||
|
# 如果发生了格式转换,记录转换后的请求体
|
||||||
|
converted_request_body = original_request_body
|
||||||
|
if format_conversion_info["converted"]:
|
||||||
|
try:
|
||||||
|
converted_request_body = format_conversion_registry.convert_video_request(
|
||||||
|
original_request_body,
|
||||||
|
self.FORMAT_ID,
|
||||||
|
format_conversion_info["provider_format"],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"[OpenAIVideoHandler] Failed to record converted request: %s",
|
||||||
|
sanitize_error_message(str(e)),
|
||||||
|
)
|
||||||
|
|
||||||
task = self._create_task_record(
|
task = self._create_task_record(
|
||||||
external_task_id=external_task_id,
|
external_task_id=external_task_id,
|
||||||
candidate=outcome.candidate,
|
candidate=outcome.candidate,
|
||||||
original_request_body=original_request_body,
|
original_request_body=original_request_body,
|
||||||
|
converted_request_body=converted_request_body,
|
||||||
internal_request=internal_request,
|
internal_request=internal_request,
|
||||||
candidate_keys=outcome.candidate_keys,
|
candidate_keys=outcome.candidate_keys,
|
||||||
original_headers=original_headers,
|
original_headers=original_headers,
|
||||||
billing_rule_snapshot=billing_rule_snapshot,
|
billing_rule_snapshot=billing_rule_snapshot,
|
||||||
|
format_converted=format_conversion_info["converted"],
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
self.db.add(task)
|
self.db.add(task)
|
||||||
@@ -174,6 +269,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
self.db.rollback()
|
self.db.rollback()
|
||||||
raise HTTPException(status_code=409, detail="Task already exists")
|
raise HTTPException(status_code=409, detail="Task already exists")
|
||||||
|
|
||||||
|
# 先构建返回给客户端的响应(OpenAI Sora 使用 UUID)
|
||||||
internal_task = InternalVideoTask(
|
internal_task = InternalVideoTask(
|
||||||
id=task.id,
|
id=task.id,
|
||||||
external_id=external_task_id,
|
external_id=external_task_id,
|
||||||
@@ -182,6 +278,42 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
original_request=internal_request,
|
original_request=internal_request,
|
||||||
)
|
)
|
||||||
response_body = self._normalizer.video_task_from_internal(internal_task)
|
response_body = self._normalizer.video_task_from_internal(internal_task)
|
||||||
|
|
||||||
|
# 提交成功后立即结算 Usage(费用暂时为 0,轮询完成后更新)
|
||||||
|
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||||
|
try:
|
||||||
|
# 构建发送给上游的请求头(脱敏)
|
||||||
|
upstream_request_headers = self._build_upstream_headers(
|
||||||
|
original_headers,
|
||||||
|
"", # key 不重要,只是用于记录
|
||||||
|
outcome.candidate.endpoint,
|
||||||
|
)
|
||||||
|
|
||||||
|
UsageService.finalize_submitted(
|
||||||
|
self.db,
|
||||||
|
request_id=self.request_id,
|
||||||
|
provider_name=outcome.candidate.provider.name,
|
||||||
|
provider_id=outcome.candidate.provider.id,
|
||||||
|
provider_endpoint_id=outcome.candidate.endpoint.id,
|
||||||
|
provider_api_key_id=outcome.candidate.key.id,
|
||||||
|
response_time_ms=response_time_ms,
|
||||||
|
status_code=outcome.upstream_status_code or 200,
|
||||||
|
endpoint_api_format=make_signature_key(
|
||||||
|
str(getattr(outcome.candidate.endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(outcome.candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
),
|
||||||
|
provider_request_headers=upstream_request_headers,
|
||||||
|
response_headers=outcome.upstream_headers,
|
||||||
|
response_body=response_body, # 使用我们转换后的响应(包含我们的 ID)
|
||||||
|
)
|
||||||
|
self.db.commit()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to finalize submitted usage for video request_id=%s: %s",
|
||||||
|
self.request_id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
|
||||||
return JSONResponse(response_body)
|
return JSONResponse(response_body)
|
||||||
|
|
||||||
async def handle_get_task(
|
async def handle_get_task(
|
||||||
@@ -206,17 +338,59 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
query_params: dict[str, str] | None = None,
|
query_params: dict[str, str] | None = None,
|
||||||
path_params: dict[str, Any] | None = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
tasks = (
|
params = query_params or {}
|
||||||
self.db.query(VideoTask)
|
|
||||||
.filter(VideoTask.user_id == self.user.id)
|
# 解析分页参数
|
||||||
.order_by(VideoTask.created_at.desc())
|
after = params.get("after")
|
||||||
.limit(100)
|
try:
|
||||||
.all()
|
limit = min(int(params.get("limit") or 20), 100) # 默认 20,最大 100
|
||||||
)
|
except (ValueError, TypeError):
|
||||||
|
limit = 20
|
||||||
|
order = params.get("order", "desc").lower()
|
||||||
|
if order not in ("asc", "desc"):
|
||||||
|
order = "desc"
|
||||||
|
|
||||||
|
# 构建查询
|
||||||
|
query = self.db.query(VideoTask).filter(VideoTask.user_id == self.user.id)
|
||||||
|
|
||||||
|
# 处理游标分页(after 参数使用 UUID)
|
||||||
|
if after:
|
||||||
|
after_task = (
|
||||||
|
self.db.query(VideoTask)
|
||||||
|
.filter(VideoTask.id == after, VideoTask.user_id == self.user.id)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
if after_task and after_task.created_at:
|
||||||
|
if order == "desc":
|
||||||
|
query = query.filter(VideoTask.created_at < after_task.created_at)
|
||||||
|
else:
|
||||||
|
query = query.filter(VideoTask.created_at > after_task.created_at)
|
||||||
|
|
||||||
|
# 排序
|
||||||
|
if order == "asc":
|
||||||
|
query = query.order_by(VideoTask.created_at.asc())
|
||||||
|
else:
|
||||||
|
query = query.order_by(VideoTask.created_at.desc())
|
||||||
|
|
||||||
|
# 获取 limit + 1 条记录以判断是否有更多数据
|
||||||
|
tasks = query.limit(limit + 1).all()
|
||||||
|
has_more = len(tasks) > limit
|
||||||
|
tasks = tasks[:limit]
|
||||||
|
|
||||||
items = [
|
items = [
|
||||||
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
|
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
|
||||||
]
|
]
|
||||||
return JSONResponse({"object": "list", "data": items})
|
|
||||||
|
response_data: dict[str, Any] = {
|
||||||
|
"object": "list",
|
||||||
|
"data": items,
|
||||||
|
"has_more": has_more,
|
||||||
|
}
|
||||||
|
# 如果有更多数据,返回最后一条的 ID 作为下一页游标
|
||||||
|
if has_more and tasks:
|
||||||
|
response_data["last_id"] = tasks[-1].id
|
||||||
|
|
||||||
|
return JSONResponse(response_data)
|
||||||
|
|
||||||
async def handle_cancel_task(
|
async def handle_cancel_task(
|
||||||
self,
|
self,
|
||||||
@@ -245,9 +419,80 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
|
|
||||||
task.status = VideoStatus.CANCELLED.value
|
task.status = VideoStatus.CANCELLED.value
|
||||||
task.updated_at = datetime.now(timezone.utc)
|
task.updated_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# 将 Usage 作废(不收费)
|
||||||
|
# 尝试 finalize_void(处理 pending)和 void_settled(处理已 settled)
|
||||||
|
try:
|
||||||
|
voided = UsageService.finalize_void(
|
||||||
|
self.db,
|
||||||
|
request_id=task.request_id,
|
||||||
|
reason="cancelled_by_user",
|
||||||
|
)
|
||||||
|
if not voided:
|
||||||
|
# pending 状态未找到,尝试处理已 settled 的记录
|
||||||
|
UsageService.void_settled(
|
||||||
|
self.db,
|
||||||
|
request_id=task.request_id,
|
||||||
|
reason="cancelled_by_user",
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to void usage for cancelled task=%s: %s",
|
||||||
|
task.id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
self.db.commit()
|
self.db.commit()
|
||||||
return JSONResponse({})
|
return JSONResponse({})
|
||||||
|
|
||||||
|
async def handle_delete_task(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
task_id: str,
|
||||||
|
http_request: Request,
|
||||||
|
original_headers: dict[str, str],
|
||||||
|
query_params: dict[str, str] | None = None,
|
||||||
|
path_params: dict[str, Any] | None = None,
|
||||||
|
) -> JSONResponse:
|
||||||
|
"""删除已完成或失败的视频及其存储资源"""
|
||||||
|
task = self._get_task(task_id)
|
||||||
|
|
||||||
|
# 只能删除已完成或失败的视频
|
||||||
|
if task.status not in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail=f"Can only delete completed or failed videos (current status: {task.status})",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 如果有 external_task_id,向上游发送删除请求
|
||||||
|
if task.external_task_id:
|
||||||
|
try:
|
||||||
|
endpoint, key = self._get_endpoint_and_key(task)
|
||||||
|
if key.api_key:
|
||||||
|
upstream_key = crypto_service.decrypt(key.api_key)
|
||||||
|
upstream_url = self._build_upstream_url(
|
||||||
|
endpoint.base_url, task.external_task_id
|
||||||
|
)
|
||||||
|
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
response = await client.delete(upstream_url, headers=headers)
|
||||||
|
if response.status_code >= 400 and response.status_code != 404:
|
||||||
|
# 404 表示上游已删除,不算错误
|
||||||
|
return self._build_error_response(response)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to delete video from upstream task=%s: %s",
|
||||||
|
task.id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
# 继续删除本地记录
|
||||||
|
|
||||||
|
# 删除本地任务记录
|
||||||
|
self.db.delete(task)
|
||||||
|
self.db.commit()
|
||||||
|
|
||||||
|
return JSONResponse({"id": task_id, "object": "video", "deleted": True})
|
||||||
|
|
||||||
async def handle_remix_task(
|
async def handle_remix_task(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -280,8 +525,13 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
)
|
)
|
||||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||||
|
|
||||||
|
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string)
|
||||||
|
request_body = original_request_body.copy()
|
||||||
|
if "seconds" in request_body and request_body["seconds"] is not None:
|
||||||
|
request_body["seconds"] = str(request_body["seconds"])
|
||||||
|
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
response = await client.post(upstream_url, headers=headers, json=original_request_body)
|
response = await client.post(upstream_url, headers=headers, json=request_body)
|
||||||
|
|
||||||
if response.status_code >= 400:
|
if response.status_code >= 400:
|
||||||
return self._build_error_response(response)
|
return self._build_error_response(response)
|
||||||
@@ -350,7 +600,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
raise HTTPException(status_code=409, detail="Task already exists")
|
raise HTTPException(status_code=409, detail="Task already exists")
|
||||||
|
|
||||||
internal_task = InternalVideoTask(
|
internal_task = InternalVideoTask(
|
||||||
id=task.id,
|
id=task.id, # OpenAI Sora 使用 UUID
|
||||||
external_id=external_task_id,
|
external_id=external_task_id,
|
||||||
status=VideoStatus.SUBMITTED,
|
status=VideoStatus.SUBMITTED,
|
||||||
created_at=task.created_at,
|
created_at=task.created_at,
|
||||||
@@ -389,15 +639,25 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
if task.status == VideoStatus.CANCELLED.value:
|
if task.status == VideoStatus.CANCELLED.value:
|
||||||
raise HTTPException(status_code=404, detail="Video task was cancelled")
|
raise HTTPException(status_code=404, detail="Video task was cancelled")
|
||||||
|
|
||||||
|
# 支持 variant 查询参数: video (默认), thumbnail, spritesheet
|
||||||
|
variant = (query_params or {}).get("variant", "video")
|
||||||
|
|
||||||
|
# 如果 video_url 是完整的 HTTP URL,直接代理该 URL(适用于不支持 /content 端点的上游如 API易)
|
||||||
|
# 保持流式代理而非重定向,确保客户端行为与官方 OpenAI 一致
|
||||||
|
if variant == "video" and task.video_url and task.video_url.startswith("http"):
|
||||||
|
logger.debug(
|
||||||
|
"[VideoDownload] Proxying direct URL task=%s url=%s",
|
||||||
|
task_id,
|
||||||
|
task.video_url,
|
||||||
|
)
|
||||||
|
return await self._proxy_direct_url(task.video_url, task_id)
|
||||||
|
|
||||||
if not task.external_task_id:
|
if not task.external_task_id:
|
||||||
raise HTTPException(status_code=500, detail="Task missing external_task_id")
|
raise HTTPException(status_code=500, detail="Task missing external_task_id")
|
||||||
endpoint, key = self._get_endpoint_and_key(task)
|
endpoint, key = self._get_endpoint_and_key(task)
|
||||||
if not key.api_key:
|
if not key.api_key:
|
||||||
raise HTTPException(status_code=500, detail="Provider key not configured")
|
raise HTTPException(status_code=500, detail="Provider key not configured")
|
||||||
upstream_key = crypto_service.decrypt(key.api_key)
|
upstream_key = crypto_service.decrypt(key.api_key)
|
||||||
|
|
||||||
# 支持 variant 查询参数: video (默认), thumbnail, spritesheet
|
|
||||||
variant = (query_params or {}).get("variant", "video")
|
|
||||||
if variant not in {"video", "thumbnail", "spritesheet"}:
|
if variant not in {"video", "thumbnail", "spritesheet"}:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=400,
|
status_code=400,
|
||||||
@@ -412,6 +672,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||||
|
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
logger.debug(
|
||||||
|
"[VideoDownload] Requesting upstream url=%s task=%s external_task_id=%s",
|
||||||
|
upstream_url,
|
||||||
|
task_id,
|
||||||
|
task.external_task_id,
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
# 使用 httpx 的 stream 方法并正确管理上下文
|
# 使用 httpx 的 stream 方法并正确管理上下文
|
||||||
# 视频下载可能较大,设置 5 分钟超时
|
# 视频下载可能较大,设置 5 分钟超时
|
||||||
@@ -419,8 +685,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
response = await client.send(request, stream=True, timeout=300.0)
|
response = await client.send(request, stream=True, timeout=300.0)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"[VideoDownload] Upstream connection failed task=%s: %s",
|
"[VideoDownload] Upstream connection failed task=%s url=%s: %s",
|
||||||
task_id,
|
task_id,
|
||||||
|
upstream_url,
|
||||||
sanitize_error_message(str(exc)),
|
sanitize_error_message(str(exc)),
|
||||||
)
|
)
|
||||||
raise HTTPException(status_code=502, detail="Upstream connection failed") from exc
|
raise HTTPException(status_code=502, detail="Upstream connection failed") from exc
|
||||||
@@ -483,6 +750,46 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
|
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
|
||||||
return upstream_key, candidate.endpoint, candidate.key
|
return upstream_key, candidate.endpoint, candidate.key
|
||||||
|
|
||||||
|
async def _proxy_direct_url(self, url: str, task_id: str) -> Response | StreamingResponse:
|
||||||
|
"""代理直接的视频 URL(如 CDN URL),保持与官方 API 一致的流式返回行为"""
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
try:
|
||||||
|
request = client.build_request("GET", url)
|
||||||
|
response = await client.send(request, stream=True, timeout=300.0)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"[VideoDownload] Direct URL connection failed task=%s url=%s: %s",
|
||||||
|
task_id,
|
||||||
|
url,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
raise HTTPException(status_code=502, detail="Video download failed") from exc
|
||||||
|
|
||||||
|
if response.status_code >= 400:
|
||||||
|
await response.aread() # consume body before closing
|
||||||
|
await response.aclose()
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=response.status_code,
|
||||||
|
content={"error": {"type": "upstream_error", "message": "Video not available"}},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _iter_bytes() -> AsyncIterator[bytes]:
|
||||||
|
try:
|
||||||
|
async for chunk in response.aiter_bytes():
|
||||||
|
yield chunk
|
||||||
|
finally:
|
||||||
|
await response.aclose()
|
||||||
|
|
||||||
|
safe_headers = {
|
||||||
|
k: v for k, v in response.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
|
||||||
|
}
|
||||||
|
return StreamingResponse(
|
||||||
|
_iter_bytes(),
|
||||||
|
status_code=response.status_code,
|
||||||
|
headers=safe_headers,
|
||||||
|
media_type=response.headers.get("content-type", "video/mp4"),
|
||||||
|
)
|
||||||
|
|
||||||
def _build_upstream_url(self, base_url: str | None, suffix: str | None = None) -> str:
|
def _build_upstream_url(self, base_url: str | None, suffix: str | None = None) -> str:
|
||||||
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
|
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
|
||||||
if base.endswith("/v1"):
|
if base.endswith("/v1"):
|
||||||
@@ -508,6 +815,42 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
endpoint_headers=extra_headers,
|
endpoint_headers=extra_headers,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Gemini format conversion helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _build_gemini_upstream_url(self, base_url: str | None, model: str) -> str:
|
||||||
|
"""构建 Gemini Veo API 的上游 URL"""
|
||||||
|
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
|
||||||
|
if base.endswith("/v1beta"):
|
||||||
|
base = base[: -len("/v1beta")]
|
||||||
|
return f"{base}/v1beta/models/{model}:predictLongRunning"
|
||||||
|
|
||||||
|
def _build_gemini_upstream_headers(
|
||||||
|
self,
|
||||||
|
original_headers: dict[str, str],
|
||||||
|
upstream_key: str,
|
||||||
|
endpoint: ProviderEndpoint,
|
||||||
|
auth_info: Any | None,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
"""构建 Gemini 格式的请求头"""
|
||||||
|
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||||
|
endpoint_sig = make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
headers = build_upstream_headers_for_endpoint(
|
||||||
|
original_headers,
|
||||||
|
endpoint_sig,
|
||||||
|
upstream_key,
|
||||||
|
endpoint_headers=extra_headers,
|
||||||
|
)
|
||||||
|
if auth_info:
|
||||||
|
# 覆盖为 OAuth2 Bearer(Vertex AI)
|
||||||
|
headers.pop("x-goog-api-key", None)
|
||||||
|
headers[auth_info.auth_header] = auth_info.auth_value
|
||||||
|
return headers
|
||||||
|
|
||||||
# _build_error_response 继承自基类 VideoHandlerBase
|
# _build_error_response 继承自基类 VideoHandlerBase
|
||||||
|
|
||||||
def _create_task_record(
|
def _create_task_record(
|
||||||
@@ -520,6 +863,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
candidate_keys: list[dict[str, Any]] | None = None,
|
candidate_keys: list[dict[str, Any]] | None = None,
|
||||||
original_headers: dict[str, str] | None = None,
|
original_headers: dict[str, str] | None = None,
|
||||||
billing_rule_snapshot: dict[str, Any] | None = None,
|
billing_rule_snapshot: dict[str, Any] | None = None,
|
||||||
|
converted_request_body: dict[str, Any] | None = None,
|
||||||
|
format_converted: bool = False,
|
||||||
) -> VideoTask:
|
) -> VideoTask:
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
size = internal_request.extra.get("original_size")
|
size = internal_request.extra.get("original_size")
|
||||||
@@ -543,8 +888,14 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
}
|
}
|
||||||
request_metadata["request_headers"] = safe_headers
|
request_metadata["request_headers"] = safe_headers
|
||||||
|
|
||||||
|
provider_api_format = make_signature_key(
|
||||||
|
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
|
||||||
return VideoTask(
|
return VideoTask(
|
||||||
id=str(uuid4()),
|
id=str(uuid4()),
|
||||||
|
request_id=self.request_id,
|
||||||
external_task_id=external_task_id,
|
external_task_id=external_task_id,
|
||||||
user_id=self.user.id,
|
user_id=self.user.id,
|
||||||
api_key_id=self.api_key.id,
|
api_key_id=self.api_key.id,
|
||||||
@@ -552,15 +903,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
endpoint_id=candidate.endpoint.id,
|
endpoint_id=candidate.endpoint.id,
|
||||||
key_id=candidate.key.id,
|
key_id=candidate.key.id,
|
||||||
client_api_format=self.FORMAT_ID,
|
client_api_format=self.FORMAT_ID,
|
||||||
provider_api_format=make_signature_key(
|
provider_api_format=provider_api_format,
|
||||||
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
format_converted=format_converted,
|
||||||
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
|
||||||
),
|
|
||||||
format_converted=False,
|
|
||||||
model=internal_request.model,
|
model=internal_request.model,
|
||||||
prompt=internal_request.prompt,
|
prompt=internal_request.prompt,
|
||||||
original_request_body=original_request_body,
|
original_request_body=original_request_body,
|
||||||
converted_request_body=original_request_body,
|
converted_request_body=converted_request_body or original_request_body,
|
||||||
duration_seconds=internal_request.duration_seconds,
|
duration_seconds=internal_request.duration_seconds,
|
||||||
resolution=internal_request.resolution,
|
resolution=internal_request.resolution,
|
||||||
aspect_ratio=internal_request.aspect_ratio,
|
aspect_ratio=internal_request.aspect_ratio,
|
||||||
@@ -580,8 +928,23 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
status = VideoStatus(task.status)
|
status = VideoStatus(task.status)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
status = VideoStatus.PENDING
|
status = VideoStatus.PENDING
|
||||||
|
|
||||||
|
# 构建 extra 字段
|
||||||
|
extra: dict[str, Any] = {
|
||||||
|
"model": task.model,
|
||||||
|
"size": task.size,
|
||||||
|
"seconds": str(task.duration_seconds) if task.duration_seconds else None,
|
||||||
|
"prompt": task.prompt,
|
||||||
|
}
|
||||||
|
|
||||||
|
# 检查是否是 remix 视频
|
||||||
|
if task.original_request_body and isinstance(task.original_request_body, dict):
|
||||||
|
remixed_from = task.original_request_body.get("remix_video_id")
|
||||||
|
if remixed_from:
|
||||||
|
extra["remixed_from_video_id"] = remixed_from
|
||||||
|
|
||||||
return InternalVideoTask(
|
return InternalVideoTask(
|
||||||
id=task.id,
|
id=task.id, # OpenAI Sora 使用 UUID
|
||||||
external_id=task.external_task_id,
|
external_id=task.external_task_id,
|
||||||
status=status,
|
status=status,
|
||||||
progress_percent=task.progress_percent or 0,
|
progress_percent=task.progress_percent or 0,
|
||||||
@@ -596,7 +959,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
expires_at=task.video_expires_at,
|
expires_at=task.video_expires_at,
|
||||||
error_code=task.error_code,
|
error_code=task.error_code,
|
||||||
error_message=task.error_message,
|
error_message=task.error_message,
|
||||||
extra={"model": task.model, "size": task.size, "seconds": task.duration_seconds},
|
extra=extra,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _record_failed_usage(
|
async def _record_failed_usage(
|
||||||
@@ -609,8 +972,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
original_headers: dict[str, str],
|
original_headers: dict[str, str],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录失败请求的使用记录(无任务记录)"""
|
"""记录失败请求的使用记录(无任务记录)"""
|
||||||
import time
|
|
||||||
|
|
||||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||||
safe_headers = {
|
safe_headers = {
|
||||||
k: v
|
k: v
|
||||||
@@ -669,8 +1030,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
candidate_keys: list[dict[str, Any]] | None = None,
|
candidate_keys: list[dict[str, Any]] | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""创建失败的任务记录和使用记录"""
|
"""创建失败的任务记录和使用记录"""
|
||||||
import time
|
|
||||||
|
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||||
|
|
||||||
@@ -696,6 +1055,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
|||||||
# 创建失败的任务记录
|
# 创建失败的任务记录
|
||||||
task = VideoTask(
|
task = VideoTask(
|
||||||
id=str(uuid4()),
|
id=str(uuid4()),
|
||||||
|
request_id=self.request_id,
|
||||||
external_task_id=None,
|
external_task_id=None,
|
||||||
user_id=self.user.id,
|
user_id=self.user.id,
|
||||||
api_key_id=self.api_key.id,
|
api_key_id=self.api_key.id,
|
||||||
|
|||||||
@@ -14,13 +14,14 @@ from .system_catalog import router as system_catalog_router
|
|||||||
from .videos import router as videos_router
|
from .videos import router as videos_router
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
# Models API 需要在最前面注册,避免被其他路由的 path 参数捕获
|
# Video API 路由需要在 Models API 之前注册,因为 Models API 有 /v1beta/models/{path} 通配符路由
|
||||||
|
# 会错误匹配 /v1beta/models/{model}/operations/{id}/content 等视频路由
|
||||||
|
router.include_router(videos_router, tags=["Video Generation"])
|
||||||
router.include_router(models_router)
|
router.include_router(models_router)
|
||||||
router.include_router(claude_router, tags=["Claude API"])
|
router.include_router(claude_router, tags=["Claude API"])
|
||||||
router.include_router(openai_router)
|
router.include_router(openai_router)
|
||||||
router.include_router(gemini_router, tags=["Gemini API"])
|
router.include_router(gemini_router, tags=["Gemini API"])
|
||||||
router.include_router(gemini_files_router, tags=["Gemini Files API"])
|
router.include_router(gemini_files_router, tags=["Gemini Files API"])
|
||||||
router.include_router(videos_router, tags=["Video Generation"])
|
|
||||||
router.include_router(system_catalog_router, tags=["System Catalog"])
|
router.include_router(system_catalog_router, tags=["System Catalog"])
|
||||||
router.include_router(catalog_router)
|
router.include_router(catalog_router)
|
||||||
router.include_router(capabilities_router)
|
router.include_router(capabilities_router)
|
||||||
|
|||||||
@@ -15,14 +15,17 @@ Gemini Files API 代理端点
|
|||||||
|
|
||||||
参考文档:
|
参考文档:
|
||||||
https://ai.google.dev/api/files
|
https://ai.google.dev/api/files
|
||||||
|
|
||||||
|
优化:HTTP 代理请求期间不持有数据库连接,避免阻塞其他请求。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict, Optional, Tuple
|
from typing import Any, Dict, Optional, Tuple
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
from fastapi import APIRouter, HTTPException, Request, Response
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
@@ -30,20 +33,30 @@ from src.clients.http_client import HTTPClientPool
|
|||||||
from src.core.api_format import get_auth_handler, get_default_auth_method_for_endpoint
|
from src.core.api_format import get_auth_handler, get_default_auth_method_for_endpoint
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.database import get_db
|
from src.database import create_session
|
||||||
from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User
|
from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User
|
||||||
from src.services.auth.service import AuthService
|
from src.services.auth.service import AuthService
|
||||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||||
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
|
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
|
||||||
from src.services.provider.transport import redact_url_for_log
|
from src.services.provider.transport import redact_url_for_log
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class UpstreamContext:
|
||||||
|
"""上游请求上下文(不依赖数据库会话)"""
|
||||||
|
|
||||||
|
upstream_key: str
|
||||||
|
base_url: str
|
||||||
|
file_key_id: str
|
||||||
|
user_id: str
|
||||||
|
|
||||||
|
|
||||||
router = APIRouter(tags=["Gemini Files API"])
|
router = APIRouter(tags=["Gemini Files API"])
|
||||||
|
|
||||||
# Gemini Files API 基础 URL
|
# Gemini Files API 基础 URL
|
||||||
GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
|
GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
|
||||||
|
|
||||||
# Gemini Files API 能力标签
|
# Gemini Files API 无能力限制(任何 Gemini key 都可用)
|
||||||
REQUIRED_CAPABILITIES = {"gemini_files_api": True}
|
|
||||||
|
|
||||||
# 需要从客户端请求中移除的头部(这些会由代理重新设置或不应转发)
|
# 需要从客户端请求中移除的头部(这些会由代理重新设置或不应转发)
|
||||||
HEADERS_TO_REMOVE = frozenset(
|
HEADERS_TO_REMOVE = frozenset(
|
||||||
@@ -184,9 +197,25 @@ async def _select_provider_candidate(
|
|||||||
db: Session,
|
db: Session,
|
||||||
user_api_key: ApiKey,
|
user_api_key: ApiKey,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
|
require_files_capability: bool = True,
|
||||||
) -> ProviderCandidate | None:
|
) -> ProviderCandidate | None:
|
||||||
"""选择支持 Files API 的 Provider/Endpoint/Key 组合"""
|
"""
|
||||||
|
选择可用的 Provider/Endpoint/Key 组合
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: 数据库会话
|
||||||
|
user_api_key: 用户 API Key
|
||||||
|
model_name: 模型名称
|
||||||
|
require_files_capability: 是否要求 gemini_files 能力(默认 True)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
匹配的候选,如果没有则返回 None
|
||||||
|
"""
|
||||||
scheduler = CacheAwareScheduler()
|
scheduler = CacheAwareScheduler()
|
||||||
|
|
||||||
|
# 要求 gemini_files 能力:只有 Google 官方 API 才支持 Files API
|
||||||
|
capability_requirements = {"gemini_files": True} if require_files_capability else None
|
||||||
|
|
||||||
candidates, _global_model_id = await scheduler.list_all_candidates(
|
candidates, _global_model_id = await scheduler.list_all_candidates(
|
||||||
db=db,
|
db=db,
|
||||||
api_format="gemini:chat",
|
api_format="gemini:chat",
|
||||||
@@ -194,7 +223,7 @@ async def _select_provider_candidate(
|
|||||||
affinity_key=str(user_api_key.id),
|
affinity_key=str(user_api_key.id),
|
||||||
user_api_key=user_api_key,
|
user_api_key=user_api_key,
|
||||||
max_candidates=10,
|
max_candidates=10,
|
||||||
capability_requirements=REQUIRED_CAPABILITIES,
|
capability_requirements=capability_requirements,
|
||||||
)
|
)
|
||||||
for candidate in candidates:
|
for candidate in candidates:
|
||||||
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
|
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
|
||||||
@@ -206,12 +235,20 @@ async def _select_provider_candidate(
|
|||||||
async def _resolve_upstream_context(
|
async def _resolve_upstream_context(
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session,
|
db: Session,
|
||||||
) -> tuple[str, str, str]:
|
) -> tuple[str, str, str, str]:
|
||||||
"""
|
"""
|
||||||
解析上游 Key 与 Base URL
|
解析上游 Key 与 Base URL(需要外部提供 db session)
|
||||||
|
|
||||||
仅允许系统 API Key,通过能力标签选择支持 Files API 的 Provider Key。
|
仅允许系统 API Key,选择可用的 Gemini Provider Key(无能力限制)。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: HTTP 请求
|
||||||
|
db: 数据库会话
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(upstream_key, base_url, key_id, user_id)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
client_key = _extract_gemini_api_key(request)
|
client_key = _extract_gemini_api_key(request)
|
||||||
if not client_key:
|
if not client_key:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -252,14 +289,19 @@ async def _resolve_upstream_context(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
candidate = await _select_provider_candidate(db, user_api_key, model_name)
|
# 选择可用的 provider candidate(要求 gemini_files 能力)
|
||||||
|
candidate = await _select_provider_candidate(
|
||||||
|
db, user_api_key, model_name, require_files_capability=True
|
||||||
|
)
|
||||||
|
|
||||||
if not candidate:
|
if not candidate:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=503,
|
status_code=503,
|
||||||
detail={
|
detail={
|
||||||
"error": {
|
"error": {
|
||||||
"code": 503,
|
"code": 503,
|
||||||
"message": "No available key with gemini_files_api capability",
|
"message": "No available Gemini key with 'gemini_files' capability. "
|
||||||
|
"Please ensure at least one Provider Key has the 'gemini_files' capability enabled.",
|
||||||
"status": "UNAVAILABLE",
|
"status": "UNAVAILABLE",
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
@@ -268,7 +310,7 @@ async def _resolve_upstream_context(
|
|||||||
try:
|
try:
|
||||||
upstream_key = crypto_service.decrypt(candidate.key.api_key)
|
upstream_key = crypto_service.decrypt(candidate.key.api_key)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error(f"Failed to decrypt provider key for Gemini Files API: {exc}")
|
logger.error("Failed to decrypt provider key for Gemini Files API: %s", exc)
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=500,
|
status_code=500,
|
||||||
detail={
|
detail={
|
||||||
@@ -281,7 +323,29 @@ async def _resolve_upstream_context(
|
|||||||
)
|
)
|
||||||
|
|
||||||
base_url = candidate.endpoint.base_url or GEMINI_FILES_BASE_URL
|
base_url = candidate.endpoint.base_url or GEMINI_FILES_BASE_URL
|
||||||
return upstream_key, base_url, str(candidate.key.id)
|
return upstream_key, base_url, str(candidate.key.id), str(user.id)
|
||||||
|
|
||||||
|
|
||||||
|
async def _resolve_upstream_context_standalone(request: Request) -> UpstreamContext:
|
||||||
|
"""
|
||||||
|
解析上游上下文(自管理数据库连接,适用于 HTTP 代理场景)
|
||||||
|
|
||||||
|
优化:在返回上下文后立即释放数据库连接,HTTP 请求期间不持有连接。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: HTTP 请求
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
UpstreamContext: 包含所有必要信息的上下文对象
|
||||||
|
"""
|
||||||
|
with create_session() as db:
|
||||||
|
upstream_key, base_url, file_key_id, user_id = await _resolve_upstream_context(request, db)
|
||||||
|
return UpstreamContext(
|
||||||
|
upstream_key=upstream_key,
|
||||||
|
base_url=base_url,
|
||||||
|
file_key_id=file_key_id,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _proxy_request(
|
async def _proxy_request(
|
||||||
@@ -291,6 +355,7 @@ async def _proxy_request(
|
|||||||
content: bytes | None = None,
|
content: bytes | None = None,
|
||||||
json_body: dict[str, Any] | None = None,
|
json_body: dict[str, Any] | None = None,
|
||||||
file_key_id: str | None = None,
|
file_key_id: str | None = None,
|
||||||
|
user_id: str | None = None,
|
||||||
) -> Response:
|
) -> Response:
|
||||||
"""
|
"""
|
||||||
代理请求到上游 Gemini API
|
代理请求到上游 Gemini API
|
||||||
@@ -302,6 +367,7 @@ async def _proxy_request(
|
|||||||
content: 原始请求体(二进制)
|
content: 原始请求体(二进制)
|
||||||
json_body: JSON 请求体
|
json_body: JSON 请求体
|
||||||
file_key_id: 上游 Provider Key ID,用于成功响应时存储 file→key 映射
|
file_key_id: 上游 Provider Key ID,用于成功响应时存储 file→key 映射
|
||||||
|
user_id: 用户 ID,用于文件映射的权限验证
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
FastAPI Response 对象
|
FastAPI Response 对象
|
||||||
@@ -338,12 +404,28 @@ async def _proxy_request(
|
|||||||
try:
|
try:
|
||||||
payload = response.json()
|
payload = response.json()
|
||||||
file_name = None
|
file_name = None
|
||||||
|
file_obj = None
|
||||||
|
|
||||||
if isinstance(payload, dict):
|
if isinstance(payload, dict):
|
||||||
|
# 单文件上传响应
|
||||||
file_name = payload.get("name")
|
file_name = payload.get("name")
|
||||||
|
file_obj = payload
|
||||||
|
|
||||||
|
# 嵌套格式:{"file": {...}}
|
||||||
if not file_name and isinstance(payload.get("file"), dict):
|
if not file_name and isinstance(payload.get("file"), dict):
|
||||||
file_name = payload["file"].get("name")
|
file_name = payload["file"].get("name")
|
||||||
if file_name:
|
file_obj = payload["file"]
|
||||||
await store_file_key_mapping(file_name, file_key_id)
|
|
||||||
|
if file_name and file_obj:
|
||||||
|
display_name = file_obj.get("displayName") or file_obj.get("display_name")
|
||||||
|
mime_type = file_obj.get("mimeType") or file_obj.get("mime_type")
|
||||||
|
await store_file_key_mapping(
|
||||||
|
file_name,
|
||||||
|
file_key_id,
|
||||||
|
user_id=user_id,
|
||||||
|
display_name=display_name,
|
||||||
|
mime_type=mime_type,
|
||||||
|
)
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Gemini file→key 映射已存储: {file_name} → key_id={file_key_id}"
|
f"Gemini file→key 映射已存储: {file_name} → key_id={file_key_id}"
|
||||||
)
|
)
|
||||||
@@ -355,14 +437,26 @@ async def _proxy_request(
|
|||||||
mapped_count = 0
|
mapped_count = 0
|
||||||
for item in files_list:
|
for item in files_list:
|
||||||
if isinstance(item, dict) and item.get("name"):
|
if isinstance(item, dict) and item.get("name"):
|
||||||
await store_file_key_mapping(item["name"], file_key_id)
|
item_display_name = item.get("displayName") or item.get(
|
||||||
|
"display_name"
|
||||||
|
)
|
||||||
|
item_mime_type = item.get("mimeType") or item.get("mime_type")
|
||||||
|
await store_file_key_mapping(
|
||||||
|
item["name"],
|
||||||
|
file_key_id,
|
||||||
|
user_id=user_id,
|
||||||
|
display_name=item_display_name,
|
||||||
|
mime_type=item_mime_type,
|
||||||
|
)
|
||||||
mapped_count += 1
|
mapped_count += 1
|
||||||
if mapped_count > 0:
|
if mapped_count > 0:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Gemini list_files 批量映射已存储: {mapped_count} 个文件 → key_id={file_key_id}"
|
"Gemini list_files 批量映射已存储: %d 个文件 → key_id=%s",
|
||||||
|
mapped_count,
|
||||||
|
file_key_id,
|
||||||
)
|
)
|
||||||
except (ValueError, KeyError) as e:
|
except (ValueError, KeyError) as e:
|
||||||
logger.debug(f"Failed to store Gemini file mapping: {e}")
|
logger.debug("Failed to store Gemini file mapping: %s", e)
|
||||||
|
|
||||||
return Response(
|
return Response(
|
||||||
content=response.content,
|
content=response.content,
|
||||||
@@ -373,7 +467,7 @@ async def _proxy_request(
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
sanitized_error = redact_url_for_log(str(e))
|
sanitized_error = redact_url_for_log(str(e))
|
||||||
logger.error(f"Gemini Files API proxy error: {sanitized_error}")
|
logger.error("Gemini Files API proxy error: %s", sanitized_error)
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
status_code=502,
|
status_code=502,
|
||||||
content={
|
content={
|
||||||
@@ -394,7 +488,6 @@ async def _proxy_request(
|
|||||||
@router.post("/upload/v1beta/files")
|
@router.post("/upload/v1beta/files")
|
||||||
async def upload_file(
|
async def upload_file(
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""
|
"""
|
||||||
上传文件到 Gemini Files API
|
上传文件到 Gemini Files API
|
||||||
@@ -421,26 +514,34 @@ async def upload_file(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
"""
|
|
||||||
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
|
|
||||||
|
|
||||||
# 读取请求体
|
优化:HTTP 代理期间不持有数据库连接
|
||||||
|
"""
|
||||||
|
# 阶段 1:解析上下文(短暂持有数据库连接)
|
||||||
|
ctx = await _resolve_upstream_context_standalone(request)
|
||||||
|
|
||||||
|
# 阶段 2:读取请求体
|
||||||
body = await request.body()
|
body = await request.body()
|
||||||
|
|
||||||
# 构建上游请求
|
# 阶段 3:代理请求(不持有数据库连接)
|
||||||
upstream_url = _build_upstream_url(
|
upstream_url = _build_upstream_url(
|
||||||
base_url,
|
ctx.base_url,
|
||||||
"/v1beta/files",
|
"/v1beta/files",
|
||||||
dict(request.query_params),
|
dict(request.query_params),
|
||||||
is_upload=True,
|
is_upload=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
|
||||||
|
|
||||||
logger.debug(f"Gemini Files upload proxy: POST {redact_url_for_log(upstream_url)}")
|
logger.debug("Gemini Files upload proxy: POST %s", redact_url_for_log(upstream_url))
|
||||||
|
|
||||||
return await _proxy_request(
|
return await _proxy_request(
|
||||||
"POST", upstream_url, headers, content=body, file_key_id=file_key_id
|
"POST",
|
||||||
|
upstream_url,
|
||||||
|
headers,
|
||||||
|
content=body,
|
||||||
|
file_key_id=ctx.file_key_id,
|
||||||
|
user_id=ctx.user_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -452,13 +553,14 @@ async def upload_file(
|
|||||||
@router.get("/v1beta/files")
|
@router.get("/v1beta/files")
|
||||||
async def list_files(
|
async def list_files(
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
|
||||||
pageSize: int | None = None,
|
pageSize: int | None = None,
|
||||||
pageToken: str | None = None,
|
pageToken: str | None = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""
|
"""
|
||||||
列出已上传的文件
|
列出已上传的文件
|
||||||
|
|
||||||
|
优化:HTTP 代理期间不持有数据库连接
|
||||||
|
|
||||||
**认证方式**:
|
**认证方式**:
|
||||||
- `x-goog-api-key` 请求头,或
|
- `x-goog-api-key` 请求头,或
|
||||||
- `?key=` URL 参数
|
- `?key=` URL 参数
|
||||||
@@ -488,22 +590,222 @@ async def list_files(
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
|
# 阶段 1:解析上下文(短暂持有数据库连接)
|
||||||
|
ctx = await _resolve_upstream_context_standalone(request)
|
||||||
|
|
||||||
# 构建查询参数
|
# 阶段 2:代理请求(不持有数据库连接)
|
||||||
query_params = dict(request.query_params)
|
query_params = dict(request.query_params)
|
||||||
if pageSize is not None:
|
if pageSize is not None:
|
||||||
query_params["pageSize"] = pageSize
|
query_params["pageSize"] = pageSize
|
||||||
if pageToken is not None:
|
if pageToken is not None:
|
||||||
query_params["pageToken"] = pageToken
|
query_params["pageToken"] = pageToken
|
||||||
|
|
||||||
# 构建上游请求
|
upstream_url = _build_upstream_url(ctx.base_url, "/v1beta/files", query_params)
|
||||||
upstream_url = _build_upstream_url(base_url, "/v1beta/files", query_params)
|
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
|
||||||
|
|
||||||
|
logger.debug("Gemini Files list proxy: GET %s", redact_url_for_log(upstream_url))
|
||||||
|
|
||||||
|
return await _proxy_request(
|
||||||
|
"GET", upstream_url, headers, file_key_id=ctx.file_key_id, user_id=ctx.user_id
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# 下载文件内容端点(用于视频等媒体文件)
|
||||||
|
# 注意:必须在 /v1beta/files/{file_name:path} 之前注册,否则会被通配符路由捕获
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
async def _find_video_task_by_id(
|
||||||
|
db: Session, short_id: str, user_id: str
|
||||||
|
) -> tuple[str | None, str | None]:
|
||||||
|
"""
|
||||||
|
通过短 ID 查找视频任务,返回其 provider key 和 video_url
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: 数据库会话
|
||||||
|
short_id: 视频任务的短 ID(VideoTask.short_id,Gemini 风格)
|
||||||
|
user_id: 用户 ID(用于权限验证)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(upstream_key, video_url) - 如果找到任务返回 key 和 url,否则返回 (None, None)
|
||||||
|
"""
|
||||||
|
from src.models.database import ProviderAPIKey, VideoTask
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"[Files Download] Searching video task: short_id=%s, user_id=%s", short_id, user_id
|
||||||
|
)
|
||||||
|
|
||||||
|
# 通过 short_id 查找,同时验证用户权限
|
||||||
|
task = (
|
||||||
|
db.query(VideoTask)
|
||||||
|
.filter(VideoTask.short_id == short_id, VideoTask.user_id == user_id)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
|
||||||
|
if not task:
|
||||||
|
logger.debug("[Files Download] No video task found: short_id=%s", short_id)
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
if not task.video_url:
|
||||||
|
logger.debug("[Files Download] Task found but no video_url: short_id=%s", short_id)
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
if not task.key_id:
|
||||||
|
logger.debug("[Files Download] Task found but no key_id: short_id=%s", short_id)
|
||||||
|
return None, task.video_url
|
||||||
|
|
||||||
|
# 获取 provider key
|
||||||
|
provider_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
|
||||||
|
if not provider_key or not provider_key.api_key:
|
||||||
|
logger.debug("[Files Download] Provider key not found: key_id=%s", task.key_id)
|
||||||
|
return None, task.video_url
|
||||||
|
|
||||||
|
try:
|
||||||
|
upstream_key = crypto_service.decrypt(provider_key.api_key)
|
||||||
|
logger.debug("[Files Download] Found key for task: short_id=%s", short_id)
|
||||||
|
return upstream_key, task.video_url
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Files Download] Failed to decrypt key: %s", e)
|
||||||
|
return None, task.video_url
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/v1beta/files/{file_id}:download")
|
||||||
|
async def download_file(
|
||||||
|
file_id: str,
|
||||||
|
request: Request,
|
||||||
|
) -> Any:
|
||||||
|
"""
|
||||||
|
下载文件(官方 Gemini API 格式)
|
||||||
|
|
||||||
|
**认证方式**:
|
||||||
|
- `x-goog-api-key` 请求头,或
|
||||||
|
- `?key=` URL 参数
|
||||||
|
|
||||||
|
**路径参数**:
|
||||||
|
- `file_id`: 文件 ID
|
||||||
|
- 以 `aev_` 开头:视频任务下载(如 `aev_sknuzqlo8sds`,Gemini 风格短 ID)
|
||||||
|
- 其他:普通 Gemini 文件下载(透传到上游)
|
||||||
|
|
||||||
|
**查询参数**:
|
||||||
|
- `alt=media`: 可选,保持与官方 API 兼容
|
||||||
|
|
||||||
|
**示例**:
|
||||||
|
```
|
||||||
|
GET /v1beta/files/aev_{short_id}:download?alt=media # 视频任务
|
||||||
|
GET /v1beta/files/{gemini_file_id}:download?alt=media # 普通文件
|
||||||
|
```
|
||||||
|
|
||||||
|
优化:HTTP 下载期间不持有数据库连接
|
||||||
|
"""
|
||||||
|
import httpx
|
||||||
|
from fastapi import HTTPException
|
||||||
|
from fastapi.responses import JSONResponse, Response
|
||||||
|
|
||||||
|
# ========== 阶段 1:数据库操作(短暂持有连接)==========
|
||||||
|
client_key = _extract_gemini_api_key(request)
|
||||||
|
if not client_key:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=401,
|
||||||
|
detail={
|
||||||
|
"error": {"code": 401, "message": "API key required", "status": "UNAUTHENTICATED"}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
# 在数据库会话内完成所有查询
|
||||||
|
with create_session() as db:
|
||||||
|
auth_result = AuthService.authenticate_api_key(db, client_key)
|
||||||
|
if not auth_result:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=401,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 401,
|
||||||
|
"message": "API key not valid",
|
||||||
|
"status": "UNAUTHENTICATED",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
user, _user_api_key = auth_result
|
||||||
|
|
||||||
|
# 根据前缀判断处理方式
|
||||||
|
if file_id.startswith("aev_"):
|
||||||
|
# 视频任务下载:使用短 ID 查找
|
||||||
|
short_id = file_id[4:] # 去掉 "aev_" 前缀
|
||||||
|
logger.debug("[Files Download] Video task: short_id=%s, user_id=%s", short_id, user.id)
|
||||||
|
upstream_key, video_url = await _find_video_task_by_id(db, short_id, user.id)
|
||||||
|
if not upstream_key or not video_url:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=404,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 404,
|
||||||
|
"message": f"Video not found or not ready: {file_id}",
|
||||||
|
"status": "NOT_FOUND",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
upstream_url = video_url
|
||||||
|
else:
|
||||||
|
# 普通文件下载:透传到 Gemini
|
||||||
|
try:
|
||||||
|
upstream_key, base_url, _file_key_id, _user_id = await _resolve_upstream_context(
|
||||||
|
request, db
|
||||||
|
)
|
||||||
|
except HTTPException:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=404,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 404,
|
||||||
|
"message": f"File not found: {file_id}",
|
||||||
|
"status": "NOT_FOUND",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
file_name = f"files/{file_id}" if not file_id.startswith("files/") else file_id
|
||||||
|
upstream_url = _build_upstream_url(
|
||||||
|
base_url,
|
||||||
|
f"/v1beta/{file_name}:download",
|
||||||
|
dict(request.query_params),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ========== 阶段 2:HTTP 下载(不持有数据库连接)==========
|
||||||
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
||||||
|
|
||||||
logger.debug(f"Gemini Files list proxy: GET {redact_url_for_log(upstream_url)}")
|
logger.debug("Gemini Files download proxy: GET %s", redact_url_for_log(upstream_url))
|
||||||
|
|
||||||
return await _proxy_request("GET", upstream_url, headers, file_key_id=file_key_id)
|
# 使用 follow_redirects=True 跟随重定向(Gemini 文件下载会重定向)
|
||||||
|
try:
|
||||||
|
async with httpx.AsyncClient(follow_redirects=True, timeout=httpx.Timeout(300.0)) as client:
|
||||||
|
response = await client.get(upstream_url, headers=headers)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Gemini Files download failed: %s", exc)
|
||||||
|
raise HTTPException(status_code=502, detail="Failed to download file")
|
||||||
|
|
||||||
|
if response.status_code >= 400:
|
||||||
|
content: dict[str, Any]
|
||||||
|
if response.headers.get("content-type", "").startswith("application/json"):
|
||||||
|
try:
|
||||||
|
content = response.json()
|
||||||
|
except Exception:
|
||||||
|
content = {"error": response.text}
|
||||||
|
else:
|
||||||
|
content = {"error": response.text}
|
||||||
|
return JSONResponse(content=content, status_code=response.status_code)
|
||||||
|
|
||||||
|
# 返回文件内容
|
||||||
|
return Response(
|
||||||
|
content=response.content,
|
||||||
|
status_code=response.status_code,
|
||||||
|
headers={
|
||||||
|
k: v
|
||||||
|
for k, v in response.headers.items()
|
||||||
|
if k.lower() not in {"transfer-encoding", "connection", "keep-alive"}
|
||||||
|
},
|
||||||
|
media_type=response.headers.get("content-type", "application/octet-stream"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
@@ -515,7 +817,6 @@ async def list_files(
|
|||||||
async def get_file(
|
async def get_file(
|
||||||
file_name: str,
|
file_name: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""
|
"""
|
||||||
获取指定文件的元数据
|
获取指定文件的元数据
|
||||||
@@ -542,24 +843,29 @@ async def get_file(
|
|||||||
"state": "ACTIVE"
|
"state": "ACTIVE"
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
"""
|
|
||||||
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
|
|
||||||
|
|
||||||
|
优化:HTTP 代理期间不持有数据库连接
|
||||||
|
"""
|
||||||
|
# 阶段 1:解析上下文(短暂持有数据库连接)
|
||||||
|
ctx = await _resolve_upstream_context_standalone(request)
|
||||||
|
|
||||||
|
# 阶段 2:代理请求(不持有数据库连接)
|
||||||
# 规范化文件名(确保以 files/ 开头)
|
# 规范化文件名(确保以 files/ 开头)
|
||||||
if not file_name.startswith("files/"):
|
if not file_name.startswith("files/"):
|
||||||
file_name = f"files/{file_name}"
|
file_name = f"files/{file_name}"
|
||||||
|
|
||||||
# 构建上游请求
|
|
||||||
upstream_url = _build_upstream_url(
|
upstream_url = _build_upstream_url(
|
||||||
base_url,
|
ctx.base_url,
|
||||||
f"/v1beta/{file_name}",
|
f"/v1beta/{file_name}",
|
||||||
dict(request.query_params),
|
dict(request.query_params),
|
||||||
)
|
)
|
||||||
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
|
||||||
|
|
||||||
logger.debug(f"Gemini Files get proxy: GET {redact_url_for_log(upstream_url)}")
|
logger.debug("Gemini Files get proxy: GET %s", redact_url_for_log(upstream_url))
|
||||||
|
|
||||||
return await _proxy_request("GET", upstream_url, headers, file_key_id=file_key_id)
|
return await _proxy_request(
|
||||||
|
"GET", upstream_url, headers, file_key_id=ctx.file_key_id, user_id=ctx.user_id
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
@@ -571,7 +877,6 @@ async def get_file(
|
|||||||
async def delete_file(
|
async def delete_file(
|
||||||
file_name: str,
|
file_name: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""
|
"""
|
||||||
删除指定文件
|
删除指定文件
|
||||||
@@ -585,29 +890,31 @@ async def delete_file(
|
|||||||
|
|
||||||
**响应格式**:
|
**响应格式**:
|
||||||
成功时返回空 JSON 对象:`{}`
|
成功时返回空 JSON 对象:`{}`
|
||||||
"""
|
|
||||||
upstream_key, base_url, _file_key_id = await _resolve_upstream_context(request, db)
|
|
||||||
|
|
||||||
|
优化:HTTP 代理期间不持有数据库连接
|
||||||
|
"""
|
||||||
|
# 阶段 1:解析上下文(短暂持有数据库连接)
|
||||||
|
ctx = await _resolve_upstream_context_standalone(request)
|
||||||
|
|
||||||
|
# 阶段 2:代理请求(不持有数据库连接)
|
||||||
# 规范化文件名(确保以 files/ 开头)
|
# 规范化文件名(确保以 files/ 开头)
|
||||||
if not file_name.startswith("files/"):
|
if not file_name.startswith("files/"):
|
||||||
file_name = f"files/{file_name}"
|
file_name = f"files/{file_name}"
|
||||||
|
|
||||||
# 构建上游请求
|
|
||||||
upstream_url = _build_upstream_url(
|
upstream_url = _build_upstream_url(
|
||||||
base_url,
|
ctx.base_url,
|
||||||
f"/v1beta/{file_name}",
|
f"/v1beta/{file_name}",
|
||||||
dict(request.query_params),
|
dict(request.query_params),
|
||||||
)
|
)
|
||||||
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
|
||||||
|
|
||||||
logger.debug(f"Gemini Files delete proxy: DELETE {redact_url_for_log(upstream_url)}")
|
logger.debug("Gemini Files delete proxy: DELETE %s", redact_url_for_log(upstream_url))
|
||||||
|
|
||||||
del _file_key_id # 显式标记:delete 端点不需要存储映射
|
|
||||||
response = await _proxy_request("DELETE", upstream_url, headers)
|
response = await _proxy_request("DELETE", upstream_url, headers)
|
||||||
if response.status_code < 300:
|
if response.status_code < 300:
|
||||||
await delete_file_key_mapping(file_name)
|
await delete_file_key_mapping(file_name)
|
||||||
else:
|
else:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Gemini Files delete failed, skip mapping cleanup: status={response.status_code}"
|
"Gemini Files delete failed, skip mapping cleanup: status=%s", response.status_code
|
||||||
)
|
)
|
||||||
return response
|
return response
|
||||||
|
|||||||
@@ -61,9 +61,10 @@ async def list_video_tasks_sora(http_request: Request, db: Session = Depends(get
|
|||||||
|
|
||||||
|
|
||||||
@router.delete("/v1/videos/{task_id}")
|
@router.delete("/v1/videos/{task_id}")
|
||||||
async def cancel_video_task_sora(
|
async def delete_video_task_sora(
|
||||||
task_id: str, http_request: Request, db: Session = Depends(get_db)
|
task_id: str, http_request: Request, db: Session = Depends(get_db)
|
||||||
) -> Any:
|
) -> Any:
|
||||||
|
"""删除已完成或失败的视频及其存储资源"""
|
||||||
adapter = OpenAIVideoAdapter()
|
adapter = OpenAIVideoAdapter()
|
||||||
return await pipeline.run(
|
return await pipeline.run(
|
||||||
adapter=adapter,
|
adapter=adapter,
|
||||||
@@ -71,7 +72,7 @@ async def cancel_video_task_sora(
|
|||||||
db=db,
|
db=db,
|
||||||
mode=adapter.mode,
|
mode=adapter.mode,
|
||||||
api_format_hint=adapter.allowed_api_formats[0],
|
api_format_hint=adapter.allowed_api_formats[0],
|
||||||
path_params={"task_id": task_id, "action": "cancel"},
|
path_params={"task_id": task_id},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -121,6 +122,47 @@ async def create_video_veo(model: str, http_request: Request, db: Session = Depe
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Gemini Veo operation routes - support both formats:
|
||||||
|
# 1. models/{model}/operations/{id} (official Gemini Veo format)
|
||||||
|
# 2. operations/{...} (legacy format for compatibility)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/v1beta/models/{model}/operations/{operation_id}")
|
||||||
|
async def get_video_veo_by_model(
|
||||||
|
model: str, operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
||||||
|
) -> Any:
|
||||||
|
"""Get video task status (Gemini Veo format: models/{model}/operations/{id})"""
|
||||||
|
adapter = GeminiVeoAdapter()
|
||||||
|
# Reconstruct full operation name
|
||||||
|
full_operation_name = f"models/{model}/operations/{operation_id}"
|
||||||
|
return await pipeline.run(
|
||||||
|
adapter=adapter,
|
||||||
|
http_request=http_request,
|
||||||
|
db=db,
|
||||||
|
mode=adapter.mode,
|
||||||
|
api_format_hint=adapter.allowed_api_formats[0],
|
||||||
|
path_params={"task_id": full_operation_name},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/v1beta/models/{model}/operations/{operation_id}:cancel")
|
||||||
|
async def cancel_video_veo_by_model(
|
||||||
|
model: str, operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
||||||
|
) -> Any:
|
||||||
|
"""Cancel video task (Gemini Veo format: models/{model}/operations/{id}:cancel)"""
|
||||||
|
adapter = GeminiVeoAdapter()
|
||||||
|
full_operation_name = f"models/{model}/operations/{operation_id}"
|
||||||
|
return await pipeline.run(
|
||||||
|
adapter=adapter,
|
||||||
|
http_request=http_request,
|
||||||
|
db=db,
|
||||||
|
mode=adapter.mode,
|
||||||
|
api_format_hint=adapter.allowed_api_formats[0],
|
||||||
|
path_params={"task_id": full_operation_name, "action": "cancel"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy routes for backward compatibility
|
||||||
@router.get("/v1beta/operations/{operation_id:path}")
|
@router.get("/v1beta/operations/{operation_id:path}")
|
||||||
async def get_video_veo(
|
async def get_video_veo(
|
||||||
operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
||||||
@@ -163,16 +205,4 @@ async def cancel_video_veo(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/v1beta/operations/{operation_id:path}/content")
|
# Video download is now handled by /v1beta/files/{task_id}:download in gemini_files.py
|
||||||
async def download_video_content_veo(
|
|
||||||
operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
|
||||||
) -> Any:
|
|
||||||
adapter = GeminiVeoAdapter()
|
|
||||||
return await pipeline.run(
|
|
||||||
adapter=adapter,
|
|
||||||
http_request=http_request,
|
|
||||||
db=db,
|
|
||||||
mode=adapter.mode,
|
|
||||||
api_format_hint=adapter.allowed_api_formats[0],
|
|
||||||
path_params={"task_id": operation_id},
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -736,6 +736,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||||
from sqlalchemy import or_
|
from sqlalchemy import or_
|
||||||
|
from sqlalchemy.orm import load_only
|
||||||
|
|
||||||
from src.models.database import ProviderEndpoint
|
from src.models.database import ProviderEndpoint
|
||||||
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
|
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
|
||||||
@@ -877,7 +878,44 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 计算总数用于分页
|
# 计算总数用于分页
|
||||||
total_records = query.count()
|
# Perf: avoid Query.count() building a subquery selecting many columns
|
||||||
|
total_records = int(query.with_entities(func.count(Usage.id)).scalar() or 0)
|
||||||
|
|
||||||
|
# Perf: do not load large request/response columns for list view
|
||||||
|
query = query.options(
|
||||||
|
load_only(
|
||||||
|
Usage.id,
|
||||||
|
Usage.user_id,
|
||||||
|
Usage.api_key_id,
|
||||||
|
Usage.provider_name,
|
||||||
|
Usage.model,
|
||||||
|
Usage.target_model,
|
||||||
|
Usage.input_tokens,
|
||||||
|
Usage.output_tokens,
|
||||||
|
Usage.total_tokens,
|
||||||
|
Usage.total_cost_usd,
|
||||||
|
Usage.response_time_ms,
|
||||||
|
Usage.first_byte_time_ms,
|
||||||
|
Usage.is_stream,
|
||||||
|
Usage.status,
|
||||||
|
Usage.created_at,
|
||||||
|
Usage.cache_creation_input_tokens,
|
||||||
|
Usage.cache_read_input_tokens,
|
||||||
|
Usage.status_code,
|
||||||
|
Usage.error_message,
|
||||||
|
Usage.api_format,
|
||||||
|
Usage.endpoint_api_format,
|
||||||
|
Usage.has_format_conversion,
|
||||||
|
Usage.input_price_per_1m,
|
||||||
|
Usage.output_price_per_1m,
|
||||||
|
Usage.cache_creation_price_per_1m,
|
||||||
|
Usage.cache_read_price_per_1m,
|
||||||
|
Usage.actual_total_cost_usd,
|
||||||
|
Usage.rate_multiplier,
|
||||||
|
),
|
||||||
|
load_only(ApiKey.id, ApiKey.name, ApiKey.key_encrypted),
|
||||||
|
load_only(ProviderEndpoint.id, ProviderEndpoint.api_format),
|
||||||
|
)
|
||||||
usage_records = (
|
usage_records = (
|
||||||
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
|
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
|
||||||
)
|
)
|
||||||
@@ -1231,7 +1269,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
|||||||
endpoint_format,
|
endpoint_format,
|
||||||
format_acceptance_config,
|
format_acceptance_config,
|
||||||
is_stream=False,
|
is_stream=False,
|
||||||
global_conversion_enabled=global_conversion_enabled,
|
effective_conversion_enabled=global_conversion_enabled,
|
||||||
)
|
)
|
||||||
if is_compatible:
|
if is_compatible:
|
||||||
provider_to_formats.setdefault(provider_id, set()).add(endpoint_format)
|
provider_to_formats.setdefault(provider_id, set()).add(endpoint_format)
|
||||||
|
|||||||
@@ -3,11 +3,20 @@
|
|||||||
|
|
||||||
用于候选筛选时判断端点是否可以处理客户端请求格式。
|
用于候选筛选时判断端点是否可以处理客户端请求格式。
|
||||||
|
|
||||||
|
三层开关优先级(从高到低):
|
||||||
|
1. 全局开关 ON → 强制允许(跳过后续检查)
|
||||||
|
2. 全局开关 OFF → 看提供商开关
|
||||||
|
- 提供商开关 ON → 强制允许(跳过端点检查)
|
||||||
|
- 提供商开关 OFF → 看端点配置
|
||||||
|
3. 端点配置(format_acceptance_config)
|
||||||
|
- enabled=true + 白名单/黑名单检查 → 允许
|
||||||
|
- enabled=false 或未配置 → 禁止
|
||||||
|
|
||||||
转换逻辑:
|
转换逻辑:
|
||||||
1. 格式完全匹配 -> 透传(无需转换)
|
1. 格式完全匹配 -> 透传(无需转换)
|
||||||
2. 格式不同 -> 需要检查全局开关 + 端点开关
|
2. 格式不同 -> 需要检查三层开关
|
||||||
- data_format_id 相同 -> 可透传(无需数据转换),但需全局开关 + 端点开关
|
- data_format_id 相同 -> 可透传(无需数据转换)
|
||||||
- data_format_id 不同 -> 需要转换,检查全局开关 + 端点配置 + 转换器能力
|
- data_format_id 不同 -> 需要转换,检查转换器能力
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -28,8 +37,10 @@ def is_format_compatible(
|
|||||||
endpoint_api_format: str,
|
endpoint_api_format: str,
|
||||||
endpoint_format_acceptance_config: dict | None,
|
endpoint_format_acceptance_config: dict | None,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
global_conversion_enabled: bool,
|
effective_conversion_enabled: bool,
|
||||||
registry: FormatConversionRegistry | None = None,
|
registry: FormatConversionRegistry | None = None,
|
||||||
|
*,
|
||||||
|
skip_endpoint_check: bool = False,
|
||||||
) -> tuple[bool, bool, str | None]:
|
) -> tuple[bool, bool, str | None]:
|
||||||
"""
|
"""
|
||||||
检查端点是否兼容客户端格式
|
检查端点是否兼容客户端格式
|
||||||
@@ -39,8 +50,9 @@ def is_format_compatible(
|
|||||||
endpoint_api_format: 端点的 API 格式
|
endpoint_api_format: 端点的 API 格式
|
||||||
endpoint_format_acceptance_config: 端点的格式接受配置
|
endpoint_format_acceptance_config: 端点的格式接受配置
|
||||||
is_stream: 是否是流式请求
|
is_stream: 是否是流式请求
|
||||||
global_conversion_enabled: 全局格式转换开关(来自环境变量 FORMAT_CONVERSION_ENABLED,默认 True)
|
effective_conversion_enabled: 有效格式转换开关(全局 OR 提供商)
|
||||||
registry: 转换器注册表(可选,默认使用全局单例)
|
registry: 转换器注册表(可选,默认使用全局单例)
|
||||||
|
skip_endpoint_check: 是否跳过端点配置检查(当全局或提供商开关为 ON 时设为 True)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(is_compatible, needs_conversion, skip_reason)
|
(is_compatible, needs_conversion, skip_reason)
|
||||||
@@ -66,30 +78,36 @@ def is_format_compatible(
|
|||||||
if provider_key == client_key:
|
if provider_key == client_key:
|
||||||
return True, False, None
|
return True, False, None
|
||||||
|
|
||||||
# 2. 格式不同 -> 需要检查全局格式转换开关
|
# 2. 格式不同 -> 需要检查格式转换开关
|
||||||
# 即使 data_format_id 相同(如 claude:chat / claude:cli),也需要全局开关启用
|
# 如果有效开关为 False(全局 OFF 且提供商 OFF),直接拒绝
|
||||||
if not global_conversion_enabled:
|
if not effective_conversion_enabled:
|
||||||
return False, False, "全局格式转换未启用(环境变量 FORMAT_CONVERSION_ENABLED=false)"
|
return False, False, "格式转换已禁用(全局和提供商开关均为关闭)"
|
||||||
|
|
||||||
# 3. 格式不同时,统一检查端点配置(核心控制)
|
# 3. 如果全局或提供商开关为 ON,跳过端点配置检查
|
||||||
if endpoint_format_acceptance_config is None:
|
if not skip_endpoint_check:
|
||||||
return False, False, "端点未配置格式接受策略"
|
# 检查端点配置(第三层开关)
|
||||||
|
if endpoint_format_acceptance_config is None:
|
||||||
|
return False, False, "端点未配置格式接受策略"
|
||||||
|
|
||||||
config = endpoint_format_acceptance_config
|
config = endpoint_format_acceptance_config
|
||||||
if not isinstance(config, dict):
|
if not isinstance(config, dict):
|
||||||
return False, False, "端点格式配置无效"
|
return False, False, "端点格式配置无效"
|
||||||
if not config.get("enabled", False):
|
if not config.get("enabled", False):
|
||||||
return False, False, "端点格式接受未启用"
|
return False, False, "端点格式接受未启用"
|
||||||
|
|
||||||
# 检查 reject_formats(优先)
|
# 检查 reject_formats(优先)
|
||||||
reject_formats = config.get("reject_formats", [])
|
reject_formats = config.get("reject_formats", [])
|
||||||
if client_key in [f.upper() for f in reject_formats]:
|
if client_key in [f.upper() for f in reject_formats]:
|
||||||
return False, False, f"端点拒绝 {client_format} 格式"
|
return False, False, f"端点拒绝 {client_format} 格式"
|
||||||
|
|
||||||
# 检查 accept_formats
|
# 检查 accept_formats
|
||||||
accept_formats = config.get("accept_formats", [])
|
accept_formats = config.get("accept_formats", [])
|
||||||
if accept_formats and client_key not in [f.upper() for f in accept_formats]:
|
if accept_formats and client_key not in [f.upper() for f in accept_formats]:
|
||||||
return False, False, f"端点不接受 {client_format} 格式"
|
return False, False, f"端点不接受 {client_format} 格式"
|
||||||
|
|
||||||
|
# 检查流式转换
|
||||||
|
if is_stream and not config.get("stream_conversion", True):
|
||||||
|
return False, False, "端点不支持流式格式转换"
|
||||||
|
|
||||||
# 4. 检查是否可以透传(data_format_id 相同)
|
# 4. 检查是否可以透传(data_format_id 相同)
|
||||||
# 例如:claude:chat / claude:cli 的 data_format_id 都是 "claude",数据格式相同可透传
|
# 例如:claude:chat / claude:cli 的 data_format_id 都是 "claude",数据格式相同可透传
|
||||||
@@ -99,11 +117,7 @@ def is_format_compatible(
|
|||||||
return True, False, None
|
return True, False, None
|
||||||
|
|
||||||
# 5. 需要数据转换的情况(data_format_id 不同)
|
# 5. 需要数据转换的情况(data_format_id 不同)
|
||||||
# 检查流式转换
|
# 检查转换器能力
|
||||||
if is_stream and not config.get("stream_conversion", True):
|
|
||||||
return False, False, "端点不支持流式格式转换"
|
|
||||||
|
|
||||||
# 6. 检查转换器能力
|
|
||||||
if not registry.can_convert_full(
|
if not registry.can_convert_full(
|
||||||
client_key,
|
client_key,
|
||||||
provider_key,
|
provider_key,
|
||||||
|
|||||||
@@ -736,10 +736,17 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
instance = instances[0] if isinstance(instances[0], dict) else {}
|
instance = instances[0] if isinstance(instances[0], dict) else {}
|
||||||
params = request.get("parameters") or {}
|
params = request.get("parameters") or {}
|
||||||
|
|
||||||
|
# 解析 image(用于 image-to-video 或第一帧)
|
||||||
|
# 官方格式: {"image": {"inlineData": {"mimeType": "image/png", "data": "base64..."}}}
|
||||||
image = instance.get("image") if isinstance(instance, dict) else None
|
image = instance.get("image") if isinstance(instance, dict) else None
|
||||||
image_ref = None
|
image_ref = None
|
||||||
if isinstance(image, dict):
|
if isinstance(image, dict):
|
||||||
image_ref = image.get("bytesBase64Encoded")
|
inline_data = image.get("inlineData", {})
|
||||||
|
if isinstance(inline_data, dict):
|
||||||
|
image_ref = inline_data.get("data")
|
||||||
|
# 兼容旧格式
|
||||||
|
if not image_ref:
|
||||||
|
image_ref = image.get("bytesBase64Encoded")
|
||||||
|
|
||||||
prompt = instance.get("prompt") if isinstance(instance, dict) else None
|
prompt = instance.get("prompt") if isinstance(instance, dict) else None
|
||||||
prompt_str = str(prompt).strip() if prompt else ""
|
prompt_str = str(prompt).strip() if prompt else ""
|
||||||
@@ -747,7 +754,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
raise ValueError("Video prompt is required")
|
raise ValueError("Video prompt is required")
|
||||||
|
|
||||||
duration_raw = params.get("durationSeconds")
|
duration_raw = params.get("durationSeconds")
|
||||||
sample_count_raw = params.get("sampleCount")
|
sample_count_raw = params.get("sampleCount") or params.get("numberOfVideos")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
duration_seconds = int(duration_raw) if duration_raw else 8
|
duration_seconds = int(duration_raw) if duration_raw else 8
|
||||||
@@ -760,6 +767,35 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
sample_count = 1
|
sample_count = 1
|
||||||
|
|
||||||
|
# 构建 extra 字段,保留所有 Veo 特有的参数
|
||||||
|
extra: dict[str, Any] = {
|
||||||
|
"personGeneration": params.get("personGeneration"),
|
||||||
|
"sampleCount": sample_count,
|
||||||
|
}
|
||||||
|
|
||||||
|
# negativePrompt - 负面提示词
|
||||||
|
if params.get("negativePrompt"):
|
||||||
|
extra["negativePrompt"] = params["negativePrompt"]
|
||||||
|
|
||||||
|
# lastFrame - 最后一帧(用于插值)
|
||||||
|
last_frame = params.get("lastFrame")
|
||||||
|
if isinstance(last_frame, dict):
|
||||||
|
extra["lastFrame"] = last_frame
|
||||||
|
|
||||||
|
# referenceImages - 参考图像(最多3张,仅 Veo 3.1)
|
||||||
|
ref_images = params.get("referenceImages")
|
||||||
|
if isinstance(ref_images, list) and ref_images:
|
||||||
|
extra["referenceImages"] = ref_images
|
||||||
|
|
||||||
|
# video - 视频扩展输入(用于视频续写)
|
||||||
|
video_input = instance.get("video") if isinstance(instance, dict) else None
|
||||||
|
if isinstance(video_input, dict):
|
||||||
|
extra["video"] = video_input
|
||||||
|
|
||||||
|
# seed - 种子值(Veo 3)
|
||||||
|
if params.get("seed") is not None:
|
||||||
|
extra["seed"] = params["seed"]
|
||||||
|
|
||||||
return InternalVideoRequest(
|
return InternalVideoRequest(
|
||||||
prompt=prompt_str,
|
prompt=prompt_str,
|
||||||
model=str(request.get("model") or "veo-3.1-generate-preview"),
|
model=str(request.get("model") or "veo-3.1-generate-preview"),
|
||||||
@@ -767,10 +803,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
aspect_ratio=str(params.get("aspectRatio") or "16:9"),
|
aspect_ratio=str(params.get("aspectRatio") or "16:9"),
|
||||||
resolution=str(params.get("resolution") or "720p"),
|
resolution=str(params.get("resolution") or "720p"),
|
||||||
reference_image_url=image_ref,
|
reference_image_url=image_ref,
|
||||||
extra={
|
extra=extra,
|
||||||
"personGeneration": params.get("personGeneration"),
|
|
||||||
"sampleCount": sample_count,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def video_request_from_internal(self, internal: InternalVideoRequest) -> dict[str, Any]:
|
def video_request_from_internal(self, internal: InternalVideoRequest) -> dict[str, Any]:
|
||||||
@@ -783,15 +816,31 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
"resolution": internal.resolution,
|
"resolution": internal.resolution,
|
||||||
"durationSeconds": internal.duration_seconds,
|
"durationSeconds": internal.duration_seconds,
|
||||||
}
|
}
|
||||||
for key in ["personGeneration", "sampleCount"]:
|
for key in ["personGeneration", "sampleCount", "negativePrompt", "seed"]:
|
||||||
if key in internal.extra:
|
if key in internal.extra:
|
||||||
parameters[key] = internal.extra[key]
|
parameters[key] = internal.extra[key]
|
||||||
|
|
||||||
return {
|
# lastFrame 和 referenceImages 需要特殊处理
|
||||||
|
if internal.extra.get("lastFrame"):
|
||||||
|
parameters["lastFrame"] = internal.extra["lastFrame"]
|
||||||
|
if internal.extra.get("referenceImages"):
|
||||||
|
parameters["referenceImages"] = internal.extra["referenceImages"]
|
||||||
|
|
||||||
|
# video 输入(视频续写)
|
||||||
|
if internal.extra.get("video"):
|
||||||
|
instance["video"] = internal.extra["video"]
|
||||||
|
|
||||||
|
result: dict[str, Any] = {
|
||||||
"instances": [instance],
|
"instances": [instance],
|
||||||
"parameters": parameters,
|
"parameters": parameters,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 模型信息(用于 URL 构建)
|
||||||
|
if internal.model:
|
||||||
|
result["model"] = internal.model
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
def video_task_to_internal(self, response: dict[str, Any]) -> InternalVideoTask:
|
def video_task_to_internal(self, response: dict[str, Any]) -> InternalVideoTask:
|
||||||
operation_name = str(response.get("name") or "")
|
operation_name = str(response.get("name") or "")
|
||||||
done = bool(response.get("done"))
|
done = bool(response.get("done"))
|
||||||
@@ -824,20 +873,30 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def video_task_from_internal(self, internal: InternalVideoTask) -> dict[str, Any]:
|
def video_task_from_internal(self, internal: InternalVideoTask) -> dict[str, Any]:
|
||||||
# 优先使用 external_id(上游返回的 operation name),否则用内部 id
|
# 从 external_id 中提取 model 名称,用于构建 operation name
|
||||||
operation_name = internal.external_id or f"operations/{internal.id}"
|
# external_id 格式: models/{model}/operations/{gemini_id}
|
||||||
if not operation_name.startswith("operations/"):
|
model_name = "unknown"
|
||||||
operation_name = f"operations/{operation_name}"
|
if internal.external_id:
|
||||||
|
parts = internal.external_id.split("/")
|
||||||
|
if len(parts) >= 2 and parts[0] == "models":
|
||||||
|
model_name = parts[1]
|
||||||
|
|
||||||
|
# 使用我们的内部 task_id 构建 operation name,不暴露 Gemini 的 operation ID
|
||||||
|
# 格式: models/{model}/operations/{our_task_id}
|
||||||
|
operation_name = f"models/{model_name}/operations/{internal.id}"
|
||||||
|
|
||||||
if internal.status == VideoStatus.COMPLETED:
|
if internal.status == VideoStatus.COMPLETED:
|
||||||
urls = internal.video_urls or ([internal.video_url] if internal.video_url else [])
|
# 使用我们的内部 task_id 构建下载 URL,不暴露真实的 Gemini file_id
|
||||||
|
# 使用 aev_ 前缀标识这是视频任务的下载链接
|
||||||
|
# 格式:/v1beta/files/aev_{task_id}:download?alt=media
|
||||||
|
proxy_download_url = f"/v1beta/files/aev_{internal.id}:download?alt=media"
|
||||||
return {
|
return {
|
||||||
"name": operation_name,
|
"name": operation_name,
|
||||||
"done": True,
|
"done": True,
|
||||||
"response": {
|
"response": {
|
||||||
"generateVideoResponse": {
|
"generateVideoResponse": {
|
||||||
"generatedSamples": [
|
"generatedSamples": [
|
||||||
{"video": {"uri": url, "mimeType": "video/mp4"}} for url in urls
|
{"video": {"uri": proxy_download_url, "mimeType": "video/mp4"}}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -752,10 +752,21 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
"message": internal.error_message,
|
"message": internal.error_message,
|
||||||
}
|
}
|
||||||
|
|
||||||
for key in ["model", "size", "seconds"]:
|
# 基本字段
|
||||||
if key in internal.extra:
|
for key in ["model", "size", "prompt"]:
|
||||||
|
if internal.extra.get(key):
|
||||||
payload[key] = internal.extra[key]
|
payload[key] = internal.extra[key]
|
||||||
|
|
||||||
|
# seconds 必须是字符串类型
|
||||||
|
seconds = internal.extra.get("seconds")
|
||||||
|
if seconds is not None:
|
||||||
|
payload["seconds"] = str(seconds)
|
||||||
|
|
||||||
|
# remix 相关字段
|
||||||
|
remixed_from = internal.extra.get("remixed_from_video_id")
|
||||||
|
if remixed_from:
|
||||||
|
payload["remixed_from_video_id"] = remixed_from
|
||||||
|
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
def video_poll_to_internal(self, response: dict[str, Any]) -> InternalVideoPollResult:
|
def video_poll_to_internal(self, response: dict[str, Any]) -> InternalVideoPollResult:
|
||||||
@@ -764,14 +775,18 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if status == "completed":
|
if status == "completed":
|
||||||
expires_at = response.get("expires_at")
|
expires_at = response.get("expires_at")
|
||||||
# 使用任务 ID 构建内容路径,由调用方拼接完整 URL
|
# 优先使用上游返回的直接 URL(某些代理如 API易 会返回 CDN URL)
|
||||||
# 如果 task_id 不存在,说明上游响应异常
|
# 回退到构建相对路径(标准 OpenAI API 通过 /content 端点下载)
|
||||||
video_url = f"videos/{task_id}/content" if task_id else None
|
direct_url = (
|
||||||
|
response.get("video_url") or response.get("url") or response.get("result_url")
|
||||||
|
)
|
||||||
|
# 使用直接 URL 或回退到相对路径
|
||||||
|
video_url = direct_url or (f"videos/{task_id}/content" if task_id else None)
|
||||||
if not video_url:
|
if not video_url:
|
||||||
return InternalVideoPollResult(
|
return InternalVideoPollResult(
|
||||||
status=VideoStatus.FAILED,
|
status=VideoStatus.FAILED,
|
||||||
error_code="missing_task_id",
|
error_code="missing_video_url",
|
||||||
error_message="Upstream response missing task id",
|
error_message="Upstream response missing video url",
|
||||||
raw_response=response,
|
raw_response=response,
|
||||||
)
|
)
|
||||||
return InternalVideoPollResult(
|
return InternalVideoPollResult(
|
||||||
|
|||||||
@@ -155,6 +155,110 @@ class FormatConversionRegistry:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||||
|
|
||||||
|
# ==================== 视频格式转换 ====================
|
||||||
|
|
||||||
|
def convert_video_request(
|
||||||
|
self,
|
||||||
|
request: dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""转换视频请求格式(OpenAI <-> Gemini)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: 原始视频请求
|
||||||
|
source_format: 源格式(如 openai:video, gemini:video)
|
||||||
|
target_format: 目标格式
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换后的视频请求
|
||||||
|
"""
|
||||||
|
# 统一使用基础格式 ID(去掉 :video 后缀)
|
||||||
|
src_base = self._video_format_to_base(source_format)
|
||||||
|
tgt_base = self._video_format_to_base(target_format)
|
||||||
|
|
||||||
|
if src_base == tgt_base:
|
||||||
|
return request
|
||||||
|
|
||||||
|
src = self._require_normalizer(src_base)
|
||||||
|
tgt = self._require_normalizer(tgt_base)
|
||||||
|
|
||||||
|
with _track_conversion_metrics(
|
||||||
|
"video_request", str(source_format).upper(), str(target_format).upper()
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
internal = src.video_request_to_internal(request)
|
||||||
|
return tgt.video_request_from_internal(internal)
|
||||||
|
except Exception as e:
|
||||||
|
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||||
|
|
||||||
|
def convert_video_task(
|
||||||
|
self,
|
||||||
|
task_response: dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""转换视频任务响应格式(OpenAI <-> Gemini)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
task_response: 原始任务响应
|
||||||
|
source_format: 源格式
|
||||||
|
target_format: 目标格式
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换后的任务响应
|
||||||
|
"""
|
||||||
|
src_base = self._video_format_to_base(source_format)
|
||||||
|
tgt_base = self._video_format_to_base(target_format)
|
||||||
|
|
||||||
|
if src_base == tgt_base:
|
||||||
|
return task_response
|
||||||
|
|
||||||
|
src = self._require_normalizer(src_base)
|
||||||
|
tgt = self._require_normalizer(tgt_base)
|
||||||
|
|
||||||
|
with _track_conversion_metrics(
|
||||||
|
"video_task", str(source_format).upper(), str(target_format).upper()
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
internal = src.video_task_to_internal(task_response)
|
||||||
|
return tgt.video_task_from_internal(internal)
|
||||||
|
except Exception as e:
|
||||||
|
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||||
|
|
||||||
|
def can_convert_video(self, source_format: str, target_format: str) -> bool:
|
||||||
|
"""检查是否支持视频格式转换"""
|
||||||
|
src_base = self._video_format_to_base(source_format)
|
||||||
|
tgt_base = self._video_format_to_base(target_format)
|
||||||
|
|
||||||
|
if src_base == tgt_base:
|
||||||
|
return True
|
||||||
|
|
||||||
|
src = self.get_normalizer(src_base)
|
||||||
|
tgt = self.get_normalizer(tgt_base)
|
||||||
|
|
||||||
|
if src is None or tgt is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 检查是否有视频转换方法
|
||||||
|
return (
|
||||||
|
hasattr(src, "video_request_to_internal")
|
||||||
|
and hasattr(src, "video_task_to_internal")
|
||||||
|
and hasattr(tgt, "video_request_from_internal")
|
||||||
|
and hasattr(tgt, "video_task_from_internal")
|
||||||
|
)
|
||||||
|
|
||||||
|
def _video_format_to_base(self, format_id: str) -> str:
|
||||||
|
"""将视频格式 ID 转换为基础格式 ID
|
||||||
|
|
||||||
|
例如: openai:video -> openai:chat, gemini:video -> gemini:chat
|
||||||
|
"""
|
||||||
|
upper = str(format_id).upper()
|
||||||
|
if upper.endswith(":VIDEO"):
|
||||||
|
base = upper[:-6] # 去掉 :VIDEO
|
||||||
|
return f"{base}:CHAT"
|
||||||
|
return upper
|
||||||
|
|
||||||
# ==================== 流式转换(严格) ====================
|
# ==================== 流式转换(严格) ====================
|
||||||
|
|
||||||
def convert_stream_chunk(
|
def convert_stream_chunk(
|
||||||
|
|||||||
@@ -248,10 +248,10 @@ register_capability(
|
|||||||
)
|
)
|
||||||
|
|
||||||
register_capability(
|
register_capability(
|
||||||
name="gemini_files_api",
|
name="gemini_files",
|
||||||
display_name="Gemini文件上传",
|
display_name="Gemini 文件 API",
|
||||||
description="支持 Gemini Files API(上传、查询、删除),第三方 Key 通常不支持",
|
description="支持 Gemini Files API(文件上传/管理),仅 Google 官方 API 支持",
|
||||||
match_mode=CapabilityMatchMode.COMPATIBLE, # 需要时选有的,不需要时都可选
|
match_mode=CapabilityMatchMode.COMPATIBLE,
|
||||||
config_mode=CapabilityConfigMode.REQUEST_PARAM, # 从请求路径检测
|
config_mode=CapabilityConfigMode.USER_CONFIGURABLE,
|
||||||
short_name="文件上传",
|
short_name="文件API",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -236,15 +236,23 @@ class ModuleRegistry:
|
|||||||
# 获取启用状态
|
# 获取启用状态
|
||||||
enabled = self.is_enabled(name, db) if available else False
|
enabled = self.is_enabled(name, db) if available else False
|
||||||
|
|
||||||
# 注意:配置验证失败时不自动禁用模块
|
# 配置验证失败时自动禁用模块
|
||||||
# 自动禁用会在查询方法中产生写操作副作用,违反幂等性原则
|
# 注意:此处故意在 get_status() 中写入,以确保模块状态与配置同步
|
||||||
# 配置验证状态通过 config_validated/config_error 字段返回,由调用方决定如何处理
|
# 场景:用户删除了模块所依赖的 Provider Key 后,模块应自动关闭
|
||||||
|
# 权衡:查询方法中的写操作副作用 vs 状态一致性保证
|
||||||
|
if enabled and not config_validated:
|
||||||
|
self.set_enabled(name, False, db)
|
||||||
|
enabled = False
|
||||||
|
|
||||||
|
# 计算激活状态:available && enabled && config_validated && 依赖模块都激活
|
||||||
|
is_active = self.is_active(name, db) if available else False
|
||||||
|
active = is_active and config_validated
|
||||||
|
|
||||||
return ModuleStatus(
|
return ModuleStatus(
|
||||||
name=name,
|
name=name,
|
||||||
available=available,
|
available=available,
|
||||||
enabled=enabled,
|
enabled=enabled,
|
||||||
active=self.is_active(name, db) if available else False,
|
active=active,
|
||||||
config_validated=config_validated,
|
config_validated=config_validated,
|
||||||
config_error=config_error,
|
config_error=config_error,
|
||||||
display_name=meta.display_name,
|
display_name=meta.display_name,
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from ..models.database import ApiKey, Base, Usage, User, UserQuota
|
from ..models.database import ApiKey, Base, Usage, User, UserQuota
|
||||||
from .database import create_session, get_db, get_db_url, init_db, log_pool_status
|
from .database import create_session, get_db, get_db_context, get_db_url, init_db, log_pool_status
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Base",
|
"Base",
|
||||||
@@ -12,6 +12,7 @@ __all__ = [
|
|||||||
"Usage",
|
"Usage",
|
||||||
"UserQuota",
|
"UserQuota",
|
||||||
"get_db",
|
"get_db",
|
||||||
|
"get_db_context",
|
||||||
"init_db",
|
"init_db",
|
||||||
"create_session",
|
"create_session",
|
||||||
"get_db_url",
|
"get_db_url",
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
|
|
||||||
import time
|
import time
|
||||||
from collections.abc import Generator
|
from collections.abc import Generator
|
||||||
|
from contextlib import contextmanager
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
from sqlalchemy import create_engine, event
|
from sqlalchemy import create_engine, event
|
||||||
@@ -271,6 +272,31 @@ def create_session() -> Session:
|
|||||||
return _SessionLocal()
|
return _SessionLocal()
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def get_db_context() -> Generator[Session, None, None]:
|
||||||
|
"""
|
||||||
|
获取数据库会话的上下文管理器
|
||||||
|
|
||||||
|
自动管理会话生命周期:创建、提交/回滚、关闭
|
||||||
|
|
||||||
|
示例:
|
||||||
|
with get_db_context() as db:
|
||||||
|
user = db.query(User).first()
|
||||||
|
# 事务在 with 块结束时自动提交或回滚
|
||||||
|
"""
|
||||||
|
_ensure_engine()
|
||||||
|
assert _SessionLocal is not None
|
||||||
|
db = _SessionLocal()
|
||||||
|
try:
|
||||||
|
yield db
|
||||||
|
db.commit()
|
||||||
|
except Exception:
|
||||||
|
db.rollback()
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
def get_db_url() -> str:
|
def get_db_url() -> str:
|
||||||
"""返回当前配置的数据库连接字符串(供脚本/测试使用)。"""
|
"""返回当前配置的数据库连接字符串(供脚本/测试使用)。"""
|
||||||
return config.database_url
|
return config.database_url
|
||||||
|
|||||||
26
src/main.py
26
src/main.py
@@ -202,14 +202,14 @@ async def lifespan(app: FastAPI) -> Any:
|
|||||||
logger.info("启动月卡额度重置调度器...")
|
logger.info("启动月卡额度重置调度器...")
|
||||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||||
|
from src.services.task.task_poller import get_task_poller
|
||||||
from src.services.usage.quota_scheduler import get_quota_scheduler
|
from src.services.usage.quota_scheduler import get_quota_scheduler
|
||||||
from src.services.video.task_poller import get_video_task_poller
|
|
||||||
from src.utils.task_coordinator import StartupTaskCoordinator
|
from src.utils.task_coordinator import StartupTaskCoordinator
|
||||||
|
|
||||||
quota_scheduler = get_quota_scheduler()
|
quota_scheduler = get_quota_scheduler()
|
||||||
maintenance_scheduler = get_maintenance_scheduler()
|
maintenance_scheduler = get_maintenance_scheduler()
|
||||||
model_fetch_scheduler = get_model_fetch_scheduler()
|
model_fetch_scheduler = get_model_fetch_scheduler()
|
||||||
video_task_poller = get_video_task_poller()
|
task_poller = get_task_poller()
|
||||||
task_coordinator = StartupTaskCoordinator(redis_client)
|
task_coordinator = StartupTaskCoordinator(redis_client)
|
||||||
|
|
||||||
# 启动额度调度器
|
# 启动额度调度器
|
||||||
@@ -238,14 +238,14 @@ async def lifespan(app: FastAPI) -> Any:
|
|||||||
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
|
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
|
||||||
model_fetch_scheduler = None # type: ignore[assignment]
|
model_fetch_scheduler = None # type: ignore[assignment]
|
||||||
|
|
||||||
# 启动视频任务轮询服务
|
# 启动异步任务轮询服务(当前仅视频)
|
||||||
video_poller_active = await task_coordinator.acquire("video_task_poller")
|
task_poller_active = await task_coordinator.acquire("task_poller:video")
|
||||||
if video_poller_active:
|
if task_poller_active:
|
||||||
logger.info("启动视频任务轮询服务...")
|
logger.info("启动 TaskPoller(video)...")
|
||||||
await video_task_poller.start()
|
await task_poller.start()
|
||||||
else:
|
else:
|
||||||
logger.info("检测到其他 worker 已运行视频任务轮询,本实例跳过")
|
logger.info("检测到其他 worker 已运行 TaskPoller(video),本实例跳过")
|
||||||
video_task_poller = None # type: ignore[assignment]
|
task_poller = None # type: ignore[assignment]
|
||||||
|
|
||||||
# 启动统一的定时任务调度器
|
# 启动统一的定时任务调度器
|
||||||
from src.services.system.scheduler import get_scheduler
|
from src.services.system.scheduler import get_scheduler
|
||||||
@@ -296,10 +296,10 @@ async def lifespan(app: FastAPI) -> Any:
|
|||||||
await model_fetch_scheduler.stop()
|
await model_fetch_scheduler.stop()
|
||||||
await task_coordinator.release("model_fetch_scheduler")
|
await task_coordinator.release("model_fetch_scheduler")
|
||||||
|
|
||||||
if video_task_poller:
|
if task_poller:
|
||||||
logger.info("停止视频任务轮询...")
|
logger.info("停止 TaskPoller(video)...")
|
||||||
await video_task_poller.stop()
|
await task_poller.stop()
|
||||||
await task_coordinator.release("video_task_poller")
|
await task_coordinator.release("task_poller:video")
|
||||||
|
|
||||||
# 停止统一的定时任务调度器
|
# 停止统一的定时任务调度器
|
||||||
logger.info("停止定时任务调度器...")
|
logger.info("停止定时任务调度器...")
|
||||||
|
|||||||
@@ -272,6 +272,14 @@ class Usage(Base):
|
|||||||
"""使用记录模型"""
|
"""使用记录模型"""
|
||||||
|
|
||||||
__tablename__ = "usage"
|
__tablename__ = "usage"
|
||||||
|
__table_args__ = (
|
||||||
|
# Composite indexes for common query patterns (analytics / list pages)
|
||||||
|
Index("idx_usage_user_created", "user_id", "created_at"),
|
||||||
|
Index("idx_usage_apikey_created", "api_key_id", "created_at"),
|
||||||
|
Index("idx_usage_provider_model_created", "provider_name", "model", "created_at"),
|
||||||
|
Index("idx_usage_provider_created", "provider_name", "created_at"),
|
||||||
|
Index("idx_usage_model_created", "model", "created_at"),
|
||||||
|
)
|
||||||
|
|
||||||
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
|
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
|
||||||
user_id = Column(String(36), ForeignKey("users.id", ondelete="SET NULL"), nullable=True)
|
user_id = Column(String(36), ForeignKey("users.id", ondelete="SET NULL"), nullable=True)
|
||||||
@@ -347,6 +355,13 @@ class Usage(Base):
|
|||||||
# cancelled: 客户端主动断开连接
|
# cancelled: 客户端主动断开连接
|
||||||
status = Column(String(20), default="completed", nullable=False, index=True)
|
status = Column(String(20), default="completed", nullable=False, index=True)
|
||||||
|
|
||||||
|
# 结算状态(与 status 解耦)
|
||||||
|
# - pending: 等待结算(任务未完成 / 流式未结束)
|
||||||
|
# - settled: 已结算(cost 已写入,可能 > 0 或 = 0)
|
||||||
|
# - void: 作废(不收费,如任务未开始就取消)
|
||||||
|
billing_status = Column(String(20), default="settled", nullable=False, index=True)
|
||||||
|
finalized_at = Column(DateTime(timezone=True), nullable=True) # 结算完成时间(可选)
|
||||||
|
|
||||||
# 完整请求和响应记录
|
# 完整请求和响应记录
|
||||||
request_headers = Column(JSON, nullable=True) # 客户端请求头
|
request_headers = Column(JSON, nullable=True) # 客户端请求头
|
||||||
request_body = Column(JSON, nullable=True) # 请求体(7天内未压缩)
|
request_body = Column(JSON, nullable=True) # 请求体(7天内未压缩)
|
||||||
@@ -654,6 +669,14 @@ class Provider(Base):
|
|||||||
# 注意:如果全局配置 KEEP_PRIORITY_ON_CONVERSION=true,此字段被忽略(所有提供商都保持优先级)
|
# 注意:如果全局配置 KEEP_PRIORITY_ON_CONVERSION=true,此字段被忽略(所有提供商都保持优先级)
|
||||||
keep_priority_on_conversion = Column(Boolean, default=False, nullable=False)
|
keep_priority_on_conversion = Column(Boolean, default=False, nullable=False)
|
||||||
|
|
||||||
|
# 是否允许格式转换(默认 True)
|
||||||
|
# - True: 该提供商可以作为格式转换的目标(如 OpenAI 客户端请求可以路由到此 Gemini 提供商)
|
||||||
|
# - False: 该提供商不接受需要格式转换的请求
|
||||||
|
# 优先级逻辑:
|
||||||
|
# - 全局开关 ON → 强制允许所有提供商的格式转换(忽略此字段)
|
||||||
|
# - 全局开关 OFF → 由此字段决定是否允许该提供商的格式转换
|
||||||
|
enable_format_conversion = Column(Boolean, default=False, nullable=False)
|
||||||
|
|
||||||
# 状态
|
# 状态
|
||||||
is_active = Column(Boolean, default=True, nullable=False)
|
is_active = Column(Boolean, default=True, nullable=False)
|
||||||
|
|
||||||
@@ -1350,12 +1373,26 @@ class ProviderAPIKey(Base):
|
|||||||
provider = relationship("Provider", back_populates="api_keys")
|
provider = relationship("Provider", back_populates="api_keys")
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_short_id(length: int = 12) -> str:
|
||||||
|
"""生成 Gemini 风格的短 ID(小写字母+数字)"""
|
||||||
|
import secrets
|
||||||
|
import string
|
||||||
|
|
||||||
|
alphabet = string.ascii_lowercase + string.digits
|
||||||
|
return "".join(secrets.choice(alphabet) for _ in range(length))
|
||||||
|
|
||||||
|
|
||||||
class VideoTask(Base):
|
class VideoTask(Base):
|
||||||
"""视频生成任务"""
|
"""视频生成任务"""
|
||||||
|
|
||||||
__tablename__ = "video_tasks"
|
__tablename__ = "video_tasks"
|
||||||
|
|
||||||
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
|
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
|
||||||
|
# Gemini 风格的短 ID,用于对外暴露(如 operations/xxx)
|
||||||
|
short_id = Column(String(16), unique=True, index=True, default=_generate_short_id)
|
||||||
|
request_id = Column(
|
||||||
|
String(100), unique=True, index=True, nullable=False
|
||||||
|
) # 关联 Usage/RequestCandidate
|
||||||
external_task_id = Column(String(200))
|
external_task_id = Column(String(200))
|
||||||
|
|
||||||
# 关联
|
# 关联
|
||||||
@@ -1865,6 +1902,7 @@ class RequestCandidate(Base):
|
|||||||
Index("idx_request_candidates_request_id", "request_id"),
|
Index("idx_request_candidates_request_id", "request_id"),
|
||||||
Index("idx_request_candidates_status", "status"),
|
Index("idx_request_candidates_status", "status"),
|
||||||
Index("idx_request_candidates_provider_id", "provider_id"),
|
Index("idx_request_candidates_provider_id", "provider_id"),
|
||||||
|
Index("idx_request_candidates_created_at", "created_at"),
|
||||||
)
|
)
|
||||||
|
|
||||||
# 关系
|
# 关系
|
||||||
@@ -2109,5 +2147,60 @@ class StatsUserDaily(Base):
|
|||||||
user = relationship("User")
|
user = relationship("User")
|
||||||
|
|
||||||
|
|
||||||
|
class GeminiFileMapping(Base):
|
||||||
|
"""
|
||||||
|
Gemini Files API 文件与 Provider Key 的映射关系
|
||||||
|
|
||||||
|
用于持久化存储 file_id → key_id 的绑定关系,
|
||||||
|
确保后续 generateContent 请求使用上传时的同一 Key。
|
||||||
|
|
||||||
|
Gemini 文件有 48 小时有效期,此表中的记录也会在过期后被清理。
|
||||||
|
"""
|
||||||
|
|
||||||
|
__tablename__ = "gemini_file_mappings"
|
||||||
|
|
||||||
|
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
|
||||||
|
|
||||||
|
# 文件名(如 files/abc123xyz)
|
||||||
|
file_name = Column(String(255), nullable=False, unique=True, index=True)
|
||||||
|
|
||||||
|
# Provider Key ID(关联到 provider_api_keys 表)
|
||||||
|
key_id = Column(
|
||||||
|
String(36),
|
||||||
|
ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 用户 ID(用于权限验证,可选)
|
||||||
|
user_id = Column(
|
||||||
|
String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True
|
||||||
|
)
|
||||||
|
|
||||||
|
# 文件元数据(可选,用于调试)
|
||||||
|
display_name = Column(String(255), nullable=True)
|
||||||
|
mime_type = Column(String(100), nullable=True)
|
||||||
|
|
||||||
|
# 源文件哈希(用于关联相同源文件的不同上传,可选)
|
||||||
|
# 当同一源文件上传到多个 Key 时,可通过此字段找到所有等效文件
|
||||||
|
source_hash = Column(String(64), nullable=True, index=True)
|
||||||
|
|
||||||
|
# 时间戳
|
||||||
|
created_at = Column(
|
||||||
|
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
|
||||||
|
)
|
||||||
|
# 过期时间(Gemini 文件 48 小时后过期)
|
||||||
|
expires_at = Column(DateTime(timezone=True), nullable=False, index=True)
|
||||||
|
|
||||||
|
# 关系
|
||||||
|
key = relationship("ProviderAPIKey")
|
||||||
|
user = relationship("User")
|
||||||
|
|
||||||
|
__table_args__ = (
|
||||||
|
Index("idx_gemini_file_mappings_expires", "expires_at"),
|
||||||
|
Index("idx_gemini_file_mappings_source_hash", "source_hash"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# 导入扩展的数据库模型
|
# 导入扩展的数据库模型
|
||||||
from .database_extensions import ApiKeyProviderMapping, ProviderUsageTracking
|
from .database_extensions import ApiKeyProviderMapping, ProviderUsageTracking
|
||||||
|
|||||||
@@ -635,6 +635,10 @@ class ProviderUpdateRequest(BaseModel):
|
|||||||
None,
|
None,
|
||||||
description="格式转换时是否保持优先级(True=保持原优先级,False=需要转换时降级)",
|
description="格式转换时是否保持优先级(True=保持原优先级,False=需要转换时降级)",
|
||||||
)
|
)
|
||||||
|
enable_format_conversion: bool | None = Field(
|
||||||
|
None,
|
||||||
|
description="是否允许格式转换(提供商级别开关)",
|
||||||
|
)
|
||||||
is_active: bool | None = None
|
is_active: bool | None = None
|
||||||
billing_type: str | None = Field(
|
billing_type: str | None = Field(
|
||||||
None, description="计费类型:monthly_quota/pay_as_you_go/free_tier"
|
None, description="计费类型:monthly_quota/pay_as_you_go/free_tier"
|
||||||
@@ -667,6 +671,10 @@ class ProviderWithEndpointsSummary(BaseModel):
|
|||||||
default=False,
|
default=False,
|
||||||
description="格式转换时是否保持优先级(True=保持原优先级,False=需要转换时降级)",
|
description="格式转换时是否保持优先级(True=保持原优先级,False=需要转换时降级)",
|
||||||
)
|
)
|
||||||
|
enable_format_conversion: bool = Field(
|
||||||
|
default=True,
|
||||||
|
description="是否允许格式转换(提供商级别开关)",
|
||||||
|
)
|
||||||
is_active: bool
|
is_active: bool
|
||||||
|
|
||||||
# 计费相关字段
|
# 计费相关字段
|
||||||
|
|||||||
@@ -7,6 +7,7 @@
|
|||||||
from src.core.modules.base import ModuleDefinition
|
from src.core.modules.base import ModuleDefinition
|
||||||
|
|
||||||
# 导入所有模块定义
|
# 导入所有模块定义
|
||||||
|
from src.modules.gemini_files import gemini_files_module
|
||||||
from src.modules.ldap import ldap_module
|
from src.modules.ldap import ldap_module
|
||||||
from src.modules.oauth import oauth_module
|
from src.modules.oauth import oauth_module
|
||||||
|
|
||||||
@@ -14,6 +15,7 @@ from src.modules.oauth import oauth_module
|
|||||||
ALL_MODULES: list[ModuleDefinition] = [
|
ALL_MODULES: list[ModuleDefinition] = [
|
||||||
ldap_module,
|
ldap_module,
|
||||||
oauth_module,
|
oauth_module,
|
||||||
|
gemini_files_module,
|
||||||
]
|
]
|
||||||
|
|
||||||
__all__ = ["ALL_MODULES"]
|
__all__ = ["ALL_MODULES"]
|
||||||
|
|||||||
96
src/modules/gemini_files/__init__.py
Normal file
96
src/modules/gemini_files/__init__.py
Normal file
@@ -0,0 +1,96 @@
|
|||||||
|
"""
|
||||||
|
Gemini Files 文件管理模块
|
||||||
|
|
||||||
|
提供 Gemini Files API 文件上传和管理功能:
|
||||||
|
- 文件上传到 Google Gemini Files API
|
||||||
|
- 文件映射管理(file_id → key_id)
|
||||||
|
- 文件列表查看和删除
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
from src.core.modules.base import (
|
||||||
|
ModuleCategory,
|
||||||
|
ModuleDefinition,
|
||||||
|
ModuleHealth,
|
||||||
|
ModuleMetadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
|
||||||
|
def _get_router() -> Any:
|
||||||
|
"""延迟导入路由"""
|
||||||
|
from src.api.admin.gemini_files import router
|
||||||
|
|
||||||
|
return router
|
||||||
|
|
||||||
|
|
||||||
|
async def _health_check() -> ModuleHealth:
|
||||||
|
"""健康检查"""
|
||||||
|
# 检查是否有可用的 gemini_files 能力的 Key
|
||||||
|
from src.database import create_session
|
||||||
|
from src.models.database import ProviderAPIKey
|
||||||
|
|
||||||
|
db = create_session()
|
||||||
|
try:
|
||||||
|
# 查找有 gemini_files 能力的 Key
|
||||||
|
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
|
||||||
|
has_capable_key = any(
|
||||||
|
key.capabilities and key.capabilities.get("gemini_files", False) for key in keys
|
||||||
|
)
|
||||||
|
if has_capable_key:
|
||||||
|
return ModuleHealth.HEALTHY
|
||||||
|
return ModuleHealth.DEGRADED
|
||||||
|
except Exception:
|
||||||
|
return ModuleHealth.UNKNOWN
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_config(db: Session) -> tuple[bool, str]:
|
||||||
|
"""
|
||||||
|
验证 Gemini Files 模块配置
|
||||||
|
|
||||||
|
检查项:
|
||||||
|
1. 至少有一个有 gemini_files 能力的 Provider Key
|
||||||
|
"""
|
||||||
|
from src.models.database import ProviderAPIKey
|
||||||
|
|
||||||
|
# 查找有 gemini_files 能力的 Key
|
||||||
|
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
|
||||||
|
capable_keys = [
|
||||||
|
key for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
|
||||||
|
]
|
||||||
|
|
||||||
|
if not capable_keys:
|
||||||
|
return (
|
||||||
|
False,
|
||||||
|
"至少启用一个具有「Gemini 文件 API」能力的 Key",
|
||||||
|
)
|
||||||
|
|
||||||
|
return True, ""
|
||||||
|
|
||||||
|
|
||||||
|
gemini_files_module = ModuleDefinition(
|
||||||
|
metadata=ModuleMetadata(
|
||||||
|
name="gemini_files",
|
||||||
|
display_name="文件缓存",
|
||||||
|
description="管理 Gemini Files API 上传的文件,支持文件上传、查看和删除",
|
||||||
|
category=ModuleCategory.INTEGRATION,
|
||||||
|
env_key="GEMINI_FILES_AVAILABLE",
|
||||||
|
default_available=True,
|
||||||
|
required_packages=[],
|
||||||
|
api_prefix="/api/admin/gemini-files",
|
||||||
|
admin_route="/admin/gemini-files",
|
||||||
|
admin_menu_icon="FileUp",
|
||||||
|
admin_menu_group="system",
|
||||||
|
admin_menu_order=60,
|
||||||
|
),
|
||||||
|
router_factory=_get_router,
|
||||||
|
health_check=_health_check,
|
||||||
|
validate_config=_validate_config,
|
||||||
|
)
|
||||||
@@ -804,8 +804,33 @@ class OAuthService:
|
|||||||
|
|
||||||
cfg = OAuthService._get_provider_config(db, provider_type)
|
cfg = OAuthService._get_provider_config(db, provider_type)
|
||||||
|
|
||||||
|
# Read all required fields first, then release DB connection before any awaits.
|
||||||
|
# This prevents holding a pooled connection while doing network I/O.
|
||||||
auth_url = provider.get_effective_authorization_url(cfg)
|
auth_url = provider.get_effective_authorization_url(cfg)
|
||||||
token_url = provider.get_effective_token_url(cfg)
|
token_url = provider.get_effective_token_url(cfg)
|
||||||
|
redirect_uri = cfg.redirect_uri
|
||||||
|
client_id = cfg.client_id
|
||||||
|
has_secret = bool(cfg.client_secret_encrypted)
|
||||||
|
client_secret = cfg.get_client_secret() if has_secret else None
|
||||||
|
|
||||||
|
# Release DB connection (safe only when session has no pending changes).
|
||||||
|
try:
|
||||||
|
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
|
||||||
|
except Exception:
|
||||||
|
has_pending_changes = False
|
||||||
|
if not has_pending_changes:
|
||||||
|
original_expire_on_commit = getattr(db, "expire_on_commit", True)
|
||||||
|
db.expire_on_commit = False
|
||||||
|
try:
|
||||||
|
if db.in_transaction():
|
||||||
|
db.commit()
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
db.rollback()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
db.expire_on_commit = original_expire_on_commit
|
||||||
|
|
||||||
async def _reachable(url: str) -> bool:
|
async def _reachable(url: str) -> bool:
|
||||||
try:
|
try:
|
||||||
@@ -823,7 +848,7 @@ class OAuthService:
|
|||||||
secret_status = "unknown"
|
secret_status = "unknown"
|
||||||
details = ""
|
details = ""
|
||||||
|
|
||||||
if cfg.client_secret_encrypted:
|
if has_secret and client_secret:
|
||||||
# 使用无效 code 做一次 token 请求(仅做粗略判定)
|
# 使用无效 code 做一次 token 请求(仅做粗略判定)
|
||||||
try:
|
try:
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
@@ -834,9 +859,9 @@ class OAuthService:
|
|||||||
data={
|
data={
|
||||||
"grant_type": "authorization_code",
|
"grant_type": "authorization_code",
|
||||||
"code": "invalid",
|
"code": "invalid",
|
||||||
"redirect_uri": cfg.redirect_uri,
|
"redirect_uri": redirect_uri,
|
||||||
"client_id": cfg.client_id,
|
"client_id": client_id,
|
||||||
"client_secret": cfg.get_client_secret(),
|
"client_secret": client_secret,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -29,6 +29,8 @@ from src.services.billing.models import (
|
|||||||
CostBreakdown,
|
CostBreakdown,
|
||||||
StandardizedUsage,
|
StandardizedUsage,
|
||||||
)
|
)
|
||||||
|
from src.services.billing.schema import BillingSnapshot, CostResult
|
||||||
|
from src.services.billing.service import BillingService
|
||||||
from src.services.billing.templates import BILLING_TEMPLATE_REGISTRY, BillingTemplates
|
from src.services.billing.templates import BILLING_TEMPLATE_REGISTRY, BillingTemplates
|
||||||
from src.services.billing.usage_mapper import UsageMapper, map_usage, map_usage_from_response
|
from src.services.billing.usage_mapper import UsageMapper, map_usage, map_usage_from_response
|
||||||
|
|
||||||
@@ -44,6 +46,10 @@ __all__ = [
|
|||||||
# 计算器
|
# 计算器
|
||||||
"BillingCalculator",
|
"BillingCalculator",
|
||||||
"calculate_request_cost",
|
"calculate_request_cost",
|
||||||
|
# 统一入口(Phase2)
|
||||||
|
"BillingService",
|
||||||
|
"BillingSnapshot",
|
||||||
|
"CostResult",
|
||||||
# 映射器
|
# 映射器
|
||||||
"UsageMapper",
|
"UsageMapper",
|
||||||
"map_usage",
|
"map_usage",
|
||||||
|
|||||||
64
src/services/billing/schema.py
Normal file
64
src/services/billing/schema.py
Normal file
@@ -0,0 +1,64 @@
|
|||||||
|
"""
|
||||||
|
Billing schema (stable contracts)
|
||||||
|
|
||||||
|
These dataclasses are meant to be stored in `Usage.request_metadata` / `Task.request_metadata`
|
||||||
|
for auditability. They are internal-only and MUST NOT be exposed to end users without sanitizing.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Literal
|
||||||
|
|
||||||
|
BILLING_SNAPSHOT_SCHEMA_VERSION = "1.0"
|
||||||
|
|
||||||
|
BillingSnapshotStatus = Literal["complete", "incomplete", "no_rule", "legacy"]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class BillingSnapshot:
|
||||||
|
"""Stable billing snapshot for audit."""
|
||||||
|
|
||||||
|
schema_version: str = BILLING_SNAPSHOT_SCHEMA_VERSION
|
||||||
|
|
||||||
|
# Rule info (optional for legacy/no_rule)
|
||||||
|
rule_id: str | None = None
|
||||||
|
rule_name: str | None = None
|
||||||
|
scope: str | None = None
|
||||||
|
|
||||||
|
# Rule expression (internal, do not expose to clients)
|
||||||
|
expression: str | None = None
|
||||||
|
|
||||||
|
# Dimensions
|
||||||
|
dimensions_used: dict[str, Any] = field(default_factory=dict)
|
||||||
|
missing_required: list[str] = field(default_factory=list)
|
||||||
|
|
||||||
|
# Result
|
||||||
|
cost: float = 0.0
|
||||||
|
status: BillingSnapshotStatus = "no_rule"
|
||||||
|
|
||||||
|
# Audit
|
||||||
|
calculated_at: str = "" # ISO 8601
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"rule_id": self.rule_id,
|
||||||
|
"rule_name": self.rule_name,
|
||||||
|
"scope": self.scope,
|
||||||
|
"expression": self.expression,
|
||||||
|
"dimensions_used": self.dimensions_used,
|
||||||
|
"missing_required": self.missing_required,
|
||||||
|
"cost": self.cost,
|
||||||
|
"status": self.status,
|
||||||
|
"calculated_at": self.calculated_at,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CostResult:
|
||||||
|
"""Billing calculation output."""
|
||||||
|
|
||||||
|
cost: float
|
||||||
|
status: BillingSnapshotStatus
|
||||||
|
snapshot: BillingSnapshot
|
||||||
145
src/services/billing/service.py
Normal file
145
src/services/billing/service.py
Normal file
@@ -0,0 +1,145 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.config.settings import config
|
||||||
|
from src.core.logger import logger
|
||||||
|
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.model.cost import ModelCostService
|
||||||
|
|
||||||
|
from .schema import BILLING_SNAPSHOT_SCHEMA_VERSION, BillingSnapshot, CostResult
|
||||||
|
|
||||||
|
|
||||||
|
class BillingService:
|
||||||
|
"""
|
||||||
|
BillingService (pure-ish application helper for billing domain).
|
||||||
|
|
||||||
|
Notes:
|
||||||
|
- This service **does not** write Usage rows.
|
||||||
|
- It may read billing rules & collectors from DB.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, db: Session):
|
||||||
|
self.db = db
|
||||||
|
self._formula_engine = FormulaEngine()
|
||||||
|
self._dimension_collector = DimensionCollectorService(db)
|
||||||
|
|
||||||
|
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]:
|
||||||
|
return self._dimension_collector.collect_dimensions(
|
||||||
|
api_format=api_format,
|
||||||
|
task_type=task_type,
|
||||||
|
request=request,
|
||||||
|
response=response,
|
||||||
|
metadata=metadata,
|
||||||
|
base_dimensions=base_dimensions,
|
||||||
|
)
|
||||||
|
|
||||||
|
def calculate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
task_type: str,
|
||||||
|
model: str,
|
||||||
|
provider_id: str,
|
||||||
|
dimensions: dict[str, Any],
|
||||||
|
strict_mode: bool | None = None,
|
||||||
|
) -> CostResult:
|
||||||
|
"""
|
||||||
|
Calculate cost for a task.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
CostResult (includes BillingSnapshot)
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
BillingIncompleteError: when strict_mode=True and required dims missing.
|
||||||
|
"""
|
||||||
|
strict = config.billing_strict_mode if strict_mode is None else bool(strict_mode)
|
||||||
|
|
||||||
|
lookup = BillingRuleService.find_rule(
|
||||||
|
self.db,
|
||||||
|
provider_id=provider_id,
|
||||||
|
model_name=model,
|
||||||
|
task_type=task_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
if lookup and lookup.rule and lookup.rule.expression:
|
||||||
|
rule = lookup.rule
|
||||||
|
result = self._formula_engine.evaluate(
|
||||||
|
expression=rule.expression,
|
||||||
|
variables=rule.variables or {},
|
||||||
|
dimensions=dimensions,
|
||||||
|
dimension_mappings=rule.dimension_mappings or {},
|
||||||
|
strict_mode=strict,
|
||||||
|
)
|
||||||
|
cost = float(result.cost) if result.status == "complete" else 0.0
|
||||||
|
snapshot = BillingSnapshot(
|
||||||
|
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
|
||||||
|
rule_id=str(rule.id),
|
||||||
|
rule_name=str(rule.name),
|
||||||
|
scope=str(getattr(lookup, "scope", None) or ""),
|
||||||
|
expression=str(rule.expression),
|
||||||
|
dimensions_used=dimensions,
|
||||||
|
missing_required=result.missing_required,
|
||||||
|
cost=cost,
|
||||||
|
status=result.status,
|
||||||
|
calculated_at=datetime.now(timezone.utc).isoformat(),
|
||||||
|
)
|
||||||
|
return CostResult(cost=cost, status=result.status, snapshot=snapshot)
|
||||||
|
|
||||||
|
# No rule fallback
|
||||||
|
if task_type in ("chat", "cli"):
|
||||||
|
input_tokens = int(dimensions.get("input_tokens") or 0)
|
||||||
|
output_tokens = int(dimensions.get("output_tokens") or 0)
|
||||||
|
cost = float(
|
||||||
|
ModelCostService.calculate_cost(
|
||||||
|
model=model,
|
||||||
|
input_tokens=input_tokens,
|
||||||
|
output_tokens=output_tokens,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
snapshot = BillingSnapshot(
|
||||||
|
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
|
||||||
|
rule_id=None,
|
||||||
|
rule_name=None,
|
||||||
|
scope=None,
|
||||||
|
expression=None,
|
||||||
|
dimensions_used=dimensions,
|
||||||
|
missing_required=[],
|
||||||
|
cost=cost,
|
||||||
|
status="legacy",
|
||||||
|
calculated_at=datetime.now(timezone.utc).isoformat(),
|
||||||
|
)
|
||||||
|
return CostResult(cost=cost, status="legacy", snapshot=snapshot)
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"No billing rule for task (task_type=%s, model=%s, provider_id=%s)",
|
||||||
|
task_type,
|
||||||
|
model,
|
||||||
|
provider_id,
|
||||||
|
)
|
||||||
|
snapshot = BillingSnapshot(
|
||||||
|
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
|
||||||
|
rule_id=None,
|
||||||
|
rule_name=None,
|
||||||
|
scope=None,
|
||||||
|
expression=None,
|
||||||
|
dimensions_used=dimensions,
|
||||||
|
missing_required=[],
|
||||||
|
cost=0.0,
|
||||||
|
status="no_rule",
|
||||||
|
calculated_at=datetime.now(timezone.utc).isoformat(),
|
||||||
|
)
|
||||||
|
return CostResult(cost=0.0, status="no_rule", snapshot=snapshot)
|
||||||
72
src/services/cache/aware_scheduler.py
vendored
72
src/services/cache/aware_scheduler.py
vendored
@@ -182,6 +182,43 @@ class CacheAwareScheduler:
|
|||||||
"last_reservation_result": None,
|
"last_reservation_result": None,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _release_db_connection_before_await(db: Session) -> None:
|
||||||
|
"""
|
||||||
|
Best-effort: end a read-only transaction before awaiting async I/O.
|
||||||
|
|
||||||
|
This scheduler does a lot of async work (cache/Redis) mixed with sync SQLAlchemy reads.
|
||||||
|
If a SELECT has already started a transaction, the pooled connection can remain checked
|
||||||
|
out while we await, causing pool pressure under concurrency.
|
||||||
|
|
||||||
|
Safety:
|
||||||
|
- Only commits when the Session has no ORM pending changes.
|
||||||
|
- Temporarily disables expire_on_commit to keep already-loaded ORM objects usable.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if db is None:
|
||||||
|
return
|
||||||
|
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
|
||||||
|
if has_pending_changes:
|
||||||
|
return
|
||||||
|
if not db.in_transaction():
|
||||||
|
return
|
||||||
|
|
||||||
|
original_expire_on_commit = getattr(db, "expire_on_commit", True)
|
||||||
|
db.expire_on_commit = False
|
||||||
|
try:
|
||||||
|
db.commit()
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
db.rollback()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
db.expire_on_commit = original_expire_on_commit
|
||||||
|
except Exception:
|
||||||
|
# Never let this optimization break scheduling
|
||||||
|
return
|
||||||
|
|
||||||
async def _ensure_initialized(self) -> None:
|
async def _ensure_initialized(self) -> None:
|
||||||
"""确保所有异步组件已初始化"""
|
"""确保所有异步组件已初始化"""
|
||||||
if self._affinity_manager is None:
|
if self._affinity_manager is None:
|
||||||
@@ -577,6 +614,8 @@ class CacheAwareScheduler:
|
|||||||
Returns:
|
Returns:
|
||||||
(候选列表, global_model_id) - global_model_id 用于缓存亲和性
|
(候选列表, global_model_id) - global_model_id 用于缓存亲和性
|
||||||
"""
|
"""
|
||||||
|
# If the caller already touched the DB, release the connection before we do async work.
|
||||||
|
self._release_db_connection_before_await(db)
|
||||||
await self._ensure_initialized()
|
await self._ensure_initialized()
|
||||||
|
|
||||||
target_format = normalize_endpoint_signature(api_format)
|
target_format = normalize_endpoint_signature(api_format)
|
||||||
@@ -648,6 +687,9 @@ class CacheAwareScheduler:
|
|||||||
provider_limit=provider_limit,
|
provider_limit=provider_limit,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Provider query starts a transaction; release connection before entering async candidate build.
|
||||||
|
self._release_db_connection_before_await(db)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"[Scheduler] Found %d active providers",
|
"[Scheduler] Found %d active providers",
|
||||||
len(providers),
|
len(providers),
|
||||||
@@ -680,8 +722,13 @@ class CacheAwareScheduler:
|
|||||||
|
|
||||||
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤)
|
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤)
|
||||||
from src.config.settings import config
|
from src.config.settings import config
|
||||||
|
from src.services.system.config import SystemConfigService
|
||||||
|
|
||||||
global_conversion_enabled = config.format_conversion_enabled
|
# 全局格式转换开关:优先使用数据库配置,回退到环境变量
|
||||||
|
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
|
||||||
|
# 如果环境变量明确禁用,则禁用(环境变量可作为强制禁用开关)
|
||||||
|
if not config.format_conversion_enabled:
|
||||||
|
global_conversion_enabled = False
|
||||||
candidates = await self._build_candidates(
|
candidates = await self._build_candidates(
|
||||||
db=db,
|
db=db,
|
||||||
providers=providers,
|
providers=providers,
|
||||||
@@ -801,6 +848,9 @@ class CacheAwareScheduler:
|
|||||||
- supported_capabilities: 模型支持的能力列表
|
- supported_capabilities: 模型支持的能力列表
|
||||||
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
|
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
|
||||||
"""
|
"""
|
||||||
|
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
|
||||||
|
self._release_db_connection_before_await(db)
|
||||||
|
|
||||||
# 使用 ModelCacheService 解析模型名称(支持映射名)
|
# 使用 ModelCacheService 解析模型名称(支持映射名)
|
||||||
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
||||||
db, model_name
|
db, model_name
|
||||||
@@ -1120,18 +1170,34 @@ class CacheAwareScheduler:
|
|||||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 计算格式转换的有效开关状态(三层优先级)
|
||||||
|
# 全局 ON → 强制允许(跳过端点检查)
|
||||||
|
# 全局 OFF → 提供商 ON → 强制允许(跳过端点检查)
|
||||||
|
# 全局 OFF → 提供商 OFF → 看端点配置
|
||||||
|
provider_allows_conversion = getattr(provider, "enable_format_conversion", True)
|
||||||
|
effective_conversion_enabled = (
|
||||||
|
global_conversion_enabled or provider_allows_conversion
|
||||||
|
)
|
||||||
|
# 如果全局或提供商开关为 ON,跳过端点配置检查
|
||||||
|
skip_endpoint_check = global_conversion_enabled or provider_allows_conversion
|
||||||
|
|
||||||
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
|
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
|
||||||
client_format_str,
|
client_format_str,
|
||||||
endpoint_format_str,
|
endpoint_format_str,
|
||||||
getattr(endpoint, "format_acceptance_config", None),
|
getattr(endpoint, "format_acceptance_config", None),
|
||||||
is_stream,
|
is_stream,
|
||||||
global_conversion_enabled,
|
effective_conversion_enabled,
|
||||||
|
skip_endpoint_check=skip_endpoint_check,
|
||||||
)
|
)
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, reason=%s",
|
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, "
|
||||||
|
"global=%s, provider=%s, skip_endpoint=%s, reason=%s",
|
||||||
client_format_str,
|
client_format_str,
|
||||||
endpoint_format_str,
|
endpoint_format_str,
|
||||||
is_compatible,
|
is_compatible,
|
||||||
|
global_conversion_enabled,
|
||||||
|
provider_allows_conversion,
|
||||||
|
skip_endpoint_check,
|
||||||
_compat_reason,
|
_compat_reason,
|
||||||
)
|
)
|
||||||
if not is_compatible:
|
if not is_compatible:
|
||||||
|
|||||||
28
src/services/candidate/__init__.py
Normal file
28
src/services/candidate/__init__.py
Normal file
@@ -0,0 +1,28 @@
|
|||||||
|
"""
|
||||||
|
Candidate domain (Phase2)
|
||||||
|
|
||||||
|
This package centralizes:
|
||||||
|
- candidate resolving (Provider/Endpoint/Key combinations)
|
||||||
|
- request_candidates recording & audit
|
||||||
|
- failover execution policies
|
||||||
|
"""
|
||||||
|
|
||||||
|
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
|
||||||
|
from src.services.candidate.schema import (
|
||||||
|
CANDIDATE_KEY_SCHEMA_VERSION,
|
||||||
|
CandidateKey,
|
||||||
|
CandidateResult,
|
||||||
|
)
|
||||||
|
from src.services.candidate.service import CandidateService
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"CandidateService",
|
||||||
|
# schema
|
||||||
|
"CANDIDATE_KEY_SCHEMA_VERSION",
|
||||||
|
"CandidateKey",
|
||||||
|
"CandidateResult",
|
||||||
|
# policies
|
||||||
|
"RetryMode",
|
||||||
|
"RetryPolicy",
|
||||||
|
"SkipPolicy",
|
||||||
|
]
|
||||||
57
src/services/candidate/failover.py
Normal file
57
src/services/candidate/failover.py
Normal file
@@ -0,0 +1,57 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any, Protocol
|
||||||
|
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||||
|
|
||||||
|
from .policy import RetryPolicy, SkipPolicy
|
||||||
|
from .schema import CandidateKey, CandidateResult
|
||||||
|
|
||||||
|
|
||||||
|
class AttemptFunc(Protocol):
|
||||||
|
async def __call__(self, candidate: ProviderCandidate) -> Any: ...
|
||||||
|
|
||||||
|
|
||||||
|
class FailoverEngine:
|
||||||
|
"""
|
||||||
|
FailoverEngine executes candidate attempts under policies.
|
||||||
|
|
||||||
|
Phase2 scaffolding: implementation will gradually replace legacy orchestrators.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, db: Session):
|
||||||
|
self.db = db
|
||||||
|
|
||||||
|
async def execute(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
candidates: list[ProviderCandidate],
|
||||||
|
attempt_func: AttemptFunc,
|
||||||
|
retry_policy: RetryPolicy,
|
||||||
|
skip_policy: SkipPolicy,
|
||||||
|
request_id: str | None = None,
|
||||||
|
max_candidates: int | None = None,
|
||||||
|
) -> CandidateResult:
|
||||||
|
# NOTE: intentionally minimal for now; legacy orchestrators still in use.
|
||||||
|
# This will be implemented when migrating video/chat flows to CandidateService.
|
||||||
|
_ = (retry_policy, skip_policy, request_id, max_candidates)
|
||||||
|
candidate_keys: list[CandidateKey] = []
|
||||||
|
for idx, cand in enumerate(candidates):
|
||||||
|
candidate_keys.append(
|
||||||
|
CandidateKey(
|
||||||
|
candidate_index=idx,
|
||||||
|
provider_id=str(cand.provider.id),
|
||||||
|
provider_name=str(cand.provider.name),
|
||||||
|
endpoint_id=str(cand.endpoint.id),
|
||||||
|
key_id=str(cand.key.id),
|
||||||
|
key_name=str(getattr(cand.key, "name", "") or ""),
|
||||||
|
auth_type=str(getattr(cand.key, "auth_type", "") or ""),
|
||||||
|
priority=int(getattr(cand.key, "priority", 0) or 0),
|
||||||
|
is_cached=bool(getattr(cand, "is_cached", False)),
|
||||||
|
status="available",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
raise NotImplementedError("FailoverEngine.execute is not implemented yet")
|
||||||
41
src/services/candidate/policy.py
Normal file
41
src/services/candidate/policy.py
Normal file
@@ -0,0 +1,41 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class RetryMode(str, Enum):
|
||||||
|
"""Retry mode for candidate attempts."""
|
||||||
|
|
||||||
|
PRE_EXPAND = "pre_expand" # pre-create retry slots (sync)
|
||||||
|
ON_DEMAND = "on_demand" # create retry record only when retry happens
|
||||||
|
DISABLED = "disabled" # no retry
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class RetryPolicy:
|
||||||
|
"""Unified retry policy."""
|
||||||
|
|
||||||
|
mode: RetryMode = RetryMode.DISABLED
|
||||||
|
max_retries: int = 1
|
||||||
|
retry_on_cached_only: bool = True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def for_sync_task(cls) -> "RetryPolicy":
|
||||||
|
return cls(mode=RetryMode.PRE_EXPAND, max_retries=2)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def for_async_task(cls) -> "RetryPolicy":
|
||||||
|
return cls(mode=RetryMode.DISABLED, max_retries=1)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def for_async_submit_with_retry(cls) -> "RetryPolicy":
|
||||||
|
return cls(mode=RetryMode.ON_DEMAND, max_retries=2)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SkipPolicy:
|
||||||
|
"""Rules for skipping unsupported candidates."""
|
||||||
|
|
||||||
|
allow_format_conversion: bool = True
|
||||||
|
supported_auth_types: set[str] | None = None
|
||||||
59
src/services/candidate/recorder.py
Normal file
59
src/services/candidate/recorder.py
Normal file
@@ -0,0 +1,59 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.models.database import RequestCandidate
|
||||||
|
|
||||||
|
from .schema import CandidateKey
|
||||||
|
|
||||||
|
|
||||||
|
class CandidateRecorder:
|
||||||
|
"""Read helpers for RequestCandidate audit data."""
|
||||||
|
|
||||||
|
def __init__(self, db: Session):
|
||||||
|
self.db = db
|
||||||
|
|
||||||
|
def get_candidate_keys(self, request_id: str) -> list[CandidateKey]:
|
||||||
|
rows: list[RequestCandidate] = (
|
||||||
|
self.db.query(RequestCandidate)
|
||||||
|
.filter(RequestCandidate.request_id == request_id)
|
||||||
|
.order_by(RequestCandidate.candidate_index.asc(), RequestCandidate.retry_index.asc())
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
|
||||||
|
result: list[CandidateKey] = []
|
||||||
|
for row in rows:
|
||||||
|
provider_name = None
|
||||||
|
if getattr(row, "provider", None) is not None:
|
||||||
|
provider_name = getattr(row.provider, "name", None)
|
||||||
|
|
||||||
|
key_name = None
|
||||||
|
auth_type = None
|
||||||
|
priority = None
|
||||||
|
if getattr(row, "key", None) is not None:
|
||||||
|
key_name = getattr(row.key, "name", None)
|
||||||
|
auth_type = getattr(row.key, "auth_type", None)
|
||||||
|
priority = getattr(row.key, "priority", None)
|
||||||
|
|
||||||
|
result.append(
|
||||||
|
CandidateKey(
|
||||||
|
candidate_index=int(row.candidate_index or 0),
|
||||||
|
retry_index=int(row.retry_index or 0),
|
||||||
|
provider_id=str(row.provider_id) if row.provider_id else None,
|
||||||
|
provider_name=str(provider_name) if provider_name else None,
|
||||||
|
endpoint_id=str(row.endpoint_id) if row.endpoint_id else None,
|
||||||
|
key_id=str(row.key_id) if row.key_id else None,
|
||||||
|
key_name=str(key_name) if key_name else None,
|
||||||
|
auth_type=str(auth_type) if auth_type else None,
|
||||||
|
priority=int(priority) if priority is not None else None,
|
||||||
|
is_cached=bool(getattr(row, "is_cached", False)),
|
||||||
|
status=str(getattr(row, "status", "") or "pending"),
|
||||||
|
skip_reason=getattr(row, "skip_reason", None),
|
||||||
|
error_type=getattr(row, "error_type", None),
|
||||||
|
error_message=getattr(row, "error_message", None),
|
||||||
|
status_code=getattr(row, "status_code", None),
|
||||||
|
latency_ms=getattr(row, "latency_ms", None),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
10
src/services/candidate/resolver.py
Normal file
10
src/services/candidate/resolver.py
Normal file
@@ -0,0 +1,10 @@
|
|||||||
|
"""
|
||||||
|
CandidateResolver facade import.
|
||||||
|
|
||||||
|
Phase2 keeps the implementation in `services/orchestration/` for compatibility,
|
||||||
|
and gradually migrates it into this package.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from src.services.orchestration.candidate_resolver import CandidateResolver
|
||||||
|
|
||||||
|
__all__ = ["CandidateResolver"]
|
||||||
71
src/services/candidate/schema.py
Normal file
71
src/services/candidate/schema.py
Normal file
@@ -0,0 +1,71 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||||
|
|
||||||
|
CANDIDATE_KEY_SCHEMA_VERSION = "1.0"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CandidateKey:
|
||||||
|
"""Stable candidate key snapshot for audit."""
|
||||||
|
|
||||||
|
schema_version: str = CANDIDATE_KEY_SCHEMA_VERSION
|
||||||
|
|
||||||
|
candidate_index: int = 0
|
||||||
|
retry_index: int = 0
|
||||||
|
|
||||||
|
provider_id: str | None = None
|
||||||
|
provider_name: str | None = None
|
||||||
|
endpoint_id: str | None = None
|
||||||
|
key_id: str | None = None
|
||||||
|
key_name: str | None = None
|
||||||
|
auth_type: str | None = None
|
||||||
|
priority: int | None = None
|
||||||
|
is_cached: bool = False
|
||||||
|
|
||||||
|
status: str = "pending" # pending/success/failed/skipped/available/...
|
||||||
|
skip_reason: str | None = None
|
||||||
|
error_type: str | None = None
|
||||||
|
error_message: str | None = None
|
||||||
|
status_code: int | None = None
|
||||||
|
latency_ms: int | None = None
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, Any]:
|
||||||
|
data: dict[str, Any] = {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"candidate_index": self.candidate_index,
|
||||||
|
"retry_index": self.retry_index,
|
||||||
|
"provider_id": self.provider_id,
|
||||||
|
"provider_name": self.provider_name,
|
||||||
|
"endpoint_id": self.endpoint_id,
|
||||||
|
"key_id": self.key_id,
|
||||||
|
"key_name": self.key_name,
|
||||||
|
"auth_type": self.auth_type,
|
||||||
|
"priority": self.priority,
|
||||||
|
"is_cached": self.is_cached,
|
||||||
|
"status": self.status,
|
||||||
|
"skip_reason": self.skip_reason,
|
||||||
|
"error_type": self.error_type,
|
||||||
|
"error_message": self.error_message,
|
||||||
|
"status_code": self.status_code,
|
||||||
|
"latency_ms": self.latency_ms,
|
||||||
|
}
|
||||||
|
# drop Nones for compact audit payload
|
||||||
|
return {k: v for k, v in data.items() if v is not None}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class CandidateResult:
|
||||||
|
"""Failover execution result."""
|
||||||
|
|
||||||
|
success: bool
|
||||||
|
selected: ProviderCandidate | None
|
||||||
|
selected_index: int | None
|
||||||
|
candidate_keys: list[CandidateKey]
|
||||||
|
|
||||||
|
external_task_id: str | None = None
|
||||||
|
error: Exception | None = None
|
||||||
|
last_status_code: int | None = None
|
||||||
515
src/services/candidate/service.py
Normal file
515
src/services/candidate/service.py
Normal file
@@ -0,0 +1,515 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from sqlalchemy import update
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.config.settings import config
|
||||||
|
from src.core.exceptions import ProviderNotAvailableException
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.models.database import ApiKey, RequestCandidate
|
||||||
|
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||||
|
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
||||||
|
from src.services.candidate.submit import (
|
||||||
|
AllCandidatesFailedError,
|
||||||
|
SubmitOutcome,
|
||||||
|
UpstreamClientRequestError,
|
||||||
|
)
|
||||||
|
from src.services.orchestration.error_classifier import ErrorClassifier
|
||||||
|
from src.services.system.config import SystemConfigService
|
||||||
|
|
||||||
|
from .recorder import CandidateRecorder
|
||||||
|
from .resolver import CandidateResolver
|
||||||
|
from .schema import CandidateKey
|
||||||
|
|
||||||
|
_SENSITIVE_PATTERN = re.compile(
|
||||||
|
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _sanitize(message: str, max_length: int = 200) -> str:
|
||||||
|
if not message:
|
||||||
|
return "request_failed"
|
||||||
|
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
|
||||||
|
|
||||||
|
|
||||||
|
class CandidateService:
|
||||||
|
"""
|
||||||
|
CandidateService (Facade).
|
||||||
|
|
||||||
|
Phase2 note: this is introduced as a new domain entrypoint. Legacy orchestrators
|
||||||
|
still exist and will be migrated gradually to use this service.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
|
||||||
|
self.db = db
|
||||||
|
self.redis = redis_client
|
||||||
|
self._cache_scheduler = None
|
||||||
|
self._resolver: CandidateResolver | None = None
|
||||||
|
self._error_classifier: ErrorClassifier | None = None
|
||||||
|
self._recorder = CandidateRecorder(db)
|
||||||
|
|
||||||
|
async def _ensure_initialized(self) -> None:
|
||||||
|
if self._cache_scheduler is not None:
|
||||||
|
return
|
||||||
|
|
||||||
|
priority_mode = SystemConfigService.get_config(
|
||||||
|
self.db,
|
||||||
|
"provider_priority_mode",
|
||||||
|
"provider",
|
||||||
|
)
|
||||||
|
scheduling_mode = SystemConfigService.get_config(
|
||||||
|
self.db,
|
||||||
|
"scheduling_mode",
|
||||||
|
"cache_affinity",
|
||||||
|
)
|
||||||
|
self._cache_scheduler = await get_cache_aware_scheduler(
|
||||||
|
self.redis,
|
||||||
|
priority_mode=priority_mode,
|
||||||
|
scheduling_mode=scheduling_mode,
|
||||||
|
)
|
||||||
|
self._resolver = CandidateResolver(db=self.db, cache_scheduler=self._cache_scheduler)
|
||||||
|
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
|
||||||
|
|
||||||
|
async def resolve(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
api_format: str,
|
||||||
|
model_name: str,
|
||||||
|
affinity_key: str,
|
||||||
|
user_api_key: ApiKey | None = None,
|
||||||
|
request_id: str | None = None,
|
||||||
|
is_stream: bool = False,
|
||||||
|
capability_requirements: dict[str, bool] | None = None,
|
||||||
|
preferred_key_ids: list[str] | None = None,
|
||||||
|
) -> tuple[list[ProviderCandidate], str]:
|
||||||
|
await self._ensure_initialized()
|
||||||
|
assert self._resolver is not None
|
||||||
|
return await self._resolver.fetch_candidates(
|
||||||
|
api_format=api_format,
|
||||||
|
model_name=model_name,
|
||||||
|
affinity_key=affinity_key,
|
||||||
|
user_api_key=user_api_key,
|
||||||
|
request_id=request_id,
|
||||||
|
is_stream=is_stream,
|
||||||
|
capability_requirements=capability_requirements,
|
||||||
|
preferred_key_ids=preferred_key_ids,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
|
||||||
|
"""
|
||||||
|
Decide whether an upstream HTTP error is a client error (no failover).
|
||||||
|
|
||||||
|
Rules:
|
||||||
|
- 401/403/429 are usually key/permission/ratelimit issues -> allow failover
|
||||||
|
- other 4xx: stop only if ErrorClassifier says it's a client error
|
||||||
|
"""
|
||||||
|
if status_code in (401, 403, 429):
|
||||||
|
return False
|
||||||
|
if 400 <= status_code < 500:
|
||||||
|
assert self._error_classifier is not None
|
||||||
|
return self._error_classifier.is_client_error(error_text)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def submit_with_failover(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
api_format: str,
|
||||||
|
model_name: str,
|
||||||
|
affinity_key: str,
|
||||||
|
user_api_key: ApiKey,
|
||||||
|
request_id: str | None,
|
||||||
|
task_type: str,
|
||||||
|
submit_func: Any,
|
||||||
|
extract_external_task_id: Any,
|
||||||
|
supported_auth_types: set[str] | None = None,
|
||||||
|
allow_format_conversion: bool = False,
|
||||||
|
capability_requirements: dict[str, bool] | None = None,
|
||||||
|
max_candidates: int | None = None,
|
||||||
|
) -> SubmitOutcome:
|
||||||
|
"""
|
||||||
|
Submit async task with failover, returning the selected candidate + external_task_id.
|
||||||
|
|
||||||
|
Phase2 submit entrypoint (replaces legacy submit orchestrator).
|
||||||
|
"""
|
||||||
|
# IMPORTANT:
|
||||||
|
# This method awaits upstream HTTP calls. If we have an open DB transaction before awaiting,
|
||||||
|
# the connection can be held for a long time (pool exhaustion under concurrency).
|
||||||
|
#
|
||||||
|
# Also note SQLAlchemy's default expire_on_commit=True would expire ORM objects and may
|
||||||
|
# trigger unexpected lazy DB loads after we commit (potentially during the await).
|
||||||
|
# We disable it temporarily to keep candidate/provider/key objects in-memory.
|
||||||
|
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
|
||||||
|
self.db.expire_on_commit = False
|
||||||
|
await self._ensure_initialized()
|
||||||
|
assert self._resolver is not None
|
||||||
|
try:
|
||||||
|
candidates, _global_model_id = await self._resolver.fetch_candidates(
|
||||||
|
api_format=api_format,
|
||||||
|
model_name=model_name,
|
||||||
|
affinity_key=affinity_key,
|
||||||
|
user_api_key=user_api_key,
|
||||||
|
request_id=request_id,
|
||||||
|
is_stream=False,
|
||||||
|
capability_requirements=capability_requirements,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not candidates:
|
||||||
|
raise ProviderNotAvailableException("No candidates available")
|
||||||
|
|
||||||
|
if max_candidates is not None and max_candidates > 0:
|
||||||
|
candidates = candidates[:max_candidates]
|
||||||
|
|
||||||
|
# Pre-create RequestCandidate records (no retry expand for async submit stage)
|
||||||
|
record_map: dict[tuple[int, int], str] = {}
|
||||||
|
if request_id:
|
||||||
|
try:
|
||||||
|
record_map = self.create_candidate_records(
|
||||||
|
candidates=candidates,
|
||||||
|
request_id=request_id,
|
||||||
|
user_api_key=user_api_key,
|
||||||
|
required_capabilities=capability_requirements,
|
||||||
|
expand_retries=False,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"[CandidateService] Failed to create candidate records: %s",
|
||||||
|
_sanitize(str(exc)),
|
||||||
|
)
|
||||||
|
record_map = {}
|
||||||
|
|
||||||
|
candidate_keys: list[dict[str, Any]] = []
|
||||||
|
eligible_count = 0
|
||||||
|
last_status_code: int | None = None
|
||||||
|
|
||||||
|
for idx, cand in enumerate(candidates):
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
|
||||||
|
|
||||||
|
candidate_info: dict[str, Any] = {
|
||||||
|
"index": idx,
|
||||||
|
"provider_id": cand.provider.id,
|
||||||
|
"provider_name": cand.provider.name,
|
||||||
|
"endpoint_id": cand.endpoint.id,
|
||||||
|
"key_id": cand.key.id,
|
||||||
|
"key_name": getattr(cand.key, "name", None),
|
||||||
|
"auth_type": auth_type,
|
||||||
|
"priority": getattr(cand.key, "priority", 0) or 0,
|
||||||
|
"is_cached": bool(getattr(cand, "is_cached", False)),
|
||||||
|
}
|
||||||
|
candidate_keys.append(candidate_info)
|
||||||
|
|
||||||
|
record_id = record_map.get((idx, 0))
|
||||||
|
|
||||||
|
# Scheduler marked skip
|
||||||
|
if getattr(cand, "is_skipped", False):
|
||||||
|
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
|
||||||
|
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||||
|
if record_id:
|
||||||
|
# record is usually already skipped, but keep it consistent
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(status="skipped", skip_reason=skip_reason)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Format conversion checks
|
||||||
|
# 优先级:全局开关 ON 强制允许,全局开关 OFF 看提供商开关
|
||||||
|
needs_conversion = bool(getattr(cand, "needs_conversion", False))
|
||||||
|
if needs_conversion:
|
||||||
|
# 1. Check handler-level switch (handler 不支持则直接跳过)
|
||||||
|
if not allow_format_conversion:
|
||||||
|
skip_reason = "format_conversion_not_supported"
|
||||||
|
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(status="skipped", skip_reason=skip_reason)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 2. Check global + provider switches
|
||||||
|
# 全局 ON → 允许;全局 OFF → 看提供商
|
||||||
|
from src.services.system.config import SystemConfigService
|
||||||
|
|
||||||
|
global_enabled = SystemConfigService.is_format_conversion_enabled(self.db)
|
||||||
|
provider_enabled = getattr(cand.provider, "enable_format_conversion", True)
|
||||||
|
effective_enabled = global_enabled or provider_enabled
|
||||||
|
|
||||||
|
if not effective_enabled:
|
||||||
|
skip_reason = "format_conversion_disabled"
|
||||||
|
candidate_info.update(
|
||||||
|
{
|
||||||
|
"skipped": True,
|
||||||
|
"skip_reason": skip_reason,
|
||||||
|
"global_conversion_enabled": global_enabled,
|
||||||
|
"provider_conversion_enabled": provider_enabled,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(status="skipped", skip_reason=skip_reason)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# auth_type filter
|
||||||
|
if supported_auth_types is not None and auth_type not in supported_auth_types:
|
||||||
|
skip_reason = f"unsupported_auth_type:{auth_type}"
|
||||||
|
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(status="skipped", skip_reason=skip_reason)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# billing rule filter
|
||||||
|
rule_lookup: BillingRuleLookupResult | None = None
|
||||||
|
has_billing_rule = True
|
||||||
|
if config.billing_require_rule:
|
||||||
|
rule_lookup = BillingRuleService.find_rule(
|
||||||
|
self.db,
|
||||||
|
provider_id=cand.provider.id,
|
||||||
|
model_name=model_name,
|
||||||
|
task_type=task_type,
|
||||||
|
)
|
||||||
|
has_billing_rule = rule_lookup is not None
|
||||||
|
if not has_billing_rule:
|
||||||
|
skip_reason = "billing_rule_missing"
|
||||||
|
candidate_info.update(
|
||||||
|
{"has_billing_rule": False, "skipped": True, "skip_reason": skip_reason}
|
||||||
|
)
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(status="skipped", skip_reason=skip_reason)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
candidate_info["has_billing_rule"] = has_billing_rule
|
||||||
|
|
||||||
|
eligible_count += 1
|
||||||
|
|
||||||
|
# Mark pending
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(status="pending", started_at=now)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Flush/commit BEFORE awaiting upstream submit to avoid holding DB connections
|
||||||
|
# during potentially slow network operations.
|
||||||
|
if self.db.in_transaction():
|
||||||
|
try:
|
||||||
|
self.db.commit()
|
||||||
|
except Exception:
|
||||||
|
self.db.rollback()
|
||||||
|
raise
|
||||||
|
|
||||||
|
# Attempt submit (upstream HTTP)
|
||||||
|
try:
|
||||||
|
response: httpx.Response = await submit_func(cand)
|
||||||
|
except Exception as exc:
|
||||||
|
finished_at = datetime.now(timezone.utc)
|
||||||
|
error_msg = _sanitize(str(exc))
|
||||||
|
candidate_info.update(
|
||||||
|
{
|
||||||
|
"attempt_status": "exception",
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"error_message": error_msg,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(
|
||||||
|
status="failed",
|
||||||
|
error_type=type(exc).__name__,
|
||||||
|
error_message=error_msg,
|
||||||
|
finished_at=finished_at,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
last_status_code = int(getattr(response, "status_code", 0) or 0)
|
||||||
|
|
||||||
|
if response.status_code >= 400:
|
||||||
|
finished_at = datetime.now(timezone.utc)
|
||||||
|
try:
|
||||||
|
error_text = response.text or ""
|
||||||
|
except Exception:
|
||||||
|
error_text = ""
|
||||||
|
error_msg = _sanitize(error_text)
|
||||||
|
candidate_info.update(
|
||||||
|
{
|
||||||
|
"attempt_status": "http_error",
|
||||||
|
"status_code": response.status_code,
|
||||||
|
"error_message": error_msg,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(
|
||||||
|
status="failed",
|
||||||
|
status_code=response.status_code,
|
||||||
|
error_type="http_error",
|
||||||
|
error_message=error_msg,
|
||||||
|
finished_at=finished_at,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if self._should_stop_on_http_error(
|
||||||
|
status_code=response.status_code, error_text=error_text
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
self.db.commit()
|
||||||
|
except Exception:
|
||||||
|
self.db.rollback()
|
||||||
|
raise UpstreamClientRequestError(
|
||||||
|
response=response,
|
||||||
|
candidate_keys=candidate_keys,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Parse JSON
|
||||||
|
payload: dict[str, Any] | None = None
|
||||||
|
try:
|
||||||
|
data = response.json()
|
||||||
|
if isinstance(data, dict):
|
||||||
|
payload = data
|
||||||
|
except Exception as exc:
|
||||||
|
finished_at = datetime.now(timezone.utc)
|
||||||
|
error_msg = _sanitize(str(exc))
|
||||||
|
candidate_info.update(
|
||||||
|
{
|
||||||
|
"attempt_status": "invalid_json",
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"error_message": error_msg,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(
|
||||||
|
status="failed",
|
||||||
|
status_code=response.status_code,
|
||||||
|
error_type="invalid_json",
|
||||||
|
error_message=error_msg,
|
||||||
|
finished_at=finished_at,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
external_task_id = extract_external_task_id(payload or {})
|
||||||
|
if not external_task_id:
|
||||||
|
finished_at = datetime.now(timezone.utc)
|
||||||
|
candidate_info.update(
|
||||||
|
{
|
||||||
|
"attempt_status": "empty_task_id",
|
||||||
|
"error_message": "Upstream returned empty task id",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(
|
||||||
|
status="failed",
|
||||||
|
status_code=response.status_code,
|
||||||
|
error_type="empty_task_id",
|
||||||
|
error_message="Upstream returned empty task id",
|
||||||
|
finished_at=finished_at,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Success
|
||||||
|
finished_at = datetime.now(timezone.utc)
|
||||||
|
candidate_info.update({"attempt_status": "success", "selected": True})
|
||||||
|
if record_id:
|
||||||
|
self.db.execute(
|
||||||
|
update(RequestCandidate)
|
||||||
|
.where(RequestCandidate.id == record_id)
|
||||||
|
.values(
|
||||||
|
status="success",
|
||||||
|
status_code=response.status_code,
|
||||||
|
finished_at=finished_at,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
self.db.commit()
|
||||||
|
except Exception:
|
||||||
|
self.db.rollback()
|
||||||
|
|
||||||
|
return SubmitOutcome(
|
||||||
|
candidate=cand,
|
||||||
|
candidate_keys=candidate_keys,
|
||||||
|
external_task_id=str(external_task_id),
|
||||||
|
rule_lookup=rule_lookup,
|
||||||
|
upstream_payload=payload,
|
||||||
|
upstream_headers=dict(response.headers),
|
||||||
|
upstream_status_code=response.status_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Persist candidate records before raising
|
||||||
|
try:
|
||||||
|
self.db.commit()
|
||||||
|
except Exception:
|
||||||
|
self.db.rollback()
|
||||||
|
|
||||||
|
if eligible_count == 0:
|
||||||
|
reason = "no_eligible_candidates"
|
||||||
|
if config.billing_require_rule:
|
||||||
|
reason = "no_candidate_with_billing_rule"
|
||||||
|
raise AllCandidatesFailedError(
|
||||||
|
reason=reason,
|
||||||
|
candidate_keys=candidate_keys,
|
||||||
|
last_status_code=last_status_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
raise AllCandidatesFailedError(
|
||||||
|
reason="all_candidates_failed",
|
||||||
|
candidate_keys=candidate_keys,
|
||||||
|
last_status_code=last_status_code,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
# Restore Session behavior for the rest of the request lifecycle.
|
||||||
|
self.db.expire_on_commit = original_expire_on_commit
|
||||||
|
|
||||||
|
def create_candidate_records(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
candidates: list[ProviderCandidate],
|
||||||
|
request_id: str,
|
||||||
|
user_api_key: ApiKey,
|
||||||
|
required_capabilities: dict[str, bool] | None = None,
|
||||||
|
expand_retries: bool = True,
|
||||||
|
) -> dict[tuple[int, int], str]:
|
||||||
|
# CandidateResolver.create_candidate_records is currently the canonical implementation
|
||||||
|
assert self._resolver is not None, "Call resolve() once before create_candidate_records()"
|
||||||
|
return self._resolver.create_candidate_records(
|
||||||
|
all_candidates=candidates,
|
||||||
|
request_id=request_id,
|
||||||
|
user_id=str(user_api_key.user_id),
|
||||||
|
user_api_key=user_api_key,
|
||||||
|
required_capabilities=required_capabilities,
|
||||||
|
expand_retries=expand_retries,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_candidate_keys(self, request_id: str) -> list["CandidateKey"]:
|
||||||
|
return self._recorder.get_candidate_keys(request_id)
|
||||||
66
src/services/candidate/submit.py
Normal file
66
src/services/candidate/submit.py
Normal file
@@ -0,0 +1,66 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||||
|
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class SubmitFunc(Protocol):
|
||||||
|
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class ExtractExternalTaskIdFunc(Protocol):
|
||||||
|
def __call__(self, payload: dict[str, Any]) -> str | None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class UpstreamClientRequestError(RuntimeError):
|
||||||
|
"""可判定为客户端请求问题(不应 failover)的上游错误。"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
response: httpx.Response,
|
||||||
|
candidate_keys: list[dict[str, Any]],
|
||||||
|
) -> None:
|
||||||
|
self.response = response
|
||||||
|
self.candidate_keys = candidate_keys
|
||||||
|
super().__init__(f"Upstream client error: HTTP {response.status_code}")
|
||||||
|
|
||||||
|
|
||||||
|
class AllCandidatesFailedError(RuntimeError):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
reason: str,
|
||||||
|
candidate_keys: list[dict[str, Any]],
|
||||||
|
last_status_code: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.reason = reason
|
||||||
|
self.candidate_keys = candidate_keys
|
||||||
|
self.last_status_code = last_status_code
|
||||||
|
super().__init__(f"All candidates failed: {reason}")
|
||||||
|
|
||||||
|
|
||||||
|
class CandidateUnsupportedError(RuntimeError):
|
||||||
|
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
|
||||||
|
|
||||||
|
|
||||||
|
class CandidateSubmissionError(RuntimeError):
|
||||||
|
"""候选提交异常(网络/解密/解析等)。"""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class SubmitOutcome:
|
||||||
|
candidate: ProviderCandidate
|
||||||
|
candidate_keys: list[dict[str, Any]]
|
||||||
|
external_task_id: str
|
||||||
|
rule_lookup: BillingRuleLookupResult | None
|
||||||
|
upstream_payload: dict[str, Any] | None = None
|
||||||
|
upstream_headers: dict[str, str] | None = None
|
||||||
|
upstream_status_code: int | None = None
|
||||||
@@ -1,21 +1,41 @@
|
|||||||
"""
|
"""
|
||||||
Gemini Files API - 文件与 Key 绑定缓存
|
Gemini Files API - 文件与 Key 绑定映射服务
|
||||||
|
|
||||||
用于在上传文件后记录 file_id -> provider_key_id,
|
用于在上传文件后记录 file_id -> provider_key_id,
|
||||||
并在后续 generateContent 请求中优先使用同一 Key。
|
并在后续 generateContent 请求中优先使用同一 Key。
|
||||||
|
|
||||||
|
存储策略:
|
||||||
|
- 数据库(持久化):主存储,支持服务重启后恢复
|
||||||
|
- Redis(缓存):加速读取,TTL=48小时
|
||||||
|
|
||||||
|
读取策略:
|
||||||
|
1. 先查 Redis 缓存
|
||||||
|
2. 缓存未命中时回查数据库
|
||||||
|
3. 从数据库读取后回填缓存
|
||||||
|
|
||||||
|
清理策略:
|
||||||
|
- 数据库中 expires_at 过期的记录由定时任务清理
|
||||||
|
- Redis 缓存由 TTL 自动过期
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Set
|
import uuid
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import delete
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.cache_service import CacheService
|
from src.core.cache_service import CacheService
|
||||||
|
from src.core.logger import logger
|
||||||
|
|
||||||
FILE_MAPPING_TTL_SECONDS = 60 * 60 * 48 # 48小时
|
FILE_MAPPING_TTL_SECONDS = 60 * 60 * 48 # 48小时
|
||||||
FILE_MAPPING_CACHE_PREFIX = "gemini_files:key"
|
FILE_MAPPING_CACHE_PREFIX = "gemini_files:key"
|
||||||
|
|
||||||
|
|
||||||
def _normalize_file_name(file_name: str) -> str:
|
def _normalize_file_name(file_name: str) -> str:
|
||||||
|
"""规范化文件名,确保以 files/ 开头"""
|
||||||
name = (file_name or "").strip()
|
name = (file_name or "").strip()
|
||||||
if not name:
|
if not name:
|
||||||
return ""
|
return ""
|
||||||
@@ -23,31 +43,294 @@ def _normalize_file_name(file_name: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def build_file_mapping_key(file_name: str) -> str:
|
def build_file_mapping_key(file_name: str) -> str:
|
||||||
|
"""构建 Redis 缓存键"""
|
||||||
normalized = _normalize_file_name(file_name)
|
normalized = _normalize_file_name(file_name)
|
||||||
return f"{FILE_MAPPING_CACHE_PREFIX}:{normalized}" if normalized else ""
|
return f"{FILE_MAPPING_CACHE_PREFIX}:{normalized}" if normalized else ""
|
||||||
|
|
||||||
|
|
||||||
async def store_file_key_mapping(file_name: str, key_id: str) -> None:
|
# =============================================================================
|
||||||
cache_key = build_file_mapping_key(file_name)
|
# 异步接口(用于请求处理流程)
|
||||||
if not cache_key or not key_id:
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
async def store_file_key_mapping(
|
||||||
|
file_name: str,
|
||||||
|
key_id: str,
|
||||||
|
user_id: str | None = None,
|
||||||
|
display_name: str | None = None,
|
||||||
|
mime_type: str | None = None,
|
||||||
|
source_hash: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
存储文件→Key 映射(同时写入 Redis 和数据库)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_name: 文件名(如 files/abc123)
|
||||||
|
key_id: Provider Key ID
|
||||||
|
user_id: 用户 ID(可选,用于权限验证)
|
||||||
|
display_name: 文件显示名(可选)
|
||||||
|
mime_type: 文件 MIME 类型(可选)
|
||||||
|
source_hash: 源文件哈希(可选,用于关联相同源文件的不同上传)
|
||||||
|
"""
|
||||||
|
normalized_name = _normalize_file_name(file_name)
|
||||||
|
if not normalized_name or not key_id:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# 1. 写入 Redis 缓存
|
||||||
|
cache_key = build_file_mapping_key(normalized_name)
|
||||||
await CacheService.set(cache_key, str(key_id), ttl_seconds=FILE_MAPPING_TTL_SECONDS)
|
await CacheService.set(cache_key, str(key_id), ttl_seconds=FILE_MAPPING_TTL_SECONDS)
|
||||||
|
|
||||||
|
# 2. 写入数据库(异步执行,不阻塞主流程)
|
||||||
|
try:
|
||||||
|
await _store_to_database(
|
||||||
|
file_name=normalized_name,
|
||||||
|
key_id=key_id,
|
||||||
|
user_id=user_id,
|
||||||
|
display_name=display_name,
|
||||||
|
mime_type=mime_type,
|
||||||
|
source_hash=source_hash,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
# 数据库写入失败只记录警告,不影响主流程
|
||||||
|
logger.warning(f"Failed to persist Gemini file mapping to database: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
async def _store_to_database(
|
||||||
|
file_name: str,
|
||||||
|
key_id: str,
|
||||||
|
user_id: str | None = None,
|
||||||
|
display_name: str | None = None,
|
||||||
|
mime_type: str | None = None,
|
||||||
|
source_hash: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""将映射写入数据库"""
|
||||||
|
from src.database import get_db_context
|
||||||
|
from src.models.database import GeminiFileMapping
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
expires_at = now + timedelta(hours=48)
|
||||||
|
|
||||||
|
with get_db_context() as db:
|
||||||
|
# 使用 upsert 逻辑:存在则更新,不存在则插入
|
||||||
|
existing = (
|
||||||
|
db.query(GeminiFileMapping).filter(GeminiFileMapping.file_name == file_name).first()
|
||||||
|
)
|
||||||
|
|
||||||
|
if existing:
|
||||||
|
# 更新现有记录
|
||||||
|
existing.key_id = key_id
|
||||||
|
existing.user_id = user_id
|
||||||
|
existing.display_name = display_name
|
||||||
|
existing.mime_type = mime_type
|
||||||
|
existing.source_hash = source_hash
|
||||||
|
existing.expires_at = expires_at
|
||||||
|
else:
|
||||||
|
# 插入新记录
|
||||||
|
mapping = GeminiFileMapping(
|
||||||
|
id=str(uuid.uuid4()),
|
||||||
|
file_name=file_name,
|
||||||
|
key_id=key_id,
|
||||||
|
user_id=user_id,
|
||||||
|
display_name=display_name,
|
||||||
|
mime_type=mime_type,
|
||||||
|
source_hash=source_hash,
|
||||||
|
created_at=now,
|
||||||
|
expires_at=expires_at,
|
||||||
|
)
|
||||||
|
db.add(mapping)
|
||||||
|
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
async def get_file_key_mapping(file_name: str) -> str | None:
|
async def get_file_key_mapping(file_name: str) -> str | None:
|
||||||
cache_key = build_file_mapping_key(file_name)
|
"""
|
||||||
if not cache_key:
|
获取文件→Key 映射
|
||||||
|
|
||||||
|
读取策略:
|
||||||
|
1. 先查 Redis 缓存
|
||||||
|
2. 缓存未命中时回查数据库
|
||||||
|
3. 从数据库读取后回填缓存
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_name: 文件名(如 files/abc123)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Provider Key ID,如果不存在或已过期则返回 None
|
||||||
|
"""
|
||||||
|
normalized_name = _normalize_file_name(file_name)
|
||||||
|
if not normalized_name:
|
||||||
return None
|
return None
|
||||||
value = await CacheService.get(cache_key)
|
|
||||||
if value:
|
cache_key = build_file_mapping_key(normalized_name)
|
||||||
return str(value)
|
|
||||||
|
# 1. 先查 Redis 缓存
|
||||||
|
cached_value = await CacheService.get(cache_key)
|
||||||
|
if cached_value:
|
||||||
|
return str(cached_value)
|
||||||
|
|
||||||
|
# 2. 缓存未命中,回查数据库
|
||||||
|
key_id = await _get_from_database(normalized_name)
|
||||||
|
|
||||||
|
if key_id:
|
||||||
|
# 3. 回填缓存(使用剩余有效期或默认 TTL)
|
||||||
|
await CacheService.set(cache_key, key_id, ttl_seconds=FILE_MAPPING_TTL_SECONDS)
|
||||||
|
logger.debug(f"Gemini file mapping cache refilled from database: {normalized_name}")
|
||||||
|
|
||||||
|
return key_id
|
||||||
|
|
||||||
|
|
||||||
|
async def _get_from_database(file_name: str) -> str | None:
|
||||||
|
"""从数据库查询映射"""
|
||||||
|
from src.database import get_db_context
|
||||||
|
from src.models.database import GeminiFileMapping
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
try:
|
||||||
|
with get_db_context() as db:
|
||||||
|
mapping = (
|
||||||
|
db.query(GeminiFileMapping)
|
||||||
|
.filter(
|
||||||
|
GeminiFileMapping.file_name == file_name,
|
||||||
|
GeminiFileMapping.expires_at > now, # 只返回未过期的
|
||||||
|
)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
|
||||||
|
if mapping:
|
||||||
|
return str(mapping.key_id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Failed to query Gemini file mapping from database: {e}")
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def get_all_key_ids_for_file(file_name: str) -> list[str]:
|
||||||
|
"""
|
||||||
|
获取支持指定文件的所有 Key ID 列表
|
||||||
|
|
||||||
|
当同一个源文件被上传到多个 Key 时,返回所有可用的 Key ID。
|
||||||
|
这允许系统在首选 Key 不可用时选择其他 Key。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_name: 文件名(如 files/abc123)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
所有支持该文件的 Key ID 列表(包括原始映射和具有相同 source_hash 的映射)
|
||||||
|
"""
|
||||||
|
from src.database import get_db_context
|
||||||
|
from src.models.database import GeminiFileMapping
|
||||||
|
|
||||||
|
normalized_name = _normalize_file_name(file_name)
|
||||||
|
if not normalized_name:
|
||||||
|
return []
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
try:
|
||||||
|
with get_db_context() as db:
|
||||||
|
# 首先获取原始映射
|
||||||
|
original_mapping = (
|
||||||
|
db.query(GeminiFileMapping)
|
||||||
|
.filter(
|
||||||
|
GeminiFileMapping.file_name == normalized_name,
|
||||||
|
GeminiFileMapping.expires_at > now,
|
||||||
|
)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
|
||||||
|
if not original_mapping:
|
||||||
|
return []
|
||||||
|
|
||||||
|
key_ids = [str(original_mapping.key_id)]
|
||||||
|
|
||||||
|
# 如果有 source_hash,查找所有具有相同 source_hash 的映射
|
||||||
|
if original_mapping.source_hash:
|
||||||
|
related_mappings = (
|
||||||
|
db.query(GeminiFileMapping)
|
||||||
|
.filter(
|
||||||
|
GeminiFileMapping.source_hash == original_mapping.source_hash,
|
||||||
|
GeminiFileMapping.expires_at > now,
|
||||||
|
GeminiFileMapping.file_name != normalized_name, # 排除原始映射
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
|
||||||
|
for mapping in related_mappings:
|
||||||
|
kid = str(mapping.key_id)
|
||||||
|
if kid not in key_ids:
|
||||||
|
key_ids.append(kid)
|
||||||
|
|
||||||
|
return key_ids
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Failed to query related Gemini file mappings: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
async def delete_file_key_mapping(file_name: str) -> None:
|
async def delete_file_key_mapping(file_name: str) -> None:
|
||||||
cache_key = build_file_mapping_key(file_name)
|
"""
|
||||||
if cache_key:
|
删除文件→Key 映射(同时从 Redis 和数据库删除)
|
||||||
await CacheService.delete(cache_key)
|
|
||||||
|
Args:
|
||||||
|
file_name: 文件名(如 files/abc123)
|
||||||
|
"""
|
||||||
|
normalized_name = _normalize_file_name(file_name)
|
||||||
|
if not normalized_name:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 1. 从 Redis 删除
|
||||||
|
cache_key = build_file_mapping_key(normalized_name)
|
||||||
|
await CacheService.delete(cache_key)
|
||||||
|
|
||||||
|
# 2. 从数据库删除
|
||||||
|
try:
|
||||||
|
await _delete_from_database(normalized_name)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Failed to delete Gemini file mapping from database: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
async def _delete_from_database(file_name: str) -> None:
|
||||||
|
"""从数据库删除映射"""
|
||||||
|
from src.database import get_db_context
|
||||||
|
from src.models.database import GeminiFileMapping
|
||||||
|
|
||||||
|
with get_db_context() as db:
|
||||||
|
db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.file_name == file_name))
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# 同步接口(用于定时任务等场景)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
def cleanup_expired_mappings(db: Session) -> int:
|
||||||
|
"""
|
||||||
|
清理过期的文件映射记录(同步方法,供定时任务调用)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: 数据库会话
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
删除的记录数
|
||||||
|
"""
|
||||||
|
from src.models.database import GeminiFileMapping
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
deleted_count = result.rowcount
|
||||||
|
if deleted_count > 0:
|
||||||
|
logger.info(f"Cleaned up {deleted_count} expired Gemini file mappings")
|
||||||
|
|
||||||
|
return deleted_count
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# 请求解析工具函数
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def _extract_file_name_from_uri(file_uri: str) -> str | None:
|
def _extract_file_name_from_uri(file_uri: str) -> str | None:
|
||||||
|
|||||||
@@ -39,6 +39,10 @@ MAX_CONCURRENT_REQUESTS = 5
|
|||||||
# 单个 Key 处理的超时时间(秒)
|
# 单个 Key 处理的超时时间(秒)
|
||||||
KEY_FETCH_TIMEOUT_SECONDS = 120
|
KEY_FETCH_TIMEOUT_SECONDS = 120
|
||||||
|
|
||||||
|
# 模型获取 HTTP 请求超时时间(秒)
|
||||||
|
# 使用较短的超时(10秒),避免不支持 /models 端点的提供商长时间阻塞
|
||||||
|
MODEL_FETCH_HTTP_TIMEOUT = 10.0
|
||||||
|
|
||||||
# 上游模型缓存 TTL(与定时任务间隔保持一致)
|
# 上游模型缓存 TTL(与定时任务间隔保持一致)
|
||||||
UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60
|
UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60
|
||||||
|
|
||||||
@@ -260,7 +264,48 @@ class ModelFetchScheduler:
|
|||||||
logger.exception(f"更新 Key {key_id} 错误信息失败")
|
logger.exception(f"更新 Key {key_id} 错误信息失败")
|
||||||
|
|
||||||
async def _fetch_models_for_key_by_id(self, key_id: str) -> str:
|
async def _fetch_models_for_key_by_id(self, key_id: str) -> str:
|
||||||
"""根据 Key ID 获取模型并更新,返回结果状态"""
|
"""
|
||||||
|
根据 Key ID 获取模型并更新,返回结果状态
|
||||||
|
|
||||||
|
优化:分两个阶段处理,HTTP 请求期间不持有数据库连接,避免阻塞其他请求
|
||||||
|
"""
|
||||||
|
# ========== 阶段 1:准备数据(短暂持有连接)==========
|
||||||
|
fetch_context = self._prepare_fetch_context(key_id)
|
||||||
|
if fetch_context is None:
|
||||||
|
return "skip"
|
||||||
|
if isinstance(fetch_context, str):
|
||||||
|
return fetch_context # "error" or "skip"
|
||||||
|
|
||||||
|
key_id, provider_id, provider_name, api_key_value, endpoint_configs = fetch_context
|
||||||
|
|
||||||
|
# ========== 阶段 2:HTTP 请求(不持有数据库连接)==========
|
||||||
|
# 使用较短的超时时间(10秒),避免长时间阻塞
|
||||||
|
all_models, errors, has_success = await fetch_models_from_endpoints(
|
||||||
|
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
|
||||||
|
)
|
||||||
|
|
||||||
|
# ========== 阶段 3:更新数据库(获取新连接)==========
|
||||||
|
return await self._update_key_after_fetch(
|
||||||
|
key_id=key_id,
|
||||||
|
provider_id=provider_id,
|
||||||
|
provider_name=provider_name,
|
||||||
|
all_models=all_models,
|
||||||
|
errors=errors,
|
||||||
|
has_success=has_success,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prepare_fetch_context(
|
||||||
|
self, key_id: str
|
||||||
|
) -> tuple[str, str, str, str, list[dict]] | str | None:
|
||||||
|
"""
|
||||||
|
准备获取模型所需的上下文数据
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
- tuple: (key_id, provider_id, provider_name, api_key_value, endpoint_configs)
|
||||||
|
- "skip": 跳过该 Key
|
||||||
|
- "error": 出错
|
||||||
|
- None: Key 不存在
|
||||||
|
"""
|
||||||
with create_session() as db:
|
with create_session() as db:
|
||||||
key = (
|
key = (
|
||||||
db.query(ProviderAPIKey)
|
db.query(ProviderAPIKey)
|
||||||
@@ -271,142 +316,154 @@ class ModelFetchScheduler:
|
|||||||
|
|
||||||
if not key:
|
if not key:
|
||||||
logger.warning(f"Key {key_id} 不存在,跳过")
|
logger.warning(f"Key {key_id} 不存在,跳过")
|
||||||
return "skip"
|
return None
|
||||||
|
|
||||||
if not key.is_active or not key.auto_fetch_models:
|
if not key.is_active or not key.auto_fetch_models:
|
||||||
logger.debug(f"Key {key_id} 已禁用或关闭自动获取,跳过")
|
logger.debug(f"Key {key_id} 已禁用或关闭自动获取,跳过")
|
||||||
return "skip"
|
return "skip"
|
||||||
|
|
||||||
try:
|
now = datetime.now(timezone.utc)
|
||||||
result = await self._fetch_models_for_key(db, key)
|
provider_id = key.provider_id
|
||||||
|
|
||||||
|
# 获取 Provider 和 Endpoints
|
||||||
|
provider = (
|
||||||
|
db.query(Provider)
|
||||||
|
.options(joinedload(Provider.endpoints))
|
||||||
|
.filter(Provider.id == provider_id)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
|
||||||
|
if not provider:
|
||||||
|
logger.warning(f"Provider {provider_id} 不存在,跳过 Key {key.id}")
|
||||||
|
key.last_models_fetch_error = "Provider not found"
|
||||||
|
key.last_models_fetch_at = now
|
||||||
db.commit()
|
db.commit()
|
||||||
return result
|
return "error"
|
||||||
|
|
||||||
|
# Vertex AI 类型不支持自动获取模型
|
||||||
|
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
|
||||||
|
if auth_type == "vertex_ai":
|
||||||
|
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key"
|
||||||
|
key.last_models_fetch_at = now
|
||||||
|
db.commit()
|
||||||
|
logger.info(f"Key {key.id} 为 Vertex AI 类型,跳过自动获取模型")
|
||||||
|
return "skip"
|
||||||
|
|
||||||
|
# 解密 API Key
|
||||||
|
if not key.api_key:
|
||||||
|
logger.warning(f"Key {key.id} 没有 API Key,跳过")
|
||||||
|
key.last_models_fetch_error = "No API key configured"
|
||||||
|
key.last_models_fetch_at = now
|
||||||
|
db.commit()
|
||||||
|
return "error"
|
||||||
|
|
||||||
|
try:
|
||||||
|
api_key_value = crypto_service.decrypt(key.api_key)
|
||||||
except Exception:
|
except Exception:
|
||||||
db.rollback()
|
logger.error(f"解密 Key {key.id} 失败")
|
||||||
raise
|
key.last_models_fetch_error = "Decrypt error"
|
||||||
|
key.last_models_fetch_at = now
|
||||||
|
db.commit()
|
||||||
|
return "error"
|
||||||
|
|
||||||
async def _fetch_models_for_key(
|
# 构建 api_format -> endpoint 映射
|
||||||
|
format_to_endpoint: dict[str, object] = {}
|
||||||
|
for endpoint in provider.endpoints: # type: ignore[attr-defined]
|
||||||
|
if endpoint.is_active:
|
||||||
|
format_to_endpoint[endpoint.api_format] = endpoint
|
||||||
|
|
||||||
|
if not format_to_endpoint:
|
||||||
|
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
|
||||||
|
key.last_models_fetch_error = "No active endpoints"
|
||||||
|
key.last_models_fetch_at = now
|
||||||
|
db.commit()
|
||||||
|
return "error"
|
||||||
|
|
||||||
|
# 使用公共函数构建所有格式的端点配置
|
||||||
|
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
return (key_id, provider_id, provider.name, api_key_value, endpoint_configs)
|
||||||
|
|
||||||
|
async def _update_key_after_fetch(
|
||||||
self,
|
self,
|
||||||
db: Session,
|
key_id: str,
|
||||||
key: ProviderAPIKey,
|
provider_id: str,
|
||||||
|
provider_name: str,
|
||||||
|
all_models: list[dict],
|
||||||
|
errors: list[str],
|
||||||
|
has_success: bool,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""为单个 Key 获取模型并更新 allowed_models,返回结果状态"""
|
"""
|
||||||
|
HTTP 请求完成后更新数据库
|
||||||
|
|
||||||
|
使用新的数据库连接来更新 Key 的 allowed_models
|
||||||
|
"""
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
provider_id = key.provider_id
|
|
||||||
|
|
||||||
# 获取 Provider 和 Endpoints
|
with create_session() as db:
|
||||||
provider = (
|
# 重新获取 Key(因为之前的连接已关闭)
|
||||||
db.query(Provider)
|
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||||
.options(joinedload(Provider.endpoints))
|
if not key:
|
||||||
.filter(Provider.id == provider_id)
|
logger.warning(f"Key {key_id} 在更新时不存在")
|
||||||
.first()
|
return "error"
|
||||||
)
|
|
||||||
|
|
||||||
if not provider:
|
# 记录获取时间
|
||||||
logger.warning(f"Provider {provider_id} 不存在,跳过 Key {key.id}")
|
|
||||||
key.last_models_fetch_error = "Provider not found"
|
|
||||||
key.last_models_fetch_at = now
|
key.last_models_fetch_at = now
|
||||||
return "error"
|
|
||||||
|
|
||||||
# Vertex AI 类型不支持自动获取模型(需要使用 Service Account 认证)
|
# 如果没有任何成功的响应,不更新 allowed_models(保留旧数据)
|
||||||
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
|
if not has_success:
|
||||||
if auth_type == "vertex_ai":
|
error_msg = "; ".join(errors) if errors else "All endpoints failed"
|
||||||
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key"
|
key.last_models_fetch_error = error_msg
|
||||||
key.last_models_fetch_at = now
|
logger.warning(
|
||||||
logger.info(f"Key {key.id} 为 Vertex AI 类型,跳过自动获取模型")
|
f"Provider {provider_name} Key {key.id} 所有端点获取失败,保留现有模型列表"
|
||||||
return "skip"
|
)
|
||||||
|
db.commit()
|
||||||
|
return "error"
|
||||||
|
|
||||||
# 解密 API Key
|
# 有成功的响应,清除错误状态
|
||||||
if not key.api_key:
|
key.last_models_fetch_error = None
|
||||||
logger.warning(f"Key {key.id} 没有 API Key,跳过")
|
|
||||||
key.last_models_fetch_error = "No API key configured"
|
|
||||||
key.last_models_fetch_at = now
|
|
||||||
return "error"
|
|
||||||
|
|
||||||
try:
|
# 去重获取模型 ID 列表
|
||||||
api_key_value = crypto_service.decrypt(key.api_key)
|
fetched_model_ids: set[str] = set()
|
||||||
except Exception:
|
for model in all_models:
|
||||||
# 不记录异常详情,避免泄露密钥信息
|
model_id = model.get("id")
|
||||||
logger.error(f"解密 Key {key.id} 失败")
|
if model_id:
|
||||||
key.last_models_fetch_error = "Decrypt error"
|
fetched_model_ids.add(model_id)
|
||||||
key.last_models_fetch_at = now
|
|
||||||
return "error"
|
|
||||||
|
|
||||||
# 构建 api_format -> endpoint 映射
|
logger.info(
|
||||||
format_to_endpoint: dict[str, object] = {}
|
f"Provider {provider_name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
|
||||||
for endpoint in provider.endpoints: # type: ignore[attr-defined]
|
|
||||||
if endpoint.is_active:
|
|
||||||
format_to_endpoint[endpoint.api_format] = endpoint
|
|
||||||
|
|
||||||
if not format_to_endpoint:
|
|
||||||
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
|
|
||||||
key.last_models_fetch_error = "No active endpoints"
|
|
||||||
key.last_models_fetch_at = now
|
|
||||||
return "error"
|
|
||||||
|
|
||||||
# 使用公共函数构建所有格式的端点配置
|
|
||||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
|
|
||||||
|
|
||||||
# 并发获取模型
|
|
||||||
all_models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
|
|
||||||
|
|
||||||
# 记录获取时间
|
|
||||||
key.last_models_fetch_at = now
|
|
||||||
|
|
||||||
# 如果没有任何成功的响应,不更新 allowed_models(保留旧数据)
|
|
||||||
if not has_success:
|
|
||||||
# 所有端点都失败时,记录错误
|
|
||||||
error_msg = "; ".join(errors) if errors else "All endpoints failed"
|
|
||||||
key.last_models_fetch_error = error_msg
|
|
||||||
logger.warning(
|
|
||||||
f"Provider {provider.name} Key {key.id} 所有端点获取失败,保留现有模型列表"
|
|
||||||
)
|
|
||||||
return "error"
|
|
||||||
|
|
||||||
# 有成功的响应,清除错误状态(部分失败不算失败)
|
|
||||||
key.last_models_fetch_error = None
|
|
||||||
|
|
||||||
# 去重获取模型 ID 列表
|
|
||||||
fetched_model_ids: set[str] = set()
|
|
||||||
for model in all_models:
|
|
||||||
model_id = model.get("id")
|
|
||||||
if model_id:
|
|
||||||
fetched_model_ids.add(model_id)
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
f"Provider {provider.name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
|
|
||||||
seen_keys: set[str] = set()
|
|
||||||
unique_models: list[dict] = []
|
|
||||||
for model in all_models:
|
|
||||||
model_id = model.get("id")
|
|
||||||
api_format = model.get("api_format", "")
|
|
||||||
unique_key = f"{model_id}:{api_format}"
|
|
||||||
if model_id and unique_key not in seen_keys:
|
|
||||||
seen_keys.add(unique_key)
|
|
||||||
unique_models.append(model)
|
|
||||||
await set_upstream_models_to_cache(
|
|
||||||
provider_id, # type: ignore[arg-type]
|
|
||||||
key.id, # type: ignore[arg-type]
|
|
||||||
unique_models,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 更新 allowed_models(保留 locked_models)
|
|
||||||
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
|
|
||||||
|
|
||||||
# 如果白名单有变化,触发缓存失效和自动关联检查
|
|
||||||
if has_changed and provider_id:
|
|
||||||
from src.services.model.global_model import on_key_allowed_models_changed
|
|
||||||
|
|
||||||
await on_key_allowed_models_changed(
|
|
||||||
db=db,
|
|
||||||
provider_id=provider_id,
|
|
||||||
allowed_models=list(key.allowed_models or []),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return "success"
|
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
|
||||||
|
seen_keys: set[str] = set()
|
||||||
|
unique_models: list[dict] = []
|
||||||
|
for model in all_models:
|
||||||
|
model_id = model.get("id")
|
||||||
|
api_format = model.get("api_format", "")
|
||||||
|
unique_key = f"{model_id}:{api_format}"
|
||||||
|
if model_id and unique_key not in seen_keys:
|
||||||
|
seen_keys.add(unique_key)
|
||||||
|
unique_models.append(model)
|
||||||
|
await set_upstream_models_to_cache(provider_id, key.id, unique_models)
|
||||||
|
|
||||||
|
# 更新 allowed_models(保留 locked_models)
|
||||||
|
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
|
||||||
|
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
# 如果白名单有变化,触发缓存失效和自动关联检查
|
||||||
|
if has_changed and provider_id:
|
||||||
|
from src.services.model.global_model import on_key_allowed_models_changed
|
||||||
|
|
||||||
|
# 使用新会话处理后续操作
|
||||||
|
with create_session() as db2:
|
||||||
|
await on_key_allowed_models_changed(
|
||||||
|
db=db2,
|
||||||
|
provider_id=provider_id,
|
||||||
|
allowed_models=list(key.allowed_models or []),
|
||||||
|
)
|
||||||
|
|
||||||
|
return "success"
|
||||||
|
|
||||||
def _update_key_allowed_models(self, key: ProviderAPIKey, fetched_model_ids: set[str]) -> bool:
|
def _update_key_allowed_models(self, key: ProviderAPIKey, fetched_model_ids: set[str]) -> bool:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -156,6 +156,8 @@ class CandidateResolver:
|
|||||||
user_id: str,
|
user_id: str,
|
||||||
user_api_key: ApiKey,
|
user_api_key: ApiKey,
|
||||||
required_capabilities: dict[str, bool] | None = None,
|
required_capabilities: dict[str, bool] | None = None,
|
||||||
|
*,
|
||||||
|
expand_retries: bool = True,
|
||||||
) -> dict[tuple[int, int], str]:
|
) -> dict[tuple[int, int], str]:
|
||||||
"""
|
"""
|
||||||
为所有候选预先创建 available 状态记录(批量插入优化)
|
为所有候选预先创建 available 状态记录(批量插入优化)
|
||||||
@@ -211,9 +213,12 @@ class CandidateResolver:
|
|||||||
candidate_record_map[(candidate_index, 0)] = record_id
|
candidate_record_map[(candidate_index, 0)] = record_id
|
||||||
else:
|
else:
|
||||||
# max_retries 已从 Endpoint 迁移到 Provider(Endpoint 仍可能保留旧字段用于兼容)
|
# max_retries 已从 Endpoint 迁移到 Provider(Endpoint 仍可能保留旧字段用于兼容)
|
||||||
max_retries_for_candidate = (
|
if not expand_retries:
|
||||||
int(provider.max_retries or 2) if candidate.is_cached else 1
|
max_retries_for_candidate = 1
|
||||||
)
|
else:
|
||||||
|
max_retries_for_candidate = (
|
||||||
|
int(provider.max_retries or 2) if candidate.is_cached else 1
|
||||||
|
)
|
||||||
|
|
||||||
for retry_index in range(max_retries_for_candidate):
|
for retry_index in range(max_retries_for_candidate):
|
||||||
record_id = str(uuid.uuid4())
|
record_id = str(uuid.uuid4())
|
||||||
|
|||||||
@@ -38,6 +38,19 @@ BALANCE_CACHE_TTL = 86400
|
|||||||
# 认证失败缓存 TTL(60 秒,避免频繁重试但允许用户修正后快速重试)
|
# 认证失败缓存 TTL(60 秒,避免频繁重试但允许用户修正后快速重试)
|
||||||
AUTH_FAILED_CACHE_TTL = 60
|
AUTH_FAILED_CACHE_TTL = 60
|
||||||
|
|
||||||
|
# 后台余额刷新并发限制(避免启动时耗尽连接池)
|
||||||
|
# 使用较小的值(3)确保不会对连接池造成过大压力
|
||||||
|
_balance_refresh_semaphore: asyncio.Semaphore | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_balance_refresh_semaphore() -> asyncio.Semaphore:
|
||||||
|
"""获取余额刷新信号量(延迟初始化)"""
|
||||||
|
global _balance_refresh_semaphore
|
||||||
|
if _balance_refresh_semaphore is None:
|
||||||
|
# 限制为 3 个并发,确保后台任务不会占用太多连接
|
||||||
|
_balance_refresh_semaphore = asyncio.Semaphore(3)
|
||||||
|
return _balance_refresh_semaphore
|
||||||
|
|
||||||
|
|
||||||
def _get_batch_balance_concurrency() -> int:
|
def _get_batch_balance_concurrency() -> int:
|
||||||
"""
|
"""
|
||||||
@@ -98,6 +111,44 @@ class ProviderOpsService:
|
|||||||
# 连接器缓存 {provider_id: ProviderConnector}
|
# 连接器缓存 {provider_id: ProviderConnector}
|
||||||
self._connectors: dict[str, ProviderConnector] = {}
|
self._connectors: dict[str, ProviderConnector] = {}
|
||||||
|
|
||||||
|
def _release_db_connection_before_await(self) -> None:
|
||||||
|
"""
|
||||||
|
Release pooled DB connection before long awaits (network/Redis).
|
||||||
|
|
||||||
|
SQLAlchemy Session will keep a connection checked out while a transaction is open,
|
||||||
|
even for read-only queries. In async code, this can exhaust the pool if we `await`
|
||||||
|
network I/O while holding that transaction.
|
||||||
|
|
||||||
|
Safety:
|
||||||
|
- Only commits when the session has no pending changes (new/dirty/deleted).
|
||||||
|
- Temporarily disables expire_on_commit to avoid unexpected lazy reloads.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
has_pending_changes = bool(self.db.new) or bool(self.db.dirty) or bool(self.db.deleted)
|
||||||
|
except Exception:
|
||||||
|
has_pending_changes = False
|
||||||
|
|
||||||
|
if has_pending_changes:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
if not self.db.in_transaction():
|
||||||
|
return
|
||||||
|
except Exception:
|
||||||
|
return
|
||||||
|
|
||||||
|
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
|
||||||
|
self.db.expire_on_commit = False
|
||||||
|
try:
|
||||||
|
self.db.commit()
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
self.db.rollback()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
self.db.expire_on_commit = original_expire_on_commit
|
||||||
|
|
||||||
# ==================== 配置管理 ====================
|
# ==================== 配置管理 ====================
|
||||||
|
|
||||||
def get_config(self, provider_id: str) -> ProviderOpsConfig | None:
|
def get_config(self, provider_id: str) -> ProviderOpsConfig | None:
|
||||||
@@ -251,6 +302,9 @@ class ProviderOpsService:
|
|||||||
if not actual_credentials:
|
if not actual_credentials:
|
||||||
return False, "未提供凭据"
|
return False, "未提供凭据"
|
||||||
|
|
||||||
|
# Avoid holding a DB connection while awaiting network I/O.
|
||||||
|
self._release_db_connection_before_await()
|
||||||
|
|
||||||
# 建立连接
|
# 建立连接
|
||||||
logger.info(
|
logger.info(
|
||||||
f"尝试连接: provider_id={provider_id}, "
|
f"尝试连接: provider_id={provider_id}, "
|
||||||
@@ -333,6 +387,8 @@ class ProviderOpsService:
|
|||||||
)
|
)
|
||||||
connector = self._connectors.get(provider_id)
|
connector = self._connectors.get(provider_id)
|
||||||
|
|
||||||
|
# Avoid holding a DB connection while awaiting authentication checks.
|
||||||
|
self._release_db_connection_before_await()
|
||||||
if not connector or not await connector.is_authenticated():
|
if not connector or not await connector.is_authenticated():
|
||||||
return ActionResult(
|
return ActionResult(
|
||||||
status=ActionStatus.AUTH_EXPIRED,
|
status=ActionStatus.AUTH_EXPIRED,
|
||||||
@@ -373,6 +429,9 @@ class ProviderOpsService:
|
|||||||
# 创建操作实例
|
# 创建操作实例
|
||||||
action = architecture.get_action(action_type, merged_config)
|
action = architecture.get_action(action_type, merged_config)
|
||||||
|
|
||||||
|
# Avoid holding a DB connection while awaiting the upstream action.
|
||||||
|
self._release_db_connection_before_await()
|
||||||
|
|
||||||
# 执行操作
|
# 执行操作
|
||||||
async with connector.get_client() as client:
|
async with connector.get_client() as client:
|
||||||
result = await action.execute(client)
|
result = await action.execute(client)
|
||||||
@@ -422,6 +481,9 @@ class ProviderOpsService:
|
|||||||
Returns:
|
Returns:
|
||||||
操作结果(可能是缓存的)
|
操作结果(可能是缓存的)
|
||||||
"""
|
"""
|
||||||
|
# Avoid holding a DB connection while awaiting Redis/cache I/O.
|
||||||
|
self._release_db_connection_before_await()
|
||||||
|
|
||||||
# 尝试从缓存获取
|
# 尝试从缓存获取
|
||||||
cached = await self._get_cached_balance(provider_id)
|
cached = await self._get_cached_balance(provider_id)
|
||||||
|
|
||||||
@@ -453,7 +515,20 @@ class ProviderOpsService:
|
|||||||
|
|
||||||
注意:这是一个后台任务,使用独立的短生命周期 session,
|
注意:这是一个后台任务,使用独立的短生命周期 session,
|
||||||
避免长时间占用连接池资源。
|
避免长时间占用连接池资源。
|
||||||
|
|
||||||
|
使用信号量限制并发数,避免启动时多个刷新任务同时运行导致连接池耗尽。
|
||||||
"""
|
"""
|
||||||
|
semaphore = _get_balance_refresh_semaphore()
|
||||||
|
|
||||||
|
# 尝试获取信号量,如果无法立即获取则跳过本次刷新
|
||||||
|
# 这样可以避免在连接池紧张时阻塞
|
||||||
|
try:
|
||||||
|
# 使用 wait_for 设置超时,避免无限等待
|
||||||
|
await asyncio.wait_for(semaphore.acquire(), timeout=5.0)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.debug(f"异步刷新余额跳过(并发限制): provider_id={provider_id}")
|
||||||
|
return
|
||||||
|
|
||||||
db = None
|
db = None
|
||||||
try:
|
try:
|
||||||
# 后台任务需要创建独立的 session,因为原请求的 session 可能已关闭
|
# 后台任务需要创建独立的 session,因为原请求的 session 可能已关闭
|
||||||
@@ -469,6 +544,8 @@ class ProviderOpsService:
|
|||||||
db.close()
|
db.close()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
# 释放信号量
|
||||||
|
semaphore.release()
|
||||||
|
|
||||||
async def _clear_balance_cache(self, provider_id: str) -> None:
|
async def _clear_balance_cache(self, provider_id: str) -> None:
|
||||||
"""清除余额缓存"""
|
"""清除余额缓存"""
|
||||||
@@ -762,6 +839,9 @@ class ProviderOpsService:
|
|||||||
if not provider_ids:
|
if not provider_ids:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
# Release the DB connection before awaiting many async cache refreshes.
|
||||||
|
self._release_db_connection_before_await()
|
||||||
|
|
||||||
# 使用信号量限制并发数,避免同时发起过多请求耗尽连接池
|
# 使用信号量限制并发数,避免同时发起过多请求耗尽连接池
|
||||||
concurrency = _get_batch_balance_concurrency()
|
concurrency = _get_batch_balance_concurrency()
|
||||||
semaphore = asyncio.Semaphore(concurrency)
|
semaphore = asyncio.Semaphore(concurrency)
|
||||||
@@ -822,6 +902,9 @@ class ProviderOpsService:
|
|||||||
|
|
||||||
from src.utils.ssl_utils import get_ssl_context
|
from src.utils.ssl_utils import get_ssl_context
|
||||||
|
|
||||||
|
# Avoid holding a DB connection while awaiting verify pre-processing / network.
|
||||||
|
self._release_db_connection_before_await()
|
||||||
|
|
||||||
# 移除 base_url 末尾的斜杠
|
# 移除 base_url 末尾的斜杠
|
||||||
base_url = base_url.rstrip("/")
|
base_url = base_url.rstrip("/")
|
||||||
|
|
||||||
|
|||||||
@@ -136,6 +136,11 @@ class SystemConfigService:
|
|||||||
"value": [],
|
"value": [],
|
||||||
"description": "邮箱后缀列表,配合 email_suffix_mode 使用",
|
"description": "邮箱后缀列表,配合 email_suffix_mode 使用",
|
||||||
},
|
},
|
||||||
|
# 格式转换开关
|
||||||
|
"enable_format_conversion": {
|
||||||
|
"value": False,
|
||||||
|
"description": "全局格式转换开关:开启时强制允许所有提供商的格式转换;关闭时由各提供商自行决定",
|
||||||
|
},
|
||||||
"audit_log_retention_days": {
|
"audit_log_retention_days": {
|
||||||
"value": 30,
|
"value": 30,
|
||||||
"description": "审计日志保留天数,超过此天数的审计日志将被自动清理",
|
"description": "审计日志保留天数,超过此天数的审计日志将被自动清理",
|
||||||
@@ -358,6 +363,11 @@ class SystemConfigService:
|
|||||||
"""获取敏感请求头列表"""
|
"""获取敏感请求头列表"""
|
||||||
return cls.get_config(db, "sensitive_headers", [])
|
return cls.get_config(db, "sensitive_headers", [])
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_format_conversion_enabled(cls, db: Session) -> bool:
|
||||||
|
"""检查全局格式转换是否启用"""
|
||||||
|
return bool(cls.get_config(db, "enable_format_conversion", True))
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def mask_sensitive_headers(cls, db: Session, headers: dict[str, Any]) -> dict[str, Any]:
|
def mask_sensitive_headers(cls, db: Session, headers: dict[str, Any]) -> dict[str, Any]:
|
||||||
"""脱敏敏感请求头"""
|
"""脱敏敏感请求头"""
|
||||||
|
|||||||
@@ -8,6 +8,7 @@
|
|||||||
- 审计日志清理:定期清理过期的审计日志
|
- 审计日志清理:定期清理过期的审计日志
|
||||||
- 连接池监控:定期检查数据库连接池状态
|
- 连接池监控:定期检查数据库连接池状态
|
||||||
- Pending 状态清理:清理异常的 Pending 状态记录
|
- Pending 状态清理:清理异常的 Pending 状态记录
|
||||||
|
- Gemini 文件映射清理:清理过期的 Gemini 文件→Key 映射
|
||||||
|
|
||||||
使用 APScheduler 进行任务调度,支持时区配置。
|
使用 APScheduler 进行任务调度,支持时区配置。
|
||||||
"""
|
"""
|
||||||
@@ -103,6 +104,14 @@ class MaintenanceScheduler:
|
|||||||
name="审计日志清理",
|
name="审计日志清理",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Gemini 文件映射清理 - 每小时执行
|
||||||
|
scheduler.add_interval_job(
|
||||||
|
self._scheduled_gemini_file_mapping_cleanup,
|
||||||
|
hours=1,
|
||||||
|
job_id="gemini_file_mapping_cleanup",
|
||||||
|
name="Gemini文件映射清理",
|
||||||
|
)
|
||||||
|
|
||||||
# Provider 签到任务 - 凌晨 1:05 执行
|
# Provider 签到任务 - 凌晨 1:05 执行
|
||||||
scheduler.add_cron_job(
|
scheduler.add_cron_job(
|
||||||
self._scheduled_provider_checkin,
|
self._scheduled_provider_checkin,
|
||||||
@@ -117,8 +126,9 @@ class MaintenanceScheduler:
|
|||||||
|
|
||||||
async def _run_startup_tasks(self) -> None:
|
async def _run_startup_tasks(self) -> None:
|
||||||
"""启动时执行的初始化任务"""
|
"""启动时执行的初始化任务"""
|
||||||
# 延迟一点执行,确保系统完全启动
|
# 延迟执行,等待系统完全启动(Redis 连接、其他后台任务稳定)
|
||||||
await asyncio.sleep(2)
|
# 增加延迟时间避免与 UsageQueueConsumer 等后台任务竞争数据库连接
|
||||||
|
await asyncio.sleep(10)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
logger.info("启动时执行首次清理任务...")
|
logger.info("启动时执行首次清理任务...")
|
||||||
@@ -170,6 +180,10 @@ class MaintenanceScheduler:
|
|||||||
"""审计日志清理任务(定时调用)"""
|
"""审计日志清理任务(定时调用)"""
|
||||||
await self._perform_audit_cleanup()
|
await self._perform_audit_cleanup()
|
||||||
|
|
||||||
|
async def _scheduled_gemini_file_mapping_cleanup(self) -> None:
|
||||||
|
"""Gemini 文件映射清理任务(定时调用)"""
|
||||||
|
await self._perform_gemini_file_mapping_cleanup()
|
||||||
|
|
||||||
async def _scheduled_provider_checkin(self) -> None:
|
async def _scheduled_provider_checkin(self) -> None:
|
||||||
"""Provider 签到任务(定时调用)"""
|
"""Provider 签到任务(定时调用)"""
|
||||||
await self._perform_provider_checkin()
|
await self._perform_provider_checkin()
|
||||||
@@ -483,6 +497,26 @@ class MaintenanceScheduler:
|
|||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
async def _perform_gemini_file_mapping_cleanup(self) -> None:
|
||||||
|
"""清理过期的 Gemini 文件映射记录"""
|
||||||
|
db = create_session()
|
||||||
|
try:
|
||||||
|
from src.services.gemini_files_mapping import cleanup_expired_mappings
|
||||||
|
|
||||||
|
deleted_count = cleanup_expired_mappings(db)
|
||||||
|
|
||||||
|
if deleted_count > 0:
|
||||||
|
logger.info(f"清理了 {deleted_count} 条过期的 Gemini 文件映射")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"Gemini 文件映射清理失败: {e}")
|
||||||
|
try:
|
||||||
|
db.rollback()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
async def _perform_provider_checkin(self) -> None:
|
async def _perform_provider_checkin(self) -> None:
|
||||||
"""执行 Provider 签到任务
|
"""执行 Provider 签到任务
|
||||||
|
|
||||||
@@ -508,8 +542,21 @@ class MaintenanceScheduler:
|
|||||||
|
|
||||||
logger.info(f"开始执行 Provider 签到,共 {len(provider_ids)} 个...")
|
logger.info(f"开始执行 Provider 签到,共 {len(provider_ids)} 个...")
|
||||||
|
|
||||||
# 创建 ProviderOpsService 并执行批量余额查询(会触发签到)
|
# 释放主 session 的连接,避免在整个签到期间占用连接池
|
||||||
service = ProviderOpsService(db)
|
# (后续每个 provider 将使用独立短生命周期 session)
|
||||||
|
try:
|
||||||
|
if db.in_transaction():
|
||||||
|
db.commit()
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
db.rollback()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
db.close()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
db = None
|
||||||
|
|
||||||
# 使用信号量限制并发,避免同时发起过多请求
|
# 使用信号量限制并发,避免同时发起过多请求
|
||||||
concurrency = 3 # 签到任务并发数
|
concurrency = 3 # 签到任务并发数
|
||||||
@@ -518,7 +565,9 @@ class MaintenanceScheduler:
|
|||||||
async def _checkin_provider(provider_id: str) -> tuple[str, bool, str]:
|
async def _checkin_provider(provider_id: str) -> tuple[str, bool, str]:
|
||||||
"""执行单个 Provider 的签到"""
|
"""执行单个 Provider 的签到"""
|
||||||
async with semaphore:
|
async with semaphore:
|
||||||
|
task_db = create_session()
|
||||||
try:
|
try:
|
||||||
|
service = ProviderOpsService(task_db)
|
||||||
# 触发余额查询(会先执行签到)
|
# 触发余额查询(会先执行签到)
|
||||||
result = await service.query_balance(provider_id)
|
result = await service.query_balance(provider_id)
|
||||||
# 检查签到结果
|
# 检查签到结果
|
||||||
@@ -537,6 +586,11 @@ class MaintenanceScheduler:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Provider {provider_id} 签到失败: {e}")
|
logger.warning(f"Provider {provider_id} 签到失败: {e}")
|
||||||
return provider_id, False, str(e)
|
return provider_id, False, str(e)
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
task_db.close()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
# 并行执行签到
|
# 并行执行签到
|
||||||
tasks = [_checkin_provider(pid) for pid in provider_ids]
|
tasks = [_checkin_provider(pid) for pid in provider_ids]
|
||||||
@@ -556,7 +610,8 @@ class MaintenanceScheduler:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Provider 签到任务执行失败: {e}")
|
logger.exception(f"Provider 签到任务执行失败: {e}")
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
if db is not None:
|
||||||
|
db.close()
|
||||||
|
|
||||||
async def _perform_cleanup(self) -> None:
|
async def _perform_cleanup(self) -> None:
|
||||||
"""执行清理任务"""
|
"""执行清理任务"""
|
||||||
|
|||||||
@@ -1,25 +1,13 @@
|
|||||||
"""
|
"""
|
||||||
异步任务服务层
|
任务服务层(Phase2)
|
||||||
|
|
||||||
提供视频/图片/音频等异步任务的:
|
统一任务框架相关的应用层入口:
|
||||||
- 提交阶段故障转移(AsyncTaskOrchestrator)
|
- 候选提交阶段:`services.candidate.CandidateService`
|
||||||
- 终态计费与 Usage 写入(VideoTelemetry 等)
|
- 终态结算:`services.task.application.TaskApplicationService`
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from .orchestrator import (
|
from .application import TaskApplicationService
|
||||||
AllCandidatesFailedError,
|
|
||||||
AsyncTaskOrchestrator,
|
|
||||||
CandidateSubmissionError,
|
|
||||||
CandidateUnsupportedError,
|
|
||||||
SubmitOutcome,
|
|
||||||
UpstreamClientRequestError,
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AsyncTaskOrchestrator",
|
"TaskApplicationService",
|
||||||
"SubmitOutcome",
|
|
||||||
"AllCandidatesFailedError",
|
|
||||||
"UpstreamClientRequestError",
|
|
||||||
"CandidateUnsupportedError",
|
|
||||||
"CandidateSubmissionError",
|
|
||||||
]
|
]
|
||||||
|
|||||||
312
src/services/task/application.py
Normal file
312
src/services/task/application.py
Normal file
@@ -0,0 +1,312 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.config.settings import config
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.models.database import ApiKey, Provider, Usage, 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.usage.service import UsageService
|
||||||
|
|
||||||
|
|
||||||
|
class TaskApplicationService:
|
||||||
|
"""
|
||||||
|
TaskApplicationService (Phase2)
|
||||||
|
|
||||||
|
当前仅先收敛"终态结算"入口,用于替代旧版 VideoTelemetry 直写 Usage 的流程。
|
||||||
|
后续将扩展 submit/cancel 并迁移候选编排逻辑。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
|
||||||
|
self.db = db
|
||||||
|
self.redis = redis_client
|
||||||
|
|
||||||
|
async def finalize_video_task(self, task: VideoTask) -> bool:
|
||||||
|
"""
|
||||||
|
更新视频任务的计费信息(轮询完成后调用)。
|
||||||
|
|
||||||
|
异步任务的计费流程:
|
||||||
|
1. 提交成功时:Usage 已结算(billing_status='settled',费用=0)
|
||||||
|
2. 轮询完成时:更新实际费用(成功则计费,失败则保持0)
|
||||||
|
|
||||||
|
返回 True 表示成功更新,False 表示无需更新(如已是最终状态)
|
||||||
|
"""
|
||||||
|
request_id = getattr(task, "request_id", None) or task.id
|
||||||
|
|
||||||
|
existing = self.db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||||
|
if not existing:
|
||||||
|
# Usage 不存在,尝试创建并结算(兜底逻辑)
|
||||||
|
logger.warning(
|
||||||
|
"Usage not found for video task, creating fallback: task_id=%s request_id=%s",
|
||||||
|
task.id,
|
||||||
|
request_id,
|
||||||
|
)
|
||||||
|
return await self._create_fallback_usage(task, request_id)
|
||||||
|
|
||||||
|
# 检查是否已有计费更新标记(避免重复计费)
|
||||||
|
metadata = existing.request_metadata or {}
|
||||||
|
if metadata.get("billing_updated_at"):
|
||||||
|
logger.debug(
|
||||||
|
"Video task billing already updated: task_id=%s request_id=%s",
|
||||||
|
task.id,
|
||||||
|
request_id,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 计算异步任务总耗时(ms)
|
||||||
|
response_time_ms: int | None = None
|
||||||
|
if task.submitted_at and task.completed_at:
|
||||||
|
delta = task.completed_at - task.submitted_at
|
||||||
|
response_time_ms = int(delta.total_seconds() * 1000)
|
||||||
|
|
||||||
|
# === 收集计费维度 ===
|
||||||
|
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,
|
||||||
|
}
|
||||||
|
|
||||||
|
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 [],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
dims = DimensionCollectorService(self.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,
|
||||||
|
)
|
||||||
|
|
||||||
|
# === 计算成本(优先使用冻结的 billing_rule_snapshot)===
|
||||||
|
rule_snapshot = None
|
||||||
|
if isinstance(task.request_metadata, dict):
|
||||||
|
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
|
||||||
|
|
||||||
|
expression = None
|
||||||
|
variables: dict[str, Any] | None = None
|
||||||
|
dimension_mappings: dict[str, dict[str, Any]] | None = 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") or {}
|
||||||
|
dimension_mappings = rule_snapshot.get("dimension_mappings") or {}
|
||||||
|
else:
|
||||||
|
lookup = BillingRuleService.find_rule(
|
||||||
|
self.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 = getattr(lookup, "scope", None)
|
||||||
|
expression = rule.expression
|
||||||
|
variables = rule.variables or {}
|
||||||
|
dimension_mappings = rule.dimension_mappings or {}
|
||||||
|
|
||||||
|
billing_snapshot: dict[str, Any] = {
|
||||||
|
"schema_version": "1.0",
|
||||||
|
"rule_id": str(rule_id) if rule_id else None,
|
||||||
|
"rule_name": str(rule_name) if rule_name else None,
|
||||||
|
"scope": str(rule_scope) if rule_scope else None,
|
||||||
|
"expression": str(expression) if expression else None,
|
||||||
|
"dimensions_used": dims,
|
||||||
|
"missing_required": [],
|
||||||
|
"cost": 0.0,
|
||||||
|
"status": "no_rule",
|
||||||
|
"calculated_at": datetime.now(timezone.utc).isoformat(),
|
||||||
|
}
|
||||||
|
|
||||||
|
cost = 0.0
|
||||||
|
# 只有任务成功时才计费
|
||||||
|
if task.status == "completed" and expression:
|
||||||
|
engine = FormulaEngine()
|
||||||
|
try:
|
||||||
|
result = engine.evaluate(
|
||||||
|
expression=str(expression),
|
||||||
|
variables=variables,
|
||||||
|
dimensions=dims,
|
||||||
|
dimension_mappings=dimension_mappings,
|
||||||
|
strict_mode=config.billing_strict_mode,
|
||||||
|
)
|
||||||
|
billing_snapshot["status"] = result.status
|
||||||
|
billing_snapshot["missing_required"] = result.missing_required
|
||||||
|
if result.status == "complete":
|
||||||
|
cost = float(result.cost)
|
||||||
|
billing_snapshot["cost"] = cost
|
||||||
|
except BillingIncompleteError as exc:
|
||||||
|
# strict_mode=true:标记任务失败并隐藏产物,避免"免费放行"
|
||||||
|
task.status = "failed"
|
||||||
|
task.error_code = "billing_incomplete"
|
||||||
|
task.error_message = f"Missing required dimensions: {exc.missing_required}"
|
||||||
|
task.video_url = None
|
||||||
|
task.video_urls = None
|
||||||
|
billing_snapshot["status"] = "incomplete"
|
||||||
|
billing_snapshot["missing_required"] = exc.missing_required
|
||||||
|
billing_snapshot["cost"] = 0.0
|
||||||
|
except Exception as exc:
|
||||||
|
billing_snapshot["status"] = "incomplete"
|
||||||
|
billing_snapshot["error"] = str(exc)
|
||||||
|
billing_snapshot["cost"] = 0.0
|
||||||
|
|
||||||
|
# 回写到 task.request_metadata 便于审计/重算
|
||||||
|
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
|
||||||
|
metadata = dict(task.request_metadata) if task.request_metadata else {}
|
||||||
|
metadata["billing_snapshot"] = billing_snapshot
|
||||||
|
task.request_metadata = metadata
|
||||||
|
|
||||||
|
# === 更新已结算的 Usage 计费信息 ===
|
||||||
|
updated = UsageService.update_settled_billing(
|
||||||
|
self.db,
|
||||||
|
request_id=request_id,
|
||||||
|
total_cost_usd=cost,
|
||||||
|
request_cost_usd=cost,
|
||||||
|
status="completed" if task.status == "completed" else "failed",
|
||||||
|
status_code=200 if task.status == "completed" else 500,
|
||||||
|
error_message=(
|
||||||
|
None
|
||||||
|
if task.status == "completed"
|
||||||
|
else (task.error_message or task.error_code or "video_task_failed")
|
||||||
|
),
|
||||||
|
response_time_ms=response_time_ms,
|
||||||
|
billing_snapshot=billing_snapshot,
|
||||||
|
extra_metadata={
|
||||||
|
"dimensions": dims,
|
||||||
|
"raw_response_ref": {
|
||||||
|
"video_task_id": task.id,
|
||||||
|
"field": "video_tasks.request_metadata.poll_raw_response",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
if updated:
|
||||||
|
logger.debug(
|
||||||
|
"Updated video task billing: task_id=%s request_id=%s cost=%.6f",
|
||||||
|
task.id,
|
||||||
|
request_id,
|
||||||
|
cost,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to update video task billing (may already be updated): "
|
||||||
|
"task_id=%s request_id=%s",
|
||||||
|
task.id,
|
||||||
|
request_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
return updated
|
||||||
|
|
||||||
|
async def _create_fallback_usage(self, task: VideoTask, request_id: str) -> bool:
|
||||||
|
"""
|
||||||
|
兜底逻辑:当 Usage 不存在时创建完整记录。
|
||||||
|
这种情况理论上不应发生(submit 阶段已创建),但保留以防万一。
|
||||||
|
"""
|
||||||
|
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
|
||||||
|
api_key_obj = (
|
||||||
|
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
|
||||||
|
if task.api_key_id
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
provider_obj = (
|
||||||
|
self.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"
|
||||||
|
|
||||||
|
# 计算响应时间
|
||||||
|
response_time_ms: int | None = None
|
||||||
|
if task.submitted_at and task.completed_at:
|
||||||
|
delta = task.completed_at - task.submitted_at
|
||||||
|
response_time_ms = int(delta.total_seconds() * 1000)
|
||||||
|
|
||||||
|
try:
|
||||||
|
await UsageService.record_usage_with_custom_cost(
|
||||||
|
db=self.db,
|
||||||
|
user=user_obj,
|
||||||
|
api_key=api_key_obj,
|
||||||
|
provider=provider_name,
|
||||||
|
model=task.model,
|
||||||
|
request_type="video",
|
||||||
|
total_cost_usd=0.0, # 兜底记录不计费
|
||||||
|
request_cost_usd=0.0,
|
||||||
|
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 == "completed" else 500,
|
||||||
|
error_message=(
|
||||||
|
None
|
||||||
|
if task.status == "completed"
|
||||||
|
else (task.error_message or task.error_code or "video_task_failed")
|
||||||
|
),
|
||||||
|
metadata={
|
||||||
|
"fallback_created": True,
|
||||||
|
"video_task_id": task.id,
|
||||||
|
},
|
||||||
|
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 == "completed" else "failed",
|
||||||
|
target_model=None,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception(
|
||||||
|
"Failed to create fallback usage for video task=%s: %s",
|
||||||
|
task.id,
|
||||||
|
str(exc),
|
||||||
|
)
|
||||||
|
return False
|
||||||
36
src/services/task/context.py
Normal file
36
src/services/task/context.py
Normal file
@@ -0,0 +1,36 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class TaskMode(str, Enum):
|
||||||
|
SYNC = "sync"
|
||||||
|
ASYNC = "async"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class TaskContext:
|
||||||
|
"""
|
||||||
|
TaskContext (pure DTO)
|
||||||
|
|
||||||
|
- Only primitive types / IDs
|
||||||
|
- Serializable & safe to pass across processes
|
||||||
|
"""
|
||||||
|
|
||||||
|
request_id: str
|
||||||
|
task_type: str # chat/cli/video/image/audio
|
||||||
|
task_mode: TaskMode
|
||||||
|
|
||||||
|
user_id: str
|
||||||
|
api_key_id: str
|
||||||
|
|
||||||
|
client_ip: str = ""
|
||||||
|
user_agent: str = ""
|
||||||
|
start_time: float = 0.0
|
||||||
|
|
||||||
|
api_format: str | None = None
|
||||||
|
model: str | None = None
|
||||||
|
mapped_model: str | None = None
|
||||||
|
|
||||||
|
capability_requirements: dict[str, bool] = field(default_factory=dict)
|
||||||
@@ -1,3 +1,7 @@
|
|||||||
"""Task telemetry implementations for concrete task types (video/image/audio)."""
|
"""Per-task-type implementations (Phase2).
|
||||||
|
|
||||||
|
Currently includes:
|
||||||
|
- video: polling adapter
|
||||||
|
"""
|
||||||
|
|
||||||
__all__ = []
|
__all__ = []
|
||||||
|
|||||||
626
src/services/task/impl/video_poller.py
Normal file
626
src/services/task/impl/video_poller.py
Normal file
@@ -0,0 +1,626 @@
|
|||||||
|
"""
|
||||||
|
Video task poller adapter.
|
||||||
|
|
||||||
|
Implements the video-specific poll/normalize/update logic used by TaskPollerService.
|
||||||
|
|
||||||
|
优化:HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
|
||||||
|
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
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 (
|
||||||
|
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 (
|
||||||
|
build_upstream_headers_for_endpoint,
|
||||||
|
get_extra_headers_from_endpoint,
|
||||||
|
make_signature_key,
|
||||||
|
)
|
||||||
|
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
||||||
|
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||||
|
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.services.task.application import TaskApplicationService
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class VideoPollContext:
|
||||||
|
"""视频轮询上下文,保存 HTTP 请求所需的数据(不依赖数据库会话)"""
|
||||||
|
|
||||||
|
task_id: str
|
||||||
|
external_task_id: str
|
||||||
|
provider_api_format: str
|
||||||
|
base_url: str
|
||||||
|
upstream_key: str
|
||||||
|
headers: dict[str, str]
|
||||||
|
# 用于更新任务的原始数据
|
||||||
|
poll_count: int
|
||||||
|
retry_count: int
|
||||||
|
poll_interval_seconds: int
|
||||||
|
max_poll_count: int
|
||||||
|
current_status: str
|
||||||
|
|
||||||
|
|
||||||
|
# 永久性错误指示词(用于降级判断,不应重试)
|
||||||
|
_PERMANENT_ERROR_INDICATORS = frozenset(
|
||||||
|
{
|
||||||
|
"not found",
|
||||||
|
"404",
|
||||||
|
"unauthorized",
|
||||||
|
"401",
|
||||||
|
"forbidden",
|
||||||
|
"403",
|
||||||
|
"invalid request",
|
||||||
|
"invalid api key",
|
||||||
|
"does not exist",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PollHTTPError(RuntimeError):
|
||||||
|
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
|
||||||
|
|
||||||
|
def __init__(self, status_code: int, message: str):
|
||||||
|
# 确保错误信息包含状态码
|
||||||
|
full_message = f"HTTP {status_code}: {message}" if message else f"HTTP {status_code}"
|
||||||
|
super().__init__(full_message)
|
||||||
|
self.status_code = status_code
|
||||||
|
self.original_message = message
|
||||||
|
|
||||||
|
|
||||||
|
class VideoTaskPollerAdapter:
|
||||||
|
task_type = "video"
|
||||||
|
|
||||||
|
# scheduler
|
||||||
|
job_id = "task_poller:video"
|
||||||
|
job_name = "视频任务轮询"
|
||||||
|
interval_seconds = config.video_poll_interval_seconds
|
||||||
|
|
||||||
|
# distributed lock
|
||||||
|
lock_key = "task_poller:video:lock"
|
||||||
|
lock_ttl = 60
|
||||||
|
|
||||||
|
# execution
|
||||||
|
batch_size = config.video_poll_batch_size
|
||||||
|
concurrency = config.video_poll_concurrency
|
||||||
|
consecutive_failure_alert_threshold = 5
|
||||||
|
max_backoff_seconds = 300
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._openai_normalizer = OpenAINormalizer()
|
||||||
|
self._gemini_normalizer = GeminiNormalizer()
|
||||||
|
|
||||||
|
def sanitize_error_message(self, message: str) -> str:
|
||||||
|
return sanitize_error_message(message)
|
||||||
|
|
||||||
|
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]:
|
||||||
|
tasks = (
|
||||||
|
db.query(VideoTask)
|
||||||
|
.filter(
|
||||||
|
VideoTask.status.in_(
|
||||||
|
[
|
||||||
|
VideoStatus.SUBMITTED.value,
|
||||||
|
VideoStatus.QUEUED.value,
|
||||||
|
VideoStatus.PROCESSING.value,
|
||||||
|
]
|
||||||
|
),
|
||||||
|
VideoTask.next_poll_at <= now,
|
||||||
|
VideoTask.poll_count < VideoTask.max_poll_count,
|
||||||
|
)
|
||||||
|
.order_by(VideoTask.next_poll_at.asc())
|
||||||
|
.limit(limit)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
return [t.id for t in tasks]
|
||||||
|
|
||||||
|
def get_task(self, db: Session, task_id: str) -> VideoTask | None:
|
||||||
|
# SQLAlchemy 1.4+ API
|
||||||
|
return db.get(VideoTask, task_id)
|
||||||
|
|
||||||
|
# ==================== 分阶段处理方法(优化数据库连接占用)====================
|
||||||
|
|
||||||
|
async def prepare_poll_context(
|
||||||
|
self, db: Session, task: VideoTask
|
||||||
|
) -> VideoPollContext | InternalVideoPollResult:
|
||||||
|
"""
|
||||||
|
阶段 1:准备轮询上下文(短暂持有数据库连接)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
VideoPollContext: 成功时返回上下文
|
||||||
|
InternalVideoPollResult: 失败时返回错误结果
|
||||||
|
"""
|
||||||
|
if not task.endpoint_id or not task.key_id:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="missing_provider_info",
|
||||||
|
error_message="Task missing endpoint_id or key_id",
|
||||||
|
)
|
||||||
|
|
||||||
|
endpoint = self._get_endpoint(db, task.endpoint_id)
|
||||||
|
key = self._get_key(db, task.key_id)
|
||||||
|
|
||||||
|
if not key.api_key:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="provider_config_error",
|
||||||
|
error_message="Provider key not properly configured",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
upstream_key = crypto_service.decrypt(key.api_key)
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Failed to decrypt provider key for task %s", task.id)
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="decryption_error",
|
||||||
|
error_message="Failed to decrypt provider key",
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_format = (task.provider_api_format or "").strip().lower()
|
||||||
|
if not provider_format:
|
||||||
|
provider_format = make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 构建请求头
|
||||||
|
if provider_format.startswith("gemini:"):
|
||||||
|
auth_info = await get_provider_auth(endpoint, key)
|
||||||
|
else:
|
||||||
|
auth_info = None
|
||||||
|
headers = self._build_headers(provider_format, upstream_key, endpoint, auth_info)
|
||||||
|
|
||||||
|
return VideoPollContext(
|
||||||
|
task_id=task.id,
|
||||||
|
external_task_id=task.external_task_id or "",
|
||||||
|
provider_api_format=provider_format,
|
||||||
|
base_url=endpoint.base_url or "",
|
||||||
|
upstream_key=upstream_key,
|
||||||
|
headers=headers,
|
||||||
|
poll_count=task.poll_count,
|
||||||
|
retry_count=task.retry_count,
|
||||||
|
poll_interval_seconds=task.poll_interval_seconds,
|
||||||
|
max_poll_count=task.max_poll_count,
|
||||||
|
current_status=task.status,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def poll_task_http(self, ctx: VideoPollContext) -> InternalVideoPollResult:
|
||||||
|
"""
|
||||||
|
阶段 2:执行 HTTP 请求(不持有数据库连接)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
ctx: 轮询上下文
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
InternalVideoPollResult: 轮询结果
|
||||||
|
"""
|
||||||
|
if not ctx.external_task_id:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="missing_external_task_id",
|
||||||
|
error_message="Task missing external_task_id",
|
||||||
|
)
|
||||||
|
|
||||||
|
if ctx.provider_api_format.startswith("gemini:"):
|
||||||
|
return await self._poll_gemini_with_context(ctx)
|
||||||
|
return await self._poll_openai_with_context(ctx)
|
||||||
|
|
||||||
|
async def update_task_after_poll(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
result: InternalVideoPollResult,
|
||||||
|
ctx: VideoPollContext | None,
|
||||||
|
redis_client: Any | None,
|
||||||
|
error_exception: Exception | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
阶段 3:更新数据库(获取新的数据库连接)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
task_id: 任务 ID
|
||||||
|
result: 轮询结果
|
||||||
|
ctx: 轮询上下文(准备阶段就失败时为 None)
|
||||||
|
redis_client: Redis 客户端
|
||||||
|
error_exception: 如果 HTTP 请求失败,传入异常对象
|
||||||
|
"""
|
||||||
|
with create_session() as db:
|
||||||
|
task = db.get(VideoTask, task_id)
|
||||||
|
if not task:
|
||||||
|
logger.warning("Task %s disappeared during poll update", task_id)
|
||||||
|
return
|
||||||
|
|
||||||
|
if error_exception is not None and ctx is not None:
|
||||||
|
# HTTP 请求失败(需要 ctx 来计算 backoff)
|
||||||
|
self._handle_poll_error(task, error_exception, ctx)
|
||||||
|
elif result.status == VideoStatus.COMPLETED:
|
||||||
|
task.status = VideoStatus.COMPLETED.value
|
||||||
|
task.video_url = result.video_url
|
||||||
|
task.video_expires_at = result.expires_at
|
||||||
|
task.completed_at = datetime.now(timezone.utc)
|
||||||
|
task.progress_percent = 100
|
||||||
|
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
|
||||||
|
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
|
||||||
|
seconds=task.poll_interval_seconds
|
||||||
|
)
|
||||||
|
|
||||||
|
# 超时检查
|
||||||
|
task.updated_at = datetime.now(timezone.utc)
|
||||||
|
if task.poll_count >= task.max_poll_count and task.status not in [
|
||||||
|
VideoStatus.COMPLETED.value,
|
||||||
|
VideoStatus.FAILED.value,
|
||||||
|
VideoStatus.CANCELLED.value,
|
||||||
|
]:
|
||||||
|
task.status = VideoStatus.FAILED.value
|
||||||
|
task.error_code = "poll_timeout"
|
||||||
|
task.error_message = f"Task timed out after {task.poll_count} polls"
|
||||||
|
task.completed_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# 终态结算
|
||||||
|
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||||
|
try:
|
||||||
|
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
|
||||||
|
task
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception(
|
||||||
|
"Failed to record video usage for task=%s: %s",
|
||||||
|
task.id,
|
||||||
|
sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
def _handle_poll_error(self, task: VideoTask, exc: Exception, ctx: VideoPollContext) -> None:
|
||||||
|
"""处理轮询错误"""
|
||||||
|
task.poll_count += 1
|
||||||
|
error_msg = sanitize_error_message(str(exc))
|
||||||
|
logger.warning("Poll error for task %s: %s", task.id, error_msg)
|
||||||
|
task.progress_message = f"Poll error: {error_msg}"
|
||||||
|
|
||||||
|
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
|
||||||
|
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
||||||
|
if is_permanent:
|
||||||
|
task.status = VideoStatus.FAILED.value
|
||||||
|
task.error_code = "poll_permanent_error"
|
||||||
|
task.error_message = error_msg
|
||||||
|
task.completed_at = datetime.now(timezone.utc)
|
||||||
|
else:
|
||||||
|
backoff = min(
|
||||||
|
ctx.poll_interval_seconds * (2 ** min(ctx.retry_count, 5)),
|
||||||
|
self.max_backoff_seconds,
|
||||||
|
)
|
||||||
|
task.retry_count += 1
|
||||||
|
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
||||||
|
|
||||||
|
async def _poll_openai_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
|
||||||
|
"""使用上下文进行 OpenAI 轮询(不需要数据库)"""
|
||||||
|
url = self._build_openai_url(ctx.base_url, ctx.external_task_id)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
response = await client.get(url, headers=ctx.headers)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
error_message = self._extract_error_message(response.text, response.status_code)
|
||||||
|
raise PollHTTPError(response.status_code, error_message)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
return self._openai_normalizer.video_poll_to_internal(payload)
|
||||||
|
|
||||||
|
async def _poll_gemini_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
|
||||||
|
"""使用上下文进行 Gemini 轮询(不需要数据库)"""
|
||||||
|
operation_name = normalize_gemini_operation_id(ctx.external_task_id)
|
||||||
|
url = self._build_gemini_url(ctx.base_url, operation_name)
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
|
||||||
|
ctx.task_id,
|
||||||
|
ctx.external_task_id,
|
||||||
|
url,
|
||||||
|
)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
response = await client.get(url, headers=ctx.headers)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
logger.warning(
|
||||||
|
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
|
||||||
|
ctx.task_id,
|
||||||
|
response.status_code,
|
||||||
|
response.text[:500] if response.text else "(empty)",
|
||||||
|
)
|
||||||
|
error_message = self._extract_error_message(response.text, response.status_code)
|
||||||
|
raise PollHTTPError(response.status_code, error_message)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
return self._gemini_normalizer.video_poll_to_internal(payload)
|
||||||
|
|
||||||
|
# ==================== 旧版方法(保留兼容性)====================
|
||||||
|
|
||||||
|
async def poll_single_task(
|
||||||
|
self, db: Session, task: VideoTask, *, redis_client: Any | None
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
旧版单任务轮询方法(保留向后兼容)
|
||||||
|
|
||||||
|
注意:此方法在 HTTP 请求期间持有数据库连接,建议使用分阶段方法。
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
result = await self._poll_task_status(db, task)
|
||||||
|
if result.status == VideoStatus.COMPLETED:
|
||||||
|
task.status = VideoStatus.COMPLETED.value
|
||||||
|
task.video_url = result.video_url
|
||||||
|
task.video_expires_at = result.expires_at
|
||||||
|
task.completed_at = datetime.now(timezone.utc)
|
||||||
|
task.progress_percent = 100
|
||||||
|
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
|
||||||
|
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
|
||||||
|
seconds=task.poll_interval_seconds
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
task.poll_count += 1
|
||||||
|
error_msg = sanitize_error_message(str(exc))
|
||||||
|
logger.warning("Poll error for task %s: %s", task.id, error_msg)
|
||||||
|
task.progress_message = f"Poll error: {error_msg}"
|
||||||
|
|
||||||
|
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
|
||||||
|
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
||||||
|
if is_permanent:
|
||||||
|
task.status = VideoStatus.FAILED.value
|
||||||
|
task.error_code = "poll_permanent_error"
|
||||||
|
task.error_message = error_msg
|
||||||
|
task.completed_at = datetime.now(timezone.utc)
|
||||||
|
else:
|
||||||
|
backoff = min(
|
||||||
|
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
|
||||||
|
self.max_backoff_seconds,
|
||||||
|
)
|
||||||
|
task.retry_count += 1
|
||||||
|
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
||||||
|
|
||||||
|
# 超时:超过最大轮询次数且未进入终态
|
||||||
|
task.updated_at = datetime.now(timezone.utc)
|
||||||
|
if task.poll_count >= task.max_poll_count and task.status not in [
|
||||||
|
VideoStatus.COMPLETED.value,
|
||||||
|
VideoStatus.FAILED.value,
|
||||||
|
VideoStatus.CANCELLED.value,
|
||||||
|
]:
|
||||||
|
task.status = VideoStatus.FAILED.value
|
||||||
|
task.error_code = "poll_timeout"
|
||||||
|
task.error_message = f"Task timed out after {task.poll_count} polls"
|
||||||
|
task.completed_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# 终态结算
|
||||||
|
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||||
|
try:
|
||||||
|
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
|
||||||
|
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
|
||||||
|
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
|
||||||
|
# (直接修改 JSON 字段内部不会自动标记为 dirty)
|
||||||
|
metadata = dict(task.request_metadata) if task.request_metadata else {}
|
||||||
|
metadata["poll_raw_response"] = result.raw_response
|
||||||
|
task.request_metadata = metadata
|
||||||
|
|
||||||
|
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
|
||||||
|
if status_code is not None:
|
||||||
|
return 400 <= status_code < 500 and status_code != 429
|
||||||
|
error_msg = str(exc).lower()
|
||||||
|
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
|
||||||
|
|
||||||
|
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
|
||||||
|
if not task.endpoint_id or not task.key_id:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="missing_provider_info",
|
||||||
|
error_message="Task missing endpoint_id or key_id",
|
||||||
|
)
|
||||||
|
endpoint = self._get_endpoint(db, task.endpoint_id)
|
||||||
|
key = self._get_key(db, task.key_id)
|
||||||
|
if not key.api_key:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="provider_config_error",
|
||||||
|
error_message="Provider key not properly configured",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
upstream_key = crypto_service.decrypt(key.api_key)
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Failed to decrypt provider key for task %s", task.id)
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="decryption_error",
|
||||||
|
error_message="Failed to decrypt provider key",
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_format = (task.provider_api_format or "").strip().lower()
|
||||||
|
if not provider_format:
|
||||||
|
provider_format = make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
|
||||||
|
if provider_format.startswith("gemini:"):
|
||||||
|
auth_info = await get_provider_auth(endpoint, key)
|
||||||
|
return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
|
||||||
|
return await self._poll_openai(task, endpoint, upstream_key)
|
||||||
|
|
||||||
|
async def _poll_openai(
|
||||||
|
self,
|
||||||
|
task: VideoTask,
|
||||||
|
endpoint: ProviderEndpoint,
|
||||||
|
upstream_key: str,
|
||||||
|
) -> InternalVideoPollResult:
|
||||||
|
if not task.external_task_id:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="missing_external_task_id",
|
||||||
|
error_message="Task missing external_task_id",
|
||||||
|
)
|
||||||
|
url = self._build_openai_url(endpoint.base_url, task.external_task_id)
|
||||||
|
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
response = await client.get(url, headers=headers)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
error_message = self._extract_error_message(response.text, response.status_code)
|
||||||
|
raise PollHTTPError(response.status_code, error_message)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
return self._openai_normalizer.video_poll_to_internal(payload)
|
||||||
|
|
||||||
|
async def _poll_gemini(
|
||||||
|
self,
|
||||||
|
task: VideoTask,
|
||||||
|
endpoint: ProviderEndpoint,
|
||||||
|
upstream_key: str,
|
||||||
|
auth_info: ProviderAuthInfo | None,
|
||||||
|
) -> InternalVideoPollResult:
|
||||||
|
if not task.external_task_id:
|
||||||
|
return InternalVideoPollResult(
|
||||||
|
status=VideoStatus.FAILED,
|
||||||
|
error_code="missing_external_task_id",
|
||||||
|
error_message="Task missing external_task_id",
|
||||||
|
)
|
||||||
|
operation_name = normalize_gemini_operation_id(task.external_task_id)
|
||||||
|
url = self._build_gemini_url(endpoint.base_url, operation_name)
|
||||||
|
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
||||||
|
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||||
|
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||||
|
)
|
||||||
|
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
|
||||||
|
task.id,
|
||||||
|
task.external_task_id,
|
||||||
|
url,
|
||||||
|
)
|
||||||
|
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
response = await client.get(url, headers=headers)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
logger.warning(
|
||||||
|
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
|
||||||
|
task.id,
|
||||||
|
response.status_code,
|
||||||
|
response.text[:500] if response.text else "(empty)",
|
||||||
|
)
|
||||||
|
error_message = self._extract_error_message(response.text, response.status_code)
|
||||||
|
raise PollHTTPError(response.status_code, error_message)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
return self._gemini_normalizer.video_poll_to_internal(payload)
|
||||||
|
|
||||||
|
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
|
||||||
|
base = (base_url or "https://api.openai.com").rstrip("/")
|
||||||
|
if base.endswith("/v1"):
|
||||||
|
return f"{base}/videos/{task_id}"
|
||||||
|
return f"{base}/v1/videos/{task_id}"
|
||||||
|
|
||||||
|
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
|
||||||
|
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
|
||||||
|
if base.endswith("/v1beta"):
|
||||||
|
base = base[: -len("/v1beta")]
|
||||||
|
return f"{base}/v1beta/{operation_name}"
|
||||||
|
|
||||||
|
def _build_headers(
|
||||||
|
self,
|
||||||
|
endpoint_sig: str,
|
||||||
|
upstream_key: str,
|
||||||
|
endpoint: ProviderEndpoint,
|
||||||
|
auth_info: ProviderAuthInfo | None = None,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||||
|
headers = build_upstream_headers_for_endpoint(
|
||||||
|
{},
|
||||||
|
endpoint_sig,
|
||||||
|
upstream_key,
|
||||||
|
endpoint_headers=extra_headers,
|
||||||
|
)
|
||||||
|
if auth_info:
|
||||||
|
headers.pop("x-goog-api-key", None)
|
||||||
|
headers[auth_info.auth_header] = auth_info.auth_value
|
||||||
|
return headers
|
||||||
|
|
||||||
|
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
|
||||||
|
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
|
||||||
|
if not endpoint:
|
||||||
|
raise RuntimeError("Provider endpoint not found")
|
||||||
|
return endpoint
|
||||||
|
|
||||||
|
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
|
||||||
|
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||||
|
if not key:
|
||||||
|
raise RuntimeError("Provider key not found")
|
||||||
|
return key
|
||||||
|
|
||||||
|
def _extract_error_message(self, response_text: str | None, status_code: int) -> str:
|
||||||
|
"""从响应中提取有意义的错误信息"""
|
||||||
|
if not response_text:
|
||||||
|
return f"Request failed with status {status_code}"
|
||||||
|
|
||||||
|
# 尝试解析 JSON 格式的错误
|
||||||
|
try:
|
||||||
|
data = json.loads(response_text)
|
||||||
|
# OpenAI 格式: {"error": {"message": "..."}}
|
||||||
|
if isinstance(data.get("error"), dict):
|
||||||
|
error_obj = data["error"]
|
||||||
|
message = error_obj.get("message") or error_obj.get("detail") or str(error_obj)
|
||||||
|
return sanitize_error_message(message)
|
||||||
|
# Gemini 格式: {"error": {"message": "...", "code": 404}}
|
||||||
|
if "message" in data:
|
||||||
|
return sanitize_error_message(data["message"])
|
||||||
|
except (json.JSONDecodeError, TypeError, KeyError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
# 回退到原始文本(截断)
|
||||||
|
return sanitize_error_message(response_text[:500])
|
||||||
@@ -1,320 +0,0 @@
|
|||||||
"""
|
|
||||||
VideoTelemetry(Phase3)
|
|
||||||
|
|
||||||
将 Video 异步任务的“终态计费 + Usage 写入 + required 缺失告警”从 poller 中抽离出来,
|
|
||||||
便于未来 Image/Audio 复用相同框架。
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from sqlalchemy.orm import Session
|
|
||||||
|
|
||||||
from src.api.handlers.base.video_handler_base import sanitize_error_message
|
|
||||||
from src.config.settings import config
|
|
||||||
from src.core.api_format.conversion.internal_video import VideoStatus
|
|
||||||
from src.core.logger import logger
|
|
||||||
from src.models.database import ApiKey, Provider, 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.usage.service import UsageService
|
|
||||||
|
|
||||||
|
|
||||||
class VideoTelemetry:
|
|
||||||
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
|
|
||||||
self.db = db
|
|
||||||
self.redis = redis_client
|
|
||||||
self._formula_engine = FormulaEngine()
|
|
||||||
|
|
||||||
async def record_terminal_usage(self, task: VideoTask) -> None:
|
|
||||||
"""
|
|
||||||
为视频任务终态写入 Usage:
|
|
||||||
- COMPLETED: 使用 FormulaEngine 计算 cost(或 no_rule / incomplete -> cost=0)
|
|
||||||
- FAILED: cost=0
|
|
||||||
|
|
||||||
该方法可能会在 strict_mode 缺失 required 维度时将任务降级为 FAILED 并隐藏产物。
|
|
||||||
"""
|
|
||||||
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(self.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: 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(
|
|
||||||
self.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}
|
|
||||||
|
|
||||||
# 将 billing_snapshot 回写到 task.request_metadata 便于对账(不会影响 usage 的单独存档)
|
|
||||||
if task.request_metadata is None:
|
|
||||||
task.request_metadata = {}
|
|
||||||
if isinstance(task.request_metadata, dict):
|
|
||||||
task.request_metadata["billing_snapshot"] = billing_snapshot
|
|
||||||
|
|
||||||
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 = self.db.query(User).filter(User.id == task.user_id).first()
|
|
||||||
api_key_obj = (
|
|
||||||
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
|
|
||||||
if task.api_key_id
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
provider_obj = (
|
|
||||||
self.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=self.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:
|
|
||||||
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)
|
|
||||||
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))
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["VideoTelemetry"]
|
|
||||||
27
src/services/task/lifecycle.py
Normal file
27
src/services/task/lifecycle.py
Normal file
@@ -0,0 +1,27 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class TaskStatus(str, Enum):
|
||||||
|
"""Generic task status (progress)."""
|
||||||
|
|
||||||
|
PENDING = "pending"
|
||||||
|
STREAMING = "streaming"
|
||||||
|
|
||||||
|
SUBMITTED = "submitted"
|
||||||
|
QUEUED = "queued"
|
||||||
|
PROCESSING = "processing"
|
||||||
|
|
||||||
|
COMPLETED = "completed"
|
||||||
|
FAILED = "failed"
|
||||||
|
CANCELLED = "cancelled"
|
||||||
|
EXPIRED = "expired"
|
||||||
|
|
||||||
|
|
||||||
|
class BillingStatus(str, Enum):
|
||||||
|
"""Billing settlement status (Usage.billing_status)."""
|
||||||
|
|
||||||
|
PENDING = "pending"
|
||||||
|
SETTLED = "settled"
|
||||||
|
VOID = "void"
|
||||||
@@ -1,635 +0,0 @@
|
|||||||
"""
|
|
||||||
AsyncTaskOrchestrator
|
|
||||||
|
|
||||||
提交阶段故障转移(多候选尝试):
|
|
||||||
- 目标:拿到 external_task_id 后锁定 provider/endpoint/key,后续轮询不再切换。
|
|
||||||
- 仅覆盖“提交阶段”;轮询阶段由各 task poller 使用已锁定的信息执行。
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import re
|
|
||||||
import uuid
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from datetime import datetime, timezone
|
|
||||||
from typing import Any, Protocol, runtime_checkable
|
|
||||||
|
|
||||||
import httpx
|
|
||||||
from redis.asyncio import Redis
|
|
||||||
from sqlalchemy.orm import Session
|
|
||||||
|
|
||||||
from src.config.settings import config
|
|
||||||
from src.core.exceptions import ProviderNotAvailableException
|
|
||||||
from src.core.logger import logger
|
|
||||||
from src.models.database import ApiKey, RequestCandidate
|
|
||||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
|
||||||
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
|
||||||
from src.services.orchestration.candidate_resolver import CandidateResolver
|
|
||||||
from src.services.orchestration.error_classifier import ErrorClassifier
|
|
||||||
from src.services.system.config import SystemConfigService
|
|
||||||
|
|
||||||
_SENSITIVE_PATTERN = re.compile(
|
|
||||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
|
||||||
re.IGNORECASE,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _sanitize(message: str, max_length: int = 200) -> str:
|
|
||||||
if not message:
|
|
||||||
return "request_failed"
|
|
||||||
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
|
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
|
||||||
class SubmitFunc(Protocol):
|
|
||||||
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
|
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
|
||||||
class ExtractExternalTaskIdFunc(Protocol):
|
|
||||||
def __call__(self, payload: dict[str, Any]) -> str | None: ...
|
|
||||||
|
|
||||||
|
|
||||||
class UpstreamClientRequestError(RuntimeError):
|
|
||||||
"""可判定为客户端请求问题(不应 failover)的上游错误。"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
response: httpx.Response,
|
|
||||||
candidate_keys: list[dict[str, Any]],
|
|
||||||
) -> None:
|
|
||||||
self.response = response
|
|
||||||
self.candidate_keys = candidate_keys
|
|
||||||
super().__init__(f"Upstream client error: HTTP {response.status_code}")
|
|
||||||
|
|
||||||
|
|
||||||
class AllCandidatesFailedError(RuntimeError):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
reason: str,
|
|
||||||
candidate_keys: list[dict[str, Any]],
|
|
||||||
last_status_code: int | None = None,
|
|
||||||
) -> None:
|
|
||||||
self.reason = reason
|
|
||||||
self.candidate_keys = candidate_keys
|
|
||||||
self.last_status_code = last_status_code
|
|
||||||
super().__init__(f"All candidates failed: {reason}")
|
|
||||||
|
|
||||||
|
|
||||||
class CandidateUnsupportedError(RuntimeError):
|
|
||||||
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
|
|
||||||
|
|
||||||
|
|
||||||
class CandidateSubmissionError(RuntimeError):
|
|
||||||
"""候选提交异常(网络/解密/解析等)。"""
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
|
||||||
class SubmitOutcome:
|
|
||||||
candidate: ProviderCandidate
|
|
||||||
candidate_keys: list[dict[str, Any]]
|
|
||||||
external_task_id: str
|
|
||||||
rule_lookup: BillingRuleLookupResult | None
|
|
||||||
upstream_payload: dict[str, Any] | None = None
|
|
||||||
|
|
||||||
|
|
||||||
class AsyncTaskOrchestrator:
|
|
||||||
"""
|
|
||||||
异步任务编排器:只负责提交阶段的候选遍历与错误处理策略。
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, db: Session, *, redis_client: Redis | None = None) -> None:
|
|
||||||
self.db = db
|
|
||||||
self.redis = redis_client
|
|
||||||
self._candidate_resolver: CandidateResolver | None = None
|
|
||||||
self._error_classifier: ErrorClassifier | None = None
|
|
||||||
|
|
||||||
self._cache_scheduler = None
|
|
||||||
# 候选记录映射:{candidate_index: RequestCandidate}
|
|
||||||
self._candidate_records: dict[int, RequestCandidate] = {}
|
|
||||||
|
|
||||||
def _create_candidate_records(
|
|
||||||
self,
|
|
||||||
candidates: list[ProviderCandidate],
|
|
||||||
request_id: str | None,
|
|
||||||
user_api_key: ApiKey,
|
|
||||||
) -> dict[int, RequestCandidate]:
|
|
||||||
"""
|
|
||||||
为所有候选预创建 RequestCandidate 记录。
|
|
||||||
|
|
||||||
Args:
|
|
||||||
candidates: 候选列表
|
|
||||||
request_id: 请求 ID
|
|
||||||
user_api_key: 用户 API Key
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
{candidate_index: RequestCandidate} 映射
|
|
||||||
"""
|
|
||||||
if not request_id:
|
|
||||||
return {}
|
|
||||||
|
|
||||||
now = datetime.now(timezone.utc)
|
|
||||||
records: dict[int, RequestCandidate] = {}
|
|
||||||
|
|
||||||
for idx, cand in enumerate(candidates):
|
|
||||||
record = RequestCandidate(
|
|
||||||
id=str(uuid.uuid4()),
|
|
||||||
request_id=request_id,
|
|
||||||
candidate_index=idx,
|
|
||||||
retry_index=0,
|
|
||||||
user_id=user_api_key.user_id if user_api_key else None,
|
|
||||||
api_key_id=user_api_key.id if user_api_key else None,
|
|
||||||
provider_id=cand.provider.id,
|
|
||||||
endpoint_id=cand.endpoint.id,
|
|
||||||
key_id=cand.key.id,
|
|
||||||
status="available",
|
|
||||||
is_cached=bool(getattr(cand, "is_cached", False)),
|
|
||||||
created_at=now,
|
|
||||||
)
|
|
||||||
self.db.add(record)
|
|
||||||
records[idx] = record
|
|
||||||
|
|
||||||
try:
|
|
||||||
self.db.flush()
|
|
||||||
except Exception as exc:
|
|
||||||
logger.warning(
|
|
||||||
"[AsyncTaskOrchestrator] Failed to create candidate records: %s",
|
|
||||||
str(exc),
|
|
||||||
)
|
|
||||||
self.db.rollback()
|
|
||||||
return {}
|
|
||||||
|
|
||||||
return records
|
|
||||||
|
|
||||||
def _update_candidate_record(
|
|
||||||
self,
|
|
||||||
idx: int,
|
|
||||||
*,
|
|
||||||
status: str,
|
|
||||||
skip_reason: str | None = None,
|
|
||||||
status_code: int | None = None,
|
|
||||||
error_type: str | None = None,
|
|
||||||
error_message: str | None = None,
|
|
||||||
started_at: datetime | None = None,
|
|
||||||
finished_at: datetime | None = None,
|
|
||||||
) -> None:
|
|
||||||
"""更新候选记录状态。"""
|
|
||||||
record = self._candidate_records.get(idx)
|
|
||||||
if not record:
|
|
||||||
return
|
|
||||||
|
|
||||||
record.status = status
|
|
||||||
if skip_reason is not None:
|
|
||||||
record.skip_reason = skip_reason
|
|
||||||
if status_code is not None:
|
|
||||||
record.status_code = status_code
|
|
||||||
if error_type is not None:
|
|
||||||
record.error_type = error_type
|
|
||||||
if error_message is not None:
|
|
||||||
record.error_message = error_message
|
|
||||||
if started_at is not None:
|
|
||||||
record.started_at = started_at
|
|
||||||
if finished_at is not None:
|
|
||||||
record.finished_at = finished_at
|
|
||||||
|
|
||||||
try:
|
|
||||||
self.db.flush()
|
|
||||||
except Exception as exc:
|
|
||||||
logger.warning(
|
|
||||||
"[AsyncTaskOrchestrator] Failed to update candidate record %d: %s",
|
|
||||||
idx,
|
|
||||||
str(exc),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _commit_candidate_records(self) -> None:
|
|
||||||
"""提交候选记录到数据库。"""
|
|
||||||
if not self._candidate_records:
|
|
||||||
return
|
|
||||||
try:
|
|
||||||
self.db.commit()
|
|
||||||
except Exception as exc:
|
|
||||||
logger.warning(
|
|
||||||
"[AsyncTaskOrchestrator] Failed to commit candidate records: %s",
|
|
||||||
str(exc),
|
|
||||||
)
|
|
||||||
self.db.rollback()
|
|
||||||
|
|
||||||
async def _ensure_initialized(self) -> None:
|
|
||||||
if self._cache_scheduler is not None:
|
|
||||||
return
|
|
||||||
|
|
||||||
# 使用 SystemConfigService 读取运行时调度策略(与 Chat/CLI 一致)
|
|
||||||
priority_mode = SystemConfigService.get_config(
|
|
||||||
self.db,
|
|
||||||
"provider_priority_mode",
|
|
||||||
"provider",
|
|
||||||
)
|
|
||||||
scheduling_mode = SystemConfigService.get_config(
|
|
||||||
self.db,
|
|
||||||
"scheduling_mode",
|
|
||||||
"cache_affinity",
|
|
||||||
)
|
|
||||||
self._cache_scheduler = await get_cache_aware_scheduler(
|
|
||||||
self.redis,
|
|
||||||
priority_mode=priority_mode,
|
|
||||||
scheduling_mode=scheduling_mode,
|
|
||||||
)
|
|
||||||
|
|
||||||
self._candidate_resolver = CandidateResolver(
|
|
||||||
db=self.db,
|
|
||||||
cache_scheduler=self._cache_scheduler,
|
|
||||||
)
|
|
||||||
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
|
|
||||||
|
|
||||||
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
|
|
||||||
"""
|
|
||||||
判断某个上游 HTTP 错误是否为“客户端错误”(不应 failover)。
|
|
||||||
|
|
||||||
规则:
|
|
||||||
- 401/403/429:一般是 key/权限/限流问题,优先 failover
|
|
||||||
- 其他 4xx:若 ErrorClassifier 判断为客户端请求错误,则停止
|
|
||||||
"""
|
|
||||||
if status_code in (401, 403, 429):
|
|
||||||
return False
|
|
||||||
if 400 <= status_code < 500:
|
|
||||||
assert self._error_classifier is not None
|
|
||||||
return self._error_classifier.is_client_error(error_text)
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def submit_with_failover(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
api_format: str,
|
|
||||||
model_name: str,
|
|
||||||
affinity_key: str,
|
|
||||||
user_api_key: ApiKey,
|
|
||||||
request_id: str | None,
|
|
||||||
task_type: str,
|
|
||||||
submit_func: SubmitFunc,
|
|
||||||
extract_external_task_id: ExtractExternalTaskIdFunc,
|
|
||||||
supported_auth_types: set[str] | None = None,
|
|
||||||
allow_format_conversion: bool = False,
|
|
||||||
capability_requirements: dict[str, bool] | None = None,
|
|
||||||
max_candidates: int | None = None,
|
|
||||||
) -> SubmitOutcome:
|
|
||||||
"""
|
|
||||||
提交异步任务并在失败时自动尝试下一个候选,直到拿到 external_task_id。
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
SubmitOutcome(包含选中的候选 + external_task_id + candidate_keys + billing rule lookup)
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
UpstreamClientRequestError: 判定为客户端请求错误(不应 failover)
|
|
||||||
ProviderNotAvailableException: 没有可用候选(调度器层面)
|
|
||||||
AllCandidatesFailedError: 有候选但全部提交失败
|
|
||||||
"""
|
|
||||||
await self._ensure_initialized()
|
|
||||||
assert self._candidate_resolver is not None
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] submit_with_failover: "
|
|
||||||
"api_format=%s, model=%s, task_type=%s, request_id=%s",
|
|
||||||
api_format,
|
|
||||||
model_name,
|
|
||||||
task_type,
|
|
||||||
request_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
candidates, _global_model_id = await self._candidate_resolver.fetch_candidates(
|
|
||||||
api_format=api_format,
|
|
||||||
model_name=model_name,
|
|
||||||
affinity_key=affinity_key,
|
|
||||||
user_api_key=user_api_key,
|
|
||||||
request_id=request_id,
|
|
||||||
is_stream=False,
|
|
||||||
capability_requirements=capability_requirements,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] fetch_candidates returned %d candidates for model=%s",
|
|
||||||
len(candidates),
|
|
||||||
model_name,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果没有候选,直接抛出异常
|
|
||||||
if not candidates:
|
|
||||||
logger.error(
|
|
||||||
"[AsyncTaskOrchestrator] No candidates returned from fetch_candidates for model=%s",
|
|
||||||
model_name,
|
|
||||||
)
|
|
||||||
raise ProviderNotAvailableException("No candidates available")
|
|
||||||
|
|
||||||
if max_candidates is not None and max_candidates > 0:
|
|
||||||
candidates = candidates[:max_candidates]
|
|
||||||
|
|
||||||
# 创建候选记录(用于链路追踪)
|
|
||||||
self._candidate_records = self._create_candidate_records(
|
|
||||||
candidates=candidates,
|
|
||||||
request_id=request_id,
|
|
||||||
user_api_key=user_api_key,
|
|
||||||
)
|
|
||||||
|
|
||||||
candidate_keys: list[dict[str, Any]] = []
|
|
||||||
eligible_count = 0
|
|
||||||
last_status_code: int | None = None
|
|
||||||
|
|
||||||
for idx, cand in enumerate(candidates):
|
|
||||||
submit_started_at = datetime.now(timezone.utc)
|
|
||||||
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
|
|
||||||
candidate_info: dict[str, Any] = {
|
|
||||||
"index": idx,
|
|
||||||
"provider_id": cand.provider.id,
|
|
||||||
"provider_name": cand.provider.name,
|
|
||||||
"endpoint_id": cand.endpoint.id,
|
|
||||||
"key_id": cand.key.id,
|
|
||||||
"key_name": cand.key.name,
|
|
||||||
"auth_type": auth_type,
|
|
||||||
"priority": getattr(cand.key, "priority", 0) or 0,
|
|
||||||
"is_cached": bool(getattr(cand, "is_cached", False)),
|
|
||||||
}
|
|
||||||
candidate_keys.append(candidate_info)
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Checking candidate %d: provider=%s, is_skipped=%s, skip_reason=%s, needs_conversion=%s, auth_type=%s",
|
|
||||||
idx,
|
|
||||||
cand.provider.name,
|
|
||||||
getattr(cand, "is_skipped", False),
|
|
||||||
getattr(cand, "skip_reason", None),
|
|
||||||
getattr(cand, "needs_conversion", False),
|
|
||||||
auth_type,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 调度器层面标记为跳过(健康/熔断/并发等)
|
|
||||||
if getattr(cand, "is_skipped", False):
|
|
||||||
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"skipped": True,
|
|
||||||
"skip_reason": skip_reason,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d skipped: is_skipped=True, reason=%s",
|
|
||||||
idx,
|
|
||||||
cand.skip_reason,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# 视频/图片等直连 upstream 的 handler 目前不支持跨格式转换
|
|
||||||
if not allow_format_conversion and bool(getattr(cand, "needs_conversion", False)):
|
|
||||||
candidate_info.update(
|
|
||||||
{"skipped": True, "skip_reason": "format_conversion_not_supported"}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx, status="skipped", skip_reason="format_conversion_not_supported"
|
|
||||||
)
|
|
||||||
logger.info("[AsyncTaskOrchestrator] Candidate %d skipped: needs_conversion", idx)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# auth_type 过滤
|
|
||||||
if supported_auth_types is not None and auth_type not in supported_auth_types:
|
|
||||||
skip_reason = f"unsupported_auth_type:{auth_type}"
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"skipped": True,
|
|
||||||
"skip_reason": skip_reason,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d skipped: unsupported_auth_type=%s",
|
|
||||||
idx,
|
|
||||||
auth_type,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# billing rule 过滤(可选)
|
|
||||||
rule_lookup: BillingRuleLookupResult | None = None
|
|
||||||
has_billing_rule = True
|
|
||||||
if config.billing_require_rule:
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Checking billing rule for candidate %d (billing_require_rule=True)",
|
|
||||||
idx,
|
|
||||||
)
|
|
||||||
rule_lookup = BillingRuleService.find_rule(
|
|
||||||
self.db,
|
|
||||||
provider_id=cand.provider.id,
|
|
||||||
model_name=model_name,
|
|
||||||
task_type=task_type,
|
|
||||||
)
|
|
||||||
has_billing_rule = rule_lookup is not None
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Billing rule lookup result: has_rule=%s",
|
|
||||||
has_billing_rule,
|
|
||||||
)
|
|
||||||
if not has_billing_rule:
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"has_billing_rule": False,
|
|
||||||
"skipped": True,
|
|
||||||
"skip_reason": "billing_rule_missing",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx, status="skipped", skip_reason="billing_rule_missing"
|
|
||||||
)
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d skipped: billing_rule_missing", idx
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
candidate_info["has_billing_rule"] = has_billing_rule
|
|
||||||
|
|
||||||
logger.info("[AsyncTaskOrchestrator] Candidate %d eligible, attempting submit", idx)
|
|
||||||
eligible_count += 1
|
|
||||||
|
|
||||||
# 更新记录为 pending 状态(开始尝试)
|
|
||||||
self._update_candidate_record(idx, status="pending", started_at=submit_started_at)
|
|
||||||
|
|
||||||
# 尝试提交
|
|
||||||
try:
|
|
||||||
response = await submit_func(cand)
|
|
||||||
except Exception as exc:
|
|
||||||
finished_at = datetime.now(timezone.utc)
|
|
||||||
logger.error(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d submit exception: %s: %s",
|
|
||||||
idx,
|
|
||||||
type(exc).__name__,
|
|
||||||
str(exc),
|
|
||||||
)
|
|
||||||
error_msg = _sanitize(str(exc))
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"attempt_status": "exception",
|
|
||||||
"error_type": type(exc).__name__,
|
|
||||||
"error_message": error_msg,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx,
|
|
||||||
status="failed",
|
|
||||||
error_type=type(exc).__name__,
|
|
||||||
error_message=error_msg,
|
|
||||||
finished_at=finished_at,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d submit response: status_code=%d",
|
|
||||||
idx,
|
|
||||||
response.status_code,
|
|
||||||
)
|
|
||||||
|
|
||||||
last_status_code = int(getattr(response, "status_code", 0) or 0)
|
|
||||||
|
|
||||||
# 上游错误:决定是否停止
|
|
||||||
if response.status_code >= 400:
|
|
||||||
finished_at = datetime.now(timezone.utc)
|
|
||||||
error_text = ""
|
|
||||||
try:
|
|
||||||
error_text = response.text or ""
|
|
||||||
except Exception:
|
|
||||||
error_text = ""
|
|
||||||
error_msg = _sanitize(error_text)
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"attempt_status": "http_error",
|
|
||||||
"status_code": response.status_code,
|
|
||||||
"error_message": error_msg,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx,
|
|
||||||
status="failed",
|
|
||||||
status_code=response.status_code,
|
|
||||||
error_type="http_error",
|
|
||||||
error_message=error_msg,
|
|
||||||
finished_at=finished_at,
|
|
||||||
)
|
|
||||||
if self._should_stop_on_http_error(
|
|
||||||
status_code=response.status_code, error_text=error_text
|
|
||||||
):
|
|
||||||
self._commit_candidate_records()
|
|
||||||
raise UpstreamClientRequestError(
|
|
||||||
response=response,
|
|
||||||
candidate_keys=candidate_keys,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# 解析任务 ID(200 但缺字段也视为失败并 failover)
|
|
||||||
payload: dict[str, Any] | None = None
|
|
||||||
try:
|
|
||||||
data = response.json()
|
|
||||||
if isinstance(data, dict):
|
|
||||||
payload = data
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d response payload: %s",
|
|
||||||
idx,
|
|
||||||
str(payload)[:500] if payload else "None",
|
|
||||||
)
|
|
||||||
except Exception as exc:
|
|
||||||
finished_at = datetime.now(timezone.utc)
|
|
||||||
logger.error(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d invalid JSON: %s",
|
|
||||||
idx,
|
|
||||||
str(exc),
|
|
||||||
)
|
|
||||||
error_msg = _sanitize(str(exc))
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"attempt_status": "invalid_json",
|
|
||||||
"error_type": type(exc).__name__,
|
|
||||||
"error_message": error_msg,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx,
|
|
||||||
status="failed",
|
|
||||||
status_code=response.status_code,
|
|
||||||
error_type="invalid_json",
|
|
||||||
error_message=error_msg,
|
|
||||||
finished_at=finished_at,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
external_task_id = extract_external_task_id(payload or {})
|
|
||||||
logger.info(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d extracted task_id: %s",
|
|
||||||
idx,
|
|
||||||
external_task_id,
|
|
||||||
)
|
|
||||||
if not external_task_id:
|
|
||||||
finished_at = datetime.now(timezone.utc)
|
|
||||||
candidate_info.update(
|
|
||||||
{
|
|
||||||
"attempt_status": "empty_task_id",
|
|
||||||
"error_message": "Upstream returned empty task id",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx,
|
|
||||||
status="failed",
|
|
||||||
status_code=response.status_code,
|
|
||||||
error_type="empty_task_id",
|
|
||||||
error_message="Upstream returned empty task id",
|
|
||||||
finished_at=finished_at,
|
|
||||||
)
|
|
||||||
logger.warning(
|
|
||||||
"[AsyncTaskOrchestrator] Candidate %d: empty task_id, payload keys: %s",
|
|
||||||
idx,
|
|
||||||
list(payload.keys()) if payload else [],
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# 成功
|
|
||||||
finished_at = datetime.now(timezone.utc)
|
|
||||||
candidate_info.update({"attempt_status": "success", "selected": True})
|
|
||||||
self._update_candidate_record(
|
|
||||||
idx,
|
|
||||||
status="success",
|
|
||||||
status_code=response.status_code,
|
|
||||||
finished_at=finished_at,
|
|
||||||
)
|
|
||||||
self._commit_candidate_records()
|
|
||||||
return SubmitOutcome(
|
|
||||||
candidate=cand,
|
|
||||||
candidate_keys=candidate_keys,
|
|
||||||
external_task_id=str(external_task_id),
|
|
||||||
rule_lookup=rule_lookup,
|
|
||||||
upstream_payload=payload,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 没有任何候选可尝试
|
|
||||||
if not candidates:
|
|
||||||
raise ProviderNotAvailableException("No candidates available")
|
|
||||||
|
|
||||||
# 提交所有候选记录
|
|
||||||
self._commit_candidate_records()
|
|
||||||
|
|
||||||
if eligible_count == 0:
|
|
||||||
reason = "no_eligible_candidates"
|
|
||||||
if config.billing_require_rule:
|
|
||||||
reason = "no_candidate_with_billing_rule"
|
|
||||||
raise AllCandidatesFailedError(
|
|
||||||
reason=reason,
|
|
||||||
candidate_keys=candidate_keys,
|
|
||||||
last_status_code=last_status_code,
|
|
||||||
)
|
|
||||||
|
|
||||||
raise AllCandidatesFailedError(
|
|
||||||
reason="all_candidates_failed",
|
|
||||||
candidate_keys=candidate_keys,
|
|
||||||
last_status_code=last_status_code,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"AsyncTaskOrchestrator",
|
|
||||||
"SubmitOutcome",
|
|
||||||
"AllCandidatesFailedError",
|
|
||||||
"UpstreamClientRequestError",
|
|
||||||
"CandidateUnsupportedError",
|
|
||||||
"CandidateSubmissionError",
|
|
||||||
]
|
|
||||||
240
src/services/task/task_poller.py
Normal file
240
src/services/task/task_poller.py
Normal file
@@ -0,0 +1,240 @@
|
|||||||
|
"""
|
||||||
|
Task poller (Phase2)
|
||||||
|
|
||||||
|
Provides a generic polling skeleton for async tasks.
|
||||||
|
Currently wired with a video poller adapter.
|
||||||
|
|
||||||
|
优化:HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
|
||||||
|
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any, Protocol, runtime_checkable
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.core.api_format.conversion.internal_video import InternalVideoPollResult
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.database import create_session
|
||||||
|
from src.services.system.scheduler import get_scheduler
|
||||||
|
from src.services.task.impl.video_poller import VideoPollContext, VideoTaskPollerAdapter
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class TaskPollerAdapter(Protocol):
|
||||||
|
task_type: str
|
||||||
|
|
||||||
|
# scheduler
|
||||||
|
job_id: str
|
||||||
|
job_name: str
|
||||||
|
interval_seconds: int
|
||||||
|
|
||||||
|
# distributed lock (optional, best-effort)
|
||||||
|
lock_key: str
|
||||||
|
lock_ttl: int
|
||||||
|
|
||||||
|
# execution
|
||||||
|
batch_size: int
|
||||||
|
concurrency: int
|
||||||
|
consecutive_failure_alert_threshold: int
|
||||||
|
|
||||||
|
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]: ...
|
||||||
|
|
||||||
|
def get_task(self, db: Session, task_id: str) -> Any | None: ...
|
||||||
|
|
||||||
|
# 分阶段处理方法(推荐使用)
|
||||||
|
async def prepare_poll_context(
|
||||||
|
self, db: Session, task: Any
|
||||||
|
) -> Any: ... # Returns context or error result
|
||||||
|
|
||||||
|
async def poll_task_http(self, ctx: Any) -> Any: ... # Returns poll result
|
||||||
|
|
||||||
|
async def update_task_after_poll(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
result: Any,
|
||||||
|
ctx: Any,
|
||||||
|
redis_client: Any | None,
|
||||||
|
error_exception: Exception | None = None,
|
||||||
|
) -> None: ...
|
||||||
|
|
||||||
|
# 旧版方法(保留兼容性)
|
||||||
|
async def poll_single_task(
|
||||||
|
self, db: Session, task: Any, *, redis_client: Any | None
|
||||||
|
) -> None: ...
|
||||||
|
|
||||||
|
def sanitize_error_message(self, message: str) -> str: ...
|
||||||
|
|
||||||
|
|
||||||
|
class TaskPollerService:
|
||||||
|
"""Generic background poller for async tasks."""
|
||||||
|
|
||||||
|
def __init__(self, adapter: TaskPollerAdapter) -> None:
|
||||||
|
self.adapter = adapter
|
||||||
|
self._lock = asyncio.Lock()
|
||||||
|
self.redis: Any | None = None
|
||||||
|
self._semaphore: asyncio.Semaphore | None = None
|
||||||
|
self._consecutive_failures = 0
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
if self._semaphore is None:
|
||||||
|
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
|
||||||
|
|
||||||
|
# lazy import to avoid redis hard dependency in local runs
|
||||||
|
from src.clients.redis_client import get_redis_client
|
||||||
|
|
||||||
|
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=self.adapter.interval_seconds,
|
||||||
|
job_id=self.adapter.job_id,
|
||||||
|
name=self.adapter.job_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
scheduler = get_scheduler()
|
||||||
|
scheduler.remove_job(self.adapter.job_id)
|
||||||
|
|
||||||
|
async def poll_pending_tasks(self) -> None:
|
||||||
|
async with self._lock:
|
||||||
|
token = await self._acquire_redis_lock()
|
||||||
|
if token is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
with create_session() as db:
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
task_ids = self.adapter.list_due_task_ids(
|
||||||
|
db, now=now, limit=self.adapter.batch_size
|
||||||
|
)
|
||||||
|
|
||||||
|
if not task_ids:
|
||||||
|
self._consecutive_failures = 0
|
||||||
|
return
|
||||||
|
|
||||||
|
poll_results: list[bool] = []
|
||||||
|
|
||||||
|
if self._semaphore is None:
|
||||||
|
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
|
||||||
|
semaphore = self._semaphore
|
||||||
|
|
||||||
|
async def poll_with_semaphore(task_id: str) -> None:
|
||||||
|
async with semaphore:
|
||||||
|
try:
|
||||||
|
# ========== 阶段 1:准备数据(短暂持有连接)==========
|
||||||
|
with create_session() as task_db:
|
||||||
|
task_obj = self.adapter.get_task(task_db, task_id)
|
||||||
|
if not task_obj:
|
||||||
|
logger.warning(
|
||||||
|
"[%s] Task %s disappeared during poll",
|
||||||
|
self.adapter.task_type,
|
||||||
|
task_id,
|
||||||
|
)
|
||||||
|
poll_results.append(True)
|
||||||
|
return
|
||||||
|
|
||||||
|
ctx_or_result = await self.adapter.prepare_poll_context(
|
||||||
|
task_db, task_obj
|
||||||
|
)
|
||||||
|
|
||||||
|
# 检查是否是错误结果(而非上下文)
|
||||||
|
if isinstance(ctx_or_result, InternalVideoPollResult):
|
||||||
|
# 准备阶段就失败了,直接更新任务状态
|
||||||
|
await self.adapter.update_task_after_poll(
|
||||||
|
task_id=task_id,
|
||||||
|
result=ctx_or_result,
|
||||||
|
ctx=None, # type: ignore[arg-type]
|
||||||
|
redis_client=self.redis,
|
||||||
|
)
|
||||||
|
poll_results.append(True)
|
||||||
|
return
|
||||||
|
|
||||||
|
ctx: VideoPollContext = ctx_or_result
|
||||||
|
|
||||||
|
# ========== 阶段 2:HTTP 请求(不持有数据库连接)==========
|
||||||
|
error_exception: Exception | None = None
|
||||||
|
try:
|
||||||
|
result = await self.adapter.poll_task_http(ctx)
|
||||||
|
except Exception as http_exc:
|
||||||
|
# HTTP 请求失败,记录异常以便后续处理
|
||||||
|
error_exception = http_exc
|
||||||
|
result = InternalVideoPollResult(
|
||||||
|
status=None, # type: ignore[arg-type]
|
||||||
|
error_message=str(http_exc),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ========== 阶段 3:更新数据库(获取新连接)==========
|
||||||
|
await self.adapter.update_task_after_poll(
|
||||||
|
task_id=task_id,
|
||||||
|
result=result,
|
||||||
|
ctx=ctx,
|
||||||
|
redis_client=self.redis,
|
||||||
|
error_exception=error_exception,
|
||||||
|
)
|
||||||
|
poll_results.append(True)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception(
|
||||||
|
"[%s] Unexpected error polling task %s: %s",
|
||||||
|
self.adapter.task_type,
|
||||||
|
task_id,
|
||||||
|
self.adapter.sanitize_error_message(str(exc)),
|
||||||
|
)
|
||||||
|
poll_results.append(False)
|
||||||
|
|
||||||
|
async with asyncio.TaskGroup() as tg:
|
||||||
|
for tid in task_ids:
|
||||||
|
tg.create_task(poll_with_semaphore(tid))
|
||||||
|
|
||||||
|
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.adapter.consecutive_failure_alert_threshold
|
||||||
|
):
|
||||||
|
logger.error(
|
||||||
|
"[ALERT] %s poller: %d consecutive batches failed.",
|
||||||
|
self.adapter.task_type,
|
||||||
|
self._consecutive_failures,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._consecutive_failures = 0
|
||||||
|
finally:
|
||||||
|
await self._release_redis_lock(token)
|
||||||
|
|
||||||
|
async def _acquire_redis_lock(self) -> str | None:
|
||||||
|
if not self.redis:
|
||||||
|
return "no_redis"
|
||||||
|
token = str(uuid4())
|
||||||
|
acquired = await self.redis.set(
|
||||||
|
self.adapter.lock_key, token, nx=True, ex=self.adapter.lock_ttl
|
||||||
|
)
|
||||||
|
return token if acquired else None
|
||||||
|
|
||||||
|
async def _release_redis_lock(self, token: str) -> None:
|
||||||
|
if not self.redis or token == "no_redis":
|
||||||
|
return
|
||||||
|
script = """
|
||||||
|
if redis.call('GET', KEYS[1]) == ARGV[1] then
|
||||||
|
return redis.call('DEL', KEYS[1])
|
||||||
|
end
|
||||||
|
return 0
|
||||||
|
"""
|
||||||
|
await self.redis.eval(script, 1, self.adapter.lock_key, token)
|
||||||
|
|
||||||
|
|
||||||
|
_task_poller: TaskPollerService | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_task_poller() -> TaskPollerService:
|
||||||
|
global _task_poller
|
||||||
|
if _task_poller is None:
|
||||||
|
_task_poller = TaskPollerService(VideoTaskPollerAdapter())
|
||||||
|
return _task_poller
|
||||||
@@ -576,6 +576,12 @@ class UsageService:
|
|||||||
"""更新已存在的 Usage 记录(内部方法)"""
|
"""更新已存在的 Usage 记录(内部方法)"""
|
||||||
# 更新关键字段
|
# 更新关键字段
|
||||||
existing_usage.provider_name = usage_params["provider_name"]
|
existing_usage.provider_name = usage_params["provider_name"]
|
||||||
|
existing_usage.model = usage_params["model"]
|
||||||
|
existing_usage.request_type = usage_params["request_type"]
|
||||||
|
existing_usage.api_format = usage_params["api_format"]
|
||||||
|
existing_usage.endpoint_api_format = usage_params["endpoint_api_format"]
|
||||||
|
existing_usage.has_format_conversion = usage_params["has_format_conversion"]
|
||||||
|
existing_usage.is_stream = usage_params["is_stream"]
|
||||||
existing_usage.status = usage_params["status"]
|
existing_usage.status = usage_params["status"]
|
||||||
existing_usage.status_code = usage_params["status_code"]
|
existing_usage.status_code = usage_params["status_code"]
|
||||||
existing_usage.error_message = usage_params["error_message"]
|
existing_usage.error_message = usage_params["error_message"]
|
||||||
@@ -621,6 +627,10 @@ class UsageService:
|
|||||||
existing_usage.provider_endpoint_id = usage_params["provider_endpoint_id"]
|
existing_usage.provider_endpoint_id = usage_params["provider_endpoint_id"]
|
||||||
existing_usage.provider_api_key_id = usage_params["provider_api_key_id"]
|
existing_usage.provider_api_key_id = usage_params["provider_api_key_id"]
|
||||||
|
|
||||||
|
# 更新元数据(如 billing_snapshot/dimensions 等)
|
||||||
|
if usage_params.get("request_metadata") is not None:
|
||||||
|
existing_usage.request_metadata = usage_params["request_metadata"]
|
||||||
|
|
||||||
# 更新模型映射信息
|
# 更新模型映射信息
|
||||||
if target_model is not None:
|
if target_model is not None:
|
||||||
existing_usage.target_model = target_model
|
existing_usage.target_model = target_model
|
||||||
@@ -1000,6 +1010,11 @@ class UsageService:
|
|||||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 结算标记:record_usage_async 写入的 Usage 通常为终态记录
|
||||||
|
if status not in ("pending", "streaming"):
|
||||||
|
usage.billing_status = "settled"
|
||||||
|
usage.finalized_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
db.commit() # 立即提交事务,释放数据库锁
|
db.commit() # 立即提交事务,释放数据库锁
|
||||||
return usage
|
return usage
|
||||||
|
|
||||||
@@ -1172,6 +1187,11 @@ class UsageService:
|
|||||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 结算标记:终态请求写入 settled + finalized_at
|
||||||
|
if status not in ("pending", "streaming"):
|
||||||
|
usage.billing_status = "settled"
|
||||||
|
usage.finalized_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
# 提交事务
|
# 提交事务
|
||||||
try:
|
try:
|
||||||
db.commit()
|
db.commit()
|
||||||
@@ -1297,17 +1317,39 @@ class UsageService:
|
|||||||
is_free_tier=is_free_tier,
|
is_free_tier=is_free_tier,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Upsert(与 record_usage 保持一致)
|
# Upsert(并发幂等:优先用 billing_status 作为结算闸门)
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
existing_usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
existing_usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||||
if existing_usage:
|
if existing_usage:
|
||||||
# 避免重复记账:若已是终态记录,直接返回(批量接口也采用该策略)
|
# 避免重复记账:若已结算/作废,直接返回(防止并发重复加计数)
|
||||||
if existing_usage.status not in ("pending", "streaming"):
|
if getattr(existing_usage, "billing_status", None) in ("settled", "void"):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"record_usage_with_custom_cost: request_id=%s already finalized (status=%s), skip",
|
"record_usage_with_custom_cost: request_id=%s already finalized (billing_status=%s), skip",
|
||||||
request_id,
|
request_id,
|
||||||
existing_usage.status,
|
getattr(existing_usage, "billing_status", None),
|
||||||
)
|
)
|
||||||
return existing_usage
|
return existing_usage
|
||||||
|
|
||||||
|
# 并发闸门:只有 billing_status='pending' 的那一次调用可以继续
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
claim = db.execute(
|
||||||
|
update(Usage)
|
||||||
|
.where(
|
||||||
|
Usage.request_id == request_id,
|
||||||
|
Usage.billing_status == "pending",
|
||||||
|
)
|
||||||
|
.values(billing_status="settled", finalized_at=now)
|
||||||
|
)
|
||||||
|
if claim.rowcount != 1:
|
||||||
|
# 已被其他 worker 抢先处理(或被 VOID)
|
||||||
|
latest = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||||
|
return latest or existing_usage
|
||||||
|
|
||||||
|
# 同步 ORM 对象(避免后续代码读到旧值)
|
||||||
|
existing_usage.billing_status = "settled"
|
||||||
|
existing_usage.finalized_at = now
|
||||||
|
|
||||||
cls._update_existing_usage(existing_usage, usage_params, target_model)
|
cls._update_existing_usage(existing_usage, usage_params, target_model)
|
||||||
usage = existing_usage
|
usage = existing_usage
|
||||||
else:
|
else:
|
||||||
@@ -1382,6 +1424,11 @@ class UsageService:
|
|||||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 结算标记:record_usage_with_custom_cost 写入/更新的 Usage 通常为终态记录
|
||||||
|
if status not in ("pending", "streaming"):
|
||||||
|
usage.billing_status = "settled"
|
||||||
|
usage.finalized_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
db.commit()
|
db.commit()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -2147,35 +2194,31 @@ class UsageService:
|
|||||||
# ========== 请求状态追踪方法 ==========
|
# ========== 请求状态追踪方法 ==========
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create_pending_usage(
|
def begin_pending_usage(
|
||||||
cls,
|
cls,
|
||||||
db: Session,
|
db: Session,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
user: User | None,
|
user: User | None,
|
||||||
api_key: ApiKey | None,
|
api_key: ApiKey | None,
|
||||||
model: str,
|
model: str,
|
||||||
|
*,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
|
request_type: str = "chat",
|
||||||
api_format: str | None = None,
|
api_format: str | None = None,
|
||||||
request_headers: dict[str, Any] | None = None,
|
request_headers: dict[str, Any] | None = None,
|
||||||
request_body: Any | None = None,
|
request_body: Any | None = None,
|
||||||
) -> Usage:
|
) -> Usage:
|
||||||
"""
|
"""
|
||||||
创建 pending 状态的使用记录(在请求开始时调用)
|
创建(或返回已有)pending Usage 记录,但**不提交事务**。
|
||||||
|
|
||||||
Args:
|
适用场景:
|
||||||
db: 数据库会话
|
- ApplicationService 在同一事务内创建 pending usage + task + candidates
|
||||||
request_id: 请求ID
|
- submit 幂等:重复调用同一 request_id 时返回已有记录
|
||||||
user: 用户对象
|
|
||||||
api_key: API Key 对象
|
|
||||||
model: 模型名称
|
|
||||||
is_stream: 是否流式请求
|
|
||||||
api_format: API 格式
|
|
||||||
request_headers: 请求头
|
|
||||||
request_body: 请求体
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
创建的 Usage 记录
|
|
||||||
"""
|
"""
|
||||||
|
existing = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||||
|
if existing:
|
||||||
|
return existing
|
||||||
|
|
||||||
# 根据配置决定是否记录请求详情
|
# 根据配置决定是否记录请求详情
|
||||||
should_log_headers = SystemConfigService.should_log_headers(db)
|
should_log_headers = SystemConfigService.should_log_headers(db)
|
||||||
should_log_body = SystemConfigService.should_log_body(db)
|
should_log_body = SystemConfigService.should_log_body(db)
|
||||||
@@ -2204,21 +2247,359 @@ class UsageService:
|
|||||||
output_tokens=0,
|
output_tokens=0,
|
||||||
total_tokens=0,
|
total_tokens=0,
|
||||||
total_cost_usd=0.0,
|
total_cost_usd=0.0,
|
||||||
request_type="chat",
|
request_type=request_type,
|
||||||
api_format=api_format,
|
api_format=api_format,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
status="pending",
|
status="pending",
|
||||||
|
billing_status="pending",
|
||||||
request_headers=processed_request_headers,
|
request_headers=processed_request_headers,
|
||||||
request_body=processed_request_body,
|
request_body=processed_request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
db.add(usage)
|
db.add(usage)
|
||||||
|
db.flush()
|
||||||
|
return usage
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create_pending_usage(
|
||||||
|
cls,
|
||||||
|
db: Session,
|
||||||
|
request_id: str,
|
||||||
|
user: User | None,
|
||||||
|
api_key: ApiKey | None,
|
||||||
|
model: str,
|
||||||
|
is_stream: bool = False,
|
||||||
|
request_type: str = "chat",
|
||||||
|
api_format: str | None = None,
|
||||||
|
request_headers: dict[str, Any] | None = None,
|
||||||
|
request_body: Any | None = None,
|
||||||
|
) -> Usage:
|
||||||
|
"""
|
||||||
|
创建 pending 状态的使用记录(在请求开始时调用)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: 数据库会话
|
||||||
|
request_id: 请求ID
|
||||||
|
user: 用户对象
|
||||||
|
api_key: API Key 对象
|
||||||
|
model: 模型名称
|
||||||
|
is_stream: 是否流式请求
|
||||||
|
api_format: API 格式
|
||||||
|
request_headers: 请求头
|
||||||
|
request_body: 请求体
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
创建的 Usage 记录
|
||||||
|
"""
|
||||||
|
usage = cls.begin_pending_usage(
|
||||||
|
db,
|
||||||
|
request_id=request_id,
|
||||||
|
user=user,
|
||||||
|
api_key=api_key,
|
||||||
|
model=model,
|
||||||
|
is_stream=is_stream,
|
||||||
|
request_type=request_type,
|
||||||
|
api_format=api_format,
|
||||||
|
request_headers=request_headers,
|
||||||
|
request_body=request_body,
|
||||||
|
)
|
||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
logger.debug(f"创建 pending 使用记录: request_id={request_id}, model={model}")
|
logger.debug(f"创建 pending 使用记录: request_id={request_id}, model={model}")
|
||||||
|
|
||||||
return usage
|
return usage
|
||||||
|
|
||||||
|
# ========== billing_status 并发幂等 finalize ==========
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def finalize_settled(
|
||||||
|
cls,
|
||||||
|
db: Session,
|
||||||
|
request_id: str,
|
||||||
|
*,
|
||||||
|
total_cost_usd: float,
|
||||||
|
request_cost_usd: float | None = None,
|
||||||
|
status: str = "completed",
|
||||||
|
status_code: int = 200,
|
||||||
|
error_message: str | None = None,
|
||||||
|
response_time_ms: int | None = None,
|
||||||
|
billing_snapshot: dict[str, Any] | None = None,
|
||||||
|
extra_metadata: dict[str, Any] | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
并发安全的幂等 finalize(settled)。
|
||||||
|
|
||||||
|
约定:
|
||||||
|
- 仅当 billing_status='pending' 时才会生效(rowcount==1)
|
||||||
|
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||||
|
"""
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
cost = float(total_cost_usd)
|
||||||
|
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
|
||||||
|
|
||||||
|
result = db.execute(
|
||||||
|
update(Usage)
|
||||||
|
.where(
|
||||||
|
Usage.request_id == request_id,
|
||||||
|
Usage.billing_status == "pending",
|
||||||
|
)
|
||||||
|
.values(
|
||||||
|
billing_status="settled",
|
||||||
|
finalized_at=now,
|
||||||
|
total_cost_usd=cost,
|
||||||
|
request_cost_usd=request_cost,
|
||||||
|
status=status,
|
||||||
|
status_code=status_code,
|
||||||
|
error_message=error_message,
|
||||||
|
response_time_ms=response_time_ms,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if result.rowcount != 1:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 写入审计快照(只在本次 finalize 生效时执行)
|
||||||
|
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||||
|
if usage:
|
||||||
|
metadata = usage.request_metadata or {}
|
||||||
|
if billing_snapshot is not None:
|
||||||
|
metadata["billing_snapshot"] = billing_snapshot
|
||||||
|
if extra_metadata:
|
||||||
|
metadata.update(extra_metadata)
|
||||||
|
usage.request_metadata = metadata
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def finalize_void(
|
||||||
|
cls,
|
||||||
|
db: Session,
|
||||||
|
request_id: str,
|
||||||
|
*,
|
||||||
|
reason: str | None = None,
|
||||||
|
status_code: int = 499,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
并发安全的幂等 finalize(void,不收费)。
|
||||||
|
|
||||||
|
约定:
|
||||||
|
- 仅当 billing_status='pending' 时才会生效(rowcount==1)
|
||||||
|
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||||
|
"""
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
result = db.execute(
|
||||||
|
update(Usage)
|
||||||
|
.where(
|
||||||
|
Usage.request_id == request_id,
|
||||||
|
Usage.billing_status == "pending",
|
||||||
|
)
|
||||||
|
.values(
|
||||||
|
billing_status="void",
|
||||||
|
finalized_at=now,
|
||||||
|
total_cost_usd=0.0,
|
||||||
|
request_cost_usd=0.0,
|
||||||
|
status="cancelled",
|
||||||
|
status_code=status_code,
|
||||||
|
error_message=reason,
|
||||||
|
response_time_ms=None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return result.rowcount == 1
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def finalize_submitted(
|
||||||
|
cls,
|
||||||
|
db: Session,
|
||||||
|
request_id: str,
|
||||||
|
*,
|
||||||
|
provider_name: str,
|
||||||
|
provider_id: str | None = None,
|
||||||
|
provider_endpoint_id: str | None = None,
|
||||||
|
provider_api_key_id: str | None = None,
|
||||||
|
response_time_ms: int | None = None,
|
||||||
|
status_code: int = 200,
|
||||||
|
endpoint_api_format: str | None = None,
|
||||||
|
provider_request_headers: dict[str, Any] | None = None,
|
||||||
|
response_headers: dict[str, Any] | None = None,
|
||||||
|
response_body: Any | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
异步任务提交成功时的幂等结算。
|
||||||
|
|
||||||
|
将 pending 使用记录标记为 settled,费用暂时为 0。
|
||||||
|
后续轮询完成后通过 update_settled_billing 更新实际费用。
|
||||||
|
|
||||||
|
约定:
|
||||||
|
- 仅当 billing_status='pending' 时才会生效(rowcount==1)
|
||||||
|
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||||
|
"""
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# 处理响应头和响应体
|
||||||
|
should_log_headers = SystemConfigService.should_log_headers(db)
|
||||||
|
should_log_body = SystemConfigService.should_log_body(db)
|
||||||
|
|
||||||
|
processed_provider_headers = None
|
||||||
|
if should_log_headers and provider_request_headers:
|
||||||
|
processed_provider_headers = SystemConfigService.mask_sensitive_headers(
|
||||||
|
db, provider_request_headers
|
||||||
|
)
|
||||||
|
|
||||||
|
processed_response_headers = None
|
||||||
|
if should_log_headers and response_headers:
|
||||||
|
processed_response_headers = dict(response_headers)
|
||||||
|
|
||||||
|
processed_response_body = None
|
||||||
|
if should_log_body and response_body:
|
||||||
|
processed_response_body = SystemConfigService.truncate_body(
|
||||||
|
db, response_body, is_request=False
|
||||||
|
)
|
||||||
|
|
||||||
|
values: dict[str, Any] = {
|
||||||
|
"billing_status": "settled",
|
||||||
|
"finalized_at": now,
|
||||||
|
"total_cost_usd": 0.0,
|
||||||
|
"request_cost_usd": 0.0,
|
||||||
|
"status": "completed",
|
||||||
|
"status_code": status_code,
|
||||||
|
"response_time_ms": response_time_ms,
|
||||||
|
"provider_name": provider_name,
|
||||||
|
"provider_id": provider_id,
|
||||||
|
"provider_endpoint_id": provider_endpoint_id,
|
||||||
|
"provider_api_key_id": provider_api_key_id,
|
||||||
|
"endpoint_api_format": endpoint_api_format,
|
||||||
|
}
|
||||||
|
|
||||||
|
if processed_provider_headers is not None:
|
||||||
|
values["provider_request_headers"] = processed_provider_headers
|
||||||
|
if processed_response_headers is not None:
|
||||||
|
values["response_headers"] = processed_response_headers
|
||||||
|
if processed_response_body is not None:
|
||||||
|
values["response_body"] = processed_response_body
|
||||||
|
|
||||||
|
result = db.execute(
|
||||||
|
update(Usage)
|
||||||
|
.where(
|
||||||
|
Usage.request_id == request_id,
|
||||||
|
Usage.billing_status == "pending",
|
||||||
|
)
|
||||||
|
.values(**values)
|
||||||
|
)
|
||||||
|
return result.rowcount == 1
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def update_settled_billing(
|
||||||
|
cls,
|
||||||
|
db: Session,
|
||||||
|
request_id: str,
|
||||||
|
*,
|
||||||
|
total_cost_usd: float,
|
||||||
|
request_cost_usd: float | None = None,
|
||||||
|
status: str = "completed",
|
||||||
|
status_code: int = 200,
|
||||||
|
error_message: str | None = None,
|
||||||
|
response_time_ms: int | None = None,
|
||||||
|
billing_snapshot: dict[str, Any] | None = None,
|
||||||
|
extra_metadata: dict[str, Any] | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
更新已结算记录的计费信息(用于异步任务轮询完成后)。
|
||||||
|
|
||||||
|
与 finalize_settled 不同:
|
||||||
|
- finalize_settled: pending -> settled(首次结算)
|
||||||
|
- update_settled_billing: settled -> settled(更新费用)
|
||||||
|
|
||||||
|
约定:
|
||||||
|
- 仅当 billing_status='settled' 时才会生效
|
||||||
|
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||||
|
"""
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
cost = float(total_cost_usd)
|
||||||
|
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
|
||||||
|
|
||||||
|
values: dict[str, Any] = {
|
||||||
|
"total_cost_usd": cost,
|
||||||
|
"request_cost_usd": request_cost,
|
||||||
|
"status": status,
|
||||||
|
"status_code": status_code,
|
||||||
|
}
|
||||||
|
if error_message is not None:
|
||||||
|
values["error_message"] = error_message
|
||||||
|
if response_time_ms is not None:
|
||||||
|
values["response_time_ms"] = response_time_ms
|
||||||
|
|
||||||
|
result = db.execute(
|
||||||
|
update(Usage)
|
||||||
|
.where(
|
||||||
|
Usage.request_id == request_id,
|
||||||
|
Usage.billing_status == "settled",
|
||||||
|
)
|
||||||
|
.values(**values)
|
||||||
|
)
|
||||||
|
if result.rowcount != 1:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 写入审计快照
|
||||||
|
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||||
|
if usage:
|
||||||
|
metadata = usage.request_metadata or {}
|
||||||
|
if billing_snapshot is not None:
|
||||||
|
metadata["billing_snapshot"] = billing_snapshot
|
||||||
|
if extra_metadata:
|
||||||
|
metadata.update(extra_metadata)
|
||||||
|
metadata["billing_updated_at"] = now.isoformat()
|
||||||
|
usage.request_metadata = metadata
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def void_settled(
|
||||||
|
cls,
|
||||||
|
db: Session,
|
||||||
|
request_id: str,
|
||||||
|
*,
|
||||||
|
reason: str | None = None,
|
||||||
|
status_code: int = 499,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
将已结算的记录作废(用于异步任务取消)。
|
||||||
|
|
||||||
|
与 finalize_void 不同:
|
||||||
|
- finalize_void: pending -> void(未结算时作废)
|
||||||
|
- void_settled: settled -> void(已结算后取消,费用归零)
|
||||||
|
|
||||||
|
约定:
|
||||||
|
- 仅当 billing_status='settled' 时才会生效
|
||||||
|
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||||
|
"""
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
result = db.execute(
|
||||||
|
update(Usage)
|
||||||
|
.where(
|
||||||
|
Usage.request_id == request_id,
|
||||||
|
Usage.billing_status == "settled",
|
||||||
|
)
|
||||||
|
.values(
|
||||||
|
billing_status="void",
|
||||||
|
finalized_at=now,
|
||||||
|
total_cost_usd=0.0,
|
||||||
|
request_cost_usd=0.0,
|
||||||
|
status="cancelled",
|
||||||
|
status_code=status_code,
|
||||||
|
error_message=reason,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return result.rowcount == 1
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def update_usage_status(
|
def update_usage_status(
|
||||||
cls,
|
cls,
|
||||||
@@ -2235,6 +2616,7 @@ class UsageService:
|
|||||||
api_format: str | None = None,
|
api_format: str | None = None,
|
||||||
endpoint_api_format: str | None = None,
|
endpoint_api_format: str | None = None,
|
||||||
has_format_conversion: bool | None = None,
|
has_format_conversion: bool | None = None,
|
||||||
|
status_code: int | None = None,
|
||||||
) -> Usage | None:
|
) -> Usage | None:
|
||||||
"""
|
"""
|
||||||
快速更新使用记录状态
|
快速更新使用记录状态
|
||||||
@@ -2253,6 +2635,7 @@ class UsageService:
|
|||||||
api_format: API 格式(可选,用于获取按格式配置的倍率)
|
api_format: API 格式(可选,用于获取按格式配置的倍率)
|
||||||
endpoint_api_format: 端点原生 API 格式(可选)
|
endpoint_api_format: 端点原生 API 格式(可选)
|
||||||
has_format_conversion: 是否发生了格式转换(可选)
|
has_format_conversion: 是否发生了格式转换(可选)
|
||||||
|
status_code: HTTP 状态码(可选)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
更新后的 Usage 记录,如果未找到则返回 None
|
更新后的 Usage 记录,如果未找到则返回 None
|
||||||
@@ -2295,6 +2678,16 @@ class UsageService:
|
|||||||
usage.endpoint_api_format = endpoint_api_format
|
usage.endpoint_api_format = endpoint_api_format
|
||||||
if has_format_conversion is not None:
|
if has_format_conversion is not None:
|
||||||
usage.has_format_conversion = has_format_conversion
|
usage.has_format_conversion = has_format_conversion
|
||||||
|
if status_code is not None:
|
||||||
|
usage.status_code = status_code
|
||||||
|
|
||||||
|
# 结算状态:当请求进入终态时,将 billing_status 标记为 settled
|
||||||
|
# 注意:取消是否应 VOID/部分结算由更高层策略决定;这里默认终态均视为已结算。
|
||||||
|
if status in ("completed", "failed", "cancelled"):
|
||||||
|
if getattr(usage, "billing_status", None) == "pending":
|
||||||
|
usage.billing_status = "settled"
|
||||||
|
if getattr(usage, "finalized_at", None) is None:
|
||||||
|
usage.finalized_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +0,0 @@
|
|||||||
"""
|
|
||||||
视频相关服务
|
|
||||||
"""
|
|
||||||
|
|
||||||
from src.services.video.task_poller import VideoTaskPollerService, get_video_task_poller
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"VideoTaskPollerService",
|
|
||||||
"get_video_task_poller",
|
|
||||||
]
|
|
||||||
@@ -1,446 +0,0 @@
|
|||||||
"""
|
|
||||||
视频任务后台轮询服务
|
|
||||||
"""
|
|
||||||
|
|
||||||
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 (
|
|
||||||
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 (
|
|
||||||
build_upstream_headers_for_endpoint,
|
|
||||||
get_extra_headers_from_endpoint,
|
|
||||||
make_signature_key,
|
|
||||||
)
|
|
||||||
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
|
||||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
|
||||||
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.services.system.scheduler import get_scheduler
|
|
||||||
from src.services.task.impl.video_telemetry import VideoTelemetry
|
|
||||||
|
|
||||||
# 永久性错误指示词(用于降级判断,不应重试)
|
|
||||||
_PERMANENT_ERROR_INDICATORS = frozenset(
|
|
||||||
{
|
|
||||||
"not found",
|
|
||||||
"404",
|
|
||||||
"unauthorized",
|
|
||||||
"401",
|
|
||||||
"forbidden",
|
|
||||||
"403",
|
|
||||||
"invalid request",
|
|
||||||
"invalid api key",
|
|
||||||
"does not exist",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class PollHTTPError(RuntimeError):
|
|
||||||
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
|
|
||||||
|
|
||||||
def __init__(self, status_code: int, message: str):
|
|
||||||
super().__init__(message)
|
|
||||||
self.status_code = status_code
|
|
||||||
|
|
||||||
|
|
||||||
class VideoTaskPollerService:
|
|
||||||
"""后台轮询视频生成任务状态"""
|
|
||||||
|
|
||||||
LOCK_KEY = "video_task_poller:lock"
|
|
||||||
LOCK_TTL = 60
|
|
||||||
MAX_BACKOFF_SECONDS = 300
|
|
||||||
# 连续失败告警阈值
|
|
||||||
CONSECUTIVE_FAILURE_ALERT_THRESHOLD = 5
|
|
||||||
|
|
||||||
def __init__(self) -> None:
|
|
||||||
self._lock = asyncio.Lock()
|
|
||||||
self.redis = None
|
|
||||||
self._openai_normalizer = OpenAINormalizer()
|
|
||||||
self._gemini_normalizer = GeminiNormalizer()
|
|
||||||
# 追踪连续失败次数(用于告警)
|
|
||||||
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=config.video_poll_interval_seconds,
|
|
||||||
job_id="video_task_poller",
|
|
||||||
name="视频任务轮询",
|
|
||||||
)
|
|
||||||
|
|
||||||
async def stop(self) -> None:
|
|
||||||
"""停止轮询服务"""
|
|
||||||
scheduler = get_scheduler()
|
|
||||||
scheduler.remove_job("video_task_poller")
|
|
||||||
|
|
||||||
async def poll_pending_tasks(self) -> None:
|
|
||||||
async with self._lock:
|
|
||||||
token = await self._acquire_redis_lock()
|
|
||||||
if token is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
try:
|
|
||||||
with create_session() as db:
|
|
||||||
now = datetime.now(timezone.utc)
|
|
||||||
tasks = (
|
|
||||||
db.query(VideoTask)
|
|
||||||
.filter(
|
|
||||||
VideoTask.status.in_(
|
|
||||||
[
|
|
||||||
VideoStatus.SUBMITTED.value,
|
|
||||||
VideoStatus.QUEUED.value,
|
|
||||||
VideoStatus.PROCESSING.value,
|
|
||||||
]
|
|
||||||
),
|
|
||||||
VideoTask.next_poll_at <= now,
|
|
||||||
VideoTask.poll_count < VideoTask.max_poll_count,
|
|
||||||
)
|
|
||||||
.order_by(VideoTask.next_poll_at.asc())
|
|
||||||
.limit(self._batch_size)
|
|
||||||
.all()
|
|
||||||
)
|
|
||||||
|
|
||||||
if not tasks:
|
|
||||||
# 无任务时重置连续失败计数
|
|
||||||
self._consecutive_failures = 0
|
|
||||||
return
|
|
||||||
|
|
||||||
# 提取任务 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:
|
|
||||||
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:
|
|
||||||
logger.exception(
|
|
||||||
"Unexpected error polling task %s: %s",
|
|
||||||
task_id,
|
|
||||||
sanitize_error_message(str(exc)),
|
|
||||||
)
|
|
||||||
poll_results.append(False)
|
|
||||||
|
|
||||||
async with asyncio.TaskGroup() as tg:
|
|
||||||
for tid in task_ids:
|
|
||||||
tg.create_task(poll_with_semaphore(tid))
|
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
async def _poll_single_task(self, db: Session, task: VideoTask) -> None:
|
|
||||||
try:
|
|
||||||
result = await self._poll_task_status(db, task)
|
|
||||||
if result.status == VideoStatus.COMPLETED:
|
|
||||||
task.status = VideoStatus.COMPLETED.value
|
|
||||||
task.video_url = result.video_url
|
|
||||||
task.video_expires_at = result.expires_at
|
|
||||||
task.completed_at = datetime.now(timezone.utc)
|
|
||||||
task.progress_percent = 100
|
|
||||||
# 存储多视频 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
|
|
||||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
|
|
||||||
seconds=task.poll_interval_seconds
|
|
||||||
)
|
|
||||||
except Exception as exc:
|
|
||||||
task.poll_count += 1
|
|
||||||
error_msg = sanitize_error_message(str(exc))
|
|
||||||
logger.warning("Poll error for task %s: %s", task.id, error_msg)
|
|
||||||
task.progress_message = f"Poll error: {error_msg}"
|
|
||||||
|
|
||||||
# 区分临时性错误和永久性错误
|
|
||||||
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
|
|
||||||
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
|
||||||
if is_permanent:
|
|
||||||
task.status = VideoStatus.FAILED.value
|
|
||||||
task.error_code = "poll_permanent_error"
|
|
||||||
task.error_message = error_msg
|
|
||||||
task.completed_at = datetime.now(timezone.utc)
|
|
||||||
else:
|
|
||||||
# 临时性错误:指数退避重试
|
|
||||||
backoff = min(
|
|
||||||
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
|
|
||||||
self.MAX_BACKOFF_SECONDS,
|
|
||||||
)
|
|
||||||
task.retry_count += 1
|
|
||||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
|
||||||
|
|
||||||
# 检查是否超过最大轮询次数(超时)
|
|
||||||
task.updated_at = datetime.now(timezone.utc)
|
|
||||||
if task.poll_count >= task.max_poll_count and task.status not in [
|
|
||||||
VideoStatus.COMPLETED.value,
|
|
||||||
VideoStatus.FAILED.value,
|
|
||||||
VideoStatus.CANCELLED.value,
|
|
||||||
]:
|
|
||||||
task.status = VideoStatus.FAILED.value
|
|
||||||
task.error_code = "poll_timeout"
|
|
||||||
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 VideoTelemetry(db, redis_client=self.redis).record_terminal_usage(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
|
|
||||||
|
|
||||||
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
|
|
||||||
"""判断是否为永久性错误(不应重试)"""
|
|
||||||
# 优先使用 HTTP 状态码判断
|
|
||||||
if status_code is not None:
|
|
||||||
# 4xx 客户端错误(除 429 限流)通常是永久性错误
|
|
||||||
return 400 <= status_code < 500 and status_code != 429
|
|
||||||
|
|
||||||
# 降级到字符串匹配
|
|
||||||
error_msg = str(exc).lower()
|
|
||||||
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
|
|
||||||
|
|
||||||
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
|
|
||||||
if not task.endpoint_id or not task.key_id:
|
|
||||||
return InternalVideoPollResult(
|
|
||||||
status=VideoStatus.FAILED,
|
|
||||||
error_code="missing_provider_info",
|
|
||||||
error_message="Task missing endpoint_id or key_id",
|
|
||||||
)
|
|
||||||
endpoint = self._get_endpoint(db, task.endpoint_id)
|
|
||||||
key = self._get_key(db, task.key_id)
|
|
||||||
if not key.api_key:
|
|
||||||
return InternalVideoPollResult(
|
|
||||||
status=VideoStatus.FAILED,
|
|
||||||
error_code="provider_config_error",
|
|
||||||
error_message="Provider key not properly configured",
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
upstream_key = crypto_service.decrypt(key.api_key)
|
|
||||||
except Exception:
|
|
||||||
logger.warning("Failed to decrypt provider key for task %s", task.id)
|
|
||||||
return InternalVideoPollResult(
|
|
||||||
status=VideoStatus.FAILED,
|
|
||||||
error_code="decryption_error",
|
|
||||||
error_message="Failed to decrypt provider key",
|
|
||||||
)
|
|
||||||
|
|
||||||
provider_format = (task.provider_api_format or "").strip().lower()
|
|
||||||
if not provider_format:
|
|
||||||
provider_format = make_signature_key(
|
|
||||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
|
||||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
|
||||||
)
|
|
||||||
|
|
||||||
if provider_format.startswith("gemini:"):
|
|
||||||
auth_info = await get_provider_auth(endpoint, key)
|
|
||||||
return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
|
|
||||||
return await self._poll_openai(task, endpoint, upstream_key)
|
|
||||||
|
|
||||||
async def _poll_openai(
|
|
||||||
self,
|
|
||||||
task: VideoTask,
|
|
||||||
endpoint: ProviderEndpoint,
|
|
||||||
upstream_key: str,
|
|
||||||
) -> InternalVideoPollResult:
|
|
||||||
if not task.external_task_id:
|
|
||||||
return InternalVideoPollResult(
|
|
||||||
status=VideoStatus.FAILED,
|
|
||||||
error_code="missing_external_task_id",
|
|
||||||
error_message="Task missing external_task_id",
|
|
||||||
)
|
|
||||||
url = self._build_openai_url(endpoint.base_url, task.external_task_id)
|
|
||||||
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
|
||||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
|
||||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
|
||||||
)
|
|
||||||
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
|
|
||||||
|
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
|
||||||
response = await client.get(url, headers=headers)
|
|
||||||
if response.status_code >= 400:
|
|
||||||
raise PollHTTPError(
|
|
||||||
response.status_code,
|
|
||||||
sanitize_error_message(response.text or "Poll error"),
|
|
||||||
)
|
|
||||||
|
|
||||||
payload = response.json()
|
|
||||||
return self._openai_normalizer.video_poll_to_internal(payload)
|
|
||||||
|
|
||||||
async def _poll_gemini(
|
|
||||||
self,
|
|
||||||
task: VideoTask,
|
|
||||||
endpoint: ProviderEndpoint,
|
|
||||||
upstream_key: str,
|
|
||||||
auth_info: ProviderAuthInfo | None,
|
|
||||||
) -> InternalVideoPollResult:
|
|
||||||
if not task.external_task_id:
|
|
||||||
return InternalVideoPollResult(
|
|
||||||
status=VideoStatus.FAILED,
|
|
||||||
error_code="missing_external_task_id",
|
|
||||||
error_message="Task missing external_task_id",
|
|
||||||
)
|
|
||||||
operation_name = normalize_gemini_operation_id(task.external_task_id)
|
|
||||||
url = self._build_gemini_url(endpoint.base_url, operation_name)
|
|
||||||
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
|
||||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
|
||||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
|
||||||
)
|
|
||||||
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
|
|
||||||
|
|
||||||
client = await HTTPClientPool.get_default_client_async()
|
|
||||||
response = await client.get(url, headers=headers)
|
|
||||||
if response.status_code >= 400:
|
|
||||||
raise PollHTTPError(
|
|
||||||
response.status_code,
|
|
||||||
sanitize_error_message(response.text or "Poll error"),
|
|
||||||
)
|
|
||||||
|
|
||||||
payload = response.json()
|
|
||||||
return self._gemini_normalizer.video_poll_to_internal(payload)
|
|
||||||
|
|
||||||
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
|
|
||||||
base = (base_url or "https://api.openai.com").rstrip("/")
|
|
||||||
if base.endswith("/v1"):
|
|
||||||
return f"{base}/videos/{task_id}"
|
|
||||||
return f"{base}/v1/videos/{task_id}"
|
|
||||||
|
|
||||||
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
|
|
||||||
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
|
|
||||||
if base.endswith("/v1beta"):
|
|
||||||
base = base[: -len("/v1beta")]
|
|
||||||
return f"{base}/v1beta/{operation_name}"
|
|
||||||
|
|
||||||
def _build_headers(
|
|
||||||
self,
|
|
||||||
endpoint_sig: str,
|
|
||||||
upstream_key: str,
|
|
||||||
endpoint: ProviderEndpoint,
|
|
||||||
auth_info: ProviderAuthInfo | None = None,
|
|
||||||
) -> dict[str, str]:
|
|
||||||
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
|
||||||
headers = build_upstream_headers_for_endpoint(
|
|
||||||
{},
|
|
||||||
endpoint_sig,
|
|
||||||
upstream_key,
|
|
||||||
endpoint_headers=extra_headers,
|
|
||||||
)
|
|
||||||
if auth_info:
|
|
||||||
headers.pop("x-goog-api-key", None)
|
|
||||||
headers[auth_info.auth_header] = auth_info.auth_value
|
|
||||||
return headers
|
|
||||||
|
|
||||||
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
|
|
||||||
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
|
|
||||||
if not endpoint:
|
|
||||||
raise RuntimeError("Provider endpoint not found")
|
|
||||||
return endpoint
|
|
||||||
|
|
||||||
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
|
|
||||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
|
||||||
if not key:
|
|
||||||
raise RuntimeError("Provider key not found")
|
|
||||||
return key
|
|
||||||
|
|
||||||
async def _acquire_redis_lock(self) -> str | None:
|
|
||||||
if not self.redis:
|
|
||||||
return "no_redis"
|
|
||||||
token = str(uuid4())
|
|
||||||
acquired = await self.redis.set(self.LOCK_KEY, token, nx=True, ex=self.LOCK_TTL)
|
|
||||||
return token if acquired else None
|
|
||||||
|
|
||||||
async def _release_redis_lock(self, token: str) -> None:
|
|
||||||
if not self.redis or token == "no_redis":
|
|
||||||
return
|
|
||||||
script = """
|
|
||||||
if redis.call('GET', KEYS[1]) == ARGV[1] then
|
|
||||||
return redis.call('DEL', KEYS[1])
|
|
||||||
end
|
|
||||||
return 0
|
|
||||||
"""
|
|
||||||
await self.redis.eval(script, 1, self.LOCK_KEY, token)
|
|
||||||
|
|
||||||
|
|
||||||
_video_task_poller: VideoTaskPollerService | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def get_video_task_poller() -> VideoTaskPollerService:
|
|
||||||
global _video_task_poller
|
|
||||||
if _video_task_poller is None:
|
|
||||||
_video_task_poller = VideoTaskPollerService()
|
|
||||||
return _video_task_poller
|
|
||||||
198
tests/core/api_format/conversion/test_video_format_conversion.py
Normal file
198
tests/core/api_format/conversion/test_video_format_conversion.py
Normal file
@@ -0,0 +1,198 @@
|
|||||||
|
"""
|
||||||
|
视频格式转换单元测试
|
||||||
|
|
||||||
|
覆盖重点:
|
||||||
|
- OpenAI <-> Gemini 视频请求格式转换
|
||||||
|
- 视频任务响应格式转换
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||||
|
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||||
|
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||||
|
|
||||||
|
|
||||||
|
def _make_registry() -> FormatConversionRegistry:
|
||||||
|
reg = FormatConversionRegistry()
|
||||||
|
reg.register(OpenAINormalizer())
|
||||||
|
reg.register(GeminiNormalizer())
|
||||||
|
return reg
|
||||||
|
|
||||||
|
|
||||||
|
class TestVideoRequestConversion:
|
||||||
|
"""视频请求格式转换测试"""
|
||||||
|
|
||||||
|
def test_openai_to_gemini_video_request(self) -> None:
|
||||||
|
"""OpenAI Sora -> Gemini Veo 请求格式转换"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
openai_request = {
|
||||||
|
"model": "sora-2",
|
||||||
|
"prompt": "A cat playing piano",
|
||||||
|
"size": "1280x720",
|
||||||
|
"seconds": 8,
|
||||||
|
}
|
||||||
|
|
||||||
|
gemini_request = reg.convert_video_request(openai_request, "openai:video", "gemini:video")
|
||||||
|
|
||||||
|
# 验证 Gemini 格式
|
||||||
|
assert "instances" in gemini_request
|
||||||
|
assert "parameters" in gemini_request
|
||||||
|
assert isinstance(gemini_request["instances"], list)
|
||||||
|
assert len(gemini_request["instances"]) > 0
|
||||||
|
|
||||||
|
instance = gemini_request["instances"][0]
|
||||||
|
assert instance["prompt"] == "A cat playing piano"
|
||||||
|
|
||||||
|
params = gemini_request["parameters"]
|
||||||
|
assert params["aspectRatio"] == "16:9"
|
||||||
|
assert params["resolution"] == "720p"
|
||||||
|
assert params["durationSeconds"] == 8
|
||||||
|
|
||||||
|
def test_gemini_to_openai_video_request(self) -> None:
|
||||||
|
"""Gemini Veo -> OpenAI Sora 请求格式转换"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
gemini_request = {
|
||||||
|
"model": "veo-3.1-generate-preview",
|
||||||
|
"instances": [{"prompt": "A beautiful sunset over mountains"}],
|
||||||
|
"parameters": {
|
||||||
|
"aspectRatio": "16:9",
|
||||||
|
"resolution": "1080p",
|
||||||
|
"durationSeconds": 5,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
openai_request = reg.convert_video_request(gemini_request, "gemini:video", "openai:video")
|
||||||
|
|
||||||
|
# 验证 OpenAI 格式
|
||||||
|
assert openai_request["prompt"] == "A beautiful sunset over mountains"
|
||||||
|
assert openai_request["model"] == "veo-3.1-generate-preview"
|
||||||
|
assert openai_request["seconds"] == 5
|
||||||
|
assert openai_request["size"] == "1920x1080"
|
||||||
|
|
||||||
|
def test_same_format_no_conversion(self) -> None:
|
||||||
|
"""相同格式不进行转换"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
openai_request = {
|
||||||
|
"model": "sora-2",
|
||||||
|
"prompt": "Test",
|
||||||
|
"size": "720x1280",
|
||||||
|
"seconds": 4,
|
||||||
|
}
|
||||||
|
|
||||||
|
result = reg.convert_video_request(openai_request, "openai:video", "openai:video")
|
||||||
|
assert result == openai_request
|
||||||
|
|
||||||
|
def test_conversion_with_reference_image(self) -> None:
|
||||||
|
"""带参考图片的请求转换"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
openai_request = {
|
||||||
|
"model": "sora-2",
|
||||||
|
"prompt": "Animate this image",
|
||||||
|
"size": "1280x720",
|
||||||
|
"seconds": 4,
|
||||||
|
"input_reference": "base64_encoded_image_data",
|
||||||
|
}
|
||||||
|
|
||||||
|
gemini_request = reg.convert_video_request(openai_request, "openai:video", "gemini:video")
|
||||||
|
|
||||||
|
instance = gemini_request["instances"][0]
|
||||||
|
assert "image" in instance
|
||||||
|
assert instance["image"]["bytesBase64Encoded"] == "base64_encoded_image_data"
|
||||||
|
|
||||||
|
|
||||||
|
class TestVideoTaskConversion:
|
||||||
|
"""视频任务响应格式转换测试"""
|
||||||
|
|
||||||
|
def test_gemini_to_openai_processing_task(self) -> None:
|
||||||
|
"""Gemini 处理中任务 -> OpenAI 格式"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
gemini_response = {
|
||||||
|
"name": "operations/12345",
|
||||||
|
"done": False,
|
||||||
|
"metadata": {"progress": 50},
|
||||||
|
}
|
||||||
|
|
||||||
|
openai_response = reg.convert_video_task(gemini_response, "gemini:video", "openai:video")
|
||||||
|
|
||||||
|
# OpenAI 格式使用 status 字段
|
||||||
|
assert openai_response["status"] in ("queued", "processing")
|
||||||
|
assert "id" in openai_response
|
||||||
|
|
||||||
|
def test_gemini_to_openai_completed_task(self) -> None:
|
||||||
|
"""Gemini 完成任务 -> OpenAI 格式"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
gemini_response = {
|
||||||
|
"name": "operations/12345",
|
||||||
|
"done": True,
|
||||||
|
"response": {
|
||||||
|
"generateVideoResponse": {
|
||||||
|
"generatedSamples": [{"video": {"uri": "https://example.com/video.mp4"}}]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
openai_response = reg.convert_video_task(gemini_response, "gemini:video", "openai:video")
|
||||||
|
|
||||||
|
assert openai_response["status"] == "completed"
|
||||||
|
assert openai_response["progress"] == 100
|
||||||
|
|
||||||
|
def test_openai_to_gemini_processing_task(self) -> None:
|
||||||
|
"""OpenAI 处理中任务 -> Gemini 格式"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
openai_response = {
|
||||||
|
"id": "task_12345",
|
||||||
|
"object": "video",
|
||||||
|
"status": "processing",
|
||||||
|
"progress": 30,
|
||||||
|
"created_at": 1700000000,
|
||||||
|
}
|
||||||
|
|
||||||
|
gemini_response = reg.convert_video_task(openai_response, "openai:video", "gemini:video")
|
||||||
|
|
||||||
|
# Gemini 格式使用 done 字段
|
||||||
|
assert gemini_response["done"] is False
|
||||||
|
assert "name" in gemini_response
|
||||||
|
|
||||||
|
def test_openai_to_gemini_completed_task(self) -> None:
|
||||||
|
"""OpenAI 完成任务 -> Gemini 格式"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
openai_response = {
|
||||||
|
"id": "task_12345",
|
||||||
|
"object": "video",
|
||||||
|
"status": "completed",
|
||||||
|
"progress": 100,
|
||||||
|
"created_at": 1700000000,
|
||||||
|
"completed_at": 1700000100,
|
||||||
|
}
|
||||||
|
|
||||||
|
gemini_response = reg.convert_video_task(openai_response, "openai:video", "gemini:video")
|
||||||
|
|
||||||
|
assert gemini_response["done"] is True
|
||||||
|
|
||||||
|
|
||||||
|
class TestVideoConversionCapabilities:
|
||||||
|
"""视频格式转换能力检查测试"""
|
||||||
|
|
||||||
|
def test_can_convert_video(self) -> None:
|
||||||
|
"""检查视频格式转换能力"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
assert reg.can_convert_video("openai:video", "gemini:video") is True
|
||||||
|
assert reg.can_convert_video("gemini:video", "openai:video") is True
|
||||||
|
assert reg.can_convert_video("openai:video", "openai:video") is True
|
||||||
|
|
||||||
|
def test_cannot_convert_unsupported_format(self) -> None:
|
||||||
|
"""不支持的格式无法转换"""
|
||||||
|
reg = _make_registry()
|
||||||
|
|
||||||
|
# 没有注册 Claude normalizer 用于视频
|
||||||
|
assert reg.can_convert_video("openai:video", "claude:video") is False
|
||||||
@@ -6,11 +6,8 @@ import httpx
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.config.settings import config
|
from src.config.settings import config
|
||||||
from src.services.task.orchestrator import (
|
from src.services.candidate.service import CandidateService
|
||||||
AllCandidatesFailedError,
|
from src.services.candidate.submit import AllCandidatesFailedError, UpstreamClientRequestError
|
||||||
AsyncTaskOrchestrator,
|
|
||||||
UpstreamClientRequestError,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _make_candidate(
|
def _make_candidate(
|
||||||
@@ -46,10 +43,10 @@ async def test_submit_with_failover_skips_http_500_then_succeeds(
|
|||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
db = MagicMock()
|
db = MagicMock()
|
||||||
orch = AsyncTaskOrchestrator(db)
|
svc = CandidateService(db)
|
||||||
|
|
||||||
# bypass init
|
# bypass init
|
||||||
orch._candidate_resolver = SimpleNamespace(
|
svc._resolver = SimpleNamespace(
|
||||||
fetch_candidates=AsyncMock(
|
fetch_candidates=AsyncMock(
|
||||||
return_value=(
|
return_value=(
|
||||||
[
|
[
|
||||||
@@ -60,8 +57,8 @@ async def test_submit_with_failover_skips_http_500_then_succeeds(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
|
svc._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
|
||||||
orch._ensure_initialized = AsyncMock(return_value=None)
|
svc._ensure_initialized = AsyncMock(return_value=None)
|
||||||
|
|
||||||
responses = [
|
responses = [
|
||||||
httpx.Response(500, text='{"error": {"message": "server"}}'),
|
httpx.Response(500, text='{"error": {"message": "server"}}'),
|
||||||
@@ -69,12 +66,12 @@ async def test_submit_with_failover_skips_http_500_then_succeeds(
|
|||||||
]
|
]
|
||||||
submit = AsyncMock(side_effect=responses)
|
submit = AsyncMock(side_effect=responses)
|
||||||
|
|
||||||
outcome = await orch.submit_with_failover(
|
outcome = await svc.submit_with_failover(
|
||||||
api_format="openai:video",
|
api_format="openai:video",
|
||||||
model_name="sora",
|
model_name="sora",
|
||||||
affinity_key="a1",
|
affinity_key="a1",
|
||||||
user_api_key=MagicMock(),
|
user_api_key=MagicMock(),
|
||||||
request_id="req-1",
|
request_id=None,
|
||||||
task_type="video",
|
task_type="video",
|
||||||
submit_func=submit,
|
submit_func=submit,
|
||||||
extract_external_task_id=lambda payload: payload.get("id"),
|
extract_external_task_id=lambda payload: payload.get("id"),
|
||||||
@@ -92,13 +89,13 @@ async def test_submit_with_failover_skips_http_500_then_succeeds(
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_submit_with_failover_stops_on_client_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
async def test_submit_with_failover_stops_on_client_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
db = MagicMock()
|
db = MagicMock()
|
||||||
orch = AsyncTaskOrchestrator(db)
|
svc = CandidateService(db)
|
||||||
|
|
||||||
orch._candidate_resolver = SimpleNamespace(
|
svc._resolver = SimpleNamespace(
|
||||||
fetch_candidates=AsyncMock(return_value=([_make_candidate()], "gm1"))
|
fetch_candidates=AsyncMock(return_value=([_make_candidate()], "gm1"))
|
||||||
)
|
)
|
||||||
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: True)
|
svc._error_classifier = SimpleNamespace(is_client_error=lambda _text: True)
|
||||||
orch._ensure_initialized = AsyncMock(return_value=None)
|
svc._ensure_initialized = AsyncMock(return_value=None)
|
||||||
|
|
||||||
response = httpx.Response(
|
response = httpx.Response(
|
||||||
400,
|
400,
|
||||||
@@ -107,12 +104,12 @@ async def test_submit_with_failover_stops_on_client_error(monkeypatch: pytest.Mo
|
|||||||
submit = AsyncMock(return_value=response)
|
submit = AsyncMock(return_value=response)
|
||||||
|
|
||||||
with pytest.raises(UpstreamClientRequestError):
|
with pytest.raises(UpstreamClientRequestError):
|
||||||
await orch.submit_with_failover(
|
await svc.submit_with_failover(
|
||||||
api_format="openai:video",
|
api_format="openai:video",
|
||||||
model_name="sora",
|
model_name="sora",
|
||||||
affinity_key="a1",
|
affinity_key="a1",
|
||||||
user_api_key=MagicMock(),
|
user_api_key=MagicMock(),
|
||||||
request_id="req-2",
|
request_id=None,
|
||||||
task_type="video",
|
task_type="video",
|
||||||
submit_func=submit,
|
submit_func=submit,
|
||||||
extract_external_task_id=lambda payload: payload.get("id"),
|
extract_external_task_id=lambda payload: payload.get("id"),
|
||||||
@@ -127,21 +124,21 @@ async def test_submit_with_failover_no_eligible_candidates_due_to_auth_type(
|
|||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
db = MagicMock()
|
db = MagicMock()
|
||||||
orch = AsyncTaskOrchestrator(db)
|
svc = CandidateService(db)
|
||||||
|
|
||||||
orch._candidate_resolver = SimpleNamespace(
|
svc._resolver = SimpleNamespace(
|
||||||
fetch_candidates=AsyncMock(return_value=([_make_candidate(auth_type="vertex_ai")], "gm1"))
|
fetch_candidates=AsyncMock(return_value=([_make_candidate(auth_type="vertex_ai")], "gm1"))
|
||||||
)
|
)
|
||||||
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
|
svc._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
|
||||||
orch._ensure_initialized = AsyncMock(return_value=None)
|
svc._ensure_initialized = AsyncMock(return_value=None)
|
||||||
|
|
||||||
with pytest.raises(AllCandidatesFailedError) as excinfo:
|
with pytest.raises(AllCandidatesFailedError) as excinfo:
|
||||||
await orch.submit_with_failover(
|
await svc.submit_with_failover(
|
||||||
api_format="openai:video",
|
api_format="openai:video",
|
||||||
model_name="sora",
|
model_name="sora",
|
||||||
affinity_key="a1",
|
affinity_key="a1",
|
||||||
user_api_key=MagicMock(),
|
user_api_key=MagicMock(),
|
||||||
request_id="req-3",
|
request_id=None,
|
||||||
task_type="video",
|
task_type="video",
|
||||||
submit_func=AsyncMock(),
|
submit_func=AsyncMock(),
|
||||||
extract_external_task_id=lambda payload: payload.get("id"),
|
extract_external_task_id=lambda payload: payload.get("id"),
|
||||||
@@ -158,9 +155,9 @@ async def test_submit_with_failover_filters_missing_billing_rule(
|
|||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
db = MagicMock()
|
db = MagicMock()
|
||||||
orch = AsyncTaskOrchestrator(db)
|
svc = CandidateService(db)
|
||||||
|
|
||||||
orch._candidate_resolver = SimpleNamespace(
|
svc._resolver = SimpleNamespace(
|
||||||
fetch_candidates=AsyncMock(
|
fetch_candidates=AsyncMock(
|
||||||
return_value=(
|
return_value=(
|
||||||
[
|
[
|
||||||
@@ -171,8 +168,8 @@ async def test_submit_with_failover_filters_missing_billing_rule(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
orch._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
|
svc._error_classifier = SimpleNamespace(is_client_error=lambda _text: False)
|
||||||
orch._ensure_initialized = AsyncMock(return_value=None)
|
svc._ensure_initialized = AsyncMock(return_value=None)
|
||||||
|
|
||||||
# enable require_rule
|
# enable require_rule
|
||||||
old = config.billing_require_rule
|
old = config.billing_require_rule
|
||||||
@@ -185,16 +182,16 @@ async def test_submit_with_failover_filters_missing_billing_rule(
|
|||||||
return None if provider_id == "p1" else object()
|
return None if provider_id == "p1" else object()
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.orchestrator.BillingRuleService.find_rule", _find_rule
|
"src.services.candidate.service.BillingRuleService.find_rule", _find_rule
|
||||||
)
|
)
|
||||||
|
|
||||||
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-999"}))
|
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-999"}))
|
||||||
outcome = await orch.submit_with_failover(
|
outcome = await svc.submit_with_failover(
|
||||||
api_format="openai:video",
|
api_format="openai:video",
|
||||||
model_name="sora",
|
model_name="sora",
|
||||||
affinity_key="a1",
|
affinity_key="a1",
|
||||||
user_api_key=MagicMock(),
|
user_api_key=MagicMock(),
|
||||||
request_id="req-4",
|
request_id=None,
|
||||||
task_type="video",
|
task_type="video",
|
||||||
submit_func=submit,
|
submit_func=submit,
|
||||||
extract_external_task_id=lambda payload: payload.get("id"),
|
extract_external_task_id=lambda payload: payload.get("id"),
|
||||||
@@ -7,12 +7,13 @@ import pytest
|
|||||||
from src.config.settings import config
|
from src.config.settings import config
|
||||||
from src.core.api_format.conversion.internal_video import VideoStatus
|
from src.core.api_format.conversion.internal_video import VideoStatus
|
||||||
from src.services.billing.formula_engine import BillingIncompleteError
|
from src.services.billing.formula_engine import BillingIncompleteError
|
||||||
from src.services.task.impl.video_telemetry import VideoTelemetry
|
from src.services.task.application import TaskApplicationService
|
||||||
|
|
||||||
|
|
||||||
def _make_task(**overrides: Any) -> SimpleNamespace:
|
def _make_task(**overrides: Any) -> SimpleNamespace:
|
||||||
task = SimpleNamespace(
|
task = SimpleNamespace(
|
||||||
id="t1",
|
id="t1",
|
||||||
|
request_id="req-1",
|
||||||
user_id="u1",
|
user_id="u1",
|
||||||
api_key_id="ak1",
|
api_key_id="ak1",
|
||||||
provider_id="p1",
|
provider_id="p1",
|
||||||
@@ -37,7 +38,7 @@ def _make_task(**overrides: Any) -> SimpleNamespace:
|
|||||||
error_code=None,
|
error_code=None,
|
||||||
error_message=None,
|
error_message=None,
|
||||||
status=VideoStatus.COMPLETED.value,
|
status=VideoStatus.COMPLETED.value,
|
||||||
request_metadata={"request_id": "req-1", "poll_raw_response": {"foo": "bar"}},
|
request_metadata={"poll_raw_response": {"foo": "bar"}},
|
||||||
)
|
)
|
||||||
for k, v in overrides.items():
|
for k, v in overrides.items():
|
||||||
setattr(task, k, v)
|
setattr(task, k, v)
|
||||||
@@ -50,7 +51,15 @@ def _make_db() -> MagicMock:
|
|||||||
user_obj = SimpleNamespace(id="u1")
|
user_obj = SimpleNamespace(id="u1")
|
||||||
api_key_obj = SimpleNamespace(id="ak1")
|
api_key_obj = SimpleNamespace(id="ak1")
|
||||||
provider_obj = SimpleNamespace(id="p1", name="prov1")
|
provider_obj = SimpleNamespace(id="p1", name="prov1")
|
||||||
|
usage_obj = SimpleNamespace(
|
||||||
|
id="usage-1",
|
||||||
|
request_id="req-1",
|
||||||
|
billing_status="settled", # finalize_submitted already settled
|
||||||
|
request_metadata=None, # no billing_updated_at yet
|
||||||
|
)
|
||||||
|
|
||||||
|
q_usage = MagicMock()
|
||||||
|
q_usage.filter.return_value.first.return_value = usage_obj
|
||||||
q_user = MagicMock()
|
q_user = MagicMock()
|
||||||
q_user.filter.return_value.first.return_value = user_obj
|
q_user.filter.return_value.first.return_value = user_obj
|
||||||
q_key = MagicMock()
|
q_key = MagicMock()
|
||||||
@@ -58,107 +67,117 @@ def _make_db() -> MagicMock:
|
|||||||
q_provider = MagicMock()
|
q_provider = MagicMock()
|
||||||
q_provider.filter.return_value.first.return_value = provider_obj
|
q_provider.filter.return_value.first.return_value = provider_obj
|
||||||
|
|
||||||
db.query.side_effect = [q_user, q_key, q_provider]
|
db.query.side_effect = [q_usage, q_user, q_key, q_provider]
|
||||||
return db
|
return db
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_video_telemetry_failed_records_cost_zero(monkeypatch: pytest.MonkeyPatch) -> None:
|
async def test_video_finalize_failed_records_cost_zero(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
db = _make_db()
|
db = _make_db()
|
||||||
task = _make_task(status=VideoStatus.FAILED.value, error_message="boom")
|
task = _make_task(status=VideoStatus.FAILED.value, error_message="boom")
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.DimensionCollectorService.collect_dimensions",
|
"src.services.billing.dimension_collector_service.DimensionCollectorService.collect_dimensions",
|
||||||
lambda _self, **_kwargs: {"duration_seconds": 4},
|
lambda _self, **_kwargs: {"duration_seconds": 4},
|
||||||
)
|
)
|
||||||
record = AsyncMock()
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.UsageService.record_usage_with_custom_cost",
|
"src.services.task.application.BillingRuleService.find_rule",
|
||||||
record,
|
lambda *_args, **_kwargs: None,
|
||||||
|
)
|
||||||
|
# Mock update_settled_billing (used by finalize_video_task)
|
||||||
|
update_settled = MagicMock(return_value=True)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.task.application.UsageService.update_settled_billing",
|
||||||
|
update_settled,
|
||||||
)
|
)
|
||||||
|
|
||||||
telemetry = VideoTelemetry(db)
|
app = TaskApplicationService(db)
|
||||||
await telemetry.record_terminal_usage(task)
|
await app.finalize_video_task(task)
|
||||||
|
|
||||||
# billing_snapshot should be written back to task.request_metadata
|
# billing_snapshot should be written back to task.request_metadata
|
||||||
assert task.request_metadata["billing_snapshot"]["billed_reason"] == "task_failed"
|
assert task.request_metadata["billing_snapshot"]["cost"] == 0.0
|
||||||
|
|
||||||
assert record.await_count == 1
|
assert update_settled.call_count == 1
|
||||||
kwargs = record.await_args.kwargs
|
kwargs = update_settled.call_args.kwargs
|
||||||
assert kwargs["total_cost_usd"] == 0.0
|
assert kwargs["total_cost_usd"] == 0.0
|
||||||
assert kwargs["status"] == "failed"
|
assert kwargs["status"] == "failed"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_video_telemetry_completed_no_rule(monkeypatch: pytest.MonkeyPatch) -> None:
|
async def test_video_finalize_completed_no_rule(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
db = _make_db()
|
db = _make_db()
|
||||||
task = _make_task(status=VideoStatus.COMPLETED.value)
|
task = _make_task(status=VideoStatus.COMPLETED.value)
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.DimensionCollectorService.collect_dimensions",
|
"src.services.billing.dimension_collector_service.DimensionCollectorService.collect_dimensions",
|
||||||
lambda _self, **_kwargs: {"duration_seconds": 4},
|
lambda _self, **_kwargs: {"duration_seconds": 4},
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.BillingRuleService.find_rule",
|
"src.services.task.application.BillingRuleService.find_rule",
|
||||||
lambda *_args, **_kwargs: None,
|
lambda *_args, **_kwargs: None,
|
||||||
)
|
)
|
||||||
record = AsyncMock()
|
# Mock update_settled_billing (used by finalize_video_task)
|
||||||
|
update_settled = MagicMock(return_value=True)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.UsageService.record_usage_with_custom_cost",
|
"src.services.task.application.UsageService.update_settled_billing",
|
||||||
record,
|
update_settled,
|
||||||
)
|
)
|
||||||
|
|
||||||
telemetry = VideoTelemetry(db)
|
app = TaskApplicationService(db)
|
||||||
await telemetry.record_terminal_usage(task)
|
await app.finalize_video_task(task)
|
||||||
|
|
||||||
assert task.request_metadata["billing_snapshot"]["status"] == "no_rule"
|
assert task.request_metadata["billing_snapshot"]["status"] == "no_rule"
|
||||||
|
|
||||||
kwargs = record.await_args.kwargs
|
kwargs = update_settled.call_args.kwargs
|
||||||
assert kwargs["total_cost_usd"] == 0.0
|
assert kwargs["total_cost_usd"] == 0.0
|
||||||
assert kwargs["status"] == "completed"
|
assert kwargs["status"] == "completed"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_video_telemetry_strict_mode_missing_required_marks_failed(
|
async def test_video_finalize_strict_mode_missing_required_marks_failed(
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
db = _make_db()
|
db = _make_db()
|
||||||
task = _make_task(status=VideoStatus.COMPLETED.value)
|
task = _make_task(
|
||||||
|
status=VideoStatus.COMPLETED.value,
|
||||||
# provide a billing rule so formula path is taken
|
request_metadata={
|
||||||
rule = SimpleNamespace(
|
"poll_raw_response": {"foo": "bar"},
|
||||||
id="r1",
|
"billing_rule_snapshot": {
|
||||||
name="video",
|
"status": "ok",
|
||||||
expression="duration_seconds",
|
"rule_id": "r1",
|
||||||
variables={},
|
"rule_name": "video",
|
||||||
dimension_mappings={},
|
"scope": "model",
|
||||||
|
"expression": "duration_seconds",
|
||||||
|
"variables": {},
|
||||||
|
"dimension_mappings": {},
|
||||||
|
},
|
||||||
|
},
|
||||||
)
|
)
|
||||||
lookup = SimpleNamespace(rule=rule, scope="model", effective_task_type="video")
|
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.DimensionCollectorService.collect_dimensions",
|
"src.services.billing.dimension_collector_service.DimensionCollectorService.collect_dimensions",
|
||||||
lambda _self, **_kwargs: {"duration_seconds": None},
|
lambda _self, **_kwargs: {"duration_seconds": None},
|
||||||
)
|
)
|
||||||
|
# Mock update_settled_billing (used by finalize_video_task)
|
||||||
|
update_settled = MagicMock(return_value=True)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.task.impl.video_telemetry.BillingRuleService.find_rule",
|
"src.services.task.application.UsageService.update_settled_billing",
|
||||||
lambda *_args, **_kwargs: lookup,
|
update_settled,
|
||||||
)
|
|
||||||
record = AsyncMock()
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"src.services.task.impl.video_telemetry.UsageService.record_usage_with_custom_cost",
|
|
||||||
record,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
old = config.billing_strict_mode
|
old = config.billing_strict_mode
|
||||||
try:
|
try:
|
||||||
config.billing_strict_mode = True
|
config.billing_strict_mode = True
|
||||||
telemetry = VideoTelemetry(db)
|
monkeypatch.setattr(
|
||||||
telemetry._formula_engine.evaluate = MagicMock(
|
"src.services.task.application.FormulaEngine.evaluate",
|
||||||
side_effect=BillingIncompleteError(
|
MagicMock(
|
||||||
"Missing required dimensions", missing_required=["duration_seconds"]
|
side_effect=BillingIncompleteError(
|
||||||
)
|
"Missing required dimensions", missing_required=["duration_seconds"]
|
||||||
|
)
|
||||||
|
),
|
||||||
)
|
)
|
||||||
await telemetry.record_terminal_usage(task)
|
app = TaskApplicationService(db)
|
||||||
|
await app.finalize_video_task(task)
|
||||||
finally:
|
finally:
|
||||||
config.billing_strict_mode = old
|
config.billing_strict_mode = old
|
||||||
|
|
||||||
@@ -167,6 +186,6 @@ async def test_video_telemetry_strict_mode_missing_required_marks_failed(
|
|||||||
assert task.video_urls is None
|
assert task.video_urls is None
|
||||||
assert "billing_incomplete" in (task.error_code or "")
|
assert "billing_incomplete" in (task.error_code or "")
|
||||||
|
|
||||||
kwargs = record.await_args.kwargs
|
kwargs = update_settled.call_args.kwargs
|
||||||
assert kwargs["total_cost_usd"] == 0.0
|
assert kwargs["total_cost_usd"] == 0.0
|
||||||
assert kwargs["status"] == "failed"
|
assert kwargs["status"] == "failed"
|
||||||
@@ -4,14 +4,14 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
||||||
from src.services.video.task_poller import VideoTaskPollerService
|
from src.services.task.impl.video_poller import VideoTaskPollerAdapter
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_poll_task_status_routes_gemini_video_to_gemini(
|
async def test_poll_task_status_routes_gemini_video_to_gemini(
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
poller = VideoTaskPollerService()
|
adapter = VideoTaskPollerAdapter()
|
||||||
|
|
||||||
task = SimpleNamespace(
|
task = SimpleNamespace(
|
||||||
endpoint_id="e1",
|
endpoint_id="e1",
|
||||||
@@ -22,24 +22,24 @@ async def test_poll_task_status_routes_gemini_video_to_gemini(
|
|||||||
endpoint = SimpleNamespace(id="e1", base_url="https://example.com", api_format="gemini:video")
|
endpoint = SimpleNamespace(id="e1", base_url="https://example.com", api_format="gemini:video")
|
||||||
key = SimpleNamespace(id="k1", api_key="enc")
|
key = SimpleNamespace(id="k1", api_key="enc")
|
||||||
|
|
||||||
monkeypatch.setattr(poller, "_get_endpoint", lambda _db, _id: endpoint)
|
monkeypatch.setattr(adapter, "_get_endpoint", lambda _db, _id: endpoint)
|
||||||
monkeypatch.setattr(poller, "_get_key", lambda _db, _id: key)
|
monkeypatch.setattr(adapter, "_get_key", lambda _db, _id: key)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.video.task_poller.crypto_service.decrypt", lambda _v: "decrypted"
|
"src.services.task.impl.video_poller.crypto_service.decrypt", lambda _v: "decrypted"
|
||||||
)
|
)
|
||||||
|
|
||||||
auth_info = SimpleNamespace(auth_header="authorization", auth_value="Bearer x")
|
auth_info = SimpleNamespace(auth_header="authorization", auth_value="Bearer x")
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"src.services.video.task_poller.get_provider_auth",
|
"src.services.task.impl.video_poller.get_provider_auth",
|
||||||
AsyncMock(return_value=auth_info),
|
AsyncMock(return_value=auth_info),
|
||||||
)
|
)
|
||||||
|
|
||||||
poll_gemini = AsyncMock(return_value=InternalVideoPollResult(status=VideoStatus.PROCESSING))
|
poll_gemini = AsyncMock(return_value=InternalVideoPollResult(status=VideoStatus.PROCESSING))
|
||||||
poll_openai = AsyncMock(return_value=InternalVideoPollResult(status=VideoStatus.PROCESSING))
|
poll_openai = AsyncMock(return_value=InternalVideoPollResult(status=VideoStatus.PROCESSING))
|
||||||
monkeypatch.setattr(poller, "_poll_gemini", poll_gemini)
|
monkeypatch.setattr(adapter, "_poll_gemini", poll_gemini)
|
||||||
monkeypatch.setattr(poller, "_poll_openai", poll_openai)
|
monkeypatch.setattr(adapter, "_poll_openai", poll_openai)
|
||||||
|
|
||||||
result = await poller._poll_task_status(MagicMock(), task)
|
result = await adapter._poll_task_status(MagicMock(), task)
|
||||||
assert result.status == VideoStatus.PROCESSING
|
assert result.status == VideoStatus.PROCESSING
|
||||||
assert poll_gemini.await_count == 1
|
assert poll_gemini.await_count == 1
|
||||||
assert poll_openai.await_count == 0
|
assert poll_openai.await_count == 0
|
||||||
|
|||||||
Reference in New Issue
Block a user