diff --git a/packages/ai/example/tutorial.ts b/packages/ai/example/tutorial.ts index 3924a57dd2..3f7be10e33 100644 --- a/packages/ai/example/tutorial.ts +++ b/packages/ai/example/tutorial.ts @@ -1,6 +1,6 @@ import { Config, Effect, Formatter, Layer, Schema, Stream } from "effect" import { LLM, LLMClient, LLMRequest, Message, ProviderID, Tool, ToolRuntime } from "@opencode-ai/ai" -import { Route, Auth, Endpoint, Framing, Protocol, RequestExecutor, WebSocketExecutor } from "@opencode-ai/ai/route" +import { Route, Auth, Endpoint, Framing, Protocol, RequestExecutor } from "@opencode-ai/ai/route" import { OpenAI } from "@opencode-ai/ai/providers" /** @@ -214,8 +214,7 @@ const FakeEcho = { // enabled at a time so the tutorial can demonstrate generate, stream, or // tool-loop behavior without spending tokens on every example. const requestExecutorLayer = RequestExecutor.fetchLayer -const llmDeps = Layer.mergeAll(requestExecutorLayer, WebSocketExecutor.layer) -const llmClientLayer = LLMClient.layer.pipe(Layer.provide(llmDeps)) +const llmClientLayer = LLMClient.layer.pipe(Layer.provide(requestExecutorLayer)) const program = Effect.gen(function* () { // yield* generateOnce @@ -223,6 +222,6 @@ const program = Effect.gen(function* () { // yield* generateStructuredObject // yield* generateDynamicObject.pipe(Effect.andThen((response) => Effect.sync(() => console.log(response.object)))) yield* streamWithTools -}).pipe(Effect.provide(Layer.mergeAll(llmDeps, llmClientLayer))) +}).pipe(Effect.provide(Layer.mergeAll(requestExecutorLayer, llmClientLayer))) Effect.runPromise(program) diff --git a/packages/ai/src/route/client.ts b/packages/ai/src/route/client.ts index 89000c15db..c6eb0476ce 100644 --- a/packages/ai/src/route/client.ts +++ b/packages/ai/src/route/client.ts @@ -1,12 +1,10 @@ import { Cause, Context, Effect, Layer, Schema, Stream } from "effect" -import * as Option from "effect/Option" import { Auth } from "./auth" import { Endpoint, type EndpointPatch } from "./endpoint" import { RequestExecutor } from "./executor" import { Framing } from "./framing" import { HttpTransport } from "./transport" -import type { HttpMiddleware, Transport, TransportRuntime } from "./transport" -import { WebSocketExecutor } from "./transport" +import type { HttpMiddleware, Transport, TransportRuntime, WebSocketChannelExecutor } from "./transport" import type { Protocol } from "./protocol" import { applyCachePolicy } from "../cache-policy" import * as ProviderShared from "../protocols/shared" @@ -58,6 +56,7 @@ export interface Route { prepared: Prepared, request: LLMRequest, runtime: TransportRuntime, + options?: StreamOptions, ) => Stream.Stream } @@ -157,6 +156,7 @@ export interface Interface { export interface StreamOptions { readonly http?: HttpMiddleware + readonly webSocket?: WebSocketChannelExecutor } export interface StreamMethod { @@ -255,13 +255,7 @@ const requireTerminalEvent = (route: string) => (events: Stream.Stream - terminal - ? Effect.void - : Effect.fail(incompleteStreamError(route)), - ), - ), + Stream.onEnd(Effect.suspend(() => (terminal ? Effect.void : Effect.fail(incompleteStreamError(route))))), ) }) @@ -321,22 +315,27 @@ function makeFromTransport( headers: routeInput.headers, middleware: options?.http, }), - streamPrepared: (prepared: Prepared, request: LLMRequest, runtime: TransportRuntime) => { + streamPrepared: (prepared: Prepared, request: LLMRequest, runtime: TransportRuntime, options?: StreamOptions) => { const route = `${request.model.provider}/${request.model.route.id}` - const events = routeInput.transport - .frames(prepared, request, runtime) - .pipe( - Stream.mapEffect(decodeEvent(route)), - protocol.stream.terminal ? Stream.takeUntil(protocol.stream.terminal) : (stream) => stream, - ) - return events.pipe( - Stream.mapAccumEffect( - () => protocol.stream.initial(request), - protocol.stream.step, - protocol.stream.onHalt ? { onHalt: protocol.stream.onHalt } : undefined, + return Stream.unwrap( + routeInput.transport.execute(prepared, request, runtime, options).pipe( + Effect.map((execution) => { + const events = execution.frames.pipe( + Stream.mapEffect(decodeEvent(route)), + protocol.stream.terminal ? Stream.takeUntil(protocol.stream.terminal) : (stream) => stream, + ) + const stream = events.pipe( + Stream.mapAccumEffect( + () => protocol.stream.initial(request), + protocol.stream.step, + protocol.stream.onHalt ? { onHalt: protocol.stream.onHalt } : undefined, + ), + Stream.catchCause((cause) => Stream.fail(streamError(route, `Failed to read ${route} stream`, cause))), + requireTerminalEvent(route), + ) + return execution.complete ? stream.pipe(Stream.onEnd(execution.complete)) : stream + }), ), - Stream.catchCause((cause) => Stream.fail(streamError(route, `Failed to read ${route} stream`, cause))), - requireTerminalEvent(route), ) }, } satisfies Route @@ -419,7 +418,7 @@ const streamRequestWith = (runtime: TransportRuntime) => (request: LLMRequest, o Stream.unwrap( Effect.gen(function* () { const compiled = yield* compile(request, options) - return compiled.route.streamPrepared(compiled.prepared, compiled.request, runtime) + return compiled.route.streamPrepared(compiled.prepared, compiled.request, runtime, options) }), ) @@ -457,7 +456,6 @@ export const layer: Layer.Layer = Layer Effect.gen(function* () { const stream = streamRequestWith({ http: yield* RequestExecutor.Service, - webSocket: Option.getOrUndefined(yield* Effect.serviceOption(WebSocketExecutor.Service)), }) return Service.of({ stream, generate: generateWith(stream) }) }), diff --git a/packages/ai/src/route/index.ts b/packages/ai/src/route/index.ts index eb9f1759c5..159f571aa8 100644 --- a/packages/ai/src/route/index.ts +++ b/packages/ai/src/route/index.ts @@ -16,11 +16,28 @@ export { AuthOptions } from "./auth-options" export { Endpoint } from "./endpoint" export { Framing } from "./framing" export { Protocol } from "./protocol" -export { HttpTransport, WebSocketExecutor, WebSocketTransport } from "./transport" +export { HttpTransport, WebSocketTransport } from "./transport" export * as Transport from "./transport" export type { Definition as AuthShape, AuthInput, Credential, CredentialError } from "./auth" export type { ApiKeyMode, AuthOverride, ProviderAuthOption } from "./auth-options" export type { Definition as EndpointFn, EndpointInput } from "./endpoint" export type { Definition as FramingDef } from "./framing" export type { Protocol as ProtocolDef } from "./protocol" -export type { HttpHandler, HttpMiddleware, Transport as TransportDef, TransportRuntime } from "./transport" +export type { + ChannelCheckpoint, + ChannelCreate, + ChannelObservation, + HttpHandler, + HttpMiddleware, + Transport as TransportDef, + TransportExecuteOptions, + TransportExecution, + TransportRuntime, + WebSocketConnection, + WebSocketChannelDriver, + WebSocketChannelExchange, + WebSocketChannelExecution, + WebSocketChannelExecutor, + WebSocketConnector, + WebSocketRequest, +} from "./transport" diff --git a/packages/ai/src/route/transport/http.ts b/packages/ai/src/route/transport/http.ts index 223afb353c..c394bcf189 100644 --- a/packages/ai/src/route/transport/http.ts +++ b/packages/ai/src/route/transport/http.ts @@ -86,26 +86,28 @@ export const httpJson = (input: HttpJsonInput): HttpJs middleware: prepareInput.middleware, } }), - frames: (prepared, request, runtime) => - Stream.unwrap( - runtime.http - .execute(prepared.request, prepared.middleware) - .pipe( - Effect.map((response) => - prepared.framing.frame( - response.stream.pipe( - Stream.mapError((error) => - ProviderShared.eventError( - `${request.model.provider}/${request.model.route.id}`, - `Failed to read ${request.model.provider}/${request.model.route.id} stream`, - ProviderShared.errorText(error), + execute: (prepared, request, runtime) => + Effect.succeed({ + frames: Stream.unwrap( + runtime.http + .execute(prepared.request, prepared.middleware) + .pipe( + Effect.map((response) => + prepared.framing.frame( + response.stream.pipe( + Stream.mapError((error) => + ProviderShared.eventError( + `${request.model.provider}/${request.model.route.id}`, + `Failed to read ${request.model.provider}/${request.model.route.id} stream`, + ProviderShared.errorText(error), + ), ), ), ), ), ), - ), - ), + ), + }), }) export const sseJson = { diff --git a/packages/ai/src/route/transport/index.ts b/packages/ai/src/route/transport/index.ts index c74e578e8b..646790e61e 100644 --- a/packages/ai/src/route/transport/index.ts +++ b/packages/ai/src/route/transport/index.ts @@ -1,19 +1,33 @@ -import type { Effect, Stream } from "effect" +import type { Effect, Scope, Stream } from "effect" import { Endpoint } from "../endpoint" import { Auth } from "../auth" import type { HttpMiddleware, Interface as RequestExecutorInterface } from "../executor" -import type { Interface as WebSocketExecutorInterface } from "./websocket" +import type { WebSocketChannelExecutor } from "./websocket-channel" import type { AIError, LLMRequest } from "../../schema" export interface TransportRuntime { readonly http: RequestExecutorInterface - readonly webSocket?: WebSocketExecutorInterface +} + +export interface TransportExecution { + readonly frames: Stream.Stream + /** Optional successful-consumption acknowledgement. HTTP leaves this absent. */ + readonly complete?: Effect.Effect +} + +export interface TransportExecuteOptions { + readonly webSocket?: WebSocketChannelExecutor } export interface Transport { readonly id: string readonly prepare: (input: TransportPrepareInput) => Effect.Effect - readonly frames: (prepared: Prepared, request: LLMRequest, runtime: TransportRuntime) => Stream.Stream + readonly execute: ( + prepared: Prepared, + request: LLMRequest, + runtime: TransportRuntime, + options?: TransportExecuteOptions, + ) => Effect.Effect, AIError, Scope.Scope> } export interface TransportPrepareInput { @@ -28,4 +42,14 @@ export interface TransportPrepareInput { export * as HttpTransport from "./http" export type { HttpHandler, HttpMiddleware } from "../executor" -export { WebSocketExecutor, WebSocketTransport } from "./websocket" +export type { + ChannelCheckpoint, + ChannelCreate, + ChannelObservation, + WebSocketChannelDriver, + WebSocketChannelExchange, + WebSocketChannelExecution, + WebSocketChannelExecutor, +} from "./websocket-channel" +export type { WebSocketConnection, WebSocketConnector, WebSocketRequest } from "./websocket" +export { WebSocketTransport } from "./websocket" diff --git a/packages/ai/src/route/transport/websocket-channel.ts b/packages/ai/src/route/transport/websocket-channel.ts new file mode 100644 index 0000000000..174676e070 --- /dev/null +++ b/packages/ai/src/route/transport/websocket-channel.ts @@ -0,0 +1,48 @@ +import type { Effect, Scope, Stream } from "effect" +import type { Headers } from "effect/unstable/http" +import type { AIError } from "../../schema" + +export interface WebSocketChannelExecutor { + readonly execute: ( + exchange: WebSocketChannelExchange, + ) => Effect.Effect +} + +export interface WebSocketChannelExecution { + readonly frames: Stream.Stream + /** Commits staged state after the decoded Route stream ends successfully. */ + readonly complete: Effect.Effect +} + +export interface WebSocketChannelExchange { + readonly id: string + readonly connect: { + readonly url: string + readonly headers: Headers.Headers + } + readonly fallback: () => Stream.Stream + readonly driver: WebSocketChannelDriver +} + +export interface WebSocketChannelDriver { + readonly create: (checkpoint: ChannelCheckpoint | undefined) => Effect.Effect + readonly observe: (create: ChannelCreate, frame: string) => Effect.Effect +} + +export interface ChannelCreate { + readonly message: string + readonly mode: "full" | "incremental" +} + +export type ChannelObservation = + | { readonly type: "frame"; readonly frame: string } + | { readonly type: "completed"; readonly frame: string; readonly checkpoint?: ChannelCheckpoint } + | { readonly type: "incomplete"; readonly frame: string } + | { readonly type: "provider-failure"; readonly error: AIError } + | { readonly type: "rejected"; readonly error: AIError; readonly recovery: "retry-full" } + | { readonly type: "rejected"; readonly error: AIError; readonly recovery: "rotate-and-retry-full" } + +export interface ChannelCheckpoint { + readonly protocol: string + readonly value: unknown +} diff --git a/packages/ai/src/route/transport/websocket.ts b/packages/ai/src/route/transport/websocket.ts index 39cdc26ee0..914ae6b682 100644 --- a/packages/ai/src/route/transport/websocket.ts +++ b/packages/ai/src/route/transport/websocket.ts @@ -1,8 +1,15 @@ -import { Cause, Context, Effect, Layer, Queue, Stream } from "effect" +import { Cause, Effect, Queue, Stream } from "effect" import { Headers } from "effect/unstable/http" +import { Socket } from "effect/unstable/socket" import { AIError, TransportReason } from "../../schema" import * as HttpTransport from "./http" import type { Transport } from "./index" +import type { + ChannelObservation, + WebSocketChannelDriver, + WebSocketChannelExchange, + WebSocketChannelExecutor, +} from "./websocket-channel" export interface WebSocketRequest { readonly url: string @@ -15,17 +22,15 @@ export interface WebSocketConnection { readonly close: Effect.Effect } -export interface Interface { +export interface WebSocketConnector { readonly open: (input: WebSocketRequest) => Effect.Effect } -type WebSocketConstructorWithHeaders = new ( +type WebSocketConstructorWithHeaders = ( url: string, options?: { readonly headers?: Headers.Headers }, ) => globalThis.WebSocket -export class Service extends Context.Service()("@opencode/AI/WebSocketExecutor") {} - const transportError = ( method: string, message: string, @@ -37,7 +42,7 @@ const transportError = ( } = {}, ) => new AIError({ - module: "WebSocketExecutor", + module: "WebSocketConnector", method, reason: new TransportReason({ message, @@ -165,19 +170,25 @@ const webSocketUrl = (value: string) => }) export const open = (input: WebSocketRequest) => - Effect.try({ - try: () => - new (globalThis.WebSocket as unknown as WebSocketConstructorWithHeaders)(input.url, { headers: input.headers }), - catch: (error) => - transportError("open", error instanceof Error ? error.message : "Failed to construct WebSocket", { - url: input.url, - kind: "open", - phase: "connect", - delivery: "not-sent", - }), - }).pipe(Effect.flatMap((ws) => fromWebSocket(ws, input))) - -export const layer: Layer.Layer = Layer.succeed(Service, Service.of({ open })) + Effect.gen(function* () { + const constructor = yield* Socket.WebSocketConstructor + const ws = yield* Effect.try({ + try: () => + // Platform implementations may extend Effect's browser-compatible constructor with handshake options. + // oxlint-disable-next-line typescript-eslint/no-unsafe-type-assertion + (constructor as unknown as WebSocketConstructorWithHeaders)(input.url, { + headers: input.headers, + }), + catch: (error) => + transportError("open", error instanceof Error ? error.message : "Failed to construct WebSocket", { + url: input.url, + kind: "open", + phase: "connect", + delivery: "not-sent", + }), + }) + return yield* fromWebSocket(ws, input) + }) export const fromWebSocket = ( ws: globalThis.WebSocket, @@ -263,6 +274,57 @@ export const fromWebSocket = ( export const messageText = (message: string | Uint8Array, decoder: TextDecoder) => typeof message === "string" ? message : decoder.decode(message) +const observationFrame = (observation: ChannelObservation) => { + if (observation.type === "frame" || observation.type === "completed" || observation.type === "incomplete") + return Effect.succeed(observation.frame) + return Effect.fail(observation.error) +} + +const observationTerminal = (observation: ChannelObservation) => observation.type !== "frame" + +export const makeDirect = (connector: WebSocketConnector): WebSocketChannelExecutor => ({ + execute: (exchange) => + Effect.gen(function* () { + const connection = yield* Effect.acquireRelease( + connector + .open(exchange.connect) + .pipe(Effect.mapError((error) => annotateTransportError(error, { phase: "connect", delivery: "not-sent" }))), + (connection) => connection.close, + ) + const create = yield* exchange.driver.create(undefined) + yield* connection.sendText(create.message) + const decoder = new TextDecoder() + let observed = false + return { + frames: connection.messages.pipe( + Stream.map((message) => { + observed = true + return messageText(message, decoder) + }), + Stream.mapError((error) => + annotateTransportError(error, { + phase: error.reason._tag === "Transport" && error.reason.phase === "close" ? "close" : "receive", + delivery: observed ? "accepted" : "ambiguous", + }), + ), + Stream.mapEffect((frame) => exchange.driver.observe(create, frame)), + Stream.takeUntil(observationTerminal), + Stream.mapEffect(observationFrame), + ), + complete: Effect.void, + } + }), +}) + +export const direct: Effect.Effect = Effect.gen( + function* () { + const constructor = yield* Socket.WebSocketConstructor + return makeDirect({ + open: (input) => open(input).pipe(Effect.provideService(Socket.WebSocketConstructor, constructor)), + }) + }, +) + export interface JsonPrepared { readonly url: string readonly headers: Headers.Headers @@ -294,11 +356,11 @@ export const json = (input: JsonInput): JsonTransp message: input.encodeMessage(yield* input.toMessage(parts.jsonBody)), } }), - frames: (prepared, _request, runtime) => { - const webSocket = runtime.webSocket + execute: (prepared, request, _runtime, options) => { + const webSocket = options?.webSocket if (!webSocket) { - return Stream.fail( - transportError("json", "WebSocket JSON transport requires WebSocketExecutor.Service", { + return Effect.fail( + transportError("json", "WebSocket JSON transport requires StreamOptions.webSocket", { url: prepared.url, kind: "websocket", phase: "prepare", @@ -306,33 +368,25 @@ export const json = (input: JsonInput): JsonTransp }), ) } - const decoder = new TextDecoder() - return Stream.unwrap( - Effect.gen(function* () { - const connection = yield* Effect.acquireRelease( - webSocket - .open({ url: prepared.url, headers: prepared.headers }) - .pipe( - Effect.mapError((error) => annotateTransportError(error, { phase: "connect", delivery: "not-sent" })), - ), - (connection) => connection.close, - ) - yield* connection.sendText(prepared.message) - let observed = false - return connection.messages.pipe( - Stream.map((message) => { - observed = true - return messageText(message, decoder) + const driver: WebSocketChannelDriver = { + create: () => Effect.succeed({ message: prepared.message, mode: "full" }), + observe: (_create, frame) => Effect.succeed({ type: "frame", frame }), + } + const exchange: WebSocketChannelExchange = { + id: request.id ?? "request", + connect: { url: prepared.url, headers: prepared.headers }, + fallback: () => + Stream.fail( + transportError("fallback", "WebSocket JSON transport does not provide HTTP fallback", { + url: prepared.url, + kind: "websocket", + phase: "fallback", + delivery: "not-sent", }), - Stream.mapError((error) => - annotateTransportError(error, { - phase: error.reason._tag === "Transport" && error.reason.phase === "close" ? "close" : "receive", - delivery: observed ? "accepted" : "ambiguous", - }), - ), - ) - }), - ) + ), + driver, + } + return webSocket.execute(exchange) }, }) @@ -341,15 +395,12 @@ export const jsonTransport = { with: json, } as const -export const WebSocketExecutor = { - Service, - layer, +export const WebSocketTransport = { + json, + jsonTransport, + direct, + makeDirect, open, fromWebSocket, messageText, } as const - -export const WebSocketTransport = { - json, - jsonTransport, -} as const diff --git a/packages/ai/test/executor.test.ts b/packages/ai/test/executor.test.ts index b4e5fe62b8..1e16efd6a8 100644 --- a/packages/ai/test/executor.test.ts +++ b/packages/ai/test/executor.test.ts @@ -1,10 +1,11 @@ import { describe, expect } from "bun:test" -import { Effect, Layer, Ref } from "effect" +import { Deferred, Effect, Fiber, Layer, Ref, Stream } from "effect" import { Headers, HttpClient, HttpClientRequest, HttpClientResponse } from "effect/unstable/http" import { LLM, AIError } from "../src" -import { LLMClient, RequestExecutor } from "../src/route" +import { LLMClient, RequestExecutor, WebSocketTransport, type WebSocketChannelExecutor } from "../src/route" import * as OpenAIChat from "../src/protocols/openai-chat" -import { dynamicResponse } from "./lib/http" +import * as OpenAI from "../src/providers/openai" +import { dynamicResponse, fixedResponse } from "./lib/http" import { deltaChunk } from "./lib/openai-chunks" import { sseRaw } from "./lib/sse" import { it } from "./lib/effect" @@ -413,3 +414,127 @@ describe("RequestExecutor", () => { }), ) }) + +describe("WebSocket channel execution", () => { + const model = OpenAI.configure({ baseURL: "https://api.openai.test/v1/", apiKey: "test" }).responsesWebSocket( + "gpt-4.1-mini", + ) + const request = LLM.request({ model, prompt: "Say hello." }) + const frames = [ + JSON.stringify({ type: "response.output_text.delta", item_id: "msg_1", delta: "Hi" }), + JSON.stringify({ type: "response.completed", response: { id: "resp_1" } }), + ] + + it.effect("runs a channel driver through the direct executor", () => + Effect.gen(function* () { + const sent = yield* Ref.make("") + const closed = yield* Ref.make(false) + const observed = yield* Ref.make(0) + const webSocket = WebSocketTransport.makeDirect({ + open: () => + Effect.succeed({ + sendText: (message) => Ref.set(sent, message), + messages: Stream.make("one", "done", "late"), + close: Ref.set(closed, true), + }), + }) + const received = yield* Effect.scoped( + Effect.gen(function* () { + const execution = yield* webSocket.execute({ + id: "exchange_1", + connect: { url: "wss://api.openai.test/v1/responses", headers: Headers.empty }, + fallback: () => Stream.empty, + driver: { + create: () => Effect.succeed({ message: "create", mode: "full" }), + observe: (_create, frame) => + Ref.update(observed, (value) => value + 1).pipe( + Effect.as( + frame === "done" ? { type: "completed" as const, frame } : { type: "frame" as const, frame }, + ), + ), + }, + }) + return yield* Stream.runCollect(execution.frames) + }), + ) + + expect(Array.from(received)).toEqual(["one", "done"]) + expect(yield* Ref.get(sent)).toBe("create") + expect(yield* Ref.get(observed)).toBe(2) + expect(yield* Ref.get(closed)).toBe(true) + }), + ) + + it.effect("requires a per-call WebSocket executor", () => + Effect.gen(function* () { + const error = yield* LLMClient.generate(request).pipe(Effect.provide(fixedResponse("")), Effect.flip) + + expect(error.reason).toMatchObject({ + _tag: "Transport", + phase: "prepare", + delivery: "not-sent", + }) + expect(error.message).toContain("StreamOptions.webSocket") + }), + ) + + it.effect("commits channel execution only after complete consumption", () => + Effect.gen(function* () { + const commits = yield* Ref.make(0) + const executor = (input: Stream.Stream): WebSocketChannelExecutor => ({ + execute: () => + Effect.succeed({ + frames: input, + complete: Ref.update(commits, (value) => value + 1), + }), + }) + + const response = yield* LLMClient.generate(request, { + webSocket: executor(Stream.fromArray(frames)), + }).pipe(Effect.provide(fixedResponse(""))) + expect(response.text).toBe("Hi") + expect(yield* Ref.get(commits)).toBe(1) + + yield* LLMClient.generate(request, { webSocket: executor(Stream.make("not-json")) }).pipe( + Effect.provide(fixedResponse("")), + Effect.flip, + ) + expect(yield* Ref.get(commits)).toBe(1) + + yield* LLMClient.stream(request, { webSocket: executor(Stream.fromArray(frames)) }).pipe( + Stream.take(1), + Stream.runDrain, + Effect.provide(fixedResponse("")), + ) + expect(yield* Ref.get(commits)).toBe(1) + }), + ) + + it.effect("does not commit interrupted channel execution", () => + Effect.gen(function* () { + const commits = yield* Ref.make(0) + const started = yield* Deferred.make() + const executor: WebSocketChannelExecutor = { + execute: () => + Effect.succeed({ + frames: Stream.fromEffect( + Deferred.succeed(started, undefined).pipe( + Effect.as(JSON.stringify({ type: "response.created", response: { id: "resp_1" } })), + ), + ).pipe(Stream.concat(Stream.never)), + complete: Ref.update(commits, (value) => value + 1), + }), + } + const fiber = yield* LLMClient.stream(request, { webSocket: executor }).pipe( + Stream.runDrain, + Effect.provide(fixedResponse("")), + Effect.forkChild({ startImmediately: true }), + ) + + yield* Deferred.await(started) + yield* Fiber.interrupt(fiber) + + expect(yield* Ref.get(commits)).toBe(0) + }), + ) +}) diff --git a/packages/ai/test/exports.test.ts b/packages/ai/test/exports.test.ts index b0003b55d7..809cc314a7 100644 --- a/packages/ai/test/exports.test.ts +++ b/packages/ai/test/exports.test.ts @@ -1,6 +1,6 @@ import { describe, expect, test } from "bun:test" import { AIError, ImageInput, LanguageModel, LLM, LLMClient, Provider } from "@opencode-ai/ai" -import { Route, Protocol } from "@opencode-ai/ai/route" +import { Route, Protocol, WebSocketTransport } from "@opencode-ai/ai/route" import { Provider as ProviderSubpath } from "@opencode-ai/ai/provider" import { CloudflareAIGateway, @@ -37,6 +37,7 @@ describe("public exports", () => { test("route barrel exposes route-authoring APIs", () => { expect(Route.make).toBeFunction() expect(Protocol.make).toBeFunction() + expect(WebSocketTransport.makeDirect).toBeFunction() }) test("provider barrels expose user-facing facades", async () => { diff --git a/packages/ai/test/lib/http.ts b/packages/ai/test/lib/http.ts index f6c600555b..cfe7e6883b 100644 --- a/packages/ai/test/lib/http.ts +++ b/packages/ai/test/lib/http.ts @@ -1,9 +1,8 @@ import { Effect, Layer, Ref } from "effect" import { HttpClient, HttpClientRequest, HttpClientResponse } from "effect/unstable/http" -import { LLMClient, RequestExecutor, WebSocketExecutor } from "../../src/route" +import { LLMClient, RequestExecutor } from "../../src/route" import type { Service as LLMClientService } from "../../src/route/client" import type { Service as RequestExecutorService } from "../../src/route/executor" -import type { Service as WebSocketExecutorService } from "../../src/route/transport/websocket" export type HandlerInput = { readonly request: HttpClientRequest.HttpClientRequest @@ -32,13 +31,12 @@ const handlerLayer = (handler: Handler): Layer.Layer => ), ) -export type RuntimeEnv = RequestExecutorService | WebSocketExecutorService | LLMClientService +export type RuntimeEnv = RequestExecutorService | LLMClientService export const runtimeLayer = (layer: Layer.Layer): Layer.Layer => { const requestExecutorLayer = RequestExecutor.layer.pipe(Layer.provide(layer)) - const deps = Layer.mergeAll(requestExecutorLayer, WebSocketExecutor.layer) - const llmClientLayer = LLMClient.layer.pipe(Layer.provide(deps)) - return Layer.mergeAll(deps, llmClientLayer) + const llmClientLayer = LLMClient.layer.pipe(Layer.provide(requestExecutorLayer)) + return Layer.mergeAll(requestExecutorLayer, llmClientLayer) } const SSE_HEADERS = { "content-type": "text/event-stream" } as const diff --git a/packages/ai/test/provider/openai-responses.test.ts b/packages/ai/test/provider/openai-responses.test.ts index 8af8459d7a..167e69737b 100644 --- a/packages/ai/test/provider/openai-responses.test.ts +++ b/packages/ai/test/provider/openai-responses.test.ts @@ -1,5 +1,5 @@ import { describe, expect } from "bun:test" -import { ConfigProvider, Effect, Layer, Stream } from "effect" +import { ConfigProvider, Effect, Layer, Ref, Stream } from "effect" import { Headers, HttpClientRequest } from "effect/unstable/http" import { LLM, @@ -14,7 +14,7 @@ import { TransportReason, Usage, } from "../../src" -import { Auth, LLMClient, RequestExecutor, WebSocketExecutor } from "../../src/route" +import { Auth, LLMClient, RequestExecutor, WebSocketTransport } from "../../src/route" import { compileRequest } from "../../src/route/client" import * as Azure from "../../src/providers/azure" import * as OpenAI from "../../src/providers/openai" @@ -239,34 +239,29 @@ describe("OpenAI Responses route", () => { const sent: string[] = [] const opened: Array<{ readonly url: string; readonly authorization: string | undefined }> = [] let closed = false - const deps = Layer.mergeAll( - Layer.succeed( - RequestExecutor.Service, - RequestExecutor.Service.of({ - execute: () => Effect.die("unexpected HTTP request"), - }), - ), - Layer.succeed( - WebSocketExecutor.Service, - WebSocketExecutor.Service.of({ - open: (input) => - Effect.succeed({ - sendText: (message) => - Effect.sync(() => { - opened.push({ url: input.url, authorization: input.headers.authorization }) - sent.push(message) - }), - messages: Stream.fromArray([ - ProviderShared.encodeJson({ type: "response.output_text.delta", item_id: "msg_1", delta: "Hi" }), - ProviderShared.encodeJson({ type: "response.completed", response: { id: "resp_ws" } }), - ]), - close: Effect.sync(() => { - closed = true - }), - }), - }), - ), + const deps = Layer.succeed( + RequestExecutor.Service, + RequestExecutor.Service.of({ + execute: () => Effect.die("unexpected HTTP request"), + }), ) + const webSocket = WebSocketTransport.makeDirect({ + open: (input) => + Effect.succeed({ + sendText: (message) => + Effect.sync(() => { + opened.push({ url: input.url, authorization: input.headers.authorization }) + sent.push(message) + }), + messages: Stream.fromArray([ + ProviderShared.encodeJson({ type: "response.output_text.delta", item_id: "msg_1", delta: "Hi" }), + ProviderShared.encodeJson({ type: "response.completed", response: { id: "resp_ws" } }), + ]), + close: Effect.sync(() => { + closed = true + }), + }), + }) const response = yield* LLMClient.generate( LLM.request({ model: OpenAI.configure({ baseURL: "https://api.openai.test/v1/", apiKey: "test" }).responsesWebSocket( @@ -274,6 +269,7 @@ describe("OpenAI Responses route", () => { ), prompt: "Say hello.", }), + { webSocket }, ).pipe(Effect.provide(LLMClient.layer.pipe(Layer.provide(deps)))) expect(response.text).toBe("Hi") @@ -289,6 +285,48 @@ describe("OpenAI Responses route", () => { }), ) + it.effect("closes a direct WebSocket execution after partial consumption", () => + Effect.gen(function* () { + const closed = yield* Ref.make(false) + const webSocket = WebSocketTransport.makeDirect({ + open: () => + Effect.succeed({ + sendText: () => Effect.void, + messages: Stream.fromArray([ + ProviderShared.encodeJson({ type: "response.output_text.delta", item_id: "msg_1", delta: "Hi" }), + ProviderShared.encodeJson({ type: "response.completed", response: { id: "resp_ws" } }), + ]), + close: Ref.set(closed, true), + }), + }) + + yield* LLMClient.stream( + LLM.request({ + model: OpenAI.configure({ baseURL: "https://api.openai.test/v1/", apiKey: "test" }).responsesWebSocket( + "gpt-4.1-mini", + ), + prompt: "Say hello.", + }), + { webSocket }, + ).pipe( + Stream.take(1), + Stream.runDrain, + Effect.provide( + LLMClient.layer.pipe( + Layer.provide( + Layer.succeed( + RequestExecutor.Service, + RequestExecutor.Service.of({ execute: () => Effect.die("unexpected HTTP request") }), + ), + ), + ), + ), + ) + + expect(yield* Ref.get(closed)).toBe(true) + }), + ) + it.effect("terminates WebSocket control events without waiting for the socket to close", () => Effect.gen(function* () { const events = [ @@ -314,26 +352,23 @@ describe("OpenAI Responses route", () => { ), prompt: "Say hello.", }), + { + webSocket: WebSocketTransport.makeDirect({ + open: () => + Effect.succeed({ + sendText: () => Effect.void, + messages: Stream.make(ProviderShared.encodeJson(event)).pipe(Stream.concat(Stream.never)), + close: Effect.void, + }), + }), + }, ).pipe( Effect.provide( LLMClient.layer.pipe( Layer.provide( - Layer.mergeAll( - Layer.succeed( - RequestExecutor.Service, - RequestExecutor.Service.of({ execute: () => Effect.die("unexpected HTTP request") }), - ), - Layer.succeed( - WebSocketExecutor.Service, - WebSocketExecutor.Service.of({ - open: () => - Effect.succeed({ - sendText: () => Effect.void, - messages: Stream.make(ProviderShared.encodeJson(event)).pipe(Stream.concat(Stream.never)), - close: Effect.void, - }), - }), - ), + Layer.succeed( + RequestExecutor.Service, + RequestExecutor.Service.of({ execute: () => Effect.die("unexpected HTTP request") }), ), ), ), @@ -362,29 +397,24 @@ describe("OpenAI Responses route", () => { Stream.fail(failure), Stream.make(ProviderShared.encodeJson({ type: "response.created" })).pipe(Stream.concat(Stream.fail(failure))), ] - const deps = Layer.mergeAll( - Layer.succeed( - RequestExecutor.Service, - RequestExecutor.Service.of({ execute: () => Effect.die("unexpected HTTP request") }), - ), - Layer.succeed( - WebSocketExecutor.Service, - WebSocketExecutor.Service.of({ - open: () => - Effect.succeed({ - sendText: () => Effect.void, - messages: streams.shift() ?? Stream.die("unexpected WebSocket open"), - close: Effect.void, - }), - }), - ), + const deps = Layer.succeed( + RequestExecutor.Service, + RequestExecutor.Service.of({ execute: () => Effect.die("unexpected HTTP request") }), ) + const webSocket = WebSocketTransport.makeDirect({ + open: () => + Effect.succeed({ + sendText: () => Effect.void, + messages: streams.shift() ?? Stream.die("unexpected WebSocket open"), + close: Effect.void, + }), + }) const model = OpenAI.configure({ baseURL: "https://api.openai.test/v1/", apiKey: "test" }).responsesWebSocket( "gpt-4.1-mini", ) const errors = yield* Effect.forEach(["first", "second"], (prompt) => - LLMClient.generate(LLM.request({ model, prompt })).pipe( + LLMClient.generate(LLM.request({ model, prompt }), { webSocket }).pipe( Effect.provide(LLMClient.layer.pipe(Layer.provide(deps))), Effect.flip, ), @@ -399,7 +429,7 @@ describe("OpenAI Responses route", () => { it.effect("fails immediately when WebSocket is already closed", () => Effect.gen(function* () { - const error = yield* WebSocketExecutor.fromWebSocket( + const error = yield* WebSocketTransport.fromWebSocket( // oxlint-disable-next-line typescript-eslint/no-unsafe-type-assertion -- fromWebSocket reads readyState before touching WebSocket methods on this branch. { readyState: globalThis.WebSocket.CLOSED } as globalThis.WebSocket, { url: "wss://api.openai.test/v1/responses", headers: Headers.empty }, diff --git a/packages/ai/test/recorded-test.ts b/packages/ai/test/recorded-test.ts index a0df87d9ca..ae2cfd7d72 100644 --- a/packages/ai/test/recorded-test.ts +++ b/packages/ai/test/recorded-test.ts @@ -2,12 +2,11 @@ import { HttpRecorder } from "@opencode-ai/http-recorder" import { Layer } from "effect" import * as path from "node:path" import { fileURLToPath } from "node:url" -import { LLMClient, RequestExecutor, WebSocketExecutor } from "../src/route" +import { LLMClient, RequestExecutor } from "../src/route" import { ImageClient } from "../src/image-client" import type { Service as ImageClientService } from "../src/image-client" import type { Service as LLMClientService } from "../src/route/client" import type { Service as RequestExecutorService } from "../src/route/executor" -import type { Service as WebSocketExecutorService } from "../src/route/transport/websocket" import { recordedEffectGroup, type RecordedCaseOptions as RunnerCaseOptions, @@ -17,7 +16,7 @@ import { const __dirname = path.dirname(fileURLToPath(import.meta.url)) const FIXTURES_DIR = path.resolve(__dirname, "fixtures", "recordings") -type RecordedEnv = RequestExecutorService | WebSocketExecutorService | LLMClientService | ImageClientService +type RecordedEnv = RequestExecutorService | LLMClientService | ImageClientService type RecordedTestsOptions = RecordedGroupOptions & { readonly options?: HttpRecorder.RecorderOptions @@ -82,11 +81,10 @@ export const recordedTests = (options: RecordedTestsOptions) => }), ), ) - const deps = Layer.mergeAll(requestExecutor, WebSocketExecutor.layer) return Layer.mergeAll( - deps, - LLMClient.layer.pipe(Layer.provide(deps)), - ImageClient.layer.pipe(Layer.provide(deps)), + requestExecutor, + LLMClient.layer.pipe(Layer.provide(requestExecutor)), + ImageClient.layer.pipe(Layer.provide(requestExecutor)), ) }, })