mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -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)
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
203
src/main.py
203
src/main.py
@@ -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 自带 prefix,api_prefix 字段仅用于日志和文档
|
# 注意:模块的 router 自带 prefix,api_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("启动 TaskPoller(video)...")
|
logger.info("启动 TaskPoller(video)...")
|
||||||
await task_poller.start()
|
await state.task_poller.start()
|
||||||
else:
|
else:
|
||||||
logger.info("检测到其他 worker 已运行 TaskPoller(video),本实例跳过")
|
logger.info("检测到其他 worker 已运行 TaskPoller(video),本实例跳过")
|
||||||
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")
|
||||||
|
|
||||||
# 停止月卡额度重置调度器,并释放分布式锁
|
# 停止月卡额度重置调度器,并释放分布式锁
|
||||||
|
if state.quota_scheduler:
|
||||||
logger.info("停止月卡额度重置调度器...")
|
logger.info("停止月卡额度重置调度器...")
|
||||||
if quota_scheduler:
|
await state.quota_scheduler.stop()
|
||||||
await quota_scheduler.stop()
|
if state.task_coordinator:
|
||||||
if task_coordinator:
|
await state.task_coordinator.release("quota_scheduler")
|
||||||
await 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("停止 TaskPoller(video)...")
|
logger.info("停止 TaskPoller(video)...")
|
||||||
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 元数据定义
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
0
src/services/task/core/__init__.py
Normal file
0
src/services/task/core/__init__.py
Normal file
0
src/services/task/execute/__init__.py
Normal file
0
src/services/task/execute/__init__.py
Normal file
512
src/services/task/execute/error_handler.py
Normal file
512
src/services/task/execute/error_handler.py
Normal 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"
|
||||||
25
src/services/task/execute/exception_classification.py
Normal file
25
src/services/task/execute/exception_classification.py
Normal 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
|
||||||
107
src/services/task/execute/failure.py
Normal file
107
src/services/task/execute/failure.py
Normal 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,
|
||||||
|
)
|
||||||
222
src/services/task/execute/pool.py
Normal file
222
src/services/task/execute/pool.py
Normal 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
|
||||||
101
src/services/task/execute/state_transition.py
Normal file
101
src/services/task/execute/state_transition.py
Normal 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,
|
||||||
|
)
|
||||||
370
src/services/task/execute/sync_execute.py
Normal file
370
src/services/task/execute/sync_execute.py
Normal 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,
|
||||||
|
)
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
"""Per-task-type implementations (Phase2).
|
|
||||||
|
|
||||||
Currently includes:
|
|
||||||
- video: polling adapter
|
|
||||||
"""
|
|
||||||
|
|
||||||
__all__ = []
|
|
||||||
0
src/services/task/polling/__init__.py
Normal file
0
src/services/task/polling/__init__.py
Normal 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
12
src/services/task/submit/__init__.py
Normal file
12
src/services/task/submit/__init__.py
Normal 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"]
|
||||||
75
src/services/task/submit/attempt.py
Normal file
75
src/services/task/submit/attempt.py
Normal 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,
|
||||||
|
)
|
||||||
104
src/services/task/submit/execute.py
Normal file
104
src/services/task/submit/execute.py
Normal 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,
|
||||||
|
)
|
||||||
137
src/services/task/submit/filter.py
Normal file
137
src/services/task/submit/filter.py
Normal 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)
|
||||||
|
)
|
||||||
60
src/services/task/submit/outcome_builder.py
Normal file
60
src/services/task/submit/outcome_builder.py
Normal 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,
|
||||||
|
)
|
||||||
121
src/services/task/submit/prepare.py
Normal file
121
src/services/task/submit/prepare.py
Normal 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)
|
||||||
67
src/services/task/submit/record.py
Normal file
67
src/services/task/submit/record.py
Normal 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),
|
||||||
|
)
|
||||||
|
)
|
||||||
201
src/services/task/submit/response.py
Normal file
201
src/services/task/submit/response.py
Normal 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,
|
||||||
|
)
|
||||||
43
src/services/task/submit/rule_decider.py
Normal file
43
src/services/task/submit/rule_decider.py
Normal 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,
|
||||||
|
)
|
||||||
152
src/services/task/submit/submit_service.py
Normal file
152
src/services/task/submit/submit_service.py
Normal 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
|
||||||
0
src/services/task/video/__init__.py
Normal file
0
src/services/task/video/__init__.py
Normal file
320
src/services/task/video/billing.py
Normal file
320
src/services/task/video/billing.py
Normal 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)
|
||||||
149
src/services/task/video/cancel.py
Normal file
149
src/services/task/video/cancel.py
Normal 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
|
||||||
38
src/services/task/video/facade.py
Normal file
38
src/services/task/video/facade.py
Normal 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)
|
||||||
134
src/services/task/video/operations.py
Normal file
134
src/services/task/video/operations.py
Normal 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)
|
||||||
@@ -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,77 +380,36 @@ class VideoTaskPollerAdapter:
|
|||||||
self, db: Session, task: VideoTask, *, redis_client: Any | None
|
self, db: Session, task: VideoTask, *, redis_client: Any | None
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
旧版单任务轮询方法(保留向后兼容)
|
兼容入口:复用三阶段轮询流程,避免维护重复逻辑。
|
||||||
|
|
||||||
注意:此方法在 HTTP 请求期间持有数据库连接,建议使用分阶段方法。
|
|
||||||
"""
|
"""
|
||||||
try:
|
ctx_or_result = await self.prepare_poll_context(db, task)
|
||||||
result = await self._poll_task_status(db, task)
|
|
||||||
if result.status == VideoStatus.COMPLETED:
|
if isinstance(ctx_or_result, InternalVideoPollResult):
|
||||||
task.status = VideoStatus.COMPLETED.value
|
await self.update_task_after_poll(
|
||||||
task.video_url = result.video_url
|
task_id=task.id,
|
||||||
task.video_expires_at = result.expires_at
|
result=ctx_or_result,
|
||||||
task.completed_at = datetime.now(timezone.utc)
|
ctx=None,
|
||||||
task.progress_percent = 100
|
redis_client=redis_client,
|
||||||
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:
|
return
|
||||||
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
|
ctx = ctx_or_result
|
||||||
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
error_exception: Exception | None = None
|
||||||
if is_permanent:
|
|
||||||
task.status = VideoStatus.FAILED.value
|
|
||||||
task.error_code = "poll_permanent_error"
|
|
||||||
task.error_message = error_msg
|
|
||||||
task.completed_at = datetime.now(timezone.utc)
|
|
||||||
else:
|
|
||||||
backoff = min(
|
|
||||||
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
|
|
||||||
self.max_backoff_seconds,
|
|
||||||
)
|
|
||||||
task.retry_count += 1
|
|
||||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
|
||||||
|
|
||||||
# 超时:超过最大轮询次数且未进入终态
|
|
||||||
task.updated_at = datetime.now(timezone.utc)
|
|
||||||
if task.poll_count >= task.max_poll_count and task.status not in [
|
|
||||||
VideoStatus.COMPLETED.value,
|
|
||||||
VideoStatus.FAILED.value,
|
|
||||||
VideoStatus.CANCELLED.value,
|
|
||||||
]:
|
|
||||||
task.status = VideoStatus.FAILED.value
|
|
||||||
task.error_code = "poll_timeout"
|
|
||||||
task.error_message = f"Task timed out after {task.poll_count} polls"
|
|
||||||
task.completed_at = datetime.now(timezone.utc)
|
|
||||||
|
|
||||||
# 终态结算
|
|
||||||
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
|
||||||
try:
|
try:
|
||||||
await TaskService(db, redis_client=redis_client).finalize_video_task(task)
|
result = await self.poll_task_http(ctx)
|
||||||
except Exception as exc:
|
except Exception as http_exc:
|
||||||
logger.exception(
|
error_exception = http_exc
|
||||||
"Failed to record video usage for task={}: {}",
|
result = InternalVideoPollResult(
|
||||||
task.id,
|
status=None, # type: ignore[arg-type]
|
||||||
sanitize_error_message(str(exc)),
|
error_message=str(http_exc),
|
||||||
|
)
|
||||||
|
|
||||||
|
await self.update_task_after_poll(
|
||||||
|
task_id=task.id,
|
||||||
|
result=result,
|
||||||
|
ctx=ctx,
|
||||||
|
redis_client=redis_client,
|
||||||
|
error_exception=error_exception,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
|
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
|
||||||
@@ -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]
|
||||||
|
|||||||
@@ -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(
|
||||||
|
Provider,
|
||||||
|
SimpleNamespace(
|
||||||
id=provider_id,
|
id=provider_id,
|
||||||
name=provider_name,
|
name=provider_name,
|
||||||
max_retries=provider_max_retries,
|
max_retries=provider_max_retries,
|
||||||
config=provider_config,
|
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": {
|
||||||
|
|||||||
132
tests/services/test_sync_execute_state_transition.py
Normal file
132
tests/services/test_sync_execute_state_transition.py
Normal 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()
|
||||||
118
tests/services/test_task_pool_ops.py
Normal file
118
tests/services/test_task_pool_ops.py
Normal 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"
|
||||||
@@ -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
|
||||||
|
|||||||
44
tests/services/test_video_ops_cancel_delegate.py
Normal file
44
tests/services/test_video_ops_cancel_delegate.py
Normal 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"
|
||||||
@@ -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(
|
||||||
|
VideoTask,
|
||||||
|
SimpleNamespace(
|
||||||
endpoint_id="e1",
|
endpoint_id="e1",
|
||||||
key_id="k1",
|
key_id="k1",
|
||||||
provider_api_format="gemini:video",
|
provider_api_format="gemini:video",
|
||||||
external_task_id="operations/123",
|
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),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user