Skip to content
Closed
Show file tree
Hide file tree
Changes from 1 commit
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
7 changes: 7 additions & 0 deletions packages/core/src/v1/config/provider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,13 @@ export const Model = Schema.Struct({
Schema.StructWithRest(
Schema.Struct({
disabled: Schema.optional(Schema.Boolean).annotate({ description: "Disable this variant for the model" }),
limit: Schema.optional(
Schema.Struct({
context: Schema.optional(Schema.Finite),
input: Schema.optional(Schema.Finite),
output: Schema.optional(Schema.Finite),
}),
).annotate({ description: "Override the model's token limits when this variant is active" }),
}),
[Schema.Record(Schema.String, Schema.Any)],
),
Expand Down
6 changes: 6 additions & 0 deletions packages/opencode/src/provider/transform.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1326,6 +1326,12 @@ export function maxOutputTokens(model: Provider.Model, outputTokenMax = OUTPUT_T
return Math.min(model.limit.output, outputTokenMax) || outputTokenMax
}

export function withVariantLimit(model: Provider.Model, variantID?: string): Provider.Model {
const limit = variantID === undefined ? undefined : model.variants?.[variantID]?.["limit"]
if (!isPlainObject(limit)) return model
return { ...model, limit: { ...model.limit, ...(limit as Partial<Provider.Model["limit"]>) } }
}

type JsonRecord = Record<string, unknown>

function isPlainObject(value: unknown): value is JsonRecord {
Expand Down
14 changes: 10 additions & 4 deletions packages/opencode/src/session/llm/request.ts
Original file line number Diff line number Diff line change
Expand Up @@ -77,18 +77,24 @@ export const prepare = Effect.fn("LLMRequestPrep.prepare")(function* (input: Pre
system.push(header, rest.join("\n"))
}

const variant =
const variant: Record<string, any> =
!input.small && input.model.variants && input.user.model.variant
? input.model.variants[input.user.model.variant]
? (input.model.variants[input.user.model.variant] ?? {})
: {}
const variantOptions = { ...variant }
delete variantOptions.limit
const model = ProviderTransform.withVariantLimit(input.model, input.small ? undefined : input.user.model.variant)
const base = input.small
? ProviderTransform.smallOptions(input.model)
: ProviderTransform.options({
model: input.model,
sessionID: input.sessionID,
providerOptions: input.provider.options,
})
const options = mergeOptions(mergeOptions(mergeOptions(base, input.model.options), input.agent.options), variant)
const options = mergeOptions(
mergeOptions(mergeOptions(base, input.model.options), input.agent.options),
variantOptions,
)
if (
input.model.api.npm === "@ai-sdk/azure" &&
(input.provider.options.useCompletionUrls || input.model.options.useCompletionUrls || options.useCompletionUrls)
Expand Down Expand Up @@ -126,7 +132,7 @@ export const prepare = Effect.fn("LLMRequestPrep.prepare")(function* (input: Pre
: undefined,
topP: input.agent.topP ?? ProviderTransform.topP(input.model),
topK: ProviderTransform.topK(input.model),
maxOutputTokens: ProviderTransform.maxOutputTokens(input.model, input.flags.outputTokenMax),
maxOutputTokens: ProviderTransform.maxOutputTokens(model, input.flags.outputTokenMax),
options,
},
)
Expand Down
6 changes: 5 additions & 1 deletion packages/opencode/src/session/prompt.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import { SessionRevert } from "./revert"
import { Session } from "./session"
import { Agent } from "../agent/agent"
import { Provider } from "@/provider/provider"
import { ProviderTransform } from "@/provider/transform"

import { type Tool as AITool, tool, jsonSchema } from "ai"
import type { JSONSchema7 } from "@ai-sdk/provider"
Expand Down Expand Up @@ -1138,7 +1139,10 @@ const layer = Layer.effect(
history: msgs,
}).pipe(Effect.ignore, Effect.forkIn(scope))

const model = yield* getModel(lastUser.model.providerID, lastUser.model.modelID, sessionID)
const model = ProviderTransform.withVariantLimit(
yield* getModel(lastUser.model.providerID, lastUser.model.modelID, sessionID),
lastUser.model.variant,
)
const task = tasks.pop()

if (task?.type === "subtask") {
Expand Down
74 changes: 74 additions & 0 deletions packages/opencode/test/session/llm.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -832,6 +832,80 @@ describe("session.llm.stream", () => {
},
)

it.instance(
"applies variant limit overrides to max output tokens",
() =>
Effect.gen(function* () {
const fixture = loadFixture(vivgridFixture.providerID, vivgridFixture.modelID)
const request = waitRequest(
"/chat/completions",
new Response(createChatStream("Hello"), {
status: 200,
headers: { "Content-Type": "text/event-stream" },
}),
)

const resolved = yield* Provider.use.getModel(
ProviderV2.ID.make(vivgridFixture.providerID),
ModelV2.ID.make(fixture.model.id),
)
const sessionID = SessionID.make("session-test-variant-limit")
const agent = {
name: "test",
mode: "primary",
options: {},
permission: [{ permission: "*", pattern: "*", action: "allow" }],
} satisfies Agent.Info

const user = {
id: MessageID.make("msg_user-variant-limit"),
sessionID,
role: "user",
time: { created: Date.now() },
agent: agent.name,
model: {
providerID: ProviderV2.ID.make(vivgridFixture.providerID),
modelID: resolved.id,
variant: "throttled",
},
} satisfies SessionV1.User

yield* drain({
user,
sessionID,
model: resolved,
agent,
system: ["You are a helpful assistant."],
messages: [{ role: "user", content: "Hello" }],
tools: {},
})

const capture = yield* Effect.promise(() => request)
const body = capture.body

const maxTokens = (body.max_tokens as number | undefined) ?? (body.max_output_tokens as number | undefined)
expect(maxTokens).toBe(4096)
expect(body.limit).toBeUndefined()
}),
{
config: () => ({
enabled_providers: [vivgridFixture.providerID],
provider: {
[vivgridFixture.providerID]: {
options: { apiKey: "test-key", baseURL: `${state.server!.url.origin}/v1` },
models: {
[vivgridFixture.modelID]: {
variants: {
throttled: { limit: { output: 4096 } },
},
},
},
},
},
}),
},
)

const alibabaQwenFixture = { providerID: "alibaba", modelID: "qwen-plus" }
it.instance(
"service stream cancellation cancels provider response body promptly",
Expand Down
15 changes: 14 additions & 1 deletion packages/sdk/js/src/v2/gen/types.gen.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1811,7 +1811,20 @@ export type ProviderConfig = {
variants?: {
[key: string]: {
disabled?: boolean
[key: string]: unknown | boolean | undefined
limit?: {
context?: number
input?: number
output?: number
}
[key: string]:
| unknown
| boolean
| {
context?: number
input?: number
output?: number
}
| undefined
}
}
}
Expand Down
15 changes: 15 additions & 0 deletions packages/sdk/openapi.json
Original file line number Diff line number Diff line change
Expand Up @@ -20944,6 +20944,21 @@
"properties": {
"disabled": {
"type": "boolean"
},
"limit": {
"type": "object",
"properties": {
"context": {
"type": "number"
},
"input": {
"type": "number"
},
"output": {
"type": "number"
}
},
"additionalProperties": false
}
},
"additionalProperties": {}
Expand Down
Loading