Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion electron/services/agent/cache.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import { createHash } from 'crypto'
import type { ProviderOptions, SystemModelMessage } from '@ai-sdk/provider-utils'
import type { ToolSet } from 'ai'
import { isArkBaseURL } from './arkContextFetch'
import { isNvidiaInferenceBaseURL } from './promptCacheCompat'
import type { AgentReasoningEffort, AgentRunInput } from './types'

export interface AgentPromptParts {
Expand Down Expand Up @@ -146,7 +147,10 @@ export function buildProviderOptions(input: AgentRunInput, promptCacheKey: strin
} else {
// @ai-sdk/openai-compatible 只校验 camelCase 标准项;厂商扩展字段要用请求体原名透传。
// 这里恢复 AI SDK 6 时代 OpenAI-compatible 服务常用的 prompt_cache_key 行为。
option.prompt_cache_key = promptCacheKey
// 英伟达推理接口(nvidia.com)不认该字段,命中时跳过注入,避免 400。见 issue #353。
if (!isNvidiaInferenceBaseURL(input.providerConfig.baseURL)) {
option.prompt_cache_key = promptCacheKey
}
}
if (Object.keys(option).length > 0) {
const keys = new Set(['openai'])
Expand Down
34 changes: 34 additions & 0 deletions electron/services/agent/promptCacheCompat.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,34 @@
/**
* prompt cache 兼容性辅助:集中处理「哪些端点不支持 `prompt_cache_key`」。
*
* 背景:英伟达的 OpenAI 兼容推理接口(integrate.api.nvidia.com / api.nvidia.com)
* 不认 OpenAI 的 `prompt_cache_key` 扩展字段,带上会直接 400
* (Unsupported parameter(s): prompt_cache_key,见 issue #353)。
*
* 该字段会从两条路径进入请求体,必须同时拦截:
* 1. provider.ts 的 transformRequestBody(直接写请求体)
* 2. cache.ts 的 buildProviderOptions(经 providerOptions 透传)
* 本模块是唯一真相来源,两处都 import 它,避免逻辑漂移。
*
* 零依赖:可被单元测试单独导入,无需安装 Electron 依赖树。
*/

/** 判断 baseURL 是否属于英伟达推理接口。 */
export function isNvidiaInferenceBaseURL(baseURL?: string): boolean {
return !!baseURL && /(^|\.)nvidia\.com(\/|$)/i.test(baseURL)
}

/**
* 仅对支持 prompt cache 的 OpenAI 兼容端点注入 `prompt_cache_key`;
* 英伟达推理接口不支持该字段,命中时跳过注入,避免 400。
* 纯函数,导出供单元测试。
*/
export function injectOpenAICompatiblePromptCacheKey(
args: Record<string, any>,
promptCacheKey?: string,
baseURL?: string,
): Record<string, any> {
if (!promptCacheKey || args.prompt_cache_key) return args
if (isNvidiaInferenceBaseURL(baseURL)) return args
return { ...args, prompt_cache_key: promptCacheKey }
}
8 changes: 2 additions & 6 deletions electron/services/agent/provider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ import { CODEX_SUBSCRIPTION_DUMMY_API_KEY, createCodexSubscriptionFetch, getCode
import { withOpenAICompatibleStreamSanitizer } from '../ai/openaiCompatibleStreamSanitizer'
import { withGoogleExplicitCache } from './googleCacheFetch'
import { isArkBaseURL, withArkContextCache } from './arkContextFetch'
import { injectOpenAICompatiblePromptCacheKey } from './promptCacheCompat'
import type { AgentProviderConfig } from './types'

export type AgentLanguageModelOptions = {
Expand Down Expand Up @@ -100,11 +101,6 @@ function withAnthropicSanitizer(baseFetch: typeof globalThis.fetch | undefined):
}) as typeof globalThis.fetch
}

function injectOpenAICompatiblePromptCacheKey(args: Record<string, any>, promptCacheKey?: string): Record<string, any> {
if (!promptCacheKey || args.prompt_cache_key) return args
return { ...args, prompt_cache_key: promptCacheKey }
}

export function createLanguageModel(config: AgentProviderConfig, options: AgentLanguageModelOptions = {}): LanguageModel {
const { providerKind, name, apiKey, baseURL, model, headers, proxyUrl } = config
const fetch = createProxyFetch(proxyUrl)
Expand Down Expand Up @@ -141,7 +137,7 @@ export function createLanguageModel(config: AgentProviderConfig, options: AgentL
includeUsage: true,
// 火山方舟端点:system 前缀自动走 context 缓存,见 arkContextFetch.ts
fetch: withOpenAICompatibleStreamSanitizer(compatibleFetch),
transformRequestBody: (args) => injectOpenAICompatiblePromptCacheKey(args, options.promptCacheKey),
transformRequestBody: (args) => injectOpenAICompatiblePromptCacheKey(args, options.promptCacheKey, baseURL),
}).chatModel(model)
}

Expand Down
33 changes: 33 additions & 0 deletions scripts/test-nvidia-prompt-cache.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
import assert from 'node:assert/strict'
import { injectOpenAICompatiblePromptCacheKey, isNvidiaInferenceBaseURL } from '../electron/services/agent/promptCacheCompat.ts'

// ---- isNvidiaInferenceBaseURL ----
assert.equal(isNvidiaInferenceBaseURL('https://integrate.api.nvidia.com/v1'), true, 'integrate.api.nvidia.com/v1 应判为英伟达')
assert.equal(isNvidiaInferenceBaseURL('https://api.nvidia.com/v1'), true, 'api.nvidia.com/v1 应判为英伟达')
assert.equal(isNvidiaInferenceBaseURL('https://integrate.api.nvidia.com'), true, '无路径的 integrate.api.nvidia.com 也应判为英伟达')
assert.equal(isNvidiaInferenceBaseURL('https://api.openai.com/v1'), false, 'OpenAI 不应判为英伟达')
assert.equal(isNvidiaInferenceBaseURL('https://api.deepseek.com'), false, 'DeepSeek 不应判为英伟达')
assert.equal(isNvidiaInferenceBaseURL('https://ark.cn/v1'), false, '火山方舟不应被误判为英伟达(方舟需保留 prompt_cache_key)')
assert.equal(isNvidiaInferenceBaseURL(undefined), false, '无 baseURL 不应判为英伟达')

// ---- injectOpenAICompatiblePromptCacheKey:NVIDIA 必须跳过,其它端点保留 ----
const baseArgs = { model: 'x', messages: [] }

const nvidiaResult = injectOpenAICompatiblePromptCacheKey(baseArgs, 'cache-key', 'https://integrate.api.nvidia.com/v1')
assert.equal('prompt_cache_key' in nvidiaResult, false, 'NVIDIA 请求体不得含 prompt_cache_key')
assert.equal(nvidiaResult, baseArgs, 'NVIDIA 命中 skip 分支时应原样返回(不新增字段、不复制)')

const openaiResult = injectOpenAICompatiblePromptCacheKey(baseArgs, 'cache-key', 'https://api.openai.com/v1')
assert.equal(openaiResult.prompt_cache_key, 'cache-key', 'OpenAI 应保留 prompt_cache_key')

const deepseekResult = injectOpenAICompatiblePromptCacheKey(baseArgs, 'cache-key', 'https://api.deepseek.com')
assert.equal(deepseekResult.prompt_cache_key, 'cache-key', 'DeepSeek 应保留 prompt_cache_key')

const arkResult = injectOpenAICompatiblePromptCacheKey(baseArgs, 'cache-key', 'https://ark.cn/v1')
assert.equal(arkResult.prompt_cache_key, 'cache-key', '火山方舟应保留 prompt_cache_key(context 缓存依赖它)')

// 未提供 key 时原样返回(不新增字段)
assert.equal(injectOpenAICompatiblePromptCacheKey(baseArgs, undefined, 'https://integrate.api.nvidia.com/v1'), baseArgs, '无 key 时原样返回')
assert.equal(injectOpenAICompatiblePromptCacheKey(baseArgs, undefined, 'https://api.openai.com/v1'), baseArgs, '无 key 时原样返回')

console.log('nvidia prompt_cache_key compat tests passed')