feat(pool): 引入多维评分调度策略与账号状态检测

- 新增 multi_score 调度模式,支持 LRU/延迟/健康度/剩余额度多维加权评分
- 新增调度预设维度系统(free_team_first, quota_balanced, recent_refresh, single_account),支持有序对象列表配置格式并兼容旧字符串列表
- 新增 account_state 模块,统一账号封禁/受限检测逻辑,替代分散在 routes 中的判断代码
- 新增 health_cache 模块和 latency 采样(redis_ops.record_latency / batch_get_latency_avgs)
- RequestDispatcher 返回 ttfb_ms,PoolManager.on_request_success 记录延迟样本
- 前端:PoolConfigDialog 替换为 PoolSchedulingDialog,支持预设维度可视化配置;号池管理页增加调度模式标签与账号异常 Badge 显示
- 提取前端 accountBlock 工具函数,ProviderDetailDrawer 复用统一判断
- scheduling_dimensions 增加 account_state 和 latency 维度评估
- 补充 account_state、health_cache、multi_score 策略、preset 维度、redis latency 等测试
This commit is contained in:
fawney19
2026-03-04 22:06:19 +08:00
parent 57b86034cf
commit b2dcf82ca8
40 changed files with 4167 additions and 660 deletions

View File

@@ -75,6 +75,20 @@ export interface PoolOverviewResponse {
items: PoolOverviewItem[] items: PoolOverviewItem[]
} }
export interface PoolPresetModeMeta {
value: string
label: string
}
export interface PoolPresetMeta {
name: string
label: string
description: string
providers: string[]
modes?: PoolPresetModeMeta[] | null
default_mode?: string | null
}
export interface PoolKeyDetail { export interface PoolKeyDetail {
key_id: string key_id: string
key_name: string key_name: string
@@ -182,6 +196,13 @@ export async function getPoolOverview(): Promise<PoolOverviewResponse> {
}) })
} }
export async function getPoolSchedulingPresets(): Promise<PoolPresetMeta[]> {
return dedupedRequest('pool:scheduling-presets', async () => {
const response = await client.get<PoolPresetMeta[]>('/api/admin/pool/scheduling-presets')
return response.data
})
}
export async function listPoolKeys( export async function listPoolKeys(
providerId: string, providerId: string,
params: PoolKeysQuery = {}, params: PoolKeysQuery = {},

View File

@@ -453,11 +453,29 @@ export interface ClaudeCodeAdvancedConfig {
cli_only_enabled?: boolean cli_only_enabled?: boolean
} }
export interface SchedulingPresetItem {
preset: string
enabled: boolean
mode?: string | null
}
export interface PoolAdvancedConfig { export interface PoolAdvancedConfig {
global_priority?: number | null 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
scheduling_mode?: 'lru' | 'multi_score' | null
// 新格式:对象列表;旧格式:字符串列表
scheduling_presets?: SchedulingPresetItem[] | string[] | null
scoring_weights?: {
lru?: number
latency?: number
health?: number
cost_remaining?: number
} | null
latency_window_seconds?: number | null
latency_sample_limit?: number | null
cost_window_seconds?: number | null cost_window_seconds?: number | null
cost_limit_per_key_tokens?: number | null cost_limit_per_key_tokens?: number | null
cost_soft_threshold_percent?: number | null cost_soft_threshold_percent?: number | null

View File

@@ -1,474 +0,0 @@
<template>
<Dialog
:model-value="modelValue"
title="号池配置"
description="调整号池调度策略和健康检查参数"
size="lg"
@update:model-value="emit('update:modelValue', $event)"
>
<form
class="space-y-5"
@submit.prevent="handleSave"
>
<!-- 调度策略 -->
<div class="space-y-3">
<h3 class="text-sm font-medium border-b pb-2">
调度策略
</h3>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">LRU 调度</span>
<p class="text-xs text-muted-foreground">
优先选择最久未用的 Key
</p>
</div>
<Switch
:model-value="form.lru_enabled"
@update:model-value="(v: boolean) => form.lru_enabled = v"
/>
</div>
<div class="grid grid-cols-2 gap-4">
<div class="space-y-1.5">
<Label>
粘性会话 TTL
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.sticky_session_ttl_seconds ?? ''"
type="number"
min="60"
max="86400"
placeholder="3600 (留空禁用)"
@update:model-value="(v) => form.sticky_session_ttl_seconds = parseNum(v)"
/>
<p class="text-xs text-muted-foreground">
同一对话始终路由到同一 Key
</p>
</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 class="space-y-3">
<h3 class="text-sm font-medium border-b pb-2">
冷却与健康
</h3>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">健康策略</span>
<p class="text-xs text-muted-foreground">
按上游错误码自动冷却/禁用 Key
</p>
</div>
<Switch
:model-value="form.health_policy_enabled"
@update:model-value="(v: boolean) => form.health_policy_enabled = v"
/>
</div>
<div class="grid grid-cols-2 gap-4">
<div class="space-y-1.5">
<Label>
429 冷却
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.rate_limit_cooldown_seconds ?? ''"
type="number"
min="10"
max="3600"
placeholder="300"
@update:model-value="(v) => form.rate_limit_cooldown_seconds = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
529 冷却
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.overload_cooldown_seconds ?? ''"
type="number"
min="5"
max="600"
placeholder="30"
@update:model-value="(v) => form.overload_cooldown_seconds = parseNum(v)"
/>
</div>
</div>
</div>
<!-- 成本控制 -->
<div class="space-y-3">
<h3 class="text-sm font-medium border-b pb-2">
成本控制
</h3>
<div class="grid grid-cols-2 gap-4">
<div class="space-y-1.5">
<Label>
成本窗口
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.cost_window_seconds ?? ''"
type="number"
min="3600"
max="86400"
placeholder="18000 (5 小时)"
@update:model-value="(v) => form.cost_window_seconds = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
Key 窗口限额
<span class="text-xs text-muted-foreground">(tokens)</span>
</Label>
<Input
:model-value="form.cost_limit_per_key_tokens ?? ''"
type="number"
min="0"
placeholder="留空 = 不限"
@update:model-value="(v) => form.cost_limit_per_key_tokens = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
软阈值
<span class="text-xs text-muted-foreground">(%)</span>
</Label>
<Input
:model-value="form.cost_soft_threshold_percent ?? ''"
type="number"
min="0"
max="100"
placeholder="80"
@update:model-value="(v) => form.cost_soft_threshold_percent = parseNum(v)"
/>
</div>
</div>
</div>
<!-- Claude Code 特有配置 -->
<template v-if="providerType === 'claude_code'">
<div class="space-y-3">
<h3 class="text-sm font-medium border-b pb-2">
Claude Code
</h3>
<div class="space-y-3 p-3 border rounded-lg bg-muted/50">
<div class="flex items-center justify-between">
<div class="space-y-0.5">
<span class="text-sm font-medium">会话数量控制</span>
<p class="text-xs text-muted-foreground">
限制同时活跃的会话数量
</p>
</div>
<Switch
:model-value="claudeForm.session_control_enabled"
@update:model-value="(v: boolean) => claudeForm.session_control_enabled = v"
/>
</div>
<div
v-if="claudeForm.session_control_enabled"
class="grid grid-cols-2 gap-3"
>
<div class="space-y-1.5">
<Label class="text-xs">最大会话数</Label>
<Input
:model-value="claudeForm.max_sessions ?? ''"
type="number"
min="1"
max="1000"
placeholder="例如 20"
@update:model-value="(v) => claudeForm.max_sessions = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label class="text-xs">会话空闲超时 (分钟)</Label>
<Input
:model-value="claudeForm.session_idle_timeout_minutes ?? ''"
type="number"
min="1"
max="1440"
placeholder="默认 5"
@update:model-value="(v) => claudeForm.session_idle_timeout_minutes = parseNum(v) ?? 5"
/>
</div>
</div>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">TLS 指纹模拟</span>
<p class="text-xs text-muted-foreground">
模拟 Node.js / Claude Code 客户端的 TLS 指纹
</p>
</div>
<Switch
:model-value="claudeForm.enable_tls_fingerprint"
@update:model-value="(v: boolean) => claudeForm.enable_tls_fingerprint = v"
/>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">会话 ID 伪装</span>
<p class="text-xs text-muted-foreground">
启用后在 15 分钟内固定 metadata.user_id 中的 session ID
</p>
</div>
<Switch
:model-value="claudeForm.session_id_masking_enabled"
@update:model-value="(v: boolean) => claudeForm.session_id_masking_enabled = v"
/>
</div>
<div class="space-y-3 p-3 border rounded-lg bg-muted/50">
<div class="flex items-center justify-between">
<div class="space-y-0.5">
<span class="text-sm font-medium">Cache TTL 统一</span>
<p class="text-xs text-muted-foreground">
强制统一所有请求的 cache_control 类型,避免多人共用时行为指纹不一致
</p>
</div>
<Switch
:model-value="claudeForm.cache_ttl_override_enabled"
@update:model-value="(v: boolean) => claudeForm.cache_ttl_override_enabled = v"
/>
</div>
<div
v-if="claudeForm.cache_ttl_override_enabled"
class="space-y-1.5"
>
<Label class="text-xs">目标 TTL 类型</Label>
<select
:value="claudeForm.cache_ttl_override_target"
class="flex h-9 w-full rounded-md border border-input bg-transparent px-3 py-1 text-sm shadow-sm transition-colors focus-visible:outline-none focus-visible:ring-1 focus-visible:ring-ring"
@change="(e) => claudeForm.cache_ttl_override_target = (e.target as HTMLSelectElement).value"
>
<option value="ephemeral">
ephemeral (5 分钟)
</option>
<option value="1h">
1h (1 小时)
</option>
</select>
</div>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">仅限 CLI 客户端</span>
<p class="text-xs text-muted-foreground">
仅允许 Claude Code CLI 客户端访问,拒绝非 CLI 流量
</p>
</div>
<Switch
:model-value="claudeForm.cli_only_enabled"
@update:model-value="(v: boolean) => claudeForm.cli_only_enabled = v"
/>
</div>
</div>
</template>
</form>
<template #footer>
<Button
variant="outline"
:disabled="loading"
@click="emit('update:modelValue', false)"
>
取消
</Button>
<Button
:disabled="loading"
@click="handleSave"
>
{{ loading ? '保存中...' : '保存' }}
</Button>
</template>
</Dialog>
</template>
<script setup lang="ts">
import { ref, watch } from 'vue'
import { Dialog, Button, Input, Label, Switch } from '@/components/ui'
import { useToast } from '@/composables/useToast'
import { parseApiError } from '@/utils/errorParser'
import { updateProvider } from '@/api/endpoints'
import type { PoolAdvancedConfig, ClaudeCodeAdvancedConfig } from '@/api/endpoints/types/provider'
const props = defineProps<{
modelValue: boolean
providerId: string
providerType?: string
currentConfig: PoolAdvancedConfig | null
currentClaudeConfig?: ClaudeCodeAdvancedConfig | null
}>()
const emit = defineEmits<{
'update:modelValue': [value: boolean]
saved: []
}>()
const { success, error: showError } = useToast()
const loading = ref(false)
const form = ref<PoolAdvancedConfig>({
global_priority: null,
sticky_session_ttl_seconds: null,
lru_enabled: true,
cost_window_seconds: null,
cost_limit_per_key_tokens: null,
cost_soft_threshold_percent: null,
rate_limit_cooldown_seconds: null,
overload_cooldown_seconds: null,
health_policy_enabled: true,
})
interface ClaudeFormState {
session_control_enabled: boolean
max_sessions: number | undefined
session_idle_timeout_minutes: number
enable_tls_fingerprint: boolean
session_id_masking_enabled: boolean
cache_ttl_override_enabled: boolean
cache_ttl_override_target: string
cli_only_enabled: boolean
}
const claudeForm = ref<ClaudeFormState>({
session_control_enabled: true,
max_sessions: undefined,
session_idle_timeout_minutes: 5,
enable_tls_fingerprint: true,
session_id_masking_enabled: true,
cache_ttl_override_enabled: false,
cache_ttl_override_target: 'ephemeral',
cli_only_enabled: false,
})
function parseNum(v: string | number): number | undefined {
if (v === '' || v === null || v === undefined) return undefined
const n = Number(v)
return isNaN(n) ? undefined : n
}
watch(() => props.modelValue, (v) => {
if (v && props.currentConfig) {
form.value = { ...props.currentConfig }
} else if (v) {
form.value = {
global_priority: null,
sticky_session_ttl_seconds: null,
lru_enabled: true,
cost_window_seconds: null,
cost_limit_per_key_tokens: null,
cost_soft_threshold_percent: null,
rate_limit_cooldown_seconds: null,
overload_cooldown_seconds: null,
health_policy_enabled: true,
}
}
// Claude Code 配置
if (v && props.providerType === 'claude_code') {
const cc = props.currentClaudeConfig
if (cc) {
// max_sessions 为 null 表示用户明确关闭了会话控制,其余情况默认开启
const sessionOff = cc.max_sessions === null
claudeForm.value = {
session_control_enabled: !sessionOff,
max_sessions: sessionOff ? undefined : (cc.max_sessions ?? undefined),
session_idle_timeout_minutes: cc.session_idle_timeout_minutes ?? 5,
enable_tls_fingerprint: cc.enable_tls_fingerprint ?? true,
session_id_masking_enabled: cc.session_id_masking_enabled ?? true,
cache_ttl_override_enabled: cc.cache_ttl_override_enabled ?? false,
cache_ttl_override_target: cc.cache_ttl_override_target ?? 'ephemeral',
cli_only_enabled: cc.cli_only_enabled ?? false,
}
} else {
// 默认值:全部开启
claudeForm.value = {
session_control_enabled: true,
max_sessions: undefined,
session_idle_timeout_minutes: 5,
enable_tls_fingerprint: true,
session_id_masking_enabled: true,
cache_ttl_override_enabled: false,
cache_ttl_override_target: 'ephemeral',
cli_only_enabled: false,
}
}
}
})
async function handleSave() {
loading.value = true
try {
const payload: Record<string, unknown> = {
pool_advanced: {
global_priority: form.value.global_priority ?? undefined,
sticky_session_ttl_seconds: form.value.sticky_session_ttl_seconds ?? undefined,
lru_enabled: form.value.lru_enabled,
cost_window_seconds: form.value.cost_window_seconds ?? undefined,
cost_limit_per_key_tokens: form.value.cost_limit_per_key_tokens ?? undefined,
cost_soft_threshold_percent: form.value.cost_soft_threshold_percent ?? undefined,
rate_limit_cooldown_seconds: form.value.rate_limit_cooldown_seconds ?? undefined,
overload_cooldown_seconds: form.value.overload_cooldown_seconds ?? undefined,
health_policy_enabled: form.value.health_policy_enabled,
},
}
// Claude Code 特有配置
if (props.providerType === 'claude_code') {
payload.claude_code_advanced = {
max_sessions: claudeForm.value.session_control_enabled
? (claudeForm.value.max_sessions ?? undefined)
: null,
session_idle_timeout_minutes: claudeForm.value.session_control_enabled
? (claudeForm.value.session_idle_timeout_minutes ?? 5)
: null,
enable_tls_fingerprint: claudeForm.value.enable_tls_fingerprint,
session_id_masking_enabled: claudeForm.value.session_id_masking_enabled,
cache_ttl_override_enabled: claudeForm.value.cache_ttl_override_enabled,
cache_ttl_override_target: claudeForm.value.cache_ttl_override_enabled
? claudeForm.value.cache_ttl_override_target
: 'ephemeral',
cli_only_enabled: claudeForm.value.cli_only_enabled,
}
}
await updateProvider(props.providerId, payload)
success('号池配置已保存')
emit('saved')
emit('update:modelValue', false)
} catch (err) {
showError(parseApiError(err))
} finally {
loading.value = false
}
}
</script>

View File

@@ -0,0 +1,886 @@
<template>
<Dialog
:model-value="modelValue"
title="号池调度"
description="拖拽排序调度维度,越靠前优先级越高"
size="lg"
@update:model-value="emit('update:modelValue', $event)"
>
<div class="space-y-5">
<!-- Preset List -->
<div class="space-y-3">
<div class="space-y-1">
<h3 class="text-sm font-medium border-b pb-2">
调度维度
</h3>
<p class="text-xs text-muted-foreground">
拖拽排序越靠前优先级越高不适用当前 Provider 类型的维度已禁用
</p>
</div>
<div class="space-y-0.5">
<div
v-for="(item, index) in presetList"
:key="item.preset"
class="group flex items-center gap-3 px-3 py-2.5 rounded-lg border transition-all duration-200"
:class="[
!item.applicable
? 'border-border/30 bg-muted/20 opacity-50'
: draggedIndex === index
? 'border-primary/50 bg-primary/5 shadow-md scale-[1.01]'
: dragOverIndex === index
? 'border-primary/30 bg-primary/5'
: 'border-border/50 bg-background hover:border-border hover:bg-muted/30'
]"
:draggable="item.applicable"
@dragstart="item.applicable && handleDragStart(index, $event)"
@dragend="handleDragEnd"
@dragover.prevent="item.applicable && handleDragOver(index)"
@dragleave="handleDragLeave"
@drop="item.applicable && handleDrop(index)"
>
<!-- Drag handle -->
<div
class="p-1 rounded transition-colors shrink-0"
:class="item.applicable
? 'cursor-grab active:cursor-grabbing text-muted-foreground/40 group-hover:text-muted-foreground'
: 'text-muted-foreground/15 cursor-default'"
>
<GripVertical class="w-4 h-4" />
</div>
<!-- Enable/disable switch -->
<Switch
:model-value="item.enabled"
:disabled="!item.applicable"
@update:model-value="(v: boolean) => togglePreset(index, v)"
/>
<!-- Info -->
<div class="flex-1 min-w-0">
<div class="flex items-center gap-2">
<span
class="text-sm font-medium"
:class="!item.applicable ? 'text-muted-foreground' : ''"
>{{ item.label }}</span>
<span
v-if="!item.applicable"
class="text-[10px] text-muted-foreground/60"
>
(不适用)
</span>
</div>
<p class="text-xs text-muted-foreground mt-0.5">
{{ item.desc }}
</p>
<!-- Mode sub-config -->
<div
v-if="item.modeOptions.length > 0 && item.enabled && item.applicable"
class="flex gap-0.5 mt-2 p-0.5 bg-muted/40 rounded-md w-fit"
>
<button
v-for="modeOpt in item.modeOptions"
:key="modeOpt.value"
type="button"
class="px-2.5 py-1 text-xs font-medium rounded transition-all"
:class="[
item.mode === modeOpt.value
? 'bg-primary text-primary-foreground shadow-sm'
: 'text-muted-foreground hover:text-foreground hover:bg-background/50'
]"
@click="setPresetMode(index, modeOpt.value)"
>
{{ modeOpt.label }}
</button>
</div>
</div>
</div>
</div>
</div>
<!-- Advanced toggle -->
<div class="pt-1">
<Button
type="button"
size="sm"
variant="ghost"
@click="showAdvanced = !showAdvanced"
>
{{ showAdvanced ? '收起高级参数' : '展开高级参数' }}
</Button>
</div>
<!-- Advanced options -->
<div
v-if="showAdvanced"
class="space-y-4"
>
<!-- Cooldown & Health -->
<div class="space-y-3">
<h3 class="text-sm font-medium border-b pb-2">
冷却与健康
</h3>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">健康策略</span>
<p class="text-xs text-muted-foreground">
按上游错误自动冷却并跳过账号
</p>
</div>
<Switch
:model-value="form.health_policy_enabled"
@update:model-value="(v: boolean) => form.health_policy_enabled = v"
/>
</div>
<div class="grid grid-cols-2 gap-4">
<div class="space-y-1.5">
<Label>
429 冷却
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.rate_limit_cooldown_seconds ?? ''"
type="number"
min="10"
max="3600"
placeholder="300"
@update:model-value="(v) => form.rate_limit_cooldown_seconds = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
529 冷却
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.overload_cooldown_seconds ?? ''"
type="number"
min="5"
max="600"
placeholder="30"
@update:model-value="(v) => form.overload_cooldown_seconds = parseNum(v)"
/>
</div>
</div>
<div class="grid grid-cols-2 gap-4">
<div class="space-y-1.5">
<Label>
粘性会话 TTL
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.sticky_session_ttl_seconds ?? ''"
type="number"
min="60"
max="86400"
placeholder="3600 (留空禁用)"
@update:model-value="(v) => form.sticky_session_ttl_seconds = parseNum(v)"
/>
</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)"
/>
</div>
</div>
</div>
<!-- Claude Code -->
<div
v-if="isClaudeCode"
class="space-y-3"
>
<h3 class="text-sm font-medium border-b pb-2">
Claude Code
</h3>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">TLS 指纹模拟</span>
<p class="text-xs text-muted-foreground">
模拟 Node.js / Claude Code 客户端指纹
</p>
</div>
<Switch
:model-value="claudeForm.enable_tls_fingerprint"
@update:model-value="(v: boolean) => claudeForm.enable_tls_fingerprint = v"
/>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">Session ID 伪装</span>
<p class="text-xs text-muted-foreground">
固定 metadata.user_id 中 session 片段
</p>
</div>
<Switch
:model-value="claudeForm.session_id_masking_enabled"
@update:model-value="(v: boolean) => claudeForm.session_id_masking_enabled = v"
/>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">仅限 CLI 客户端</span>
<p class="text-xs text-muted-foreground">
仅允许 Claude Code CLI 格式请求
</p>
</div>
<Switch
:model-value="claudeForm.cli_only_enabled"
@update:model-value="(v: boolean) => claudeForm.cli_only_enabled = v"
/>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">Cache TTL 统一</span>
<p class="text-xs text-muted-foreground">
强制所有 cache_control 使用相同 TTL 类型
</p>
</div>
<Switch
:model-value="claudeForm.cache_ttl_override_enabled"
@update:model-value="(v: boolean) => claudeForm.cache_ttl_override_enabled = v"
/>
</div>
<div
v-if="claudeForm.cache_ttl_override_enabled"
class="pl-3"
>
<div class="space-y-1.5">
<Label>TTL 类型</Label>
<div class="flex gap-0.5 p-0.5 bg-muted/40 rounded-md w-fit">
<button
v-for="opt in ['ephemeral']"
:key="opt"
type="button"
class="px-2.5 py-1 text-xs font-medium rounded transition-all"
:class="[
claudeForm.cache_ttl_override_target === opt
? 'bg-primary text-primary-foreground shadow-sm'
: 'text-muted-foreground hover:text-foreground hover:bg-background/50'
]"
@click="claudeForm.cache_ttl_override_target = opt"
>
{{ opt }}
</button>
</div>
</div>
</div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5">
<span class="text-sm font-medium">会话数量控制</span>
<p class="text-xs text-muted-foreground">
限制单 Key 同时活跃会话数
</p>
</div>
<Switch
:model-value="claudeForm.session_control_enabled"
@update:model-value="(v: boolean) => claudeForm.session_control_enabled = v"
/>
</div>
<div
v-if="claudeForm.session_control_enabled"
class="grid grid-cols-2 gap-4"
>
<div class="space-y-1.5">
<Label>
最大会话数
</Label>
<Input
:model-value="claudeForm.max_sessions ?? ''"
type="number"
min="1"
max="100"
placeholder="留空 = 不限"
@update:model-value="(v) => claudeForm.max_sessions = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
空闲超时
<span class="text-xs text-muted-foreground">(分钟)</span>
</Label>
<Input
:model-value="claudeForm.session_idle_timeout_minutes ?? ''"
type="number"
min="1"
max="1440"
placeholder="5"
@update:model-value="(v) => claudeForm.session_idle_timeout_minutes = parseNum(v) ?? 5"
/>
</div>
</div>
</div>
<!-- Cost Control -->
<div class="space-y-3">
<h3 class="text-sm font-medium border-b pb-2">
成本控制
</h3>
<div class="grid grid-cols-2 gap-4">
<div class="space-y-1.5">
<Label>
成本窗口
<span class="text-xs text-muted-foreground">(秒)</span>
</Label>
<Input
:model-value="form.cost_window_seconds ?? ''"
type="number"
min="3600"
max="86400"
placeholder="18000 (5 小时)"
@update:model-value="(v) => form.cost_window_seconds = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
Key 窗口限额
<span class="text-xs text-muted-foreground">(tokens)</span>
</Label>
<Input
:model-value="form.cost_limit_per_key_tokens ?? ''"
type="number"
min="0"
placeholder="留空 = 不限"
@update:model-value="(v) => form.cost_limit_per_key_tokens = parseNum(v)"
/>
</div>
<div class="space-y-1.5">
<Label>
软阈值
<span class="text-xs text-muted-foreground">(%)</span>
</Label>
<Input
:model-value="form.cost_soft_threshold_percent ?? ''"
type="number"
min="0"
max="100"
placeholder="80"
@update:model-value="(v) => form.cost_soft_threshold_percent = parseNum(v)"
/>
</div>
</div>
</div>
</div>
</div>
<template #footer>
<Button
variant="outline"
:disabled="loading"
@click="emit('update:modelValue', false)"
>
取消
</Button>
<Button
:disabled="loading"
@click="handleSave"
>
{{ loading ? '保存中...' : '保存' }}
</Button>
</template>
</Dialog>
</template>
<script setup lang="ts">
import { computed, ref, watch } from 'vue'
import { GripVertical } from 'lucide-vue-next'
import { Dialog, Button, Input, Label, Switch } from '@/components/ui'
import { useToast } from '@/composables/useToast'
import { parseApiError } from '@/utils/errorParser'
import { updateProvider } from '@/api/endpoints'
import { getPoolSchedulingPresets } from '@/api/endpoints/pool'
import type { PoolPresetMeta } from '@/api/endpoints/pool'
import type { PoolAdvancedConfig, ClaudeCodeAdvancedConfig, SchedulingPresetItem } from '@/api/endpoints/types/provider'
interface PresetModeOption {
value: string
label: string
}
interface PresetListItem {
preset: string
label: string
desc: string
enabled: boolean
mode: string | null
modeOptions: PresetModeOption[]
applicable: boolean
}
const props = defineProps<{
modelValue: boolean
providerId: string
providerType?: string
currentConfig: PoolAdvancedConfig | null
currentClaudeConfig?: ClaudeCodeAdvancedConfig | null
}>()
const emit = defineEmits<{
'update:modelValue': [value: boolean]
saved: []
}>()
const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
{
name: 'lru',
label: 'LRU 轮转',
description: '最久未使用的 Key 优先',
providers: [],
modes: null,
default_mode: null,
},
{
name: 'free_team_first',
label: 'Free/Team 优先',
description: '优先消耗低档账号(依赖 plan_type',
providers: ['codex', 'kiro'],
modes: [
{ value: 'free_only', label: 'Free' },
{ value: 'team_only', label: 'Team' },
{ value: 'both', label: '全部' },
],
default_mode: 'both',
},
{
name: 'quota_balanced',
label: '额度平均',
description: '优先选额度消耗最少的账号',
providers: [],
modes: null,
default_mode: null,
},
{
name: 'recent_refresh',
label: '额度刷新优先',
description: '优先选即将刷新额度的账号',
providers: ['codex', 'kiro'],
modes: null,
default_mode: null,
},
{
name: 'single_account',
label: '单号优先',
description: '集中使用同一账号(反向 LRU',
providers: [],
modes: null,
default_mode: null,
},
]
const DEFAULT_ENABLED_PRESETS = new Set(['lru', 'quota_balanced'])
const { success, error: showError } = useToast()
const loading = ref(false)
const showAdvanced = ref(false)
const presetDefs = ref<PoolPresetMeta[]>([])
const presetDefsLoaded = ref(false)
const loadingPresetDefs = ref(false)
const draggedIndex = ref<number | null>(null)
const dragOverIndex = ref<number | null>(null)
const presetList = ref<PresetListItem[]>([])
const form = ref({
global_priority: null as number | null | undefined,
sticky_session_ttl_seconds: null as number | null | undefined,
health_policy_enabled: true,
rate_limit_cooldown_seconds: null as number | null | undefined,
overload_cooldown_seconds: null as number | null | undefined,
cost_window_seconds: null as number | null | undefined,
cost_limit_per_key_tokens: null as number | null | undefined,
cost_soft_threshold_percent: null as number | null | undefined,
})
const isClaudeCode = computed(() => normalizeProviderType(props.providerType) === 'claude_code')
interface ClaudeFormState {
session_control_enabled: boolean
max_sessions: number | undefined
session_idle_timeout_minutes: number
enable_tls_fingerprint: boolean
session_id_masking_enabled: boolean
cache_ttl_override_enabled: boolean
cache_ttl_override_target: string
cli_only_enabled: boolean
}
const claudeForm = ref<ClaudeFormState>({
session_control_enabled: true,
max_sessions: undefined,
session_idle_timeout_minutes: 5,
enable_tls_fingerprint: true,
session_id_masking_enabled: true,
cache_ttl_override_enabled: false,
cache_ttl_override_target: 'ephemeral',
cli_only_enabled: false,
})
function parseNum(v: string | number): number | undefined {
if (v === '' || v === null || v === undefined) return undefined
const n = Number(v)
return Number.isNaN(n) ? undefined : n
}
function normalizeProviderType(value: string | undefined): string {
return (value || '').trim().toLowerCase()
}
function normalizePresetName(value: unknown): string {
return String(value ?? '').trim().toLowerCase()
}
function normalizeMode(value: unknown): string | null {
const normalized = String(value ?? '').trim().toLowerCase()
return normalized || null
}
function normalizePresetDefs(defs: PoolPresetMeta[]): PoolPresetMeta[] {
const ordered: PoolPresetMeta[] = []
const seen = new Set<string>()
for (const raw of defs) {
const name = normalizePresetName(raw.name)
if (!name || seen.has(name)) continue
seen.add(name)
const providers = Array.isArray(raw.providers)
? raw.providers.map(p => normalizeProviderType(p)).filter(Boolean)
: []
const modes = Array.isArray(raw.modes)
? raw.modes
.map(mode => ({
value: normalizePresetName(mode.value),
label: String(mode.label ?? '').trim() || String(mode.value ?? '').trim(),
}))
.filter(mode => Boolean(mode.value))
: null
const defaultMode = normalizeMode(raw.default_mode)
ordered.push({
name,
label: String(raw.label ?? '').trim() || name,
description: String(raw.description ?? '').trim(),
providers,
modes: modes && modes.length > 0 ? modes : null,
default_mode: defaultMode,
})
}
return ordered
}
function getPresetDefs(): PoolPresetMeta[] {
if (presetDefs.value.length > 0) {
return presetDefs.value
}
return FALLBACK_PRESET_DEFS
}
async function ensurePresetDefsLoaded(): Promise<void> {
if (presetDefsLoaded.value || loadingPresetDefs.value) return
loadingPresetDefs.value = true
try {
const remoteDefs = await getPoolSchedulingPresets()
const normalized = normalizePresetDefs(Array.isArray(remoteDefs) ? remoteDefs : [])
if (normalized.length > 0) {
presetDefs.value = normalized
}
} catch (err) {
showError(parseApiError(err))
} finally {
presetDefsLoaded.value = true
loadingPresetDefs.value = false
}
}
function isApplicablePreset(def: PoolPresetMeta): boolean {
const providerType = normalizeProviderType(props.providerType)
const providers = Array.isArray(def.providers) ? def.providers : []
if (providers.length === 0) return true
if (!providerType) return true
return providers.includes(providerType)
}
function getModeOptions(def: PoolPresetMeta): PresetModeOption[] {
const modes = Array.isArray(def.modes) ? def.modes : []
return modes
.map(mode => ({
value: normalizePresetName(mode.value),
label: String(mode.label ?? '').trim() || String(mode.value ?? '').trim(),
}))
.filter(mode => Boolean(mode.value))
}
function defaultModeForPreset(def: PoolPresetMeta): string | null {
const options = getModeOptions(def)
if (options.length === 0) return null
const normalizedDefault = normalizeMode(def.default_mode)
if (normalizedDefault && options.some(option => option.value === normalizedDefault)) {
return normalizedDefault
}
return options[0].value
}
function buildDefaultPresetList(): PresetListItem[] {
return getPresetDefs().map(def => ({
preset: def.name,
label: def.label,
desc: def.description,
enabled: DEFAULT_ENABLED_PRESETS.has(def.name),
mode: defaultModeForPreset(def),
modeOptions: getModeOptions(def),
applicable: isApplicablePreset(def),
}))
}
function isNewFormatPresetItem(item: unknown): item is SchedulingPresetItem {
return typeof item === 'object' && item !== null && 'preset' in item
}
function resolveMode(def: PoolPresetMeta, mode: unknown): string | null {
const options = getModeOptions(def)
if (options.length === 0) return null
const normalized = normalizeMode(mode)
if (normalized && options.some(option => option.value === normalized)) {
return normalized
}
return defaultModeForPreset(def)
}
function loadFromConfig(cfg: PoolAdvancedConfig | null): PresetListItem[] {
const defs = getPresetDefs()
const defsByName = new Map(defs.map(def => [def.name, def]))
const defaults = buildDefaultPresetList()
if (!cfg) return defaults
const rawPresets = cfg.scheduling_presets
if (!Array.isArray(rawPresets) || rawPresets.length === 0) {
if (cfg.scheduling_mode === 'lru' || (!cfg.scheduling_mode && cfg.lru_enabled !== false)) {
return defaults.map(item => ({
...item,
enabled: item.preset === 'lru',
}))
}
return defaults
}
const first = rawPresets[0]
if (isNewFormatPresetItem(first)) {
const configItems = rawPresets as SchedulingPresetItem[]
const ordered: PresetListItem[] = []
const seen = new Set<string>()
for (const ci of configItems) {
const presetName = normalizePresetName(ci.preset)
const def = defsByName.get(presetName)
if (!def || seen.has(presetName)) continue
seen.add(presetName)
ordered.push({
preset: presetName,
label: def.label,
desc: def.description,
enabled: ci.enabled !== false,
mode: resolveMode(def, ci.mode),
modeOptions: getModeOptions(def),
applicable: isApplicablePreset(def),
})
}
for (const def of defs) {
if (seen.has(def.name)) continue
ordered.push({
preset: def.name,
label: def.label,
desc: def.description,
enabled: false,
mode: defaultModeForPreset(def),
modeOptions: getModeOptions(def),
applicable: isApplicablePreset(def),
})
}
return ordered
}
const legacyPresets = rawPresets as string[]
const lruEnabled = cfg.lru_enabled !== false
const ordered: PresetListItem[] = []
const seen = new Set<string>()
const lruDef = defsByName.get('lru')
if (lruDef) {
ordered.push({
preset: 'lru',
label: lruDef.label,
desc: lruDef.description,
enabled: lruEnabled,
mode: null,
modeOptions: [],
applicable: isApplicablePreset(lruDef),
})
seen.add('lru')
}
for (const name of legacyPresets) {
const presetName = normalizePresetName(name)
const def = defsByName.get(presetName)
if (!def || seen.has(presetName)) continue
seen.add(presetName)
ordered.push({
preset: presetName,
label: def.label,
desc: def.description,
enabled: true,
mode: resolveMode(def, undefined),
modeOptions: getModeOptions(def),
applicable: isApplicablePreset(def),
})
}
for (const def of defs) {
if (seen.has(def.name)) continue
ordered.push({
preset: def.name,
label: def.label,
desc: def.description,
enabled: false,
mode: defaultModeForPreset(def),
modeOptions: getModeOptions(def),
applicable: isApplicablePreset(def),
})
}
return ordered
}
function togglePreset(index: number, enabled: boolean) {
presetList.value[index].enabled = enabled
}
function setPresetMode(index: number, mode: string) {
presetList.value[index].mode = mode
}
function handleDragStart(index: number, event: DragEvent) {
draggedIndex.value = index
if (event.dataTransfer) {
event.dataTransfer.effectAllowed = 'move'
event.dataTransfer.setData('text/html', '')
}
}
function handleDragEnd() {
draggedIndex.value = null
dragOverIndex.value = null
}
function handleDragOver(index: number) {
dragOverIndex.value = index
}
function handleDragLeave() {
dragOverIndex.value = null
}
function handleDrop(dropIndex: number) {
if (draggedIndex.value === null || draggedIndex.value === dropIndex) {
draggedIndex.value = null
dragOverIndex.value = null
return
}
const items = [...presetList.value]
const [draggedItem] = items.splice(draggedIndex.value, 1)
items.splice(dropIndex, 0, draggedItem)
presetList.value = items
draggedIndex.value = null
dragOverIndex.value = null
}
watch(() => props.modelValue, async (open) => {
if (!open) return
showAdvanced.value = false
await ensurePresetDefsLoaded()
presetList.value = loadFromConfig(props.currentConfig)
const cfg = props.currentConfig
form.value = {
global_priority: cfg?.global_priority ?? null,
sticky_session_ttl_seconds: cfg?.sticky_session_ttl_seconds ?? null,
health_policy_enabled: cfg?.health_policy_enabled !== false,
rate_limit_cooldown_seconds: cfg?.rate_limit_cooldown_seconds ?? null,
overload_cooldown_seconds: cfg?.overload_cooldown_seconds ?? null,
cost_window_seconds: cfg?.cost_window_seconds ?? null,
cost_limit_per_key_tokens: cfg?.cost_limit_per_key_tokens ?? null,
cost_soft_threshold_percent: cfg?.cost_soft_threshold_percent ?? null,
}
const cc = props.currentClaudeConfig
claudeForm.value = {
session_control_enabled: cc?.max_sessions !== null,
max_sessions: cc?.max_sessions ?? undefined,
session_idle_timeout_minutes: cc?.session_idle_timeout_minutes ?? 5,
enable_tls_fingerprint: cc?.enable_tls_fingerprint !== false,
session_id_masking_enabled: cc?.session_id_masking_enabled !== false,
cache_ttl_override_enabled: cc?.cache_ttl_override_enabled ?? false,
cache_ttl_override_target: cc?.cache_ttl_override_target ?? 'ephemeral',
cli_only_enabled: cc?.cli_only_enabled ?? false,
}
})
async function handleSave() {
loading.value = true
try {
const schedulingPresets: SchedulingPresetItem[] = presetList.value.map(item => {
const result: SchedulingPresetItem = {
preset: item.preset,
enabled: item.enabled && item.applicable,
}
if (item.modeOptions.length > 0 && item.mode) {
result.mode = item.mode
}
return result
})
const payload: Parameters<typeof updateProvider>[1] = {
pool_advanced: {
global_priority: form.value.global_priority ?? undefined,
sticky_session_ttl_seconds: form.value.sticky_session_ttl_seconds ?? undefined,
scheduling_presets: schedulingPresets,
scoring_weights: undefined,
latency_window_seconds: undefined,
latency_sample_limit: undefined,
cost_window_seconds: form.value.cost_window_seconds ?? undefined,
cost_limit_per_key_tokens: form.value.cost_limit_per_key_tokens ?? undefined,
cost_soft_threshold_percent: form.value.cost_soft_threshold_percent ?? undefined,
rate_limit_cooldown_seconds: form.value.rate_limit_cooldown_seconds ?? undefined,
overload_cooldown_seconds: form.value.overload_cooldown_seconds ?? undefined,
health_policy_enabled: form.value.health_policy_enabled,
},
}
if (isClaudeCode.value) {
const cf = claudeForm.value
payload.claude_code_advanced = {
max_sessions: cf.session_control_enabled ? (cf.max_sessions ?? null) : null,
session_idle_timeout_minutes: cf.session_control_enabled ? cf.session_idle_timeout_minutes : null,
enable_tls_fingerprint: cf.enable_tls_fingerprint,
session_id_masking_enabled: cf.session_id_masking_enabled,
cache_ttl_override_enabled: cf.cache_ttl_override_enabled,
cache_ttl_override_target: cf.cache_ttl_override_enabled ? cf.cache_ttl_override_target : undefined,
cli_only_enabled: cf.cli_only_enabled,
}
}
await updateProvider(props.providerId, payload)
success('号池调度已保存')
emit('saved')
emit('update:modelValue', false)
} catch (err) {
showError(parseApiError(err))
} finally {
loading.value = false
}
}
</script>

View File

@@ -264,7 +264,7 @@
:class="key.is_pool_aggregate :class="key.is_pool_aggregate
? 'text-muted-foreground/60 cursor-not-allowed' ? 'text-muted-foreground/60 cursor-not-allowed'
: 'text-muted-foreground cursor-pointer hover:bg-primary/10 hover:text-primary'" : 'text-muted-foreground cursor-pointer hover:bg-primary/10 hover:text-primary'"
:title="key.is_pool_aggregate ? '号池优先级请在号池配置中调整' : '点击编辑优先级'" :title="key.is_pool_aggregate ? '号池优先级请在号池调度中调整' : '点击编辑优先级'"
@click.stop="!key.is_pool_aggregate && startEditKeyPriority(format, key)" @click.stop="!key.is_pool_aggregate && startEditKeyPriority(format, key)"
> >
{{ key.priority }} {{ key.priority }}

View File

@@ -1101,6 +1101,7 @@ import {
import type { UpstreamMetadata, AntigravityModelQuota } from '@/api/endpoints/types' import type { UpstreamMetadata, AntigravityModelQuota } from '@/api/endpoints/types'
import { formatApiFormat } from '@/api/endpoints/types/api-format' import { formatApiFormat } from '@/api/endpoints/types/api-format'
import { isOAuthAccountProviderType, isKeyManagedProviderType } from '../utils/providerTypeUtils' import { isOAuthAccountProviderType, isKeyManagedProviderType } from '../utils/providerTypeUtils'
import { isAccountLevelBlockReason } from '@/utils/accountBlock'
// 扩展端点类型,包含密钥列表 // 扩展端点类型,包含密钥列表
interface ProviderEndpointWithKeys extends ProviderEndpoint { interface ProviderEndpointWithKeys extends ProviderEndpoint {
@@ -1568,8 +1569,7 @@ async function handleRefreshOAuth(key: EndpointAPIKey) {
// 判断是否为账号级别的封禁(刷新 token 无法修复) // 判断是否为账号级别的封禁(刷新 token 无法修复)
function isAccountLevelBlock(key: EndpointAPIKey): boolean { function isAccountLevelBlock(key: EndpointAPIKey): boolean {
if (!key.oauth_invalid_reason) return false return isAccountLevelBlockReason(key.oauth_invalid_reason)
return key.oauth_invalid_reason.startsWith('[ACCOUNT_BLOCK]')
} }
// 清除 OAuth 失效标记 // 清除 OAuth 失效标记

View File

@@ -0,0 +1,34 @@
// 账号级别封禁/异常的关键词匹配(用于判断 oauth_invalid_reason 是否属于账号封禁)
const ACCOUNT_BLOCK_REASON_KEYWORDS = [
'account_block',
'account blocked',
'account has been disabled',
'account disabled',
'organization has been disabled',
'organization_disabled',
'validation_required',
'verify your account',
'suspended',
// Kiro quota refresher 写入的确切文本
'账户已封禁',
// Antigravity quota refresher 写入的确切文本
'账户访问被禁止',
'封禁',
'封号',
'被封',
'访问被禁止',
'账号异常',
]
export function isAccountLevelBlockReason(reason: string | null | undefined): boolean {
if (!reason) return false
const text = reason.trim()
if (!text) return false
if (text.startsWith('[ACCOUNT_BLOCK]')) return true
const lowered = text.toLowerCase()
return ACCOUNT_BLOCK_REASON_KEYWORDS.some(keyword => lowered.includes(keyword))
}
export function cleanAccountBlockReason(reason: string): string {
return reason.replace(/^\[ACCOUNT_BLOCK\]\s*/i, '').trim()
}

View File

@@ -34,13 +34,14 @@
</Button> </Button>
<Button <Button
v-if="selectedProviderId" v-if="selectedProviderId"
variant="ghost" variant="outline"
size="icon" size="sm"
class="h-8 w-8" class="h-8 px-2 text-xs gap-1"
title="号池配置" title="调整号池调度"
@click="showConfigDialog = true" @click="showSchedulingDialog = true"
> >
<Settings class="w-3.5 h-3.5" /> 调度
<ChevronDown class="w-3 h-3 text-muted-foreground" />
</Button> </Button>
<Button <Button
v-if="selectedProviderId" v-if="selectedProviderId"
@@ -53,7 +54,8 @@
<Ban class="w-3.5 h-3.5" /> <Ban class="w-3.5 h-3.5" />
</Button> </Button>
<RefreshButton <RefreshButton
:loading="keysLoading" :loading="refreshCurrentPageLoading"
:title="refreshButtonTitle"
@click="refreshCurrentPage" @click="refreshCurrentPage"
/> />
</div> </div>
@@ -197,16 +199,16 @@
> >
<Upload class="w-3.5 h-3.5" /> <Upload class="w-3.5 h-3.5" />
</Button> </Button>
<Button <button
v-if="selectedProviderId" v-if="selectedProviderId"
variant="ghost" class="group inline-flex items-center gap-1.5 px-2.5 h-8 rounded-md border border-border/50 bg-muted/20 hover:bg-muted/40 hover:border-primary/40 transition-all duration-200 text-xs"
size="icon" title="点击调整号池调度"
class="h-8 w-8" @click="showSchedulingDialog = true"
title="号池配置"
@click="showConfigDialog = true"
> >
<Settings class="w-3.5 h-3.5" /> <span class="text-muted-foreground/80 hidden lg:inline">调度:</span>
</Button> <span class="font-medium text-foreground/90">{{ poolSchedulingLabel }}</span>
<ChevronDown class="w-3 h-3 text-muted-foreground/70 group-hover:text-foreground transition-colors" />
</button>
<Button <Button
v-if="selectedProviderId" v-if="selectedProviderId"
variant="ghost" variant="ghost"
@@ -218,7 +220,8 @@
<Ban class="w-3.5 h-3.5" /> <Ban class="w-3.5 h-3.5" />
</Button> </Button>
<RefreshButton <RefreshButton
:loading="keysLoading" :loading="refreshCurrentPageLoading"
:title="refreshButtonTitle"
@click="refreshCurrentPage" @click="refreshCurrentPage"
/> />
</div> </div>
@@ -370,6 +373,14 @@
{{ getKeyOAuthExpires(key)?.text }} {{ getKeyOAuthExpires(key)?.text }}
</span> </span>
</template> </template>
<Badge
v-if="getAccountAlertLabel(key)"
variant="destructive"
class="text-[9px] px-1 py-0 h-4 shrink-0"
:title="getAccountAlertTitle(key)"
>
{{ getAccountAlertLabel(key) }}
</Badge>
<Badge <Badge
v-if="key.oauth_plan_type" v-if="key.oauth_plan_type"
variant="outline" variant="outline"
@@ -753,6 +764,14 @@
{{ getKeyOAuthExpires(key)?.text }} {{ getKeyOAuthExpires(key)?.text }}
</span> </span>
</template> </template>
<Badge
v-if="getAccountAlertLabel(key)"
variant="destructive"
class="text-[9px] px-1 py-0 h-4 shrink-0"
:title="getAccountAlertTitle(key)"
>
{{ getAccountAlertLabel(key) }}
</Badge>
<Badge <Badge
v-if="key.oauth_plan_type" v-if="key.oauth_plan_type"
variant="outline" variant="outline"
@@ -1092,13 +1111,13 @@
@close="showImportDialog = false" @close="showImportDialog = false"
@saved="handleAccountDialogSaved" @saved="handleAccountDialogSaved"
/> />
<PoolConfigDialog <PoolSchedulingDialog
v-if="selectedProviderId" v-if="selectedProviderId"
v-model="showConfigDialog" v-model="showSchedulingDialog"
:provider-id="selectedProviderId" :provider-id="selectedProviderId"
:provider-type="selectedProviderData?.provider_type" :provider-type="selectedProviderType"
:current-config="selectedProviderConfig" :current-config="selectedProviderConfig"
:current-claude-config="selectedProviderData?.claude_code_advanced" :current-claude-config="selectedProviderClaudeConfig"
@saved="loadOverview" @saved="loadOverview"
/> />
<KeyFormDialog <KeyFormDialog
@@ -1132,7 +1151,7 @@ import { ref, computed, watch, onMounted, onBeforeUnmount } from 'vue'
import { import {
Search, Search,
Upload, Upload,
Settings, ChevronDown,
RefreshCw, RefreshCw,
Power, Power,
Database, Database,
@@ -1176,6 +1195,7 @@ import { useConfirm } from '@/composables/useConfirm'
import { parseApiError } from '@/utils/errorParser' import { parseApiError } from '@/utils/errorParser'
import { import {
getPoolOverview, getPoolOverview,
getPoolSchedulingPresets,
listPoolKeys, listPoolKeys,
clearPoolCooldown, clearPoolCooldown,
cleanupBannedPoolKeys, cleanupBannedPoolKeys,
@@ -1193,18 +1213,20 @@ import type {
PoolOverviewItem, PoolOverviewItem,
PoolKeyDetail, PoolKeyDetail,
PoolKeysPageResponse, PoolKeysPageResponse,
PoolPresetMeta,
} from '@/api/endpoints/pool' } from '@/api/endpoints/pool'
import type { EndpointAPIKey, PoolAdvancedConfig, ProviderWithEndpointsSummary } from '@/api/endpoints/types/provider' import type { ClaudeCodeAdvancedConfig, EndpointAPIKey, PoolAdvancedConfig, ProviderWithEndpointsSummary } from '@/api/endpoints/types/provider'
import { getProvider } from '@/api/endpoints' import { getProvider } from '@/api/endpoints'
import { useProxyNodesStore } from '@/stores/proxy-nodes' import { useProxyNodesStore } from '@/stores/proxy-nodes'
import PoolConfigDialog from '@/features/pool/components/PoolConfigDialog.vue' import PoolSchedulingDialog from '@/features/pool/components/PoolSchedulingDialog.vue'
import KeyAllowedModelsEditDialog from '@/features/providers/components/KeyAllowedModelsEditDialog.vue' import KeyAllowedModelsEditDialog from '@/features/providers/components/KeyAllowedModelsEditDialog.vue'
import KeyFormDialog from '@/features/providers/components/KeyFormDialog.vue' import KeyFormDialog from '@/features/providers/components/KeyFormDialog.vue'
import OAuthKeyEditDialog from '@/features/providers/components/OAuthKeyEditDialog.vue' import OAuthKeyEditDialog from '@/features/providers/components/OAuthKeyEditDialog.vue'
import OAuthAccountDialog from '@/features/providers/components/OAuthAccountDialog.vue' import OAuthAccountDialog from '@/features/providers/components/OAuthAccountDialog.vue'
import ProxyNodeSelect from '@/features/providers/components/ProxyNodeSelect.vue' import ProxyNodeSelect from '@/features/providers/components/ProxyNodeSelect.vue'
import { isAccountLevelBlockReason, cleanAccountBlockReason } from '@/utils/accountBlock'
const { success, error: showError } = useToast() const { success, error: showError, warning: showWarning } = useToast()
const { confirm } = useConfirm() const { confirm } = useConfirm()
const { copyToClipboard } = useClipboard() const { copyToClipboard } = useClipboard()
const { tick: countdownTick, start: startCountdownTimer } = useCountdownTimer() const { tick: countdownTick, start: startCountdownTimer } = useCountdownTimer()
@@ -1273,6 +1295,81 @@ const selectedProviderConfig = computed<PoolAdvancedConfig | null>(() => {
return (selectedProviderData.value as Record<string, unknown> | null)?.pool_advanced as PoolAdvancedConfig | null ?? null return (selectedProviderData.value as Record<string, unknown> | null)?.pool_advanced as PoolAdvancedConfig | null ?? null
}) })
const selectedProviderClaudeConfig = computed(() => {
return (selectedProviderData.value as Record<string, unknown> | null)?.claude_code_advanced as ClaudeCodeAdvancedConfig | null ?? null
})
const DEFAULT_PRESET_LABELS: Record<string, string> = {
lru: 'LRU',
free_team_first: 'Free/Team',
recent_refresh: '刷新优先',
quota_balanced: '额度均衡',
single_account: '单号优先',
}
const presetLabelsByName = ref<Record<string, string>>({ ...DEFAULT_PRESET_LABELS })
function normalizePresetName(value: unknown): string {
return String(value ?? '').trim().toLowerCase()
}
async function loadSchedulingPresetMetas(): Promise<void> {
try {
const metas = await getPoolSchedulingPresets()
const next: Record<string, string> = {}
for (const meta of metas as PoolPresetMeta[]) {
const name = normalizePresetName(meta.name)
if (!name) continue
const label = String(meta.label ?? '').trim()
next[name] = label || name
}
if (Object.keys(next).length > 0) {
presetLabelsByName.value = next
}
} catch {
presetLabelsByName.value = { ...DEFAULT_PRESET_LABELS }
}
}
const poolSchedulingLabel = computed(() => {
const cfg = selectedProviderConfig.value
const presets = Array.isArray(cfg?.scheduling_presets) ? cfg.scheduling_presets : []
const presetLabels = presetLabelsByName.value
if (presets.length > 0) {
// New format: object list with { preset, enabled }
const first = presets[0]
if (typeof first === 'object' && first !== null && 'preset' in first) {
const enabledLabels = (presets as Array<{ preset: string; enabled?: boolean }>)
.filter(p => p.enabled !== false)
.map(p => presetLabels[normalizePresetName(p.preset)])
.filter(Boolean)
return enabledLabels.length > 0 ? enabledLabels.join('+') : '无启用维度'
}
// Legacy string list format
if (typeof first === 'string') {
const labels = (presets as string[])
.map(p => presetLabels[normalizePresetName(p)])
.filter(Boolean)
if (labels.length > 0) return labels.join('+')
}
}
// Fallback: legacy scheduling_mode field
if (cfg?.scheduling_mode === 'multi_score') {
return '多维评分'
}
const lruEnabled = cfg?.lru_enabled !== false
const stickyTtl = Number(cfg?.sticky_session_ttl_seconds ?? 3600)
const stickyEnabled = Number.isFinite(stickyTtl) && stickyTtl > 0
if (lruEnabled && stickyEnabled) return 'LRU + 粘性'
if (lruEnabled) return 'LRU'
if (stickyEnabled) return '粘性'
return '随机'
})
const selectedProviderType = computed(() => { const selectedProviderType = computed(() => {
const fromDetail = String(selectedProviderData.value?.provider_type || '').trim().toLowerCase() const fromDetail = String(selectedProviderData.value?.provider_type || '').trim().toLowerCase()
if (fromDetail) return fromDetail if (fromDetail) return fromDetail
@@ -1330,11 +1427,11 @@ async function refresh() {
const keyPage = ref<PoolKeysPageResponse>({ total: 0, page: 1, page_size: 50, keys: [] }) const keyPage = ref<PoolKeysPageResponse>({ total: 0, page: 1, page_size: 50, keys: [] })
const keysLoading = ref(false) const keysLoading = ref(false)
const refreshingCurrentPageQuota = ref(false) const refreshingCurrentPageQuota = ref(false)
const queuedCurrentPageQuotaRefresh = ref(false)
const searchQuery = ref('') const searchQuery = ref('')
const statusFilter = ref('all') const statusFilter = ref('all')
const currentPage = ref(1) const currentPage = ref(1)
const pageSize = ref(50) const pageSize = ref(50)
const MANUAL_QUOTA_REFRESH_COOLDOWN_SECONDS = 5 * 60
const refreshingOAuthKeyId = ref<string | null>(null) const refreshingOAuthKeyId = ref<string | null>(null)
const revealedKeys = ref<Map<string, string>>(new Map()) const revealedKeys = ref<Map<string, string>>(new Map())
const recoveringHealthKeyId = ref<string | null>(null) const recoveringHealthKeyId = ref<string | null>(null)
@@ -1372,58 +1469,124 @@ const quotaRefreshSupported = computed(() => {
|| selectedProviderType.value === 'antigravity' || selectedProviderType.value === 'antigravity'
}) })
function getCurrentPageQuotaKeyIds(): string[] { const refreshCurrentPageLoading = computed(() => {
const ids: string[] = [] return keysLoading.value || refreshingCurrentPageQuota.value
})
function normalizeQuotaUpdatedAt(raw: number | null | undefined): number | null {
const value = Number(raw ?? 0)
if (!Number.isFinite(value) || value <= 0) return null
if (value > 1_000_000_000_000) {
return Math.floor(value / 1000)
}
return Math.floor(value)
}
const currentPageQuotaRefreshStats = computed(() => {
void countdownTick.value
const seen = new Set<string>() const seen = new Set<string>()
const eligibleIds: string[] = []
let cooledDownCount = 0
let minRemainingSeconds = 0
const nowSeconds = Math.floor(Date.now() / 1000)
for (const key of keyPage.value.keys) { for (const key of keyPage.value.keys) {
const id = String(key.key_id || '').trim() const id = String(key.key_id || '').trim()
if (!id || seen.has(id)) continue if (!id || seen.has(id)) continue
seen.add(id) seen.add(id)
ids.push(id) const updatedAt = normalizeQuotaUpdatedAt(key.quota_updated_at ?? null)
if (updatedAt == null) {
eligibleIds.push(id)
continue
} }
return ids const remaining = MANUAL_QUOTA_REFRESH_COOLDOWN_SECONDS - (nowSeconds - updatedAt)
if (remaining > 0) {
cooledDownCount += 1
if (minRemainingSeconds <= 0 || remaining < minRemainingSeconds) {
minRemainingSeconds = remaining
} }
continue
}
eligibleIds.push(id)
}
return {
total: seen.size,
eligibleIds,
cooledDownCount,
minRemainingSeconds,
}
})
async function refreshCurrentPageQuotaInBackground(options: { silent?: boolean } = {}) { async function refreshCurrentPageQuotaInBackground(
if (!selectedProviderId.value || !quotaRefreshSupported.value) return options: { silent?: boolean; reloadAfter?: boolean } = {},
): Promise<boolean> {
if (!selectedProviderId.value || !quotaRefreshSupported.value) return false
const providerId = selectedProviderId.value const providerId = selectedProviderId.value
const keyIds = getCurrentPageQuotaKeyIds() const quotaStats = currentPageQuotaRefreshStats.value
if (keyIds.length === 0) return if (quotaStats.eligibleIds.length === 0) {
if (!options.silent && quotaStats.total > 0 && quotaStats.cooledDownCount > 0) {
const waitText = quotaStats.minRemainingSeconds > 0
? formatTTL(quotaStats.minRemainingSeconds)
: '稍后'
showWarning(`当前页额度均在冷却中,请 ${waitText} 后再试`)
}
return false
}
if (refreshingCurrentPageQuota.value) { if (refreshingCurrentPageQuota.value) {
queuedCurrentPageQuotaRefresh.value = true return false
return
} }
refreshingCurrentPageQuota.value = true refreshingCurrentPageQuota.value = true
try { try {
const result = await refreshProviderQuota(providerId, keyIds) const result = await refreshProviderQuota(providerId, quotaStats.eligibleIds)
const successCount = Number(result.success || 0) const successCount = Number(result.success || 0)
const failedCount = Number(result.failed || 0) const failedCount = Number(result.failed || 0)
const skippedCount = Math.max(quotaStats.total - quotaStats.eligibleIds.length, 0)
// 刷新当前页数据,展示最新额度与状态 // 刷新当前页数据,展示最新额度与状态
if (selectedProviderId.value === providerId) { if (selectedProviderId.value === providerId && options.reloadAfter !== false) {
await loadKeys() await loadKeys()
} }
if (!options.silent) { if (!options.silent) {
success(`当前页额度刷新完成:成功 ${successCount},失败 ${failedCount}`) const skippedText = skippedCount > 0 ? `,冷却跳过 ${skippedCount}` : ''
success(`当前页额度刷新完成:成功 ${successCount},失败 ${failedCount}${skippedText}`)
} }
return true
} catch (err) { } catch (err) {
showError(parseApiError(err, '刷新当前页额度失败')) showError(parseApiError(err, '刷新当前页额度失败'))
return false
} finally { } finally {
refreshingCurrentPageQuota.value = false refreshingCurrentPageQuota.value = false
if (queuedCurrentPageQuotaRefresh.value) {
queuedCurrentPageQuotaRefresh.value = false
void refreshCurrentPageQuotaInBackground(options)
}
} }
} }
const refreshButtonTitle = computed(() => {
if (refreshCurrentPageLoading.value) return '刷新中...'
if (!selectedProviderId.value) return '刷新'
if (!quotaRefreshSupported.value) return '刷新数据'
const quotaStats = currentPageQuotaRefreshStats.value
if (quotaStats.total === 0) return '刷新数据和额度'
if (quotaStats.eligibleIds.length === 0 && quotaStats.cooledDownCount > 0) {
const waitText = quotaStats.minRemainingSeconds > 0
? formatTTL(quotaStats.minRemainingSeconds)
: '稍后'
return `刷新数据(额度冷却 ${waitText}`
}
if (quotaStats.cooledDownCount > 0) {
return `刷新数据和额度(可刷新 ${quotaStats.eligibleIds.length}/${quotaStats.total}`
}
return '刷新数据和额度'
})
async function refreshCurrentPage() { async function refreshCurrentPage() {
const quotaDidReload = await refreshCurrentPageQuotaInBackground({ reloadAfter: true })
if (!quotaDidReload) {
await refresh() await refresh()
} }
}
async function loadKeys() { async function loadKeys() {
if (!selectedProviderId.value) return if (!selectedProviderId.value) return
@@ -1807,11 +1970,13 @@ async function handleCleanupBannedKeys() {
// --- Dialogs --- // --- Dialogs ---
const showImportDialog = ref(false) const showImportDialog = ref(false)
const showConfigDialog = ref(false) const showSchedulingDialog = ref(false)
async function handleAccountDialogSaved() { async function handleAccountDialogSaved() {
showImportDialog.value = false showImportDialog.value = false
await Promise.all([loadKeys(), loadOverview()]) await Promise.all([loadKeys(), loadOverview()])
// 导入账号后补一次静默额度刷新,避免新账号在列表里暂无额度信息
await refreshCurrentPageQuotaInBackground({ silent: true })
} }
// --- Formatting --- // --- Formatting ---
@@ -1831,6 +1996,8 @@ function formatCooldownReason(reason: string): string {
type PoolStatusVariant = 'default' | 'secondary' | 'destructive' | 'outline' | 'success' | 'warning' | 'dark' type PoolStatusVariant = 'default' | 'secondary' | 'destructive' | 'outline' | 'success' | 'warning' | 'dark'
function getSchedulingStatus(key: PoolKeyDetail): 'available' | 'degraded' | 'blocked' { function getSchedulingStatus(key: PoolKeyDetail): 'available' | 'degraded' | 'blocked' {
if (getAccountAlertLabel(key)) return 'blocked'
const status = key.scheduling_status const status = key.scheduling_status
if (status === 'available' || status === 'degraded' || status === 'blocked') { if (status === 'available' || status === 'degraded' || status === 'blocked') {
return status return status
@@ -1845,6 +2012,9 @@ function getSchedulingStatus(key: PoolKeyDetail): 'available' | 'degraded' | 'bl
} }
function getSchedulingBadgeLabel(key: PoolKeyDetail): string { function getSchedulingBadgeLabel(key: PoolKeyDetail): string {
const accountAlert = getAccountAlertLabel(key)
if (accountAlert) return accountAlert
const rawLabel = String(key.scheduling_label || '').trim() const rawLabel = String(key.scheduling_label || '').trim()
if (rawLabel) { if (rawLabel) {
if (rawLabel === '禁用' || rawLabel === '停用') return '禁用' if (rawLabel === '禁用' || rawLabel === '停用') return '禁用'
@@ -1861,6 +2031,8 @@ function getSchedulingBadgeLabel(key: PoolKeyDetail): string {
} }
function getSchedulingBadgeVariant(key: PoolKeyDetail): PoolStatusVariant { function getSchedulingBadgeVariant(key: PoolKeyDetail): PoolStatusVariant {
if (getAccountAlertLabel(key)) return 'destructive'
const reason = key.scheduling_reason const reason = key.scheduling_reason
if (reason === 'manual_disabled') return 'dark' if (reason === 'manual_disabled') return 'dark'
if (reason === 'cooldown' || reason === 'circuit_open' || reason === 'cost_exhausted') return 'destructive' if (reason === 'cooldown' || reason === 'circuit_open' || reason === 'cost_exhausted') return 'destructive'
@@ -1875,6 +2047,9 @@ function getSchedulingBadgeVariant(key: PoolKeyDetail): PoolStatusVariant {
} }
function getSchedulingTitle(key: PoolKeyDetail): string { function getSchedulingTitle(key: PoolKeyDetail): string {
const accountAlertTitle = getAccountAlertTitle(key)
if (accountAlertTitle) return accountAlertTitle
if (key.scheduling_dimensions && key.scheduling_dimensions.length > 0) { if (key.scheduling_dimensions && key.scheduling_dimensions.length > 0) {
return key.scheduling_dimensions.map((item) => { return key.scheduling_dimensions.map((item) => {
const ttl = item.ttl_seconds && item.ttl_seconds > 0 ? ` (${formatTTL(item.ttl_seconds)})` : '' const ttl = item.ttl_seconds && item.ttl_seconds > 0 ? ` (${formatTTL(item.ttl_seconds)})` : ''
@@ -2048,6 +2223,41 @@ function getOAuthStatusTitle(key: PoolKeyDetail): string {
return `Token 剩余有效期: ${status.text}` return `Token 剩余有效期: ${status.text}`
} }
const _accountAlertCache = new WeakMap<PoolKeyDetail, string | null>()
function getAccountAlertLabel(key: PoolKeyDetail): string | null {
const cached = _accountAlertCache.get(key)
if (cached !== undefined) return cached
let result: string | null = null
const quotaText = String(key.account_quota || '').trim()
// 后端 _build_account_quota 返回的确切文本: "账号已封禁" / "访问受限"
if (quotaText === '账号已封禁' || quotaText === '封禁') result = '账号封禁'
else if (quotaText === '访问受限') result = '访问受限'
else if (isAccountLevelBlockReason(key.oauth_invalid_reason)) result = '账号异常'
_accountAlertCache.set(key, result)
return result
}
function getAccountAlertTitle(key: PoolKeyDetail): string {
const label = getAccountAlertLabel(key)
if (!label) return ''
const reason = String(key.oauth_invalid_reason || '').trim()
if (reason) {
if (isAccountLevelBlockReason(reason)) {
const cleaned = cleanAccountBlockReason(reason)
return cleaned ? `${label}: ${cleaned}` : label
}
return `${label}: ${reason}`
}
const quotaText = String(key.account_quota || '').trim()
if (quotaText) return `${label}: ${quotaText}`
return label
}
function normalizeQuotaLabel(label: string): string { function normalizeQuotaLabel(label: string): string {
const normalized = label.trim() const normalized = label.trim()
if (!normalized) return '额度' if (!normalized) return '额度'
@@ -2214,8 +2424,7 @@ function formatRelativeTime(isoStr: string): string {
// --- Init --- // --- Init ---
onMounted(async () => { onMounted(async () => {
startCountdownTimer() startCountdownTimer()
await loadOverview() await Promise.all([loadSchedulingPresetMetas(), loadOverview()])
void refreshCurrentPageQuotaInBackground({ silent: true })
}) })
onBeforeUnmount(() => { onBeforeUnmount(() => {

View File

@@ -28,7 +28,9 @@ from src.core.logger import logger
from src.database import get_db from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, Usage from src.models.database import Provider, ProviderAPIKey, Usage
from src.services.provider.pool import redis_ops as pool_redis from src.services.provider.pool import redis_ops as pool_redis
from src.services.provider.pool.account_state import resolve_pool_account_state
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.dimensions import get_preset_dimension_metas
from src.services.provider.pool.scheduling_dimensions import ( from src.services.provider.pool.scheduling_dimensions import (
PoolSchedulingSnapshot, PoolSchedulingSnapshot,
evaluate_pool_scheduling_dimensions, evaluate_pool_scheduling_dimensions,
@@ -47,6 +49,8 @@ from .schemas import (
PoolOverviewResponse, PoolOverviewResponse,
PoolSchedulingDimension, PoolSchedulingDimension,
PoolSchedulingReason, PoolSchedulingReason,
PresetDimensionMetaResponse,
PresetModeMetaResponse,
) )
router = APIRouter(prefix="/api/admin/pool", tags=["pool-management"]) router = APIRouter(prefix="/api/admin/pool", tags=["pool-management"])
@@ -68,6 +72,31 @@ async def pool_overview(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode) return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# ---------------------------------------------------------------------------
# GET /api/admin/pool/scheduling-presets
# ---------------------------------------------------------------------------
def _preset_mode_label(mode: str) -> str:
mapping = {
"free_only": "Free",
"team_only": "Team",
"both": "全部",
}
return mapping.get(mode, mode)
@router.get("/scheduling-presets", response_model=list[PresetDimensionMetaResponse])
async def list_scheduling_presets(
request: Request,
db: Session = Depends(get_db),
) -> list[PresetDimensionMetaResponse]:
"""Return scheduling preset definitions for frontend rendering."""
adapter = AdminListSchedulingPresetsAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# GET /api/admin/pool/{provider_id}/keys # GET /api/admin/pool/{provider_id}/keys
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -126,24 +155,6 @@ _COOLDOWN_REASON_LABELS: dict[str, str] = {
"server_error_500": "500 错误", "server_error_500": "500 错误",
} }
_ACCOUNT_BLOCK_REASON_KEYWORDS: tuple[str, ...] = (
"account_block",
"account blocked",
"account has been disabled",
"account disabled",
"organization has been disabled",
"organization_disabled",
"validation_required",
"verify your account",
"forbidden",
"suspended",
"封禁",
"封号",
"被封",
"访问被禁止",
"账号异常",
)
def _to_float(value: Any) -> float | None: def _to_float(value: Any) -> float | None:
if isinstance(value, bool): if isinstance(value, bool):
@@ -161,65 +172,15 @@ def _to_float(value: Any) -> float | None:
return None return None
def _is_truthy_flag(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, (int, float)):
return value != 0
if isinstance(value, str):
normalized = value.strip().lower()
return normalized in {"1", "true", "yes", "y"}
return False
def _is_known_banned_reason(reason: str | None) -> bool:
if not reason:
return False
text = str(reason).strip()
if not text:
return False
lowered = text.lower()
# 结构化账号级别封禁标记(如 [ACCOUNT_BLOCK] ...
try:
from src.services.provider.oauth_token import is_account_level_block
if is_account_level_block(text):
return True
except Exception:
pass
return any(keyword in lowered for keyword in _ACCOUNT_BLOCK_REASON_KEYWORDS)
def _is_known_banned_key(key: ProviderAPIKey, provider_type: str) -> bool: def _is_known_banned_key(key: ProviderAPIKey, provider_type: str) -> bool:
upstream_metadata = getattr(key, "upstream_metadata", None) from src.services.provider.pool.account_state import resolve_pool_account_state
normalized_provider = provider_type.strip().lower()
provider_bucket: dict[str, Any] | None = None
if isinstance(upstream_metadata, dict):
maybe_bucket = upstream_metadata.get(normalized_provider)
if isinstance(maybe_bucket, dict):
provider_bucket = maybe_bucket
if normalized_provider == "kiro" and provider_bucket: state = resolve_pool_account_state(
if _is_truthy_flag(provider_bucket.get("is_banned")): provider_type=provider_type,
return True upstream_metadata=getattr(key, "upstream_metadata", None),
if normalized_provider == "antigravity" and provider_bucket: oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
if _is_truthy_flag(provider_bucket.get("is_forbidden")): )
return True return state.blocked
for source in (provider_bucket, upstream_metadata):
if not isinstance(source, dict):
continue
if _is_truthy_flag(source.get("is_banned")):
return True
if _is_truthy_flag(source.get("is_forbidden")):
return True
if _is_truthy_flag(source.get("account_disabled")):
return True
return _is_known_banned_reason(getattr(key, "oauth_invalid_reason", None))
def _format_percent(value: float) -> str: def _format_percent(value: float) -> str:
@@ -536,6 +497,10 @@ def _format_cooldown_detail(raw: str | None) -> str | None:
def _build_pool_scheduling_state( def _build_pool_scheduling_state(
*, *,
is_active: bool, is_active: bool,
account_blocked: bool,
account_block_label: str | None,
account_block_reason: str | None,
latency_avg_ms: float | None,
cooldown_reason: str | None, cooldown_reason: str | None,
cooldown_ttl_seconds: int | None, cooldown_ttl_seconds: int | None,
circuit_breaker_open: bool, circuit_breaker_open: bool,
@@ -557,6 +522,10 @@ def _build_pool_scheduling_state(
"""Build unified scheduling state for frontend display.""" """Build unified scheduling state for frontend display."""
snapshot = PoolSchedulingSnapshot( snapshot = PoolSchedulingSnapshot(
is_active=is_active, is_active=is_active,
account_blocked=account_blocked,
account_block_label=account_block_label,
account_block_reason=account_block_reason,
latency_avg_ms=latency_avg_ms,
cooldown_reason=cooldown_reason, cooldown_reason=cooldown_reason,
cooldown_ttl_seconds=cooldown_ttl_seconds, cooldown_ttl_seconds=cooldown_ttl_seconds,
circuit_breaker_open=circuit_breaker_open, circuit_breaker_open=circuit_breaker_open,
@@ -652,6 +621,39 @@ async def cleanup_banned_keys(
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class AdminListSchedulingPresetsAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
items: list[PresetDimensionMetaResponse] = [
PresetDimensionMetaResponse(
name="lru",
label="LRU 轮转",
description="最久未使用的 Key 优先",
providers=[],
modes=None,
default_mode=None,
)
]
for meta in get_preset_dimension_metas():
modes = None
if meta.modes:
modes = [
PresetModeMetaResponse(value=mode, label=_preset_mode_label(mode))
for mode in meta.modes
]
items.append(
PresetDimensionMetaResponse(
name=meta.name,
label=meta.label,
description=meta.description,
providers=list(meta.providers),
modes=modes,
default_mode=meta.default_mode,
)
)
return items
class AdminPoolOverviewAdapter(AdminApiAdapter): class AdminPoolOverviewAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db db = context.db
@@ -813,20 +815,34 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
if pcfg and pcfg.lru_enabled if pcfg and pcfg.lru_enabled
else asyncio.sleep(0, result={}) else asyncio.sleep(0, result={})
) )
_latency_coro = (
pool_redis.batch_get_latency_avgs(pid, key_ids, pcfg.latency_window_seconds)
if pcfg and pcfg.scheduling_mode == "multi_score"
else asyncio.sleep(0, result={})
)
_cost_coro = ( _cost_coro = (
pool_redis.batch_get_cost_totals(pid, key_ids, pcfg.cost_window_seconds) pool_redis.batch_get_cost_totals(pid, key_ids, pcfg.cost_window_seconds)
if pcfg if pcfg
else asyncio.sleep(0, result={}) else asyncio.sleep(0, result={})
) )
cooldowns, cooldown_ttls, lru_scores, cost_totals, sticky_counts = await asyncio.gather( (
cooldowns,
cooldown_ttls,
lru_scores,
latency_avgs,
cost_totals,
sticky_counts,
) = await asyncio.gather(
pool_redis.batch_get_cooldowns(pid, key_ids), pool_redis.batch_get_cooldowns(pid, key_ids),
pool_redis.batch_get_cooldown_ttls(pid, key_ids), pool_redis.batch_get_cooldown_ttls(pid, key_ids),
_lru_coro, _lru_coro,
_latency_coro,
_cost_coro, _cost_coro,
pool_redis.batch_get_key_sticky_counts(pid, key_ids), pool_redis.batch_get_key_sticky_counts(pid, key_ids),
) )
else: else:
cooldowns, cooldown_ttls, lru_scores, cost_totals, sticky_counts = ( cooldowns, cooldown_ttls, lru_scores, latency_avgs, cost_totals, sticky_counts = (
{},
{}, {},
{}, {},
{}, {},
@@ -874,6 +890,13 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
) )
cost_usage = int(cost_totals.get(kid, 0) or 0) cost_usage = int(cost_totals.get(kid, 0) or 0)
cost_limit = pcfg.cost_limit_per_key_tokens if pcfg else None cost_limit = pcfg.cost_limit_per_key_tokens if pcfg else None
latency_avg_raw = latency_avgs.get(kid)
latency_avg_ms = float(latency_avg_raw) if latency_avg_raw is not None else None
account_state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(k, "upstream_metadata", None),
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None),
)
( (
scheduling_status, scheduling_status,
scheduling_reason, scheduling_reason,
@@ -886,6 +909,10 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
scheduling_dimensions, scheduling_dimensions,
) = _build_pool_scheduling_state( ) = _build_pool_scheduling_state(
is_active=bool(k.is_active), is_active=bool(k.is_active),
account_blocked=account_state.blocked,
account_block_label=account_state.label,
account_block_reason=account_state.reason,
latency_avg_ms=latency_avg_ms,
cooldown_reason=cd_reason, cooldown_reason=cd_reason,
cooldown_ttl_seconds=cd_ttl, cooldown_ttl_seconds=cd_ttl,
circuit_breaker_open=any_circuit_open, circuit_breaker_open=any_circuit_open,
@@ -955,9 +982,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
key_name=k.name or "", key_name=k.name or "",
is_active=bool(k.is_active), is_active=bool(k.is_active),
auth_type=str(getattr(k, "auth_type", "api_key") or "api_key"), auth_type=str(getattr(k, "auth_type", "api_key") or "api_key"),
oauth_expires_at=_derive_oauth_expires_at( oauth_expires_at=_derive_oauth_expires_at(k, auth_config=oauth_auth_config),
k, auth_config=oauth_auth_config
),
oauth_invalid_at=( oauth_invalid_at=(
int(k.oauth_invalid_at.timestamp()) int(k.oauth_invalid_at.timestamp())
if getattr(k, "oauth_invalid_at", None) if getattr(k, "oauth_invalid_at", None)

View File

@@ -29,6 +29,25 @@ class PoolOverviewResponse(BaseModel):
items: list[PoolOverviewItem] = Field(default_factory=list) items: list[PoolOverviewItem] = Field(default_factory=list)
# ---------------------------------------------------------------------------
# Scheduling presets metadata
# ---------------------------------------------------------------------------
class PresetModeMetaResponse(BaseModel):
value: str
label: str
class PresetDimensionMetaResponse(BaseModel):
name: str
label: str
description: str
providers: list[str] = Field(default_factory=list)
modes: list[PresetModeMetaResponse] | None = None
default_mode: str | None = None
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Paginated key list # Paginated key list
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------

View File

@@ -127,6 +127,80 @@ class FailoverRulesConfig(BaseModel):
) )
class ScoringWeightsConfig(BaseModel):
"""多维评分权重配置。"""
lru: float = Field(0.3, ge=0.0, le=1.0)
latency: float = Field(0.25, ge=0.0, le=1.0)
health: float = Field(0.2, ge=0.0, le=1.0)
cost_remaining: float = Field(0.25, ge=0.0, le=1.0)
def _allowed_pool_preset_names() -> set[str]:
from src.services.provider.pool.dimensions import get_preset_names
return get_preset_names() | {"lru"}
def _preset_mode_meta(name: str) -> tuple[set[str], str | None]:
from src.services.provider.pool.dimensions import get_preset_dimension
dim = get_preset_dimension(name)
if dim is None or not dim.modes:
return set(), None
ordered_modes = [str(mode).strip().lower() for mode in dim.modes if str(mode).strip()]
if not ordered_modes:
return set(), None
modes = set(ordered_modes)
default_mode = str(dim.default_mode or "").strip().lower()
if not default_mode or default_mode not in modes:
default_mode = ordered_modes[0]
return modes, default_mode
class SchedulingPresetItem(BaseModel):
"""调度预设条目(新格式:有序对象列表)。"""
preset: str
enabled: bool = True
mode: str | None = None
@field_validator("preset")
@classmethod
def validate_preset(cls, v: str) -> str:
normalized = v.strip().lower()
allowed = _allowed_pool_preset_names()
if normalized not in allowed:
raise ValueError(f"无效的 preset: {normalized}")
return normalized
@field_validator("mode")
@classmethod
def normalize_mode(cls, v: str | None) -> str | None:
if v is None:
return None
normalized = v.strip().lower()
return normalized or None
@model_validator(mode="after")
def validate_mode(self) -> "SchedulingPresetItem":
allowed_modes, default_mode = _preset_mode_meta(self.preset)
if not allowed_modes:
self.mode = None
return self
if self.mode is None:
self.mode = default_mode
return self
if self.mode not in allowed_modes:
raise ValueError(
f"preset={self.preset} 的 mode 必须是: {', '.join(sorted(allowed_modes))}"
)
return self
class PoolAdvancedConfig(BaseModel): class PoolAdvancedConfig(BaseModel):
"""通用号池配置(适用于所有 Provider 类型)。""" """通用号池配置(适用于所有 Provider 类型)。"""
@@ -148,7 +222,33 @@ class PoolAdvancedConfig(BaseModel):
le=100, le=100,
description="负载率阈值(%),超过时该 Key 被降权。默认 80", description="负载率阈值(%),超过时该 Key 被降权。默认 80",
) )
# 保留旧字段供向后兼容(新客户端不再发送)
lru_enabled: bool = Field(True, description="LRU 调度(优先选择最久未用的 Key") lru_enabled: bool = Field(True, description="LRU 调度(优先选择最久未用的 Key")
scheduling_mode: str | None = Field(
None,
pattern="^(lru|multi_score)$",
description="号池调度模式lru 或 multi_score",
)
scheduling_presets: list[SchedulingPresetItem] | list[str] | None = Field(
None,
description=(
"调度预设列表(新格式:对象列表 [{preset, enabled, mode}]"
"旧格式:字符串列表 ['quota_balanced', ...]"
),
)
scoring_weights: ScoringWeightsConfig | None = Field(None, description="多维评分权重")
latency_window_seconds: int | None = Field(
None,
ge=300,
le=86400,
description="延迟窗口(秒),仅 multi_score 生效",
)
latency_sample_limit: int | None = Field(
None,
ge=10,
le=200,
description="每个 Key 的延迟样本上限,仅 multi_score 生效",
)
cost_window_seconds: int | None = Field( cost_window_seconds: int | None = Field(
None, None,
ge=3600, ge=3600,

View File

@@ -60,7 +60,7 @@ class RequestDispatcher:
attempt_counter: int, attempt_counter: int,
max_attempts: int, max_attempts: int,
is_stream: bool = False, is_stream: bool = False,
) -> tuple[Any, str, str, str, str, str]: ) -> tuple[Any, str, str, str, str, str, int | None]:
""" """
执行请求并返回结果 执行请求并返回结果
@@ -81,7 +81,7 @@ class RequestDispatcher:
is_stream: 是否为流式请求 is_stream: 是否为流式请求
Returns: Returns:
(response, provider_name, candidate_record_id, provider_id, endpoint_id, key_id) (response, provider_name, candidate_record_id, provider_id, endpoint_id, key_id, ttfb_ms)
Raises: Raises:
ExecutionError: 执行失败时 ExecutionError: 执行失败时
@@ -144,6 +144,19 @@ class RequestDispatcher:
logger.debug(f" [{request_id}] 请求成功: Provider={provider_name}, 耗时={elapsed_ms}ms") logger.debug(f" [{request_id}] 请求成功: Provider={provider_name}, 耗时={elapsed_ms}ms")
# Non-stream requests don't have first-byte telemetry in this path.
# Use elapsed latency as a conservative fallback for pool latency sampling.
ttfb_ms: int | None = None
if not is_stream:
raw_ttfb = getattr(execution_result.response, "first_byte_time_ms", None)
try:
if raw_ttfb is not None:
ttfb_ms = max(int(raw_ttfb), 0)
except (TypeError, ValueError):
ttfb_ms = None
if ttfb_ms is None and elapsed_ms >= 0:
ttfb_ms = int(elapsed_ms)
return ( return (
execution_result.response, execution_result.response,
provider_name, provider_name,
@@ -151,4 +164,5 @@ class RequestDispatcher:
provider_id, provider_id,
endpoint_id, endpoint_id,
key_id, key_id,
ttfb_ms,
) )

View File

@@ -3,12 +3,18 @@
Re-exports the main public API for convenience. Re-exports the main public API for convenience.
""" """
from src.services.provider.pool.config import PoolConfig, UnschedulableRule, parse_pool_config from src.services.provider.pool.config import (
PoolConfig,
ScoringWeights,
UnschedulableRule,
parse_pool_config,
)
from src.services.provider.pool.manager import PoolManager from src.services.provider.pool.manager import PoolManager
__all__ = [ __all__ = [
"PoolConfig", "PoolConfig",
"PoolManager", "PoolManager",
"ScoringWeights",
"UnschedulableRule", "UnschedulableRule",
"parse_pool_config", "parse_pool_config",
] ]

View File

@@ -0,0 +1,189 @@
"""Pool account state helpers.
Provides a shared way to classify account-level hard-block states
from upstream metadata and OAuth invalid reasons.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
OAUTH_ACCOUNT_BLOCK_PREFIX = "[ACCOUNT_BLOCK] "
ACCOUNT_BLOCK_REASON_KEYWORDS: tuple[str, ...] = (
"account_block",
"account blocked",
"account has been disabled",
"account disabled",
"organization has been disabled",
"organization_disabled",
"validation_required",
"verify your account",
"suspended",
# Kiro quota refresher 写入的确切文本
"账户已封禁",
# Antigravity quota refresher 写入的确切文本
"账户访问被禁止",
"封禁",
"封号",
"被封",
"访问被禁止",
"账号异常",
)
@dataclass(frozen=True, slots=True)
class PoolAccountState:
"""Resolved account-level state for one key."""
blocked: bool
code: str | None = None # account_banned / account_forbidden / account_blocked
label: str | None = None
reason: str | None = None
def _is_truthy_flag(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, (int, float)):
return value != 0
if isinstance(value, str):
normalized = value.strip().lower()
return normalized in {"1", "true", "yes", "y"}
return False
def _clean_text(value: Any) -> str | None:
if not isinstance(value, str):
return None
text = value.strip()
return text or None
def _extract_reason(source: dict[str, Any] | None, *fields: str) -> str | None:
if not isinstance(source, dict):
return None
for field in fields:
text = _clean_text(source.get(field))
if text:
return text
return None
def _resolve_from_metadata(
provider_type: str | None,
upstream_metadata: Any,
) -> PoolAccountState | None:
if not isinstance(upstream_metadata, dict):
return None
normalized_provider = str(provider_type or "").strip().lower()
provider_bucket: dict[str, Any] | None = None
if normalized_provider:
maybe_bucket = upstream_metadata.get(normalized_provider)
if isinstance(maybe_bucket, dict):
provider_bucket = maybe_bucket
if (
normalized_provider == "kiro"
and provider_bucket
and _is_truthy_flag(provider_bucket.get("is_banned"))
):
reason = _extract_reason(provider_bucket, "ban_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_banned",
label="账号封禁",
reason=reason or "Kiro 账号已封禁",
)
if (
normalized_provider == "antigravity"
and provider_bucket
and _is_truthy_flag(provider_bucket.get("is_forbidden"))
):
reason = _extract_reason(provider_bucket, "forbidden_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_forbidden",
label="访问受限",
reason=reason or "Antigravity 账户访问受限",
)
for source in (provider_bucket, upstream_metadata):
if not isinstance(source, dict):
continue
if _is_truthy_flag(source.get("is_banned")):
reason = _extract_reason(source, "ban_reason", "forbidden_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_banned",
label="账号封禁",
reason=reason or "账号已封禁",
)
if _is_truthy_flag(source.get("is_forbidden")) or _is_truthy_flag(
source.get("account_disabled")
):
reason = _extract_reason(source, "forbidden_reason", "ban_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_forbidden",
label="访问受限",
reason=reason or "账号访问受限",
)
return None
def _resolve_from_oauth_invalid_reason(reason: str | None) -> PoolAccountState | None:
text = _clean_text(reason)
if not text:
return None
if text.startswith(OAUTH_ACCOUNT_BLOCK_PREFIX):
cleaned = text[len(OAUTH_ACCOUNT_BLOCK_PREFIX) :].strip()
return PoolAccountState(
blocked=True,
code="account_blocked",
label="账号异常",
reason=cleaned or "账号异常",
)
lowered = text.lower()
if any(keyword in lowered for keyword in ACCOUNT_BLOCK_REASON_KEYWORDS):
return PoolAccountState(
blocked=True,
code="account_blocked",
label="账号异常",
reason=text,
)
return None
def resolve_pool_account_state(
*,
provider_type: str | None,
upstream_metadata: Any,
oauth_invalid_reason: str | None,
) -> PoolAccountState:
"""Resolve account-level hard-block state for pool scheduling."""
from_metadata = _resolve_from_metadata(provider_type, upstream_metadata)
if from_metadata is not None:
return from_metadata
from_oauth = _resolve_from_oauth_invalid_reason(oauth_invalid_reason)
if from_oauth is not None:
return from_oauth
return PoolAccountState(blocked=False)
__all__ = [
"ACCOUNT_BLOCK_REASON_KEYWORDS",
"OAUTH_ACCOUNT_BLOCK_PREFIX",
"PoolAccountState",
"resolve_pool_account_state",
]

View File

@@ -6,6 +6,26 @@ from dataclasses import dataclass, field
from typing import Any from typing import Any
from src.core.logger import logger from src.core.logger import logger
from src.services.provider.pool.dimensions import get_preset_dimension, get_preset_names
@dataclass(frozen=True, slots=True)
class ScoringWeights:
"""Weights used by multi-score scheduling."""
lru: float = 0.3
latency: float = 0.25
health: float = 0.2
cost_remaining: float = 0.25
@dataclass(frozen=True, slots=True)
class SchedulingPreset:
"""Single scheduling preset item with enable/disable and optional sub-config."""
preset: str
enabled: bool = True
mode: str | None = None
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -32,8 +52,17 @@ class PoolConfig:
# -- Load-Aware Selection ------------------------------------------------- # -- Load-Aware Selection -------------------------------------------------
load_threshold_percent: int = 80 load_threshold_percent: int = 80
# -- LRU ------------------------------------------------------------------ # -- Scheduling (unified preset list) -------------------------------------
scheduling_presets: tuple[SchedulingPreset, ...] = (
SchedulingPreset(preset="lru", enabled=True),
)
# Derived from scheduling_presets at parse time (backward compat for consumers)
lru_enabled: bool = True lru_enabled: bool = True
scheduling_mode: str = "lru" # lru | multi_score
scoring_weights: ScoringWeights = field(default_factory=ScoringWeights)
latency_window_seconds: int = 3600
latency_sample_limit: int = 50
# -- Rolling-Window Cost Tracking ----------------------------------------- # -- Rolling-Window Cost Tracking -----------------------------------------
cost_window_seconds: int = 18000 # 5 hours cost_window_seconds: int = 18000 # 5 hours
@@ -122,11 +151,35 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
scoring_weights = _parse_scoring_weights(raw_advanced.get("scoring_weights"))
# Parse scheduling presets (new object-list format or legacy string-list)
presets = _parse_scheduling_presets_v2(
raw_advanced.get("scheduling_presets"),
legacy_mode=raw_advanced.get("scheduling_mode"),
legacy_lru=raw_advanced.get("lru_enabled"),
)
# Derive scheduling_mode and lru_enabled from the presets list
enabled = [p for p in presets if p.enabled]
lru_enabled = any(p.preset == "lru" for p in enabled)
non_lru_enabled = [p for p in enabled if p.preset != "lru"]
scheduling_mode = "multi_score" if non_lru_enabled else "lru"
strategies = list(_parse_strategies(raw_advanced.get("strategies")))
if scheduling_mode == "multi_score" and "multi_score" not in strategies:
strategies.append("multi_score")
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"), 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), scheduling_presets=presets,
lru_enabled=lru_enabled,
scheduling_mode=scheduling_mode,
scoring_weights=scoring_weights,
latency_window_seconds=_int_or("latency_window_seconds", 3600),
latency_sample_limit=_int_or("latency_sample_limit", 50),
cost_window_seconds=_int_or("cost_window_seconds", 18000), cost_window_seconds=_int_or("cost_window_seconds", 18000),
cost_limit_per_key_tokens=_opt_int("cost_limit_per_key_tokens"), cost_limit_per_key_tokens=_opt_int("cost_limit_per_key_tokens"),
cost_soft_threshold_percent=_int_or("cost_soft_threshold_percent", 80), cost_soft_threshold_percent=_int_or("cost_soft_threshold_percent", 80),
@@ -138,12 +191,134 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
stream_timeout_threshold=_int_or("stream_timeout_threshold", 3), stream_timeout_threshold=_int_or("stream_timeout_threshold", 3),
stream_timeout_window_seconds=_int_or("stream_timeout_window_seconds", 1800), stream_timeout_window_seconds=_int_or("stream_timeout_window_seconds", 1800),
stream_timeout_cooldown_seconds=_int_or("stream_timeout_cooldown_seconds", 300), stream_timeout_cooldown_seconds=_int_or("stream_timeout_cooldown_seconds", 300),
strategies=_parse_strategies(raw_advanced.get("strategies")), strategies=tuple(strategies),
) )
# ---------------------------------------------------------------------------
# Internal parsers
# ---------------------------------------------------------------------------
def _allowed_preset_names() -> set[str]:
return get_preset_names() | {"lru"}
def _get_preset_mode_meta(preset_name: str) -> tuple[tuple[str, ...], str | None]:
dim = get_preset_dimension(preset_name)
if dim is None or not dim.modes:
return (), None
modes = tuple(str(mode).strip().lower() for mode in dim.modes if str(mode).strip())
if not modes:
return (), None
raw_default = str(dim.default_mode or "").strip().lower()
default_mode = raw_default if raw_default in modes else modes[0]
return modes, default_mode
def _parse_strategies(raw: Any) -> tuple[str, ...]: def _parse_strategies(raw: Any) -> tuple[str, ...]:
"""Parse strategy names from config (list[str] -> tuple[str, ...]).""" """Parse strategy names from config (list[str] -> tuple[str, ...])."""
if not isinstance(raw, list): if not isinstance(raw, list):
return () return ()
return tuple(str(s) for s in raw if isinstance(s, str) and s) return tuple(str(s) for s in raw if isinstance(s, str) and s)
def _parse_scoring_weights(raw: Any) -> ScoringWeights:
"""Parse scoring weights with graceful fallback."""
if not isinstance(raw, dict):
return ScoringWeights()
def _float_or(value: Any, default: float) -> float:
try:
parsed = float(value)
except (TypeError, ValueError):
return default
return max(0.0, min(parsed, 1.0))
return ScoringWeights(
lru=_float_or(raw.get("lru"), 0.3),
latency=_float_or(raw.get("latency"), 0.25),
health=_float_or(raw.get("health"), 0.2),
cost_remaining=_float_or(raw.get("cost_remaining"), 0.25),
)
def _parse_scheduling_presets_v2(
raw: Any,
*,
legacy_mode: Any = None,
legacy_lru: Any = None,
) -> tuple[SchedulingPreset, ...]:
"""Parse scheduling presets, supporting both new and legacy formats.
New format::
[{"preset": "lru", "enabled": true},
{"preset": "free_team_first", "enabled": true, "mode": "free_only"},
...]
Legacy format::
["free_team_first", "recent_refresh"] (with separate scheduling_mode / lru_enabled)
"""
if isinstance(raw, list) and raw:
first = raw[0]
if isinstance(first, dict):
return _parse_preset_object_list(raw)
if isinstance(first, str):
return _convert_legacy_string_list(raw, legacy_mode, legacy_lru)
# No presets at all: derive from legacy fields
return _build_from_legacy_fields(legacy_mode, legacy_lru)
def _parse_preset_object_list(raw: list[Any]) -> tuple[SchedulingPreset, ...]:
"""Parse new-format object list into SchedulingPreset tuple."""
allowed = _allowed_preset_names()
ordered: list[SchedulingPreset] = []
seen: set[str] = set()
for item in raw:
if not isinstance(item, dict):
continue
name = str(item.get("preset", "")).strip().lower()
if name not in allowed or name in seen:
continue
seen.add(name)
enabled = bool(item.get("enabled", True))
mode: str | None = None
modes, default_mode = _get_preset_mode_meta(name)
if modes:
raw_mode = str(item.get("mode", default_mode) or "").strip().lower()
mode = raw_mode if raw_mode in modes else default_mode
ordered.append(SchedulingPreset(preset=name, enabled=enabled, mode=mode))
return tuple(ordered) if ordered else (SchedulingPreset(preset="lru", enabled=True),)
def _convert_legacy_string_list(
raw: list[Any],
legacy_mode: Any,
legacy_lru: Any,
) -> tuple[SchedulingPreset, ...]:
"""Convert legacy string list + mode/lru fields to new format."""
lru_enabled = legacy_lru if isinstance(legacy_lru, bool) else True
allowed_non_lru = _allowed_preset_names() - {"lru"}
items: list[SchedulingPreset] = [SchedulingPreset(preset="lru", enabled=lru_enabled)]
seen: set[str] = {"lru"}
for p in raw:
if not isinstance(p, str):
continue
name = p.strip().lower()
if name not in allowed_non_lru or name in seen:
continue
seen.add(name)
items.append(SchedulingPreset(preset=name, enabled=True))
return tuple(items)
def _build_from_legacy_fields(legacy_mode: Any, legacy_lru: Any) -> tuple[SchedulingPreset, ...]:
"""Build presets from legacy scheduling_mode / lru_enabled only."""
lru_enabled = legacy_lru if isinstance(legacy_lru, bool) else True
return (SchedulingPreset(preset="lru", enabled=lru_enabled),)

View File

@@ -0,0 +1,30 @@
"""Pool scheduling preset dimensions.
Importing this package registers all built-in preset dimensions.
"""
from __future__ import annotations
from . import free_team_first # noqa: F401
from . import quota_balanced # noqa: F401
from . import recent_refresh # noqa: F401
from . import single_account # noqa: F401
from .registry import (
PresetDimensionBase,
PresetDimensionMeta,
get_all_preset_dimensions,
get_preset_dimension,
get_preset_dimension_metas,
get_preset_names,
register_preset_dimension,
)
__all__ = [
"PresetDimensionBase",
"PresetDimensionMeta",
"get_all_preset_dimensions",
"get_preset_dimension",
"get_preset_dimension_metas",
"get_preset_names",
"register_preset_dimension",
]

View File

@@ -0,0 +1,227 @@
"""Shared helpers for pool preset dimensions."""
from __future__ import annotations
import math
import time
from typing import Any
def safe_float(value: Any) -> float | None:
try:
parsed = float(value)
except (TypeError, ValueError):
return None
if math.isnan(parsed) or math.isinf(parsed):
return None
return parsed
def safe_metadata(key_obj: Any) -> dict[str, Any]:
raw = getattr(key_obj, "upstream_metadata", None)
return raw if isinstance(raw, dict) else {}
def normalize_plan(value: Any) -> str | None:
if not isinstance(value, str):
return None
normalized = value.strip().lower()
return normalized or None
def rank_ascending(key_id: str, scores: dict[str, float], all_ids: list[str]) -> float:
"""Rank score within all IDs; lower value means better rank."""
if not all_ids:
return 0.0
valid_count = sum(1 for kid in all_ids if safe_float(scores.get(kid)) is not None)
if valid_count <= 0:
return 0.5
decorated: list[tuple[int, float, int, str]] = []
for idx, kid in enumerate(all_ids):
score_raw = safe_float(scores.get(kid))
if score_raw is None:
decorated.append((1, float("inf"), idx, kid))
else:
decorated.append((0, score_raw, idx, kid))
decorated.sort(key=lambda item: (item[0], item[1], item[2]))
rank_idx = 0
for idx, (_missing, _value, _order, kid) in enumerate(decorated):
if kid == key_id:
rank_idx = idx
break
n = len(all_ids)
if n <= 1:
return 0.0
return rank_idx / float(n - 1)
def rank_descending(key_id: str, scores: dict[str, float], all_ids: list[str]) -> float:
"""Rank score within all IDs; higher value means better rank."""
if not all_ids:
return 0.0
valid_count = sum(1 for kid in all_ids if safe_float(scores.get(kid)) is not None)
if valid_count <= 0:
return 0.5
decorated: list[tuple[int, float, int, str]] = []
for idx, kid in enumerate(all_ids):
score_raw = safe_float(scores.get(kid))
if score_raw is None:
decorated.append((1, float("inf"), idx, kid))
else:
# 排序时取负值使分值越大排名越靠前rank 越小)
decorated.append((0, -score_raw, idx, kid))
decorated.sort(key=lambda item: (item[0], item[1], item[2]))
rank_idx = 0
for idx, (_missing, _value, _order, kid) in enumerate(decorated):
if kid == key_id:
rank_idx = idx
break
n = len(all_ids)
if n <= 1:
return 0.0
return rank_idx / float(n - 1)
def extract_plan_type(key_obj: Any) -> str | None:
direct = normalize_plan(getattr(key_obj, "oauth_plan_type", None))
if direct:
return direct
metadata = safe_metadata(key_obj)
codex = metadata.get("codex")
if isinstance(codex, dict):
codex_plan = normalize_plan(codex.get("plan_type"))
if codex_plan:
return codex_plan
kiro = metadata.get("kiro")
if isinstance(kiro, dict):
subscription_title = normalize_plan(kiro.get("subscription_title"))
if subscription_title:
# Normalize common Kiro labels into free/team buckets used by free_team_first.
if "team" in subscription_title:
return "team"
if "free" in subscription_title:
return "free"
if "pro" in subscription_title:
return "pro"
if "plus" in subscription_title:
return "plus"
return subscription_title
return None
def extract_reset_seconds(key_obj: Any) -> float | None:
metadata = safe_metadata(key_obj)
candidates: list[float] = []
codex = metadata.get("codex")
if isinstance(codex, dict):
for field in ("secondary_reset_seconds", "primary_reset_seconds"):
parsed = safe_float(codex.get(field))
if parsed is None or parsed < 0:
continue
candidates.append(parsed)
kiro = metadata.get("kiro")
if isinstance(kiro, dict):
next_reset_at = safe_float(kiro.get("next_reset_at"))
if next_reset_at is not None and next_reset_at > 0:
candidates.append(max(0.0, next_reset_at - time.time()))
if not candidates:
return None
return min(candidates)
def extract_usage_ratio(key_obj: Any) -> float | None:
metadata = safe_metadata(key_obj)
codex = metadata.get("codex")
if isinstance(codex, dict):
codex_values: list[float] = []
for field in ("primary_used_percent", "secondary_used_percent"):
parsed = safe_float(codex.get(field))
if parsed is None:
continue
codex_values.append(max(0.0, min(parsed, 100.0)) / 100.0)
if codex_values:
return sum(codex_values) / len(codex_values)
kiro = metadata.get("kiro")
if isinstance(kiro, dict):
parsed = safe_float(kiro.get("usage_percentage"))
if parsed is not None:
return max(0.0, min(parsed, 100.0)) / 100.0
antigravity = metadata.get("antigravity")
if isinstance(antigravity, dict):
quota_by_model = antigravity.get("quota_by_model")
if isinstance(quota_by_model, dict):
usage_values: list[float] = []
for model_info in quota_by_model.values():
if not isinstance(model_info, dict):
continue
used_percent = safe_float(model_info.get("used_percent"))
if used_percent is None:
remaining_fraction = safe_float(model_info.get("remaining_fraction"))
if remaining_fraction is not None:
used_percent = (1.0 - remaining_fraction) * 100.0
if used_percent is None:
continue
usage_values.append(max(0.0, min(used_percent, 100.0)) / 100.0)
if usage_values:
return sum(usage_values) / len(usage_values)
return None
def plan_priority_score(plan_type: str | None, mode: str | None = None) -> float:
"""Score a key based on plan type and free_team_first mode."""
effective_mode = (mode or "both").strip().lower()
if effective_mode == "free_only":
if plan_type == "free":
return 0.0
if plan_type == "team":
return 0.5
elif effective_mode == "team_only":
if plan_type == "team":
return 0.0
if plan_type == "free":
return 0.5
else:
# "both" or unrecognized -> original behavior
if plan_type in {"free", "team"}:
return 0.0
if plan_type in {"enterprise", "business"}:
return 0.2
if plan_type in {"plus", "pro"}:
return 0.6
if plan_type:
return 0.7
return 0.8
__all__ = [
"extract_plan_type",
"extract_reset_seconds",
"extract_usage_ratio",
"normalize_plan",
"plan_priority_score",
"rank_ascending",
"rank_descending",
"safe_float",
"safe_metadata",
]

View File

@@ -0,0 +1,52 @@
"""free_team_first preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import extract_plan_type, plan_priority_score, rank_ascending
from .registry import PresetDimensionBase, register_preset_dimension
class FreeTeamFirstDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "free_team_first"
@property
def label(self) -> str:
return "Free/Team 优先"
@property
def description(self) -> str:
return "优先消耗低档账号(依赖 plan_type"
@property
def providers(self) -> tuple[str, ...]:
return ("codex", "kiro")
@property
def modes(self) -> tuple[str, ...] | None:
return ("free_only", "team_only", "both")
@property
def default_mode(self) -> str | None:
return "both"
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
plan_scores = {
kid: plan_priority_score(extract_plan_type(keys_by_id.get(kid)), mode)
for kid in all_key_ids
}
return rank_ascending(key_id, plan_scores, all_key_ids)
register_preset_dimension(FreeTeamFirstDimension())

View File

@@ -0,0 +1,41 @@
"""quota_balanced preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import extract_usage_ratio, rank_ascending
from .registry import PresetDimensionBase, register_preset_dimension
class QuotaBalancedDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "quota_balanced"
@property
def label(self) -> str:
return "额度平均"
@property
def description(self) -> str:
return "优先选额度消耗最少的账号"
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
usage_scores: dict[str, float] = {}
for kid in all_key_ids:
usage_ratio = extract_usage_ratio(keys_by_id.get(kid))
if usage_ratio is not None:
usage_scores[kid] = usage_ratio
return rank_ascending(key_id, usage_scores, all_key_ids)
register_preset_dimension(QuotaBalancedDimension())

View File

@@ -0,0 +1,45 @@
"""recent_refresh preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import extract_reset_seconds, rank_ascending
from .registry import PresetDimensionBase, register_preset_dimension
class RecentRefreshDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "recent_refresh"
@property
def label(self) -> str:
return "额度刷新优先"
@property
def description(self) -> str:
return "优先选即将刷新额度的账号"
@property
def providers(self) -> tuple[str, ...]:
return ("codex", "kiro")
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
reset_scores: dict[str, float] = {}
for kid in all_key_ids:
reset_seconds = extract_reset_seconds(keys_by_id.get(kid))
if reset_seconds is not None:
reset_scores[kid] = reset_seconds
return rank_ascending(key_id, reset_scores, all_key_ids)
register_preset_dimension(RecentRefreshDimension())

View File

@@ -0,0 +1,217 @@
"""Preset dimension registry for pool multi-score scheduling."""
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass
from threading import RLock
from typing import Any
@dataclass(frozen=True, slots=True)
class PresetDimensionMeta:
"""Serializable metadata for one preset dimension."""
name: str
label: str
description: str
providers: tuple[str, ...]
modes: tuple[str, ...] | None
default_mode: str | None
class PresetDimensionBase(ABC):
"""Base class of one pool scheduling preset dimension."""
@property
@abstractmethod
def name(self) -> str:
"""Stable preset key, e.g. ``free_team_first``."""
@property
@abstractmethod
def label(self) -> str:
"""User-facing label."""
@property
@abstractmethod
def description(self) -> str:
"""User-facing description."""
@property
def providers(self) -> tuple[str, ...]:
"""Supported provider types.
Empty tuple means the dimension is universal and applies to all providers.
"""
return ()
@property
def modes(self) -> tuple[str, ...] | None:
"""Optional sub-modes for this dimension."""
return None
@property
def default_mode(self) -> str | None:
"""Default mode when mode is omitted."""
return None
@abstractmethod
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
"""Compute normalized metric in [0, 1], lower is better."""
def is_applicable(self, provider_type: str) -> bool:
"""Return whether this dimension applies to the given provider type."""
if not self.providers:
return True
normalized = _normalize_name(provider_type)
return normalized in self.providers
def _normalize_name(value: Any) -> str:
if not isinstance(value, str):
return ""
return value.strip().lower()
def _normalize_names(values: tuple[str, ...] | list[str]) -> tuple[str, ...]:
normalized = [_normalize_name(item) for item in values]
return tuple(item for item in normalized if item)
_registry_lock = RLock()
_registry: dict[str, PresetDimensionBase] = {}
def register_preset_dimension(dim: PresetDimensionBase) -> None:
"""Register or replace one preset dimension by name."""
name = _normalize_name(dim.name)
if not name:
raise ValueError("preset dimension name must be a non-empty string")
providers = _normalize_names(dim.providers)
modes = _normalize_names(dim.modes or ())
default_mode = _normalize_name(dim.default_mode)
if modes and default_mode and default_mode not in modes:
raise ValueError(f"default_mode must be one of modes for preset '{name}'")
class _NormalizedDimension(PresetDimensionBase):
# Lightweight wrapper to keep normalized metadata while preserving compute logic.
def __init__(self, wrapped: PresetDimensionBase) -> None:
self._wrapped = wrapped
@property
def name(self) -> str:
return name
@property
def label(self) -> str:
return self._wrapped.label
@property
def description(self) -> str:
return self._wrapped.description
@property
def providers(self) -> tuple[str, ...]:
return providers
@property
def modes(self) -> tuple[str, ...] | None:
return modes or None
@property
def default_mode(self) -> str | None:
if not modes:
return None
if default_mode:
return default_mode
return modes[0]
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
return self._wrapped.compute_metric(
key_id=key_id,
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
mode=mode,
)
normalized = _NormalizedDimension(dim)
with _registry_lock:
_registry[name] = normalized
def get_preset_dimension(name: str) -> PresetDimensionBase | None:
"""Get one registered preset dimension by name."""
key = _normalize_name(name)
if not key:
return None
with _registry_lock:
return _registry.get(key)
def get_all_preset_dimensions() -> list[PresetDimensionBase]:
"""Get all registered preset dimensions in registration order."""
with _registry_lock:
return list(_registry.values())
def get_preset_names() -> set[str]:
"""Get all registered preset names."""
with _registry_lock:
return set(_registry.keys())
def get_preset_dimension_metas() -> list[PresetDimensionMeta]:
"""Get serializable metadata for all preset dimensions."""
metas: list[PresetDimensionMeta] = []
for dim in get_all_preset_dimensions():
metas.append(
PresetDimensionMeta(
name=dim.name,
label=dim.label,
description=dim.description,
providers=dim.providers,
modes=dim.modes,
default_mode=dim.default_mode,
)
)
return metas
__all__ = [
"PresetDimensionBase",
"PresetDimensionMeta",
"get_all_preset_dimensions",
"get_preset_dimension",
"get_preset_dimension_metas",
"get_preset_names",
"register_preset_dimension",
]

View File

@@ -0,0 +1,36 @@
"""single_account preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import rank_descending
from .registry import PresetDimensionBase, register_preset_dimension
class SingleAccountDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "single_account"
@property
def label(self) -> str:
return "单号优先"
@property
def description(self) -> str:
return "集中使用同一账号(反向 LRU"
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
return rank_descending(key_id, lru_scores, all_key_ids)
register_preset_dimension(SingleAccountDimension())

View File

@@ -0,0 +1,91 @@
"""In-process pool health score cache.
This cache avoids recomputing per-key health aggregation on every request.
It does not replace persistent health storage; source data still comes from
``ProviderAPIKey.health_by_format`` carried on key objects.
"""
from __future__ import annotations
import threading
import time
from typing import Any
_TTL_SECONDS = 30.0
_LOCK = threading.Lock()
_CACHE: dict[str, tuple[float, dict[str, float]]] = {}
def aggregate_health_score(health_by_format: Any) -> float:
"""Aggregate health score from ``health_by_format`` (lower-bound strategy)."""
if not isinstance(health_by_format, dict) or not health_by_format:
return 1.0
scores: list[float] = []
for item in health_by_format.values():
if not isinstance(item, dict):
continue
try:
score = float(item.get("health_score") or 1.0)
except (TypeError, ValueError):
score = 1.0
scores.append(max(0.0, min(score, 1.0)))
if not scores:
return 1.0
return min(scores)
def get_health_scores(provider_id: str, keys: list[Any]) -> dict[str, float]:
"""Return key health scores with per-provider TTL cache.
Uses incremental merge: if the cache is still valid but missing some keys,
only the missing keys are computed and merged into the existing cache entry.
"""
now = time.monotonic()
keys_by_id: dict[str, Any] = {}
for k in keys:
kid = str(getattr(k, "id", "") or "")
if kid:
keys_by_id[kid] = k
if not keys_by_id:
return {}
with _LOCK:
cached = _CACHE.get(provider_id)
if cached is not None:
expires_at, payload = cached
if now < expires_at:
missing_ids = [kid for kid in keys_by_id if kid not in payload]
if not missing_ids:
return {kid: payload[kid] for kid in keys_by_id}
# Compute only for missing keys, merge into existing cache
for kid in missing_ids:
payload[kid] = aggregate_health_score(
getattr(keys_by_id[kid], "health_by_format", None)
)
return {kid: payload[kid] for kid in keys_by_id}
fresh: dict[str, float] = {}
for kid, key in keys_by_id.items():
fresh[kid] = aggregate_health_score(getattr(key, "health_by_format", None))
with _LOCK:
_CACHE[provider_id] = (now + _TTL_SECONDS, fresh)
return dict(fresh)
def invalidate_provider_health_scores(provider_id: str) -> None:
"""Invalidate health-score cache for one provider."""
with _LOCK:
_CACHE.pop(provider_id, None)
def _clear_cache_for_tests() -> None:
with _LOCK:
_CACHE.clear()
__all__ = [
"aggregate_health_score",
"get_health_scores",
"invalidate_provider_health_scores",
]

View File

@@ -20,7 +20,9 @@ from typing import TYPE_CHECKING, Any, TypeVar
from src.core.logger import logger from src.core.logger import logger
from src.services.provider.pool import redis_ops from src.services.provider.pool import redis_ops
from src.services.provider.pool.account_state import resolve_pool_account_state
from src.services.provider.pool.config import PoolConfig from src.services.provider.pool.config import PoolConfig
from src.services.provider.pool.health_cache import get_health_scores
from src.services.provider.pool.trace import PoolCandidateTrace, PoolSchedulingTrace from src.services.provider.pool.trace import PoolCandidateTrace, PoolSchedulingTrace
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -31,11 +33,17 @@ if TYPE_CHECKING:
class PoolManager: class PoolManager:
"""Coordinate pool-level scheduling for a single Provider.""" """Coordinate pool-level scheduling for a single Provider."""
__slots__ = ("provider_id", "config") __slots__ = ("provider_id", "config", "provider_type")
def __init__(self, provider_id: str, config: PoolConfig) -> None: def __init__(
self,
provider_id: str,
config: PoolConfig,
provider_type: str | None = None,
) -> None:
self.provider_id = provider_id self.provider_id = provider_id
self.config = config self.config = config
self.provider_type = str(provider_type or "").strip().lower() or None
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Core scheduling: reorder candidate list for pool-aware selection # Core scheduling: reorder candidate list for pool-aware selection
@@ -53,7 +61,7 @@ class PoolManager:
1. **Sticky session hit** -- if the session is already bound to a key 1. **Sticky session hit** -- if the session is already bound to a key
and that key appears in *candidates* and is not in cooldown, move it and that key appears in *candidates* and is not in cooldown, move it
to position 0. to position 0.
2. **Filter** out keys in cooldown or cost-exhausted state (mark 2. **Filter** out keys in account-blocked / cooldown / cost-exhausted state (mark
``is_skipped``). ``is_skipped``).
3. **LRU sort** -- among remaining candidates at the same priority 3. **LRU sort** -- among remaining candidates at the same priority
level, sort by least-recently-used. level, sort by least-recently-used.
@@ -102,6 +110,13 @@ class PoolManager:
pid, session_uuid, self.config.sticky_session_ttl_seconds pid, session_uuid, self.config.sticky_session_ttl_seconds
) )
provider_type = self.provider_type
if provider_type is None and candidates:
first_provider = getattr(candidates[0], "provider", None)
provider_type = str(getattr(first_provider, "provider_type", "") or "").strip().lower()
if not provider_type:
provider_type = None
# --- 2. Batch fetch pool state (parallel) --------------------- # --- 2. Batch fetch pool state (parallel) ---------------------
all_key_ids = [str(c.key.id) for c in candidates] all_key_ids = [str(c.key.id) for c in candidates]
@@ -113,17 +128,26 @@ class PoolManager:
else None else None
) )
_lru_coro = redis_ops.get_lru_scores(pid, all_key_ids) if self.config.lru_enabled else None _lru_coro = redis_ops.get_lru_scores(pid, all_key_ids) if self.config.lru_enabled else None
_latency_coro = (
redis_ops.batch_get_latency_avgs(pid, all_key_ids, self.config.latency_window_seconds)
if self.config.scheduling_mode == "multi_score"
else None
)
# Gather all non-None coroutines in parallel. # Gather all non-None coroutines in parallel.
coros: list[Any] = [_cooldown_coro] coros: list[Any] = [_cooldown_coro]
_cost_idx = -1 _cost_idx = -1
_lru_idx = -1 _lru_idx = -1
_latency_idx = -1
if _cost_coro is not None: if _cost_coro is not None:
_cost_idx = len(coros) _cost_idx = len(coros)
coros.append(_cost_coro) coros.append(_cost_coro)
if _lru_coro is not None: if _lru_coro is not None:
_lru_idx = len(coros) _lru_idx = len(coros)
coros.append(_lru_coro) coros.append(_lru_coro)
if _latency_coro is not None:
_latency_idx = len(coros)
coros.append(_latency_coro)
gathered = await asyncio.gather(*coros) gathered = await asyncio.gather(*coros)
@@ -158,7 +182,29 @@ class PoolManager:
if _lru_idx >= 0: if _lru_idx >= 0:
lru_scores = gathered[_lru_idx] lru_scores = gathered[_lru_idx]
# Latency averages
latency_avgs: dict[str, float] = {}
if _latency_idx >= 0:
latency_avgs = gathered[_latency_idx]
# Health scores (TTL cached, no Redis round-trip) -- only needed for multi_score
health_scores: dict[str, float] = {}
if self.config.scheduling_mode == "multi_score":
health_scores = get_health_scores(pid, [c.key for c in candidates])
strategy_context.update(
{
"all_key_ids": all_key_ids,
"lru_scores": lru_scores,
"cost_totals": cost_totals,
"latency_avgs": latency_avgs,
"health_scores": health_scores,
"keys_by_id": {str(c.key.id): c.key for c in candidates},
}
)
# --- Strategy: compute_score ---------------------------------- # --- Strategy: compute_score ----------------------------------
custom_scores: dict[str, float] = {}
for strategy in strategies: for strategy in strategies:
if hasattr(strategy, "compute_score"): if hasattr(strategy, "compute_score"):
for kid in all_key_ids: for kid in all_key_ids:
@@ -169,7 +215,8 @@ class PoolManager:
context=strategy_context, context=strategy_context,
) )
if custom is not None: if custom is not None:
lru_scores[kid] = custom custom_scores[kid] = float(custom)
lru_scores[kid] = float(custom)
except Exception: except Exception:
pass pass
@@ -181,6 +228,11 @@ class PoolManager:
for c in candidates: for c in candidates:
kid = str(c.key.id) kid = str(c.key.id)
ct = PoolCandidateTrace(key_id=kid) ct = PoolCandidateTrace(key_id=kid)
ct.scoring_mode = self.config.scheduling_mode
ct.latency_avg_ms = float(latency_avgs.get(kid, 0.0) or 0.0)
ct.health_score = float(health_scores.get(kid, 1.0) or 1.0)
if kid in custom_scores:
ct.composite_score = float(custom_scores[kid])
# Already skipped upstream? # Already skipped upstream?
if c.is_skipped: if c.is_skipped:
@@ -191,6 +243,25 @@ class PoolManager:
continue continue
# Cooldown? # Cooldown?
account_state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(c.key, "upstream_metadata", None),
oauth_invalid_reason=getattr(c.key, "oauth_invalid_reason", None),
)
if account_state.blocked:
c.is_skipped = True
skip_reason = account_state.reason or account_state.label or "account blocked"
c.skip_reason = f"pool account blocked: {skip_reason}"
skipped.append(c)
ct.skipped = True
ct.skip_type = "account_blocked"
ct.account_block_code = account_state.code
ct.account_block_label = account_state.label
ct.account_block_reason = account_state.reason
_attach_pool_extra(c, ct)
trace.candidate_traces[kid] = ct
continue
cd_reason = cooldowns.get(kid) cd_reason = cooldowns.get(kid)
if cd_reason is not None: if cd_reason is not None:
c.is_skipped = True c.is_skipped = True
@@ -225,6 +296,9 @@ class PoolManager:
trace.sticky_session_used = True trace.sticky_session_used = True
else: else:
available.append(c) available.append(c)
if kid in custom_scores and self.config.scheduling_mode == "multi_score":
ct.reason = "multi_score"
else:
ct.reason = "lru" if lru_scores.get(kid, 0) > 0 else "random" ct.reason = "lru" if lru_scores.get(kid, 0) > 0 else "random"
ct.lru_score = lru_scores.get(kid, 0.0) ct.lru_score = lru_scores.get(kid, 0.0)
@@ -363,7 +437,7 @@ class PoolManager:
:class:`ProviderAPIKey` objects instead of candidates: :class:`ProviderAPIKey` objects instead of candidates:
1. Sticky session hit (if bound and still healthy). 1. Sticky session hit (if bound and still healthy).
2. Filter out keys in cooldown or cost-exhausted. 2. Filter out keys in account-blocked / cooldown / cost-exhausted.
3. LRU sort among remaining keys. 3. LRU sort among remaining keys.
4. Random tiebreak for identical LRU scores. 4. Random tiebreak for identical LRU scores.
5. Return the first available key, or ``None``. 5. Return the first available key, or ``None``.
@@ -390,22 +464,32 @@ class PoolManager:
else None else None
) )
_lru_coro = redis_ops.get_lru_scores(pid, key_ids) if self.config.lru_enabled else None _lru_coro = redis_ops.get_lru_scores(pid, key_ids) if self.config.lru_enabled else None
_latency_coro = (
redis_ops.batch_get_latency_avgs(pid, key_ids, self.config.latency_window_seconds)
if self.config.scheduling_mode == "multi_score"
else None
)
coros_sk: list[Any] = [_cooldown_coro] coros_sk: list[Any] = [_cooldown_coro]
_cost_idx_sk = -1 _cost_idx_sk = -1
_lru_idx_sk = -1 _lru_idx_sk = -1
_latency_idx_sk = -1
if _cost_coro is not None: if _cost_coro is not None:
_cost_idx_sk = len(coros_sk) _cost_idx_sk = len(coros_sk)
coros_sk.append(_cost_coro) coros_sk.append(_cost_coro)
if _lru_coro is not None: if _lru_coro is not None:
_lru_idx_sk = len(coros_sk) _lru_idx_sk = len(coros_sk)
coros_sk.append(_lru_coro) coros_sk.append(_lru_coro)
if _latency_coro is not None:
_latency_idx_sk = len(coros_sk)
coros_sk.append(_latency_coro)
gathered_sk = await asyncio.gather(*coros_sk) gathered_sk = await asyncio.gather(*coros_sk)
cooldowns = gathered_sk[0] cooldowns = gathered_sk[0]
cost_exhausted: set[str] = set() cost_exhausted: set[str] = set()
cost_totals: dict[str, int] = {}
if _cost_idx_sk >= 0: if _cost_idx_sk >= 0:
cost_totals = gathered_sk[_cost_idx_sk] cost_totals = gathered_sk[_cost_idx_sk]
for kid, total in cost_totals.items(): for kid, total in cost_totals.items():
@@ -416,9 +500,25 @@ class PoolManager:
if _lru_idx_sk >= 0: if _lru_idx_sk >= 0:
lru_scores = gathered_sk[_lru_idx_sk] lru_scores = gathered_sk[_lru_idx_sk]
latency_avgs: dict[str, float] = {}
if _latency_idx_sk >= 0:
latency_avgs = gathered_sk[_latency_idx_sk]
health_scores: dict[str, float] = {}
if self.config.scheduling_mode == "multi_score":
health_scores = get_health_scores(pid, keys)
# --- Strategy: compute_score ------------------------------------------ # --- Strategy: compute_score ------------------------------------------
strategies = _get_active_strategies(self.config) strategies = _get_active_strategies(self.config)
strategy_context: dict[str, Any] = {"session_uuid": session_uuid} strategy_context: dict[str, Any] = {
"session_uuid": session_uuid,
"all_key_ids": key_ids,
"lru_scores": lru_scores,
"cost_totals": cost_totals if _cost_idx_sk >= 0 else {},
"latency_avgs": latency_avgs,
"health_scores": health_scores,
"keys_by_id": {str(k.id): k for k in keys},
}
for strategy in strategies: for strategy in strategies:
if hasattr(strategy, "compute_score"): if hasattr(strategy, "compute_score"):
for kid in key_ids: for kid in key_ids:
@@ -429,7 +529,7 @@ class PoolManager:
context=strategy_context, context=strategy_context,
) )
if custom is not None: if custom is not None:
lru_scores[kid] = custom lru_scores[kid] = float(custom)
except Exception: except Exception:
pass pass
@@ -440,6 +540,14 @@ class PoolManager:
for k in keys: for k in keys:
kid = str(k.id) kid = str(k.id)
account_state = resolve_pool_account_state(
provider_type=self.provider_type,
upstream_metadata=getattr(k, "upstream_metadata", None),
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None),
)
if account_state.blocked:
continue
if cooldowns.get(kid) is not None: if cooldowns.get(kid) is not None:
continue continue
if kid in cost_exhausted: if kid in cost_exhausted:
@@ -480,6 +588,7 @@ class PoolManager:
session_uuid: str | None, session_uuid: str | None,
key_id: str, key_id: str,
tokens_used: int = 0, tokens_used: int = 0,
ttfb_ms: int | None = None,
) -> None: ) -> None:
"""Called after a successful upstream request.""" """Called after a successful upstream request."""
pid = self.provider_id pid = self.provider_id
@@ -500,6 +609,16 @@ class PoolManager:
pid, key_id, tokens_used, self.config.cost_window_seconds pid, key_id, tokens_used, self.config.cost_window_seconds
) )
# Record latency sample for multi-score scheduling.
if self.config.scheduling_mode == "multi_score" and ttfb_ms is not None and ttfb_ms >= 0:
await redis_ops.record_latency(
pid,
key_id,
ttfb_ms,
self.config.latency_window_seconds,
self.config.latency_sample_limit,
)
async def on_request_error( async def on_request_error(
self, self,
*, *,
@@ -570,6 +689,8 @@ def _get_active_strategies(config: PoolConfig) -> list[Any]:
if not config.strategies: if not config.strategies:
return [] return []
try: try:
# Import triggers built-in strategy registration via module-level side effects.
from src.services.provider.pool import strategies as _builtin_strategies # noqa: F401
from src.services.provider.pool.strategy import get_active_strategies from src.services.provider.pool.strategy import get_active_strategies
return get_active_strategies(config.strategies) return get_active_strategies(config.strategies)

View File

@@ -10,6 +10,7 @@ ap:{pid}:sticky:{session_uuid} STRING -> key_id (TTL: config)
ap:{pid}:lru ZSET member=key_id, score=unix_ts ap:{pid}:lru ZSET member=key_id, score=unix_ts
ap:{pid}:cooldown:{key_id} STRING -> reason (TTL: error-specific) ap:{pid}:cooldown:{key_id} STRING -> reason (TTL: error-specific)
ap:{pid}:cost:{key_id} ZSET member=req_id, score=unix_ts ap:{pid}:cost:{key_id} ZSET member=req_id, score=unix_ts
ap:{pid}:latency:{key_id} ZSET member=req_id:ttfb_ms, score=unix_ts
provider_oauth_token_cache:{key_id} STRING -> access_token (TTL: expires - 60) provider_oauth_token_cache:{key_id} STRING -> access_token (TTL: expires - 60)
""" """
@@ -44,6 +45,10 @@ def _cost_key(provider_id: str, key_id: str) -> str:
return f"{PREFIX}:{provider_id}:cost:{key_id}" return f"{PREFIX}:{provider_id}:cost:{key_id}"
def _latency_key(provider_id: str, key_id: str) -> str:
return f"{PREFIX}:{provider_id}:latency:{key_id}"
def _oauth_cache_key(key_id: str) -> str: def _oauth_cache_key(key_id: str) -> str:
return f"provider_oauth_token_cache:{key_id}" return f"provider_oauth_token_cache:{key_id}"
@@ -91,6 +96,32 @@ end
return total return total
""" """
# Latency window cleanup + average in a single round-trip.
# KEYS[1] = latency zset key, ARGV[1] = window_start timestamp
# Returns nil when there are no samples, or avg(ms) as number.
_LATENCY_WINDOW_AVG_LUA = """
local key = KEYS[1]
local window_start = tonumber(ARGV[1])
redis.call("ZREMRANGEBYSCORE", key, "-inf", window_start)
local members = redis.call("ZRANGEBYSCORE", key, window_start, "+inf")
local total = 0
local count = 0
for _, m in ipairs(members) do
local colon = string.find(m, ":", 1, true)
if colon then
local n = tonumber(string.sub(m, colon + 1))
if n then
total = total + n
count = count + 1
end
end
end
if count == 0 then
return nil
end
return total / count
"""
async def _get_redis() -> "aioredis.Redis | None": async def _get_redis() -> "aioredis.Redis | None":
return await get_redis_client(require_redis=False) return await get_redis_client(require_redis=False)
@@ -337,6 +368,64 @@ async def batch_get_cost_totals(
return {k: 0 for k in key_ids} return {k: 0 for k in key_ids}
async def record_latency(
provider_id: str,
key_id: str,
ttfb_ms: int,
window_seconds: int,
sample_limit: int,
) -> None:
"""Record one TTFB sample with rolling-window cleanup."""
redis = await _get_redis()
if redis is None:
return
try:
now = time.time()
latency_k = _latency_key(provider_id, key_id)
sample = max(int(ttfb_ms), 0)
member = f"{uuid.uuid4().hex}:{sample}"
pipe = redis.pipeline()
pipe.zadd(latency_k, {member: now})
window_start = now - max(int(window_seconds), 1)
pipe.zremrangebyscore(latency_k, "-inf", window_start)
capped_limit = max(int(sample_limit), 1)
pipe.zremrangebyrank(latency_k, 0, -(capped_limit + 1))
pipe.expire(latency_k, max(int(window_seconds), 1) + 600)
await pipe.execute()
except Exception:
logger.debug("Pool: latency ADD failed for key {}", key_id[:8])
async def batch_get_latency_avgs(
provider_id: str,
key_ids: list[str],
window_seconds: int,
) -> dict[str, float]:
"""Batch-fetch latency averages (ms) for keys in a rolling window."""
redis = await _get_redis()
if redis is None:
return {}
try:
now = time.time()
window_start = now - max(int(window_seconds), 1)
pipe = redis.pipeline()
for kid in key_ids:
pipe.eval(_LATENCY_WINDOW_AVG_LUA, 1, _latency_key(provider_id, kid), str(window_start))
results = await pipe.execute()
out: dict[str, float] = {}
for kid, val in zip(key_ids, results):
if val is None:
continue
try:
out[kid] = float(val)
except Exception:
continue
return out
except Exception:
logger.debug("Pool: batch latency AVG failed for provider {}", provider_id[:8])
return {}
async def clear_cost(provider_id: str, key_id: str) -> None: async def clear_cost(provider_id: str, key_id: str) -> None:
redis = await _get_redis() redis = await _get_redis()
if redis is None: if redis is None:

View File

@@ -25,6 +25,10 @@ class PoolSchedulingSnapshot:
cost_limit: int | None cost_limit: int | None
cost_soft_threshold_percent: int = 80 cost_soft_threshold_percent: int = 80
health_score: float = 1.0 health_score: float = 1.0
latency_avg_ms: float | None = None
account_blocked: bool = False
account_block_label: str | None = None
account_block_reason: str | None = None
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -67,6 +71,44 @@ class PoolSchedulingDimension(Protocol):
"""Evaluate one dimension from snapshot.""" """Evaluate one dimension from snapshot."""
@dataclass(frozen=True, slots=True)
class _AccountStateDimension:
code: str = "account_state"
label: str = "账号状态"
source: str = "policy"
weight: int = 10
def evaluate(self, snapshot: PoolSchedulingSnapshot) -> PoolSchedulingDimensionResult:
if not snapshot.account_blocked:
return PoolSchedulingDimensionResult(
code=self.code,
label=self.label,
source=self.source,
weight=self.weight,
status="ok",
score=1.0,
)
blocked_label = snapshot.account_block_label or "账号异常"
if blocked_label == "账号封禁":
blocked_code = "account_banned"
elif blocked_label == "访问受限":
blocked_code = "account_forbidden"
else:
blocked_code = "account_blocked"
return PoolSchedulingDimensionResult(
code=blocked_code,
label=blocked_label,
source=self.source,
weight=self.weight,
status="blocked",
blocking=True,
score=0.0,
detail=snapshot.account_block_reason or blocked_label,
)
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class _ManualEnableDimension: class _ManualEnableDimension:
code: str = "manual_disabled" code: str = "manual_disabled"
@@ -264,6 +306,59 @@ class _HealthDimension:
) )
@dataclass(frozen=True, slots=True)
class _LatencyDimension:
code: str = "latency"
label: str = "延迟"
source: str = "runtime"
weight: int = 3
def evaluate(self, snapshot: PoolSchedulingSnapshot) -> PoolSchedulingDimensionResult:
latency = snapshot.latency_avg_ms
if latency is None:
return PoolSchedulingDimensionResult(
code=self.code,
label=self.label,
source=self.source,
weight=self.weight,
status="ok",
score=1.0,
detail="-",
)
value = max(float(latency), 0.0)
detail = f"{value:.0f}ms"
if value >= 3000:
return PoolSchedulingDimensionResult(
code="latency_high",
label="延迟偏高",
source=self.source,
weight=self.weight,
status="degraded",
score=0.5,
detail=detail,
)
if value >= 1200:
return PoolSchedulingDimensionResult(
code="latency_slow",
label="延迟较慢",
source=self.source,
weight=self.weight,
status="degraded",
score=0.72,
detail=detail,
)
return PoolSchedulingDimensionResult(
code=self.code,
label=self.label,
source=self.source,
weight=self.weight,
status="ok",
score=1.0,
detail=detail,
)
_POOL_DIMENSION_REGISTRY: dict[str, PoolSchedulingDimension] = {} _POOL_DIMENSION_REGISTRY: dict[str, PoolSchedulingDimension] = {}
_POOL_DIMENSION_ORDER: list[str] = [] _POOL_DIMENSION_ORDER: list[str] = []
@@ -363,10 +458,12 @@ def summarize_pool_scheduling_dimensions(
def _register_default_dimensions() -> None: def _register_default_dimensions() -> None:
register_pool_scheduling_dimension("account_state", _AccountStateDimension())
register_pool_scheduling_dimension("manual", _ManualEnableDimension()) register_pool_scheduling_dimension("manual", _ManualEnableDimension())
register_pool_scheduling_dimension("cooldown", _CooldownDimension()) register_pool_scheduling_dimension("cooldown", _CooldownDimension())
register_pool_scheduling_dimension("circuit", _CircuitBreakerDimension()) register_pool_scheduling_dimension("circuit", _CircuitBreakerDimension())
register_pool_scheduling_dimension("cost", _CostDimension()) register_pool_scheduling_dimension("cost", _CostDimension())
register_pool_scheduling_dimension("latency", _LatencyDimension())
register_pool_scheduling_dimension("health", _HealthDimension()) register_pool_scheduling_dimension("health", _HealthDimension())

View File

@@ -0,0 +1,8 @@
"""Built-in pool strategies."""
# Import side effects: register built-in strategies.
import src.services.provider.pool.dimensions # noqa: F401
from . import multi_score # noqa: F401
__all__ = ["multi_score"]

View File

@@ -0,0 +1,178 @@
"""Multi-dimension pool scoring strategy."""
from __future__ import annotations
from typing import Any
import src.services.provider.pool.dimensions # noqa: F401
from src.services.provider.pool.dimensions import get_preset_dimension, get_preset_names
from src.services.provider.pool.dimensions._helpers import rank_ascending, safe_float
from src.services.provider.pool.strategy import register_pool_strategy
# When LRU is enabled alongside presets, this fraction of the final score
# comes from the LRU rank (tiebreaker to avoid same-score collisions).
_LRU_BLEND_FACTOR = 0.04
# Positional weight decay factor: weight = 1 / (1 + DECAY * index).
_POSITIONAL_DECAY = 0.6
def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None], ...]:
"""Extract enabled (preset_name, mode) tuples from config.scheduling_presets.
Supports both new SchedulingPreset objects and legacy string lists.
Excludes ``lru`` since LRU is handled separately as a blend factor.
"""
raw = getattr(config, "scheduling_presets", ())
if not isinstance(raw, (list, tuple)):
return ()
allowed = get_preset_names() | {"lru"}
ordered: list[tuple[str, str | None]] = []
seen: set[str] = set()
for item in raw:
preset_name: str | None = None
enabled = True
mode: str | None = None
if hasattr(item, "preset"):
preset_name = str(getattr(item, "preset", "")).strip().lower()
enabled = bool(getattr(item, "enabled", True))
raw_mode = getattr(item, "mode", None)
if isinstance(raw_mode, str):
mode = raw_mode.strip().lower() or None
elif isinstance(item, str):
preset_name = item.strip().lower()
else:
continue
if not preset_name or preset_name not in allowed or preset_name in seen:
continue
if not enabled:
continue
if preset_name == "lru":
continue
seen.add(preset_name)
ordered.append((preset_name, mode))
return tuple(ordered)
class MultiScoreStrategy:
name = "multi_score"
def compute_score(
self,
*,
key_id: str,
config: Any,
context: dict[str, Any],
) -> float | None:
mode = str(getattr(config, "scheduling_mode", "lru") or "lru").strip().lower()
if mode != "multi_score":
return None
all_key_ids = [str(k) for k in (context.get("all_key_ids") or []) if str(k)]
if not all_key_ids:
return None
lru_scores = context.get("lru_scores", {})
if not isinstance(lru_scores, dict):
lru_scores = {}
latency_avgs = context.get("latency_avgs", {})
if not isinstance(latency_avgs, dict):
latency_avgs = {}
health_scores = context.get("health_scores", {})
if not isinstance(health_scores, dict):
health_scores = {}
cost_totals = context.get("cost_totals", {})
if not isinstance(cost_totals, dict):
cost_totals = {}
keys_by_id = context.get("keys_by_id", {})
if not isinstance(keys_by_id, dict):
keys_by_id = {}
presets = _normalize_presets_from_config(config)
lru_enabled = bool(getattr(config, "lru_enabled", True))
if presets:
return self._compute_preset_score(
key_id=key_id,
all_key_ids=all_key_ids,
presets=presets,
lru_enabled=lru_enabled,
lru_scores=lru_scores,
keys_by_id=keys_by_id,
)
weights = getattr(config, "scoring_weights", None)
w_lru = safe_float(getattr(weights, "lru", 0.3)) or 0.0
w_latency = safe_float(getattr(weights, "latency", 0.25)) or 0.0
w_health = safe_float(getattr(weights, "health", 0.2)) or 0.0
w_cost = safe_float(getattr(weights, "cost_remaining", 0.25)) or 0.0
lru_rank = rank_ascending(key_id, lru_scores, all_key_ids)
latency_rank = rank_ascending(key_id, latency_avgs, all_key_ids)
health_raw = safe_float(health_scores.get(key_id))
if health_raw is None:
health_raw = 1.0
health_norm = 1.0 - max(0.0, min(health_raw, 1.0))
cost_limit = getattr(config, "cost_limit_per_key_tokens", None)
used = safe_float(cost_totals.get(key_id)) or 0.0
if cost_limit is None or int(cost_limit) <= 0:
cost_norm = 0.0
else:
cost_norm = max(0.0, min(used / float(cost_limit), 1.0))
return (
w_lru * lru_rank
+ w_latency * latency_rank
+ w_health * health_norm
+ w_cost * cost_norm
)
def _compute_preset_score(
self,
*,
key_id: str,
all_key_ids: list[str],
presets: tuple[tuple[str, str | None], ...],
lru_enabled: bool,
lru_scores: dict[str, Any],
keys_by_id: dict[str, Any],
) -> float:
lru_rank_asc = rank_ascending(key_id, lru_scores, all_key_ids)
weighted_sum = 0.0
weight_sum = 0.0
for idx, (preset_name, mode) in enumerate(presets):
metric = 0.5
dim = get_preset_dimension(preset_name)
if dim is not None:
metric = dim.compute_metric(
key_id=key_id,
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
mode=mode,
)
weight = 1.0 / (1.0 + _POSITIONAL_DECAY * idx)
weighted_sum += metric * weight
weight_sum += weight
if weight_sum <= 0:
return lru_rank_asc
lru_blend = _LRU_BLEND_FACTOR if lru_enabled else 0.0
preset_blend = 1.0 - lru_blend
blended = (weighted_sum / weight_sum) * preset_blend + lru_rank_asc * lru_blend
return max(0.0, min(blended, 1.0))
register_pool_strategy("multi_score", MultiScoreStrategy())
__all__ = [
"MultiScoreStrategy",
]

View File

@@ -23,9 +23,16 @@ class PoolCandidateTrace:
cost_limit: int | None = None cost_limit: int | None = None
cost_soft_threshold: bool = False cost_soft_threshold: bool = False
skipped: bool = False skipped: bool = False
skip_type: str | None = None # cooldown / cost_exhausted skip_type: str | None = None # cooldown / cost_exhausted / account_blocked / upstream
cooldown_reason: str | None = None cooldown_reason: str | None = None
cooldown_ttl: int | None = None cooldown_ttl: int | None = None
account_block_code: str | None = None
account_block_label: str | None = None
account_block_reason: str | None = None
latency_avg_ms: float = 0.0
health_score: float = 1.0
composite_score: float = 0.0
scoring_mode: str = "lru"
def to_extra_data(self) -> dict[str, Any]: def to_extra_data(self) -> dict[str, Any]:
"""Build dict to merge into ``RequestCandidate.extra_data``.""" """Build dict to merge into ``RequestCandidate.extra_data``."""
@@ -35,8 +42,16 @@ class PoolCandidateTrace:
skip_info["cooldown_reason"] = self.cooldown_reason skip_info["cooldown_reason"] = self.cooldown_reason
if self.cooldown_ttl is not None: if self.cooldown_ttl is not None:
skip_info["cooldown_ttl"] = self.cooldown_ttl skip_info["cooldown_ttl"] = self.cooldown_ttl
if self.account_block_code is not None:
skip_info["account_block_code"] = self.account_block_code
if self.account_block_label is not None:
skip_info["account_block_label"] = self.account_block_label
if self.account_block_reason is not None:
skip_info["account_block_reason"] = self.account_block_reason
if self.cost_window_usage: if self.cost_window_usage:
skip_info["cost_window_usage"] = self.cost_window_usage skip_info["cost_window_usage"] = self.cost_window_usage
if self.scoring_mode:
skip_info["scoring_mode"] = self.scoring_mode
return {"pool_skip": skip_info} return {"pool_skip": skip_info}
sel: dict[str, Any] = {"reason": self.reason} sel: dict[str, Any] = {"reason": self.reason}
@@ -50,6 +65,14 @@ class PoolCandidateTrace:
sel["cost_limit"] = self.cost_limit sel["cost_limit"] = self.cost_limit
if self.cost_soft_threshold: if self.cost_soft_threshold:
sel["cost_soft_threshold"] = True sel["cost_soft_threshold"] = True
if self.latency_avg_ms > 0:
sel["latency_avg_ms"] = round(self.latency_avg_ms, 2)
if self.health_score < 1.0:
sel["health_score"] = round(self.health_score, 4)
if self.reason == "multi_score":
sel["composite_score"] = round(self.composite_score, 6)
if self.scoring_mode:
sel["scoring_mode"] = self.scoring_mode
return {"pool_selection": sel} return {"pool_selection": sel}
@@ -72,6 +95,7 @@ class PoolSchedulingTrace:
"""Build compact dict for ``Usage.request_metadata["pool_summary"]``.""" """Build compact dict for ``Usage.request_metadata["pool_summary"]``."""
skipped_cooldown = 0 skipped_cooldown = 0
skipped_cost = 0 skipped_cost = 0
skipped_account_blocked = 0
attempted = 0 attempted = 0
for t in self.candidate_traces.values(): for t in self.candidate_traces.values():
if t.skipped: if t.skipped:
@@ -79,6 +103,8 @@ class PoolSchedulingTrace:
skipped_cooldown += 1 skipped_cooldown += 1
elif t.skip_type == "cost_exhausted": elif t.skip_type == "cost_exhausted":
skipped_cost += 1 skipped_cost += 1
elif t.skip_type == "account_blocked":
skipped_account_blocked += 1
if attempted_key_ids is None: if attempted_key_ids is None:
# Backward-compatible behavior: count all schedulable keys. # Backward-compatible behavior: count all schedulable keys.
@@ -101,6 +127,7 @@ class PoolSchedulingTrace:
"attempted": attempted, "attempted": attempted,
"skipped_cooldown": skipped_cooldown, "skipped_cooldown": skipped_cooldown,
"skipped_cost": skipped_cost, "skipped_cost": skipped_cost,
"skipped_account_blocked": skipped_account_blocked,
"sticky_session": self.sticky_session_used, "sticky_session": self.sticky_session_used,
} }
if success_key_id: if success_key_id:

View File

@@ -235,7 +235,7 @@ class TaskService:
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)
manager = PoolManager(provider_id, pool_cfg) manager = PoolManager(provider_id, pool_cfg, provider_type=provider_type)
candidate_keys = list(candidate.pool_keys or []) candidate_keys = list(candidate.pool_keys or [])
if not candidate_keys and getattr(candidate, "key", None) is not None: if not candidate_keys and getattr(candidate, "key", None) is not None:
@@ -334,6 +334,7 @@ class TaskService:
async def _pool_on_success( async def _pool_on_success(
candidate: Any, candidate: Any,
request_body: dict[str, Any] | None, request_body: dict[str, Any] | None,
ttfb_ms: int | None = None,
) -> None: ) -> None:
"""Notify the pool manager about a successful request (sticky + LRU).""" """Notify the pool manager about a successful request (sticky + LRU)."""
try: try:
@@ -354,10 +355,11 @@ class TaskService:
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(provider_id, pool_cfg) mgr = PoolManager(provider_id, pool_cfg, provider_type=provider_type)
await mgr.on_request_success( await mgr.on_request_success(
session_uuid=session_uuid, session_uuid=session_uuid,
key_id=key_id, key_id=key_id,
ttfb_ms=ttfb_ms,
) )
except Exception: except Exception:
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)") logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
@@ -571,8 +573,15 @@ class TaskService:
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
response, _provider_name, attempt_id, _provider_id, _endpoint_id, _key_id = ( (
await request_dispatcher.dispatch( response,
_provider_name,
attempt_id,
_provider_id,
_endpoint_id,
_key_id,
_first_byte_time_ms,
) = await request_dispatcher.dispatch(
candidate=candidate, candidate=candidate,
candidate_index=candidate_index, candidate_index=candidate_index,
retry_index=retry_index, retry_index=retry_index,
@@ -588,11 +597,14 @@ class TaskService:
max_attempts=max_attempts_local, max_attempts=max_attempts_local,
is_stream=is_stream, is_stream=is_stream,
) )
)
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id) _ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
# Account Pool: on success, update sticky binding + LRU. # Account Pool: on success, update sticky binding + LRU.
await self._pool_on_success(candidate, request_body) await self._pool_on_success(
candidate,
request_body,
ttfb_ms=_first_byte_time_ms,
)
if is_stream: if is_stream:
return AttemptResult( return AttemptResult(

View File

@@ -0,0 +1,92 @@
"""Tests for pool account-state resolution helpers."""
from __future__ import annotations
from src.services.provider.pool.account_state import resolve_pool_account_state
def test_resolve_from_kiro_banned_metadata() -> None:
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": True, "ban_reason": "account suspended"}},
oauth_invalid_reason=None,
)
assert state.blocked is True
assert state.code == "account_banned"
assert state.label == "账号封禁"
assert state.reason == "account suspended"
def test_resolve_from_antigravity_forbidden_metadata() -> None:
state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={"antigravity": {"is_forbidden": True, "forbidden_reason": "403"}},
oauth_invalid_reason=None,
)
assert state.blocked is True
assert state.code == "account_forbidden"
assert state.label == "访问受限"
assert state.reason == "403"
def test_resolve_from_structured_oauth_reason() -> None:
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata=None,
oauth_invalid_reason="[ACCOUNT_BLOCK] Google requires verification",
)
assert state.blocked is True
assert state.code == "account_blocked"
assert state.label == "账号异常"
assert state.reason == "Google requires verification"
def test_resolve_from_keyword_oauth_reason() -> None:
state = resolve_pool_account_state(
provider_type=None,
upstream_metadata={},
oauth_invalid_reason="organization has been disabled by admin",
)
assert state.blocked is True
assert state.code == "account_blocked"
assert state.label == "账号异常"
def test_resolve_healthy_state() -> None:
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata={"codex": {"primary_used_percent": 30}},
oauth_invalid_reason="Token expired",
)
assert state.blocked is False
assert state.code is None
def test_bare_forbidden_not_treated_as_account_block() -> None:
"""HTTP 403 'Forbidden' from token refresh should not be misclassified."""
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata={},
oauth_invalid_reason="Forbidden",
)
assert state.blocked is False
def test_kiro_oauth_reason_text_detected_as_block() -> None:
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={},
oauth_invalid_reason="账户已封禁: Terms of Service violation",
)
assert state.blocked is True
assert state.code == "account_blocked"
def test_antigravity_oauth_reason_text_detected_as_block() -> None:
state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={},
oauth_invalid_reason="账户访问被禁止: 403 Forbidden",
)
assert state.blocked is True
assert state.code == "account_blocked"

View File

@@ -4,6 +4,8 @@ from __future__ import annotations
from src.services.provider.pool.config import ( from src.services.provider.pool.config import (
PoolConfig, PoolConfig,
SchedulingPreset,
ScoringWeights,
UnschedulableRule, UnschedulableRule,
parse_pool_config, parse_pool_config,
) )
@@ -21,6 +23,14 @@ def test_parse_pool_config_returns_defaults_for_empty_advanced() -> None:
assert cfg.sticky_session_ttl_seconds == 3600 assert cfg.sticky_session_ttl_seconds == 3600
assert cfg.load_threshold_percent == 80 assert cfg.load_threshold_percent == 80
assert cfg.lru_enabled is True assert cfg.lru_enabled is True
assert cfg.scheduling_mode == "lru"
assert cfg.scoring_weights == ScoringWeights()
# Default: only LRU preset enabled
assert len(cfg.scheduling_presets) == 1
assert cfg.scheduling_presets[0].preset == "lru"
assert cfg.scheduling_presets[0].enabled is True
assert cfg.latency_window_seconds == 3600
assert cfg.latency_sample_limit == 50
assert cfg.cost_window_seconds == 18000 assert cfg.cost_window_seconds == 18000
assert cfg.cost_limit_per_key_tokens is None assert cfg.cost_limit_per_key_tokens is None
assert cfg.cost_soft_threshold_percent == 80 assert cfg.cost_soft_threshold_percent == 80
@@ -31,13 +41,24 @@ def test_parse_pool_config_returns_defaults_for_empty_advanced() -> None:
assert cfg.unschedulable_rules == [] assert cfg.unschedulable_rules == []
def test_parse_pool_config_overrides_values() -> None: def test_parse_pool_config_overrides_values_legacy_string_list() -> None:
"""Legacy string-list format with scheduling_mode/lru_enabled."""
cfg = parse_pool_config( cfg = parse_pool_config(
{ {
"pool_advanced": { "pool_advanced": {
"sticky_session_ttl_seconds": 7200, "sticky_session_ttl_seconds": 7200,
"load_threshold_percent": 90, "load_threshold_percent": 90,
"lru_enabled": False, "lru_enabled": False,
"scheduling_mode": "multi_score",
"scheduling_presets": ["free_team_first", "recent_refresh", "free_team_first"],
"scoring_weights": {
"lru": 0.1,
"latency": 0.5,
"health": 0.2,
"cost_remaining": 0.2,
},
"latency_window_seconds": 7200,
"latency_sample_limit": 80,
"cost_window_seconds": 36000, "cost_window_seconds": 36000,
"cost_limit_per_key_tokens": 100000, "cost_limit_per_key_tokens": 100000,
"cost_soft_threshold_percent": 70, "cost_soft_threshold_percent": 70,
@@ -52,6 +73,20 @@ def test_parse_pool_config_overrides_values() -> None:
assert cfg.sticky_session_ttl_seconds == 7200 assert cfg.sticky_session_ttl_seconds == 7200
assert cfg.load_threshold_percent == 90 assert cfg.load_threshold_percent == 90
assert cfg.lru_enabled is False assert cfg.lru_enabled is False
assert cfg.scheduling_mode == "multi_score"
# Legacy string list → SchedulingPreset objects, deduped
preset_names = tuple(p.preset for p in cfg.scheduling_presets)
assert "free_team_first" in preset_names
assert "recent_refresh" in preset_names
assert cfg.scoring_weights == ScoringWeights(
lru=0.1,
latency=0.5,
health=0.2,
cost_remaining=0.2,
)
assert cfg.latency_window_seconds == 7200
assert cfg.latency_sample_limit == 80
assert "multi_score" in cfg.strategies
assert cfg.cost_window_seconds == 36000 assert cfg.cost_window_seconds == 36000
assert cfg.cost_limit_per_key_tokens == 100000 assert cfg.cost_limit_per_key_tokens == 100000
assert cfg.cost_soft_threshold_percent == 70 assert cfg.cost_soft_threshold_percent == 70
@@ -61,6 +96,114 @@ def test_parse_pool_config_overrides_values() -> None:
assert cfg.health_policy_enabled is False assert cfg.health_policy_enabled is False
def test_parse_pool_config_new_object_list_format() -> None:
"""New object-list format: [{preset, enabled, mode}]."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": True},
{"preset": "free_team_first", "enabled": True, "mode": "free_only"},
{"preset": "quota_balanced", "enabled": False},
{"preset": "recent_refresh", "enabled": True},
],
}
}
)
assert cfg is not None
assert len(cfg.scheduling_presets) == 4
lru = cfg.scheduling_presets[0]
assert lru.preset == "lru"
assert lru.enabled is True
ftf = cfg.scheduling_presets[1]
assert ftf.preset == "free_team_first"
assert ftf.enabled is True
assert ftf.mode == "free_only"
qb = cfg.scheduling_presets[2]
assert qb.preset == "quota_balanced"
assert qb.enabled is False
rr = cfg.scheduling_presets[3]
assert rr.preset == "recent_refresh"
assert rr.enabled is True
# Derived fields: lru enabled, non-lru enabled → multi_score
assert cfg.lru_enabled is True
assert cfg.scheduling_mode == "multi_score"
def test_parse_pool_config_new_format_lru_only() -> None:
"""When only LRU is enabled, scheduling_mode should be 'lru'."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": True},
{"preset": "quota_balanced", "enabled": False},
],
}
}
)
assert cfg is not None
assert cfg.lru_enabled is True
assert cfg.scheduling_mode == "lru"
def test_parse_pool_config_new_format_lru_disabled() -> None:
"""LRU disabled, other presets enabled → multi_score + lru_enabled=False."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": False},
{"preset": "quota_balanced", "enabled": True},
],
}
}
)
assert cfg is not None
assert cfg.lru_enabled is False
assert cfg.scheduling_mode == "multi_score"
def test_parse_pool_config_new_format_free_team_mode_validation() -> None:
"""Invalid mode falls back to 'both'."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "free_team_first", "enabled": True, "mode": "invalid_mode"},
],
}
}
)
assert cfg is not None
ftf = [p for p in cfg.scheduling_presets if p.preset == "free_team_first"][0]
assert ftf.mode == "both"
def test_parse_pool_config_new_format_dedup_presets() -> None:
"""Duplicate presets in object list should be deduplicated."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": True},
{"preset": "lru", "enabled": False},
{"preset": "quota_balanced", "enabled": True},
],
}
}
)
assert cfg is not None
lru_presets = [p for p in cfg.scheduling_presets if p.preset == "lru"]
assert len(lru_presets) == 1
assert lru_presets[0].enabled is True # first occurrence wins
def test_parse_pool_config_parses_unschedulable_rules() -> None: def test_parse_pool_config_parses_unschedulable_rules() -> None:
cfg = parse_pool_config( cfg = parse_pool_config(
{ {
@@ -97,6 +240,47 @@ def test_parse_pool_config_handles_invalid_types_gracefully() -> None:
assert cfg.cost_limit_per_key_tokens is None # default for opt_int assert cfg.cost_limit_per_key_tokens is None # default for opt_int
def test_parse_pool_config_invalid_scheduling_mode_falls_back_to_lru() -> None:
cfg = parse_pool_config({"pool_advanced": {"scheduling_mode": "unknown"}})
assert cfg is not None
assert cfg.scheduling_mode == "lru"
def test_parse_pool_config_scoring_weights_invalid_values_are_clamped() -> None:
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_mode": "multi_score",
"scoring_weights": {
"lru": 2.0,
"latency": -1.0,
"health": "bad",
},
}
}
)
assert cfg is not None
assert cfg.scoring_weights.lru == 1.0
assert cfg.scoring_weights.latency == 0.0
assert cfg.scoring_weights.health == 0.2
def test_parse_pool_config_invalid_scheduling_presets_are_ignored() -> None:
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_mode": "multi_score",
"scheduling_presets": ["quota_balanced", "unknown", 123, "single_account"],
}
}
)
assert cfg is not None
preset_names = tuple(p.preset for p in cfg.scheduling_presets if p.preset != "lru")
assert "quota_balanced" in preset_names
assert "single_account" in preset_names
assert "unknown" not in preset_names
def test_pool_config_is_frozen() -> None: def test_pool_config_is_frozen() -> None:
cfg = PoolConfig() cfg = PoolConfig()
try: try:
@@ -104,3 +288,12 @@ def test_pool_config_is_frozen() -> None:
assert False, "Should have raised FrozenInstanceError" assert False, "Should have raised FrozenInstanceError"
except AttributeError: except AttributeError:
pass pass
def test_scheduling_preset_is_frozen() -> None:
preset = SchedulingPreset(preset="lru", enabled=True)
try:
preset.enabled = False # type: ignore[misc]
assert False, "Should have raised FrozenInstanceError"
except AttributeError:
pass

View File

@@ -0,0 +1,58 @@
"""Tests for pool health cache helpers."""
from __future__ import annotations
from types import SimpleNamespace
from src.services.provider.pool import health_cache
def setup_function() -> None:
health_cache._clear_cache_for_tests()
def teardown_function() -> None:
health_cache._clear_cache_for_tests()
def test_aggregate_health_score_uses_lowest_format_score() -> None:
score = health_cache.aggregate_health_score(
{
"openai:chat": {"health_score": 0.92},
"openai:responses": {"health_score": 0.61},
}
)
assert score == 0.61
def test_get_health_scores_uses_cache_for_same_provider() -> None:
key = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.7}})
first = health_cache.get_health_scores("p1", [key])
assert first["k1"] == 0.7
key.health_by_format = {"f1": {"health_score": 0.2}}
second = health_cache.get_health_scores("p1", [key])
assert second["k1"] == 0.7
def test_get_health_scores_merges_missing_keys_into_cache() -> None:
k1 = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.7}})
first = health_cache.get_health_scores("p1", [k1])
assert first == {"k1": 0.7}
# Request with a new key k2 -- k1 should come from cache, k2 freshly computed
k1_stale = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.1}})
k2 = SimpleNamespace(id="k2", health_by_format={"f1": {"health_score": 0.5}})
second = health_cache.get_health_scores("p1", [k1_stale, k2])
assert second["k1"] == 0.7 # cached, not recomputed
assert second["k2"] == 0.5 # freshly computed
def test_invalidate_provider_health_scores_clears_cache_entry() -> None:
key = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.8}})
_ = health_cache.get_health_scores("p1", [key])
health_cache.invalidate_provider_health_scores("p1")
key.health_by_format = {"f1": {"health_score": 0.3}}
refreshed = health_cache.get_health_scores("p1", [key])
assert refreshed["k1"] == 0.3

View File

@@ -7,13 +7,23 @@ from unittest.mock import AsyncMock, patch
import pytest import pytest
from src.services.provider.pool.config import PoolConfig from src.services.provider.pool.config import PoolConfig, SchedulingPreset, ScoringWeights
from src.services.provider.pool.manager import PoolManager from src.services.provider.pool.manager import PoolManager
def _make_candidate(key_id: str, *, is_skipped: bool = False) -> SimpleNamespace: def _make_candidate(
key_id: str,
*,
is_skipped: bool = False,
upstream_metadata: dict | None = None,
oauth_invalid_reason: str | None = None,
) -> SimpleNamespace:
return SimpleNamespace( return SimpleNamespace(
key=SimpleNamespace(id=key_id), key=SimpleNamespace(
id=key_id,
upstream_metadata=upstream_metadata,
oauth_invalid_reason=oauth_invalid_reason,
),
is_skipped=is_skipped, is_skipped=is_skipped,
skip_reason=None, skip_reason=None,
) )
@@ -127,6 +137,125 @@ async def test_reorder_cost_exhausted_keys_are_skipped() -> None:
assert result[0].key.id == "key-2" assert result[0].key.id == "key-2"
@pytest.mark.asyncio
async def test_reorder_account_blocked_keys_are_skipped() -> None:
pool = PoolManager("provider-1", PoolConfig(), provider_type="kiro")
c1 = _make_candidate(
"key-1",
upstream_metadata={"kiro": {"is_banned": True, "ban_reason": "account suspended"}},
)
c2 = _make_candidate("key-2")
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value=None,
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert c1.is_skipped is True
assert "account blocked" in (c1.skip_reason or "")
assert result[0].key.id == "key-2"
assert c1._pool_extra_data["pool_skip"]["type"] == "account_blocked"
assert c1._pool_extra_data["pool_skip"]["account_block_label"] == "账号封禁"
@pytest.mark.asyncio
async def test_reorder_multi_score_uses_composite_score() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(
scheduling_mode="multi_score",
strategies=("multi_score",),
scoring_weights=ScoringWeights(lru=0.0, latency=1.0, health=0.0, cost_remaining=0.0),
),
)
c1 = _make_candidate("key-1")
c2 = _make_candidate("key-2")
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value=None,
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 10.0, "key-2": 20.0},
),
patch(
"src.services.provider.pool.redis_ops.batch_get_latency_avgs",
new_callable=AsyncMock,
return_value={"key-1": 600.0, "key-2": 120.0},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert result[0].key.id == "key-2"
assert result[0]._pool_extra_data["pool_selection"]["reason"] == "multi_score"
assert result[0]._pool_extra_data["pool_selection"]["scoring_mode"] == "multi_score"
@pytest.mark.asyncio
async def test_reorder_multi_score_presets_free_team_first() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(
scheduling_mode="multi_score",
strategies=("multi_score",),
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=True, mode="both"),
),
),
)
c1 = _make_candidate("key-1", upstream_metadata={"codex": {"plan_type": "plus"}})
c2 = _make_candidate("key-2", upstream_metadata={"codex": {"plan_type": "team"}})
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value=None,
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 100.0, "key-2": 100.0},
),
patch(
"src.services.provider.pool.redis_ops.batch_get_latency_avgs",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert result[0].key.id == "key-2"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_reorder_lru_sorts_least_recently_used_first( async def test_reorder_lru_sorts_least_recently_used_first(
pool: PoolManager, pool: PoolManager,
@@ -216,6 +345,41 @@ async def test_on_request_success_records_cost_when_configured() -> None:
mock_cost.assert_called_once_with("provider-1", "key-1", 500, 18000) mock_cost.assert_called_once_with("provider-1", "key-1", 500, 18000)
@pytest.mark.asyncio
async def test_on_request_success_records_latency_when_multi_score_enabled() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(
scheduling_mode="multi_score",
latency_window_seconds=7200,
latency_sample_limit=80,
),
)
with (
patch(
"src.services.provider.pool.redis_ops.set_sticky_binding",
new_callable=AsyncMock,
),
patch(
"src.services.provider.pool.redis_ops.touch_lru",
new_callable=AsyncMock,
),
patch(
"src.services.provider.pool.redis_ops.record_latency",
new_callable=AsyncMock,
) as mock_latency,
):
await pool.on_request_success(
session_uuid=None,
key_id="key-1",
tokens_used=0,
ttfb_ms=321,
)
mock_latency.assert_called_once_with("provider-1", "key-1", 321, 7200, 80)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# on_request_error # on_request_error
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------

View File

@@ -0,0 +1,233 @@
"""Tests for built-in multi-score pool strategy."""
from __future__ import annotations
from types import SimpleNamespace
from src.services.provider.pool.config import PoolConfig, SchedulingPreset, ScoringWeights
from src.services.provider.pool.strategies.multi_score import MultiScoreStrategy
from src.services.provider.pool.strategy import get_pool_strategy, register_pool_strategy
def _context() -> dict:
return {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 200.0, "k3": 300.0},
"latency_avgs": {"k1": 120.0, "k2": 300.0, "k3": 600.0},
"health_scores": {"k1": 0.95, "k2": 0.7, "k3": 0.4},
"cost_totals": {"k1": 100, "k2": 300, "k3": 900},
}
def _key_with_metadata(metadata: dict) -> SimpleNamespace:
return SimpleNamespace(upstream_metadata=metadata)
def test_multi_score_returns_none_when_mode_not_enabled() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(scheduling_mode="lru")
score = strategy.compute_score(key_id="k1", config=cfg, context=_context())
assert score is None
def test_multi_score_prefers_low_latency_when_latency_weight_is_high() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scoring_weights=ScoringWeights(lru=0.0, latency=1.0, health=0.0, cost_remaining=0.0),
)
ctx = _context()
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
assert s1 < s2 < s3
def test_multi_score_combines_health_and_cost() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scoring_weights=ScoringWeights(lru=0.0, latency=0.0, health=0.5, cost_remaining=0.5),
cost_limit_per_key_tokens=1000,
)
ctx = _context()
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s3 is not None
assert s1 < s3
def test_multi_score_strategy_is_registered() -> None:
register_pool_strategy("multi_score", MultiScoreStrategy())
registered = get_pool_strategy("multi_score")
assert registered is not None
def test_multi_score_preset_free_team_first_prefers_free_or_team() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(SchedulingPreset(preset="free_team_first", enabled=True, mode="both"),),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 100.0, "k3": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"plan_type": "plus"}}),
"k2": _key_with_metadata({"codex": {"plan_type": "free"}}),
"k3": _key_with_metadata({"codex": {"plan_type": "team"}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
assert s2 < s1
assert s3 < s1
def test_multi_score_preset_free_team_first_free_only_mode() -> None:
"""free_only mode: free is preferred, team is mid-priority."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=True, mode="free_only"),
),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 100.0, "k3": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"plan_type": "plus"}}),
"k2": _key_with_metadata({"codex": {"plan_type": "free"}}),
"k3": _key_with_metadata({"codex": {"plan_type": "team"}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
# free < team < plus
assert s2 < s3 < s1
def test_multi_score_preset_free_team_first_team_only_mode() -> None:
"""team_only mode: team is preferred, free is mid-priority."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=True, mode="team_only"),
),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 100.0, "k3": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"plan_type": "plus"}}),
"k2": _key_with_metadata({"codex": {"plan_type": "free"}}),
"k3": _key_with_metadata({"codex": {"plan_type": "team"}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
# team < free < plus
assert s3 < s2 < s1
def test_multi_score_preset_recent_refresh_prefers_nearer_reset() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(SchedulingPreset(preset="recent_refresh", enabled=True),),
)
ctx = {
"all_key_ids": ["k1", "k2"],
"lru_scores": {"k1": 100.0, "k2": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"primary_reset_seconds": 600}}),
"k2": _key_with_metadata({"codex": {"primary_reset_seconds": 120}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
assert s1 is not None and s2 is not None
assert s2 < s1
def test_multi_score_preset_single_account_prefers_latest_used() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(SchedulingPreset(preset="single_account", enabled=True),),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 900.0, "k3": 400.0},
"keys_by_id": {},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
assert s2 < s3 < s1
def test_multi_score_lru_disabled_no_blend() -> None:
"""When lru_enabled=False, LRU blend factor should be 0."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
lru_enabled=False,
scheduling_presets=(SchedulingPreset(preset="quota_balanced", enabled=True),),
)
ctx = {
"all_key_ids": ["k1", "k2"],
"lru_scores": {"k1": 100.0, "k2": 200.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"primary_used_percent": 80}}),
"k2": _key_with_metadata({"codex": {"primary_used_percent": 20}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
assert s1 is not None and s2 is not None
assert s2 < s1
def test_multi_score_disabled_presets_are_skipped() -> None:
"""Disabled presets should not affect scoring."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=False),
SchedulingPreset(preset="quota_balanced", enabled=True),
),
)
ctx = {
"all_key_ids": ["k1", "k2"],
"lru_scores": {"k1": 100.0, "k2": 100.0},
"keys_by_id": {
"k1": _key_with_metadata(
{
"codex": {"plan_type": "free", "primary_used_percent": 80},
}
),
"k2": _key_with_metadata(
{
"codex": {"plan_type": "plus", "primary_used_percent": 20},
}
),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
assert s1 is not None and s2 is not None
# quota_balanced only: k2 (20%) should score lower (better) than k1 (80%)
# free_team_first is disabled so plan_type should not matter
assert s2 < s1

View File

@@ -0,0 +1,83 @@
"""Tests for pool preset dimension registry and built-in dimensions."""
from __future__ import annotations
from types import SimpleNamespace
import src.services.provider.pool.dimensions # noqa: F401
from src.services.provider.pool.dimensions import get_preset_dimension, get_preset_names
def _key(metadata: dict, *, plan_type: str | None = None) -> SimpleNamespace:
return SimpleNamespace(upstream_metadata=metadata, oauth_plan_type=plan_type)
def test_registry_discovers_builtin_dimensions() -> None:
names = get_preset_names()
assert {"free_team_first", "recent_refresh", "quota_balanced", "single_account"}.issubset(names)
def test_universal_dimensions_are_applicable_to_any_provider() -> None:
for name in ("quota_balanced", "single_account"):
dim = get_preset_dimension(name)
assert dim is not None
assert dim.is_applicable("openai") is True
assert dim.is_applicable("codex") is True
assert dim.is_applicable("unknown_provider") is True
def test_provider_specific_dimensions_are_filtered() -> None:
for name in ("free_team_first", "recent_refresh"):
dim = get_preset_dimension(name)
assert dim is not None
assert dim.is_applicable("codex") is True
assert dim.is_applicable("kiro") is True
assert dim.is_applicable("openai") is False
def test_builtin_dimensions_compute_metric_in_range() -> None:
all_key_ids = ["k1", "k2", "k3"]
lru_scores = {"k1": 100.0, "k2": 800.0, "k3": 300.0}
keys_by_id = {
"k1": _key(
{
"codex": {
"plan_type": "plus",
"primary_reset_seconds": 900,
"primary_used_percent": 60,
}
}
),
"k2": _key(
{
"codex": {
"plan_type": "free",
"primary_reset_seconds": 120,
"primary_used_percent": 20,
}
}
),
"k3": _key(
{
"kiro": {
"next_reset_at": 4102444800,
"usage_percentage": 45,
"subscription_title": "Kiro Team",
}
},
plan_type="team",
),
}
for name in get_preset_names():
dim = get_preset_dimension(name)
assert dim is not None
mode = dim.default_mode
metric = dim.compute_metric(
key_id="k1",
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
mode=mode,
)
assert 0.0 <= metric <= 1.0

View File

@@ -0,0 +1,88 @@
"""Tests for pool redis latency operations."""
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool import redis_ops
class _FakePipe:
def __init__(self, execute_result: list[object] | None = None) -> None:
self.ops: list[tuple] = []
self._execute_result = execute_result or []
def zadd(self, key: str, mapping: dict[str, float]):
self.ops.append(("zadd", key, mapping))
return self
def zremrangebyscore(self, key: str, start: str, stop: float):
self.ops.append(("zremrangebyscore", key, start, stop))
return self
def zremrangebyrank(self, key: str, start: int, stop: int):
self.ops.append(("zremrangebyrank", key, start, stop))
return self
def expire(self, key: str, ttl: int):
self.ops.append(("expire", key, ttl))
return self
def eval(self, script: str, numkeys: int, key: str, window_start: str):
self.ops.append(("eval", script, numkeys, key, window_start))
return self
async def execute(self):
return self._execute_result
class _FakeRedis:
def __init__(self, pipe: _FakePipe) -> None:
self._pipe = pipe
def pipeline(self) -> _FakePipe:
return self._pipe
@pytest.mark.asyncio
async def test_record_latency_writes_sample_and_trims() -> None:
pipe = _FakePipe()
fake_redis = _FakeRedis(pipe)
with patch(
"src.services.provider.pool.redis_ops._get_redis",
new_callable=AsyncMock,
return_value=fake_redis,
):
await redis_ops.record_latency(
provider_id="prov-1",
key_id="key-1",
ttfb_ms=250,
window_seconds=3600,
sample_limit=50,
)
op_names = [item[0] for item in pipe.ops]
assert op_names == ["zadd", "zremrangebyscore", "zremrangebyrank", "expire"]
assert pipe.ops[0][1] == "ap:prov-1:latency:key-1"
assert pipe.ops[2][2:] == (0, -51)
@pytest.mark.asyncio
async def test_batch_get_latency_avgs_returns_numeric_results_only() -> None:
pipe = _FakePipe(execute_result=[120.5, None, "330"])
fake_redis = _FakeRedis(pipe)
with patch(
"src.services.provider.pool.redis_ops._get_redis",
new_callable=AsyncMock,
return_value=fake_redis,
):
result = await redis_ops.batch_get_latency_avgs(
provider_id="prov-1",
key_ids=["k1", "k2", "k3"],
window_seconds=3600,
)
assert result == {"k1": 120.5, "k3": 330.0}
assert len([item for item in pipe.ops if item[0] == "eval"]) == 3

View File

@@ -28,7 +28,7 @@ def _snapshot(**overrides: object) -> PoolSchedulingSnapshot:
def test_default_dimension_registry_contains_core_dimensions() -> None: def test_default_dimension_registry_contains_core_dimensions() -> None:
names = list_pool_scheduling_dimensions() names = list_pool_scheduling_dimensions()
assert names == ("manual", "cooldown", "circuit", "cost", "health") assert names == ("account_state", "manual", "cooldown", "circuit", "cost", "latency", "health")
def test_summary_available_when_all_dimensions_ok() -> None: def test_summary_available_when_all_dimensions_ok() -> None:
@@ -54,6 +54,37 @@ def test_summary_blocked_when_manual_disabled() -> None:
assert summary.score < 100.0 assert summary.score < 100.0
def test_summary_blocked_when_account_state_blocked() -> None:
dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(
account_blocked=True,
account_block_label="账号封禁",
account_block_reason="account suspended",
)
)
summary = summarize_pool_scheduling_dimensions(dimensions)
assert summary.status == "blocked"
assert summary.reason == "account_banned"
assert summary.candidate_eligible is False
assert summary.blocked_count >= 1
def test_account_state_takes_priority_over_manual_disabled() -> None:
dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(
is_active=False,
account_blocked=True,
account_block_label="访问受限",
account_block_reason="forbidden",
)
)
summary = summarize_pool_scheduling_dimensions(dimensions)
assert summary.status == "blocked"
assert summary.reason == "account_forbidden"
def test_summary_degraded_when_cost_reaches_soft_threshold() -> None: def test_summary_degraded_when_cost_reaches_soft_threshold() -> None:
dimensions = evaluate_pool_scheduling_dimensions( dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(cost_window_usage=8200, cost_limit=10000, cost_soft_threshold_percent=80) _snapshot(cost_window_usage=8200, cost_limit=10000, cost_soft_threshold_percent=80)
@@ -92,3 +123,11 @@ def test_dimension_result_keeps_degraded_health_details() -> None:
assert isinstance(health, PoolSchedulingDimensionResult) assert isinstance(health, PoolSchedulingDimensionResult)
assert health.status == "degraded" assert health.status == "degraded"
assert health.detail == "0.65" assert health.detail == "0.65"
def test_latency_dimension_degraded_when_latency_high() -> None:
dimensions = evaluate_pool_scheduling_dimensions(_snapshot(latency_avg_ms=3200))
latency = next((item for item in dimensions if item.code == "latency_high"), None)
assert isinstance(latency, PoolSchedulingDimensionResult)
assert latency.status == "degraded"
assert latency.detail == "3200ms"

View File

@@ -48,7 +48,25 @@ class TestPoolCandidateTraceExtraData:
ct = PoolCandidateTrace(key_id="k4", reason="random") ct = PoolCandidateTrace(key_id="k4", reason="random")
data = ct.to_extra_data() data = ct.to_extra_data()
sel = data["pool_selection"] sel = data["pool_selection"]
assert sel == {"reason": "random"} assert sel["reason"] == "random"
assert sel["scoring_mode"] == "lru"
def test_selected_multi_score_fields(self) -> None:
ct = PoolCandidateTrace(
key_id="k4b",
reason="multi_score",
scoring_mode="multi_score",
latency_avg_ms=245.7,
health_score=0.82,
composite_score=0.372156,
)
data = ct.to_extra_data()
sel = data["pool_selection"]
assert sel["reason"] == "multi_score"
assert sel["scoring_mode"] == "multi_score"
assert sel["latency_avg_ms"] == 245.7
assert sel["health_score"] == 0.82
assert sel["composite_score"] == 0.372156
def test_skipped_cooldown(self) -> None: def test_skipped_cooldown(self) -> None:
ct = PoolCandidateTrace( ct = PoolCandidateTrace(
@@ -78,11 +96,28 @@ class TestPoolCandidateTraceExtraData:
assert skip["cost_window_usage"] == 2000 assert skip["cost_window_usage"] == 2000
assert "cooldown_reason" not in skip assert "cooldown_reason" not in skip
def test_skipped_account_blocked(self) -> None:
ct = PoolCandidateTrace(
key_id="k6b",
skipped=True,
skip_type="account_blocked",
account_block_code="account_banned",
account_block_label="账号封禁",
account_block_reason="account suspended",
)
data = ct.to_extra_data()
skip = data["pool_skip"]
assert skip["type"] == "account_blocked"
assert skip["account_block_code"] == "account_banned"
assert skip["account_block_label"] == "账号封禁"
assert skip["account_block_reason"] == "account suspended"
def test_skipped_minimal(self) -> None: def test_skipped_minimal(self) -> None:
ct = PoolCandidateTrace(key_id="k7", skipped=True, skip_type="upstream") ct = PoolCandidateTrace(key_id="k7", skipped=True, skip_type="upstream")
data = ct.to_extra_data() data = ct.to_extra_data()
skip = data["pool_skip"] skip = data["pool_skip"]
assert skip == {"type": "upstream"} assert skip["type"] == "upstream"
assert skip["scoring_mode"] == "lru"
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -113,6 +148,7 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 2 assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 2 assert summary["skipped_cooldown"] == 2
assert summary["skipped_cost"] == 1 assert summary["skipped_cost"] == 1
assert summary["skipped_account_blocked"] == 0
assert summary["sticky_session"] is True assert summary["sticky_session"] is True
assert summary["success_key_id"] == "k1"[:8] assert summary["success_key_id"] == "k1"[:8]
assert summary["success_reason"] == "sticky" assert summary["success_reason"] == "sticky"
@@ -129,6 +165,7 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 2 assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 0 assert summary["skipped_cooldown"] == 0
assert summary["skipped_cost"] == 0 assert summary["skipped_cost"] == 0
assert summary["skipped_account_blocked"] == 0
assert "success_key_id" not in summary assert "success_key_id" not in summary
assert "success_reason" not in summary assert "success_reason" not in summary
@@ -143,3 +180,15 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 0 assert summary["attempted"] == 0
assert summary["skipped_cooldown"] == 1 assert summary["skipped_cooldown"] == 1
assert summary["skipped_cost"] == 1 assert summary["skipped_cost"] == 1
assert summary["skipped_account_blocked"] == 0
def test_summary_with_account_blocked(self) -> None:
trace = PoolSchedulingTrace(provider_id="prov-4", total_keys=2)
trace.candidate_traces = {
"k1": PoolCandidateTrace(key_id="k1", skipped=True, skip_type="account_blocked"),
"k2": PoolCandidateTrace(key_id="k2", reason="lru"),
}
summary = trace.build_summary(success_key_id="k2")
assert summary["attempted"] == 1
assert summary["skipped_account_blocked"] == 1

View File

@@ -2,13 +2,8 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from src.api.admin.pool.routes import _build_pool_scheduling_state
from src.services.provider.pool.account_state import resolve_pool_account_state
from src.api.admin.pool.routes import (
_build_pool_scheduling_state,
_is_known_banned_key,
_is_known_banned_reason,
)
def test_pool_scheduling_state_manual_disabled_is_blocked() -> None: def test_pool_scheduling_state_manual_disabled_is_blocked() -> None:
@@ -24,6 +19,10 @@ def test_pool_scheduling_state_manual_disabled_is_blocked() -> None:
dimensions, dimensions,
) = _build_pool_scheduling_state( ) = _build_pool_scheduling_state(
is_active=False, is_active=False,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason=None, cooldown_reason=None,
cooldown_ttl_seconds=None, cooldown_ttl_seconds=None,
circuit_breaker_open=False, circuit_breaker_open=False,
@@ -55,6 +54,10 @@ def test_pool_scheduling_state_cooldown_detail_is_mapped() -> None:
dimensions, dimensions,
) = _build_pool_scheduling_state( ) = _build_pool_scheduling_state(
is_active=True, is_active=True,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason="rate_limited_429", cooldown_reason="rate_limited_429",
cooldown_ttl_seconds=180, cooldown_ttl_seconds=180,
circuit_breaker_open=False, circuit_breaker_open=False,
@@ -84,6 +87,10 @@ def test_pool_scheduling_state_cost_soft_is_degraded() -> None:
_dimensions, _dimensions,
) = _build_pool_scheduling_state( ) = _build_pool_scheduling_state(
is_active=True, is_active=True,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason=None, cooldown_reason=None,
cooldown_ttl_seconds=None, cooldown_ttl_seconds=None,
circuit_breaker_open=False, circuit_breaker_open=False,
@@ -101,28 +108,36 @@ def test_pool_scheduling_state_cost_soft_is_degraded() -> None:
def test_known_banned_reason_account_block_prefix() -> None: def test_known_banned_reason_account_block_prefix() -> None:
assert _is_known_banned_reason("[ACCOUNT_BLOCK] Google 要求验证账号") is True state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata={},
oauth_invalid_reason="[ACCOUNT_BLOCK] Google 要求验证账号",
)
assert state.blocked is True
def test_known_banned_key_detects_kiro_banned_metadata() -> None: def test_known_banned_key_detects_kiro_banned_metadata() -> None:
key = SimpleNamespace( state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": True}}, upstream_metadata={"kiro": {"is_banned": True}},
oauth_invalid_reason=None, oauth_invalid_reason=None,
) )
assert _is_known_banned_key(key, "kiro") is True assert state.blocked is True
def test_known_banned_key_detects_reason_keywords() -> None: def test_known_banned_key_detects_reason_keywords() -> None:
key = SimpleNamespace( state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={}, upstream_metadata={},
oauth_invalid_reason="AWS account temporarily suspended", oauth_invalid_reason="AWS account temporarily suspended",
) )
assert _is_known_banned_key(key, "antigravity") is True assert state.blocked is True
def test_known_banned_key_does_not_treat_token_expired_as_banned() -> None: def test_known_banned_key_does_not_treat_token_expired_as_banned() -> None:
key = SimpleNamespace( state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": False}}, upstream_metadata={"kiro": {"is_banned": False}},
oauth_invalid_reason="access token expired", oauth_invalid_reason="access token expired",
) )
assert _is_known_banned_key(key, "kiro") is False assert state.blocked is False