refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构

- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1735 additions and 2433 deletions

View File

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

View File

@@ -0,0 +1,36 @@
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
class TaskMode(str, Enum):
SYNC = "sync"
ASYNC = "async"
@dataclass(slots=True)
class TaskContext:
"""
TaskContext (pure DTO)
- Only primitive types / IDs
- Serializable & safe to pass across processes
"""
request_id: str
task_type: str # chat/cli/video/image/audio
task_mode: TaskMode
user_id: str
api_key_id: str
client_ip: str = ""
user_agent: str = ""
start_time: float = 0.0
api_format: str | None = None
model: str | None = None
mapped_model: str | None = None
capability_requirements: dict[str, bool] = field(default_factory=dict)

View File

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

View File

@@ -0,0 +1,27 @@
from __future__ import annotations
from enum import Enum
class TaskStatus(str, Enum):
"""Generic task status (progress)."""
PENDING = "pending"
STREAMING = "streaming"
SUBMITTED = "submitted"
QUEUED = "queued"
PROCESSING = "processing"
COMPLETED = "completed"
FAILED = "failed"
CANCELLED = "cancelled"
EXPIRED = "expired"
class BillingStatus(str, Enum):
"""Billing settlement status (Usage.billing_status)."""
PENDING = "pending"
SETTLED = "settled"
VOID = "void"

View File

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

View File

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

View File

@@ -0,0 +1,520 @@
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.request.model_test_debug import (
get_candidate_model_test_debug,
merge_model_test_debug,
)
from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.request_state import RequestBodyState
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,
request_id: str | None,
candidate_record_id: str,
elapsed_ms: int,
captured_key_concurrent: int | None,
serializable_extra_data: dict[str, Any],
request_body_state: RequestBodyState | 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_state 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 = request_body_state.rectify_stage()
if stage <= 0 and request_body_state.is_rectified():
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_state.current_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_state.mark_rectified(rectified_body, 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,
).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_state: RequestBodyState | 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_async,
)
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 = await resolve_proxy_info_async(_eff_proxy)
_proxy_extra: dict[str, Any] | None = {"proxy": _proxy_info} if _proxy_info else None
_proxy_extra = merge_model_test_debug(
_proxy_extra, get_candidate_model_test_debug(candidate)
)
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
serializable_extra_data = (
merge_model_test_debug(
serializable_extra_data,
get_candidate_model_test_debug(candidate),
)
or serializable_extra_data
)
if isinstance(converted_error, ThinkingSignatureException):
action = self.handle_thinking_signature_error(
converted_error=converted_error,
provider_type=str(getattr(provider, "provider_type", "") or "").lower(),
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_state=request_body_state,
)
if action == "continue":
return "continue"
if isinstance(converted_error, UpstreamClientException):
logger.warning(
" [{}] 客户端请求错误,继续转移: {}",
request_id,
str(converted_error.message),
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="UpstreamClientException",
error_message=converted_error.message,
status_code=status_code,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=serializable_extra_data,
)
return "break"
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="HTTPStatusError",
error_message=extract_error_message(cause, status_code),
status_code=status_code,
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=serializable_extra_data,
)
return "continue" if has_retry_left else "break"
if isinstance(cause, error_classifier.RETRIABLE_ERRORS):
await error_classifier.handle_retriable_error(
error=cause,
provider=provider,
endpoint=endpoint,
key=key,
affinity_key=affinity_key,
api_format=api_format,
global_model_id=global_model_id,
captured_key_concurrent=captured_key_concurrent,
elapsed_ms=elapsed_ms,
request_id=request_id,
attempt=attempt,
max_attempts=max_attempts,
)
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type=type(cause).__name__,
error_message=extract_error_message(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "continue" if has_retry_left else "break"
if isinstance(cause, FormatConversionError):
logger.warning(" [{}] 格式转换失败,切换候选: {}", request_id, str(cause))
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type="FormatConversionError",
error_message=str(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "break"
RequestCandidateService.mark_candidate_failed(
db=self.db,
candidate_id=candidate_record_id,
error_type=type(cause).__name__,
error_message=extract_error_message(cause),
latency_ms=elapsed_ms,
concurrent_requests=captured_key_concurrent,
extra_data=_proxy_extra,
)
return "break"

View File

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

View File

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

View File

@@ -0,0 +1,265 @@
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]
# 构造延迟可用性检查回调(从 CandidateBuilder 打包的参数)
checker = None
deferred_params = candidate._deferred_check_params
if deferred_params is not None:
checker = self._build_availability_checker(deferred_params)
ordered_keys, trace = await manager.select_pool_keys(
session_uuid,
candidate_keys,
availability_checker=checker,
)
# 移除 deferred key未检查的避免为其创建 DB 记录
candidate.pool_keys = [
k for k in ordered_keys if getattr(k, "_pool_skip_reason", None) != "deferred"
]
selected_key_index = 0
selected_key = None
for idx, pool_key in enumerate(candidate.pool_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
@staticmethod
def _build_availability_checker(
params: dict[str, Any],
) -> Any:
"""Construct a key availability checker from deferred check params."""
from src.services.scheduling.candidate_builder import CandidateBuilder
endpoint_format = params.get("endpoint_format")
model_name = params.get("model_name", "")
capability_requirements = params.get("capability_requirements")
model_mappings = params.get("model_mappings")
candidate_models = params.get("candidate_models")
provider_type = params.get("provider_type")
# _check_key_availability 不依赖 _sorter传 None 安全
builder = CandidateBuilder(candidate_sorter=None) # type: ignore[arg-type]
def _checker(key: Any) -> tuple[bool, str | None, str | None]:
return builder._check_key_availability(
key,
endpoint_format,
model_name,
capability_requirements,
model_mappings=model_mappings,
candidate_models=candidate_models,
provider_type=provider_type,
)
return _checker

View File

@@ -0,0 +1,99 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from src.services.candidate.policy import FailoverAction
from src.services.task.execute.exception_classification import CandidateErrorAction
from src.services.task.request_state import RequestBodyState
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_state: RequestBodyState | None
last_error: Exception | None = None
last_candidate: Any | None = None
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_state:
return None
if not self.request_body_state.consume_rectified_this_turn():
return None
return max(max_retries_for_candidate, retry_index + 2)
def raise_classified_error(
self,
*,
fallback_error: Any,
failure_ops: TaskFailureOperationsService,
model_name: str,
api_format: str,
) -> None:
if self.last_error is not None:
failure_ops.attach_metadata_to_error(
self.last_error,
self.last_candidate,
model_name,
api_format,
)
raise self.last_error
if isinstance(fallback_error, Exception):
raise fallback_error
raise RuntimeError("execution_error_handler requested raise without exception context")
def resolve_execution_error_transition(
*,
action: CandidateErrorAction,
state: SyncExecutionState,
max_retries_for_candidate: int,
retry_index: int,
) -> ExecutionErrorTransition:
"""根据异常动作分类,返回 FailoverEngine 可消费的状态流转结果。"""
if action == CandidateErrorAction.RETRY_CURRENT:
return ExecutionErrorTransition(
failover_action=FailoverAction.RETRY,
max_retries=state.consume_rectify_retry_extension(
max_retries_for_candidate=max_retries_for_candidate,
retry_index=retry_index,
),
)
return ExecutionErrorTransition(
failover_action=FailoverAction.CONTINUE,
max_retries=None,
)

View File

@@ -0,0 +1,421 @@
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.task.request_state import RequestBodyState
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_state: RequestBodyState | None,
request_headers: dict[str, Any] | None,
request_body: dict[str, Any] | None,
create_pending_usage: bool = True,
) -> 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())
# IMPORTANT:
# This SYNC path awaits upstream HTTP work while the failover engine may commit
# candidate audit rows between attempts. SQLAlchemy's default
# expire_on_commit=True would expire provider/endpoint/key ORM objects and can
# trigger an unexpected lazy DB reload later in error handling (for example when
# reading candidate.provider.config for failover_rules after a timeout).
#
# Keep already-loaded candidate objects resident in memory for the duration of
# the request, mirroring the async submit path.
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
self.db.expire_on_commit = False
try:
# 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)
user: User | None = None
username_snapshot = None
api_key_name_snapshot = getattr(user_api_key, "name", None)
try:
user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
username_snapshot = getattr(user, "username", None) if user else None
except Exception as exc:
# username 仅用于审计快照,不应阻塞主请求链路。
logger.warning("查询用户快照失败: {}", str(exc))
# 默认由 TaskService 创建 pending 使用记录;已预创建的调用方可关闭。
if create_pending_usage:
try:
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,
)
# 号池排序涉及大量 Redis 操作,提前释放 DB 连接避免连接池压力
from src.services.scheduling.utils import release_db_connection_before_await
release_db_connection_before_await(self.db)
# 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 = await candidate_resolver.create_candidate_records_async(
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_state=request_body_state,
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),
username=username_snapshot,
api_key_name=api_key_name_snapshot,
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,
_first_byte_time_ms,
) = 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,
user_id=user_id,
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,
_first_byte_time_ms,
)
# 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_state=request_body_state,
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),
username=username_snapshot,
api_key_name=api_key_name_snapshot,
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,
)
finally:
self.db.expire_on_commit = original_expire_on_commit

View File

@@ -0,0 +1,262 @@
"""
Task poller (Phase2)
Provides a generic polling skeleton for async tasks.
Currently wired with a video poller adapter.
优化HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
"""
from __future__ import annotations
import asyncio
from datetime import datetime, timezone
from typing import Any, Protocol, runtime_checkable
from uuid import uuid4
from sqlalchemy.orm import Session
from src.core.api_format.conversion.internal_video import InternalVideoPollResult
from src.core.logger import logger
from src.database import create_session
from src.services.system.scheduler import get_scheduler
from src.services.task.video.poller_adapter import VideoPollContext, VideoTaskPollerAdapter
@runtime_checkable
class TaskPollerAdapter(Protocol):
task_type: str
# scheduler
job_id: str
job_name: str
interval_seconds: int
# distributed lock (optional, best-effort)
lock_key: str
lock_ttl: int
# execution
batch_size: int
concurrency: int
consecutive_failure_alert_threshold: int
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]: ...
def get_task(self, db: Session, task_id: str) -> Any | None: ...
# 分阶段处理方法(推荐使用)
async def prepare_poll_context(
self, db: Session, task: Any
) -> Any: ... # Returns context or error result
async def poll_task_http(self, ctx: Any) -> Any: ... # Returns poll result
async def update_task_after_poll(
self,
task_id: str,
result: Any,
ctx: Any,
redis_client: Any | None,
error_exception: Exception | None = None,
) -> None: ...
# 旧版方法(保留兼容性)
async def poll_single_task(
self, db: Session, task: Any, *, redis_client: Any | None
) -> None: ...
def sanitize_error_message(self, message: str) -> str: ...
class TaskPollerService:
"""Generic background poller for async tasks."""
def __init__(self, adapter: TaskPollerAdapter) -> None:
self.adapter = adapter
self._lock = asyncio.Lock()
self.redis: Any | None = None
self._semaphore: asyncio.Semaphore | None = None
self._consecutive_failures = 0
async def start(self) -> None:
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
# lazy import to avoid redis hard dependency in local runs
from src.clients.redis_client import get_redis_client
if self.redis is None:
self.redis = await get_redis_client(require_redis=False)
scheduler = get_scheduler()
scheduler.add_interval_job(
self.poll_pending_tasks,
seconds=self.adapter.interval_seconds,
job_id=self.adapter.job_id,
name=self.adapter.job_name,
)
async def stop(self) -> None:
scheduler = get_scheduler()
scheduler.remove_job(self.adapter.job_id)
async def poll_pending_tasks(self) -> None:
try:
await self._do_poll()
except asyncio.CancelledError:
logger.debug("[{}] poll_pending_tasks cancelled (shutdown?)", self.adapter.task_type)
return
async def _do_poll(self) -> None:
async with self._lock:
token = await self._acquire_redis_lock()
if token is None:
return
try:
with create_session() as db:
now = datetime.now(timezone.utc)
task_ids = self.adapter.list_due_task_ids(
db, now=now, limit=self.adapter.batch_size
)
if not task_ids:
self._consecutive_failures = 0
return
poll_results: list[bool] = []
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
semaphore = self._semaphore
async def poll_with_semaphore(task_id: str) -> None:
async with semaphore:
try:
# ========== 阶段 1准备数据短暂持有连接==========
with create_session() as task_db:
task_obj = self.adapter.get_task(task_db, task_id)
if not task_obj:
logger.warning(
"[{}] Task {} disappeared during poll",
self.adapter.task_type,
task_id,
)
poll_results.append(True)
return
ctx_or_result = await self.adapter.prepare_poll_context(
task_db, task_obj
)
# 检查是否是错误结果(而非上下文)
if isinstance(ctx_or_result, InternalVideoPollResult):
# 准备阶段就失败了,直接更新任务状态
await self.adapter.update_task_after_poll(
task_id=task_id,
result=ctx_or_result,
ctx=None, # type: ignore[arg-type]
redis_client=self.redis,
)
poll_results.append(True)
return
ctx: VideoPollContext = ctx_or_result
# ========== 阶段 2HTTP 请求(不持有数据库连接)==========
error_exception: Exception | None = None
try:
result = await self.adapter.poll_task_http(ctx)
except Exception as http_exc:
# HTTP 请求失败,记录异常以便后续处理
error_exception = http_exc
result = InternalVideoPollResult(
status=None, # type: ignore[arg-type]
error_message=str(http_exc),
)
# ========== 阶段 3更新数据库获取新连接==========
await self.adapter.update_task_after_poll(
task_id=task_id,
result=result,
ctx=ctx,
redis_client=self.redis,
error_exception=error_exception,
)
poll_results.append(True)
except Exception as exc:
logger.exception(
"[{}] Unexpected error polling task {}: {}",
self.adapter.task_type,
task_id,
self.adapter.sanitize_error_message(str(exc)),
)
poll_results.append(False)
async with asyncio.TaskGroup() as tg:
for tid in task_ids:
tg.create_task(poll_with_semaphore(tid))
batch_failures = sum(1 for r in poll_results if r is False)
if batch_failures == len(task_ids):
self._consecutive_failures += 1
if (
self._consecutive_failures
>= self.adapter.consecutive_failure_alert_threshold
):
logger.error(
"[ALERT] {} poller: {} consecutive batches failed.",
self.adapter.task_type,
self._consecutive_failures,
)
else:
self._consecutive_failures = 0
finally:
await self._release_redis_lock(token)
async def _acquire_redis_lock(self) -> str | None:
if not self.redis:
return "no_redis"
token = str(uuid4())
try:
acquired = await self.redis.set(
self.adapter.lock_key, token, nx=True, ex=self.adapter.lock_ttl
)
except Exception as exc:
logger.warning(
"[{}] Redis lock acquire failed (best-effort skip): {}",
self.adapter.task_type,
exc,
)
return "no_redis"
return token if acquired else None
async def _release_redis_lock(self, token: str) -> None:
if not self.redis or token == "no_redis":
return
script = """
if redis.call('GET', KEYS[1]) == ARGV[1] then
return redis.call('DEL', KEYS[1])
end
return 0
"""
try:
await self.redis.eval(script, 1, self.adapter.lock_key, token)
except Exception as exc:
logger.warning(
"[{}] Redis lock release failed (will expire via TTL): {}",
self.adapter.task_type,
exc,
)
_task_poller: TaskPollerService | None = None
def get_task_poller() -> TaskPollerService:
global _task_poller
if _task_poller is None:
_task_poller = TaskPollerService(VideoTaskPollerAdapter())
return _task_poller

View File

@@ -0,0 +1,64 @@
from __future__ import annotations
import copy
from dataclasses import dataclass, field
from typing import Any, Protocol
class RequestBodyState(Protocol):
@property
def current_body(self) -> dict[str, Any]: ...
def build_attempt_body(self) -> dict[str, Any]: ...
def is_rectified(self) -> bool: ...
def rectify_stage(self) -> int: ...
def mark_rectified(self, body: dict[str, Any], *, stage: int) -> None: ...
def consume_rectified_this_turn(self) -> bool: ...
@dataclass(slots=True)
class MutableRequestBodyState:
"""Owns the mutable working request body used across retries."""
original_body: dict[str, Any]
_current_body: dict[str, Any] = field(init=False, repr=False)
_rectified: bool = field(default=False, init=False, repr=False)
_rectified_this_turn: bool = field(default=False, init=False, repr=False)
_rectify_stage: int = field(default=0, init=False, repr=False)
def __post_init__(self) -> None:
# Attempts already deep-copy per dispatch, and rectification paths clone before
# rewriting. Keep the initial working body as a direct view to avoid an eager copy.
self._current_body = self.original_body
@property
def current_body(self) -> dict[str, Any]:
return self._current_body
def build_attempt_body(self) -> dict[str, Any]:
return copy.deepcopy(self._current_body)
def is_rectified(self) -> bool:
return self._rectified
def rectify_stage(self) -> int:
return self._rectify_stage
def mark_rectified(self, body: dict[str, Any], *, stage: int) -> None:
self._current_body = body
self._rectified = True
self._rectified_this_turn = True
self._rectify_stage = stage
def consume_rectified_this_turn(self) -> bool:
if not self._rectified_this_turn:
return False
self._rectified_this_turn = False
return True
__all__ = ["MutableRequestBodyState", "RequestBodyState"]

View File

@@ -0,0 +1,691 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from types import SimpleNamespace
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import ApiKey
from src.services.candidate.recorder import CandidateRecorder
from src.services.task.core.context import TaskMode
from src.services.task.core.protocol import AttemptKind, AttemptResult
from src.services.task.core.schema import ExecutionResult, TaskStatusResult
from src.services.task.execute.error_handler import TaskErrorOperationsService
from src.services.task.execute.failure import TaskFailureOperationsService
from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.execute.sync_execute import SyncTaskExecutionService
from src.services.task.request_state import RequestBodyState
from src.services.task.submit.submit_service import AsyncTaskSubmitService
from src.services.task.video.facade import TaskVideoFacadeService
from src.services.task.video.operations import VideoTaskOperationsService
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
elif isinstance(getattr(cause, "error_message", None), str):
error_text = str(getattr(cause, "error_message", "") or "")[:4000]
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
class TaskService:
"""
Unified task service facade (Phase 3).
Phase 3.1 scope:
- Provide a single entrypoint for SYNC tasks.
- Keep behavior consistent with the pre-Phase-3 implementation.
- Return a structured `ExecutionResult` for downstream compatibility.
"""
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
self._candidate_recorder = CandidateRecorder(db)
# 兼容历史注入点_execute_facade_ops/_submit_facade_ops
# 不再依赖独立门面类,默认直接绑定 TaskService 内部实现。
self._execute_facade_ops = SimpleNamespace(
execute=self._execute_internal,
_get_candidate_keys=self._candidate_recorder.get_candidate_keys,
)
pool_ops = TaskPoolOperationsService()
error_ops = TaskErrorOperationsService(db, pool_ops=pool_ops)
failure_ops = TaskFailureOperationsService()
self._sync_ops = SyncTaskExecutionService(
db,
redis_client,
recorder=self._candidate_recorder,
pool_ops=pool_ops,
error_ops=error_ops,
failure_ops=failure_ops,
)
self._submit_ops = AsyncTaskSubmitService(
db,
redis_client,
apply_pool_reorder=pool_ops.apply_pool_reorder,
expand_pool_candidates_for_async_submit=pool_ops.expand_pool_candidates_for_async_submit,
)
self._video_ops = VideoTaskOperationsService(db, redis_client)
self._submit_facade_ops = self._submit_ops
self._video_facade_ops = TaskVideoFacadeService(self._video_ops)
async def execute(
self,
*,
task_type: str, # chat/cli/video/image
task_mode: TaskMode,
api_format: str,
model_name: str,
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | None = None,
request_body_state: RequestBodyState | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
# ASYNC-only (video submit)
extract_external_task_id: Any | None = None,
supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False,
max_candidates: int | None = None,
create_pending_usage: bool = True,
) -> ExecutionResult:
"""兼容入口:默认绑定到 TaskService 内部执行路由。"""
return await self._execute_facade_ops.execute(
task_type=task_type,
task_mode=task_mode,
api_format=api_format,
model_name=model_name,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body_state=request_body_state,
request_headers=request_headers,
request_body=request_body,
extract_external_task_id=extract_external_task_id,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
max_candidates=max_candidates,
create_pending_usage=create_pending_usage,
)
async def _execute_internal(
self,
*,
task_type: str, # chat/cli/video/image
task_mode: TaskMode,
api_format: str,
model_name: str,
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | None = None,
request_body_state: RequestBodyState | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
extract_external_task_id: Any | None = None,
supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False,
max_candidates: int | None = None,
create_pending_usage: bool = True,
) -> ExecutionResult:
if task_mode == TaskMode.ASYNC:
if extract_external_task_id is None:
raise ValueError("extract_external_task_id is required for task_mode=ASYNC")
outcome = await self.submit_with_failover(
api_format=api_format,
model_name=model_name,
affinity_key=str(user_api_key.id),
user_api_key=user_api_key,
request_id=request_id,
task_type=task_type,
submit_func=request_func,
extract_external_task_id=extract_external_task_id,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
capability_requirements=capability_requirements,
max_candidates=max_candidates,
request_body=request_body,
)
candidate_keys = []
if request_id:
try:
candidate_keys = self._execute_facade_ops._get_candidate_keys(request_id)
except Exception:
candidate_keys = []
selected_idx = -1
if candidate_keys:
for ck in candidate_keys:
if str(getattr(ck, "status", "")) == "success":
idx_val = getattr(ck, "candidate_index", -1)
selected_idx = int(idx_val) if idx_val is not None else -1
break
attempt_count = 0
if candidate_keys:
attempt_count = sum(
1
for ck in candidate_keys
if str(getattr(ck, "status", ""))
in {"pending", "success", "failed", "cancelled"}
)
attempt_result = AttemptResult(
kind=AttemptKind.ASYNC_SUBMIT,
http_status=int(outcome.upstream_status_code or 200),
http_headers=dict(outcome.upstream_headers or {}),
provider_task_id=str(outcome.external_task_id),
response_body=outcome.upstream_payload,
)
return ExecutionResult(
success=True,
attempt_result=attempt_result,
candidate=outcome.candidate,
candidate_index=selected_idx,
retry_index=0,
provider_id=str(outcome.candidate.provider.id),
provider_name=str(outcome.candidate.provider.name),
endpoint_id=str(outcome.candidate.endpoint.id),
key_id=str(outcome.candidate.key.id),
candidate_keys=candidate_keys,
attempt_count=attempt_count,
request_candidate_id=None,
)
_ = task_type # reserved for future routing (chat/cli/video/image)
return await self._sync_ops.execute_sync_unified(
api_format=api_format,
model_name=model_name,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body_state=request_body_state,
request_headers=request_headers,
request_body=request_body,
create_pending_usage=create_pending_usage,
)
async def execute_sync_candidates(
self,
*,
api_format: str,
model_name: str,
candidates: list[Any],
request_func: Callable[..., Any],
request_id: str | None = None,
current_user: Any | None = None,
user_api_key: ApiKey | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
request_body_state: RequestBodyState | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
affinity_key: str | None = None,
create_pending_usage: bool = False,
enable_cache_affinity: bool = False,
is_cancelled: Callable[[], Awaitable[bool]] | None = None,
) -> ExecutionResult:
"""Execute a pre-built candidate set through the unified SYNC runtime."""
from uuid import uuid4
from src.models.database import User
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import 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.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.request.candidate import RequestCandidateService
from src.services.request.executor import RequestExecutor
from src.services.scheduling.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.system.config import SystemConfigService
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,
)
from src.services.usage.service import UsageService
if not request_id:
request_id = str(uuid4())
api_format_norm = normalize_endpoint_signature(api_format)
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,
)
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 if enable_cache_affinity else None,
)
pool_ops = self._sync_ops._pool_ops
error_ops = self._sync_ops._error_ops
failure_ops = self._sync_ops._failure_ops
resolved_user = current_user
if resolved_user is None and user_api_key is not None:
try:
resolved_user = user_api_key.user if hasattr(user_api_key, "user") else None
except Exception:
resolved_user = None
if resolved_user is None and getattr(user_api_key, "user_id", None):
resolved_user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
user_id: str | None = None
if resolved_user is not None and getattr(resolved_user, "id", None):
user_id = str(resolved_user.id)
elif user_api_key is not None and getattr(user_api_key, "user_id", None):
user_id = str(user_api_key.user_id)
username_snapshot = getattr(resolved_user, "username", None) if resolved_user else None
api_key_name_snapshot = getattr(user_api_key, "name", None) if user_api_key else None
resolved_affinity_key = affinity_key
if not resolved_affinity_key:
api_key_id = getattr(user_api_key, "id", None) if user_api_key is not None else None
resolved_affinity_key = str(api_key_id) if api_key_id else f"internal-test:{request_id}"
if create_pending_usage:
try:
UsageService.create_pending_usage(
db=self.db,
request_id=request_id,
user=resolved_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:
from src.core.logger import logger as _logger
_logger.warning("创建 pending 使用记录失败: {}", str(exc))
all_candidates = list(candidates)
# 号池排序涉及大量 Redis 操作,提前释放 DB 连接避免连接池压力
from src.services.scheduling.utils import release_db_connection_before_await
release_db_connection_before_await(self.db)
all_candidates, pool_traces = await pool_ops.apply_pool_reorder(
all_candidates, request_body=request_body
)
candidate_record_map = await candidate_resolver.create_candidate_records_async(
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)
execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map,
request_body_state=request_body_state,
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))
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
)
candidate_record = 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=(getattr(user_api_key, "id", None) if user_api_key else None),
username=username_snapshot,
api_key_name=api_key_name_snapshot,
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,
)
self.db.flush()
candidate_record_id = str(candidate_record.id)
execution_state.candidate_record_map[(candidate_index, retry_index)] = (
candidate_record_id
)
setattr(candidate, "_utf_candidate_record_id", candidate_record_id)
(
response,
_provider_name,
attempt_id,
_provider_id,
_endpoint_id,
_key_id,
_first_byte_time_ms,
) = 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,
user_id=user_id,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=resolved_affinity_key,
global_model_id=model_name,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
await 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[Any, int | None]:
execution_state.track_execution_error(exec_err=exec_err, candidate=candidate)
candidate_record_id = execution_state.resolve_candidate_record_id(
candidate_index=candidate_index,
record_id=record_id,
)
raw_action = await 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=resolved_affinity_key,
api_format=api_format_norm,
global_model_id=model_name,
request_id=request_id,
attempt=attempt_count,
max_attempts=int(max_attempts or 0),
request_body_state=request_body_state,
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=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._candidate_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) if user_api_key and getattr(user_api_key, "id", None) else None
),
username=username_snapshot,
api_key_name=api_key_name_snapshot,
candidate_record_map=candidate_record_map,
max_attempts=max_attempts,
execution_error_handler=_handle_exec_err,
is_cancelled=is_cancelled,
)
if result.success:
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
failure_ops.raise_all_failed_exception(
request_id,
max_attempts,
execution_state.last_candidate,
model_name,
api_format_norm,
execution_state.last_error,
)
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,
) -> Any:
"""
Unified ASYNC submit entrypoint (Phase 3.2).
兼容入口:默认直接绑定 AsyncTaskSubmitService。
"""
return await self._submit_facade_ops.submit_with_failover(
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
task_type=task_type,
submit_func=submit_func,
extract_external_task_id=extract_external_task_id,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
capability_requirements=capability_requirements,
max_candidates=max_candidates,
request_body=request_body,
)
# ====================
# Phase 3.1: Async task helpers (poll/finalize)
# ====================
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
return await self._video_facade_ops.poll(task_id, user_id=user_id)
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
return await self._video_facade_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_facade_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_facade_ops.finalize_video_task(task)
async def finalize(self, task_id: str) -> bool:
return await self._video_facade_ops.finalize(task_id)

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -0,0 +1,126 @@
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,
)
# 号池排序涉及大量 Redis 操作,提前释放 DB 连接避免连接池压力
from src.services.scheduling.utils import release_db_connection_before_await
release_db_connection_before_await(self.db)
# 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 = await resolver.create_candidate_records_async(
all_candidates=candidates,
request_id=request_id,
user_id=str(user_api_key.user_id),
user_api_key=user_api_key,
required_capabilities=capability_requirements,
expand_retries=False,
)
except Exception as exc:
logger.warning(
"[TaskService] Failed to create candidate records: {}",
self._sanitize(str(exc)),
)
record_map = {}
return PreparedSubmitCandidates(candidates=candidates, record_map=record_map)

View File

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

View File

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

View File

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

View File

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

View File

@@ -0,0 +1,324 @@
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,
finalized_at=getattr(task, "completed_at", 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
# Advisory check无锁仅用于快速跳过实际状态转换由 update_settled_billing 的
# with_for_update() 保证原子性。
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)
if getattr(existing, "billing_status", None) != "pending":
logger.debug(
"Skip video task billing finalize because Usage is already terminal: task_id={} request_id={} billing_status={}",
getattr(task, "id", None),
request_id,
getattr(existing, "billing_status", None),
)
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",
},
},
finalized_at=getattr(task, "completed_at", None),
)
if updated:
logger.debug(
"Updated video task billing: task_id={} request_id={} cost={:.6f}",
getattr(task, "id", None),
request_id,
cost,
)
else:
logger.warning(
"Failed to update video task billing (may already be updated): "
"task_id={} request_id={}",
getattr(task, "id", None),
request_id,
)
return bool(updated)

View File

@@ -0,0 +1,266 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.logger import logger
from src.models.database import 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)
"""
import json
from datetime import datetime, timezone
import httpx
from fastapi import HTTPException
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
current_status = str(getattr(task, "status", "") or "")
non_cancellable_statuses = {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}
if current_status in non_cancellable_statuses:
raise HTTPException(
status_code=409,
detail=f"Task cannot be cancelled in status: {current_status}",
)
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),
)
async def _try_rust_cancel_response(
*,
method: str,
url: str,
request_headers: dict[str, str],
body: Any,
content_type: str | None = None,
) -> httpx.Response:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
build_execution_plan_body,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return httpx.Response(
status_code=503,
request=httpx.Request(method, url, headers=request_headers),
json={"error": {"message": "Video 取消仅支持 Rust executor"}},
)
final_headers = dict(request_headers)
if (
body is not None
and content_type
and not any(str(key).lower() == "content-type" for key in final_headers)
):
final_headers["content-type"] = content_type
try:
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=str(getattr(task, "request_id", "") or task_id),
candidate_id=None,
provider_name=provider_format_norm.split(":", 1)[0],
provider_id=str(getattr(endpoint, "provider_id", "") or ""),
endpoint_id=str(getattr(endpoint, "id", "") or ""),
key_id=str(getattr(key, "id", "") or ""),
method=method,
url=url,
headers=final_headers,
body=build_execution_plan_body(body, content_type=content_type),
stream=False,
provider_api_format=provider_format,
client_api_format=provider_format,
model_name=str(getattr(task, "model", "") or ""),
content_type=content_type,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=300_000,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[VideoCancel] Rust executor unavailable task={} method={} url={}: {}",
getattr(task, "id", task_id),
method,
url,
str(exc),
)
return httpx.Response(
status_code=503,
request=httpx.Request(method, url, headers=final_headers),
json={"error": {"message": "执行器暂时不可用,请稍后重试"}},
)
response_headers = dict(result.headers)
if result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(result.response_json, ensure_ascii=False).encode("utf-8")
elif result.response_body_bytes is not None:
response_body = result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=result.status_code,
request=httpx.Request(method, url, headers=final_headers),
headers=response_headers,
content=response_body,
)
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 _try_rust_cancel_response(
method="DELETE",
url=upstream_url,
request_headers=headers,
body=None,
)
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 _try_rust_cancel_response(
method="POST",
url=upstream_url,
request_headers=headers,
body={},
content_type="application/json",
)
if response.status_code >= 400:
return response
else:
raise HTTPException(
status_code=400,
detail=f"Cancel not supported for provider format: {provider_format}",
)
now = datetime.now(timezone.utc)
task.status = VideoStatus.CANCELLED.value
task.completed_at = getattr(task, "completed_at", None) or now
task.updated_at = now
# Void Usage (no charge)
try:
voided = UsageService.finalize_void(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
finalized_at=task.completed_at,
)
if not voided:
logger.warning(
"Skip voiding video usage because billing is already terminal: task_id={} request_id={}",
getattr(task, "id", task_id),
getattr(task, "request_id", None),
)
except Exception as exc:
logger.warning(
"Failed to void usage for cancelled task={}: {}",
getattr(task, "id", task_id),
str(exc),
)
self.db.commit()
return None

View File

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

View File

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

View File

@@ -0,0 +1,656 @@
"""
Video task poller adapter.
Implements the video-specific poll/normalize/update logic used by TaskPollerService.
优化HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.api_format import (
build_upstream_headers_for_endpoint,
get_extra_headers_from_endpoint,
make_signature_key,
)
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.core.provider_auth_types import ProviderAuthInfo
from src.core.video_utils import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.provider.auth import get_provider_auth
from src.services.provider.provider_context import resolve_provider_proxy
@dataclass(slots=True)
class VideoPollContext:
"""视频轮询上下文,保存 HTTP 请求所需的数据(不依赖数据库会话)"""
task_id: str
external_task_id: str
provider_api_format: str
base_url: str
upstream_key: str
headers: dict[str, str]
# 用于更新任务的原始数据
poll_count: int
retry_count: int
poll_interval_seconds: int
max_poll_count: int
current_status: str
proxy_config: dict[str, Any] | None = None
delegate_config: dict[str, Any] | None = None
proxy_snapshot: Any = None
# 永久性错误指示词(用于降级判断,不应重试)
_PERMANENT_ERROR_INDICATORS = frozenset(
{
"not found",
"404",
"unauthorized",
"401",
"forbidden",
"403",
"invalid request",
"invalid api key",
"does not exist",
}
)
class PollHTTPError(RuntimeError):
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
def __init__(self, status_code: int, message: str):
# 确保错误信息包含状态码
full_message = f"HTTP {status_code}: {message}" if message else f"HTTP {status_code}"
super().__init__(full_message)
self.status_code = status_code
self.original_message = message
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:
task_type = "video"
# scheduler
job_id = "task_poller:video"
job_name = "视频任务轮询"
interval_seconds = config.video_poll_interval_seconds
# distributed lock
lock_key = "task_poller:video:lock"
lock_ttl = 60
# execution
batch_size = config.video_poll_batch_size
concurrency = config.video_poll_concurrency
consecutive_failure_alert_threshold = 5
max_backoff_seconds = 300
def __init__(self, finalize_video_task_fn: VideoTaskFinalizeFn | None = None) -> None:
self._openai_normalizer = OpenAINormalizer()
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:
return sanitize_error_message(message)
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]:
tasks = (
db.query(VideoTask)
.filter(
VideoTask.status.in_(
[
VideoStatus.SUBMITTED.value,
VideoStatus.QUEUED.value,
VideoStatus.PROCESSING.value,
]
),
VideoTask.next_poll_at <= now,
VideoTask.poll_count < VideoTask.max_poll_count,
)
.order_by(VideoTask.next_poll_at.asc())
.limit(limit)
.all()
)
return [t.id for t in tasks]
def get_task(self, db: Session, task_id: str) -> VideoTask | None:
# SQLAlchemy 1.4+ API
return db.get(VideoTask, task_id)
# ==================== 分阶段处理方法(优化数据库连接占用)====================
async def prepare_poll_context(
self, db: Session, task: VideoTask
) -> VideoPollContext | InternalVideoPollResult:
"""
阶段 1准备轮询上下文短暂持有数据库连接
Returns:
VideoPollContext: 成功时返回上下文
InternalVideoPollResult: 失败时返回错误结果
"""
if not task.endpoint_id or not task.key_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_provider_info",
error_message="Task missing endpoint_id or key_id",
)
endpoint = self._get_endpoint(db, task.endpoint_id)
key = self._get_key(db, task.key_id)
if not key.api_key:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="provider_config_error",
error_message="Provider key not properly configured",
)
try:
upstream_key = crypto_service.decrypt(key.api_key)
except Exception:
logger.warning("Failed to decrypt provider key for task {}", task.id)
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="decryption_error",
error_message="Failed to decrypt provider key",
)
provider_format = (task.provider_api_format or "").strip().lower()
if not provider_format:
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 构建请求头
if provider_format.startswith("gemini:"):
auth_info = await get_provider_auth(endpoint, key)
else:
auth_info = None
headers = self._build_headers(provider_format, upstream_key, endpoint, auth_info)
proxy_config, delegate_config, proxy_snapshot = await self._build_transport_context(
endpoint=endpoint,
key=key,
)
return VideoPollContext(
task_id=task.id,
external_task_id=task.external_task_id or "",
provider_api_format=provider_format,
base_url=endpoint.base_url or "",
upstream_key=upstream_key,
headers=headers,
poll_count=task.poll_count,
retry_count=task.retry_count,
poll_interval_seconds=task.poll_interval_seconds,
max_poll_count=task.max_poll_count,
current_status=task.status,
proxy_config=proxy_config,
delegate_config=delegate_config,
proxy_snapshot=proxy_snapshot,
)
async def poll_task_http(self, ctx: VideoPollContext) -> InternalVideoPollResult:
"""
阶段 2执行 HTTP 请求(不持有数据库连接)
Args:
ctx: 轮询上下文
Returns:
InternalVideoPollResult: 轮询结果
"""
if not ctx.external_task_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
if ctx.provider_api_format.startswith("gemini:"):
return await self._poll_gemini_with_context(ctx)
return await self._poll_openai_with_context(ctx)
async def update_task_after_poll(
self,
task_id: str,
result: InternalVideoPollResult,
ctx: VideoPollContext | None,
redis_client: Any | None,
error_exception: Exception | None = None,
) -> None:
"""
阶段 3更新数据库获取新的数据库连接
Args:
task_id: 任务 ID
result: 轮询结果
ctx: 轮询上下文(准备阶段就失败时为 None
redis_client: Redis 客户端
error_exception: 如果 HTTP 请求失败,传入异常对象
"""
with create_session() as db:
task = db.get(VideoTask, task_id)
if not task:
logger.warning("Task {} disappeared during poll update", task_id)
return
if task.status in {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}:
logger.debug(
"Skip poll update for terminal task {} with status {}",
task_id,
task.status,
)
return
if error_exception is not None and ctx is not None:
# HTTP 请求失败(需要 ctx 来计算 backoff
self._handle_poll_error(task, error_exception, ctx)
elif result.status == VideoStatus.COMPLETED:
task.status = VideoStatus.COMPLETED.value
task.video_url = result.video_url
task.video_expires_at = result.expires_at
task.completed_at = datetime.now(timezone.utc)
task.progress_percent = 100
if result.video_urls:
task.video_urls = result.video_urls
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
)
# 超时检查
task.updated_at = datetime.now(timezone.utc)
if task.poll_count >= task.max_poll_count and task.status not in [
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
]:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_timeout"
task.error_message = f"Task timed out after {task.poll_count} polls"
task.completed_at = datetime.now(timezone.utc)
# 终态结算
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await self._finalize_video_task(db, task, redis_client)
except Exception as exc:
logger.exception(
"Failed to record video usage for task={}: {}",
task.id,
sanitize_error_message(str(exc)),
)
db.commit()
def _handle_poll_error(self, task: VideoTask, exc: Exception, ctx: VideoPollContext) -> None:
"""处理轮询错误"""
task.poll_count += 1
error_msg = sanitize_error_message(str(exc))
logger.warning("Poll error for task {}: {}", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
is_permanent = self._is_permanent_error(exc, status_code=status_code)
if is_permanent:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_permanent_error"
task.error_message = error_msg
task.completed_at = datetime.now(timezone.utc)
else:
backoff = min(
ctx.poll_interval_seconds * (2 ** min(ctx.retry_count, 5)),
self.max_backoff_seconds,
)
task.retry_count += 1
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
async def _poll_openai_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
"""使用上下文进行 OpenAI 轮询(不需要数据库)"""
url = self._build_openai_url(ctx.base_url, ctx.external_task_id)
payload = await self._try_rust_poll_payload(ctx=ctx, url=url)
if payload is None:
raise PollHTTPError(503, "Video 轮询仅支持 Rust executor")
return self._openai_normalizer.video_poll_to_internal(payload)
async def _poll_gemini_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
"""使用上下文进行 Gemini 轮询(不需要数据库)"""
operation_name = normalize_gemini_operation_id(ctx.external_task_id)
url = self._build_gemini_url(ctx.base_url, operation_name)
logger.debug(
"[VideoPoller] Gemini poll: task={} external_id={} url={}",
ctx.task_id,
ctx.external_task_id,
url,
)
payload = await self._try_rust_poll_payload(ctx=ctx, url=url)
if payload is None:
raise PollHTTPError(503, "Video 轮询仅支持 Rust executor")
return self._gemini_normalizer.video_poll_to_internal(payload)
async def _try_rust_poll_payload(
self,
*,
ctx: VideoPollContext,
url: str,
) -> dict[str, Any] | None:
import httpx
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return None
try:
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=f"video-poll-{ctx.task_id}",
candidate_id=None,
provider_name=ctx.provider_api_format.split(":", 1)[0],
provider_id="",
endpoint_id="",
key_id="",
method="GET",
url=url,
headers=dict(ctx.headers),
body=ExecutionPlanBody(),
stream=False,
provider_api_format=ctx.provider_api_format,
client_api_format=ctx.provider_api_format,
model_name="video-poll",
proxy=ctx.proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=300_000,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError, ValueError) as exc:
logger.warning(
"[VideoPoller] Rust poll fallback task={} url={} error={}",
ctx.task_id,
url,
sanitize_error_message(str(exc)),
)
return None
if result.status_code >= 400:
response_text = ""
if result.response_json is not None:
response_text = json.dumps(result.response_json, ensure_ascii=False)
elif result.response_body_bytes is not None:
response_text = result.response_body_bytes.decode("utf-8", errors="replace")
raise PollHTTPError(
result.status_code,
self._extract_error_message(response_text, result.status_code),
)
if isinstance(result.response_json, dict):
return result.response_json
if result.response_body_bytes is not None:
try:
payload = json.loads(result.response_body_bytes.decode("utf-8"))
if isinstance(payload, dict):
return payload
except Exception:
logger.warning(
"[VideoPoller] Rust poll returned non-json body task={} url={}",
ctx.task_id,
url,
)
return None
# ==================== 旧版方法(保留兼容性)====================
async def poll_single_task(
self, db: Session, task: VideoTask, *, redis_client: Any | None
) -> None:
"""
兼容入口:复用三阶段轮询流程,避免维护重复逻辑。
"""
if task.status in {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}:
logger.debug(
"Skip legacy poll for terminal task {} with status {}", task.id, task.status
)
return
ctx_or_result = await self.prepare_poll_context(db, task)
if isinstance(ctx_or_result, InternalVideoPollResult):
await self.update_task_after_poll(
task_id=task.id,
result=ctx_or_result,
ctx=None,
redis_client=redis_client,
)
return
ctx = ctx_or_result
error_exception: Exception | None = None
try:
result = await self.poll_task_http(ctx)
except Exception as http_exc:
error_exception = http_exc
result = InternalVideoPollResult(
status=None, # type: ignore[arg-type]
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:
if not result.raw_response:
return
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
# (直接修改 JSON 字段内部不会自动标记为 dirty
metadata = dict(task.request_metadata) if task.request_metadata else {}
metadata["poll_raw_response"] = result.raw_response
task.request_metadata = metadata
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
if status_code is not None:
return 400 <= status_code < 500 and status_code != 429
error_msg = str(exc).lower()
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
ctx_or_result = await self.prepare_poll_context(db, task)
if isinstance(ctx_or_result, InternalVideoPollResult):
return ctx_or_result
return await self.poll_task_http(ctx_or_result)
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
base = (base_url or "https://api.openai.com").rstrip("/")
if base.endswith("/v1"):
return f"{base}/videos/{task_id}"
return f"{base}/v1/videos/{task_id}"
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/{operation_name}"
def _build_headers(
self,
endpoint_sig: str,
upstream_key: str,
endpoint: ProviderEndpoint,
auth_info: ProviderAuthInfo | None = None,
) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint)
headers = build_upstream_headers_for_endpoint(
{},
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
if auth_info:
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
return headers
async def _build_transport_context(
self,
*,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
) -> tuple[dict[str, Any] | None, dict[str, Any] | None, Any]:
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import ExecutionProxySnapshot
try:
effective_proxy = resolve_effective_proxy(
resolve_provider_proxy(endpoint=endpoint, key=key),
getattr(key, "proxy", None),
)
if not effective_proxy or not effective_proxy.get("enabled", True):
effective_proxy = await get_system_proxy_config_async()
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
proxy_url: str | None = None
if effective_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(effective_proxy)
proxy_info = await resolve_proxy_info_async(effective_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
return effective_proxy, delegate_cfg, proxy_snapshot
except Exception as exc:
logger.warning(
"[VideoPoller] Failed to build transport context endpoint={} key={}: {}",
getattr(endpoint, "id", None),
getattr(key, "id", None),
sanitize_error_message(str(exc)),
)
return None, None, None
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
if not endpoint:
raise RuntimeError("Provider endpoint not found")
return endpoint
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise RuntimeError("Provider key not found")
return key
def _extract_error_message(self, response_text: str | None, status_code: int) -> str:
"""从响应中提取有意义的错误信息"""
if not response_text:
return f"Request failed with status {status_code}"
# 尝试解析 JSON 格式的错误
try:
data = json.loads(response_text)
# OpenAI 格式: {"error": {"message": "..."}}
if isinstance(data.get("error"), dict):
error_obj = data["error"]
message = error_obj.get("message") or error_obj.get("detail") or str(error_obj)
return sanitize_error_message(message)
# Gemini 格式: {"error": {"message": "...", "code": 404}}
if "message" in data:
return sanitize_error_message(data["message"])
except (json.JSONDecodeError, TypeError, KeyError):
pass
# 回退到原始文本(截断)
return sanitize_error_message(response_text[:500])