refactor(session): move llm stream into layer (#22358)

This commit is contained in:
Kit Langton
2026-04-13 19:53:30 -04:00
committed by GitHub
parent 43b37346b6
commit e8471256f2
2 changed files with 364 additions and 448 deletions
+61 -52
View File
@@ -1,7 +1,6 @@
import { Provider } from "@/provider/provider" import { Provider } from "@/provider/provider"
import { Log } from "@/util/log" import { Log } from "@/util/log"
import { Cause, Effect, Layer, Record, Context } from "effect" import { Context, Effect, Layer, Record } from "effect"
import * as Queue from "effect/Queue"
import * as Stream from "effect/Stream" import * as Stream from "effect/Stream"
import { streamText, wrapLanguageModel, type ModelMessage, type Tool, tool, jsonSchema } from "ai" import { streamText, wrapLanguageModel, type ModelMessage, type Tool, tool, jsonSchema } from "ai"
import { mergeDeep, pipe } from "remeda" import { mergeDeep, pipe } from "remeda"
@@ -21,11 +20,13 @@ import { Wildcard } from "@/util/wildcard"
import { SessionID } from "@/session/schema" import { SessionID } from "@/session/schema"
import { Auth } from "@/auth" import { Auth } from "@/auth"
import { Installation } from "@/installation" import { Installation } from "@/installation"
import { AppRuntime } from "@/effect/app-runtime" import { makeRuntime } from "@/effect/run-service"
export namespace LLM { export namespace LLM {
const log = Log.create({ service: "llm" }) const log = Log.create({ service: "llm" })
const perms = makeRuntime(Permission.Service, Permission.defaultLayer)
export const OUTPUT_TOKEN_MAX = ProviderTransform.OUTPUT_TOKEN_MAX export const OUTPUT_TOKEN_MAX = ProviderTransform.OUTPUT_TOKEN_MAX
type Result = Awaited<ReturnType<typeof streamText>>
export type StreamInput = { export type StreamInput = {
user: MessageV2.User user: MessageV2.User
@@ -46,7 +47,7 @@ export namespace LLM {
abort: AbortSignal abort: AbortSignal
} }
export type Event = Awaited<ReturnType<typeof stream>>["fullStream"] extends AsyncIterable<infer T> ? T : never export type Event = Result["fullStream"] extends AsyncIterable<infer T> ? T : never
export interface Interface { export interface Interface {
readonly stream: (input: StreamInput) => Stream.Stream<Event, unknown> readonly stream: (input: StreamInput) => Stream.Stream<Event, unknown>
@@ -54,35 +55,16 @@ export namespace LLM {
export class Service extends Context.Service<Service, Interface>()("@opencode/LLM") {} export class Service extends Context.Service<Service, Interface>()("@opencode/LLM") {}
export const layer = Layer.effect( export const layer: Layer.Layer<Service, never, Auth.Service | Config.Service | Provider.Service | Plugin.Service> =
Layer.effect(
Service, Service,
Effect.gen(function* () { Effect.gen(function* () {
return Service.of({ const auth = yield* Auth.Service
stream(input) { const config = yield* Config.Service
return Stream.scoped( const provider = yield* Provider.Service
Stream.unwrap( const plugin = yield* Plugin.Service
Effect.gen(function* () {
const ctrl = yield* Effect.acquireRelease(
Effect.sync(() => new AbortController()),
(ctrl) => Effect.sync(() => ctrl.abort()),
)
const result = yield* Effect.promise(() => LLM.stream({ ...input, abort: ctrl.signal })) const run = Effect.fn("LLM.run")(function* (input: StreamRequest) {
return Stream.fromAsyncIterable(result.fullStream, (e) =>
e instanceof Error ? e : new Error(String(e)),
)
}),
),
)
},
})
}),
)
export const defaultLayer = layer
export async function stream(input: StreamRequest) {
const l = log const l = log
.clone() .clone()
.tag("providerID", input.model.providerID) .tag("providerID", input.model.providerID)
@@ -95,24 +77,19 @@ export namespace LLM {
modelID: input.model.id, modelID: input.model.id,
providerID: input.model.providerID, providerID: input.model.providerID,
}) })
const [language, cfg, provider, info] = await Effect.runPromise(
Effect.gen(function* () { const [language, cfg, item, info] = yield* Effect.all(
const auth = yield* Auth.Service
const cfg = yield* Config.Service
const provider = yield* Provider.Service
return yield* Effect.all(
[ [
provider.getLanguage(input.model), provider.getLanguage(input.model),
cfg.get(), config.get(),
provider.getProvider(input.model.providerID), provider.getProvider(input.model.providerID),
auth.get(input.model.providerID), auth.get(input.model.providerID),
], ],
{ concurrency: "unbounded" }, { concurrency: "unbounded" },
) )
}).pipe(Effect.provide(Layer.mergeAll(Auth.defaultLayer, Config.defaultLayer, Provider.defaultLayer))),
)
// TODO: move this to a proper hook // TODO: move this to a proper hook
const isOpenaiOauth = provider.id === "openai" && info?.type === "oauth" const isOpenaiOauth = item.id === "openai" && info?.type === "oauth"
const system: string[] = [] const system: string[] = []
system.push( system.push(
@@ -129,7 +106,7 @@ export namespace LLM {
) )
const header = system[0] const header = system[0]
await Plugin.trigger( yield* plugin.trigger(
"experimental.chat.system.transform", "experimental.chat.system.transform",
{ sessionID: input.sessionID, model: input.model }, { sessionID: input.sessionID, model: input.model },
{ system }, { system },
@@ -150,7 +127,7 @@ export namespace LLM {
: ProviderTransform.options({ : ProviderTransform.options({
model: input.model, model: input.model,
sessionID: input.sessionID, sessionID: input.sessionID,
providerOptions: provider.options, providerOptions: item.options,
}) })
const options: Record<string, any> = pipe( const options: Record<string, any> = pipe(
base, base,
@@ -177,13 +154,13 @@ export namespace LLM {
...input.messages, ...input.messages,
] ]
const params = await Plugin.trigger( const params = yield* plugin.trigger(
"chat.params", "chat.params",
{ {
sessionID: input.sessionID, sessionID: input.sessionID,
agent: input.agent.name, agent: input.agent.name,
model: input.model, model: input.model,
provider, provider: item,
message: input.user, message: input.user,
}, },
{ {
@@ -197,13 +174,13 @@ export namespace LLM {
}, },
) )
const { headers } = await Plugin.trigger( const { headers } = yield* plugin.trigger(
"chat.headers", "chat.headers",
{ {
sessionID: input.sessionID, sessionID: input.sessionID,
agent: input.agent.name, agent: input.agent.name,
model: input.model, model: input.model,
provider, provider: item,
message: input.user, message: input.user,
}, },
{ {
@@ -220,7 +197,7 @@ export namespace LLM {
// 1. Providers with "litellm" in their ID or API ID (auto-detected) // 1. Providers with "litellm" in their ID or API ID (auto-detected)
// 2. Providers with explicit "litellmProxy: true" option (opt-in for custom gateways) // 2. Providers with explicit "litellmProxy: true" option (opt-in for custom gateways)
const isLiteLLMProxy = const isLiteLLMProxy =
provider.options?.["litellmProxy"] === true || item.options?.["litellmProxy"] === true ||
input.model.providerID.toLowerCase().includes("litellm") || input.model.providerID.toLowerCase().includes("litellm") ||
input.model.api.id.toLowerCase().includes("litellm") input.model.api.id.toLowerCase().includes("litellm")
@@ -306,8 +283,7 @@ export namespace LLM {
} }
}) })
const uniquePatterns = [...new Set(toolPatterns)] as string[] const uniquePatterns = [...new Set(toolPatterns)] as string[]
await AppRuntime.runPromise( await perms.runPromise((svc) =>
Permission.Service.use((svc) =>
svc.ask({ svc.ask({
id, id,
sessionID: SessionID.make(input.sessionID), sessionID: SessionID.make(input.sessionID),
@@ -317,10 +293,12 @@ export namespace LLM {
always: uniquePatterns, always: uniquePatterns,
ruleset: [], ruleset: [],
}), }),
),
) )
for (const name of uniqueNames) approvedToolsForSession.add(name) for (const name of uniqueNames) approvedToolsForSession.add(name)
workflowModel.sessionPreapprovedTools = [...(workflowModel.sessionPreapprovedTools ?? []), ...uniqueNames] workflowModel.sessionPreapprovedTools = [
...(workflowModel.sessionPreapprovedTools ?? []),
...uniqueNames,
]
return { approved: true } return { approved: true }
} catch { } catch {
return { approved: false } return { approved: false }
@@ -407,7 +385,38 @@ export namespace LLM {
}, },
}, },
}) })
} })
const stream: Interface["stream"] = (input) =>
Stream.scoped(
Stream.unwrap(
Effect.gen(function* () {
const ctrl = yield* Effect.acquireRelease(
Effect.sync(() => new AbortController()),
(ctrl) => Effect.sync(() => ctrl.abort()),
)
const result = yield* run({ ...input, abort: ctrl.signal })
return Stream.fromAsyncIterable(result.fullStream, (e) =>
e instanceof Error ? e : new Error(String(e)),
)
}),
),
)
return Service.of({ stream })
}),
)
export const defaultLayer = Layer.suspend(() =>
layer.pipe(
Layer.provide(Auth.defaultLayer),
Layer.provide(Config.defaultLayer),
Layer.provide(Provider.defaultLayer),
Layer.provide(Plugin.defaultLayer),
),
)
function resolveTools(input: Pick<StreamInput, "tools" | "agent" | "permission" | "user">) { function resolveTools(input: Pick<StreamInput, "tools" | "agent" | "permission" | "user">) {
const disabled = Permission.disabled( const disabled = Permission.disabled(
+13 -106
View File
@@ -26,6 +26,12 @@ async function getModel(providerID: ProviderID, modelID: ModelID) {
) )
} }
const llm = makeRuntime(LLM.Service, LLM.defaultLayer)
async function drain(input: LLM.StreamInput) {
return llm.runPromise((svc) => svc.stream(input).pipe(Stream.runDrain))
}
describe("session.llm.hasToolCalls", () => { describe("session.llm.hasToolCalls", () => {
test("returns false for empty messages array", () => { test("returns false for empty messages array", () => {
expect(LLM.hasToolCalls([])).toBe(false) expect(LLM.hasToolCalls([])).toBe(false)
@@ -355,20 +361,16 @@ describe("session.llm.stream", () => {
model: { providerID: ProviderID.make(providerID), modelID: resolved.id, variant: "high" }, model: { providerID: ProviderID.make(providerID), modelID: resolved.id, variant: "high" },
} satisfies MessageV2.User } satisfies MessageV2.User
const stream = await LLM.stream({ await drain({
user, user,
sessionID, sessionID,
model: resolved, model: resolved,
agent, agent,
system: ["You are a helpful assistant."], system: ["You are a helpful assistant."],
abort: new AbortController().signal,
messages: [{ role: "user", content: "Hello" }], messages: [{ role: "user", content: "Hello" }],
tools: {}, tools: {},
}) })
for await (const _ of stream.fullStream) {
}
const capture = await request const capture = await request
const body = capture.body const body = capture.body
const headers = capture.headers const headers = capture.headers
@@ -393,80 +395,6 @@ describe("session.llm.stream", () => {
}) })
}) })
test("raw stream abort signal cancels provider response body promptly", async () => {
const server = state.server
if (!server) throw new Error("Server not initialized")
const providerID = "alibaba"
const modelID = "qwen-plus"
const fixture = await loadFixture(providerID, modelID)
const model = fixture.model
const pending = waitStreamingRequest("/chat/completions")
await using tmp = await tmpdir({
init: async (dir) => {
await Bun.write(
path.join(dir, "opencode.json"),
JSON.stringify({
$schema: "https://opencode.ai/config.json",
enabled_providers: [providerID],
provider: {
[providerID]: {
options: {
apiKey: "test-key",
baseURL: `${server.url.origin}/v1`,
},
},
},
}),
)
},
})
await Instance.provide({
directory: tmp.path,
fn: async () => {
const resolved = await getModel(ProviderID.make(providerID), ModelID.make(model.id))
const sessionID = SessionID.make("session-test-raw-abort")
const agent = {
name: "test",
mode: "primary",
options: {},
permission: [{ permission: "*", pattern: "*", action: "allow" }],
} satisfies Agent.Info
const user = {
id: MessageID.make("user-raw-abort"),
sessionID,
role: "user",
time: { created: Date.now() },
agent: agent.name,
model: { providerID: ProviderID.make(providerID), modelID: resolved.id },
} satisfies MessageV2.User
const ctrl = new AbortController()
const result = await LLM.stream({
user,
sessionID,
model: resolved,
agent,
system: ["You are a helpful assistant."],
abort: ctrl.signal,
messages: [{ role: "user", content: "Hello" }],
tools: {},
})
const iter = result.fullStream[Symbol.asyncIterator]()
await pending.request
await iter.next()
ctrl.abort()
await Promise.race([pending.responseCanceled, timeout(500)])
await Promise.race([pending.requestAborted, timeout(500)]).catch(() => undefined)
await iter.return?.()
},
})
})
test("service stream cancellation cancels provider response body promptly", async () => { test("service stream cancellation cancels provider response body promptly", async () => {
const server = state.server const server = state.server
if (!server) throw new Error("Server not initialized") if (!server) throw new Error("Server not initialized")
@@ -518,8 +446,7 @@ describe("session.llm.stream", () => {
} satisfies MessageV2.User } satisfies MessageV2.User
const ctrl = new AbortController() const ctrl = new AbortController()
const { runPromiseExit } = makeRuntime(LLM.Service, LLM.defaultLayer) const run = llm.runPromiseExit(
const run = runPromiseExit(
(svc) => (svc) =>
svc svc
.stream({ .stream({
@@ -610,14 +537,13 @@ describe("session.llm.stream", () => {
tools: { question: true }, tools: { question: true },
} satisfies MessageV2.User } satisfies MessageV2.User
const stream = await LLM.stream({ await drain({
user, user,
sessionID, sessionID,
model: resolved, model: resolved,
agent, agent,
permission: [{ permission: "question", pattern: "*", action: "allow" }], permission: [{ permission: "question", pattern: "*", action: "allow" }],
system: ["You are a helpful assistant."], system: ["You are a helpful assistant."],
abort: new AbortController().signal,
messages: [{ role: "user", content: "Hello" }], messages: [{ role: "user", content: "Hello" }],
tools: { tools: {
question: tool({ question: tool({
@@ -628,9 +554,6 @@ describe("session.llm.stream", () => {
}, },
}) })
for await (const _ of stream.fullStream) {
}
const capture = await request const capture = await request
const tools = capture.body.tools as Array<{ function?: { name?: string } }> | undefined const tools = capture.body.tools as Array<{ function?: { name?: string } }> | undefined
expect(tools?.some((item) => item.function?.name === "question")).toBe(true) expect(tools?.some((item) => item.function?.name === "question")).toBe(true)
@@ -728,20 +651,16 @@ describe("session.llm.stream", () => {
model: { providerID: ProviderID.make("openai"), modelID: resolved.id, variant: "high" }, model: { providerID: ProviderID.make("openai"), modelID: resolved.id, variant: "high" },
} satisfies MessageV2.User } satisfies MessageV2.User
const stream = await LLM.stream({ await drain({
user, user,
sessionID, sessionID,
model: resolved, model: resolved,
agent, agent,
system: ["You are a helpful assistant."], system: ["You are a helpful assistant."],
abort: new AbortController().signal,
messages: [{ role: "user", content: "Hello" }], messages: [{ role: "user", content: "Hello" }],
tools: {}, tools: {},
}) })
for await (const _ of stream.fullStream) {
}
const capture = await request const capture = await request
const body = capture.body const body = capture.body
@@ -847,13 +766,12 @@ describe("session.llm.stream", () => {
model: { providerID: ProviderID.make("openai"), modelID: resolved.id }, model: { providerID: ProviderID.make("openai"), modelID: resolved.id },
} satisfies MessageV2.User } satisfies MessageV2.User
const stream = await LLM.stream({ await drain({
user, user,
sessionID, sessionID,
model: resolved, model: resolved,
agent, agent,
system: ["You are a helpful assistant."], system: ["You are a helpful assistant."],
abort: new AbortController().signal,
messages: [ messages: [
{ {
role: "user", role: "user",
@@ -871,9 +789,6 @@ describe("session.llm.stream", () => {
tools: {}, tools: {},
}) })
for await (const _ of stream.fullStream) {
}
const capture = await request const capture = await request
expect(capture.url.pathname.endsWith("/responses")).toBe(true) expect(capture.url.pathname.endsWith("/responses")).toBe(true)
}, },
@@ -972,20 +887,16 @@ describe("session.llm.stream", () => {
model: { providerID: ProviderID.make("minimax"), modelID: ModelID.make("MiniMax-M2.5") }, model: { providerID: ProviderID.make("minimax"), modelID: ModelID.make("MiniMax-M2.5") },
} satisfies MessageV2.User } satisfies MessageV2.User
const stream = await LLM.stream({ await drain({
user, user,
sessionID, sessionID,
model: resolved, model: resolved,
agent, agent,
system: ["You are a helpful assistant."], system: ["You are a helpful assistant."],
abort: new AbortController().signal,
messages: [{ role: "user", content: "Hello" }], messages: [{ role: "user", content: "Hello" }],
tools: {}, tools: {},
}) })
for await (const _ of stream.fullStream) {
}
const capture = await request const capture = await request
const body = capture.body const body = capture.body
@@ -1073,20 +984,16 @@ describe("session.llm.stream", () => {
model: { providerID: ProviderID.make(providerID), modelID: resolved.id }, model: { providerID: ProviderID.make(providerID), modelID: resolved.id },
} satisfies MessageV2.User } satisfies MessageV2.User
const stream = await LLM.stream({ await drain({
user, user,
sessionID, sessionID,
model: resolved, model: resolved,
agent, agent,
system: ["You are a helpful assistant."], system: ["You are a helpful assistant."],
abort: new AbortController().signal,
messages: [{ role: "user", content: "Hello" }], messages: [{ role: "user", content: "Hello" }],
tools: {}, tools: {},
}) })
for await (const _ of stream.fullStream) {
}
const capture = await request const capture = await request
const body = capture.body const body = capture.body
const config = body.generationConfig as const config = body.generationConfig as