mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
feat(provider): 增加 Claude Code 适配器、高级配置能力与 OAuth 账号类型统一解析
- 新增 Claude Code provider adapter (context, envelope, plugin, constants) - 扩展 provider admin 路由,支持 Claude Code 高级配置 (CRUD) - 统一 OAuth 账号类型解析逻辑,前后端对齐 - 重构 BatchAssignModelsDialog / ModelMappingDialog,简化组件逻辑 - handler 基类增强: request_builder 支持 Claude Code 信封格式 - CLI stream/sync mixin 适配 Claude Code 流式与同步模式 - 扩展 candidate builder / failover / scheduler 对 Claude Code 的支持 - 前端增加请求时间线可视化 (HorizontalRequestTimeline) - 补充 Claude Code envelope / runtime controls / distributed sessions 等测试 Closes #183 Closes #185 Co-Authored-By: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
@@ -7,3 +7,4 @@ export * from './health'
|
|||||||
export * from './models'
|
export * from './models'
|
||||||
export * from './adaptive'
|
export * from './adaptive'
|
||||||
export * from './global-models'
|
export * from './global-models'
|
||||||
|
export * from './pool'
|
||||||
|
|||||||
@@ -1,5 +1,10 @@
|
|||||||
import client from '../client'
|
import client from '../client'
|
||||||
import type { ProviderWithEndpointsSummary, ProxyConfig } from './types'
|
import type {
|
||||||
|
ClaudeCodeAdvancedConfig,
|
||||||
|
PoolAdvancedConfig,
|
||||||
|
ProviderWithEndpointsSummary,
|
||||||
|
ProxyConfig,
|
||||||
|
} from './types'
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 获取 Providers 摘要(包含 Endpoints 统计)
|
* 获取 Providers 摘要(包含 Endpoints 统计)
|
||||||
@@ -41,6 +46,8 @@ export async function updateProvider(
|
|||||||
max_probe_interval_minutes: number
|
max_probe_interval_minutes: number
|
||||||
enable_format_conversion: boolean // 是否允许格式转换(提供商级别开关)
|
enable_format_conversion: boolean // 是否允许格式转换(提供商级别开关)
|
||||||
is_active: boolean
|
is_active: boolean
|
||||||
|
claude_code_advanced: ClaudeCodeAdvancedConfig | null
|
||||||
|
pool_advanced: PoolAdvancedConfig | null
|
||||||
}>
|
}>
|
||||||
): Promise<ProviderWithEndpointsSummary> {
|
): Promise<ProviderWithEndpointsSummary> {
|
||||||
const response = await client.patch(`/api/admin/providers/${providerId}`, data)
|
const response = await client.patch(`/api/admin/providers/${providerId}`, data)
|
||||||
@@ -68,6 +75,8 @@ export async function createProvider(
|
|||||||
stream_first_byte_timeout?: number | null
|
stream_first_byte_timeout?: number | null
|
||||||
request_timeout?: number | null
|
request_timeout?: number | null
|
||||||
proxy?: ProxyConfig | null
|
proxy?: ProxyConfig | null
|
||||||
|
claude_code_advanced?: ClaudeCodeAdvancedConfig | null
|
||||||
|
pool_advanced?: PoolAdvancedConfig | null
|
||||||
}
|
}
|
||||||
): Promise<{ id: string; name: string; message?: string }> {
|
): Promise<{ id: string; name: string; message?: string }> {
|
||||||
const response = await client.post('/api/admin/providers/', data)
|
const response = await client.post('/api/admin/providers/', data)
|
||||||
|
|||||||
@@ -442,6 +442,30 @@ export interface PublicEndpointStatusMonitorResponse {
|
|||||||
|
|
||||||
export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro'
|
export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro'
|
||||||
|
|
||||||
|
export interface ClaudeCodeAdvancedConfig {
|
||||||
|
// 会话数量控制:null/undefined 表示不限制
|
||||||
|
max_sessions?: number | null
|
||||||
|
session_idle_timeout_minutes?: number | null
|
||||||
|
// TLS 指纹模拟(模拟 Node.js/Claude Code 客户端指纹)
|
||||||
|
enable_tls_fingerprint?: boolean
|
||||||
|
// 会话 ID 伪装(固定 metadata.user_id 中 session 片段)
|
||||||
|
session_id_masking_enabled?: boolean
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface PoolAdvancedConfig {
|
||||||
|
sticky_session_ttl_seconds?: number | null
|
||||||
|
load_threshold_percent?: number | null
|
||||||
|
lru_enabled?: boolean
|
||||||
|
cost_window_seconds?: number | null
|
||||||
|
cost_limit_per_key_tokens?: number | null
|
||||||
|
cost_soft_threshold_percent?: number | null
|
||||||
|
rate_limit_cooldown_seconds?: number | null
|
||||||
|
overload_cooldown_seconds?: number | null
|
||||||
|
proactive_refresh_seconds?: number | null
|
||||||
|
health_policy_enabled?: boolean
|
||||||
|
unschedulable_rules?: Array<Record<string, unknown>> | null
|
||||||
|
}
|
||||||
|
|
||||||
export interface ProviderWithEndpointsSummary {
|
export interface ProviderWithEndpointsSummary {
|
||||||
id: string
|
id: string
|
||||||
name: string
|
name: string
|
||||||
@@ -475,6 +499,8 @@ export interface ProviderWithEndpointsSummary {
|
|||||||
unhealthy_endpoints: number
|
unhealthy_endpoints: number
|
||||||
api_formats: string[]
|
api_formats: string[]
|
||||||
endpoint_health_details: EndpointHealthDetail[]
|
endpoint_health_details: EndpointHealthDetail[]
|
||||||
|
claude_code_advanced?: ClaudeCodeAdvancedConfig | null
|
||||||
|
pool_advanced?: PoolAdvancedConfig | null
|
||||||
ops_configured: boolean // 是否配置了扩展操作(余额监控等)
|
ops_configured: boolean // 是否配置了扩展操作(余额监控等)
|
||||||
ops_architecture_id?: string // 扩展操作使用的架构 ID(如 cubence, anyrouter)
|
ops_architecture_id?: string // 扩展操作使用的架构 ID(如 cubence, anyrouter)
|
||||||
created_at: string
|
created_at: string
|
||||||
|
|||||||
@@ -19,35 +19,9 @@
|
|||||||
class="pl-8 h-9"
|
class="pl-8 h-9"
|
||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
<button
|
|
||||||
v-if="upstreamModelsLoaded"
|
|
||||||
type="button"
|
|
||||||
class="p-2 hover:bg-muted rounded-md transition-colors shrink-0"
|
|
||||||
title="刷新上游模型"
|
|
||||||
:disabled="fetchingUpstreamModels"
|
|
||||||
@click="fetchUpstreamModels(true)"
|
|
||||||
>
|
|
||||||
<RefreshCw
|
|
||||||
class="w-4 h-4"
|
|
||||||
:class="{ 'animate-spin': fetchingUpstreamModels }"
|
|
||||||
/>
|
|
||||||
</button>
|
|
||||||
<button
|
|
||||||
v-else-if="!fetchingUpstreamModels"
|
|
||||||
type="button"
|
|
||||||
class="p-2 hover:bg-muted rounded-md transition-colors shrink-0"
|
|
||||||
title="从提供商获取模型"
|
|
||||||
@click="fetchUpstreamModels()"
|
|
||||||
>
|
|
||||||
<Zap class="w-4 h-4" />
|
|
||||||
</button>
|
|
||||||
<Loader2
|
|
||||||
v-else
|
|
||||||
class="w-4 h-4 animate-spin text-muted-foreground shrink-0"
|
|
||||||
/>
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 单列模型列表 -->
|
<!-- 模型列表 -->
|
||||||
<div class="border rounded-lg overflow-hidden">
|
<div class="border rounded-lg overflow-hidden">
|
||||||
<div class="max-h-96 overflow-y-auto">
|
<div class="max-h-96 overflow-y-auto">
|
||||||
<div
|
<div
|
||||||
@@ -58,17 +32,12 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<template v-else>
|
<template v-else>
|
||||||
<!-- 全局模型组 -->
|
<!-- 全局模型列表 -->
|
||||||
<div v-if="filteredGlobalModels.length > 0 || !upstreamModelsLoaded">
|
<div v-if="filteredGlobalModels.length > 0">
|
||||||
<div
|
<div
|
||||||
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-10 cursor-pointer hover:bg-muted/80 transition-colors"
|
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-10"
|
||||||
@click="toggleGroupCollapse('global')"
|
|
||||||
>
|
>
|
||||||
<div class="flex items-center gap-2">
|
<div class="flex items-center gap-2">
|
||||||
<ChevronDown
|
|
||||||
class="w-4 h-4 transition-transform shrink-0"
|
|
||||||
:class="collapsedGroups.has('global') ? '-rotate-90' : ''"
|
|
||||||
/>
|
|
||||||
<span class="text-xs font-medium">全局模型</span>
|
<span class="text-xs font-medium">全局模型</span>
|
||||||
<span class="text-xs text-muted-foreground">({{ filteredGlobalModels.length }})</span>
|
<span class="text-xs text-muted-foreground">({{ filteredGlobalModels.length }})</span>
|
||||||
</div>
|
</div>
|
||||||
@@ -81,16 +50,7 @@
|
|||||||
{{ isAllGlobalModelsSelected ? '取消全选' : '全选' }}
|
{{ isAllGlobalModelsSelected ? '取消全选' : '全选' }}
|
||||||
</button>
|
</button>
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div class="space-y-1 p-2">
|
||||||
v-show="!collapsedGroups.has('global')"
|
|
||||||
class="space-y-1 p-2"
|
|
||||||
>
|
|
||||||
<div
|
|
||||||
v-if="filteredGlobalModels.length === 0"
|
|
||||||
class="py-4 text-center text-xs text-muted-foreground"
|
|
||||||
>
|
|
||||||
暂无可用全局模型
|
|
||||||
</div>
|
|
||||||
<div
|
<div
|
||||||
v-for="model in filteredGlobalModels"
|
v-for="model in filteredGlobalModels"
|
||||||
:key="model.id"
|
:key="model.id"
|
||||||
@@ -118,85 +78,17 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 上游模型组 -->
|
|
||||||
<div v-if="filteredUpstreamModels.length > 0">
|
|
||||||
<div
|
|
||||||
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-10 cursor-pointer hover:bg-muted/80 transition-colors"
|
|
||||||
@click="toggleGroupCollapse('upstream')"
|
|
||||||
>
|
|
||||||
<div class="flex items-center gap-2">
|
|
||||||
<ChevronDown
|
|
||||||
class="w-4 h-4 transition-transform shrink-0"
|
|
||||||
:class="collapsedGroups.has('upstream') ? '-rotate-90' : ''"
|
|
||||||
/>
|
|
||||||
<span class="text-xs font-medium">上游模型</span>
|
|
||||||
<span class="text-xs text-muted-foreground">({{ filteredUpstreamModels.length }})</span>
|
|
||||||
</div>
|
|
||||||
<button
|
|
||||||
type="button"
|
|
||||||
class="text-xs text-primary hover:underline shrink-0"
|
|
||||||
@click.stop="toggleAllUpstreamModels"
|
|
||||||
>
|
|
||||||
{{ isAllUpstreamModelsSelected ? '取消全选' : '全选' }}
|
|
||||||
</button>
|
|
||||||
</div>
|
|
||||||
<div
|
|
||||||
v-show="!collapsedGroups.has('upstream')"
|
|
||||||
class="space-y-1 p-2"
|
|
||||||
>
|
|
||||||
<div
|
|
||||||
v-for="model in filteredUpstreamModels"
|
|
||||||
:key="model.id"
|
|
||||||
class="flex items-center gap-2 px-2 py-1.5 rounded hover:bg-muted cursor-pointer"
|
|
||||||
@click="toggleUpstreamModelSelection(model.id)"
|
|
||||||
>
|
|
||||||
<div
|
|
||||||
class="w-4 h-4 border rounded flex items-center justify-center shrink-0"
|
|
||||||
:class="isUpstreamModelSelected(model.id) ? 'bg-primary border-primary' : ''"
|
|
||||||
>
|
|
||||||
<Check
|
|
||||||
v-if="isUpstreamModelSelected(model.id)"
|
|
||||||
class="w-3 h-3 text-primary-foreground"
|
|
||||||
/>
|
|
||||||
</div>
|
|
||||||
<div class="flex-1 min-w-0">
|
|
||||||
<div class="flex items-center gap-1.5">
|
|
||||||
<p class="text-sm font-medium truncate">
|
|
||||||
{{ model.id }}
|
|
||||||
</p>
|
|
||||||
<span
|
|
||||||
v-for="fmt in model.api_formats"
|
|
||||||
:key="fmt"
|
|
||||||
class="text-[10px] px-1 py-0.5 rounded bg-muted text-muted-foreground shrink-0"
|
|
||||||
>
|
|
||||||
{{ formatApiFormat(fmt) }}
|
|
||||||
</span>
|
|
||||||
</div>
|
|
||||||
<p
|
|
||||||
v-if="model.owned_by"
|
|
||||||
class="text-xs text-muted-foreground truncate"
|
|
||||||
>
|
|
||||||
{{ model.owned_by }}
|
|
||||||
</p>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<!-- 空状态 -->
|
<!-- 空状态 -->
|
||||||
<div
|
<div
|
||||||
v-if="filteredGlobalModels.length === 0 && filteredUpstreamModels.length === 0"
|
v-if="filteredGlobalModels.length === 0"
|
||||||
class="flex flex-col items-center justify-center py-12 text-muted-foreground"
|
class="flex flex-col items-center justify-center py-12 text-muted-foreground"
|
||||||
>
|
>
|
||||||
<Layers class="w-10 h-10 mb-2 opacity-30" />
|
<Layers class="w-10 h-10 mb-2 opacity-30" />
|
||||||
<p class="text-sm">
|
<p class="text-sm">
|
||||||
{{ searchQuery ? '无匹配结果' : '暂无可用模型' }}
|
{{ searchQuery ? '无匹配结果' : '暂无可用全局模型' }}
|
||||||
</p>
|
</p>
|
||||||
<p
|
<p class="text-xs mt-1">
|
||||||
v-if="!upstreamModelsLoaded"
|
请先前往"模型目录"页面创建全局模型
|
||||||
class="text-xs mt-1"
|
|
||||||
>
|
|
||||||
点击上方按钮从上游获取模型
|
|
||||||
</p>
|
</p>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
@@ -234,7 +126,7 @@
|
|||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, computed, watch } from 'vue'
|
import { ref, computed, watch } from 'vue'
|
||||||
import { Layers, Loader2, ChevronDown, Zap, RefreshCw, Search, Check } from 'lucide-vue-next'
|
import { Layers, Loader2, Search, Check } from 'lucide-vue-next'
|
||||||
import Dialog from '@/components/ui/dialog/Dialog.vue'
|
import Dialog from '@/components/ui/dialog/Dialog.vue'
|
||||||
import Button from '@/components/ui/button.vue'
|
import Button from '@/components/ui/button.vue'
|
||||||
import Input from '@/components/ui/input.vue'
|
import Input from '@/components/ui/input.vue'
|
||||||
@@ -249,11 +141,8 @@ import {
|
|||||||
getProviderModels,
|
getProviderModels,
|
||||||
batchAssignModelsToProvider,
|
batchAssignModelsToProvider,
|
||||||
deleteModel,
|
deleteModel,
|
||||||
importModelsFromUpstream,
|
|
||||||
type Model
|
type Model
|
||||||
} from '@/api/endpoints'
|
} from '@/api/endpoints'
|
||||||
import { formatApiFormat } from '@/api/endpoints/types/api-format'
|
|
||||||
import { useUpstreamModelsCache, type UpstreamModel } from '../composables/useUpstreamModelsCache'
|
|
||||||
|
|
||||||
const props = defineProps<{
|
const props = defineProps<{
|
||||||
open: boolean
|
open: boolean
|
||||||
@@ -266,32 +155,22 @@ const emit = defineEmits<{
|
|||||||
'changed': []
|
'changed': []
|
||||||
}>()
|
}>()
|
||||||
|
|
||||||
const { fetchModels: fetchCachedModels } = useUpstreamModelsCache()
|
|
||||||
|
|
||||||
const { error: showError, success } = useToast()
|
const { error: showError, success } = useToast()
|
||||||
const { confirmWarning } = useConfirm()
|
const { confirmWarning } = useConfirm()
|
||||||
|
|
||||||
// 状态
|
// 状态
|
||||||
const loadingGlobalModels = ref(false)
|
const loadingGlobalModels = ref(false)
|
||||||
const fetchingUpstreamModels = ref(false)
|
|
||||||
const upstreamModelsLoaded = ref(false)
|
|
||||||
const saving = ref(false)
|
const saving = ref(false)
|
||||||
|
|
||||||
// 数据
|
// 数据
|
||||||
const allGlobalModels = ref<GlobalModelResponse[]>([])
|
const allGlobalModels = ref<GlobalModelResponse[]>([])
|
||||||
const existingModels = ref<Model[]>([])
|
const existingModels = ref<Model[]>([])
|
||||||
const upstreamModels = ref<UpstreamModel[]>([])
|
|
||||||
|
|
||||||
// 选择状态(本地状态,保存时才提交)
|
// 选择状态(本地状态,保存时才提交)
|
||||||
const selectedGlobalModelIds = ref<Set<string>>(new Set())
|
const selectedGlobalModelIds = ref<Set<string>>(new Set())
|
||||||
const selectedUpstreamModelIds = ref<Set<string>>(new Set())
|
|
||||||
|
|
||||||
// 初始状态(用于计算变更)
|
// 初始状态(用于计算变更)
|
||||||
const initialGlobalModelIds = ref<Set<string>>(new Set())
|
const initialGlobalModelIds = ref<Set<string>>(new Set())
|
||||||
const initialUpstreamModelNames = ref<Set<string>>(new Set())
|
|
||||||
|
|
||||||
// 折叠状态
|
|
||||||
const collapsedGroups = ref<Set<string>>(new Set())
|
|
||||||
|
|
||||||
// 搜索状态
|
// 搜索状态
|
||||||
const searchQuery = ref('')
|
const searchQuery = ref('')
|
||||||
@@ -305,18 +184,6 @@ const existingGlobalModelIds = computed(() => {
|
|||||||
)
|
)
|
||||||
})
|
})
|
||||||
|
|
||||||
// 已关联的上游模型名称集合
|
|
||||||
const existingUpstreamModelNames = computed(() => {
|
|
||||||
const names = new Set<string>()
|
|
||||||
for (const m of existingModels.value) {
|
|
||||||
names.add(m.provider_model_name)
|
|
||||||
for (const mapping of m.provider_model_mappings ?? []) {
|
|
||||||
if (mapping.name) names.add(mapping.name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return names
|
|
||||||
})
|
|
||||||
|
|
||||||
// 过滤后的全局模型
|
// 过滤后的全局模型
|
||||||
const filteredGlobalModels = computed(() => {
|
const filteredGlobalModels = computed(() => {
|
||||||
const query = searchQuery.value.toLowerCase().trim()
|
const query = searchQuery.value.toLowerCase().trim()
|
||||||
@@ -328,59 +195,17 @@ const filteredGlobalModels = computed(() => {
|
|||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
// 过滤后的上游模型(后端已按 id 聚合)
|
|
||||||
const filteredUpstreamModels = computed(() => {
|
|
||||||
if (!upstreamModelsLoaded.value) return []
|
|
||||||
|
|
||||||
const query = searchQuery.value.toLowerCase().trim()
|
|
||||||
let models = upstreamModels.value
|
|
||||||
|
|
||||||
if (query) {
|
|
||||||
models = models.filter(m => m.id.toLowerCase().includes(query))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 按 id 排序
|
|
||||||
return [...models].sort((a, b) => a.id.localeCompare(b.id))
|
|
||||||
})
|
|
||||||
|
|
||||||
// 上游模型是否全选
|
|
||||||
const isAllUpstreamModelsSelected = computed(() => {
|
|
||||||
if (filteredUpstreamModels.value.length === 0) return false
|
|
||||||
return filteredUpstreamModels.value.every(m => selectedUpstreamModelIds.value.has(m.id))
|
|
||||||
})
|
|
||||||
|
|
||||||
// 全选/取消全选上游模型
|
|
||||||
function toggleAllUpstreamModels() {
|
|
||||||
const allIds = filteredUpstreamModels.value.map(m => m.id)
|
|
||||||
if (isAllUpstreamModelsSelected.value) {
|
|
||||||
// 取消全选
|
|
||||||
for (const id of allIds) {
|
|
||||||
selectedUpstreamModelIds.value.delete(id)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// 全选
|
|
||||||
for (const id of allIds) {
|
|
||||||
selectedUpstreamModelIds.value.add(id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 检查全局模型是否已选中
|
|
||||||
function isGlobalModelSelected(globalModelId: string): boolean {
|
|
||||||
return selectedGlobalModelIds.value.has(globalModelId)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 检查上游模型是否已选中
|
|
||||||
function isUpstreamModelSelected(modelId: string): boolean {
|
|
||||||
return selectedUpstreamModelIds.value.has(modelId)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 全局模型是否全选
|
// 全局模型是否全选
|
||||||
const isAllGlobalModelsSelected = computed(() => {
|
const isAllGlobalModelsSelected = computed(() => {
|
||||||
if (filteredGlobalModels.value.length === 0) return false
|
if (filteredGlobalModels.value.length === 0) return false
|
||||||
return filteredGlobalModels.value.every(m => isGlobalModelSelected(m.id))
|
return filteredGlobalModels.value.every(m => isGlobalModelSelected(m.id))
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// 检查全局模型是否已选中
|
||||||
|
function isGlobalModelSelected(globalModelId: string): boolean {
|
||||||
|
return selectedGlobalModelIds.value.has(globalModelId)
|
||||||
|
}
|
||||||
|
|
||||||
// 计算待添加的全局模型
|
// 计算待添加的全局模型
|
||||||
const globalModelsToAdd = computed(() => {
|
const globalModelsToAdd = computed(() => {
|
||||||
const toAdd: string[] = []
|
const toAdd: string[] = []
|
||||||
@@ -403,42 +228,16 @@ const globalModelsToRemove = computed(() => {
|
|||||||
return toRemove
|
return toRemove
|
||||||
})
|
})
|
||||||
|
|
||||||
// 计算待添加的上游模型
|
|
||||||
const upstreamModelsToAdd = computed(() => {
|
|
||||||
const toAdd: string[] = []
|
|
||||||
for (const id of selectedUpstreamModelIds.value) {
|
|
||||||
if (!initialUpstreamModelNames.value.has(id)) {
|
|
||||||
toAdd.push(id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return toAdd
|
|
||||||
})
|
|
||||||
|
|
||||||
// 计算待移除的上游模型
|
|
||||||
const upstreamModelsToRemove = computed(() => {
|
|
||||||
const toRemove: string[] = []
|
|
||||||
for (const id of initialUpstreamModelNames.value) {
|
|
||||||
if (!selectedUpstreamModelIds.value.has(id)) {
|
|
||||||
toRemove.push(id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return toRemove
|
|
||||||
})
|
|
||||||
|
|
||||||
// 是否有变更
|
// 是否有变更
|
||||||
const hasChanges = computed(() => {
|
const hasChanges = computed(() => {
|
||||||
return globalModelsToAdd.value.length > 0 ||
|
return globalModelsToAdd.value.length > 0 ||
|
||||||
globalModelsToRemove.value.length > 0 ||
|
globalModelsToRemove.value.length > 0
|
||||||
upstreamModelsToAdd.value.length > 0 ||
|
|
||||||
upstreamModelsToRemove.value.length > 0
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// 待变更数量
|
// 待变更数量
|
||||||
const pendingChangesCount = computed(() => {
|
const pendingChangesCount = computed(() => {
|
||||||
return globalModelsToAdd.value.length +
|
return globalModelsToAdd.value.length +
|
||||||
globalModelsToRemove.value.length +
|
globalModelsToRemove.value.length
|
||||||
upstreamModelsToAdd.value.length +
|
|
||||||
upstreamModelsToRemove.value.length
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// 切换全局模型选择
|
// 切换全局模型选择
|
||||||
@@ -451,26 +250,14 @@ function toggleGlobalModelSelection(id: string) {
|
|||||||
selectedGlobalModelIds.value = new Set(selectedGlobalModelIds.value)
|
selectedGlobalModelIds.value = new Set(selectedGlobalModelIds.value)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 切换上游模型选择
|
|
||||||
function toggleUpstreamModelSelection(id: string) {
|
|
||||||
if (selectedUpstreamModelIds.value.has(id)) {
|
|
||||||
selectedUpstreamModelIds.value.delete(id)
|
|
||||||
} else {
|
|
||||||
selectedUpstreamModelIds.value.add(id)
|
|
||||||
}
|
|
||||||
selectedUpstreamModelIds.value = new Set(selectedUpstreamModelIds.value)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 全选/取消全选全局模型
|
// 全选/取消全选全局模型
|
||||||
function toggleAllGlobalModels() {
|
function toggleAllGlobalModels() {
|
||||||
const allIds = filteredGlobalModels.value.map(m => m.id)
|
const allIds = filteredGlobalModels.value.map(m => m.id)
|
||||||
if (isAllGlobalModelsSelected.value) {
|
if (isAllGlobalModelsSelected.value) {
|
||||||
// 取消全选
|
|
||||||
for (const id of allIds) {
|
for (const id of allIds) {
|
||||||
selectedGlobalModelIds.value.delete(id)
|
selectedGlobalModelIds.value.delete(id)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// 全选
|
|
||||||
for (const id of allIds) {
|
for (const id of allIds) {
|
||||||
selectedGlobalModelIds.value.add(id)
|
selectedGlobalModelIds.value.add(id)
|
||||||
}
|
}
|
||||||
@@ -478,16 +265,6 @@ function toggleAllGlobalModels() {
|
|||||||
selectedGlobalModelIds.value = new Set(selectedGlobalModelIds.value)
|
selectedGlobalModelIds.value = new Set(selectedGlobalModelIds.value)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 切换折叠状态
|
|
||||||
function toggleGroupCollapse(group: string) {
|
|
||||||
if (collapsedGroups.value.has(group)) {
|
|
||||||
collapsedGroups.value.delete(group)
|
|
||||||
} else {
|
|
||||||
collapsedGroups.value.add(group)
|
|
||||||
}
|
|
||||||
collapsedGroups.value = new Set(collapsedGroups.value)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 处理关闭
|
// 处理关闭
|
||||||
async function handleClose() {
|
async function handleClose() {
|
||||||
if (hasChanges.value) {
|
if (hasChanges.value) {
|
||||||
@@ -530,23 +307,6 @@ async function handleSave() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 移除上游模型
|
|
||||||
for (const modelId of upstreamModelsToRemove.value) {
|
|
||||||
const existingModel = existingModels.value.find(m =>
|
|
||||||
m.provider_model_name === modelId ||
|
|
||||||
m.provider_model_mappings?.some(mapping => mapping.name === modelId)
|
|
||||||
)
|
|
||||||
if (existingModel) {
|
|
||||||
hasAnyOperation = true
|
|
||||||
try {
|
|
||||||
await deleteModel(props.providerId, existingModel.id)
|
|
||||||
totalSuccess++
|
|
||||||
} catch (err: unknown) {
|
|
||||||
allErrors.push(parseApiError(err, '移除失败'))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 添加全局模型
|
// 添加全局模型
|
||||||
if (globalModelsToAdd.value.length > 0) {
|
if (globalModelsToAdd.value.length > 0) {
|
||||||
hasAnyOperation = true
|
hasAnyOperation = true
|
||||||
@@ -561,20 +321,6 @@ async function handleSave() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 添加上游模型
|
|
||||||
if (upstreamModelsToAdd.value.length > 0) {
|
|
||||||
hasAnyOperation = true
|
|
||||||
try {
|
|
||||||
const result = await importModelsFromUpstream(props.providerId, upstreamModelsToAdd.value)
|
|
||||||
totalSuccess += result.success.length
|
|
||||||
if (result.errors.length > 0) {
|
|
||||||
allErrors.push(...result.errors.map(e => e.error))
|
|
||||||
}
|
|
||||||
} catch (err: unknown) {
|
|
||||||
allErrors.push(parseApiError(err, '导入上游模型失败'))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if (totalSuccess > 0) {
|
if (totalSuccess > 0) {
|
||||||
success(`成功处理 ${totalSuccess} 个模型`)
|
success(`成功处理 ${totalSuccess} 个模型`)
|
||||||
}
|
}
|
||||||
@@ -587,7 +333,6 @@ async function handleSave() {
|
|||||||
emit('update:open', false)
|
emit('update:open', false)
|
||||||
} catch (err: unknown) {
|
} catch (err: unknown) {
|
||||||
showError(parseApiError(err, '保存失败'), '错误')
|
showError(parseApiError(err, '保存失败'), '错误')
|
||||||
// 即使出错,如果已执行过操作,也通知父组件刷新数据
|
|
||||||
if (hasAnyOperation) {
|
if (hasAnyOperation) {
|
||||||
emit('changed')
|
emit('changed')
|
||||||
}
|
}
|
||||||
@@ -596,52 +341,28 @@ async function handleSave() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 从已有数据同步选择状态(全局模型)
|
// 从已有数据同步选择状态
|
||||||
function syncGlobalModelSelection() {
|
function syncGlobalModelSelection() {
|
||||||
const globalIds = [...existingGlobalModelIds.value].filter((id): id is string => id !== undefined)
|
const globalIds = [...existingGlobalModelIds.value].filter((id): id is string => id !== undefined)
|
||||||
selectedGlobalModelIds.value = new Set(globalIds)
|
selectedGlobalModelIds.value = new Set(globalIds)
|
||||||
initialGlobalModelIds.value = new Set(globalIds)
|
initialGlobalModelIds.value = new Set(globalIds)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 从已有数据同步选择状态(上游模型)
|
|
||||||
function syncUpstreamModelSelection() {
|
|
||||||
// 只同步当前已加载的上游模型中,与已关联模型匹配的部分
|
|
||||||
const selected = new Set<string>()
|
|
||||||
for (const model of upstreamModels.value) {
|
|
||||||
if (existingUpstreamModelNames.value.has(model.id)) {
|
|
||||||
selected.add(model.id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
selectedUpstreamModelIds.value = selected
|
|
||||||
initialUpstreamModelNames.value = new Set(selected)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 监听打开状态
|
// 监听打开状态
|
||||||
watch(() => props.open, async (isOpen) => {
|
watch(() => props.open, async (isOpen) => {
|
||||||
if (isOpen && props.providerId) {
|
if (isOpen && props.providerId) {
|
||||||
await loadData()
|
await loadData()
|
||||||
} else {
|
} else {
|
||||||
// 重置状态
|
|
||||||
upstreamModels.value = []
|
|
||||||
upstreamModelsLoaded.value = false
|
|
||||||
collapsedGroups.value = new Set()
|
|
||||||
searchQuery.value = ''
|
searchQuery.value = ''
|
||||||
selectedGlobalModelIds.value = new Set()
|
selectedGlobalModelIds.value = new Set()
|
||||||
selectedUpstreamModelIds.value = new Set()
|
|
||||||
initialGlobalModelIds.value = new Set()
|
initialGlobalModelIds.value = new Set()
|
||||||
initialUpstreamModelNames.value = new Set()
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
// 加载数据
|
// 加载数据
|
||||||
async function loadData() {
|
async function loadData() {
|
||||||
await Promise.all([loadGlobalModels(), loadExistingModels()])
|
await Promise.all([loadGlobalModels(), loadExistingModels()])
|
||||||
|
|
||||||
// 同步全局模型选择状态
|
|
||||||
syncGlobalModelSelection()
|
syncGlobalModelSelection()
|
||||||
|
|
||||||
// 初始折叠状态
|
|
||||||
collapsedGroups.value = new Set()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 加载全局模型列表
|
// 加载全局模型列表
|
||||||
@@ -665,26 +386,4 @@ async function loadExistingModels() {
|
|||||||
showError(parseApiError(err, '加载已关联模型失败'), '错误')
|
showError(parseApiError(err, '加载已关联模型失败'), '错误')
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 从提供商获取模型
|
|
||||||
async function fetchUpstreamModels(forceRefresh = false) {
|
|
||||||
try {
|
|
||||||
fetchingUpstreamModels.value = true
|
|
||||||
const result = await fetchCachedModels(props.providerId, undefined, forceRefresh)
|
|
||||||
if (result) {
|
|
||||||
if (result.error) {
|
|
||||||
showError(result.error, '错误')
|
|
||||||
} else {
|
|
||||||
upstreamModels.value = result.models
|
|
||||||
upstreamModelsLoaded.value = true
|
|
||||||
// 同步上游模型选择状态
|
|
||||||
syncUpstreamModelSelection()
|
|
||||||
// 全部折叠
|
|
||||||
collapsedGroups.value = new Set(['global', 'upstream'])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} finally {
|
|
||||||
fetchingUpstreamModels.value = false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
</script>
|
</script>
|
||||||
|
|||||||
@@ -63,33 +63,6 @@
|
|||||||
>
|
>
|
||||||
已选 {{ selectedNames.length }} 个
|
已选 {{ selectedNames.length }} 个
|
||||||
</span>
|
</span>
|
||||||
<!-- 刷新上游模型按钮 -->
|
|
||||||
<button
|
|
||||||
v-if="upstreamModelsLoaded"
|
|
||||||
type="button"
|
|
||||||
class="p-2 hover:bg-muted rounded-md transition-colors shrink-0"
|
|
||||||
:disabled="fetchingUpstreamModels"
|
|
||||||
title="刷新上游模型"
|
|
||||||
@click="fetchUpstreamModels()"
|
|
||||||
>
|
|
||||||
<RefreshCw
|
|
||||||
class="w-4 h-4"
|
|
||||||
:class="{ 'animate-spin': fetchingUpstreamModels }"
|
|
||||||
/>
|
|
||||||
</button>
|
|
||||||
<button
|
|
||||||
v-else-if="!fetchingUpstreamModels"
|
|
||||||
type="button"
|
|
||||||
class="p-2 hover:bg-muted rounded-md transition-colors shrink-0"
|
|
||||||
title="从提供商获取模型"
|
|
||||||
@click="fetchUpstreamModels()"
|
|
||||||
>
|
|
||||||
<Zap class="w-4 h-4" />
|
|
||||||
</button>
|
|
||||||
<Loader2
|
|
||||||
v-else
|
|
||||||
class="w-4 h-4 animate-spin text-muted-foreground shrink-0"
|
|
||||||
/>
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 模型列表 -->
|
<!-- 模型列表 -->
|
||||||
@@ -124,22 +97,14 @@
|
|||||||
<!-- 自定义映射名称 -->
|
<!-- 自定义映射名称 -->
|
||||||
<div v-if="customNames.length > 0">
|
<div v-if="customNames.length > 0">
|
||||||
<div
|
<div
|
||||||
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-20 cursor-pointer hover:bg-muted/80 transition-colors"
|
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-20"
|
||||||
@click="toggleGroupCollapse('custom')"
|
|
||||||
>
|
>
|
||||||
<div class="flex items-center gap-2">
|
<div class="flex items-center gap-2">
|
||||||
<ChevronDown
|
|
||||||
class="w-4 h-4 transition-transform shrink-0"
|
|
||||||
:class="collapsedGroups.has('custom') ? '-rotate-90' : ''"
|
|
||||||
/>
|
|
||||||
<span class="text-xs font-medium">自定义模型</span>
|
<span class="text-xs font-medium">自定义模型</span>
|
||||||
<span class="text-xs text-muted-foreground">({{ customNames.length }})</span>
|
<span class="text-xs text-muted-foreground">({{ customNames.length }})</span>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div class="space-y-1 p-2">
|
||||||
v-show="!collapsedGroups.has('custom')"
|
|
||||||
class="space-y-1 p-2"
|
|
||||||
>
|
|
||||||
<div
|
<div
|
||||||
v-for="name in sortedCustomNames"
|
v-for="name in sortedCustomNames"
|
||||||
:key="name"
|
:key="name"
|
||||||
@@ -160,52 +125,6 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 上游模型 -->
|
|
||||||
<template v-if="filteredUpstreamModels.length > 0">
|
|
||||||
<div
|
|
||||||
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-20 cursor-pointer hover:bg-muted/80 transition-colors"
|
|
||||||
@click="toggleGroupCollapse('upstream')"
|
|
||||||
>
|
|
||||||
<div class="flex items-center gap-2">
|
|
||||||
<ChevronDown
|
|
||||||
class="w-4 h-4 transition-transform shrink-0"
|
|
||||||
:class="collapsedGroups.has('upstream') ? '-rotate-90' : ''"
|
|
||||||
/>
|
|
||||||
<span class="text-xs font-medium">上游模型</span>
|
|
||||||
<span class="text-xs text-muted-foreground">({{ upstreamModelNames.length }})</span>
|
|
||||||
</div>
|
|
||||||
<button
|
|
||||||
type="button"
|
|
||||||
class="text-xs text-primary hover:underline"
|
|
||||||
@click.stop="toggleAllUpstreamModels"
|
|
||||||
>
|
|
||||||
{{ isAllUpstreamModelsSelected ? '取消全选' : '全选' }}
|
|
||||||
</button>
|
|
||||||
</div>
|
|
||||||
<div
|
|
||||||
v-show="!collapsedGroups.has('upstream')"
|
|
||||||
class="space-y-1 p-2"
|
|
||||||
>
|
|
||||||
<div
|
|
||||||
v-for="name in filteredUpstreamModels"
|
|
||||||
:key="name"
|
|
||||||
class="flex items-center gap-2 px-2 py-1.5 rounded hover:bg-muted cursor-pointer"
|
|
||||||
@click="toggleName(name)"
|
|
||||||
>
|
|
||||||
<div
|
|
||||||
class="w-4 h-4 border rounded flex items-center justify-center shrink-0"
|
|
||||||
:class="selectedNames.includes(name) ? 'bg-primary border-primary' : ''"
|
|
||||||
>
|
|
||||||
<Check
|
|
||||||
v-if="selectedNames.includes(name)"
|
|
||||||
class="w-3 h-3 text-primary-foreground"
|
|
||||||
/>
|
|
||||||
</div>
|
|
||||||
<span class="text-sm font-mono truncate flex-1">{{ name }}</span>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
</template>
|
|
||||||
|
|
||||||
<!-- 空状态 -->
|
<!-- 空状态 -->
|
||||||
<div
|
<div
|
||||||
v-if="showEmptyState"
|
v-if="showEmptyState"
|
||||||
@@ -248,7 +167,7 @@
|
|||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, computed, watch } from 'vue'
|
import { ref, computed, watch } from 'vue'
|
||||||
import { Tag, Loader2, Plus, Search, Check, ChevronDown, RefreshCw, Zap } from 'lucide-vue-next'
|
import { Tag, Loader2, Plus, Search, Check } from 'lucide-vue-next'
|
||||||
import {
|
import {
|
||||||
Button,
|
Button,
|
||||||
Input,
|
Input,
|
||||||
@@ -265,16 +184,14 @@ import { parseApiError } from '@/utils/errorParser'
|
|||||||
import {
|
import {
|
||||||
type Model,
|
type Model,
|
||||||
type ProviderModelAlias,
|
type ProviderModelAlias,
|
||||||
type UpstreamModel,
|
|
||||||
} from '@/api/endpoints'
|
} from '@/api/endpoints'
|
||||||
import { updateModel } from '@/api/endpoints/models'
|
import { updateModel } from '@/api/endpoints/models'
|
||||||
import { useUpstreamModelsCache } from '../composables/useUpstreamModelsCache'
|
|
||||||
|
|
||||||
export interface AliasGroup {
|
export interface AliasGroup {
|
||||||
model: Model
|
model: Model
|
||||||
/** @deprecated 作用域功能已废弃,将在后续版本移除 */
|
/** @deprecated */
|
||||||
apiFormatsKey: string
|
apiFormatsKey: string
|
||||||
/** @deprecated 作用域功能已废弃,将在后续版本移除 */
|
/** @deprecated */
|
||||||
apiFormats: string[]
|
apiFormats: string[]
|
||||||
aliases: ProviderModelAlias[]
|
aliases: ProviderModelAlias[]
|
||||||
}
|
}
|
||||||
@@ -282,12 +199,11 @@ export interface AliasGroup {
|
|||||||
const props = defineProps<{
|
const props = defineProps<{
|
||||||
open: boolean
|
open: boolean
|
||||||
providerId: string
|
providerId: string
|
||||||
/** @deprecated 作用域功能已废弃,此 prop 将在后续版本移除 */
|
/** @deprecated */
|
||||||
providerApiFormats?: string[]
|
providerApiFormats?: string[]
|
||||||
models: Model[]
|
models: Model[]
|
||||||
editingGroup?: AliasGroup | null
|
editingGroup?: AliasGroup | null
|
||||||
preselectedModelId?: string | null
|
preselectedModelId?: string | null
|
||||||
hasAutoFetchKey?: boolean
|
|
||||||
}>()
|
}>()
|
||||||
|
|
||||||
const emit = defineEmits<{
|
const emit = defineEmits<{
|
||||||
@@ -296,23 +212,14 @@ const emit = defineEmits<{
|
|||||||
}>()
|
}>()
|
||||||
|
|
||||||
const { error: showError, success: showSuccess } = useToast()
|
const { error: showError, success: showSuccess } = useToast()
|
||||||
const { fetchModels: fetchCachedModels } = useUpstreamModelsCache()
|
|
||||||
|
|
||||||
// 状态
|
// 状态
|
||||||
const submitting = ref(false)
|
const submitting = ref(false)
|
||||||
const loadingModels = ref(false)
|
const loadingModels = ref(false)
|
||||||
const fetchingUpstreamModels = ref(false)
|
|
||||||
const upstreamModelsLoaded = ref(false)
|
|
||||||
|
|
||||||
// 搜索
|
// 搜索
|
||||||
const searchQuery = ref('')
|
const searchQuery = ref('')
|
||||||
|
|
||||||
// 折叠状态
|
|
||||||
const collapsedGroups = ref<Set<string>>(new Set())
|
|
||||||
|
|
||||||
// 上游模型
|
|
||||||
const upstreamModels = ref<UpstreamModel[]>([])
|
|
||||||
|
|
||||||
// 表单数据
|
// 表单数据
|
||||||
const formData = ref<{
|
const formData = ref<{
|
||||||
modelId: string
|
modelId: string
|
||||||
@@ -326,26 +233,9 @@ const selectedNames = ref<string[]>([])
|
|||||||
// 自定义名称列表(手动添加的)
|
// 自定义名称列表(手动添加的)
|
||||||
const allCustomNames = ref<string[]>([])
|
const allCustomNames = ref<string[]>([])
|
||||||
|
|
||||||
// 所有已知名称集合
|
// 自定义名称列表
|
||||||
const allKnownNames = computed(() => {
|
|
||||||
const set = new Set<string>()
|
|
||||||
upstreamModels.value.forEach(m => set.add(m.id))
|
|
||||||
return set
|
|
||||||
})
|
|
||||||
|
|
||||||
// 上游模型名称列表(去重后)
|
|
||||||
const upstreamModelNames = computed(() => {
|
|
||||||
const names = new Set<string>()
|
|
||||||
upstreamModels.value.forEach(m => {
|
|
||||||
names.add(m.id)
|
|
||||||
})
|
|
||||||
return Array.from(names).sort()
|
|
||||||
})
|
|
||||||
|
|
||||||
// 自定义名称列表(排除上游模型中已有的)
|
|
||||||
const customNames = computed(() => {
|
const customNames = computed(() => {
|
||||||
const upstreamSet = new Set(upstreamModelNames.value)
|
return allCustomNames.value
|
||||||
return allCustomNames.value.filter(name => !upstreamSet.has(name))
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// 排序后的自定义名称
|
// 排序后的自定义名称
|
||||||
@@ -369,31 +259,14 @@ const sortedCustomNames = computed(() => {
|
|||||||
const canAddAsCustom = computed(() => {
|
const canAddAsCustom = computed(() => {
|
||||||
const search = searchQuery.value.trim()
|
const search = searchQuery.value.trim()
|
||||||
if (!search) return false
|
if (!search) return false
|
||||||
// 已经选中了就不显示
|
|
||||||
if (selectedNames.value.includes(search)) return false
|
if (selectedNames.value.includes(search)) return false
|
||||||
// 已经在自定义列表中就不显示
|
|
||||||
if (allCustomNames.value.includes(search)) return false
|
if (allCustomNames.value.includes(search)) return false
|
||||||
// 精确匹配上游模型就不显示
|
|
||||||
if (upstreamModelNames.value.includes(search)) return false
|
|
||||||
return true
|
return true
|
||||||
})
|
})
|
||||||
|
|
||||||
// 过滤后的上游模型
|
|
||||||
const filteredUpstreamModels = computed(() => {
|
|
||||||
if (!searchQuery.value.trim()) return upstreamModelNames.value
|
|
||||||
const query = searchQuery.value.toLowerCase()
|
|
||||||
return upstreamModelNames.value.filter(name => name.toLowerCase().includes(query))
|
|
||||||
})
|
|
||||||
|
|
||||||
// 空状态判断
|
// 空状态判断
|
||||||
const showEmptyState = computed(() => {
|
const showEmptyState = computed(() => {
|
||||||
return filteredUpstreamModels.value.length === 0 && customNames.value.length === 0
|
return customNames.value.length === 0
|
||||||
})
|
|
||||||
|
|
||||||
// 上游模型是否全选
|
|
||||||
const isAllUpstreamModelsSelected = computed(() => {
|
|
||||||
if (filteredUpstreamModels.value.length === 0) return false
|
|
||||||
return filteredUpstreamModels.value.every(name => selectedNames.value.includes(name))
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// 切换名称选中状态
|
// 切换名称选中状态
|
||||||
@@ -411,73 +284,17 @@ function addCustomName() {
|
|||||||
const name = searchQuery.value.trim()
|
const name = searchQuery.value.trim()
|
||||||
if (name && !selectedNames.value.includes(name)) {
|
if (name && !selectedNames.value.includes(name)) {
|
||||||
selectedNames.value.push(name)
|
selectedNames.value.push(name)
|
||||||
if (!allKnownNames.value.has(name) && !allCustomNames.value.includes(name)) {
|
if (!allCustomNames.value.includes(name)) {
|
||||||
allCustomNames.value.push(name)
|
allCustomNames.value.push(name)
|
||||||
}
|
}
|
||||||
searchQuery.value = ''
|
searchQuery.value = ''
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 全选/取消全选上游模型
|
|
||||||
function toggleAllUpstreamModels() {
|
|
||||||
const allNames = filteredUpstreamModels.value
|
|
||||||
if (isAllUpstreamModelsSelected.value) {
|
|
||||||
selectedNames.value = selectedNames.value.filter(name => !allNames.includes(name))
|
|
||||||
} else {
|
|
||||||
allNames.forEach(name => {
|
|
||||||
if (!selectedNames.value.includes(name)) {
|
|
||||||
selectedNames.value.push(name)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 切换折叠状态
|
|
||||||
function toggleGroupCollapse(group: string) {
|
|
||||||
if (collapsedGroups.value.has(group)) {
|
|
||||||
collapsedGroups.value.delete(group)
|
|
||||||
} else {
|
|
||||||
collapsedGroups.value.add(group)
|
|
||||||
}
|
|
||||||
collapsedGroups.value = new Set(collapsedGroups.value)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 从提供商获取模型(使用缓存)
|
|
||||||
async function fetchUpstreamModels() {
|
|
||||||
if (!props.providerId) return
|
|
||||||
try {
|
|
||||||
loadingModels.value = true
|
|
||||||
fetchingUpstreamModels.value = true
|
|
||||||
const result = await fetchCachedModels(props.providerId)
|
|
||||||
if (result.models.length > 0) {
|
|
||||||
upstreamModels.value = result.models
|
|
||||||
upstreamModelsLoaded.value = true
|
|
||||||
// 获取上游模型后,将不在上游列表中的已选名称添加到自定义列表
|
|
||||||
const upstreamIds = new Set(result.models.map(m => m.id))
|
|
||||||
const customFromSelected = selectedNames.value.filter(name => !upstreamIds.has(name))
|
|
||||||
// 合并现有自定义名称和从已选中提取的自定义名称
|
|
||||||
const mergedCustom = new Set([...allCustomNames.value, ...customFromSelected])
|
|
||||||
allCustomNames.value = Array.from(mergedCustom).filter(name => !upstreamIds.has(name))
|
|
||||||
}
|
|
||||||
if (result.error) {
|
|
||||||
showError(result.error, '获取上游模型失败')
|
|
||||||
}
|
|
||||||
} catch (err: unknown) {
|
|
||||||
showError(parseApiError(err, '获取上游模型列表失败'), '错误')
|
|
||||||
} finally {
|
|
||||||
loadingModels.value = false
|
|
||||||
fetchingUpstreamModels.value = false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 监听打开状态
|
// 监听打开状态
|
||||||
watch(() => props.open, async (isOpen) => {
|
watch(() => props.open, async (isOpen) => {
|
||||||
if (isOpen) {
|
if (isOpen) {
|
||||||
initForm()
|
initForm()
|
||||||
// 只有在有 key 配置了自动获取时才自动加载上游模型
|
|
||||||
if (props.hasAutoFetchKey) {
|
|
||||||
await fetchUpstreamModels()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -489,7 +306,6 @@ function initForm() {
|
|||||||
}
|
}
|
||||||
const existingNames = props.editingGroup.aliases.map(a => a.name)
|
const existingNames = props.editingGroup.aliases.map(a => a.name)
|
||||||
selectedNames.value = [...existingNames]
|
selectedNames.value = [...existingNames]
|
||||||
// 将已有映射名称添加到自定义列表,使其在列表中可见(可取消选中来移除映射)
|
|
||||||
allCustomNames.value = [...existingNames]
|
allCustomNames.value = [...existingNames]
|
||||||
} else {
|
} else {
|
||||||
formData.value = {
|
formData.value = {
|
||||||
@@ -499,10 +315,6 @@ function initForm() {
|
|||||||
allCustomNames.value = []
|
allCustomNames.value = []
|
||||||
}
|
}
|
||||||
searchQuery.value = ''
|
searchQuery.value = ''
|
||||||
upstreamModels.value = []
|
|
||||||
upstreamModelsLoaded.value = false
|
|
||||||
// 默认展开所有分组
|
|
||||||
collapsedGroups.value = new Set()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 处理模型选择变更
|
// 处理模型选择变更
|
||||||
@@ -532,7 +344,6 @@ async function handleSubmit() {
|
|||||||
const currentAliases = targetModel.provider_model_mappings || []
|
const currentAliases = targetModel.provider_model_mappings || []
|
||||||
let newAliases: ProviderModelAlias[]
|
let newAliases: ProviderModelAlias[]
|
||||||
|
|
||||||
// 为每个选中的名称创建映射(所有映射使用相同优先级,实现同级负载均衡)
|
|
||||||
const buildAliases = (names: string[]): ProviderModelAlias[] => {
|
const buildAliases = (names: string[]): ProviderModelAlias[] => {
|
||||||
return names.map((name) => ({
|
return names.map((name) => ({
|
||||||
name: name.trim(),
|
name: name.trim(),
|
||||||
|
|||||||
@@ -159,7 +159,7 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="space-y-1.5">
|
<div class="space-y-1.5">
|
||||||
<label class="text-xs font-medium text-muted-foreground">TOTP Secret (可选)</label>
|
<label class="text-xs font-medium text-muted-foreground">TOTP Secret (可选, 2FA认证)</label>
|
||||||
<input
|
<input
|
||||||
v-model="device.totp_secret"
|
v-model="device.totp_secret"
|
||||||
type="text"
|
type="text"
|
||||||
|
|||||||
@@ -51,7 +51,7 @@
|
|||||||
</Button>
|
</Button>
|
||||||
</span>
|
</span>
|
||||||
<Popover
|
<Popover
|
||||||
v-if="provider.provider_type !== 'custom'"
|
v-if="provider.pool_advanced"
|
||||||
:open="providerProxyPopoverOpen"
|
:open="providerProxyPopoverOpen"
|
||||||
@update:open="handleProviderProxyPopoverToggle"
|
@update:open="handleProviderProxyPopoverToggle"
|
||||||
>
|
>
|
||||||
@@ -427,9 +427,9 @@
|
|||||||
>
|
>
|
||||||
<Shield class="w-3.5 h-3.5" />
|
<Shield class="w-3.5 h-3.5" />
|
||||||
</Button>
|
</Button>
|
||||||
<!-- 代理节点配置(仅非 custom 类型显示) -->
|
<!-- 代理节点配置(仅号池模式显示) -->
|
||||||
<Popover
|
<Popover
|
||||||
v-if="provider.provider_type !== 'custom'"
|
v-if="provider.pool_advanced"
|
||||||
:open="proxyPopoverOpenKeyId === key.id"
|
:open="proxyPopoverOpenKeyId === key.id"
|
||||||
@update:open="(v: boolean) => handleProxyPopoverToggle(key.id, v)"
|
@update:open="(v: boolean) => handleProxyPopoverToggle(key.id, v)"
|
||||||
>
|
>
|
||||||
|
|||||||
@@ -37,11 +37,14 @@
|
|||||||
<SelectValue placeholder="请选择" />
|
<SelectValue placeholder="请选择" />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
<SelectContent>
|
<SelectContent>
|
||||||
<!-- 新建模式:允许自定义、Codex、Kiro 和 Antigravity -->
|
<!-- 新建模式:允许自定义及各反代类型 -->
|
||||||
<template v-if="!isEditMode">
|
<template v-if="!isEditMode">
|
||||||
<SelectItem value="custom">
|
<SelectItem value="custom">
|
||||||
自定义
|
自定义
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
|
<SelectItem value="claude_code">
|
||||||
|
ClaudeCode
|
||||||
|
</SelectItem>
|
||||||
<SelectItem value="codex">
|
<SelectItem value="codex">
|
||||||
Codex
|
Codex
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
@@ -216,14 +219,15 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 格式转换配置 -->
|
<!-- 功能开关 -->
|
||||||
<div class="space-y-3">
|
<div class="space-y-3">
|
||||||
<h3 class="text-sm font-medium border-b pb-2">
|
<h3 class="text-sm font-medium border-b pb-2">
|
||||||
格式转换
|
功能开关
|
||||||
</h3>
|
</h3>
|
||||||
|
|
||||||
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
|
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
|
||||||
<div class="space-y-0.5">
|
<div class="space-y-0.5">
|
||||||
<span class="text-sm font-medium">保持优先级</span>
|
<span class="text-sm font-medium">格式转换保持优先级</span>
|
||||||
<p class="text-xs text-muted-foreground">
|
<p class="text-xs text-muted-foreground">
|
||||||
跨格式请求时保持原优先级排名,不降级到格式匹配的提供商之后
|
跨格式请求时保持原优先级排名,不降级到格式匹配的提供商之后
|
||||||
</p>
|
</p>
|
||||||
@@ -233,30 +237,17 @@
|
|||||||
@update:model-value="(v: boolean) => form.keep_priority_on_conversion = v"
|
@update:model-value="(v: boolean) => form.keep_priority_on_conversion = v"
|
||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
|
||||||
|
|
||||||
<!-- 代理配置 -->
|
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
|
||||||
<div class="space-y-3">
|
<div class="space-y-0.5">
|
||||||
<div class="flex items-center justify-between">
|
<span class="text-sm font-medium">号池调度模式</span>
|
||||||
<h3 class="text-sm font-medium">
|
<p class="text-xs text-muted-foreground">
|
||||||
代理配置
|
启用后该提供商的密钥将由号池统一调度
|
||||||
</h3>
|
</p>
|
||||||
<div class="flex items-center gap-2">
|
|
||||||
<Switch
|
|
||||||
:model-value="form.proxy_enabled"
|
|
||||||
@update:model-value="handleProxyToggle"
|
|
||||||
/>
|
|
||||||
<span class="text-sm text-muted-foreground">启用代理</span>
|
|
||||||
</div>
|
</div>
|
||||||
</div>
|
<Switch
|
||||||
<div
|
:model-value="form.pool_mode_enabled"
|
||||||
v-if="form.proxy_enabled"
|
@update:model-value="(v: boolean) => form.pool_mode_enabled = v"
|
||||||
class="space-y-1.5 p-3 border rounded-lg bg-muted/50"
|
|
||||||
>
|
|
||||||
<Label class="text-xs">代理节点 *</Label>
|
|
||||||
<ProxyNodeSelect
|
|
||||||
ref="proxyNodeSelectRef"
|
|
||||||
v-model="form.proxy_node_id"
|
|
||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -282,7 +273,7 @@
|
|||||||
</template>
|
</template>
|
||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, computed } from 'vue'
|
import { ref, computed, watch } from 'vue'
|
||||||
import {
|
import {
|
||||||
Dialog,
|
Dialog,
|
||||||
Button,
|
Button,
|
||||||
@@ -301,8 +292,6 @@ import { useFormDialog } from '@/composables/useFormDialog'
|
|||||||
import { createProvider, updateProvider, type ProviderWithEndpointsSummary } from '@/api/endpoints'
|
import { createProvider, updateProvider, type ProviderWithEndpointsSummary } from '@/api/endpoints'
|
||||||
import { parseApiError } from '@/utils/errorParser'
|
import { parseApiError } from '@/utils/errorParser'
|
||||||
import { parseNumberInput } from '@/utils/form'
|
import { parseNumberInput } from '@/utils/form'
|
||||||
import ProxyNodeSelect from './ProxyNodeSelect.vue'
|
|
||||||
import { useProxyNodesStore } from '@/stores/proxy-nodes'
|
|
||||||
|
|
||||||
const props = defineProps<{
|
const props = defineProps<{
|
||||||
modelValue: boolean
|
modelValue: boolean
|
||||||
@@ -318,16 +307,6 @@ const emit = defineEmits<{
|
|||||||
|
|
||||||
const { success, error: showError } = useToast()
|
const { success, error: showError } = useToast()
|
||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
const proxyNodeSelectRef = ref<InstanceType<typeof ProxyNodeSelect> | null>(null)
|
|
||||||
const proxyNodesStore = useProxyNodesStore()
|
|
||||||
|
|
||||||
/** 启用代理时懒加载节点列表 */
|
|
||||||
function handleProxyToggle(v: boolean) {
|
|
||||||
form.value.proxy_enabled = v
|
|
||||||
if (v) {
|
|
||||||
proxyNodesStore.ensureLoaded()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 内部状态
|
// 内部状态
|
||||||
const internalOpen = computed(() => props.modelValue)
|
const internalOpen = computed(() => props.modelValue)
|
||||||
@@ -363,9 +342,8 @@ const form = ref({
|
|||||||
// 超时配置(秒)
|
// 超时配置(秒)
|
||||||
stream_first_byte_timeout: undefined as number | undefined,
|
stream_first_byte_timeout: undefined as number | undefined,
|
||||||
request_timeout: undefined as number | undefined,
|
request_timeout: undefined as number | undefined,
|
||||||
// 代理配置
|
// 号池模式
|
||||||
proxy_enabled: false,
|
pool_mode_enabled: false,
|
||||||
proxy_node_id: '',
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// 重置表单
|
// 重置表单
|
||||||
@@ -390,9 +368,8 @@ function resetForm() {
|
|||||||
// 超时配置
|
// 超时配置
|
||||||
stream_first_byte_timeout: undefined,
|
stream_first_byte_timeout: undefined,
|
||||||
request_timeout: undefined,
|
request_timeout: undefined,
|
||||||
// 代理配置
|
// 号池模式
|
||||||
proxy_enabled: false,
|
pool_mode_enabled: false,
|
||||||
proxy_node_id: '',
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -400,7 +377,6 @@ function resetForm() {
|
|||||||
function loadProviderData() {
|
function loadProviderData() {
|
||||||
if (!props.provider) return
|
if (!props.provider) return
|
||||||
|
|
||||||
const proxy = props.provider.proxy
|
|
||||||
form.value = {
|
form.value = {
|
||||||
name: props.provider.name,
|
name: props.provider.name,
|
||||||
provider_type: props.provider.provider_type || 'custom',
|
provider_type: props.provider.provider_type || 'custom',
|
||||||
@@ -423,14 +399,8 @@ function loadProviderData() {
|
|||||||
// 超时配置
|
// 超时配置
|
||||||
stream_first_byte_timeout: props.provider.stream_first_byte_timeout ?? undefined,
|
stream_first_byte_timeout: props.provider.stream_first_byte_timeout ?? undefined,
|
||||||
request_timeout: props.provider.request_timeout ?? undefined,
|
request_timeout: props.provider.request_timeout ?? undefined,
|
||||||
// 代理配置
|
// 号池模式
|
||||||
proxy_enabled: proxy?.enabled ?? false,
|
pool_mode_enabled: !!props.provider.pool_advanced,
|
||||||
proxy_node_id: proxy?.node_id || '',
|
|
||||||
}
|
|
||||||
|
|
||||||
// 如果有代理配置,确保加载节点列表(直接调用 store,避免 ref 未挂载时静默失败)
|
|
||||||
if (proxy?.enabled) {
|
|
||||||
proxyNodesStore.ensureLoaded()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -444,6 +414,13 @@ const { isEditMode, handleDialogUpdate, handleCancel } = useFormDialog({
|
|||||||
resetForm,
|
resetForm,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// 新建模式下切换 provider_type 时自动设置号池模式:非自定义类型默认开启
|
||||||
|
watch(() => form.value.provider_type, (newType) => {
|
||||||
|
if (!isEditMode.value) {
|
||||||
|
form.value.pool_mode_enabled = newType !== 'custom'
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
// 提交表单
|
// 提交表单
|
||||||
const handleSubmit = async () => {
|
const handleSubmit = async () => {
|
||||||
// 月卡类型必须设置周期开始时间
|
// 月卡类型必须设置周期开始时间
|
||||||
@@ -452,20 +429,8 @@ const handleSubmit = async () => {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 启用代理时的验证
|
|
||||||
if (form.value.proxy_enabled && !form.value.proxy_node_id) {
|
|
||||||
showError('请选择代理节点', '验证失败')
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
loading.value = true
|
loading.value = true
|
||||||
try {
|
try {
|
||||||
// 构建代理配置
|
|
||||||
const proxy = form.value.proxy_enabled && form.value.proxy_node_id ? {
|
|
||||||
node_id: form.value.proxy_node_id,
|
|
||||||
enabled: true,
|
|
||||||
} : null
|
|
||||||
|
|
||||||
const payload = {
|
const payload = {
|
||||||
name: form.value.name,
|
name: form.value.name,
|
||||||
provider_type: form.value.provider_type,
|
provider_type: form.value.provider_type,
|
||||||
@@ -484,7 +449,9 @@ const handleSubmit = async () => {
|
|||||||
// 超时配置(null 表示清除,使用全局配置)
|
// 超时配置(null 表示清除,使用全局配置)
|
||||||
stream_first_byte_timeout: form.value.stream_first_byte_timeout ?? null,
|
stream_first_byte_timeout: form.value.stream_first_byte_timeout ?? null,
|
||||||
request_timeout: form.value.request_timeout ?? null,
|
request_timeout: form.value.request_timeout ?? null,
|
||||||
proxy,
|
pool_advanced: form.value.pool_mode_enabled
|
||||||
|
? (props.provider?.pool_advanced ?? {})
|
||||||
|
: null,
|
||||||
}
|
}
|
||||||
|
|
||||||
if (isEditMode.value && props.provider) {
|
if (isEditMode.value && props.provider) {
|
||||||
|
|||||||
@@ -13,3 +13,4 @@ export { default as OAuthKeyEditDialog } from './OAuthKeyEditDialog.vue'
|
|||||||
|
|
||||||
export { default as ModelsTab } from './provider-tabs/ModelsTab.vue'
|
export { default as ModelsTab } from './provider-tabs/ModelsTab.vue'
|
||||||
export { default as ProviderAuthDialog } from './ProviderAuthDialog.vue'
|
export { default as ProviderAuthDialog } from './ProviderAuthDialog.vue'
|
||||||
|
export { default as PoolStatusCard } from './PoolStatusCard.vue'
|
||||||
|
|||||||
@@ -325,7 +325,6 @@
|
|||||||
:models="models"
|
:models="models"
|
||||||
:editing-group="editingGroup"
|
:editing-group="editingGroup"
|
||||||
:preselected-model-id="preselectedModelId"
|
:preselected-model-id="preselectedModelId"
|
||||||
:has-auto-fetch-key="hasAutoFetchKey"
|
|
||||||
@saved="onDialogSaved"
|
@saved="onDialogSaved"
|
||||||
/>
|
/>
|
||||||
|
|
||||||
@@ -424,11 +423,6 @@ const providerKeysState = computed(() => props.providerKeys ?? [])
|
|||||||
// 展开状态
|
// 展开状态
|
||||||
const expandedItems = ref<Set<string>>(new Set())
|
const expandedItems = ref<Set<string>>(new Set())
|
||||||
|
|
||||||
// 是否有 key 配置了自动获取上游模型
|
|
||||||
const hasAutoFetchKey = computed(() => {
|
|
||||||
return providerKeysState.value.some(k => k.auto_fetch_models)
|
|
||||||
})
|
|
||||||
|
|
||||||
// 生成作用域唯一键
|
// 生成作用域唯一键
|
||||||
function getApiFormatsKey(formats: string[] | undefined): string {
|
function getApiFormatsKey(formats: string[] | undefined): string {
|
||||||
if (!formats || formats.length === 0) return ''
|
if (!formats || formats.length === 0) return ''
|
||||||
|
|||||||
@@ -252,6 +252,62 @@
|
|||||||
</span>
|
</span>
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
|
<div
|
||||||
|
v-if="currentAttempt.extra_data?.pool_selection"
|
||||||
|
class="info-item"
|
||||||
|
>
|
||||||
|
<span class="info-label">号池调度</span>
|
||||||
|
<span class="info-value info-value-stacked">
|
||||||
|
<span class="pool-reason">
|
||||||
|
<span
|
||||||
|
class="pool-reason-tag"
|
||||||
|
:class="'pool-' + currentAttempt.extra_data.pool_selection.reason"
|
||||||
|
>
|
||||||
|
{{ poolSelectionLabel(currentAttempt.extra_data.pool_selection.reason) }}
|
||||||
|
</span>
|
||||||
|
<span
|
||||||
|
v-if="currentAttempt.extra_data.pool_selection.cost_soft_threshold"
|
||||||
|
class="pool-cost-warn"
|
||||||
|
>接近限额</span>
|
||||||
|
</span>
|
||||||
|
<span
|
||||||
|
v-if="currentAttempt.extra_data.pool_selection.cost_window_usage"
|
||||||
|
class="text-xs text-muted-foreground"
|
||||||
|
>
|
||||||
|
{{ formatNumber(currentAttempt.extra_data.pool_selection.cost_window_usage) }}
|
||||||
|
<template v-if="currentAttempt.extra_data.pool_selection.cost_limit">
|
||||||
|
/ {{ formatNumber(currentAttempt.extra_data.pool_selection.cost_limit) }}
|
||||||
|
</template>
|
||||||
|
tokens
|
||||||
|
</span>
|
||||||
|
</span>
|
||||||
|
</div>
|
||||||
|
<div
|
||||||
|
v-if="currentAttempt.extra_data?.pool_skip"
|
||||||
|
class="info-item"
|
||||||
|
>
|
||||||
|
<span class="info-label">号池跳过</span>
|
||||||
|
<span class="info-value info-value-stacked">
|
||||||
|
<span class="pool-skip-type">
|
||||||
|
{{ poolSkipLabel(currentAttempt.extra_data.pool_skip.type) }}
|
||||||
|
</span>
|
||||||
|
<span
|
||||||
|
v-if="currentAttempt.extra_data.pool_skip.cooldown_reason"
|
||||||
|
class="text-xs text-muted-foreground"
|
||||||
|
>
|
||||||
|
{{ currentAttempt.extra_data.pool_skip.cooldown_reason }}
|
||||||
|
<template v-if="currentAttempt.extra_data.pool_skip.cooldown_ttl != null">
|
||||||
|
({{ currentAttempt.extra_data.pool_skip.cooldown_ttl }}s)
|
||||||
|
</template>
|
||||||
|
</span>
|
||||||
|
<span
|
||||||
|
v-if="currentAttempt.extra_data.pool_skip.cost_window_usage"
|
||||||
|
class="text-xs text-muted-foreground"
|
||||||
|
>
|
||||||
|
{{ formatNumber(currentAttempt.extra_data.pool_skip.cost_window_usage) }} tokens
|
||||||
|
</span>
|
||||||
|
</span>
|
||||||
|
</div>
|
||||||
<div
|
<div
|
||||||
v-if="mergedCapabilities.length > 0"
|
v-if="mergedCapabilities.length > 0"
|
||||||
class="info-item"
|
class="info-item"
|
||||||
@@ -764,6 +820,25 @@ const formatCapabilityLabel = (cap: string): string => {
|
|||||||
return labels[cap] || cap
|
return labels[cap] || cap
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const poolSelectionLabel = (reason: string): string => {
|
||||||
|
const labels: Record<string, string> = {
|
||||||
|
sticky: '粘性会话',
|
||||||
|
lru: 'LRU',
|
||||||
|
random: '随机',
|
||||||
|
tiebreak: '随机 (平分)',
|
||||||
|
}
|
||||||
|
return labels[reason] || reason
|
||||||
|
}
|
||||||
|
|
||||||
|
const poolSkipLabel = (type: string): string => {
|
||||||
|
const labels: Record<string, string> = {
|
||||||
|
cooldown: '冷却中',
|
||||||
|
cost_exhausted: '额度耗尽',
|
||||||
|
upstream: '上游跳过',
|
||||||
|
}
|
||||||
|
return labels[type] || type
|
||||||
|
}
|
||||||
|
|
||||||
// 检查组是否被悬浮
|
// 检查组是否被悬浮
|
||||||
const isGroupHovered = (groupIndex: number) => {
|
const isGroupHovered = (groupIndex: number) => {
|
||||||
return hoveredGroupIndex.value === groupIndex
|
return hoveredGroupIndex.value === groupIndex
|
||||||
@@ -1485,6 +1560,52 @@ const getStatusColorClass = (status: string) => {
|
|||||||
gap: 0.375rem;
|
gap: 0.375rem;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/* 号池调度 */
|
||||||
|
.pool-reason {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 0.375rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
.pool-reason-tag {
|
||||||
|
display: inline-flex;
|
||||||
|
align-items: center;
|
||||||
|
padding: 0.15rem 0.5rem;
|
||||||
|
font-size: 0.7rem;
|
||||||
|
font-weight: 500;
|
||||||
|
border-radius: 4px;
|
||||||
|
white-space: nowrap;
|
||||||
|
border: 1px solid hsl(var(--border));
|
||||||
|
}
|
||||||
|
|
||||||
|
.pool-reason-tag.pool-sticky {
|
||||||
|
color: hsl(var(--chart-4));
|
||||||
|
border-color: hsl(var(--chart-4) / 0.3);
|
||||||
|
background: hsl(var(--chart-4) / 0.08);
|
||||||
|
}
|
||||||
|
|
||||||
|
.pool-reason-tag.pool-lru {
|
||||||
|
color: hsl(var(--chart-2));
|
||||||
|
border-color: hsl(var(--chart-2) / 0.3);
|
||||||
|
background: hsl(var(--chart-2) / 0.08);
|
||||||
|
}
|
||||||
|
|
||||||
|
.pool-reason-tag.pool-random,
|
||||||
|
.pool-reason-tag.pool-tiebreak {
|
||||||
|
color: hsl(var(--muted-foreground));
|
||||||
|
}
|
||||||
|
|
||||||
|
.pool-cost-warn {
|
||||||
|
font-size: 0.65rem;
|
||||||
|
color: hsl(var(--chart-5));
|
||||||
|
font-weight: 500;
|
||||||
|
}
|
||||||
|
|
||||||
|
.pool-skip-type {
|
||||||
|
font-weight: 500;
|
||||||
|
color: hsl(var(--muted-foreground));
|
||||||
|
}
|
||||||
|
|
||||||
/* 能力标签 */
|
/* 能力标签 */
|
||||||
.capability-tags {
|
.capability-tags {
|
||||||
display: flex;
|
display: flex;
|
||||||
|
|||||||
@@ -370,6 +370,51 @@
|
|||||||
</div>
|
</div>
|
||||||
</Card>
|
</Card>
|
||||||
|
|
||||||
|
<!-- 号池调度摘要 -->
|
||||||
|
<Card v-if="poolSummary">
|
||||||
|
<div class="p-3 sm:p-4">
|
||||||
|
<div class="text-xs text-muted-foreground mb-2 font-medium">
|
||||||
|
号池调度
|
||||||
|
</div>
|
||||||
|
<div class="flex items-center gap-3 flex-wrap text-sm">
|
||||||
|
<span class="font-mono">
|
||||||
|
{{ poolSummary.total_keys }} 候选
|
||||||
|
</span>
|
||||||
|
<span class="text-muted-foreground">|</span>
|
||||||
|
<span class="font-mono">
|
||||||
|
{{ poolSummary.attempted }} 尝试
|
||||||
|
</span>
|
||||||
|
<template v-if="poolSummary.skipped_cooldown > 0">
|
||||||
|
<span class="text-muted-foreground">|</span>
|
||||||
|
<span class="font-mono text-amber-600 dark:text-amber-400">
|
||||||
|
{{ poolSummary.skipped_cooldown }} 冷却跳过
|
||||||
|
</span>
|
||||||
|
</template>
|
||||||
|
<template v-if="poolSummary.skipped_cost > 0">
|
||||||
|
<span class="text-muted-foreground">|</span>
|
||||||
|
<span class="font-mono text-orange-600 dark:text-orange-400">
|
||||||
|
{{ poolSummary.skipped_cost }} 成本跳过
|
||||||
|
</span>
|
||||||
|
</template>
|
||||||
|
<template v-if="poolSummary.sticky_session">
|
||||||
|
<span class="text-muted-foreground">|</span>
|
||||||
|
<Badge
|
||||||
|
variant="outline"
|
||||||
|
class="text-[10px] px-1.5 py-0 h-4"
|
||||||
|
>
|
||||||
|
粘性会话
|
||||||
|
</Badge>
|
||||||
|
</template>
|
||||||
|
<template v-if="poolSummary.success_reason">
|
||||||
|
<span class="text-muted-foreground">|</span>
|
||||||
|
<span class="text-xs text-muted-foreground">
|
||||||
|
{{ poolSummary.success_reason }}
|
||||||
|
</span>
|
||||||
|
</template>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</Card>
|
||||||
|
|
||||||
<!-- 请求链路追踪卡片 -->
|
<!-- 请求链路追踪卡片 -->
|
||||||
<div v-if="detail.request_id || detail.id">
|
<div v-if="detail.request_id || detail.id">
|
||||||
<HorizontalRequestTimeline
|
<HorizontalRequestTimeline
|
||||||
@@ -733,6 +778,13 @@ const isDark = computed(() => {
|
|||||||
return document.documentElement.classList.contains('dark')
|
return document.documentElement.classList.contains('dark')
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// 号池调度摘要
|
||||||
|
const poolSummary = computed(() => {
|
||||||
|
const ps = detail.value?.metadata?.pool_summary as Record<string, unknown> | undefined
|
||||||
|
if (!ps || !ps.enabled) return null
|
||||||
|
return ps
|
||||||
|
})
|
||||||
|
|
||||||
// 检测是否有提供商请求头
|
// 检测是否有提供商请求头
|
||||||
const hasProviderHeaders = computed(() => {
|
const hasProviderHeaders = computed(() => {
|
||||||
return !!(detail.value?.provider_request_headers &&
|
return !!(detail.value?.provider_request_headers &&
|
||||||
|
|||||||
@@ -357,6 +357,7 @@ import {
|
|||||||
Gauge,
|
Gauge,
|
||||||
Layers,
|
Layers,
|
||||||
FolderTree,
|
FolderTree,
|
||||||
|
Database,
|
||||||
Box,
|
Box,
|
||||||
LogOut,
|
LogOut,
|
||||||
SunMoon,
|
SunMoon,
|
||||||
@@ -554,6 +555,7 @@ const navigation = computed(() => {
|
|||||||
items: [
|
items: [
|
||||||
{ name: '用户管理', href: '/admin/users', icon: Users },
|
{ name: '用户管理', href: '/admin/users', icon: Users },
|
||||||
{ name: '提供商', href: '/admin/providers', icon: FolderTree },
|
{ name: '提供商', href: '/admin/providers', icon: FolderTree },
|
||||||
|
{ name: '号池管理', href: '/admin/pool', icon: Database },
|
||||||
{ name: '模型管理', href: '/admin/models', icon: Layers },
|
{ name: '模型管理', href: '/admin/models', icon: Layers },
|
||||||
{ name: '独立密钥', href: '/admin/keys', icon: Key },
|
{ name: '独立密钥', href: '/admin/keys', icon: Key },
|
||||||
{ name: '异步任务', href: '/admin/async-tasks', icon: Zap },
|
{ name: '异步任务', href: '/admin/async-tasks', icon: Zap },
|
||||||
|
|||||||
@@ -160,6 +160,11 @@ const routes: RouteRecordRaw[] = [
|
|||||||
name: 'ProviderManagement',
|
name: 'ProviderManagement',
|
||||||
component: () => importWithRetry(() => import('@/views/admin/ProviderManagement.vue'))
|
component: () => importWithRetry(() => import('@/views/admin/ProviderManagement.vue'))
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
path: 'pool',
|
||||||
|
name: 'PoolManagement',
|
||||||
|
component: () => importWithRetry(() => import('@/views/admin/PoolManagement.vue'))
|
||||||
|
},
|
||||||
{
|
{
|
||||||
path: 'models',
|
path: 'models',
|
||||||
name: 'ModelManagement',
|
name: 'ModelManagement',
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from .endpoints import router as endpoints_router
|
|||||||
from .models import router as models_router
|
from .models import router as models_router
|
||||||
from .modules import router as modules_router
|
from .modules import router as modules_router
|
||||||
from .monitoring import router as monitoring_router
|
from .monitoring import router as monitoring_router
|
||||||
|
from .pool import router as pool_router
|
||||||
from .provider_oauth import router as provider_oauth_router
|
from .provider_oauth import router as provider_oauth_router
|
||||||
from .provider_ops import router as provider_ops_router
|
from .provider_ops import router as provider_ops_router
|
||||||
from .provider_query import router as provider_query_router
|
from .provider_query import router as provider_query_router
|
||||||
@@ -38,6 +39,7 @@ router.include_router(security_router)
|
|||||||
router.include_router(stats_router)
|
router.include_router(stats_router)
|
||||||
router.include_router(provider_query_router)
|
router.include_router(provider_query_router)
|
||||||
router.include_router(modules_router)
|
router.include_router(modules_router)
|
||||||
|
router.include_router(pool_router)
|
||||||
router.include_router(provider_ops_router)
|
router.include_router(provider_ops_router)
|
||||||
router.include_router(video_tasks_router)
|
router.include_router(video_tasks_router)
|
||||||
|
|
||||||
|
|||||||
@@ -357,6 +357,43 @@ def _build_kiro_key_name(
|
|||||||
return f"{base} ({method})"
|
return f"{base} ({method})"
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_codex_plan_group(plan_type: Any) -> str | None:
|
||||||
|
"""将 Codex plan_type 归一化到判重分组。
|
||||||
|
|
||||||
|
分组规则:
|
||||||
|
- free
|
||||||
|
- team/plus/enterprise(同组)
|
||||||
|
"""
|
||||||
|
if not isinstance(plan_type, str):
|
||||||
|
return None
|
||||||
|
normalized = plan_type.strip().lower()
|
||||||
|
if not normalized:
|
||||||
|
return None
|
||||||
|
if normalized == "free":
|
||||||
|
return "free"
|
||||||
|
if normalized in {"team", "plus", "enterprise"}:
|
||||||
|
return "team_plus_enterprise"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _is_codex_cross_plan_group_non_duplicate(
|
||||||
|
*,
|
||||||
|
new_provider_type: Any,
|
||||||
|
existing_provider_type: Any,
|
||||||
|
new_plan_type: Any,
|
||||||
|
existing_plan_type: Any,
|
||||||
|
) -> bool:
|
||||||
|
"""Codex 账号在 free 与 Team/Plus/Enterprise 之间不判重。"""
|
||||||
|
new_pt = str(new_provider_type or "").strip().lower()
|
||||||
|
existing_pt = str(existing_provider_type or "").strip().lower()
|
||||||
|
if new_pt != ProviderType.CODEX.value and existing_pt != ProviderType.CODEX.value:
|
||||||
|
return False
|
||||||
|
|
||||||
|
new_group = _normalize_codex_plan_group(new_plan_type)
|
||||||
|
existing_group = _normalize_codex_plan_group(existing_plan_type)
|
||||||
|
return bool(new_group and existing_group and new_group != existing_group)
|
||||||
|
|
||||||
|
|
||||||
def _check_duplicate_oauth_account(
|
def _check_duplicate_oauth_account(
|
||||||
db: Session,
|
db: Session,
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
@@ -368,6 +405,7 @@ def _check_duplicate_oauth_account(
|
|||||||
|
|
||||||
通过以下字段判断重复:
|
通过以下字段判断重复:
|
||||||
- user_id: Codex 等使用用户级别 ID(同 team 下不同成员共享 account_id 但 user_id 不同)
|
- user_id: Codex 等使用用户级别 ID(同 team 下不同成员共享 account_id 但 user_id 不同)
|
||||||
|
对 Codex 额外按账号类型分组:free 与 Team/Plus/Enterprise 互不判重
|
||||||
- email + auth_method: Kiro 使用 email + auth_method 组合判断
|
- email + auth_method: Kiro 使用 email + auth_method 组合判断
|
||||||
(同一邮箱可能通过 Social 和 IdC 两种方式登录,视为不同账号)
|
(同一邮箱可能通过 Social 和 IdC 两种方式登录,视为不同账号)
|
||||||
- email: 其他 OAuth Provider 使用邮箱判断
|
- email: 其他 OAuth Provider 使用邮箱判断
|
||||||
@@ -383,6 +421,7 @@ def _check_duplicate_oauth_account(
|
|||||||
new_user_id = auth_config.get("user_id")
|
new_user_id = auth_config.get("user_id")
|
||||||
new_auth_method = auth_config.get("auth_method") # Kiro: social / idc
|
new_auth_method = auth_config.get("auth_method") # Kiro: social / idc
|
||||||
new_provider_type = auth_config.get("provider_type")
|
new_provider_type = auth_config.get("provider_type")
|
||||||
|
new_plan_type = auth_config.get("plan_type")
|
||||||
|
|
||||||
# 如果没有可用于识别的字段,跳过检查
|
# 如果没有可用于识别的字段,跳过检查
|
||||||
if not new_email and not new_user_id:
|
if not new_email and not new_user_id:
|
||||||
@@ -409,12 +448,19 @@ def _check_duplicate_oauth_account(
|
|||||||
existing_user_id = decrypted_config.get("user_id")
|
existing_user_id = decrypted_config.get("user_id")
|
||||||
existing_auth_method = decrypted_config.get("auth_method")
|
existing_auth_method = decrypted_config.get("auth_method")
|
||||||
existing_provider_type = decrypted_config.get("provider_type")
|
existing_provider_type = decrypted_config.get("provider_type")
|
||||||
|
existing_plan_type = decrypted_config.get("plan_type")
|
||||||
|
|
||||||
is_duplicate = False
|
is_duplicate = False
|
||||||
|
|
||||||
# user_id 相同即重复(Codex 等,同一 team 下不同成员共享 account_id 但 user_id 不同)
|
# user_id 相同即重复(Codex 等,同一 team 下不同成员共享 account_id 但 user_id 不同)
|
||||||
if new_user_id and existing_user_id and new_user_id == existing_user_id:
|
if new_user_id and existing_user_id and new_user_id == existing_user_id:
|
||||||
is_duplicate = True
|
if not _is_codex_cross_plan_group_non_duplicate(
|
||||||
|
new_provider_type=new_provider_type,
|
||||||
|
existing_provider_type=existing_provider_type,
|
||||||
|
new_plan_type=new_plan_type,
|
||||||
|
existing_plan_type=existing_plan_type,
|
||||||
|
):
|
||||||
|
is_duplicate = True
|
||||||
|
|
||||||
# email 判断
|
# email 判断
|
||||||
if not is_duplicate and new_email and existing_email and new_email == existing_email:
|
if not is_duplicate and new_email and existing_email and new_email == existing_email:
|
||||||
@@ -428,7 +474,13 @@ def _check_duplicate_oauth_account(
|
|||||||
):
|
):
|
||||||
is_duplicate = True
|
is_duplicate = True
|
||||||
else:
|
else:
|
||||||
is_duplicate = True
|
if not _is_codex_cross_plan_group_non_duplicate(
|
||||||
|
new_provider_type=new_provider_type,
|
||||||
|
existing_provider_type=existing_provider_type,
|
||||||
|
new_plan_type=new_plan_type,
|
||||||
|
existing_plan_type=existing_plan_type,
|
||||||
|
):
|
||||||
|
is_duplicate = True
|
||||||
|
|
||||||
if is_duplicate:
|
if is_duplicate:
|
||||||
# 失效账号允许覆盖
|
# 失效账号允许覆盖
|
||||||
|
|||||||
@@ -37,6 +37,84 @@ MAPPING_PREVIEW_MAX_MODELS = 500
|
|||||||
MAPPING_PREVIEW_TIMEOUT_SECONDS = 10.0
|
MAPPING_PREVIEW_TIMEOUT_SECONDS = 10.0
|
||||||
|
|
||||||
|
|
||||||
|
def _should_enable_format_conversion_by_default(provider_type: str | None) -> bool:
|
||||||
|
"""固定类型 Provider 默认是否开启格式转换。"""
|
||||||
|
pt = (provider_type or "custom").strip().lower()
|
||||||
|
envelope_provider_types = {
|
||||||
|
ProviderType.ANTIGRAVITY.value,
|
||||||
|
ProviderType.CLAUDE_CODE.value,
|
||||||
|
ProviderType.CODEX.value,
|
||||||
|
ProviderType.KIRO.value,
|
||||||
|
}
|
||||||
|
return pt in envelope_provider_types
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_provider_type(provider_type: str | None) -> str:
|
||||||
|
return (provider_type or "custom").strip().lower()
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_pool_advanced_config(
|
||||||
|
*,
|
||||||
|
provider_config: dict[str, Any] | None,
|
||||||
|
pool_advanced: dict[str, Any] | None,
|
||||||
|
pool_advanced_in_payload: bool,
|
||||||
|
) -> tuple[dict[str, Any] | None, bool]:
|
||||||
|
"""合并 pool_advanced 到 provider.config(任何 provider_type 均可使用)。"""
|
||||||
|
merged_config = dict(provider_config or {})
|
||||||
|
config_changed = False
|
||||||
|
|
||||||
|
if not pool_advanced_in_payload:
|
||||||
|
return merged_config or None, config_changed
|
||||||
|
|
||||||
|
if pool_advanced is None:
|
||||||
|
if "pool_advanced" in merged_config:
|
||||||
|
merged_config.pop("pool_advanced", None)
|
||||||
|
config_changed = True
|
||||||
|
else:
|
||||||
|
next_value = dict(pool_advanced)
|
||||||
|
if merged_config.get("pool_advanced") != next_value:
|
||||||
|
merged_config["pool_advanced"] = next_value
|
||||||
|
config_changed = True
|
||||||
|
|
||||||
|
return merged_config or None, config_changed
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_claude_code_advanced_config(
|
||||||
|
*,
|
||||||
|
provider_type: str | None,
|
||||||
|
provider_config: dict[str, Any] | None,
|
||||||
|
claude_code_advanced: dict[str, Any] | None,
|
||||||
|
claude_advanced_in_payload: bool,
|
||||||
|
) -> tuple[dict[str, Any] | None, bool]:
|
||||||
|
"""合并并规范 claude_code_advanced,确保仅在 claude_code 下保留。"""
|
||||||
|
normalized_provider_type = _normalize_provider_type(provider_type)
|
||||||
|
merged_config = dict(provider_config or {})
|
||||||
|
config_changed = False
|
||||||
|
|
||||||
|
if normalized_provider_type != ProviderType.CLAUDE_CODE.value:
|
||||||
|
if claude_advanced_in_payload and claude_code_advanced is not None:
|
||||||
|
raise InvalidRequestException("claude_code_advanced 仅适用于 provider_type=claude_code")
|
||||||
|
if "claude_code_advanced" in merged_config:
|
||||||
|
merged_config.pop("claude_code_advanced", None)
|
||||||
|
config_changed = True
|
||||||
|
return merged_config or None, config_changed
|
||||||
|
|
||||||
|
if not claude_advanced_in_payload:
|
||||||
|
return merged_config or None, config_changed
|
||||||
|
|
||||||
|
if claude_code_advanced is None:
|
||||||
|
if "claude_code_advanced" in merged_config:
|
||||||
|
merged_config.pop("claude_code_advanced", None)
|
||||||
|
config_changed = True
|
||||||
|
else:
|
||||||
|
next_value = dict(claude_code_advanced)
|
||||||
|
if merged_config.get("claude_code_advanced") != next_value:
|
||||||
|
merged_config["claude_code_advanced"] = next_value
|
||||||
|
config_changed = True
|
||||||
|
|
||||||
|
return merged_config or None, config_changed
|
||||||
|
|
||||||
|
|
||||||
# ========== Response Models ==========
|
# ========== Response Models ==========
|
||||||
|
|
||||||
|
|
||||||
@@ -161,7 +239,7 @@ async def create_provider(request: Request, db: Session = Depends(get_db)) -> An
|
|||||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||||
|
|
||||||
|
|
||||||
@router.put("/{provider_id}")
|
@router.patch("/{provider_id}")
|
||||||
async def update_provider(
|
async def update_provider(
|
||||||
provider_id: str, request: Request, db: Session = Depends(get_db)
|
provider_id: str, request: Request, db: Session = Depends(get_db)
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -290,15 +368,29 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
|
|||||||
else ProviderBillingType.PAY_AS_YOU_GO
|
else ProviderBillingType.PAY_AS_YOU_GO
|
||||||
)
|
)
|
||||||
|
|
||||||
# 有 envelope 包装的 Provider 类型(如 Antigravity、Codex)需要格式转换来正确
|
# 有 envelope 包装的 Provider 类型(如 ClaudeCode、Antigravity、Codex)需要
|
||||||
# 解包上游响应,创建时默认开启 enable_format_conversion。
|
# 格式转换来正确解包上游响应,创建时默认开启 enable_format_conversion。
|
||||||
pt = (validated_data.provider_type or "custom").strip()
|
pt = _normalize_provider_type(validated_data.provider_type)
|
||||||
envelope_provider_types = {
|
default_enable_format_conversion = _should_enable_format_conversion_by_default(pt)
|
||||||
ProviderType.ANTIGRAVITY,
|
provider_config, _ = _merge_claude_code_advanced_config(
|
||||||
ProviderType.CODEX,
|
provider_type=pt,
|
||||||
ProviderType.KIRO,
|
provider_config=validated_data.config,
|
||||||
}
|
claude_code_advanced=(
|
||||||
default_enable_format_conversion = pt in envelope_provider_types
|
validated_data.claude_code_advanced.model_dump(exclude_none=True)
|
||||||
|
if validated_data.claude_code_advanced is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
claude_advanced_in_payload=validated_data.claude_code_advanced is not None,
|
||||||
|
)
|
||||||
|
provider_config, _pool_changed = _merge_pool_advanced_config(
|
||||||
|
provider_config=provider_config,
|
||||||
|
pool_advanced=(
|
||||||
|
validated_data.pool_advanced.model_dump(exclude_none=True)
|
||||||
|
if validated_data.pool_advanced is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
pool_advanced_in_payload=validated_data.pool_advanced is not None,
|
||||||
|
)
|
||||||
|
|
||||||
# 创建 Provider 对象
|
# 创建 Provider 对象
|
||||||
provider = Provider(
|
provider = Provider(
|
||||||
@@ -319,7 +411,7 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
|
|||||||
# 超时配置
|
# 超时配置
|
||||||
stream_first_byte_timeout=validated_data.stream_first_byte_timeout,
|
stream_first_byte_timeout=validated_data.stream_first_byte_timeout,
|
||||||
request_timeout=validated_data.request_timeout,
|
request_timeout=validated_data.request_timeout,
|
||||||
config=validated_data.config,
|
config=provider_config or None,
|
||||||
# 有 envelope 的反代类型默认开启格式转换
|
# 有 envelope 的反代类型默认开启格式转换
|
||||||
enable_format_conversion=default_enable_format_conversion,
|
enable_format_conversion=default_enable_format_conversion,
|
||||||
)
|
)
|
||||||
@@ -411,6 +503,45 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
|
|||||||
try:
|
try:
|
||||||
# 更新字段(只更新非 None 的字段)
|
# 更新字段(只更新非 None 的字段)
|
||||||
update_data = validated_data.model_dump(exclude_unset=True)
|
update_data = validated_data.model_dump(exclude_unset=True)
|
||||||
|
config_in_payload = "config" in update_data
|
||||||
|
claude_advanced_in_payload = "claude_code_advanced" in update_data
|
||||||
|
pool_advanced_in_payload = "pool_advanced" in update_data
|
||||||
|
provider_config = (
|
||||||
|
dict(update_data.pop("config") or {})
|
||||||
|
if config_in_payload
|
||||||
|
else dict(provider.config or {})
|
||||||
|
)
|
||||||
|
claude_advanced = (
|
||||||
|
update_data.pop("claude_code_advanced") if claude_advanced_in_payload else None
|
||||||
|
)
|
||||||
|
pool_advanced = update_data.pop("pool_advanced") if pool_advanced_in_payload else None
|
||||||
|
target_provider_type = (
|
||||||
|
update_data.get("provider_type")
|
||||||
|
or getattr(provider, "provider_type", None)
|
||||||
|
or "custom"
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_config, config_changed_by_claude = _merge_claude_code_advanced_config(
|
||||||
|
provider_type=target_provider_type,
|
||||||
|
provider_config=provider_config,
|
||||||
|
claude_code_advanced=claude_advanced,
|
||||||
|
claude_advanced_in_payload=claude_advanced_in_payload,
|
||||||
|
)
|
||||||
|
provider_config, config_changed_by_pool = _merge_pool_advanced_config(
|
||||||
|
provider_config=provider_config,
|
||||||
|
pool_advanced=pool_advanced,
|
||||||
|
pool_advanced_in_payload=pool_advanced_in_payload,
|
||||||
|
)
|
||||||
|
|
||||||
|
config_touched = (
|
||||||
|
config_in_payload
|
||||||
|
or claude_advanced_in_payload
|
||||||
|
or config_changed_by_claude
|
||||||
|
or pool_advanced_in_payload
|
||||||
|
or config_changed_by_pool
|
||||||
|
)
|
||||||
|
if config_touched:
|
||||||
|
update_data["config"] = provider_config
|
||||||
|
|
||||||
for field, value in update_data.items():
|
for field, value in update_data.items():
|
||||||
if field == "billing_type" and value is not None:
|
if field == "billing_type" and value is not None:
|
||||||
@@ -719,3 +850,165 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
|||||||
truncated_keys=truncated_keys,
|
truncated_keys=truncated_keys,
|
||||||
truncated_models=truncated_models,
|
truncated_models=truncated_models,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ========== Claude Code Pool Management ==========
|
||||||
|
|
||||||
|
|
||||||
|
class PoolKeyStatus(BaseModel):
|
||||||
|
"""Single key's pool status."""
|
||||||
|
|
||||||
|
key_id: str
|
||||||
|
key_name: str
|
||||||
|
is_active: bool
|
||||||
|
cooldown_reason: str | None = None
|
||||||
|
cooldown_ttl_seconds: int | None = None
|
||||||
|
cost_window_usage: int = 0
|
||||||
|
cost_limit: int | None = None
|
||||||
|
sticky_sessions: int = 0
|
||||||
|
lru_score: float | None = None
|
||||||
|
|
||||||
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
|
|
||||||
|
class PoolStatusResponse(BaseModel):
|
||||||
|
"""Pool status for a Provider with pool config."""
|
||||||
|
|
||||||
|
provider_id: str
|
||||||
|
provider_name: str
|
||||||
|
pool_enabled: bool = False
|
||||||
|
total_keys: int = 0
|
||||||
|
total_sticky_sessions: int = 0
|
||||||
|
keys: list[PoolKeyStatus] = Field(default_factory=list)
|
||||||
|
|
||||||
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{provider_id}/pool-status", response_model=PoolStatusResponse)
|
||||||
|
async def get_pool_status(
|
||||||
|
request: Request,
|
||||||
|
provider_id: str,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> PoolStatusResponse:
|
||||||
|
"""获取 Provider 的号池状态。"""
|
||||||
|
provider = db.query(Provider).filter(Provider.id == provider_id).first()
|
||||||
|
if not provider:
|
||||||
|
raise NotFoundException("提供商不存在", "provider")
|
||||||
|
|
||||||
|
from src.services.provider.pool import redis_ops as pool_redis
|
||||||
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
|
||||||
|
pcfg = parse_pool_config(provider.config)
|
||||||
|
if pcfg is None:
|
||||||
|
return PoolStatusResponse(
|
||||||
|
provider_id=provider.id,
|
||||||
|
provider_name=provider.name,
|
||||||
|
pool_enabled=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == provider_id).all()
|
||||||
|
|
||||||
|
key_ids = [str(k.id) for k in keys]
|
||||||
|
pid = str(provider.id)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
# Batch fetch pool state (parallel)
|
||||||
|
lru_coro = (
|
||||||
|
pool_redis.get_lru_scores(pid, key_ids) if pcfg.lru_enabled else asyncio.sleep(0, result={})
|
||||||
|
)
|
||||||
|
cooldowns, cooldown_ttls, lru_scores, cost_totals, total_sticky = await asyncio.gather(
|
||||||
|
pool_redis.batch_get_cooldowns(pid, key_ids),
|
||||||
|
pool_redis.batch_get_cooldown_ttls(pid, key_ids),
|
||||||
|
lru_coro,
|
||||||
|
pool_redis.batch_get_cost_totals(pid, key_ids, pcfg.cost_window_seconds),
|
||||||
|
pool_redis.get_sticky_session_count(pid),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Sticky count per key requires SCAN+MGET; batch with gather.
|
||||||
|
sticky_counts: dict[str, int] = {}
|
||||||
|
if key_ids:
|
||||||
|
counts = await asyncio.gather(
|
||||||
|
*(pool_redis.get_key_sticky_count(pid, kid) for kid in key_ids)
|
||||||
|
)
|
||||||
|
sticky_counts = dict(zip(key_ids, counts))
|
||||||
|
|
||||||
|
key_statuses: list[PoolKeyStatus] = []
|
||||||
|
for k in keys:
|
||||||
|
kid = str(k.id)
|
||||||
|
cd_reason = cooldowns.get(kid)
|
||||||
|
|
||||||
|
key_statuses.append(
|
||||||
|
PoolKeyStatus(
|
||||||
|
key_id=kid,
|
||||||
|
key_name=k.name or "",
|
||||||
|
is_active=bool(k.is_active),
|
||||||
|
cooldown_reason=cd_reason,
|
||||||
|
cooldown_ttl_seconds=cooldown_ttls.get(kid) if cd_reason else None,
|
||||||
|
cost_window_usage=cost_totals.get(kid, 0),
|
||||||
|
cost_limit=pcfg.cost_limit_per_key_tokens,
|
||||||
|
sticky_sessions=sticky_counts.get(kid, 0),
|
||||||
|
lru_score=lru_scores.get(kid),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return PoolStatusResponse(
|
||||||
|
provider_id=provider.id,
|
||||||
|
provider_name=provider.name,
|
||||||
|
pool_enabled=True,
|
||||||
|
total_keys=len(keys),
|
||||||
|
total_sticky_sessions=total_sticky,
|
||||||
|
keys=key_statuses,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{provider_id}/pool/clear-cooldown/{key_id}")
|
||||||
|
async def clear_pool_cooldown(
|
||||||
|
request: Request,
|
||||||
|
provider_id: str,
|
||||||
|
key_id: str,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> dict[str, str]:
|
||||||
|
"""手动清除指定 Key 的号池冷却状态。"""
|
||||||
|
provider = db.query(Provider).filter(Provider.id == provider_id).first()
|
||||||
|
if not provider:
|
||||||
|
raise NotFoundException("提供商不存在", "provider")
|
||||||
|
|
||||||
|
key = (
|
||||||
|
db.query(ProviderAPIKey)
|
||||||
|
.filter(ProviderAPIKey.id == key_id, ProviderAPIKey.provider_id == provider_id)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
if not key:
|
||||||
|
raise NotFoundException("密钥不存在", "key")
|
||||||
|
|
||||||
|
from src.services.provider.pool import redis_ops as pool_redis
|
||||||
|
|
||||||
|
await pool_redis.clear_cooldown(str(provider.id), str(key.id))
|
||||||
|
return {"message": f"已清除 Key {key.name or key_id} 的冷却状态"}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{provider_id}/pool/reset-cost/{key_id}")
|
||||||
|
async def reset_pool_cost(
|
||||||
|
request: Request,
|
||||||
|
provider_id: str,
|
||||||
|
key_id: str,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
) -> dict[str, str]:
|
||||||
|
"""重置指定 Key 的号池成本窗口。"""
|
||||||
|
provider = db.query(Provider).filter(Provider.id == provider_id).first()
|
||||||
|
if not provider:
|
||||||
|
raise NotFoundException("提供商不存在", "provider")
|
||||||
|
|
||||||
|
key = (
|
||||||
|
db.query(ProviderAPIKey)
|
||||||
|
.filter(ProviderAPIKey.id == key_id, ProviderAPIKey.provider_id == provider_id)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
if not key:
|
||||||
|
raise NotFoundException("密钥不存在", "key")
|
||||||
|
|
||||||
|
from src.services.provider.pool import redis_ops as pool_redis
|
||||||
|
|
||||||
|
await pool_redis.clear_cost(str(provider.id), str(key.id))
|
||||||
|
return {"message": f"已重置 Key {key.name or key_id} 的成本窗口"}
|
||||||
|
|||||||
@@ -15,9 +15,10 @@ from src.api.base.context import ApiRequestContext
|
|||||||
from src.api.base.models_service import invalidate_models_list_cache
|
from src.api.base.models_service import invalidate_models_list_cache
|
||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.core.enums import ProviderBillingType
|
from src.core.enums import ProviderBillingType
|
||||||
from src.core.exceptions import NotFoundException
|
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
|
from src.models.admin_requests import ClaudeCodeAdvancedConfig, PoolAdvancedConfig
|
||||||
from src.models.database import (
|
from src.models.database import (
|
||||||
Model,
|
Model,
|
||||||
Provider,
|
Provider,
|
||||||
@@ -213,6 +214,74 @@ async def update_provider_settings(
|
|||||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_pool_advanced_from_config(
|
||||||
|
provider_config: dict[str, Any] | None,
|
||||||
|
*,
|
||||||
|
provider_id: str,
|
||||||
|
) -> PoolAdvancedConfig | None:
|
||||||
|
"""从 Provider.config 中安全提取通用号池配置。
|
||||||
|
|
||||||
|
优先查找 ``pool_advanced``,回退查找 ``claude_code_advanced`` 中的号池字段。
|
||||||
|
"""
|
||||||
|
cfg = provider_config or {}
|
||||||
|
raw = cfg.get("pool_advanced")
|
||||||
|
if raw is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if isinstance(raw, PoolAdvancedConfig):
|
||||||
|
return raw
|
||||||
|
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
logger.warning(
|
||||||
|
"Provider {} 的 pool_advanced 类型无效: {},已忽略",
|
||||||
|
provider_id,
|
||||||
|
type(raw).__name__,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
return PoolAdvancedConfig.model_validate(raw)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Provider {} 的 pool_advanced 配置无效,已忽略: {}",
|
||||||
|
provider_id,
|
||||||
|
str(exc),
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_claude_code_advanced_from_config(
|
||||||
|
provider_config: dict[str, Any] | None,
|
||||||
|
*,
|
||||||
|
provider_id: str,
|
||||||
|
) -> ClaudeCodeAdvancedConfig | None:
|
||||||
|
"""从 Provider.config 中安全提取 Claude Code 高级配置。"""
|
||||||
|
raw_config = (provider_config or {}).get("claude_code_advanced")
|
||||||
|
if raw_config is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if isinstance(raw_config, ClaudeCodeAdvancedConfig):
|
||||||
|
return raw_config
|
||||||
|
|
||||||
|
if not isinstance(raw_config, dict):
|
||||||
|
logger.warning(
|
||||||
|
"Provider {} 的 claude_code_advanced 类型无效: {},已忽略",
|
||||||
|
provider_id,
|
||||||
|
type(raw_config).__name__,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
return ClaudeCodeAdvancedConfig.model_validate(raw_config)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Provider {} 的 claude_code_advanced 配置无效,已忽略: {}",
|
||||||
|
provider_id,
|
||||||
|
str(exc),
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndpointsSummary:
|
def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndpointsSummary:
|
||||||
endpoints = db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
|
endpoints = db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
|
||||||
|
|
||||||
@@ -311,12 +380,29 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
|
|||||||
for e in endpoints
|
for e in endpoints
|
||||||
]
|
]
|
||||||
|
|
||||||
|
provider_config_raw = provider.config
|
||||||
|
provider_config = provider_config_raw if isinstance(provider_config_raw, dict) else {}
|
||||||
|
if provider_config_raw is not None and not isinstance(provider_config_raw, dict):
|
||||||
|
logger.warning(
|
||||||
|
"Provider {} 的 config 类型无效: {},按空配置处理",
|
||||||
|
provider.id,
|
||||||
|
type(provider_config_raw).__name__,
|
||||||
|
)
|
||||||
|
|
||||||
# 检查是否配置了 Provider Ops(余额监控等)
|
# 检查是否配置了 Provider Ops(余额监控等)
|
||||||
provider_ops_config = (provider.config or {}).get("provider_ops")
|
provider_ops_config = provider_config.get("provider_ops")
|
||||||
ops_configured = bool(provider_ops_config)
|
ops_configured = bool(provider_ops_config)
|
||||||
ops_architecture_id = (
|
ops_architecture_id = (
|
||||||
provider_ops_config.get("architecture_id") if provider_ops_config else None
|
provider_ops_config.get("architecture_id") if provider_ops_config else None
|
||||||
)
|
)
|
||||||
|
claude_code_advanced = _extract_claude_code_advanced_from_config(
|
||||||
|
provider_config,
|
||||||
|
provider_id=str(provider.id),
|
||||||
|
)
|
||||||
|
pool_advanced = _extract_pool_advanced_from_config(
|
||||||
|
provider_config,
|
||||||
|
provider_id=str(provider.id),
|
||||||
|
)
|
||||||
|
|
||||||
return ProviderWithEndpointsSummary(
|
return ProviderWithEndpointsSummary(
|
||||||
id=provider.id,
|
id=provider.id,
|
||||||
@@ -338,6 +424,8 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
|
|||||||
proxy=provider.proxy,
|
proxy=provider.proxy,
|
||||||
stream_first_byte_timeout=provider.stream_first_byte_timeout,
|
stream_first_byte_timeout=provider.stream_first_byte_timeout,
|
||||||
request_timeout=provider.request_timeout,
|
request_timeout=provider.request_timeout,
|
||||||
|
claude_code_advanced=claude_code_advanced,
|
||||||
|
pool_advanced=pool_advanced,
|
||||||
total_endpoints=total_endpoints,
|
total_endpoints=total_endpoints,
|
||||||
active_endpoints=active_endpoints,
|
active_endpoints=active_endpoints,
|
||||||
total_keys=total_keys,
|
total_keys=total_keys,
|
||||||
@@ -496,6 +584,30 @@ class AdminUpdateProviderSettingsAdapter(AdminApiAdapter):
|
|||||||
raise NotFoundException("Provider not found", "provider")
|
raise NotFoundException("Provider not found", "provider")
|
||||||
|
|
||||||
update_dict = self.update_data.model_dump(exclude_unset=True)
|
update_dict = self.update_data.model_dump(exclude_unset=True)
|
||||||
|
if "claude_code_advanced" in update_dict:
|
||||||
|
claude_advanced = update_dict.pop("claude_code_advanced")
|
||||||
|
provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
|
||||||
|
if claude_advanced is not None and provider_type != "claude_code":
|
||||||
|
raise InvalidRequestException(
|
||||||
|
"claude_code_advanced 仅适用于 provider_type=claude_code"
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_config = dict(provider.config or {})
|
||||||
|
if claude_advanced is None:
|
||||||
|
provider_config.pop("claude_code_advanced", None)
|
||||||
|
else:
|
||||||
|
provider_config["claude_code_advanced"] = dict(claude_advanced)
|
||||||
|
update_dict["config"] = provider_config or None
|
||||||
|
|
||||||
|
if "pool_advanced" in update_dict:
|
||||||
|
pool_advanced = update_dict.pop("pool_advanced")
|
||||||
|
provider_config = dict(update_dict.get("config") or provider.config or {})
|
||||||
|
if pool_advanced is None:
|
||||||
|
provider_config.pop("pool_advanced", None)
|
||||||
|
else:
|
||||||
|
provider_config["pool_advanced"] = dict(pool_advanced)
|
||||||
|
update_dict["config"] = provider_config or None
|
||||||
|
|
||||||
if "billing_type" in update_dict and update_dict["billing_type"] is not None:
|
if "billing_type" in update_dict and update_dict["billing_type"] is not None:
|
||||||
update_dict["billing_type"] = ProviderBillingType(update_dict["billing_type"])
|
update_dict["billing_type"] = ProviderBillingType(update_dict["billing_type"])
|
||||||
|
|
||||||
|
|||||||
@@ -102,6 +102,7 @@ class ProviderRequestResult:
|
|||||||
provider_api_format: str = ""
|
provider_api_format: str = ""
|
||||||
client_api_format: str = ""
|
client_api_format: str = ""
|
||||||
auth_info: Any = None
|
auth_info: Any = None
|
||||||
|
tls_profile: str | None = None
|
||||||
|
|
||||||
|
|
||||||
class ChatHandlerBase(BaseMessageHandler, ABC):
|
class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||||
@@ -579,6 +580,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
ctx.provider_id = provider_id
|
ctx.provider_id = provider_id
|
||||||
ctx.endpoint_id = endpoint_id
|
ctx.endpoint_id = endpoint_id
|
||||||
ctx.key_id = key_id
|
ctx.key_id = key_id
|
||||||
|
if getattr(exec_result, "pool_summary", None):
|
||||||
|
ctx.pool_summary = exec_result.pool_summary
|
||||||
# 同步整流状态(如果请求体被整流过)
|
# 同步整流状态(如果请求体被整流过)
|
||||||
ctx.rectified = request_body_ref.get("_rectified", False)
|
ctx.rectified = request_body_ref.get("_rectified", False)
|
||||||
|
|
||||||
@@ -714,6 +717,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
policy=upstream_policy,
|
policy=upstream_policy,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Envelope lifecycle: prepare_context (pre-wrap hook).
|
||||||
|
envelope_tls_profile: str | None = None
|
||||||
|
if envelope and hasattr(envelope, "prepare_context"):
|
||||||
|
envelope_tls_profile = envelope.prepare_context(
|
||||||
|
provider_config=getattr(provider, "config", None),
|
||||||
|
key_id=str(getattr(key, "id", "") or ""),
|
||||||
|
is_stream=upstream_is_stream,
|
||||||
|
provider_id=str(getattr(provider, "id", "") or ""),
|
||||||
|
)
|
||||||
|
|
||||||
# 跨格式:先做请求体转换(失败触发 failover)
|
# 跨格式:先做请求体转换(失败触发 failover)
|
||||||
registry = get_format_converter_registry()
|
registry = get_format_converter_registry()
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
@@ -776,6 +789,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
url_model=url_model,
|
url_model=url_model,
|
||||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||||
)
|
)
|
||||||
|
# Envelope lifecycle: post_wrap_request (post-wrap hook).
|
||||||
|
if hasattr(envelope, "post_wrap_request"):
|
||||||
|
await envelope.post_wrap_request(request_body)
|
||||||
|
|
||||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||||
extra_headers: dict[str, str] = {}
|
extra_headers: dict[str, str] = {}
|
||||||
@@ -793,6 +809,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
provider_api_format=provider_api_format,
|
provider_api_format=provider_api_format,
|
||||||
client_api_format=client_api_format,
|
client_api_format=client_api_format,
|
||||||
auth_info=auth_info,
|
auth_info=auth_info,
|
||||||
|
tls_profile=envelope_tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _execute_stream_request(
|
async def _execute_stream_request(
|
||||||
@@ -853,6 +870,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
envelope = prep.envelope
|
envelope = prep.envelope
|
||||||
upstream_is_stream = prep.upstream_is_stream
|
upstream_is_stream = prep.upstream_is_stream
|
||||||
auth_info = prep.auth_info
|
auth_info = prep.auth_info
|
||||||
|
tls_profile = prep.tls_profile
|
||||||
|
|
||||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||||
provider_payload, provider_headers = self._request_builder.build(
|
provider_payload, provider_headers = self._request_builder.build(
|
||||||
@@ -863,6 +881,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
is_stream=upstream_is_stream,
|
is_stream=upstream_is_stream,
|
||||||
extra_headers=prep.extra_headers if prep.extra_headers else None,
|
extra_headers=prep.extra_headers if prep.extra_headers else None,
|
||||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||||
|
envelope=envelope,
|
||||||
)
|
)
|
||||||
if upstream_is_stream:
|
if upstream_is_stream:
|
||||||
from src.core.api_format.headers import set_accept_if_absent
|
from src.core.api_format.headers import set_accept_if_absent
|
||||||
@@ -908,7 +927,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||||
http_client = await HTTPClientPool.get_upstream_client(
|
http_client = await HTTPClientPool.get_upstream_client(
|
||||||
delegate_cfg, proxy_config=effective_proxy
|
delegate_cfg,
|
||||||
|
proxy_config=effective_proxy,
|
||||||
|
tls_profile=tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1104,7 +1125,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||||
http_client = await HTTPClientPool.get_upstream_client(
|
http_client = await HTTPClientPool.get_upstream_client(
|
||||||
delegate_cfg, proxy_config=effective_proxy
|
delegate_cfg,
|
||||||
|
proxy_config=effective_proxy,
|
||||||
|
tls_profile=tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
||||||
|
|||||||
@@ -193,6 +193,8 @@ class ChatSyncExecutor:
|
|||||||
request_metadata = handler._build_request_metadata() or {}
|
request_metadata = handler._build_request_metadata() or {}
|
||||||
if ctx.sync_proxy_info:
|
if ctx.sync_proxy_info:
|
||||||
request_metadata["proxy"] = ctx.sync_proxy_info
|
request_metadata["proxy"] = ctx.sync_proxy_info
|
||||||
|
if getattr(exec_result, "pool_summary", None):
|
||||||
|
request_metadata["pool_summary"] = exec_result.pool_summary
|
||||||
total_cost = await handler.telemetry.record_success( # noqa: F841
|
total_cost = await handler.telemetry.record_success( # noqa: F841
|
||||||
provider=ctx.provider_name,
|
provider=ctx.provider_name,
|
||||||
model=model,
|
model=model,
|
||||||
@@ -425,6 +427,7 @@ class ChatSyncExecutor:
|
|||||||
envelope = prep.envelope
|
envelope = prep.envelope
|
||||||
upstream_is_stream = prep.upstream_is_stream
|
upstream_is_stream = prep.upstream_is_stream
|
||||||
auth_info = prep.auth_info
|
auth_info = prep.auth_info
|
||||||
|
tls_profile = prep.tls_profile
|
||||||
|
|
||||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||||
provider_payload, provider_hdrs = handler._request_builder.build(
|
provider_payload, provider_hdrs = handler._request_builder.build(
|
||||||
@@ -435,6 +438,7 @@ class ChatSyncExecutor:
|
|||||||
is_stream=upstream_is_stream,
|
is_stream=upstream_is_stream,
|
||||||
extra_headers=prep.extra_headers if prep.extra_headers else None,
|
extra_headers=prep.extra_headers if prep.extra_headers else None,
|
||||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||||
|
envelope=envelope,
|
||||||
)
|
)
|
||||||
if upstream_is_stream:
|
if upstream_is_stream:
|
||||||
from src.core.api_format.headers import set_accept_if_absent
|
from src.core.api_format.headers import set_accept_if_absent
|
||||||
@@ -496,7 +500,9 @@ class ChatSyncExecutor:
|
|||||||
|
|
||||||
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
||||||
http_client = await HTTPClientPool.get_upstream_client(
|
http_client = await HTTPClientPool.get_upstream_client(
|
||||||
delegate_cfg, proxy_config=_effective_proxy
|
delegate_cfg,
|
||||||
|
proxy_config=_effective_proxy,
|
||||||
|
tls_profile=tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||||
|
|||||||
@@ -342,6 +342,8 @@ class CliPrefetchMixin:
|
|||||||
last_data_time = time.time()
|
last_data_time = time.time()
|
||||||
buffer = b""
|
buffer = b""
|
||||||
output_state = {"first_yield": True, "streaming_updated": False}
|
output_state = {"first_yield": True, "streaming_updated": False}
|
||||||
|
_sample_lines: list[str] = [] # 采集前几行原始内容,用于空流诊断
|
||||||
|
_MAX_SAMPLE_LINES = 5
|
||||||
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
|
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
|
||||||
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
|
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
|
||||||
|
|
||||||
@@ -409,6 +411,8 @@ class CliPrefetchMixin:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
ctx.chunk_count += 1
|
ctx.chunk_count += 1
|
||||||
|
if len(_sample_lines) < _MAX_SAMPLE_LINES:
|
||||||
|
_sample_lines.append(normalized_line[:200])
|
||||||
|
|
||||||
# 格式转换或直接透传
|
# 格式转换或直接透传
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
@@ -468,6 +472,8 @@ class CliPrefetchMixin:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
ctx.chunk_count += 1
|
ctx.chunk_count += 1
|
||||||
|
if len(_sample_lines) < _MAX_SAMPLE_LINES:
|
||||||
|
_sample_lines.append(normalized_line[:200])
|
||||||
|
|
||||||
# 空流检测:超过阈值且无数据,发送错误事件并结束
|
# 空流检测:超过阈值且无数据,发送错误事件并结束
|
||||||
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
|
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
|
||||||
@@ -524,9 +530,10 @@ class CliPrefetchMixin:
|
|||||||
# 检查是否收到数据
|
# 检查是否收到数据
|
||||||
if ctx.data_count == 0:
|
if ctx.data_count == 0:
|
||||||
# 空流通常意味着配置错误(如 base_url 指向了网页而非 API)
|
# 空流通常意味着配置错误(如 base_url 指向了网页而非 API)
|
||||||
|
sample_info = f", 前几行内容: {_sample_lines!r}" if _sample_lines else ""
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Provider '{ctx.provider_name}' 返回空流式响应 (收到 {ctx.chunk_count} 个非数据行), "
|
f"Provider '{ctx.provider_name}' 返回空流式响应 (收到 {ctx.chunk_count} 个非数据行), "
|
||||||
f"可能是 endpoint base_url 配置错误"
|
f"可能是 endpoint base_url 配置错误{sample_info}"
|
||||||
)
|
)
|
||||||
# 设置错误状态用于后续记录
|
# 设置错误状态用于后续记录
|
||||||
ctx.status_code = 503
|
ctx.status_code = 503
|
||||||
|
|||||||
@@ -188,6 +188,8 @@ class CliStreamMixin:
|
|||||||
ctx.endpoint_id = endpoint_id
|
ctx.endpoint_id = endpoint_id
|
||||||
if not ctx.key_id:
|
if not ctx.key_id:
|
||||||
ctx.key_id = key_id
|
ctx.key_id = key_id
|
||||||
|
if getattr(exec_result, "pool_summary", None):
|
||||||
|
ctx.pool_summary = exec_result.pool_summary
|
||||||
# 同步整流状态(如果请求体被整流过)
|
# 同步整流状态(如果请求体被整流过)
|
||||||
ctx.rectified = request_body_ref.get("_rectified", False)
|
ctx.rectified = request_body_ref.get("_rectified", False)
|
||||||
|
|
||||||
@@ -319,6 +321,15 @@ class CliStreamMixin:
|
|||||||
client_is_stream=True,
|
client_is_stream=True,
|
||||||
policy=upstream_policy,
|
policy=upstream_policy,
|
||||||
)
|
)
|
||||||
|
# Envelope lifecycle: prepare_context (pre-wrap hook).
|
||||||
|
envelope_tls_profile: str | None = None
|
||||||
|
if envelope and hasattr(envelope, "prepare_context"):
|
||||||
|
envelope_tls_profile = envelope.prepare_context(
|
||||||
|
provider_config=getattr(provider, "config", None),
|
||||||
|
key_id=str(getattr(key, "id", "") or ""),
|
||||||
|
is_stream=upstream_is_stream,
|
||||||
|
provider_id=str(getattr(provider, "id", "") or ""),
|
||||||
|
)
|
||||||
|
|
||||||
# 跨格式:先做请求体转换(失败触发 failover)
|
# 跨格式:先做请求体转换(失败触发 failover)
|
||||||
if needs_conversion and provider_api_format:
|
if needs_conversion and provider_api_format:
|
||||||
@@ -374,6 +385,9 @@ class CliStreamMixin:
|
|||||||
url_model=url_model,
|
url_model=url_model,
|
||||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||||
)
|
)
|
||||||
|
# Envelope lifecycle: post_wrap_request (post-wrap hook).
|
||||||
|
if hasattr(envelope, "post_wrap_request"):
|
||||||
|
await envelope.post_wrap_request(request_body)
|
||||||
|
|
||||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||||
extra_headers: dict[str, str] = {}
|
extra_headers: dict[str, str] = {}
|
||||||
@@ -391,6 +405,7 @@ class CliStreamMixin:
|
|||||||
is_stream=upstream_is_stream,
|
is_stream=upstream_is_stream,
|
||||||
extra_headers=extra_headers if extra_headers else None,
|
extra_headers=extra_headers if extra_headers else None,
|
||||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||||
|
envelope=envelope,
|
||||||
)
|
)
|
||||||
if upstream_is_stream:
|
if upstream_is_stream:
|
||||||
from src.core.api_format.headers import set_accept_if_absent
|
from src.core.api_format.headers import set_accept_if_absent
|
||||||
@@ -429,7 +444,9 @@ class CliStreamMixin:
|
|||||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||||
http_client = await HTTPClientPool.get_upstream_client(
|
http_client = await HTTPClientPool.get_upstream_client(
|
||||||
delegate_cfg, proxy_config=effective_proxy
|
delegate_cfg,
|
||||||
|
proxy_config=effective_proxy,
|
||||||
|
tls_profile=envelope_tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -623,7 +640,9 @@ class CliStreamMixin:
|
|||||||
|
|
||||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||||
http_client = await HTTPClientPool.get_upstream_client(
|
http_client = await HTTPClientPool.get_upstream_client(
|
||||||
delegate_cfg, proxy_config=effective_proxy
|
delegate_cfg,
|
||||||
|
proxy_config=effective_proxy,
|
||||||
|
tls_profile=envelope_tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
||||||
@@ -798,6 +817,8 @@ class CliStreamMixin:
|
|||||||
last_data_time = time.time()
|
last_data_time = time.time()
|
||||||
buffer = b""
|
buffer = b""
|
||||||
output_state = {"first_yield": True, "streaming_updated": False}
|
output_state = {"first_yield": True, "streaming_updated": False}
|
||||||
|
_sample_lines: list[str] = [] # 采集前几行原始内容,用于空流诊断
|
||||||
|
_MAX_SAMPLE_LINES = 5
|
||||||
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
|
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
|
||||||
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
|
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
|
||||||
|
|
||||||
@@ -864,6 +885,8 @@ class CliStreamMixin:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
ctx.chunk_count += 1
|
ctx.chunk_count += 1
|
||||||
|
if len(_sample_lines) < _MAX_SAMPLE_LINES:
|
||||||
|
_sample_lines.append(normalized_line[:200])
|
||||||
|
|
||||||
# 空流检测:超过阈值且无数据,发送错误事件并结束
|
# 空流检测:超过阈值且无数据,发送错误事件并结束
|
||||||
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
|
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
|
||||||
@@ -922,7 +945,12 @@ class CliStreamMixin:
|
|||||||
|
|
||||||
# 检查是否收到数据
|
# 检查是否收到数据
|
||||||
if ctx.data_count == 0:
|
if ctx.data_count == 0:
|
||||||
logger.warning("Provider '{}' 返回空流式响应", ctx.provider_name)
|
sample_info = f", 前几行内容: {_sample_lines!r}" if _sample_lines else ""
|
||||||
|
logger.warning(
|
||||||
|
"Provider '{}' 返回空流式响应{}",
|
||||||
|
ctx.provider_name,
|
||||||
|
sample_info,
|
||||||
|
)
|
||||||
ctx.status_code = 503
|
ctx.status_code = 503
|
||||||
ctx.error_message = "上游服务返回了空的流式响应"
|
ctx.error_message = "上游服务返回了空的流式响应"
|
||||||
ctx.upstream_response = (
|
ctx.upstream_response = (
|
||||||
|
|||||||
@@ -147,6 +147,15 @@ class CliSyncMixin:
|
|||||||
client_is_stream=False,
|
client_is_stream=False,
|
||||||
policy=upstream_policy,
|
policy=upstream_policy,
|
||||||
)
|
)
|
||||||
|
# Envelope lifecycle: prepare_context (pre-wrap hook).
|
||||||
|
envelope_tls_profile: str | None = None
|
||||||
|
if envelope and hasattr(envelope, "prepare_context"):
|
||||||
|
envelope_tls_profile = envelope.prepare_context(
|
||||||
|
provider_config=getattr(provider, "config", None),
|
||||||
|
key_id=str(getattr(key, "id", "") or ""),
|
||||||
|
is_stream=upstream_is_stream,
|
||||||
|
provider_id=str(getattr(provider, "id", "") or ""),
|
||||||
|
)
|
||||||
|
|
||||||
# 跨格式:先做请求体转换(失败触发 failover)
|
# 跨格式:先做请求体转换(失败触发 failover)
|
||||||
if needs_conversion and provider_api_format:
|
if needs_conversion and provider_api_format:
|
||||||
@@ -202,6 +211,9 @@ class CliSyncMixin:
|
|||||||
url_model=url_model,
|
url_model=url_model,
|
||||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||||
)
|
)
|
||||||
|
# Envelope lifecycle: post_wrap_request (post-wrap hook).
|
||||||
|
if hasattr(envelope, "post_wrap_request"):
|
||||||
|
await envelope.post_wrap_request(request_body)
|
||||||
|
|
||||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||||
extra_headers: dict[str, str] = {}
|
extra_headers: dict[str, str] = {}
|
||||||
@@ -219,6 +231,7 @@ class CliSyncMixin:
|
|||||||
is_stream=upstream_is_stream,
|
is_stream=upstream_is_stream,
|
||||||
extra_headers=extra_headers if extra_headers else None,
|
extra_headers=extra_headers if extra_headers else None,
|
||||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||||
|
envelope=envelope,
|
||||||
)
|
)
|
||||||
if upstream_is_stream:
|
if upstream_is_stream:
|
||||||
from src.core.api_format.headers import set_accept_if_absent
|
from src.core.api_format.headers import set_accept_if_absent
|
||||||
@@ -274,7 +287,9 @@ class CliSyncMixin:
|
|||||||
|
|
||||||
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
||||||
http_client = await HTTPClientPool.get_upstream_client(
|
http_client = await HTTPClientPool.get_upstream_client(
|
||||||
delegate_cfg, proxy_config=_effective_proxy
|
delegate_cfg,
|
||||||
|
proxy_config=_effective_proxy,
|
||||||
|
tls_profile=envelope_tls_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||||
@@ -521,6 +536,8 @@ class CliSyncMixin:
|
|||||||
request_metadata = self._build_request_metadata() or {}
|
request_metadata = self._build_request_metadata() or {}
|
||||||
if sync_proxy_info:
|
if sync_proxy_info:
|
||||||
request_metadata["proxy"] = sync_proxy_info
|
request_metadata["proxy"] = sync_proxy_info
|
||||||
|
if getattr(exec_result, "pool_summary", None):
|
||||||
|
request_metadata["pool_summary"] = exec_result.pool_summary
|
||||||
total_cost = await self.telemetry.record_success(
|
total_cost = await self.telemetry.record_success(
|
||||||
provider=provider_name,
|
provider=provider_name,
|
||||||
model=model,
|
model=model,
|
||||||
|
|||||||
@@ -26,9 +26,9 @@ from src.core.api_format import (
|
|||||||
make_signature_key,
|
make_signature_key,
|
||||||
)
|
)
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.provider_auth_types import ProviderAuthInfo
|
|
||||||
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
|
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
|
||||||
from src.services.provider.auth import get_provider_auth
|
from src.services.provider.auth import get_provider_auth # noqa: F401
|
||||||
|
from src.services.provider.envelope import ProviderEnvelope
|
||||||
|
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
# 统一的头部配置常量
|
# 统一的头部配置常量
|
||||||
@@ -1036,6 +1036,7 @@ class RequestBuilder(ABC):
|
|||||||
*,
|
*,
|
||||||
extra_headers: dict[str, str] | None = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
pre_computed_auth: tuple[str, str] | None = None,
|
pre_computed_auth: tuple[str, str] | None = None,
|
||||||
|
envelope: ProviderEnvelope | None = None,
|
||||||
) -> dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""构建请求头"""
|
"""构建请求头"""
|
||||||
pass
|
pass
|
||||||
@@ -1051,6 +1052,7 @@ class RequestBuilder(ABC):
|
|||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
extra_headers: dict[str, str] | None = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
pre_computed_auth: tuple[str, str] | None = None,
|
pre_computed_auth: tuple[str, str] | None = None,
|
||||||
|
envelope: ProviderEnvelope | None = None,
|
||||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||||
"""
|
"""
|
||||||
构建完整的请求(请求体 + 请求头)
|
构建完整的请求(请求体 + 请求头)
|
||||||
@@ -1085,6 +1087,7 @@ class RequestBuilder(ABC):
|
|||||||
key,
|
key,
|
||||||
extra_headers=extra_headers,
|
extra_headers=extra_headers,
|
||||||
pre_computed_auth=pre_computed_auth,
|
pre_computed_auth=pre_computed_auth,
|
||||||
|
envelope=envelope,
|
||||||
)
|
)
|
||||||
return payload, headers
|
return payload, headers
|
||||||
|
|
||||||
@@ -1114,6 +1117,70 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
"""
|
"""
|
||||||
return dict(original_body)
|
return dict(original_body)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _merge_comma_header_values(primary: str, secondary: str) -> str:
|
||||||
|
"""合并逗号分隔 header 值并去重,保持 primary 在前。"""
|
||||||
|
seen: set[str] = set()
|
||||||
|
merged: list[str] = []
|
||||||
|
|
||||||
|
def _append(raw: str) -> None:
|
||||||
|
for token in str(raw or "").split(","):
|
||||||
|
token = token.strip()
|
||||||
|
if not token or token in seen:
|
||||||
|
continue
|
||||||
|
seen.add(token)
|
||||||
|
merged.append(token)
|
||||||
|
|
||||||
|
_append(primary)
|
||||||
|
_append(secondary)
|
||||||
|
return ",".join(merged)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _drop_beta_token(cls, value: str, token: str) -> str:
|
||||||
|
"""从逗号分隔 header 中移除指定 token。"""
|
||||||
|
if not value or token not in value:
|
||||||
|
return value
|
||||||
|
return cls._merge_comma_header_values(
|
||||||
|
",".join(p.strip() for p in str(value).split(",") if p.strip() and p.strip() != token),
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _merge_extra_headers_with_original(
|
||||||
|
cls,
|
||||||
|
original_headers: dict[str, str],
|
||||||
|
extra_headers: dict[str, str] | None,
|
||||||
|
*,
|
||||||
|
envelope: ProviderEnvelope | None = None,
|
||||||
|
) -> dict[str, str] | None:
|
||||||
|
"""合并 extra_headers 与原始头部中的特定字段。"""
|
||||||
|
if not extra_headers:
|
||||||
|
return None
|
||||||
|
|
||||||
|
merged_extra = dict(extra_headers)
|
||||||
|
beta_extra_key = next((k for k in merged_extra if k.lower() == "anthropic-beta"), None)
|
||||||
|
if beta_extra_key is None:
|
||||||
|
return merged_extra
|
||||||
|
|
||||||
|
incoming_beta = next(
|
||||||
|
(v for k, v in original_headers.items() if k.lower() == "anthropic-beta"),
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
merged_beta = str(merged_extra.get(beta_extra_key) or "")
|
||||||
|
if incoming_beta:
|
||||||
|
merged_beta = cls._merge_comma_header_values(
|
||||||
|
merged_beta,
|
||||||
|
str(incoming_beta),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 由 envelope 声明需要排除的 beta token(如 Claude Code OAuth 的 context-1m)。
|
||||||
|
if envelope and hasattr(envelope, "excluded_beta_tokens"):
|
||||||
|
for token in envelope.excluded_beta_tokens():
|
||||||
|
merged_beta = cls._drop_beta_token(merged_beta, token)
|
||||||
|
|
||||||
|
merged_extra[beta_extra_key] = merged_beta
|
||||||
|
return merged_extra
|
||||||
|
|
||||||
def build_headers(
|
def build_headers(
|
||||||
self,
|
self,
|
||||||
original_headers: dict[str, str],
|
original_headers: dict[str, str],
|
||||||
@@ -1122,6 +1189,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
*,
|
*,
|
||||||
extra_headers: dict[str, str] | None = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
pre_computed_auth: tuple[str, str] | None = None,
|
pre_computed_auth: tuple[str, str] | None = None,
|
||||||
|
envelope: ProviderEnvelope | None = None,
|
||||||
) -> dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||||
@@ -1177,8 +1245,13 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
builder.apply_rules(header_rules, protected_keys)
|
builder.apply_rules(header_rules, protected_keys)
|
||||||
|
|
||||||
# 4. 添加额外头部
|
# 4. 添加额外头部
|
||||||
if extra_headers:
|
effective_extra_headers = self._merge_extra_headers_with_original(
|
||||||
builder.add_many(extra_headers)
|
original_headers,
|
||||||
|
extra_headers,
|
||||||
|
envelope=envelope,
|
||||||
|
)
|
||||||
|
if effective_extra_headers:
|
||||||
|
builder.add_many(effective_extra_headers)
|
||||||
|
|
||||||
# 5. 设置认证头(最高优先级,上游始终使用 header 认证)
|
# 5. 设置认证头(最高优先级,上游始终使用 header 认证)
|
||||||
builder.add(auth_header, auth_value)
|
builder.add(auth_header, auth_value)
|
||||||
|
|||||||
@@ -135,6 +135,9 @@ class StreamContext:
|
|||||||
# 代理信息(用于 usage 记录和日志,含 ttfb_ms)
|
# 代理信息(用于 usage 记录和日志,含 ttfb_ms)
|
||||||
proxy_info: dict[str, Any] | None = None
|
proxy_info: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
# 号池调度摘要(来自 ExecutionResult.pool_summary)
|
||||||
|
pool_summary: dict[str, Any] | None = None
|
||||||
|
|
||||||
# 流式格式转换状态(跨 chunk 追踪)
|
# 流式格式转换状态(跨 chunk 追踪)
|
||||||
stream_conversion_state: StreamState | None = None
|
stream_conversion_state: StreamState | None = None
|
||||||
stream_conversion_event_count: int = 0 # 流式转换成功的 event 计数
|
stream_conversion_event_count: int = 0 # 流式转换成功的 event 计数
|
||||||
|
|||||||
@@ -208,6 +208,8 @@ class StreamTelemetryRecorder:
|
|||||||
metadata["perf"] = ctx.perf_metrics
|
metadata["perf"] = ctx.perf_metrics
|
||||||
if ctx.proxy_info:
|
if ctx.proxy_info:
|
||||||
metadata["proxy"] = ctx.proxy_info
|
metadata["proxy"] = ctx.proxy_info
|
||||||
|
if ctx.pool_summary:
|
||||||
|
metadata["pool_summary"] = ctx.pool_summary
|
||||||
|
|
||||||
await writer.record_success(
|
await writer.record_success(
|
||||||
provider=ctx.provider_name or "unknown",
|
provider=ctx.provider_name or "unknown",
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from src.services.proxy_node.resolver import (
|
|||||||
get_system_proxy_config,
|
get_system_proxy_config,
|
||||||
make_proxy_param,
|
make_proxy_param,
|
||||||
)
|
)
|
||||||
from src.utils.ssl_utils import get_ssl_context
|
from src.utils.ssl_utils import get_ssl_context, get_ssl_context_for_profile
|
||||||
|
|
||||||
# 模块级锁,避免类属性延迟初始化的竞态条件
|
# 模块级锁,避免类属性延迟初始化的竞态条件
|
||||||
_proxy_clients_lock = asyncio.Lock()
|
_proxy_clients_lock = asyncio.Lock()
|
||||||
@@ -186,6 +186,7 @@ class HTTPClientPool:
|
|||||||
async def get_proxy_client(
|
async def get_proxy_client(
|
||||||
cls,
|
cls,
|
||||||
proxy_config: dict[str, Any] | None = None,
|
proxy_config: dict[str, Any] | None = None,
|
||||||
|
tls_profile: str | None = None,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
"""
|
"""
|
||||||
获取代理客户端(带缓存复用)
|
获取代理客户端(带缓存复用)
|
||||||
@@ -212,6 +213,9 @@ class HTTPClientPool:
|
|||||||
return await cls._get_tunnel_client(delegate_cfg["node_id"])
|
return await cls._get_tunnel_client(delegate_cfg["node_id"])
|
||||||
|
|
||||||
cache_key = compute_proxy_cache_key(proxy_config)
|
cache_key = compute_proxy_cache_key(proxy_config)
|
||||||
|
tls_profile_key = str(tls_profile or "").strip().lower()
|
||||||
|
if tls_profile_key:
|
||||||
|
cache_key = f"{cache_key}::tls:{tls_profile_key}"
|
||||||
|
|
||||||
# 无代理时返回默认客户端
|
# 无代理时返回默认客户端
|
||||||
if cache_key == "__no_proxy__":
|
if cache_key == "__no_proxy__":
|
||||||
@@ -229,6 +233,10 @@ class HTTPClientPool:
|
|||||||
else:
|
else:
|
||||||
# 更新最后使用时间
|
# 更新最后使用时间
|
||||||
cls._proxy_clients[cache_key] = (client, time.time())
|
cls._proxy_clients[cache_key] = (client, time.time())
|
||||||
|
if tls_profile_key:
|
||||||
|
logger.debug(
|
||||||
|
"复用代理客户端 TLS profile={} key={}", tls_profile_key, cache_key
|
||||||
|
)
|
||||||
return client
|
return client
|
||||||
|
|
||||||
# 淘汰旧客户端(如果超过上限)
|
# 淘汰旧客户端(如果超过上限)
|
||||||
@@ -237,7 +245,7 @@ class HTTPClientPool:
|
|||||||
# 创建新客户端(使用默认超时,请求时可覆盖)
|
# 创建新客户端(使用默认超时,请求时可覆盖)
|
||||||
client_config: dict[str, Any] = {
|
client_config: dict[str, Any] = {
|
||||||
"http2": False,
|
"http2": False,
|
||||||
"verify": get_ssl_context(),
|
"verify": get_ssl_context_for_profile(tls_profile),
|
||||||
"follow_redirects": True,
|
"follow_redirects": True,
|
||||||
"limits": httpx.Limits(
|
"limits": httpx.Limits(
|
||||||
max_connections=config.http_max_connections,
|
max_connections=config.http_max_connections,
|
||||||
@@ -269,6 +277,8 @@ class HTTPClientPool:
|
|||||||
logger.debug(
|
logger.debug(
|
||||||
"创建代理客户端(缓存): {}, 缓存数量: {}", proxy_label, len(cls._proxy_clients)
|
"创建代理客户端(缓存): {}, 缓存数量: {}", proxy_label, len(cls._proxy_clients)
|
||||||
)
|
)
|
||||||
|
if tls_profile_key:
|
||||||
|
logger.debug("创建代理客户端 TLS profile={} key={}", tls_profile_key, cache_key)
|
||||||
|
|
||||||
return client
|
return client
|
||||||
|
|
||||||
@@ -342,6 +352,7 @@ class HTTPClientPool:
|
|||||||
cls,
|
cls,
|
||||||
proxy_config: dict[str, Any] | None = None,
|
proxy_config: dict[str, Any] | None = None,
|
||||||
timeout: httpx.Timeout | None = None,
|
timeout: httpx.Timeout | None = None,
|
||||||
|
tls_profile: str | None = None,
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
"""
|
"""
|
||||||
@@ -359,7 +370,7 @@ class HTTPClientPool:
|
|||||||
"""
|
"""
|
||||||
client_config: dict[str, Any] = {
|
client_config: dict[str, Any] = {
|
||||||
"http2": False,
|
"http2": False,
|
||||||
"verify": get_ssl_context(),
|
"verify": get_ssl_context_for_profile(tls_profile),
|
||||||
"follow_redirects": True,
|
"follow_redirects": True,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -392,6 +403,7 @@ class HTTPClientPool:
|
|||||||
cls,
|
cls,
|
||||||
delegate_cfg: dict[str, Any] | None,
|
delegate_cfg: dict[str, Any] | None,
|
||||||
proxy_config: dict[str, Any] | None = None,
|
proxy_config: dict[str, Any] | None = None,
|
||||||
|
tls_profile: str | None = None,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
"""
|
"""
|
||||||
获取可复用的上游请求客户端(自动选择 tunnel/代理模式)
|
获取可复用的上游请求客户端(自动选择 tunnel/代理模式)
|
||||||
@@ -401,7 +413,7 @@ class HTTPClientPool:
|
|||||||
"""
|
"""
|
||||||
if delegate_cfg and delegate_cfg.get("tunnel"):
|
if delegate_cfg and delegate_cfg.get("tunnel"):
|
||||||
return await cls._get_tunnel_client(delegate_cfg["node_id"])
|
return await cls._get_tunnel_client(delegate_cfg["node_id"])
|
||||||
return await cls.get_proxy_client(proxy_config=proxy_config)
|
return await cls.get_proxy_client(proxy_config=proxy_config, tls_profile=tls_profile)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def create_upstream_stream_client(
|
async def create_upstream_stream_client(
|
||||||
@@ -409,6 +421,7 @@ class HTTPClientPool:
|
|||||||
delegate_cfg: dict[str, Any] | None,
|
delegate_cfg: dict[str, Any] | None,
|
||||||
proxy_config: dict[str, Any] | None = None,
|
proxy_config: dict[str, Any] | None = None,
|
||||||
timeout: httpx.Timeout | None = None,
|
timeout: httpx.Timeout | None = None,
|
||||||
|
tls_profile: str | None = None,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
"""
|
"""
|
||||||
创建上游流式请求客户端(自动选择 tunnel/代理模式)
|
创建上游流式请求客户端(自动选择 tunnel/代理模式)
|
||||||
@@ -417,7 +430,11 @@ class HTTPClientPool:
|
|||||||
"""
|
"""
|
||||||
if delegate_cfg and delegate_cfg.get("tunnel"):
|
if delegate_cfg and delegate_cfg.get("tunnel"):
|
||||||
return await cls._get_tunnel_client(delegate_cfg["node_id"], timeout=timeout)
|
return await cls._get_tunnel_client(delegate_cfg["node_id"], timeout=timeout)
|
||||||
return cls.create_client_with_proxy(proxy_config=proxy_config, timeout=timeout)
|
return cls.create_client_with_proxy(
|
||||||
|
proxy_config=proxy_config,
|
||||||
|
timeout=timeout,
|
||||||
|
tls_profile=tls_profile,
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def _get_tunnel_client(
|
async def _get_tunnel_client(
|
||||||
|
|||||||
@@ -520,6 +520,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
message_raw = chunk.get("message")
|
message_raw = chunk.get("message")
|
||||||
message: dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
message: dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
||||||
msg_id = str(message.get("id") or "")
|
msg_id = str(message.get("id") or "")
|
||||||
|
usage_info = self._claude_usage_to_internal(message.get("usage"))
|
||||||
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||||
model = state.model or str(message.get("model") or "")
|
model = state.model or str(message.get("model") or "")
|
||||||
state.message_id = msg_id or state.message_id
|
state.message_id = msg_id or state.message_id
|
||||||
@@ -527,7 +528,9 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
state.model = model
|
state.model = model
|
||||||
ss["message_started"] = True
|
ss["message_started"] = True
|
||||||
ss.setdefault("block_index_to_tool_id", {})
|
ss.setdefault("block_index_to_tool_id", {})
|
||||||
events.append(MessageStartEvent(message_id=msg_id, model=model))
|
if usage_info is not None:
|
||||||
|
ss["usage"] = message.get("usage")
|
||||||
|
events.append(MessageStartEvent(message_id=msg_id, model=model, usage=usage_info))
|
||||||
return events
|
return events
|
||||||
|
|
||||||
if event_type == "content_block_start":
|
if event_type == "content_block_start":
|
||||||
|
|||||||
@@ -80,6 +80,87 @@ class ProxyConfig(BaseModel):
|
|||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class PoolAdvancedConfig(BaseModel):
|
||||||
|
"""通用号池配置(适用于所有 Provider 类型)。"""
|
||||||
|
|
||||||
|
sticky_session_ttl_seconds: int | None = Field(
|
||||||
|
None,
|
||||||
|
ge=60,
|
||||||
|
le=86400,
|
||||||
|
description="粘性会话 TTL(秒),同一对话始终路由到同一 Key。None = 禁用",
|
||||||
|
)
|
||||||
|
load_threshold_percent: int | None = Field(
|
||||||
|
None,
|
||||||
|
ge=10,
|
||||||
|
le=100,
|
||||||
|
description="负载率阈值(%),超过时该 Key 被降权。默认 80",
|
||||||
|
)
|
||||||
|
lru_enabled: bool = Field(True, description="LRU 调度(优先选择最久未用的 Key)")
|
||||||
|
cost_window_seconds: int | None = Field(
|
||||||
|
None,
|
||||||
|
ge=3600,
|
||||||
|
le=86400,
|
||||||
|
description="滚动成本窗口(秒)。默认 18000(5 小时)",
|
||||||
|
)
|
||||||
|
cost_limit_per_key_tokens: int | None = Field(
|
||||||
|
None, ge=0, description="每个 Key 在窗口内的最大 token 用量。None = 不限"
|
||||||
|
)
|
||||||
|
cost_soft_threshold_percent: int | None = Field(
|
||||||
|
None,
|
||||||
|
ge=0,
|
||||||
|
le=100,
|
||||||
|
description="成本软阈值(%),超过时优先选用其他 Key。默认 80",
|
||||||
|
)
|
||||||
|
rate_limit_cooldown_seconds: int | None = Field(
|
||||||
|
None, ge=10, le=3600, description="429 冷却时间(秒)。默认 300"
|
||||||
|
)
|
||||||
|
overload_cooldown_seconds: int | None = Field(
|
||||||
|
None, ge=5, le=600, description="529 冷却时间(秒)。默认 30"
|
||||||
|
)
|
||||||
|
proactive_refresh_seconds: int | None = Field(
|
||||||
|
None,
|
||||||
|
ge=60,
|
||||||
|
le=600,
|
||||||
|
description="OAuth Token 提前刷新秒数。默认 180(3 分钟)",
|
||||||
|
)
|
||||||
|
health_policy_enabled: bool = Field(
|
||||||
|
True, description="启用号池健康策略(按上游错误码自动冷却/禁用 Key)"
|
||||||
|
)
|
||||||
|
unschedulable_rules: list[dict] | None = Field(
|
||||||
|
None,
|
||||||
|
description="关键词临时不可调度规则: [{'keyword': '...', 'duration_minutes': 5}]",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ClaudeCodeAdvancedConfig(BaseModel):
|
||||||
|
"""Claude Code 特有配置。"""
|
||||||
|
|
||||||
|
max_sessions: int | None = Field(
|
||||||
|
None, ge=1, le=1000, description="最大活跃会话数(为空表示不限制)"
|
||||||
|
)
|
||||||
|
session_idle_timeout_minutes: int | None = Field(
|
||||||
|
None, ge=1, le=1440, description="会话空闲超时(分钟)"
|
||||||
|
)
|
||||||
|
enable_tls_fingerprint: bool = Field(
|
||||||
|
False, description="是否启用 TLS 指纹模拟(模拟 Node.js/Claude Code 客户端)"
|
||||||
|
)
|
||||||
|
session_id_masking_enabled: bool = Field(
|
||||||
|
False, description="是否启用会话 ID 伪装(固定 metadata.user_id 中 session 片段)"
|
||||||
|
)
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def normalize_session_control(self) -> "ClaudeCodeAdvancedConfig":
|
||||||
|
# 未启用会话限制时,不保留超时配置,避免产生误导。
|
||||||
|
if self.max_sessions is None:
|
||||||
|
self.session_idle_timeout_minutes = None
|
||||||
|
return self
|
||||||
|
|
||||||
|
# 启用会话限制但未设置超时时,回落到 5 分钟默认值。
|
||||||
|
if self.session_idle_timeout_minutes is None:
|
||||||
|
self.session_idle_timeout_minutes = 5
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class CreateProviderRequest(BaseModel):
|
class CreateProviderRequest(BaseModel):
|
||||||
"""创建 Provider 请求"""
|
"""创建 Provider 请求"""
|
||||||
|
|
||||||
@@ -148,6 +229,12 @@ class CreateProviderRequest(BaseModel):
|
|||||||
request_timeout: float | None = Field(
|
request_timeout: float | None = Field(
|
||||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||||
)
|
)
|
||||||
|
pool_advanced: PoolAdvancedConfig | None = Field(
|
||||||
|
None, description="号池高级配置(适用于所有 Provider 类型)"
|
||||||
|
)
|
||||||
|
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
|
||||||
|
None, description="Claude Code 特有配置"
|
||||||
|
)
|
||||||
config: dict[str, Any] | None = Field(None, description="其他配置")
|
config: dict[str, Any] | None = Field(None, description="其他配置")
|
||||||
|
|
||||||
@field_validator("provider_type")
|
@field_validator("provider_type")
|
||||||
@@ -214,6 +301,13 @@ class CreateProviderRequest(BaseModel):
|
|||||||
valid_types = [t.value for t in ProviderBillingType]
|
valid_types = [t.value for t in ProviderBillingType]
|
||||||
raise ValueError(f"无效的计费类型,有效值为: {', '.join(valid_types)}")
|
raise ValueError(f"无效的计费类型,有效值为: {', '.join(valid_types)}")
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_claude_code_advanced_scope(self) -> "CreateProviderRequest":
|
||||||
|
provider_type = (self.provider_type or "custom").strip()
|
||||||
|
if self.claude_code_advanced is not None and provider_type != "claude_code":
|
||||||
|
raise ValueError("claude_code_advanced 仅适用于 provider_type=claude_code")
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class UpdateProviderRequest(BaseModel):
|
class UpdateProviderRequest(BaseModel):
|
||||||
"""更新 Provider 请求"""
|
"""更新 Provider 请求"""
|
||||||
@@ -244,6 +338,12 @@ class UpdateProviderRequest(BaseModel):
|
|||||||
request_timeout: float | None = Field(
|
request_timeout: float | None = Field(
|
||||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||||
)
|
)
|
||||||
|
pool_advanced: PoolAdvancedConfig | None = Field(
|
||||||
|
None, description="号池高级配置(适用于所有 Provider 类型)"
|
||||||
|
)
|
||||||
|
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
|
||||||
|
None, description="Claude Code 特有配置"
|
||||||
|
)
|
||||||
config: dict[str, Any] | None = None
|
config: dict[str, Any] | None = None
|
||||||
|
|
||||||
# 复用相同的验证器
|
# 复用相同的验证器
|
||||||
@@ -258,6 +358,15 @@ class UpdateProviderRequest(BaseModel):
|
|||||||
CreateProviderRequest.validate_provider_type.__func__
|
CreateProviderRequest.validate_provider_type.__func__
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_claude_code_advanced_scope(self) -> "UpdateProviderRequest":
|
||||||
|
# 更新场景下 provider_type 可能不在 payload 中,最终校验由路由层结合数据库值完成。
|
||||||
|
if self.claude_code_advanced is not None and self.provider_type is not None:
|
||||||
|
provider_type = (self.provider_type or "custom").strip()
|
||||||
|
if provider_type != "claude_code":
|
||||||
|
raise ValueError("claude_code_advanced 仅适用于 provider_type=claude_code")
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class CreateEndpointRequest(BaseModel):
|
class CreateEndpointRequest(BaseModel):
|
||||||
"""创建 Endpoint 请求"""
|
"""创建 Endpoint 请求"""
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from typing import Any, Literal
|
|||||||
|
|
||||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||||
|
|
||||||
from src.models.admin_requests import ProxyConfig
|
from src.models.admin_requests import ClaudeCodeAdvancedConfig, PoolAdvancedConfig, ProxyConfig
|
||||||
|
|
||||||
# ========== Header Rule 类型定义 ==========
|
# ========== Header Rule 类型定义 ==========
|
||||||
# 请求头规则支持三种操作:
|
# 请求头规则支持三种操作:
|
||||||
@@ -934,6 +934,10 @@ class ProviderUpdateRequest(BaseModel):
|
|||||||
request_timeout: float | None = Field(
|
request_timeout: float | None = Field(
|
||||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||||
)
|
)
|
||||||
|
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
|
||||||
|
None, description="Claude Code 高级配置"
|
||||||
|
)
|
||||||
|
pool_advanced: PoolAdvancedConfig | None = Field(None, description="通用号池配置")
|
||||||
|
|
||||||
|
|
||||||
class ProviderWithEndpointsSummary(BaseModel):
|
class ProviderWithEndpointsSummary(BaseModel):
|
||||||
@@ -974,6 +978,10 @@ class ProviderWithEndpointsSummary(BaseModel):
|
|||||||
default=None, description="流式请求首字节超时(秒)"
|
default=None, description="流式请求首字节超时(秒)"
|
||||||
)
|
)
|
||||||
request_timeout: float | None = Field(default=None, description="非流式请求整体超时(秒)")
|
request_timeout: float | None = Field(default=None, description="非流式请求整体超时(秒)")
|
||||||
|
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
|
||||||
|
default=None, description="Claude Code 高级配置"
|
||||||
|
)
|
||||||
|
pool_advanced: PoolAdvancedConfig | None = Field(default=None, description="通用号池配置")
|
||||||
|
|
||||||
# Endpoint 统计
|
# Endpoint 统计
|
||||||
total_endpoints: int = Field(default=0, description="总 Endpoint 数量")
|
total_endpoints: int = Field(default=0, description="总 Endpoint 数量")
|
||||||
|
|||||||
@@ -513,6 +513,10 @@ class FailoverEngine:
|
|||||||
api_key_id: str | None,
|
api_key_id: str | None,
|
||||||
) -> str:
|
) -> str:
|
||||||
# Create "available" record, then caller will mark pending.
|
# Create "available" record, then caller will mark pending.
|
||||||
|
extra: dict = {}
|
||||||
|
pool_extra = getattr(candidate, "_pool_extra_data", None)
|
||||||
|
if pool_extra:
|
||||||
|
extra.update(pool_extra)
|
||||||
row = RequestCandidateService.create_candidate(
|
row = RequestCandidateService.create_candidate(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
@@ -525,7 +529,7 @@ class FailoverEngine:
|
|||||||
key_id=str(candidate.key.id),
|
key_id=str(candidate.key.id),
|
||||||
status="available",
|
status="available",
|
||||||
is_cached=bool(getattr(candidate, "is_cached", False)),
|
is_cached=bool(getattr(candidate, "is_cached", False)),
|
||||||
extra_data={},
|
extra_data=extra,
|
||||||
)
|
)
|
||||||
return str(row.id)
|
return str(row.id)
|
||||||
|
|
||||||
@@ -539,6 +543,10 @@ class FailoverEngine:
|
|||||||
api_key_id: str | None,
|
api_key_id: str | None,
|
||||||
skip_reason: str | None,
|
skip_reason: str | None,
|
||||||
) -> str:
|
) -> str:
|
||||||
|
extra: dict = {}
|
||||||
|
pool_extra = getattr(candidate, "_pool_extra_data", None)
|
||||||
|
if pool_extra:
|
||||||
|
extra.update(pool_extra)
|
||||||
row = RequestCandidateService.create_candidate(
|
row = RequestCandidateService.create_candidate(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
@@ -552,7 +560,7 @@ class FailoverEngine:
|
|||||||
status="skipped",
|
status="skipped",
|
||||||
skip_reason=skip_reason,
|
skip_reason=skip_reason,
|
||||||
is_cached=bool(getattr(candidate, "is_cached", False)),
|
is_cached=bool(getattr(candidate, "is_cached", False)),
|
||||||
extra_data={},
|
extra_data=extra,
|
||||||
)
|
)
|
||||||
# ensure visible for subsequent recorder reads
|
# ensure visible for subsequent recorder reads
|
||||||
if self.db.in_transaction():
|
if self.db.in_transaction():
|
||||||
|
|||||||
@@ -52,6 +52,7 @@ class CandidateResolver:
|
|||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: dict[str, bool] | None = None,
|
capability_requirements: dict[str, bool] | None = None,
|
||||||
preferred_key_ids: list[str] | None = None,
|
preferred_key_ids: list[str] | None = None,
|
||||||
|
request_body: dict | None = None,
|
||||||
) -> tuple[list[ProviderCandidate], str]:
|
) -> tuple[list[ProviderCandidate], str]:
|
||||||
"""
|
"""
|
||||||
获取所有可用候选
|
获取所有可用候选
|
||||||
@@ -96,6 +97,7 @@ class CandidateResolver:
|
|||||||
provider_limit=provider_batch_size,
|
provider_limit=provider_batch_size,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
capability_requirements=capability_requirements,
|
capability_requirements=capability_requirements,
|
||||||
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
3
src/services/provider/adapters/claude_code/__init__.py
Normal file
3
src/services/provider/adapters/claude_code/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
"""Claude Code provider adapter."""
|
||||||
|
|
||||||
|
__all__ = []
|
||||||
52
src/services/provider/adapters/claude_code/constants.py
Normal file
52
src/services/provider/adapters/claude_code/constants.py
Normal file
@@ -0,0 +1,52 @@
|
|||||||
|
"""Claude Code adapter constants."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
CLAUDE_MESSAGES_PATH = "/v1/messages"
|
||||||
|
DEFAULT_ANTHROPIC_VERSION = "2023-06-01"
|
||||||
|
DEFAULT_ACCEPT = "application/json"
|
||||||
|
STREAM_HELPER_METHOD = "stream"
|
||||||
|
SESSION_ID_MASKING_TTL_SECONDS = 15 * 60
|
||||||
|
# 仅代表“启用 Claude Code TLS 配置”的客户端 profile 标识(best-effort)。
|
||||||
|
TLS_PROFILE_CLAUDE_CODE = "claude_code_nodejs"
|
||||||
|
|
||||||
|
# Claude Code OAuth required betas.
|
||||||
|
BETA_CLAUDE_CODE = "claude-code-20250219"
|
||||||
|
BETA_OAUTH = "oauth-2025-04-20"
|
||||||
|
BETA_INTERLEAVED_THINKING = "interleaved-thinking-2025-05-14"
|
||||||
|
BETA_CONTEXT_1M = "context-1m-2025-08-07"
|
||||||
|
|
||||||
|
CLAUDE_CODE_REQUIRED_BETA_TOKENS: tuple[str, ...] = (
|
||||||
|
BETA_CLAUDE_CODE,
|
||||||
|
BETA_OAUTH,
|
||||||
|
BETA_INTERLEAVED_THINKING,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Mimic headers observed from Claude Code traffic.
|
||||||
|
CLAUDE_CODE_DEFAULT_HEADERS: dict[str, str] = {
|
||||||
|
"X-Stainless-Lang": "js",
|
||||||
|
"X-Stainless-Package-Version": "0.70.0",
|
||||||
|
"X-Stainless-OS": "Linux",
|
||||||
|
"X-Stainless-Arch": "arm64",
|
||||||
|
"X-Stainless-Runtime": "node",
|
||||||
|
"X-Stainless-Runtime-Version": "v24.13.0",
|
||||||
|
"X-Stainless-Retry-Count": "0",
|
||||||
|
"X-Stainless-Timeout": "600",
|
||||||
|
"X-App": "cli",
|
||||||
|
"Anthropic-Dangerous-Direct-Browser-Access": "true",
|
||||||
|
}
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"BETA_CLAUDE_CODE",
|
||||||
|
"BETA_CONTEXT_1M",
|
||||||
|
"BETA_INTERLEAVED_THINKING",
|
||||||
|
"BETA_OAUTH",
|
||||||
|
"CLAUDE_CODE_DEFAULT_HEADERS",
|
||||||
|
"CLAUDE_CODE_REQUIRED_BETA_TOKENS",
|
||||||
|
"CLAUDE_MESSAGES_PATH",
|
||||||
|
"DEFAULT_ACCEPT",
|
||||||
|
"DEFAULT_ANTHROPIC_VERSION",
|
||||||
|
"SESSION_ID_MASKING_TTL_SECONDS",
|
||||||
|
"STREAM_HELPER_METHOD",
|
||||||
|
"TLS_PROFILE_CLAUDE_CODE",
|
||||||
|
]
|
||||||
141
src/services/provider/adapters/claude_code/context.py
Normal file
141
src/services/provider/adapters/claude_code/context.py
Normal file
@@ -0,0 +1,141 @@
|
|||||||
|
"""Claude Code request context using contextvars."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.models.admin_requests import ClaudeCodeAdvancedConfig
|
||||||
|
from src.services.provider.adapters.claude_code.constants import TLS_PROFILE_CLAUDE_CODE
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from src.services.provider.pool.config import PoolConfig
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ClaudeCodeRequestContext:
|
||||||
|
is_stream: bool = False
|
||||||
|
# 使用 Key 级别作用域,确保会话限制按 OAuth 账号隔离。
|
||||||
|
scope_key: str | None = None
|
||||||
|
key_id: str | None = None
|
||||||
|
max_sessions: int | None = None
|
||||||
|
session_idle_timeout_minutes: int = 5
|
||||||
|
enable_tls_fingerprint: bool = False
|
||||||
|
session_id_masking_enabled: bool = False
|
||||||
|
# Account Pool fields
|
||||||
|
provider_id: str | None = None
|
||||||
|
pool_config: PoolConfig | None = None
|
||||||
|
session_uuid: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
_claude_code_request_context: contextvars.ContextVar[ClaudeCodeRequestContext | None] = (
|
||||||
|
contextvars.ContextVar(
|
||||||
|
"claude_code_request_context",
|
||||||
|
default=None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def set_claude_code_request_context(ctx: ClaudeCodeRequestContext | None) -> None:
|
||||||
|
_claude_code_request_context.set(ctx)
|
||||||
|
|
||||||
|
|
||||||
|
def get_claude_code_request_context() -> ClaudeCodeRequestContext | None:
|
||||||
|
return _claude_code_request_context.get()
|
||||||
|
|
||||||
|
|
||||||
|
def build_claude_code_request_context(
|
||||||
|
*,
|
||||||
|
provider_config: Any,
|
||||||
|
key_id: str | None,
|
||||||
|
is_stream: bool,
|
||||||
|
provider_id: str | None = None,
|
||||||
|
) -> ClaudeCodeRequestContext:
|
||||||
|
"""根据 Provider.config 构建 Claude Code 请求上下文。"""
|
||||||
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
|
||||||
|
normalized_key_id = str(key_id or "").strip() or None
|
||||||
|
|
||||||
|
advanced_config: ClaudeCodeAdvancedConfig | None = None
|
||||||
|
provider_config_dict = provider_config if isinstance(provider_config, dict) else {}
|
||||||
|
raw_advanced = provider_config_dict.get("claude_code_advanced")
|
||||||
|
|
||||||
|
if raw_advanced is not None:
|
||||||
|
try:
|
||||||
|
if isinstance(raw_advanced, ClaudeCodeAdvancedConfig):
|
||||||
|
advanced_config = raw_advanced
|
||||||
|
elif isinstance(raw_advanced, dict):
|
||||||
|
advanced_config = ClaudeCodeAdvancedConfig.model_validate(raw_advanced)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Claude Code advanced config 类型无效: {},已忽略",
|
||||||
|
type(raw_advanced).__name__,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Claude Code advanced config 解析失败,已忽略: {}", str(exc))
|
||||||
|
|
||||||
|
max_sessions = advanced_config.max_sessions if advanced_config else None
|
||||||
|
idle_timeout_minutes = (
|
||||||
|
advanced_config.session_idle_timeout_minutes
|
||||||
|
if advanced_config and advanced_config.session_idle_timeout_minutes is not None
|
||||||
|
else 5
|
||||||
|
)
|
||||||
|
enable_tls_fingerprint = (
|
||||||
|
bool(advanced_config.enable_tls_fingerprint) if advanced_config else False
|
||||||
|
)
|
||||||
|
session_id_masking_enabled = (
|
||||||
|
bool(advanced_config.session_id_masking_enabled) if advanced_config else False
|
||||||
|
)
|
||||||
|
|
||||||
|
# Parse pool config (None = non-pool provider, keep as None for semantic consistency)
|
||||||
|
pool_cfg = parse_pool_config(provider_config_dict)
|
||||||
|
|
||||||
|
return ClaudeCodeRequestContext(
|
||||||
|
is_stream=bool(is_stream),
|
||||||
|
scope_key=f"key:{normalized_key_id}" if normalized_key_id else None,
|
||||||
|
key_id=normalized_key_id,
|
||||||
|
max_sessions=max_sessions,
|
||||||
|
session_idle_timeout_minutes=idle_timeout_minutes,
|
||||||
|
enable_tls_fingerprint=enable_tls_fingerprint,
|
||||||
|
session_id_masking_enabled=session_id_masking_enabled,
|
||||||
|
provider_id=str(provider_id or "").strip() or None,
|
||||||
|
pool_config=pool_cfg,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_claude_code_tls_profile(
|
||||||
|
ctx: ClaudeCodeRequestContext | None,
|
||||||
|
) -> str | None:
|
||||||
|
if ctx and ctx.enable_tls_fingerprint:
|
||||||
|
return TLS_PROFILE_CLAUDE_CODE
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def build_and_set_claude_code_request_context(
|
||||||
|
*,
|
||||||
|
provider_config: Any,
|
||||||
|
key_id: str | None,
|
||||||
|
is_stream: bool,
|
||||||
|
provider_id: str | None = None,
|
||||||
|
) -> tuple[ClaudeCodeRequestContext, str | None]:
|
||||||
|
"""构建并写入 Claude Code 上下文,同时返回对应 TLS profile。"""
|
||||||
|
ctx = build_claude_code_request_context(
|
||||||
|
provider_config=provider_config,
|
||||||
|
key_id=key_id,
|
||||||
|
is_stream=is_stream,
|
||||||
|
provider_id=provider_id,
|
||||||
|
)
|
||||||
|
set_claude_code_request_context(ctx)
|
||||||
|
return ctx, resolve_claude_code_tls_profile(ctx)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"build_and_set_claude_code_request_context",
|
||||||
|
"build_claude_code_request_context",
|
||||||
|
"ClaudeCodeRequestContext",
|
||||||
|
"get_claude_code_request_context",
|
||||||
|
"resolve_claude_code_tls_profile",
|
||||||
|
"set_claude_code_request_context",
|
||||||
|
]
|
||||||
526
src/services/provider/adapters/claude_code/envelope.py
Normal file
526
src/services/provider/adapters/claude_code/envelope.py
Normal file
@@ -0,0 +1,526 @@
|
|||||||
|
"""Claude Code upstream envelope hooks."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from dataclasses import replace
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from src.clients.redis_client import get_redis_client, get_redis_client_sync
|
||||||
|
from src.config.settings import config
|
||||||
|
from src.core.exceptions import ConcurrencyLimitError
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.services.provider.adapters.claude_code.constants import (
|
||||||
|
BETA_CONTEXT_1M,
|
||||||
|
CLAUDE_CODE_DEFAULT_HEADERS,
|
||||||
|
CLAUDE_CODE_REQUIRED_BETA_TOKENS,
|
||||||
|
DEFAULT_ACCEPT,
|
||||||
|
DEFAULT_ANTHROPIC_VERSION,
|
||||||
|
SESSION_ID_MASKING_TTL_SECONDS,
|
||||||
|
STREAM_HELPER_METHOD,
|
||||||
|
)
|
||||||
|
from src.services.provider.adapters.claude_code.context import (
|
||||||
|
ClaudeCodeRequestContext,
|
||||||
|
get_claude_code_request_context,
|
||||||
|
set_claude_code_request_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
_SESSION_MARKER = "_session_"
|
||||||
|
_DUMMY_THINKING_SIGNATURE = "skip_thought_signature_validator"
|
||||||
|
_session_runtime_lock = threading.Lock()
|
||||||
|
# key: scope_key -> {session_id -> last_seen_monotonic}
|
||||||
|
_active_sessions: dict[str, dict[str, float]] = {}
|
||||||
|
# key: scope_key -> (masked_session_uuid, expire_at_monotonic)
|
||||||
|
_masked_sessions: dict[str, tuple[str, float]] = {}
|
||||||
|
_REDIS_SESSION_KEY_PREFIX = "claude_code:sessions"
|
||||||
|
_REDIS_SESSION_RESERVE_LUA = """
|
||||||
|
local key = KEYS[1]
|
||||||
|
local sid = ARGV[1]
|
||||||
|
local now = tonumber(ARGV[2])
|
||||||
|
local expire_before = tonumber(ARGV[3])
|
||||||
|
local max_sessions = tonumber(ARGV[4])
|
||||||
|
local ttl_seconds = tonumber(ARGV[5])
|
||||||
|
|
||||||
|
redis.call("ZREMRANGEBYSCORE", key, "-inf", expire_before)
|
||||||
|
|
||||||
|
local exists = redis.call("ZSCORE", key, sid)
|
||||||
|
if exists then
|
||||||
|
redis.call("ZADD", key, now, sid)
|
||||||
|
redis.call("EXPIRE", key, ttl_seconds)
|
||||||
|
return {1, redis.call("ZCARD", key)}
|
||||||
|
end
|
||||||
|
|
||||||
|
local active = redis.call("ZCARD", key)
|
||||||
|
if active >= max_sessions then
|
||||||
|
return {0, active}
|
||||||
|
end
|
||||||
|
|
||||||
|
redis.call("ZADD", key, now, sid)
|
||||||
|
redis.call("EXPIRE", key, ttl_seconds)
|
||||||
|
return {1, active + 1}
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def merge_anthropic_beta_tokens(
|
||||||
|
incoming: str | None,
|
||||||
|
*,
|
||||||
|
required: tuple[str, ...] = CLAUDE_CODE_REQUIRED_BETA_TOKENS,
|
||||||
|
) -> str:
|
||||||
|
"""Merge required beta tokens and incoming anthropic-beta with deduplication."""
|
||||||
|
seen: set[str] = set()
|
||||||
|
merged: list[str] = []
|
||||||
|
|
||||||
|
def _append(token: str) -> None:
|
||||||
|
token = token.strip()
|
||||||
|
if not token or token in seen:
|
||||||
|
return
|
||||||
|
seen.add(token)
|
||||||
|
merged.append(token)
|
||||||
|
|
||||||
|
for token in required:
|
||||||
|
_append(token)
|
||||||
|
for token in str(incoming or "").split(","):
|
||||||
|
_append(token)
|
||||||
|
|
||||||
|
return ",".join(merged)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_stream_flag(raw_stream: Any) -> bool:
|
||||||
|
if isinstance(raw_stream, bool):
|
||||||
|
return raw_stream
|
||||||
|
return str(raw_stream).strip().lower() in {"1", "true", "yes", "on"}
|
||||||
|
|
||||||
|
|
||||||
|
def _get_metadata_user_id(request_body: dict[str, Any]) -> str | None:
|
||||||
|
metadata = request_body.get("metadata")
|
||||||
|
if not isinstance(metadata, dict):
|
||||||
|
return None
|
||||||
|
user_id = metadata.get("user_id")
|
||||||
|
if not isinstance(user_id, str):
|
||||||
|
return None
|
||||||
|
text = user_id.strip()
|
||||||
|
return text or None
|
||||||
|
|
||||||
|
|
||||||
|
def _set_metadata_user_id(request_body: dict[str, Any], user_id: str) -> None:
|
||||||
|
metadata = request_body.get("metadata")
|
||||||
|
if not isinstance(metadata, dict):
|
||||||
|
metadata = {}
|
||||||
|
request_body["metadata"] = metadata
|
||||||
|
metadata["user_id"] = user_id
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_session_id(user_id: str) -> str | None:
|
||||||
|
idx = user_id.rfind(_SESSION_MARKER)
|
||||||
|
if idx == -1:
|
||||||
|
return None
|
||||||
|
session_id = user_id[idx + len(_SESSION_MARKER) :].strip()
|
||||||
|
return session_id or None
|
||||||
|
|
||||||
|
|
||||||
|
def _is_thinking_enabled(request_body: dict[str, Any]) -> bool:
|
||||||
|
thinking = request_body.get("thinking")
|
||||||
|
if not isinstance(thinking, dict):
|
||||||
|
return False
|
||||||
|
thinking_type = str(thinking.get("type") or "").strip().lower()
|
||||||
|
return thinking_type in {"enabled", "adaptive"}
|
||||||
|
|
||||||
|
|
||||||
|
def _sanitize_thinking_blocks(request_body: dict[str, Any]) -> None:
|
||||||
|
"""过滤可能导致 Claude Code 400 的无效 thinking 块。"""
|
||||||
|
messages = request_body.get("messages")
|
||||||
|
if not isinstance(messages, list) or not messages:
|
||||||
|
return
|
||||||
|
|
||||||
|
thinking_enabled = _is_thinking_enabled(request_body)
|
||||||
|
filtered_messages = 0
|
||||||
|
filtered_blocks = 0
|
||||||
|
|
||||||
|
for message in messages:
|
||||||
|
if not isinstance(message, dict):
|
||||||
|
continue
|
||||||
|
|
||||||
|
role = str(message.get("role") or "")
|
||||||
|
content = message.get("content")
|
||||||
|
if not isinstance(content, list):
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_content: list[Any] = []
|
||||||
|
modified = False
|
||||||
|
for block in content:
|
||||||
|
if not isinstance(block, dict):
|
||||||
|
new_content.append(block)
|
||||||
|
continue
|
||||||
|
|
||||||
|
block_type = str(block.get("type") or "")
|
||||||
|
if block_type in {"thinking", "redacted_thinking"}:
|
||||||
|
keep = False
|
||||||
|
# 仅保留 assistant 且带真实 signature 的 thinking 块。
|
||||||
|
if thinking_enabled and role == "assistant":
|
||||||
|
signature = str(block.get("signature") or "").strip()
|
||||||
|
keep = bool(signature and signature != _DUMMY_THINKING_SIGNATURE)
|
||||||
|
if keep:
|
||||||
|
new_content.append(block)
|
||||||
|
else:
|
||||||
|
modified = True
|
||||||
|
filtered_blocks += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 兼容无 type 但带 thinking 字段的历史块,直接移除。
|
||||||
|
if not block_type and "thinking" in block:
|
||||||
|
modified = True
|
||||||
|
filtered_blocks += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_content.append(block)
|
||||||
|
|
||||||
|
if modified:
|
||||||
|
message["content"] = new_content
|
||||||
|
filtered_messages += 1
|
||||||
|
|
||||||
|
if filtered_blocks:
|
||||||
|
logger.info(
|
||||||
|
"Claude Code thinking 预过滤: messages={}, blocks={}, thinking_enabled={}",
|
||||||
|
filtered_messages,
|
||||||
|
filtered_blocks,
|
||||||
|
thinking_enabled,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_or_create_masked_session(scope_key: str) -> str:
|
||||||
|
now = time.monotonic()
|
||||||
|
with _session_runtime_lock:
|
||||||
|
existing = _masked_sessions.get(scope_key)
|
||||||
|
if existing and existing[1] > now:
|
||||||
|
masked_session_id = existing[0]
|
||||||
|
else:
|
||||||
|
masked_session_id = str(uuid.uuid4())
|
||||||
|
_masked_sessions[scope_key] = (
|
||||||
|
masked_session_id,
|
||||||
|
now + SESSION_ID_MASKING_TTL_SECONDS,
|
||||||
|
)
|
||||||
|
return masked_session_id
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_session_id_masking(request_body: dict[str, Any], *, scope_key: str) -> None:
|
||||||
|
user_id = _get_metadata_user_id(request_body)
|
||||||
|
if not user_id:
|
||||||
|
return
|
||||||
|
idx = user_id.rfind(_SESSION_MARKER)
|
||||||
|
if idx == -1:
|
||||||
|
return
|
||||||
|
masked_session_id = _get_or_create_masked_session(scope_key)
|
||||||
|
_set_metadata_user_id(
|
||||||
|
request_body,
|
||||||
|
user_id[: idx + len(_SESSION_MARKER)] + masked_session_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _register_or_reject_session(
|
||||||
|
*,
|
||||||
|
scope_key: str,
|
||||||
|
session_id: str,
|
||||||
|
max_sessions: int,
|
||||||
|
idle_timeout_minutes: int,
|
||||||
|
) -> tuple[bool, int]:
|
||||||
|
now = time.monotonic()
|
||||||
|
idle_seconds = max(60, int(idle_timeout_minutes * 60))
|
||||||
|
|
||||||
|
with _session_runtime_lock:
|
||||||
|
bucket = _active_sessions.setdefault(scope_key, {})
|
||||||
|
|
||||||
|
# 先清理过期会话,避免误判占用。
|
||||||
|
expired = [sid for sid, last_seen in bucket.items() if now - last_seen > idle_seconds]
|
||||||
|
for sid in expired:
|
||||||
|
bucket.pop(sid, None)
|
||||||
|
|
||||||
|
if session_id in bucket:
|
||||||
|
bucket[session_id] = now
|
||||||
|
return True, len(bucket)
|
||||||
|
|
||||||
|
if len(bucket) >= max_sessions:
|
||||||
|
return False, len(bucket)
|
||||||
|
|
||||||
|
bucket[session_id] = now
|
||||||
|
return True, len(bucket)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_session_limit_error(
|
||||||
|
*,
|
||||||
|
max_sessions: int,
|
||||||
|
active_count: int,
|
||||||
|
key_id: str | None,
|
||||||
|
) -> ConcurrencyLimitError:
|
||||||
|
return ConcurrencyLimitError(
|
||||||
|
message=(f"Claude Code 活跃会话数已达上限({max_sessions})。当前活跃会话: {active_count}"),
|
||||||
|
key_id=key_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _redis_session_key(scope_key: str) -> str:
|
||||||
|
return f"{_REDIS_SESSION_KEY_PREFIX}:{scope_key}"
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_redis_session_result(raw: Any) -> tuple[bool, int] | None:
|
||||||
|
if not isinstance(raw, (list, tuple)) or len(raw) < 2:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
allowed = int(raw[0]) == 1
|
||||||
|
active_count = int(raw[1])
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
return allowed, active_count
|
||||||
|
|
||||||
|
|
||||||
|
def _enforce_session_controls(
|
||||||
|
request_body: dict[str, Any],
|
||||||
|
ctx: ClaudeCodeRequestContext,
|
||||||
|
*,
|
||||||
|
enforce_max_sessions: bool = True,
|
||||||
|
) -> None:
|
||||||
|
"""同步执行会话限制 + masking。
|
||||||
|
|
||||||
|
当 ``enforce_max_sessions=False``(将由 ``enforce_distributed_session_controls``
|
||||||
|
异步接管)时仅做 masking;masking 始终在会话限制检查之后执行,避免
|
||||||
|
用被伪装后的 session_id 做计数。
|
||||||
|
"""
|
||||||
|
scope_key = str(ctx.scope_key or "").strip()
|
||||||
|
if not scope_key:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 先基于真实 session_id 做会话限制检查。
|
||||||
|
if enforce_max_sessions and ctx.max_sessions and ctx.max_sessions > 0:
|
||||||
|
user_id = _get_metadata_user_id(request_body)
|
||||||
|
if user_id:
|
||||||
|
session_id = _extract_session_id(user_id) or user_id
|
||||||
|
allowed, active_count = _register_or_reject_session(
|
||||||
|
scope_key=scope_key,
|
||||||
|
session_id=session_id,
|
||||||
|
max_sessions=ctx.max_sessions,
|
||||||
|
idle_timeout_minutes=ctx.session_idle_timeout_minutes,
|
||||||
|
)
|
||||||
|
if not allowed:
|
||||||
|
raise _build_session_limit_error(
|
||||||
|
max_sessions=ctx.max_sessions,
|
||||||
|
active_count=active_count,
|
||||||
|
key_id=ctx.key_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not enforce_max_sessions:
|
||||||
|
# 分布式模式下 masking 延迟到 enforce_distributed_session_controls 中执行,
|
||||||
|
# 避免 wrap_request 提前改写 user_id 导致分布式检查拿到伪装后的 session_id。
|
||||||
|
return
|
||||||
|
|
||||||
|
# 仅在本地模式下立即 masking。
|
||||||
|
if ctx.session_id_masking_enabled:
|
||||||
|
_apply_session_id_masking(request_body, scope_key=scope_key)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_distributed_session_control_available() -> bool:
|
||||||
|
try:
|
||||||
|
return get_redis_client_sync() is not None
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
async def enforce_distributed_session_controls(
|
||||||
|
request_body: dict[str, Any],
|
||||||
|
ctx: ClaudeCodeRequestContext | None,
|
||||||
|
) -> None:
|
||||||
|
"""异步执行会话限制 + masking。
|
||||||
|
|
||||||
|
优先使用 Redis(多实例共享);Redis 不可用时回退到进程内计数。
|
||||||
|
masking 在会话限制检查通过后执行,确保计数使用真实 session_id。
|
||||||
|
"""
|
||||||
|
if ctx is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
scope_key = str(ctx.scope_key or "").strip()
|
||||||
|
if not scope_key:
|
||||||
|
# 即使无 scope_key 也无法做 masking(需要 scope_key 作为 key),直接返回。
|
||||||
|
return
|
||||||
|
|
||||||
|
if not ctx.max_sessions or ctx.max_sessions <= 0:
|
||||||
|
# 无会话限制,仅做 masking。
|
||||||
|
if ctx.session_id_masking_enabled:
|
||||||
|
_apply_session_id_masking(request_body, scope_key=scope_key)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 基于真实 user_id 提取 session_id 做限制检查。
|
||||||
|
user_id = _get_metadata_user_id(request_body)
|
||||||
|
if not user_id:
|
||||||
|
return
|
||||||
|
session_id = _extract_session_id(user_id) or user_id
|
||||||
|
|
||||||
|
idle_seconds = max(60, int(ctx.session_idle_timeout_minutes * 60))
|
||||||
|
redis_ttl = idle_seconds + 300
|
||||||
|
now = int(time.time())
|
||||||
|
expire_before = now - idle_seconds
|
||||||
|
|
||||||
|
redis_client = await get_redis_client(require_redis=False)
|
||||||
|
if redis_client is not None:
|
||||||
|
try:
|
||||||
|
raw_result = await redis_client.eval(
|
||||||
|
_REDIS_SESSION_RESERVE_LUA,
|
||||||
|
1,
|
||||||
|
_redis_session_key(scope_key),
|
||||||
|
session_id,
|
||||||
|
str(now),
|
||||||
|
str(expire_before),
|
||||||
|
str(ctx.max_sessions),
|
||||||
|
str(redis_ttl),
|
||||||
|
)
|
||||||
|
parsed = _parse_redis_session_result(raw_result)
|
||||||
|
if parsed is None:
|
||||||
|
raise ValueError(f"invalid redis eval result: {raw_result!r}")
|
||||||
|
|
||||||
|
allowed, active_count = parsed
|
||||||
|
if not allowed:
|
||||||
|
raise _build_session_limit_error(
|
||||||
|
max_sessions=ctx.max_sessions,
|
||||||
|
active_count=active_count,
|
||||||
|
key_id=ctx.key_id,
|
||||||
|
)
|
||||||
|
# 会话限制通过后再做 masking。
|
||||||
|
if ctx.session_id_masking_enabled:
|
||||||
|
_apply_session_id_masking(request_body, scope_key=scope_key)
|
||||||
|
return
|
||||||
|
except ConcurrencyLimitError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Claude Code 分布式会话控制失败,回退本地计数: {}", str(exc))
|
||||||
|
|
||||||
|
allowed, active_count = _register_or_reject_session(
|
||||||
|
scope_key=scope_key,
|
||||||
|
session_id=session_id,
|
||||||
|
max_sessions=ctx.max_sessions,
|
||||||
|
idle_timeout_minutes=ctx.session_idle_timeout_minutes,
|
||||||
|
)
|
||||||
|
if not allowed:
|
||||||
|
raise _build_session_limit_error(
|
||||||
|
max_sessions=ctx.max_sessions,
|
||||||
|
active_count=active_count,
|
||||||
|
key_id=ctx.key_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 会话限制通过后再做 masking。
|
||||||
|
if ctx.session_id_masking_enabled:
|
||||||
|
_apply_session_id_masking(request_body, scope_key=scope_key)
|
||||||
|
|
||||||
|
|
||||||
|
class ClaudeCodeEnvelope:
|
||||||
|
"""Provider envelope hooks for Claude Code OAuth upstream."""
|
||||||
|
|
||||||
|
name = "claude:cli"
|
||||||
|
|
||||||
|
def extra_headers(self) -> dict[str, str] | None:
|
||||||
|
ctx = get_claude_code_request_context()
|
||||||
|
is_stream = bool(ctx.is_stream) if ctx else False
|
||||||
|
|
||||||
|
headers = dict(CLAUDE_CODE_DEFAULT_HEADERS)
|
||||||
|
headers["Accept"] = DEFAULT_ACCEPT
|
||||||
|
headers["anthropic-version"] = DEFAULT_ANTHROPIC_VERSION
|
||||||
|
headers["anthropic-beta"] = merge_anthropic_beta_tokens(None)
|
||||||
|
if is_stream:
|
||||||
|
headers["x-stainless-helper-method"] = STREAM_HELPER_METHOD
|
||||||
|
|
||||||
|
ua = str(getattr(config, "internal_user_agent_claude_cli", "") or "").strip()
|
||||||
|
if ua:
|
||||||
|
headers["User-Agent"] = ua
|
||||||
|
|
||||||
|
return headers
|
||||||
|
|
||||||
|
def wrap_request(
|
||||||
|
self,
|
||||||
|
request_body: dict[str, Any],
|
||||||
|
*,
|
||||||
|
model: str, # noqa: ARG002
|
||||||
|
url_model: str | None,
|
||||||
|
decrypted_auth_config: dict[str, Any] | None, # noqa: ARG002
|
||||||
|
) -> tuple[dict[str, Any], str | None]:
|
||||||
|
raw_stream = request_body.get("stream", False)
|
||||||
|
is_stream = _parse_stream_flag(raw_stream)
|
||||||
|
|
||||||
|
ctx = get_claude_code_request_context()
|
||||||
|
if ctx is None:
|
||||||
|
ctx = ClaudeCodeRequestContext()
|
||||||
|
# Extract session_uuid from metadata.user_id for pool sticky session.
|
||||||
|
session_uuid: str | None = None
|
||||||
|
user_id = _get_metadata_user_id(request_body)
|
||||||
|
if user_id:
|
||||||
|
session_uuid = _extract_session_id(user_id)
|
||||||
|
ctx = replace(ctx, is_stream=is_stream, session_uuid=session_uuid)
|
||||||
|
set_claude_code_request_context(ctx)
|
||||||
|
|
||||||
|
_sanitize_thinking_blocks(request_body)
|
||||||
|
|
||||||
|
_enforce_session_controls(
|
||||||
|
request_body,
|
||||||
|
ctx,
|
||||||
|
enforce_max_sessions=not _is_distributed_session_control_available(),
|
||||||
|
)
|
||||||
|
return request_body, url_model
|
||||||
|
|
||||||
|
def unwrap_response(self, data: Any) -> Any:
|
||||||
|
return data
|
||||||
|
|
||||||
|
def postprocess_unwrapped_response(self, *, model: str, data: Any) -> None: # noqa: ARG002
|
||||||
|
return
|
||||||
|
|
||||||
|
def capture_selected_base_url(self) -> str | None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def on_http_status(self, *, base_url: str | None, status_code: int) -> None: # noqa: ARG002
|
||||||
|
return
|
||||||
|
|
||||||
|
def on_connection_error(self, *, base_url: str | None, exc: Exception) -> None: # noqa: ARG002
|
||||||
|
return
|
||||||
|
|
||||||
|
def force_stream_rewrite(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Optional lifecycle hooks
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def prepare_context(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
provider_config: Any,
|
||||||
|
key_id: str,
|
||||||
|
is_stream: bool,
|
||||||
|
provider_id: str | None = None,
|
||||||
|
) -> str | None:
|
||||||
|
from src.services.provider.adapters.claude_code.context import (
|
||||||
|
build_and_set_claude_code_request_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
_ctx, tls_profile = build_and_set_claude_code_request_context(
|
||||||
|
provider_config=provider_config,
|
||||||
|
key_id=key_id,
|
||||||
|
is_stream=is_stream,
|
||||||
|
provider_id=provider_id,
|
||||||
|
)
|
||||||
|
return tls_profile
|
||||||
|
|
||||||
|
async def post_wrap_request(self, request_body: dict[str, Any]) -> None:
|
||||||
|
await enforce_distributed_session_controls(
|
||||||
|
request_body,
|
||||||
|
get_claude_code_request_context(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def excluded_beta_tokens(self) -> frozenset[str]:
|
||||||
|
return frozenset({BETA_CONTEXT_1M})
|
||||||
|
|
||||||
|
|
||||||
|
claude_code_envelope = ClaudeCodeEnvelope()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"ClaudeCodeEnvelope",
|
||||||
|
"claude_code_envelope",
|
||||||
|
"enforce_distributed_session_controls",
|
||||||
|
"merge_anthropic_beta_tokens",
|
||||||
|
]
|
||||||
62
src/services/provider/adapters/claude_code/plugin.py
Normal file
62
src/services/provider/adapters/claude_code/plugin.py
Normal file
@@ -0,0 +1,62 @@
|
|||||||
|
"""Claude Code provider plugin."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
|
from src.services.provider.adapters.claude_code.constants import CLAUDE_MESSAGES_PATH
|
||||||
|
from src.services.provider.preset_models import create_preset_models_fetcher
|
||||||
|
|
||||||
|
fetch_models_claude_code = create_preset_models_fetcher("claude_code")
|
||||||
|
|
||||||
|
|
||||||
|
def build_claude_code_url(
|
||||||
|
endpoint: Any,
|
||||||
|
*,
|
||||||
|
is_stream: bool,
|
||||||
|
effective_query_params: dict[str, Any],
|
||||||
|
) -> str:
|
||||||
|
"""Build Claude Code upstream URL and avoid duplicate /v1/messages suffix."""
|
||||||
|
_ = is_stream
|
||||||
|
|
||||||
|
base = str(getattr(endpoint, "base_url", "") or "").rstrip("/")
|
||||||
|
if base.endswith(CLAUDE_MESSAGES_PATH) or base.endswith("/messages"):
|
||||||
|
url = base
|
||||||
|
elif base.endswith("/v1"):
|
||||||
|
url = f"{base}/messages"
|
||||||
|
else:
|
||||||
|
url = f"{base}{CLAUDE_MESSAGES_PATH}"
|
||||||
|
|
||||||
|
if effective_query_params:
|
||||||
|
query_string = urlencode(effective_query_params, doseq=True)
|
||||||
|
if query_string:
|
||||||
|
url = f"{url}?{query_string}"
|
||||||
|
|
||||||
|
return url
|
||||||
|
|
||||||
|
|
||||||
|
def register_all() -> None:
|
||||||
|
"""Register Claude Code hooks into shared registries."""
|
||||||
|
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
|
||||||
|
from src.services.provider.adapters.claude_code.envelope import claude_code_envelope
|
||||||
|
from src.services.provider.envelope import register_envelope
|
||||||
|
from src.services.provider.transport import register_transport_hook
|
||||||
|
|
||||||
|
register_envelope("claude_code", "claude:cli", claude_code_envelope)
|
||||||
|
register_envelope("claude_code", "", claude_code_envelope)
|
||||||
|
|
||||||
|
register_transport_hook("claude_code", "claude:cli", build_claude_code_url)
|
||||||
|
|
||||||
|
UpstreamModelsFetcherRegistry.register(
|
||||||
|
provider_types=["claude_code"],
|
||||||
|
fetcher=fetch_models_claude_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
from src.services.provider.adapters.claude_code.pool_hook import claude_code_pool_hook
|
||||||
|
from src.services.provider.pool.hooks import register_pool_hook
|
||||||
|
|
||||||
|
register_pool_hook("claude_code", claude_code_pool_hook)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["build_claude_code_url", "fetch_models_claude_code", "register_all"]
|
||||||
@@ -236,6 +236,7 @@ async def get_provider_auth(
|
|||||||
key: "ProviderAPIKey",
|
key: "ProviderAPIKey",
|
||||||
*,
|
*,
|
||||||
force_refresh: bool = False,
|
force_refresh: bool = False,
|
||||||
|
refresh_skew: int | None = None,
|
||||||
) -> ProviderAuthInfo | None:
|
) -> ProviderAuthInfo | None:
|
||||||
"""
|
"""
|
||||||
获取 Provider 的认证信息
|
获取 Provider 的认证信息
|
||||||
@@ -261,7 +262,12 @@ async def get_provider_auth(
|
|||||||
if auth_type == "oauth":
|
if auth_type == "oauth":
|
||||||
# OAuth token 保存在 key.api_key(加密),refresh_token/expires_at 等在 auth_config(加密 JSON)中。
|
# OAuth token 保存在 key.api_key(加密),refresh_token/expires_at 等在 auth_config(加密 JSON)中。
|
||||||
# 在请求前做一次懒刷新:接近过期时刷新 access_token,并用 Redis lock 避免并发风暴。
|
# 在请求前做一次懒刷新:接近过期时刷新 access_token,并用 Redis lock 避免并发风暴。
|
||||||
|
|
||||||
encrypted_auth_config = getattr(key, "auth_config", None)
|
encrypted_auth_config = getattr(key, "auth_config", None)
|
||||||
|
|
||||||
|
# 先解密 auth_config -- 下游 build_provider_url 等依赖 decrypted_auth_config
|
||||||
|
# 中的 provider_type / project_id / region 等元数据,即使 access_token 命中缓存
|
||||||
|
# 也不能跳过。
|
||||||
if encrypted_auth_config:
|
if encrypted_auth_config:
|
||||||
try:
|
try:
|
||||||
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
||||||
@@ -271,16 +277,51 @@ async def get_provider_auth(
|
|||||||
else:
|
else:
|
||||||
token_meta = {}
|
token_meta = {}
|
||||||
|
|
||||||
|
decrypted_auth_config: dict[str, Any] | None = (
|
||||||
|
token_meta if isinstance(token_meta, dict) and token_meta else None
|
||||||
|
)
|
||||||
|
|
||||||
|
# 快路径:查 Redis token 缓存,命中则跳过 refresh 和 api_key 解密。
|
||||||
|
# 注意:token_meta/decrypted_auth_config 已在上方解密,此处只是跳过后续刷新逻辑。
|
||||||
|
if not force_refresh and encrypted_auth_config:
|
||||||
|
try:
|
||||||
|
from src.services.provider.pool.oauth_cache import get_cached_token
|
||||||
|
|
||||||
|
_cached = await get_cached_token(str(key.id))
|
||||||
|
if _cached:
|
||||||
|
return ProviderAuthInfo(
|
||||||
|
auth_header="Authorization",
|
||||||
|
auth_value=f"Bearer {_cached}",
|
||||||
|
decrypted_auth_config=decrypted_auth_config,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.debug("OAuth token cache lookup failed for key {}", str(key.id)[:8])
|
||||||
|
|
||||||
expires_at = token_meta.get("expires_at")
|
expires_at = token_meta.get("expires_at")
|
||||||
refresh_token = token_meta.get("refresh_token")
|
refresh_token = token_meta.get("refresh_token")
|
||||||
provider_type = str(token_meta.get("provider_type") or "")
|
provider_type = str(token_meta.get("provider_type") or "")
|
||||||
cached_access_token = str(token_meta.get("access_token") or "").strip()
|
cached_access_token = str(token_meta.get("access_token") or "").strip()
|
||||||
|
|
||||||
# 120s skew (or force refresh when upstream returns 401)
|
# Refresh skew: providers with pool config use configurable
|
||||||
|
# proactive_refresh_seconds (default 180 s), others use 120 s.
|
||||||
|
# Prefer the caller-supplied value to avoid ORM lazy-load on key.provider.
|
||||||
|
_refresh_skew = refresh_skew if refresh_skew is not None else 120
|
||||||
|
if refresh_skew is None:
|
||||||
|
try:
|
||||||
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
|
||||||
|
provider_obj = getattr(key, "provider", None)
|
||||||
|
pcfg = getattr(provider_obj, "config", None) if provider_obj else None
|
||||||
|
pool_cfg = parse_pool_config(pcfg) if pcfg else None
|
||||||
|
if pool_cfg is not None:
|
||||||
|
_refresh_skew = pool_cfg.proactive_refresh_seconds
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
should_refresh = False
|
should_refresh = False
|
||||||
try:
|
try:
|
||||||
if expires_at is not None:
|
if expires_at is not None:
|
||||||
should_refresh = int(time.time()) >= int(expires_at) - 120
|
should_refresh = int(time.time()) >= int(expires_at) - _refresh_skew
|
||||||
except Exception:
|
except Exception:
|
||||||
should_refresh = False
|
should_refresh = False
|
||||||
|
|
||||||
@@ -294,6 +335,7 @@ async def get_provider_auth(
|
|||||||
elif crypto_service.decrypt(key.api_key) == "__placeholder__":
|
elif crypto_service.decrypt(key.api_key) == "__placeholder__":
|
||||||
should_refresh = True
|
should_refresh = True
|
||||||
|
|
||||||
|
_refreshed = False
|
||||||
if should_refresh and refresh_token and provider_type:
|
if should_refresh and refresh_token and provider_type:
|
||||||
try:
|
try:
|
||||||
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
||||||
@@ -313,6 +355,7 @@ async def get_provider_auth(
|
|||||||
token_meta = await _refresh_generic_oauth_token(
|
token_meta = await _refresh_generic_oauth_token(
|
||||||
key, endpoint, template, provider_type, refresh_token, token_meta
|
key, endpoint, template, provider_type, refresh_token, token_meta
|
||||||
)
|
)
|
||||||
|
_refreshed = True
|
||||||
finally:
|
finally:
|
||||||
if got_lock:
|
if got_lock:
|
||||||
await _release_refresh_lock(redis, key.id)
|
await _release_refresh_lock(redis, key.id)
|
||||||
@@ -328,7 +371,20 @@ async def get_provider_auth(
|
|||||||
else:
|
else:
|
||||||
effective_token = crypto_service.decrypt(key.api_key)
|
effective_token = crypto_service.decrypt(key.api_key)
|
||||||
|
|
||||||
decrypted_auth_config: dict[str, Any] | None = None
|
# 刷新成功后写入 Redis token 缓存(所有 OAuth key 均可受益)
|
||||||
|
if _refreshed and effective_token:
|
||||||
|
try:
|
||||||
|
from src.services.provider.pool.oauth_cache import cache_token
|
||||||
|
|
||||||
|
new_expires_at = token_meta.get("expires_at")
|
||||||
|
if new_expires_at is not None:
|
||||||
|
remaining = int(new_expires_at) - int(time.time())
|
||||||
|
if remaining > 0:
|
||||||
|
await cache_token(str(key.id), effective_token, remaining)
|
||||||
|
except Exception:
|
||||||
|
logger.debug("OAuth token cache write failed for key {}", str(key.id)[:8])
|
||||||
|
|
||||||
|
# 刷新可能更新了 token_meta,同步 decrypted_auth_config
|
||||||
if isinstance(token_meta, dict) and token_meta:
|
if isinstance(token_meta, dict) and token_meta:
|
||||||
decrypted_auth_config = token_meta
|
decrypted_auth_config = token_meta
|
||||||
|
|
||||||
|
|||||||
@@ -50,6 +50,39 @@ class ProviderEnvelope(Protocol):
|
|||||||
def force_stream_rewrite(self) -> bool:
|
def force_stream_rewrite(self) -> bool:
|
||||||
"""Whether streaming should always go through the rewrite/conversion path."""
|
"""Whether streaming should always go through the rewrite/conversion path."""
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Optional lifecycle hooks (checked via hasattr before calling)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def prepare_context(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
provider_config: Any,
|
||||||
|
key_id: str,
|
||||||
|
is_stream: bool,
|
||||||
|
provider_id: str | None = None,
|
||||||
|
) -> str | None:
|
||||||
|
"""Pre-wrap hook: build provider-specific request context.
|
||||||
|
|
||||||
|
Called before wrap_request(). Returns tls_profile (or None).
|
||||||
|
Implementations typically set contextvars that wrap_request()
|
||||||
|
and extra_headers() will read.
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def post_wrap_request(self, request_body: dict[str, Any]) -> None:
|
||||||
|
"""Post-wrap hook: async processing after wrap_request().
|
||||||
|
|
||||||
|
Called after wrap_request() completes. Use for async operations
|
||||||
|
like distributed session control that cannot run in sync wrap_request().
|
||||||
|
"""
|
||||||
|
|
||||||
|
def excluded_beta_tokens(self) -> frozenset[str]:
|
||||||
|
"""Beta tokens to strip from the merged anthropic-beta header.
|
||||||
|
|
||||||
|
Called by the request builder after merging envelope extra_headers
|
||||||
|
with client original headers. Return an empty frozenset to keep all.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Envelope Registry
|
# Envelope Registry
|
||||||
@@ -119,10 +152,14 @@ def ensure_providers_bootstrapped() -> None:
|
|||||||
from src.services.provider.adapters.antigravity.plugin import (
|
from src.services.provider.adapters.antigravity.plugin import (
|
||||||
register_all as _reg_antigravity,
|
register_all as _reg_antigravity,
|
||||||
)
|
)
|
||||||
|
from src.services.provider.adapters.claude_code.plugin import (
|
||||||
|
register_all as _reg_claude_code,
|
||||||
|
)
|
||||||
from src.services.provider.adapters.codex.plugin import register_all as _reg_codex
|
from src.services.provider.adapters.codex.plugin import register_all as _reg_codex
|
||||||
from src.services.provider.adapters.kiro.plugin import register_all as _reg_kiro
|
from src.services.provider.adapters.kiro.plugin import register_all as _reg_kiro
|
||||||
|
|
||||||
_reg_antigravity()
|
_reg_antigravity()
|
||||||
|
_reg_claude_code()
|
||||||
_reg_codex()
|
_reg_codex()
|
||||||
_reg_kiro()
|
_reg_kiro()
|
||||||
|
|
||||||
|
|||||||
@@ -60,6 +60,39 @@ PRESET_MODELS: dict[str, list[dict[str, Any]]] = {
|
|||||||
"display_name": "Claude Haiku 4.5",
|
"display_name": "Claude Haiku 4.5",
|
||||||
},
|
},
|
||||||
],
|
],
|
||||||
|
# Claude Code (Claude CLI OAuth 反代)
|
||||||
|
"claude_code": [
|
||||||
|
{
|
||||||
|
"id": "claude-opus-4-5-20251101",
|
||||||
|
"object": "model",
|
||||||
|
"owned_by": "anthropic",
|
||||||
|
"display_name": "Claude Opus 4.5",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "claude-opus-4-6",
|
||||||
|
"object": "model",
|
||||||
|
"owned_by": "anthropic",
|
||||||
|
"display_name": "Claude Opus 4.6",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "claude-sonnet-4-6",
|
||||||
|
"object": "model",
|
||||||
|
"owned_by": "anthropic",
|
||||||
|
"display_name": "Claude Sonnet 4.6",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "claude-sonnet-4-5-20250929",
|
||||||
|
"object": "model",
|
||||||
|
"owned_by": "anthropic",
|
||||||
|
"display_name": "Claude Sonnet 4.5",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "claude-haiku-4-5-20251001",
|
||||||
|
"object": "model",
|
||||||
|
"owned_by": "anthropic",
|
||||||
|
"display_name": "Claude Haiku 4.5",
|
||||||
|
},
|
||||||
|
],
|
||||||
# Codex (OpenAI CLI 反代)
|
# Codex (OpenAI CLI 反代)
|
||||||
"codex": [
|
"codex": [
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -361,6 +361,7 @@ class CacheAwareScheduler:
|
|||||||
max_candidates: int | None = None,
|
max_candidates: int | None = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: dict[str, bool] | None = None,
|
capability_requirements: dict[str, bool] | None = None,
|
||||||
|
request_body: dict | None = None,
|
||||||
) -> tuple[list[ProviderCandidate], str, int]:
|
) -> tuple[list[ProviderCandidate], str, int]:
|
||||||
"""
|
"""
|
||||||
预先获取所有可用的 Provider/Endpoint/Key 组合
|
预先获取所有可用的 Provider/Endpoint/Key 组合
|
||||||
@@ -519,6 +520,7 @@ class CacheAwareScheduler:
|
|||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
capability_requirements=capability_requirements,
|
capability_requirements=capability_requirements,
|
||||||
global_conversion_enabled=global_conversion_enabled,
|
global_conversion_enabled=global_conversion_enabled,
|
||||||
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 3. 应用优先级模式排序 + 调度模式排序
|
# 3. 应用优先级模式排序 + 调度模式排序
|
||||||
|
|||||||
@@ -35,12 +35,20 @@ from src.services.scheduling.utils import release_db_connection_before_await
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.models.database import GlobalModel
|
from src.models.database import GlobalModel
|
||||||
|
from src.services.provider.pool.config import PoolConfig
|
||||||
from src.services.scheduling.protocols import CandidateSorterProtocol
|
from src.services.scheduling.protocols import CandidateSorterProtocol
|
||||||
from src.services.scheduling.schemas import ProviderCandidate
|
from src.services.scheduling.schemas import ProviderCandidate
|
||||||
|
|
||||||
from src.services.cache.model_cache import ModelCacheService
|
from src.services.cache.model_cache import ModelCacheService
|
||||||
|
|
||||||
|
|
||||||
|
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
|
||||||
|
"""Return parsed PoolConfig if the provider has pool enabled, else None."""
|
||||||
|
from src.services.provider.pool.config import PoolConfig, parse_pool_config
|
||||||
|
|
||||||
|
return parse_pool_config(getattr(provider, "config", None))
|
||||||
|
|
||||||
|
|
||||||
def _sort_endpoints_by_family_priority(
|
def _sort_endpoints_by_family_priority(
|
||||||
eps: Sequence[ProviderEndpoint],
|
eps: Sequence[ProviderEndpoint],
|
||||||
) -> list[ProviderEndpoint]:
|
) -> list[ProviderEndpoint]:
|
||||||
@@ -359,6 +367,7 @@ class CandidateBuilder:
|
|||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: dict[str, bool] | None = None,
|
capability_requirements: dict[str, bool] | None = None,
|
||||||
global_conversion_enabled: bool = True,
|
global_conversion_enabled: bool = True,
|
||||||
|
request_body: dict | None = None,
|
||||||
) -> "list[ProviderCandidate]":
|
) -> "list[ProviderCandidate]":
|
||||||
"""
|
"""
|
||||||
构建候选列表
|
构建候选列表
|
||||||
@@ -417,6 +426,8 @@ class CandidateBuilder:
|
|||||||
] = {}
|
] = {}
|
||||||
exact_candidates: list[ProviderCandidate] = []
|
exact_candidates: list[ProviderCandidate] = []
|
||||||
convertible_candidates: list[ProviderCandidate] = []
|
convertible_candidates: list[ProviderCandidate] = []
|
||||||
|
pool_has_usable = False
|
||||||
|
pool_cfg = _get_pool_config(provider)
|
||||||
|
|
||||||
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
|
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
|
||||||
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
|
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
|
||||||
@@ -553,21 +564,34 @@ class CandidateBuilder:
|
|||||||
if not active_keys:
|
if not active_keys:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# 检查是否所有 Key 都是 TTL=0(轮换模式)
|
# --- Pool branch: select a single key internally ------
|
||||||
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
|
if pool_cfg is not None:
|
||||||
if use_random and len(active_keys) > 1:
|
selected_key = await self._pool_select_key(
|
||||||
logger.debug(
|
db, provider, pool_cfg, active_keys, request_body
|
||||||
" Provider {} 启用 Key 轮换模式 (endpoint_format={}, {} keys)",
|
)
|
||||||
provider.name,
|
if selected_key is None:
|
||||||
endpoint_format_str,
|
logger.debug(
|
||||||
len(active_keys),
|
"Pool[{}]: no schedulable key for endpoint {}",
|
||||||
|
str(provider.id)[:8],
|
||||||
|
endpoint_format_str,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
keys_to_check: list[ProviderAPIKey] = [selected_key]
|
||||||
|
else:
|
||||||
|
# --- Normal branch: check all keys ----
|
||||||
|
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
|
||||||
|
if use_random and len(active_keys) > 1:
|
||||||
|
logger.debug(
|
||||||
|
" Provider {} 启用 Key 轮换模式 (endpoint_format={}, {} keys)",
|
||||||
|
provider.name,
|
||||||
|
endpoint_format_str,
|
||||||
|
len(active_keys),
|
||||||
|
)
|
||||||
|
keys_to_check = self._sorter.shuffle_keys_by_internal_priority(
|
||||||
|
active_keys, affinity_key, use_random
|
||||||
)
|
)
|
||||||
|
|
||||||
keys = self._sorter.shuffle_keys_by_internal_priority(
|
for key in keys_to_check:
|
||||||
active_keys, affinity_key, use_random
|
|
||||||
)
|
|
||||||
|
|
||||||
for key in keys:
|
|
||||||
# Key 级别检查(健康度/熔断按 provider_format bucket)
|
# Key 级别检查(健康度/熔断按 provider_format bucket)
|
||||||
# 传入 provider_model_names 作为 candidate_models,
|
# 传入 provider_model_names 作为 candidate_models,
|
||||||
# 用于检查 Key 的 allowed_models 是否支持 Provider 定义的模型名称
|
# 用于检查 Key 的 allowed_models 是否支持 Provider 定义的模型名称
|
||||||
@@ -600,6 +624,13 @@ class CandidateBuilder:
|
|||||||
else:
|
else:
|
||||||
exact_candidates.append(candidate)
|
exact_candidates.append(candidate)
|
||||||
|
|
||||||
|
if is_available:
|
||||||
|
pool_has_usable = True
|
||||||
|
|
||||||
|
# Pool mode: stop after the first endpoint that produced a usable candidate.
|
||||||
|
if pool_cfg is not None and pool_has_usable:
|
||||||
|
break
|
||||||
|
|
||||||
candidates.extend(exact_candidates)
|
candidates.extend(exact_candidates)
|
||||||
candidates.extend(convertible_candidates)
|
candidates.extend(convertible_candidates)
|
||||||
|
|
||||||
@@ -608,3 +639,22 @@ class CandidateBuilder:
|
|||||||
candidates = candidates[:max_candidates]
|
candidates = candidates[:max_candidates]
|
||||||
|
|
||||||
return candidates
|
return candidates
|
||||||
|
|
||||||
|
async def _pool_select_key(
|
||||||
|
self,
|
||||||
|
db: Session,
|
||||||
|
provider: Provider,
|
||||||
|
pool_cfg: "PoolConfig",
|
||||||
|
active_keys: list[ProviderAPIKey],
|
||||||
|
request_body: dict | None,
|
||||||
|
) -> ProviderAPIKey | None:
|
||||||
|
"""Select a single key via pool scheduling (sticky -> cooldown/cost -> LRU)."""
|
||||||
|
from src.services.provider.pool.hooks import get_pool_hook
|
||||||
|
from src.services.provider.pool.manager import PoolManager
|
||||||
|
|
||||||
|
provider_type = str(getattr(provider, "provider_type", "") or "")
|
||||||
|
hook = get_pool_hook(provider_type)
|
||||||
|
session_uuid = hook.extract_session_uuid(request_body) if hook and request_body else None
|
||||||
|
mgr = PoolManager(str(provider.id), pool_cfg)
|
||||||
|
release_db_connection_before_await(db)
|
||||||
|
return await mgr.select_key(session_uuid, active_keys)
|
||||||
|
|||||||
@@ -33,6 +33,9 @@ class ExecutionResult:
|
|||||||
attempt_count: int = 0
|
attempt_count: int = 0
|
||||||
request_candidate_id: str | None = None
|
request_candidate_id: str | None = None
|
||||||
|
|
||||||
|
# pool scheduling summary (populated when pool mode is active)
|
||||||
|
pool_summary: dict[str, Any] | None = None
|
||||||
|
|
||||||
# failure
|
# failure
|
||||||
error_type: str | None = None
|
error_type: str | None = None
|
||||||
error_message: str | None = None
|
error_message: str | None = None
|
||||||
|
|||||||
@@ -186,6 +186,151 @@ class TaskService:
|
|||||||
request_body=request_body,
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_session_uuid(
|
||||||
|
provider_type: str, request_body: dict[str, Any] | None
|
||||||
|
) -> str | None:
|
||||||
|
"""Extract a session UUID from the request body (provider-type aware)."""
|
||||||
|
if not isinstance(request_body, dict):
|
||||||
|
return None
|
||||||
|
from src.services.provider.pool.hooks import get_pool_hook
|
||||||
|
|
||||||
|
hook = get_pool_hook(provider_type)
|
||||||
|
if hook is not None:
|
||||||
|
return hook.extract_session_uuid(request_body)
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _apply_pool_reorder(
|
||||||
|
candidates: list[Any],
|
||||||
|
request_body: dict[str, Any] | None,
|
||||||
|
) -> tuple[list[Any], list[Any]]:
|
||||||
|
"""Apply Account Pool reordering when applicable.
|
||||||
|
|
||||||
|
Groups candidates by provider_id and applies pool reordering
|
||||||
|
independently per provider, then reassembles in original group order.
|
||||||
|
Non-pool providers are left in their original order.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (reordered_candidates, pool_traces) where pool_traces
|
||||||
|
is a list of :class:`PoolSchedulingTrace` objects (one per
|
||||||
|
pooled provider group, may be empty).
|
||||||
|
"""
|
||||||
|
if not candidates:
|
||||||
|
return candidates, []
|
||||||
|
|
||||||
|
pool_traces: list[Any] = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
from collections import OrderedDict
|
||||||
|
|
||||||
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
from src.services.provider.pool.manager import PoolManager
|
||||||
|
|
||||||
|
# Group candidates by provider_id while preserving order.
|
||||||
|
groups: OrderedDict[str, list[Any]] = OrderedDict()
|
||||||
|
for c in candidates:
|
||||||
|
pid = str(getattr(c.provider, "id", "") or "")
|
||||||
|
groups.setdefault(pid, []).append(c)
|
||||||
|
|
||||||
|
result: list[Any] = []
|
||||||
|
for pid, group in groups.items():
|
||||||
|
provider = group[0].provider
|
||||||
|
pool_cfg = parse_pool_config(getattr(provider, "config", None))
|
||||||
|
if pool_cfg is None or not pid:
|
||||||
|
result.extend(group)
|
||||||
|
continue
|
||||||
|
|
||||||
|
provider_type = str(getattr(provider, "provider_type", "") or "")
|
||||||
|
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
|
||||||
|
mgr = PoolManager(pid, pool_cfg)
|
||||||
|
reordered = await mgr.reorder_candidates(session_uuid, group)
|
||||||
|
result.extend(reordered)
|
||||||
|
|
||||||
|
# Extract trace attached by PoolManager.reorder_candidates
|
||||||
|
if reordered:
|
||||||
|
trace = getattr(reordered[0], "_pool_scheduling_trace", None)
|
||||||
|
if trace is not None:
|
||||||
|
pool_traces.append(trace)
|
||||||
|
|
||||||
|
return result, pool_traces
|
||||||
|
except Exception:
|
||||||
|
from src.core.logger import logger
|
||||||
|
|
||||||
|
logger.opt(exception=True).debug("Pool reorder failed, using original order")
|
||||||
|
return candidates, []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _pool_on_success(
|
||||||
|
candidate: Any,
|
||||||
|
request_body: dict[str, Any] | None,
|
||||||
|
) -> None:
|
||||||
|
"""Notify the pool manager about a successful request (sticky + LRU)."""
|
||||||
|
try:
|
||||||
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
from src.services.provider.pool.manager import PoolManager
|
||||||
|
|
||||||
|
provider = candidate.provider
|
||||||
|
provider_config = getattr(provider, "config", None)
|
||||||
|
pool_cfg = parse_pool_config(provider_config)
|
||||||
|
if pool_cfg is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
provider_id = str(getattr(provider, "id", "") or "")
|
||||||
|
key_id = str(getattr(candidate.key, "id", "") or "")
|
||||||
|
if not provider_id or not key_id:
|
||||||
|
return
|
||||||
|
|
||||||
|
provider_type = str(getattr(provider, "provider_type", "") or "")
|
||||||
|
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
|
||||||
|
|
||||||
|
mgr = PoolManager(provider_id, pool_cfg)
|
||||||
|
await mgr.on_request_success(
|
||||||
|
session_uuid=session_uuid,
|
||||||
|
key_id=key_id,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _pool_on_error(
|
||||||
|
provider: Any,
|
||||||
|
key: Any,
|
||||||
|
status_code: int,
|
||||||
|
cause: Any,
|
||||||
|
) -> None:
|
||||||
|
"""Notify the pool manager about an upstream error (health policy)."""
|
||||||
|
try:
|
||||||
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
from src.services.provider.pool.health_policy import apply_health_policy
|
||||||
|
|
||||||
|
pool_cfg = parse_pool_config(getattr(provider, "config", None))
|
||||||
|
if pool_cfg is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
error_text = ""
|
||||||
|
resp_headers: dict[str, str] = {}
|
||||||
|
if getattr(cause, "response", None) is not None:
|
||||||
|
try:
|
||||||
|
error_text = (cause.response.text or "")[:4000]
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
resp_headers = dict(cause.response.headers)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
await apply_health_policy(
|
||||||
|
provider_id=str(provider.id),
|
||||||
|
key_id=str(key.id),
|
||||||
|
status_code=status_code,
|
||||||
|
error_body=error_text,
|
||||||
|
response_headers=resp_headers,
|
||||||
|
config=pool_cfg,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
async def _execute_sync_unified(
|
async def _execute_sync_unified(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -287,6 +432,12 @@ class TaskService:
|
|||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
capability_requirements=capability_requirements,
|
capability_requirements=capability_requirements,
|
||||||
preferred_key_ids=preferred_key_ids,
|
preferred_key_ids=preferred_key_ids,
|
||||||
|
request_body=request_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Account Pool: reorder candidates for claude_code providers.
|
||||||
|
all_candidates, pool_traces = await self._apply_pool_reorder(
|
||||||
|
all_candidates, request_body=request_body
|
||||||
)
|
)
|
||||||
|
|
||||||
candidate_record_map = candidate_resolver.create_candidate_records(
|
candidate_record_map = candidate_resolver.create_candidate_records(
|
||||||
@@ -350,6 +501,9 @@ class TaskService:
|
|||||||
)
|
)
|
||||||
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
|
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
|
||||||
|
|
||||||
|
# Account Pool: on success, update sticky binding + LRU.
|
||||||
|
await self._pool_on_success(candidate, request_body)
|
||||||
|
|
||||||
if is_stream:
|
if is_stream:
|
||||||
return AttemptResult(
|
return AttemptResult(
|
||||||
kind=AttemptKind.STREAM,
|
kind=AttemptKind.STREAM,
|
||||||
@@ -441,6 +595,16 @@ class TaskService:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if result.success:
|
if result.success:
|
||||||
|
# Build pool scheduling summary from traces collected during reorder.
|
||||||
|
if pool_traces and result.key_id:
|
||||||
|
try:
|
||||||
|
for pt in pool_traces:
|
||||||
|
summary = pt.build_summary(result.key_id)
|
||||||
|
if summary:
|
||||||
|
result.pool_summary = summary
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
return result
|
return result
|
||||||
|
|
||||||
self._raise_all_failed_exception(
|
self._raise_all_failed_exception(
|
||||||
@@ -933,6 +1097,9 @@ class TaskService:
|
|||||||
attempt=attempt,
|
attempt=attempt,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Account Pool: apply health policy (cooldown/disable).
|
||||||
|
await self._pool_on_error(provider, key, status_code, cause)
|
||||||
|
|
||||||
converted_error = extra_data.get("converted_error")
|
converted_error = extra_data.get("converted_error")
|
||||||
serializable_extra_data = {
|
serializable_extra_data = {
|
||||||
k: v for k, v in extra_data.items() if k != "converted_error"
|
k: v for k, v in extra_data.items() if k != "converted_error"
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
|
|||||||
"billing_snapshot",
|
"billing_snapshot",
|
||||||
"billing_updated_at",
|
"billing_updated_at",
|
||||||
"perf",
|
"perf",
|
||||||
|
"pool_summary",
|
||||||
"_metadata_truncated",
|
"_metadata_truncated",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -7,14 +7,23 @@ import ssl
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
try:
|
|
||||||
import certifi
|
|
||||||
|
|
||||||
_SSL_CONTEXT = ssl.create_default_context(cafile=certifi.where())
|
def _create_default_ssl_context() -> ssl.SSLContext:
|
||||||
except ImportError:
|
try:
|
||||||
|
import certifi
|
||||||
|
|
||||||
|
return ssl.create_default_context(cafile=certifi.where())
|
||||||
|
except ImportError:
|
||||||
|
return ssl.create_default_context()
|
||||||
|
|
||||||
|
|
||||||
|
try:
|
||||||
|
_SSL_CONTEXT = _create_default_ssl_context()
|
||||||
|
except Exception:
|
||||||
_SSL_CONTEXT = ssl.create_default_context()
|
_SSL_CONTEXT = ssl.create_default_context()
|
||||||
|
|
||||||
_PROXY_SSL_CONTEXT: ssl.SSLContext | None = None
|
_PROXY_SSL_CONTEXT: ssl.SSLContext | None = None
|
||||||
|
_PROFILE_SSL_CONTEXTS: dict[str, ssl.SSLContext] = {}
|
||||||
|
|
||||||
|
|
||||||
def get_ssl_context() -> ssl.SSLContext:
|
def get_ssl_context() -> ssl.SSLContext:
|
||||||
@@ -55,3 +64,54 @@ def get_proxy_ssl_context(expected_fingerprint: str | None = None) -> ssl.SSLCon
|
|||||||
_PROXY_SSL_CONTEXT = ctx
|
_PROXY_SSL_CONTEXT = ctx
|
||||||
# TODO: 实现基于 expected_fingerprint 的证书指纹校验
|
# TODO: 实现基于 expected_fingerprint 的证书指纹校验
|
||||||
return _PROXY_SSL_CONTEXT
|
return _PROXY_SSL_CONTEXT
|
||||||
|
|
||||||
|
|
||||||
|
def _build_claude_code_ssl_context() -> ssl.SSLContext:
|
||||||
|
"""构建 Claude Code best-effort TLS 配置。
|
||||||
|
|
||||||
|
说明:Python/OpenSSL 无法完整模拟 Node.js ClientHello。
|
||||||
|
这里仅做可控项的尽力对齐(ALPN/TLS 版本/常见 cipher 偏好)。
|
||||||
|
"""
|
||||||
|
ctx = _create_default_ssl_context()
|
||||||
|
ctx.minimum_version = ssl.TLSVersion.TLSv1_2
|
||||||
|
try:
|
||||||
|
ctx.maximum_version = ssl.TLSVersion.TLSv1_3
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
ctx.set_alpn_protocols(["h2", "http/1.1"])
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
ctx.set_ciphers(
|
||||||
|
"ECDHE-ECDSA-AES128-GCM-SHA256:"
|
||||||
|
"ECDHE-RSA-AES128-GCM-SHA256:"
|
||||||
|
"ECDHE-ECDSA-AES256-GCM-SHA384:"
|
||||||
|
"ECDHE-RSA-AES256-GCM-SHA384:"
|
||||||
|
"ECDHE-ECDSA-CHACHA20-POLY1305:"
|
||||||
|
"ECDHE-RSA-CHACHA20-POLY1305"
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return ctx
|
||||||
|
|
||||||
|
|
||||||
|
def get_ssl_context_for_profile(tls_profile: str | None = None) -> ssl.SSLContext:
|
||||||
|
"""按 profile 返回 SSL 上下文。"""
|
||||||
|
profile = str(tls_profile or "").strip().lower()
|
||||||
|
if not profile:
|
||||||
|
return get_ssl_context()
|
||||||
|
|
||||||
|
if profile in _PROFILE_SSL_CONTEXTS:
|
||||||
|
logger.debug("复用 TLS profile SSL context: {}", profile)
|
||||||
|
return _PROFILE_SSL_CONTEXTS[profile]
|
||||||
|
|
||||||
|
if profile == "claude_code_nodejs":
|
||||||
|
logger.info("启用 TLS profile: {}(best-effort)", profile)
|
||||||
|
ctx = _build_claude_code_ssl_context()
|
||||||
|
else:
|
||||||
|
logger.warning("未知 TLS profile: {},回退默认 SSL context", profile)
|
||||||
|
ctx = get_ssl_context()
|
||||||
|
|
||||||
|
_PROFILE_SSL_CONTEXTS[profile] = ctx
|
||||||
|
return ctx
|
||||||
|
|||||||
84
tests/api/handlers/base/test_upstream_stream_bridge.py
Normal file
84
tests/api/handlers/base/test_upstream_stream_bridge.py
Normal file
@@ -0,0 +1,84 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.handlers.base.upstream_stream_bridge import (
|
||||||
|
aggregate_upstream_stream_to_internal_response,
|
||||||
|
)
|
||||||
|
from src.core.api_format.conversion import register_default_normalizers
|
||||||
|
from src.core.api_format.conversion.internal import TextBlock
|
||||||
|
|
||||||
|
|
||||||
|
async def _iter_stream_lines(lines: list[str]) -> AsyncIterator[bytes]:
|
||||||
|
for line in lines:
|
||||||
|
yield line.encode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_aggregate_claude_stream_uses_message_start_usage_when_message_delta_absent() -> None:
|
||||||
|
register_default_normalizers()
|
||||||
|
|
||||||
|
lines = [
|
||||||
|
"data: "
|
||||||
|
+ json.dumps(
|
||||||
|
{
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {
|
||||||
|
"id": "msg_bridge_usage",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": "claude-sonnet-4-5",
|
||||||
|
"content": [],
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": 120,
|
||||||
|
"output_tokens": 0,
|
||||||
|
"cache_read_input_tokens": 11,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ensure_ascii=False,
|
||||||
|
)
|
||||||
|
+ "\n",
|
||||||
|
"data: "
|
||||||
|
+ json.dumps(
|
||||||
|
{
|
||||||
|
"type": "content_block_start",
|
||||||
|
"index": 0,
|
||||||
|
"content_block": {"type": "text", "text": ""},
|
||||||
|
},
|
||||||
|
ensure_ascii=False,
|
||||||
|
)
|
||||||
|
+ "\n",
|
||||||
|
"data: "
|
||||||
|
+ json.dumps(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": "hello"},
|
||||||
|
},
|
||||||
|
ensure_ascii=False,
|
||||||
|
)
|
||||||
|
+ "\n",
|
||||||
|
"data: "
|
||||||
|
+ json.dumps({"type": "content_block_stop", "index": 0}, ensure_ascii=False)
|
||||||
|
+ "\n",
|
||||||
|
]
|
||||||
|
|
||||||
|
internal = await aggregate_upstream_stream_to_internal_response(
|
||||||
|
_iter_stream_lines(lines),
|
||||||
|
provider_api_format="claude:cli",
|
||||||
|
provider_name="claude_code",
|
||||||
|
model="claude-sonnet-4-5",
|
||||||
|
request_id="req_bridge_usage",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert internal.usage is not None
|
||||||
|
assert internal.usage.input_tokens == 120
|
||||||
|
assert internal.usage.output_tokens == 0
|
||||||
|
assert internal.usage.cache_read_tokens == 11
|
||||||
|
assert len(internal.content) == 1
|
||||||
|
assert isinstance(internal.content[0], TextBlock)
|
||||||
|
assert internal.content[0].text == "hello"
|
||||||
@@ -283,6 +283,40 @@ def test_claude_stream_chunk_and_event_roundtrip_basic() -> None:
|
|||||||
assert out_events[-1]["type"] == "message_stop"
|
assert out_events[-1]["type"] == "message_stop"
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_stream_message_start_preserves_usage() -> None:
|
||||||
|
n = ClaudeNormalizer()
|
||||||
|
state = StreamState(model="claude-3-sonnet")
|
||||||
|
|
||||||
|
chunk = {
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {
|
||||||
|
"id": "msg_usage_start",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": "claude-3-sonnet",
|
||||||
|
"content": [],
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": 9,
|
||||||
|
"output_tokens": 0,
|
||||||
|
"cache_read_input_tokens": 3,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
events = n.stream_chunk_to_internal(chunk, state)
|
||||||
|
assert len(events) == 1
|
||||||
|
start_event = events[0]
|
||||||
|
assert isinstance(start_event, MessageStartEvent)
|
||||||
|
assert start_event.usage is not None
|
||||||
|
assert start_event.usage.input_tokens == 9
|
||||||
|
assert start_event.usage.cache_read_tokens == 3
|
||||||
|
|
||||||
|
out = n.stream_event_from_internal(start_event, StreamState(model="claude-3-sonnet"))
|
||||||
|
assert out[0]["type"] == "message_start"
|
||||||
|
assert out[0]["message"]["usage"]["input_tokens"] == 9
|
||||||
|
assert out[0]["message"]["usage"]["cache_read_input_tokens"] == 3
|
||||||
|
|
||||||
|
|
||||||
def test_claude_error_conversion() -> None:
|
def test_claude_error_conversion() -> None:
|
||||||
n = ClaudeNormalizer()
|
n = ClaudeNormalizer()
|
||||||
err_resp = {
|
err_resp = {
|
||||||
|
|||||||
39
tests/services/test_admin_provider_defaults.py
Normal file
39
tests/services/test_admin_provider_defaults.py
Normal file
@@ -0,0 +1,39 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.admin.providers.routes import (
|
||||||
|
_merge_claude_code_advanced_config,
|
||||||
|
_should_enable_format_conversion_by_default,
|
||||||
|
)
|
||||||
|
from src.core.exceptions import InvalidRequestException
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_defaults_format_conversion_enabled() -> None:
|
||||||
|
assert _should_enable_format_conversion_by_default("claude_code") is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_custom_defaults_format_conversion_disabled() -> None:
|
||||||
|
assert _should_enable_format_conversion_by_default("custom") is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_claude_code_advanced_clears_stale_config_for_non_claude_provider() -> None:
|
||||||
|
merged, changed = _merge_claude_code_advanced_config(
|
||||||
|
provider_type="custom",
|
||||||
|
provider_config={"foo": "bar", "claude_code_advanced": {"max_sessions": 9}},
|
||||||
|
claude_code_advanced=None,
|
||||||
|
claude_advanced_in_payload=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert merged == {"foo": "bar"}
|
||||||
|
assert changed is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_claude_code_advanced_rejects_non_claude_payload() -> None:
|
||||||
|
with pytest.raises(InvalidRequestException):
|
||||||
|
_merge_claude_code_advanced_config(
|
||||||
|
provider_type="custom",
|
||||||
|
provider_config={},
|
||||||
|
claude_code_advanced={"max_sessions": 9},
|
||||||
|
claude_advanced_in_payload=True,
|
||||||
|
)
|
||||||
@@ -30,6 +30,7 @@ class _FakeScheduler:
|
|||||||
max_candidates: int | None = None,
|
max_candidates: int | None = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: dict[str, bool] | None = None,
|
capability_requirements: dict[str, bool] | None = None,
|
||||||
|
request_body: dict | None = None,
|
||||||
) -> tuple[list[Any], str, int]:
|
) -> tuple[list[Any], str, int]:
|
||||||
_ = (
|
_ = (
|
||||||
db,
|
db,
|
||||||
|
|||||||
166
tests/services/test_claude_code_distributed_sessions.py
Normal file
166
tests/services/test_claude_code_distributed_sessions.py
Normal file
@@ -0,0 +1,166 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.core.exceptions import ConcurrencyLimitError
|
||||||
|
from src.services.provider.adapters.claude_code.context import (
|
||||||
|
ClaudeCodeRequestContext,
|
||||||
|
set_claude_code_request_context,
|
||||||
|
)
|
||||||
|
from src.services.provider.adapters.claude_code.envelope import (
|
||||||
|
claude_code_envelope,
|
||||||
|
enforce_distributed_session_controls,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _StubRedis:
|
||||||
|
def __init__(self, *, result=None, exc: Exception | None = None) -> None:
|
||||||
|
self._result = result
|
||||||
|
self._exc = exc
|
||||||
|
self.calls: list[tuple] = []
|
||||||
|
|
||||||
|
async def eval(self, *args):
|
||||||
|
self.calls.append(args)
|
||||||
|
if self._exc is not None:
|
||||||
|
raise self._exc
|
||||||
|
return self._result
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _reset_claude_code_context():
|
||||||
|
set_claude_code_request_context(None)
|
||||||
|
yield
|
||||||
|
set_claude_code_request_context(None)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_distributed_session_controls_accept_when_redis_allows(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
|
stub = _StubRedis(result=[1, 1])
|
||||||
|
|
||||||
|
async def _fake_get_redis_client(*, require_redis: bool = False):
|
||||||
|
_ = require_redis
|
||||||
|
return stub
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.provider.adapters.claude_code.envelope.get_redis_client",
|
||||||
|
_fake_get_redis_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx = ClaudeCodeRequestContext(
|
||||||
|
scope_key=f"key:test-dist-ok-{uuid.uuid4()}",
|
||||||
|
key_id="key-ok",
|
||||||
|
max_sessions=1,
|
||||||
|
session_idle_timeout_minutes=5,
|
||||||
|
)
|
||||||
|
request_body = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_11111111-1111-1111-1111-111111111111"}
|
||||||
|
}
|
||||||
|
|
||||||
|
await enforce_distributed_session_controls(request_body, ctx)
|
||||||
|
assert len(stub.calls) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_distributed_session_controls_reject_when_redis_denies(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
|
stub = _StubRedis(result=[0, 1])
|
||||||
|
|
||||||
|
async def _fake_get_redis_client(*, require_redis: bool = False):
|
||||||
|
_ = require_redis
|
||||||
|
return stub
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.provider.adapters.claude_code.envelope.get_redis_client",
|
||||||
|
_fake_get_redis_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx = ClaudeCodeRequestContext(
|
||||||
|
scope_key=f"key:test-dist-deny-{uuid.uuid4()}",
|
||||||
|
key_id="key-deny",
|
||||||
|
max_sessions=1,
|
||||||
|
session_idle_timeout_minutes=5,
|
||||||
|
)
|
||||||
|
request_body = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_22222222-2222-2222-2222-222222222222"}
|
||||||
|
}
|
||||||
|
|
||||||
|
with pytest.raises(ConcurrencyLimitError):
|
||||||
|
await enforce_distributed_session_controls(request_body, ctx)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_distributed_session_controls_fallback_to_local_when_redis_error(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
|
stub = _StubRedis(exc=RuntimeError("redis unavailable"))
|
||||||
|
|
||||||
|
async def _fake_get_redis_client(*, require_redis: bool = False):
|
||||||
|
_ = require_redis
|
||||||
|
return stub
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.provider.adapters.claude_code.envelope.get_redis_client",
|
||||||
|
_fake_get_redis_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx = ClaudeCodeRequestContext(
|
||||||
|
scope_key=f"key:test-dist-fallback-{uuid.uuid4()}",
|
||||||
|
key_id="key-fallback",
|
||||||
|
max_sessions=1,
|
||||||
|
session_idle_timeout_minutes=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
first = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"}
|
||||||
|
}
|
||||||
|
second = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb"}
|
||||||
|
}
|
||||||
|
|
||||||
|
await enforce_distributed_session_controls(first, ctx)
|
||||||
|
with pytest.raises(ConcurrencyLimitError):
|
||||||
|
await enforce_distributed_session_controls(second, ctx)
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrap_request_skips_local_limit_when_distributed_store_available(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.provider.adapters.claude_code.envelope.get_redis_client_sync",
|
||||||
|
lambda: object(),
|
||||||
|
)
|
||||||
|
scope_key = f"key:test-wrap-skip-{uuid.uuid4()}"
|
||||||
|
set_claude_code_request_context(
|
||||||
|
ClaudeCodeRequestContext(
|
||||||
|
is_stream=False,
|
||||||
|
scope_key=scope_key,
|
||||||
|
key_id="key-wrap",
|
||||||
|
max_sessions=1,
|
||||||
|
session_idle_timeout_minutes=5,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
first = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_11111111-1111-1111-1111-111111111111"}
|
||||||
|
}
|
||||||
|
second = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_22222222-2222-2222-2222-222222222222"}
|
||||||
|
}
|
||||||
|
|
||||||
|
claude_code_envelope.wrap_request(
|
||||||
|
first,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
claude_code_envelope.wrap_request(
|
||||||
|
second,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
178
tests/services/test_claude_code_envelope.py
Normal file
178
tests/services/test_claude_code_envelope.py
Normal file
@@ -0,0 +1,178 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.handlers.base.request_builder import PassthroughRequestBuilder
|
||||||
|
from src.config.settings import config
|
||||||
|
from src.services.provider.adapters.claude_code.constants import (
|
||||||
|
BETA_CLAUDE_CODE,
|
||||||
|
BETA_CONTEXT_1M,
|
||||||
|
BETA_INTERLEAVED_THINKING,
|
||||||
|
BETA_OAUTH,
|
||||||
|
CLAUDE_CODE_REQUIRED_BETA_TOKENS,
|
||||||
|
DEFAULT_ACCEPT,
|
||||||
|
DEFAULT_ANTHROPIC_VERSION,
|
||||||
|
STREAM_HELPER_METHOD,
|
||||||
|
)
|
||||||
|
from src.services.provider.adapters.claude_code.context import set_claude_code_request_context
|
||||||
|
from src.services.provider.adapters.claude_code.envelope import (
|
||||||
|
claude_code_envelope,
|
||||||
|
merge_anthropic_beta_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _reset_claude_code_context():
|
||||||
|
set_claude_code_request_context(None)
|
||||||
|
yield
|
||||||
|
set_claude_code_request_context(None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_anthropic_beta_tokens_adds_required_and_deduplicates() -> None:
|
||||||
|
merged = merge_anthropic_beta_tokens(
|
||||||
|
"context-1m-2025-08-07,oauth-2025-04-20,custom-beta,claude-code-20250219"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert merged.split(",") == [
|
||||||
|
BETA_CLAUDE_CODE,
|
||||||
|
BETA_OAUTH,
|
||||||
|
BETA_INTERLEAVED_THINKING,
|
||||||
|
"context-1m-2025-08-07",
|
||||||
|
"custom-beta",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_envelope_extra_headers_include_required_defaults(
|
||||||
|
monkeypatch,
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.setattr(config, "internal_user_agent_claude_cli", "claude-code/test")
|
||||||
|
|
||||||
|
headers = claude_code_envelope.extra_headers() or {}
|
||||||
|
|
||||||
|
assert headers.get("anthropic-version") == DEFAULT_ANTHROPIC_VERSION
|
||||||
|
assert headers.get("anthropic-beta") == ",".join(CLAUDE_CODE_REQUIRED_BETA_TOKENS)
|
||||||
|
assert headers.get("Accept") == DEFAULT_ACCEPT
|
||||||
|
assert headers.get("X-App") == "cli"
|
||||||
|
assert headers.get("X-Stainless-Lang") == "js"
|
||||||
|
assert headers.get("Anthropic-Dangerous-Direct-Browser-Access") == "true"
|
||||||
|
assert headers.get("User-Agent") == "claude-code/test"
|
||||||
|
assert "x-stainless-helper-method" not in headers
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_envelope_adds_stream_helper_header_for_stream_request() -> None:
|
||||||
|
_, _ = claude_code_envelope.wrap_request(
|
||||||
|
{"stream": True},
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
headers = claude_code_envelope.extra_headers() or {}
|
||||||
|
assert headers.get("Accept") == DEFAULT_ACCEPT
|
||||||
|
assert headers.get("x-stainless-helper-method") == STREAM_HELPER_METHOD
|
||||||
|
|
||||||
|
|
||||||
|
def test_passthrough_request_builder_drops_context_1m_for_claude_code_oauth() -> None:
|
||||||
|
builder = PassthroughRequestBuilder()
|
||||||
|
endpoint = SimpleNamespace(header_rules=None)
|
||||||
|
key = SimpleNamespace(api_key="unused")
|
||||||
|
|
||||||
|
headers = builder.build_headers(
|
||||||
|
original_headers={"anthropic-beta": BETA_CONTEXT_1M},
|
||||||
|
endpoint=endpoint,
|
||||||
|
key=key,
|
||||||
|
extra_headers={"anthropic-beta": ",".join(CLAUDE_CODE_REQUIRED_BETA_TOKENS)},
|
||||||
|
pre_computed_auth=("Authorization", "Bearer test-token"),
|
||||||
|
envelope=claude_code_envelope,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert headers.get("anthropic-beta") == ",".join(CLAUDE_CODE_REQUIRED_BETA_TOKENS)
|
||||||
|
|
||||||
|
|
||||||
|
def test_passthrough_request_builder_keeps_context_1m_for_non_claude_code_provider() -> None:
|
||||||
|
builder = PassthroughRequestBuilder()
|
||||||
|
endpoint = SimpleNamespace(header_rules=None)
|
||||||
|
key = SimpleNamespace(api_key="unused")
|
||||||
|
|
||||||
|
headers = builder.build_headers(
|
||||||
|
original_headers={"anthropic-beta": BETA_CONTEXT_1M},
|
||||||
|
endpoint=endpoint,
|
||||||
|
key=key,
|
||||||
|
extra_headers={"anthropic-beta": ",".join(CLAUDE_CODE_REQUIRED_BETA_TOKENS)},
|
||||||
|
pre_computed_auth=("Authorization", "Bearer test-token"),
|
||||||
|
envelope=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert BETA_CONTEXT_1M in str(headers.get("anthropic-beta") or "")
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_envelope_filters_invalid_thinking_blocks_when_enabled() -> None:
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "enabled"},
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "keep", "signature": "sig_valid"},
|
||||||
|
{"type": "thinking", "thinking": "drop-empty-signature", "signature": ""},
|
||||||
|
{
|
||||||
|
"type": "thinking",
|
||||||
|
"thinking": "drop-dummy-signature",
|
||||||
|
"signature": "skip_thought_signature_validator",
|
||||||
|
},
|
||||||
|
{"type": "redacted_thinking", "data": "keep", "signature": "sig_redacted"},
|
||||||
|
{"type": "redacted_thinking", "data": "drop-no-signature"},
|
||||||
|
{"thinking": "drop-no-type"},
|
||||||
|
{"type": "text", "text": "ok"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
wrapped, _ = claude_code_envelope.wrap_request(
|
||||||
|
body,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
content = wrapped["messages"][0]["content"]
|
||||||
|
assert content == [
|
||||||
|
{"type": "thinking", "thinking": "keep", "signature": "sig_valid"},
|
||||||
|
{"type": "redacted_thinking", "data": "keep", "signature": "sig_redacted"},
|
||||||
|
{"type": "text", "text": "ok"},
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_envelope_drops_all_thinking_blocks_when_disabled() -> None:
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "disabled"},
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "remove", "signature": "sig_valid"},
|
||||||
|
{"type": "redacted_thinking", "data": "remove", "signature": "sig_redacted"},
|
||||||
|
{"type": "text", "text": "keep"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "remove-user", "signature": "sig_user"},
|
||||||
|
{"type": "text", "text": "keep-user"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
wrapped, _ = claude_code_envelope.wrap_request(
|
||||||
|
body,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert wrapped["messages"][0]["content"] == [{"type": "text", "text": "keep"}]
|
||||||
|
assert wrapped["messages"][1]["content"] == [{"type": "text", "text": "keep-user"}]
|
||||||
192
tests/services/test_claude_code_runtime_controls.py
Normal file
192
tests/services/test_claude_code_runtime_controls.py
Normal file
@@ -0,0 +1,192 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.core.exceptions import ConcurrencyLimitError
|
||||||
|
from src.services.provider.adapters.claude_code.constants import TLS_PROFILE_CLAUDE_CODE
|
||||||
|
from src.services.provider.adapters.claude_code.context import (
|
||||||
|
ClaudeCodeRequestContext,
|
||||||
|
build_and_set_claude_code_request_context,
|
||||||
|
build_claude_code_request_context,
|
||||||
|
get_claude_code_request_context,
|
||||||
|
resolve_claude_code_tls_profile,
|
||||||
|
set_claude_code_request_context,
|
||||||
|
)
|
||||||
|
from src.services.provider.adapters.claude_code.envelope import claude_code_envelope
|
||||||
|
from src.utils.ssl_utils import get_ssl_context, get_ssl_context_for_profile
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _reset_claude_code_context():
|
||||||
|
set_claude_code_request_context(None)
|
||||||
|
yield
|
||||||
|
set_claude_code_request_context(None)
|
||||||
|
|
||||||
|
|
||||||
|
def _session_tail(user_id: str) -> str:
|
||||||
|
return user_id.split("_session_")[-1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_context_reads_claude_code_advanced_from_provider_config() -> None:
|
||||||
|
ctx = build_claude_code_request_context(
|
||||||
|
provider_config={
|
||||||
|
"claude_code_advanced": {
|
||||||
|
"max_sessions": 3,
|
||||||
|
"session_idle_timeout_minutes": 7,
|
||||||
|
"enable_tls_fingerprint": True,
|
||||||
|
"session_id_masking_enabled": True,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
key_id="key-123",
|
||||||
|
is_stream=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert ctx.scope_key == "key:key-123"
|
||||||
|
assert ctx.key_id == "key-123"
|
||||||
|
assert ctx.max_sessions == 3
|
||||||
|
assert ctx.session_idle_timeout_minutes == 7
|
||||||
|
assert ctx.enable_tls_fingerprint is True
|
||||||
|
assert ctx.session_id_masking_enabled is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_and_set_context_returns_tls_profile_when_enabled() -> None:
|
||||||
|
ctx, tls_profile = build_and_set_claude_code_request_context(
|
||||||
|
provider_config={"claude_code_advanced": {"enable_tls_fingerprint": True}},
|
||||||
|
key_id="key-tls",
|
||||||
|
is_stream=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert ctx.key_id == "key-tls"
|
||||||
|
assert get_claude_code_request_context() == ctx
|
||||||
|
assert tls_profile == TLS_PROFILE_CLAUDE_CODE
|
||||||
|
assert resolve_claude_code_tls_profile(ctx) == TLS_PROFILE_CLAUDE_CODE
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_ssl_context_for_claude_code_profile_is_cached() -> None:
|
||||||
|
first = get_ssl_context_for_profile(TLS_PROFILE_CLAUDE_CODE)
|
||||||
|
second = get_ssl_context_for_profile(TLS_PROFILE_CLAUDE_CODE)
|
||||||
|
|
||||||
|
assert first is second
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_ssl_context_for_unknown_profile_falls_back_to_default() -> None:
|
||||||
|
assert get_ssl_context_for_profile("unknown_profile") is get_ssl_context()
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrap_request_masks_session_id_when_enabled() -> None:
|
||||||
|
scope_key = f"key:test-mask-{uuid.uuid4()}"
|
||||||
|
set_claude_code_request_context(
|
||||||
|
ClaudeCodeRequestContext(
|
||||||
|
is_stream=False,
|
||||||
|
scope_key=scope_key,
|
||||||
|
key_id="key-mask",
|
||||||
|
session_id_masking_enabled=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
body1 = {
|
||||||
|
"metadata": {
|
||||||
|
"user_id": "user_client_account_main_session_11111111-1111-1111-1111-111111111111"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
body2 = {
|
||||||
|
"metadata": {
|
||||||
|
"user_id": "user_client_account_main_session_22222222-2222-2222-2222-222222222222"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
wrapped1, _ = claude_code_envelope.wrap_request(
|
||||||
|
body1,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
wrapped2, _ = claude_code_envelope.wrap_request(
|
||||||
|
body2,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
tail1 = _session_tail(wrapped1["metadata"]["user_id"])
|
||||||
|
tail2 = _session_tail(wrapped2["metadata"]["user_id"])
|
||||||
|
|
||||||
|
assert tail1 != "11111111-1111-1111-1111-111111111111"
|
||||||
|
assert tail1 == tail2
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrap_request_enforces_max_sessions() -> None:
|
||||||
|
scope_key = f"key:test-limit-{uuid.uuid4()}"
|
||||||
|
set_claude_code_request_context(
|
||||||
|
ClaudeCodeRequestContext(
|
||||||
|
is_stream=False,
|
||||||
|
scope_key=scope_key,
|
||||||
|
key_id="key-limit",
|
||||||
|
max_sessions=1,
|
||||||
|
session_idle_timeout_minutes=5,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
first = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"}
|
||||||
|
}
|
||||||
|
second = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb"}
|
||||||
|
}
|
||||||
|
|
||||||
|
claude_code_envelope.wrap_request(
|
||||||
|
first,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ConcurrencyLimitError):
|
||||||
|
claude_code_envelope.wrap_request(
|
||||||
|
second,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrap_request_releases_expired_sessions(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
scope_key = f"key:test-expire-{uuid.uuid4()}"
|
||||||
|
set_claude_code_request_context(
|
||||||
|
ClaudeCodeRequestContext(
|
||||||
|
is_stream=False,
|
||||||
|
scope_key=scope_key,
|
||||||
|
key_id="key-expire",
|
||||||
|
max_sessions=1,
|
||||||
|
session_idle_timeout_minutes=1,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
ticks = iter([0.0, 61.0])
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.provider.adapters.claude_code.envelope.time.monotonic",
|
||||||
|
lambda: next(ticks),
|
||||||
|
)
|
||||||
|
|
||||||
|
first = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_11111111-1111-1111-1111-111111111111"}
|
||||||
|
}
|
||||||
|
second = {
|
||||||
|
"metadata": {"user_id": "user_a_account_b_session_22222222-2222-2222-2222-222222222222"}
|
||||||
|
}
|
||||||
|
|
||||||
|
claude_code_envelope.wrap_request(
|
||||||
|
first,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
|
# 61s > 1min idle timeout,旧会话应过期,第二个会话可进入。
|
||||||
|
claude_code_envelope.wrap_request(
|
||||||
|
second,
|
||||||
|
model="claude-sonnet-4-5-20250929",
|
||||||
|
url_model=None,
|
||||||
|
decrypted_auth_config=None,
|
||||||
|
)
|
||||||
30
tests/services/test_provider_summary_claude_code_advanced.py
Normal file
30
tests/services/test_provider_summary_claude_code_advanced.py
Normal file
@@ -0,0 +1,30 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from src.api.admin.providers.summary import _extract_claude_code_advanced_from_config
|
||||||
|
from src.models.admin_requests import ClaudeCodeAdvancedConfig
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_claude_code_advanced_valid_dict() -> None:
|
||||||
|
config = {"claude_code_advanced": {"max_sessions": 12}}
|
||||||
|
|
||||||
|
parsed = _extract_claude_code_advanced_from_config(config, provider_id="provider-1")
|
||||||
|
|
||||||
|
assert isinstance(parsed, ClaudeCodeAdvancedConfig)
|
||||||
|
assert parsed.max_sessions == 12
|
||||||
|
assert parsed.session_idle_timeout_minutes == 5
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_claude_code_advanced_invalid_type_returns_none() -> None:
|
||||||
|
config = {"claude_code_advanced": "not-a-dict"}
|
||||||
|
|
||||||
|
parsed = _extract_claude_code_advanced_from_config(config, provider_id="provider-1")
|
||||||
|
|
||||||
|
assert parsed is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_claude_code_advanced_invalid_payload_returns_none() -> None:
|
||||||
|
config = {"claude_code_advanced": {"max_sessions": 0}}
|
||||||
|
|
||||||
|
parsed = _extract_claude_code_advanced_from_config(config, provider_id="provider-1")
|
||||||
|
|
||||||
|
assert parsed is None
|
||||||
84
tests/services/test_provider_transport_claude_code.py
Normal file
84
tests/services/test_provider_transport_claude_code.py
Normal file
@@ -0,0 +1,84 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
|
||||||
|
from src.services.provider.adapters.claude_code.plugin import register_all
|
||||||
|
from src.services.provider.transport import build_provider_url
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _DummyEndpoint:
|
||||||
|
base_url: str
|
||||||
|
api_format: str
|
||||||
|
custom_path: str | None = None
|
||||||
|
provider: object | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_claude_cli_uses_messages_path() -> None:
|
||||||
|
endpoint = _DummyEndpoint(
|
||||||
|
base_url="https://api.anthropic.com",
|
||||||
|
api_format="claude:cli",
|
||||||
|
provider=SimpleNamespace(provider_type="claude_code"),
|
||||||
|
)
|
||||||
|
|
||||||
|
url = build_provider_url(
|
||||||
|
endpoint, # type: ignore[arg-type]
|
||||||
|
path_params={"model": "ignored"},
|
||||||
|
is_stream=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert url == "https://api.anthropic.com/v1/messages"
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_claude_cli_does_not_duplicate_messages_suffix() -> None:
|
||||||
|
endpoint = _DummyEndpoint(
|
||||||
|
base_url="https://api.anthropic.com/v1/messages",
|
||||||
|
api_format="claude:cli",
|
||||||
|
provider=SimpleNamespace(provider_type="claude_code"),
|
||||||
|
)
|
||||||
|
|
||||||
|
url = build_provider_url(
|
||||||
|
endpoint, # type: ignore[arg-type]
|
||||||
|
path_params={"model": "ignored"},
|
||||||
|
is_stream=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert url == "https://api.anthropic.com/v1/messages"
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_code_claude_cli_appends_query_params() -> None:
|
||||||
|
endpoint = _DummyEndpoint(
|
||||||
|
base_url="https://api.anthropic.com/v1",
|
||||||
|
api_format="claude:cli",
|
||||||
|
provider=SimpleNamespace(provider_type="claude_code"),
|
||||||
|
)
|
||||||
|
|
||||||
|
url = build_provider_url(
|
||||||
|
endpoint, # type: ignore[arg-type]
|
||||||
|
query_params={"beta": "true"},
|
||||||
|
path_params={"model": "ignored"},
|
||||||
|
is_stream=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert url == "https://api.anthropic.com/v1/messages?beta=true"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_claude_code_registers_preset_model_fetcher() -> None:
|
||||||
|
register_all()
|
||||||
|
|
||||||
|
fetcher = UpstreamModelsFetcherRegistry.get("claude_code")
|
||||||
|
assert fetcher is not None
|
||||||
|
|
||||||
|
models, errors, has_success, upstream_metadata = await fetcher(SimpleNamespace(), 1.0)
|
||||||
|
model_ids = {m["id"] for m in models}
|
||||||
|
|
||||||
|
assert "claude-sonnet-4-5-20250929" in model_ids
|
||||||
|
assert "claude-opus-4-6" in model_ids
|
||||||
|
assert errors == []
|
||||||
|
assert has_success is True
|
||||||
|
assert upstream_metadata is None
|
||||||
Reference in New Issue
Block a user