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.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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: {
|
||||
|
||||
@@ -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<string, unknown>;
|
||||
@@ -117,6 +124,7 @@ export interface AuthGatewayFormatModule {
|
||||
events: AssistantMessageEventStream,
|
||||
requestedModelId: string,
|
||||
options?: AuthGatewayParsedRequestOptions,
|
||||
control?: AuthGatewayStreamControl,
|
||||
): ReadableStream<Uint8Array>;
|
||||
/**
|
||||
* Emit a protocol-specific error envelope. OpenAI returns
|
||||
|
||||
@@ -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<string, unknown> = {
|
||||
export function encodeStream(
|
||||
events: AssistantMessageEventStream,
|
||||
requestedModelId: string,
|
||||
_options?: ParsedRequest["options"],
|
||||
control?: AuthGatewayStreamControl,
|
||||
): ReadableStream<Uint8Array> {
|
||||
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();
|
||||
},
|
||||
});
|
||||
|
||||
@@ -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<Uint8Array> {
|
||||
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<string, unknown>, finishReason: string | null) => ({
|
||||
id,
|
||||
@@ -518,7 +524,7 @@ export function encodeStream(
|
||||
});
|
||||
|
||||
const writeSse = (controller: ReadableStreamDefaultController<Uint8Array>, 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<Uint8Array>, 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);
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
@@ -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<Uint8Array> {
|
||||
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<Uint8Array>({
|
||||
async start(controller) {
|
||||
const emit = (name: string, data: Record<string, unknown>) => {
|
||||
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);
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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<Uint8Array> {
|
||||
export function encodeStream(
|
||||
events: AssistantMessageEventStream,
|
||||
_requestedModelId?: string,
|
||||
_options?: AuthGatewayParsedRequestOptions,
|
||||
control?: AuthGatewayStreamControl,
|
||||
): ReadableStream<Uint8Array> {
|
||||
let cancelled = control?.signal?.aborted === true;
|
||||
const markCancelled = () => {
|
||||
cancelled = true;
|
||||
};
|
||||
control?.signal?.addEventListener("abort", markCancelled, { once: true });
|
||||
return new ReadableStream<Uint8Array>({
|
||||
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);
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
@@ -343,4 +343,30 @@ describe("auth-gateway openai-chat: encodeStream", () => {
|
||||
const payloads = lines.map(parseSseLine) as Array<Record<string, unknown>>;
|
||||
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<AssistantMessage> }).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");
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user