mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
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:
@@ -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 = {},
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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>
|
||||
886
frontend/src/features/pool/components/PoolSchedulingDialog.vue
Normal file
886
frontend/src/features/pool/components/PoolSchedulingDialog.vue
Normal 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>
|
||||
@@ -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 }}
|
||||
|
||||
@@ -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 失效标记
|
||||
|
||||
34
frontend/src/utils/accountBlock.ts
Normal file
34
frontend/src/utils/accountBlock.ts
Normal 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()
|
||||
}
|
||||
@@ -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(() => {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
189
src/services/provider/pool/account_state.py
Normal file
189
src/services/provider/pool/account_state.py
Normal 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",
|
||||
]
|
||||
@@ -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),)
|
||||
|
||||
30
src/services/provider/pool/dimensions/__init__.py
Normal file
30
src/services/provider/pool/dimensions/__init__.py
Normal 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",
|
||||
]
|
||||
227
src/services/provider/pool/dimensions/_helpers.py
Normal file
227
src/services/provider/pool/dimensions/_helpers.py
Normal 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",
|
||||
]
|
||||
52
src/services/provider/pool/dimensions/free_team_first.py
Normal file
52
src/services/provider/pool/dimensions/free_team_first.py
Normal 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())
|
||||
41
src/services/provider/pool/dimensions/quota_balanced.py
Normal file
41
src/services/provider/pool/dimensions/quota_balanced.py
Normal 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())
|
||||
45
src/services/provider/pool/dimensions/recent_refresh.py
Normal file
45
src/services/provider/pool/dimensions/recent_refresh.py
Normal 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())
|
||||
217
src/services/provider/pool/dimensions/registry.py
Normal file
217
src/services/provider/pool/dimensions/registry.py
Normal 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",
|
||||
]
|
||||
36
src/services/provider/pool/dimensions/single_account.py
Normal file
36
src/services/provider/pool/dimensions/single_account.py
Normal 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())
|
||||
91
src/services/provider/pool/health_cache.py
Normal file
91
src/services/provider/pool/health_cache.py
Normal 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",
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
|
||||
8
src/services/provider/pool/strategies/__init__.py
Normal file
8
src/services/provider/pool/strategies/__init__.py
Normal 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"]
|
||||
178
src/services/provider/pool/strategies/multi_score.py
Normal file
178
src/services/provider/pool/strategies/multi_score.py
Normal 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",
|
||||
]
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
92
tests/services/test_pool_account_state.py
Normal file
92
tests/services/test_pool_account_state.py
Normal 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"
|
||||
@@ -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
|
||||
|
||||
58
tests/services/test_pool_health_cache.py
Normal file
58
tests/services/test_pool_health_cache.py
Normal 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
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
233
tests/services/test_pool_multi_score_strategy.py
Normal file
233
tests/services/test_pool_multi_score_strategy.py
Normal 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
|
||||
83
tests/services/test_pool_preset_dimensions.py
Normal file
83
tests/services/test_pool_preset_dimensions.py
Normal 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
|
||||
88
tests/services/test_pool_redis_latency_ops.py
Normal file
88
tests/services/test_pool_redis_latency_ops.py
Normal 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
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user