mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +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:
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
|
||||
Reference in New Issue
Block a user