fix(ai): add completions session affinity compat closes #3430

This commit is contained in:
Mario Zechner
2026-04-20 16:55:16 +02:00
parent 2f4f283cc2
commit 4b2caf43a8
4 changed files with 107 additions and 38 deletions

View File

@@ -851,6 +851,7 @@ interface OpenAICompletionsCompat {
supportsReasoningEffort?: boolean; // Whether provider supports `reasoning_effort` (default: true) supportsReasoningEffort?: boolean; // Whether provider supports `reasoning_effort` (default: true)
supportsUsageInStreaming?: boolean; // Whether provider supports `stream_options: { include_usage: true }` (default: true) supportsUsageInStreaming?: boolean; // Whether provider supports `stream_options: { include_usage: true }` (default: true)
supportsStrictMode?: boolean; // Whether provider supports `strict` in tool definitions (default: true) supportsStrictMode?: boolean; // Whether provider supports `strict` in tool definitions (default: true)
sendSessionAffinityHeaders?: boolean; // Whether to send `session_id`, `x-client-request-id`, and `x-session-affinity` from `sessionId` when caching is enabled (default: false)
maxTokensField?: 'max_completion_tokens' | 'max_tokens'; // Which field name to use (default: max_completion_tokens) maxTokensField?: 'max_completion_tokens' | 'max_tokens'; // Which field name to use (default: max_completion_tokens)
requiresToolResultName?: boolean; // Whether tool results require the `name` field (default: false) requiresToolResultName?: boolean; // Whether tool results require the `name` field (default: false)
requiresAssistantAfterToolResult?: boolean; // Whether tool results must be followed by an assistant message (default: false) requiresAssistantAfterToolResult?: boolean; // Whether tool results must be followed by an assistant message (default: false)

View File

@@ -291,6 +291,8 @@ export interface OpenAICompletionsCompat {
zaiToolStream?: boolean; zaiToolStream?: boolean;
/** Whether the provider supports the `strict` field in tool definitions. Default: true. */ /** Whether the provider supports the `strict` field in tool definitions. Default: true. */
supportsStrictMode?: boolean; supportsStrictMode?: boolean;
/** Whether to send known session-affinity headers (`session_id`, `x-client-request-id`, `x-session-affinity`) from `options.sessionId` when caching is enabled. Default: false. */
sendSessionAffinityHeaders?: boolean;
} }
/** Compatibility settings for OpenAI Responses APIs. */ /** Compatibility settings for OpenAI Responses APIs. */

View File

@@ -1,16 +1,30 @@
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import { getModel } from "../src/models.js"; import { getModel } from "../src/models.js";
import { streamOpenAICompletions } from "../src/providers/openai-completions.js"; import { streamOpenAICompletions } from "../src/providers/openai-completions.js";
import type { Model } from "../src/types.js";
interface FakeOpenAIClientOptions {
apiKey: string;
baseURL: string;
dangerouslyAllowBrowser: boolean;
defaultHeaders?: Record<string, string>;
}
interface CapturedCompletionsPayload {
prompt_cache_key?: string;
prompt_cache_retention?: "24h" | "in-memory" | null;
}
const mockState = vi.hoisted(() => ({ const mockState = vi.hoisted(() => ({
lastParams: undefined as unknown, lastParams: undefined as CapturedCompletionsPayload | undefined,
lastClientOptions: undefined as FakeOpenAIClientOptions | undefined,
})); }));
vi.mock("openai", () => { vi.mock("openai", () => {
class FakeOpenAI { class FakeOpenAI {
chat = { chat = {
completions: { completions: {
create: (params: unknown) => { create: (params: CapturedCompletionsPayload) => {
mockState.lastParams = params; mockState.lastParams = params;
const stream = { const stream = {
async *[Symbol.asyncIterator]() { async *[Symbol.asyncIterator]() {
@@ -39,6 +53,10 @@ vi.mock("openai", () => {
}, },
}, },
}; };
constructor(options: FakeOpenAIClientOptions) {
mockState.lastClientOptions = options;
}
} }
return { default: FakeOpenAI }; return { default: FakeOpenAI };
@@ -49,6 +67,7 @@ describe("openai-completions prompt caching", () => {
beforeEach(() => { beforeEach(() => {
mockState.lastParams = undefined; mockState.lastParams = undefined;
mockState.lastClientOptions = undefined;
delete process.env.PI_CACHE_RETENTION; delete process.env.PI_CACHE_RETENTION;
}); });
@@ -60,10 +79,23 @@ describe("openai-completions prompt caching", () => {
} }
}); });
async function capturePayload(options?: { cacheRetention?: "none" | "short" | "long"; sessionId?: string }) { function createModel(overrides: Partial<Model<"openai-completions">> = {}): Model<"openai-completions"> {
const { compat: _compat, ...baseModel } = getModel("openai", "gpt-4o-mini"); const { compat: _compat, ...baseModel } = getModel("openai", "gpt-4o-mini");
const model = { ...baseModel, api: "openai-completions" } as const; return {
...(baseModel as Omit<Model<"openai-completions">, "api">),
api: "openai-completions",
...overrides,
};
}
async function captureRequest(
options?: {
cacheRetention?: "none" | "short" | "long";
sessionId?: string;
headers?: Record<string, string>;
},
model: Model<"openai-completions"> = createModel(),
) {
await streamOpenAICompletions( await streamOpenAICompletions(
model, model,
{ {
@@ -73,61 +105,94 @@ describe("openai-completions prompt caching", () => {
{ apiKey: "test-key", ...options }, { apiKey: "test-key", ...options },
).result(); ).result();
return mockState.lastParams as { prompt_cache_key?: string; prompt_cache_retention?: "24h" | "in-memory" | null }; return {
payload: mockState.lastParams,
headers: mockState.lastClientOptions?.defaultHeaders ?? {},
};
} }
it("sets prompt_cache_key for direct OpenAI requests when caching is enabled", async () => { it("sets prompt_cache_key for direct OpenAI requests when caching is enabled", async () => {
const payload = await capturePayload({ sessionId: "session-123" }); const { payload } = await captureRequest({ sessionId: "session-123" });
expect(payload.prompt_cache_key).toBe("session-123"); expect(payload?.prompt_cache_key).toBe("session-123");
expect(payload.prompt_cache_retention).toBeUndefined(); expect(payload?.prompt_cache_retention).toBeUndefined();
}); });
it("sets prompt_cache_retention to 24h for direct OpenAI requests when cacheRetention is long", async () => { it("sets prompt_cache_retention to 24h for direct OpenAI requests when cacheRetention is long", async () => {
const payload = await capturePayload({ cacheRetention: "long", sessionId: "session-456" }); const { payload } = await captureRequest({ cacheRetention: "long", sessionId: "session-456" });
expect(payload.prompt_cache_key).toBe("session-456"); expect(payload?.prompt_cache_key).toBe("session-456");
expect(payload.prompt_cache_retention).toBe("24h"); expect(payload?.prompt_cache_retention).toBe("24h");
}); });
it("omits prompt cache fields when cacheRetention is none", async () => { it("omits prompt cache fields when cacheRetention is none", async () => {
const payload = await capturePayload({ cacheRetention: "none", sessionId: "session-789" }); const { payload } = await captureRequest({ cacheRetention: "none", sessionId: "session-789" });
expect(payload.prompt_cache_key).toBeUndefined(); expect(payload?.prompt_cache_key).toBeUndefined();
expect(payload.prompt_cache_retention).toBeUndefined(); expect(payload?.prompt_cache_retention).toBeUndefined();
}); });
it("omits prompt cache fields for non-OpenAI base URLs", async () => { it("omits prompt cache fields for non-OpenAI base URLs", async () => {
const { compat: _compat, ...baseModel } = getModel("openai", "gpt-4o-mini"); const model = createModel({
const model = {
...baseModel,
api: "openai-completions",
baseUrl: "https://proxy.example.com/v1", baseUrl: "https://proxy.example.com/v1",
} as const; });
const { payload } = await captureRequest({ cacheRetention: "long", sessionId: "session-proxy" }, model);
await streamOpenAICompletions( expect(payload?.prompt_cache_key).toBeUndefined();
model, expect(payload?.prompt_cache_retention).toBeUndefined();
{
systemPrompt: "sys",
messages: [{ role: "user", content: "hi", timestamp: Date.now() }],
},
{ apiKey: "test-key", cacheRetention: "long", sessionId: "session-proxy" },
).result();
const payload = mockState.lastParams as {
prompt_cache_key?: string;
prompt_cache_retention?: "24h" | "in-memory" | null;
};
expect(payload.prompt_cache_key).toBeUndefined();
expect(payload.prompt_cache_retention).toBeUndefined();
}); });
it("uses PI_CACHE_RETENTION for direct OpenAI requests", async () => { it("uses PI_CACHE_RETENTION for direct OpenAI requests", async () => {
process.env.PI_CACHE_RETENTION = "long"; process.env.PI_CACHE_RETENTION = "long";
const payload = await capturePayload({ sessionId: "session-env" }); const { payload } = await captureRequest({ sessionId: "session-env" });
expect(payload.prompt_cache_key).toBe("session-env"); expect(payload?.prompt_cache_key).toBe("session-env");
expect(payload.prompt_cache_retention).toBe("24h"); expect(payload?.prompt_cache_retention).toBe("24h");
});
it("sends known session-affinity headers when compat.sendSessionAffinityHeaders is enabled", async () => {
const model = createModel({
baseUrl: "https://proxy.example.com/v1",
compat: { sendSessionAffinityHeaders: true },
});
const { headers } = await captureRequest({ sessionId: "session-affinity" }, model);
expect(headers.session_id).toBe("session-affinity");
expect(headers["x-client-request-id"]).toBe("session-affinity");
expect(headers["x-session-affinity"]).toBe("session-affinity");
});
it("omits session-affinity headers when cacheRetention is none", async () => {
const model = createModel({
baseUrl: "https://proxy.example.com/v1",
compat: { sendSessionAffinityHeaders: true },
});
const { headers } = await captureRequest({ cacheRetention: "none", sessionId: "session-affinity" }, model);
expect(headers.session_id).toBeUndefined();
expect(headers["x-client-request-id"]).toBeUndefined();
expect(headers["x-session-affinity"]).toBeUndefined();
});
it("lets explicit headers override generated session-affinity headers", async () => {
const model = createModel({
baseUrl: "https://proxy.example.com/v1",
compat: { sendSessionAffinityHeaders: true },
});
const { headers } = await captureRequest(
{
sessionId: "session-affinity",
headers: {
session_id: "override-session",
"x-client-request-id": "override-request",
"x-session-affinity": "override-affinity",
},
},
model,
);
expect(headers.session_id).toBe("override-session");
expect(headers["x-client-request-id"]).toBe("override-request");
expect(headers["x-session-affinity"]).toBe("override-affinity");
}); });
}); });

View File

@@ -34,6 +34,7 @@ const compat: Required<OpenAICompletionsCompat> = {
vercelGatewayRouting: {}, vercelGatewayRouting: {},
zaiToolStream: false, zaiToolStream: false,
supportsStrictMode: true, supportsStrictMode: true,
sendSessionAffinityHeaders: false,
}; };
function buildToolResult(toolCallId: string, timestamp: number): ToolResultMessage { function buildToolResult(toolCallId: string, timestamp: number): ToolResultMessage {