Merge pull request #3914 from code-yeongyu/fix/max-output-tokens
fix: prevent non-positive maxOutputTokens from leaking to SDK
This commit is contained in:
@@ -5,7 +5,7 @@ import { join } from "node:path"
|
|||||||
|
|
||||||
import { createChatParamsHandler, type ChatParamsOutput } from "./chat-params"
|
import { createChatParamsHandler, type ChatParamsOutput } from "./chat-params"
|
||||||
import * as dataPathModule from "../shared/data-path"
|
import * as dataPathModule from "../shared/data-path"
|
||||||
import { writeProviderModelsCache } from "../shared"
|
import * as sharedModule from "../shared"
|
||||||
import {
|
import {
|
||||||
clearSessionPromptParams,
|
clearSessionPromptParams,
|
||||||
getSessionPromptParams,
|
getSessionPromptParams,
|
||||||
@@ -21,13 +21,13 @@ describe("createChatParamsHandler", () => {
|
|||||||
getCacheDirSpy = spyOn(dataPathModule, "getOmoOpenCodeCacheDir").mockReturnValue(
|
getCacheDirSpy = spyOn(dataPathModule, "getOmoOpenCodeCacheDir").mockReturnValue(
|
||||||
join(tempCacheRoot, "oh-my-opencode"),
|
join(tempCacheRoot, "oh-my-opencode"),
|
||||||
)
|
)
|
||||||
writeProviderModelsCache({ connected: [], models: {} })
|
sharedModule.writeProviderModelsCache({ connected: [], models: {} })
|
||||||
})
|
})
|
||||||
|
|
||||||
afterEach(() => {
|
afterEach(() => {
|
||||||
clearSessionPromptParams("ses_chat_params")
|
clearSessionPromptParams("ses_chat_params")
|
||||||
clearSessionPromptParams("ses_chat_params_temperature")
|
clearSessionPromptParams("ses_chat_params_temperature")
|
||||||
writeProviderModelsCache({ connected: [], models: {} })
|
sharedModule.writeProviderModelsCache({ connected: [], models: {} })
|
||||||
getCacheDirSpy?.mockRestore()
|
getCacheDirSpy?.mockRestore()
|
||||||
if (tempCacheRoot) {
|
if (tempCacheRoot) {
|
||||||
rmSync(tempCacheRoot, { recursive: true, force: true })
|
rmSync(tempCacheRoot, { recursive: true, force: true })
|
||||||
@@ -101,7 +101,7 @@ describe("createChatParamsHandler", () => {
|
|||||||
|
|
||||||
test("applies stored prompt params for the session", async () => {
|
test("applies stored prompt params for the session", async () => {
|
||||||
//#given
|
//#given
|
||||||
writeProviderModelsCache({
|
sharedModule.writeProviderModelsCache({
|
||||||
connected: ["openai"],
|
connected: ["openai"],
|
||||||
models: {
|
models: {
|
||||||
openai: [
|
openai: [
|
||||||
@@ -253,4 +253,74 @@ describe("createChatParamsHandler", () => {
|
|||||||
options: {},
|
options: {},
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
|
test("falls back to default maxOutputTokens when stored and compatibility tokens are non-positive", async () => {
|
||||||
|
//#given
|
||||||
|
const logSpy = spyOn(sharedModule, "log").mockImplementation(() => undefined)
|
||||||
|
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: "custom-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(4096)
|
||||||
|
expect(logSpy).toHaveBeenCalledWith(
|
||||||
|
"[plugin] maxOutputTokens=0 is non-positive; using safe fallback 4096",
|
||||||
|
)
|
||||||
|
|
||||||
|
logSpy.mockRestore()
|
||||||
|
})
|
||||||
|
|
||||||
|
test("uses safe fallback instead of model max when stored maxOutputTokens is non-positive", async () => {
|
||||||
|
//#given
|
||||||
|
setSessionPromptParams("ses_chat_params", {
|
||||||
|
maxOutputTokens: -1,
|
||||||
|
})
|
||||||
|
|
||||||
|
const handler = createChatParamsHandler({
|
||||||
|
anthropicEffort: null,
|
||||||
|
})
|
||||||
|
|
||||||
|
const input = {
|
||||||
|
sessionID: "ses_chat_params",
|
||||||
|
agent: { name: "oracle" },
|
||||||
|
model: { providerID: "openai", modelID: "gpt-5.4" },
|
||||||
|
provider: { id: "openai" },
|
||||||
|
message: {},
|
||||||
|
}
|
||||||
|
|
||||||
|
const output: ChatParamsOutput = {
|
||||||
|
topP: 1,
|
||||||
|
topK: 1,
|
||||||
|
maxOutputTokens: -1,
|
||||||
|
options: {},
|
||||||
|
}
|
||||||
|
|
||||||
|
//#when
|
||||||
|
await handler(input, output)
|
||||||
|
|
||||||
|
//#then
|
||||||
|
expect(output.maxOutputTokens).toBe(4096)
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
import { getSessionPromptParams } from "../shared/session-prompt-params-state"
|
import { getSessionPromptParams } from "../shared/session-prompt-params-state"
|
||||||
import { getModelCapabilities, resolveCompatibleModelSettings } from "../shared"
|
import { getModelCapabilities, log, resolveCompatibleModelSettings } from "../shared"
|
||||||
|
|
||||||
|
const SAFE_MAX_OUTPUT_TOKENS_FALLBACK = 4096
|
||||||
|
|
||||||
export type ChatParamsInput = {
|
export type ChatParamsInput = {
|
||||||
sessionID: string
|
sessionID: string
|
||||||
@@ -96,7 +98,10 @@ export function createChatParamsHandler(args: {
|
|||||||
if (storedPromptParams.topP !== undefined) {
|
if (storedPromptParams.topP !== undefined) {
|
||||||
output.topP = storedPromptParams.topP
|
output.topP = storedPromptParams.topP
|
||||||
}
|
}
|
||||||
if (storedPromptParams.maxOutputTokens !== undefined) {
|
if (
|
||||||
|
typeof storedPromptParams.maxOutputTokens === "number" &&
|
||||||
|
storedPromptParams.maxOutputTokens > 0
|
||||||
|
) {
|
||||||
(output as Record<string, unknown>).maxOutputTokens = storedPromptParams.maxOutputTokens
|
(output as Record<string, unknown>).maxOutputTokens = storedPromptParams.maxOutputTokens
|
||||||
}
|
}
|
||||||
if (storedPromptParams.options) {
|
if (storedPromptParams.options) {
|
||||||
@@ -162,10 +167,18 @@ export function createChatParamsHandler(args: {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if ("maxTokens" in compatibility) {
|
if ("maxTokens" in compatibility) {
|
||||||
if (compatibility.maxTokens !== undefined) {
|
if (compatibility.maxTokens !== undefined && compatibility.maxTokens > 0) {
|
||||||
output.maxOutputTokens = compatibility.maxTokens
|
output.maxOutputTokens = compatibility.maxTokens
|
||||||
} else {
|
} else {
|
||||||
delete output.maxOutputTokens
|
const originalMaxOutputTokens = typeof output.maxOutputTokens === "number"
|
||||||
|
? output.maxOutputTokens
|
||||||
|
: compatibility.maxTokens
|
||||||
|
output.maxOutputTokens = SAFE_MAX_OUTPUT_TOKENS_FALLBACK
|
||||||
|
if (typeof originalMaxOutputTokens === "number" && originalMaxOutputTokens <= 0) {
|
||||||
|
log(
|
||||||
|
`[plugin] maxOutputTokens=${originalMaxOutputTokens} is non-positive; using safe fallback ${SAFE_MAX_OUTPUT_TOKENS_FALLBACK}`,
|
||||||
|
)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -615,6 +615,18 @@ describe("resolveCompatibleModelSettings", () => {
|
|||||||
expect(result.changes).toEqual([])
|
expect(result.changes).toEqual([])
|
||||||
})
|
})
|
||||||
|
|
||||||
|
test("#given desired.maxTokens is 0 #then maxTokens is dropped", () => {
|
||||||
|
const result = resolveCompatibleModelSettings({
|
||||||
|
providerID: "openai",
|
||||||
|
modelID: "gpt-5.4",
|
||||||
|
desired: { maxTokens: 0 },
|
||||||
|
capabilities: { maxOutputTokens: 128_000 },
|
||||||
|
})
|
||||||
|
|
||||||
|
expect(result.maxTokens).toBeUndefined()
|
||||||
|
expect(result.changes).toEqual([])
|
||||||
|
})
|
||||||
|
|
||||||
// Passthrough: undefined desired values produce no changes
|
// Passthrough: undefined desired values produce no changes
|
||||||
test("no-op when desired settings are empty", () => {
|
test("no-op when desired settings are empty", () => {
|
||||||
const result = resolveCompatibleModelSettings({
|
const result = resolveCompatibleModelSettings({
|
||||||
|
|||||||
@@ -175,6 +175,10 @@ export function resolveCompatibleModelSettings(
|
|||||||
}
|
}
|
||||||
|
|
||||||
let maxTokens = input.desired.maxTokens
|
let maxTokens = input.desired.maxTokens
|
||||||
|
if (maxTokens !== undefined && maxTokens <= 0) {
|
||||||
|
maxTokens = undefined
|
||||||
|
}
|
||||||
|
|
||||||
if (
|
if (
|
||||||
maxTokens !== undefined &&
|
maxTokens !== undefined &&
|
||||||
input.capabilities?.maxOutputTokens !== undefined &&
|
input.capabilities?.maxOutputTokens !== undefined &&
|
||||||
|
|||||||
Reference in New Issue
Block a user