mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat(pool): 号池候选重构为 PoolCandidate 单候选模式与池内 key 故障转移
- 新增 PoolCandidate 子类,排序阶段作为单候选参与,执行阶段在 pool_keys 内部选择/切换 key - FailoverEngine 新增 _execute_pool_candidate 方法,支持池内 key 级别故障转移与重试 - 提取 _execute_attempt / _attach_attempt_context / _classify_attempt_error 公共方法 - CandidateBuilder 对号池 Provider 构建单个 PoolCandidate(包含所有可用 key) - CandidateSorter 支持 PoolCandidate 独立优先级分组(global_priority / pool_priority) - CandidateResolver PRE_EXPAND 模式按 pool_keys 展开预创建记录,附加 pool_group_id - TaskService._apply_pool_reorder 改为对 PoolCandidate 调用 select_pool_keys - PoolManager 新增 select_pool_keys 方法,复用 reorder_candidates 逻辑 - 新增 global_priority 号池配置字段(前后端同步) - 前端 Timeline 支持按 pool_group_id 分组显示多号池尝试 - 异步提交路径新增 _expand_pool_candidates_for_async_submit 展开逻辑
This commit is contained in:
@@ -454,6 +454,7 @@ export interface ClaudeCodeAdvancedConfig {
|
|||||||
}
|
}
|
||||||
|
|
||||||
export interface PoolAdvancedConfig {
|
export interface PoolAdvancedConfig {
|
||||||
|
global_priority?: number | null
|
||||||
sticky_session_ttl_seconds?: number | null
|
sticky_session_ttl_seconds?: number | null
|
||||||
load_threshold_percent?: number | null
|
load_threshold_percent?: number | null
|
||||||
lru_enabled?: boolean
|
lru_enabled?: boolean
|
||||||
|
|||||||
@@ -47,6 +47,23 @@
|
|||||||
同一对话始终路由到同一 Key
|
同一对话始终路由到同一 Key
|
||||||
</p>
|
</p>
|
||||||
</div>
|
</div>
|
||||||
|
<div class="space-y-1.5">
|
||||||
|
<Label>
|
||||||
|
全局优先级
|
||||||
|
<span class="text-xs text-muted-foreground">(global_key)</span>
|
||||||
|
</Label>
|
||||||
|
<Input
|
||||||
|
:model-value="form.global_priority ?? ''"
|
||||||
|
type="number"
|
||||||
|
min="0"
|
||||||
|
max="999999"
|
||||||
|
placeholder="留空回退 provider_priority"
|
||||||
|
@update:model-value="(v) => form.global_priority = parseNum(v)"
|
||||||
|
/>
|
||||||
|
<p class="text-xs text-muted-foreground">
|
||||||
|
global_key 模式下号池整体排序值(越小越优先)
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
@@ -320,6 +337,7 @@ const { success, error: showError } = useToast()
|
|||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
|
|
||||||
const form = ref<PoolAdvancedConfig>({
|
const form = ref<PoolAdvancedConfig>({
|
||||||
|
global_priority: null,
|
||||||
sticky_session_ttl_seconds: null,
|
sticky_session_ttl_seconds: null,
|
||||||
lru_enabled: true,
|
lru_enabled: true,
|
||||||
cost_window_seconds: null,
|
cost_window_seconds: null,
|
||||||
@@ -363,6 +381,7 @@ watch(() => props.modelValue, (v) => {
|
|||||||
form.value = { ...props.currentConfig }
|
form.value = { ...props.currentConfig }
|
||||||
} else if (v) {
|
} else if (v) {
|
||||||
form.value = {
|
form.value = {
|
||||||
|
global_priority: null,
|
||||||
sticky_session_ttl_seconds: null,
|
sticky_session_ttl_seconds: null,
|
||||||
lru_enabled: true,
|
lru_enabled: true,
|
||||||
cost_window_seconds: null,
|
cost_window_seconds: null,
|
||||||
@@ -411,6 +430,7 @@ async function handleSave() {
|
|||||||
try {
|
try {
|
||||||
const payload: Record<string, unknown> = {
|
const payload: Record<string, unknown> = {
|
||||||
pool_advanced: {
|
pool_advanced: {
|
||||||
|
global_priority: form.value.global_priority ?? undefined,
|
||||||
sticky_session_ttl_seconds: form.value.sticky_session_ttl_seconds ?? undefined,
|
sticky_session_ttl_seconds: form.value.sticky_session_ttl_seconds ?? undefined,
|
||||||
lru_enabled: form.value.lru_enabled,
|
lru_enabled: form.value.lru_enabled,
|
||||||
cost_window_seconds: form.value.cost_window_seconds ?? undefined,
|
cost_window_seconds: form.value.cost_window_seconds ?? undefined,
|
||||||
|
|||||||
@@ -666,7 +666,23 @@ const schedulingAudit = computed<Record<string, unknown> | null>(() => {
|
|||||||
return raw as Record<string, unknown>
|
return raw as Record<string, unknown>
|
||||||
})
|
})
|
||||||
|
|
||||||
|
const extractPoolGroupId = (candidate: CandidateRecord): string | null => {
|
||||||
|
const extra = candidate.extra_data
|
||||||
|
if (!extra || typeof extra !== 'object' || Array.isArray(extra)) return null
|
||||||
|
const value = (extra as Record<string, unknown>).pool_group_id
|
||||||
|
if (typeof value !== 'string') return null
|
||||||
|
const text = value.trim()
|
||||||
|
return text || null
|
||||||
|
}
|
||||||
|
|
||||||
const poolAttemptCandidates = computed<CandidateRecord[]>(() => {
|
const poolAttemptCandidates = computed<CandidateRecord[]>(() => {
|
||||||
|
// 新链路:优先使用后端写入的 extra_data.pool_group_id。
|
||||||
|
const fromTrace = rawTimeline.value.filter((candidate) => extractPoolGroupId(candidate) !== null)
|
||||||
|
if (fromTrace.length > 0) {
|
||||||
|
return fromTrace
|
||||||
|
}
|
||||||
|
|
||||||
|
// 兼容旧链路:回退到 request_metadata.scheduling_audit.attempts。
|
||||||
const audit = schedulingAudit.value
|
const audit = schedulingAudit.value
|
||||||
if (!audit) return []
|
if (!audit) return []
|
||||||
const attempts = audit.attempts
|
const attempts = audit.attempts
|
||||||
@@ -711,11 +727,37 @@ const poolAttemptCandidates = computed<CandidateRecord[]>(() => {
|
|||||||
if (typeof raw.key_name === 'string') merged.key_name = raw.key_name
|
if (typeof raw.key_name === 'string') merged.key_name = raw.key_name
|
||||||
if (typeof raw.status_code === 'number') merged.status_code = raw.status_code
|
if (typeof raw.status_code === 'number') merged.status_code = raw.status_code
|
||||||
if (typeof raw.error_type === 'string') merged.error_type = raw.error_type
|
if (typeof raw.error_type === 'string') merged.error_type = raw.error_type
|
||||||
|
const rawPoolGroupId = typeof raw.pool_group_id === 'string' ? raw.pool_group_id.trim() : ''
|
||||||
|
const fallbackPoolGroupId = typeof raw.provider_id === 'string' ? raw.provider_id.trim() : ''
|
||||||
|
const finalPoolGroupId = rawPoolGroupId || fallbackPoolGroupId
|
||||||
|
if (finalPoolGroupId) {
|
||||||
|
merged.extra_data = {
|
||||||
|
...(merged.extra_data || {}),
|
||||||
|
pool_group_id: finalPoolGroupId,
|
||||||
|
}
|
||||||
|
}
|
||||||
return merged
|
return merged
|
||||||
})
|
})
|
||||||
.filter((item): item is CandidateRecord => item !== null)
|
.filter((item): item is CandidateRecord => item !== null)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
const poolAttemptsByGroup = computed<Map<string, CandidateRecord[]>>(() => {
|
||||||
|
const grouped = new Map<string, CandidateRecord[]>()
|
||||||
|
for (const attempt of poolAttemptCandidates.value) {
|
||||||
|
const groupId =
|
||||||
|
extractPoolGroupId(attempt)
|
||||||
|
|| String(attempt.provider_id || '').trim()
|
||||||
|
|| '__pool_group__'
|
||||||
|
const existing = grouped.get(groupId)
|
||||||
|
if (existing) {
|
||||||
|
existing.push(attempt)
|
||||||
|
} else {
|
||||||
|
grouped.set(groupId, [attempt])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return grouped
|
||||||
|
})
|
||||||
|
|
||||||
const poolAttemptKeySet = computed<Set<string>>(() => {
|
const poolAttemptKeySet = computed<Set<string>>(() => {
|
||||||
return new Set(
|
return new Set(
|
||||||
poolAttemptCandidates.value.map((item) => makeAttemptKey(item.candidate_index, item.retry_index)),
|
poolAttemptCandidates.value.map((item) => makeAttemptKey(item.candidate_index, item.retry_index)),
|
||||||
@@ -805,45 +847,56 @@ const buildProviderGroups = (items: CandidateRecord[]): NodeGroup[] => {
|
|||||||
// 将相同 Provider 的所有请求合并为组(同提供商的 Key 放在子节点)
|
// 将相同 Provider 的所有请求合并为组(同提供商的 Key 放在子节点)
|
||||||
const groupedTimeline = computed<NodeGroup[]>(() => {
|
const groupedTimeline = computed<NodeGroup[]>(() => {
|
||||||
const providerGroups = buildProviderGroups(timeline.value)
|
const providerGroups = buildProviderGroups(timeline.value)
|
||||||
const poolAttempts = poolAttemptCandidates.value
|
if (poolAttemptsByGroup.value.size === 0) {
|
||||||
if (poolAttempts.length === 0) {
|
|
||||||
return providerGroups
|
return providerGroups
|
||||||
}
|
}
|
||||||
|
|
||||||
const poolPrimaryStatus = poolAttempts.reduce((best, current) => {
|
const poolProviderIds = new Set<string>()
|
||||||
|
const poolProviderNames = new Set<string>()
|
||||||
|
const poolGroups: NodeGroup[] = []
|
||||||
|
|
||||||
|
for (const [groupId, attemptsRaw] of poolAttemptsByGroup.value.entries()) {
|
||||||
|
const attempts = [...attemptsRaw].sort((a, b) => {
|
||||||
|
if (a.candidate_index !== b.candidate_index) {
|
||||||
|
return a.candidate_index - b.candidate_index
|
||||||
|
}
|
||||||
|
return a.retry_index - b.retry_index
|
||||||
|
})
|
||||||
|
if (attempts.length === 0) continue
|
||||||
|
|
||||||
|
const poolPrimaryStatus = attempts.reduce((best, current) => {
|
||||||
const bestPriority = STATUS_PRIORITY[best] ?? 0
|
const bestPriority = STATUS_PRIORITY[best] ?? 0
|
||||||
const currentPriority = STATUS_PRIORITY[current.status] ?? 0
|
const currentPriority = STATUS_PRIORITY[current.status] ?? 0
|
||||||
return currentPriority > bestPriority ? current.status : best
|
return currentPriority > bestPriority ? current.status : best
|
||||||
}, poolAttempts[0].status)
|
}, attempts[0].status)
|
||||||
|
|
||||||
const successAttempt = poolAttempts.find((item) => item.status === 'success')
|
const successAttempt = attempts.find((item) => item.status === 'success')
|
||||||
const poolPrimary = successAttempt || poolAttempts[poolAttempts.length - 1] || poolAttempts[0]
|
const poolPrimary = successAttempt || attempts[attempts.length - 1] || attempts[0]
|
||||||
|
const startIndex = Math.min(...attempts.map(item => item.candidate_index))
|
||||||
|
const endIndex = Math.max(...attempts.map(item => item.candidate_index))
|
||||||
|
|
||||||
const poolGroup: NodeGroup = {
|
poolGroups.push({
|
||||||
id: '__pool_group__',
|
id: `pool:${groupId}`,
|
||||||
providerName: getProviderDisplayName(poolPrimary),
|
providerName: getProviderDisplayName(poolPrimary),
|
||||||
primary: poolPrimary,
|
primary: poolPrimary,
|
||||||
primaryStatus: poolPrimaryStatus,
|
primaryStatus: poolPrimaryStatus,
|
||||||
allAttempts: poolAttempts,
|
allAttempts: attempts,
|
||||||
retryCount: Math.max(0, poolAttempts.length - 1),
|
retryCount: Math.max(0, attempts.length - 1),
|
||||||
totalLatency: poolAttempts.reduce((sum, item) => sum + (item.latency_ms || 0), 0),
|
totalLatency: attempts.reduce((sum, item) => sum + (item.latency_ms || 0), 0),
|
||||||
startIndex: 0,
|
startIndex,
|
||||||
endIndex: poolAttempts.length - 1,
|
endIndex,
|
||||||
hasConversion: poolAttempts.some((item) => item.extra_data?.needs_conversion === true),
|
hasConversion: attempts.some((item) => item.extra_data?.needs_conversion === true),
|
||||||
providerApiFormat: null,
|
providerApiFormat: null,
|
||||||
isPoolGroup: true,
|
isPoolGroup: true,
|
||||||
}
|
})
|
||||||
|
|
||||||
const poolProviderIds = new Set(
|
for (const attempt of attempts) {
|
||||||
poolAttempts
|
const providerId = String(attempt.provider_id || '').trim()
|
||||||
.map(item => String(item.provider_id || '').trim())
|
if (providerId) poolProviderIds.add(providerId)
|
||||||
.filter(Boolean),
|
const providerName = normalizeProviderIdentity(attempt.provider_name)
|
||||||
)
|
if (providerName) poolProviderNames.add(providerName)
|
||||||
const poolProviderNames = new Set(
|
}
|
||||||
poolAttempts
|
}
|
||||||
.map(item => normalizeProviderIdentity(item.provider_name))
|
|
||||||
.filter(Boolean),
|
|
||||||
)
|
|
||||||
|
|
||||||
const dedupedProviderGroups = providerGroups.filter((group) => {
|
const dedupedProviderGroups = providerGroups.filter((group) => {
|
||||||
const sameProviderById = group.allAttempts.some((attempt) => {
|
const sameProviderById = group.allAttempts.some((attempt) => {
|
||||||
@@ -858,7 +911,8 @@ const groupedTimeline = computed<NodeGroup[]>(() => {
|
|||||||
return true
|
return true
|
||||||
})
|
})
|
||||||
|
|
||||||
return [poolGroup, ...dedupedProviderGroups]
|
poolGroups.sort((a, b) => a.startIndex - b.startIndex)
|
||||||
|
return [...poolGroups, ...dedupedProviderGroups]
|
||||||
})
|
})
|
||||||
|
|
||||||
// 格式转换分界点索引(首个 hasConversion=true 的 group index)
|
// 格式转换分界点索引(首个 hasConversion=true 的 group index)
|
||||||
|
|||||||
@@ -130,6 +130,12 @@ class FailoverRulesConfig(BaseModel):
|
|||||||
class PoolAdvancedConfig(BaseModel):
|
class PoolAdvancedConfig(BaseModel):
|
||||||
"""通用号池配置(适用于所有 Provider 类型)。"""
|
"""通用号池配置(适用于所有 Provider 类型)。"""
|
||||||
|
|
||||||
|
global_priority: int | None = Field(
|
||||||
|
None,
|
||||||
|
ge=0,
|
||||||
|
le=999999,
|
||||||
|
description="global_key 模式下号池整体优先级(数字越小越优先)",
|
||||||
|
)
|
||||||
sticky_session_ttl_seconds: int | None = Field(
|
sticky_session_ttl_seconds: int | None = Field(
|
||||||
None,
|
None,
|
||||||
ge=60,
|
ge=60,
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from src.core.logger import logger
|
|||||||
from src.models.database import RequestCandidate
|
from src.models.database import RequestCandidate
|
||||||
from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
|
from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
|
||||||
from src.services.request.candidate import RequestCandidateService
|
from src.services.request.candidate import RequestCandidateService
|
||||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
|
||||||
from src.services.task.exceptions import StreamProbeError
|
from src.services.task.exceptions import StreamProbeError
|
||||||
from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult
|
from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult
|
||||||
from src.services.task.schema import ExecutionResult
|
from src.services.task.schema import ExecutionResult
|
||||||
@@ -141,6 +141,26 @@ class FailoverEngine:
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
pool_result, attempt_count, last_status_code = await self._execute_pool_candidate(
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
attempt_func=attempt_func,
|
||||||
|
retry_policy=retry_policy,
|
||||||
|
request_id=request_id,
|
||||||
|
user_id=user_id,
|
||||||
|
api_key_id=api_key_id,
|
||||||
|
candidate_record_map=candidate_record_map,
|
||||||
|
candidate_keys_fallback=candidate_keys_fallback,
|
||||||
|
candidates=candidates,
|
||||||
|
attempt_count=attempt_count,
|
||||||
|
max_attempts=max_attempts,
|
||||||
|
execution_error_handler=execution_error_handler,
|
||||||
|
)
|
||||||
|
if pool_result is not None:
|
||||||
|
return pool_result
|
||||||
|
continue
|
||||||
|
|
||||||
max_retries = self._get_max_retries(candidate, retry_policy)
|
max_retries = self._get_max_retries(candidate, retry_policy)
|
||||||
retry_index = 0
|
retry_index = 0
|
||||||
while retry_index < max_retries:
|
while retry_index < max_retries:
|
||||||
@@ -163,16 +183,9 @@ class FailoverEngine:
|
|||||||
api_key_id=api_key_id,
|
api_key_id=api_key_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Attach per-attempt context onto candidate for attempt_func (keeps AttemptFunc signature stable).
|
self._attach_attempt_context(
|
||||||
try:
|
candidate, candidate_index, retry_index, record_id, attempt_count, max_attempts
|
||||||
setattr(candidate, "_utf_candidate_index", candidate_index)
|
)
|
||||||
setattr(candidate, "_utf_retry_index", retry_index)
|
|
||||||
setattr(candidate, "_utf_candidate_record_id", record_id)
|
|
||||||
setattr(candidate, "_utf_attempt_count", attempt_count)
|
|
||||||
setattr(candidate, "_utf_max_attempts", max_attempts)
|
|
||||||
except Exception:
|
|
||||||
# Best-effort only; attempt_func may not rely on these attributes.
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Mark pending
|
# Mark pending
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
@@ -187,42 +200,12 @@ class FailoverEngine:
|
|||||||
self._commit_before_await()
|
self._commit_before_await()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
attempt_result = await attempt_func(candidate)
|
attempt_result = await self._execute_attempt(
|
||||||
last_status_code = int(getattr(attempt_result, "http_status", 0) or 0)
|
|
||||||
|
|
||||||
# Stream: probe first chunk, failover only before first chunk
|
|
||||||
if attempt_result.kind == AttemptKind.STREAM:
|
|
||||||
attempt_result = await self._probe_stream_first_chunk(
|
|
||||||
attempt_result=attempt_result,
|
|
||||||
record_id=record_id,
|
|
||||||
candidate=candidate,
|
candidate=candidate,
|
||||||
|
record_id=record_id,
|
||||||
|
attempt_func=attempt_func,
|
||||||
)
|
)
|
||||||
|
last_status_code = int(getattr(attempt_result, "http_status", 0) or 0)
|
||||||
# Sync: check success_failover_patterns on response body
|
|
||||||
if attempt_result.kind == AttemptKind.SYNC_RESPONSE:
|
|
||||||
body = getattr(attempt_result, "response_body", None)
|
|
||||||
if body:
|
|
||||||
if isinstance(body, bytes):
|
|
||||||
body_text = body.decode("utf-8", errors="replace")
|
|
||||||
elif isinstance(body, (dict, list)):
|
|
||||||
body_text = json.dumps(body, ensure_ascii=False)
|
|
||||||
else:
|
|
||||||
body_text = str(body)
|
|
||||||
rule_action = self._check_provider_failover_rules(
|
|
||||||
candidate, is_success=True, response_text=body_text
|
|
||||||
)
|
|
||||||
if rule_action == FailoverAction.CONTINUE:
|
|
||||||
self._record_attempt_failure(
|
|
||||||
record_id,
|
|
||||||
Exception("success_failover_pattern matched"),
|
|
||||||
200,
|
|
||||||
)
|
|
||||||
raise StreamProbeError(
|
|
||||||
"Success failover pattern matched",
|
|
||||||
http_status=200,
|
|
||||||
)
|
|
||||||
|
|
||||||
self._record_attempt_success(record_id, attempt_result)
|
|
||||||
|
|
||||||
# PRE_EXPAND: mark unused slots after request ends (success)
|
# PRE_EXPAND: mark unused slots after request ends (success)
|
||||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||||
@@ -316,6 +299,307 @@ class FailoverEngine:
|
|||||||
attempt_count=attempt_count,
|
attempt_count=attempt_count,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def _execute_pool_candidate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
candidate: PoolCandidate,
|
||||||
|
candidate_index: int,
|
||||||
|
attempt_func: AttemptFunc,
|
||||||
|
retry_policy: RetryPolicy,
|
||||||
|
request_id: str | None,
|
||||||
|
user_id: str | None,
|
||||||
|
api_key_id: str | None,
|
||||||
|
candidate_record_map: dict[tuple[int, int], str] | None,
|
||||||
|
candidate_keys_fallback: list[CandidateKey],
|
||||||
|
candidates: list[ProviderCandidate],
|
||||||
|
attempt_count: int,
|
||||||
|
max_attempts: int | None,
|
||||||
|
execution_error_handler: Any,
|
||||||
|
) -> tuple[ExecutionResult | None, int, int | None]:
|
||||||
|
"""Execute a PoolCandidate with in-pool key failover."""
|
||||||
|
last_status_code: int | None = None
|
||||||
|
retry_slots_per_key = self._get_pool_key_max_retries(candidate, retry_policy)
|
||||||
|
|
||||||
|
for key_index, pool_key in enumerate(candidate.pool_keys or []):
|
||||||
|
base_retry_index = key_index * retry_slots_per_key
|
||||||
|
candidate.key = pool_key
|
||||||
|
candidate._pool_key_index = key_index
|
||||||
|
candidate.mapping_matched_model = getattr(pool_key, "_pool_mapping_matched_model", None)
|
||||||
|
|
||||||
|
if bool(getattr(pool_key, "_pool_skipped", False)):
|
||||||
|
skip_reason = str(
|
||||||
|
getattr(pool_key, "_pool_skip_reason", None)
|
||||||
|
or getattr(candidate, "skip_reason", None)
|
||||||
|
or "pool_skipped"
|
||||||
|
)
|
||||||
|
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||||
|
self._mark_retry_indices_status(
|
||||||
|
candidate_record_map=candidate_record_map,
|
||||||
|
candidate_idx=candidate_index,
|
||||||
|
retry_indices=range(
|
||||||
|
base_retry_index, base_retry_index + retry_slots_per_key
|
||||||
|
),
|
||||||
|
status="skipped",
|
||||||
|
skip_reason=skip_reason,
|
||||||
|
)
|
||||||
|
elif request_id:
|
||||||
|
await self._create_skipped_record(
|
||||||
|
request_id=request_id,
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
retry_index=base_retry_index,
|
||||||
|
user_id=user_id,
|
||||||
|
api_key_id=api_key_id,
|
||||||
|
skip_reason=skip_reason,
|
||||||
|
)
|
||||||
|
candidate_keys_fallback.append(
|
||||||
|
self._make_candidate_key(
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
retry_index=base_retry_index,
|
||||||
|
status="skipped",
|
||||||
|
skip_reason=skip_reason,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
max_retries_for_key = retry_slots_per_key
|
||||||
|
retry_index = 0
|
||||||
|
while retry_index < max_retries_for_key:
|
||||||
|
attempt_count += 1
|
||||||
|
composite_retry_index = base_retry_index + retry_index
|
||||||
|
|
||||||
|
record_id = None
|
||||||
|
if candidate_record_map:
|
||||||
|
record_id = candidate_record_map.get((candidate_index, composite_retry_index))
|
||||||
|
if record_id is None:
|
||||||
|
# Rectify may extend retries beyond pre-created range.
|
||||||
|
record_id = candidate_record_map.get((candidate_index, base_retry_index))
|
||||||
|
if record_id is None and request_id and retry_policy.mode != RetryMode.PRE_EXPAND:
|
||||||
|
record_id = await self._ensure_record_exists(
|
||||||
|
request_id=request_id,
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
retry_index=composite_retry_index,
|
||||||
|
user_id=user_id,
|
||||||
|
api_key_id=api_key_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._attach_attempt_context(
|
||||||
|
candidate,
|
||||||
|
candidate_index,
|
||||||
|
composite_retry_index,
|
||||||
|
record_id,
|
||||||
|
attempt_count,
|
||||||
|
max_attempts,
|
||||||
|
)
|
||||||
|
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
if record_id:
|
||||||
|
self._update_record(
|
||||||
|
record_id,
|
||||||
|
status="pending",
|
||||||
|
started_at=now,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._commit_before_await()
|
||||||
|
|
||||||
|
try:
|
||||||
|
attempt_result = await self._execute_attempt(
|
||||||
|
candidate=candidate,
|
||||||
|
record_id=record_id,
|
||||||
|
attempt_func=attempt_func,
|
||||||
|
)
|
||||||
|
last_status_code = int(getattr(attempt_result, "http_status", 0) or 0)
|
||||||
|
|
||||||
|
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||||
|
self._mark_remaining_slots_unused(
|
||||||
|
candidate_record_map=candidate_record_map,
|
||||||
|
candidates=candidates,
|
||||||
|
success_candidate_idx=candidate_index,
|
||||||
|
success_retry_idx=composite_retry_index,
|
||||||
|
retry_policy=retry_policy,
|
||||||
|
)
|
||||||
|
|
||||||
|
return (
|
||||||
|
ExecutionResult(
|
||||||
|
success=True,
|
||||||
|
attempt_result=attempt_result,
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
retry_index=composite_retry_index,
|
||||||
|
provider_id=str(candidate.provider.id),
|
||||||
|
provider_name=str(candidate.provider.name),
|
||||||
|
endpoint_id=str(candidate.endpoint.id),
|
||||||
|
key_id=str(candidate.key.id),
|
||||||
|
candidate_keys=self._get_candidate_keys(
|
||||||
|
request_id=request_id,
|
||||||
|
fallback=candidate_keys_fallback,
|
||||||
|
candidates=candidates,
|
||||||
|
),
|
||||||
|
attempt_count=attempt_count,
|
||||||
|
request_candidate_id=record_id,
|
||||||
|
),
|
||||||
|
attempt_count,
|
||||||
|
last_status_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
except StreamProbeError as exc:
|
||||||
|
last_status_code = exc.http_status
|
||||||
|
self._record_attempt_failure(record_id, exc, exc.http_status)
|
||||||
|
action = FailoverAction.CONTINUE
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
outcome = await self._handle_pool_attempt_error(
|
||||||
|
exc,
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
key_retry_index=retry_index,
|
||||||
|
composite_retry_index=composite_retry_index,
|
||||||
|
max_retries=max_retries_for_key,
|
||||||
|
record_id=record_id,
|
||||||
|
attempt_count=attempt_count,
|
||||||
|
max_attempts=max_attempts,
|
||||||
|
execution_error_handler=execution_error_handler,
|
||||||
|
)
|
||||||
|
action = outcome.action
|
||||||
|
last_status_code = outcome.last_status_code
|
||||||
|
max_retries_for_key = min(outcome.max_retries, retry_slots_per_key)
|
||||||
|
|
||||||
|
if action == FailoverAction.CONTINUE:
|
||||||
|
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||||
|
# max_retries_for_key may have been shrunk by error handler;
|
||||||
|
# mark unused up to the *original* retry_slots_per_key to cover
|
||||||
|
# all pre-created records.
|
||||||
|
self._mark_retry_indices_status(
|
||||||
|
candidate_record_map=candidate_record_map,
|
||||||
|
candidate_idx=candidate_index,
|
||||||
|
retry_indices=range(
|
||||||
|
composite_retry_index + 1,
|
||||||
|
base_retry_index + retry_slots_per_key,
|
||||||
|
),
|
||||||
|
status="unused",
|
||||||
|
)
|
||||||
|
break
|
||||||
|
if action == FailoverAction.RETRY:
|
||||||
|
retry_index += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
# STOP: only stop this pool candidate; outer candidate traversal continues.
|
||||||
|
# Rationale: pool-internal STOP (from error_stop_patterns on a non-ExecutionError)
|
||||||
|
# should not terminate the entire request because other providers may still succeed.
|
||||||
|
# When handler_used=True, TaskService raises directly for true STOP semantics.
|
||||||
|
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||||
|
self._mark_candidate_remaining_retries_unused(
|
||||||
|
candidate_record_map=candidate_record_map,
|
||||||
|
candidate_idx=candidate_index,
|
||||||
|
from_retry_idx=composite_retry_index + 1,
|
||||||
|
retry_policy=retry_policy,
|
||||||
|
)
|
||||||
|
return None, attempt_count, last_status_code
|
||||||
|
|
||||||
|
return None, attempt_count, last_status_code
|
||||||
|
|
||||||
|
async def _handle_pool_attempt_error(
|
||||||
|
self,
|
||||||
|
exc: Exception,
|
||||||
|
*,
|
||||||
|
candidate: ProviderCandidate,
|
||||||
|
candidate_index: int,
|
||||||
|
key_retry_index: int,
|
||||||
|
composite_retry_index: int,
|
||||||
|
max_retries: int,
|
||||||
|
record_id: str | None,
|
||||||
|
attempt_count: int,
|
||||||
|
max_attempts: int | None,
|
||||||
|
execution_error_handler: Any,
|
||||||
|
) -> AttemptErrorOutcome:
|
||||||
|
"""Handle pool attempt errors without forcing outer STOP semantics.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key_retry_index: key 内部的重试索引 (用于判断 has_retry_left)
|
||||||
|
composite_retry_index: 全局维度的重试索引 (传给 execution_error_handler,
|
||||||
|
与 candidate_record_map 对齐)
|
||||||
|
"""
|
||||||
|
return await self._classify_attempt_error(
|
||||||
|
exc,
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
retry_index=composite_retry_index,
|
||||||
|
has_retry_left=key_retry_index + 1 < max_retries,
|
||||||
|
max_retries=max_retries,
|
||||||
|
record_id=record_id,
|
||||||
|
attempt_count=attempt_count,
|
||||||
|
max_attempts=max_attempts,
|
||||||
|
execution_error_handler=execution_error_handler,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _attach_attempt_context(
|
||||||
|
candidate: ProviderCandidate,
|
||||||
|
candidate_index: int,
|
||||||
|
retry_index: int,
|
||||||
|
record_id: str | None,
|
||||||
|
attempt_count: int,
|
||||||
|
max_attempts: int | None,
|
||||||
|
) -> None:
|
||||||
|
"""Attach per-attempt context onto candidate for attempt_func (best-effort)."""
|
||||||
|
try:
|
||||||
|
setattr(candidate, "_utf_candidate_index", candidate_index)
|
||||||
|
setattr(candidate, "_utf_retry_index", retry_index)
|
||||||
|
setattr(candidate, "_utf_candidate_record_id", record_id)
|
||||||
|
setattr(candidate, "_utf_attempt_count", attempt_count)
|
||||||
|
setattr(candidate, "_utf_max_attempts", max_attempts)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def _execute_attempt(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
candidate: ProviderCandidate,
|
||||||
|
record_id: str | None,
|
||||||
|
attempt_func: AttemptFunc,
|
||||||
|
) -> AttemptResult:
|
||||||
|
"""Run attempt_func with stream probe and sync failover-pattern checks.
|
||||||
|
|
||||||
|
On success records the attempt; raises StreamProbeError on failover-pattern
|
||||||
|
match or stream probe failure so the caller can handle retries uniformly.
|
||||||
|
"""
|
||||||
|
attempt_result = await attempt_func(candidate)
|
||||||
|
|
||||||
|
if attempt_result.kind == AttemptKind.STREAM:
|
||||||
|
attempt_result = await self._probe_stream_first_chunk(
|
||||||
|
attempt_result=attempt_result,
|
||||||
|
record_id=record_id,
|
||||||
|
candidate=candidate,
|
||||||
|
)
|
||||||
|
|
||||||
|
if attempt_result.kind == AttemptKind.SYNC_RESPONSE:
|
||||||
|
body = getattr(attempt_result, "response_body", None)
|
||||||
|
if body:
|
||||||
|
if isinstance(body, bytes):
|
||||||
|
body_text = body.decode("utf-8", errors="replace")
|
||||||
|
elif isinstance(body, (dict, list)):
|
||||||
|
body_text = json.dumps(body, ensure_ascii=False)
|
||||||
|
else:
|
||||||
|
body_text = str(body)
|
||||||
|
rule_action = self._check_provider_failover_rules(
|
||||||
|
candidate, is_success=True, response_text=body_text
|
||||||
|
)
|
||||||
|
if rule_action == FailoverAction.CONTINUE:
|
||||||
|
self._record_attempt_failure(
|
||||||
|
record_id,
|
||||||
|
Exception("success_failover_pattern matched"),
|
||||||
|
200,
|
||||||
|
)
|
||||||
|
raise StreamProbeError(
|
||||||
|
"Success failover pattern matched",
|
||||||
|
http_status=200,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._record_attempt_success(record_id, attempt_result)
|
||||||
|
return attempt_result
|
||||||
|
|
||||||
def _record_attempt_success(self, record_id: str | None, attempt_result: AttemptResult) -> None:
|
def _record_attempt_success(self, record_id: str | None, attempt_result: AttemptResult) -> None:
|
||||||
"""Mark attempt record as success/streaming."""
|
"""Mark attempt record as success/streaming."""
|
||||||
if not record_id:
|
if not record_id:
|
||||||
@@ -375,9 +659,62 @@ class FailoverEngine:
|
|||||||
Returns:
|
Returns:
|
||||||
AttemptErrorOutcome; stop_result is non-None only when action==STOP.
|
AttemptErrorOutcome; stop_result is non-None only when action==STOP.
|
||||||
"""
|
"""
|
||||||
has_retry_left = retry_index + 1 < max_retries
|
outcome = await self._classify_attempt_error(
|
||||||
|
exc,
|
||||||
|
candidate=candidate,
|
||||||
|
candidate_index=candidate_index,
|
||||||
|
retry_index=retry_index,
|
||||||
|
has_retry_left=retry_index + 1 < max_retries,
|
||||||
|
max_retries=max_retries,
|
||||||
|
record_id=record_id,
|
||||||
|
attempt_count=attempt_count,
|
||||||
|
max_attempts=max_attempts,
|
||||||
|
execution_error_handler=execution_error_handler,
|
||||||
|
)
|
||||||
|
|
||||||
# If caller provides an execution_error_handler, prefer it for ExecutionError.
|
if outcome.action == FailoverAction.STOP:
|
||||||
|
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||||
|
self._mark_remaining_slots_unused(
|
||||||
|
candidate_record_map=candidate_record_map,
|
||||||
|
candidates=candidates,
|
||||||
|
success_candidate_idx=candidate_index,
|
||||||
|
success_retry_idx=retry_index,
|
||||||
|
retry_policy=retry_policy,
|
||||||
|
)
|
||||||
|
outcome.stop_result = ExecutionResult(
|
||||||
|
success=False,
|
||||||
|
error_type=type(exc).__name__,
|
||||||
|
error_message=self._sanitize(str(exc)),
|
||||||
|
last_status_code=outcome.last_status_code or None,
|
||||||
|
candidate_keys=self._get_candidate_keys(
|
||||||
|
request_id=request_id,
|
||||||
|
fallback=candidate_keys_fallback,
|
||||||
|
candidates=candidates,
|
||||||
|
),
|
||||||
|
attempt_count=attempt_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
return outcome
|
||||||
|
|
||||||
|
async def _classify_attempt_error(
|
||||||
|
self,
|
||||||
|
exc: Exception,
|
||||||
|
*,
|
||||||
|
candidate: ProviderCandidate,
|
||||||
|
candidate_index: int,
|
||||||
|
retry_index: int,
|
||||||
|
has_retry_left: bool,
|
||||||
|
max_retries: int,
|
||||||
|
record_id: str | None,
|
||||||
|
attempt_count: int,
|
||||||
|
max_attempts: int | None,
|
||||||
|
execution_error_handler: Any,
|
||||||
|
) -> AttemptErrorOutcome:
|
||||||
|
"""Classify an attempt error: delegate to external handler or internal classifier.
|
||||||
|
|
||||||
|
Returns a base AttemptErrorOutcome (without stop_result). Callers add
|
||||||
|
STOP-specific logic (e.g. PRE_EXPAND cleanup, stop_result construction) as needed.
|
||||||
|
"""
|
||||||
handler_used = False
|
handler_used = False
|
||||||
action = FailoverAction.CONTINUE
|
action = FailoverAction.CONTINUE
|
||||||
if execution_error_handler is not None:
|
if execution_error_handler is not None:
|
||||||
@@ -408,40 +745,11 @@ class FailoverEngine:
|
|||||||
candidate=candidate,
|
candidate=candidate,
|
||||||
has_retry_left=has_retry_left,
|
has_retry_left=has_retry_left,
|
||||||
)
|
)
|
||||||
|
|
||||||
last_status_code = int(getattr(exc, "status_code", 0) or 0) or int(
|
last_status_code = int(getattr(exc, "status_code", 0) or 0) or int(
|
||||||
getattr(exc, "http_status", 0) or 0
|
getattr(exc, "http_status", 0) or 0
|
||||||
)
|
)
|
||||||
|
|
||||||
self._record_attempt_failure(record_id, exc, last_status_code or None)
|
self._record_attempt_failure(record_id, exc, last_status_code or None)
|
||||||
|
|
||||||
if action == FailoverAction.STOP:
|
|
||||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
|
||||||
self._mark_remaining_slots_unused(
|
|
||||||
candidate_record_map=candidate_record_map,
|
|
||||||
candidates=candidates,
|
|
||||||
success_candidate_idx=candidate_index,
|
|
||||||
success_retry_idx=retry_index,
|
|
||||||
retry_policy=retry_policy,
|
|
||||||
)
|
|
||||||
return AttemptErrorOutcome(
|
|
||||||
action=action,
|
|
||||||
last_status_code=last_status_code,
|
|
||||||
max_retries=max_retries,
|
|
||||||
stop_result=ExecutionResult(
|
|
||||||
success=False,
|
|
||||||
error_type=type(exc).__name__,
|
|
||||||
error_message=self._sanitize(str(exc)),
|
|
||||||
last_status_code=last_status_code or None,
|
|
||||||
candidate_keys=self._get_candidate_keys(
|
|
||||||
request_id=request_id,
|
|
||||||
fallback=candidate_keys_fallback,
|
|
||||||
candidates=candidates,
|
|
||||||
),
|
|
||||||
attempt_count=attempt_count,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
return AttemptErrorOutcome(
|
return AttemptErrorOutcome(
|
||||||
action=action,
|
action=action,
|
||||||
last_status_code=last_status_code,
|
last_status_code=last_status_code,
|
||||||
@@ -538,10 +846,7 @@ class FailoverEngine:
|
|||||||
api_key_id: str | None,
|
api_key_id: str | None,
|
||||||
) -> str:
|
) -> str:
|
||||||
# Create "available" record, then caller will mark pending.
|
# Create "available" record, then caller will mark pending.
|
||||||
extra: dict = {}
|
extra = self._build_pool_extra_data(candidate)
|
||||||
pool_extra = getattr(candidate, "_pool_extra_data", None)
|
|
||||||
if pool_extra:
|
|
||||||
extra.update(pool_extra)
|
|
||||||
row = RequestCandidateService.create_candidate(
|
row = RequestCandidateService.create_candidate(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
@@ -564,19 +869,17 @@ class FailoverEngine:
|
|||||||
request_id: str,
|
request_id: str,
|
||||||
candidate: ProviderCandidate,
|
candidate: ProviderCandidate,
|
||||||
candidate_index: int,
|
candidate_index: int,
|
||||||
|
retry_index: int = 0,
|
||||||
user_id: str | None,
|
user_id: str | None,
|
||||||
api_key_id: str | None,
|
api_key_id: str | None,
|
||||||
skip_reason: str | None,
|
skip_reason: str | None,
|
||||||
) -> str:
|
) -> str:
|
||||||
extra: dict = {}
|
extra = self._build_pool_extra_data(candidate)
|
||||||
pool_extra = getattr(candidate, "_pool_extra_data", None)
|
|
||||||
if pool_extra:
|
|
||||||
extra.update(pool_extra)
|
|
||||||
row = RequestCandidateService.create_candidate(
|
row = RequestCandidateService.create_candidate(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
candidate_index=candidate_index,
|
candidate_index=candidate_index,
|
||||||
retry_index=0,
|
retry_index=retry_index,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
api_key_id=api_key_id,
|
api_key_id=api_key_id,
|
||||||
provider_id=str(candidate.provider.id),
|
provider_id=str(candidate.provider.id),
|
||||||
@@ -592,6 +895,20 @@ class FailoverEngine:
|
|||||||
self.db.commit()
|
self.db.commit()
|
||||||
return str(row.id)
|
return str(row.id)
|
||||||
|
|
||||||
|
def _build_pool_extra_data(self, candidate: ProviderCandidate) -> dict[str, Any]:
|
||||||
|
extra: dict[str, Any] = {}
|
||||||
|
pool_extra = getattr(candidate, "_pool_extra_data", None)
|
||||||
|
if isinstance(pool_extra, dict):
|
||||||
|
extra.update(pool_extra)
|
||||||
|
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
extra["pool_group_id"] = str(getattr(candidate.provider, "id", "") or "")
|
||||||
|
extra["pool_key_index"] = int(getattr(candidate, "_pool_key_index", 0) or 0)
|
||||||
|
key_extra = getattr(candidate.key, "_pool_extra_data", None)
|
||||||
|
if isinstance(key_extra, dict):
|
||||||
|
extra.update(key_extra)
|
||||||
|
return extra
|
||||||
|
|
||||||
def _should_skip(
|
def _should_skip(
|
||||||
self, candidate: ProviderCandidate, skip_policy: SkipPolicy
|
self, candidate: ProviderCandidate, skip_policy: SkipPolicy
|
||||||
) -> tuple[bool, str | None]:
|
) -> tuple[bool, str | None]:
|
||||||
@@ -612,6 +929,15 @@ class FailoverEngine:
|
|||||||
return False, None
|
return False, None
|
||||||
|
|
||||||
def _get_max_retries(self, candidate: ProviderCandidate, retry_policy: RetryPolicy) -> int:
|
def _get_max_retries(self, candidate: ProviderCandidate, retry_policy: RetryPolicy) -> int:
|
||||||
|
per_key_retries = self._get_pool_key_max_retries(candidate, retry_policy)
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
key_count = len(candidate.pool_keys or []) or 1
|
||||||
|
return max(1, key_count * per_key_retries)
|
||||||
|
return per_key_retries
|
||||||
|
|
||||||
|
def _get_pool_key_max_retries(
|
||||||
|
self, candidate: ProviderCandidate, retry_policy: RetryPolicy
|
||||||
|
) -> int:
|
||||||
if retry_policy.mode == RetryMode.DISABLED:
|
if retry_policy.mode == RetryMode.DISABLED:
|
||||||
return 1
|
return 1
|
||||||
if retry_policy.retry_on_cached_only and not bool(getattr(candidate, "is_cached", False)):
|
if retry_policy.retry_on_cached_only and not bool(getattr(candidate, "is_cached", False)):
|
||||||
@@ -941,6 +1267,26 @@ class FailoverEngine:
|
|||||||
self._update_record(record_id, status="unused", finished_at=now)
|
self._update_record(record_id, status="unused", finished_at=now)
|
||||||
self.db.commit()
|
self.db.commit()
|
||||||
|
|
||||||
|
def _mark_retry_indices_status(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
candidate_record_map: dict[tuple[int, int], str],
|
||||||
|
candidate_idx: int,
|
||||||
|
retry_indices: range,
|
||||||
|
status: str,
|
||||||
|
skip_reason: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
for retry_idx in retry_indices:
|
||||||
|
record_id = candidate_record_map.get((candidate_idx, retry_idx))
|
||||||
|
if not record_id:
|
||||||
|
continue
|
||||||
|
values: dict[str, Any] = {"status": status, "finished_at": now}
|
||||||
|
if status == "skipped":
|
||||||
|
values["skip_reason"] = skip_reason
|
||||||
|
self._update_record(record_id, **values)
|
||||||
|
self.db.commit()
|
||||||
|
|
||||||
def _mark_all_remaining_available_unused(
|
def _mark_all_remaining_available_unused(
|
||||||
self, candidate_record_map: dict[tuple[int, int], str]
|
self, candidate_record_map: dict[tuple[int, int], str]
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|||||||
@@ -14,7 +14,8 @@ from src.core.exceptions import ProviderNotAvailableException
|
|||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import ApiKey
|
from src.models.database import ApiKey
|
||||||
from src.services.provider.format import normalize_endpoint_signature
|
from src.services.provider.format import normalize_endpoint_signature
|
||||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||||
|
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
|
||||||
|
|
||||||
|
|
||||||
class CandidateResolver:
|
class CandidateResolver:
|
||||||
@@ -153,14 +154,27 @@ class CandidateResolver:
|
|||||||
if preferred_key_ids:
|
if preferred_key_ids:
|
||||||
preferred_set = {str(kid) for kid in preferred_key_ids if kid}
|
preferred_set = {str(kid) for kid in preferred_key_ids if kid}
|
||||||
if preferred_set:
|
if preferred_set:
|
||||||
preferred_candidates = [
|
|
||||||
c for c in all_candidates if c.key and str(c.key.id) in preferred_set
|
def _is_preferred_candidate(c: ProviderCandidate) -> bool:
|
||||||
]
|
if c.key and str(c.key.id) in preferred_set:
|
||||||
other_candidates = [
|
return True
|
||||||
c for c in all_candidates if not (c.key and str(c.key.id) in preferred_set)
|
if isinstance(c, PoolCandidate):
|
||||||
]
|
return any(str(pk.id) in preferred_set for pk in (c.pool_keys or []))
|
||||||
|
return False
|
||||||
|
|
||||||
|
preferred_candidates = [c for c in all_candidates if _is_preferred_candidate(c)]
|
||||||
|
other_candidates = [c for c in all_candidates if not _is_preferred_candidate(c)]
|
||||||
if preferred_candidates:
|
if preferred_candidates:
|
||||||
matched_key_ids = [str(c.key.id) for c in preferred_candidates if c.key]
|
matched_key_ids: list[str] = []
|
||||||
|
for candidate in preferred_candidates:
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
matched_key_ids.extend(
|
||||||
|
str(pk.id)
|
||||||
|
for pk in (candidate.pool_keys or [])
|
||||||
|
if str(pk.id) in preferred_set
|
||||||
|
)
|
||||||
|
elif candidate.key:
|
||||||
|
matched_key_ids.append(str(candidate.key.id))
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f" [{request_id}] 优先候选命中: {len(preferred_candidates)} 个 "
|
f" [{request_id}] 优先候选命中: {len(preferred_candidates)} 个 "
|
||||||
f"(key_ids={matched_key_ids[:3]}{'...' if len(matched_key_ids) > 3 else ''})"
|
f"(key_ids={matched_key_ids[:3]}{'...' if len(matched_key_ids) > 3 else ''})"
|
||||||
@@ -210,6 +224,11 @@ class CandidateResolver:
|
|||||||
if not active_capabilities:
|
if not active_capabilities:
|
||||||
active_capabilities = None
|
active_capabilities = None
|
||||||
|
|
||||||
|
def _retry_slots_for_candidate(candidate: ProviderCandidate) -> int:
|
||||||
|
if not expand_retries:
|
||||||
|
return 1
|
||||||
|
return int(candidate.provider.max_retries or 2) if candidate.is_cached else 1
|
||||||
|
|
||||||
for candidate_index, candidate in enumerate(all_candidates):
|
for candidate_index, candidate in enumerate(all_candidates):
|
||||||
provider = candidate.provider
|
provider = candidate.provider
|
||||||
endpoint = candidate.endpoint
|
endpoint = candidate.endpoint
|
||||||
@@ -219,6 +238,71 @@ class CandidateResolver:
|
|||||||
if isinstance(getattr(candidate, "_pool_extra_data", None), dict)
|
if isinstance(getattr(candidate, "_pool_extra_data", None), dict)
|
||||||
else {}
|
else {}
|
||||||
)
|
)
|
||||||
|
base_extra = {
|
||||||
|
"needs_conversion": candidate.needs_conversion,
|
||||||
|
"provider_api_format": candidate.provider_api_format or None,
|
||||||
|
"mapping_matched_model": candidate.mapping_matched_model or None,
|
||||||
|
**pool_extra,
|
||||||
|
}
|
||||||
|
|
||||||
|
if isinstance(candidate, PoolCandidate) and candidate.pool_keys:
|
||||||
|
retry_slots = _retry_slots_for_candidate(candidate)
|
||||||
|
for key_idx, pool_key in enumerate(candidate.pool_keys):
|
||||||
|
key_id = str(pool_key.id)
|
||||||
|
key_pool_extra = (
|
||||||
|
getattr(pool_key, "_pool_extra_data", None)
|
||||||
|
if isinstance(getattr(pool_key, "_pool_extra_data", None), dict)
|
||||||
|
else {}
|
||||||
|
)
|
||||||
|
key_skipped = candidate.is_skipped or bool(
|
||||||
|
getattr(pool_key, "_pool_skipped", False)
|
||||||
|
)
|
||||||
|
key_skip_reason_raw = (
|
||||||
|
getattr(pool_key, "_pool_skip_reason", None) if key_skipped else None
|
||||||
|
)
|
||||||
|
key_skip_reason = (
|
||||||
|
str(key_skip_reason_raw)
|
||||||
|
if key_skip_reason_raw
|
||||||
|
else (candidate.skip_reason if key_skipped else None)
|
||||||
|
)
|
||||||
|
mapping_model = getattr(pool_key, "_pool_mapping_matched_model", None)
|
||||||
|
extra_data = {
|
||||||
|
**base_extra,
|
||||||
|
"mapping_matched_model": (
|
||||||
|
mapping_model
|
||||||
|
if mapping_model
|
||||||
|
else base_extra.get("mapping_matched_model")
|
||||||
|
),
|
||||||
|
"pool_group_id": str(provider.id),
|
||||||
|
"pool_key_index": key_idx,
|
||||||
|
**key_pool_extra,
|
||||||
|
}
|
||||||
|
|
||||||
|
for retry_in_key in range(retry_slots):
|
||||||
|
retry_index = key_idx * retry_slots + retry_in_key
|
||||||
|
status = "skipped" if key_skipped else "available"
|
||||||
|
record_id = str(uuid.uuid4())
|
||||||
|
candidate_records_to_insert.append(
|
||||||
|
{
|
||||||
|
"id": record_id,
|
||||||
|
"request_id": request_id,
|
||||||
|
"candidate_index": candidate_index,
|
||||||
|
"retry_index": retry_index,
|
||||||
|
"user_id": user_id,
|
||||||
|
"api_key_id": user_api_key.id if user_api_key else None,
|
||||||
|
"provider_id": provider.id,
|
||||||
|
"endpoint_id": endpoint.id,
|
||||||
|
"key_id": key_id,
|
||||||
|
"status": status,
|
||||||
|
"skip_reason": key_skip_reason if key_skipped else None,
|
||||||
|
"is_cached": candidate.is_cached,
|
||||||
|
"extra_data": extra_data,
|
||||||
|
"required_capabilities": active_capabilities,
|
||||||
|
"created_at": datetime.now(timezone.utc),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
candidate_record_map[(candidate_index, retry_index)] = record_id
|
||||||
|
continue
|
||||||
|
|
||||||
if candidate.is_skipped:
|
if candidate.is_skipped:
|
||||||
record_id = str(uuid.uuid4())
|
record_id = str(uuid.uuid4())
|
||||||
@@ -236,25 +320,14 @@ class CandidateResolver:
|
|||||||
"status": "skipped",
|
"status": "skipped",
|
||||||
"skip_reason": candidate.skip_reason,
|
"skip_reason": candidate.skip_reason,
|
||||||
"is_cached": candidate.is_cached,
|
"is_cached": candidate.is_cached,
|
||||||
"extra_data": {
|
"extra_data": base_extra,
|
||||||
"needs_conversion": candidate.needs_conversion,
|
|
||||||
"provider_api_format": candidate.provider_api_format or None,
|
|
||||||
"mapping_matched_model": candidate.mapping_matched_model or None,
|
|
||||||
**pool_extra,
|
|
||||||
},
|
|
||||||
"required_capabilities": active_capabilities,
|
"required_capabilities": active_capabilities,
|
||||||
"created_at": datetime.now(timezone.utc),
|
"created_at": datetime.now(timezone.utc),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
candidate_record_map[(candidate_index, 0)] = record_id
|
candidate_record_map[(candidate_index, 0)] = record_id
|
||||||
else:
|
else:
|
||||||
# max_retries 已从 Endpoint 迁移到 Provider(Endpoint 仍可能保留旧字段用于兼容)
|
max_retries_for_candidate = _retry_slots_for_candidate(candidate)
|
||||||
if not expand_retries:
|
|
||||||
max_retries_for_candidate = 1
|
|
||||||
else:
|
|
||||||
max_retries_for_candidate = (
|
|
||||||
int(provider.max_retries or 2) if candidate.is_cached else 1
|
|
||||||
)
|
|
||||||
|
|
||||||
for retry_index in range(max_retries_for_candidate):
|
for retry_index in range(max_retries_for_candidate):
|
||||||
record_id = str(uuid.uuid4())
|
record_id = str(uuid.uuid4())
|
||||||
@@ -271,12 +344,7 @@ class CandidateResolver:
|
|||||||
"key_id": key.id,
|
"key_id": key.id,
|
||||||
"status": "available",
|
"status": "available",
|
||||||
"is_cached": candidate.is_cached,
|
"is_cached": candidate.is_cached,
|
||||||
"extra_data": {
|
"extra_data": base_extra,
|
||||||
"needs_conversion": candidate.needs_conversion,
|
|
||||||
"provider_api_format": candidate.provider_api_format or None,
|
|
||||||
"mapping_matched_model": candidate.mapping_matched_model or None,
|
|
||||||
**pool_extra,
|
|
||||||
},
|
|
||||||
"required_capabilities": active_capabilities,
|
"required_capabilities": active_capabilities,
|
||||||
"created_at": datetime.now(timezone.utc),
|
"created_at": datetime.now(timezone.utc),
|
||||||
}
|
}
|
||||||
@@ -326,7 +394,22 @@ class CandidateResolver:
|
|||||||
total = 0
|
total = 0
|
||||||
for candidate in all_candidates:
|
for candidate in all_candidates:
|
||||||
if not candidate.is_skipped:
|
if not candidate.is_skipped:
|
||||||
provider = candidate.provider
|
retries_per_slot = (
|
||||||
max_retries = int(provider.max_retries or 2) if candidate.is_cached else 1
|
int(candidate.provider.max_retries or 2) if candidate.is_cached else 1
|
||||||
total += max_retries
|
)
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
schedulable_keys = [
|
||||||
|
k
|
||||||
|
for k in (candidate.pool_keys or [])
|
||||||
|
if not bool(getattr(k, "_pool_skipped", False))
|
||||||
|
]
|
||||||
|
if schedulable_keys:
|
||||||
|
total += len(schedulable_keys) * retries_per_slot
|
||||||
|
elif candidate.pool_keys:
|
||||||
|
# 兜底:尚未附加 _pool_skipped 标记时,按 key 数估算。
|
||||||
|
total += len(candidate.pool_keys) * retries_per_slot
|
||||||
|
else:
|
||||||
|
total += retries_per_slot
|
||||||
|
else:
|
||||||
|
total += retries_per_slot
|
||||||
return total
|
return total
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ class PoolConfig:
|
|||||||
|
|
||||||
# -- Sticky Session -------------------------------------------------------
|
# -- Sticky Session -------------------------------------------------------
|
||||||
sticky_session_ttl_seconds: int = 3600 # 1 hour
|
sticky_session_ttl_seconds: int = 3600 # 1 hour
|
||||||
|
# Key 优先模式下号池整体优先级(None 时回退 provider_priority)
|
||||||
|
global_priority: int | None = None
|
||||||
|
|
||||||
# -- Load-Aware Selection -------------------------------------------------
|
# -- Load-Aware Selection -------------------------------------------------
|
||||||
load_threshold_percent: int = 80
|
load_threshold_percent: int = 80
|
||||||
@@ -122,6 +124,7 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
|
|||||||
|
|
||||||
return PoolConfig(
|
return PoolConfig(
|
||||||
sticky_session_ttl_seconds=_int_or("sticky_session_ttl_seconds", 3600),
|
sticky_session_ttl_seconds=_int_or("sticky_session_ttl_seconds", 3600),
|
||||||
|
global_priority=_opt_int("global_priority"),
|
||||||
load_threshold_percent=_int_or("load_threshold_percent", 80),
|
load_threshold_percent=_int_or("load_threshold_percent", 80),
|
||||||
lru_enabled=_bool_or("lru_enabled", True),
|
lru_enabled=_bool_or("lru_enabled", True),
|
||||||
cost_window_seconds=_int_or("cost_window_seconds", 18000),
|
cost_window_seconds=_int_or("cost_window_seconds", 18000),
|
||||||
|
|||||||
@@ -280,6 +280,66 @@ class PoolManager:
|
|||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
async def select_pool_keys(
|
||||||
|
self,
|
||||||
|
session_uuid: str | None,
|
||||||
|
keys: list[ProviderAPIKey],
|
||||||
|
) -> tuple[list[ProviderAPIKey], PoolSchedulingTrace]:
|
||||||
|
"""Select and order pool keys with trace output.
|
||||||
|
|
||||||
|
Reuses :meth:`reorder_candidates` logic by adapting keys to lightweight
|
||||||
|
candidate-like wrappers, then propagates skip/trace metadata back onto
|
||||||
|
each key object for downstream execution/recording.
|
||||||
|
"""
|
||||||
|
if not keys:
|
||||||
|
return (
|
||||||
|
[],
|
||||||
|
PoolSchedulingTrace(
|
||||||
|
provider_id=self.provider_id,
|
||||||
|
total_keys=0,
|
||||||
|
session_uuid=session_uuid[:8] if session_uuid else None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
class _KeyCandidate:
|
||||||
|
__slots__ = ("key", "is_skipped", "skip_reason")
|
||||||
|
|
||||||
|
def __init__(self, key: ProviderAPIKey) -> None:
|
||||||
|
self.key = key
|
||||||
|
self.is_skipped = False
|
||||||
|
self.skip_reason: str | None = None
|
||||||
|
|
||||||
|
wrappers = [_KeyCandidate(k) for k in keys]
|
||||||
|
reordered_wrappers = await self.reorder_candidates(session_uuid, wrappers) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
trace: PoolSchedulingTrace | None = None
|
||||||
|
if reordered_wrappers:
|
||||||
|
maybe_trace = getattr(reordered_wrappers[0], "_pool_scheduling_trace", None)
|
||||||
|
if isinstance(maybe_trace, PoolSchedulingTrace):
|
||||||
|
trace = maybe_trace
|
||||||
|
if trace is None:
|
||||||
|
trace = PoolSchedulingTrace(
|
||||||
|
provider_id=self.provider_id,
|
||||||
|
total_keys=len(keys),
|
||||||
|
session_uuid=session_uuid[:8] if session_uuid else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
ordered_keys: list[ProviderAPIKey] = []
|
||||||
|
for order_idx, wrapped in enumerate(reordered_wrappers):
|
||||||
|
key = wrapped.key
|
||||||
|
is_skipped = bool(getattr(wrapped, "is_skipped", False))
|
||||||
|
skip_reason = str(getattr(wrapped, "skip_reason", "") or "")
|
||||||
|
setattr(key, "_pool_skipped", is_skipped)
|
||||||
|
setattr(key, "_pool_skip_reason", skip_reason if skip_reason else None)
|
||||||
|
setattr(key, "_pool_order_index", order_idx)
|
||||||
|
pool_extra = getattr(wrapped, "_pool_extra_data", None)
|
||||||
|
setattr(
|
||||||
|
key, "_pool_extra_data", dict(pool_extra) if isinstance(pool_extra, dict) else {}
|
||||||
|
)
|
||||||
|
ordered_keys.append(key)
|
||||||
|
|
||||||
|
return ordered_keys, trace
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Single-key selection (used by CandidateBuilder for pooled providers)
|
# Single-key selection (used by CandidateBuilder for pooled providers)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|||||||
@@ -76,6 +76,7 @@ from src.services.scheduling.concurrency_checker import ConcurrencyChecker
|
|||||||
from src.services.scheduling.restriction_checker import get_effective_restrictions
|
from src.services.scheduling.restriction_checker import get_effective_restrictions
|
||||||
from src.services.scheduling.scheduling_config import SchedulingConfig
|
from src.services.scheduling.scheduling_config import SchedulingConfig
|
||||||
from src.services.scheduling.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
|
from src.services.scheduling.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
|
||||||
|
from src.services.scheduling.schemas import PoolCandidate as PoolCandidate # re-export
|
||||||
from src.services.scheduling.schemas import ProviderCandidate as ProviderCandidate # re-export
|
from src.services.scheduling.schemas import ProviderCandidate as ProviderCandidate # re-export
|
||||||
from src.services.scheduling.utils import affinity_hash as _affinity_hash # re-export compat
|
from src.services.scheduling.utils import affinity_hash as _affinity_hash # re-export compat
|
||||||
from src.services.scheduling.utils import (
|
from src.services.scheduling.utils import (
|
||||||
@@ -652,11 +653,20 @@ class CacheAwareScheduler:
|
|||||||
endpoint = candidate.endpoint
|
endpoint = candidate.endpoint
|
||||||
key = candidate.key
|
key = candidate.key
|
||||||
|
|
||||||
if (
|
is_pool_candidate = isinstance(candidate, PoolCandidate)
|
||||||
provider.id == affinity.provider_id
|
pool_matched = (
|
||||||
|
is_pool_candidate
|
||||||
|
and provider.id == affinity.provider_id
|
||||||
|
and endpoint.id == affinity.endpoint_id
|
||||||
|
)
|
||||||
|
key_matched = (
|
||||||
|
(not is_pool_candidate)
|
||||||
|
and provider.id == affinity.provider_id
|
||||||
and endpoint.id == affinity.endpoint_id
|
and endpoint.id == affinity.endpoint_id
|
||||||
and key.id == affinity.key_id
|
and key.id == affinity.key_id
|
||||||
):
|
)
|
||||||
|
|
||||||
|
if pool_matched or key_matched:
|
||||||
candidate.is_cached = True
|
candidate.is_cached = True
|
||||||
matched_candidate = candidate
|
matched_candidate = candidate
|
||||||
matched = True
|
matched = True
|
||||||
|
|||||||
@@ -398,7 +398,7 @@ class CandidateBuilder:
|
|||||||
Returns:
|
Returns:
|
||||||
候选列表
|
候选列表
|
||||||
"""
|
"""
|
||||||
from src.services.scheduling.schemas import ProviderCandidate
|
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
|
||||||
|
|
||||||
candidates: list[ProviderCandidate] = []
|
candidates: list[ProviderCandidate] = []
|
||||||
client_format_str = normalize_endpoint_signature(client_format)
|
client_format_str = normalize_endpoint_signature(client_format)
|
||||||
@@ -438,7 +438,6 @@ class CandidateBuilder:
|
|||||||
] = {}
|
] = {}
|
||||||
exact_candidates: list[ProviderCandidate] = []
|
exact_candidates: list[ProviderCandidate] = []
|
||||||
convertible_candidates: list[ProviderCandidate] = []
|
convertible_candidates: list[ProviderCandidate] = []
|
||||||
pool_has_usable = False
|
|
||||||
pool_cfg = _get_pool_config(provider)
|
pool_cfg = _get_pool_config(provider)
|
||||||
|
|
||||||
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
|
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
|
||||||
@@ -576,8 +575,6 @@ class CandidateBuilder:
|
|||||||
if not active_keys:
|
if not active_keys:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Pool provider should still expose all key candidates here.
|
|
||||||
# Runtime pool scheduling/failover is handled later by TaskService._apply_pool_reorder.
|
|
||||||
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
|
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
|
||||||
if pool_cfg is not None:
|
if pool_cfg is not None:
|
||||||
use_random = False
|
use_random = False
|
||||||
@@ -592,6 +589,84 @@ class CandidateBuilder:
|
|||||||
active_keys, affinity_key, use_random
|
active_keys, affinity_key, use_random
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if pool_cfg is not None:
|
||||||
|
# 号池 Provider 仅构建一个 PoolCandidate,内部 key 选择延迟到执行阶段。
|
||||||
|
pool_keys: list[ProviderAPIKey] = []
|
||||||
|
pool_miss_counts: list[int] = []
|
||||||
|
pool_mapping: dict[str, str | None] = {}
|
||||||
|
for key in keys_to_check:
|
||||||
|
is_available, _key_skip_reason, mapping_matched_model = (
|
||||||
|
self._check_key_availability(
|
||||||
|
key,
|
||||||
|
endpoint_format_str,
|
||||||
|
model_name,
|
||||||
|
capability_requirements,
|
||||||
|
model_mappings=model_mappings,
|
||||||
|
candidate_models=provider_model_names,
|
||||||
|
provider_type=getattr(provider, "provider_type", None),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if not is_available:
|
||||||
|
continue
|
||||||
|
|
||||||
|
pool_keys.append(key)
|
||||||
|
pool_miss_counts.append(
|
||||||
|
compute_capability_score(
|
||||||
|
key.capabilities or {},
|
||||||
|
capability_requirements,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
pool_mapping[str(key.id)] = mapping_matched_model
|
||||||
|
|
||||||
|
if not pool_keys:
|
||||||
|
continue
|
||||||
|
|
||||||
|
provider_priority_raw = getattr(provider, "provider_priority", None)
|
||||||
|
try:
|
||||||
|
provider_priority = (
|
||||||
|
int(provider_priority_raw)
|
||||||
|
if provider_priority_raw is not None
|
||||||
|
else 999999
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
provider_priority = 999999
|
||||||
|
try:
|
||||||
|
pool_priority = (
|
||||||
|
int(pool_cfg.global_priority)
|
||||||
|
if pool_cfg.global_priority is not None
|
||||||
|
else provider_priority
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pool_priority = provider_priority
|
||||||
|
|
||||||
|
pool_candidate = PoolCandidate(
|
||||||
|
provider=provider,
|
||||||
|
endpoint=endpoint,
|
||||||
|
key=pool_keys[0],
|
||||||
|
pool_keys=pool_keys,
|
||||||
|
pool_config=pool_cfg,
|
||||||
|
pool_priority=pool_priority,
|
||||||
|
mapping_matched_model=pool_mapping.get(str(pool_keys[0].id)),
|
||||||
|
needs_conversion=needs_conversion,
|
||||||
|
provider_api_format=str(endpoint_format_str or ""),
|
||||||
|
output_limit=output_limit,
|
||||||
|
capability_miss_count=min(pool_miss_counts) if pool_miss_counts else 0,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 在 key 对象上附加映射结果,供 PoolCandidate 运行时切 key 后同步模型名。
|
||||||
|
for pool_key in pool_keys:
|
||||||
|
setattr(
|
||||||
|
pool_key,
|
||||||
|
"_pool_mapping_matched_model",
|
||||||
|
pool_mapping.get(str(pool_key.id)),
|
||||||
|
)
|
||||||
|
|
||||||
|
if needs_conversion:
|
||||||
|
convertible_candidates.append(pool_candidate)
|
||||||
|
else:
|
||||||
|
exact_candidates.append(pool_candidate)
|
||||||
|
break
|
||||||
|
|
||||||
for key in keys_to_check:
|
for key in keys_to_check:
|
||||||
# Key 级别检查(健康度/熔断按 provider_format bucket)
|
# Key 级别检查(健康度/熔断按 provider_format bucket)
|
||||||
# 传入 provider_model_names 作为 candidate_models,
|
# 传入 provider_model_names 作为 candidate_models,
|
||||||
@@ -634,13 +709,6 @@ class CandidateBuilder:
|
|||||||
else:
|
else:
|
||||||
exact_candidates.append(candidate)
|
exact_candidates.append(candidate)
|
||||||
|
|
||||||
if is_available:
|
|
||||||
pool_has_usable = True
|
|
||||||
|
|
||||||
# Pool mode: stop after the first endpoint that produced a usable candidate.
|
|
||||||
if pool_cfg is not None and pool_has_usable:
|
|
||||||
break
|
|
||||||
|
|
||||||
candidates.extend(exact_candidates)
|
candidates.extend(exact_candidates)
|
||||||
candidates.extend(convertible_candidates)
|
candidates.extend(convertible_candidates)
|
||||||
|
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from collections import defaultdict
|
|||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from src.services.scheduling.scheduling_config import SchedulingConfig
|
from src.services.scheduling.scheduling_config import SchedulingConfig
|
||||||
|
from src.services.scheduling.schemas import PoolCandidate
|
||||||
from src.services.scheduling.utils import affinity_hash
|
from src.services.scheduling.utils import affinity_hash
|
||||||
from src.services.system.config import SystemConfigService
|
from src.services.system.config import SystemConfigService
|
||||||
|
|
||||||
@@ -149,6 +150,8 @@ class CandidateSorter:
|
|||||||
|
|
||||||
def get_priority(candidate: ProviderCandidate) -> int:
|
def get_priority(candidate: ProviderCandidate) -> int:
|
||||||
"""获取候选的优先级"""
|
"""获取候选的优先级"""
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
return int(getattr(candidate, "pool_priority", 999999) or 999999)
|
||||||
if not candidate.key:
|
if not candidate.key:
|
||||||
return 999999
|
return 999999
|
||||||
priority_by_format = candidate.key.global_priority_by_format or {}
|
priority_by_format = candidate.key.global_priority_by_format or {}
|
||||||
@@ -170,8 +173,11 @@ class CandidateSorter:
|
|||||||
# 同优先级内哈希分散负载均衡
|
# 同优先级内哈希分散负载均衡
|
||||||
scored_candidates = []
|
scored_candidates = []
|
||||||
for candidate in group:
|
for candidate in group:
|
||||||
key_id = candidate.key.id if candidate.key else ""
|
if isinstance(candidate, PoolCandidate):
|
||||||
hash_value = affinity_hash(affinity_key, key_id)
|
hash_id = str(getattr(candidate.provider, "id", "") or "")
|
||||||
|
else:
|
||||||
|
hash_id = candidate.key.id if candidate.key else ""
|
||||||
|
hash_value = affinity_hash(affinity_key, hash_id)
|
||||||
scored_candidates.append((hash_value, candidate))
|
scored_candidates.append((hash_value, candidate))
|
||||||
|
|
||||||
# 按哈希值排序
|
# 按哈希值排序
|
||||||
@@ -181,11 +187,16 @@ class CandidateSorter:
|
|||||||
# 单个候选或没有 affinity_key,按次要排序条件排序
|
# 单个候选或没有 affinity_key,按次要排序条件排序
|
||||||
def secondary_sort(c: ProviderCandidate) -> tuple[int, int, str]:
|
def secondary_sort(c: ProviderCandidate) -> tuple[int, int, str]:
|
||||||
pp = c.provider.provider_priority
|
pp = c.provider.provider_priority
|
||||||
|
if isinstance(c, PoolCandidate):
|
||||||
|
ip = int(getattr(c, "pool_priority", 999999) or 999999)
|
||||||
|
key_id = str(getattr(c.provider, "id", "") or "")
|
||||||
|
else:
|
||||||
ip = c.key.internal_priority if c.key else None
|
ip = c.key.internal_priority if c.key else None
|
||||||
|
key_id = c.key.id if c.key else ""
|
||||||
return (
|
return (
|
||||||
pp if pp is not None else 999999,
|
pp if pp is not None else 999999,
|
||||||
ip if ip is not None else 999999,
|
ip if ip is not None else 999999,
|
||||||
c.key.id if c.key else "",
|
key_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
result.extend(sorted(group, key=secondary_sort))
|
result.extend(sorted(group, key=secondary_sort))
|
||||||
@@ -222,16 +233,26 @@ class CandidateSorter:
|
|||||||
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
|
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
|
||||||
# 全局 Key 优先模式:按格式特定优先级分组
|
# 全局 Key 优先模式:按格式特定优先级分组
|
||||||
for candidate in candidates:
|
for candidate in candidates:
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
priority = int(getattr(candidate, "pool_priority", 999999) or 999999)
|
||||||
|
# -1 使号池候选独立成组,不与普通 key 候选 (0) 混组打乱
|
||||||
|
priority_groups[(priority, -1)].append(candidate)
|
||||||
|
continue
|
||||||
|
else:
|
||||||
priority = 999999
|
priority = 999999
|
||||||
if candidate.key:
|
if candidate.key:
|
||||||
priority_by_format = candidate.key.global_priority_by_format or {}
|
priority_by_format = candidate.key.global_priority_by_format or {}
|
||||||
if api_format and api_format in priority_by_format:
|
if api_format and api_format in priority_by_format:
|
||||||
priority = priority_by_format[api_format]
|
priority = priority_by_format[api_format]
|
||||||
priority_groups[(priority,)].append(candidate)
|
priority_groups[(priority, 0)].append(candidate)
|
||||||
else:
|
else:
|
||||||
# 提供商优先模式:按 (provider_priority, internal_priority) 分组
|
# 提供商优先模式:按 (provider_priority, internal_priority) 分组
|
||||||
for candidate in candidates:
|
for candidate in candidates:
|
||||||
pp = candidate.provider.provider_priority
|
pp = candidate.provider.provider_priority
|
||||||
|
if isinstance(candidate, PoolCandidate):
|
||||||
|
# 号池候选独立成组,不与普通 key 候选混组打乱。
|
||||||
|
ip = -1
|
||||||
|
else:
|
||||||
ip = candidate.key.internal_priority if candidate.key else None
|
ip = candidate.key.internal_priority if candidate.key else None
|
||||||
key = (
|
key = (
|
||||||
pp if pp is not None else 999999,
|
pp if pp is not None else 999999,
|
||||||
|
|||||||
@@ -6,7 +6,8 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, field
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from src.models.database import (
|
from src.models.database import (
|
||||||
Provider,
|
Provider,
|
||||||
@@ -14,6 +15,9 @@ from src.models.database import (
|
|||||||
ProviderEndpoint,
|
ProviderEndpoint,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from src.services.provider.pool.config import PoolConfig
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ProviderCandidate:
|
class ProviderCandidate:
|
||||||
@@ -69,6 +73,19 @@ class ProviderCandidate:
|
|||||||
return self._stable_order_key() < other._stable_order_key()
|
return self._stable_order_key() < other._stable_order_key()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PoolCandidate(ProviderCandidate):
|
||||||
|
"""号池候选。
|
||||||
|
|
||||||
|
排序阶段作为单个候选参与;执行阶段再在 pool_keys 内部选择/切换 key。
|
||||||
|
"""
|
||||||
|
|
||||||
|
pool_keys: list[ProviderAPIKey] = field(default_factory=list)
|
||||||
|
pool_config: PoolConfig | None = None
|
||||||
|
pool_priority: int = 999999
|
||||||
|
_pool_key_index: int = 0
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ConcurrencySnapshot:
|
class ConcurrencySnapshot:
|
||||||
key_current: int
|
key_current: int
|
||||||
|
|||||||
@@ -206,61 +206,130 @@ class TaskService:
|
|||||||
candidates: list[Any],
|
candidates: list[Any],
|
||||||
request_body: dict[str, Any] | None,
|
request_body: dict[str, Any] | None,
|
||||||
) -> tuple[list[Any], list[Any]]:
|
) -> tuple[list[Any], list[Any]]:
|
||||||
"""Apply Account Pool reordering when applicable.
|
"""Apply pool key ordering for PoolCandidate objects."""
|
||||||
|
|
||||||
Groups candidates by provider_id and applies pool reordering
|
|
||||||
independently per provider, then reassembles in original group order.
|
|
||||||
Non-pool providers are left in their original order.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (reordered_candidates, pool_traces) where pool_traces
|
|
||||||
is a list of :class:`PoolSchedulingTrace` objects (one per
|
|
||||||
pooled provider group, may be empty).
|
|
||||||
"""
|
|
||||||
if not candidates:
|
if not candidates:
|
||||||
return candidates, []
|
return candidates, []
|
||||||
|
|
||||||
pool_traces: list[Any] = []
|
pool_traces: list[Any] = []
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from collections import OrderedDict
|
|
||||||
|
|
||||||
from src.services.provider.pool.config import parse_pool_config
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
from src.services.provider.pool.manager import PoolManager
|
from src.services.provider.pool.manager import PoolManager
|
||||||
|
from src.services.scheduling.schemas import PoolCandidate
|
||||||
|
|
||||||
# Group candidates by provider_id while preserving order.
|
for candidate in candidates:
|
||||||
groups: OrderedDict[str, list[Any]] = OrderedDict()
|
if not isinstance(candidate, PoolCandidate):
|
||||||
for c in candidates:
|
|
||||||
pid = str(getattr(c.provider, "id", "") or "")
|
|
||||||
groups.setdefault(pid, []).append(c)
|
|
||||||
|
|
||||||
result: list[Any] = []
|
|
||||||
for pid, group in groups.items():
|
|
||||||
provider = group[0].provider
|
|
||||||
pool_cfg = parse_pool_config(getattr(provider, "config", None))
|
|
||||||
if pool_cfg is None or not pid:
|
|
||||||
result.extend(group)
|
|
||||||
continue
|
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 "")
|
provider_type = str(getattr(provider, "provider_type", "") or "")
|
||||||
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
|
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
|
||||||
mgr = PoolManager(pid, pool_cfg)
|
manager = PoolManager(provider_id, pool_cfg)
|
||||||
reordered = await mgr.reorder_candidates(session_uuid, group)
|
|
||||||
result.extend(reordered)
|
candidate_keys = list(candidate.pool_keys or [])
|
||||||
|
if not candidate_keys and getattr(candidate, "key", None) is not None:
|
||||||
|
candidate_keys = [candidate.key]
|
||||||
|
|
||||||
|
ordered_keys, trace = await manager.select_pool_keys(session_uuid, candidate_keys)
|
||||||
|
candidate.pool_keys = ordered_keys
|
||||||
|
|
||||||
|
selected_key_index = 0
|
||||||
|
selected_key = None
|
||||||
|
for idx, pool_key in enumerate(ordered_keys):
|
||||||
|
if not bool(getattr(pool_key, "_pool_skipped", False)):
|
||||||
|
selected_key = pool_key
|
||||||
|
selected_key_index = idx
|
||||||
|
break
|
||||||
|
|
||||||
|
if selected_key is not None:
|
||||||
|
candidate.key = selected_key
|
||||||
|
candidate._pool_key_index = selected_key_index
|
||||||
|
candidate.mapping_matched_model = getattr(
|
||||||
|
selected_key, "_pool_mapping_matched_model", None
|
||||||
|
)
|
||||||
|
candidate.is_skipped = False
|
||||||
|
candidate.skip_reason = None
|
||||||
|
else:
|
||||||
|
candidate.is_skipped = True
|
||||||
|
candidate.skip_reason = "pool: all keys unavailable"
|
||||||
|
|
||||||
# Extract trace attached by PoolManager.reorder_candidates
|
|
||||||
if reordered:
|
|
||||||
trace = getattr(reordered[0], "_pool_scheduling_trace", None)
|
|
||||||
if trace is not None:
|
if trace is not None:
|
||||||
pool_traces.append(trace)
|
pool_traces.append(trace)
|
||||||
|
|
||||||
return result, pool_traces
|
return candidates, pool_traces
|
||||||
except Exception:
|
except Exception:
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
|
||||||
logger.opt(exception=True).debug("Pool reorder failed, using original order")
|
logger.opt(exception=True).debug("Pool reorder failed, using original order")
|
||||||
return candidates, []
|
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
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _pool_on_success(
|
async def _pool_on_success(
|
||||||
candidate: Any,
|
candidate: Any,
|
||||||
@@ -466,6 +535,25 @@ class TaskService:
|
|||||||
|
|
||||||
# Safety net: if record_id missing, create an "available" record on-demand.
|
# Safety net: if record_id missing, create an "available" record on-demand.
|
||||||
if not candidate_record_id:
|
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(
|
created = RequestCandidateService.create_candidate(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
@@ -478,6 +566,7 @@ class TaskService:
|
|||||||
key_id=str(candidate.key.id),
|
key_id=str(candidate.key.id),
|
||||||
status="available",
|
status="available",
|
||||||
is_cached=bool(getattr(candidate, "is_cached", False)),
|
is_cached=bool(getattr(candidate, "is_cached", False)),
|
||||||
|
extra_data=extra_data,
|
||||||
)
|
)
|
||||||
candidate_record_id = str(created.id)
|
candidate_record_id = str(created.id)
|
||||||
candidate_record_map[(candidate_index, retry_index)] = candidate_record_id
|
candidate_record_map[(candidate_index, retry_index)] = candidate_record_id
|
||||||
@@ -1331,6 +1420,7 @@ class TaskService:
|
|||||||
candidates, _pool_traces = await self._apply_pool_reorder(
|
candidates, _pool_traces = await self._apply_pool_reorder(
|
||||||
candidates, request_body=request_body
|
candidates, request_body=request_body
|
||||||
)
|
)
|
||||||
|
candidates = self._expand_pool_candidates_for_async_submit(candidates)
|
||||||
|
|
||||||
if max_candidates is not None and max_candidates > 0:
|
if max_candidates is not None and max_candidates > 0:
|
||||||
candidates = candidates[:max_candidates]
|
candidates = candidates[:max_candidates]
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||||
|
from src.services.scheduling.schemas import PoolCandidate
|
||||||
|
|
||||||
|
|
||||||
def _mock_key(key_id: str, api_formats: list[str]) -> MagicMock:
|
def _mock_key(key_id: str, api_formats: list[str]) -> MagicMock:
|
||||||
@@ -27,7 +28,7 @@ def _mock_endpoint(api_format: str) -> MagicMock:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_pool_provider_enumerates_all_key_candidates() -> None:
|
async def test_pool_provider_builds_single_pool_candidate() -> None:
|
||||||
scheduler = CacheAwareScheduler()
|
scheduler = CacheAwareScheduler()
|
||||||
builder = scheduler._candidate_builder
|
builder = scheduler._candidate_builder
|
||||||
builder._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[method-assign]
|
builder._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[method-assign]
|
||||||
@@ -53,5 +54,8 @@ async def test_pool_provider_enumerates_all_key_candidates() -> None:
|
|||||||
global_conversion_enabled=True,
|
global_conversion_enabled=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert len(candidates) == 2
|
assert len(candidates) == 1
|
||||||
assert {str(c.key.id) for c in candidates} == {"k1", "k2"}
|
pool_candidate = candidates[0]
|
||||||
|
assert isinstance(pool_candidate, PoolCandidate)
|
||||||
|
assert str(pool_candidate.key.id) == "k1"
|
||||||
|
assert {str(k.id) for k in pool_candidate.pool_keys} == {"k1", "k2"}
|
||||||
|
|||||||
Reference in New Issue
Block a user