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

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

View File

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

View File

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

View File

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

View File

@@ -0,0 +1,265 @@
from __future__ import annotations
from typing import Any
from src.core.logger import logger
class TaskPoolOperationsService:
"""任务池化相关操作(重排、展开、健康回写)。"""
def extract_session_uuid(
self,
provider_type: str,
request_body: dict[str, Any] | None,
) -> str | None:
"""Extract a session UUID from the request body (provider-type aware)."""
if not isinstance(request_body, dict):
return None
from src.services.provider.pool.hooks import get_pool_hook
hook = get_pool_hook(provider_type)
if hook is not None:
return hook.extract_session_uuid(request_body)
return None
async def apply_pool_reorder(
self,
candidates: list[Any],
request_body: dict[str, Any] | None,
) -> tuple[list[Any], list[Any]]:
"""Apply pool key ordering for PoolCandidate objects."""
if not candidates:
return candidates, []
pool_traces: list[Any] = []
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.manager import PoolManager
from src.services.scheduling.schemas import PoolCandidate
for candidate in candidates:
if not isinstance(candidate, PoolCandidate):
continue
provider = candidate.provider
provider_id = str(getattr(provider, "id", "") or "")
if not provider_id:
continue
pool_cfg = candidate.pool_config or parse_pool_config(
getattr(provider, "config", None)
)
if pool_cfg is None:
continue
candidate.pool_config = pool_cfg
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = self.extract_session_uuid(provider_type, request_body)
manager = PoolManager(provider_id, pool_cfg)
candidate_keys = list(candidate.pool_keys or [])
if not candidate_keys and getattr(candidate, "key", None) is not None:
candidate_keys = [candidate.key]
# 构造延迟可用性检查回调(从 CandidateBuilder 打包的参数)
checker = None
deferred_params = candidate._deferred_check_params
if deferred_params is not None:
checker = self._build_availability_checker(deferred_params)
ordered_keys, trace = await manager.select_pool_keys(
session_uuid,
candidate_keys,
availability_checker=checker,
)
# 移除 deferred key未检查的避免为其创建 DB 记录
candidate.pool_keys = [
k for k in ordered_keys if getattr(k, "_pool_skip_reason", None) != "deferred"
]
selected_key_index = 0
selected_key = None
for idx, pool_key in enumerate(candidate.pool_keys):
if not bool(getattr(pool_key, "_pool_skipped", False)):
selected_key = pool_key
selected_key_index = idx
break
if selected_key is not None:
candidate.key = selected_key
candidate._pool_key_index = selected_key_index
candidate.mapping_matched_model = getattr(
selected_key, "_pool_mapping_matched_model", None
)
candidate.is_skipped = False
candidate.skip_reason = None
else:
candidate.is_skipped = True
candidate.skip_reason = "pool: all keys unavailable"
if trace is not None:
pool_traces.append(trace)
return candidates, pool_traces
except Exception:
logger.opt(exception=True).debug("Pool reorder failed, using original order")
return candidates, []
@staticmethod
def expand_pool_candidates_for_async_submit(candidates: list[Any]) -> list[Any]:
"""Expand PoolCandidate to key-level candidates for async submit traversal."""
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
expanded: list[Any] = []
for candidate in candidates:
if not isinstance(candidate, PoolCandidate):
expanded.append(candidate)
continue
pool_keys = list(candidate.pool_keys or [])
if not pool_keys:
expanded.append(candidate)
continue
for key_index, pool_key in enumerate(pool_keys):
key_skipped = bool(getattr(pool_key, "_pool_skipped", False))
key_skip_reason = (
str(getattr(pool_key, "_pool_skip_reason", "") or "") or candidate.skip_reason
)
key_extra = (
getattr(pool_key, "_pool_extra_data", None)
if isinstance(getattr(pool_key, "_pool_extra_data", None), dict)
else {}
)
key_candidate = ProviderCandidate(
provider=candidate.provider,
endpoint=candidate.endpoint,
key=pool_key,
is_cached=candidate.is_cached,
is_skipped=bool(candidate.is_skipped) or key_skipped,
skip_reason=(
key_skip_reason if (bool(candidate.is_skipped) or key_skipped) else None
),
mapping_matched_model=getattr(pool_key, "_pool_mapping_matched_model", None)
or candidate.mapping_matched_model,
needs_conversion=candidate.needs_conversion,
provider_api_format=candidate.provider_api_format,
output_limit=candidate.output_limit,
capability_miss_count=candidate.capability_miss_count,
)
setattr(
key_candidate,
"_pool_extra_data",
{
"pool_group_id": str(candidate.provider.id),
"pool_key_index": key_index,
**key_extra,
},
)
expanded.append(key_candidate)
return expanded
async def pool_on_success(
self,
candidate: Any,
request_body: dict[str, Any] | None,
) -> None:
"""Notify the pool manager about a successful request (sticky + LRU)."""
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.manager import PoolManager
provider = candidate.provider
provider_config = getattr(provider, "config", None)
pool_cfg = parse_pool_config(provider_config)
if pool_cfg is None:
return
provider_id = str(getattr(provider, "id", "") or "")
key_id = str(getattr(candidate.key, "id", "") or "")
if not provider_id or not key_id:
return
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = self.extract_session_uuid(provider_type, request_body)
mgr = PoolManager(provider_id, pool_cfg)
await mgr.on_request_success(
session_uuid=session_uuid,
key_id=key_id,
)
except Exception:
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
@staticmethod
async def pool_on_error(
provider: Any,
key: Any,
status_code: int,
cause: Any,
) -> None:
"""Notify the pool manager about an upstream error (health policy)."""
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.health_policy import apply_health_policy
pool_cfg = parse_pool_config(getattr(provider, "config", None))
if pool_cfg is None:
return
error_text = ""
resp_headers: dict[str, str] = {}
if getattr(cause, "response", None) is not None:
try:
error_text = (cause.response.text or "")[:4000]
except Exception:
pass
try:
resp_headers = dict(cause.response.headers)
except Exception:
pass
await apply_health_policy(
provider_id=str(provider.id),
key_id=str(key.id),
status_code=status_code,
error_body=error_text,
response_headers=resp_headers,
config=pool_cfg,
)
except Exception:
pass
@staticmethod
def _build_availability_checker(
params: dict[str, Any],
) -> Any:
"""Construct a key availability checker from deferred check params."""
from src.services.scheduling.candidate_builder import CandidateBuilder
endpoint_format = params.get("endpoint_format")
model_name = params.get("model_name", "")
capability_requirements = params.get("capability_requirements")
model_mappings = params.get("model_mappings")
candidate_models = params.get("candidate_models")
provider_type = params.get("provider_type")
# _check_key_availability 不依赖 _sorter传 None 安全
builder = CandidateBuilder(candidate_sorter=None) # type: ignore[arg-type]
def _checker(key: Any) -> tuple[bool, str | None, str | None]:
return builder._check_key_availability(
key,
endpoint_format,
model_name,
capability_requirements,
model_mappings=model_mappings,
candidate_models=candidate_models,
provider_type=provider_type,
)
return _checker

View File

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

View File

@@ -0,0 +1,421 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
from uuid import uuid4
from src.core.logger import logger
from src.models.database import ApiKey, User
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import FailoverAction, RetryPolicy, SkipPolicy
from src.services.candidate.resolver import CandidateResolver
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.orchestration.request_dispatcher import RequestDispatcher
from src.services.provider.format import normalize_endpoint_signature
from src.services.request.candidate import RequestCandidateService
from src.services.scheduling.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.system.config import SystemConfigService
from src.services.task.core.protocol import AttemptKind, AttemptResult
from src.services.task.core.schema import ExecutionResult
from src.services.task.execute.error_handler import TaskErrorOperationsService
from src.services.task.execute.exception_classification import (
CandidateErrorAction,
classify_candidate_error_action,
)
from src.services.task.execute.failure import TaskFailureOperationsService
from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.execute.state_transition import (
SyncExecutionState,
resolve_execution_error_transition,
)
from src.services.task.request_state import RequestBodyState
from src.services.usage.service import UsageService
class SyncTaskExecutionService:
"""同步任务执行服务(候选遍历 + 错误处理 + 结果聚合)。"""
def __init__(
self,
db: Any,
redis_client: Any | None,
*,
recorder: Any,
pool_ops: TaskPoolOperationsService,
error_ops: TaskErrorOperationsService,
failure_ops: TaskFailureOperationsService,
) -> None:
self.db = db
self.redis = redis_client
self._recorder = recorder
self._pool_ops = pool_ops
self._error_ops = error_ops
self._failure_ops = failure_ops
async def execute_sync_unified(
self,
*,
api_format: str,
model_name: str,
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None,
is_stream: bool,
capability_requirements: dict[str, bool] | None,
preferred_key_ids: list[str] | None,
request_body_state: RequestBodyState | None,
request_headers: dict[str, Any] | None,
request_body: dict[str, Any] | None,
create_pending_usage: bool = True,
) -> ExecutionResult:
"""
Unified candidate traversal loop for SYNC.
This intentionally reuses existing components for parity:
- CandidateResolver fetch + record creation
- RequestDispatcher execution
- Error classification/rectify logic ported from the previous SYNC implementation
"""
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.request.executor import RequestExecutor
if not request_id:
request_id = str(uuid4())
# IMPORTANT:
# This SYNC path awaits upstream HTTP work while the failover engine may commit
# candidate audit rows between attempts. SQLAlchemy's default
# expire_on_commit=True would expire provider/endpoint/key ORM objects and can
# trigger an unexpected lazy DB reload later in error handling (for example when
# reading candidate.provider.config for failover_rules after a timeout).
#
# Keep already-loaded candidate objects resident in memory for the duration of
# the request, mirroring the async submit path.
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
self.db.expire_on_commit = False
try:
# Build execution components (mirrors pre-Phase-3 initialization)
priority_mode = SystemConfigService.get_config(
self.db,
"provider_priority_mode",
CacheAwareScheduler.PRIORITY_MODE_PROVIDER,
)
scheduling_mode = SystemConfigService.get_config(
self.db,
"scheduling_mode",
CacheAwareScheduler.SCHEDULING_MODE_CACHE_AFFINITY,
)
cache_scheduler = await get_cache_aware_scheduler(
self.redis,
priority_mode=priority_mode,
scheduling_mode=scheduling_mode,
)
# Ensure cache_scheduler inner state is ready
await cache_scheduler._ensure_initialized()
concurrency_manager = await get_concurrency_manager()
adaptive_manager = get_adaptive_rpm_manager()
request_executor = RequestExecutor(
db=self.db,
concurrency_manager=concurrency_manager,
adaptive_manager=adaptive_manager,
)
candidate_resolver = CandidateResolver(
db=self.db,
cache_scheduler=cache_scheduler,
)
error_classifier = ErrorClassifier(
db=self.db,
cache_scheduler=cache_scheduler,
adaptive_manager=adaptive_manager,
)
request_dispatcher = RequestDispatcher(
db=self.db,
request_executor=request_executor,
cache_scheduler=cache_scheduler,
)
affinity_key = str(user_api_key.id)
user_id = str(user_api_key.user_id)
api_format_norm = normalize_endpoint_signature(api_format)
user: User | None = None
username_snapshot = None
api_key_name_snapshot = getattr(user_api_key, "name", None)
try:
user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
username_snapshot = getattr(user, "username", None) if user else None
except Exception as exc:
# username 仅用于审计快照,不应阻塞主请求链路。
logger.warning("查询用户快照失败: {}", str(exc))
# 默认由 TaskService 创建 pending 使用记录;已预创建的调用方可关闭。
if create_pending_usage:
try:
UsageService.create_pending_usage(
db=self.db,
request_id=request_id,
user=user,
api_key=user_api_key,
model=model_name,
is_stream=is_stream,
api_format=api_format_norm,
request_headers=request_headers,
request_body=request_body,
)
except Exception as exc:
logger.warning("创建 pending 使用记录失败: {}", str(exc))
all_candidates, global_model_id = await candidate_resolver.fetch_candidates(
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body=request_body,
)
# 号池排序涉及大量 Redis 操作,提前释放 DB 连接避免连接池压力
from src.services.scheduling.utils import release_db_connection_before_await
release_db_connection_before_await(self.db)
# Account Pool: reorder candidates for claude_code providers.
all_candidates, pool_traces = await self._pool_ops.apply_pool_reorder(
all_candidates, request_body=request_body
)
candidate_record_map = await candidate_resolver.create_candidate_records_async(
all_candidates=all_candidates,
request_id=request_id,
user_id=user_id,
user_api_key=user_api_key,
required_capabilities=capability_requirements,
)
max_attempts = candidate_resolver.count_total_attempts(all_candidates)
# Keep behavior consistent with previous behavior: last_candidate is updated even if skipped.
execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map,
request_body_state=request_body_state,
last_candidate=all_candidates[-1] if all_candidates else None,
)
async def _attempt(candidate: Any) -> AttemptResult:
execution_state.touch_candidate(candidate)
candidate_index = int(getattr(candidate, "_utf_candidate_index", -1))
retry_index = int(getattr(candidate, "_utf_retry_index", 0))
candidate_record_id = str(getattr(candidate, "_utf_candidate_record_id", "") or "")
attempt_counter = int(getattr(candidate, "_utf_attempt_count", 0))
max_attempts_local = int(getattr(candidate, "_utf_max_attempts", max_attempts))
# Safety net: if record_id missing, create an "available" record on-demand.
if not candidate_record_id:
from src.services.scheduling.schemas import PoolCandidate
pool_extra = (
getattr(candidate.key, "_pool_extra_data", None)
if isinstance(getattr(candidate.key, "_pool_extra_data", None), dict)
else {}
)
extra_data: dict[str, Any] = {
"needs_conversion": bool(getattr(candidate, "needs_conversion", False)),
"provider_api_format": getattr(candidate, "provider_api_format", None)
or None,
"mapping_matched_model": getattr(candidate, "mapping_matched_model", None)
or None,
**pool_extra,
}
if isinstance(candidate, PoolCandidate):
extra_data["pool_group_id"] = str(candidate.provider.id)
extra_data["pool_key_index"] = int(
getattr(candidate, "_pool_key_index", 0) or 0
)
created = RequestCandidateService.create_candidate(
db=self.db,
request_id=request_id,
candidate_index=candidate_index,
retry_index=retry_index,
user_id=user_id,
api_key_id=str(user_api_key.id),
username=username_snapshot,
api_key_name=api_key_name_snapshot,
provider_id=str(candidate.provider.id),
endpoint_id=str(candidate.endpoint.id),
key_id=str(candidate.key.id),
status="available",
is_cached=bool(getattr(candidate, "is_cached", False)),
extra_data=extra_data,
)
candidate_record_id = str(created.id)
execution_state.candidate_record_map[(candidate_index, retry_index)] = (
candidate_record_id
)
(
response,
_provider_name,
attempt_id,
_provider_id,
_endpoint_id,
_key_id,
_first_byte_time_ms,
) = await request_dispatcher.dispatch(
candidate=candidate,
candidate_index=candidate_index,
retry_index=retry_index,
candidate_record_id=candidate_record_id,
user_api_key=user_api_key,
user_id=user_id,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
_ = (
attempt_id,
_provider_name,
_provider_id,
_endpoint_id,
_key_id,
_first_byte_time_ms,
)
# Account Pool: on success, update sticky binding + LRU.
await self._pool_ops.pool_on_success(candidate, request_body)
if is_stream:
return AttemptResult(
kind=AttemptKind.STREAM,
http_status=200,
http_headers={},
stream_iterator=response,
)
return AttemptResult(
kind=AttemptKind.SYNC_RESPONSE,
http_status=200,
http_headers={},
response_body=response,
)
async def _handle_exec_err(
*,
exec_err: Any,
candidate: Any,
candidate_index: int,
retry_index: int,
max_retries_for_candidate: int,
record_id: str | None,
attempt_count: int,
max_attempts: int | None,
) -> tuple[FailoverAction, int | None]:
execution_state.track_execution_error(exec_err=exec_err, candidate=candidate)
# Fall back to retry 0 record if needed (rectify may extend retries).
candidate_record_id = execution_state.resolve_candidate_record_id(
candidate_index=candidate_index,
record_id=record_id,
)
raw_action = await self._error_ops.handle_candidate_error(
exec_err=exec_err,
candidate=candidate,
candidate_record_id=candidate_record_id,
retry_index=retry_index,
max_retries_for_candidate=max_retries_for_candidate,
affinity_key=affinity_key,
api_format=api_format_norm,
global_model_id=global_model_id,
request_id=request_id,
attempt=attempt_count,
max_attempts=int(max_attempts or 0),
request_body_state=request_body_state,
error_classifier=error_classifier,
)
action = classify_candidate_error_action(raw_action)
if action == CandidateErrorAction.RAISE_ERROR:
execution_state.raise_classified_error(
fallback_error=exec_err,
failure_ops=self._failure_ops,
model_name=model_name,
api_format=api_format_norm,
)
return resolve_execution_error_transition(
action=action,
state=execution_state,
max_retries_for_candidate=max_retries_for_candidate,
retry_index=retry_index,
).as_failover_tuple()
engine = FailoverEngine(
self.db,
error_classifier=error_classifier,
recorder=self._recorder,
)
result = await engine.execute(
candidates=all_candidates,
attempt_func=_attempt,
retry_policy=RetryPolicy.for_sync_task(),
skip_policy=SkipPolicy(),
request_id=request_id,
user_id=user_id,
api_key_id=str(user_api_key.id),
username=username_snapshot,
api_key_name=api_key_name_snapshot,
candidate_record_map=candidate_record_map,
max_attempts=max_attempts,
execution_error_handler=_handle_exec_err,
)
if result.success:
# Build pool scheduling summary from traces collected during reorder.
if pool_traces and result.key_id:
try:
attempted_key_ids: set[str] = set()
for ck in result.candidate_keys or []:
status = str(getattr(ck, "status", "") or "").strip().lower()
if status in {"", "available", "pending", "skipped", "unused"}:
continue
kid = getattr(ck, "key_id", None)
if isinstance(kid, str) and kid:
attempted_key_ids.add(kid)
if not attempted_key_ids:
attempted_key_ids.add(str(result.key_id))
for pt in pool_traces:
summary = pt.build_summary(
result.key_id,
attempted_key_ids=attempted_key_ids,
)
if summary:
result.pool_summary = summary
break
except Exception:
pass
return result
self._failure_ops.raise_all_failed_exception(
request_id,
max_attempts,
execution_state.last_candidate,
model_name,
api_format_norm,
execution_state.last_error,
)
finally:
self.db.expire_on_commit = original_expire_on_commit