返回 CodeWhale
tool_cancellation.rs
根目录 / crates / tui / src / core / engine / tests / tool_cancellation.rs
1 use super::*;
2 use crate::llm_client::mock::{MockLlmClient, canned};
3 use crate::tools::spec::{ToolCapability, ToolContext, ToolSpec};
4
5 #[cfg(unix)]
6 #[tokio::test]
7 #[allow(clippy::await_holding_lock)]
8 async fn engine_cancel_stops_started_foreground_descendants_and_preserves_background() {
9 let _env = lock_test_env();
10 let workspace = tempdir().expect("workspace");
11 let pid_file = workspace.path().join("descendant.pid");
12 let command = format!(
13 "printf 'foreground-started\\n'; CODEWHALE_SHELL_DESCENDANT_HELPER=1 \
14 CODEWHALE_SHELL_DESCENDANT_PID_FILE={} {} --exact \
15 tools::shell::tests::shell_descendant_helper_process --nocapture",
16 shell_words::quote(&pid_file.display().to_string()),
17 shell_words::quote(&std::env::current_exe().unwrap().display().to_string()),
18 );
19 let arguments = json!({"command": command, "timeout": 3600}).to_string();
20 let mock = Arc::new(MockLlmClient::new(vec![
21 tool_batch_turn(&[
22 ("call-foreground", "bash", &arguments),
23 (
24 "call-skipped",
25 "bash",
26 r#"{"command":"touch must-not-start"}"#,
27 ),
28 ]),
29 canned::simple_text_turn("Next user turn completed."),
30 ]));
31 let config = Config::default();
32 let (engine, handle) = Engine::new_with_model_client(
33 deterministic_engine_config(workspace.path()),
34 &config,
35 mock.clone(),
36 );
37 let shell_manager = engine.shell_manager.clone();
38 let session_id = engine.session.id.clone();
39 let background_id = shell_manager
40 .lock()
41 .unwrap()
42 .execute_with_options_env_for_session(
43 "sleep 30",
44 None,
45 30_000,
46 true,
47 None,
48 false,
49 None,
50 HashMap::new(),
51 &session_id,
52 )
53 .expect("start intentionally backgrounded control")
54 .task_id
55 .unwrap();
56 let task = tokio::spawn(engine.run());
57 let mut op = external_user_message_op("Run the foreground fixture.", AppMode::Agent, &config);
58 if let Op::SendMessage(TurnSpec {
59 trust_mode,
60 auto_approve,
61 approval_mode,
62 ..
63 }) = &mut op
64 {
65 *trust_mode = true;
66 *auto_approve = true;
67 *approval_mode = ApprovalMode::Bypass;
68 }
69 handle.send(op).await.expect("dispatch real shell tool");
70 let descendant: libc::pid_t = tokio::time::timeout(Duration::from_secs(10), async {
71 loop {
72 if let Ok(raw) = fs::read_to_string(&pid_file)
73 && let Ok(pid) = raw.trim().parse()
74 {
75 break pid;
76 }
77 tokio::time::sleep(Duration::from_millis(10)).await;
78 }
79 })
80 .await
81 .expect("the actual descendant must start before cancellation");
82 handle.cancel();
83
84 let mut receipts = Vec::new();
85 let mut foreground_execution_id = None;
86 let mut skipped = false;
87 tokio::time::timeout(Duration::from_secs(10), async {
88 let mut events = handle.rx_event.write().await;
89 loop {
90 match events.recv().await.expect("engine remains alive") {
91 Event::ToolCallComplete {
92 id,
93 result,
94 model_call: Some(model_call),
95 ..
96 } if model_call.provider_id == "call-foreground" => {
97 foreground_execution_id = Some(id);
98 receipts.push(result.expect("model-visible cancellation receipt"));
99 }
100 Event::ToolCallComplete {
101 model_call: Some(model_call),
102 result,
103 ..
104 } if model_call.provider_id == "call-skipped" => {
105 let result = result.unwrap();
106 assert_eq!(result.metadata.unwrap()["executed"], false);
107 assert!(result.content.contains("before this tool ran"));
108 skipped = true;
109 }
110 Event::TurnComplete { status, error, .. } => {
111 assert_eq!(status, TurnOutcomeStatus::Interrupted, "{error:?}");
112 break;
113 }
114 _ => {}
115 }
116 }
117 })
118 .await
119 .expect("the cancelled turn must settle");
120 tokio::time::timeout(Duration::from_secs(5), async {
121 loop {
122 // This is the PID written by the isolated helper, not an ambient process.
123 if unsafe { libc::kill(descendant, 0) } == -1
124 && std::io::Error::last_os_error().raw_os_error() == Some(libc::ESRCH)
125 {
126 break;
127 }
128 tokio::time::sleep(Duration::from_millis(10)).await;
129 }
130 })
131 .await
132 .expect("foreground descendant must be gone before recovery");
133 assert_eq!(receipts.len(), 1);
134 assert!(
135 skipped,
136 "the unstarted call must have its own truthful receipt"
137 );
138 assert!(!workspace.path().join("must-not-start").exists());
139 assert!(!receipts[0].success);
140 assert!(receipts[0].content.contains("after shell work started"));
141 assert!(receipts[0].content.contains("Killed"));
142 assert!(!receipts[0].content.contains("before this tool ran"));
143 assert_eq!(receipts[0].metadata.as_ref().unwrap()["executed"], true);
144 {
145 let mut manager = shell_manager.lock().unwrap();
146 let jobs = manager.list_jobs_for_session(&session_id);
147 let foreground = jobs
148 .iter()
149 .find(|job| job.origin_tool_call_id.as_deref() == foreground_execution_id.as_deref())
150 .expect("the exact foreground owner remains inspectable");
151 assert_eq!(foreground.status, crate::tools::shell::ShellStatus::Killed);
152 assert!(foreground.stdout_tail.contains("foreground-started"));
153 assert_eq!(
154 manager.inspect_job(&background_id).unwrap().snapshot.status,
155 crate::tools::shell::ShellStatus::Running
156 );
157 assert!(
158 !manager.has_finished_unreported_jobs_for_session(&session_id),
159 "cancelled foreground work must not wake an unsolicited model turn"
160 );
161 }
162 assert_eq!(mock.call_count(), 1);
163 let snapshot = tokio::time::timeout(Duration::from_secs(10), handle.get_session_snapshot())
164 .await
165 .expect("cancelled session must remain inspectable")
166 .unwrap();
167 let manager =
168 crate::session_manager::SessionManager::new(workspace.path().join("checkpoints")).unwrap();
169 let saved = crate::session_manager::create_saved_session_with_id_and_mode(
170 session_id.clone(),
171 &snapshot.messages,
172 "mock-model",
173 workspace.path(),
174 0,
175 None,
176 Some("agent"),
177 );
178 manager.save_checkpoint(&saved).unwrap();
179 let restored = manager
180 .load_session_checkpoint(&session_id)
181 .unwrap()
182 .unwrap();
183 assert!(restored.messages.iter().flat_map(|message| &message.content).any(|block| matches!(
184 block, ContentBlock::ToolResult { tool_use_id, content, is_error: Some(true), .. }
185 if tool_use_id == "call-foreground" && content.contains("after shell work started")
186 )));
187
188 handle
189 .send(external_user_message_op(
190 "Continue after cancellation.",
191 AppMode::Agent,
192 &config,
193 ))
194 .await
195 .unwrap();
196 tokio::time::timeout(Duration::from_secs(10), handle.get_session_snapshot())
197 .await
198 .unwrap()
199 .unwrap();
200 assert_eq!(
201 mock.call_count(),
202 2,
203 "only the next explicit user turn may resume"
204 );
205 assert_eq!(
206 guardian_tool_results(&mock.captured_requests()[1], "call-foreground")[0].1,
207 Some(true)
208 );
209 shell_manager.lock().unwrap().kill(&background_id).unwrap();
210 handle.send(Op::Shutdown).await.unwrap();
211 tokio::time::timeout(Duration::from_secs(10), task)
212 .await
213 .expect("engine must stop after shutdown")
214 .unwrap();
215 }
216
217 struct ReturnedResultTool;
218
219 #[async_trait::async_trait]
220 impl ToolSpec for ReturnedResultTool {
221 fn name(&self) -> &str {
222 "returned_result"
223 }
224 fn description(&self) -> &str {
225 "Return the selected success or failure fixture."
226 }
227 fn input_schema(&self) -> Value {
228 json!({"type":"object","properties":{"fail":{"type":"boolean"}},"required":["fail"]})
229 }
230 fn capabilities(&self) -> Vec<ToolCapability> {
231 vec![ToolCapability::ReadOnly]
232 }
233 async fn execute(&self, input: Value, _context: &ToolContext) -> Result<ToolResult, ToolError> {
234 Ok(if input["fail"] == true {
235 ToolResult::error("fixture failure")
236 } else {
237 ToolResult::success("fixture success")
238 })
239 }
240 }
241
242 #[tokio::test]
243 async fn returned_tool_failure_reaches_next_model_request_as_error() {
244 let workspace = tempdir().unwrap();
245 let mock = Arc::new(MockLlmClient::new(vec![
246 tool_batch_turn(&[
247 ("call-failure", "returned_result", r#"{"fail":true}"#),
248 ("call-success", "returned_result", r#"{"fail":false}"#),
249 ]),
250 canned::simple_text_turn("Both tool results received."),
251 ]));
252 let (mut engine, handle) = Engine::new_with_model_client(
253 deterministic_engine_config(workspace.path()),
254 &Config::default(),
255 mock.clone(),
256 );
257 let mut registry = crate::tools::ToolRegistry::new(ToolContext::new(workspace.path()));
258 registry.register(Arc::new(ReturnedResultTool));
259 let tools = Some(registry.to_api_tools_with_cache(true));
260 let surface = test_tool_surface(&engine, registry, tools, AppMode::Agent);
261 let mut turn = TurnContext::new(4);
262 let (status, error) = engine.run_turn(&mut turn, surface, None, None).await;
263 assert_eq!(status, TurnOutcomeStatus::Completed, "{error:?}");
264 let requests = mock.captured_requests();
265 assert_eq!(requests.len(), 2);
266 assert_eq!(
267 guardian_tool_results(&requests[1], "call-failure"),
268 vec![("fixture failure", Some(true))]
269 );
270 assert_eq!(
271 guardian_tool_results(&requests[1], "call-success"),
272 vec![("fixture success", None)]
273 );
274 let mut events = handle.rx_event.write().await;
275 let results = std::iter::from_fn(|| events.try_recv().ok())
276 .filter_map(|event| match event {
277 Event::ToolCallComplete {
278 model_call: Some(model_call),
279 result,
280 ..
281 } => Some((model_call.provider_id, result)),
282 _ => None,
283 })
284 .collect::<Vec<_>>();
285 assert!(
286 results.iter().any(|(id, result)| id == "call-failure"
287 && result.as_ref().is_ok_and(|output| !output.success)),
288 "this must cover Ok(ToolResult::error), not Err(ToolError)"
289 );
290 }
291
291 lines RUST