fix(imagegen): preserve selected provider routing

This commit is contained in:
程序员阿江(Relakkes)
2026-08-05 01:57:22 +08:00
parent d11a7410d3
commit 563003ec95
8 changed files with 347 additions and 70 deletions
+191
View File
@@ -30,6 +30,7 @@ import {
import { SessionService, sessionService } from '../services/sessionService.js'
import { ProviderService } from '../services/providerService.js'
import { resetTerminalShellEnvironmentCacheForTests } from '../../utils/terminalShellEnvironment.js'
import * as openAIModelCatalog from '../../services/openaiAuth/modelCatalog.js'
async function rmWithRetry(targetPath: string): Promise<void> {
const attempts = process.platform === 'win32' ? 5 : 1
@@ -3644,6 +3645,196 @@ describe('WebSocket Chat Integration', () => {
}
}, 20_000)
it('should not send a turn to the custom runtime while OpenAI runtime validation is pending', async () => {
const providerService = new ProviderService()
const customProvider = await providerService.addProvider({
presetId: 'custom',
name: 'Custom Images Before OpenAI',
apiKey: 'custom-chat-key',
baseUrl: 'https://custom-chat.example.test',
apiFormat: 'anthropic',
models: {
main: 'custom-main',
haiku: 'custom-main',
sonnet: 'custom-main',
opus: 'custom-main',
},
imageGeneration: {
model: 'custom-image-model',
baseUrl: 'https://custom-images.example.test/v1',
apiKey: 'custom-image-key',
},
})
await providerService.activateProvider(customProvider.id)
const createRes = await fetch(`${baseUrl}/api/sessions`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ workDir: process.cwd() }),
})
expect(createRes.status).toBe(201)
const { sessionId } = await createRes.json() as { sessionId: string }
const originalStartSession = conversationService.startSession.bind(conversationService)
const originalSendMessage = conversationService.sendMessage.bind(conversationService)
const startCalls: Array<{
providerId: string | null | undefined
model: string | undefined
imageProviderKind: string | undefined
}> = []
const sendCalls: Array<{
content: string
activeProviderId: string | null | undefined
}> = []
conversationService.startSession = (async function patchedStartSession(
sid: string,
workDir: string,
sdkUrl: string,
options?: { permissionMode?: string; model?: string; effort?: string; thinking?: 'enabled' | 'adaptive' | 'disabled'; providerId?: string | null },
) {
const env = await (conversationService as any).buildChildEnv(
workDir,
sdkUrl,
options,
) as Record<string, string>
startCalls.push({
providerId: options?.providerId,
model: options?.model,
imageProviderKind: env.CC_HAHA_IMAGE_PROVIDER_KIND,
})
return originalStartSession(sid, workDir, sdkUrl, options)
}) as typeof conversationService.startSession
conversationService.sendMessage = (function patchedSendMessage(
sid: string,
content: string,
attachments?: any,
) {
sendCalls.push({
content,
activeProviderId: startCalls.at(-1)?.providerId,
})
return originalSendMessage(sid, content, attachments)
}) as typeof conversationService.sendMessage
let markValidationStarted!: () => void
const validationStarted = new Promise<void>((resolve) => {
markValidationStarted = resolve
})
let releaseValidation!: () => void
const validationGate = new Promise<Awaited<ReturnType<typeof openAIModelCatalog.getOpenAICodexModelCatalog>>>((resolve) => {
releaseValidation = () => resolve([{
value: 'gpt-5.6-sol',
label: 'GPT-5.6 Sol',
description: 'Test model',
defaultReasoningEffort: 'medium',
supportedReasoningEfforts: ['low', 'medium', 'high'],
}])
})
const modelCatalogSpy = spyOn(
openAIModelCatalog,
'getOpenAICodexModelCatalog',
).mockImplementation(() => {
markValidationStarted()
return validationGate
})
const ws = new WebSocket(`${wsUrl}/ws/${sessionId}`)
try {
await new Promise<void>((resolve, reject) => {
const timeout = setTimeout(() => {
reject(new Error(`Timed out connecting runtime validation session ${sessionId}`))
}, 5_000)
ws.onmessage = (event) => {
const msg = JSON.parse(event.data as string)
if (msg.type === 'connected') {
clearTimeout(timeout)
ws.send(JSON.stringify({ type: 'prewarm_session' }))
resolve()
}
}
ws.onerror = () => {
clearTimeout(timeout)
reject(new Error(`WebSocket error for runtime validation session ${sessionId}`))
}
})
await waitUntil(
() => startCalls.length === 1 && conversationService.hasSession(sessionId),
`prewarmed custom runtime for ${sessionId}`,
)
expect(startCalls[0]).toEqual({
providerId: customProvider.id,
model: undefined,
imageProviderKind: 'openai_images',
})
const completion = new Promise<void>((resolve, reject) => {
const timeout = setTimeout(() => {
reject(new Error(`Timed out waiting for OpenAI runtime turn ${sessionId}`))
}, 15_000)
ws.onmessage = (event) => {
const msg = JSON.parse(event.data as string)
if (msg.type === 'error') {
clearTimeout(timeout)
reject(new Error(msg.message))
return
}
if (msg.type === 'message_complete') {
clearTimeout(timeout)
resolve()
}
}
})
ws.send(JSON.stringify({
type: 'set_runtime_config',
providerId: 'openai-official',
modelId: 'gpt-5.6-sol',
effortLevel: 'low',
}))
ws.send(JSON.stringify({
type: 'set_runtime_config',
providerId: 'openai-official',
modelId: 'gpt-5.6-sol',
effortLevel: 'low',
}))
ws.send(JSON.stringify({
type: 'user_message',
content: 'generate only after OpenAI runtime validation',
}))
await validationStarted
await new Promise((resolve) => setTimeout(resolve, 50))
expect(startCalls).toHaveLength(1)
expect(sendCalls).toHaveLength(0)
releaseValidation()
await completion
expect(startCalls).toHaveLength(2)
expect(startCalls[1]).toEqual({
providerId: 'openai-official',
model: 'gpt-5.6-sol',
imageProviderKind: 'openai_oauth',
})
expect(sendCalls).toEqual([{
content: 'generate only after OpenAI runtime validation',
activeProviderId: 'openai-official',
}])
} finally {
releaseValidation()
modelCatalogSpy.mockRestore()
ws.close()
conversationService.startSession = originalStartSession
conversationService.sendMessage = originalSendMessage
conversationService.stopSession(sessionId)
await providerService.activateOfficial()
await providerService.deleteProvider(customProvider.id)
}
}, 20_000)
it('should keep the session idle in the UI while applying a runtime-only model switch', async () => {
const providerService = new ProviderService()
const provider = await providerService.addProvider({
+69 -66
View File
@@ -1426,8 +1426,8 @@ async function handleSetRuntimeConfig(
message: Extract<ClientMessage, { type: 'set_runtime_config' }>
) {
const { sessionId } = ws.data
let modelId = typeof message.modelId === 'string' ? message.modelId.trim() : ''
if (!modelId) {
const requestedModelId = typeof message.modelId === 'string' ? message.modelId.trim() : ''
if (!requestedModelId) {
sendMessage(ws, {
type: 'error',
message: 'Runtime model selection is invalid.',
@@ -1435,70 +1435,73 @@ async function handleSetRuntimeConfig(
})
return
}
if (isGrokOfficialProviderId(message.providerId)) {
modelId = (await getGrokReasoningEfforts(modelId)).modelId
}
const effortLevel =
const requestedEffort =
typeof message.effortLevel === 'string' ? message.effortLevel.trim() : undefined
const effortResolution = effortLevel === undefined
? { valid: true, effort: undefined }
: await resolveRuntimeEffort(message.providerId, modelId, effortLevel)
if (!effortResolution.valid) {
sendMessage(ws, {
type: 'error',
message: 'Runtime effort selection is invalid.',
code: 'RUNTIME_CONFIG_INVALID',
})
return
}
const nextOverride = {
providerId: message.providerId ?? null,
modelId,
...(effortResolution.effort ? { effort: effortResolution.effort } : {}),
}
const prevOverride = runtimeOverrides.get(sessionId)
if (
prevOverride &&
prevOverride.providerId === nextOverride.providerId &&
prevOverride.modelId === nextOverride.modelId &&
prevOverride.effort === nextOverride.effort
) {
return
}
runtimeOverrides.set(sessionId, nextOverride)
runtimeOverrideVersions.set(
sessionId,
(runtimeOverrideVersions.get(sessionId) ?? 0) + 1,
)
if (shouldDeferRuntimeRestartForActiveTurn(sessionId)) {
deferredRuntimeRestarts.set(sessionId, nextOverride)
await persistSessionRuntimeConfig(sessionId, nextOverride)
return
}
if (conversationService.hasSession(sessionId)) {
await enqueueRuntimeTransition(sessionId, async () => {
await persistSessionRuntimeConfig(sessionId, nextOverride)
await restartSessionWithRuntimeConfig(ws, sessionId)
})
return
}
const pendingStartup = sessionStartupPromises.get(sessionId)
if (pendingStartup) {
const startupRuntimeVersion = sessionStartupRuntimeVersions.get(sessionId) ?? 0
const currentRuntimeVersion = runtimeOverrideVersions.get(sessionId) ?? 0
if (startupRuntimeVersion >= currentRuntimeVersion) {
await persistSessionRuntimeConfig(sessionId, nextOverride)
await pendingStartup
broadcastAppliedRuntimeConfig(sessionId)
// Register the transition before remote model-catalog or provider validation.
// A user message arriving in that async admission window must wait for the
// selected runtime instead of entering the previous provider's CLI process.
await enqueueRuntimeTransition(sessionId, async () => {
let modelId = requestedModelId
if (isGrokOfficialProviderId(message.providerId)) {
modelId = (await getGrokReasoningEfforts(modelId)).modelId
}
const effortResolution = requestedEffort === undefined
? { valid: true, effort: undefined }
: await resolveRuntimeEffort(message.providerId, modelId, requestedEffort)
if (!effortResolution.valid) {
sendMessage(ws, {
type: 'error',
message: 'Runtime effort selection is invalid.',
code: 'RUNTIME_CONFIG_INVALID',
})
return
}
await enqueueRuntimeTransition(sessionId, async () => {
const nextOverride = {
providerId: message.providerId ?? null,
modelId,
...(effortResolution.effort ? { effort: effortResolution.effort } : {}),
}
const prevOverride = runtimeOverrides.get(sessionId)
if (
prevOverride &&
prevOverride.providerId === nextOverride.providerId &&
prevOverride.modelId === nextOverride.modelId &&
prevOverride.effort === nextOverride.effort
) {
return
}
runtimeOverrides.set(sessionId, nextOverride)
runtimeOverrideVersions.set(
sessionId,
(runtimeOverrideVersions.get(sessionId) ?? 0) + 1,
)
if (shouldDeferRuntimeRestartForActiveTurn(sessionId)) {
deferredRuntimeRestarts.set(sessionId, nextOverride)
await persistSessionRuntimeConfig(sessionId, nextOverride)
return
}
if (conversationService.hasSession(sessionId)) {
await persistSessionRuntimeConfig(sessionId, nextOverride)
await restartSessionWithRuntimeConfig(ws, sessionId)
return
}
const pendingStartup = sessionStartupPromises.get(sessionId)
if (pendingStartup) {
const startupRuntimeVersion = sessionStartupRuntimeVersions.get(sessionId) ?? 0
const currentRuntimeVersion = runtimeOverrideVersions.get(sessionId) ?? 0
if (startupRuntimeVersion >= currentRuntimeVersion) {
await persistSessionRuntimeConfig(sessionId, nextOverride)
await pendingStartup
broadcastAppliedRuntimeConfig(sessionId)
return
}
await persistSessionRuntimeConfig(sessionId, nextOverride)
await pendingStartup.catch(() => undefined)
const currentOverride = runtimeOverrides.get(sessionId)
@@ -1511,12 +1514,12 @@ async function handleSetRuntimeConfig(
return
}
await restartSessionWithRuntimeConfig(ws, sessionId)
})
return
}
return
}
await persistSessionRuntimeConfig(sessionId, nextOverride)
broadcastAppliedRuntimeConfig(sessionId)
await persistSessionRuntimeConfig(sessionId, nextOverride)
broadcastAppliedRuntimeConfig(sessionId)
})
}
async function restartSessionWithPermissionMode(
+2
View File
@@ -59,6 +59,8 @@ describe('bundled imagegen skill', () => {
expect(text).toContain('one call per image')
expect(text).toContain('omit <code>input_images</code> entirely')
expect(text).toContain('Never pass <code>/dev/null</code>')
expect(text).toContain('otherwise omit the field')
expect(text).toContain('Never pass <code>default</code>')
expect(text).not.toContain('CC_HAHA_IMAGE_API_KEY')
})
})
+1 -1
View File
@@ -19,7 +19,7 @@ Use the built-in `ImageGen` tool. Provider authentication, model routing, output
- For multi-turn editing, use the latest selected output as the next turn's `edit_target`. Repeat all identity, layout, text, and unchanged-region constraints on every turn so edits do not drift.
- To edit several images independently, make one call per image. Put multiple images in one call only when the user wants them combined or used together as references. A single call accepts at most three source images.
- Prefer a useful default composition when the user leaves details open. Do not invent branding, logos, or people they did not request.
- Respect an explicitly requested provider model by passing `model`; otherwise omit it so the configured default is used.
- Respect an explicitly requested concrete provider model ID by passing `model`; otherwise omit the field so the configured model is used. Never pass `default` or another placeholder model name.
- If the provider or tool returns an error, do not retry `ImageGen` automatically. Explain the failure and let the user decide whether to retry or change providers.
## Build the prompt
@@ -97,5 +97,7 @@ describe('ImageGenTool', () => {
)
expect(prompt).toContain('omit input_images entirely')
expect(prompt).toContain('never pass /dev/null')
expect(prompt).toContain('Omit model unless')
expect(prompt).toContain('never pass "default"')
})
})
+2 -2
View File
@@ -59,7 +59,7 @@ const inputSchema = lazySchema(() =>
.string()
.min(1)
.optional()
.describe('Override the configured image model only when the user asks for one'),
.describe('Concrete image model ID to override the configured model only when the user explicitly asks; otherwise omit this field and never use "default" as a placeholder'),
aspect_ratio: z
.enum(ASPECT_RATIOS)
.optional()
@@ -110,7 +110,7 @@ export const ImageGenTool = buildTool({
return 'Generate one or more images with the image provider configured for this desktop session.'
},
async prompt() {
return `Use this tool when the user asks to generate or edit an image. For a brand-new image, omit input_images entirely; never pass /dev/null or another placeholder path. For edits, pass ordered input_images using only paths surfaced by [Image source: ...] in the current conversation or returned by a prior ImageGen call; repeat preservation constraints in every edit prompt. One call represents one distinct prompt; use count only for variations of that same prompt. The tool saves finished raster images locally and returns their absolute paths. If a provider call fails, do not retry ImageGen automatically; explain the error and wait for the user to decide.`
return `Use this tool when the user asks to generate or edit an image. For a brand-new image, omit input_images entirely; never pass /dev/null or another placeholder path. Omit model unless the user explicitly requests a concrete image model ID; never pass "default" as a placeholder. For edits, pass ordered input_images using only paths surfaced by [Image source: ...] in the current conversation or returned by a prior ImageGen call; repeat preservation constraints in every edit prompt. One call represents one distinct prompt; use count only for variations of that same prompt. The tool saves finished raster images locally and returns their absolute paths. If a provider call fails, do not retry ImageGen automatically; explain the error and wait for the user to decide.`
},
get inputSchema(): InputSchema {
return inputSchema()
+75
View File
@@ -3,6 +3,7 @@ import { mkdtemp, readFile, rm, writeFile } from 'fs/promises'
import { tmpdir } from 'os'
import { join } from 'path'
import { OPENAI_CODEX_OAUTH_FILE_ENV_KEY } from '../../services/openaiAuth/storage.js'
import type { ImageGenerationRuntimeConfig } from '../../services/imageGeneration/config.js'
import {
buildChatGPTRequestBody,
@@ -268,6 +269,80 @@ describe('ImageGen backend', () => {
expect(result.inputImageCount).toBe(0)
})
test('uses the configured image model for a model-supplied default placeholder', async () => {
outputDir = await mkdtemp(join(tmpdir(), 'imagegen-output-'))
let requestBody: Record<string, unknown> | undefined
const fetchImpl = async (_input: string | URL | Request, init?: RequestInit) => {
requestBody = JSON.parse(String(init?.body))
return Response.json({
data: [{ b64_json: PNG_BYTES.toString('base64') }],
})
}
const result = await generateImages({
prompt: 'A paper-cut fox poster',
count: 1,
model: 'default',
}, customConfig, { fetchImpl, outputDir })
expect(requestBody).toMatchObject({ model: 'relay-image-model' })
expect(result.model).toBe('relay-image-model')
})
test('uses the configured ChatGPT OAuth image model for a default placeholder', async () => {
outputDir = await mkdtemp(join(tmpdir(), 'imagegen-openai-oauth-'))
const tokenPath = join(outputDir, 'openai-oauth.json')
const previousTokenPath = process.env[OPENAI_CODEX_OAUTH_FILE_ENV_KEY]
await writeFile(tokenPath, JSON.stringify({
accessToken: 'test-openai-access-token',
refreshToken: 'test-openai-refresh-token',
expiresAt: 4_100_000_000_000,
accountId: 'test-openai-account',
}))
process.env[OPENAI_CODEX_OAUTH_FILE_ENV_KEY] = tokenPath
let requestBody: Record<string, any> | undefined
const fetchImpl = async (_input: string | URL | Request, init?: RequestInit) => {
requestBody = JSON.parse(String(init?.body))
const event = {
type: 'response.output_item.done',
item: {
type: 'image_generation_call',
result: PNG_BYTES.toString('base64'),
},
}
return new Response(`data: ${JSON.stringify(event)}\n\ndata: [DONE]\n\n`)
}
try {
const result = await generateImages({
prompt: 'A paper-cut fox poster',
count: 1,
model: 'default',
}, {
kind: 'openai_oauth',
providerId: 'openai-official',
model: 'gpt-image-2',
}, { fetchImpl, outputDir })
expect(requestBody?.tools?.[0]).toMatchObject({
type: 'image_generation',
model: 'gpt-image-2',
})
expect(result).toMatchObject({
providerId: 'openai-official',
providerKind: 'openai_oauth',
model: 'gpt-image-2',
})
} finally {
if (previousTokenPath === undefined) {
delete process.env[OPENAI_CODEX_OAUTH_FILE_ENV_KEY]
} else {
process.env[OPENAI_CODEX_OAUTH_FILE_ENV_KEY] = previousTokenPath
}
}
})
test('rejects edit paths outside the session upload and generated-image roots', async () => {
outputDir = await mkdtemp(join(tmpdir(), 'imagegen-edit-root-'))
const outsideDir = await mkdtemp(join(tmpdir(), 'imagegen-edit-outside-'))
+5 -1
View File
@@ -94,7 +94,11 @@ export async function generateImages(
options: GenerateOptions = {},
): Promise<ImageGenerationOutput> {
const startedAt = Date.now()
const model = input.model?.trim() || config.model
const requestedModel = input.model?.trim()
// Some tool-calling models use "default" to mean no model override.
const model = requestedModel && requestedModel.toLowerCase() !== 'default'
? requestedModel
: config.model
const fetchImpl = options.fetchImpl ?? fetch
const { signal, cleanup } = createCombinedAbortSignal(options.signal, {
timeoutMs: IMAGE_REQUEST_TIMEOUT_MS,