返回 CodeWhale
wait-for.test.mjs
根目录 / crates / tui / plugins / computer-use / tests / wait-for.test.mjs
1 // wait_for polling, element-targeted type/key, and the persistent agent
2 // channel: real MCP server over stdio with the injected fake backend, plus
3 // the remote agent's --serve mode driven as a local child process.
4 import { hostKeysLine, attest, attestParams } from "./fixtures/host-decision.mjs";
5 import { test, before, after } from "node:test";
6 import assert from "node:assert/strict";
7 import { spawn } from "node:child_process";
8 import fs from "node:fs";
9 import os from "node:os";
10 import path from "node:path";
11 import url from "node:url";
12 import { ensureSshChannel, channelRequest, b64 } from "../src/transport.mjs";
13
14 const __dirname = path.dirname(url.fileURLToPath(import.meta.url));
15 const ROOT = path.resolve(__dirname, "..");
16 const FAKE = path.join(__dirname, "fixtures", "fake-backend.mjs");
17
18 const stateDir = fs.mkdtempSync(path.join(os.tmpdir(), "cu-wait-state-"));
19 const recDir = fs.mkdtempSync(path.join(os.tmpdir(), "cu-wait-rec-"));
20 const callsFile = path.join(fs.mkdtempSync(path.join(os.tmpdir(), "cu-wait-")), "calls.jsonl");
21 const controlFile = callsFile + ".control.json";
22
23 let server;
24 let buf = "";
25 const pending = new Map();
26 let nextId = 1;
27
28 function rpc(method, params, timeoutMs = 30_000) {
29 const id = nextId++;
30 return new Promise((resolve, reject) => {
31 const t = setTimeout(() => { pending.delete(id); reject(new Error(`timeout: ${method}`)); }, timeoutMs);
32 pending.set(id, (msg) => { clearTimeout(t); resolve(msg); });
33 server.stdin.write(JSON.stringify({ jsonrpc: "2.0", id, method, params: attestParams(method, params) }) + "\n");
34 });
35 }
36
37 async function tool(name, args = {}) {
38 const res = await rpc("tools/call", { name, arguments: args });
39 assert.ok(res.result, `${name}: protocol error ${JSON.stringify(res.error ?? {})}`);
40 return JSON.parse(res.result.content[0].text);
41 }
42
43 function calls(method) {
44 if (!fs.existsSync(callsFile)) return [];
45 return fs.readFileSync(callsFile, "utf8").split("\n").filter(Boolean).map((l) => JSON.parse(l)).filter((c) => c.method === method);
46 }
47
48 function setControl(obj) {
49 if (obj == null) fs.rmSync(controlFile, { force: true });
50 else fs.writeFileSync(controlFile, JSON.stringify(obj));
51 }
52
53 before(async () => {
54 server = spawn("node", [path.join(ROOT, "mcp", "server.mjs")], {
55 env: {
56 ...process.env,
57 CODEWHALE_CU_APP: "off",
58 CODEWHALE_CU_STATE_DIR: stateDir,
59 CODEWHALE_CU_RECORDINGS_DIR: recDir,
60 CODEWHALE_CU_TEST_BACKEND: FAKE,
61 FAKE_BACKEND_CALLS: callsFile,
62 FAKE_BACKEND_CONTROL: controlFile,
63 },
64 stdio: ["pipe", "pipe", "pipe"],
65 });
66 server.stdin.write(hostKeysLine());
67 server.stderr.on("data", (d) => process.stderr.write(`[server] ${d}`));
68 server.stdout.setEncoding("utf8");
69 server.stdout.on("data", (d) => {
70 buf += d;
71 let i;
72 while ((i = buf.indexOf("\n")) !== -1) {
73 const line = buf.slice(0, i).trim();
74 buf = buf.slice(i + 1);
75 if (!line) continue;
76 try {
77 const msg = JSON.parse(line);
78 if (msg.id && pending.has(msg.id)) { pending.get(msg.id)(msg); pending.delete(msg.id); }
79 } catch {}
80 }
81 });
82 const init = await rpc("initialize", { protocolVersion: "2025-06-18" });
83 assert.equal(init.result.serverInfo.name, "codewhale-cu");
84 // The local consent ledger gates app-targeted calls; record the fixture
85 // app's decision up front, as a real session would.
86 const c = await tool("consent", { action: "allow", app: "FakeApp" });
87 assert.equal(c.ok, true, JSON.stringify(c));
88 });
89
90 after(() => {
91 server?.kill("SIGTERM");
92 for (const d of [stateDir, recDir, path.dirname(callsFile)]) { try { fs.rmSync(d, { recursive: true, force: true }); } catch {} }
93 });
94
95 test("wait_for returns matched elements bound to a fresh targetable state_id", async () => {
96 const r = await tool("wait_for", { query: "OK", role: "AXButton", timeout: 5 });
97 assert.equal(r.ok, true, JSON.stringify(r));
98 assert.equal(r.matched, true);
99 assert.equal(r.matched_count, 1);
100 assert.equal(r.elements[0].label, "OK");
101 assert.ok(r.state_id, "the satisfying observation is bound");
102 // The returned state_id is targetable: the button resolves and presses.
103 const press = await tool("perform_action", { target: { type: "element", state_id: r.state_id, index: r.elements[0].index }, action: "AXPress" });
104 assert.equal(press.ok, true, JSON.stringify(press.error));
105 });
106
107 test("wait_for absent is satisfied immediately when nothing matches", async () => {
108 const r = await tool("wait_for", { query: "no such element anywhere", state: "absent", timeout: 5 });
109 assert.equal(r.ok, true, JSON.stringify(r));
110 assert.equal(r.matched, true);
111 assert.equal(r.matched_count, 0);
112 assert.equal(r.timed_out, undefined);
113 });
114
115 test("wait_for times out honestly when the predicate never holds", async () => {
116 const before = calls("get_app_state").length;
117 const r = await tool("wait_for", { query: "never-present-label", timeout: 0.6, interval: 150 });
118 assert.equal(r.ok, true, JSON.stringify(r));
119 assert.equal(r.matched, false);
120 assert.equal(r.timed_out, true);
121 assert.ok(r.polls >= 2, `expected several polls, got ${r.polls}`);
122 assert.equal(r.state_id, undefined, "a timed-out wait binds nothing");
123 assert.ok(calls("get_app_state").length - before >= r.polls - 1, "ephemeral polls reached the backend");
124 });
125
126 test("ephemeral wait_for polls do not evict earlier states", async () => {
127 const st = await tool("get_app_state", { app_ref: { name: "FakeApp" } });
128 assert.ok(st.state_id);
129 // ~30 polls: more than the 24-state cache cap. If polls were cached, st
130 // would be evicted; ephemeral polling keeps it targetable.
131 const w = await tool("wait_for", { query: "never-present-label", timeout: 3, interval: 100 });
132 assert.equal(w.timed_out, true);
133 assert.ok(w.polls > 24, `expected >24 polls to prove non-caching, got ${w.polls}`);
134 const focus = await tool("focus", { target: { type: "element", state_id: st.state_id, index: 1 } });
135 assert.equal(focus.ok, true, JSON.stringify(focus.error));
136 });
137
138 test("wait_for validates its arguments", async () => {
139 const bad = async (args, match) => {
140 const res = await rpc("tools/call", { name: "wait_for", arguments: args });
141 assert.match(res.error?.message ?? "", match, JSON.stringify(res));
142 };
143 await bad({}, /needs a query and\/or role/);
144 await bad({ query: "x", state: "bogus" }, /state must be "present" or "absent"/);
145 await bad({ query: "x", timeout: 999 }, /timeout/);
146 await bad({ query: "x", interval: 5 }, /interval/);
147 });
148
149 test("type with an element target focuses first, then types", async () => {
150 const st = await tool("get_app_state", { app_ref: { name: "FakeApp" } });
151 const beforeFocus = calls("focus").length;
152 const beforeType = calls("type").length;
153 setControl({ found: true, element: { role: "AXTextField", position: { x: 10, y: 60 }, size: { w: 150, h: 25 } }, reason: null });
154 try {
155 const r = await tool("type", { text: "hello", target: { type: "element", state_id: st.state_id, index: 8 } });
156 assert.equal(r.ok, true, JSON.stringify(r.error));
157 const focusCalls = calls("focus");
158 const typeCalls = calls("type");
159 assert.equal(focusCalls.length, beforeFocus + 1);
160 assert.equal(typeCalls.length, beforeType + 1);
161 assert.deepEqual(focusCalls.at(-1).args.target.path, [0, 2], "focus received the semantic element target");
162 assert.equal(typeCalls.at(-1).args.text, "hello");
163 } finally { setControl(null); }
164 });
165
166 test("key with an element target focuses first; a coordinate target is refused", async () => {
167 const st = await tool("get_app_state", { app_ref: { name: "FakeApp" } });
168 setControl({ found: true, element: { role: "AXTextField", position: { x: 10, y: 60 }, size: { w: 150, h: 25 } }, reason: null });
169 try {
170 const r = await tool("key", { text: "return", target: { type: "element", state_id: st.state_id, index: 8 } });
171 assert.equal(r.ok, true, JSON.stringify(r.error));
172 } finally { setControl(null); }
173 const refused = await tool("type", { text: "x", target: { type: "coordinate", x: 1, y: 1 } });
174 assert.equal(refused.ok, false);
175 assert.equal(refused.error.code, "bad_target");
176 });
177
178 test("type target fails closed when the element went stale before any text is sent", async () => {
179 const st = await tool("get_app_state", { app_ref: { name: "FakeApp" } });
180 setControl({ found: false, element: null, reason: "element_gone" });
181 const before = calls("type").length;
182 try {
183 const r = await tool("type", { text: "should never land", target: { type: "element", state_id: st.state_id, index: 1 } });
184 assert.equal(r.ok, false);
185 assert.equal(r.error.code, "element_stale");
186 assert.equal(r.stage, "focus");
187 assert.equal(calls("type").length, before, "no keystrokes after the failed focus");
188 } finally {
189 setControl(null);
190 }
191 });
192
193 // ---------- persistent remote agent ----------
194
195 test("agent --serve keeps its backend binding across requests", async t => {
196 const child = spawn("node", [path.join(ROOT, "agent.mjs"), "--serve"], {
197 env: { ...process.env, CODEWHALE_CU_TEST_BACKEND: FAKE, FAKE_BACKEND_CALLS: callsFile },
198 stdio: ["pipe", "pipe", "pipe"],
199 });
200 t.after(() => child.kill("SIGKILL"));
201 const replies = new Map();
202 let rbuf = "";
203 child.stdout.setEncoding("utf8").on("data", (d) => {
204 rbuf += d;
205 let i;
206 while ((i = rbuf.indexOf("\n")) !== -1) {
207 const msg = JSON.parse(rbuf.slice(0, i));
208 rbuf = rbuf.slice(i + 1);
209 replies.set(msg.id, msg);
210 }
211 });
212 const send = async (id, toolName, args = {}) => {
213 child.stdin.write(b64({ id, tool: toolName, args }) + "\n");
214 const deadline = Date.now() + 10_000;
215 while (!replies.has(id) && Date.now() < deadline) await new Promise((r) => setTimeout(r, 10));
216 assert.ok(replies.has(id), `no reply for ${toolName}`);
217 return replies.get(id);
218 };
219 assert.equal((await send(1, "platform")).ok, true);
220 const open = await send(2, "open_application", { name: "FakeApp" });
221 assert.equal(open.ok, true, JSON.stringify(open));
222 const typed = await send(3, "type", { text: "hi" });
223 assert.equal(typed.ok, true);
224 assert.equal(typed.data.bound_app, "FakeApp", "the second request reused the first request's binding");
225 // Held input is no longer blanket-refused in persistent mode: it reaches
226 // the backend, which the fixture answers.
227 const held = await send(4, "left_mouse_down", { target: { type: "coordinate", x: 5, y: 5 } });
228 assert.equal(held.ok, true, JSON.stringify(held));
229 const refused = await send(5, "shell_escape");
230 assert.equal(refused.ok, false);
231 assert.equal(refused.error.code, "tool_not_allowed");
232 });
233
234 test("one-shot agent still refuses operations that need a persistent session", async () => {
235 const out = spawn("node", [path.join(ROOT, "agent.mjs"), b64({ tool: "left_mouse_down", args: {} })], {
236 env: { ...process.env, CODEWHALE_CU_TEST_BACKEND: FAKE },
237 stdio: ["ignore", "pipe", "pipe"],
238 });
239 let stdout = "";
240 out.stdout.setEncoding("utf8").on("data", (d) => { stdout += d; });
241 await new Promise((resolve) => out.once("close", resolve));
242 const reply = JSON.parse(stdout.trim().split("\n").pop());
243 assert.equal(reply.ok, false);
244 assert.equal(reply.error.code, "persistent_session_required");
245 });
246
247 test("channelRequest correlates replies and survives out-of-order completion", async t => {
248 const binding = {};
249 const ch = ensureSshChannel(binding, ["node", path.join(ROOT, "agent.mjs"), "--serve"]);
250 t.after(() => { try { ch.proc?.kill("SIGKILL"); } catch {} });
251 // Override spawn is not needed here: the channel factory takes argv, so a
252 // local agent process stands in for `ssh host node agent --serve`.
253 assert.equal(ch.alive, true);
254 const p1 = channelRequest(ch, { tool: "platform" }, 10_000);
255 const p2 = channelRequest(ch, { tool: "not_a_tool" }, 10_000);
256 const [r1, r2] = await Promise.all([p1, p2]);
257 assert.equal(r1.ok, true);
258 assert.equal(r2.ok, false);
259 assert.equal(r2.error.code, "tool_not_allowed");
260 ch.proc.kill("SIGKILL");
261 await assert.rejects(channelRequest(ch, { tool: "platform" }, 500), (e) => e.code === "remote_session_lost");
262 });
263
263 lines Plain Text