diff --git a/packages/opencode/src/provider/provider.ts b/packages/opencode/src/provider/provider.ts index 6f3cfccc51..42829eb14a 100644 --- a/packages/opencode/src/provider/provider.ts +++ b/packages/opencode/src/provider/provider.ts @@ -2,7 +2,7 @@ import os from "os" import fuzzysort from "fuzzysort" import { Config } from "@/config/config" import { mapValues, mergeDeep, omit, pickBy, sortBy } from "remeda" -import { NoSuchModelError, type Provider as SDK } from "ai" +import { NoSuchModelError, type ModelMessage, type Provider as SDK } from "ai" import * as Log from "@opencode-ai/core/util/log" import { Npm } from "@opencode-ai/core/npm" import { Hash } from "@opencode-ai/core/util/hash" @@ -28,6 +28,8 @@ import * as ProviderTransform from "./transform" import { ModelID, ProviderID } from "./schema" import { ModelStatus } from "./model-status" import { RuntimeFlags } from "@/effect/runtime-flags" +import type { Agent } from "@/agent/agent" +import type { MessageV2 } from "@/session/message-v2" import { ProviderError } from "./error" const log = Log.create({ service: "provider" }) @@ -955,6 +957,18 @@ export const Info = Schema.Struct({ }).annotate({ identifier: "Provider" }) export type Info = Types.DeepMutable> +export type LanguageModelRequest = { + sessionID: string + parentSessionID?: string + agent: Agent.Info + message: MessageV2.User + messages: ModelMessage[] + system: string[] + headers: Record + tools: string[] + small?: boolean +} + const DefaultModelIDs = Schema.Record(Schema.String, Schema.String) export const ListResult = Schema.Struct({ @@ -1025,7 +1039,7 @@ export interface Interface { readonly list: () => Effect.Effect> readonly getProvider: (providerID: ProviderID) => Effect.Effect readonly getModel: (providerID: ProviderID, modelID: ModelID) => Effect.Effect - readonly getLanguage: (model: Model) => Effect.Effect + readonly getLanguage: (model: Model, request?: LanguageModelRequest) => Effect.Effect readonly closest: ( providerID: ProviderID, query: string[], @@ -1542,7 +1556,12 @@ export const layer = Layer.effect( const list = Effect.fn("Provider.list")(() => InstanceState.use(state, (s) => s.providers)) - async function resolveSDK(model: Model, s: State, envs: Record) { + async function resolveSDK( + model: Model, + s: State, + envs: Record, + request?: LanguageModelRequest, + ) { try { using _ = log.time("getSDK", { providerID: model.providerID, @@ -1607,7 +1626,7 @@ export const layer = Layer.effect( }), ) const existing = s.sdk.get(key) - if (existing) return existing + if (existing && !request) return existing const customFetch = options["fetch"] const chunkTimeout = options["chunkTimeout"] @@ -1650,11 +1669,21 @@ export const layer = Layer.effect( } } - const res = await fetchFn(input, { - ...opts, - // @ts-ignore see here: https://github.com/oven-sh/bun/issues/16682 - timeout: false, - }).finally(() => headerTimeoutCtl?.clear()) + const res = await fetchFn( + input, + { + ...opts, + // @ts-ignore see here: https://github.com/oven-sh/bun/issues/16682 + timeout: false, + }, + request + ? { + ...request, + model, + provider, + } + : undefined, + ).finally(() => headerTimeoutCtl?.clear()) if (!chunkAbortCtl) return res return wrapSSE(res, chunkTimeout, chunkAbortCtl) @@ -1671,7 +1700,7 @@ export const layer = Layer.effect( name: model.providerID, ...options, }) - s.sdk.set(key, loaded) + if (!request) s.sdk.set(key, loaded) return loaded as SDK } @@ -1695,7 +1724,7 @@ export const layer = Layer.effect( name: model.providerID, ...options, }) - s.sdk.set(key, loaded) + if (!request) s.sdk.set(key, loaded) return loaded as SDK } catch (e) { throw new InitError({ providerID: model.providerID, cause: e }) @@ -1730,23 +1759,29 @@ export const layer = Layer.effect( return info }) - const getLanguage = Effect.fn("Provider.getLanguage")(function* (model: Model) { + const getLanguage = Effect.fn("Provider.getLanguage")(function* (model: Model, request?: LanguageModelRequest) { const s = yield* InstanceState.get(state) const envs = yield* env.all() const key = `${model.providerID}/${model.id}` - if (s.models.has(key)) return s.models.get(key)! - const provider = s.providers[model.providerID] + const requestFetch = + request && + typeof provider.options.fetch === "function" && + !(model.providerID === "google-vertex" && !model.api.npm.includes("@ai-sdk/openai-compatible")) + ? request + : undefined + if (!requestFetch && s.models.has(key)) return s.models.get(key)! + return yield* EffectPromise.refineRejection( async () => { - const sdk = await resolveSDK(model, s, envs) + const sdk = await resolveSDK(model, s, envs, requestFetch) const language = s.modelLoaders[model.providerID] ? await s.modelLoaders[model.providerID](sdk, model.api.id, { ...provider.options, ...model.options, }) : sdk.languageModel(model.api.id) - s.models.set(key, language) + if (!requestFetch) s.models.set(key, language) return language }, (cause) => @@ -1860,7 +1895,15 @@ export const layer = Layer.effect( } }) - return Service.of({ list, getProvider, getModel, getLanguage, closest, getSmallModel, defaultModel }) + return Service.of({ + list, + getProvider, + getModel, + getLanguage, + closest, + getSmallModel, + defaultModel, + }) }), ) diff --git a/packages/opencode/src/session/llm.ts b/packages/opencode/src/session/llm.ts index ea2efc99d0..f89ea610f1 100644 --- a/packages/opencode/src/session/llm.ts +++ b/packages/opencode/src/session/llm.ts @@ -111,12 +111,23 @@ const live: Layer.Layer< flags, isWorkflow, }) + const requestLanguage = yield* provider.getLanguage(input.model, { + sessionID: input.sessionID, + parentSessionID: input.parentSessionID, + agent: input.agent, + message: input.user, + messages: prepared.messages, + system: prepared.system, + headers: prepared.headers, + tools: Object.keys(prepared.tools), + small: input.small, + }) // Wire up toolExecutor for DWS workflow models so that tool calls // from the workflow service are executed via opencode's tool system // and results sent back over the WebSocket. - if (language instanceof GitLabWorkflowLanguageModel) { - const workflowModel = language as GitLabWorkflowLanguageModel & { + if (requestLanguage instanceof GitLabWorkflowLanguageModel) { + const workflowModel = requestLanguage as GitLabWorkflowLanguageModel & { sessionID?: string sessionPreapprovedTools?: string[] approvalHandler?: (approvalTools: { name: string; args: string }[]) => Promise<{ approved: boolean }> @@ -309,7 +320,7 @@ const live: Layer.Layer< maxRetries: input.retries ?? 0, messages: prepared.messages, model: wrapLanguageModel({ - model: language, + model: requestLanguage, middleware: [ { specificationVersion: "v3" as const, diff --git a/packages/plugin/src/index.ts b/packages/plugin/src/index.ts index 3c710d076a..fc6b5640a6 100644 --- a/packages/plugin/src/index.ts +++ b/packages/plugin/src/index.ts @@ -79,6 +79,30 @@ export type PluginModule = { tui?: never } +export type ExperimentalFetchContext = { + sessionID: string + parentSessionID?: string + agent: { + name: string + mode: string + [key: string]: unknown + } + model: ModelV2 + provider: ProviderV2 + message: UserMessage + messages: unknown[] + system: string[] + headers: Record + tools: string[] + small?: boolean +} + +export type ExperimentalFetch = ( + input: RequestInfo | URL, + init?: RequestInit, + context?: ExperimentalFetchContext, +) => Promise + type Rule = { key: string op: "eq" | "neq"