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:
fawney19
2026-02-02 03:16:52 +08:00
parent feb7484fda
commit 9e31efe26c
75 changed files with 7511 additions and 2068 deletions

View File

@@ -1,25 +1,13 @@
"""
异步任务服务层
任务服务层Phase2
提供视频/图片/音频等异步任务的
- 提交阶段故障转移AsyncTaskOrchestrator
- 终态计费与 Usage 写入VideoTelemetry 等)
统一任务框架相关的应用层入口
- 候选提交阶段`services.candidate.CandidateService`
- 终态结算:`services.task.application.TaskApplicationService`
"""
from .orchestrator import (
AllCandidatesFailedError,
AsyncTaskOrchestrator,
CandidateSubmissionError,
CandidateUnsupportedError,
SubmitOutcome,
UpstreamClientRequestError,
)
from .application import TaskApplicationService
__all__ = [
"AsyncTaskOrchestrator",
"SubmitOutcome",
"AllCandidatesFailedError",
"UpstreamClientRequestError",
"CandidateUnsupportedError",
"CandidateSubmissionError",
"TaskApplicationService",
]

View 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

View 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)

View File

@@ -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__ = []

View 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])

View File

@@ -1,320 +0,0 @@
"""
VideoTelemetryPhase3
将 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"]

View 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"

View File

@@ -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
# 解析任务 ID200 但缺字段也视为失败并 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",
]

View 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
# ========== 阶段 2HTTP 请求(不持有数据库连接)==========
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