Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion apps/desktop/src/ai/hooks/useLLMConnection.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import type { AIProviderStorage } from "@anlg/store";

import { createAppleFoundationModel } from "../apple-foundation-model";
import { createAuthFetch } from "../auth-fetch";
import { streamOnlyGenerationMiddleware } from "../stream-only-generation";
import { createTracedFetch, tracedFetch } from "../traced-fetch";

import { useAuth } from "~/auth";
Expand Down Expand Up @@ -291,7 +292,15 @@ const createLanguageModel = (
baseURL: oauth ? CHATGPT_API_BASE_URL : conn.baseUrl,
apiKey: oauth ? "oauth" : conn.apiKey,
});
return wrapWithThinkingMiddleware(provider.responses(conn.modelId));
const model = provider.responses(conn.modelId);
return wrapWithThinkingMiddleware(
oauth
? wrapLanguageModel({
model,
middleware: streamOnlyGenerationMiddleware,
})
: model,
);
}

case "grok":
Expand Down
118 changes: 118 additions & 0 deletions apps/desktop/src/ai/stream-only-generation.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
import { wrapLanguageModel } from "ai";
import { describe, expect, test, vi } from "vitest";

import { streamOnlyGenerationMiddleware } from "./stream-only-generation";

type LanguageModel = Parameters<typeof wrapLanguageModel>[0]["model"];
type StreamPart =
Awaited<
ReturnType<LanguageModel["doStream"]>
>["stream"] extends ReadableStream<infer Part>
? Part
: never;

const usage = {
inputTokens: {
total: 4,
noCache: 4,
cacheRead: 0,
cacheWrite: 0,
},
outputTokens: { total: 2, text: 2, reasoning: 0 },
};

describe("streamOnlyGenerationMiddleware", () => {
test("uses streaming for generate calls and collects the result", async () => {
const doGenerate = vi.fn(async () => {
throw new Error("non-streaming request used");
});
const model = createModel(doGenerate, [
{ type: "stream-start", warnings: [] },
{
type: "response-metadata",
id: "response-1",
modelId: "gpt-test",
timestamp: new Date("2026-08-25T00:00:00.000Z"),
},
{ type: "text-start", id: "message-1" },
{ type: "text-delta", id: "message-1", delta: "Hello" },
{ type: "text-delta", id: "message-1", delta: " world" },
{ type: "text-end", id: "message-1" },
{
type: "tool-call",
toolCallId: "tool-1",
toolName: "lookup",
input: '{"query":"test"}',
},
{
type: "finish",
finishReason: { unified: "stop", raw: "completed" },
usage,
},
]);
const wrapped = wrapLanguageModel({
model,
middleware: streamOnlyGenerationMiddleware,
});

const result = await wrapped.doGenerate({ prompt: [] });

expect(doGenerate).not.toHaveBeenCalled();
expect(result.content).toEqual([
{ type: "text", text: "Hello world" },
{
type: "tool-call",
toolCallId: "tool-1",
toolName: "lookup",
input: '{"query":"test"}',
},
]);
expect(result.finishReason).toEqual({ unified: "stop", raw: "completed" });
expect(result.usage).toEqual(usage);
expect(result.response).toMatchObject({
headers: { "x-request-id": "request-1" },
id: "response-1",
modelId: "gpt-test",
});
});

test("rejects provider errors from the response stream", async () => {
const model = createModel(vi.fn(), [
{ type: "stream-start", warnings: [] },
{ type: "error", error: new Error("upstream failed") },
]);
const wrapped = wrapLanguageModel({
model,
middleware: streamOnlyGenerationMiddleware,
});

await expect(wrapped.doGenerate({ prompt: [] })).rejects.toThrow(
"upstream failed",
);
});
});

function createModel(
doGenerate: LanguageModel["doGenerate"],
parts: StreamPart[],
): LanguageModel {
return {
specificationVersion: "v3",
provider: "test",
modelId: "gpt-test",
supportedUrls: {},
doGenerate,
doStream: async () => ({
stream: new ReadableStream<StreamPart>({
start(controller) {
for (const part of parts) {
controller.enqueue(part);
}
controller.close();
},
}),
request: { body: { stream: true } },
response: { headers: { "x-request-id": "request-1" } },
}),
};
}
158 changes: 158 additions & 0 deletions apps/desktop/src/ai/stream-only-generation.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
import type { LanguageModelMiddleware } from "ai";

type WrapGenerate = NonNullable<LanguageModelMiddleware["wrapGenerate"]>;
type GenerateResult = Awaited<
ReturnType<Parameters<WrapGenerate>[0]["doGenerate"]>
>;
type StreamResult = Awaited<
ReturnType<Parameters<WrapGenerate>[0]["doStream"]>
>;
type StreamPart =
StreamResult["stream"] extends ReadableStream<infer Part> ? Part : never;
type Content = GenerateResult["content"][number];
type TextContent = Extract<Content, { type: "reasoning" | "text" }>;

export const streamOnlyGenerationMiddleware: LanguageModelMiddleware = {
specificationVersion: "v3",
wrapGenerate: async ({ doStream }) => collectStream(await doStream()),
};

async function collectStream(result: StreamResult): Promise<GenerateResult> {
const content: Content[] = [];
const openBlocks = new Map<string, TextContent>();
let warnings: GenerateResult["warnings"] = [];
let responseMetadata: NonNullable<GenerateResult["response"]> = {};
let finish: Extract<StreamPart, { type: "finish" }> | undefined;

const reader = result.stream.getReader();
try {
while (true) {
const { done, value } = await reader.read();
if (done) break;

switch (value.type) {
case "text-start":
case "reasoning-start": {
createBlock(value, content, openBlocks);
break;
}
case "text-delta":
case "reasoning-delta": {
const block = createBlock(value, content, openBlocks);
block.text += value.delta;
if (value.providerMetadata) {
block.providerMetadata = value.providerMetadata;
}
break;
}
case "text-end":
case "reasoning-end": {
const block = createBlock(value, content, openBlocks);
if (value.providerMetadata) {
block.providerMetadata = value.providerMetadata;
}
openBlocks.delete(blockKey(value));
break;
}
case "file":
case "source":
case "tool-approval-request":
case "tool-call":
case "tool-result":
content.push(value);
break;
case "stream-start":
warnings = value.warnings;
break;
case "response-metadata":
responseMetadata = {
...responseMetadata,
id: value.id,
timestamp: value.timestamp,
modelId: value.modelId,
};
break;
case "finish":
finish = value;
break;
case "error":
throw value.error;
case "raw":
case "tool-input-start":
case "tool-input-delta":
case "tool-input-end":
break;
}
}
} finally {
reader.releaseLock();
}

if (!finish) {
throw new Error("ChatGPT response stream ended without a finish event");
}

const hasResponseMetadata =
result.response !== undefined ||
responseMetadata.id !== undefined ||
responseMetadata.timestamp !== undefined ||
responseMetadata.modelId !== undefined;

return {
content,
finishReason: finish.finishReason,
usage: finish.usage,
providerMetadata: finish.providerMetadata,
request: result.request,
response: hasResponseMetadata
? { ...responseMetadata, ...result.response }
: undefined,
warnings,
};
}

function createBlock(
part: Extract<
StreamPart,
{
type:
| "reasoning-delta"
| "reasoning-end"
| "reasoning-start"
| "text-delta"
| "text-end"
| "text-start";
}
>,
content: Content[],
openBlocks: Map<string, TextContent>,
): TextContent {
const key = blockKey(part);
const existing = openBlocks.get(key);
if (existing) {
return existing;
}

const block: TextContent = {
type: part.type.startsWith("reasoning") ? "reasoning" : "text",
text: "",
providerMetadata: part.providerMetadata,
};
openBlocks.set(key, block);
content.push(block);
return block;
}

function blockKey(part: {
id: string;
type:
| "reasoning-delta"
| "reasoning-end"
| "reasoning-start"
| "text-delta"
| "text-end"
| "text-start";
}): string {
const type = part.type.startsWith("reasoning") ? "reasoning" : "text";
return `${type}:${part.id}`;
}
74 changes: 74 additions & 0 deletions apps/desktop/src/settings/ai/llm/subscriptions/models.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
import { Effect } from "effect";
import { beforeEach, describe, expect, test, vi } from "vitest";

const mocks = vi.hoisted(() => ({
fetchJson: vi.fn(),
resolveSubscriptionAccess: vi.fn(),
}));

vi.mock("./access", () => ({
resolveSubscriptionAccess: mocks.resolveSubscriptionAccess,
}));

vi.mock("~/settings/ai/shared/list-common", async (importOriginal) => ({
...(await importOriginal()),
fetchJson: mocks.fetchJson,
}));

import { listSubscriptionModels } from "./models";

describe("ChatGPT subscription models", () => {
beforeEach(() => {
mocks.fetchJson.mockReset();
mocks.resolveSubscriptionAccess.mockReset();
mocks.resolveSubscriptionAccess.mockResolvedValue({
token: "access-token",
credential: { accountId: "account-1" },
});
});

test("parses the Codex catalog and omits hidden models", async () => {
mocks.fetchJson.mockReturnValue(
Effect.succeed({
models: [
{ slug: "gpt-5.6-sol", visibility: "list" },
{ slug: "codex-auto-review", visibility: "hide" },
{
slug: "gpt-5.3-codex-spark",
visibility: "list",
supported_in_api: false,
},
],
}),
);

await expect(
listSubscriptionModels(
"chatgpt",
"https://api.openai.com/v1",
"stored-credential",
),
).resolves.toMatchObject({
models: ["gpt-5.6-sol", "gpt-5.3-codex-spark"],
});
expect(mocks.fetchJson).toHaveBeenCalledWith(
"https://chatgpt.com/backend-api/codex/models?client_version=0.145.0",
expect.objectContaining({
Authorization: "Bearer access-token",
"ChatGPT-Account-ID": "account-1",
}),
);
});

test("does not offer stale fallback models when discovery fails", async () => {
mocks.fetchJson.mockReturnValue(Effect.fail(new Error("unavailable")));

await expect(
listSubscriptionModels(
"chatgpt",
"https://api.openai.com/v1",
"stored-credential",
),
).resolves.toEqual({ models: [], ignored: [], metadata: {} });
});
});
Loading
Loading