mirror of
https://github.com/Kilo-Org/kilocode.git
synced 2026-09-24 16:02:55 +08:00
Remove wrappers for Provider service (#10658)
* refactor(opencode): remove legacy Provider promise wrappers in favor of Effect service usage Replace static `Provider.list`, `Provider.getModel`, `Provider.getLanguage`, `Provider.getSmallModel`, and `Provider.defaultModel` promise helpers with direct `Provider.Service.use()` calls through `AppRuntime.runPromise`. This eliminates the `makeRuntime` import and the wrapper functions that bypassed the Effect dependency injection system. - Remove legacy promise helpers from provider.ts - Update kilo-sessions, roll-call, commit-message, enhance-prompt, and task tool to use Provider.Service via AppRuntime directly - Thread Provider.Interface into KiloTask.select for proper DI - Add Provider.defaultLayer to test layers that exercise TaskTool - Update commit-message tests to spy on CommitMessageRuntime.model() * docs(sdk): regenerate v2 SDK types and update edit endpoint description
This commit is contained in:
@@ -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<Response>
|
||||
@@ -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",
|
||||
|
||||
@@ -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<Omit<Result, "model">> {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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<string> {
|
||||
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,
|
||||
|
||||
@@ -98,6 +98,7 @@ export namespace KiloTask {
|
||||
agent: Pick<Agent.Info, "model" | "variant">
|
||||
config: Pick<Config.Info, "subagent_model" | "subagent_variant">
|
||||
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,
|
||||
|
||||
@@ -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<T extends { id: string }>(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 {
|
||||
|
||||
@@ -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<typeof Parameters>,
|
||||
@@ -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
|
||||
|
||||
@@ -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: "" }
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user