diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index d1722f2ec..bc642b3d0 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -2,6 +2,10 @@ ## [Unreleased] +### Fixed + +- Fixed the SSH tool to reject `cwd` values of `~` and `~/...` before sending guaranteed-bad quoted tilde paths to remote POSIX shells. ([#4002](https://github.com/can1357/oh-my-pi/issues/4002)) + ## [16.2.12] - 2026-07-01 ### Breaking Changes diff --git a/packages/coding-agent/src/prompts/tools/ssh.md b/packages/coding-agent/src/prompts/tools/ssh.md index 150342864..f8e351d05 100644 --- a/packages/coding-agent/src/prompts/tools/ssh.md +++ b/packages/coding-agent/src/prompts/tools/ssh.md @@ -19,4 +19,5 @@ Runs commands on remote hosts. You MUST verify the shell type from "Available hosts" and use matching commands. +You SHOULD omit `cwd` unless required. `cwd` MUST be an explicit remote path; NEVER use `~` or `~/…`. diff --git a/packages/coding-agent/src/tools/ssh.ts b/packages/coding-agent/src/tools/ssh.ts index f36c6e3d2..2b06a20d9 100644 --- a/packages/coding-agent/src/tools/ssh.ts +++ b/packages/coding-agent/src/tools/ssh.ts @@ -26,7 +26,7 @@ import { clampTimeout } from "./tool-timeouts"; const sshSchema = type({ host: type("string").describe("ssh host"), command: type("string").describe("remote command"), - "cwd?": type("string").describe("remote working directory"), + "cwd?": type("string").describe("remote working directory; omit unless required, never ~ or ~/..."), "timeout?": type("number").describe("timeout in seconds"), }); @@ -89,6 +89,12 @@ function quoteCmdPath(value: string): string { const escaped = value.replace(/"/g, '""'); return `"${escaped}"`; } +function assertValidSshCwd(cwd: string | undefined): void { + if (!cwd) return; + if (cwd === "~" || cwd.startsWith("~/")) { + throw new ToolError("SSH cwd must be an absolute remote path; omit cwd instead of using ~."); + } +} function buildRemoteCommand(command: string, cwd: string | undefined, info: SSHHostInfo): string { if (!cwd) return command; @@ -177,6 +183,7 @@ export class SshTool implements AgentTool { if (!hostConfig) { throw new ToolError(`SSH host not loaded: ${host}`); } + assertValidSshCwd(cwd); const hostInfo = await ensureHostInfo(hostConfig); const remoteCommand = buildRemoteCommand(command, cwd, hostInfo); diff --git a/packages/coding-agent/test/tools/ssh-description.test.ts b/packages/coding-agent/test/tools/ssh-description.test.ts index 673ea3180..e46b5bb10 100644 --- a/packages/coding-agent/test/tools/ssh-description.test.ts +++ b/packages/coding-agent/test/tools/ssh-description.test.ts @@ -2,8 +2,11 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import type { SSHHost } from "@oh-my-pi/pi-coding-agent/capability/ssh"; import type { SourceMeta } from "@oh-my-pi/pi-coding-agent/capability/types"; import * as discovery from "@oh-my-pi/pi-coding-agent/discovery"; +import type { SSHHostInfo } from "@oh-my-pi/pi-coding-agent/ssh/connection-manager"; +import * as connectionManager from "@oh-my-pi/pi-coding-agent/ssh/connection-manager"; +import * as sshExecutor from "@oh-my-pi/pi-coding-agent/ssh/ssh-executor"; import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; -import { loadSshTool } from "@oh-my-pi/pi-coding-agent/tools"; +import { loadSshTool } from "@oh-my-pi/pi-coding-agent/tools/ssh"; const SOURCE: SourceMeta = { provider: "test", @@ -16,6 +19,8 @@ const SOURCE: SourceMeta = { const RUN_ID = `${Date.now()}-${process.pid}`; const HOST_A: SSHHost = { name: `a-omp-test-${RUN_ID}`, host: "alpha.example.com", _source: SOURCE }; const HOST_B: SSHHost = { name: `b-omp-test-${RUN_ID}`, host: "beta.example.com", _source: SOURCE }; +const LINUX_BASH_INFO: SSHHostInfo = { version: 4, os: "linux", shell: "bash", compatEnabled: false }; +const WINDOWS_CMD_INFO: SSHHostInfo = { version: 4, os: "windows", shell: "cmd", compatEnabled: false }; function mockHosts(hosts: SSHHost[]): void { vi.spyOn(discovery, "loadCapability").mockResolvedValue({ @@ -30,6 +35,33 @@ function createSession(): ToolSession { return { cwd: "/tmp" } as unknown as ToolSession; } +async function loadTestTool(hosts: SSHHost[] = [HOST_A]) { + mockHosts(hosts); + const tool = await loadSshTool(createSession()); + if (!tool) { + throw new Error("expected SSH tool"); + } + return tool; +} + +function stubSshExecute() { + return vi.spyOn(sshExecutor, "executeSSH").mockResolvedValue({ + output: "ok", + exitCode: 0, + cancelled: false, + truncated: false, + totalLines: 1, + totalBytes: 2, + outputLines: 1, + outputBytes: 2, + }); +} + +function stubSshRun(info: SSHHostInfo) { + vi.spyOn(connectionManager, "ensureHostInfo").mockResolvedValue(info); + return stubSshExecute(); +} + describe("loadSshTool description", () => { afterEach(() => { vi.restoreAllMocks(); @@ -45,6 +77,7 @@ describe("loadSshTool description", () => { const tool = await loadSshTool(createSession()); expect(tool).not.toBeNull(); expect(tool?.description.startsWith("Runs commands on remote hosts.")).toBe(true); + expect(tool?.description).toContain("NEVER use `~` or `~/…`"); expect( tool?.description.endsWith( `\n\nAvailable hosts:\n- ${HOST_A.name} (${HOST_A.host}) | detecting...\n- ${HOST_B.name} (${HOST_B.host}) | detecting...`, @@ -52,3 +85,51 @@ describe("loadSshTool description", () => { ).toBe(true); }); }); + +describe("SshTool cwd handling", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("rejects tilde cwd values before probing or executing the host", async () => { + const ensureSpy = vi.spyOn(connectionManager, "ensureHostInfo").mockResolvedValue(LINUX_BASH_INFO); + const executeSpy = stubSshExecute(); + const tool = await loadTestTool(); + + await expect(tool.execute("call-tilde", { host: HOST_A.name, command: "pwd", cwd: "~" })).rejects.toThrow( + "SSH cwd must be an absolute remote path", + ); + await expect( + tool.execute("call-tilde-path", { host: HOST_A.name, command: "pwd", cwd: "~/src" }), + ).rejects.toThrow("SSH cwd must be an absolute remote path"); + + expect(ensureSpy).not.toHaveBeenCalled(); + expect(executeSpy).not.toHaveBeenCalled(); + }); + + it("quotes valid POSIX absolute cwd values in the remote command", async () => { + const executeSpy = stubSshRun(LINUX_BASH_INFO); + const tool = await loadTestTool(); + + await tool.execute("call-absolute", { host: HOST_A.name, command: "pwd", cwd: "/srv/app" }); + + expect(executeSpy).toHaveBeenCalledWith( + HOST_A, + "cd -- '/srv/app' && pwd", + expect.objectContaining({ compatEnabled: false, timeout: 60000 }), + ); + }); + + it("preserves native Windows cwd command generation", async () => { + const executeSpy = stubSshRun(WINDOWS_CMD_INFO); + const tool = await loadTestTool(); + + await tool.execute("call-windows", { host: HOST_A.name, command: "dir", cwd: "C:\\Users\\me" }); + + expect(executeSpy).toHaveBeenCalledWith( + HOST_A, + 'cd /d "C:\\Users\\me" && dir', + expect.objectContaining({ compatEnabled: false, timeout: 60000 }), + ); + }); +});