返回 CodeWhale
revert_turn.rs
根目录 / crates / tui / src / tools / revert_turn.rs
1 //! `revert_turn` — agent-callable tool that rewinds the workspace to a
2 //! prior pre-turn snapshot.
3 //!
4 //! The model invokes this when the user says something like "undo the
5 //! last edit" or "roll back". It mirrors `/restore` but speaks JSON and
6 //! takes a turn-offset (default 1 = previous turn) instead of a list
7 //! index, so the model doesn't have to count entries.
8 //!
9 //! Approval is `Required` because this mutates the workspace.
10 //!
11 //! Known limit: like `/restore`, this is a whole-tree rollback. It restores
12 //! every path to the pre-turn snapshot, so an edit made after that turn (the
13 //! user's, or another session's) is overwritten too; it does not have the
14 //! path scoping, changed-since refusal, uncapped lookup or fork-inherited
15 //! restore points of `/undo` (#6644).
16
17 use async_trait::async_trait;
18 use serde_json::{Value, json};
19
20 use super::spec::{
21 ApprovalRequirement, ToolCapability, ToolContext, ToolError, ToolResult, ToolSpec, optional_u64,
22 };
23 use crate::snapshot::SnapshotRepo;
24
25 /// Default offset: revert the most-recent turn (i.e. the last `pre-turn:*`
26 /// snapshot in history).
27 const DEFAULT_OFFSET: u64 = 1;
28 /// Hard cap so the model can't ask to roll back to the dawn of time.
29 const MAX_OFFSET: u64 = 50;
30
31 pub struct RevertTurnTool;
32
33 #[async_trait]
34 impl ToolSpec for RevertTurnTool {
35 fn name(&self) -> &str {
36 "revert_turn"
37 }
38
39 fn description(&self) -> &str {
40 "Roll back the workspace files to the snapshot taken before a recent turn. \
41 Use when the user explicitly asks to undo, revert, or roll back the most recent edits. \
42 `turn_offset` is 1-based: 1 reverts the most recent turn (max 50). Conversation \
43 history is NOT modified. The whole workspace is restored, so later edits, including \
44 the user's own, are overwritten."
45 }
46
47 fn input_schema(&self) -> Value {
48 json!({
49 "type": "object",
50 "properties": {
51 "turn_offset": {
52 "type": "integer",
53 "minimum": 1,
54 "maximum": MAX_OFFSET,
55 "description": "How many turns back to revert (default 1)."
56 }
57 },
58 "additionalProperties": false
59 })
60 }
61
62 fn capabilities(&self) -> Vec<ToolCapability> {
63 vec![
64 ToolCapability::WritesFiles,
65 ToolCapability::RequiresApproval,
66 ]
67 }
68
69 fn approval_requirement(&self) -> ApprovalRequirement {
70 ApprovalRequirement::Required
71 }
72
73 async fn execute(&self, input: Value, context: &ToolContext) -> Result<ToolResult, ToolError> {
74 let offset = optional_u64(&input, "turn_offset", DEFAULT_OFFSET)?;
75 if offset == 0 || offset > MAX_OFFSET {
76 return Err(ToolError::invalid_input(format!(
77 "turn_offset must be between 1 and {MAX_OFFSET}; got {offset}",
78 )));
79 }
80
81 let workspace = context.workspace.clone();
82 let label = format!("revert_turn(offset={offset})");
83 let session = context.state_namespace.clone();
84 #[cfg(test)]
85 let env_scope = crate::test_support::env_scope_ticket();
86 let result = tokio::task::spawn_blocking(move || -> Result<String, String> {
87 #[cfg(test)]
88 let _env_scope = crate::test_support::join_env_scope(env_scope);
89 let repo = SnapshotRepo::open_or_init(&workspace)
90 .map_err(|e| format!("Snapshot repo init failed: {e}"))?;
91 // Find pre-turn:* snapshots only — those mark the start of
92 // each turn, which is the right rollback target. We pull a
93 // generous list and filter so the model's `turn_offset` is
94 // counted in turns, not raw snapshots.
95 let snapshots = repo
96 .list((MAX_OFFSET as usize).saturating_mul(2) + 16)
97 .map_err(|e| format!("Snapshot list failed: {e}"))?;
98 let pre_turns: Vec<_> = snapshots
99 .into_iter()
100 .filter(|s| s.label.starts_with("pre-turn:"))
101 // Session ownership must be exact. Legacy snapshots are not
102 // safe automatic targets because they may belong to another
103 // conversation that used this workspace.
104 .filter(|s| s.session_id.as_deref() == Some(session.as_str()))
105 .collect();
106 let target = pre_turns
107 .get((offset - 1) as usize)
108 .ok_or_else(|| {
109 format!(
110 "Only {} current-session pre-turn snapshot(s) exist; turn_offset={offset} is out of range.",
111 pre_turns.len(),
112 )
113 })?
114 .clone();
115 if repo
116 .work_tree_matches_snapshot(&target.id)
117 .map_err(|e| format!("Snapshot comparison failed: {e}"))?
118 {
119 return Err(format!(
120 "NoSnapshotForTurn: target '{}' ({}) already matches the current workspace. \
121 Revert operates at completed turn boundaries; there is no distinct later snapshot to restore.",
122 target.label,
123 short_sha(target.id.as_str()),
124 ));
125 }
126 repo.restore(&target.id)
127 .map_err(|e| format!("Restore failed: {e}"))?;
128 Ok(format!(
129 "{label}: restored '{}' ({}). Workspace files reverted; conversation unchanged.",
130 target.label,
131 short_sha(target.id.as_str()),
132 ))
133 })
134 .await
135 .map_err(|e| ToolError::execution_failed(format!("revert_turn join failed: {e}")))?;
136
137 match result {
138 Ok(msg) => Ok(ToolResult::success(msg)),
139 Err(e) => Ok(ToolResult::error(e)),
140 }
141 }
142 }
143
144 fn short_sha(sha: &str) -> &str {
145 &sha[..sha.len().min(8)]
146 }
147
148 #[cfg(test)]
149 mod tests {
150 use super::*;
151 use tempfile::tempdir;
152
153 /// Seals the user's home onto `home` for the duration of the test. Pinning
154 /// `HOME` alone left an ambient `CODEWHALE_HOME` in charge of the snapshot
155 /// store.
156 fn scoped_home(home: &std::path::Path) -> crate::test_support::SealedHome {
157 crate::test_support::SealedHome::at(home)
158 }
159
160 #[tokio::test]
161 async fn revert_turn_default_offset_restores_pre_turn_one() {
162 let tmp = tempdir().unwrap();
163 let workspace = tmp.path().join("ws");
164 std::fs::create_dir_all(&workspace).unwrap();
165 let _guard = scoped_home(tmp.path());
166
167 // Setup: create pre-turn:1, post-turn:1 with file modifications.
168 let repo = SnapshotRepo::open_or_init(&workspace).unwrap();
169 std::fs::write(workspace.join("a.txt"), b"original").unwrap();
170 repo.snapshot_with_session("pre-turn:1", Some("workspace"))
171 .unwrap();
172 std::fs::write(workspace.join("a.txt"), b"modified").unwrap();
173 repo.snapshot_with_session("post-turn:1", Some("workspace"))
174 .unwrap();
175
176 let tool = RevertTurnTool;
177 let ctx = ToolContext::new(workspace.clone());
178 let r = tool.execute(json!({}), &ctx).await.expect("execute");
179 assert!(r.success, "expected success: {r:?}");
180
181 let content = std::fs::read_to_string(workspace.join("a.txt")).unwrap();
182 assert_eq!(content, "original");
183 }
184
185 #[tokio::test]
186 async fn revert_turn_invalid_offset_rejected() {
187 let tmp = tempdir().unwrap();
188 let workspace = tmp.path().join("ws");
189 std::fs::create_dir_all(&workspace).unwrap();
190 let _guard = scoped_home(tmp.path());
191
192 let tool = RevertTurnTool;
193 let ctx = ToolContext::new(workspace);
194 let r = tool.execute(json!({"turn_offset": 0}), &ctx).await;
195 assert!(r.is_err());
196 }
197
198 #[tokio::test]
199 async fn revert_turn_rejects_snapshot_matching_current_workspace() {
200 let tmp = tempdir().unwrap();
201 let workspace = tmp.path().join("ws");
202 std::fs::create_dir_all(&workspace).unwrap();
203 let _guard = scoped_home(tmp.path());
204
205 let repo = SnapshotRepo::open_or_init(&workspace).unwrap();
206 std::fs::write(workspace.join("a.txt"), b"unchanged").unwrap();
207 repo.snapshot_with_session("pre-turn:1", Some("workspace"))
208 .unwrap();
209
210 let tool = RevertTurnTool;
211 let ctx = ToolContext::new(workspace);
212 let r = tool.execute(json!({}), &ctx).await.expect("execute");
213 assert!(!r.success);
214 assert!(r.content.contains("NoSnapshotForTurn"), "{}", r.content);
215 }
216
217 #[tokio::test]
218 async fn revert_turn_no_snapshots_returns_error_result() {
219 let tmp = tempdir().unwrap();
220 let workspace = tmp.path().join("ws");
221 std::fs::create_dir_all(&workspace).unwrap();
222 let _guard = scoped_home(tmp.path());
223
224 let tool = RevertTurnTool;
225 let ctx = ToolContext::new(workspace);
226 let r = tool.execute(json!({}), &ctx).await.expect("execute");
227 assert!(!r.success);
228 assert!(r.content.contains("out of range"));
229 }
230
231 #[tokio::test]
232 async fn revert_turn_rejects_legacy_and_foreign_session_snapshots() {
233 let tmp = tempdir().unwrap();
234 let workspace = tmp.path().join("ws");
235 std::fs::create_dir_all(&workspace).unwrap();
236 let _guard = scoped_home(tmp.path());
237
238 let repo = SnapshotRepo::open_or_init(&workspace).unwrap();
239 std::fs::write(workspace.join("a.txt"), b"legacy").unwrap();
240 repo.snapshot("pre-turn:legacy").unwrap();
241 std::fs::write(workspace.join("a.txt"), b"foreign").unwrap();
242 repo.snapshot_with_session("pre-turn:foreign", Some("other-session"))
243 .unwrap();
244 std::fs::write(workspace.join("a.txt"), b"current").unwrap();
245
246 let tool = RevertTurnTool;
247 let ctx = ToolContext::new(workspace.clone());
248 let r = tool.execute(json!({}), &ctx).await.expect("execute");
249 assert!(!r.success);
250 assert!(
251 r.content.contains("Only 0 current-session"),
252 "{}",
253 r.content
254 );
255 assert_eq!(
256 std::fs::read_to_string(workspace.join("a.txt")).unwrap(),
257 "current"
258 );
259 }
260 }
261
261 lines RUST