From 60a46f37dbbc150317a07013612e6147269e4fcd Mon Sep 17 00:00:00 2001 From: can1357 Date: Sat, 13 Jun 2026 17:53:53 +0200 Subject: [PATCH] fix(ai): fixed SSE cancellation to stop upstream requests on client stream cancel - Extended gateway stream control to pass abort signals and onCancel into encodeStream. - Added optional cancellation control parameters to provider encodeStream handlers. - Stopped provider stream loops on cancellation and suppressed SSE completion/error output after abort. - Added a regression test verifying reader.cancel triggers onCancel and aborts upstream request. --- packages/ai/CHANGELOG.md | 2 + packages/ai/src/auth-gateway/server.ts | 9 +- packages/ai/src/auth-gateway/types.ts | 8 ++ .../providers/anthropic-messages-server.ts | 41 ++++++-- .../ai/src/providers/openai-chat-server.ts | 38 ++++++-- .../src/providers/openai-responses-server.ts | 93 ++++++++++--------- packages/ai/src/providers/pi-native-server.ts | 57 +++++++++--- .../ai/test/auth-gateway-openai-chat.test.ts | 26 ++++++ .../src/tools/browser/cmux/cmux-tab.ts | 2 +- 9 files changed, 198 insertions(+), 78 deletions(-) diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index ada2c000a..4dd8aee9e 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## [Unreleased] + ### Added - Added `GITLAB_CLIENT_ID` and `GITLAB_REDIRECT_URI` env-var overrides for the GitLab Duo OAuth login flow so users running with their own GitLab OAuth application can replace the bundled credentials when GitLab rejects the bundled `client_id`'s redirect URI. Setting `GITLAB_REDIRECT_URI` also disables the random-port fallback (strict OAuth providers reject mismatched URIs anyway). ([#2424](https://github.com/can1357/oh-my-pi/issues/2424)) @@ -12,6 +13,7 @@ ### Fixed +- Fixed streaming providers to cancel upstream model requests when the client closes the response body, so interrupted SSE sessions stop instead of continuing in the background - Fixed: provider request builders treat unknown `model.maxTokens` (`null`) as "no model cap" instead of coercing to `0` via `Math.min`; Anthropic falls back to the 64k Claude-Code cap for its required `max_tokens`. - Fixed transient stream failures on OpenAI-compatible providers by retrying HTTP 408/429/5xx responses and transient network errors with Retry-After/quota-hint aware backoff - Fixed SSE stream handling for OpenAI-compatible responses by parsing wire-level JSON frames directly and honoring `[DONE]` termination diff --git a/packages/ai/src/auth-gateway/server.ts b/packages/ai/src/auth-gateway/server.ts index 19299c112..0cbbe0f6a 100644 --- a/packages/ai/src/auth-gateway/server.ts +++ b/packages/ai/src/auth-gateway/server.ts @@ -568,7 +568,14 @@ async function handleFormatEndpoint( } if (controller.signal.aborted) return clientClosedResponse(route); - const sseStream = route.module.encodeStream(events, parsed.modelId, parsed.options); + const sseStream = route.module.encodeStream(events, parsed.modelId, parsed.options, { + signal: controller.signal, + onCancel: reason => { + if (!controller.signal.aborted) { + controller.abort(reason instanceof Error ? reason : new Error("client closed request")); + } + }, + }); return new Response(sseStream, { status: 200, headers: { diff --git a/packages/ai/src/auth-gateway/types.ts b/packages/ai/src/auth-gateway/types.ts index 333205366..6e296d434 100644 --- a/packages/ai/src/auth-gateway/types.ts +++ b/packages/ai/src/auth-gateway/types.ts @@ -110,6 +110,13 @@ export interface AuthGatewayParsedRequest { options: AuthGatewayParsedRequestOptions; } +export interface AuthGatewayStreamControl { + /** Gateway request signal. Encoders stop producing frames when it aborts. */ + signal?: AbortSignal; + /** Called when the HTTP response body is cancelled by the client. */ + onCancel?: (reason?: unknown) => void; +} + export interface AuthGatewayFormatModule { parseRequest(body: unknown, headers?: Headers): AuthGatewayParsedRequest; encodeResponse(message: AssistantMessage, requestedModelId: string): Record; @@ -117,6 +124,7 @@ export interface AuthGatewayFormatModule { events: AssistantMessageEventStream, requestedModelId: string, options?: AuthGatewayParsedRequestOptions, + control?: AuthGatewayStreamControl, ): ReadableStream; /** * Emit a protocol-specific error envelope. OpenAI returns diff --git a/packages/ai/src/providers/anthropic-messages-server.ts b/packages/ai/src/providers/anthropic-messages-server.ts index 34e388c78..f276b543e 100644 --- a/packages/ai/src/providers/anthropic-messages-server.ts +++ b/packages/ai/src/providers/anthropic-messages-server.ts @@ -30,7 +30,7 @@ import { * omp AssistantMessage[Stream] → Anthropic-shaped JSON / SSE. */ -import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { AuthGatewayStreamControl, AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; export type { ParsedRequest }; @@ -503,8 +503,15 @@ const ZERO_WIRE_USAGE: Record = { export function encodeStream( events: AssistantMessageEventStream, requestedModelId: string, + _options?: ParsedRequest["options"], + control?: AuthGatewayStreamControl, ): ReadableStream { let pingTimer: NodeJS.Timeout | undefined; + let cancelled = control?.signal?.aborted === true; + const markCancelled = () => { + cancelled = true; + }; + control?.signal?.addEventListener("abort", markCancelled, { once: true }); const stopPings = () => { if (pingTimer !== undefined) { clearInterval(pingTimer); @@ -548,6 +555,10 @@ export function encodeStream( pingTimer = setInterval(() => { try { + if (cancelled) { + stopPings(); + return; + } controller.enqueue(sseFrame("ping", { type: "ping" })); } catch { // Controller already closed/errored (client gone); stop the timer. @@ -556,8 +567,12 @@ export function encodeStream( }, STREAM_PING_INTERVAL_MS); try { + if (cancelled) { + controller.close(); + return; + } for await (const ev of events) { - if ("partial" in ev) lastPartial = ev.partial; + if (cancelled) return; switch (ev.type) { case "start": ensureStart(ev.partial); @@ -691,18 +706,24 @@ export function encodeStream( controller.enqueue(sseFrame("message_stop", { type: "message_stop" })); controller.close(); } catch (err) { - controller.enqueue( - sseFrame("error", { - type: "error", - error: { type: "api_error", message: err instanceof Error ? err.message : String(err) }, - }), - ); - controller.close(); + if (!cancelled) { + controller.enqueue( + sseFrame("error", { + type: "error", + error: { type: "api_error", message: err instanceof Error ? err.message : String(err) }, + }), + ); + controller.close(); + } } finally { + control?.signal?.removeEventListener("abort", markCancelled); stopPings(); } }, - cancel() { + cancel(reason) { + cancelled = true; + control?.signal?.removeEventListener("abort", markCancelled); + control?.onCancel?.(reason); stopPings(); }, }); diff --git a/packages/ai/src/providers/openai-chat-server.ts b/packages/ai/src/providers/openai-chat-server.ts index 49a9d3217..8e62e048c 100644 --- a/packages/ai/src/providers/openai-chat-server.ts +++ b/packages/ai/src/providers/openai-chat-server.ts @@ -4,7 +4,7 @@ import { resolvePromptCacheKey } from "../auth-gateway/http"; * Parsed inbound OpenAI chat-completions request, ready to feed into pi-ai * `stream(model, context, options)`. */ -import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { AuthGatewayStreamControl, AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; import type { AssistantMessage, AssistantMessageEventStream, @@ -501,11 +501,17 @@ export function encodeStream( events: AssistantMessageEventStream, requestedModelId: string, options?: ParsedRequest["options"], + control?: AuthGatewayStreamControl, ): ReadableStream { const encoder = new TextEncoder(); const id = makeId(); const created = Math.floor(Date.now() / 1000); const includeUsage = options?.extra?.includeStreamingUsage === true; + let cancelled = control?.signal?.aborted === true; + const markCancelled = () => { + cancelled = true; + }; + control?.signal?.addEventListener("abort", markCancelled, { once: true }); const baseChunk = (delta: Record, finishReason: string | null) => ({ id, @@ -518,7 +524,7 @@ export function encodeStream( }); const writeSse = (controller: ReadableStreamDefaultController, payload: unknown): void => { - controller.enqueue(encoder.encode(`data: ${JSON.stringify(payload)}\n\n`)); + if (!cancelled) controller.enqueue(encoder.encode(`data: ${JSON.stringify(payload)}\n\n`)); }; const writeUsage = (controller: ReadableStreamDefaultController, message: AssistantMessage): void => { @@ -545,10 +551,15 @@ export function encodeStream( let finishReason: string = "stop"; try { + if (cancelled) { + controller.close(); + return; + } // Initial role chunk. writeSse(controller, baseChunk({ role: "assistant" }, null)); for await (const event of events) { + if (cancelled) return; switch (event.type) { case "text_delta": if (event.delta.length > 0) { @@ -662,15 +673,26 @@ export function encodeStream( } // Stream ended without a terminal `done` (defensive). Close gracefully. - writeSse(controller, baseChunk({}, hasToolCalls ? "tool_calls" : "stop")); - controller.enqueue(encoder.encode("data: [DONE]\n\n")); - controller.close(); + if (!cancelled) { + writeSse(controller, baseChunk({}, hasToolCalls ? "tool_calls" : "stop")); + controller.enqueue(encoder.encode("data: [DONE]\n\n")); + controller.close(); + } } catch (err) { - const msg = err instanceof Error ? err.message : String(err); - writeSse(controller, { error: { message: msg, type: "upstream_error" } }); - controller.close(); + if (!cancelled) { + const msg = err instanceof Error ? err.message : String(err); + writeSse(controller, { error: { message: msg, type: "upstream_error" } }); + controller.close(); + } + } finally { + control?.signal?.removeEventListener("abort", markCancelled); } }, + cancel(reason) { + cancelled = true; + control?.signal?.removeEventListener("abort", markCancelled); + control?.onCancel?.(reason); + }, }); } diff --git a/packages/ai/src/providers/openai-responses-server.ts b/packages/ai/src/providers/openai-responses-server.ts index 457713c79..57fe87a1b 100644 --- a/packages/ai/src/providers/openai-responses-server.ts +++ b/packages/ai/src/providers/openai-responses-server.ts @@ -11,7 +11,7 @@ import { logger } from "@oh-my-pi/pi-utils"; import { resolvePromptCacheKey } from "../auth-gateway/http"; -import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { AuthGatewayStreamControl, AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; import type { AssistantMessage, AssistantMessageEventStream, @@ -725,18 +725,28 @@ function sseEvent(name: string, data: unknown): string { export function encodeStream( events: AssistantMessageEventStream, requestedModelId: string, + _options?: ParsedRequest["options"], + control?: AuthGatewayStreamControl, ): ReadableStream { const encoder = new TextEncoder(); const responseId = makeRespId(); let sequenceNumber = 0; + let cancelled = control?.signal?.aborted === true; + const markCancelled = () => { + cancelled = true; + }; + control?.signal?.addEventListener("abort", markCancelled, { once: true }); const seq = () => sequenceNumber++; return new ReadableStream({ async start(controller) { const emit = (name: string, data: Record) => { - controller.enqueue(encoder.encode(sseEvent(name, { type: name, sequence_number: seq(), ...data }))); + if (!cancelled) + controller.enqueue(encoder.encode(sseEvent(name, { type: name, sequence_number: seq(), ...data }))); + }; + const emitDone = () => { + if (!cancelled) controller.enqueue(encoder.encode("data: [DONE]\n\n")); }; - const emitDone = () => controller.enqueue(encoder.encode("data: [DONE]\n\n")); let createdAt = Math.floor(Date.now() / 1000); let outputIndex = 0; @@ -939,35 +949,23 @@ export function encodeStream( if (byIndex) return byIndex; return state.open?.kind === "function_call" ? state.open : undefined; }; + let finalMessage: AssistantMessage | undefined; + let failureMessage: AssistantMessage | undefined; try { - let finalMessage: AssistantMessage | null = null; - let failureMessage: AssistantMessage | null = null; - + if (cancelled) { + controller.close(); + return; + } for await (const ev of events) { + if (cancelled) return; switch (ev.type) { case "start": { createdAt = Math.floor((ev.partial.timestamp || Date.now()) / 1000); // response.created — initial envelope. - controller.enqueue( - encoder.encode( - sseEvent("response.created", { - type: "response.created", - sequence_number: seq(), - response: responseSnapshot("in_progress", []), - }), - ), - ); + emit("response.created", { response: responseSnapshot("in_progress", []) }); // response.in_progress — mirrors real OpenAI; some clients gate // on it before reading items. - controller.enqueue( - encoder.encode( - sseEvent("response.in_progress", { - type: "response.in_progress", - sequence_number: seq(), - response: responseSnapshot("in_progress", []), - }), - ), - ); + emit("response.in_progress", { response: responseSnapshot("in_progress", []) }); break; } case "text_start": { @@ -1196,26 +1194,35 @@ export function encodeStream( emitDone(); controller.close(); } catch (err) { - controller.enqueue( - encoder.encode( - sseEvent("response.failed", { - type: "response.failed", - sequence_number: seq(), - response: { - id: responseId, - object: "response", - created_at: Math.floor(Date.now() / 1000), - status: "failed", - model: requestedModelId, - output: [], - error: { message: err instanceof Error ? err.message : String(err) }, - }, - }), - ), - ); - emitDone(); - controller.close(); + if (!cancelled) { + controller.enqueue( + encoder.encode( + sseEvent("response.failed", { + type: "response.failed", + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: Math.floor(Date.now() / 1000), + status: "failed", + model: requestedModelId, + output: [], + error: { message: err instanceof Error ? err.message : String(err) }, + }, + }), + ), + ); + emitDone(); + controller.close(); + } + } finally { + control?.signal?.removeEventListener("abort", markCancelled); } }, + cancel(reason) { + cancelled = true; + control?.signal?.removeEventListener("abort", markCancelled); + control?.onCancel?.(reason); + }, }); } diff --git a/packages/ai/src/providers/pi-native-server.ts b/packages/ai/src/providers/pi-native-server.ts index eaf3e3865..7f28ded5e 100644 --- a/packages/ai/src/providers/pi-native-server.ts +++ b/packages/ai/src/providers/pi-native-server.ts @@ -25,6 +25,7 @@ * 200 JSON (stream=false): { message: AssistantMessage } * 4xx/5xx: { error: { type, message } } */ +import type { AuthGatewayParsedRequestOptions, AuthGatewayStreamControl } from "../auth-gateway/types"; import type { AssistantMessageEventStream, Context, SimpleStreamOptions } from "../types"; export interface PiNativeParsedRequest { @@ -161,31 +162,57 @@ const SSE_DONE = SSE_ENCODER.encode("data: [DONE]\n\n"); * and the client gets to feed the events straight into its existing * `AssistantMessageEventStream.push()` plumbing with zero translation. */ -export function encodeStream(events: AssistantMessageEventStream): ReadableStream { +export function encodeStream( + events: AssistantMessageEventStream, + _requestedModelId?: string, + _options?: AuthGatewayParsedRequestOptions, + control?: AuthGatewayStreamControl, +): ReadableStream { + let cancelled = control?.signal?.aborted === true; + const markCancelled = () => { + cancelled = true; + }; + control?.signal?.addEventListener("abort", markCancelled, { once: true }); return new ReadableStream({ async start(controller) { try { + if (cancelled) { + controller.close(); + return; + } for await (const event of events) { + if (cancelled) return; controller.enqueue(SSE_ENCODER.encode(`data: ${JSON.stringify(event)}\n\n`)); if (event.type === "done" || event.type === "error") break; } - controller.enqueue(SSE_DONE); - controller.close(); + if (!cancelled) { + controller.enqueue(SSE_DONE); + controller.close(); + } } catch (err) { - // Best-effort error envelope so the client iterator resolves - // instead of hanging on the dropped connection. Shape matches the - // canonical `error` event minus the unrecoverable `error: - // AssistantMessage` payload (we don't have a usable one here). - const message = err instanceof Error ? err.message : String(err); - controller.enqueue( - SSE_ENCODER.encode( - `data: ${JSON.stringify({ type: "error", reason: "error", errorMessage: message })}\n\n`, - ), - ); - controller.enqueue(SSE_DONE); - controller.close(); + if (!cancelled) { + // Best-effort error envelope so the client iterator resolves + // instead of hanging on the dropped connection. Shape matches the + // canonical `error` event minus the unrecoverable `error: + // AssistantMessage` payload (we don't have a usable one here). + const message = err instanceof Error ? err.message : String(err); + controller.enqueue( + SSE_ENCODER.encode( + `data: ${JSON.stringify({ type: "error", reason: "error", errorMessage: message })}\n\n`, + ), + ); + controller.enqueue(SSE_DONE); + controller.close(); + } + } finally { + control?.signal?.removeEventListener("abort", markCancelled); } }, + cancel(reason) { + cancelled = true; + control?.signal?.removeEventListener("abort", markCancelled); + control?.onCancel?.(reason); + }, }); } diff --git a/packages/ai/test/auth-gateway-openai-chat.test.ts b/packages/ai/test/auth-gateway-openai-chat.test.ts index 9ce5a92be..80bb27eb2 100644 --- a/packages/ai/test/auth-gateway-openai-chat.test.ts +++ b/packages/ai/test/auth-gateway-openai-chat.test.ts @@ -343,4 +343,30 @@ describe("auth-gateway openai-chat: encodeStream", () => { const payloads = lines.map(parseSseLine) as Array>; expect(payloads[1]).toEqual({ error: { message: "upstream went away", type: "upstream_error" } }); }); + + it("aborts the upstream gateway request when the client cancels the response body", async () => { + const aborted: unknown[] = []; + async function* neverEndingEvents() { + await new Promise(() => {}); + } + const events = neverEndingEvents() as unknown as AssistantMessageEventStream; + (events as { result(): Promise }).result = async () => emptyAssistant(); + const requestController = new AbortController(); + const stream = encodeStream(events, "gpt-test", undefined, { + signal: requestController.signal, + onCancel(reason) { + aborted.push(reason); + requestController.abort(reason); + }, + }); + const reader = stream.getReader(); + + const firstChunk = await reader.read(); + expect(firstChunk.done).toBe(false); + await reader.cancel("client timeout"); + + expect(aborted).toEqual(["client timeout"]); + expect(requestController.signal.aborted).toBe(true); + expect(requestController.signal.reason).toBe("client timeout"); + }); }); diff --git a/packages/coding-agent/src/tools/browser/cmux/cmux-tab.ts b/packages/coding-agent/src/tools/browser/cmux/cmux-tab.ts index f5b77d7d8..c0564a22f 100644 --- a/packages/coding-agent/src/tools/browser/cmux/cmux-tab.ts +++ b/packages/coding-agent/src/tools/browser/cmux/cmux-tab.ts @@ -94,7 +94,7 @@ interface ViewportOptions { deviceScaleFactor?: number; } -const PAGE_SELECTOR_HELPERS = String.raw` +const PAGE_SELECTOR_HELPERS = ` const isVisible = element => { const style = getComputedStyle(element); const rect = element.getBoundingClientRect();