refactor(task): 拆分 TaskService 并重构任务生命周期

- 将 task 公共协议/上下文/异常/schema 下沉到 core,并迁移 polling 目录
- 新增 execute/submit/video 子模块,拆分同步执行、异步提交流程、错误处理与视频任务操作
- 收敛 TaskService 为门面编排,内部委派到 SyncTaskExecutionService、AsyncTaskSubmitService、VideoTaskOperationsService
- 重构 main 生命周期管理:引入 LifecycleState,拆分启动与关闭流程
- 删除未使用的 TaskExecuteFacadeService 与 TaskSubmitFacadeService

Closes #201

Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-03-07 15:33:29 +08:00
parent 4cd6e0d10f
commit 239238fe47
47 changed files with 3988 additions and 2447 deletions

View File

@@ -1956,11 +1956,14 @@ async def test_model_failover(
- direct: 直接测试 provider_model_name在当前 Provider 内多 Key 故障转移 - direct: 直接测试 provider_model_name在当前 Provider 内多 Key 故障转移
""" """
from src.core.exceptions import ProviderNotAvailableException from src.core.exceptions import ProviderNotAvailableException
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
from src.services.candidate.recorder import CandidateRecorder from src.services.candidate.recorder import CandidateRecorder
from src.services.scheduling.candidate_builder import CandidateBuilder from src.services.scheduling.candidate_builder import CandidateBuilder
from src.services.scheduling.candidate_sorter import CandidateSorter from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.scheduling_config import SchedulingConfig from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.task import TaskService from src.services.task import TaskService
from src.services.task.core.protocol import AttemptKind, AttemptResult
provider = ( provider = (
db.query(Provider) db.query(Provider)

View File

@@ -557,7 +557,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 统一入口:总是通过 TaskService # 统一入口:总是通过 TaskService
from src.services.task import TaskService from src.services.task import TaskService
from src.services.task.context import TaskMode from src.services.task.core.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute( exec_result = await TaskService(self.db, self.redis).execute(
task_type="chat", task_type="chat",

View File

@@ -161,7 +161,7 @@ class ChatSyncExecutor:
# 统一入口:总是通过 TaskService # 统一入口:总是通过 TaskService
from src.services.task import TaskService from src.services.task import TaskService
from src.services.task.context import TaskMode from src.services.task.core.context import TaskMode
exec_result = await TaskService(handler.db, handler.redis).execute( exec_result = await TaskService(handler.db, handler.redis).execute(
task_type="chat", task_type="chat",

View File

@@ -162,7 +162,7 @@ class CliStreamMixin:
# 统一入口:总是通过 TaskService # 统一入口:总是通过 TaskService
from src.services.task import TaskService from src.services.task import TaskService
from src.services.task.context import TaskMode from src.services.task.core.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute( exec_result = await TaskService(self.db, self.redis).execute(
task_type="cli", task_type="cli",

View File

@@ -482,7 +482,7 @@ class CliSyncMixin:
# 统一入口:总是通过 TaskService # 统一入口:总是通过 TaskService
from src.services.task import TaskService from src.services.task import TaskService
from src.services.task.context import TaskMode from src.services.task.core.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute( exec_result = await TaskService(self.db, self.redis).execute(
task_type="cli", task_type="cli",

View File

@@ -6,7 +6,8 @@
from __future__ import annotations from __future__ import annotations
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Any from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
import uvicorn import uvicorn
from fastapi import FastAPI, HTTPException from fastapi import FastAPI, HTTPException
@@ -32,6 +33,20 @@ from src.database import init_db
from src.middleware.plugin_middleware import PluginMiddleware from src.middleware.plugin_middleware import PluginMiddleware
from src.plugins.manager import get_plugin_manager from src.plugins.manager import get_plugin_manager
if TYPE_CHECKING:
from redis.asyncio.client import Redis
from src.core.modules.base import ModuleDefinition
from src.plugins.manager import PluginManager
from src.services.model.fetch_scheduler import ModelFetchScheduler
from src.services.provider_keys.pool_quota_probe_scheduler import PoolQuotaProbeScheduler
from src.services.rate_limit.concurrency_manager import ConcurrencyManager
from src.services.system.maintenance_scheduler import MaintenanceScheduler
from src.services.system.scheduler import TaskScheduler
from src.services.task.polling.task_poller import TaskPollerService
from src.services.usage.quota_scheduler import QuotaScheduler
from src.utils.task_coordinator import StartupTaskCoordinator
async def initialize_providers() -> None: async def initialize_providers() -> None:
"""从数据库初始化提供商(仅用于日志记录)""" """从数据库初始化提供商(仅用于日志记录)"""
@@ -76,23 +91,41 @@ async def initialize_providers() -> None:
logger.exception("从数据库初始化提供商失败") logger.exception("从数据库初始化提供商失败")
@asynccontextmanager @dataclass
async def lifespan(app: FastAPI) -> Any: class LifecycleState:
"""应用生命周期管理""" """应用生命周期阶段共享的运行时状态。"""
# 禁用uvicorn的access日志(在子进程中执行)
redis_client: Redis | None = None
concurrency_manager: ConcurrencyManager | None = None
plugin_manager: PluginManager | None = None
available_modules: list[ModuleDefinition] = field(default_factory=list)
task_coordinator: StartupTaskCoordinator | None = None
quota_scheduler: QuotaScheduler | None = None
maintenance_scheduler: MaintenanceScheduler | None = None
model_fetch_scheduler: ModelFetchScheduler | None = None
pool_quota_probe_scheduler: PoolQuotaProbeScheduler | None = None
task_poller: TaskPollerService | None = None
task_scheduler: TaskScheduler | None = None
def _configure_uvicorn_access_log() -> None:
"""禁用 uvicorn access 日志(在子进程中执行)。"""
import logging import logging
logging.getLogger("uvicorn.access").setLevel(logging.CRITICAL) logging.getLogger("uvicorn.access").setLevel(logging.CRITICAL)
logging.getLogger("uvicorn.access").disabled = True logging.getLogger("uvicorn.access").disabled = True
# 启动时执行
def _log_startup_banner() -> None:
logger.info("=" * 60) logger.info("=" * 60)
from src import __version__ from src import __version__
logger.info(f"AI Proxy v{__version__} - GlobalModel Architecture") logger.info(f"AI Proxy v{__version__} - GlobalModel Architecture")
logger.info("=" * 60) logger.info("=" * 60)
# 安全配置验证(生产环境会阻止启动)
def _validate_security_or_raise() -> None:
"""启动前安全配置校验。"""
security_errors = config.validate_security_config() security_errors = config.validate_security_config()
if security_errors: if security_errors:
for error in security_errors: for error in security_errors:
@@ -104,6 +137,9 @@ async def lifespan(app: FastAPI) -> Any:
+ "\n".join(f" - {e}" for e in security_errors) + "\n".join(f" - {e}" for e in security_errors)
) )
async def _initialize_core_infrastructure(state: LifecycleState) -> None:
"""初始化数据库、缓存、并发与基础后台组件。"""
# 记录启动警告密码、连接池、JWT 等) # 记录启动警告密码、连接池、JWT 等)
config.log_startup_warnings() config.log_startup_warnings()
@@ -122,10 +158,9 @@ async def lifespan(app: FastAPI) -> Any:
logger.info("初始化全局Redis客户端...") logger.info("初始化全局Redis客户端...")
from src.clients.redis_client import get_redis_client from src.clients.redis_client import get_redis_client
redis_client = None
try: try:
redis_client = await get_redis_client(require_redis=config.require_redis) state.redis_client = await get_redis_client(require_redis=config.require_redis)
if redis_client: if state.redis_client:
logger.info("[OK] Redis客户端初始化成功缓存亲和性功能已启用") logger.info("[OK] Redis客户端初始化成功缓存亲和性功能已启用")
else: else:
logger.warning( logger.warning(
@@ -136,13 +171,13 @@ async def lifespan(app: FastAPI) -> Any:
logger.exception("[ERROR] Redis连接失败应用启动中止") logger.exception("[ERROR] Redis连接失败应用启动中止")
raise raise
logger.warning(f"Redis连接失败但配置允许降级将继续使用内存模式: {e}") logger.warning(f"Redis连接失败但配置允许降级将继续使用内存模式: {e}")
redis_client = None state.redis_client = None
# 初始化并发管理器内部会使用Redis # 初始化并发管理器内部会使用Redis
logger.info("初始化并发管理器...") logger.info("初始化并发管理器...")
from src.services.rate_limit.concurrency_manager import get_concurrency_manager from src.services.rate_limit.concurrency_manager import get_concurrency_manager
concurrency_manager = await get_concurrency_manager() state.concurrency_manager = await get_concurrency_manager()
# 初始化批量提交器(提升数据库并发能力) # 初始化批量提交器(提升数据库并发能力)
logger.info("初始化批量提交器...") logger.info("初始化批量提交器...")
@@ -167,10 +202,13 @@ async def lifespan(app: FastAPI) -> Any:
await start_usage_queue_consumer() await start_usage_queue_consumer()
async def _initialize_plugins_and_modules(app: FastAPI, state: LifecycleState) -> None:
"""初始化插件系统、模块系统,并注册路由与钩子。"""
# 初始化插件系统 # 初始化插件系统
logger.info("初始化插件系统...") logger.info("初始化插件系统...")
plugin_manager = get_plugin_manager() state.plugin_manager = get_plugin_manager()
init_results = await plugin_manager.initialize_all() init_results = await state.plugin_manager.initialize_all()
successful = sum(1 for success in init_results.values() if success) successful = sum(1 for success in init_results.values() if success)
logger.info(f"插件初始化完成: {successful}/{len(init_results)} 个插件成功启动") logger.info(f"插件初始化完成: {successful}/{len(init_results)} 个插件成功启动")
@@ -204,8 +242,8 @@ async def lifespan(app: FastAPI) -> Any:
# 注册可用模块的路由 # 注册可用模块的路由
# 注意:模块的 router 自带 prefixapi_prefix 字段仅用于日志和文档 # 注意:模块的 router 自带 prefixapi_prefix 字段仅用于日志和文档
available_modules = module_registry.get_available_modules() state.available_modules = module_registry.get_available_modules()
for module in available_modules: for module in state.available_modules:
if module.router_factory: if module.router_factory:
router = module.router_factory() router = module.router_factory()
app.include_router(router) app.include_router(router)
@@ -216,7 +254,7 @@ async def lifespan(app: FastAPI) -> Any:
if module.on_startup: if module.on_startup:
await module.on_startup() await module.on_startup()
logger.info(f"功能模块初始化完成: {len(available_modules)}/{len(ALL_MODULES)} 个模块可用") logger.info(f"功能模块初始化完成: {len(state.available_modules)}/{len(ALL_MODULES)} 个模块可用")
# 显式 bootstrap provider plugins注册 envelope/enricher 等) # 显式 bootstrap provider plugins注册 envelope/enricher 等)
# 使 core/provider_oauth_utils 不需要在运行时 lazy import services 层 # 使 core/provider_oauth_utils 不需要在运行时 lazy import services 层
@@ -229,9 +267,9 @@ async def lifespan(app: FastAPI) -> Any:
register_default_parsers() register_default_parsers()
logger.info(f"服务启动成功: http://{config.host}:{config.port}")
logger.info("=" * 60)
async def _start_background_services(state: LifecycleState) -> None:
"""启动调度器与后台轮询服务。"""
# 启动月卡额度重置调度器(仅一个 worker 执行) # 启动月卡额度重置调度器(仅一个 worker 执行)
logger.info("启动月卡额度重置调度器...") logger.info("启动月卡额度重置调度器...")
from src.services.model.fetch_scheduler import get_model_fetch_scheduler from src.services.model.fetch_scheduler import get_model_fetch_scheduler
@@ -239,75 +277,92 @@ async def lifespan(app: FastAPI) -> Any:
get_pool_quota_probe_scheduler, get_pool_quota_probe_scheduler,
) )
from src.services.system.maintenance_scheduler import get_maintenance_scheduler from src.services.system.maintenance_scheduler import get_maintenance_scheduler
from src.services.task.task_poller import get_task_poller from src.services.task.polling.task_poller import get_task_poller
from src.services.usage.quota_scheduler import get_quota_scheduler from src.services.usage.quota_scheduler import get_quota_scheduler
from src.utils.task_coordinator import StartupTaskCoordinator from src.utils.task_coordinator import StartupTaskCoordinator
quota_scheduler = get_quota_scheduler() state.quota_scheduler = get_quota_scheduler()
maintenance_scheduler = get_maintenance_scheduler() state.maintenance_scheduler = get_maintenance_scheduler()
model_fetch_scheduler = get_model_fetch_scheduler() state.model_fetch_scheduler = get_model_fetch_scheduler()
pool_quota_probe_scheduler = get_pool_quota_probe_scheduler() state.pool_quota_probe_scheduler = get_pool_quota_probe_scheduler()
task_poller = get_task_poller() state.task_poller = get_task_poller()
task_coordinator = StartupTaskCoordinator(redis_client) state.task_coordinator = StartupTaskCoordinator(state.redis_client)
# 启动额度调度器 # 启动额度调度器
quota_scheduler_active = await task_coordinator.acquire("quota_scheduler") quota_scheduler_active = await state.task_coordinator.acquire("quota_scheduler")
if quota_scheduler_active: if quota_scheduler_active:
await quota_scheduler.start() await state.quota_scheduler.start()
else: else:
logger.info("检测到其他 worker 已运行额度调度器,本实例跳过") logger.info("检测到其他 worker 已运行额度调度器,本实例跳过")
quota_scheduler = None # type: ignore[assignment] state.quota_scheduler = None
# 启动维护调度器 # 启动维护调度器
maintenance_scheduler_active = await task_coordinator.acquire("maintenance_scheduler") maintenance_scheduler_active = await state.task_coordinator.acquire("maintenance_scheduler")
if maintenance_scheduler_active: if maintenance_scheduler_active:
logger.info("启动系统维护调度器...") logger.info("启动系统维护调度器...")
await maintenance_scheduler.start() await state.maintenance_scheduler.start()
else: else:
logger.info("检测到其他 worker 已运行维护调度器,本实例跳过") logger.info("检测到其他 worker 已运行维护调度器,本实例跳过")
maintenance_scheduler = None # type: ignore[assignment] state.maintenance_scheduler = None
# 启动模型自动获取调度器 # 启动模型自动获取调度器
model_fetch_scheduler_active = await task_coordinator.acquire("model_fetch_scheduler") model_fetch_scheduler_active = await state.task_coordinator.acquire("model_fetch_scheduler")
if model_fetch_scheduler_active: if model_fetch_scheduler_active:
logger.info("启动模型自动获取调度器...") logger.info("启动模型自动获取调度器...")
await model_fetch_scheduler.start() await state.model_fetch_scheduler.start()
else: else:
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过") logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
model_fetch_scheduler = None # type: ignore[assignment] state.model_fetch_scheduler = None
# 启动号池额度主动探测调度器 # 启动号池额度主动探测调度器
pool_quota_probe_scheduler_active = await task_coordinator.acquire("pool_quota_probe_scheduler") pool_quota_probe_scheduler_active = await state.task_coordinator.acquire("pool_quota_probe_scheduler")
if pool_quota_probe_scheduler_active: if pool_quota_probe_scheduler_active:
logger.info("启动号池额度主动探测调度器...") logger.info("启动号池额度主动探测调度器...")
await pool_quota_probe_scheduler.start() await state.pool_quota_probe_scheduler.start()
else: else:
logger.info("检测到其他 worker 已运行号池额度主动探测调度器,本实例跳过") logger.info("检测到其他 worker 已运行号池额度主动探测调度器,本实例跳过")
pool_quota_probe_scheduler = None # type: ignore[assignment] state.pool_quota_probe_scheduler = None
# 启动异步任务轮询服务(当前仅视频) # 启动异步任务轮询服务(当前仅视频)
task_poller_active = await task_coordinator.acquire("task_poller:video") task_poller_active = await state.task_coordinator.acquire("task_poller:video")
if task_poller_active: if task_poller_active:
logger.info("启动 TaskPollervideo...") logger.info("启动 TaskPollervideo...")
await task_poller.start() await state.task_poller.start()
else: else:
logger.info("检测到其他 worker 已运行 TaskPollervideo本实例跳过") logger.info("检测到其他 worker 已运行 TaskPollervideo本实例跳过")
task_poller = None # type: ignore[assignment] state.task_poller = None
# 启动统一的定时任务调度器 # 启动统一的定时任务调度器
from src.services.system.scheduler import get_scheduler from src.services.system.scheduler import get_scheduler
task_scheduler = get_scheduler() state.task_scheduler = get_scheduler()
task_scheduler.start() state.task_scheduler.start()
# 启动缓存预热(后台任务,不阻塞启动) # 启动缓存预热(后台任务,不阻塞启动)
from src.services.system.cache_warmup import start_cache_warmup from src.services.system.cache_warmup import start_cache_warmup
await start_cache_warmup() await start_cache_warmup()
yield # 应用运行期间
# 关闭时执行 async def _run_startup(app: FastAPI) -> LifecycleState:
"""执行完整启动流程并返回生命周期状态。"""
_configure_uvicorn_access_log()
_log_startup_banner()
_validate_security_or_raise()
state = LifecycleState()
await _initialize_core_infrastructure(state)
await _initialize_plugins_and_modules(app, state)
logger.info(f"服务启动成功: http://{config.host}:{config.port}")
logger.info("=" * 60)
await _start_background_services(state)
return state
async def _run_shutdown(state: LifecycleState) -> None:
"""执行完整关闭流程。"""
logger.info("正在关闭服务...") logger.info("正在关闭服务...")
# 停止 Codex 配额异步同步器(停止前会 flush 待同步事件) # 停止 Codex 配额异步同步器(停止前会 flush 待同步事件)
@@ -334,52 +389,58 @@ async def lifespan(app: FastAPI) -> Any:
await stop_usage_queue_consumer() await stop_usage_queue_consumer()
# 停止维护调度器 # 停止维护调度器
if maintenance_scheduler: if state.maintenance_scheduler:
logger.info("停止系统维护调度器...") logger.info("停止系统维护调度器...")
await maintenance_scheduler.stop() await state.maintenance_scheduler.stop()
await task_coordinator.release("maintenance_scheduler") if state.task_coordinator:
await state.task_coordinator.release("maintenance_scheduler")
# 停止月卡额度重置调度器,并释放分布式锁 # 停止月卡额度重置调度器,并释放分布式锁
logger.info("停止月卡额度重置调度器...") if state.quota_scheduler:
if quota_scheduler: logger.info("停止月卡额度重置调度器...")
await quota_scheduler.stop() await state.quota_scheduler.stop()
if task_coordinator: if state.task_coordinator:
await task_coordinator.release("quota_scheduler") await state.task_coordinator.release("quota_scheduler")
# 停止模型自动获取调度器 # 停止模型自动获取调度器
if model_fetch_scheduler: if state.model_fetch_scheduler:
logger.info("停止模型自动获取调度器...") logger.info("停止模型自动获取调度器...")
await model_fetch_scheduler.stop() await state.model_fetch_scheduler.stop()
await task_coordinator.release("model_fetch_scheduler") if state.task_coordinator:
await state.task_coordinator.release("model_fetch_scheduler")
if pool_quota_probe_scheduler: if state.pool_quota_probe_scheduler:
logger.info("停止号池额度主动探测调度器...") logger.info("停止号池额度主动探测调度器...")
await pool_quota_probe_scheduler.stop() await state.pool_quota_probe_scheduler.stop()
await task_coordinator.release("pool_quota_probe_scheduler") if state.task_coordinator:
await state.task_coordinator.release("pool_quota_probe_scheduler")
if task_poller: if state.task_poller:
logger.info("停止 TaskPollervideo...") logger.info("停止 TaskPollervideo...")
await task_poller.stop() await state.task_poller.stop()
await task_coordinator.release("task_poller:video") if state.task_coordinator:
await state.task_coordinator.release("task_poller:video")
# 停止统一的定时任务调度器 # 停止统一的定时任务调度器
logger.info("停止定时任务调度器...") logger.info("停止定时任务调度器...")
task_scheduler.stop() if state.task_scheduler:
state.task_scheduler.stop()
# 关闭插件系统 # 关闭插件系统
logger.info("关闭插件系统...") logger.info("关闭插件系统...")
await plugin_manager.shutdown_all() if state.plugin_manager:
await state.plugin_manager.shutdown_all()
# 关闭功能模块 # 关闭功能模块
logger.info("关闭功能模块...") logger.info("关闭功能模块...")
for module in available_modules: for module in state.available_modules:
if module.on_shutdown: if module.on_shutdown:
await module.on_shutdown() await module.on_shutdown()
# 关闭并发管理器 # 关闭并发管理器
logger.info("关闭并发管理器...") logger.info("关闭并发管理器...")
if concurrency_manager: if state.concurrency_manager:
await concurrency_manager.close() await state.concurrency_manager.close()
# 关闭全局Redis客户端 # 关闭全局Redis客户端
logger.info("关闭全局Redis客户端...") logger.info("关闭全局Redis客户端...")
@@ -394,6 +455,16 @@ async def lifespan(app: FastAPI) -> Any:
logger.info("服务已关闭") logger.info("服务已关闭")
@asynccontextmanager
async def lifespan(app: FastAPI) -> Any:
"""应用生命周期管理"""
state = await _run_startup(app)
try:
yield # 应用运行期间
finally:
await _run_shutdown(state)
from src import __version__ as app_version from src import __version__ as app_version
# OpenAPI Tags 元数据定义 # OpenAPI Tags 元数据定义

View File

@@ -16,9 +16,9 @@ from src.models.database import RequestCandidate
from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
from src.services.request.candidate import RequestCandidateService from src.services.request.candidate import RequestCandidateService
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
from src.services.task.exceptions import StreamProbeError from src.services.task.core.exceptions import StreamProbeError
from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult from src.services.task.core.protocol import AttemptFunc, AttemptKind, AttemptResult
from src.services.task.schema import ExecutionResult from src.services.task.core.schema import ExecutionResult
from .policy import FailoverAction, RetryMode, RetryPolicy, SkipPolicy from .policy import FailoverAction, RetryMode, RetryPolicy, SkipPolicy
from .recorder import CandidateRecorder from .recorder import CandidateRecorder

View File

View File

View File

@@ -0,0 +1,512 @@
from __future__ import annotations
from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.error_utils import extract_error_message
from src.core.exceptions import (
ConcurrencyLimitError,
EmbeddedErrorException,
ProxyNodeUnavailableError,
ThinkingSignatureException,
UpstreamClientException,
)
from src.core.logger import logger
from src.core.provider_types import ProviderType
from src.services.request.candidate import RequestCandidateService
from src.services.task.execute.pool import TaskPoolOperationsService
class TaskErrorOperationsService:
"""任务执行错误处理服务(候选失败分类、整流与状态回写)。"""
def __init__(self, db: Session, *, pool_ops: TaskPoolOperationsService) -> None:
self.db = db
self._pool_ops = pool_ops
def mark_thinking_error_failed(
self,
candidate_record_id: str,
error: Any,
elapsed_ms: int,
captured_key_concurrent: int | None,
extra_data: dict[str, Any],
) -> None:
"""Mark ThinkingSignatureException as failed for the candidate."""
if not isinstance(error, ThinkingSignatureException):
return
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="ThinkingSignatureException",
error_message=str(error),
status_code=400,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=extra_data,
)
def handle_thinking_signature_error(
self,
*,
converted_error: Any,
provider_type: str | None,
model_name: str | None,
request_id: str | None,
candidate_record_id: str,
elapsed_ms: int,
captured_key_concurrent: int | None,
serializable_extra_data: dict[str, Any],
request_body_ref: dict[str, Any] | None,
) -> str:
"""Try to rectify thinking signature errors and request a retry."""
from src.services.message.thinking_rectifier import ThinkingRectifier
if not isinstance(converted_error, ThinkingSignatureException):
raise converted_error
if not config.thinking_rectifier_enabled:
logger.info(" [{}] Thinking 错误:整流器已禁用,终止重试", request_id)
self.mark_thinking_error_failed(
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
serializable_extra_data,
)
raise converted_error
if request_body_ref is None:
logger.warning(" [{}] Thinking 错误:无法获取请求体引用,终止重试", request_id)
self.mark_thinking_error_failed(
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
serializable_extra_data,
)
raise converted_error
provider_type_norm = str(provider_type or "").lower()
# Rectification may have multiple stages (Antigravity only).
stage_raw = request_body_ref.get("_rectify_stage", 0)
try:
stage = int(stage_raw or 0)
except Exception:
stage = 0
if stage <= 0 and request_body_ref.get("_rectified", False):
stage = 1
if stage >= 2 or (stage >= 1 and provider_type_norm != ProviderType.ANTIGRAVITY):
logger.warning(" [{}] Thinking 错误:已整流仍失败,终止重试", request_id)
self.mark_thinking_error_failed(
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
{**serializable_extra_data, "rectified": True, "rectify_stage": stage},
)
raise converted_error
request_body = request_body_ref.get("body", {})
stage_label = "thinking_only"
next_stage = 1
if stage == 0:
rectified_body, modified = ThinkingRectifier.rectify(request_body)
stage_label = "thinking_only"
next_stage = 1
else:
# Stage 2 only applies to Antigravity.
rectified_body, modified = ThinkingRectifier.rectify_signature_sensitive_blocks(
request_body
)
stage_label = "thinking_and_tools"
next_stage = 2
if modified:
request_body_ref["body"] = rectified_body
request_body_ref["_rectified"] = True
request_body_ref["_rectified_this_turn"] = True
request_body_ref["_rectify_stage"] = next_stage
if provider_type_norm == ProviderType.ANTIGRAVITY:
try:
from src.core.metrics import antigravity_degradation_total
antigravity_degradation_total.labels(
stage=stage_label,
model=str(model_name or "unknown"),
).inc()
except Exception:
pass
logger.info(
" [{}] 请求已整流(stage={}),在当前候选上重试",
request_id,
next_stage,
)
self.mark_thinking_error_failed(
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
{
**serializable_extra_data,
"rectified": True,
"rectify_stage": next_stage,
"rectify_stage_label": stage_label,
},
)
return "continue"
logger.warning(" [{}] Thinking 错误:无可整流内容", request_id)
self.mark_thinking_error_failed(
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
serializable_extra_data,
)
raise converted_error
async def handle_candidate_error(
self,
*,
exec_err: Any,
candidate: Any,
candidate_record_id: str,
retry_index: int,
max_retries_for_candidate: int,
affinity_key: str,
api_format: str,
global_model_id: str,
request_id: str | None,
attempt: int,
max_attempts: int,
error_classifier: Any,
request_body_ref: dict[str, Any] | None = None,
) -> str:
"""
Handle an execution error for a candidate.
Returns:
- "continue": retry current candidate
- "break": move to next candidate
- "raise": raise the underlying exception
"""
from src.core.api_format.conversion.exceptions import FormatConversionError
from src.services.proxy_node.resolver import resolve_effective_proxy, resolve_proxy_info
from src.services.request.executor import ExecutionError
# 提前解析代理信息,写入候选记录的 extra_data用于链路追踪展示
_eff_proxy = resolve_effective_proxy(
getattr(candidate.provider, "proxy", None),
getattr(candidate.key, "proxy", None),
)
_proxy_info = resolve_proxy_info(_eff_proxy)
_proxy_extra: dict[str, Any] | None = {"proxy": _proxy_info} if _proxy_info else None
if not isinstance(exec_err, ExecutionError):
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type=type(exec_err).__name__,
error_message=str(exec_err),
extra_data=_proxy_extra,
)
return "break"
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
context = exec_err.context
captured_key_concurrent = context.concurrent_requests
elapsed_ms = context.elapsed_ms
cause = exec_err.cause
has_retry_left = retry_index < (max_retries_for_candidate - 1)
if isinstance(cause, ConcurrencyLimitError):
rpm_current = context.rpm_current
if rpm_current is None:
rpm_current = captured_key_concurrent
rpm_limit = context.rpm_limit
rpm_available_for_new = context.rpm_available_for_new
reservation_ratio = context.reservation_ratio
reservation_phase = context.reservation_phase or "unknown"
reservation_confidence = context.reservation_confidence
reservation_load_factor = context.reservation_load_factor
reason_code = "unknown"
if rpm_limit is not None and rpm_current is not None:
if context.is_cached_user:
if rpm_current >= rpm_limit:
reason_code = "total_limit"
else:
if rpm_available_for_new is not None and rpm_current >= rpm_available_for_new:
reason_code = (
"reserved_for_cached" if rpm_current < rpm_limit else "total_limit"
)
elif rpm_current >= rpm_limit:
reason_code = "total_limit"
reason_text = "并发限制"
if reason_code == "reserved_for_cached":
reason_text = "并发限制: 新用户配额已满(预留给缓存用户)"
elif reason_code == "total_limit":
reason_text = "并发限制: 总配额已满"
parts: list[str] = []
if rpm_current is not None:
parts.append(f"current={rpm_current}")
if rpm_limit is not None:
parts.append(f"limit={rpm_limit}")
if rpm_available_for_new is not None and not context.is_cached_user:
parts.append(f"new={rpm_available_for_new}")
if reservation_ratio is not None:
parts.append(f"reserve={reservation_ratio:.0%}")
if reservation_phase:
parts.append(f"phase={reservation_phase}")
skip_reason = reason_text
if parts:
skip_reason = f"{reason_text} ({', '.join(parts)})"
logger.warning(
" [{}] 并发限制 (attempt={}/{}): provider={}, key={}, cached={}, reason={}, {}",
request_id,
attempt,
max_attempts,
provider.name,
str(key.id)[:8],
bool(context.is_cached_user),
reason_code,
", ".join(parts) if parts else "N/A",
)
extra_data: dict[str, Any] = {
"concurrency_denied": True,
"concurrency_reason": reason_code,
"rpm_current": rpm_current,
"rpm_limit": rpm_limit,
"rpm_available_for_new": rpm_available_for_new,
"reservation_ratio": reservation_ratio,
"reservation_phase": reservation_phase,
"reservation_confidence": reservation_confidence,
"reservation_load_factor": reservation_load_factor,
"attempt": attempt,
"max_attempts": max_attempts,
}
extra_data = {k: v for k, v in extra_data.items() if v is not None}
if _proxy_extra:
extra_data = {**_proxy_extra, **extra_data}
try:
from src.core.metrics import scheduler_concurrency_denied_total
scheduler_concurrency_denied_total.labels(
is_cached_user=str(bool(context.is_cached_user)).lower(),
reason=reason_code,
reservation_phase=str(reservation_phase or "unknown"),
).inc()
except Exception:
pass
RequestCandidateService.mark_candidate_skipped(
db=self.db,
candidate_id=candidate_record_id,
skip_reason=skip_reason,
status_code=429,
concurrent_requests=rpm_current,
extra_data=extra_data,
)
return "break"
if isinstance(cause, ProxyNodeUnavailableError):
# ProxyNode 不可用属于"配置明确指定但不可达/不可用"的情况,
# 在当前候选上重试通常没有意义,直接切换到下一个候选更合理。
node_id = cause.details.get("proxy_node_id") if cause.details else None
logger.warning(
" [{}] 代理节点不可用 (node_id={}),切换候选: {}",
request_id,
node_id or "unknown",
str(cause),
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type=type(cause).__name__,
error_message=extract_error_message(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "break"
if isinstance(cause, EmbeddedErrorException):
error_message = cause.error_message or ""
embedded_status = cause.error_code or 200
if error_classifier.is_client_error(error_message):
logger.warning(
" [{}] 嵌入式客户端错误,继续转移: {}",
request_id,
error_message[:200],
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="UpstreamClientException",
error_message=error_message,
status_code=embedded_status,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "break"
logger.warning(
" [{}] 嵌入式服务端错误,尝试重试: {}",
request_id,
error_message[:200],
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="EmbeddedErrorException",
error_message=error_message,
status_code=embedded_status,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "continue" if has_retry_left else "break"
if isinstance(cause, httpx.HTTPStatusError):
status_code = cause.response.status_code
extra_data = await error_classifier.handle_http_error(
http_error=cause,
provider=provider,
endpoint=endpoint,
key=key,
affinity_key=affinity_key,
api_format=api_format,
global_model_id=global_model_id,
request_id=request_id,
captured_key_concurrent=captured_key_concurrent,
elapsed_ms=elapsed_ms,
max_attempts=max_attempts,
attempt=attempt,
)
# Account Pool: apply health policy (cooldown/disable).
await self._pool_ops.pool_on_error(provider, key, status_code, cause)
converted_error = extra_data.get("converted_error")
serializable_extra_data = {
k: v for k, v in extra_data.items() if k != "converted_error"
}
if _proxy_info:
serializable_extra_data["proxy"] = _proxy_info
if isinstance(converted_error, ThinkingSignatureException):
action = self.handle_thinking_signature_error(
converted_error=converted_error,
provider_type=str(getattr(provider, "provider_type", "") or "").lower(),
model_name=str(global_model_id or ""),
request_id=request_id,
candidate_record_id=candidate_record_id,
elapsed_ms=elapsed_ms,
captured_key_concurrent=captured_key_concurrent,
serializable_extra_data=serializable_extra_data,
request_body_ref=request_body_ref,
)
if action == "continue":
return "continue"
if isinstance(converted_error, UpstreamClientException):
logger.warning(
" [{}] 客户端请求错误,继续转移: {}",
request_id,
str(converted_error.message),
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="UpstreamClientException",
error_message=converted_error.message,
status_code=status_code,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=serializable_extra_data,
)
return "break"
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="HTTPStatusError",
error_message=extract_error_message(cause, status_code),
status_code=status_code,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=serializable_extra_data,
)
return "continue" if has_retry_left else "break"
if isinstance(cause, error_classifier.RETRIABLE_ERRORS):
await error_classifier.handle_retriable_error(
error=cause,
provider=provider,
endpoint=endpoint,
key=key,
affinity_key=affinity_key,
api_format=api_format,
global_model_id=global_model_id,
captured_key_concurrent=captured_key_concurrent,
elapsed_ms=elapsed_ms,
request_id=request_id,
attempt=attempt,
max_attempts=max_attempts,
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type=type(cause).__name__,
error_message=extract_error_message(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "continue" if has_retry_left else "break"
if isinstance(cause, FormatConversionError):
logger.warning(" [{}] 格式转换失败,切换候选: {}", request_id, str(cause))
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="FormatConversionError",
error_message=str(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "break"
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type=type(cause).__name__,
error_message=extract_error_message(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "break"

View File

@@ -0,0 +1,25 @@
from __future__ import annotations
from enum import Enum
from typing import Any
class CandidateErrorAction(str, Enum):
"""同步执行中,候选错误处理后的标准动作。"""
RETRY_CURRENT = "retry_current"
NEXT_CANDIDATE = "next_candidate"
RAISE_ERROR = "raise_error"
def classify_candidate_error_action(action: Any) -> CandidateErrorAction:
"""将 error_handler 的字符串动作分类为可控枚举。"""
action_norm = str(action or "").strip().lower()
if action_norm == "continue":
return CandidateErrorAction.RETRY_CURRENT
if action_norm == "raise":
return CandidateErrorAction.RAISE_ERROR
if action_norm == "break":
return CandidateErrorAction.NEXT_CANDIDATE
# Fail-safe未知动作默认切到下一个候选避免卡死在当前候选。
return CandidateErrorAction.NEXT_CANDIDATE

View File

@@ -0,0 +1,107 @@
from __future__ import annotations
from typing import Any
import httpx
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.services.request.result import RequestMetadata
class TaskFailureOperationsService:
"""任务失败收敛相关操作(异常元数据与统一抛错)。"""
@staticmethod
def attach_metadata_to_error(
error: Exception | None,
candidate: Any | None,
model_name: str,
api_format: str,
) -> None:
"""Attach candidate metadata onto exception for usage recording."""
if not error or not candidate:
return
existing_metadata = getattr(error, "request_metadata", None)
if existing_metadata and getattr(existing_metadata, "api_format", None):
return
metadata = RequestMetadata(
provider_request_headers=(
getattr(existing_metadata, "provider_request_headers", {})
if existing_metadata
else {}
),
provider=getattr(existing_metadata, "provider", None) or str(candidate.provider.name),
model=getattr(existing_metadata, "model", None) or model_name,
provider_id=getattr(existing_metadata, "provider_id", None)
or str(candidate.provider.id),
provider_endpoint_id=(
getattr(existing_metadata, "provider_endpoint_id", None)
or str(candidate.endpoint.id)
),
provider_api_key_id=(
getattr(existing_metadata, "provider_api_key_id", None) or str(candidate.key.id)
),
api_format=api_format,
)
setattr(error, "request_metadata", metadata)
@staticmethod
def raise_all_failed_exception(
request_id: str | None,
max_attempts: int,
last_candidate: Any | None,
model_name: str,
api_format: str,
last_error: Exception | None = None,
) -> None:
"""Raise a unified 'all candidates failed' exception."""
logger.error(" [{}] 所有 {} 个组合均失败", request_id, max_attempts)
request_metadata = None
if last_candidate:
request_metadata = {
"provider": last_candidate.provider.name,
"model": model_name,
"provider_id": str(last_candidate.provider.id),
"provider_endpoint_id": str(last_candidate.endpoint.id),
"provider_api_key_id": str(last_candidate.key.id),
"api_format": api_format,
}
upstream_status: int | None = None
upstream_response: str | None = None
if last_error:
if isinstance(last_error, httpx.HTTPStatusError):
upstream_status = last_error.response.status_code
upstream_response = getattr(last_error, "upstream_response", None)
if not upstream_response:
try:
upstream_response = last_error.response.text
except Exception:
pass
else:
upstream_status = getattr(last_error, "upstream_status", None)
upstream_response = getattr(last_error, "upstream_response", None)
if (
not upstream_response
or not upstream_response.strip()
or upstream_response.startswith("Unable to read")
):
upstream_response = str(last_error)
friendly_message = "服务暂时不可用,请稍后重试"
if last_error:
last_error_message = getattr(last_error, "message", None)
if last_error_message and isinstance(last_error_message, str):
friendly_message = last_error_message
raise ProviderNotAvailableException(
friendly_message,
request_metadata=request_metadata,
upstream_status=upstream_status,
upstream_response=upstream_response,
)

View File

@@ -0,0 +1,222 @@
from __future__ import annotations
from typing import Any
from src.core.logger import logger
class TaskPoolOperationsService:
"""任务池化相关操作(重排、展开、健康回写)。"""
def extract_session_uuid(
self,
provider_type: str,
request_body: dict[str, Any] | None,
) -> str | None:
"""Extract a session UUID from the request body (provider-type aware)."""
if not isinstance(request_body, dict):
return None
from src.services.provider.pool.hooks import get_pool_hook
hook = get_pool_hook(provider_type)
if hook is not None:
return hook.extract_session_uuid(request_body)
return None
async def apply_pool_reorder(
self,
candidates: list[Any],
request_body: dict[str, Any] | None,
) -> tuple[list[Any], list[Any]]:
"""Apply pool key ordering for PoolCandidate objects."""
if not candidates:
return candidates, []
pool_traces: list[Any] = []
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.manager import PoolManager
from src.services.scheduling.schemas import PoolCandidate
for candidate in candidates:
if not isinstance(candidate, PoolCandidate):
continue
provider = candidate.provider
provider_id = str(getattr(provider, "id", "") or "")
if not provider_id:
continue
pool_cfg = candidate.pool_config or parse_pool_config(
getattr(provider, "config", None)
)
if pool_cfg is None:
continue
candidate.pool_config = pool_cfg
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = self.extract_session_uuid(provider_type, request_body)
manager = PoolManager(provider_id, pool_cfg)
candidate_keys = list(candidate.pool_keys or [])
if not candidate_keys and getattr(candidate, "key", None) is not None:
candidate_keys = [candidate.key]
ordered_keys, trace = await manager.select_pool_keys(session_uuid, candidate_keys)
candidate.pool_keys = ordered_keys
selected_key_index = 0
selected_key = None
for idx, pool_key in enumerate(ordered_keys):
if not bool(getattr(pool_key, "_pool_skipped", False)):
selected_key = pool_key
selected_key_index = idx
break
if selected_key is not None:
candidate.key = selected_key
candidate._pool_key_index = selected_key_index
candidate.mapping_matched_model = getattr(
selected_key, "_pool_mapping_matched_model", None
)
candidate.is_skipped = False
candidate.skip_reason = None
else:
candidate.is_skipped = True
candidate.skip_reason = "pool: all keys unavailable"
if trace is not None:
pool_traces.append(trace)
return candidates, pool_traces
except Exception:
logger.opt(exception=True).debug("Pool reorder failed, using original order")
return candidates, []
@staticmethod
def expand_pool_candidates_for_async_submit(candidates: list[Any]) -> list[Any]:
"""Expand PoolCandidate to key-level candidates for async submit traversal."""
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
expanded: list[Any] = []
for candidate in candidates:
if not isinstance(candidate, PoolCandidate):
expanded.append(candidate)
continue
pool_keys = list(candidate.pool_keys or [])
if not pool_keys:
expanded.append(candidate)
continue
for key_index, pool_key in enumerate(pool_keys):
key_skipped = bool(getattr(pool_key, "_pool_skipped", False))
key_skip_reason = (
str(getattr(pool_key, "_pool_skip_reason", "") or "") or candidate.skip_reason
)
key_extra = (
getattr(pool_key, "_pool_extra_data", None)
if isinstance(getattr(pool_key, "_pool_extra_data", None), dict)
else {}
)
key_candidate = ProviderCandidate(
provider=candidate.provider,
endpoint=candidate.endpoint,
key=pool_key,
is_cached=candidate.is_cached,
is_skipped=bool(candidate.is_skipped) or key_skipped,
skip_reason=(
key_skip_reason if (bool(candidate.is_skipped) or key_skipped) else None
),
mapping_matched_model=getattr(pool_key, "_pool_mapping_matched_model", None)
or candidate.mapping_matched_model,
needs_conversion=candidate.needs_conversion,
provider_api_format=candidate.provider_api_format,
output_limit=candidate.output_limit,
capability_miss_count=candidate.capability_miss_count,
)
setattr(
key_candidate,
"_pool_extra_data",
{
"pool_group_id": str(candidate.provider.id),
"pool_key_index": key_index,
**key_extra,
},
)
expanded.append(key_candidate)
return expanded
async def pool_on_success(
self,
candidate: Any,
request_body: dict[str, Any] | None,
) -> None:
"""Notify the pool manager about a successful request (sticky + LRU)."""
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.manager import PoolManager
provider = candidate.provider
provider_config = getattr(provider, "config", None)
pool_cfg = parse_pool_config(provider_config)
if pool_cfg is None:
return
provider_id = str(getattr(provider, "id", "") or "")
key_id = str(getattr(candidate.key, "id", "") or "")
if not provider_id or not key_id:
return
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = self.extract_session_uuid(provider_type, request_body)
mgr = PoolManager(provider_id, pool_cfg)
await mgr.on_request_success(
session_uuid=session_uuid,
key_id=key_id,
)
except Exception:
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
@staticmethod
async def pool_on_error(
provider: Any,
key: Any,
status_code: int,
cause: Any,
) -> None:
"""Notify the pool manager about an upstream error (health policy)."""
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.health_policy import apply_health_policy
pool_cfg = parse_pool_config(getattr(provider, "config", None))
if pool_cfg is None:
return
error_text = ""
resp_headers: dict[str, str] = {}
if getattr(cause, "response", None) is not None:
try:
error_text = (cause.response.text or "")[:4000]
except Exception:
pass
try:
resp_headers = dict(cause.response.headers)
except Exception:
pass
await apply_health_policy(
provider_id=str(provider.id),
key_id=str(key.id),
status_code=status_code,
error_body=error_text,
response_headers=resp_headers,
config=pool_cfg,
)
except Exception:
pass

View File

@@ -0,0 +1,101 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from src.services.candidate.policy import FailoverAction
from src.services.task.execute.exception_classification import CandidateErrorAction
if TYPE_CHECKING:
from src.services.task.execute.failure import TaskFailureOperationsService
@dataclass(slots=True)
class ExecutionErrorTransition:
"""执行异常后的状态流转决策。"""
failover_action: FailoverAction
max_retries: int | None = None
def as_failover_tuple(self) -> tuple[FailoverAction, int | None]:
return (self.failover_action, self.max_retries)
@dataclass(slots=True)
class SyncExecutionState:
"""同步执行阶段状态容器(候选上下文 + 异常上下文)。"""
candidate_record_map: dict[tuple[int, int], str]
request_body_ref: dict[str, Any] | None
last_error: Exception | None = None
last_candidate: Any | None = None
_rectify_flag_key: str = field(default="_rectified_this_turn", repr=False)
def touch_candidate(self, candidate: Any) -> None:
self.last_candidate = candidate
def track_execution_error(self, *, exec_err: Any, candidate: Any) -> None:
self.last_candidate = candidate
cause = getattr(exec_err, "cause", None)
self.last_error = cause if isinstance(cause, Exception) else None
def resolve_candidate_record_id(self, *, candidate_index: int, record_id: str | None) -> str:
if record_id:
return str(record_id)
return str(self.candidate_record_map.get((candidate_index, 0), "") or "")
def consume_rectify_retry_extension(
self, *, max_retries_for_candidate: int, retry_index: int
) -> int | None:
if not self.request_body_ref:
return None
if not self.request_body_ref.get(self._rectify_flag_key, False):
return None
self.request_body_ref[self._rectify_flag_key] = False
return max(max_retries_for_candidate, retry_index + 2)
def raise_classified_error(
self,
*,
fallback_error: Any,
failure_ops: TaskFailureOperationsService,
model_name: str,
api_format: str,
) -> None:
if self.last_error is not None:
failure_ops.attach_metadata_to_error(
self.last_error,
self.last_candidate,
model_name,
api_format,
)
raise self.last_error
if isinstance(fallback_error, Exception):
raise fallback_error
raise RuntimeError("execution_error_handler requested raise without exception context")
def resolve_execution_error_transition(
*,
action: CandidateErrorAction,
state: SyncExecutionState,
max_retries_for_candidate: int,
retry_index: int,
) -> ExecutionErrorTransition:
"""根据异常动作分类,返回 FailoverEngine 可消费的状态流转结果。"""
if action == CandidateErrorAction.RETRY_CURRENT:
return ExecutionErrorTransition(
failover_action=FailoverAction.RETRY,
max_retries=state.consume_rectify_retry_extension(
max_retries_for_candidate=max_retries_for_candidate,
retry_index=retry_index,
),
)
return ExecutionErrorTransition(
failover_action=FailoverAction.CONTINUE,
max_retries=None,
)

View File

@@ -0,0 +1,370 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
from uuid import uuid4
from src.core.logger import logger
from src.models.database import ApiKey, User
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import FailoverAction, RetryPolicy, SkipPolicy
from src.services.candidate.resolver import CandidateResolver
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.orchestration.request_dispatcher import RequestDispatcher
from src.services.provider.format import normalize_endpoint_signature
from src.services.request.candidate import RequestCandidateService
from src.services.scheduling.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.system.config import SystemConfigService
from src.services.task.core.protocol import AttemptKind, AttemptResult
from src.services.task.core.schema import ExecutionResult
from src.services.task.execute.error_handler import TaskErrorOperationsService
from src.services.task.execute.exception_classification import (
CandidateErrorAction,
classify_candidate_error_action,
)
from src.services.task.execute.failure import TaskFailureOperationsService
from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.execute.state_transition import (
SyncExecutionState,
resolve_execution_error_transition,
)
from src.services.usage.service import UsageService
class SyncTaskExecutionService:
"""同步任务执行服务(候选遍历 + 错误处理 + 结果聚合)。"""
def __init__(
self,
db: Any,
redis_client: Any | None,
*,
recorder: Any,
pool_ops: TaskPoolOperationsService,
error_ops: TaskErrorOperationsService,
failure_ops: TaskFailureOperationsService,
) -> None:
self.db = db
self.redis = redis_client
self._recorder = recorder
self._pool_ops = pool_ops
self._error_ops = error_ops
self._failure_ops = failure_ops
async def execute_sync_unified(
self,
*,
api_format: str,
model_name: str,
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None,
is_stream: bool,
capability_requirements: dict[str, bool] | None,
preferred_key_ids: list[str] | None,
request_body_ref: dict[str, Any] | None,
request_headers: dict[str, Any] | None,
request_body: dict[str, Any] | None,
) -> ExecutionResult:
"""
Unified candidate traversal loop for SYNC.
This intentionally reuses existing components for parity:
- CandidateResolver fetch + record creation
- RequestDispatcher execution
- Error classification/rectify logic ported from the previous SYNC implementation
"""
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.request.executor import RequestExecutor
if not request_id:
request_id = str(uuid4())
# Build execution components (mirrors pre-Phase-3 initialization)
priority_mode = SystemConfigService.get_config(
self.db,
"provider_priority_mode",
CacheAwareScheduler.PRIORITY_MODE_PROVIDER,
)
scheduling_mode = SystemConfigService.get_config(
self.db,
"scheduling_mode",
CacheAwareScheduler.SCHEDULING_MODE_CACHE_AFFINITY,
)
cache_scheduler = await get_cache_aware_scheduler(
self.redis,
priority_mode=priority_mode,
scheduling_mode=scheduling_mode,
)
# Ensure cache_scheduler inner state is ready
await cache_scheduler._ensure_initialized()
concurrency_manager = await get_concurrency_manager()
adaptive_manager = get_adaptive_rpm_manager()
request_executor = RequestExecutor(
db=self.db,
concurrency_manager=concurrency_manager,
adaptive_manager=adaptive_manager,
)
candidate_resolver = CandidateResolver(
db=self.db,
cache_scheduler=cache_scheduler,
)
error_classifier = ErrorClassifier(
db=self.db,
cache_scheduler=cache_scheduler,
adaptive_manager=adaptive_manager,
)
request_dispatcher = RequestDispatcher(
db=self.db,
request_executor=request_executor,
cache_scheduler=cache_scheduler,
)
affinity_key = str(user_api_key.id)
user_id = str(user_api_key.user_id)
api_format_norm = normalize_endpoint_signature(api_format)
# Keep pending usage creation behavior consistent with previous behavior
try:
user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
UsageService.create_pending_usage(
db=self.db,
request_id=request_id,
user=user,
api_key=user_api_key,
model=model_name,
is_stream=is_stream,
api_format=api_format_norm,
request_headers=request_headers,
request_body=request_body,
)
except Exception as exc:
logger.warning("创建 pending 使用记录失败: {}", str(exc))
all_candidates, global_model_id = await candidate_resolver.fetch_candidates(
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body=request_body,
)
# Account Pool: reorder candidates for claude_code providers.
all_candidates, pool_traces = await self._pool_ops.apply_pool_reorder(
all_candidates, request_body=request_body
)
candidate_record_map = candidate_resolver.create_candidate_records(
all_candidates=all_candidates,
request_id=request_id,
user_id=user_id,
user_api_key=user_api_key,
required_capabilities=capability_requirements,
)
max_attempts = candidate_resolver.count_total_attempts(all_candidates)
# Keep behavior consistent with previous behavior: last_candidate is updated even if skipped.
execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map,
request_body_ref=request_body_ref,
last_candidate=all_candidates[-1] if all_candidates else None,
)
async def _attempt(candidate: Any) -> AttemptResult:
execution_state.touch_candidate(candidate)
candidate_index = int(getattr(candidate, "_utf_candidate_index", -1))
retry_index = int(getattr(candidate, "_utf_retry_index", 0))
candidate_record_id = str(getattr(candidate, "_utf_candidate_record_id", "") or "")
attempt_counter = int(getattr(candidate, "_utf_attempt_count", 0))
max_attempts_local = int(getattr(candidate, "_utf_max_attempts", max_attempts))
# Safety net: if record_id missing, create an "available" record on-demand.
if not candidate_record_id:
from src.services.scheduling.schemas import PoolCandidate
pool_extra = (
getattr(candidate.key, "_pool_extra_data", None)
if isinstance(getattr(candidate.key, "_pool_extra_data", None), dict)
else {}
)
extra_data: dict[str, Any] = {
"needs_conversion": bool(getattr(candidate, "needs_conversion", False)),
"provider_api_format": getattr(candidate, "provider_api_format", None) or None,
"mapping_matched_model": getattr(candidate, "mapping_matched_model", None)
or None,
**pool_extra,
}
if isinstance(candidate, PoolCandidate):
extra_data["pool_group_id"] = str(candidate.provider.id)
extra_data["pool_key_index"] = int(
getattr(candidate, "_pool_key_index", 0) or 0
)
created = RequestCandidateService.create_candidate(
db=self.db,
request_id=request_id,
candidate_index=candidate_index,
retry_index=retry_index,
user_id=user_id,
api_key_id=str(user_api_key.id),
provider_id=str(candidate.provider.id),
endpoint_id=str(candidate.endpoint.id),
key_id=str(candidate.key.id),
status="available",
is_cached=bool(getattr(candidate, "is_cached", False)),
extra_data=extra_data,
)
candidate_record_id = str(created.id)
execution_state.candidate_record_map[(candidate_index, retry_index)] = (
candidate_record_id
)
response, _provider_name, attempt_id, _provider_id, _endpoint_id, _key_id = (
await request_dispatcher.dispatch(
candidate=candidate,
candidate_index=candidate_index,
retry_index=retry_index,
candidate_record_id=candidate_record_id,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
)
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
# Account Pool: on success, update sticky binding + LRU.
await self._pool_ops.pool_on_success(candidate, request_body)
if is_stream:
return AttemptResult(
kind=AttemptKind.STREAM,
http_status=200,
http_headers={},
stream_iterator=response,
)
return AttemptResult(
kind=AttemptKind.SYNC_RESPONSE,
http_status=200,
http_headers={},
response_body=response,
)
async def _handle_exec_err(
*,
exec_err: Any,
candidate: Any,
candidate_index: int,
retry_index: int,
max_retries_for_candidate: int,
record_id: str | None,
attempt_count: int,
max_attempts: int | None,
) -> tuple[FailoverAction, int | None]:
execution_state.track_execution_error(exec_err=exec_err, candidate=candidate)
# Fall back to retry 0 record if needed (rectify may extend retries).
candidate_record_id = execution_state.resolve_candidate_record_id(
candidate_index=candidate_index,
record_id=record_id,
)
raw_action = await self._error_ops.handle_candidate_error(
exec_err=exec_err,
candidate=candidate,
candidate_record_id=candidate_record_id,
retry_index=retry_index,
max_retries_for_candidate=max_retries_for_candidate,
affinity_key=affinity_key,
api_format=api_format_norm,
global_model_id=global_model_id,
request_id=request_id,
attempt=attempt_count,
max_attempts=int(max_attempts or 0),
request_body_ref=request_body_ref,
error_classifier=error_classifier,
)
action = classify_candidate_error_action(raw_action)
if action == CandidateErrorAction.RAISE_ERROR:
execution_state.raise_classified_error(
fallback_error=exec_err,
failure_ops=self._failure_ops,
model_name=model_name,
api_format=api_format_norm,
)
return resolve_execution_error_transition(
action=action,
state=execution_state,
max_retries_for_candidate=max_retries_for_candidate,
retry_index=retry_index,
).as_failover_tuple()
engine = FailoverEngine(
self.db,
error_classifier=error_classifier,
recorder=self._recorder,
)
result = await engine.execute(
candidates=all_candidates,
attempt_func=_attempt,
retry_policy=RetryPolicy.for_sync_task(),
skip_policy=SkipPolicy(),
request_id=request_id,
user_id=user_id,
api_key_id=str(user_api_key.id),
candidate_record_map=candidate_record_map,
max_attempts=max_attempts,
execution_error_handler=_handle_exec_err,
)
if result.success:
# Build pool scheduling summary from traces collected during reorder.
if pool_traces and result.key_id:
try:
attempted_key_ids: set[str] = set()
for ck in result.candidate_keys or []:
status = str(getattr(ck, "status", "") or "").strip().lower()
if status in {"", "available", "pending", "skipped", "unused"}:
continue
kid = getattr(ck, "key_id", None)
if isinstance(kid, str) and kid:
attempted_key_ids.add(kid)
if not attempted_key_ids:
attempted_key_ids.add(str(result.key_id))
for pt in pool_traces:
summary = pt.build_summary(
result.key_id,
attempted_key_ids=attempted_key_ids,
)
if summary:
result.pool_summary = summary
break
except Exception:
pass
return result
self._failure_ops.raise_all_failed_exception(
request_id,
max_attempts,
execution_state.last_candidate,
model_name,
api_format_norm,
execution_state.last_error,
)

View File

@@ -1,7 +0,0 @@
"""Per-task-type implementations (Phase2).
Currently includes:
- video: polling adapter
"""
__all__ = []

View File

View File

@@ -21,7 +21,7 @@ from src.core.api_format.conversion.internal_video import InternalVideoPollResul
from src.core.logger import logger from src.core.logger import logger
from src.database import create_session from src.database import create_session
from src.services.system.scheduler import get_scheduler from src.services.system.scheduler import get_scheduler
from src.services.task.impl.video_poller import VideoPollContext, VideoTaskPollerAdapter from src.services.task.video.poller_adapter import VideoPollContext, VideoTaskPollerAdapter
@runtime_checkable @runtime_checkable

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,12 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import Any
ApplyPoolReorderFn = Callable[
[list[Any], dict[str, Any] | None],
Awaitable[tuple[list[Any], list[Any]]],
]
ExpandPoolCandidatesFn = Callable[[list[Any]], list[Any]]
__all__ = ["ApplyPoolReorderFn", "ExpandPoolCandidatesFn"]

View File

@@ -0,0 +1,75 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.candidate.submit import SubmitOutcome
from src.services.task.submit.record import AsyncSubmitRecordService
from src.services.task.submit.response import AsyncSubmitResponseService
class AsyncSubmitAttemptService:
"""异步提交单候选执行编排服务。"""
def __init__(
self,
db: Session,
*,
sanitize: Callable[[str], str],
extract_response_text: Callable[[httpx.Response], str],
match_provider_failover_rule: Callable[..., str | None],
) -> None:
self.db = db
self._record_ops = AsyncSubmitRecordService(db)
self._response_ops = AsyncSubmitResponseService(
db,
record_ops=self._record_ops,
sanitize=sanitize,
extract_response_text=extract_response_text,
match_provider_failover_rule=match_provider_failover_rule,
)
async def submit_candidate(
self,
*,
candidate: Any,
record_id: str | None,
candidate_info: dict[str, Any],
candidate_keys: list[dict[str, Any]],
rule_lookup: BillingRuleLookupResult | None,
submit_func: Any,
extract_external_task_id: Any,
) -> tuple[SubmitOutcome | None, int | None]:
self._record_ops.mark_pending(record_id=record_id)
# Flush/commit BEFORE awaiting upstream submit to avoid holding DB connections.
if self.db.in_transaction():
try:
self.db.commit()
except Exception:
self.db.rollback()
raise
# Attempt submit (upstream HTTP)
try:
response: httpx.Response = await submit_func(candidate)
except Exception as exc:
return self._response_ops.handle_submit_exception(
record_id=record_id,
candidate_info=candidate_info,
exc=exc,
)
return self._response_ops.handle_submit_response(
candidate=candidate,
record_id=record_id,
candidate_info=candidate_info,
candidate_keys=candidate_keys,
rule_lookup=rule_lookup,
response=response,
extract_external_task_id=extract_external_task_id,
)

View File

@@ -0,0 +1,104 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.config.settings import config
from src.services.candidate.submit import AllCandidatesFailedError, SubmitOutcome
from src.services.task.submit.attempt import AsyncSubmitAttemptService
from src.services.task.submit.filter import AsyncSubmitFilterService
class AsyncSubmitExecutionService:
"""异步提交候选执行编排服务。"""
def __init__(
self,
db: Session,
*,
sanitize: Callable[[str], str],
extract_response_text: Callable[[httpx.Response], str],
match_provider_failover_rule: Callable[..., str | None],
) -> None:
self.db = db
self._filter_ops = AsyncSubmitFilterService(db)
self._attempt_ops = AsyncSubmitAttemptService(
db,
sanitize=sanitize,
extract_response_text=extract_response_text,
match_provider_failover_rule=match_provider_failover_rule,
)
async def execute_submit_loop(
self,
*,
candidates: list[Any],
record_map: dict[tuple[int, int], str],
task_type: str,
model_name: str,
submit_func: Any,
extract_external_task_id: Any,
supported_auth_types: set[str] | None,
allow_format_conversion: bool,
) -> SubmitOutcome:
candidate_keys: list[dict[str, Any]] = []
eligible_count = 0
last_status_code: int | None = None
for idx, cand in enumerate(candidates):
candidate_info = self._filter_ops.build_candidate_info(idx=idx, candidate=cand)
candidate_keys.append(candidate_info)
attempt_plan = self._filter_ops.prepare_candidate_for_attempt(
idx=idx,
candidate=cand,
record_map=record_map,
candidate_info=candidate_info,
task_type=task_type,
model_name=model_name,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
)
if attempt_plan is None:
continue
eligible_count += 1
outcome, status_code = await self._attempt_ops.submit_candidate(
candidate=cand,
record_id=attempt_plan.record_id,
candidate_info=candidate_info,
candidate_keys=candidate_keys,
rule_lookup=attempt_plan.rule_lookup,
submit_func=submit_func,
extract_external_task_id=extract_external_task_id,
)
if status_code is not None:
last_status_code = status_code
if outcome is not None:
return outcome
# Persist candidate records before raising.
try:
self.db.commit()
except Exception:
self.db.rollback()
if eligible_count == 0:
reason = "no_eligible_candidates"
if config.billing_require_rule:
reason = "no_candidate_with_billing_rule"
raise AllCandidatesFailedError(
reason=reason,
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
raise AllCandidatesFailedError(
reason="all_candidates_failed",
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)

View File

@@ -0,0 +1,137 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from sqlalchemy import update
from sqlalchemy.orm import Session
from src.config.settings import config
from src.models.database import RequestCandidate
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.system.config import SystemConfigService
@dataclass(slots=True)
class CandidateAttemptPlan:
"""候选尝试计划(通过过滤后可进入提交阶段)。"""
record_id: str | None
rule_lookup: BillingRuleLookupResult | None
class AsyncSubmitFilterService:
"""异步提交候选过滤服务。"""
def __init__(self, db: Session) -> None:
self.db = db
@staticmethod
def build_candidate_info(*, idx: int, candidate: Any) -> dict[str, Any]:
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
return {
"index": idx,
"provider_id": candidate.provider.id,
"provider_name": candidate.provider.name,
"endpoint_id": candidate.endpoint.id,
"key_id": candidate.key.id,
"key_name": getattr(candidate.key, "name", None),
"auth_type": auth_type,
"priority": getattr(candidate.key, "priority", 0) or 0,
"is_cached": bool(getattr(candidate, "is_cached", False)),
}
def prepare_candidate_for_attempt(
self,
*,
idx: int,
candidate: Any,
record_map: dict[tuple[int, int], str],
candidate_info: dict[str, Any],
task_type: str,
model_name: str,
supported_auth_types: set[str] | None,
allow_format_conversion: bool,
) -> CandidateAttemptPlan | None:
record_id = record_map.get((idx, 0))
auth_type = candidate_info.get("auth_type", "api_key")
# Scheduler marked skip
if getattr(candidate, "is_skipped", False):
skip_reason = getattr(candidate, "skip_reason", None) or "skipped"
self._mark_skip(
record_id=record_id, candidate_info=candidate_info, skip_reason=skip_reason
)
return None
# Format conversion checks
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
if needs_conversion:
# 1. handler-level switch
if not allow_format_conversion:
self._mark_skip(
record_id=record_id,
candidate_info=candidate_info,
skip_reason="format_conversion_not_supported",
)
return None
# 2. global switch (from database config)
if not SystemConfigService.is_format_conversion_enabled(self.db):
self._mark_skip(
record_id=record_id,
candidate_info=candidate_info,
skip_reason="format_conversion_disabled",
extra_info={"format_conversion_enabled": False},
)
return None
# auth_type filter
if supported_auth_types is not None and auth_type not in supported_auth_types:
self._mark_skip(
record_id=record_id,
candidate_info=candidate_info,
skip_reason=f"unsupported_auth_type:{auth_type}",
)
return None
# billing rule filter
rule_lookup: BillingRuleLookupResult | None = None
has_billing_rule = True
if config.billing_require_rule:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=model_name,
task_type=task_type,
)
has_billing_rule = rule_lookup is not None
if not has_billing_rule:
self._mark_skip(
record_id=record_id,
candidate_info=candidate_info,
skip_reason="billing_rule_missing",
extra_info={"has_billing_rule": False},
)
return None
candidate_info["has_billing_rule"] = has_billing_rule
return CandidateAttemptPlan(record_id=record_id, rule_lookup=rule_lookup)
def _mark_skip(
self,
*,
record_id: str | None,
candidate_info: dict[str, Any],
skip_reason: str,
extra_info: dict[str, Any] | None = None,
) -> None:
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
if extra_info:
candidate_info.update(extra_info)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="skipped", skip_reason=skip_reason)
)

View File

@@ -0,0 +1,60 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
import httpx
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.candidate.submit import SubmitOutcome
@dataclass(slots=True)
class SubmitPayloadParseResult:
"""提交响应解析结果。"""
payload: dict[str, Any] | None
error_type: str | None = None
error_message: str | None = None
class AsyncSubmitOutcomeBuilderService:
"""异步提交结果构建服务。"""
def __init__(self, *, sanitize: Callable[[str], str]) -> None:
self._sanitize = sanitize
def parse_payload(self, *, response: httpx.Response) -> SubmitPayloadParseResult:
payload: dict[str, Any] | None = None
try:
data = response.json()
if isinstance(data, dict):
payload = data
except Exception as exc:
return SubmitPayloadParseResult(
payload=None,
error_type=type(exc).__name__,
error_message=self._sanitize(str(exc)),
)
return SubmitPayloadParseResult(payload=payload)
@staticmethod
def build_success_outcome(
*,
candidate: Any,
candidate_keys: list[dict[str, Any]],
external_task_id: str,
rule_lookup: BillingRuleLookupResult | None,
payload: dict[str, Any] | None,
response: httpx.Response,
) -> SubmitOutcome:
return SubmitOutcome(
candidate=candidate,
candidate_keys=candidate_keys,
external_task_id=external_task_id,
rule_lookup=rule_lookup,
upstream_payload=payload,
upstream_headers=dict(response.headers),
upstream_status_code=response.status_code,
)

View File

@@ -0,0 +1,121 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey
from src.services.candidate.resolver import CandidateResolver
from src.services.candidate.submit import AllCandidatesFailedError
from src.services.scheduling.aware_scheduler import get_cache_aware_scheduler
from src.services.system.config import SystemConfigService
from src.services.task.submit import ApplyPoolReorderFn, ExpandPoolCandidatesFn
@dataclass(slots=True)
class PreparedSubmitCandidates:
"""异步提交前的候选准备结果。"""
candidates: list[Any]
record_map: dict[tuple[int, int], str]
class AsyncSubmitPreparationService:
"""异步提交候选准备服务。"""
def __init__(
self,
db: Session,
redis_client: Any | None,
*,
sanitize: Callable[[str], str],
) -> None:
self.db = db
self.redis = redis_client
self._sanitize = sanitize
async def prepare_candidates(
self,
*,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey,
request_id: str | None,
capability_requirements: dict[str, bool] | None,
request_body: dict[str, Any] | None,
max_candidates: int | None,
apply_pool_reorder: ApplyPoolReorderFn,
expand_pool_candidates_for_async_submit: ExpandPoolCandidatesFn,
) -> PreparedSubmitCandidates:
priority_mode = SystemConfigService.get_config(
self.db,
"provider_priority_mode",
"provider",
)
scheduling_mode = SystemConfigService.get_config(
self.db,
"scheduling_mode",
"cache_affinity",
)
cache_scheduler = await get_cache_aware_scheduler(
self.redis,
priority_mode=priority_mode,
scheduling_mode=scheduling_mode,
)
resolver = CandidateResolver(db=self.db, cache_scheduler=cache_scheduler)
candidates, _global_model_id = await 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,
request_body=request_body,
)
_ = _global_model_id
if not candidates:
raise AllCandidatesFailedError(
reason="no_candidates",
candidate_keys=[],
last_status_code=None,
)
# Account Pool: keep internal key failover order/skip behavior
# consistent with the SYNC path.
candidates, _pool_traces = await apply_pool_reorder(
candidates,
request_body=request_body,
)
_ = _pool_traces
candidates = expand_pool_candidates_for_async_submit(candidates)
if max_candidates is not None and max_candidates > 0:
candidates = candidates[:max_candidates]
# Pre-create RequestCandidate records (no retry expand for async submit stage)
record_map: dict[tuple[int, int], str] = {}
if request_id:
try:
record_map = resolver.create_candidate_records(
all_candidates=candidates,
request_id=request_id,
user_id=str(user_api_key.user_id),
user_api_key=user_api_key,
required_capabilities=capability_requirements,
expand_retries=False,
)
except Exception as exc:
logger.warning(
"[TaskService] Failed to create candidate records: {}",
self._sanitize(str(exc)),
)
record_map = {}
return PreparedSubmitCandidates(candidates=candidates, record_map=record_map)

View File

@@ -0,0 +1,67 @@
from __future__ import annotations
from datetime import datetime, timezone
from sqlalchemy import update
from sqlalchemy.orm import Session
from src.models.database import RequestCandidate
class AsyncSubmitRecordService:
"""异步提交阶段的 RequestCandidate 落库服务。"""
def __init__(self, db: Session) -> None:
self.db = db
def mark_pending(self, *, record_id: str | None) -> None:
if not record_id:
return
started_at = datetime.now(timezone.utc)
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="pending", started_at=started_at)
)
def mark_failed(
self,
*,
record_id: str | None,
error_type: str,
error_message: str,
status_code: int | None = None,
) -> None:
if not record_id:
return
values: dict[str, object] = {
"status": "failed",
"error_type": error_type,
"error_message": error_message,
"finished_at": datetime.now(timezone.utc),
}
if status_code is not None:
values["status_code"] = status_code
self.db.execute(
update(RequestCandidate).where(RequestCandidate.id == record_id).values(**values)
)
def mark_success(
self,
*,
record_id: str | None,
status_code: int,
) -> None:
if not record_id:
return
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(
status="success",
status_code=status_code,
finished_at=datetime.now(timezone.utc),
)
)

View File

@@ -0,0 +1,201 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.candidate.submit import SubmitOutcome, UpstreamClientRequestError
from src.services.task.submit.outcome_builder import (
AsyncSubmitOutcomeBuilderService,
)
from src.services.task.submit.record import AsyncSubmitRecordService
from src.services.task.submit.rule_decider import AsyncSubmitRuleDeciderService
class AsyncSubmitResponseService:
"""异步提交响应判定服务。"""
def __init__(
self,
db: Session,
*,
record_ops: AsyncSubmitRecordService,
sanitize: Callable[[str], str],
extract_response_text: Callable[[httpx.Response], str],
match_provider_failover_rule: Callable[..., str | None],
) -> None:
self.db = db
self._record_ops = record_ops
self._sanitize = sanitize
self._extract_response_text = extract_response_text
self._rule_decider = AsyncSubmitRuleDeciderService(
match_provider_failover_rule=match_provider_failover_rule
)
self._outcome_builder = AsyncSubmitOutcomeBuilderService(sanitize=sanitize)
def handle_submit_exception(
self,
*,
record_id: str | None,
candidate_info: dict[str, Any],
exc: Exception,
) -> tuple[None, None]:
error_type = type(exc).__name__
error_msg = self._sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "exception",
"error_type": error_type,
"error_message": error_msg,
}
)
self._record_ops.mark_failed(
record_id=record_id,
error_type=error_type,
error_message=error_msg,
)
return None, None
def handle_submit_response(
self,
*,
candidate: Any,
record_id: str | None,
candidate_info: dict[str, Any],
candidate_keys: list[dict[str, Any]],
rule_lookup: BillingRuleLookupResult | None,
response: httpx.Response,
extract_external_task_id: Any,
) -> tuple[SubmitOutcome | None, int | None]:
last_status_code = int(getattr(response, "status_code", 0) or 0)
if response.status_code >= 400:
error_text = self._extract_response_text(response)
error_msg = self._sanitize(error_text)
candidate_info.update(
{
"attempt_status": "http_error",
"status_code": response.status_code,
"error_message": error_msg,
}
)
self._record_ops.mark_failed(
record_id=record_id,
status_code=response.status_code,
error_type="http_error",
error_message=error_msg,
)
stop_pattern = self._rule_decider.detect_error_stop_pattern(
candidate=candidate,
response_text=error_text,
status_code=response.status_code,
)
if stop_pattern:
logger.info(
"[TaskService] 错误终止规则命中: pattern={}, status_code={}, provider={}",
stop_pattern,
response.status_code,
candidate.provider.name,
)
candidate_info["stop_rule_pattern"] = stop_pattern
try:
self.db.commit()
except Exception:
self.db.rollback()
raise UpstreamClientRequestError(
response=response,
candidate_keys=candidate_keys,
)
return None, last_status_code
success_text = self._extract_response_text(response)
success_continue_pattern = self._rule_decider.detect_success_failover_pattern(
candidate=candidate,
response_text=success_text,
status_code=response.status_code,
)
if success_continue_pattern:
logger.info(
"[TaskService] 成功转移规则命中: pattern={}, status_code={}, provider={}",
success_continue_pattern,
response.status_code,
candidate.provider.name,
)
failover_reason = f"success_failover_rule_matched:{success_continue_pattern}"
candidate_info.update(
{
"attempt_status": "success_failover",
"status_code": response.status_code,
"error_message": failover_reason,
"success_rule_pattern": success_continue_pattern,
}
)
self._record_ops.mark_failed(
record_id=record_id,
status_code=response.status_code,
error_type="success_failover_pattern",
error_message=failover_reason,
)
return None, last_status_code
parse_result = self._outcome_builder.parse_payload(response=response)
if parse_result.error_type:
candidate_info.update(
{
"attempt_status": "invalid_json",
"error_type": parse_result.error_type,
"error_message": parse_result.error_message,
}
)
self._record_ops.mark_failed(
record_id=record_id,
status_code=response.status_code,
error_type="invalid_json",
error_message=parse_result.error_message or "invalid_json",
)
return None, last_status_code
payload = parse_result.payload
external_task_id = extract_external_task_id(payload or {})
if not external_task_id:
candidate_info.update(
{
"attempt_status": "empty_task_id",
"error_message": "Upstream returned empty task id",
}
)
self._record_ops.mark_failed(
record_id=record_id,
status_code=response.status_code,
error_type="empty_task_id",
error_message="Upstream returned empty task id",
)
return None, last_status_code
# Success
candidate_info.update({"attempt_status": "success", "selected": True})
self._record_ops.mark_success(
record_id=record_id,
status_code=response.status_code,
)
try:
self.db.commit()
except Exception:
self.db.rollback()
return (
self._outcome_builder.build_success_outcome(
candidate=candidate,
candidate_keys=candidate_keys,
external_task_id=str(external_task_id),
rule_lookup=rule_lookup,
payload=payload,
response=response,
),
last_status_code,
)

View File

@@ -0,0 +1,43 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
class AsyncSubmitRuleDeciderService:
"""异步提交故障转移规则判定服务。"""
def __init__(
self,
*,
match_provider_failover_rule: Callable[..., str | None],
) -> None:
self._match_provider_failover_rule = match_provider_failover_rule
def detect_error_stop_pattern(
self,
*,
candidate: Any,
response_text: str,
status_code: int,
) -> str | None:
return self._match_provider_failover_rule(
candidate,
is_success=False,
response_text=response_text,
status_code=status_code,
)
def detect_success_failover_pattern(
self,
*,
candidate: Any,
response_text: str,
status_code: int,
) -> str | None:
return self._match_provider_failover_rule(
candidate,
is_success=True,
response_text=response_text,
status_code=status_code,
)

View File

@@ -0,0 +1,152 @@
from __future__ import annotations
import re
from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.models.database import ApiKey
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.submit import SubmitOutcome
from src.services.task.submit import ApplyPoolReorderFn, ExpandPoolCandidatesFn
from src.services.task.submit.execute import AsyncSubmitExecutionService
from src.services.task.submit.prepare import AsyncSubmitPreparationService
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
re.IGNORECASE,
)
class AsyncTaskSubmitService:
"""异步任务提交应用服务(候选选择 + 故障转移)。"""
def __init__(
self,
db: Session,
redis_client: Any | None,
*,
apply_pool_reorder: ApplyPoolReorderFn,
expand_pool_candidates_for_async_submit: ExpandPoolCandidatesFn,
) -> None:
self.db = db
self.redis = redis_client
self._apply_pool_reorder = apply_pool_reorder
self._expand_pool_candidates_for_async_submit = expand_pool_candidates_for_async_submit
self._prepare_ops = AsyncSubmitPreparationService(
db,
redis_client,
sanitize=self._sanitize,
)
self._execute_ops = AsyncSubmitExecutionService(
db,
sanitize=self._sanitize,
extract_response_text=self._extract_response_text,
match_provider_failover_rule=self._match_provider_failover_rule,
)
@staticmethod
def _sanitize(message: str, max_length: int = 200) -> str:
if not message:
return "request_failed"
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
@staticmethod
def _extract_response_text(response: httpx.Response) -> str:
try:
return response.text or ""
except Exception:
return ""
@staticmethod
def _match_provider_failover_rule(
candidate: Any,
*,
is_success: bool,
response_text: str,
status_code: int | None = None,
) -> str | None:
provider_config = getattr(candidate.provider, "config", None) or {}
rules = provider_config.get("failover_rules")
if not rules or not isinstance(rules, dict):
return None
compiled = FailoverEngine._get_compiled_patterns(rules)
key = "success" if is_success else "error"
for regex, rule in compiled.get(key, []):
if not is_success:
rule_status_codes = rule.get("status_codes")
if rule_status_codes and status_code not in rule_status_codes:
continue
if regex.search(response_text):
return rule.get("pattern", "")
return None
async def submit_with_failover(
self,
*,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey,
request_id: str | None,
task_type: str,
submit_func: Any,
extract_external_task_id: Any,
supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False,
capability_requirements: dict[str, bool] | None = None,
max_candidates: int | None = None,
request_body: dict[str, Any] | None = None,
) -> SubmitOutcome:
"""
异步提交入口。
行为保持与原 TaskService.submit_with_failover 一致:
- 按候选顺序依次尝试(提交阶段不做单候选重试)
- 记录 RequestCandidate 审计行
- 命中 error_stop_patterns 时立即停止并抛出上游错误
- 命中 success_failover_patterns 时继续尝试下一个候选
"""
# IMPORTANT:
# This method awaits upstream HTTP calls. If we have an open DB transaction before awaiting,
# the connection can be held for a long time (pool exhaustion under concurrency).
#
# Also note SQLAlchemy's default expire_on_commit=True would expire ORM objects and may
# trigger unexpected lazy DB loads after we commit (potentially during the await).
# We disable it temporarily to keep candidate/provider/key objects in-memory.
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
self.db.expire_on_commit = False
try:
prepared = await self._prepare_ops.prepare_candidates(
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
capability_requirements=capability_requirements,
request_body=request_body,
max_candidates=max_candidates,
apply_pool_reorder=self._apply_pool_reorder,
expand_pool_candidates_for_async_submit=(
self._expand_pool_candidates_for_async_submit
),
)
return await self._execute_ops.execute_submit_loop(
candidates=prepared.candidates,
record_map=prepared.record_map,
task_type=task_type,
model_name=model_name,
submit_func=submit_func,
extract_external_task_id=extract_external_task_id,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
)
finally:
# Restore Session behavior for the rest of the request lifecycle.
self.db.expire_on_commit = original_expire_on_commit

View File

View File

@@ -0,0 +1,320 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
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, Usage, User
from src.services.usage.service import UsageService
class VideoTaskBillingService:
"""视频任务计费/结算服务。"""
def __init__(self, db: Session) -> None:
self.db = db
async def _create_fallback_usage_for_video_task(self, task: Any, request_id: str) -> bool:
"""
Fallback: create a Usage row if it's missing (should be rare).
This keeps behavior compatible with the old Phase2 finalize logic.
"""
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 getattr(task, "api_key_id", None)
else None
)
provider_obj = (
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
if getattr(task, "provider_id", None)
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
response_time_ms: int | None = None
if getattr(task, "submitted_at", None) and getattr(task, "completed_at", None):
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
request_headers: dict[str, Any] | None = None
if isinstance(getattr(task, "request_metadata", None), dict):
task_meta = task.request_metadata
for header_key in ("request_headers", "headers", "original_headers"):
raw_headers = task_meta.get(header_key)
if isinstance(raw_headers, dict):
request_headers = dict(raw_headers)
break
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(getattr(task, "format_converted", False)),
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=request_headers,
request_body=getattr(task, "original_request_body", None),
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=getattr(task, "provider_id", None),
provider_endpoint_id=getattr(task, "endpoint_id", None),
provider_api_key_id=getattr(task, "key_id", None),
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={}: {}",
task.id,
str(exc),
)
return False
async def finalize_video_task(self, task: Any) -> bool:
"""
Update billing/usage for a completed/failed video task.
Async video billing flow:
- Submit success: Usage is already settled with cost=0
- Poll completion: update actual cost (success -> bill, failure -> keep 0)
Returns True when updated, False when skipped (already finalized).
"""
from datetime import datetime, timezone
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
request_id = getattr(task, "request_id", None) or getattr(task, "id", None)
if not request_id:
return False
existing = self.db.query(Usage).filter(Usage.request_id == request_id).first()
if not existing:
logger.warning(
"Usage not found for video task, creating fallback: task_id={} request_id={}",
getattr(task, "id", None),
request_id,
)
return await self._create_fallback_usage_for_video_task(task, request_id)
metadata = existing.request_metadata or {}
if metadata.get("billing_updated_at"):
logger.debug(
"Video task billing already updated: task_id={} request_id={}",
getattr(task, "id", None),
request_id,
)
return False
response_time_ms: int | None = None
if getattr(task, "submitted_at", None) and getattr(task, "completed_at", None):
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
base_dimensions: dict[str, Any] = {
"duration_seconds": getattr(task, "duration_seconds", None),
"resolution": getattr(task, "resolution", None),
"aspect_ratio": getattr(task, "aspect_ratio", None),
"size": getattr(task, "size", None) or "",
"retry_count": getattr(task, "retry_count", 0),
}
collector_metadata: dict[str, Any] = {
"task": {
"id": getattr(task, "id", None),
"external_task_id": getattr(task, "external_task_id", None),
"model": getattr(task, "model", None),
"duration_seconds": getattr(task, "duration_seconds", None),
"resolution": getattr(task, "resolution", None),
"aspect_ratio": getattr(task, "aspect_ratio", None),
"size": getattr(task, "size", None),
"retry_count": getattr(task, "retry_count", 0),
"video_size_bytes": getattr(task, "video_size_bytes", None),
},
"result": {
"video_url": getattr(task, "video_url", None),
"video_urls": getattr(task, "video_urls", None) or [],
},
}
poll_raw = None
if isinstance(getattr(task, "request_metadata", None), dict):
poll_raw = task.request_metadata.get("poll_raw_response")
dims = DimensionCollectorService(self.db).collect_dimensions(
api_format=getattr(task, "provider_api_format", None),
task_type="video",
request=getattr(task, "original_request_body", None) or {},
response=poll_raw if isinstance(poll_raw, dict) else None,
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# Prefer frozen rule snapshot from submit stage.
rule_snapshot = None
if isinstance(getattr(task, "request_metadata", None), 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=getattr(task, "provider_id", None),
model_name=getattr(task, "model", None),
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = getattr(rule, "id", None)
rule_name = getattr(rule, "name", None)
rule_scope = getattr(lookup, "scope", None)
expression = getattr(rule, "expression", None)
variables = getattr(rule, "variables", None) or {}
dimension_mappings = getattr(rule, "dimension_mappings", None) 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
is_success = str(getattr(task, "status", "")) in {
VideoStatus.COMPLETED.value,
"completed",
}
if is_success 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: mark task failed and hide artifacts (avoid free pass)
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
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["cost"] = 0.0
cost = 0.0
except Exception as exc:
billing_snapshot["status"] = "incomplete"
billing_snapshot["error"] = str(exc)
billing_snapshot["cost"] = 0.0
cost = 0.0
# Write back to task.request_metadata for audit/recalc.
task_meta = dict(task.request_metadata) if getattr(task, "request_metadata", None) else {}
task_meta["billing_snapshot"] = billing_snapshot
task.request_metadata = task_meta
updated = UsageService.update_settled_billing(
self.db,
request_id=request_id,
total_cost_usd=cost,
request_cost_usd=cost,
status="completed" if str(getattr(task, "status", "")) == "completed" else "failed",
status_code=200 if str(getattr(task, "status", "")) == "completed" else 500,
error_message=(
None
if str(getattr(task, "status", "")) == "completed"
else (
getattr(task, "error_message", None)
or getattr(task, "error_code", None)
or "video_task_failed"
)
),
response_time_ms=response_time_ms,
billing_snapshot=billing_snapshot,
extra_metadata={
"dimensions": dims,
"raw_response_ref": {
"video_task_id": getattr(task, "id", None),
"field": "video_tasks.request_metadata.poll_raw_response",
},
},
)
if updated:
logger.debug(
"Updated video task billing: task_id={} request_id={} cost={:.6f}",
getattr(task, "id", None),
request_id,
cost,
)
else:
logger.warning(
"Failed to update video task billing (may already be updated): "
"task_id={} request_id={}",
getattr(task, "id", None),
request_id,
)
return bool(updated)

View File

@@ -0,0 +1,149 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ProviderAPIKey, ProviderEndpoint
from src.services.usage.service import UsageService
class VideoTaskCancelService:
"""视频任务取消服务(上游取消 + 本地状态与计费回写)。"""
def __init__(self, db: Session) -> None:
self.db = db
async def cancel_task(
self,
*,
task: Any,
task_id: str,
original_headers: dict[str, str] | None = None,
) -> Any:
"""
Cancel a video task (best-effort) and void its Usage (no charge).
Returns:
- None on success
- upstream httpx.Response when upstream returns an error (status >= 400)
"""
from datetime import datetime, timezone
from fastapi import HTTPException
from src.clients.http_client import HTTPClientPool
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 VideoStatus
from src.core.crypto import crypto_service
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
external_task_id = getattr(task, "external_task_id", None)
if not external_task_id:
raise HTTPException(status_code=500, detail="Task missing external_task_id")
endpoint = (
self.db.query(ProviderEndpoint).filter(ProviderEndpoint.id == task.endpoint_id).first()
)
key = self.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
if not endpoint or not key:
raise HTTPException(status_code=500, detail="Provider endpoint or key not found")
if not getattr(key, "api_key", None):
raise HTTPException(status_code=500, detail="Provider key not configured")
upstream_key = crypto_service.decrypt(key.api_key)
extra_headers = get_extra_headers_from_endpoint(endpoint)
raw_family = str(getattr(endpoint, "api_family", "") or "").strip().lower()
raw_kind = str(getattr(endpoint, "endpoint_kind", "") or "").strip().lower()
provider_format = (
make_signature_key(raw_family, raw_kind)
if raw_family and raw_kind
else str(
getattr(endpoint, "api_format", "")
or getattr(task, "provider_api_format", "")
or ""
)
)
provider_format_norm = provider_format.strip().lower()
headers = build_upstream_headers_for_endpoint(
original_headers or {},
provider_format,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
client = await HTTPClientPool.get_default_client_async()
if provider_format_norm.startswith("openai:"):
upstream_url = build_provider_url(endpoint, is_stream=False, key=key)
upstream_url = f"{upstream_url.rstrip('/')}/{str(external_task_id).lstrip('/')}"
response = await client.delete(upstream_url, headers=headers)
if response.status_code >= 400:
return response
elif provider_format_norm.startswith("gemini:"):
# Gemini cancel endpoint supports both:
# - operations/{id}:cancel
# - models/{model}/operations/{id}:cancel
operation_name = str(external_task_id)
if not (
operation_name.startswith("operations/") or operation_name.startswith("models/")
):
operation_name = f"operations/{operation_name}"
base = (
getattr(endpoint, "base_url", None) or "https://generativelanguage.googleapis.com"
).rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
upstream_url = f"{base}/v1beta/{operation_name}:cancel"
auth_info = await get_provider_auth(endpoint, key)
if auth_info:
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
response = await client.post(upstream_url, headers=headers, json={})
if response.status_code >= 400:
return response
else:
raise HTTPException(
status_code=400,
detail=f"Cancel not supported for provider format: {provider_format}",
)
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
# Void Usage (no charge)
try:
voided = UsageService.finalize_void(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
)
if not voided:
UsageService.void_settled(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
)
except Exception as exc:
logger.warning(
"Failed to void usage for cancelled task={}: {}",
getattr(task, "id", task_id),
str(exc),
)
self.db.commit()
return None

View File

@@ -0,0 +1,38 @@
from __future__ import annotations
from typing import Any
from src.services.task.core.schema import TaskStatusResult
from src.services.task.video.operations import VideoTaskOperationsService
class TaskVideoFacadeService:
"""视频任务门面服务(向后兼容 TaskService 的视频公开方法)。"""
def __init__(self, video_ops: VideoTaskOperationsService) -> None:
self._video_ops = video_ops
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
return await self._video_ops.poll(task_id, user_id=user_id)
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
return await self._video_ops.poll_now(task_id, user_id=user_id)
async def cancel(
self,
task_id: str,
*,
user_id: str,
original_headers: dict[str, str] | None = None,
) -> Any:
return await self._video_ops.cancel(
task_id,
user_id=user_id,
original_headers=original_headers,
)
async def finalize_video_task(self, task: Any) -> bool:
return await self._video_ops.finalize_video_task(task)
async def finalize(self, task_id: str) -> bool:
return await self._video_ops.finalize(task_id)

View File

@@ -0,0 +1,134 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import VideoTask
from src.services.task.core.exceptions import TaskNotFoundError
from src.services.task.core.schema import TaskStatusResult
from src.services.task.video.billing import VideoTaskBillingService
from src.services.task.video.cancel import VideoTaskCancelService
class VideoTaskOperationsService:
"""视频任务相关应用服务(轮询/取消/终态结算)。"""
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
self._billing_ops = VideoTaskBillingService(db)
self._cancel_ops = VideoTaskCancelService(db)
def _extract_short_id(self, task_id: str) -> str:
# Keep the parsing rule consistent with handlers:
# - models/{model}/operations/{short_id}
# - operations/{short_id}
# - {short_id}
return task_id.rsplit("/", 1)[-1] if "/" in task_id else task_id
def _get_video_task_for_user(self, task_id: str, *, user_id: str) -> Any:
"""
Resolve a video task by:
- internal UUID (VideoTask.id)
- external operation id (VideoTask.short_id)
"""
task = (
self.db.query(VideoTask)
.filter(VideoTask.id == task_id, VideoTask.user_id == user_id)
.first()
)
if task:
return task
short_id = self._extract_short_id(task_id)
task = (
self.db.query(VideoTask)
.filter(VideoTask.short_id == short_id, VideoTask.user_id == user_id)
.first()
)
if not task:
raise TaskNotFoundError(task_id)
return task
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
"""Read task status from DB (does not trigger polling)."""
task = self._get_video_task_for_user(task_id, user_id=user_id)
result_url = None
if getattr(task, "status", None) == "completed":
result_url = getattr(task, "video_url", None)
error_message = None
if getattr(task, "status", None) == "failed":
error_message = getattr(task, "error_message", None) or getattr(
task, "error_code", None
)
return TaskStatusResult(
task_id=str(getattr(task, "id", task_id)),
status=str(getattr(task, "status", "unknown")),
progress_percent=int(getattr(task, "progress_percent", 0) or 0),
result_url=result_url,
error_message=str(error_message) if error_message else None,
provider_id=(
str(getattr(task, "provider_id", None))
if getattr(task, "provider_id", None)
else None
),
provider_name=(
str(getattr(task, "provider_name", None))
if getattr(task, "provider_name", None)
else None
),
endpoint_id=(
str(getattr(task, "endpoint_id", None))
if getattr(task, "endpoint_id", None)
else None
),
key_id=str(getattr(task, "key_id", None)) if getattr(task, "key_id", None) else None,
)
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
"""
Trigger a single polling attempt (best-effort), then return latest DB status.
Note: this uses the poller adapter's single-task method and may hold a DB
connection during the upstream HTTP request; keep usage low.
"""
from src.services.task.video.poller_adapter import VideoTaskPollerAdapter
task = self._get_video_task_for_user(task_id, user_id=user_id)
adapter = VideoTaskPollerAdapter()
await adapter.poll_single_task(self.db, task, redis_client=self.redis)
self.db.commit()
return await self.poll(task_id, user_id=user_id)
async def cancel(
self,
task_id: str,
*,
user_id: str,
original_headers: dict[str, str] | None = None,
) -> Any:
from fastapi import HTTPException
try:
task = self._get_video_task_for_user(task_id, user_id=user_id)
except TaskNotFoundError:
raise HTTPException(status_code=404, detail="Video task not found")
return await self._cancel_ops.cancel_task(
task=task,
task_id=task_id,
original_headers=original_headers,
)
async def finalize_video_task(self, task: Any) -> bool:
return await self._billing_ops.finalize_video_task(task)
async def finalize(self, task_id: str) -> bool:
"""Finalize a task by internal id (best-effort)."""
task = self.db.query(VideoTask).filter(VideoTask.id == task_id).first()
if not task:
return False
return await self.finalize_video_task(task)

View File

@@ -10,6 +10,7 @@ Implements the video-specific poll/normalize/update logic used by TaskPollerServ
from __future__ import annotations from __future__ import annotations
import json import json
from collections.abc import Awaitable, Callable
from dataclasses import dataclass from dataclasses import dataclass
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import Any from typing import Any
@@ -36,7 +37,6 @@ from src.core.video_utils import (
from src.database import create_session from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.provider.auth import get_provider_auth from src.services.provider.auth import get_provider_auth
from src.services.task.service import TaskService
@dataclass(slots=True) @dataclass(slots=True)
@@ -84,6 +84,20 @@ class PollHTTPError(RuntimeError):
self.original_message = message self.original_message = message
VideoTaskFinalizeFn = Callable[[Session, VideoTask, Any | None], Awaitable[None]]
async def _default_finalize_video_task(
db: Session,
task: VideoTask,
redis_client: Any | None,
) -> None:
"""默认终态结算逻辑(延迟导入,避免 task 模块循环依赖)。"""
from src.services.task.video.operations import VideoTaskOperationsService
await VideoTaskOperationsService(db, redis_client=redis_client).finalize_video_task(task)
class VideoTaskPollerAdapter: class VideoTaskPollerAdapter:
task_type = "video" task_type = "video"
@@ -102,9 +116,10 @@ class VideoTaskPollerAdapter:
consecutive_failure_alert_threshold = 5 consecutive_failure_alert_threshold = 5
max_backoff_seconds = 300 max_backoff_seconds = 300
def __init__(self) -> None: def __init__(self, finalize_video_task_fn: VideoTaskFinalizeFn | None = None) -> None:
self._openai_normalizer = OpenAINormalizer() self._openai_normalizer = OpenAINormalizer()
self._gemini_normalizer = GeminiNormalizer() self._gemini_normalizer = GeminiNormalizer()
self._finalize_video_task = finalize_video_task_fn or _default_finalize_video_task
def sanitize_error_message(self, message: str) -> str: def sanitize_error_message(self, message: str) -> str:
return sanitize_error_message(message) return sanitize_error_message(message)
@@ -287,7 +302,7 @@ class VideoTaskPollerAdapter:
# 终态结算 # 终态结算
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value): if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try: try:
await TaskService(db, redis_client=redis_client).finalize_video_task(task) await self._finalize_video_task(db, task, redis_client)
except Exception as exc: except Exception as exc:
logger.exception( logger.exception(
"Failed to record video usage for task={}: {}", "Failed to record video usage for task={}: {}",
@@ -365,78 +380,37 @@ class VideoTaskPollerAdapter:
self, db: Session, task: VideoTask, *, redis_client: Any | None self, db: Session, task: VideoTask, *, redis_client: Any | None
) -> None: ) -> None:
""" """
旧版单任务轮询方法保留向后兼容 兼容入口复用三阶段轮询流程避免维护重复逻辑
注意此方法在 HTTP 请求期间持有数据库连接建议使用分阶段方法
""" """
ctx_or_result = await self.prepare_poll_context(db, task)
if isinstance(ctx_or_result, InternalVideoPollResult):
await self.update_task_after_poll(
task_id=task.id,
result=ctx_or_result,
ctx=None,
redis_client=redis_client,
)
return
ctx = ctx_or_result
error_exception: Exception | None = None
try: try:
result = await self._poll_task_status(db, task) result = await self.poll_task_http(ctx)
if result.status == VideoStatus.COMPLETED: except Exception as http_exc:
task.status = VideoStatus.COMPLETED.value error_exception = http_exc
task.video_url = result.video_url result = InternalVideoPollResult(
task.video_expires_at = result.expires_at status=None, # type: ignore[arg-type]
task.completed_at = datetime.now(timezone.utc) error_message=str(http_exc),
task.progress_percent = 100 )
if result.video_urls:
task.video_urls = result.video_urls
if result.video_duration_seconds is not None:
task.video_duration_seconds = result.video_duration_seconds
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 {}: {}", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None await self.update_task_after_poll(
is_permanent = self._is_permanent_error(exc, status_code=status_code) task_id=task.id,
if is_permanent: result=result,
task.status = VideoStatus.FAILED.value ctx=ctx,
task.error_code = "poll_permanent_error" redis_client=redis_client,
task.error_message = error_msg error_exception=error_exception,
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 TaskService(db, redis_client=redis_client).finalize_video_task(task)
except Exception as exc:
logger.exception(
"Failed to record video usage for task={}: {}",
task.id,
sanitize_error_message(str(exc)),
)
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None: def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
if not result.raw_response: if not result.raw_response:

View File

@@ -6,8 +6,13 @@ import httpx
import pytest import pytest
from src.config.settings import config from src.config.settings import config
from src.services.candidate.submit import AllCandidatesFailedError, UpstreamClientRequestError from src.services.candidate.submit import (
AllCandidatesFailedError,
SubmitOutcome,
UpstreamClientRequestError,
)
from src.services.task.service import TaskService from src.services.task.service import TaskService
from src.services.task.submit.outcome_builder import SubmitPayloadParseResult
def _make_candidate( def _make_candidate(
@@ -356,7 +361,7 @@ async def test_submit_with_failover_applies_pool_reorder_before_submit(
reordered = [pool_candidate_b, pool_candidate_a] reordered = [pool_candidate_b, pool_candidate_a]
apply_pool_reorder = AsyncMock(return_value=(reordered, [])) apply_pool_reorder = AsyncMock(return_value=(reordered, []))
monkeypatch.setattr(svc, "_apply_pool_reorder", apply_pool_reorder) monkeypatch.setattr(svc._submit_ops, "_apply_pool_reorder", apply_pool_reorder) # type: ignore[attr-defined]
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-pooled"})) submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-pooled"}))
body = {"session_id": "sid-123"} body = {"session_id": "sid-123"}
@@ -379,8 +384,245 @@ async def test_submit_with_failover_applies_pool_reorder_before_submit(
assert outcome.candidate.key.id == "k-b" assert outcome.candidate.key.id == "k-b"
assert submit.await_count == 1 assert submit.await_count == 1
fetch_candidates.assert_awaited_once() fetch_candidates.assert_awaited_once()
assert fetch_candidates.await_args.kwargs.get("request_body") == body await_args = fetch_candidates.await_args
assert await_args is not None
assert await_args.kwargs.get("request_body") == body
apply_pool_reorder.assert_awaited_once_with( apply_pool_reorder.assert_awaited_once_with(
[pool_candidate_a, pool_candidate_b], [pool_candidate_a, pool_candidate_b],
request_body=body, request_body=body,
) )
@pytest.mark.asyncio
async def test_submit_with_failover_orchestrates_prepare_and_execute_layers() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p2", endpoint_id="e2", key_id="k2")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-1"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p2", "selected": True}],
external_task_id="task-layered",
rule_lookup=None,
upstream_payload={"id": "task-layered"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops.execute_submit_loop = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=outcome
)
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-1",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-layered"
svc._submit_ops._prepare_ops.prepare_candidates.assert_awaited_once() # type: ignore[attr-defined]
svc._submit_ops._execute_ops.execute_submit_loop.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_submit_with_failover_execute_layer_orchestrates_filter_and_attempt() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p2", endpoint_id="e2", key_id="k2")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-1"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p2", "selected": True}],
external_task_id="task-inner",
rule_lookup=None,
upstream_payload={"id": "task-inner"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops._filter_ops.build_candidate_info = MagicMock( # type: ignore[attr-defined, method-assign]
return_value={
"index": 0,
"provider_id": "p2",
"provider_name": "prov",
"endpoint_id": "e2",
"key_id": "k2",
"key_name": "key",
"auth_type": "api_key",
"priority": 0,
"is_cached": False,
}
)
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SimpleNamespace(record_id="rc-1", rule_lookup=None)
)
svc._submit_ops._execute_ops._attempt_ops.submit_candidate = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=(outcome, 200)
)
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-2",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-inner"
svc._submit_ops._execute_ops._filter_ops.build_candidate_info.assert_called_once() # type: ignore[attr-defined]
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt.assert_called_once() # type: ignore[attr-defined]
svc._submit_ops._execute_ops._attempt_ops.submit_candidate.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_submit_with_failover_attempt_layer_delegates_to_response_ops() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p3", endpoint_id="e3", key_id="k3")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-3"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p3", "selected": True}],
external_task_id="task-attempt",
rule_lookup=None,
upstream_payload={"id": "task-attempt"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops._filter_ops.build_candidate_info = MagicMock( # type: ignore[attr-defined, method-assign]
return_value={
"index": 0,
"provider_id": "p3",
"provider_name": "prov3",
"endpoint_id": "e3",
"key_id": "k3",
"key_name": "key3",
"auth_type": "api_key",
"priority": 0,
"is_cached": False,
}
)
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SimpleNamespace(record_id="rc-3", rule_lookup=None)
)
svc._submit_ops._execute_ops._attempt_ops._record_ops.mark_pending = MagicMock() # type: ignore[attr-defined, method-assign]
svc._submit_ops._execute_ops._attempt_ops._response_ops.handle_submit_response = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=(outcome, 200)
)
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-attempt"}))
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-3",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-attempt"
svc._submit_ops._execute_ops._attempt_ops._record_ops.mark_pending.assert_called_once() # type: ignore[attr-defined]
submit.assert_awaited_once_with(candidate)
svc._submit_ops._execute_ops._attempt_ops._response_ops.handle_submit_response.assert_called_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_submit_with_failover_response_layer_delegates_to_decider_and_builder() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p4", endpoint_id="e4", key_id="k4")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-4"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p4", "selected": True}],
external_task_id="task-response",
rule_lookup=None,
upstream_payload={"id": "task-response"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops._filter_ops.build_candidate_info = MagicMock( # type: ignore[attr-defined, method-assign]
return_value={
"index": 0,
"provider_id": "p4",
"provider_name": "prov4",
"endpoint_id": "e4",
"key_id": "k4",
"key_name": "key4",
"auth_type": "api_key",
"priority": 0,
"is_cached": False,
}
)
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SimpleNamespace(record_id="rc-4", rule_lookup=None)
)
response_ops = svc._submit_ops._execute_ops._attempt_ops._response_ops # type: ignore[attr-defined]
response_ops._rule_decider.detect_error_stop_pattern = MagicMock(return_value=None) # type: ignore[attr-defined, method-assign]
response_ops._rule_decider.detect_success_failover_pattern = MagicMock(return_value=None) # type: ignore[attr-defined, method-assign]
response_ops._outcome_builder.parse_payload = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SubmitPayloadParseResult(payload={"id": "task-response"})
)
response_ops._outcome_builder.build_success_outcome = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=outcome
)
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-response"}))
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-4",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-response"
response_ops._rule_decider.detect_error_stop_pattern.assert_not_called() # type: ignore[attr-defined]
response_ops._rule_decider.detect_success_failover_pattern.assert_called_once() # type: ignore[attr-defined]
response_ops._outcome_builder.parse_payload.assert_called_once() # type: ignore[attr-defined]
response_ops._outcome_builder.build_success_outcome.assert_called_once() # type: ignore[attr-defined]

View File

@@ -1,16 +1,17 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, AsyncIterator from typing import Any, AsyncIterator, cast
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.candidate.failover import FailoverEngine from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
from src.services.orchestration.error_classifier import ErrorAction from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
from src.services.scheduling.schemas import PoolCandidate from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
from src.services.task.protocol import AttemptKind, AttemptResult from src.services.task.core.protocol import AttemptKind, AttemptResult
def _make_candidate( def _make_candidate(
@@ -28,16 +29,22 @@ def _make_candidate(
needs_conversion: bool = False, needs_conversion: bool = False,
provider_max_retries: int | None = None, provider_max_retries: int | None = None,
provider_config: dict[str, Any] | None = None, provider_config: dict[str, Any] | None = None,
) -> SimpleNamespace: ) -> ProviderCandidate:
provider = SimpleNamespace( provider = cast(
id=provider_id, Provider,
name=provider_name, SimpleNamespace(
max_retries=provider_max_retries, id=provider_id,
config=provider_config, name=provider_name,
max_retries=provider_max_retries,
config=provider_config,
),
) )
endpoint = SimpleNamespace(id=endpoint_id) endpoint = cast(ProviderEndpoint, SimpleNamespace(id=endpoint_id))
key = SimpleNamespace(id=key_id, name=key_name, auth_type=auth_type, priority=priority) key = cast(
return SimpleNamespace( ProviderAPIKey,
SimpleNamespace(id=key_id, name=key_name, auth_type=auth_type, priority=priority),
)
return ProviderCandidate(
provider=provider, provider=provider,
endpoint=endpoint, endpoint=endpoint,
key=key, key=key,
@@ -62,10 +69,14 @@ class _StubErrorClassifier:
return self._action return self._action
def _stub_classifier(*, action: ErrorAction, client_error: bool = False) -> ErrorClassifier:
return cast(ErrorClassifier, _StubErrorClassifier(action=action, client_error=client_error))
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_failover_engine_success_first_candidate() -> None: async def test_failover_engine_success_first_candidate() -> None:
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")] candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
@@ -97,7 +108,7 @@ async def test_failover_engine_success_first_candidate() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_failover_engine_continue_to_next_candidate_on_error() -> None: async def test_failover_engine_continue_to_next_candidate_on_error() -> None:
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")] candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
@@ -132,7 +143,7 @@ async def test_failover_engine_continue_to_next_candidate_on_error() -> None:
async def test_failover_engine_retry_same_candidate_when_classifier_says_continue() -> None: async def test_failover_engine_retry_same_candidate_when_classifier_says_continue() -> None:
db = MagicMock() db = MagicMock()
# ErrorAction.CONTINUE => retry current candidate (mapped to FailoverAction.RETRY) # ErrorAction.CONTINUE => retry current candidate (mapped to FailoverAction.RETRY)
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.CONTINUE)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.CONTINUE))
candidates = [_make_candidate(provider_id="p1", is_cached=True, provider_max_retries=2)] candidates = [_make_candidate(provider_id="p1", is_cached=True, provider_max_retries=2)]
@@ -167,7 +178,7 @@ async def test_failover_engine_continues_when_classifier_raises() -> None:
"""After the 'default failover' change, RAISE no longer stops failover. """After the 'default failover' change, RAISE no longer stops failover.
All candidates should be attempted.""" All candidates should be attempted."""
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.RAISE)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.RAISE))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")] candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
attempt = AsyncMock(side_effect=RuntimeError("client-ish")) attempt = AsyncMock(side_effect=RuntimeError("client-ish"))
@@ -200,7 +211,7 @@ async def _empty_stream() -> AsyncIterator[bytes]:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_failover_engine_stream_probe_wraps_first_chunk() -> None: async def test_failover_engine_stream_probe_wraps_first_chunk() -> None:
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1")] candidates = [_make_candidate(provider_id="p1")]
attempt = AsyncMock( attempt = AsyncMock(
@@ -234,7 +245,7 @@ async def test_failover_engine_stream_probe_wraps_first_chunk() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_failover_engine_stream_probe_empty_triggers_failover() -> None: async def test_failover_engine_stream_probe_empty_triggers_failover() -> None:
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")] candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
@@ -273,7 +284,7 @@ async def test_failover_engine_pre_expand_marks_unused_slots_on_success(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
# patch low-level record updater to observe unused marking # patch low-level record updater to observe unused marking
engine._update_record = MagicMock() # type: ignore[method-assign] engine._update_record = MagicMock() # type: ignore[method-assign]
@@ -333,7 +344,7 @@ class _HttpError(Exception):
async def test_error_stop_pattern_with_matching_status_code_stops_failover() -> None: async def test_error_stop_pattern_with_matching_status_code_stops_failover() -> None:
"""When status_codes is set and matches, failover should stop.""" """When status_codes is set and matches, failover should stop."""
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
config = { config = {
"failover_rules": { "failover_rules": {
@@ -376,7 +387,7 @@ async def test_error_stop_pattern_with_matching_status_code_stops_failover() ->
async def test_error_stop_pattern_with_non_matching_status_code_continues() -> None: async def test_error_stop_pattern_with_non_matching_status_code_continues() -> None:
"""When status_codes is set but doesn't match, the rule is skipped and failover continues.""" """When status_codes is set but doesn't match, the rule is skipped and failover continues."""
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
config = { config = {
"failover_rules": { "failover_rules": {
@@ -420,7 +431,7 @@ async def test_error_stop_pattern_with_non_matching_status_code_continues() -> N
async def test_error_stop_pattern_without_status_codes_matches_any() -> None: async def test_error_stop_pattern_without_status_codes_matches_any() -> None:
"""When status_codes is not set, the rule matches any status code (existing behavior).""" """When status_codes is not set, the rule matches any status code (existing behavior)."""
db = MagicMock() db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK)) engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
config = { config = {
"failover_rules": { "failover_rules": {

View File

@@ -0,0 +1,132 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from src.services.candidate.policy import FailoverAction
from src.services.task.execute.exception_classification import (
CandidateErrorAction,
classify_candidate_error_action,
)
from src.services.task.execute.state_transition import (
SyncExecutionState,
resolve_execution_error_transition,
)
def _make_candidate() -> SimpleNamespace:
return SimpleNamespace(
provider=SimpleNamespace(id="p1", name="provider-1"),
endpoint=SimpleNamespace(id="e1"),
key=SimpleNamespace(id="k1"),
)
@pytest.mark.parametrize(
("raw_action", "expected"),
[
("continue", CandidateErrorAction.RETRY_CURRENT),
("break", CandidateErrorAction.NEXT_CANDIDATE),
("raise", CandidateErrorAction.RAISE_ERROR),
("unexpected", CandidateErrorAction.NEXT_CANDIDATE),
(None, CandidateErrorAction.NEXT_CANDIDATE),
],
)
def test_classify_candidate_error_action(
raw_action: str | None, expected: CandidateErrorAction
) -> None:
assert classify_candidate_error_action(raw_action) == expected
def test_resolve_execution_error_transition_retry_and_consume_rectify_flag() -> None:
request_body_ref = {"_rectified_this_turn": True}
state = SyncExecutionState(
candidate_record_map={},
request_body_ref=request_body_ref,
)
transition = resolve_execution_error_transition(
action=CandidateErrorAction.RETRY_CURRENT,
state=state,
max_retries_for_candidate=2,
retry_index=1,
)
assert transition.failover_action == FailoverAction.RETRY
assert transition.max_retries == 3
assert request_body_ref["_rectified_this_turn"] is False
def test_resolve_execution_error_transition_next_candidate() -> None:
state = SyncExecutionState(
candidate_record_map={},
request_body_ref={"_rectified_this_turn": True},
)
transition = resolve_execution_error_transition(
action=CandidateErrorAction.NEXT_CANDIDATE,
state=state,
max_retries_for_candidate=2,
retry_index=0,
)
assert transition.failover_action == FailoverAction.CONTINUE
assert transition.max_retries is None
assert state.request_body_ref == {"_rectified_this_turn": True}
def test_sync_execution_state_resolve_candidate_record_id_fallback() -> None:
state = SyncExecutionState(
candidate_record_map={(2, 0): "r20"},
request_body_ref=None,
)
assert state.resolve_candidate_record_id(candidate_index=2, record_id=None) == "r20"
assert state.resolve_candidate_record_id(candidate_index=2, record_id="r22") == "r22"
def test_sync_execution_state_raise_classified_error_uses_last_error() -> None:
err = ValueError("boom")
candidate = _make_candidate()
state = SyncExecutionState(
candidate_record_map={},
request_body_ref=None,
last_error=err,
last_candidate=candidate,
)
failure_ops = MagicMock()
with pytest.raises(ValueError, match="boom"):
state.raise_classified_error(
fallback_error=RuntimeError("fallback"),
failure_ops=failure_ops,
model_name="gpt-4.1",
api_format="openai_chat",
)
failure_ops.attach_metadata_to_error.assert_called_once_with(
err,
candidate,
"gpt-4.1",
"openai_chat",
)
def test_sync_execution_state_raise_classified_error_fallback_error() -> None:
state = SyncExecutionState(
candidate_record_map={},
request_body_ref=None,
)
failure_ops = MagicMock()
with pytest.raises(RuntimeError, match="fallback"):
state.raise_classified_error(
fallback_error=RuntimeError("fallback"),
failure_ops=failure_ops,
model_name="gpt-4.1",
api_format="openai_chat",
)
failure_ops.attach_metadata_to_error.assert_not_called()

View File

@@ -0,0 +1,118 @@
from __future__ import annotations
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
from src.services.task.execute.pool import TaskPoolOperationsService
def _provider(provider_id: str) -> Provider:
return cast(Provider, SimpleNamespace(id=provider_id))
def _endpoint(endpoint_id: str) -> ProviderEndpoint:
return cast(ProviderEndpoint, SimpleNamespace(id=endpoint_id))
def _key(key_id: str) -> ProviderAPIKey:
return cast(ProviderAPIKey, SimpleNamespace(id=key_id))
def test_extract_session_uuid_returns_none_for_non_dict_request_body() -> None:
svc = TaskPoolOperationsService()
assert svc.extract_session_uuid("openai", None) is None
assert svc.extract_session_uuid("openai", cast(Any, "not-dict")) is None
def test_extract_session_uuid_uses_pool_hook(monkeypatch: pytest.MonkeyPatch) -> None:
from src.services.provider.pool import hooks
hook = SimpleNamespace(extract_session_uuid=lambda body: f"sid:{body.get('session')}")
def _get_pool_hook(_provider_type: str) -> Any:
return hook
monkeypatch.setattr(hooks, "get_pool_hook", _get_pool_hook)
svc = TaskPoolOperationsService()
session_id = svc.extract_session_uuid("claude_code", {"session": "abc"})
assert session_id == "sid:abc"
def test_expand_pool_candidates_for_async_submit_keeps_non_pool_candidate() -> None:
provider = _provider("p1")
endpoint = _endpoint("e1")
key = _key("k1")
candidate = ProviderCandidate(provider=provider, endpoint=endpoint, key=key)
svc = TaskPoolOperationsService()
expanded = svc.expand_pool_candidates_for_async_submit([candidate])
assert len(expanded) == 1
assert expanded[0] is candidate
def test_expand_pool_candidates_for_async_submit_expands_pool_keys() -> None:
provider = _provider("p1")
endpoint = _endpoint("e1")
key = _key("k0")
pool_key_1 = cast(
ProviderAPIKey,
SimpleNamespace(
id="k1",
_pool_skipped=False,
_pool_mapping_matched_model="mapped-model",
_pool_extra_data={"source": "warm"},
),
)
pool_key_2 = cast(
ProviderAPIKey,
SimpleNamespace(
id="k2",
_pool_skipped=True,
_pool_skip_reason="cooldown",
_pool_extra_data={"reason_code": "429"},
),
)
pool_candidate = PoolCandidate(
provider=provider,
endpoint=endpoint,
key=key,
is_cached=True,
is_skipped=False,
skip_reason=None,
mapping_matched_model="fallback-model",
needs_conversion=True,
provider_api_format="openai:chat",
output_limit=1024,
capability_miss_count=1,
pool_keys=[pool_key_1, pool_key_2],
)
svc = TaskPoolOperationsService()
expanded = svc.expand_pool_candidates_for_async_submit([pool_candidate])
assert len(expanded) == 2
first, second = expanded
assert first.key is pool_key_1
assert first.is_skipped is False
assert first.mapping_matched_model == "mapped-model"
first_extra = getattr(first, "_pool_extra_data")
assert first_extra["pool_group_id"] == "p1"
assert first_extra["pool_key_index"] == 0
assert first_extra["source"] == "warm"
assert second.key is pool_key_2
assert second.is_skipped is True
assert second.skip_reason == "cooldown"
assert second.mapping_matched_model == "fallback-model"
second_extra = getattr(second, "_pool_extra_data")
assert second_extra["pool_group_id"] == "p1"
assert second_extra["pool_key_index"] == 1
assert second_extra["reason_code"] == "429"

View File

@@ -1,4 +1,4 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
@@ -8,10 +8,9 @@ import pytest
from src.core.exceptions import EmbeddedErrorException from src.core.exceptions import EmbeddedErrorException
from src.services.candidate.schema import CandidateKey from src.services.candidate.schema import CandidateKey
from src.services.candidate.submit import SubmitOutcome from src.services.candidate.submit import SubmitOutcome
from src.services.request.executor import ExecutionContext, ExecutionError from src.services.task.core.context import TaskMode
from src.services.task import service as task_service_module from src.services.task.core.protocol import AttemptKind
from src.services.task.context import TaskMode from src.services.task.service import pool_on_error
from src.services.task.protocol import AttemptKind
from src.services.task.service import TaskService from src.services.task.service import TaskService
@@ -51,7 +50,7 @@ async def test_task_service_execute_async_returns_execution_result() -> None:
) )
svc.submit_with_failover = AsyncMock(return_value=outcome) # type: ignore[method-assign] svc.submit_with_failover = AsyncMock(return_value=outcome) # type: ignore[method-assign]
svc._recorder.get_candidate_keys = MagicMock( # type: ignore[attr-defined, method-assign] svc._execute_facade_ops._get_candidate_keys = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=[ return_value=[
CandidateKey(candidate_index=0, retry_index=0, status="success", provider_id="p1") CandidateKey(candidate_index=0, retry_index=0, status="success", provider_id="p1")
] ]
@@ -82,7 +81,7 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
db = MagicMock() db = MagicMock()
svc = TaskService(db) svc = TaskService(db)
sentinel_result = object() sentinel_result = object()
svc._execute_sync_unified = AsyncMock( # type: ignore[method-assign] svc._sync_ops.execute_sync_unified = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=sentinel_result return_value=sentinel_result
) )
@@ -105,95 +104,92 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
) )
assert result is sentinel_result assert result is sentinel_result
svc._execute_sync_unified.assert_awaited_once() # type: ignore[attr-defined] svc._sync_ops.execute_sync_unified.assert_awaited_once() # type: ignore[attr-defined]
kwargs = svc._execute_sync_unified.await_args.kwargs # type: ignore[attr-defined, union-attr] kwargs = svc._sync_ops.execute_sync_unified.await_args.kwargs # type: ignore[attr-defined, union-attr]
assert kwargs["request_headers"] == request_headers assert kwargs["request_headers"] == request_headers
assert kwargs["request_body"] == request_body assert kwargs["request_body"] == request_body
assert kwargs["request_body_ref"] == request_body_ref assert kwargs["request_body_ref"] == request_body_ref
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize( async def test_task_service_execute_delegates_to_execute_facade_ops() -> None:
("is_client_error", "retry_index", "max_retries_for_candidate", "expected_action"), svc = TaskService(MagicMock())
[ sentinel = object()
(True, 0, 1, "break"), svc._execute_facade_ops.execute = AsyncMock(return_value=sentinel) # type: ignore[attr-defined, method-assign]
(False, 0, 2, "continue"),
],
)
async def test_task_service_embedded_error_branch_applies_pool_health_policy(
monkeypatch: pytest.MonkeyPatch,
is_client_error: bool,
retry_index: int,
max_retries_for_candidate: int,
expected_action: str,
) -> None:
db = MagicMock()
svc = TaskService(db)
monkeypatch.setattr( result = await svc.execute(
task_service_module.RequestCandidateService, "mark_candidate_failed", MagicMock() task_type="chat",
) task_mode=TaskMode.SYNC,
monkeypatch.setattr(
"src.services.proxy_node.resolver.resolve_effective_proxy",
lambda provider_proxy, key_proxy: provider_proxy or key_proxy,
)
monkeypatch.setattr("src.services.proxy_node.resolver.resolve_proxy_info", lambda _proxy: None)
pool_on_error = AsyncMock()
monkeypatch.setattr(svc, "_pool_on_error", pool_on_error)
candidate = SimpleNamespace(
provider=SimpleNamespace(id="p1", name="prov", proxy=None),
endpoint=SimpleNamespace(id="e1"),
key=SimpleNamespace(id="k1", proxy=None),
)
cause = EmbeddedErrorException(
provider_name="prov",
error_code=429,
error_message="usage_limit_reached",
error_status="RESOURCE_EXHAUSTED",
)
context = ExecutionContext(
candidate_id="cid-1",
candidate_index=0,
provider_id="p1",
endpoint_id="e1",
key_id="k1",
user_id=None,
api_key_id=None,
is_cached_user=False,
elapsed_ms=12,
concurrent_requests=3,
)
exec_err = ExecutionError(cause, context)
classifier = SimpleNamespace(is_client_error=lambda _text: is_client_error)
action = await svc._handle_candidate_error(
exec_err=exec_err,
candidate=candidate,
candidate_record_id="cand-1",
retry_index=retry_index,
max_retries_for_candidate=max_retries_for_candidate,
affinity_key="provider-test:p1",
api_format="openai:chat", api_format="openai:chat",
global_model_id="gpt-4o-mini", model_name="m",
request_id="req-1", user_api_key=MagicMock(id="u", user_id="user"),
attempt=1, request_func=AsyncMock(),
max_attempts=3, request_id="rid",
error_classifier=classifier,
) )
assert action == expected_action assert result is sentinel
pool_on_error.assert_awaited_once_with(candidate.provider, candidate.key, 429, cause) svc._execute_facade_ops.execute.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_task_service_pool_on_error_uses_embedded_error_message_fallback( async def test_task_service_submit_with_failover_delegates_to_submit_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._submit_facade_ops.submit_with_failover = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=sentinel
)
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(id="u", user_id="user"),
request_id="rid",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
)
assert result is sentinel
svc._submit_facade_ops.submit_with_failover.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_task_service_poll_delegates_to_video_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._video_facade_ops.poll = AsyncMock(return_value=sentinel) # type: ignore[attr-defined, method-assign]
result = await svc.poll("task-1", user_id="user-1")
assert result is sentinel
svc._video_facade_ops.poll.assert_awaited_once_with("task-1", user_id="user-1") # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_task_service_cancel_delegates_to_video_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._video_facade_ops.cancel = AsyncMock(return_value=sentinel) # type: ignore[attr-defined, method-assign]
result = await svc.cancel(
"task-1",
user_id="user-1",
original_headers={"x-test": "1"},
)
assert result is sentinel
svc._video_facade_ops.cancel.assert_awaited_once_with( # type: ignore[attr-defined]
"task-1",
user_id="user-1",
original_headers={"x-test": "1"},
)
@pytest.mark.asyncio
async def test_pool_on_error_uses_embedded_error_message_fallback(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
db = MagicMock()
svc = TaskService(db)
parsed_pool_cfg = object() parsed_pool_cfg = object()
apply_health_policy = AsyncMock() apply_health_policy = AsyncMock()
monkeypatch.setattr( monkeypatch.setattr(
@@ -212,7 +208,7 @@ async def test_task_service_pool_on_error_uses_embedded_error_message_fallback(
error_message="usage_limit_reached", error_message="usage_limit_reached",
) )
await svc._pool_on_error(provider, key, 429, cause) await pool_on_error(provider, key, 429, cause)
apply_health_policy.assert_awaited_once() apply_health_policy.assert_awaited_once()
kwargs = apply_health_policy.await_args.kwargs kwargs = apply_health_policy.await_args.kwargs

View File

@@ -0,0 +1,44 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from src.services.task.core.exceptions import TaskNotFoundError
from src.services.task.video.operations import VideoTaskOperationsService
@pytest.mark.asyncio
async def test_video_ops_cancel_delegates_to_cancel_ops() -> None:
svc = VideoTaskOperationsService(MagicMock())
task = SimpleNamespace(id="task-1")
svc._get_video_task_for_user = MagicMock(return_value=task) # type: ignore[attr-defined, method-assign]
svc._cancel_ops.cancel_task = AsyncMock(return_value={"ok": True}) # type: ignore[attr-defined, method-assign]
result = await svc.cancel(
"task-1",
user_id="user-1",
original_headers={"x-test": "1"},
)
assert result == {"ok": True}
svc._cancel_ops.cancel_task.assert_awaited_once_with( # type: ignore[attr-defined]
task=task,
task_id="task-1",
original_headers={"x-test": "1"},
)
@pytest.mark.asyncio
async def test_video_ops_cancel_maps_not_found_to_http_404() -> None:
svc = VideoTaskOperationsService(MagicMock())
svc._get_video_task_for_user = MagicMock(side_effect=TaskNotFoundError("missing")) # type: ignore[attr-defined, method-assign]
with pytest.raises(HTTPException) as excinfo:
await svc.cancel("missing", user_id="user-1")
assert excinfo.value.status_code == 404
assert excinfo.value.detail == "Video task not found"

View File

@@ -1,10 +1,12 @@
from types import SimpleNamespace from types import SimpleNamespace
from typing import cast
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.services.task.impl.video_poller import VideoTaskPollerAdapter from src.models.database import VideoTask
from src.services.task.video.poller_adapter import VideoTaskPollerAdapter
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -13,11 +15,14 @@ async def test_poll_task_status_routes_gemini_video_to_gemini(
) -> None: ) -> None:
adapter = VideoTaskPollerAdapter() adapter = VideoTaskPollerAdapter()
task = SimpleNamespace( task = cast(
endpoint_id="e1", VideoTask,
key_id="k1", SimpleNamespace(
provider_api_format="gemini:video", endpoint_id="e1",
external_task_id="operations/123", key_id="k1",
provider_api_format="gemini:video",
external_task_id="operations/123",
),
) )
endpoint = SimpleNamespace(id="e1", base_url="https://example.com", api_format="gemini:video") endpoint = SimpleNamespace(id="e1", base_url="https://example.com", api_format="gemini:video")
key = SimpleNamespace(id="k1", api_key="enc") key = SimpleNamespace(id="k1", api_key="enc")
@@ -25,12 +30,12 @@ async def test_poll_task_status_routes_gemini_video_to_gemini(
monkeypatch.setattr(adapter, "_get_endpoint", lambda _db, _id: endpoint) monkeypatch.setattr(adapter, "_get_endpoint", lambda _db, _id: endpoint)
monkeypatch.setattr(adapter, "_get_key", lambda _db, _id: key) monkeypatch.setattr(adapter, "_get_key", lambda _db, _id: key)
monkeypatch.setattr( monkeypatch.setattr(
"src.services.task.impl.video_poller.crypto_service.decrypt", lambda _v: "decrypted" "src.services.task.video.poller_adapter.crypto_service.decrypt", lambda _v: "decrypted"
) )
auth_info = SimpleNamespace(auth_header="authorization", auth_value="Bearer x") auth_info = SimpleNamespace(auth_header="authorization", auth_value="Bearer x")
monkeypatch.setattr( monkeypatch.setattr(
"src.services.task.impl.video_poller.get_provider_auth", "src.services.task.video.poller_adapter.get_provider_auth",
AsyncMock(return_value=auth_info), AsyncMock(return_value=auth_info),
) )