diff --git a/crates/worker/src/worker.rs b/crates/worker/src/worker.rs index ae8d6aba..df165038 100644 --- a/crates/worker/src/worker.rs +++ b/crates/worker/src/worker.rs @@ -2889,15 +2889,23 @@ impl Worker { }), compacted_from: None, }; + let mut initial_entries = vec![entry.clone()]; + if let Some(checkpoint) = + active_run_checkpoint_entry(w.active_run_turn_count(), w.turn_count()) + { + initial_entries.push(checkpoint); + } self.store - .create_segment(loc.session_id, fork_segment_id, &[entry.clone()]) + .create_segment(loc.session_id, fork_segment_id, &initial_entries) .map_err(WorkerError::from)?; self.segment_state.set_location(SegmentLocation { session_id: loc.session_id, segment_id: fork_segment_id, }); - self.segment_state.set_entries_written(1); - self.sink.reset_with_initial(entry); + self.segment_state + .set_entries_written(initial_entries.len()); + self.sink + .reset_with_initial_entries(initial_entries.clone()); if self.scope_allocation.is_some() { worker_allocation::update_segment(&self.manifest.worker.name, fork_segment_id)?; } @@ -7031,6 +7039,65 @@ mod build_summary_prompt_tests { assert_eq!(restored.active_run_turn_count, None); } + #[tokio::test] + async fn auto_fork_checkpoints_interrupted_run_budget_for_restore() { + let dir = tempfile::tempdir().unwrap(); + let store = session_store::FsStore::new(dir.path().join("sessions")).unwrap(); + let mut worker = Worker::new( + minimal_manifest(), + Engine::new(NoopClient), + store, + WorkerWorkspaceContext::no_workspace(), + WorkerFilesystemAuthority::None, + Scope::empty(), + ) + .await + .unwrap(); + worker.ensure_segment_head().unwrap(); + worker.engine_mut().set_turn_count(7); + worker.engine_mut().set_last_run_interrupted(true); + worker.engine_mut().set_active_run_turn_count(Some(3)); + + let session_id = worker.session_id(); + let source_segment_id = worker.segment_id(); + worker + .store() + .append( + session_id, + source_segment_id, + &LogEntry::Extension { + ts: segment_log::now_millis(), + domain: "test.auto_fork_drift".into(), + payload: serde_json::json!({}), + }, + ) + .unwrap(); + + worker.ensure_segment_head().unwrap(); + + let fork_segment_id = worker.segment_id(); + assert_ne!(fork_segment_id, source_segment_id); + let fork_entries = worker + .store() + .read_all(session_id, fork_segment_id) + .unwrap(); + assert!(matches!( + fork_entries.as_slice(), + [ + LogEntry::SegmentStart { .. }, + LogEntry::ActiveRunCheckpoint { + active_turn_count: 3, + total_turn_count: 7, + .. + } + ] + )); + let restored = session_store::collect_state(&fork_entries); + assert!(restored.last_run_interrupted); + assert_eq!(restored.turn_count, 7); + assert_eq!(restored.active_run_turn_count, Some(3)); + } + #[tokio::test] async fn flow_transition_feature_installs_runtime_local_coordinator() { let dir = tempfile::tempdir().unwrap();