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[]
}
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 {
key_id: 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(
providerId: string,
params: PoolKeysQuery = {},

View File

@@ -453,11 +453,29 @@ export interface ClaudeCodeAdvancedConfig {
cli_only_enabled?: boolean
}
export interface SchedulingPresetItem {
preset: string
enabled: boolean
mode?: string | null
}
export interface PoolAdvancedConfig {
global_priority?: number | null
sticky_session_ttl_seconds?: number | null
load_threshold_percent?: number | null
// 旧字段(兼容读取)
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_limit_per_key_tokens?: 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
? 'text-muted-foreground/60 cursor-not-allowed'
: '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)"
>
{{ key.priority }}

View File

@@ -1101,6 +1101,7 @@ import {
import type { UpstreamMetadata, AntigravityModelQuota } from '@/api/endpoints/types'
import { formatApiFormat } from '@/api/endpoints/types/api-format'
import { isOAuthAccountProviderType, isKeyManagedProviderType } from '../utils/providerTypeUtils'
import { isAccountLevelBlockReason } from '@/utils/accountBlock'
// 扩展端点类型,包含密钥列表
interface ProviderEndpointWithKeys extends ProviderEndpoint {
@@ -1568,8 +1569,7 @@ async function handleRefreshOAuth(key: EndpointAPIKey) {
// 判断是否为账号级别的封禁(刷新 token 无法修复)
function isAccountLevelBlock(key: EndpointAPIKey): boolean {
if (!key.oauth_invalid_reason) return false
return key.oauth_invalid_reason.startsWith('[ACCOUNT_BLOCK]')
return isAccountLevelBlockReason(key.oauth_invalid_reason)
}
// 清除 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
v-if="selectedProviderId"
variant="ghost"
size="icon"
class="h-8 w-8"
title="号池配置"
@click="showConfigDialog = true"
variant="outline"
size="sm"
class="h-8 px-2 text-xs gap-1"
title="调整号池调度"
@click="showSchedulingDialog = true"
>
<Settings class="w-3.5 h-3.5" />
调度
<ChevronDown class="w-3 h-3 text-muted-foreground" />
</Button>
<Button
v-if="selectedProviderId"
@@ -53,7 +54,8 @@
<Ban class="w-3.5 h-3.5" />
</Button>
<RefreshButton
:loading="keysLoading"
:loading="refreshCurrentPageLoading"
:title="refreshButtonTitle"
@click="refreshCurrentPage"
/>
</div>
@@ -197,16 +199,16 @@
>
<Upload class="w-3.5 h-3.5" />
</Button>
<Button
<button
v-if="selectedProviderId"
variant="ghost"
size="icon"
class="h-8 w-8"
title="号池配置"
@click="showConfigDialog = true"
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"
title="点击调整号池调度"
@click="showSchedulingDialog = true"
>
<Settings class="w-3.5 h-3.5" />
</Button>
<span class="text-muted-foreground/80 hidden lg:inline">调度:</span>
<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
v-if="selectedProviderId"
variant="ghost"
@@ -218,7 +220,8 @@
<Ban class="w-3.5 h-3.5" />
</Button>
<RefreshButton
:loading="keysLoading"
:loading="refreshCurrentPageLoading"
:title="refreshButtonTitle"
@click="refreshCurrentPage"
/>
</div>
@@ -370,6 +373,14 @@
{{ getKeyOAuthExpires(key)?.text }}
</span>
</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
v-if="key.oauth_plan_type"
variant="outline"
@@ -753,6 +764,14 @@
{{ getKeyOAuthExpires(key)?.text }}
</span>
</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
v-if="key.oauth_plan_type"
variant="outline"
@@ -1092,13 +1111,13 @@
@close="showImportDialog = false"
@saved="handleAccountDialogSaved"
/>
<PoolConfigDialog
<PoolSchedulingDialog
v-if="selectedProviderId"
v-model="showConfigDialog"
v-model="showSchedulingDialog"
:provider-id="selectedProviderId"
:provider-type="selectedProviderData?.provider_type"
:provider-type="selectedProviderType"
:current-config="selectedProviderConfig"
:current-claude-config="selectedProviderData?.claude_code_advanced"
:current-claude-config="selectedProviderClaudeConfig"
@saved="loadOverview"
/>
<KeyFormDialog
@@ -1132,7 +1151,7 @@ import { ref, computed, watch, onMounted, onBeforeUnmount } from 'vue'
import {
Search,
Upload,
Settings,
ChevronDown,
RefreshCw,
Power,
Database,
@@ -1176,6 +1195,7 @@ import { useConfirm } from '@/composables/useConfirm'
import { parseApiError } from '@/utils/errorParser'
import {
getPoolOverview,
getPoolSchedulingPresets,
listPoolKeys,
clearPoolCooldown,
cleanupBannedPoolKeys,
@@ -1193,18 +1213,20 @@ import type {
PoolOverviewItem,
PoolKeyDetail,
PoolKeysPageResponse,
PoolPresetMeta,
} 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 { 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 KeyFormDialog from '@/features/providers/components/KeyFormDialog.vue'
import OAuthKeyEditDialog from '@/features/providers/components/OAuthKeyEditDialog.vue'
import OAuthAccountDialog from '@/features/providers/components/OAuthAccountDialog.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 { copyToClipboard } = useClipboard()
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
})
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 fromDetail = String(selectedProviderData.value?.provider_type || '').trim().toLowerCase()
if (fromDetail) return fromDetail
@@ -1330,11 +1427,11 @@ async function refresh() {
const keyPage = ref<PoolKeysPageResponse>({ total: 0, page: 1, page_size: 50, keys: [] })
const keysLoading = ref(false)
const refreshingCurrentPageQuota = ref(false)
const queuedCurrentPageQuotaRefresh = ref(false)
const searchQuery = ref('')
const statusFilter = ref('all')
const currentPage = ref(1)
const pageSize = ref(50)
const MANUAL_QUOTA_REFRESH_COOLDOWN_SECONDS = 5 * 60
const refreshingOAuthKeyId = ref<string | null>(null)
const revealedKeys = ref<Map<string, string>>(new Map())
const recoveringHealthKeyId = ref<string | null>(null)
@@ -1372,57 +1469,123 @@ const quotaRefreshSupported = computed(() => {
|| selectedProviderType.value === 'antigravity'
})
function getCurrentPageQuotaKeyIds(): string[] {
const ids: string[] = []
const refreshCurrentPageLoading = computed(() => {
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 eligibleIds: string[] = []
let cooledDownCount = 0
let minRemainingSeconds = 0
const nowSeconds = Math.floor(Date.now() / 1000)
for (const key of keyPage.value.keys) {
const id = String(key.key_id || '').trim()
if (!id || seen.has(id)) continue
seen.add(id)
ids.push(id)
const updatedAt = normalizeQuotaUpdatedAt(key.quota_updated_at ?? null)
if (updatedAt == null) {
eligibleIds.push(id)
continue
}
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 ids
}
return {
total: seen.size,
eligibleIds,
cooledDownCount,
minRemainingSeconds,
}
})
async function refreshCurrentPageQuotaInBackground(options: { silent?: boolean } = {}) {
if (!selectedProviderId.value || !quotaRefreshSupported.value) return
async function refreshCurrentPageQuotaInBackground(
options: { silent?: boolean; reloadAfter?: boolean } = {},
): Promise<boolean> {
if (!selectedProviderId.value || !quotaRefreshSupported.value) return false
const providerId = selectedProviderId.value
const keyIds = getCurrentPageQuotaKeyIds()
if (keyIds.length === 0) return
const quotaStats = currentPageQuotaRefreshStats.value
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) {
queuedCurrentPageQuotaRefresh.value = true
return
return false
}
refreshingCurrentPageQuota.value = true
try {
const result = await refreshProviderQuota(providerId, keyIds)
const result = await refreshProviderQuota(providerId, quotaStats.eligibleIds)
const successCount = Number(result.success || 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()
}
if (!options.silent) {
success(`当前页额度刷新完成:成功 ${successCount},失败 ${failedCount}`)
const skippedText = skippedCount > 0 ? `,冷却跳过 ${skippedCount}` : ''
success(`当前页额度刷新完成:成功 ${successCount},失败 ${failedCount}${skippedText}`)
}
return true
} catch (err) {
showError(parseApiError(err, '刷新当前页额度失败'))
return false
} finally {
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() {
await refresh()
const quotaDidReload = await refreshCurrentPageQuotaInBackground({ reloadAfter: true })
if (!quotaDidReload) {
await refresh()
}
}
async function loadKeys() {
@@ -1807,11 +1970,13 @@ async function handleCleanupBannedKeys() {
// --- Dialogs ---
const showImportDialog = ref(false)
const showConfigDialog = ref(false)
const showSchedulingDialog = ref(false)
async function handleAccountDialogSaved() {
showImportDialog.value = false
await Promise.all([loadKeys(), loadOverview()])
// 导入账号后补一次静默额度刷新,避免新账号在列表里暂无额度信息
await refreshCurrentPageQuotaInBackground({ silent: true })
}
// --- Formatting ---
@@ -1831,6 +1996,8 @@ function formatCooldownReason(reason: string): string {
type PoolStatusVariant = 'default' | 'secondary' | 'destructive' | 'outline' | 'success' | 'warning' | 'dark'
function getSchedulingStatus(key: PoolKeyDetail): 'available' | 'degraded' | 'blocked' {
if (getAccountAlertLabel(key)) return 'blocked'
const status = key.scheduling_status
if (status === 'available' || status === 'degraded' || status === 'blocked') {
return status
@@ -1845,6 +2012,9 @@ function getSchedulingStatus(key: PoolKeyDetail): 'available' | 'degraded' | 'bl
}
function getSchedulingBadgeLabel(key: PoolKeyDetail): string {
const accountAlert = getAccountAlertLabel(key)
if (accountAlert) return accountAlert
const rawLabel = String(key.scheduling_label || '').trim()
if (rawLabel) {
if (rawLabel === '禁用' || rawLabel === '停用') return '禁用'
@@ -1861,6 +2031,8 @@ function getSchedulingBadgeLabel(key: PoolKeyDetail): string {
}
function getSchedulingBadgeVariant(key: PoolKeyDetail): PoolStatusVariant {
if (getAccountAlertLabel(key)) return 'destructive'
const reason = key.scheduling_reason
if (reason === 'manual_disabled') return 'dark'
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 {
const accountAlertTitle = getAccountAlertTitle(key)
if (accountAlertTitle) return accountAlertTitle
if (key.scheduling_dimensions && key.scheduling_dimensions.length > 0) {
return key.scheduling_dimensions.map((item) => {
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}`
}
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 {
const normalized = label.trim()
if (!normalized) return '额度'
@@ -2214,8 +2424,7 @@ function formatRelativeTime(isoStr: string): string {
// --- Init ---
onMounted(async () => {
startCountdownTimer()
await loadOverview()
void refreshCurrentPageQuotaInBackground({ silent: true })
await Promise.all([loadSchedulingPresetMetas(), loadOverview()])
})
onBeforeUnmount(() => {

View File

@@ -28,7 +28,9 @@ from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, Usage
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.dimensions import get_preset_dimension_metas
from src.services.provider.pool.scheduling_dimensions import (
PoolSchedulingSnapshot,
evaluate_pool_scheduling_dimensions,
@@ -47,6 +49,8 @@ from .schemas import (
PoolOverviewResponse,
PoolSchedulingDimension,
PoolSchedulingReason,
PresetDimensionMetaResponse,
PresetModeMetaResponse,
)
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)
# ---------------------------------------------------------------------------
# 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
# ---------------------------------------------------------------------------
@@ -126,24 +155,6 @@ _COOLDOWN_REASON_LABELS: dict[str, str] = {
"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:
if isinstance(value, bool):
@@ -161,65 +172,15 @@ def _to_float(value: Any) -> float | 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:
upstream_metadata = getattr(key, "upstream_metadata", None)
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
from src.services.provider.pool.account_state import resolve_pool_account_state
if normalized_provider == "kiro" and provider_bucket:
if _is_truthy_flag(provider_bucket.get("is_banned")):
return True
if normalized_provider == "antigravity" and provider_bucket:
if _is_truthy_flag(provider_bucket.get("is_forbidden")):
return True
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))
state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(key, "upstream_metadata", None),
oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
)
return state.blocked
def _format_percent(value: float) -> str:
@@ -536,6 +497,10 @@ def _format_cooldown_detail(raw: str | None) -> str | None:
def _build_pool_scheduling_state(
*,
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_ttl_seconds: int | None,
circuit_breaker_open: bool,
@@ -557,6 +522,10 @@ def _build_pool_scheduling_state(
"""Build unified scheduling state for frontend display."""
snapshot = PoolSchedulingSnapshot(
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_ttl_seconds=cooldown_ttl_seconds,
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):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
@@ -813,20 +815,34 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
if pcfg and pcfg.lru_enabled
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 = (
pool_redis.batch_get_cost_totals(pid, key_ids, pcfg.cost_window_seconds)
if pcfg
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_cooldown_ttls(pid, key_ids),
_lru_coro,
_latency_coro,
_cost_coro,
pool_redis.batch_get_key_sticky_counts(pid, key_ids),
)
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_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_reason,
@@ -886,6 +909,10 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
scheduling_dimensions,
) = _build_pool_scheduling_state(
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_ttl_seconds=cd_ttl,
circuit_breaker_open=any_circuit_open,
@@ -955,9 +982,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
key_name=k.name or "",
is_active=bool(k.is_active),
auth_type=str(getattr(k, "auth_type", "api_key") or "api_key"),
oauth_expires_at=_derive_oauth_expires_at(
k, auth_config=oauth_auth_config
),
oauth_expires_at=_derive_oauth_expires_at(k, auth_config=oauth_auth_config),
oauth_invalid_at=(
int(k.oauth_invalid_at.timestamp())
if getattr(k, "oauth_invalid_at", None)

View File

@@ -29,6 +29,25 @@ class PoolOverviewResponse(BaseModel):
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
# ---------------------------------------------------------------------------

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):
"""通用号池配置(适用于所有 Provider 类型)。"""
@@ -148,7 +222,33 @@ class PoolAdvancedConfig(BaseModel):
le=100,
description="负载率阈值(%),超过时该 Key 被降权。默认 80",
)
# 保留旧字段供向后兼容(新客户端不再发送)
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(
None,
ge=3600,

View File

@@ -60,7 +60,7 @@ class RequestDispatcher:
attempt_counter: int,
max_attempts: int,
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: 是否为流式请求
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:
ExecutionError: 执行失败时
@@ -144,6 +144,19 @@ class RequestDispatcher:
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 (
execution_result.response,
provider_name,
@@ -151,4 +164,5 @@ class RequestDispatcher:
provider_id,
endpoint_id,
key_id,
ttfb_ms,
)

View File

@@ -3,12 +3,18 @@
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
__all__ = [
"PoolConfig",
"PoolManager",
"ScoringWeights",
"UnschedulableRule",
"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 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)
@@ -32,8 +52,17 @@ class PoolConfig:
# -- Load-Aware Selection -------------------------------------------------
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
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 -----------------------------------------
cost_window_seconds: int = 18000 # 5 hours
@@ -122,11 +151,35 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
except (TypeError, ValueError):
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(
sticky_session_ttl_seconds=_int_or("sticky_session_ttl_seconds", 3600),
global_priority=_opt_int("global_priority"),
load_threshold_percent=_int_or("load_threshold_percent", 80),
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_limit_per_key_tokens=_opt_int("cost_limit_per_key_tokens"),
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_window_seconds=_int_or("stream_timeout_window_seconds", 1800),
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, ...]:
"""Parse strategy names from config (list[str] -> tuple[str, ...])."""
if not isinstance(raw, list):
return ()
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.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.health_cache import get_health_scores
from src.services.provider.pool.trace import PoolCandidateTrace, PoolSchedulingTrace
if TYPE_CHECKING:
@@ -31,11 +33,17 @@ if TYPE_CHECKING:
class PoolManager:
"""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.config = config
self.provider_type = str(provider_type or "").strip().lower() or None
# ------------------------------------------------------------------
# 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
and that key appears in *candidates* and is not in cooldown, move it
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``).
3. **LRU sort** -- among remaining candidates at the same priority
level, sort by least-recently-used.
@@ -102,6 +110,13 @@ class PoolManager:
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) ---------------------
all_key_ids = [str(c.key.id) for c in candidates]
@@ -113,17 +128,26 @@ class PoolManager:
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.
coros: list[Any] = [_cooldown_coro]
_cost_idx = -1
_lru_idx = -1
_latency_idx = -1
if _cost_coro is not None:
_cost_idx = len(coros)
coros.append(_cost_coro)
if _lru_coro is not None:
_lru_idx = len(coros)
coros.append(_lru_coro)
if _latency_coro is not None:
_latency_idx = len(coros)
coros.append(_latency_coro)
gathered = await asyncio.gather(*coros)
@@ -158,7 +182,29 @@ class PoolManager:
if _lru_idx >= 0:
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 ----------------------------------
custom_scores: dict[str, float] = {}
for strategy in strategies:
if hasattr(strategy, "compute_score"):
for kid in all_key_ids:
@@ -169,7 +215,8 @@ class PoolManager:
context=strategy_context,
)
if custom is not None:
lru_scores[kid] = custom
custom_scores[kid] = float(custom)
lru_scores[kid] = float(custom)
except Exception:
pass
@@ -181,6 +228,11 @@ class PoolManager:
for c in candidates:
kid = str(c.key.id)
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?
if c.is_skipped:
@@ -191,6 +243,25 @@ class PoolManager:
continue
# 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)
if cd_reason is not None:
c.is_skipped = True
@@ -225,7 +296,10 @@ class PoolManager:
trace.sticky_session_used = True
else:
available.append(c)
ct.reason = "lru" if lru_scores.get(kid, 0) > 0 else "random"
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.lru_score = lru_scores.get(kid, 0.0)
ct.cost_window_usage = cost_totals.get(kid, 0)
@@ -363,7 +437,7 @@ class PoolManager:
:class:`ProviderAPIKey` objects instead of candidates:
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.
4. Random tiebreak for identical LRU scores.
5. Return the first available key, or ``None``.
@@ -390,22 +464,32 @@ class PoolManager:
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]
_cost_idx_sk = -1
_lru_idx_sk = -1
_latency_idx_sk = -1
if _cost_coro is not None:
_cost_idx_sk = len(coros_sk)
coros_sk.append(_cost_coro)
if _lru_coro is not None:
_lru_idx_sk = len(coros_sk)
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)
cooldowns = gathered_sk[0]
cost_exhausted: set[str] = set()
cost_totals: dict[str, int] = {}
if _cost_idx_sk >= 0:
cost_totals = gathered_sk[_cost_idx_sk]
for kid, total in cost_totals.items():
@@ -416,9 +500,25 @@ class PoolManager:
if _lru_idx_sk >= 0:
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 ------------------------------------------
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:
if hasattr(strategy, "compute_score"):
for kid in key_ids:
@@ -429,7 +529,7 @@ class PoolManager:
context=strategy_context,
)
if custom is not None:
lru_scores[kid] = custom
lru_scores[kid] = float(custom)
except Exception:
pass
@@ -440,6 +540,14 @@ class PoolManager:
for k in keys:
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:
continue
if kid in cost_exhausted:
@@ -480,6 +588,7 @@ class PoolManager:
session_uuid: str | None,
key_id: str,
tokens_used: int = 0,
ttfb_ms: int | None = None,
) -> None:
"""Called after a successful upstream request."""
pid = self.provider_id
@@ -500,6 +609,16 @@ class PoolManager:
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(
self,
*,
@@ -570,6 +689,8 @@ def _get_active_strategies(config: PoolConfig) -> list[Any]:
if not config.strategies:
return []
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
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}:cooldown:{key_id} STRING -> reason (TTL: error-specific)
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)
"""
@@ -44,6 +45,10 @@ def _cost_key(provider_id: str, key_id: str) -> str:
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:
return f"provider_oauth_token_cache:{key_id}"
@@ -91,6 +96,32 @@ end
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":
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}
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:
redis = await _get_redis()
if redis is None:

View File

@@ -25,6 +25,10 @@ class PoolSchedulingSnapshot:
cost_limit: int | None
cost_soft_threshold_percent: int = 80
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)
@@ -67,6 +71,44 @@ class PoolSchedulingDimension(Protocol):
"""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)
class _ManualEnableDimension:
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_ORDER: list[str] = []
@@ -363,10 +458,12 @@ def summarize_pool_scheduling_dimensions(
def _register_default_dimensions() -> None:
register_pool_scheduling_dimension("account_state", _AccountStateDimension())
register_pool_scheduling_dimension("manual", _ManualEnableDimension())
register_pool_scheduling_dimension("cooldown", _CooldownDimension())
register_pool_scheduling_dimension("circuit", _CircuitBreakerDimension())
register_pool_scheduling_dimension("cost", _CostDimension())
register_pool_scheduling_dimension("latency", _LatencyDimension())
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_soft_threshold: 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_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]:
"""Build dict to merge into ``RequestCandidate.extra_data``."""
@@ -35,8 +42,16 @@ class PoolCandidateTrace:
skip_info["cooldown_reason"] = self.cooldown_reason
if self.cooldown_ttl is not None:
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:
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}
sel: dict[str, Any] = {"reason": self.reason}
@@ -50,6 +65,14 @@ class PoolCandidateTrace:
sel["cost_limit"] = self.cost_limit
if self.cost_soft_threshold:
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}
@@ -72,6 +95,7 @@ class PoolSchedulingTrace:
"""Build compact dict for ``Usage.request_metadata["pool_summary"]``."""
skipped_cooldown = 0
skipped_cost = 0
skipped_account_blocked = 0
attempted = 0
for t in self.candidate_traces.values():
if t.skipped:
@@ -79,6 +103,8 @@ class PoolSchedulingTrace:
skipped_cooldown += 1
elif t.skip_type == "cost_exhausted":
skipped_cost += 1
elif t.skip_type == "account_blocked":
skipped_account_blocked += 1
if attempted_key_ids is None:
# Backward-compatible behavior: count all schedulable keys.
@@ -101,6 +127,7 @@ class PoolSchedulingTrace:
"attempted": attempted,
"skipped_cooldown": skipped_cooldown,
"skipped_cost": skipped_cost,
"skipped_account_blocked": skipped_account_blocked,
"sticky_session": self.sticky_session_used,
}
if success_key_id:

View File

@@ -235,7 +235,7 @@ class TaskService:
provider_type = str(getattr(provider, "provider_type", "") or "")
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 [])
if not candidate_keys and getattr(candidate, "key", None) is not None:
@@ -334,6 +334,7 @@ class TaskService:
async def _pool_on_success(
candidate: Any,
request_body: dict[str, Any] | None,
ttfb_ms: int | None = None,
) -> None:
"""Notify the pool manager about a successful request (sticky + LRU)."""
try:
@@ -354,10 +355,11 @@ class TaskService:
provider_type = str(getattr(provider, "provider_type", "") or "")
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(
session_uuid=session_uuid,
key_id=key_id,
ttfb_ms=ttfb_ms,
)
except Exception:
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
@@ -571,28 +573,38 @@ class TaskService:
candidate_record_id = str(created.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(
candidate=candidate,
candidate_index=candidate_index,
retry_index=retry_index,
candidate_record_id=candidate_record_id,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
(
response,
_provider_name,
attempt_id,
_provider_id,
_endpoint_id,
_key_id,
_first_byte_time_ms,
) = await request_dispatcher.dispatch(
candidate=candidate,
candidate_index=candidate_index,
retry_index=retry_index,
candidate_record_id=candidate_record_id,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
# 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:
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 (
PoolConfig,
SchedulingPreset,
ScoringWeights,
UnschedulableRule,
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.load_threshold_percent == 80
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_limit_per_key_tokens is None
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 == []
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(
{
"pool_advanced": {
"sticky_session_ttl_seconds": 7200,
"load_threshold_percent": 90,
"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_limit_per_key_tokens": 100000,
"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.load_threshold_percent == 90
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_limit_per_key_tokens == 100000
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
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:
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
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:
cfg = PoolConfig()
try:
@@ -104,3 +288,12 @@ def test_pool_config_is_frozen() -> None:
assert False, "Should have raised FrozenInstanceError"
except AttributeError:
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
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
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(
key=SimpleNamespace(id=key_id),
key=SimpleNamespace(
id=key_id,
upstream_metadata=upstream_metadata,
oauth_invalid_reason=oauth_invalid_reason,
),
is_skipped=is_skipped,
skip_reason=None,
)
@@ -127,6 +137,125 @@ async def test_reorder_cost_exhausted_keys_are_skipped() -> None:
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
async def test_reorder_lru_sorts_least_recently_used_first(
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)
@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
# ---------------------------------------------------------------------------

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:
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:
@@ -54,6 +54,37 @@ def test_summary_blocked_when_manual_disabled() -> None:
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:
dimensions = evaluate_pool_scheduling_dimensions(
_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 health.status == "degraded"
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")
data = ct.to_extra_data()
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:
ct = PoolCandidateTrace(
@@ -78,11 +96,28 @@ class TestPoolCandidateTraceExtraData:
assert skip["cost_window_usage"] == 2000
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:
ct = PoolCandidateTrace(key_id="k7", skipped=True, skip_type="upstream")
data = ct.to_extra_data()
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["skipped_cooldown"] == 2
assert summary["skipped_cost"] == 1
assert summary["skipped_account_blocked"] == 0
assert summary["sticky_session"] is True
assert summary["success_key_id"] == "k1"[:8]
assert summary["success_reason"] == "sticky"
@@ -129,6 +165,7 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 0
assert summary["skipped_cost"] == 0
assert summary["skipped_account_blocked"] == 0
assert "success_key_id" not in summary
assert "success_reason" not in summary
@@ -143,3 +180,15 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 0
assert summary["skipped_cooldown"] == 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 types import SimpleNamespace
from src.api.admin.pool.routes import (
_build_pool_scheduling_state,
_is_known_banned_key,
_is_known_banned_reason,
)
from src.api.admin.pool.routes import _build_pool_scheduling_state
from src.services.provider.pool.account_state import resolve_pool_account_state
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,
) = _build_pool_scheduling_state(
is_active=False,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason=None,
cooldown_ttl_seconds=None,
circuit_breaker_open=False,
@@ -55,6 +54,10 @@ def test_pool_scheduling_state_cooldown_detail_is_mapped() -> None:
dimensions,
) = _build_pool_scheduling_state(
is_active=True,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason="rate_limited_429",
cooldown_ttl_seconds=180,
circuit_breaker_open=False,
@@ -84,6 +87,10 @@ def test_pool_scheduling_state_cost_soft_is_degraded() -> None:
_dimensions,
) = _build_pool_scheduling_state(
is_active=True,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason=None,
cooldown_ttl_seconds=None,
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:
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:
key = SimpleNamespace(
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": True}},
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:
key = SimpleNamespace(
state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={},
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:
key = SimpleNamespace(
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": False}},
oauth_invalid_reason="access token expired",
)
assert _is_known_banned_key(key, "kiro") is False
assert state.blocked is False