1use sha2::{Digest, Sha256};
12use std::collections::HashSet;
13use std::fs;
14use std::io::{BufRead, Write};
15use std::path::{Path, PathBuf};
16
17pub const MAX_REWIND_BLOB_BYTES: u64 = 32 * 1024 * 1024;
19
20pub const PROMPT_PREVIEW_CHARS: usize = 120;
22
23#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
25pub struct RewindFileEntry {
26 pub rel_path: String,
28 #[serde(default, skip_serializing_if = "Option::is_none")]
30 pub blob_id: Option<String>,
31 pub existed: bool,
33}
34
35#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
37pub struct RewindPointMeta {
38 pub prompt_index: usize,
40 pub created_at: u64,
41 #[serde(default)]
43 pub prompt_preview: String,
44 #[serde(default)]
46 pub files: Vec<RewindFileEntry>,
47 #[serde(default)]
49 pub created_paths: Vec<String>,
50}
51
52#[derive(Debug, Clone, Default, PartialEq, Eq)]
54pub struct RestoreSummary {
55 pub restored: usize,
56 pub deleted: usize,
57 pub skipped: usize,
58 pub errors: Vec<String>,
59}
60
61impl RestoreSummary {
62 pub fn total_changes(&self) -> usize {
63 self.restored + self.deleted
64 }
65
66 pub fn is_empty(&self) -> bool {
67 self.restored == 0 && self.deleted == 0
68 }
69}
70
71#[derive(Debug, Clone)]
73pub struct RewindStore {
74 root: PathBuf,
76 project_root: PathBuf,
77 dirty: HashSet<PathBuf>,
79 turn_written: HashSet<PathBuf>,
81}
82
83impl RewindStore {
84 pub fn new(
85 data_dir: impl AsRef<Path>,
86 session_id: &str,
87 project_root: impl AsRef<Path>,
88 ) -> Self {
89 let root = data_dir
90 .as_ref()
91 .join("sessions")
92 .join(session_id)
93 .join("rewind");
94 Self {
95 root,
96 project_root: project_root.as_ref().to_path_buf(),
97 dirty: HashSet::new(),
98 turn_written: HashSet::new(),
99 }
100 }
101
102 pub fn root(&self) -> &Path {
103 &self.root
104 }
105
106 fn blobs_dir(&self) -> PathBuf {
107 self.root.join("blobs")
108 }
109
110 fn points_path(&self) -> PathBuf {
111 self.root.join("points.jsonl")
112 }
113
114 fn ensure_dirs(&self) -> std::io::Result<()> {
115 fs::create_dir_all(self.blobs_dir())
116 }
117
118 pub fn note_written_paths(&mut self, paths: impl IntoIterator<Item = PathBuf>) {
120 for p in paths {
121 let abs = if p.is_absolute() {
122 p
123 } else {
124 self.project_root.join(&p)
125 };
126 self.dirty.insert(abs.clone());
127 self.turn_written.insert(abs);
128 }
129 }
130
131 pub fn clear_turn_written(&mut self) {
133 self.turn_written.clear();
134 }
135
136 pub fn is_dirty(&self, abs: &Path) -> bool {
138 let abs = if abs.is_absolute() {
139 abs.to_path_buf()
140 } else {
141 self.project_root.join(abs)
142 };
143 self.dirty.contains(&abs)
144 }
145
146 pub fn next_prompt_index(&self) -> usize {
148 self.load_points().len()
149 }
150
151 pub fn ensure_pre_write_capture(
157 &mut self,
158 paths: impl IntoIterator<Item = PathBuf>,
159 ) -> std::io::Result<()> {
160 let mut points = self.load_points();
161 let Some(point) = points.last_mut() else {
162 self.note_written_paths(paths);
164 return Ok(());
165 };
166
167 let mut changed = false;
168 for p in paths {
169 let abs = if p.is_absolute() {
170 p
171 } else {
172 self.project_root.join(&p)
173 };
174 if self.dirty.contains(&abs) {
175 self.turn_written.insert(abs);
176 continue;
177 }
178 let Some(rel) = rel_path_for(&self.project_root, &abs) else {
179 continue;
180 };
181 if point.files.iter().any(|f| f.rel_path == rel) {
183 self.dirty.insert(abs.clone());
184 self.turn_written.insert(abs);
185 continue;
186 }
187 let _ = self.ensure_dirs();
188 let entry = capture_file_entry(&abs, &rel, &self.blobs_dir())?;
189 point.files.push(entry);
190 self.dirty.insert(abs.clone());
191 self.turn_written.insert(abs);
192 changed = true;
193 }
194 if changed {
195 rewrite_points(&self.points_path(), &points)?;
196 }
197 Ok(())
198 }
199
200 pub fn capture_point(
202 &mut self,
203 prompt_index: usize,
204 prompt_text: &str,
205 created_at: u64,
206 ) -> std::io::Result<RewindPointMeta> {
207 self.ensure_dirs()?;
208 let mut files = Vec::new();
209 let dirty: Vec<PathBuf> = self.dirty.iter().cloned().collect();
210 for abs in dirty {
211 let Some(rel) = rel_path_for(&self.project_root, &abs) else {
212 continue;
213 };
214 let entry = capture_file_entry(&abs, &rel, &self.blobs_dir())?;
215 files.push(entry);
216 }
217
218 let point = RewindPointMeta {
221 prompt_index,
222 created_at,
223 prompt_preview: truncate_preview(prompt_text, PROMPT_PREVIEW_CHARS),
224 files,
225 created_paths: Vec::new(),
226 };
227 append_point(&self.points_path(), &point)?;
228 Ok(point)
229 }
230
231 pub fn finalize_turn_created_paths(&mut self, prompt_index: usize) -> std::io::Result<()> {
234 let mut points = self.load_points();
235 let Some(point) = points.iter_mut().find(|p| p.prompt_index == prompt_index) else {
236 return Ok(());
237 };
238 let pre_existed: HashSet<&str> = point
239 .files
240 .iter()
241 .filter(|f| f.existed)
242 .map(|f| f.rel_path.as_str())
243 .collect();
244 let mut created = Vec::new();
245 for abs in &self.turn_written {
246 let Some(rel) = rel_path_for(&self.project_root, abs) else {
247 continue;
248 };
249 if !pre_existed.contains(rel.as_str()) {
250 created.push(rel);
251 }
252 }
253 created.sort();
254 created.dedup();
255 point.created_paths = created;
256 rewrite_points(&self.points_path(), &points)?;
257 self.clear_turn_written();
258 Ok(())
259 }
260
261 pub fn load_points(&self) -> Vec<RewindPointMeta> {
262 let path = self.points_path();
263 let Ok(file) = fs::File::open(&path) else {
264 return Vec::new();
265 };
266 let reader = std::io::BufReader::new(file);
267 let mut points = Vec::new();
268 for line in reader.lines() {
269 let Ok(line) = line else {
270 continue;
271 };
272 let line = line.trim();
273 if line.is_empty() {
274 continue;
275 }
276 if let Ok(p) = serde_json::from_str::<RewindPointMeta>(line) {
277 points.push(p);
278 }
279 }
280 points
281 }
282
283 pub fn restore_to(&mut self, prompt_index: usize) -> RestoreSummary {
287 let points = self.load_points();
288 let mut summary = RestoreSummary::default();
289
290 let point = points.iter().find(|p| p.prompt_index == prompt_index);
291 if let Some(point) = point {
292 for entry in &point.files {
293 let abs = self.project_root.join(normalize_rel(&entry.rel_path));
294 if entry.existed {
295 let Some(blob_id) = entry.blob_id.as_ref() else {
296 summary.skipped += 1;
297 summary.errors.push(format!(
298 "skip restore `{}`: no blob (too large or missing)",
299 entry.rel_path
300 ));
301 continue;
302 };
303 let blob_path = self.blobs_dir().join(blob_id);
304 match fs::read(&blob_path) {
305 Ok(bytes) => {
306 if let Some(parent) = abs.parent() {
307 let _ = fs::create_dir_all(parent);
308 }
309 match fs::write(&abs, &bytes) {
310 Ok(()) => summary.restored += 1,
311 Err(e) => summary
312 .errors
313 .push(format!("write `{}`: {e}", entry.rel_path)),
314 }
315 }
316 Err(e) => summary
317 .errors
318 .push(format!("read blob for `{}`: {e}", entry.rel_path)),
319 }
320 } else if abs.is_file() {
321 match fs::remove_file(&abs) {
322 Ok(()) => summary.deleted += 1,
323 Err(e) => summary
324 .errors
325 .push(format!("delete `{}`: {e}", entry.rel_path)),
326 }
327 }
328 }
329 } else {
330 tracing::debug!(
332 prompt_index,
333 "rewind: no snapshot for index; history-only restore for files"
334 );
335 }
336
337 let mut to_delete: HashSet<String> = HashSet::new();
339 for p in points.iter().filter(|p| p.prompt_index >= prompt_index) {
340 for c in &p.created_paths {
341 to_delete.insert(c.clone());
342 }
343 }
344 let pre_paths: HashSet<String> = point
346 .map(|p| {
347 p.files
348 .iter()
349 .map(|f| f.rel_path.clone())
350 .collect::<HashSet<_>>()
351 })
352 .unwrap_or_default();
353 for p in points.iter().filter(|p| p.prompt_index > prompt_index) {
354 for f in &p.files {
355 if !pre_paths.contains(&f.rel_path) {
356 to_delete.insert(f.rel_path.clone());
357 }
358 }
359 }
360 for rel in to_delete {
361 let abs = self.project_root.join(normalize_rel(&rel));
362 if abs.is_file() {
363 match fs::remove_file(&abs) {
364 Ok(()) => summary.deleted += 1,
365 Err(e) => summary.errors.push(format!("delete created `{rel}`: {e}")),
366 }
367 }
368 }
369
370 let kept: Vec<RewindPointMeta> = points
372 .into_iter()
373 .filter(|p| p.prompt_index < prompt_index)
374 .collect();
375 let _ = rewrite_points(&self.points_path(), &kept);
376
377 let mut still = HashSet::new();
379 for p in &kept {
380 for f in &p.files {
381 still.insert(self.project_root.join(normalize_rel(&f.rel_path)));
382 }
383 for c in &p.created_paths {
384 still.insert(self.project_root.join(normalize_rel(c)));
385 }
386 }
387 self.dirty = still;
388 self.turn_written.clear();
389
390 summary
391 }
392
393 pub fn clear_all(&mut self) -> std::io::Result<()> {
395 if self.root.exists() {
396 fs::remove_dir_all(&self.root)?;
397 }
398 self.dirty.clear();
399 self.turn_written.clear();
400 Ok(())
401 }
402}
403
404fn truncate_preview(text: &str, max_chars: usize) -> String {
405 let t = text.trim().replace('\n', " ");
406 if t.chars().count() <= max_chars {
407 return t;
408 }
409 let mut out = String::new();
410 for (i, ch) in t.chars().enumerate() {
411 if i >= max_chars.saturating_sub(1) {
412 break;
413 }
414 out.push(ch);
415 }
416 out.push('…');
417 out
418}
419
420fn normalize_rel(rel: &str) -> String {
421 rel.replace('\\', "/")
422}
423
424fn rel_path_for(project_root: &Path, abs: &Path) -> Option<String> {
425 let abs = abs.canonicalize().unwrap_or_else(|_| abs.to_path_buf());
426 let root = project_root
427 .canonicalize()
428 .unwrap_or_else(|_| project_root.to_path_buf());
429 abs.strip_prefix(&root)
430 .ok()
431 .map(|p| p.to_string_lossy().replace('\\', "/"))
432}
433
434fn capture_file_entry(abs: &Path, rel: &str, blobs_dir: &Path) -> std::io::Result<RewindFileEntry> {
435 if !abs.is_file() {
436 return Ok(RewindFileEntry {
437 rel_path: rel.to_string(),
438 blob_id: None,
439 existed: false,
440 });
441 }
442 let meta = fs::metadata(abs)?;
443 if meta.len() > MAX_REWIND_BLOB_BYTES {
444 tracing::warn!(
445 path = %abs.display(),
446 size = meta.len(),
447 "rewind: skip file larger than cap"
448 );
449 return Ok(RewindFileEntry {
451 rel_path: rel.to_string(),
452 blob_id: None,
453 existed: true,
454 });
455 }
456 let bytes = fs::read(abs)?;
457 let mut hasher = Sha256::new();
458 hasher.update(&bytes);
459 let id = hex::encode(hasher.finalize());
460 let blob_path = blobs_dir.join(&id);
461 if !blob_path.exists() {
462 fs::write(&blob_path, &bytes)?;
463 }
464 Ok(RewindFileEntry {
465 rel_path: rel.to_string(),
466 blob_id: Some(id),
467 existed: true,
468 })
469}
470
471fn append_point(path: &Path, point: &RewindPointMeta) -> std::io::Result<()> {
472 if let Some(parent) = path.parent() {
473 fs::create_dir_all(parent)?;
474 }
475 let mut f = fs::OpenOptions::new()
476 .create(true)
477 .append(true)
478 .open(path)?;
479 let line = serde_json::to_string(point)
480 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
481 writeln!(f, "{line}")?;
482 Ok(())
483}
484
485fn rewrite_points(path: &Path, points: &[RewindPointMeta]) -> std::io::Result<()> {
486 if let Some(parent) = path.parent() {
487 fs::create_dir_all(parent)?;
488 }
489 let tmp = path.with_extension("jsonl.tmp");
490 {
491 let mut f = fs::File::create(&tmp)?;
492 for p in points {
493 let line = serde_json::to_string(p)
494 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
495 writeln!(f, "{line}")?;
496 }
497 f.sync_all()?;
498 }
499 fs::rename(&tmp, path)?;
500 Ok(())
501}
502
503#[cfg(test)]
504mod tests {
505 use super::*;
506 use tempfile::tempdir;
507
508 #[test]
509 fn capture_and_restore_modified_file() {
510 let tmp = tempdir().unwrap();
511 let project = tmp.path().join("proj");
512 let data = tmp.path().join("data");
513 fs::create_dir_all(&project).unwrap();
514 let file = project.join("src/a.rs");
515 fs::create_dir_all(file.parent().unwrap()).unwrap();
516 fs::write(&file, b"v1").unwrap();
517
518 let mut store = RewindStore::new(&data, "session-1", &project);
519 store.note_written_paths([file.clone()]);
520 store.capture_point(0, "please edit a.rs", 100).unwrap();
522
523 fs::write(&file, b"v2-agent").unwrap();
525 store.note_written_paths([file.clone()]);
526 store.finalize_turn_created_paths(0).unwrap();
527
528 store.capture_point(1, "again", 200).unwrap();
530 let summary = store.restore_to(0);
531 assert!(summary.errors.is_empty(), "{:?}", summary.errors);
532 assert_eq!(fs::read_to_string(&file).unwrap(), "v1");
533 assert_eq!(store.load_points().len(), 0);
534 }
535
536 #[test]
537 fn restore_deletes_files_created_after_point() {
538 let tmp = tempdir().unwrap();
539 let project = tmp.path().join("proj");
540 let data = tmp.path().join("data");
541 fs::create_dir_all(&project).unwrap();
542
543 let mut store = RewindStore::new(&data, "session-2", &project);
544 store.capture_point(0, "create b.rs", 100).unwrap();
546 let new_file = project.join("b.rs");
547 fs::write(&new_file, b"new").unwrap();
548 store.note_written_paths([new_file.clone()]);
549 store.finalize_turn_created_paths(0).unwrap();
550
551 assert!(new_file.exists());
552 let summary = store.restore_to(0);
553 assert!(
554 !new_file.exists(),
555 "created file must be deleted on restore"
556 );
557 assert!(summary.deleted >= 1);
558 }
559
560 #[test]
561 fn binary_blob_roundtrip() {
562 let tmp = tempdir().unwrap();
563 let project = tmp.path().join("proj");
564 let data = tmp.path().join("data");
565 fs::create_dir_all(&project).unwrap();
566 let file = project.join("img.bin");
567 let bytes = vec![0u8, 159, 146, 150, 255, 1, 2, 3];
568 fs::write(&file, &bytes).unwrap();
569
570 let mut store = RewindStore::new(&data, "session-3", &project);
571 store.note_written_paths([file.clone()]);
572 store.capture_point(0, "binary", 1).unwrap();
573 fs::write(&file, b"changed").unwrap();
574 store.restore_to(0);
575 assert_eq!(fs::read(&file).unwrap(), bytes);
576 }
577
578 #[test]
579 fn prompt_preview_truncates() {
580 let long = "a".repeat(200);
581 let p = truncate_preview(&long, 20);
582 assert!(p.ends_with('…'));
583 assert!(p.chars().count() <= 20);
584 }
585
586 #[test]
587 fn first_touch_pre_write_capture_restores_existing_file() {
588 let tmp = tempdir().unwrap();
589 let project = tmp.path().join("proj");
590 let data = tmp.path().join("data");
591 fs::create_dir_all(&project).unwrap();
592 let file = project.join("touched.rs");
593 fs::write(&file, b"original").unwrap();
594
595 let mut store = RewindStore::new(&data, "session-4", &project);
596 store.capture_point(0, "edit touched.rs", 1).unwrap();
598 store.ensure_pre_write_capture([file.clone()]).unwrap();
600 fs::write(&file, b"mutated").unwrap();
601 store.note_written_paths([file.clone()]);
602 store.finalize_turn_created_paths(0).unwrap();
603
604 let summary = store.restore_to(0);
605 assert!(summary.errors.is_empty(), "{:?}", summary.errors);
606 assert_eq!(fs::read_to_string(&file).unwrap(), "original");
607 }
608}