refactor(ai): extracted OpenAI compat logic into dedicated module

- Extracted OpenAI compatibility detection and resolution logic into dedicated `openai-completions-compat` module.
- Refactored `detectCompat()` and `getCompat()` to delegate to new compat module functions with simplified conditional logic.
- Fixed OAuth redirect URI validation to preserve exact configured values without trailing slash normalization.
- Improved session deletion to return boolean status and display error messages in UI instead of silently failing.
- Added `/session delete` command with Delete key support and confirmation dialogs for session management.
This commit is contained in:
can1357
2026-03-17 15:39:59 +01:00
parent 0c1586f747
commit afc2276a12
17 changed files with 403 additions and 214 deletions
+3
View File
@@ -1,6 +1,9 @@
# Changelog
## [Unreleased]
### Changed
- Extracted OpenAI compatibility detection and resolution logic into dedicated `openai-completions-compat` module for improved maintainability and reusability
### Fixed
+3 -3
View File
@@ -1,3 +1,4 @@
import { resolveOpenAICompat } from "./providers/openai-completions-compat";
import type { Api, Model as ApiModel, ThinkingConfig } from "./types";
/** User-facing thinking levels, ordered least to most intensive. */
@@ -387,9 +388,8 @@ function inferFallbackEfforts<TApi extends Api>(model: ApiModel<TApi>): readonly
return DEFAULT_REASONING_EFFORTS;
}
if (model.api === "openai-completions") {
const compat = model.compat;
const usesDiscreteEfforts = compat?.thinkingFormat === undefined || compat.thinkingFormat === "openai";
if (usesDiscreteEfforts && compat?.supportsReasoningEffort !== false) {
const compat = resolveOpenAICompat(model as ApiModel<"openai-completions">);
if (compat.thinkingFormat === "openai" && compat.supportsReasoningEffort) {
return DEFAULT_REASONING_EFFORTS_WITH_XHIGH;
}
return DEFAULT_REASONING_EFFORTS;
@@ -0,0 +1,137 @@
import type { Model, OpenAICompat } from "../types";
type OpenAIReasoningEffort = "minimal" | "low" | "medium" | "high" | "xhigh";
export type ResolvedOpenAICompat = Required<
Omit<OpenAICompat, "openRouterRouting" | "vercelGatewayRouting" | "extraBody">
> & {
openRouterRouting?: OpenAICompat["openRouterRouting"];
vercelGatewayRouting?: OpenAICompat["vercelGatewayRouting"];
extraBody?: OpenAICompat["extraBody"];
};
function detectStrictModeSupport(provider: string, baseUrl: string): boolean {
if (
provider === "openai" ||
provider === "cerebras" ||
provider === "together" ||
provider === "github-copilot" ||
provider === "zenmux"
) {
return true;
}
const normalizedBaseUrl = baseUrl.toLowerCase();
return (
normalizedBaseUrl.includes("api.openai.com") ||
normalizedBaseUrl.includes(".openai.azure.com") ||
normalizedBaseUrl.includes("models.inference.ai.azure.com") ||
normalizedBaseUrl.includes("api.cerebras.ai") ||
normalizedBaseUrl.includes("api.together.xyz") ||
normalizedBaseUrl.includes("api.deepseek.com") ||
normalizedBaseUrl.includes("deepseek.com")
);
}
/**
* Detect compatibility settings from provider and baseUrl for known providers.
* Provider takes precedence over URL-based detection since it's explicitly configured.
*/
export function detectOpenAICompat(model: Model<"openai-completions">): ResolvedOpenAICompat {
const provider = model.provider;
const baseUrl = model.baseUrl;
const isCerebras = provider === "cerebras" || baseUrl.includes("cerebras.ai");
const isZai = provider === "zai" || baseUrl.includes("api.z.ai");
const isOpenRouterKimi = provider === "openrouter" && model.id.includes("moonshotai/kimi");
const isAlibaba = provider === "alibaba-coding-plan" || baseUrl.includes("dashscope");
const isQwen = model.id.toLowerCase().includes("qwen");
const isNonStandard =
isCerebras ||
provider === "xai" ||
baseUrl.includes("api.x.ai") ||
provider === "mistral" ||
baseUrl.includes("mistral.ai") ||
baseUrl.includes("chutes.ai") ||
baseUrl.includes("deepseek.com") ||
isAlibaba ||
isZai ||
isQwen ||
provider === "opencode-zen" ||
provider === "opencode-go" ||
baseUrl.includes("opencode.ai");
const useMaxTokens = provider === "mistral" || baseUrl.includes("mistral.ai") || baseUrl.includes("chutes.ai");
const isGrok = provider === "xai" || baseUrl.includes("api.x.ai");
const isMistral = provider === "mistral" || baseUrl.includes("mistral.ai");
const reasoningEffortMap: NonNullable<OpenAICompat["reasoningEffortMap"]> =
provider === "groq" && model.id === "qwen/qwen3-32b"
? ({
minimal: "default",
low: "default",
medium: "default",
high: "default",
xhigh: "default",
} satisfies Partial<Record<OpenAIReasoningEffort, string>>)
: {};
return {
supportsStore: !isNonStandard,
supportsDeveloperRole: !isNonStandard,
supportsReasoningEffort: !isGrok && !isZai,
reasoningEffortMap,
supportsUsageInStreaming: !isCerebras,
supportsToolChoice: true,
maxTokensField: useMaxTokens ? "max_tokens" : "max_completion_tokens",
requiresToolResultName: isMistral,
requiresAssistantAfterToolResult: false,
requiresThinkingAsText: isMistral,
requiresMistralToolIds: isMistral,
thinkingFormat: isZai ? "zai" : isAlibaba || isQwen ? "qwen" : "openai",
reasoningContentField: "reasoning_content",
requiresReasoningContentForToolCalls: isOpenRouterKimi,
requiresAssistantContentForToolCalls: isOpenRouterKimi,
openRouterRouting: undefined,
vercelGatewayRouting: undefined,
supportsStrictMode: detectStrictModeSupport(provider, baseUrl),
extraBody: undefined,
};
}
/**
* Resolve compatibility settings by layering explicit model.compat overrides onto
* the detected defaults. This is the canonical compat view for both metadata and transport.
*/
export function resolveOpenAICompat(model: Model<"openai-completions">): ResolvedOpenAICompat {
const detected = detectOpenAICompat(model);
if (!model.compat) {
return detected;
}
return {
supportsStore: model.compat.supportsStore ?? detected.supportsStore,
supportsDeveloperRole: model.compat.supportsDeveloperRole ?? detected.supportsDeveloperRole,
supportsReasoningEffort: model.compat.supportsReasoningEffort ?? detected.supportsReasoningEffort,
reasoningEffortMap: model.compat.reasoningEffortMap ?? detected.reasoningEffortMap,
supportsUsageInStreaming: model.compat.supportsUsageInStreaming ?? detected.supportsUsageInStreaming,
supportsToolChoice: model.compat.supportsToolChoice ?? detected.supportsToolChoice,
maxTokensField: model.compat.maxTokensField ?? detected.maxTokensField,
requiresToolResultName: model.compat.requiresToolResultName ?? detected.requiresToolResultName,
requiresAssistantAfterToolResult:
model.compat.requiresAssistantAfterToolResult ?? detected.requiresAssistantAfterToolResult,
requiresThinkingAsText: model.compat.requiresThinkingAsText ?? detected.requiresThinkingAsText,
requiresMistralToolIds: model.compat.requiresMistralToolIds ?? detected.requiresMistralToolIds,
thinkingFormat: model.compat.thinkingFormat ?? detected.thinkingFormat,
reasoningContentField: model.compat.reasoningContentField ?? detected.reasoningContentField,
requiresReasoningContentForToolCalls:
model.compat.requiresReasoningContentForToolCalls ?? detected.requiresReasoningContentForToolCalls,
requiresAssistantContentForToolCalls:
model.compat.requiresAssistantContentForToolCalls ?? detected.requiresAssistantContentForToolCalls,
openRouterRouting: model.compat.openRouterRouting ?? detected.openRouterRouting,
vercelGatewayRouting: model.compat.vercelGatewayRouting ?? detected.vercelGatewayRouting,
supportsStrictMode: model.compat.supportsStrictMode ?? detected.supportsStrictMode,
extraBody: model.compat.extraBody,
};
}
+3 -117
View File
@@ -18,7 +18,6 @@ import {
type Message,
type MessageAttribution,
type Model,
type OpenAICompat,
type ServiceTier,
type StopReason,
type StreamFunction,
@@ -42,6 +41,7 @@ import {
hasCopilotVisionInput,
resolveGitHubCopilotBaseUrl,
} from "./github-copilot-headers";
import { detectOpenAICompat, type ResolvedOpenAICompat, resolveOpenAICompat } from "./openai-completions-compat";
import { transformMessages } from "./transform-messages";
/**
@@ -87,12 +87,6 @@ function serializeToolArguments(value: unknown): string {
return "{}";
}
type ResolvedOpenAICompat = Required<Omit<OpenAICompat, "openRouterRouting" | "vercelGatewayRouting" | "extraBody">> & {
openRouterRouting?: OpenAICompat["openRouterRouting"];
vercelGatewayRouting?: OpenAICompat["vercelGatewayRouting"];
extraBody?: OpenAICompat["extraBody"];
};
/**
* Check if conversation messages contain tool calls or tool results.
* This is needed because Anthropic (via proxy) requires the tools param
@@ -1085,95 +1079,13 @@ function mapStopReason(reason: ChatCompletionChunk.Choice["finish_reason"] | str
}
}
function detectStrictModeSupport(provider: string, baseUrl: string): boolean {
if (
provider === "openai" ||
provider === "cerebras" ||
provider === "together" ||
provider === "github-copilot" ||
provider === "zenmux"
)
return true;
const normalizedBaseUrl = baseUrl.toLowerCase();
return (
normalizedBaseUrl.includes("api.openai.com") ||
normalizedBaseUrl.includes(".openai.azure.com") ||
normalizedBaseUrl.includes("models.inference.ai.azure.com") ||
normalizedBaseUrl.includes("api.cerebras.ai") ||
normalizedBaseUrl.includes("api.together.xyz") ||
normalizedBaseUrl.includes("api.deepseek.com") ||
normalizedBaseUrl.includes("deepseek.com")
);
}
/**
* Detect compatibility settings from provider and baseUrl for known providers.
* Provider takes precedence over URL-based detection since it's explicitly configured.
* Returns a fully resolved OpenAICompat object with all fields set.
*/
export function detectCompat(model: Model<"openai-completions">): ResolvedOpenAICompat {
const provider = model.provider;
const baseUrl = model.baseUrl;
const isCerebras = provider === "cerebras" || baseUrl.includes("cerebras.ai");
const isZai = provider === "zai" || baseUrl.includes("api.z.ai");
const isOpenRouterKimi = provider === "openrouter" && model.id.includes("moonshotai/kimi");
const isAlibaba = provider === "alibaba-coding-plan" || baseUrl.includes("dashscope");
const isQwen = model.id.toLowerCase().includes("qwen");
const isNonStandard =
isCerebras ||
provider === "xai" ||
baseUrl.includes("api.x.ai") ||
provider === "mistral" ||
baseUrl.includes("mistral.ai") ||
baseUrl.includes("chutes.ai") ||
baseUrl.includes("deepseek.com") ||
isAlibaba ||
isZai ||
isQwen ||
provider === "opencode-zen" ||
provider === "opencode-go" ||
baseUrl.includes("opencode.ai");
const useMaxTokens = provider === "mistral" || baseUrl.includes("mistral.ai") || baseUrl.includes("chutes.ai");
const isGrok = provider === "xai" || baseUrl.includes("api.x.ai");
const isMistral = provider === "mistral" || baseUrl.includes("mistral.ai");
const reasoningEffortMap =
provider === "groq" && model.id === "qwen/qwen3-32b"
? ({
minimal: "default",
low: "default",
medium: "default",
high: "default",
xhigh: "default",
} satisfies Partial<Record<NonNullable<OpenAICompletionsOptions["reasoning"]>, string>>)
: {};
return {
supportsStore: !isNonStandard,
supportsDeveloperRole: !isNonStandard,
supportsReasoningEffort: !isGrok && !isZai,
reasoningEffortMap,
supportsUsageInStreaming: !isCerebras,
supportsToolChoice: true,
maxTokensField: useMaxTokens ? "max_tokens" : "max_completion_tokens",
requiresToolResultName: isMistral,
requiresAssistantAfterToolResult: false, // Mistral no longer requires this as of Dec 2024
requiresThinkingAsText: isMistral,
requiresMistralToolIds: isMistral,
thinkingFormat: isZai ? "zai" : isAlibaba || isQwen ? "qwen" : "openai",
reasoningContentField: "reasoning_content",
requiresReasoningContentForToolCalls: isOpenRouterKimi,
requiresAssistantContentForToolCalls: isOpenRouterKimi,
openRouterRouting: undefined,
vercelGatewayRouting: undefined,
supportsStrictMode: detectStrictModeSupport(provider, baseUrl),
};
return detectOpenAICompat(model);
}
/**
@@ -1181,31 +1093,5 @@ export function detectCompat(model: Model<"openai-completions">): ResolvedOpenAI
* Uses explicit model.compat if provided, otherwise auto-detects from provider/URL.
*/
function getCompat(model: Model<"openai-completions">): ResolvedOpenAICompat {
const detected = detectCompat(model);
if (!model.compat) return detected;
return {
supportsStore: model.compat.supportsStore ?? detected.supportsStore,
supportsDeveloperRole: model.compat.supportsDeveloperRole ?? detected.supportsDeveloperRole,
supportsReasoningEffort: model.compat.supportsReasoningEffort ?? detected.supportsReasoningEffort,
reasoningEffortMap: model.compat.reasoningEffortMap ?? detected.reasoningEffortMap,
supportsUsageInStreaming: model.compat.supportsUsageInStreaming ?? detected.supportsUsageInStreaming,
supportsToolChoice: model.compat.supportsToolChoice ?? detected.supportsToolChoice,
maxTokensField: model.compat.maxTokensField ?? detected.maxTokensField,
requiresToolResultName: model.compat.requiresToolResultName ?? detected.requiresToolResultName,
requiresAssistantAfterToolResult:
model.compat.requiresAssistantAfterToolResult ?? detected.requiresAssistantAfterToolResult,
requiresThinkingAsText: model.compat.requiresThinkingAsText ?? detected.requiresThinkingAsText,
requiresMistralToolIds: model.compat.requiresMistralToolIds ?? detected.requiresMistralToolIds,
thinkingFormat: model.compat.thinkingFormat ?? detected.thinkingFormat,
reasoningContentField: model.compat.reasoningContentField ?? detected.reasoningContentField,
requiresReasoningContentForToolCalls:
model.compat.requiresReasoningContentForToolCalls ?? detected.requiresReasoningContentForToolCalls,
requiresAssistantContentForToolCalls:
model.compat.requiresAssistantContentForToolCalls ?? detected.requiresAssistantContentForToolCalls,
openRouterRouting: model.compat.openRouterRouting ?? detected.openRouterRouting,
vercelGatewayRouting: model.compat.vercelGatewayRouting ?? detected.vercelGatewayRouting,
supportsStrictMode: model.compat.supportsStrictMode ?? detected.supportsStrictMode,
extraBody: model.compat.extraBody,
};
return resolveOpenAICompat(model);
}
+28
View File
@@ -277,6 +277,34 @@ describe("model thinking runtime helpers", () => {
);
});
it("derives binary-thinking fallback from resolved compat when catalog compat is partial", () => {
const model = enrichModelThinking({
id: "qwen/qwen3-32b",
name: "Qwen 3 32B",
api: "openai-completions",
provider: "openrouter",
baseUrl: "https://openrouter.ai/api/v1",
reasoning: true,
compat: {
supportsToolChoice: true,
},
input: ["text"],
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
contextWindow: 128000,
maxTokens: 32000,
} satisfies Model<"openai-completions">);
expect(model.thinking).toEqual({
mode: "effort",
minLevel: Effort.Minimal,
maxLevel: Effort.High,
});
expect(requireSupportedEffort(model, Effort.High)).toBe(Effort.High);
expect(() => requireSupportedEffort(model, Effort.XHigh)).toThrow(
/Supported efforts: minimal, low, medium, high/,
);
});
it("enables xhigh for openai-responses and openai-codex-responses APIs", () => {
const responsesModel = createModel({
id: "custom-responses",
+7 -1
View File
@@ -1,13 +1,19 @@
# Changelog
## [Unreleased]
### Added
- Added `/session delete` command to delete current session with confirmation and return to session selector
- Added session deletion in session selector via Delete key with confirmation dialog
### Changed
- Changed session deletion callback to return a boolean indicating success, allowing callers to distinguish between failed deletions and upstream cancellations
### Fixed
- Fixed OAuth redirect URI validation to preserve exact configured values without adding trailing slashes
- Fixed session deletion error handling to display error messages in the session selector UI instead of silently failing
- Added `oauth.redirectUri`, `oauth.clientSecret`, and `oauth.callbackPath` support for MCP server OAuth config so providers can use exact registered redirect URIs while preserving local callback listener settings ([#445](https://github.com/can1357/oh-my-pi/issues/445))
## [13.12.8] - 2026-03-16
@@ -36,13 +36,8 @@ export async function selectSession(sessions: SessionInfo[]): Promise<string | n
},
async (session: SessionInfo) => {
// Delete handler - SessionList will show confirmation internally
try {
await storage.deleteSessionWithArtifacts(session.path);
} catch (err) {
const errorMsg = err instanceof Error ? err.message : String(err);
console.error(`Failed to delete session: ${errorMsg}`);
throw err; // Re-throw so SessionList knows deletion failed
}
await storage.deleteSessionWithArtifacts(session.path);
return true;
},
);
return selector;
+7 -3
View File
@@ -17,14 +17,18 @@ function isLoopbackHostname(hostname: string): boolean {
}
function resolveRedirectUri(redirectUri: string | undefined): string | undefined {
const trimmed = redirectUri?.trim();
const configured = redirectUri;
const trimmed = configured?.trim();
if (!trimmed) return undefined;
if (trimmed !== configured) {
throw new Error("OAuth redirect URI must not include surrounding whitespace");
}
const parsed = new URL(trimmed);
const parsed = new URL(configured);
if (parsed.protocol !== "http:" && parsed.protocol !== "https:") {
throw new Error("OAuth redirect URI must use http or https");
}
return parsed.toString();
return configured;
}
function parseRedirectUri(redirectUri: string | undefined): URL | undefined {
@@ -4,6 +4,7 @@ import {
Input,
matchesKey,
padding,
replaceTabs,
Spacer,
Text,
truncateToWidth,
@@ -241,7 +242,8 @@ class SessionList implements Component {
export class SessionSelectorComponent extends Container {
#sessionList: SessionList;
#confirmationDialog: HookSelectorComponent | null = null;
#onDelete?: (session: SessionInfo) => Promise<void>;
#messageContainer: Container;
#onDelete?: (session: SessionInfo) => Promise<boolean>;
#onRequestRender?: () => void;
constructor(
@@ -249,19 +251,19 @@ export class SessionSelectorComponent extends Container {
onSelect: (sessionPath: string) => void,
onCancel: () => void,
onExit: () => void,
onDelete?: (session: SessionInfo) => Promise<void>,
onDelete?: (session: SessionInfo) => Promise<boolean>,
) {
super();
this.#messageContainer = new Container();
this.#onDelete = onDelete;
// Add header
this.addChild(new Spacer(1));
this.addChild(new Text(theme.bold("Resume Session"), 1, 0));
this.addChild(new Spacer(1));
this.addChild(new DynamicBorder());
this.addChild(new Spacer(1));
this.addChild(this.#messageContainer);
// Create session list
this.#sessionList = new SessionList(sessions);
this.#sessionList.onSelect = onSelect;
@@ -281,6 +283,16 @@ export class SessionSelectorComponent extends Container {
this.#onRequestRender = callback;
}
#clearError(): void {
this.#messageContainer.clear();
}
#showError(message: string): void {
this.#messageContainer.clear();
this.#messageContainer.addChild(new Text(theme.fg("error", `Error: ${replaceTabs(message)}`), 1, 0));
this.#messageContainer.addChild(new Spacer(1));
}
#showDeleteConfirmation(session: SessionInfo): void {
const displayName = session.title || session.firstMessage.slice(0, 40) || session.id;
this.#confirmationDialog = new HookSelectorComponent(
@@ -288,11 +300,14 @@ export class SessionSelectorComponent extends Container {
["Yes", "No"],
async (option: string) => {
if (option === "Yes" && this.#onDelete) {
this.#clearError();
try {
await this.#onDelete(session);
this.#sessionList.removeSession(session.path);
} catch {
// Error already shown by caller
const deleted = await this.#onDelete(session);
if (deleted) {
this.#sessionList.removeSession(session.path);
}
} catch (err) {
this.#showError(err instanceof Error ? err.message : String(err));
}
}
// Close confirmation dialog
@@ -607,15 +607,16 @@ export class SelectorController {
},
async (session: SessionInfo) => {
if (!(await this.#detachActiveSessionBeforeDeletion(session.path))) {
return;
return false;
}
const storage = new FileSessionStorage();
try {
await storage.deleteSessionWithArtifacts(session.path);
return true;
} catch (err) {
const errorMsg = err instanceof Error ? err.message : String(err);
this.ctx.showError(`Failed to delete session: ${errorMsg}`);
throw err;
throw new Error(`Failed to delete session: ${err instanceof Error ? err.message : String(err)}`, {
cause: err,
});
}
},
);
@@ -243,7 +243,7 @@ const BUILTIN_SLASH_COMMAND_REGISTRY: ReadonlyArray<BuiltinSlashCommandSpec> = [
return;
}
// Default: show session info
void runtime.ctx.handleSessionCommand();
await runtime.ctx.handleSessionCommand();
runtime.ctx.editor.setText("");
},
},
@@ -3,17 +3,10 @@ import { getBundledModel } from "@oh-my-pi/pi-ai";
import { runCommitAgentSession } from "../src/commit/agentic/agent";
import * as toolsModule from "../src/commit/agentic/tools";
import { Settings } from "../src/config/settings";
import type { CreateAgentSessionResult } from "../src/sdk";
import * as sdkModule from "../src/sdk";
import type { PromptOptions } from "../src/session/agent-session";
vi.mock("../src/sdk", () => ({
createAgentSession: vi.fn(),
}));
vi.mock("../src/commit/agentic/tools", () => ({
createCommitTools: vi.fn(() => []),
}));
describe("commit agent prompt attribution", () => {
afterEach(() => {
vi.restoreAllMocks();
@@ -29,10 +22,8 @@ describe("commit agent prompt attribution", () => {
dispose: async () => {},
};
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
});
(toolsModule.createCommitTools as unknown as { mockReturnValue: (value: unknown) => void }).mockReturnValue([]);
vi.spyOn(sdkModule, "createAgentSession").mockResolvedValue({ session } as unknown as CreateAgentSessionResult);
vi.spyOn(toolsModule, "createCommitTools").mockReturnValue([]);
const model = getBundledModel("anthropic", "claude-sonnet-4-5");
if (!model) {
@@ -133,6 +133,10 @@ function createContext(currentSessionFile: string): {
};
}
function renderText(selector: SessionSelectorComponent): string {
return selector.render(120).join("\n");
}
beforeAll(() => {
initTheme();
});
@@ -192,6 +196,36 @@ describe("SelectorController session deletion", () => {
expect(ctx.sessionManager.getSessionFile()).toBe("/tmp/project/sessions/detached.jsonl");
});
it("shows inline selector errors when session deletion fails after detach", async () => {
const activeSession = makeSessionInfo("/tmp/project/sessions/active.jsonl");
const { ctx, newSession } = createContext(activeSession.path);
vi.spyOn(SessionManager, "list").mockResolvedValue([activeSession]);
const deleteSessionWithArtifacts = vi
.spyOn(FileSessionStorage.prototype, "deleteSessionWithArtifacts")
.mockRejectedValue(new Error("disk failed"));
const controller = new SelectorController(ctx);
await controller.showSessionSelector();
const selector = ctx.editorContainer.children[0];
if (!(selector instanceof SessionSelectorComponent)) {
throw new Error("Expected session selector component");
}
const sessionList = selector.getSessionList() as unknown as {
onDeleteRequest?: (session: SessionInfo) => void;
};
sessionList.onDeleteRequest?.(activeSession);
selector.handleInput("\n");
await Bun.sleep(0);
expect(newSession).toHaveBeenCalledTimes(1);
expect(deleteSessionWithArtifacts).toHaveBeenCalledWith(activeSession.path);
expect(ctx.showError).not.toHaveBeenCalled();
expect(ctx.sessionManager.getSessionFile()).toBe("/tmp/project/sessions/detached.jsonl");
expect(renderText(selector)).toContain("Error: Failed to delete session: disk failed");
expect(renderText(selector)).toContain("Active session");
});
it("creates a fresh session before deleting via slash command and then shows the selector", async () => {
const activeSessionPath = "/tmp/project/sessions/active.jsonl";
const { ctx, calls, showHookConfirm, newSession } = createContext(activeSessionPath);
@@ -33,7 +33,7 @@ function createSession(id: string, title: string): SessionInfo {
};
}
function createSelector(onDelete: (session: SessionInfo) => Promise<void>): SessionSelectorComponent {
function createSelector(onDelete: (session: SessionInfo) => Promise<boolean>): SessionSelectorComponent {
return new SessionSelectorComponent(
[createSession("session-a", "Alpha"), createSession("session-b", "Beta")],
() => {},
@@ -48,7 +48,7 @@ function renderText(selector: SessionSelectorComponent): string {
}
describe("SessionSelectorComponent delete confirmation", () => {
it("keeps the session visible when delete fails after confirmation", async () => {
it("keeps the session visible and shows the error when delete fails after confirmation", async () => {
const onDelete = vi.fn(async () => {
throw new Error("disk failed");
});
@@ -63,13 +63,29 @@ describe("SessionSelectorComponent delete confirmation", () => {
const rendered = renderText(selector);
expect(onDelete).toHaveBeenCalledTimes(1);
expect(rendered).toContain("Error: disk failed");
expect(rendered).toContain("Alpha");
expect(rendered).toContain("Beta");
expect(rendered).not.toContain("Delete session?");
});
it("keeps the session visible when delete is canceled upstream", async () => {
const onDelete = vi.fn(async () => false);
const selector = createSelector(onDelete);
selector.handleInput("\x1b[3~");
selector.handleInput("\n");
await Bun.sleep(0);
const rendered = renderText(selector);
expect(onDelete).toHaveBeenCalledTimes(1);
expect(rendered).toContain("Alpha");
expect(rendered).toContain("Beta");
expect(rendered).not.toContain("Error:");
});
it("removes the session row after a successful delete", async () => {
const onDelete = vi.fn(async () => {});
const onDelete = vi.fn(async () => true);
const selector = createSelector(onDelete);
selector.handleInput("\x1b[3~");
@@ -163,6 +163,59 @@ describe("mcp oauth flow", () => {
});
});
it("preserves root redirectUri values without adding a trailing slash", async () => {
let observedRedirectUri = "";
let tokenRequestBody = "";
using _hook = hookFetch((input, init) => {
const url = String(input);
if (url === "https://provider.example/token") {
tokenRequestBody = String(init?.body ?? "");
return new Response(
JSON.stringify({
access_token: "access-token",
refresh_token: "refresh-token",
expires_in: 3600,
}),
{ status: 200, headers: { "Content-Type": "application/json" } },
);
}
throw new Error(`Unexpected fetch: ${url}`);
});
const flow = new MCPOAuthFlow(
{
authorizationUrl: "https://provider.example/authorize",
tokenUrl: "https://provider.example/token",
clientId: "client-id",
redirectUri: "https://public.example",
callbackPort: 14571,
},
{
onAuth: info => {
const authUrl = new URL(info.url);
observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? "";
const state = authUrl.searchParams.get("state") ?? "";
queueMicrotask(() => {
void originalFetch(`http://localhost:14571/?code=test-code&state=${state}`);
});
},
signal: AbortSignal.timeout(1_000),
},
);
const credentials = await flow.login();
const tokenParams = new URLSearchParams(tokenRequestBody);
expect(observedRedirectUri).toBe("https://public.example");
expect(tokenParams.get("redirect_uri")).toBe("https://public.example");
expect(credentials).toMatchObject({
access: "access-token",
refresh: "refresh-token",
});
});
it("supports https loopback redirectUri values behind a separate local callback port", async () => {
let observedRedirectUri = "";
let tokenRequestBody = "";
@@ -3,21 +3,29 @@ import type { InteractiveModeContext } from "@oh-my-pi/pi-coding-agent/modes/typ
import { executeBuiltinSlashCommand } from "@oh-my-pi/pi-coding-agent/slash-commands/builtin-registry";
function createRuntimeHarness(options?: {
handleSessionCommand?: InteractiveModeContext["handleSessionCommand"];
handleSessionDeleteCommand?: InteractiveModeContext["handleSessionDeleteCommand"];
}) {
const setText = vi.fn();
const handleSessionCommand =
options?.handleSessionCommand ??
vi.fn(async () => {
return;
});
const handleSessionDeleteCommand =
options?.handleSessionDeleteCommand ??
(async () => {
vi.fn(async () => {
return;
});
return {
setText,
handleSessionCommand,
handleSessionDeleteCommand,
runtime: {
ctx: {
editor: { setText } as unknown as InteractiveModeContext["editor"],
handleSessionCommand,
handleSessionDeleteCommand,
} as InteractiveModeContext,
handleBackgroundCommand: () => {},
@@ -25,7 +33,45 @@ function createRuntimeHarness(options?: {
};
}
describe("/session delete slash command", () => {
describe("/session slash command", () => {
it("awaits session info before resolving the default command", async () => {
const deferred = Promise.withResolvers<void>();
const handleSessionCommand = vi.fn(() => deferred.promise);
const harness = createRuntimeHarness({ handleSessionCommand });
let settled = false;
const execution = executeBuiltinSlashCommand("/session", harness.runtime).then(result => {
settled = true;
return result;
});
await Promise.resolve();
expect(handleSessionCommand).toHaveBeenCalledTimes(1);
expect(harness.handleSessionDeleteCommand).not.toHaveBeenCalled();
expect(harness.setText).not.toHaveBeenCalled();
expect(settled).toBe(false);
deferred.resolve();
expect(await execution).toBe(true);
expect(settled).toBe(true);
expect(harness.setText).toHaveBeenCalledWith("");
});
it("propagates session info failures through executeBuiltinSlashCommand", async () => {
const infoError = new Error("info failed");
const handleSessionCommand = vi.fn(async () => {
throw infoError;
});
const harness = createRuntimeHarness({ handleSessionCommand });
await expect(executeBuiltinSlashCommand("/session info", harness.runtime)).rejects.toBe(infoError);
expect(handleSessionCommand).toHaveBeenCalledTimes(1);
expect(harness.handleSessionDeleteCommand).not.toHaveBeenCalled();
expect(harness.setText).not.toHaveBeenCalled();
});
it("awaits session deletion before resolving the builtin command", async () => {
const deferred = Promise.withResolvers<void>();
const handleSessionDeleteCommand = vi.fn(() => deferred.promise);
@@ -2,17 +2,13 @@ import { afterEach, describe, expect, it, vi } from "bun:test";
import { type AssistantMessage, Effort } from "@oh-my-pi/pi-ai";
import { Settings } from "../../src/config/settings";
import type { LoadExtensionsResult } from "../../src/extensibility/extensions/types";
import type { CreateAgentSessionResult } from "../../src/sdk";
import * as sdkModule from "../../src/sdk";
import type { AgentSession, AgentSessionEvent, PromptOptions } from "../../src/session/agent-session";
import type { AuthStorage } from "../../src/session/auth-storage";
import { runSubprocess, SUBAGENT_WARNING_MISSING_SUBMIT_RESULT } from "../../src/task/executor";
import type { AgentDefinition } from "../../src/task/types";
vi.mock("../../src/sdk", () => ({
createAgentSession: vi.fn(),
discoverAuthStorage: vi.fn(async () => ({})),
}));
function createAssistantStopMessage(text: string): AssistantMessage {
return {
role: "assistant",
@@ -80,6 +76,18 @@ function createMockSession(
return session as unknown as AgentSession;
}
function createSessionResult(session: AgentSession): CreateAgentSessionResult {
return {
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
};
}
function mockCreateAgentSession(session: AgentSession) {
return vi.spyOn(sdkModule, "createAgentSession").mockResolvedValue(createSessionResult(session));
}
describe("runSubprocess submit_result reminders", () => {
afterEach(() => {
vi.restoreAllMocks();
@@ -128,11 +136,7 @@ describe("runSubprocess submit_result reminders", () => {
});
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
});
mockCreateAgentSession(session);
const result = await runSubprocess(baseOptions);
expect(prompts.length).toBe(2);
@@ -164,11 +168,7 @@ describe("runSubprocess submit_result reminders", () => {
});
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
});
mockCreateAgentSession(session);
const result = await runSubprocess({ ...baseOptions, id: "subagent-2" });
expect(result.output).toContain("SYSTEM WARNING: Subagent called submit_result with null data.");
@@ -206,11 +206,7 @@ describe("runSubprocess submit_result reminders", () => {
});
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
});
mockCreateAgentSession(session);
const result = await runSubprocess({ ...baseOptions, id: "subagent-err-then-success" });
expect(prompts).toHaveLength(2);
@@ -232,11 +228,7 @@ describe("runSubprocess submit_result reminders", () => {
});
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
});
const createAgentSessionSpy = mockCreateAgentSession(session);
const modelRegistry = {
refresh: async () => {},
@@ -251,11 +243,8 @@ describe("runSubprocess submit_result reminders", () => {
modelRegistry,
});
const createAgentSessionMock = sdkModule.createAgentSession as unknown as {
mock: { calls: Array<[Record<string, unknown>]> };
};
expect(createAgentSessionMock.mock.calls).toHaveLength(1);
expect(createAgentSessionMock.mock.calls[0]?.[0]?.thinkingLevel).toBe("high");
expect(createAgentSessionSpy).toHaveBeenCalledTimes(1);
expect(createAgentSessionSpy.mock.calls[0]?.[0]?.thinkingLevel).toBe(Effort.High);
});
it("prefers explicit modelOverride thinking suffix over provided thinking level, including off", async () => {
@@ -270,6 +259,8 @@ describe("runSubprocess submit_result reminders", () => {
{ modelOverride: "openai/gpt-4o:off", expectedThinkingLevel: "off" },
] as const;
const createAgentSessionSpy = vi.spyOn(sdkModule, "createAgentSession");
for (const [index, testCase] of cases.entries()) {
const session = createMockSession(({ emit }) => {
emit({
@@ -284,13 +275,7 @@ describe("runSubprocess submit_result reminders", () => {
});
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue(
{
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
},
);
createAgentSessionSpy.mockResolvedValue(createSessionResult(session));
await runSubprocess({
...baseOptions,
@@ -301,12 +286,9 @@ describe("runSubprocess submit_result reminders", () => {
});
}
const createAgentSessionMock = sdkModule.createAgentSession as unknown as {
mock: { calls: Array<[Record<string, unknown>]> };
};
expect(createAgentSessionMock.mock.calls).toHaveLength(2);
expect(createAgentSessionMock.mock.calls[0]?.[0]?.thinkingLevel).toBe(cases[0].expectedThinkingLevel);
expect(createAgentSessionMock.mock.calls[1]?.[0]?.thinkingLevel).toBe(cases[1].expectedThinkingLevel);
expect(createAgentSessionSpy).toHaveBeenCalledTimes(2);
expect(createAgentSessionSpy.mock.calls[0]?.[0]?.thinkingLevel).toBe(cases[0].expectedThinkingLevel);
expect(createAgentSessionSpy.mock.calls[1]?.[0]?.thinkingLevel).toBe(cases[1].expectedThinkingLevel);
});
it("aborts after 3 reminders when submit_result is never called", async () => {
const prompts: string[] = [];
@@ -317,11 +299,7 @@ describe("runSubprocess submit_result reminders", () => {
emit({ type: "message_end", message: assistant });
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
});
mockCreateAgentSession(session);
const result = await runSubprocess({
...baseOptions,
@@ -354,11 +332,7 @@ describe("runSubprocess submit_result reminders", () => {
});
});
(sdkModule.createAgentSession as unknown as { mockResolvedValue: (value: unknown) => void }).mockResolvedValue({
session,
extensionsResult: {} as unknown as LoadExtensionsResult,
setToolUIContext: () => {},
});
mockCreateAgentSession(session);
const result = await runSubprocess({ ...baseOptions, id: "subagent-aborted-submit-result" });
expect(result.aborted).toBe(true);