mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
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:
@@ -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__)
|
||||
|
||||
@@ -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
|
||||
24
src/services/task/exceptions.py
Normal file
24
src/services/task/exceptions.py
Normal 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
|
||||
@@ -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)",
|
||||
|
||||
51
src/services/task/protocol.py
Normal file
51
src/services/task/protocol.py
Normal 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: ...
|
||||
72
src/services/task/schema.py
Normal file
72
src/services/task/schema.py
Normal 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
1832
src/services/task/service.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user