feat(vertex-ai): 重构 Vertex AI 为插件化 adapter,支持 service_account 认证与动态路由

将 Vertex AI 从 transport.py 的硬编码逻辑重构为独立的 plugin adapter,
支持 service_account/oauth 认证类型、模型格式自动识别、区域路由和 URL 构建。
前端新增 Key 认证类型选择和 Service Account 配置表单。

Co-authored-by: NyaDoo <65238336+NyaDoo@users.noreply.github.com>
Closes #194
This commit is contained in:
fawney19
2026-03-01 23:32:48 +08:00
parent a137601728
commit 4bf3a453e7
42 changed files with 1855 additions and 546 deletions

View File

@@ -0,0 +1,263 @@
"""vertex_ai_provider_type
Migrate legacy Vertex auth_type/provider_type into the new model:
- provider_type=vertex_ai
- auth_type=service_account (legacy vertex_ai renamed)
- fixed Vertex endpoints: gemini:chat + claude:chat
Revision ID: 2a624af8dd3a
Revises: 00b9161b8729
Create Date: 2026-02-28 15:00:00.000000+00:00
"""
from __future__ import annotations
import uuid
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "2a624af8dd3a"
down_revision = "00b9161b8729"
branch_labels = None
depends_on = None
_VERTEX_BASE_URL = "https://aiplatform.googleapis.com"
_VERTEX_ENDPOINTS: tuple[tuple[str, str, str], ...] = (
("gemini:chat", "gemini", "chat"),
("claude:chat", "claude", "chat"),
)
_VERTEX_KEY_FORMATS_SA = '["gemini:chat","claude:chat"]'
_VERTEX_KEY_FORMATS_API_KEY = '["gemini:chat"]'
def _select_vertex_provider_ids(conn: sa.Connection) -> list[str]:
"""Collect providers that should be treated as Vertex after migration."""
rows = conn.execute(sa.text("""
SELECT DISTINCT p.id
FROM providers p
LEFT JOIN provider_api_keys pak ON pak.provider_id = p.id
WHERE lower(COALESCE(p.provider_type, '')) = 'vertex_ai'
OR pak.auth_type = 'vertex_ai'
"""))
return [str(row[0]) for row in rows if row[0]]
def _ensure_fixed_vertex_endpoints(conn: sa.Connection, provider_ids: list[str]) -> None:
"""Ensure every Vertex provider has fixed gemini:chat + claude:chat endpoints."""
for provider_id in provider_ids:
provider_max_retries = (
conn.execute(
sa.text("""
SELECT COALESCE(max_retries, 2)
FROM providers
WHERE id = :provider_id
"""),
{"provider_id": provider_id},
).scalar()
or 2
)
for api_format, api_family, endpoint_kind in _VERTEX_ENDPOINTS:
# Normalize existing fixed endpoint fields.
conn.execute(
sa.text("""
UPDATE provider_endpoints
SET
api_family = :api_family,
endpoint_kind = :endpoint_kind,
base_url = :base_url,
custom_path = NULL,
is_active = TRUE,
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND api_format = :api_format
"""),
{
"provider_id": provider_id,
"api_format": api_format,
"api_family": api_family,
"endpoint_kind": endpoint_kind,
"base_url": _VERTEX_BASE_URL,
},
)
exists = conn.execute(
sa.text("""
SELECT 1
FROM provider_endpoints
WHERE provider_id = :provider_id
AND api_format = :api_format
LIMIT 1
"""),
{"provider_id": provider_id, "api_format": api_format},
).first()
if not exists:
conn.execute(
sa.text("""
INSERT INTO provider_endpoints (
id,
provider_id,
api_format,
api_family,
endpoint_kind,
base_url,
custom_path,
header_rules,
body_rules,
max_retries,
is_active,
config,
format_acceptance_config,
proxy,
created_at,
updated_at
)
VALUES (
:id,
:provider_id,
:api_format,
:api_family,
:endpoint_kind,
:base_url,
NULL,
NULL,
NULL,
:max_retries,
TRUE,
NULL,
NULL,
NULL,
CURRENT_TIMESTAMP,
CURRENT_TIMESTAMP
)
"""),
{
"id": str(uuid.uuid4()),
"provider_id": provider_id,
"api_format": api_format,
"api_family": api_family,
"endpoint_kind": endpoint_kind,
"base_url": _VERTEX_BASE_URL,
"max_retries": int(provider_max_retries),
},
)
# Vertex fixed-provider model: disable non-fixed endpoints.
conn.execute(
sa.text("""
UPDATE provider_endpoints
SET
is_active = FALSE,
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND api_format NOT IN ('gemini:chat', 'claude:chat')
"""),
{"provider_id": provider_id},
)
def _normalize_vertex_key_formats(conn: sa.Connection, provider_ids: list[str]) -> None:
"""Normalize key.api_formats for Vertex keys by auth type."""
for provider_id in provider_ids:
# Service Account (and legacy vertex_ai) keys: allow Gemini + Claude models.
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET
api_formats = CAST(:api_formats AS json),
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND auth_type IN ('service_account', 'vertex_ai')
"""),
{
"provider_id": provider_id,
"api_formats": _VERTEX_KEY_FORMATS_SA,
},
)
# API Key mode on Vertex 仅支持 GeminiGoogle publisher
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET
api_formats = CAST(:api_formats AS json),
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND auth_type = 'api_key'
"""),
{
"provider_id": provider_id,
"api_formats": _VERTEX_KEY_FORMATS_API_KEY,
},
)
def upgrade() -> None:
conn = op.get_bind()
# 1) 收集目标 Provider兼容重复执行先识别 legacy/new 两种来源)。
provider_ids = _select_vertex_provider_ids(conn)
# 2) 先重命名 auth_typelegacy vertex_ai -> service_account
conn.execute(sa.text("""
UPDATE provider_api_keys
SET auth_type = 'service_account'
WHERE auth_type = 'vertex_ai'
"""))
if not provider_ids:
return
# 3) 归一 provider_type并启用格式转换Vertex 同时承载 Gemini/Claude
for provider_id in provider_ids:
conn.execute(
sa.text("""
UPDATE providers
SET
provider_type = 'vertex_ai',
enable_format_conversion = TRUE
WHERE id = :provider_id
"""),
{"provider_id": provider_id},
)
# 4) 固定端点落地gemini:chat + claude:chat。
_ensure_fixed_vertex_endpoints(conn, provider_ids)
# 5) 归一 key 的 api_formats避免调度命中旧格式。
_normalize_vertex_key_formats(conn, provider_ids)
def downgrade() -> None:
conn = op.get_bind()
provider_rows = conn.execute(sa.text("""
SELECT id
FROM providers
WHERE lower(COALESCE(provider_type, '')) = 'vertex_ai'
"""))
provider_ids = [str(row[0]) for row in provider_rows if row[0]]
if provider_ids:
for provider_id in provider_ids:
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET auth_type = 'vertex_ai'
WHERE provider_id = :provider_id
AND auth_type = 'service_account'
"""),
{"provider_id": provider_id},
)
conn.execute(
sa.text("""
UPDATE providers
SET provider_type = 'custom'
WHERE id = :provider_id
"""),
{"provider_id": provider_id},
)

View File

@@ -56,7 +56,7 @@ export async function getModelCapabilities(modelName: string): Promise<ModelCapa
* 获取完整的 API Key用于查看和复制 * 获取完整的 API Key用于查看和复制
*/ */
export interface RevealKeyResult { export interface RevealKeyResult {
auth_type: 'api_key' | 'vertex_ai' | 'oauth' auth_type: 'api_key' | 'service_account' | 'oauth'
api_key?: string api_key?: string
refresh_token?: string refresh_token?: string
auth_config?: string | Record<string, unknown> auth_config?: string | Record<string, unknown>
@@ -119,7 +119,7 @@ export async function addProviderKey(
data: { data: {
api_formats: string[] // 支持的 API 格式列表(必填) api_formats: string[] // 支持的 API 格式列表(必填)
api_key: string api_key: string
auth_type?: 'api_key' | 'vertex_ai' | 'oauth' // 认证类型 auth_type?: 'api_key' | 'service_account' | 'oauth' // 认证类型
auth_config?: Record<string, unknown> // 认证配置Vertex AI Service Account JSON auth_config?: Record<string, unknown> // 认证配置Vertex AI Service Account JSON
name: string name: string
rate_multipliers?: Record<string, number> | null // 按 API 格式的成本倍率 rate_multipliers?: Record<string, number> | null // 按 API 格式的成本倍率
@@ -147,7 +147,7 @@ export async function updateProviderKey(
data: Partial<{ data: Partial<{
api_formats: string[] // 支持的 API 格式列表 api_formats: string[] // 支持的 API 格式列表
api_key: string api_key: string
auth_type: 'api_key' | 'vertex_ai' | 'oauth' // 认证类型 auth_type: 'api_key' | 'service_account' | 'oauth' // 认证类型
auth_config: Record<string, unknown> // 认证配置Vertex AI Service Account JSON auth_config: Record<string, unknown> // 认证配置Vertex AI Service Account JSON
name: string name: string
rate_multipliers: Record<string, number> | null // 按 API 格式的成本倍率 rate_multipliers: Record<string, number> | null // 按 API 格式的成本倍率

View File

@@ -30,7 +30,7 @@ export async function updateProvider(
providerId: string, providerId: string,
data: Partial<{ data: Partial<{
name: string name: string
provider_type: 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro' provider_type: 'custom' | 'vertex_ai' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro'
description: string | null description: string | null
website: string website: string
provider_priority: number provider_priority: number
@@ -62,7 +62,7 @@ export async function updateProvider(
export async function createProvider( export async function createProvider(
data: { data: {
name: string name: string
provider_type?: 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro' provider_type?: 'custom' | 'vertex_ai' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro'
description?: string description?: string
website?: string website?: string
billing_type?: 'monthly_quota' | 'pay_as_you_go' | 'free_tier' billing_type?: 'monthly_quota' | 'pay_as_you_go' | 'free_tier'

View File

@@ -204,7 +204,7 @@ export interface EndpointAPIKey {
api_formats: string[] // 支持的 endpoint signature 列表(如 "openai:chat" api_formats: string[] // 支持的 endpoint signature 列表(如 "openai:chat"
api_key_masked: string api_key_masked: string
api_key_plain?: string | null api_key_plain?: string | null
auth_type: 'api_key' | 'vertex_ai' | 'oauth' // 认证类型(必返回) auth_type: 'api_key' | 'service_account' | 'oauth' // 认证类型(必返回)
name: string // 密钥名称(必填,用于识别) name: string // 密钥名称(必填,用于识别)
rate_multipliers?: Record<string, number> | null // 按 endpoint signature 的成本倍率 rate_multipliers?: Record<string, number> | null // 按 endpoint signature 的成本倍率
internal_priority: number // Key 内部优先级 internal_priority: number // Key 内部优先级
@@ -347,7 +347,7 @@ export interface EndpointAPIKeyUpdate {
api_formats?: string[] // 支持的 API 格式列表 api_formats?: string[] // 支持的 API 格式列表
name?: string name?: string
api_key?: string // 仅在需要更新时提供 api_key?: string // 仅在需要更新时提供
auth_type?: 'api_key' | 'vertex_ai' | 'oauth' // 认证类型 auth_type?: 'api_key' | 'service_account' | 'oauth' // 认证类型
auth_config?: Record<string, unknown> // 认证配置Vertex AI Service Account JSON auth_config?: Record<string, unknown> // 认证配置Vertex AI Service Account JSON
rate_multipliers?: Record<string, number> | null // 按 API 格式的成本倍率 rate_multipliers?: Record<string, number> | null // 按 API 格式的成本倍率
internal_priority?: number internal_priority?: number
@@ -436,7 +436,7 @@ export interface PublicEndpointStatusMonitorResponse {
formats: PublicEndpointStatusMonitor[] formats: PublicEndpointStatusMonitor[]
} }
export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro' export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro' | 'vertex_ai'
export interface ClaudeCodeAdvancedConfig { export interface ClaudeCodeAdvancedConfig {
// 会话数量控制null/undefined 表示不限制 // 会话数量控制null/undefined 表示不限制

View File

@@ -13,7 +13,7 @@ export interface CandidateRecord {
key_id?: string key_id?: string
key_name?: string // 密钥名称 key_name?: string // 密钥名称
key_preview?: string // 密钥脱敏预览(如 sk-***abcOAuth 类型不返回 key_preview?: string // 密钥脱敏预览(如 sk-***abcOAuth 类型不返回
key_auth_type?: string // 密钥认证类型api_key, oauth, vertex_ai 等) key_auth_type?: string // 密钥认证类型api_key, service_account, oauth 等)
key_oauth_plan_type?: string // OAuth 账号套餐类型free/plus/team/enterprise key_oauth_plan_type?: string // OAuth 账号套餐类型free/plus/team/enterprise
key_capabilities?: Record<string, boolean> | null // Key 支持的能力 key_capabilities?: Record<string, boolean> | null // Key 支持的能力
required_capabilities?: Record<string, boolean> | null // 请求实际需要的能力标签 required_capabilities?: Record<string, boolean> | null // 请求实际需要的能力标签

View File

@@ -32,7 +32,7 @@
data-1p-ignore="true" data-1p-ignore="true"
/> />
</div> </div>
<div> <div v-if="providerType === 'vertex_ai'">
<Label :for="authTypeSelectId">认证类型</Label> <Label :for="authTypeSelectId">认证类型</Label>
<Select <Select
v-model="form.auth_type" v-model="form.auth_type"
@@ -44,8 +44,8 @@
<SelectItem value="api_key"> <SelectItem value="api_key">
API Key API Key
</SelectItem> </SelectItem>
<SelectItem value="vertex_ai"> <SelectItem value="service_account">
Vertex AI Service Account
</SelectItem> </SelectItem>
</SelectContent> </SelectContent>
</Select> </Select>
@@ -55,10 +55,10 @@
<!-- API 密钥 / Service Account JSON --> <!-- API 密钥 / Service Account JSON -->
<div> <div>
<Label :for="apiKeyInputId"> <Label :for="apiKeyInputId">
{{ form.auth_type === 'vertex_ai' ? 'Service Account JSON' : 'API 密钥' }} {{ form.auth_type === 'service_account' ? 'Service Account JSON' : 'API 密钥' }}
{{ editingKey ? '' : '*' }} {{ editingKey ? '' : '*' }}
</Label> </Label>
<template v-if="form.auth_type === 'vertex_ai'"> <template v-if="form.auth_type === 'service_account'">
<Textarea <Textarea
:id="apiKeyInputId" :id="apiKeyInputId"
v-model="form.auth_config_text" v-model="form.auth_config_text"
@@ -101,11 +101,11 @@
</div> </div>
<!-- API 格式选择 --> <!-- API 格式选择 -->
<div v-if="sortedApiFormats.length > 0"> <div v-if="visibleApiFormats.length > 0">
<Label class="mb-1.5 block">支持的 API 格式 *</Label> <Label class="mb-1.5 block">支持的 API 格式 *</Label>
<div class="grid grid-cols-2 gap-2"> <div class="grid grid-cols-2 gap-2">
<div <div
v-for="format in sortedApiFormats" v-for="format in visibleApiFormats"
:key="format" :key="format"
class="flex items-center justify-between rounded-md border px-2 py-1.5 transition-colors cursor-pointer" class="flex items-center justify-between rounded-md border px-2 py-1.5 transition-colors cursor-pointer"
:class="form.api_formats.includes(format) :class="form.api_formats.includes(format)
@@ -309,7 +309,7 @@
</template> </template>
<script setup lang="ts"> <script setup lang="ts">
import { ref, computed, onMounted } from 'vue' import { ref, computed, onMounted, watch } from 'vue'
import { Dialog, Button, Input, Label, Switch, Select, SelectTrigger, SelectValue, SelectContent, SelectItem, Textarea } from '@/components/ui' import { Dialog, Button, Input, Label, Switch, Select, SelectTrigger, SelectValue, SelectContent, SelectItem, Textarea } from '@/components/ui'
import { Key, SquarePen } from 'lucide-vue-next' import { Key, SquarePen } from 'lucide-vue-next'
import { useToast } from '@/composables/useToast' import { useToast } from '@/composables/useToast'
@@ -346,8 +346,26 @@ const emit = defineEmits<{
const { success, error: showError } = useToast() const { success, error: showError } = useToast()
// 排序后的可用 API 格式列表 function getVertexAllowedFormatsByAuth(authType: 'api_key' | 'service_account'): Set<string> {
const sortedApiFormats = computed(() => sortApiFormats(props.availableApiFormats)) if (authType === 'api_key') {
return new Set(['gemini:chat'])
}
return new Set(['gemini:chat', 'claude:chat'])
}
function normalizeApiFormat(format: string): string {
return String(format || '').trim().toLowerCase()
}
// 按 provider/auth_type 过滤后的可用 API 格式列表
const visibleApiFormats = computed(() => {
const sorted = sortApiFormats(props.availableApiFormats)
if (props.providerType !== 'vertex_ai') {
return sorted
}
const allowed = getVertexAllowedFormatsByAuth(form.value.auth_type)
return sorted.filter(fmt => allowed.has(normalizeApiFormat(fmt)))
})
// 默认认证类型 // 默认认证类型
const defaultAuthType = 'api_key' as const const defaultAuthType = 'api_key' as const
@@ -370,8 +388,8 @@ const showAutoFetchWarning = computed(() => {
// 检查是否正在切换认证类型 // 检查是否正在切换认证类型
const switchingToVertexAI = computed(() => const switchingToVertexAI = computed(() =>
!!props.editingKey && !!props.editingKey &&
props.editingKey.auth_type !== 'vertex_ai' && props.editingKey.auth_type !== 'service_account' &&
form.value.auth_type === 'vertex_ai' form.value.auth_type === 'service_account'
) )
const switchingToApiKey = computed(() => const switchingToApiKey = computed(() =>
!!props.editingKey && !!props.editingKey &&
@@ -386,7 +404,7 @@ const canSave = computed(() => {
// 新增模式下根据认证类型判断必填字段 // 新增模式下根据认证类型判断必填字段
if (!props.editingKey) { if (!props.editingKey) {
if (form.value.auth_type === 'api_key' && !form.value.api_key.trim()) return false if (form.value.auth_type === 'api_key' && !form.value.api_key.trim()) return false
if (form.value.auth_type === 'vertex_ai' && !form.value.auth_config_text.trim()) return false if (form.value.auth_type === 'service_account' && !form.value.auth_config_text.trim()) return false
} else { } else {
// 编辑模式下切换认证类型时,必须填写对应字段 // 编辑模式下切换认证类型时,必须填写对应字段
if (switchingToApiKey.value && !form.value.api_key.trim()) return false if (switchingToApiKey.value && !form.value.api_key.trim()) return false
@@ -409,15 +427,13 @@ const apiKeyFieldName = computed(() => `api-key-field-${formNonce.value}`)
// 可用的能力列表 // 可用的能力列表
const availableCapabilities = ref<CapabilityDefinition[]>([]) const availableCapabilities = ref<CapabilityDefinition[]>([])
// 非 custom 提供商默认开启自动获取上游模型 // 新增密钥时默认不自动开启上游模型获取
const defaultAutoFetchModels = computed(() => const defaultAutoFetchModels = computed(() => false)
!!props.providerType && props.providerType !== 'custom'
)
const form = ref({ const form = ref({
name: '', name: '',
api_key: '', // 标准 API Key api_key: '', // 标准 API Key
auth_type: 'api_key' as 'api_key' | 'vertex_ai', // 认证类型 auth_type: 'api_key' as 'api_key' | 'service_account', // 认证类型
auth_config_text: '', // Service Account JSON 文本(用于表单输入) auth_config_text: '', // Service Account JSON 文本(用于表单输入)
api_formats: [] as string[], // 支持的 API 格式列表 api_formats: [] as string[], // 支持的 API 格式列表
rate_multipliers: {} as Record<string, number>, // 按 API 格式的成本倍率 rate_multipliers: {} as Record<string, number>, // 按 API 格式的成本倍率
@@ -433,6 +449,21 @@ const form = ref({
model_exclude_patterns_text: '' // 排除规则文本(逗号分隔) model_exclude_patterns_text: '' // 排除规则文本(逗号分隔)
}) })
watch(
[() => form.value.auth_type, () => props.providerType, () => props.availableApiFormats],
() => {
if (props.providerType !== 'vertex_ai') {
return
}
const allowed = getVertexAllowedFormatsByAuth(form.value.auth_type)
const filtered = form.value.api_formats.filter(fmt => allowed.has(normalizeApiFormat(fmt)))
if (filtered.length !== form.value.api_formats.length) {
form.value.api_formats = [...filtered]
}
},
{ immediate: true }
)
// 加载能力列表 // 加载能力列表
async function loadCapabilities() { async function loadCapabilities() {
try { try {
@@ -517,7 +548,7 @@ function loadKeyData() {
form.value = { form.value = {
name: props.editingKey.name, name: props.editingKey.name,
api_key: '', api_key: '',
auth_type: props.editingKey.auth_type === 'vertex_ai' ? 'vertex_ai' : 'api_key', auth_type: props.editingKey.auth_type === 'service_account' ? 'service_account' : 'api_key',
auth_config_text: '', // auth_config 不返回给前端,编辑时需要重新输入 auth_config_text: '', // auth_config 不返回给前端,编辑时需要重新输入
api_formats: props.editingKey.api_formats?.length > 0 api_formats: props.editingKey.api_formats?.length > 0
? [...props.editingKey.api_formats] ? [...props.editingKey.api_formats]
@@ -564,7 +595,7 @@ function parsePatternText(text: string): string[] {
// 解析 Service Account JSON 文本 // 解析 Service Account JSON 文本
function parseAuthConfig(): Record<string, unknown> | null { function parseAuthConfig(): Record<string, unknown> | null {
if (form.value.auth_type !== 'vertex_ai') return null if (form.value.auth_type !== 'service_account') return null
const text = form.value.auth_config_text.trim() const text = form.value.auth_config_text.trim()
if (!text) return null if (!text) return null
try { try {
@@ -588,7 +619,7 @@ async function handleSave() {
showError('请输入 API 密钥', '验证失败') showError('请输入 API 密钥', '验证失败')
return return
} }
} else if (form.value.auth_type === 'vertex_ai') { } else if (form.value.auth_type === 'service_account') {
if (!props.editingKey && !form.value.auth_config_text.trim()) { if (!props.editingKey && !form.value.auth_config_text.trim()) {
showError('请输入 Service Account JSON', '验证失败') showError('请输入 Service Account JSON', '验证失败')
return return
@@ -664,7 +695,7 @@ async function handleSave() {
if (form.value.auth_type === 'api_key' && form.value.api_key.trim()) { if (form.value.auth_type === 'api_key' && form.value.api_key.trim()) {
updateData.api_key = form.value.api_key updateData.api_key = form.value.api_key
} }
if (form.value.auth_type === 'vertex_ai' && authConfig) { if (form.value.auth_type === 'service_account' && authConfig) {
updateData.auth_config = authConfig updateData.auth_config = authConfig
} }

View File

@@ -224,7 +224,7 @@
<div class="p-4 border-b border-border/60"> <div class="p-4 border-b border-border/60">
<div class="flex items-center justify-between"> <div class="flex items-center justify-between">
<h3 class="text-sm font-semibold"> <h3 class="text-sm font-semibold">
{{ provider.provider_type === 'custom' ? '密钥管理' : '账号管理' }} {{ isKeyManagedProviderType(provider.provider_type) ? '密钥管理' : '账号管理' }}
</h3> </h3>
<div class="flex items-center gap-2"> <div class="flex items-center gap-2">
<Button <Button
@@ -235,7 +235,7 @@
@click="handleAddKeyToFirstEndpoint" @click="handleAddKeyToFirstEndpoint"
> >
<Plus class="w-3.5 h-3.5 mr-1.5" /> <Plus class="w-3.5 h-3.5 mr-1.5" />
{{ provider.provider_type === 'custom' ? '添加密钥' : '添加账号' }} {{ isKeyManagedProviderType(provider.provider_type) ? '添加密钥' : '添加账号' }}
</Button> </Button>
</div> </div>
</div> </div>
@@ -301,7 +301,7 @@
</div> </div>
<div class="flex items-center gap-1"> <div class="flex items-center gap-1">
<span class="text-[11px] font-mono text-muted-foreground"> <span class="text-[11px] font-mono text-muted-foreground">
{{ key.auth_type === 'oauth' ? '[Refresh Token]' : (key.auth_type === 'vertex_ai' ? 'Vertex AI' : key.api_key_masked) }} {{ key.auth_type === 'oauth' ? '[Refresh Token]' : (key.auth_type === 'service_account' ? '[Service Account]' : key.api_key_masked) }}
</span> </span>
<Button <Button
v-if="key.auth_type === 'oauth'" v-if="key.auth_type === 'oauth'"
@@ -844,7 +844,7 @@
v-if="shouldPaginateKeys" v-if="shouldPaginateKeys"
class="px-4 py-2 flex items-center justify-between text-xs text-muted-foreground mt-auto" class="px-4 py-2 flex items-center justify-between text-xs text-muted-foreground mt-auto"
> >
<span>共 {{ allKeys.length }} 个{{ provider.provider_type === 'custom' ? '密钥' : '账号' }}</span> <span>共 {{ allKeys.length }} 个{{ isKeyManagedProviderType(provider.provider_type) ? '密钥' : '账号' }}</span>
<div class="flex items-center gap-1.5"> <div class="flex items-center gap-1.5">
<Button <Button
variant="ghost" variant="ghost"
@@ -876,11 +876,11 @@
> >
<Key class="w-12 h-12 mx-auto mb-3 opacity-50" /> <Key class="w-12 h-12 mx-auto mb-3 opacity-50" />
<p class="text-sm"> <p class="text-sm">
{{ provider.provider_type === 'custom' ? '暂无密钥配置' : '暂无账号配置' }} {{ isKeyManagedProviderType(provider.provider_type) ? '暂无密钥配置' : '暂无账号配置' }}
</p> </p>
<p class="text-xs mt-1"> <p class="text-xs mt-1">
{{ endpoints.length > 0 {{ endpoints.length > 0
? (provider.provider_type === 'custom' ? '点击上方"添加密钥"按钮创建第一个密钥' : '点击上方"添加账号"按钮添加第一个账号') ? (isKeyManagedProviderType(provider.provider_type) ? '点击上方"添加密钥"按钮创建第一个密钥' : '点击上方"添加账号"按钮添加第一个账号')
: '请先添加端点,然后再添加密钥' }} : '请先添加端点,然后再添加密钥' }}
</p> </p>
</div> </div>
@@ -1100,6 +1100,7 @@ import {
} from '@/api/endpoints' } from '@/api/endpoints'
import type { UpstreamMetadata, AntigravityModelQuota } from '@/api/endpoints/types' import type { UpstreamMetadata, AntigravityModelQuota } from '@/api/endpoints/types'
import { formatApiFormat } from '@/api/endpoints/types/api-format' import { formatApiFormat } from '@/api/endpoints/types/api-format'
import { isOAuthAccountProviderType, isKeyManagedProviderType } from '../utils/providerTypeUtils'
// 扩展端点类型,包含密钥列表 // 扩展端点类型,包含密钥列表
interface ProviderEndpointWithKeys extends ProviderEndpoint { interface ProviderEndpointWithKeys extends ProviderEndpoint {
@@ -1417,11 +1418,11 @@ function handleAddKey(endpoint: ProviderEndpoint) {
function handleAddKeyToFirstEndpoint() { function handleAddKeyToFirstEndpoint() {
if (endpoints.value.length === 0) return if (endpoints.value.length === 0) return
// 非自定义提供商:打开 OAuth 账号对话框 // OAuth 账号型提供商:打开 OAuth 账号对话框
if (provider.value?.provider_type !== 'custom') { if (isOAuthAccountProviderType(provider.value?.provider_type)) {
oauthAccountDialogOpen.value = true oauthAccountDialogOpen.value = true
} else { } else {
// 自定义提供商:打开密钥表单对话框 // 密钥型提供商custom/vertex_ai:打开密钥表单对话框
handleAddKey(endpoints.value[0]) handleAddKey(endpoints.value[0])
} }
} }
@@ -1455,8 +1456,8 @@ async function copyFullKey(key: EndpointAPIKey) {
const result = await revealEndpointKey(key.id) const result = await revealEndpointKey(key.id)
let textToCopy: string let textToCopy: string
if (result.auth_type === 'vertex_ai' && result.auth_config) { if (result.auth_type === 'service_account' && result.auth_config) {
// Vertex AI 类型:复制 auth_config JSON // Service Account 类型:复制 auth_config JSON
textToCopy = typeof result.auth_config === 'string' textToCopy = typeof result.auth_config === 'string'
? result.auth_config ? result.auth_config
: JSON.stringify(result.auth_config, null, 2) : JSON.stringify(result.auth_config, null, 2)

View File

@@ -42,6 +42,9 @@
<SelectItem value="custom"> <SelectItem value="custom">
自定义 自定义
</SelectItem> </SelectItem>
<SelectItem value="vertex_ai">
Vertex AI
</SelectItem>
<SelectItem value="claude_code"> <SelectItem value="claude_code">
ClaudeCode ClaudeCode
</SelectItem> </SelectItem>
@@ -60,6 +63,9 @@
<SelectItem value="custom"> <SelectItem value="custom">
自定义 自定义
</SelectItem> </SelectItem>
<SelectItem value="vertex_ai">
Vertex AI
</SelectItem>
<SelectItem value="claude_code"> <SelectItem value="claude_code">
ClaudeCode ClaudeCode
</SelectItem> </SelectItem>
@@ -322,7 +328,7 @@ const defaultPriority = computed(() => {
// 表单数据 // 表单数据
const form = ref({ const form = ref({
name: '', name: '',
provider_type: 'custom' as 'custom' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro', provider_type: 'custom' as 'custom' | 'vertex_ai' | 'claude_code' | 'codex' | 'gemini_cli' | 'antigravity' | 'kiro',
description: '', description: '',
website: '', website: '',
// 计费配置 // 计费配置
@@ -414,10 +420,10 @@ const { isEditMode, handleDialogUpdate, handleCancel } = useFormDialog({
resetForm, resetForm,
}) })
// 新建模式下切换 provider_type 时自动设置号池模式:非自定义类型默认开启 // 新建模式下切换 provider_type 时自动开启号池模式
watch(() => form.value.provider_type, (newType) => { watch(() => form.value.provider_type, () => {
if (!isEditMode.value) { if (!isEditMode.value) {
form.value.pool_mode_enabled = newType !== 'custom' form.value.pool_mode_enabled = false
} }
}) })

View File

@@ -239,6 +239,7 @@ import Badge from '@/components/ui/badge.vue'
import { type ProviderWithEndpointsSummary, API_FORMAT_SHORT } from '@/api/endpoints' import { type ProviderWithEndpointsSummary, API_FORMAT_SHORT } from '@/api/endpoints'
import { formatBillingType } from '@/utils/format' import { formatBillingType } from '@/utils/format'
import { sortEndpoints, isEndpointAvailable, getEndpointDotColor, getEndpointTooltip } from '@/features/providers/composables/useEndpointStatus' import { sortEndpoints, isEndpointAvailable, getEndpointDotColor, getEndpointTooltip } from '@/features/providers/composables/useEndpointStatus'
import { isKeyManagedProviderType } from '../utils/providerTypeUtils'
const props = defineProps<{ const props = defineProps<{
provider: ProviderWithEndpointsSummary provider: ProviderWithEndpointsSummary
@@ -304,7 +305,6 @@ function handleDescriptionKeydown(event: KeyboardEvent) {
} }
function getCredentialLabel(provider: ProviderWithEndpointsSummary): '账号' | '密钥' { function getCredentialLabel(provider: ProviderWithEndpointsSummary): '账号' | '密钥' {
const providerType = String(provider.provider_type || '').trim().toLowerCase() return isKeyManagedProviderType(provider.provider_type) ? '密钥' : '账号'
return providerType && providerType !== 'custom' ? '账号' : '密钥'
} }
</script> </script>

View File

@@ -775,14 +775,14 @@ const isCapabilityUsed = (cap: string): boolean => {
// 判断是否为 OAuth 类型provider_type 为具体值时也算 OAuth // 判断是否为 OAuth 类型provider_type 为具体值时也算 OAuth
const isOAuthType = (authType?: string): boolean => { const isOAuthType = (authType?: string): boolean => {
if (!authType) return false if (!authType) return false
return !['api_key', 'vertex_ai'].includes(authType) return !['api_key', 'service_account'].includes(authType)
} }
// 格式化认证类型(合并 plan 信息,避免冗余) // 格式化认证类型(合并 plan 信息,避免冗余)
const formatAuthTypeWithPlan = (authType: string, planType?: string): string => { const formatAuthTypeWithPlan = (authType: string, planType?: string): string => {
const labels: Record<string, string> = { const labels: Record<string, string> = {
'oauth': 'OAuth', 'oauth': 'OAuth',
'vertex_ai': 'Vertex AI', 'service_account': 'Service Account',
'kiro': 'Kiro', 'kiro': 'Kiro',
'codex': 'Codex', 'codex': 'Codex',
'antigravity': 'Antigravity', 'antigravity': 'Antigravity',

View File

@@ -131,17 +131,22 @@ export class GeminiParser implements ApiFormatParser {
apiFormat: 'gemini', apiFormat: 'gemini',
} }
// Gemini 响应格式: { candidates: [{ content: { parts: [...] } }] } // Gemini 原生响应格式: { candidates: [{ content: { parts: [...] } }] }
const candidates = body.candidates as RawObject[] | undefined const candidateParts = this.getCandidateParts(body)
const candidate = candidates?.[0] as RawObject | undefined if (candidateParts.length > 0) {
const candidateContent = candidate?.content as RawObject | undefined const contentBlocks = this.parseParts(candidateParts)
if (candidateContent?.parts && Array.isArray(candidateContent.parts)) {
const contentBlocks = this.parseParts(candidateContent.parts as RawObject[])
if (contentBlocks.length > 0) { if (contentBlocks.length > 0) {
result.messages.push(createMessage('assistant', contentBlocks)) result.messages.push(createMessage('assistant', contentBlocks))
return result
} }
} }
// 统一响应格式(例如 status/output/output_text
const unifiedMessages = this.parseUnifiedOutputMessages(body)
if (unifiedMessages.length > 0) {
result.messages.push(...unifiedMessages)
}
return result return result
} catch (e) { } catch (e) {
return createEmptyConversation('gemini', `解析失败: ${e}`) return createEmptyConversation('gemini', `解析失败: ${e}`)
@@ -363,12 +368,10 @@ export class GeminiParser implements ApiFormatParser {
const body = responseBody as RawObject const body = responseBody as RawObject
const blocks: RenderBlock[] = [] const blocks: RenderBlock[] = []
// Gemini 响应格式: { candidates: [{ content: { parts: [...] } }] } // Gemini 原生响应格式: { candidates: [{ content: { parts: [...] } }] }
const candidates = body.candidates as RawObject[] | undefined const candidateParts = this.getCandidateParts(body)
const candidate = candidates?.[0] as RawObject | undefined if (candidateParts.length > 0) {
const candidateContent = candidate?.content as RawObject | undefined const parts = candidateParts
if (candidateContent?.parts && Array.isArray(candidateContent.parts)) {
const parts = candidateContent.parts as RawObject[]
const contentBlocks = this.renderParts(parts) const contentBlocks = this.renderParts(parts)
if (contentBlocks.length > 0) { if (contentBlocks.length > 0) {
const badges = this.getBadgesForParts(parts) const badges = this.getBadgesForParts(parts)
@@ -376,9 +379,13 @@ export class GeminiParser implements ApiFormatParser {
roleLabel: 'Assistant', roleLabel: 'Assistant',
badges: badges.length > 0 ? badges : undefined, badges: badges.length > 0 ? badges : undefined,
})) }))
return { blocks, isStream: false }
} }
} }
// 统一响应格式(例如 status/output/output_text
blocks.push(...this.renderUnifiedOutputMessages(body))
return { blocks, isStream: false } return { blocks, isStream: false }
} catch (e) { } catch (e) {
return createEmptyRenderResult(`渲染失败: ${e}`) return createEmptyRenderResult(`渲染失败: ${e}`)
@@ -582,6 +589,117 @@ export class GeminiParser implements ApiFormatParser {
return badges return badges
} }
/**
* 提取 Gemini candidates 中的 parts
*/
private getCandidateParts(body: RawObject): RawObject[] {
const candidates = body.candidates as RawObject[] | undefined
const candidate = candidates?.[0] as RawObject | undefined
const candidateContent = candidate?.content as RawObject | undefined
if (Array.isArray(candidateContent?.parts)) {
return candidateContent.parts as RawObject[]
}
return []
}
/**
* 解析统一响应格式中的 output 消息
*/
private parseUnifiedOutputMessages(body: RawObject): ParsedMessage[] {
const output = body.output
if (!Array.isArray(output)) {
return []
}
const messages: ParsedMessage[] = []
for (const rawItem of output) {
const item = rawItem as RawObject
if (item?.type !== 'message') {
continue
}
const texts = this.extractUnifiedOutputTexts(item.content)
if (texts.length === 0) {
continue
}
messages.push(createMessage(
this.mapUnifiedOutputRole(item.role),
texts.map(text => createTextBlock(text))
))
}
return messages
}
/**
* 渲染统一响应格式中的 output 消息
*/
private renderUnifiedOutputMessages(body: RawObject): RenderBlock[] {
const output = body.output
if (!Array.isArray(output)) {
return []
}
const blocks: RenderBlock[] = []
for (const rawItem of output) {
const item = rawItem as RawObject
if (item?.type !== 'message') {
continue
}
const texts = this.extractUnifiedOutputTexts(item.content)
if (texts.length === 0) {
continue
}
const role = this.mapUnifiedOutputRole(item.role)
blocks.push(createMessageBlock(
role,
texts.map(text => createTextRenderBlock(text)),
{ roleLabel: this.getRoleLabel(role) }
))
}
return blocks
}
/**
* 提取统一响应格式的文本内容
*/
private extractUnifiedOutputTexts(content: unknown): string[] {
if (typeof content === 'string') {
return content ? [content] : []
}
if (!Array.isArray(content)) {
return []
}
const texts: string[] = []
for (const rawPart of content) {
const part = rawPart as RawObject
if (
(part.type === 'output_text' || part.type === 'text' || part.type === 'input_text') &&
typeof part.text === 'string'
) {
texts.push(part.text)
}
}
return texts
}
/**
* 映射统一 output 结构的角色
*/
private mapUnifiedOutputRole(role: unknown): MessageRole {
if (role === 'assistant' || role === 'model') return 'assistant'
if (role === 'user') return 'user'
if (role === 'system') return 'system'
if (role === 'tool') return 'tool'
return 'assistant'
}
/** /**
* 获取已解析内容的徽章 * 获取已解析内容的徽章
*/ */

View File

@@ -21,6 +21,7 @@ from src.api.base.pipeline import ApiRequestPipeline
from src.core.api_format.signature import parse_signature_key from src.core.api_format.signature import parse_signature_key
from src.core.exceptions import InvalidRequestException, NotFoundException from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger from src.core.logger import logger
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
from src.core.provider_types import ProviderType from src.core.provider_types import ProviderType
from src.database import get_db from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
@@ -29,6 +30,7 @@ from src.models.endpoint_models import (
ProviderEndpointResponse, ProviderEndpointResponse,
ProviderEndpointUpdate, ProviderEndpointUpdate,
) )
from src.services.provider.stream_policy import UpstreamStreamPolicy, parse_upstream_stream_policy
router = APIRouter(tags=["Endpoint Management"]) router = APIRouter(tags=["Endpoint Management"])
pipeline = ApiRequestPipeline() pipeline = ApiRequestPipeline()
@@ -44,6 +46,17 @@ def mask_proxy_password(proxy_config: dict | None) -> dict | None:
return masked return masked
def _is_fixed_provider(provider_type: str | None) -> bool:
"""Whether this provider_type is managed by fixed-provider templates."""
normalized = (provider_type or "custom").strip().lower()
if normalized == ProviderType.CUSTOM.value:
return False
try:
return ProviderType(normalized) in FIXED_PROVIDERS
except Exception:
return False
@router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse]) @router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse])
async def list_provider_endpoints( async def list_provider_endpoints(
provider_id: str, provider_id: str,
@@ -283,8 +296,8 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
raise NotFoundException(f"Provider {self.provider_id} 不存在") raise NotFoundException(f"Provider {self.provider_id} 不存在")
# 固定类型 Provider禁止通过该接口新增 Endpoints端点由模板自动创建并锁定 # 固定类型 Provider禁止通过该接口新增 Endpoints端点由模板自动创建并锁定
provider_type = (getattr(provider, "provider_type", "custom") or "custom").strip() provider_type = getattr(provider, "provider_type", "custom")
if provider_type != ProviderType.CUSTOM: if _is_fixed_provider(provider_type):
raise InvalidRequestException("固定类型 Provider 不允许手动新增 Endpoint") raise InvalidRequestException("固定类型 Provider 不允许手动新增 Endpoint")
if self.endpoint_data.provider_id != self.provider_id: if self.endpoint_data.provider_id != self.provider_id:
@@ -424,12 +437,43 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
# 固定类型 Provider 的 endpoint锁定 base_url/custom_path前端禁用仅是 UX后端必须强校验 # 固定类型 Provider 的 endpoint锁定 base_url/custom_path前端禁用仅是 UX后端必须强校验
provider = db.query(Provider).filter(Provider.id == endpoint.provider_id).first() provider = db.query(Provider).filter(Provider.id == endpoint.provider_id).first()
if provider: if provider:
provider_type = (getattr(provider, "provider_type", "custom") or "custom").strip() provider_type = getattr(provider, "provider_type", "custom")
if provider_type != ProviderType.CUSTOM: if _is_fixed_provider(provider_type):
if "base_url" in update_data or "custom_path" in update_data: if "base_url" in update_data or "custom_path" in update_data:
raise InvalidRequestException( raise InvalidRequestException(
"固定类型 Provider 的 Endpoint 不允许修改 base_url/custom_path" "固定类型 Provider 的 Endpoint 不允许修改 base_url/custom_path"
) )
normalized_provider_type = str(provider_type or "custom").strip().lower()
endpoint_sig = str(getattr(endpoint, "api_format", "") or "").strip().lower()
if (
normalized_provider_type == ProviderType.CODEX.value
and endpoint_sig == "openai:cli"
):
has_config_in_payload = "config" in update_data
cfg_payload = (
update_data.get("config")
if has_config_in_payload
else getattr(endpoint, "config", None)
)
cfg = dict(cfg_payload) if isinstance(cfg_payload, dict) else {}
requested = (
cfg.get("upstream_stream_policy")
or cfg.get("upstreamStreamPolicy")
or cfg.get("upstream_stream")
)
if (
has_config_in_payload
and requested is not None
and parse_upstream_stream_policy(requested)
!= UpstreamStreamPolicy.FORCE_STREAM
):
raise InvalidRequestException(
"Codex OpenAI CLI 端点固定为强制流式,不允许修改"
)
cfg.pop("upstreamStreamPolicy", None)
cfg.pop("upstream_stream", None)
cfg["upstream_stream_policy"] = "force_stream"
update_data["config"] = cfg
# 把 proxy 转换为 dict 存储,支持显式设置为 None 清除代理 # 把 proxy 转换为 dict 存储,支持显式设置为 None 清除代理
if "proxy" in update_data: if "proxy" in update_data:

View File

@@ -40,7 +40,7 @@ class CandidateResponse(BaseModel):
key_id: str | None = None key_id: str | None = None
key_name: str | None = None # 密钥名称 key_name: str | None = None # 密钥名称
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abcOAuth 类型不返回 key_preview: str | None = None # 密钥脱敏预览(如 sk-***abcOAuth 类型不返回
key_auth_type: str | None = None # 密钥认证类型api_key, oauth, vertex_ai 等 key_auth_type: str | None = None # 密钥认证类型api_key, service_account, oauth
key_oauth_plan_type: str | None = None # OAuth 账号套餐类型free/plus/team/enterprise key_oauth_plan_type: str | None = None # OAuth 账号套餐类型free/plus/team/enterprise
key_capabilities: dict | None = None # Key 支持的能力 key_capabilities: dict | None = None # Key 支持的能力
required_capabilities: dict | None = None # 请求实际需要的能力标签 required_capabilities: dict | None = None # 请求实际需要的能力标签

View File

@@ -160,13 +160,42 @@ class ProviderCompleteOAuthResponse(BaseModel):
# ============================================================================== # ==============================================================================
def _get_fixed_template(provider_type: str) -> Any | None:
try:
return FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
return None
def _supports_oauth(template: Any | None) -> bool:
if not template:
return False
oauth = getattr(template, "oauth", None)
if oauth is None:
return False
return bool(
str(getattr(oauth, "authorize_url", "") or "").strip()
and str(getattr(oauth, "token_url", "") or "").strip()
and str(getattr(oauth, "client_id", "") or "").strip()
)
def _require_fixed_provider(provider: Provider) -> str: def _require_fixed_provider(provider: Provider) -> str:
provider_type = (getattr(provider, "provider_type", "custom") or "custom").strip() provider_type = str(getattr(provider, "provider_type", "custom") or "custom").strip().lower()
if provider_type == ProviderType.CUSTOM: if not _get_fixed_template(provider_type):
raise InvalidRequestException("该 Provider 不是固定类型,无法使用 provider-oauth") raise InvalidRequestException("该 Provider 不是固定类型,无法使用 provider-oauth")
return provider_type return provider_type
def _require_oauth_template(provider_type: str) -> Any:
template = _get_fixed_template(provider_type)
if not template:
raise InvalidRequestException("不支持的 provider_type")
if not _supports_oauth(template):
raise InvalidRequestException("该 Provider 不支持 OAuth 授权")
return template
def _resolve_proxy_for_oauth( def _resolve_proxy_for_oauth(
provider_proxy: dict[str, Any] | None, provider_proxy: dict[str, Any] | None,
proxy_node_id: str | None, proxy_node_id: str | None,
@@ -239,7 +268,7 @@ def _create_oauth_key(
api_formats: list[str], api_formats: list[str],
flush_only: bool = False, flush_only: bool = False,
proxy: dict[str, Any] | None = None, proxy: dict[str, Any] | None = None,
auto_fetch_models: bool = True, auto_fetch_models: bool = False,
) -> "ProviderAPIKey": ) -> "ProviderAPIKey":
"""创建 OAuth Key 记录并持久化。 """创建 OAuth Key 记录并持久化。
@@ -247,7 +276,7 @@ def _create_oauth_key(
flush_only: True 时仅 flush批量导入场景False 时 commit + refresh。 flush_only: True 时仅 flush批量导入场景False 时 commit + refresh。
proxy: Key 级别代理配置(如 {"node_id": "xxx", "enabled": True} proxy: Key 级别代理配置(如 {"node_id": "xxx", "enabled": True}
创建时设置后,后续 token 刷新、额度刷新等操作立即走代理,避免 IP 污染。 创建时设置后,后续 token 刷新、额度刷新等操作立即走代理,避免 IP 污染。
auto_fetch_models: 是否启用自动获取上游模型,非 custom 提供商默认开启 auto_fetch_models: 是否启用自动获取上游模型,默认关闭
""" """
from src.models.database import ProviderAPIKey as ProviderAPIKeyModel from src.models.database import ProviderAPIKey as ProviderAPIKeyModel
@@ -301,24 +330,6 @@ def _update_existing_oauth_key(
return existing_key return existing_key
async def _trigger_auto_fetch_models(key_ids: list[str]) -> None:
"""为启用了 auto_fetch_models 的新建 Key 触发模型获取。"""
if not key_ids:
return
try:
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
scheduler = get_model_fetch_scheduler()
for key_id in key_ids:
logger.info("[AUTO_FETCH] OAuth Key {} 默认开启自动获取模型,触发模型获取", key_id)
try:
await scheduler._fetch_models_for_key_by_id(key_id)
except Exception as e:
logger.error(f"[AUTO_FETCH] Key {key_id} 触发模型获取失败: {e}")
except Exception as e:
logger.error(f"[AUTO_FETCH] 获取 ModelFetchScheduler 失败: {e}")
async def _fetch_kiro_email( async def _fetch_kiro_email(
auth_config: dict[str, Any], auth_config: dict[str, Any],
proxy_config: dict[str, Any] | None = None, proxy_config: dict[str, Any] | None = None,
@@ -517,6 +528,8 @@ async def supported_types(_: User = Depends(require_admin)) -> list[dict[str, An
# 不返回 client_secret # 不返回 client_secret
result: list[dict[str, Any]] = [] result: list[dict[str, Any]] = []
for provider_type, template in FIXED_PROVIDERS.items(): for provider_type, template in FIXED_PROVIDERS.items():
if not _supports_oauth(template):
continue
result.append( result.append(
{ {
"provider_type": ( "provider_type": (
@@ -553,12 +566,7 @@ async def start_oauth(
raise NotFoundException("Provider 不存在", "provider") raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider) provider_type = _require_fixed_provider(provider)
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException("不支持的 provider_type")
redis = await get_redis_client(require_redis=True) redis = await get_redis_client(require_redis=True)
assert redis is not None assert redis is not None
@@ -646,12 +654,7 @@ async def complete_oauth(
raise NotFoundException("Provider 不存在", "provider") raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider) provider_type = _require_fixed_provider(provider)
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException("不支持的 provider_type")
# exchange token # exchange token
token_url = template.oauth.token_url token_url = template.oauth.token_url
@@ -815,12 +818,7 @@ async def refresh_oauth(
email=None, email=None,
) )
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException("不支持的 provider_type")
encrypted_auth_config = getattr(key, "auth_config", None) encrypted_auth_config = getattr(key, "auth_config", None)
if not encrypted_auth_config: if not encrypted_auth_config:
@@ -983,12 +981,7 @@ async def start_provider_oauth(
if provider_type == ProviderType.KIRO.value: if provider_type == ProviderType.KIRO.value:
raise InvalidRequestException("Kiro 不支持 OAuth 授权,请使用导入授权。") raise InvalidRequestException("Kiro 不支持 OAuth 授权,请使用导入授权。")
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException("不支持的 provider_type")
redis = await get_redis_client(require_redis=True) redis = await get_redis_client(require_redis=True)
assert redis is not None assert redis is not None
@@ -1073,12 +1066,7 @@ async def complete_provider_oauth(
if provider_type == ProviderType.KIRO.value: if provider_type == ProviderType.KIRO.value:
raise InvalidRequestException("Kiro 不支持 OAuth 授权,请使用导入授权。") raise InvalidRequestException("Kiro 不支持 OAuth 授权,请使用导入授权。")
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException("不支持的 provider_type")
# exchange token # exchange token
token_url = template.oauth.token_url token_url = template.oauth.token_url
@@ -1190,9 +1178,6 @@ async def complete_provider_oauth(
proxy=key_proxy, proxy=key_proxy,
) )
# 默认开启了 auto_fetch_models触发模型获取
await _trigger_auto_fetch_models([str(new_key.id)])
return ProviderCompleteOAuthResponse( return ProviderCompleteOAuthResponse(
key_id=str(new_key.id), key_id=str(new_key.id),
provider_type=provider_type, provider_type=provider_type,
@@ -1441,9 +1426,6 @@ async def import_refresh_token(
proxy=key_proxy, proxy=key_proxy,
) )
# 默认开启了 auto_fetch_models触发模型获取
await _trigger_auto_fetch_models([str(new_key.id)])
return ProviderCompleteOAuthResponse( return ProviderCompleteOAuthResponse(
key_id=str(new_key.id), key_id=str(new_key.id),
provider_type=provider_type, provider_type=provider_type,
@@ -1453,12 +1435,7 @@ async def import_refresh_token(
replaced=replaced, replaced=replaced,
) )
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException("不支持的 provider_type")
# 用 refresh_token 换取 access_token # 用 refresh_token 换取 access_token
refresh_token = payload.refresh_token.strip() refresh_token = payload.refresh_token.strip()
@@ -1573,9 +1550,6 @@ async def import_refresh_token(
proxy=key_proxy, proxy=key_proxy,
) )
# 默认开启了 auto_fetch_models触发模型获取
await _trigger_auto_fetch_models([str(new_key.id)])
return ProviderCompleteOAuthResponse( return ProviderCompleteOAuthResponse(
key_id=str(new_key.id), key_id=str(new_key.id),
provider_type=provider_type, provider_type=provider_type,
@@ -1635,12 +1609,7 @@ async def batch_import_oauth(
) )
# 标准 OAuth ProviderCodex、Antigravity、GeminiCli、ClaudeCode # 标准 OAuth ProviderCodex、Antigravity、GeminiCli、ClaudeCode
try: template = _require_oauth_template(provider_type)
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException(f"不支持的 provider_type: {provider_type}")
# 解析 Token 列表 # 解析 Token 列表
tokens = _parse_tokens_input(payload.credentials) tokens = _parse_tokens_input(payload.credentials)
@@ -1860,11 +1829,6 @@ async def batch_import_oauth(
if success_count > 0: if success_count > 0:
db.commit() db.commit()
# 批量导入完成后,触发所有成功 Key 的模型获取
success_key_ids = [r.key_id for r in results if r.status == "success" and r.key_id]
if success_key_ids:
await _trigger_auto_fetch_models(success_key_ids)
logger.info( logger.info(
"[BATCH_IMPORT] Provider {} ({}): 成功 {}/{}, 失败 {}", "[BATCH_IMPORT] Provider {} ({}): 成功 {}/{}, 失败 {}",
provider_id, provider_id,
@@ -2014,11 +1978,6 @@ async def _batch_import_kiro_internal(
if success_count > 0: if success_count > 0:
db.commit() db.commit()
# 批量导入完成后,触发所有成功 Key 的模型获取
success_key_ids = [r.key_id for r in results if r.status == "success" and r.key_id]
if success_key_ids:
await _trigger_auto_fetch_models(success_key_ids)
logger.info( logger.info(
"[KIRO_BATCH_IMPORT] Provider {}: 成功 {}/{}, 失败 {}", "[KIRO_BATCH_IMPORT] Provider {}: 成功 {}/{}, 失败 {}",
provider_id, provider_id,
@@ -2427,8 +2386,6 @@ async def device_poll(
proxy=key_proxy, proxy=key_proxy,
) )
await _trigger_auto_fetch_models([str(new_key.id)])
# 更新 Redis session 为已完成(短 TTL 让前端最后一次轮询能拿到结果) # 更新 Redis session 为已完成(短 TTL 让前端最后一次轮询能拿到结果)
session["status"] = "authorized" session["status"] = "authorized"
session["key_id"] = str(new_key.id) session["key_id"] = str(new_key.id)

View File

@@ -882,6 +882,8 @@ async def test_model(
auth_type=auth_type, auth_type=auth_type,
provider_type=p_type if p_type else None, provider_type=p_type if p_type else None,
decrypted_auth_config=oauth_meta if oauth_meta else None, decrypted_auth_config=oauth_meta if oauth_meta else None,
provider_endpoint=endpoint,
provider_api_key=api_key,
proxy_config=test_proxy, proxy_config=test_proxy,
) )
@@ -903,12 +905,58 @@ async def test_model(
return True return True
return False return False
def _extract_error_message(resp: dict) -> str:
"""从 check 响应中提取错误信息(用于判断是否值得回退)。"""
resp_data = resp.get("response", {}) if isinstance(resp, dict) else {}
body = resp_data.get("response_body", {})
parsed = body
if isinstance(body, str):
try:
parsed = json.loads(body)
except (json.JSONDecodeError, ValueError):
parsed = body
if isinstance(parsed, dict):
err = parsed.get("error")
if isinstance(err, dict):
msg = err.get("message")
if isinstance(msg, str):
return msg
if isinstance(err, str):
return err
err_raw = resp.get("error")
if isinstance(err_raw, str):
return err_raw
if isinstance(err_raw, dict):
msg = err_raw.get("message")
if isinstance(msg, str):
return msg
return ""
def _should_fallback_to_non_stream(resp: dict) -> bool:
"""仅在“流式特有失败”时回退到非流式,避免 429/鉴权错误的无效重试。"""
status = int(resp.get("status_code") or 0)
if status in {404, 405, 415, 501}:
return True
if status == 400:
msg = _extract_error_message(resp).lower()
stream_markers = ("stream", "sse", "streamgeneratecontent")
unsupported_markers = ("not support", "unsupported", "invalid argument")
if any(k in msg for k in stream_markers) and any(
k in msg for k in unsupported_markers
):
return True
return False
# 策略:优先流式,若失败回退到非流式 # 策略:优先流式,若失败回退到非流式
used_stream = True used_stream = True
logger.debug("[test-model] 尝试流式请求...") logger.debug("[test-model] 尝试流式请求...")
response = await _do_check(check_request) response = await _do_check(check_request)
if _response_has_error(response): if _response_has_error(response) and _should_fallback_to_non_stream(response):
logger.info( logger.info(
"[test-model] 流式请求失败 (status={}),回退到非流式请求", "[test-model] 流式请求失败 (status={}),回退到非流式请求",
response.get("status_code", "?"), response.get("status_code", "?"),
@@ -984,21 +1032,32 @@ async def test_model(
else: else:
logger.warning("[test-model] Key {} 因 403 verify 已标记为异常", api_key.id) logger.warning("[test-model] Key {} 因 403 verify 已标记为异常", api_key.id)
upstream_status = int(
response.get("status_code", 0) or error_obj.get("code", 0) or 500
)
if not (400 <= upstream_status <= 599):
upstream_status = 500
raise HTTPException( raise HTTPException(
status_code=500, status_code=upstream_status,
detail=str(error_message)[:500] if error_message else "Provider error", detail=str(error_message)[:500] if error_message else "Provider error",
) )
else: else:
logger.debug(f"[test-model] Error: {error_obj}") logger.debug(f"[test-model] Error: {error_obj}")
# error_obj 可能是字符串,截断以避免泄露过多上游信息 # error_obj 可能是字符串,截断以避免泄露过多上游信息
upstream_status = int(response.get("status_code", 0) or 500)
if not (400 <= upstream_status <= 599):
upstream_status = 500
raise HTTPException( raise HTTPException(
status_code=500, status_code=upstream_status,
detail=str(error_obj)[:500] if error_obj else "Provider error", detail=str(error_obj)[:500] if error_obj else "Provider error",
) )
elif "error" in response: elif "error" in response:
logger.debug(f"[test-model] Error: {response['error']}") logger.debug(f"[test-model] Error: {response['error']}")
upstream_status = int(response.get("status_code", 0) or 500)
if not (400 <= upstream_status <= 599):
upstream_status = 500
raise HTTPException( raise HTTPException(
status_code=500, status_code=upstream_status,
detail=str(response["error"])[:500], detail=str(response["error"])[:500],
) )
else: else:
@@ -1315,6 +1374,8 @@ async def test_model_failover(
auth_type=auth_type, auth_type=auth_type,
provider_type=p_type if p_type else None, provider_type=p_type if p_type else None,
decrypted_auth_config=oauth_meta if oauth_meta else None, decrypted_auth_config=oauth_meta if oauth_meta else None,
provider_endpoint=endpoint,
provider_api_key=key,
proxy_config=effective_proxy, proxy_config=effective_proxy,
) )
@@ -1343,9 +1404,7 @@ async def test_model_failover(
if not error_msg and isinstance(parsed, dict) and "error" in parsed: if not error_msg and isinstance(parsed, dict) and "error" in parsed:
err_val = parsed["error"] err_val = parsed["error"]
error_msg = str( error_msg = str(
err_val.get("message", err_val) err_val.get("message", err_val) if isinstance(err_val, dict) else err_val
if isinstance(err_val, dict)
else err_val
)[:300] )[:300]
attempts.append( attempts.append(
TestAttemptDetail( TestAttemptDetail(

View File

@@ -48,6 +48,7 @@ def _should_enable_format_conversion_by_default(provider_type: str | None) -> bo
ProviderType.CLAUDE_CODE.value, ProviderType.CLAUDE_CODE.value,
ProviderType.CODEX.value, ProviderType.CODEX.value,
ProviderType.KIRO.value, ProviderType.KIRO.value,
ProviderType.VERTEX_AI.value,
} }
return pt in envelope_provider_types return pt in envelope_provider_types
@@ -56,6 +57,15 @@ def _normalize_provider_type(provider_type: str | None) -> str:
return (provider_type or "custom").strip().lower() return (provider_type or "custom").strip().lower()
def _get_fixed_provider_template(provider_type: str | None) -> Any | None:
"""Return fixed-provider template when provider_type is managed by FIXED_PROVIDERS."""
normalized = _normalize_provider_type(provider_type)
try:
return FIXED_PROVIDERS.get(ProviderType(normalized))
except Exception:
return None
def _merge_pool_advanced_config( def _merge_pool_advanced_config(
*, *,
provider_config: dict[str, Any] | None, provider_config: dict[str, Any] | None,
@@ -458,33 +468,31 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
db.flush() # flush 获取 ID但不提交保持在同一事务中 db.flush() # flush 获取 ID但不提交保持在同一事务中
# 固定类型 Provider自动创建并锁定预置 Endpoints同一事务 # 固定类型 Provider自动创建并锁定预置 Endpoints同一事务
provider_type = (provider.provider_type or "custom").strip() template = _get_fixed_provider_template(provider.provider_type)
if provider_type != ProviderType.CUSTOM: if template:
try: now = datetime.now(timezone.utc)
template = FIXED_PROVIDERS.get(ProviderType(provider_type)) for sig in template.endpoint_signatures:
except Exception: endpoint_config: dict[str, str] | None = None
template = None if provider.provider_type == ProviderType.CODEX.value and sig == "openai:cli":
if template: endpoint_config = {"upstream_stream_policy": "force_stream"}
now = datetime.now(timezone.utc) endpoint = ProviderEndpoint(
for sig in template.endpoint_signatures: id=str(uuid.uuid4()),
endpoint = ProviderEndpoint( provider_id=provider.id,
id=str(uuid.uuid4()), api_format=sig,
provider_id=provider.id, api_family=sig.split(":", 1)[0],
api_format=sig, endpoint_kind=sig.split(":", 1)[1],
api_family=sig.split(":", 1)[0], base_url=template.api_base_url,
endpoint_kind=sig.split(":", 1)[1], custom_path=None,
base_url=template.api_base_url, header_rules=None,
custom_path=None, max_retries=provider.max_retries or 2,
header_rules=None, is_active=True,
max_retries=provider.max_retries or 2, config=endpoint_config,
is_active=True, proxy=None,
config=None, format_acceptance_config=None,
proxy=None, created_at=now,
format_acceptance_config=None, updated_at=now,
created_at=now, )
updated_at=now, db.add(endpoint)
)
db.add(endpoint)
db.commit() db.commit()
db.refresh(provider) db.refresh(provider)

View File

@@ -1576,7 +1576,7 @@ async def _resolve_provider_auth(
if account_id: if account_id:
auth_headers["chatgpt-account-id"] = str(account_id) auth_headers["chatgpt-account-id"] = str(account_id)
elif auth_type == "vertex_ai": elif auth_type in ("service_account", "vertex_ai"):
from src.api.handlers.base.request_builder import get_provider_auth from src.api.handlers.base.request_builder import get_provider_auth
auth_info = await get_provider_auth(endpoint, provider_key) auth_info = await get_provider_auth(endpoint, provider_key)

View File

@@ -13,7 +13,9 @@ from src.api.handlers.base.utils import get_format_converter_registry
from src.core.exceptions import ThinkingSignatureException, UpstreamClientException from src.core.exceptions import ThinkingSignatureException, UpstreamClientException
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ProviderAPIKey from src.models.database import ProviderAPIKey
from src.services.provider.transport import get_vertex_ai_effective_format from src.services.provider.adapters.vertex_ai.transport import (
get_effective_format as get_vertex_ai_effective_format,
)
from src.services.scheduling.aware_scheduler import ProviderCandidate from src.services.scheduling.aware_scheduler import ProviderCandidate
@@ -23,7 +25,7 @@ def _get_error_status_code(e: Exception, default: int = 400) -> int:
return code if isinstance(code, int) and code > 0 else default return code if isinstance(code, int) and code > 0 else default
def _resolve_vertex_ai_format( def _resolve_dynamic_format(
key: ProviderAPIKey, key: ProviderAPIKey,
auth_info: Any, auth_info: Any,
model: str, model: str,
@@ -32,9 +34,9 @@ def _resolve_vertex_ai_format(
candidate: ProviderCandidate | None, candidate: ProviderCandidate | None,
) -> tuple[str, bool]: ) -> tuple[str, bool]:
""" """
解析 Vertex AI 动态格式并计算 needs_conversion 解析动态格式并计算 needs_conversion
当 auth_type=vertex_ai 时,同一个 GCP 项目可以访问 Gemini 和 Claude 对于 Vertex AI 等跨格式 Provider同一个项目可以访问 Gemini 和 Claude
但它们的请求/响应格式不同,需要根据模型名动态选择。 但它们的请求/响应格式不同,需要根据模型名动态选择。
用户可通过 auth_config.model_format_mapping 配置自定义映射。 用户可通过 auth_config.model_format_mapping 配置自定义映射。
@@ -49,9 +51,13 @@ def _resolve_vertex_ai_format(
Returns: Returns:
(effective_provider_format, needs_conversion) 元组 (effective_provider_format, needs_conversion) 元组
""" """
key_auth_type = getattr(key, "auth_type", "api_key") from src.core.provider_types import ProviderType
if key_auth_type == "vertex_ai": # 判断是否为 Vertex AI provider基于 provider_type 而非 auth_type
provider = getattr(key, "provider", None)
provider_type = getattr(provider, "provider_type", None) if provider else None
if provider_type == ProviderType.VERTEX_AI:
vertex_auth_config = auth_info.decrypted_auth_config if auth_info else None vertex_auth_config = auth_info.decrypted_auth_config if auth_info else None
effective_format = get_vertex_ai_effective_format(model, vertex_auth_config) effective_format = get_vertex_ai_effective_format(model, vertex_auth_config)
if effective_format.upper() != provider_api_format.upper(): if effective_format.upper() != provider_api_format.upper():

View File

@@ -41,7 +41,7 @@ from src.api.handlers.base.base_handler import (
from src.api.handlers.base.chat_error_utils import ( from src.api.handlers.base.chat_error_utils import (
_build_error_json_payload, _build_error_json_payload,
_get_error_status_code, _get_error_status_code,
_resolve_vertex_ai_format, _resolve_dynamic_format,
) )
from src.api.handlers.base.parsers import get_parser_for_format from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get_provider_auth from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get_provider_auth
@@ -681,11 +681,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
流式和非流式请求共享此逻辑,唯一差异是 client_is_stream 参数。 流式和非流式请求共享此逻辑,唯一差异是 client_is_stream 参数。
""" """
# 提前获取认证信息(Vertex AI 格式判断需要使用 auth_config # 提前获取认证信息(动态格式判断需要使用 auth_config
auth_info = await get_provider_auth(endpoint, key) auth_info = await get_provider_auth(endpoint, key)
# 解析 Vertex AI 动态格式并计算 needs_conversion # 解析动态格式并计算 needs_conversionVertex AI 等跨格式 Provider
provider_api_format, needs_conversion = _resolve_vertex_ai_format( provider_api_format, needs_conversion = _resolve_dynamic_format(
key, auth_info, model, provider_api_format, client_api_format, candidate key, auth_info, model, provider_api_format, client_api_format, candidate
) )

View File

@@ -73,6 +73,7 @@ async def run_endpoint_check(
db: Any | None = None, # Session对象需要时才导入 db: Any | None = None, # Session对象需要时才导入
user: Any | None = None, # User对象 user: Any | None = None, # User对象
proxy_config: dict[str, Any] | None = None, # 原始代理配置(支持 tunnel 模式) proxy_config: dict[str, Any] | None = None, # 原始代理配置(支持 tunnel 模式)
is_stream: bool | None = None, # 显式流式标记(优先于 body/url 推断)
) -> dict[str, Any]: ) -> dict[str, Any]:
""" """
执行端点检查(重构版本,使用新的架构): 执行端点检查(重构版本,使用新的架构):
@@ -95,6 +96,7 @@ async def run_endpoint_check(
user=user, user=user,
request_id=str(uuid.uuid4())[:8], request_id=str(uuid.uuid4())[:8],
proxy_config=proxy_config, proxy_config=proxy_config,
is_stream=is_stream,
) )
# 使用协调器执行检查 # 使用协调器执行检查
@@ -567,6 +569,7 @@ class EndpointCheckRequest:
request_id: str | None = None request_id: str | None = None
timeout: float = 30.0 timeout: float = 30.0
proxy_config: dict[str, Any] | None = None # 原始代理配置(支持 tunnel 模式) proxy_config: dict[str, Any] | None = None # 原始代理配置(支持 tunnel 模式)
is_stream: bool | None = None # 显式流式标记(优先于 body/url 推断)
@dataclass @dataclass
@@ -593,8 +596,22 @@ class HttpRequestExecutor:
start_time = time.time() start_time = time.time()
request_id = request.request_id or str(uuid.uuid4())[:8] request_id = request.request_id or str(uuid.uuid4())[:8]
# 检查是否是流式请求 # 检查是否是流式请求(优先显式参数,其次 body最后 URL 推断)
is_stream = request.json_body.get("stream", False) if request.json_body else False if request.is_stream is not None:
is_stream = bool(request.is_stream)
else:
is_stream = request.json_body.get("stream", False) if request.json_body else False
if not is_stream:
lowered_url = (request.url or "").lower()
if any(
marker in lowered_url
for marker in (
":streamgeneratecontent",
"/stream",
"stream=true",
)
):
is_stream = True
try: try:
from src.services.proxy_node.resolver import build_proxy_client_kwargs from src.services.proxy_node.resolver import build_proxy_client_kwargs

View File

@@ -338,6 +338,8 @@ class HandlerAdapterBase(ApiAdapter):
auth_type: str | None = None, auth_type: str | None = None,
provider_type: str | None = None, provider_type: str | None = None,
decrypted_auth_config: dict[str, Any] | None = None, decrypted_auth_config: dict[str, Any] | None = None,
provider_endpoint: Any | None = None,
provider_api_key: Any | None = None,
# 代理配置 # 代理配置
proxy_config: dict[str, Any] | None = None, proxy_config: dict[str, Any] | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -352,8 +354,10 @@ class HandlerAdapterBase(ApiAdapter):
from src.core.provider_types import ProviderType from src.core.provider_types import ProviderType
is_antigravity = provider_type == ProviderType.ANTIGRAVITY is_antigravity = provider_type == ProviderType.ANTIGRAVITY
is_vertex = provider_type == ProviderType.VERTEX_AI
is_kiro = provider_type == ProviderType.KIRO is_kiro = provider_type == ProviderType.KIRO
is_oauth = auth_type == "oauth" is_oauth = auth_type == "oauth"
vertex_auth_info: Any | None = None
# ---- URL ---- # ---- URL ----
if is_kiro: if is_kiro:
@@ -381,6 +385,28 @@ class HandlerAdapterBase(ApiAdapter):
effective_base_url = ordered_urls[0] if ordered_urls else base_url effective_base_url = ordered_urls[0] if ordered_urls else base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent") path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(effective_base_url).rstrip('/')}{path}" url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
vertex_auth_info = await get_provider_auth(provider_endpoint, provider_api_key)
effective_auth_config = (
vertex_auth_info.decrypted_auth_config
if vertex_auth_info
else decrypted_auth_config
)
if effective_auth_config:
decrypted_auth_config = effective_auth_config
effective_model_name = model_name or request_data.get("model", "")
path_params = {"model": effective_model_name} if effective_model_name else None
url = build_provider_url(
provider_endpoint,
path_params=path_params,
is_stream=bool(request_data.get("stream", False)),
key=provider_api_key,
decrypted_auth_config=effective_auth_config,
)
else: else:
url = cls.build_endpoint_url(base_url, request_data, model_name) url = cls.build_endpoint_url(base_url, request_data, model_name)
@@ -412,15 +438,24 @@ class HandlerAdapterBase(ApiAdapter):
) )
merged_extra.update(kiro_headers) merged_extra.update(kiro_headers)
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None) if is_vertex and provider_endpoint is not None and provider_api_key is not None:
headers = dict(merged_extra)
if (
vertex_auth_info
and getattr(vertex_auth_info, "auth_header", None)
and getattr(vertex_auth_info, "auth_value", None)
):
headers[str(vertex_auth_info.auth_header)] = str(vertex_auth_info.auth_value)
else:
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
if is_oauth: if is_oauth:
from src.core.api_format import get_auth_config_for_endpoint from src.core.api_format import get_auth_config_for_endpoint
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID) default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
if default_auth_header.lower() != "authorization": if default_auth_header.lower() != "authorization":
headers.pop(default_auth_header, None) headers.pop(default_auth_header, None)
headers["Authorization"] = f"Bearer {api_key}" headers["Authorization"] = f"Bearer {api_key}"
# ---- Body ---- # ---- Body ----
body = cls.build_request_body(request_data, base_url=base_url) body = cls.build_request_body(request_data, base_url=base_url)
@@ -464,6 +499,8 @@ class HandlerAdapterBase(ApiAdapter):
else: else:
ep_auth_header, _ = _get_auth_cfg(cls.FORMAT_ID) ep_auth_header, _ = _get_auth_cfg(cls.FORMAT_ID)
protected_keys = {ep_auth_header.lower(), "content-type"} protected_keys = {ep_auth_header.lower(), "content-type"}
if vertex_auth_info and getattr(vertex_auth_info, "auth_header", None):
protected_keys.add(str(vertex_auth_info.auth_header).lower())
header_builder = HeaderBuilder() header_builder = HeaderBuilder()
header_builder.add_many(headers) header_builder.add_many(headers)
@@ -479,6 +516,7 @@ class HandlerAdapterBase(ApiAdapter):
headers=headers, headers=headers,
json_body=body, json_body=body,
api_format=cls.FORMAT_ID, api_format=cls.FORMAT_ID,
is_stream=bool(request_data.get("stream", False)),
db=db, db=db,
user=user, user=user,
provider_name=provider_name, provider_name=provider_name,

View File

@@ -296,6 +296,8 @@ class GeminiChatAdapter(ChatAdapterBase):
auth_type: str | None = None, auth_type: str | None = None,
provider_type: str | None = None, provider_type: str | None = None,
decrypted_auth_config: dict[str, Any] | None = None, decrypted_auth_config: dict[str, Any] | None = None,
provider_endpoint: Any | None = None,
provider_api_key: Any | None = None,
# 代理配置 # 代理配置
proxy_config: dict[str, Any] | None = None, proxy_config: dict[str, Any] | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -313,7 +315,9 @@ class GeminiChatAdapter(ChatAdapterBase):
} }
is_antigravity = provider_type and provider_type.lower() == "antigravity" is_antigravity = provider_type and provider_type.lower() == "antigravity"
is_vertex = provider_type and provider_type.lower() == "vertex_ai"
is_oauth = auth_type == "oauth" is_oauth = auth_type == "oauth"
vertex_auth_info: Any | None = None
# Antigravity provider 使用 v1internal 路径,而非标准 Gemini API 路径 # Antigravity provider 使用 v1internal 路径,而非标准 Gemini API 路径
if is_antigravity: if is_antigravity:
@@ -327,6 +331,27 @@ class GeminiChatAdapter(ChatAdapterBase):
ag_base = ordered_urls[0] if ordered_urls else base_url ag_base = ordered_urls[0] if ordered_urls else base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent") path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(ag_base).rstrip('/')}{path}" url = f"{str(ag_base).rstrip('/')}{path}"
elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
# Vertex AI: test-model 必须走统一 provider transport/auth
# 否则会错误命中普通 Gemini URL导致 404
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
vertex_auth_info = await get_provider_auth(provider_endpoint, provider_api_key)
effective_auth_config = (
vertex_auth_info.decrypted_auth_config
if vertex_auth_info
else decrypted_auth_config
)
if effective_auth_config:
decrypted_auth_config = effective_auth_config
url = build_provider_url(
provider_endpoint,
path_params={"model": effective_model_name},
is_stream=bool(request_data.get("stream", False)),
key=provider_api_key,
decrypted_auth_config=effective_auth_config,
)
else: else:
# 使用基类配置方法但重写URL构建逻辑 # 使用基类配置方法但重写URL构建逻辑
base_url_resolved = cls.build_endpoint_url(base_url) base_url_resolved = cls.build_endpoint_url(base_url)
@@ -337,16 +362,25 @@ class GeminiChatAdapter(ChatAdapterBase):
merged_extra = dict(extra_headers) if extra_headers else {} merged_extra = dict(extra_headers) if extra_headers else {}
if is_antigravity: if is_antigravity:
merged_extra.update(get_v1internal_extra_headers()) merged_extra.update(get_v1internal_extra_headers())
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None) if is_vertex and provider_endpoint is not None and provider_api_key is not None:
headers = dict(merged_extra)
if (
vertex_auth_info
and getattr(vertex_auth_info, "auth_header", None)
and getattr(vertex_auth_info, "auth_value", None)
):
headers[str(vertex_auth_info.auth_header)] = str(vertex_auth_info.auth_value)
else:
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
# OAuth 统一处理替换端点默认认证头x-goog-api-key为 Authorization: Bearer # OAuth 统一处理替换端点默认认证头x-goog-api-key为 Authorization: Bearer
if is_oauth: if is_oauth:
from src.core.api_format import get_auth_config_for_endpoint from src.core.api_format import get_auth_config_for_endpoint
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID) default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
if default_auth_header.lower() != "authorization": if default_auth_header.lower() != "authorization":
headers.pop(default_auth_header, None) headers.pop(default_auth_header, None)
headers["Authorization"] = f"Bearer {api_key}" headers["Authorization"] = f"Bearer {api_key}"
body = cls.build_request_body(request_data) body = cls.build_request_body(request_data)
@@ -373,6 +407,8 @@ class GeminiChatAdapter(ChatAdapterBase):
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID) auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
protected_keys = {auth_header.lower(), "content-type"} protected_keys = {auth_header.lower(), "content-type"}
if vertex_auth_info and getattr(vertex_auth_info, "auth_header", None):
protected_keys.add(str(vertex_auth_info.auth_header).lower())
header_builder = HeaderBuilder() header_builder = HeaderBuilder()
header_builder.add_many(headers) header_builder.add_many(headers)
@@ -385,6 +421,7 @@ class GeminiChatAdapter(ChatAdapterBase):
headers=headers, headers=headers,
json_body=body, json_body=body,
api_format=cls.FORMAT_ID, api_format=cls.FORMAT_ID,
is_stream=bool(request_data.get("stream", False)),
# 用量计算参数(现在强制记录) # 用量计算参数(现在强制记录)
db=db, db=db,
user=user, user=user,

View File

@@ -212,7 +212,7 @@ class GeminiVeoHandler(VideoHandlerBase):
task_type="video", task_type="video",
submit_func=_submit, submit_func=_submit,
extract_external_task_id=_extract_task_id, extract_external_task_id=_extract_task_id,
supported_auth_types={"api_key", "vertex_ai"}, supported_auth_types={"api_key", "service_account", "vertex_ai"},
allow_format_conversion=True, allow_format_conversion=True,
max_candidates=10, max_candidates=10,
) )

View File

@@ -198,7 +198,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
task_type="video", task_type="video",
submit_func=_submit, submit_func=_submit,
extract_external_task_id=_extract_task_id, extract_external_task_id=_extract_task_id,
supported_auth_types={"api_key", "vertex_ai"}, supported_auth_types={"api_key", "service_account", "vertex_ai"},
allow_format_conversion=True, allow_format_conversion=True,
max_candidates=10, max_candidates=10,
) )

View File

@@ -72,7 +72,7 @@ FIXED_PROVIDERS: dict[ProviderType, FixedProviderTemplate] = {
provider_type=ProviderType.CODEX, provider_type=ProviderType.CODEX,
display_name="Codex", display_name="Codex",
api_base_url="https://chatgpt.com/backend-api/codex", api_base_url="https://chatgpt.com/backend-api/codex",
endpoint_signatures=["openai:cli"], endpoint_signatures=["openai:cli", "openai:compact"],
oauth=FixedProviderOAuth( oauth=FixedProviderOAuth(
authorize_url="https://auth.openai.com/oauth/authorize", authorize_url="https://auth.openai.com/oauth/authorize",
token_url="https://auth.openai.com/oauth/token", token_url="https://auth.openai.com/oauth/token",
@@ -121,6 +121,23 @@ FIXED_PROVIDERS: dict[ProviderType, FixedProviderTemplate] = {
use_pkce=False, use_pkce=False,
), ),
), ),
ProviderType.VERTEX_AI: FixedProviderTemplate(
provider_type=ProviderType.VERTEX_AI,
display_name="Vertex AI",
# Vertex uses fixed global base URL; concrete upstream path is selected by transport hook.
api_base_url="https://aiplatform.googleapis.com",
endpoint_signatures=["gemini:chat", "claude:chat"],
# Vertex does not use this OAuth flow (it uses API Key / Service Account).
oauth=FixedProviderOAuth(
authorize_url="",
token_url="",
client_id="",
client_secret="",
scopes=[],
redirect_uri="",
use_pkce=False,
),
),
ProviderType.ANTIGRAVITY: FixedProviderTemplate( ProviderType.ANTIGRAVITY: FixedProviderTemplate(
provider_type=ProviderType.ANTIGRAVITY, provider_type=ProviderType.ANTIGRAVITY,
display_name="Antigravity", display_name="Antigravity",

View File

@@ -21,6 +21,7 @@ class ProviderType(str, Enum):
CODEX = "codex" CODEX = "codex"
GEMINI_CLI = "gemini_cli" GEMINI_CLI = "gemini_cli"
ANTIGRAVITY = "antigravity" ANTIGRAVITY = "antigravity"
VERTEX_AI = "vertex_ai"
# 所有有效 provider_type 值的集合(用于校验) # 所有有效 provider_type 值的集合(用于校验)

View File

@@ -1550,19 +1550,19 @@ class ProviderAPIKey(ExportMixin, Base):
# 认证类型 # 认证类型
# - "api_key": 标准 API Key 认证(默认) # - "api_key": 标准 API Key 认证(默认)
# - "vertex_ai": Google Vertex AI 认证(Service Account JSON # - "service_account": GCP Service Account JSON 认证
# - 未来可扩展oauth2, azure_ad, aws_iam 等 # - "oauth": OAuth access_token / refresh_token 认证
auth_type = Column(String(20), default="api_key", nullable=False) auth_type = Column(String(20), default="api_key", nullable=False)
# API密钥加密存储 # API密钥加密存储
# - auth_type="api_key" 时:存储 API Key 字符串 # - auth_type="api_key" 时:存储 API Key 字符串
# - auth_type="vertex_ai" :可为,敏感凭证存在 auth_config 中 # - auth_type="service_account"/"oauth" :可为占位符,敏感凭证存在 auth_config 中
api_key = Column(Text, nullable=False) # 使用 Text 支持加密后的 OAuth token api_key = Column(Text, nullable=False) # 使用 Text 支持加密后的 OAuth token
# 认证配置(加密存储) # 认证配置(加密存储)
# - auth_type="api_key" 时:可为空 # - auth_type="api_key" 时:可为空
# - auth_type="vertex_ai" 时:存储加密后的 Service Account JSON # - auth_type="service_account" 时:存储加密后的 Service Account JSON
# - auth_type="oauth2" 时:存储加密后的 {client_id, client_secret, token_url, scope} # - auth_type="oauth" 时:存储加密后的 {refresh_token, expires_at, ...}
auth_config = Column(Text, nullable=True) auth_config = Column(Text, nullable=True)
name = Column(String(100), nullable=False) # 密钥名称(必填,用于识别) name = Column(String(100), nullable=False) # 密钥名称(必填,用于识别)
note = Column(String(500), nullable=True) # 备注说明(可选) note = Column(String(500), nullable=True) # 备注说明(可选)

View File

@@ -261,7 +261,7 @@ class ProviderEndpointCreate(BaseModel):
api_format: str = Field( api_format: str = Field(
..., ...,
description=( description=(
"Endpoint signature例如: claude:chat/claude:cli, openai:chat/openai:cli/openai:video, gemini:chat/gemini:cli/gemini:video" "Endpoint signature例如: claude:chat/claude:cli, openai:chat/openai:cli/openai:compact/openai:video, gemini:chat/gemini:cli/gemini:video"
), ),
) )
base_url: str = Field(..., min_length=1, max_length=500, description="API 基础 URL") base_url: str = Field(..., min_length=1, max_length=500, description="API 基础 URL")
@@ -435,14 +435,14 @@ class EndpointAPIKeyCreate(BaseModel):
api_key: str = Field( api_key: str = Field(
default="", max_length=10000, description="API Key标准认证时必填将自动加密" default="", max_length=10000, description="API Key标准认证时必填将自动加密"
) )
auth_type: Literal["api_key", "vertex_ai", "oauth"] = Field( auth_type: Literal["api_key", "service_account", "oauth"] = Field(
default="api_key", default="api_key",
description="认证类型api_key标准 API Key/ vertex_aiVertex AI Service Account/ oauthOAuth access_token", description="认证类型api_key标准 API Key/ service_accountGCP Service Account/ oauthOAuth access_token",
) )
auth_config: dict[str, Any] | None = Field( auth_config: dict[str, Any] | None = Field(
default=None, default=None,
description=( description=(
"认证配置JSONvertex_ai 时存储完整 Service Account JSON" "认证配置JSONservice_account 时存储完整 Service Account JSON"
"oauth 时存储 token/refresh/expires_at 等(后端加密存储,不在响应中返回)" "oauth 时存储 token/refresh/expires_at 等(后端加密存储,不在响应中返回)"
), ),
) )
@@ -590,14 +590,14 @@ class EndpointAPIKeyUpdate(BaseModel):
max_length=10000, max_length=10000,
description="API Key标准认证时使用将自动加密", description="API Key标准认证时使用将自动加密",
) )
auth_type: Literal["api_key", "vertex_ai", "oauth"] | None = Field( auth_type: Literal["api_key", "service_account", "oauth"] | None = Field(
default=None, default=None,
description="认证类型api_key标准 API Key/ vertex_aiVertex AI Service Account/ oauthOAuth access_token", description="认证类型api_key标准 API Key/ service_accountGCP Service Account/ oauthOAuth access_token",
) )
auth_config: dict[str, Any] | None = Field( auth_config: dict[str, Any] | None = Field(
default=None, default=None,
description=( description=(
"认证配置JSONvertex_ai 时存储完整 Service Account JSON" "认证配置JSONservice_account 时存储完整 Service Account JSON"
"oauth 时存储 token/refresh/expires_at 等(后端加密存储,不在响应中返回)" "oauth 时存储 token/refresh/expires_at 等(后端加密存储,不在响应中返回)"
), ),
) )
@@ -719,7 +719,9 @@ class EndpointAPIKeyResponse(BaseModel):
# Key 信息(脱敏) # Key 信息(脱敏)
api_key_masked: str = Field(..., description="脱敏后的 Key") api_key_masked: str = Field(..., description="脱敏后的 Key")
api_key_plain: str | None = Field(default=None, description="完整的 Key") api_key_plain: str | None = Field(default=None, description="完整的 Key")
auth_type: str = Field(default="api_key", description="认证类型api_key 或 vertex_ai") auth_type: str = Field(
default="api_key", description="认证类型api_key / service_account / oauth"
)
# auth_config 不在响应中返回(包含敏感信息),前端通过 auth_type 判断类型 # auth_config 不在响应中返回(包含敏感信息),前端通过 auth_type 判断类型
name: str = Field(..., description="密钥名称") name: str = Field(..., description="密钥名称")

View File

@@ -408,13 +408,15 @@ class ModelFetchScheduler:
db.commit() db.commit()
return "error" return "error"
# Vertex AI 类型不支持自动获取模型 # Service Account 类型不支持自动获取模型Vertex AI SA / 旧 vertex_ai auth_type
auth_type = getattr(key, "auth_type", "api_key") or "api_key" auth_type = getattr(key, "auth_type", "api_key") or "api_key"
if auth_type == "vertex_ai": if auth_type in ("service_account", "vertex_ai"):
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key" key.last_models_fetch_error = (
"auto_fetch_models 暂不支持 Service Account 类型的 Key"
)
key.last_models_fetch_at = now key.last_models_fetch_at = now
db.commit() db.commit()
logger.info(f"Key {key.id}Vertex AI 类型,跳过自动获取模型") logger.info(f"Key {key.id}Service Account 类型,跳过自动获取模型")
return "skip" return "skip"
# 基础校验:必须有 api_keyOAuth: 加密 access_tokenAPI Key: 加密 key # 基础校验:必须有 api_keyOAuth: 加密 access_tokenAPI Key: 加密 key

View File

@@ -31,6 +31,7 @@ def build_antigravity_url(
*, *,
is_stream: bool, is_stream: bool,
effective_query_params: dict[str, Any], effective_query_params: dict[str, Any],
**_kwargs: Any,
) -> str: ) -> str:
"""构建 Antigravity v1internal URL。 """构建 Antigravity v1internal URL。

View File

@@ -16,6 +16,7 @@ def build_claude_code_url(
*, *,
is_stream: bool, is_stream: bool,
effective_query_params: dict[str, Any], effective_query_params: dict[str, Any],
**_kwargs: Any,
) -> str: ) -> str:
"""Build Claude Code upstream URL and avoid duplicate /v1/messages suffix.""" """Build Claude Code upstream URL and avoid duplicate /v1/messages suffix."""
_ = is_stream _ = is_stream

View File

@@ -37,6 +37,7 @@ def build_codex_url(
*, *,
is_stream: bool, is_stream: bool,
effective_query_params: dict[str, Any], effective_query_params: dict[str, Any],
**_kwargs: Any,
) -> str: ) -> str:
"""构建 Codex OAuth URL。 """构建 Codex OAuth URL。
@@ -48,7 +49,8 @@ def build_codex_url(
from src.services.provider.adapters.codex.context import get_codex_request_context from src.services.provider.adapters.codex.context import get_codex_request_context
ctx = get_codex_request_context() ctx = get_codex_request_context()
is_compact = ctx.is_compact if ctx else False endpoint_sig = str(getattr(endpoint, "api_format", "") or "").strip().lower()
is_compact = bool((ctx.is_compact if ctx else False) or endpoint_sig == "openai:compact")
base = str(endpoint.base_url).rstrip("/") base = str(endpoint.base_url).rstrip("/")
# 如果用户已在 base_url 中包含了 /responses不要重复追加 # 如果用户已在 base_url 中包含了 /responses不要重复追加
@@ -110,10 +112,12 @@ def register_all() -> None:
# Envelope # Envelope
register_envelope("codex", "openai:cli", codex_oauth_envelope) register_envelope("codex", "openai:cli", codex_oauth_envelope)
register_envelope("codex", "openai:compact", codex_oauth_envelope)
register_envelope("codex", "", codex_oauth_envelope) register_envelope("codex", "", codex_oauth_envelope)
# Transport # Transport
register_transport_hook("codex", "openai:cli", build_codex_url) register_transport_hook("codex", "openai:cli", build_codex_url)
register_transport_hook("codex", "openai:compact", build_codex_url)
# Auth # Auth
register_auth_enricher("codex", enrich_codex) register_auth_enricher("codex", enrich_codex)

View File

@@ -42,6 +42,7 @@ def build_kiro_url(
*, *,
is_stream: bool, is_stream: bool,
effective_query_params: dict[str, Any], effective_query_params: dict[str, Any],
**_kwargs: Any,
) -> str: ) -> str:
"""Build Kiro generateAssistantResponse URL. """Build Kiro generateAssistantResponse URL.

View File

@@ -0,0 +1,60 @@
"""Vertex AI 认证处理。
- Service Account: GCP SA JSON → JWT → Access Token → Bearer header
- API Key: 通过 URL ?key= 查询参数认证auth 层返回 None由 transport hook 处理)
"""
from __future__ import annotations
import json
from typing import Any
from src.core.provider_auth_types import ProviderAuthInfo
async def _auth_service_account(key: Any) -> ProviderAuthInfo:
"""Service Account 认证SA JSON → JWT → Access Token。"""
from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException
from src.core.vertex_auth import VertexAuthError, VertexAuthService
try:
# 优先从 auth_config 读取,兼容从 api_key 读取(过渡期)
encrypted_auth_config = getattr(key, "auth_config", None)
if encrypted_auth_config:
if isinstance(encrypted_auth_config, dict):
sa_json = encrypted_auth_config
else:
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
sa_json = json.loads(decrypted_config)
else:
# 兼容旧数据:从 api_key 读取
decrypted_key = crypto_service.decrypt(key.api_key)
if decrypted_key == "__placeholder__":
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
sa_json = json.loads(decrypted_key)
if not isinstance(sa_json, dict):
raise InvalidRequestException("Service Account JSON 无效,请重新添加该密钥。")
# 获取 Access Token注入代理配置
from src.services.proxy_node.resolver import build_proxy_client_kwargs
service = VertexAuthService(sa_json)
access_token = await service.get_access_token(
httpx_client_kwargs=build_proxy_client_kwargs(timeout=30),
)
return ProviderAuthInfo(
auth_header="Authorization",
auth_value=f"Bearer {access_token}",
decrypted_auth_config=sa_json,
)
except InvalidRequestException:
raise
except VertexAuthError as e:
raise InvalidRequestException(f"Vertex AI 认证失败:{e}")
except json.JSONDecodeError:
raise InvalidRequestException("Service Account JSON 格式无效,请重新添加该密钥。")
except Exception:
raise InvalidRequestException("Vertex AI 认证失败,请检查 Key 的 auth_config")

View File

@@ -0,0 +1,45 @@
"""Vertex AI 常量配置。
从 transport.py 迁移,集中管理 Vertex AI 模型格式映射和 region 配置。
"""
from __future__ import annotations
# Vertex AI 模型前缀到 API 格式的映射
# 用于 provider_type=vertex_ai 时,根据模型名动态确定实际的请求/响应格式
# 格式:前缀 -> endpoint signaturefamily:kind
MODEL_FORMAT_MAPPING: dict[str, str] = {
"claude-": "claude:chat", # Anthropic Claude 模型
"gemini-": "gemini:chat", # Google Gemini 模型
"imagen-": "gemini:chat", # Google Imagen 模型(使用 Gemini chat 格式)
}
# Vertex AI 默认 endpoint signature当模型前缀不匹配时
DEFAULT_FORMAT: str = "gemini:chat"
# Vertex AI 模型默认 region 映射
# 用户可以通过 auth_config.model_regions 覆盖
DEFAULT_MODEL_REGIONS: dict[str, str] = {
# Gemini 3 系列(使用 global
"gemini-3.1-pro-preview": "global",
"gemini-3-pro-image-preview": "global",
# Gemini 2.0 系列
"gemini-2.0-flash": "us-central1",
"gemini-2.0-flash-exp": "us-central1",
"gemini-2.0-flash-001": "us-central1",
"gemini-2.0-pro-exp": "us-central1",
"gemini-2.0-flash-exp-image-generation": "us-central1",
# Gemini 1.5 系列
"gemini-1.5-pro": "us-central1",
"gemini-1.5-pro-001": "us-central1",
"gemini-1.5-pro-002": "us-central1",
"gemini-1.5-flash": "us-central1",
"gemini-1.5-flash-001": "us-central1",
"gemini-1.5-flash-002": "us-central1",
# Imagen 系列
"imagen-3.0-generate-001": "us-central1",
"imagen-3.0-fast-generate-001": "us-central1",
}
# API Key 认证的全局端点
API_KEY_BASE_URL = "https://aiplatform.googleapis.com"

View File

@@ -0,0 +1,474 @@
"""Vertex AI provider plugin — 统一注册入口。
注册 Vertex AI 对各通用 registry 的 hooks
- Transport Hook (URL 构建,支持 API Key / Service Account 双策略)
- Model Fetcher (专用上游模型获取链路,不走通用 /v1beta/models / /v1/models)
- Behavior Variants (跨格式支持:同一 Provider 同时访问 Gemini 和 Claude 模型)
"""
from __future__ import annotations
from typing import Any
import httpx
from src.core.logger import logger
from src.core.vertex_auth import VertexAuthError, VertexAuthService
from src.services.provider.adapters.vertex_ai.transport import get_effective_format
# Vertex AI 公共 API 根
_VERTEX_API_BASE = "https://aiplatform.googleapis.com"
# Gemini Developer APIAPI Key 场景兜底)
_GEMINI_DEV_BASE = "https://generativelanguage.googleapis.com"
_MODEL_PAGE_SIZE = 100
_MODEL_MAX_PAGES = 20
def _normalize_extra_headers(raw: Any) -> dict[str, str]:
if not isinstance(raw, dict):
return {}
return {str(k): str(v) for k, v in raw.items() if k and v is not None}
def _looks_like_service_account(auth_config: dict[str, Any] | None) -> bool:
if not isinstance(auth_config, dict):
return False
return all(
isinstance(auth_config.get(k), str) and str(auth_config.get(k)).strip()
for k in ("client_email", "private_key", "project_id")
)
def _extract_model_id(raw_name: str) -> str:
name = str(raw_name or "").strip()
if not name:
return ""
if "/models/" in name:
return name.split("/models/", 1)[-1].strip()
if name.startswith("models/"):
return name.split("models/", 1)[-1].strip()
return name
def _extract_publisher(item: dict[str, Any], fallback: str | None = None) -> str | None:
publisher = item.get("publisher")
if isinstance(publisher, str) and publisher.strip():
return publisher.strip()
raw_name = item.get("name")
if isinstance(raw_name, str) and "/publishers/" in raw_name:
try:
after = raw_name.split("/publishers/", 1)[1]
candidate = after.split("/", 1)[0].strip()
if candidate:
return candidate
except Exception:
pass
return fallback
def _extract_items(data: Any) -> list[dict[str, Any]]:
if isinstance(data, list):
return [item for item in data if isinstance(item, dict)]
if not isinstance(data, dict):
return []
for key in ("publisherModels", "models", "data", "items"):
value = data.get(key)
if isinstance(value, list):
return [item for item in value if isinstance(item, dict)]
return []
def _parse_models_payload(
data: Any,
*,
auth_config: dict[str, Any] | None,
fallback_publisher: str | None = None,
) -> list[dict[str, Any]]:
models: list[dict[str, Any]] = []
for item in _extract_items(data):
raw_name = item.get("id") or item.get("name") or item.get("model")
if not isinstance(raw_name, str):
continue
model_id = _extract_model_id(raw_name)
if not model_id:
continue
display_name_raw = (
item.get("displayName") or item.get("display_name") or item.get("title") or model_id
)
display_name = (
str(display_name_raw).strip() if isinstance(display_name_raw, str) else model_id
)
if not display_name:
display_name = model_id
models.append(
{
"id": model_id,
"owned_by": _extract_publisher(item, fallback=fallback_publisher),
"display_name": display_name,
"api_format": get_effective_format(model_id, auth_config),
}
)
return models
def _build_google_publisher_list_url(base_url: str) -> str:
base = str(base_url or "").rstrip("/")
if not base:
base = _VERTEX_API_BASE
if base.endswith("/v1"):
return f"{base}/publishers/google/models"
if base.endswith("/v1beta"):
return f"{base}/publishers/google/models"
return f"{base}/v1/publishers/google/models"
def _iter_endpoint_base_urls(ctx: Any) -> list[str]:
seen: set[str] = set()
urls: list[str] = []
for cfg in (ctx.format_to_endpoint or {}).values():
base_url = str(getattr(cfg, "base_url", "") or "").strip()
if not base_url:
continue
norm = base_url.rstrip("/")
if norm in seen:
continue
seen.add(norm)
urls.append(norm)
if _VERTEX_API_BASE not in seen:
urls.append(_VERTEX_API_BASE)
return urls
def _get_endpoint_headers(ctx: Any, api_format: str) -> dict[str, str]:
cfg = (ctx.format_to_endpoint or {}).get(api_format)
return _normalize_extra_headers(getattr(cfg, "extra_headers", None))
def _dedupe_models(models: list[dict[str, Any]]) -> list[dict[str, Any]]:
seen: set[str] = set()
result: list[dict[str, Any]] = []
for model in models:
model_id = str(model.get("id", "")).strip()
api_format = str(model.get("api_format", "")).strip()
if not model_id:
continue
unique_key = f"{model_id}:{api_format}"
if unique_key in seen:
continue
seen.add(unique_key)
result.append(model)
return result
def _is_soft_not_found(error: str) -> bool:
return str(error).strip().startswith("HTTP 404:")
def _iter_regions(auth_config: dict[str, Any] | None) -> list[str]:
seen: set[str] = set()
regions: list[str] = []
def _add(raw: Any) -> None:
if not isinstance(raw, str):
return
region = raw.strip()
if not region or region in seen:
return
seen.add(region)
regions.append(region)
if isinstance(auth_config, dict):
_add(auth_config.get("region"))
model_regions = auth_config.get("model_regions")
if isinstance(model_regions, dict):
for region in model_regions.values():
_add(region)
_add("global")
_add("us-central1")
return regions
async def _fetch_models_from_url(
client: httpx.AsyncClient,
*,
url: str,
headers: dict[str, str],
params: dict[str, Any],
auth_config: dict[str, Any] | None,
fallback_publisher: str | None = None,
) -> tuple[list[dict[str, Any]], str | None, bool]:
all_models: list[dict[str, Any]] = []
next_page_token: str | None = None
has_success = False
for _ in range(_MODEL_MAX_PAGES):
req_params = dict(params)
if next_page_token:
req_params["pageToken"] = next_page_token
try:
resp = await client.get(url, headers=headers, params=req_params)
except httpx.TimeoutException:
return [], "timeout", has_success
except Exception as exc:
return [], f"request error: {exc}", has_success
if resp.status_code != 200:
body = resp.text[:500] if resp.text else "(empty)"
return [], f"HTTP {resp.status_code}: {body}", has_success
has_success = True
try:
payload = resp.json()
except Exception:
body = resp.text[:500] if resp.text else "(empty)"
return [], f"invalid json body: {body}", has_success
all_models.extend(
_parse_models_payload(
payload,
auth_config=auth_config,
fallback_publisher=fallback_publisher,
)
)
if not isinstance(payload, dict):
break
token = payload.get("nextPageToken")
next_page_token = str(token).strip() if isinstance(token, str) else None
if not next_page_token:
break
return all_models, None, has_success
async def _fetch_models_vertex_api_key(
client: httpx.AsyncClient,
*,
ctx: Any,
auth_config: dict[str, Any] | None,
) -> tuple[list[dict[str, Any]], list[str], bool]:
api_key = str(ctx.api_key_value or "").strip()
if not api_key or api_key == "__placeholder__":
return [], ["vertex_ai(api_key): missing api key"], False
all_models: list[dict[str, Any]] = []
hard_errors: list[str] = []
soft_errors: list[str] = []
has_success = False
endpoint_headers = _get_endpoint_headers(ctx, "gemini:chat")
vertex_list_urls = [
_build_google_publisher_list_url(base) for base in _iter_endpoint_base_urls(ctx)
]
# 1) Vertex API list (publisher=google)
for url in vertex_list_urls:
headers = {"Accept": "application/json", **endpoint_headers}
models, err, success = await _fetch_models_from_url(
client,
url=url,
headers=headers,
params={"key": api_key, "pageSize": _MODEL_PAGE_SIZE},
auth_config=auth_config,
fallback_publisher="google",
)
if success:
has_success = True
if err:
labeled = f"{url}: {err}"
if _is_soft_not_found(err):
soft_errors.append(labeled)
else:
hard_errors.append(labeled)
continue
all_models.extend(models)
# 2) 兜底Gemini Developer API
if not all_models:
fallback_url = f"{_GEMINI_DEV_BASE}/v1beta/models"
headers = {"Accept": "application/json", **endpoint_headers}
models, err, success = await _fetch_models_from_url(
client,
url=fallback_url,
headers=headers,
params={"key": api_key, "pageSize": _MODEL_PAGE_SIZE},
auth_config=auth_config,
fallback_publisher="google",
)
if success:
has_success = True
if err:
labeled = f"{fallback_url}: {err}"
if _is_soft_not_found(err):
soft_errors.append(labeled)
else:
hard_errors.append(labeled)
else:
all_models.extend(models)
deduped = _dedupe_models(all_models)
if deduped:
return deduped, hard_errors, has_success or True
if hard_errors:
return [], hard_errors, has_success
if soft_errors:
return [], [soft_errors[0]], has_success
return [], [], has_success
async def _fetch_models_vertex_service_account(
client: httpx.AsyncClient,
*,
ctx: Any,
auth_config: dict[str, Any] | None,
client_kwargs: dict[str, Any],
) -> tuple[list[dict[str, Any]], list[str], bool]:
if not isinstance(auth_config, dict):
return [], ["vertex_ai(service_account): missing auth_config"], False
try:
auth_service = VertexAuthService(auth_config)
access_token = await auth_service.get_access_token(httpx_client_kwargs=client_kwargs)
except VertexAuthError as exc:
return [], [f"vertex_ai(service_account): auth failed: {exc}"], False
except Exception as exc:
return [], [f"vertex_ai(service_account): auth failed: {exc}"], False
project_id = str(auth_config.get("project_id") or "").strip()
if not project_id:
return [], ["vertex_ai(service_account): missing project_id"], False
all_models: list[dict[str, Any]] = []
hard_errors: list[str] = []
soft_errors: list[str] = []
has_success = False
gemini_headers = {"Accept": "application/json", **_get_endpoint_headers(ctx, "gemini:chat")}
claude_headers = {"Accept": "application/json", **_get_endpoint_headers(ctx, "claude:chat")}
gemini_headers["Authorization"] = f"Bearer {access_token}"
claude_headers["Authorization"] = f"Bearer {access_token}"
for region in _iter_regions(auth_config):
base = (
_VERTEX_API_BASE
if region == "global"
else f"https://{region}-aiplatform.googleapis.com"
)
requests = [
(
"google",
f"{base}/v1/projects/{project_id}/locations/{region}/publishers/google/models",
gemini_headers,
),
(
"anthropic",
f"{base}/v1/projects/{project_id}/locations/{region}/publishers/anthropic/models",
claude_headers,
),
]
for publisher, url, headers in requests:
models, err, success = await _fetch_models_from_url(
client,
url=url,
headers=headers,
params={"pageSize": _MODEL_PAGE_SIZE},
auth_config=auth_config,
fallback_publisher=publisher,
)
if success:
has_success = True
if err:
labeled = f"{url}: {err}"
if _is_soft_not_found(err):
soft_errors.append(labeled)
else:
hard_errors.append(labeled)
continue
all_models.extend(models)
deduped = _dedupe_models(all_models)
if deduped:
return deduped, hard_errors, has_success or True
if hard_errors:
return [], hard_errors, has_success
if soft_errors:
return [], [soft_errors[0]], has_success
return [], [], has_success
async def fetch_models_vertex_ai(
ctx: Any,
timeout_seconds: float,
) -> tuple[list[dict], list[str], bool, dict[str, Any] | None]:
"""Vertex AI 专用模型获取链路。
- API Key: 优先请求 Vertex publisher models失败时兜底 Gemini Developer API
- Service Account: 使用 SA 凭证换取 Bearer Token按 region + publisher 查询
"""
from src.services.proxy_node.resolver import build_proxy_client_kwargs
auth_config = ctx.auth_config if isinstance(ctx.auth_config, dict) else None
is_service_account = _looks_like_service_account(auth_config)
client_kwargs = build_proxy_client_kwargs(ctx.proxy_config, timeout=timeout_seconds)
async with httpx.AsyncClient(**client_kwargs) as client:
if is_service_account:
models, errors, has_success = await _fetch_models_vertex_service_account(
client,
ctx=ctx,
auth_config=auth_config,
client_kwargs=client_kwargs,
)
else:
models, errors, has_success = await _fetch_models_vertex_api_key(
client,
ctx=ctx,
auth_config=auth_config,
)
if not models and errors:
logger.warning("Vertex 模型获取失败: {}", "; ".join(errors))
return models, errors, has_success, None
def register_all() -> None:
"""一次性注册 Vertex AI 的所有 hooks 到各通用 registry。"""
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
from src.services.provider.adapters.vertex_ai.transport import build_vertex_ai_url
from src.services.provider.behavior import register_behavior_variant
from src.services.provider.transport import register_transport_hook
# Transport: Vertex AI 同时支持 gemini:chat 和 claude:chat 格式
register_transport_hook("vertex_ai", "gemini:chat", build_vertex_ai_url)
register_transport_hook("vertex_ai", "claude:chat", build_vertex_ai_url)
# Model Fetcher: Vertex 走专用模型获取链路
UpstreamModelsFetcherRegistry.register(
provider_types=["vertex_ai"],
fetcher=fetch_models_vertex_ai,
)
# Behavior: 跨格式支持(同一 Vertex AI Provider 可同时访问 Gemini 和 Claude 模型)
register_behavior_variant("vertex_ai", cross_format=True)
__all__ = ["fetch_models_vertex_ai", "register_all"]

View File

@@ -0,0 +1,263 @@
"""Vertex AI URL 构建Transport Hook
根据 auth_type 选择两种完全不同的 URL 构建策略:
- API Key: 全局端点,简化路径
https://aiplatform.googleapis.com/v1/publishers/google/models/{model}:{action}?key={API_KEY}
- Service Account: 区域端点,完整路径
https://{region}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{region}/publishers/{publisher}/models/{model}:{action}
"""
from __future__ import annotations
import json
from typing import Any
from urllib.parse import urlencode
from src.core.logger import logger
from src.services.provider.adapters.vertex_ai.constants import (
API_KEY_BASE_URL,
DEFAULT_FORMAT,
DEFAULT_MODEL_REGIONS,
MODEL_FORMAT_MAPPING,
)
from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.transport import redact_url_for_log
def get_effective_format(
model: str,
auth_config: dict[str, Any] | None = None,
) -> str:
"""获取 Vertex AI 模式下模型的实际 API 格式。
优先级:
1. auth_config.model_format_mapping 中的精确匹配
2. auth_config.model_format_mapping 中的前缀匹配
3. 内置 MODEL_FORMAT_MAPPING 前缀匹配
4. auth_config.default_format
5. 内置 DEFAULT_FORMAT
"""
user_format_mapping: dict[str, str] = {}
user_default_format: str | None = None
if auth_config:
user_format_mapping = auth_config.get("model_format_mapping", {})
user_default_format = auth_config.get("default_format")
# 1. 用户配置:精确匹配
if model in user_format_mapping:
try:
return normalize_endpoint_signature(user_format_mapping[model])
except Exception:
logger.warning(
"Invalid vertex_ai model_format_mapping value for model '{}': {!r}",
model,
user_format_mapping[model],
)
# 2. 用户配置:前缀匹配
for prefix, api_format in user_format_mapping.items():
if prefix.endswith("-") and model.startswith(prefix):
try:
return normalize_endpoint_signature(api_format)
except Exception:
logger.warning(
"Invalid vertex_ai model_format_mapping value for prefix '{}': {!r}",
prefix,
api_format,
)
break
# 3. 内置配置:前缀匹配
for prefix, api_format in MODEL_FORMAT_MAPPING.items():
if model.startswith(prefix):
return normalize_endpoint_signature(api_format)
# 4. 用户默认格式
if user_default_format:
try:
return normalize_endpoint_signature(user_default_format)
except Exception:
logger.warning("Invalid vertex_ai default_format: {!r}", user_default_format)
# 5. 内置默认格式
return DEFAULT_FORMAT
def build_vertex_ai_url(
endpoint: Any,
*,
is_stream: bool,
effective_query_params: dict[str, Any],
path_params: dict[str, Any] | None = None,
key: Any = None,
decrypted_auth_config: dict[str, Any] | None = None,
) -> str:
"""Vertex AI transport hook — 统一 URL 构建入口。
根据 key.auth_type 分派到 API Key 或 Service Account 两种策略。
"""
auth_type = getattr(key, "auth_type", "api_key") if key else "api_key"
if auth_type == "api_key":
return _build_api_key_url(
key=key,
path_params=path_params,
query_params=effective_query_params,
is_stream=is_stream,
)
else:
# service_account以及向后兼容旧的 "vertex_ai" auth_type
return _build_service_account_url(
key=key,
path_params=path_params,
query_params=effective_query_params,
is_stream=is_stream,
decrypted_auth_config=decrypted_auth_config,
)
def _build_api_key_url(
key: Any,
*,
path_params: dict[str, Any] | None = None,
query_params: dict[str, Any] | None = None,
is_stream: bool = False,
) -> str:
"""构建 API Key 认证的全局端点 URL。
格式: https://aiplatform.googleapis.com/v1/publishers/google/models/{model}:{action}?key={API_KEY}
"""
from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException
model = (path_params or {}).get("model", "")
if not model:
raise InvalidRequestException("Vertex AI 请求缺少 model 参数")
if str(model).startswith("claude-"):
raise InvalidRequestException(
"Vertex API Key 不支持 Claude 模型,请改用 Service Account 认证。"
)
action = "streamGenerateContent" if is_stream else "generateContent"
path = f"/v1/publishers/google/models/{model}:{action}"
url = f"{API_KEY_BASE_URL}{path}"
# 构建查询参数
params = dict(query_params) if query_params else {}
# 附加 API Key
api_key_value = crypto_service.decrypt(key.api_key) if key else ""
if api_key_value:
params["key"] = api_key_value
# Gemini 流式请求使用 SSE
if is_stream:
params.setdefault("alt", "sse")
params.pop("beta", None)
if params:
query_string = urlencode(params, doseq=True)
if query_string:
url = f"{url}?{query_string}"
logger.debug("Vertex AI (API Key) URL: {}", redact_url_for_log(url))
return url
def _build_service_account_url(
key: Any,
*,
path_params: dict[str, Any] | None = None,
query_params: dict[str, Any] | None = None,
is_stream: bool = False,
decrypted_auth_config: dict[str, Any] | None = None,
) -> str:
"""构建 Service Account 认证的区域端点 URL。
格式: https://{region}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{region}/publishers/{publisher}/models/{model}:{action}
"""
from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException
# 优先使用传入的已解密配置,避免重复解密
auth_config: dict[str, Any] = {}
if decrypted_auth_config:
auth_config = decrypted_auth_config
else:
# 兜底:从 key.auth_config 解密(理论上不应走到这里)
raw_auth_config = getattr(key, "auth_config", None) if key else None
if raw_auth_config:
try:
if isinstance(raw_auth_config, dict):
auth_config = raw_auth_config
else:
decrypted_config = crypto_service.decrypt(raw_auth_config)
auth_config = json.loads(decrypted_config)
except Exception as e:
logger.error("解密 Vertex AI auth_config 失败: {}", e)
auth_config = {}
# 获取必需的配置
project_id = auth_config.get("project_id")
if not project_id:
raise InvalidRequestException(
"Vertex AI 配置缺少 project_id请在 Key 的 auth_config 中提供)"
)
# 获取模型名
model = (path_params or {}).get("model", "")
if not model:
raise InvalidRequestException("Vertex AI 请求缺少 model 参数")
# 确定 region优先级用户配置 > 内置默认 > 用户默认 > 兜底)
user_model_regions = auth_config.get("model_regions", {})
user_default_region = auth_config.get("region")
if model in user_model_regions:
region = user_model_regions[model]
elif model in DEFAULT_MODEL_REGIONS:
region = DEFAULT_MODEL_REGIONS[model]
elif user_default_region:
region = user_default_region
else:
region = "global"
# 判断是 Claude 还是 Gemini 模型
is_claude_model = model.startswith("claude-")
# 根据模型类型确定 publisher 和 action
if is_claude_model:
publisher = "anthropic"
action = "streamRawPredict" if is_stream else "rawPredict"
else:
publisher = "google"
action = "streamGenerateContent" if is_stream else "generateContent"
# 构建 URLglobal region 使用不同的 URL 格式)
if region == "global":
base_url = "https://aiplatform.googleapis.com"
else:
base_url = f"https://{region}-aiplatform.googleapis.com"
path = f"/v1/projects/{project_id}/locations/{region}/publishers/{publisher}/models/{model}:{action}"
url = f"{base_url}{path}"
# 添加查询参数
effective_query_params = dict(query_params) if query_params else {}
# Gemini 流式请求使用 SSE 格式Claude 不需要
if is_stream and not is_claude_model:
effective_query_params.setdefault("alt", "sse")
# 移除不适用于 Vertex AI 的参数
effective_query_params.pop("beta", None)
if effective_query_params:
query_string = urlencode(effective_query_params, doseq=True)
if query_string:
url = f"{url}?{query_string}"
logger.debug("Vertex AI (SA) URL: {} (region={})", redact_url_for_log(url), region)
return url

View File

@@ -393,58 +393,12 @@ async def get_provider_auth(
auth_value=f"Bearer {effective_token}", auth_value=f"Bearer {effective_token}",
decrypted_auth_config=decrypted_auth_config, decrypted_auth_config=decrypted_auth_config,
) )
if auth_type == "vertex_ai": if auth_type in ("service_account", "vertex_ai"):
from src.core.vertex_auth import VertexAuthError, VertexAuthService # service_account: GCP Service Account JSON → JWT → Access Token
# "vertex_ai" 保留为向后兼容(迁移期间旧数据可能仍使用该值)
from src.services.provider.adapters.vertex_ai.auth import _auth_service_account
try: return await _auth_service_account(key)
# 优先从 auth_config 读取,兼容从 api_key 读取(过渡期)
encrypted_auth_config = getattr(key, "auth_config", None)
if encrypted_auth_config:
# auth_config 可能是加密字符串或未加密的 dict
if isinstance(encrypted_auth_config, dict):
# 已经是 dict直接使用兼容未加密存储的情况
sa_json = encrypted_auth_config
else:
# 是加密字符串,需要解密
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
sa_json = json.loads(decrypted_config)
else:
# 兼容旧数据:从 api_key 读取
decrypted_key = crypto_service.decrypt(key.api_key)
# 检查是否是占位符(表示 auth_config 丢失)
if decrypted_key == "__placeholder__":
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
sa_json = json.loads(decrypted_key)
if not isinstance(sa_json, dict):
raise InvalidRequestException("Service Account JSON 无效,请重新添加该密钥。")
# 获取 Access Token注入代理配置core 层不依赖 services
from src.services.proxy_node.resolver import build_proxy_client_kwargs
service = VertexAuthService(sa_json)
access_token = await service.get_access_token(
httpx_client_kwargs=build_proxy_client_kwargs(timeout=30),
)
# Vertex AI 使用 Bearer token
return ProviderAuthInfo(
auth_header="Authorization",
auth_value=f"Bearer {access_token}",
decrypted_auth_config=sa_json,
)
except InvalidRequestException:
raise
except VertexAuthError as e:
raise InvalidRequestException(f"Vertex AI 认证失败:{e}")
except json.JSONDecodeError:
raise InvalidRequestException("Service Account JSON 格式无效,请重新添加该密钥。")
except Exception:
raise InvalidRequestException("Vertex AI 认证失败,请检查 Key 的 auth_config")
# 其他认证类型可在此扩展
# elif auth_type == "oauth2":
# ...
# 标准 API Key返回 None由 build_headers 处理 # 标准 API Key返回 None由 build_headers 处理
return None return None

View File

@@ -157,11 +157,13 @@ def ensure_providers_bootstrapped() -> None:
) )
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
from src.services.provider.adapters.vertex_ai.plugin import register_all as _reg_vertex_ai
_reg_antigravity() _reg_antigravity()
_reg_claude_code() _reg_claude_code()
_reg_codex() _reg_codex()
_reg_kiro() _reg_kiro()
_reg_vertex_ai()
__all__ = [ __all__ = [

View File

@@ -4,7 +4,7 @@
负责: 负责:
- 根据 API 格式或端点配置生成请求 URL - 根据 API 格式或端点配置生成请求 URL
- URL 脱敏(用于日志记录) - URL 脱敏(用于日志记录)
- Vertex AI URL 自动构建 - Provider transport hook 路由
""" """
from __future__ import annotations from __future__ import annotations
@@ -154,7 +154,7 @@ def build_provider_url(
根据 endpoint 配置生成请求 URL 根据 endpoint 配置生成请求 URL
优先级: 优先级:
1. Vertex AI 自动构建 - 当 key.auth_type == "vertex_ai" 1. Provider transport hook - 如有注册的 hook 则委托处理
2. endpoint.custom_path - 自定义路径(支持模板变量如 {model} 2. endpoint.custom_path - 自定义路径(支持模板变量如 {model}
3. API 格式默认路径 - 根据 api_format 自动选择 3. API 格式默认路径 - 根据 api_format 自动选择
@@ -169,17 +169,6 @@ def build_provider_url(
# 默认清理,避免上一次请求的 selected_base_url 泄漏到其他请求 # 默认清理,避免上一次请求的 selected_base_url 泄漏到其他请求
set_selected_base_url(None) set_selected_base_url(None)
# 检查是否为 Vertex AI 认证类型
auth_type = getattr(key, "auth_type", "api_key") if key else "api_key"
if auth_type == "vertex_ai":
return _build_vertex_ai_url(
key=key,
path_params=path_params,
query_params=query_params,
is_stream=is_stream,
decrypted_auth_config=decrypted_auth_config,
)
# endpoint signature新模式 # endpoint signature新模式
raw_family = getattr(endpoint, "api_family", None) raw_family = getattr(endpoint, "api_family", None)
raw_kind = getattr(endpoint, "endpoint_kind", None) raw_kind = getattr(endpoint, "endpoint_kind", None)
@@ -217,6 +206,9 @@ def build_provider_url(
endpoint, endpoint,
is_stream=is_stream, is_stream=is_stream,
effective_query_params=effective_query_params, effective_query_params=effective_query_params,
path_params=path_params,
key=key,
decrypted_auth_config=decrypted_auth_config,
) )
# 非 hook 路径:清除 contextvar避免跨请求污染 # 非 hook 路径:清除 contextvar避免跨请求污染
@@ -248,8 +240,8 @@ def build_provider_url(
path = _resolve_default_path(endpoint_sig) path = _resolve_default_path(endpoint_sig)
# Codex OAuth 端点chatgpt.com/backend-api/codex使用 /responses 而非 /v1/responses # Codex OAuth 端点chatgpt.com/backend-api/codex使用 /responses 而非 /v1/responses
base_url = getattr(endpoint, "base_url", "") or "" base_url = getattr(endpoint, "base_url", "") or ""
if endpoint_sig == "openai:cli" and is_codex_url(base_url): if endpoint_sig in {"openai:cli", "openai:compact"} and is_codex_url(base_url):
path = "/responses" path = "/responses/compact" if endpoint_sig == "openai:compact" else "/responses"
if effective_path_params: if effective_path_params:
try: try:
path = path.format(**effective_path_params) path = path.format(**effective_path_params)
@@ -286,248 +278,3 @@ def _resolve_default_path(endpoint_sig: str | None) -> str:
except Exception: except Exception:
logger.warning(f"Unknown endpoint signature '{endpoint_sig}' for endpoint, fallback to '/'") logger.warning(f"Unknown endpoint signature '{endpoint_sig}' for endpoint, fallback to '/'")
return "/" return "/"
# ==============================================================================
# Vertex AI 配置
# ==============================================================================
# Vertex AI 模型前缀到 API 格式的映射
# 用于 auth_type=vertex_ai 时,根据模型名动态确定实际的请求/响应格式
# 格式:前缀 -> endpoint signaturefamily:kind
VERTEX_AI_MODEL_FORMAT_MAPPING: dict[str, str] = {
"claude-": "claude:chat", # Anthropic Claude 模型
"gemini-": "gemini:chat", # Google Gemini 模型
"imagen-": "gemini:chat", # Google Imagen 模型(使用 Gemini chat 格式)
}
# Vertex AI 默认 endpoint signature当模型前缀不匹配时
VERTEX_AI_DEFAULT_FORMAT: str = "gemini:chat"
def get_vertex_ai_effective_format(
model: str,
auth_config: dict[str, Any] | None = None,
) -> str:
"""
获取 Vertex AI 模式下模型的实际 API 格式
优先级:
1. auth_config.model_format_mapping 中的精确匹配
2. auth_config.model_format_mapping 中的前缀匹配
3. 内置 VERTEX_AI_MODEL_FORMAT_MAPPING 前缀匹配
4. auth_config.default_format
5. 内置 VERTEX_AI_DEFAULT_FORMAT
auth_config 配置示例::
{
"project_id": "your-gcp-project-id",
"model_format_mapping": {
"claude-": "CLAUDE", # 前缀匹配
"my-custom-model": "OPENAI" # 精确匹配
},
"default_format": "GEMINI"
}
Args:
model: 模型名称
auth_config: 解密后的认证配置(可选),可包含 model_format_mapping 和 default_format
Returns:
实际应使用的 endpoint signature"claude:chat", "gemini:chat"
"""
# 用户配置的模型-格式映射
user_format_mapping: dict[str, str] = {}
user_default_format: str | None = None
if auth_config:
user_format_mapping = auth_config.get("model_format_mapping", {})
user_default_format = auth_config.get("default_format")
# 1. 用户配置:精确匹配
if model in user_format_mapping:
try:
return normalize_endpoint_signature(user_format_mapping[model])
except Exception:
logger.warning(
"Invalid vertex_ai model_format_mapping value for model '{}': {!r}",
model,
user_format_mapping[model],
)
# 2. 用户配置:前缀匹配
for prefix, api_format in user_format_mapping.items():
if prefix.endswith("-") and model.startswith(prefix):
try:
return normalize_endpoint_signature(api_format)
except Exception:
logger.warning(
"Invalid vertex_ai model_format_mapping value for prefix '{}': {!r}",
prefix,
api_format,
)
break
# 3. 内置配置:前缀匹配
for prefix, api_format in VERTEX_AI_MODEL_FORMAT_MAPPING.items():
if model.startswith(prefix):
return normalize_endpoint_signature(api_format)
# 4. 用户默认格式
if user_default_format:
try:
return normalize_endpoint_signature(user_default_format)
except Exception:
logger.warning("Invalid vertex_ai default_format: {!r}", user_default_format)
# 5. 内置默认格式
return VERTEX_AI_DEFAULT_FORMAT
# Vertex AI 模型默认 region 映射
# 用户可以通过 auth_config.model_regions 覆盖
VERTEX_AI_DEFAULT_MODEL_REGIONS: dict[str, str] = {
# Gemini 3 系列(使用 global
"gemini-3-pro-image-preview": "global",
# Gemini 2.0 系列
"gemini-2.0-flash": "us-central1",
"gemini-2.0-flash-exp": "us-central1",
"gemini-2.0-flash-001": "us-central1",
"gemini-2.0-pro-exp": "us-central1",
"gemini-2.0-flash-exp-image-generation": "us-central1",
# Gemini 1.5 系列
"gemini-1.5-pro": "us-central1",
"gemini-1.5-pro-001": "us-central1",
"gemini-1.5-pro-002": "us-central1",
"gemini-1.5-flash": "us-central1",
"gemini-1.5-flash-001": "us-central1",
"gemini-1.5-flash-002": "us-central1",
# Imagen 系列
"imagen-3.0-generate-001": "us-central1",
"imagen-3.0-fast-generate-001": "us-central1",
}
def _build_vertex_ai_url(
key: "ProviderAPIKey",
*,
path_params: dict[str, Any] | None = None,
query_params: dict[str, Any] | None = None,
is_stream: bool = False,
decrypted_auth_config: dict[str, Any] | None = None,
) -> str:
"""
构建 Vertex AI URL
Vertex AI URL 格式:
- Gemini: https://{region}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{region}/publishers/google/models/{model}:{action}
- Claude: https://{region}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{region}/publishers/anthropic/models/{model}:{action}
从 auth_config 中读取:
- project_id: GCP 项目 ID必需
- region: 默认 GCP 区域(覆盖内置默认值)
- model_regions: 模型到区域的映射(可选),覆盖内置和默认配置
Region 优先级:
1. auth_config.model_regions[model] - 用户为该模型指定的区域
2. VERTEX_AI_DEFAULT_MODEL_REGIONS[model] - 内置的模型默认区域
3. auth_config.region - 用户配置的默认区域
4. global - 最终兜底
Args:
key: Provider API Key包含 auth_config
path_params: 路径参数(需要 model
query_params: 查询参数
is_stream: 是否为流式请求
decrypted_auth_config: 已解密的认证配置(由 get_provider_auth 提供,避免重复解密)
Returns:
完整的 Vertex AI URL
"""
import json
from src.core.crypto import crypto_service
# 优先使用传入的已解密配置,避免重复解密
auth_config: dict[str, Any] = {}
if decrypted_auth_config:
auth_config = decrypted_auth_config
else:
# 兜底:从 key.auth_config 解密(理论上不应走到这里)
raw_auth_config = getattr(key, "auth_config", None)
if raw_auth_config:
try:
# auth_config 可能是加密字符串或未加密的 dict
if isinstance(raw_auth_config, dict):
auth_config = raw_auth_config
else:
decrypted_config = crypto_service.decrypt(raw_auth_config)
auth_config = json.loads(decrypted_config)
except Exception as e:
logger.error(f"解密 Vertex AI auth_config 失败: {e}")
auth_config = {}
from src.core.exceptions import InvalidRequestException
# 获取必需的配置
project_id = auth_config.get("project_id")
if not project_id:
raise InvalidRequestException(
"Vertex AI 配置缺少 project_id请在 Key 的 auth_config 中提供)"
)
# 获取模型名
model = (path_params or {}).get("model", "")
if not model:
raise InvalidRequestException("Vertex AI 请求缺少 model 参数")
# 确定 region优先级用户配置 > 内置默认 > 用户默认 > 兜底)
user_model_regions = auth_config.get("model_regions", {})
user_default_region = auth_config.get("region")
if model in user_model_regions:
region = user_model_regions[model]
elif model in VERTEX_AI_DEFAULT_MODEL_REGIONS:
region = VERTEX_AI_DEFAULT_MODEL_REGIONS[model]
elif user_default_region:
region = user_default_region
else:
region = "global"
# 判断是 Claude 还是 Gemini 模型
is_claude_model = model.startswith("claude-")
# 根据模型类型确定 publisher 和 action
if is_claude_model:
# Claude 模型使用 Anthropic publisher
publisher = "anthropic"
action = "streamRawPredict" if is_stream else "rawPredict"
else:
# Gemini 模型使用 Google publisher
publisher = "google"
action = "streamGenerateContent" if is_stream else "generateContent"
# 构建 URLglobal region 使用不同的 URL 格式)
if region == "global":
base_url = "https://aiplatform.googleapis.com"
else:
base_url = f"https://{region}-aiplatform.googleapis.com"
path = f"/v1/projects/{project_id}/locations/{region}/publishers/{publisher}/models/{model}:{action}"
url = f"{base_url}{path}"
# 添加查询参数
effective_query_params = dict(query_params) if query_params else {}
# Gemini 流式请求使用 SSE 格式Claude 不需要
if is_stream and not is_claude_model:
effective_query_params.setdefault("alt", "sse")
# 移除不适用于 Vertex AI 的参数
effective_query_params.pop("beta", None)
if effective_query_params:
query_string = urlencode(effective_query_params, doseq=True)
if query_string:
url = f"{url}?{query_string}"
logger.debug(f"Vertex AI URL: {redact_url_for_log(url)} (region={region})")
return url

View File

@@ -0,0 +1,150 @@
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from src.core.vertex_auth import VertexAuthService
from src.services.model.upstream_fetcher import (
EndpointFetchConfig,
UpstreamModelsFetchContext,
fetch_models_for_key,
)
@pytest.mark.asyncio
async def test_fetch_models_for_key_vertex_api_key_custom_fetcher() -> None:
ctx = UpstreamModelsFetchContext(
provider_type="vertex_ai",
api_key_value="test-api-key",
format_to_endpoint={
"gemini:chat": EndpointFetchConfig(base_url="https://aiplatform.googleapis.com"),
},
proxy_config=None,
auth_config={},
)
mocked_models = [
{
"id": "gemini-2.5-pro",
"owned_by": "google",
"display_name": "Gemini 2.5 Pro",
"api_format": "gemini:chat",
}
]
with (
patch(
"src.services.provider.adapters.vertex_ai.plugin._fetch_models_from_url",
AsyncMock(return_value=(mocked_models, None, True)),
),
patch(
"src.services.proxy_node.resolver.build_proxy_client_kwargs",
return_value={"timeout": 1.0},
),
):
models, errors, ok, meta = await fetch_models_for_key(ctx, timeout_seconds=1.0)
assert ok is True
assert errors == []
assert meta is None
assert [m.get("id") for m in models] == ["gemini-2.5-pro"]
@pytest.mark.asyncio
async def test_fetch_models_for_key_vertex_service_account_ignores_soft_404_when_success() -> None:
auth_config = {
"project_id": "demo-project",
"client_email": "svc@example.iam.gserviceaccount.com",
"private_key": "-----BEGIN PRIVATE KEY-----\nTEST\n-----END PRIVATE KEY-----\n",
"region": "global",
}
ctx = UpstreamModelsFetchContext(
provider_type="vertex_ai",
api_key_value="__placeholder__",
format_to_endpoint={
"gemini:chat": EndpointFetchConfig(base_url="https://aiplatform.googleapis.com"),
"claude:chat": EndpointFetchConfig(base_url="https://aiplatform.googleapis.com"),
},
proxy_config=None,
auth_config=auth_config,
)
fetch_side_effect = [
(
[
{
"id": "gemini-2.0-flash",
"owned_by": "google",
"display_name": "Gemini 2.0 Flash",
"api_format": "gemini:chat",
}
],
None,
True,
),
([], "HTTP 404: not found", False),
]
with (
patch.object(
VertexAuthService,
"get_access_token",
AsyncMock(return_value="ya29.test-token"),
),
patch(
"src.services.provider.adapters.vertex_ai.plugin._iter_regions",
return_value=["global"],
),
patch(
"src.services.provider.adapters.vertex_ai.plugin._fetch_models_from_url",
AsyncMock(side_effect=fetch_side_effect),
),
patch(
"src.services.proxy_node.resolver.build_proxy_client_kwargs",
return_value={"timeout": 1.0},
),
):
models, errors, ok, meta = await fetch_models_for_key(ctx, timeout_seconds=1.0)
assert ok is True
assert errors == []
assert meta is None
ids = {m.get("id") for m in models}
assert "gemini-2.0-flash" in ids
@pytest.mark.asyncio
async def test_fetch_models_for_key_vertex_api_key_returns_soft_404_when_all_failed() -> None:
ctx = UpstreamModelsFetchContext(
provider_type="vertex_ai",
api_key_value="test-api-key",
format_to_endpoint={
"gemini:chat": EndpointFetchConfig(base_url="https://aiplatform.googleapis.com"),
},
proxy_config=None,
auth_config={},
)
with (
patch(
"src.services.provider.adapters.vertex_ai.plugin._fetch_models_from_url",
AsyncMock(
side_effect=[
([], "HTTP 404: not found", False),
([], "HTTP 404: not found", False),
]
),
),
patch(
"src.services.proxy_node.resolver.build_proxy_client_kwargs",
return_value={"timeout": 1.0},
),
):
models, errors, ok, meta = await fetch_models_for_key(ctx, timeout_seconds=1.0)
assert ok is False
assert models == []
assert meta is None
assert errors
assert "HTTP 404" in errors[0]