mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
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:
32
_deprecated_py_src/services/task/__init__.py
Normal file
32
_deprecated_py_src/services/task/__init__.py
Normal 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__)
|
||||
0
_deprecated_py_src/services/task/core/__init__.py
Normal file
0
_deprecated_py_src/services/task/core/__init__.py
Normal file
36
_deprecated_py_src/services/task/core/context.py
Normal file
36
_deprecated_py_src/services/task/core/context.py
Normal 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)
|
||||
24
_deprecated_py_src/services/task/core/exceptions.py
Normal file
24
_deprecated_py_src/services/task/core/exceptions.py
Normal file
@@ -0,0 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class StreamProbeError(RuntimeError):
|
||||
"""Streaming probe failed before first chunk (eligible for failover)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
http_status: int,
|
||||
original_exception: Exception | None = None,
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
self.http_status = http_status
|
||||
self.original_exception = original_exception
|
||||
|
||||
|
||||
class TaskNotFoundError(LookupError):
|
||||
"""Task not found (by internal id or external id)."""
|
||||
|
||||
def __init__(self, task_id: str) -> None:
|
||||
super().__init__(f"Task not found: {task_id}")
|
||||
self.task_id = task_id
|
||||
27
_deprecated_py_src/services/task/core/lifecycle.py
Normal file
27
_deprecated_py_src/services/task/core/lifecycle.py
Normal 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"
|
||||
51
_deprecated_py_src/services/task/core/protocol.py
Normal file
51
_deprecated_py_src/services/task/core/protocol.py
Normal file
@@ -0,0 +1,51 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any, AsyncIterator, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
|
||||
from src.services.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: ...
|
||||
75
_deprecated_py_src/services/task/core/schema.py
Normal file
75
_deprecated_py_src/services/task/core/schema.py
Normal 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
|
||||
520
_deprecated_py_src/services/task/execute/error_handler.py
Normal file
520
_deprecated_py_src/services/task/execute/error_handler.py
Normal 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"
|
||||
@@ -0,0 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
|
||||
class CandidateErrorAction(str, Enum):
|
||||
"""同步执行中,候选错误处理后的标准动作。"""
|
||||
|
||||
RETRY_CURRENT = "retry_current"
|
||||
NEXT_CANDIDATE = "next_candidate"
|
||||
RAISE_ERROR = "raise_error"
|
||||
|
||||
|
||||
def classify_candidate_error_action(action: Any) -> CandidateErrorAction:
|
||||
"""将 error_handler 的字符串动作分类为可控枚举。"""
|
||||
action_norm = str(action or "").strip().lower()
|
||||
if action_norm == "continue":
|
||||
return CandidateErrorAction.RETRY_CURRENT
|
||||
if action_norm == "raise":
|
||||
return CandidateErrorAction.RAISE_ERROR
|
||||
if action_norm == "break":
|
||||
return CandidateErrorAction.NEXT_CANDIDATE
|
||||
# Fail-safe:未知动作默认切到下一个候选,避免卡死在当前候选。
|
||||
return CandidateErrorAction.NEXT_CANDIDATE
|
||||
107
_deprecated_py_src/services/task/execute/failure.py
Normal file
107
_deprecated_py_src/services/task/execute/failure.py
Normal file
@@ -0,0 +1,107 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.services.request.result import RequestMetadata
|
||||
|
||||
|
||||
class TaskFailureOperationsService:
|
||||
"""任务失败收敛相关操作(异常元数据与统一抛错)。"""
|
||||
|
||||
@staticmethod
|
||||
def attach_metadata_to_error(
|
||||
error: Exception | None,
|
||||
candidate: Any | None,
|
||||
model_name: str,
|
||||
api_format: str,
|
||||
) -> None:
|
||||
"""Attach candidate metadata onto exception for usage recording."""
|
||||
if not error or not candidate:
|
||||
return
|
||||
|
||||
existing_metadata = getattr(error, "request_metadata", None)
|
||||
if existing_metadata and getattr(existing_metadata, "api_format", None):
|
||||
return
|
||||
|
||||
metadata = RequestMetadata(
|
||||
provider_request_headers=(
|
||||
getattr(existing_metadata, "provider_request_headers", {})
|
||||
if existing_metadata
|
||||
else {}
|
||||
),
|
||||
provider=getattr(existing_metadata, "provider", None) or str(candidate.provider.name),
|
||||
model=getattr(existing_metadata, "model", None) or model_name,
|
||||
provider_id=getattr(existing_metadata, "provider_id", None)
|
||||
or str(candidate.provider.id),
|
||||
provider_endpoint_id=(
|
||||
getattr(existing_metadata, "provider_endpoint_id", None)
|
||||
or str(candidate.endpoint.id)
|
||||
),
|
||||
provider_api_key_id=(
|
||||
getattr(existing_metadata, "provider_api_key_id", None) or str(candidate.key.id)
|
||||
),
|
||||
api_format=api_format,
|
||||
)
|
||||
setattr(error, "request_metadata", metadata)
|
||||
|
||||
@staticmethod
|
||||
def raise_all_failed_exception(
|
||||
request_id: str | None,
|
||||
max_attempts: int,
|
||||
last_candidate: Any | None,
|
||||
model_name: str,
|
||||
api_format: str,
|
||||
last_error: Exception | None = None,
|
||||
) -> None:
|
||||
"""Raise a unified 'all candidates failed' exception."""
|
||||
logger.error(" [{}] 所有 {} 个组合均失败", request_id, max_attempts)
|
||||
|
||||
request_metadata = None
|
||||
if last_candidate:
|
||||
request_metadata = {
|
||||
"provider": last_candidate.provider.name,
|
||||
"model": model_name,
|
||||
"provider_id": str(last_candidate.provider.id),
|
||||
"provider_endpoint_id": str(last_candidate.endpoint.id),
|
||||
"provider_api_key_id": str(last_candidate.key.id),
|
||||
"api_format": api_format,
|
||||
}
|
||||
|
||||
upstream_status: int | None = None
|
||||
upstream_response: str | None = None
|
||||
if last_error:
|
||||
if isinstance(last_error, httpx.HTTPStatusError):
|
||||
upstream_status = last_error.response.status_code
|
||||
upstream_response = getattr(last_error, "upstream_response", None)
|
||||
if not upstream_response:
|
||||
try:
|
||||
upstream_response = last_error.response.text
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
upstream_status = getattr(last_error, "upstream_status", None)
|
||||
upstream_response = getattr(last_error, "upstream_response", None)
|
||||
|
||||
if (
|
||||
not upstream_response
|
||||
or not upstream_response.strip()
|
||||
or upstream_response.startswith("Unable to read")
|
||||
):
|
||||
upstream_response = str(last_error)
|
||||
|
||||
friendly_message = "服务暂时不可用,请稍后重试"
|
||||
if last_error:
|
||||
last_error_message = getattr(last_error, "message", None)
|
||||
if last_error_message and isinstance(last_error_message, str):
|
||||
friendly_message = last_error_message
|
||||
|
||||
raise ProviderNotAvailableException(
|
||||
friendly_message,
|
||||
request_metadata=request_metadata,
|
||||
upstream_status=upstream_status,
|
||||
upstream_response=upstream_response,
|
||||
)
|
||||
265
_deprecated_py_src/services/task/execute/pool.py
Normal file
265
_deprecated_py_src/services/task/execute/pool.py
Normal 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
|
||||
99
_deprecated_py_src/services/task/execute/state_transition.py
Normal file
99
_deprecated_py_src/services/task/execute/state_transition.py
Normal 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,
|
||||
)
|
||||
421
_deprecated_py_src/services/task/execute/sync_execute.py
Normal file
421
_deprecated_py_src/services/task/execute/sync_execute.py
Normal 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
|
||||
262
_deprecated_py_src/services/task/polling/task_poller.py
Normal file
262
_deprecated_py_src/services/task/polling/task_poller.py
Normal 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
|
||||
|
||||
# ========== 阶段 2:HTTP 请求(不持有数据库连接)==========
|
||||
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
|
||||
64
_deprecated_py_src/services/task/request_state.py
Normal file
64
_deprecated_py_src/services/task/request_state.py
Normal 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"]
|
||||
691
_deprecated_py_src/services/task/service.py
Normal file
691
_deprecated_py_src/services/task/service.py
Normal 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)
|
||||
12
_deprecated_py_src/services/task/submit/__init__.py
Normal file
12
_deprecated_py_src/services/task/submit/__init__.py
Normal file
@@ -0,0 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
ApplyPoolReorderFn = Callable[
|
||||
[list[Any], dict[str, Any] | None],
|
||||
Awaitable[tuple[list[Any], list[Any]]],
|
||||
]
|
||||
ExpandPoolCandidatesFn = Callable[[list[Any]], list[Any]]
|
||||
|
||||
__all__ = ["ApplyPoolReorderFn", "ExpandPoolCandidatesFn"]
|
||||
75
_deprecated_py_src/services/task/submit/attempt.py
Normal file
75
_deprecated_py_src/services/task/submit/attempt.py
Normal file
@@ -0,0 +1,75 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
from src.services.candidate.submit import SubmitOutcome
|
||||
from src.services.task.submit.record import AsyncSubmitRecordService
|
||||
from src.services.task.submit.response import AsyncSubmitResponseService
|
||||
|
||||
|
||||
class AsyncSubmitAttemptService:
|
||||
"""异步提交单候选执行编排服务。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
*,
|
||||
sanitize: Callable[[str], str],
|
||||
extract_response_text: Callable[[httpx.Response], str],
|
||||
match_provider_failover_rule: Callable[..., str | None],
|
||||
) -> None:
|
||||
self.db = db
|
||||
self._record_ops = AsyncSubmitRecordService(db)
|
||||
self._response_ops = AsyncSubmitResponseService(
|
||||
db,
|
||||
record_ops=self._record_ops,
|
||||
sanitize=sanitize,
|
||||
extract_response_text=extract_response_text,
|
||||
match_provider_failover_rule=match_provider_failover_rule,
|
||||
)
|
||||
|
||||
async def submit_candidate(
|
||||
self,
|
||||
*,
|
||||
candidate: Any,
|
||||
record_id: str | None,
|
||||
candidate_info: dict[str, Any],
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
rule_lookup: BillingRuleLookupResult | None,
|
||||
submit_func: Any,
|
||||
extract_external_task_id: Any,
|
||||
) -> tuple[SubmitOutcome | None, int | None]:
|
||||
self._record_ops.mark_pending(record_id=record_id)
|
||||
|
||||
# Flush/commit BEFORE awaiting upstream submit to avoid holding DB connections.
|
||||
if self.db.in_transaction():
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
|
||||
# Attempt submit (upstream HTTP)
|
||||
try:
|
||||
response: httpx.Response = await submit_func(candidate)
|
||||
except Exception as exc:
|
||||
return self._response_ops.handle_submit_exception(
|
||||
record_id=record_id,
|
||||
candidate_info=candidate_info,
|
||||
exc=exc,
|
||||
)
|
||||
|
||||
return self._response_ops.handle_submit_response(
|
||||
candidate=candidate,
|
||||
record_id=record_id,
|
||||
candidate_info=candidate_info,
|
||||
candidate_keys=candidate_keys,
|
||||
rule_lookup=rule_lookup,
|
||||
response=response,
|
||||
extract_external_task_id=extract_external_task_id,
|
||||
)
|
||||
104
_deprecated_py_src/services/task/submit/execute.py
Normal file
104
_deprecated_py_src/services/task/submit/execute.py
Normal file
@@ -0,0 +1,104 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.services.candidate.submit import AllCandidatesFailedError, SubmitOutcome
|
||||
from src.services.task.submit.attempt import AsyncSubmitAttemptService
|
||||
from src.services.task.submit.filter import AsyncSubmitFilterService
|
||||
|
||||
|
||||
class AsyncSubmitExecutionService:
|
||||
"""异步提交候选执行编排服务。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
*,
|
||||
sanitize: Callable[[str], str],
|
||||
extract_response_text: Callable[[httpx.Response], str],
|
||||
match_provider_failover_rule: Callable[..., str | None],
|
||||
) -> None:
|
||||
self.db = db
|
||||
self._filter_ops = AsyncSubmitFilterService(db)
|
||||
self._attempt_ops = AsyncSubmitAttemptService(
|
||||
db,
|
||||
sanitize=sanitize,
|
||||
extract_response_text=extract_response_text,
|
||||
match_provider_failover_rule=match_provider_failover_rule,
|
||||
)
|
||||
|
||||
async def execute_submit_loop(
|
||||
self,
|
||||
*,
|
||||
candidates: list[Any],
|
||||
record_map: dict[tuple[int, int], str],
|
||||
task_type: str,
|
||||
model_name: str,
|
||||
submit_func: Any,
|
||||
extract_external_task_id: Any,
|
||||
supported_auth_types: set[str] | None,
|
||||
allow_format_conversion: bool,
|
||||
) -> SubmitOutcome:
|
||||
candidate_keys: list[dict[str, Any]] = []
|
||||
eligible_count = 0
|
||||
last_status_code: int | None = None
|
||||
|
||||
for idx, cand in enumerate(candidates):
|
||||
candidate_info = self._filter_ops.build_candidate_info(idx=idx, candidate=cand)
|
||||
candidate_keys.append(candidate_info)
|
||||
|
||||
attempt_plan = self._filter_ops.prepare_candidate_for_attempt(
|
||||
idx=idx,
|
||||
candidate=cand,
|
||||
record_map=record_map,
|
||||
candidate_info=candidate_info,
|
||||
task_type=task_type,
|
||||
model_name=model_name,
|
||||
supported_auth_types=supported_auth_types,
|
||||
allow_format_conversion=allow_format_conversion,
|
||||
)
|
||||
if attempt_plan is None:
|
||||
continue
|
||||
|
||||
eligible_count += 1
|
||||
outcome, status_code = await self._attempt_ops.submit_candidate(
|
||||
candidate=cand,
|
||||
record_id=attempt_plan.record_id,
|
||||
candidate_info=candidate_info,
|
||||
candidate_keys=candidate_keys,
|
||||
rule_lookup=attempt_plan.rule_lookup,
|
||||
submit_func=submit_func,
|
||||
extract_external_task_id=extract_external_task_id,
|
||||
)
|
||||
|
||||
if status_code is not None:
|
||||
last_status_code = status_code
|
||||
if outcome is not None:
|
||||
return outcome
|
||||
|
||||
# Persist candidate records before raising.
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
|
||||
if eligible_count == 0:
|
||||
reason = "no_eligible_candidates"
|
||||
if config.billing_require_rule:
|
||||
reason = "no_candidate_with_billing_rule"
|
||||
raise AllCandidatesFailedError(
|
||||
reason=reason,
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
|
||||
raise AllCandidatesFailedError(
|
||||
reason="all_candidates_failed",
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
137
_deprecated_py_src/services/task/submit/filter.py
Normal file
137
_deprecated_py_src/services/task/submit/filter.py
Normal file
@@ -0,0 +1,137 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import update
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.models.database import RequestCandidate
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class CandidateAttemptPlan:
|
||||
"""候选尝试计划(通过过滤后可进入提交阶段)。"""
|
||||
|
||||
record_id: str | None
|
||||
rule_lookup: BillingRuleLookupResult | None
|
||||
|
||||
|
||||
class AsyncSubmitFilterService:
|
||||
"""异步提交候选过滤服务。"""
|
||||
|
||||
def __init__(self, db: Session) -> None:
|
||||
self.db = db
|
||||
|
||||
@staticmethod
|
||||
def build_candidate_info(*, idx: int, candidate: Any) -> dict[str, Any]:
|
||||
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
|
||||
return {
|
||||
"index": idx,
|
||||
"provider_id": candidate.provider.id,
|
||||
"provider_name": candidate.provider.name,
|
||||
"endpoint_id": candidate.endpoint.id,
|
||||
"key_id": candidate.key.id,
|
||||
"key_name": getattr(candidate.key, "name", None),
|
||||
"auth_type": auth_type,
|
||||
"priority": getattr(candidate.key, "priority", 0) or 0,
|
||||
"is_cached": bool(getattr(candidate, "is_cached", False)),
|
||||
}
|
||||
|
||||
def prepare_candidate_for_attempt(
|
||||
self,
|
||||
*,
|
||||
idx: int,
|
||||
candidate: Any,
|
||||
record_map: dict[tuple[int, int], str],
|
||||
candidate_info: dict[str, Any],
|
||||
task_type: str,
|
||||
model_name: str,
|
||||
supported_auth_types: set[str] | None,
|
||||
allow_format_conversion: bool,
|
||||
) -> CandidateAttemptPlan | None:
|
||||
record_id = record_map.get((idx, 0))
|
||||
auth_type = candidate_info.get("auth_type", "api_key")
|
||||
|
||||
# Scheduler marked skip
|
||||
if getattr(candidate, "is_skipped", False):
|
||||
skip_reason = getattr(candidate, "skip_reason", None) or "skipped"
|
||||
self._mark_skip(
|
||||
record_id=record_id, candidate_info=candidate_info, skip_reason=skip_reason
|
||||
)
|
||||
return None
|
||||
|
||||
# Format conversion checks
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
if needs_conversion:
|
||||
# 1. handler-level switch
|
||||
if not allow_format_conversion:
|
||||
self._mark_skip(
|
||||
record_id=record_id,
|
||||
candidate_info=candidate_info,
|
||||
skip_reason="format_conversion_not_supported",
|
||||
)
|
||||
return None
|
||||
|
||||
# 2. global switch (from database config)
|
||||
if not SystemConfigService.is_format_conversion_enabled(self.db):
|
||||
self._mark_skip(
|
||||
record_id=record_id,
|
||||
candidate_info=candidate_info,
|
||||
skip_reason="format_conversion_disabled",
|
||||
extra_info={"format_conversion_enabled": False},
|
||||
)
|
||||
return None
|
||||
|
||||
# auth_type filter
|
||||
if supported_auth_types is not None and auth_type not in supported_auth_types:
|
||||
self._mark_skip(
|
||||
record_id=record_id,
|
||||
candidate_info=candidate_info,
|
||||
skip_reason=f"unsupported_auth_type:{auth_type}",
|
||||
)
|
||||
return None
|
||||
|
||||
# billing rule filter
|
||||
rule_lookup: BillingRuleLookupResult | None = None
|
||||
has_billing_rule = True
|
||||
if config.billing_require_rule:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=candidate.provider.id,
|
||||
model_name=model_name,
|
||||
task_type=task_type,
|
||||
)
|
||||
has_billing_rule = rule_lookup is not None
|
||||
if not has_billing_rule:
|
||||
self._mark_skip(
|
||||
record_id=record_id,
|
||||
candidate_info=candidate_info,
|
||||
skip_reason="billing_rule_missing",
|
||||
extra_info={"has_billing_rule": False},
|
||||
)
|
||||
return None
|
||||
candidate_info["has_billing_rule"] = has_billing_rule
|
||||
|
||||
return CandidateAttemptPlan(record_id=record_id, rule_lookup=rule_lookup)
|
||||
|
||||
def _mark_skip(
|
||||
self,
|
||||
*,
|
||||
record_id: str | None,
|
||||
candidate_info: dict[str, Any],
|
||||
skip_reason: str,
|
||||
extra_info: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||
if extra_info:
|
||||
candidate_info.update(extra_info)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="skipped", skip_reason=skip_reason)
|
||||
)
|
||||
60
_deprecated_py_src/services/task/submit/outcome_builder.py
Normal file
60
_deprecated_py_src/services/task/submit/outcome_builder.py
Normal file
@@ -0,0 +1,60 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
from src.services.candidate.submit import SubmitOutcome
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SubmitPayloadParseResult:
|
||||
"""提交响应解析结果。"""
|
||||
|
||||
payload: dict[str, Any] | None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
|
||||
class AsyncSubmitOutcomeBuilderService:
|
||||
"""异步提交结果构建服务。"""
|
||||
|
||||
def __init__(self, *, sanitize: Callable[[str], str]) -> None:
|
||||
self._sanitize = sanitize
|
||||
|
||||
def parse_payload(self, *, response: httpx.Response) -> SubmitPayloadParseResult:
|
||||
payload: dict[str, Any] | None = None
|
||||
try:
|
||||
data = response.json()
|
||||
if isinstance(data, dict):
|
||||
payload = data
|
||||
except Exception as exc:
|
||||
return SubmitPayloadParseResult(
|
||||
payload=None,
|
||||
error_type=type(exc).__name__,
|
||||
error_message=self._sanitize(str(exc)),
|
||||
)
|
||||
return SubmitPayloadParseResult(payload=payload)
|
||||
|
||||
@staticmethod
|
||||
def build_success_outcome(
|
||||
*,
|
||||
candidate: Any,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
external_task_id: str,
|
||||
rule_lookup: BillingRuleLookupResult | None,
|
||||
payload: dict[str, Any] | None,
|
||||
response: httpx.Response,
|
||||
) -> SubmitOutcome:
|
||||
return SubmitOutcome(
|
||||
candidate=candidate,
|
||||
candidate_keys=candidate_keys,
|
||||
external_task_id=external_task_id,
|
||||
rule_lookup=rule_lookup,
|
||||
upstream_payload=payload,
|
||||
upstream_headers=dict(response.headers),
|
||||
upstream_status_code=response.status_code,
|
||||
)
|
||||
126
_deprecated_py_src/services/task/submit/prepare.py
Normal file
126
_deprecated_py_src/services/task/submit/prepare.py
Normal 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)
|
||||
67
_deprecated_py_src/services/task/submit/record.py
Normal file
67
_deprecated_py_src/services/task/submit/record.py
Normal file
@@ -0,0 +1,67 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import update
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import RequestCandidate
|
||||
|
||||
|
||||
class AsyncSubmitRecordService:
|
||||
"""异步提交阶段的 RequestCandidate 落库服务。"""
|
||||
|
||||
def __init__(self, db: Session) -> None:
|
||||
self.db = db
|
||||
|
||||
def mark_pending(self, *, record_id: str | None) -> None:
|
||||
if not record_id:
|
||||
return
|
||||
started_at = datetime.now(timezone.utc)
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="pending", started_at=started_at)
|
||||
)
|
||||
|
||||
def mark_failed(
|
||||
self,
|
||||
*,
|
||||
record_id: str | None,
|
||||
error_type: str,
|
||||
error_message: str,
|
||||
status_code: int | None = None,
|
||||
) -> None:
|
||||
if not record_id:
|
||||
return
|
||||
|
||||
values: dict[str, object] = {
|
||||
"status": "failed",
|
||||
"error_type": error_type,
|
||||
"error_message": error_message,
|
||||
"finished_at": datetime.now(timezone.utc),
|
||||
}
|
||||
if status_code is not None:
|
||||
values["status_code"] = status_code
|
||||
|
||||
self.db.execute(
|
||||
update(RequestCandidate).where(RequestCandidate.id == record_id).values(**values)
|
||||
)
|
||||
|
||||
def mark_success(
|
||||
self,
|
||||
*,
|
||||
record_id: str | None,
|
||||
status_code: int,
|
||||
) -> None:
|
||||
if not record_id:
|
||||
return
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(
|
||||
status="success",
|
||||
status_code=status_code,
|
||||
finished_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
201
_deprecated_py_src/services/task/submit/response.py
Normal file
201
_deprecated_py_src/services/task/submit/response.py
Normal file
@@ -0,0 +1,201 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
from src.services.candidate.submit import SubmitOutcome, UpstreamClientRequestError
|
||||
from src.services.task.submit.outcome_builder import (
|
||||
AsyncSubmitOutcomeBuilderService,
|
||||
)
|
||||
from src.services.task.submit.record import AsyncSubmitRecordService
|
||||
from src.services.task.submit.rule_decider import AsyncSubmitRuleDeciderService
|
||||
|
||||
|
||||
class AsyncSubmitResponseService:
|
||||
"""异步提交响应判定服务。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
*,
|
||||
record_ops: AsyncSubmitRecordService,
|
||||
sanitize: Callable[[str], str],
|
||||
extract_response_text: Callable[[httpx.Response], str],
|
||||
match_provider_failover_rule: Callable[..., str | None],
|
||||
) -> None:
|
||||
self.db = db
|
||||
self._record_ops = record_ops
|
||||
self._sanitize = sanitize
|
||||
self._extract_response_text = extract_response_text
|
||||
self._rule_decider = AsyncSubmitRuleDeciderService(
|
||||
match_provider_failover_rule=match_provider_failover_rule
|
||||
)
|
||||
self._outcome_builder = AsyncSubmitOutcomeBuilderService(sanitize=sanitize)
|
||||
|
||||
def handle_submit_exception(
|
||||
self,
|
||||
*,
|
||||
record_id: str | None,
|
||||
candidate_info: dict[str, Any],
|
||||
exc: Exception,
|
||||
) -> tuple[None, None]:
|
||||
error_type = type(exc).__name__
|
||||
error_msg = self._sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "exception",
|
||||
"error_type": error_type,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._record_ops.mark_failed(
|
||||
record_id=record_id,
|
||||
error_type=error_type,
|
||||
error_message=error_msg,
|
||||
)
|
||||
return None, None
|
||||
|
||||
def handle_submit_response(
|
||||
self,
|
||||
*,
|
||||
candidate: Any,
|
||||
record_id: str | None,
|
||||
candidate_info: dict[str, Any],
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
rule_lookup: BillingRuleLookupResult | None,
|
||||
response: httpx.Response,
|
||||
extract_external_task_id: Any,
|
||||
) -> tuple[SubmitOutcome | None, int | None]:
|
||||
last_status_code = int(getattr(response, "status_code", 0) or 0)
|
||||
|
||||
if response.status_code >= 400:
|
||||
error_text = self._extract_response_text(response)
|
||||
error_msg = self._sanitize(error_text)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "http_error",
|
||||
"status_code": response.status_code,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._record_ops.mark_failed(
|
||||
record_id=record_id,
|
||||
status_code=response.status_code,
|
||||
error_type="http_error",
|
||||
error_message=error_msg,
|
||||
)
|
||||
|
||||
stop_pattern = self._rule_decider.detect_error_stop_pattern(
|
||||
candidate=candidate,
|
||||
response_text=error_text,
|
||||
status_code=response.status_code,
|
||||
)
|
||||
if stop_pattern:
|
||||
logger.info(
|
||||
"[TaskService] 错误终止规则命中: pattern={}, status_code={}, provider={}",
|
||||
stop_pattern,
|
||||
response.status_code,
|
||||
candidate.provider.name,
|
||||
)
|
||||
candidate_info["stop_rule_pattern"] = stop_pattern
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise UpstreamClientRequestError(
|
||||
response=response,
|
||||
candidate_keys=candidate_keys,
|
||||
)
|
||||
return None, last_status_code
|
||||
|
||||
success_text = self._extract_response_text(response)
|
||||
success_continue_pattern = self._rule_decider.detect_success_failover_pattern(
|
||||
candidate=candidate,
|
||||
response_text=success_text,
|
||||
status_code=response.status_code,
|
||||
)
|
||||
if success_continue_pattern:
|
||||
logger.info(
|
||||
"[TaskService] 成功转移规则命中: pattern={}, status_code={}, provider={}",
|
||||
success_continue_pattern,
|
||||
response.status_code,
|
||||
candidate.provider.name,
|
||||
)
|
||||
failover_reason = f"success_failover_rule_matched:{success_continue_pattern}"
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "success_failover",
|
||||
"status_code": response.status_code,
|
||||
"error_message": failover_reason,
|
||||
"success_rule_pattern": success_continue_pattern,
|
||||
}
|
||||
)
|
||||
self._record_ops.mark_failed(
|
||||
record_id=record_id,
|
||||
status_code=response.status_code,
|
||||
error_type="success_failover_pattern",
|
||||
error_message=failover_reason,
|
||||
)
|
||||
return None, last_status_code
|
||||
|
||||
parse_result = self._outcome_builder.parse_payload(response=response)
|
||||
if parse_result.error_type:
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "invalid_json",
|
||||
"error_type": parse_result.error_type,
|
||||
"error_message": parse_result.error_message,
|
||||
}
|
||||
)
|
||||
self._record_ops.mark_failed(
|
||||
record_id=record_id,
|
||||
status_code=response.status_code,
|
||||
error_type="invalid_json",
|
||||
error_message=parse_result.error_message or "invalid_json",
|
||||
)
|
||||
return None, last_status_code
|
||||
|
||||
payload = parse_result.payload
|
||||
external_task_id = extract_external_task_id(payload or {})
|
||||
if not external_task_id:
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "empty_task_id",
|
||||
"error_message": "Upstream returned empty task id",
|
||||
}
|
||||
)
|
||||
self._record_ops.mark_failed(
|
||||
record_id=record_id,
|
||||
status_code=response.status_code,
|
||||
error_type="empty_task_id",
|
||||
error_message="Upstream returned empty task id",
|
||||
)
|
||||
return None, last_status_code
|
||||
|
||||
# Success
|
||||
candidate_info.update({"attempt_status": "success", "selected": True})
|
||||
self._record_ops.mark_success(
|
||||
record_id=record_id,
|
||||
status_code=response.status_code,
|
||||
)
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
|
||||
return (
|
||||
self._outcome_builder.build_success_outcome(
|
||||
candidate=candidate,
|
||||
candidate_keys=candidate_keys,
|
||||
external_task_id=str(external_task_id),
|
||||
rule_lookup=rule_lookup,
|
||||
payload=payload,
|
||||
response=response,
|
||||
),
|
||||
last_status_code,
|
||||
)
|
||||
43
_deprecated_py_src/services/task/submit/rule_decider.py
Normal file
43
_deprecated_py_src/services/task/submit/rule_decider.py
Normal file
@@ -0,0 +1,43 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
|
||||
class AsyncSubmitRuleDeciderService:
|
||||
"""异步提交故障转移规则判定服务。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
match_provider_failover_rule: Callable[..., str | None],
|
||||
) -> None:
|
||||
self._match_provider_failover_rule = match_provider_failover_rule
|
||||
|
||||
def detect_error_stop_pattern(
|
||||
self,
|
||||
*,
|
||||
candidate: Any,
|
||||
response_text: str,
|
||||
status_code: int,
|
||||
) -> str | None:
|
||||
return self._match_provider_failover_rule(
|
||||
candidate,
|
||||
is_success=False,
|
||||
response_text=response_text,
|
||||
status_code=status_code,
|
||||
)
|
||||
|
||||
def detect_success_failover_pattern(
|
||||
self,
|
||||
*,
|
||||
candidate: Any,
|
||||
response_text: str,
|
||||
status_code: int,
|
||||
) -> str | None:
|
||||
return self._match_provider_failover_rule(
|
||||
candidate,
|
||||
is_success=True,
|
||||
response_text=response_text,
|
||||
status_code=status_code,
|
||||
)
|
||||
152
_deprecated_py_src/services/task/submit/submit_service.py
Normal file
152
_deprecated_py_src/services/task/submit/submit_service.py
Normal file
@@ -0,0 +1,152 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import ApiKey
|
||||
from src.services.candidate.failover import FailoverEngine
|
||||
from src.services.candidate.submit import SubmitOutcome
|
||||
from src.services.task.submit import ApplyPoolReorderFn, ExpandPoolCandidatesFn
|
||||
from src.services.task.submit.execute import AsyncSubmitExecutionService
|
||||
from src.services.task.submit.prepare import AsyncSubmitPreparationService
|
||||
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
class AsyncTaskSubmitService:
|
||||
"""异步任务提交应用服务(候选选择 + 故障转移)。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
redis_client: Any | None,
|
||||
*,
|
||||
apply_pool_reorder: ApplyPoolReorderFn,
|
||||
expand_pool_candidates_for_async_submit: ExpandPoolCandidatesFn,
|
||||
) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._apply_pool_reorder = apply_pool_reorder
|
||||
self._expand_pool_candidates_for_async_submit = expand_pool_candidates_for_async_submit
|
||||
self._prepare_ops = AsyncSubmitPreparationService(
|
||||
db,
|
||||
redis_client,
|
||||
sanitize=self._sanitize,
|
||||
)
|
||||
self._execute_ops = AsyncSubmitExecutionService(
|
||||
db,
|
||||
sanitize=self._sanitize,
|
||||
extract_response_text=self._extract_response_text,
|
||||
match_provider_failover_rule=self._match_provider_failover_rule,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _sanitize(message: str, max_length: int = 200) -> str:
|
||||
if not message:
|
||||
return "request_failed"
|
||||
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_text(response: httpx.Response) -> str:
|
||||
try:
|
||||
return response.text or ""
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _match_provider_failover_rule(
|
||||
candidate: Any,
|
||||
*,
|
||||
is_success: bool,
|
||||
response_text: str,
|
||||
status_code: int | None = None,
|
||||
) -> str | None:
|
||||
provider_config = getattr(candidate.provider, "config", None) or {}
|
||||
rules = provider_config.get("failover_rules")
|
||||
if not rules or not isinstance(rules, dict):
|
||||
return None
|
||||
|
||||
compiled = FailoverEngine._get_compiled_patterns(rules)
|
||||
key = "success" if is_success else "error"
|
||||
|
||||
for regex, rule in compiled.get(key, []):
|
||||
if not is_success:
|
||||
rule_status_codes = rule.get("status_codes")
|
||||
if rule_status_codes and status_code not in rule_status_codes:
|
||||
continue
|
||||
if regex.search(response_text):
|
||||
return rule.get("pattern", "")
|
||||
|
||||
return None
|
||||
|
||||
async def submit_with_failover(
|
||||
self,
|
||||
*,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey,
|
||||
request_id: str | None,
|
||||
task_type: str,
|
||||
submit_func: Any,
|
||||
extract_external_task_id: Any,
|
||||
supported_auth_types: set[str] | None = None,
|
||||
allow_format_conversion: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
max_candidates: int | None = None,
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> SubmitOutcome:
|
||||
"""
|
||||
异步提交入口。
|
||||
|
||||
行为保持与原 TaskService.submit_with_failover 一致:
|
||||
- 按候选顺序依次尝试(提交阶段不做单候选重试)
|
||||
- 记录 RequestCandidate 审计行
|
||||
- 命中 error_stop_patterns 时立即停止并抛出上游错误
|
||||
- 命中 success_failover_patterns 时继续尝试下一个候选
|
||||
"""
|
||||
# IMPORTANT:
|
||||
# This method awaits upstream HTTP calls. If we have an open DB transaction before awaiting,
|
||||
# the connection can be held for a long time (pool exhaustion under concurrency).
|
||||
#
|
||||
# Also note SQLAlchemy's default expire_on_commit=True would expire ORM objects and may
|
||||
# trigger unexpected lazy DB loads after we commit (potentially during the await).
|
||||
# We disable it temporarily to keep candidate/provider/key objects in-memory.
|
||||
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
|
||||
self.db.expire_on_commit = False
|
||||
|
||||
try:
|
||||
prepared = await self._prepare_ops.prepare_candidates(
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
user_api_key=user_api_key,
|
||||
request_id=request_id,
|
||||
capability_requirements=capability_requirements,
|
||||
request_body=request_body,
|
||||
max_candidates=max_candidates,
|
||||
apply_pool_reorder=self._apply_pool_reorder,
|
||||
expand_pool_candidates_for_async_submit=(
|
||||
self._expand_pool_candidates_for_async_submit
|
||||
),
|
||||
)
|
||||
|
||||
return await self._execute_ops.execute_submit_loop(
|
||||
candidates=prepared.candidates,
|
||||
record_map=prepared.record_map,
|
||||
task_type=task_type,
|
||||
model_name=model_name,
|
||||
submit_func=submit_func,
|
||||
extract_external_task_id=extract_external_task_id,
|
||||
supported_auth_types=supported_auth_types,
|
||||
allow_format_conversion=allow_format_conversion,
|
||||
)
|
||||
finally:
|
||||
# Restore Session behavior for the rest of the request lifecycle.
|
||||
self.db.expire_on_commit = original_expire_on_commit
|
||||
0
_deprecated_py_src/services/task/video/__init__.py
Normal file
0
_deprecated_py_src/services/task/video/__init__.py
Normal file
324
_deprecated_py_src/services/task/video/billing.py
Normal file
324
_deprecated_py_src/services/task/video/billing.py
Normal 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)
|
||||
266
_deprecated_py_src/services/task/video/cancel.py
Normal file
266
_deprecated_py_src/services/task/video/cancel.py
Normal 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
|
||||
38
_deprecated_py_src/services/task/video/facade.py
Normal file
38
_deprecated_py_src/services/task/video/facade.py
Normal file
@@ -0,0 +1,38 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.services.task.core.schema import TaskStatusResult
|
||||
from src.services.task.video.operations import VideoTaskOperationsService
|
||||
|
||||
|
||||
class TaskVideoFacadeService:
|
||||
"""视频任务门面服务(向后兼容 TaskService 的视频公开方法)。"""
|
||||
|
||||
def __init__(self, video_ops: VideoTaskOperationsService) -> None:
|
||||
self._video_ops = video_ops
|
||||
|
||||
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
|
||||
return await self._video_ops.poll(task_id, user_id=user_id)
|
||||
|
||||
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
|
||||
return await self._video_ops.poll_now(task_id, user_id=user_id)
|
||||
|
||||
async def cancel(
|
||||
self,
|
||||
task_id: str,
|
||||
*,
|
||||
user_id: str,
|
||||
original_headers: dict[str, str] | None = None,
|
||||
) -> Any:
|
||||
return await self._video_ops.cancel(
|
||||
task_id,
|
||||
user_id=user_id,
|
||||
original_headers=original_headers,
|
||||
)
|
||||
|
||||
async def finalize_video_task(self, task: Any) -> bool:
|
||||
return await self._video_ops.finalize_video_task(task)
|
||||
|
||||
async def finalize(self, task_id: str) -> bool:
|
||||
return await self._video_ops.finalize(task_id)
|
||||
134
_deprecated_py_src/services/task/video/operations.py
Normal file
134
_deprecated_py_src/services/task/video/operations.py
Normal file
@@ -0,0 +1,134 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import VideoTask
|
||||
from src.services.task.core.exceptions import TaskNotFoundError
|
||||
from src.services.task.core.schema import TaskStatusResult
|
||||
from src.services.task.video.billing import VideoTaskBillingService
|
||||
from src.services.task.video.cancel import VideoTaskCancelService
|
||||
|
||||
|
||||
class VideoTaskOperationsService:
|
||||
"""视频任务相关应用服务(轮询/取消/终态结算)。"""
|
||||
|
||||
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._billing_ops = VideoTaskBillingService(db)
|
||||
self._cancel_ops = VideoTaskCancelService(db)
|
||||
|
||||
def _extract_short_id(self, task_id: str) -> str:
|
||||
# Keep the parsing rule consistent with handlers:
|
||||
# - models/{model}/operations/{short_id}
|
||||
# - operations/{short_id}
|
||||
# - {short_id}
|
||||
return task_id.rsplit("/", 1)[-1] if "/" in task_id else task_id
|
||||
|
||||
def _get_video_task_for_user(self, task_id: str, *, user_id: str) -> Any:
|
||||
"""
|
||||
Resolve a video task by:
|
||||
- internal UUID (VideoTask.id)
|
||||
- external operation id (VideoTask.short_id)
|
||||
"""
|
||||
task = (
|
||||
self.db.query(VideoTask)
|
||||
.filter(VideoTask.id == task_id, VideoTask.user_id == user_id)
|
||||
.first()
|
||||
)
|
||||
if task:
|
||||
return task
|
||||
|
||||
short_id = self._extract_short_id(task_id)
|
||||
task = (
|
||||
self.db.query(VideoTask)
|
||||
.filter(VideoTask.short_id == short_id, VideoTask.user_id == user_id)
|
||||
.first()
|
||||
)
|
||||
if not task:
|
||||
raise TaskNotFoundError(task_id)
|
||||
return task
|
||||
|
||||
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
|
||||
"""Read task status from DB (does not trigger polling)."""
|
||||
task = self._get_video_task_for_user(task_id, user_id=user_id)
|
||||
|
||||
result_url = None
|
||||
if getattr(task, "status", None) == "completed":
|
||||
result_url = getattr(task, "video_url", None)
|
||||
|
||||
error_message = None
|
||||
if getattr(task, "status", None) == "failed":
|
||||
error_message = getattr(task, "error_message", None) or getattr(
|
||||
task, "error_code", None
|
||||
)
|
||||
|
||||
return TaskStatusResult(
|
||||
task_id=str(getattr(task, "id", task_id)),
|
||||
status=str(getattr(task, "status", "unknown")),
|
||||
progress_percent=int(getattr(task, "progress_percent", 0) or 0),
|
||||
result_url=result_url,
|
||||
error_message=str(error_message) if error_message else None,
|
||||
provider_id=(
|
||||
str(getattr(task, "provider_id", None))
|
||||
if getattr(task, "provider_id", None)
|
||||
else None
|
||||
),
|
||||
provider_name=(
|
||||
str(getattr(task, "provider_name", None))
|
||||
if getattr(task, "provider_name", None)
|
||||
else None
|
||||
),
|
||||
endpoint_id=(
|
||||
str(getattr(task, "endpoint_id", None))
|
||||
if getattr(task, "endpoint_id", None)
|
||||
else None
|
||||
),
|
||||
key_id=str(getattr(task, "key_id", None)) if getattr(task, "key_id", None) else None,
|
||||
)
|
||||
|
||||
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
|
||||
"""
|
||||
Trigger a single polling attempt (best-effort), then return latest DB status.
|
||||
|
||||
Note: this uses the poller adapter's single-task method and may hold a DB
|
||||
connection during the upstream HTTP request; keep usage low.
|
||||
"""
|
||||
from src.services.task.video.poller_adapter import VideoTaskPollerAdapter
|
||||
|
||||
task = self._get_video_task_for_user(task_id, user_id=user_id)
|
||||
adapter = VideoTaskPollerAdapter()
|
||||
await adapter.poll_single_task(self.db, task, redis_client=self.redis)
|
||||
self.db.commit()
|
||||
return await self.poll(task_id, user_id=user_id)
|
||||
|
||||
async def cancel(
|
||||
self,
|
||||
task_id: str,
|
||||
*,
|
||||
user_id: str,
|
||||
original_headers: dict[str, str] | None = None,
|
||||
) -> Any:
|
||||
from fastapi import HTTPException
|
||||
|
||||
try:
|
||||
task = self._get_video_task_for_user(task_id, user_id=user_id)
|
||||
except TaskNotFoundError:
|
||||
raise HTTPException(status_code=404, detail="Video task not found")
|
||||
return await self._cancel_ops.cancel_task(
|
||||
task=task,
|
||||
task_id=task_id,
|
||||
original_headers=original_headers,
|
||||
)
|
||||
|
||||
async def finalize_video_task(self, task: Any) -> bool:
|
||||
return await self._billing_ops.finalize_video_task(task)
|
||||
|
||||
async def finalize(self, task_id: str) -> bool:
|
||||
"""Finalize a task by internal id (best-effort)."""
|
||||
task = self.db.query(VideoTask).filter(VideoTask.id == task_id).first()
|
||||
if not task:
|
||||
return False
|
||||
return await self.finalize_video_task(task)
|
||||
656
_deprecated_py_src/services/task/video/poller_adapter.py
Normal file
656
_deprecated_py_src/services/task/video/poller_adapter.py
Normal 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])
|
||||
Reference in New Issue
Block a user