diff --git a/packages/opencode/src/kilo-sessions/kilo-sessions.ts b/packages/opencode/src/kilo-sessions/kilo-sessions.ts index a6dbe1e929a..a4f5b80155e 100644 --- a/packages/opencode/src/kilo-sessions/kilo-sessions.ts +++ b/packages/opencode/src/kilo-sessions/kilo-sessions.ts @@ -111,6 +111,20 @@ export namespace KiloSessions { }) } + async function model(providerID: ProviderID, modelID: ModelID) { + const { AppRuntime } = await import("@/effect/app-runtime") + return AppRuntime.runPromise(Provider.Service.use((svc) => svc.getModel(providerID, modelID))) + } + + async function models(refs: Array<{ providerID: string; modelID: string }>) { + const { AppRuntime } = await import("@/effect/app-runtime") + return AppRuntime.runPromise( + Provider.Service.use((svc) => + Effect.all(refs.map((ref) => svc.getModel(ProviderID.make(ref.providerID), ModelID.make(ref.modelID)))), + ), + ) + } + type Client = { url: string fetch: (input: RequestInfo | URL, init?: RequestInit) => Promise @@ -233,11 +247,8 @@ export namespace KiloSessions { yield* watch(MessageV2.Event.Updated, async (evt) => { await ingest.sync(evt.properties.info.sessionID, [{ type: "message", data: evt.properties.info }]) if (evt.properties.info.role !== "user") return - const model = await Provider.getModel( - evt.properties.info.model.providerID, - evt.properties.info.model.modelID, - ) - await ingest.sync(evt.properties.info.sessionID, [{ type: "model", data: [model] }]) + const mdl = await model(evt.properties.info.model.providerID, evt.properties.info.model.modelID) + await ingest.sync(evt.properties.info.sessionID, [{ type: "model", data: [mdl] }]) }) yield* watch(MessageV2.Event.PartUpdated, (evt) => ingest.sync(evt.properties.part.sessionID, [{ type: "part", data: evt.properties.part }]), @@ -640,11 +651,8 @@ export namespace KiloSessions { ) const messages = await Array.fromAsync(MessageV2.stream(SessionID.make(sessionId))) messages.reverse() - const models = await Promise.all( - messages - .filter((m) => m.info.role === "user") - .map((m) => (m.info as SDK.UserMessage).model) - .map((m) => Provider.getModel(ProviderID.make(m.providerID), ModelID.make(m.modelID)).then((m) => m)), + const mdls = await models( + messages.filter((m) => m.info.role === "user").map((m) => (m.info as SDK.UserMessage).model), ) await ingest.sync(sessionId, [ @@ -667,7 +675,7 @@ export namespace KiloSessions { }, { type: "model", - data: models, + data: mdls, }, { type: "session_status", diff --git a/packages/opencode/src/kilocode/cli/cmd/roll-call.ts b/packages/opencode/src/kilocode/cli/cmd/roll-call.ts index 03ffc225027..cf025fdef6b 100644 --- a/packages/opencode/src/kilocode/cli/cmd/roll-call.ts +++ b/packages/opencode/src/kilocode/cli/cmd/roll-call.ts @@ -4,6 +4,7 @@ import { Provider } from "../../../provider/provider" import { ProviderTransform } from "../../../provider/transform" import { cmd } from "../../../cli/cmd/cmd" import { UI } from "../../../cli/ui" +import { AppRuntime } from "../../../effect/app-runtime" import { generateText } from "ai" import { randomUUID } from "crypto" @@ -125,8 +126,16 @@ interface Result { errorMessage: string | null } +function list() { + return AppRuntime.runPromise(Provider.Service.use((svc) => svc.list())) +} + +function lang(model: Provider.Model) { + return AppRuntime.runPromise(Provider.Service.use((svc) => svc.getLanguage(model))) +} + export async function handle(args: ArgumentsCamelCase) { - const list = args.list ?? Provider.list + const load = args.list ?? list if (args.parallel < 1) { UI.error("--parallel must be at least 1") @@ -161,7 +170,7 @@ export async function handle(args: ArgumentsCamelCase) { await WithInstance.provide({ directory: process.cwd(), async fn() { - const providers = await list() + const providers = await load() const regex = (() => { try { return new RegExp(args.filter, "i") @@ -272,7 +281,7 @@ async function call( start: number, ): Promise> { try { - const language = await Provider.getLanguage(model) + const language = await lang(model) const sessionID = randomUUID() const options = ProviderTransform.options({ model, sessionID }) const providerOptions = ProviderTransform.providerOptions(model, options) @@ -342,5 +351,5 @@ type ArgumentsCamelCase = { output: "table" | "json" | "md" verbose: boolean quiet: boolean - list?: typeof Provider.list + list?: typeof list } diff --git a/packages/opencode/src/kilocode/commit-message/generate.ts b/packages/opencode/src/kilocode/commit-message/generate.ts index 6f8ab2b71e3..79380d40cd2 100644 --- a/packages/opencode/src/kilocode/commit-message/generate.ts +++ b/packages/opencode/src/kilocode/commit-message/generate.ts @@ -11,6 +11,16 @@ import { getGitContext } from "./git-context" const log = Log.create({ service: "commit-message" }) export const CommitMessageRuntime = { + model() { + return AppRuntime.runPromise( + Provider.Service.use((svc) => + Effect.gen(function* () { + const ref = yield* svc.defaultModel() + return (yield* svc.getSmallModel(ref.providerID)) ?? (yield* svc.getModel(ref.providerID, ref.modelID)) + }), + ), + ) + }, generate(input: LLM.StreamInput, signal: AbortSignal) { // runPromise is needed until generateCommitMessage() uses Effect return AppRuntime.runPromise( @@ -141,10 +151,7 @@ export async function generateCommitMessage(request: CommitMessageRequest): Prom files: ctx.files.length, }) - const defaultModel = await Provider.defaultModel() - const model = - (await Provider.getSmallModel(defaultModel.providerID)) ?? - (await Provider.getModel(defaultModel.providerID, defaultModel.modelID)) + const model = await CommitMessageRuntime.model() const agent: Agent.Info = { name: "commit-message", diff --git a/packages/opencode/src/kilocode/enhance-prompt.ts b/packages/opencode/src/kilocode/enhance-prompt.ts index 252ef87ea4d..3924e073398 100644 --- a/packages/opencode/src/kilocode/enhance-prompt.ts +++ b/packages/opencode/src/kilocode/enhance-prompt.ts @@ -2,6 +2,8 @@ import { generateText } from "ai" import { mergeDeep } from "remeda" import { Provider } from "@/provider/provider" import { ProviderTransform } from "@/provider/transform" +import { AppRuntime } from "@/effect/app-runtime" +import { Effect } from "effect" import * as Log from "@opencode-ai/core/util/log" const log = Log.create({ service: "enhance-prompt" }) @@ -28,19 +30,23 @@ export function clean(text: string) { export async function enhancePrompt(text: string): Promise { log.info("enhancing", { length: text.length }) - const defaultModel = await Provider.defaultModel() - const model = - (await Provider.getSmallModel(defaultModel.providerID)) ?? - (await Provider.getModel(defaultModel.providerID, defaultModel.modelID)) - - const language = await Provider.getLanguage(model) + const resolved = await AppRuntime.runPromise( + Provider.Service.use((svc) => + Effect.gen(function* () { + const ref = yield* svc.defaultModel() + const model = (yield* svc.getSmallModel(ref.providerID)) ?? (yield* svc.getModel(ref.providerID, ref.modelID)) + const language = yield* svc.getLanguage(model) + return { model, language } + }), + ), + ) const result = await generateText({ - model: language, - temperature: model.capabilities.temperature ? 0.7 : undefined, + model: resolved.language, + temperature: resolved.model.capabilities.temperature ? 0.7 : undefined, providerOptions: ProviderTransform.providerOptions( - model, - mergeDeep(ProviderTransform.smallOptions(model), model.options), + resolved.model, + mergeDeep(ProviderTransform.smallOptions(resolved.model), resolved.model.options), ), maxRetries: 3, system: INSTRUCTION, diff --git a/packages/opencode/src/kilocode/tool/task.ts b/packages/opencode/src/kilocode/tool/task.ts index 8397735c620..90514a80e05 100644 --- a/packages/opencode/src/kilocode/tool/task.ts +++ b/packages/opencode/src/kilocode/tool/task.ts @@ -98,6 +98,7 @@ export namespace KiloTask { agent: Pick config: Pick parent: Model + provider: Provider.Interface }) { const state = yield* saved(input.name) const cfg = parse(input.config.subagent_model) @@ -116,10 +117,8 @@ export namespace KiloTask { for (const choice of choices) { if (!choice) continue if (choice.direct) return { model: choice.model, variant: choice.variant } - const full = yield* Effect.tryPromise(() => - Provider.getModel(choice.model.providerID, choice.model.modelID), - ).pipe( - Effect.catch((err) => + const full = yield* input.provider.getModel(choice.model.providerID, choice.model.modelID).pipe( + Effect.catchDefect((err) => Effect.sync(() => { log.debug("skipping unavailable task subagent model", { providerID: choice.model.providerID, diff --git a/packages/opencode/src/provider/provider.ts b/packages/opencode/src/provider/provider.ts index 33ef3b785d6..93a513f5e71 100644 --- a/packages/opencode/src/provider/provider.ts +++ b/packages/opencode/src/provider/provider.ts @@ -7,7 +7,6 @@ import * as Log from "@opencode-ai/core/util/log" import { Npm } from "@opencode-ai/core/npm" import { Hash } from "@opencode-ai/core/util/hash" import { Plugin } from "../plugin" -import { makeRuntime } from "@/effect/run-service" // kilocode_change import { type LanguageModelV3 } from "@ai-sdk/provider" import * as ModelsDev from "./models" import { Auth } from "../auth" @@ -1809,17 +1808,6 @@ export function sort(models: T[]) { ) } -// kilocode_change start - legacy promise helpers for Kilo callsites -const { runPromise: runProviderPromise } = makeRuntime(Service, defaultLayer) -export const list = () => runProviderPromise((svc) => svc.list()) -export const getModel = (providerID: ProviderID, modelID: ModelID) => - runProviderPromise((svc) => svc.getModel(providerID, modelID)) -export const getProvider = (providerID: ProviderID) => runProviderPromise((svc) => svc.getProvider(providerID)) -export const getLanguage = (model: Model) => runProviderPromise((svc) => svc.getLanguage(model)) -export const getSmallModel = (providerID: ProviderID) => runProviderPromise((svc) => svc.getSmallModel(providerID)) -export const defaultModel = () => runProviderPromise((svc) => svc.defaultModel()) -// kilocode_change end - export function parseModel(model: string) { const [providerID, ...rest] = model.split("/") return { diff --git a/packages/opencode/src/tool/task.ts b/packages/opencode/src/tool/task.ts index ad709f1d614..c198c8ea499 100644 --- a/packages/opencode/src/tool/task.ts +++ b/packages/opencode/src/tool/task.ts @@ -6,6 +6,7 @@ import { MessageV2 } from "../session/message-v2" import { Agent } from "../agent/agent" import type { SessionPrompt } from "../session/prompt" import { Config } from "@/config/config" +import { Provider } from "@/provider/provider" // kilocode_change import { KiloTask } from "../kilocode/tool/task" // kilocode_change import { KiloCostPropagation } from "../kilocode/session/cost-propagation" // kilocode_change import { KiloSessionProcessor } from "../kilocode/session/processor" // kilocode_change @@ -38,6 +39,7 @@ export const TaskTool = Tool.define( const agent = yield* Agent.Service const config = yield* Config.Service const sessions = yield* Session.Service + const provider = yield* Provider.Service // kilocode_change const run = Effect.fn("TaskTool.execute")(function* ( params: Schema.Schema.Type, @@ -128,6 +130,7 @@ export const TaskTool = Tool.define( modelID: msg.info.modelID, providerID: msg.info.providerID, }, + provider, }) const model = selected.model const variant = selected.variant diff --git a/packages/opencode/test/kilocode/commit-message/generate.test.ts b/packages/opencode/test/kilocode/commit-message/generate.test.ts index 8b030061b05..de9a1701cd3 100644 --- a/packages/opencode/test/kilocode/commit-message/generate.test.ts +++ b/packages/opencode/test/kilocode/commit-message/generate.test.ts @@ -1,5 +1,6 @@ import { describe, expect, test, mock, beforeEach, spyOn } from "bun:test" import type { GitContext } from "@/kilocode/commit-message/types" +import type { Provider } from "@/provider/provider" // Mock dependencies before importing the module under test. // IMPORTANT: Bun's mock.module() is process-wide and permanent. To avoid @@ -7,7 +8,6 @@ import type { GitContext } from "@/kilocode/commit-message/types" // this test needs. const realLog = await import("@opencode-ai/core/util/log") -const realProvider = await import("@/provider/provider") const realAgent = await import("@/agent/agent") const realGitContext = await import("@/kilocode/commit-message/git-context") @@ -36,19 +36,6 @@ mock.module("@/kilocode/commit-message/git-context", () => ({ }, })) -mock.module("@/provider/provider", () => ({ - ...realProvider, - Provider: { - ...realProvider.Provider, - defaultModel: async () => ({ providerID: "test", modelID: "test-model" }), - getSmallModel: async () => ({ - providerID: "test", - id: "test-small-model", - }), - getModel: async () => ({ providerID: "test", id: "test-model" }), - }, -})) - mock.module("@/agent/agent", () => ({ ...realAgent, Agent: {}, @@ -67,10 +54,14 @@ mock.module("@opencode-ai/core/util/log", () => ({ import { CommitMessageRuntime, generateCommitMessage } from "../../../src/kilocode/commit-message/generate" const stream = spyOn(CommitMessageRuntime, "generate").mockImplementation(async () => mockStreamText) +const model = spyOn(CommitMessageRuntime, "model").mockImplementation( + async () => ({ providerID: "test", id: "test-small-model" }) as Provider.Model, +) describe("commit-message.generate", () => { beforeEach(() => { stream.mockImplementation(async () => mockStreamText) + model.mockImplementation(async () => ({ providerID: "test", id: "test-small-model" }) as Provider.Model) mockStreamText = "feat(src): add hello world logging" mockGitContext = { ...defaultGitContext } captured = { path: "" } diff --git a/packages/opencode/test/kilocode/task-nesting.test.ts b/packages/opencode/test/kilocode/task-nesting.test.ts index 6b034ced935..803088daed0 100644 --- a/packages/opencode/test/kilocode/task-nesting.test.ts +++ b/packages/opencode/test/kilocode/task-nesting.test.ts @@ -8,6 +8,7 @@ import { MessageV2 } from "../../src/session/message-v2" import type { SessionPrompt } from "../../src/session/prompt" import { MessageID, PartID } from "../../src/session/schema" import { ModelID, ProviderID } from "../../src/provider/schema" +import { Provider } from "../../src/provider/provider" import { TaskTool, type TaskPromptOps } from "../../src/tool/task" import { Truncate } from "../../src/tool/truncate" import { ToolRegistry } from "../../src/tool/registry" @@ -26,6 +27,7 @@ const it = testEffect( CrossSpawnSpawner.defaultLayer, Session.defaultLayer, Truncate.defaultLayer, + Provider.defaultLayer, ToolRegistry.defaultLayer, ), ) diff --git a/packages/opencode/test/kilocode/tool-task-model.test.ts b/packages/opencode/test/kilocode/tool-task-model.test.ts index bafe1e19cd6..58d69033976 100644 --- a/packages/opencode/test/kilocode/tool-task-model.test.ts +++ b/packages/opencode/test/kilocode/tool-task-model.test.ts @@ -12,6 +12,7 @@ import { MessageV2 } from "../../src/session/message-v2" import type { SessionPrompt } from "../../src/session/prompt" import { MessageID, PartID } from "../../src/session/schema" import { ModelID, ProviderID } from "../../src/provider/schema" +import { Provider } from "../../src/provider/provider" import { TaskTool, type TaskPromptOps } from "../../src/tool/task" import { Truncate } from "../../src/tool/truncate" import { ToolRegistry } from "../../src/tool/registry" @@ -94,6 +95,7 @@ const it = testEffect( CrossSpawnSpawner.defaultLayer, Session.defaultLayer, Truncate.defaultLayer, + Provider.defaultLayer, ToolRegistry.defaultLayer, ), ) diff --git a/packages/opencode/test/tool/task.test.ts b/packages/opencode/test/tool/task.test.ts index 758e27c2e32..df991a56af9 100644 --- a/packages/opencode/test/tool/task.test.ts +++ b/packages/opencode/test/tool/task.test.ts @@ -8,6 +8,7 @@ import { MessageV2 } from "../../src/session/message-v2" import type { SessionPrompt } from "../../src/session/prompt" import { MessageID, PartID, SessionID } from "../../src/session/schema" // kilocode_change - SessionID used by cost propagation tests import { ModelID, ProviderID } from "../../src/provider/schema" +import { Provider } from "../../src/provider/provider" // kilocode_change import { TaskTool, type TaskPromptOps } from "../../src/tool/task" import { Truncate } from "@/tool/truncate" import { ToolRegistry } from "@/tool/registry" @@ -30,6 +31,7 @@ const it = testEffect( CrossSpawnSpawner.defaultLayer, Session.defaultLayer, Truncate.defaultLayer, + Provider.defaultLayer, // kilocode_change ToolRegistry.defaultLayer, ), )