simulated-provider.test.ts 9.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254
  1. import { expect, test } from "bun:test"
  2. import { Deferred, Effect, Fiber, Queue, Stream } from "effect"
  3. import type { Scope } from "effect/Scope"
  4. import { SimulatedProvider } from "../src/backend/simulated-provider"
  5. import { availableEndpoint, connect } from "./fixture/websocket"
  6. test("streams a Drive-controlled provider response and removes the finished invocation", async () => {
  7. await runProvider((provider, socket, messages) =>
  8. Effect.gen(function* () {
  9. socket.send(
  10. JSON.stringify({
  11. jsonrpc: "2.0",
  12. id: 0,
  13. method: "simulation.handshake",
  14. params: {
  15. client: { name: "test", version: "test" },
  16. expectedRole: "backend",
  17. offeredVersions: [1],
  18. requiredCapabilities: ["llm.attach", "llm.request"],
  19. optionalCapabilities: [],
  20. },
  21. }),
  22. )
  23. expect(yield* Queue.take(messages)).toMatchObject({
  24. id: 0,
  25. result: {
  26. protocolVersion: 1,
  27. role: "backend",
  28. server: { name: "opencode", version: expect.any(String) },
  29. capabilities: expect.arrayContaining(["llm.attach", "llm.request"]),
  30. },
  31. })
  32. socket.send("{")
  33. expect(yield* Queue.take(messages)).toMatchObject({ id: null, error: { code: -32000 } })
  34. yield* attach(socket, messages)
  35. const response = yield* provider.stream(request).pipe(Stream.runCollect, Effect.forkScoped)
  36. const opened = yield* takeInvocation(messages)
  37. expect(opened).toMatchObject({
  38. method: "llm.request",
  39. params: {
  40. url: "https://api.openai.com/v1/chat/completions",
  41. body: { model: "gpt-5" },
  42. },
  43. })
  44. const params = requireRecord(opened.params)
  45. if (typeof params.id !== "string") throw new Error("llm.request did not contain an invocation id")
  46. expect(response.pollUnsafe()).toBeUndefined()
  47. socket.send(
  48. JSON.stringify({
  49. jsonrpc: "2.0",
  50. id: 2,
  51. method: "llm.chunk",
  52. params: { id: params.id, items: [{ type: "textDelta", text: "Hello from Drive" }] },
  53. }),
  54. )
  55. expect(yield* Queue.take(messages)).toMatchObject({ id: 2, result: { ok: true } })
  56. socket.send(
  57. JSON.stringify({
  58. jsonrpc: "2.0",
  59. id: 3,
  60. method: "llm.finish",
  61. params: { id: params.id, reason: "stop" },
  62. }),
  63. )
  64. expect(yield* Queue.take(messages)).toMatchObject({ id: 3, result: { ok: true } })
  65. expect(Array.from(yield* Fiber.join(response))).toEqual([
  66. { type: "textDelta", text: "Hello from Drive" },
  67. { type: "finish", reason: "stop" },
  68. ])
  69. socket.send(JSON.stringify({ jsonrpc: "2.0", id: 4, method: "llm.pending" }))
  70. expect(yield* Queue.take(messages)).toMatchObject({ id: 4, result: { invocations: [] } })
  71. }),
  72. )
  73. })
  74. test("replays an invocation to a controller that attaches after it opens", async () => {
  75. await runProvider((provider, socket, messages) =>
  76. Effect.gen(function* () {
  77. const response = yield* provider.stream(request).pipe(Stream.runCollect, Effect.forkScoped)
  78. socket.send(JSON.stringify({ jsonrpc: "2.0", id: 1, method: "llm.attach" }))
  79. const received = [requireRecord(yield* Queue.take(messages)), requireRecord(yield* Queue.take(messages))]
  80. expect(received).toContainEqual(expect.objectContaining({ id: 1, result: { attached: true } }))
  81. const opened = received.find((message) => message.method === "llm.request")
  82. if (!opened) throw new Error("The pending invocation was not replayed")
  83. const params = requireRecord(opened.params)
  84. if (typeof params.id !== "string") throw new Error("llm.request did not contain an invocation id")
  85. socket.send(
  86. JSON.stringify({ jsonrpc: "2.0", id: 2, method: "llm.finish", params: { id: params.id, reason: "stop" } }),
  87. )
  88. expect(yield* Queue.take(messages)).toMatchObject({ id: 2, result: { ok: true } })
  89. expect(Array.from(yield* Fiber.join(response))).toEqual([{ type: "finish", reason: "stop" }])
  90. }),
  91. )
  92. })
  93. test("replaces the previous attached controller", async () => {
  94. const endpoint = availableEndpoint()
  95. await Effect.runPromise(
  96. Effect.gen(function* () {
  97. const provider = yield* SimulatedProvider.Service
  98. const first = yield* connect(endpoint)
  99. const second = yield* connect(endpoint)
  100. const firstMessages = yield* messagesFrom(first)
  101. const secondMessages = yield* messagesFrom(second)
  102. yield* attach(first, firstMessages)
  103. yield* attach(second, secondMessages)
  104. const response = yield* provider.stream(request).pipe(Stream.runCollect, Effect.forkScoped)
  105. const opened = yield* takeInvocation(secondMessages)
  106. expect(yield* Queue.size(firstMessages)).toBe(0)
  107. const params = requireRecord(opened.params)
  108. if (typeof params.id !== "string") throw new Error("llm.request did not contain an invocation id")
  109. second.send(
  110. JSON.stringify({ jsonrpc: "2.0", id: 2, method: "llm.finish", params: { id: params.id, reason: "stop" } }),
  111. )
  112. expect(yield* Queue.take(secondMessages)).toMatchObject({ id: 2, result: { ok: true } })
  113. expect(Array.from(yield* Fiber.join(response))).toEqual([{ type: "finish", reason: "stop" }])
  114. }).pipe(Effect.provide(SimulatedProvider.layerDrive({ endpoint })), Effect.scoped),
  115. )
  116. })
  117. test("removes an invocation when its response stream is interrupted", async () => {
  118. await runProvider((provider, socket, messages) =>
  119. Effect.gen(function* () {
  120. yield* attach(socket, messages)
  121. const response = yield* provider.stream(request).pipe(Stream.runDrain, Effect.forkScoped)
  122. yield* takeInvocation(messages)
  123. yield* Fiber.interrupt(response)
  124. socket.send(JSON.stringify({ jsonrpc: "2.0", id: 2, method: "llm.pending" }))
  125. expect(yield* Queue.take(messages)).toMatchObject({ id: 2, result: { invocations: [] } })
  126. }),
  127. )
  128. })
  129. test("releases a backpressured response when its consumer is interrupted", async () => {
  130. await runProvider((provider, socket, messages) =>
  131. Effect.gen(function* () {
  132. yield* attach(socket, messages)
  133. const started = yield* Deferred.make<void>()
  134. const response = yield* provider.stream(request).pipe(
  135. Stream.runForEach(() => Deferred.succeed(started, void 0).pipe(Effect.andThen(Effect.never))),
  136. Effect.forkScoped,
  137. )
  138. const opened = yield* takeInvocation(messages)
  139. const params = requireRecord(opened.params)
  140. if (typeof params.id !== "string") throw new Error("llm.request did not contain an invocation id")
  141. socket.send(
  142. JSON.stringify({
  143. jsonrpc: "2.0",
  144. id: 2,
  145. method: "llm.chunk",
  146. params: {
  147. id: params.id,
  148. items: Array.from({ length: 300 }, (_, index) => ({ type: "textDelta", text: String(index) })),
  149. },
  150. }),
  151. )
  152. const result = yield* Queue.take(messages).pipe(Effect.forkScoped)
  153. yield* Deferred.await(started)
  154. expect(result.pollUnsafe()).toBeUndefined()
  155. yield* Fiber.interrupt(response)
  156. expect(yield* Fiber.join(result)).toMatchObject({ id: 2 })
  157. }),
  158. )
  159. })
  160. test("fails the provider stream when Drive disconnects the invocation", async () => {
  161. await runProvider((provider, socket, messages) =>
  162. Effect.gen(function* () {
  163. yield* attach(socket, messages)
  164. const response = yield* provider.stream(request).pipe(Stream.runCollect, Effect.flip, Effect.forkScoped)
  165. const opened = yield* takeInvocation(messages)
  166. const params = requireRecord(opened.params)
  167. if (typeof params.id !== "string") throw new Error("llm.request did not contain an invocation id")
  168. socket.send(JSON.stringify({ jsonrpc: "2.0", id: 2, method: "llm.disconnect", params: { id: params.id } }))
  169. expect(yield* Queue.take(messages)).toMatchObject({ id: 2, result: { ok: true } })
  170. expect(yield* Fiber.join(response)).toBeInstanceOf(SimulatedProvider.ProviderDisconnectedError)
  171. }),
  172. )
  173. })
  174. const request: SimulatedProvider.ProviderRequest = {
  175. url: "https://api.openai.com/v1/chat/completions",
  176. body: { model: "gpt-5", messages: [{ role: "user", content: "Hello" }] },
  177. }
  178. function runProvider<E>(
  179. body: (
  180. provider: SimulatedProvider.Interface,
  181. socket: WebSocket,
  182. messages: Queue.Queue<unknown>,
  183. ) => Effect.Effect<void, E, Scope>,
  184. ) {
  185. const endpoint = availableEndpoint()
  186. return Effect.runPromise(
  187. Effect.gen(function* () {
  188. const provider = yield* SimulatedProvider.Service
  189. const socket = yield* connect(endpoint)
  190. const messages = yield* messagesFrom(socket)
  191. yield* body(provider, socket, messages)
  192. }).pipe(Effect.provide(SimulatedProvider.layerDrive({ endpoint })), Effect.scoped),
  193. )
  194. }
  195. function messagesFrom(socket: WebSocket) {
  196. return Effect.gen(function* () {
  197. const messages = yield* Queue.unbounded<unknown>()
  198. socket.addEventListener("message", (event) => {
  199. Queue.offerUnsafe(messages, JSON.parse(String(event.data)))
  200. })
  201. return messages
  202. })
  203. }
  204. function attach(socket: WebSocket, messages: Queue.Queue<unknown>) {
  205. return Effect.gen(function* () {
  206. socket.send(JSON.stringify({ jsonrpc: "2.0", id: 1, method: "llm.attach" }))
  207. expect(yield* Queue.take(messages)).toMatchObject({ id: 1, result: { attached: true } })
  208. })
  209. }
  210. function takeInvocation(messages: Queue.Queue<unknown>) {
  211. return Queue.take(messages).pipe(
  212. Effect.map((message) => {
  213. const opened = requireRecord(message)
  214. if (opened.method !== "llm.request") throw new Error("Expected an llm.request notification")
  215. return opened
  216. }),
  217. )
  218. }
  219. function requireRecord(value: unknown): Record<string, unknown> {
  220. if (!isRecord(value)) throw new Error("Expected an object")
  221. return value
  222. }
  223. function isRecord(value: unknown): value is Record<string, unknown> {
  224. return typeof value === "object" && value !== null
  225. }