refactor: 统一任务框架 Phase 3 - 用 TaskService/FailoverEngine 替代 FallbackOrchestrator

核心重构:
- 移除 FallbackOrchestrator,用 TaskService + FailoverEngine 替代
- TaskService 作为统一入口,支持 SYNC/ASYNC 两种任务模式
- FailoverEngine 实现候选遍历、重试、故障转移逻辑
- 新增 AttemptFunc/AttemptResult 协议,统一尝试结果表示

功能改进:
- 流式响应首字节探测(30s 超时,空流触发故障转移)
- 流式取消归因优化(区分客户端断连 vs 服务端中断)
- 新增 OpenAI Sora 视频取消路由 POST /v1/videos/{task_id}/cancel
- OpenAI 流式请求自动添加 stream_options.include_usage

代码规范:
- 修复 loguru 日志格式(%s → {})
- 新增 FORMAT_CONVERSION_ENABLED 环境变量说明

测试覆盖:
- test_failover_engine.py: FailoverEngine 单元测试
- test_task_service_async_execute.py: TaskService ASYNC 模式测试
- test_video_cancel_e2e.py: 视频取消端到端测试
This commit is contained in:
fawney19
2026-02-02 21:16:28 +08:00
parent ebe1a8d3e3
commit ed68aebfb0
53 changed files with 3966 additions and 2256 deletions

View File

@@ -1,13 +1,32 @@
"""
任务服务层Phase2
任务服务层Phase2/Phase3
统一任务框架相关的应用层入口:
- 候选提交阶段`services.candidate.CandidateService`
- 终态结算:`services.task.application.TaskApplicationService`
- 候选域能力`services.candidate.CandidateService`resolve/record 等)
- 终态结算:`services.task.service.TaskService.finalize_video_task`
- 统一门面:`services.task.service.TaskService`
"""
from .application import TaskApplicationService
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING: # pragma: no cover
# NOTE: keep package import lightweight; avoid importing heavy modules on submodule imports
from .service import TaskService as TaskService
__all__ = [
"TaskApplicationService",
"TaskService",
]
def __getattr__(name: str) -> type: # pragma: no cover
if name == "TaskService":
from .service import TaskService
return TaskService
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
def __dir__() -> list[str]: # pragma: no cover
return sorted(__all__)

View File

@@ -1,312 +0,0 @@
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,24 @@
from __future__ import annotations
class StreamProbeError(RuntimeError):
"""Streaming probe failed before first chunk (eligible for failover)."""
def __init__(
self,
message: str,
*,
http_status: int,
original_exception: Exception | None = None,
) -> None:
super().__init__(message)
self.http_status = http_status
self.original_exception = original_exception
class TaskNotFoundError(LookupError):
"""Task not found (by internal id or external id)."""
def __init__(self, task_id: str) -> None:
super().__init__(f"Task not found: {task_id}")
self.task_id = task_id

View File

@@ -35,7 +35,7 @@ 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
from src.services.task.service import TaskService
@dataclass(slots=True)
@@ -164,7 +164,7 @@ class VideoTaskPollerAdapter:
try:
upstream_key = crypto_service.decrypt(key.api_key)
except Exception:
logger.warning("Failed to decrypt provider key for task %s", task.id)
logger.warning("Failed to decrypt provider key for task {}", task.id)
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="decryption_error",
@@ -241,7 +241,7 @@ class VideoTaskPollerAdapter:
with create_session() as db:
task = db.get(VideoTask, task_id)
if not task:
logger.warning("Task %s disappeared during poll update", task_id)
logger.warning("Task {} disappeared during poll update", task_id)
return
if error_exception is not None and ctx is not None:
@@ -284,12 +284,10 @@ class VideoTaskPollerAdapter:
# 终态结算
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
task
)
await TaskService(db, redis_client=redis_client).finalize_video_task(task)
except Exception as exc:
logger.exception(
"Failed to record video usage for task=%s: %s",
"Failed to record video usage for task={}: {}",
task.id,
sanitize_error_message(str(exc)),
)
@@ -300,7 +298,7 @@ class VideoTaskPollerAdapter:
"""处理轮询错误"""
task.poll_count += 1
error_msg = sanitize_error_message(str(exc))
logger.warning("Poll error for task %s: %s", task.id, error_msg)
logger.warning("Poll error for task {}: {}", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
@@ -337,7 +335,7 @@ class VideoTaskPollerAdapter:
url = self._build_gemini_url(ctx.base_url, operation_name)
logger.debug(
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
"[VideoPoller] Gemini poll: task={} external_id={} url={}",
ctx.task_id,
ctx.external_task_id,
url,
@@ -347,7 +345,7 @@ class VideoTaskPollerAdapter:
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",
"[VideoPoller] Gemini poll failed: task={} status={} response={}",
ctx.task_id,
response.status_code,
response.text[:500] if response.text else "(empty)",
@@ -394,7 +392,7 @@ class VideoTaskPollerAdapter:
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)
logger.warning("Poll error for task {}: {}", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
@@ -427,12 +425,10 @@ class VideoTaskPollerAdapter:
# 终态结算
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
task
)
await TaskService(db, redis_client=redis_client).finalize_video_task(task)
except Exception as exc:
logger.exception(
"Failed to record video usage for task=%s: %s",
"Failed to record video usage for task={}: {}",
task.id,
sanitize_error_message(str(exc)),
)
@@ -470,7 +466,7 @@ class VideoTaskPollerAdapter:
try:
upstream_key = crypto_service.decrypt(key.api_key)
except Exception:
logger.warning("Failed to decrypt provider key for task %s", task.id)
logger.warning("Failed to decrypt provider key for task {}", task.id)
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="decryption_error",
@@ -539,7 +535,7 @@ class VideoTaskPollerAdapter:
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
logger.debug(
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
"[VideoPoller] Gemini poll: task={} external_id={} url={}",
task.id,
task.external_task_id,
url,
@@ -549,7 +545,7 @@ class VideoTaskPollerAdapter:
response = await client.get(url, headers=headers)
if response.status_code >= 400:
logger.warning(
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
"[VideoPoller] Gemini poll failed: task={} status={} response={}",
task.id,
response.status_code,
response.text[:500] if response.text else "(empty)",

View File

@@ -0,0 +1,51 @@
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from typing import Any, AsyncIterator, Protocol, runtime_checkable
import httpx
from src.services.cache.aware_scheduler import ProviderCandidate
class AttemptKind(str, Enum):
"""`attempt_func` return kind."""
SYNC_RESPONSE = "sync_response"
STREAM = "stream"
ASYNC_SUBMIT = "async_submit"
@dataclass(slots=True)
class AttemptResult:
"""
Unified attempt result returned by `AttemptFunc`.
Notes:
- `http_status` / `http_headers` MUST be filled for all kinds (for audit/classification).
- The payload fields are filled depending on `kind`.
"""
kind: AttemptKind
# HTTP meta (always filled)
http_status: int
http_headers: dict[str, str]
# SYNC_RESPONSE
response_body: Any = None
# STREAM
stream_iterator: AsyncIterator[bytes] | None = None
# ASYNC_SUBMIT
provider_task_id: str | None = None
# Raw response reference (optional, for audit/debugging)
raw_response: httpx.Response | None = None
@runtime_checkable
class AttemptFunc(Protocol):
async def __call__(self, candidate: ProviderCandidate) -> AttemptResult: ...

View File

@@ -0,0 +1,72 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.candidate.schema import CandidateKey
from .protocol import AttemptKind, AttemptResult
@dataclass(slots=True)
class ExecutionResult:
"""FailoverEngine.execute() unified result."""
success: bool
# payload (filled based on AttemptKind)
attempt_result: AttemptResult | None = None
# selected candidate
candidate: ProviderCandidate | None = None
candidate_index: int = -1
retry_index: int = 0
provider_id: str | None = None
provider_name: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
# audit
candidate_keys: list[CandidateKey] = field(default_factory=list)
attempt_count: int = 0
request_candidate_id: str | None = None
# failure
error_type: str | None = None
error_message: str | None = None
last_status_code: int | None = None
@property
def response(self) -> Any:
"""Compatibility accessor: returns response body or stream iterator."""
if not self.attempt_result:
return None
if self.attempt_result.kind == AttemptKind.STREAM:
return self.attempt_result.stream_iterator
return self.attempt_result.response_body
@property
def provider_task_id(self) -> str | None:
if self.attempt_result and self.attempt_result.kind == AttemptKind.ASYNC_SUBMIT:
return self.attempt_result.provider_task_id
return None
@dataclass(slots=True)
class TaskStatusResult:
"""Generic task status payload returned by TaskService.poll()."""
task_id: str
status: str
progress_percent: int | None = None
result_url: str | None = None
error_message: str | None = None
# optional metadata (best-effort)
provider_id: str | None = None
provider_name: str | None = None
endpoint_id: str | None = None
key_id: str | None = None

1832
src/services/task/service.py Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -133,7 +133,7 @@ class TaskPollerService:
task_obj = self.adapter.get_task(task_db, task_id)
if not task_obj:
logger.warning(
"[%s] Task %s disappeared during poll",
"[{}] Task {} disappeared during poll",
self.adapter.task_type,
task_id,
)
@@ -181,7 +181,7 @@ class TaskPollerService:
poll_results.append(True)
except Exception as exc:
logger.exception(
"[%s] Unexpected error polling task %s: %s",
"[{}] Unexpected error polling task {}: {}",
self.adapter.task_type,
task_id,
self.adapter.sanitize_error_message(str(exc)),
@@ -200,7 +200,7 @@ class TaskPollerService:
>= self.adapter.consecutive_failure_alert_threshold
):
logger.error(
"[ALERT] %s poller: %d consecutive batches failed.",
"[ALERT] {} poller: {} consecutive batches failed.",
self.adapter.task_type,
self._consecutive_failures,
)