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:
fawney19
2026-02-27 13:54:46 +08:00
parent f2f2a2dbc4
commit 579b5e4623
55 changed files with 3106 additions and 652 deletions

View File

@@ -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'

View File

@@ -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)

View File

@@ -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

View File

@@ -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>

View File

@@ -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(),

View File

@@ -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"

View File

@@ -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)"
> >

View File

@@ -37,11 +37,14 @@
<SelectValue placeholder="请选择" /> <SelectValue placeholder="请选择" />
</SelectTrigger> </SelectTrigger>
<SelectContent> <SelectContent>
<!-- 新建模式允许自定义CodexKiro 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) {

View File

@@ -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'

View File

@@ -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 ''

View File

@@ -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;

View File

@@ -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 &&

View File

@@ -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 },

View File

@@ -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',

View File

@@ -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)

View File

@@ -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:
# 失效账号允许覆盖 # 失效账号允许覆盖

View File

@@ -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} 的成本窗口"}

View File

@@ -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"])

View File

@@ -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 使用)

View File

@@ -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因为复用的客户端不应该被关闭

View File

@@ -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

View File

@@ -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 = (

View File

@@ -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,

View File

@@ -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)

View File

@@ -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 计数

View File

@@ -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",

View File

@@ -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(

View File

@@ -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":

View File

@@ -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="滚动成本窗口(秒)。默认 180005 小时)",
)
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 提前刷新秒数。默认 1803 分钟)",
)
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 请求"""

View File

@@ -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 数量")

View File

@@ -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():

View File

@@ -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,
) )
) )

View File

@@ -0,0 +1,3 @@
"""Claude Code provider adapter."""
__all__ = []

View 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",
]

View 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",
]

View 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``
异步接管)时仅做 maskingmasking 始终在会话限制检查之后执行,避免
用被伪装后的 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",
]

View 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"]

View File

@@ -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

View File

@@ -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()

View File

@@ -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": [
{ {

View File

@@ -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. 应用优先级模式排序 + 调度模式排序

View File

@@ -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)

View File

@@ -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

View File

@@ -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"

View File

@@ -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",
} }
) )

View File

@@ -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

View 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"

View File

@@ -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 = {

View 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,
)

View File

@@ -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,

View 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,
)

View 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"}]

View 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,
)

View 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

View 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