diff --git a/src/plugin/chat-params.test.ts b/src/plugin/chat-params.test.ts index 736e36d21..20bb9ddfb 100644 --- a/src/plugin/chat-params.test.ts +++ b/src/plugin/chat-params.test.ts @@ -323,4 +323,48 @@ describe("createChatParamsHandler", () => { //#then expect(output.maxOutputTokens).toBe(4096) }) + + test("uses model max when non-positive stored maxOutputTokens fallback exceeds model limit", async () => { + //#given + sharedModule.writeProviderModelsCache({ + connected: ["custom-provider"], + models: { + "custom-provider": [ + { + id: "tiny-model", + name: "tiny-model", + limit: { output: 512 }, + }, + ], + }, + }) + setSessionPromptParams("ses_chat_params", { + maxOutputTokens: 0, + }) + + const handler = createChatParamsHandler({ + anthropicEffort: null, + }) + + const input = { + sessionID: "ses_chat_params", + agent: { name: "oracle" }, + model: { providerID: "custom-provider", modelID: "tiny-model" }, + provider: { id: "custom-provider" }, + message: {}, + } + + const output: ChatParamsOutput = { + topP: 1, + topK: 1, + maxOutputTokens: 0, + options: {}, + } + + //#when + await handler(input, output) + + //#then + expect(output.maxOutputTokens).toBe(512) + }) }) diff --git a/src/plugin/chat-params.ts b/src/plugin/chat-params.ts index 26f35d03d..16a02234f 100644 --- a/src/plugin/chat-params.ts +++ b/src/plugin/chat-params.ts @@ -3,6 +3,14 @@ import { getModelCapabilities, log, resolveCompatibleModelSettings } from "../sh const SAFE_MAX_OUTPUT_TOKENS_FALLBACK = 4096 +function resolveSafeMaxOutputTokensFallback(capabilitiesMaxOutputTokens: number | undefined): number { + if (typeof capabilitiesMaxOutputTokens !== "number" || capabilitiesMaxOutputTokens <= 0) { + return SAFE_MAX_OUTPUT_TOKENS_FALLBACK + } + + return Math.min(SAFE_MAX_OUTPUT_TOKENS_FALLBACK, capabilitiesMaxOutputTokens) +} + export type ChatParamsInput = { sessionID: string agent: { name?: string } @@ -173,10 +181,10 @@ export function createChatParamsHandler(args: { const originalMaxOutputTokens = typeof output.maxOutputTokens === "number" ? output.maxOutputTokens : compatibility.maxTokens - output.maxOutputTokens = SAFE_MAX_OUTPUT_TOKENS_FALLBACK + output.maxOutputTokens = resolveSafeMaxOutputTokensFallback(capabilities?.maxOutputTokens) if (typeof originalMaxOutputTokens === "number" && originalMaxOutputTokens <= 0) { log( - `[plugin] maxOutputTokens=${originalMaxOutputTokens} is non-positive; using safe fallback ${SAFE_MAX_OUTPUT_TOKENS_FALLBACK}`, + `[plugin] maxOutputTokens=${originalMaxOutputTokens} is non-positive; using safe fallback ${output.maxOutputTokens}`, ) } }