diff --git a/src/server/__tests__/conversations.test.ts b/src/server/__tests__/conversations.test.ts index b15ef1ce..86bc8f96 100644 --- a/src/server/__tests__/conversations.test.ts +++ b/src/server/__tests__/conversations.test.ts @@ -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 { 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 + 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((resolve) => { + markValidationStarted = resolve + }) + let releaseValidation!: () => void + const validationGate = new Promise>>((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((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((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({ diff --git a/src/server/ws/handler.ts b/src/server/ws/handler.ts index 0649660a..7737ad0f 100644 --- a/src/server/ws/handler.ts +++ b/src/server/ws/handler.ts @@ -1426,8 +1426,8 @@ async function handleSetRuntimeConfig( message: Extract ) { 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( diff --git a/src/skills/bundled/imagegen.test.ts b/src/skills/bundled/imagegen.test.ts index bce9605d..57b7daab 100644 --- a/src/skills/bundled/imagegen.test.ts +++ b/src/skills/bundled/imagegen.test.ts @@ -59,6 +59,8 @@ describe('bundled imagegen skill', () => { expect(text).toContain('one call per image') expect(text).toContain('omit input_images entirely') expect(text).toContain('Never pass /dev/null') + expect(text).toContain('otherwise omit the field') + expect(text).toContain('Never pass default') expect(text).not.toContain('CC_HAHA_IMAGE_API_KEY') }) }) diff --git a/src/skills/bundled/imagegen/SKILL.md b/src/skills/bundled/imagegen/SKILL.md index 005cda08..e71fd4cb 100644 --- a/src/skills/bundled/imagegen/SKILL.md +++ b/src/skills/bundled/imagegen/SKILL.md @@ -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 diff --git a/src/tools/ImageGenTool/ImageGenTool.test.ts b/src/tools/ImageGenTool/ImageGenTool.test.ts index c170fff2..40d8ab36 100644 --- a/src/tools/ImageGenTool/ImageGenTool.test.ts +++ b/src/tools/ImageGenTool/ImageGenTool.test.ts @@ -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"') }) }) diff --git a/src/tools/ImageGenTool/ImageGenTool.ts b/src/tools/ImageGenTool/ImageGenTool.ts index 7d3fd35e..ab17403b 100644 --- a/src/tools/ImageGenTool/ImageGenTool.ts +++ b/src/tools/ImageGenTool/ImageGenTool.ts @@ -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() diff --git a/src/tools/ImageGenTool/backend.test.ts b/src/tools/ImageGenTool/backend.test.ts index cd3ea81c..d6037a23 100644 --- a/src/tools/ImageGenTool/backend.test.ts +++ b/src/tools/ImageGenTool/backend.test.ts @@ -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 | 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 | 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-')) diff --git a/src/tools/ImageGenTool/backend.ts b/src/tools/ImageGenTool/backend.ts index 88bff1d6..df3a230d 100644 --- a/src/tools/ImageGenTool/backend.ts +++ b/src/tools/ImageGenTool/backend.ts @@ -94,7 +94,11 @@ export async function generateImages( options: GenerateOptions = {}, ): Promise { 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,