prodigy/cook/execution/mapreduce/checkpoint/
environment.rs1use super::effects::storage::CheckpointStorageEnv;
7use super::pure::triggers::CheckpointTriggerConfig;
8use super::{CheckpointStorage, MapReduceCheckpoint};
9use chrono::{DateTime, Utc};
10use std::path::PathBuf;
11use std::sync::atomic::{AtomicUsize, Ordering};
12use std::sync::Arc;
13use stillwater::{asks, local, Effect};
14use tokio::sync::RwLock;
15
16#[derive(Clone)]
21pub struct CheckpointEnv {
22 pub job_id: String,
24 pub storage: Arc<dyn CheckpointStorage>,
26 pub current_checkpoint: Arc<RwLock<Option<MapReduceCheckpoint>>>,
28 pub storage_path: PathBuf,
30 pub trigger_config: CheckpointTriggerConfig,
32 pub items_since_checkpoint: Arc<AtomicUsize>,
34 pub last_checkpoint_time: Arc<RwLock<DateTime<Utc>>>,
36 pub enabled: bool,
38}
39
40impl CheckpointEnv {
41 pub fn new(
43 job_id: String,
44 storage: Arc<dyn CheckpointStorage>,
45 storage_path: PathBuf,
46 trigger_config: CheckpointTriggerConfig,
47 ) -> Self {
48 Self {
49 job_id,
50 storage,
51 current_checkpoint: Arc::new(RwLock::new(None)),
52 storage_path,
53 trigger_config,
54 items_since_checkpoint: Arc::new(AtomicUsize::new(0)),
55 last_checkpoint_time: Arc::new(RwLock::new(Utc::now())),
56 enabled: true,
57 }
58 }
59
60 pub fn disabled() -> Self {
62 use super::FileCheckpointStorage;
63
64 let temp_path = std::env::temp_dir().join("prodigy_disabled_checkpoints");
65 let storage: Arc<dyn CheckpointStorage> =
66 Arc::new(FileCheckpointStorage::new(temp_path.clone(), false));
67
68 Self {
69 job_id: "disabled".to_string(),
70 storage,
71 current_checkpoint: Arc::new(RwLock::new(None)),
72 storage_path: temp_path,
73 trigger_config: CheckpointTriggerConfig::none(),
74 items_since_checkpoint: Arc::new(AtomicUsize::new(0)),
75 last_checkpoint_time: Arc::new(RwLock::new(Utc::now())),
76 enabled: false,
77 }
78 }
79
80 pub fn increment_items(&self, count: usize) {
82 self.items_since_checkpoint
83 .fetch_add(count, Ordering::SeqCst);
84 }
85
86 pub fn reset_items(&self) {
88 self.items_since_checkpoint.store(0, Ordering::SeqCst);
89 }
90
91 pub fn get_items(&self) -> usize {
93 self.items_since_checkpoint.load(Ordering::Acquire)
94 }
95}
96
97impl CheckpointStorageEnv for CheckpointEnv {
99 fn storage(&self) -> Arc<dyn CheckpointStorage> {
100 Arc::clone(&self.storage)
101 }
102
103 fn current_checkpoint(&self) -> Arc<RwLock<Option<MapReduceCheckpoint>>> {
104 Arc::clone(&self.current_checkpoint)
105 }
106
107 fn storage_path(&self) -> PathBuf {
108 self.storage_path.clone()
109 }
110}
111
112#[derive(Debug, Clone, thiserror::Error)]
118pub enum CheckpointError {
119 #[error("Checkpointing is disabled")]
120 Disabled,
121
122 #[error("No checkpoint available")]
123 NoCheckpoint,
124
125 #[error("Storage error: {0}")]
126 Storage(String),
127
128 #[error("Validation error: {0}")]
129 Validation(String),
130}
131
132pub fn get_checkpoint_job_id(
134) -> impl Effect<Output = String, Error = CheckpointError, Env = CheckpointEnv> {
135 asks(|env: &CheckpointEnv| env.job_id.clone())
136}
137
138pub fn get_trigger_config(
140) -> impl Effect<Output = CheckpointTriggerConfig, Error = CheckpointError, Env = CheckpointEnv> {
141 asks(|env: &CheckpointEnv| env.trigger_config.clone())
142}
143
144pub fn get_checkpoint_storage(
146) -> impl Effect<Output = Arc<dyn CheckpointStorage>, Error = CheckpointError, Env = CheckpointEnv>
147{
148 asks(|env: &CheckpointEnv| env.storage.clone())
149}
150
151pub fn get_items_since_checkpoint(
153) -> impl Effect<Output = usize, Error = CheckpointError, Env = CheckpointEnv> {
154 asks(|env: &CheckpointEnv| env.items_since_checkpoint.load(Ordering::Acquire))
155}
156
157pub fn is_checkpointing_enabled(
159) -> impl Effect<Output = bool, Error = CheckpointError, Env = CheckpointEnv> {
160 asks(|env: &CheckpointEnv| env.enabled)
161}
162
163pub fn get_checkpoint_storage_path(
165) -> impl Effect<Output = PathBuf, Error = CheckpointError, Env = CheckpointEnv> {
166 asks(|env: &CheckpointEnv| env.storage_path.clone())
167}
168
169pub fn with_checkpointing_disabled<E>(
175 effect: E,
176) -> impl Effect<Output = E::Output, Error = CheckpointError, Env = CheckpointEnv>
177where
178 E: Effect<Error = CheckpointError, Env = CheckpointEnv>,
179{
180 local(
181 |env: &CheckpointEnv| CheckpointEnv {
182 enabled: false,
183 ..env.clone()
184 },
185 effect,
186 )
187}
188
189pub fn with_trigger_config<E>(
191 config: CheckpointTriggerConfig,
192 effect: E,
193) -> impl Effect<Output = E::Output, Error = CheckpointError, Env = CheckpointEnv>
194where
195 E: Effect<Error = CheckpointError, Env = CheckpointEnv>,
196{
197 local(
198 move |env: &CheckpointEnv| CheckpointEnv {
199 trigger_config: config.clone(),
200 ..env.clone()
201 },
202 effect,
203 )
204}
205
206#[derive(Clone)]
212pub struct MockCheckpointEnvBuilder {
213 job_id: String,
214 trigger_config: CheckpointTriggerConfig,
215 enabled: bool,
216 initial_checkpoint: Option<MapReduceCheckpoint>,
217}
218
219impl Default for MockCheckpointEnvBuilder {
220 fn default() -> Self {
221 Self::new()
222 }
223}
224
225impl MockCheckpointEnvBuilder {
226 pub fn new() -> Self {
228 Self {
229 job_id: "mock-job-123".to_string(),
230 trigger_config: CheckpointTriggerConfig::default(),
231 enabled: true,
232 initial_checkpoint: None,
233 }
234 }
235
236 pub fn with_job_id(mut self, job_id: impl Into<String>) -> Self {
238 self.job_id = job_id.into();
239 self
240 }
241
242 pub fn with_trigger_config(mut self, config: CheckpointTriggerConfig) -> Self {
244 self.trigger_config = config;
245 self
246 }
247
248 pub fn disabled(mut self) -> Self {
250 self.enabled = false;
251 self
252 }
253
254 pub fn with_checkpoint(mut self, checkpoint: MapReduceCheckpoint) -> Self {
256 self.initial_checkpoint = Some(checkpoint);
257 self
258 }
259
260 pub fn build(self) -> CheckpointEnv {
262 use super::FileCheckpointStorage;
263
264 let temp_dir = std::env::temp_dir().join(format!("prodigy_mock_{}", self.job_id));
265 let _ = std::fs::create_dir_all(&temp_dir);
266 let storage: Arc<dyn CheckpointStorage> =
267 Arc::new(FileCheckpointStorage::new(temp_dir.clone(), true));
268
269 let mut env = CheckpointEnv::new(self.job_id, storage, temp_dir, self.trigger_config);
270 env.enabled = self.enabled;
271
272 if let Some(checkpoint) = self.initial_checkpoint {
273 let current_checkpoint = Arc::clone(&env.current_checkpoint);
274 tokio::task::block_in_place(|| {
275 tokio::runtime::Handle::current().block_on(async {
276 *current_checkpoint.write().await = Some(checkpoint);
277 })
278 });
279 }
280
281 env
282 }
283}
284
285#[cfg(test)]
286mod tests {
287 use super::*;
288
289 #[tokio::test]
290 async fn test_get_checkpoint_job_id() {
291 let env = MockCheckpointEnvBuilder::new()
292 .with_job_id("my-test-job")
293 .build();
294
295 let effect = get_checkpoint_job_id();
296 let result = effect.run(&env).await;
297
298 assert!(result.is_ok());
299 assert_eq!(result.unwrap(), "my-test-job");
300 }
301
302 #[tokio::test]
303 async fn test_get_trigger_config() {
304 let config = CheckpointTriggerConfig::item_interval(10);
305 let env = MockCheckpointEnvBuilder::new()
306 .with_trigger_config(config)
307 .build();
308
309 let effect = get_trigger_config();
310 let result = effect.run(&env).await;
311
312 assert!(result.is_ok());
313 assert_eq!(result.unwrap().agent_completion_interval, Some(10));
314 }
315
316 #[tokio::test]
317 async fn test_is_checkpointing_enabled() {
318 let enabled_env = MockCheckpointEnvBuilder::new().build();
319 let disabled_env = MockCheckpointEnvBuilder::new().disabled().build();
320
321 let effect = is_checkpointing_enabled();
323 assert!(effect.run(&enabled_env).await.unwrap());
324
325 let effect = is_checkpointing_enabled();
327 assert!(!effect.run(&disabled_env).await.unwrap());
328 }
329
330 #[tokio::test]
331 async fn test_with_checkpointing_disabled() {
332 let env = MockCheckpointEnvBuilder::new().build();
333
334 let effect = is_checkpointing_enabled();
336 assert!(effect.run(&env).await.unwrap());
337
338 let effect = with_checkpointing_disabled(is_checkpointing_enabled());
340 assert!(!effect.run(&env).await.unwrap());
341
342 let effect = is_checkpointing_enabled();
344 assert!(effect.run(&env).await.unwrap());
345 }
346
347 #[tokio::test]
348 async fn test_with_trigger_config_override() {
349 let env = MockCheckpointEnvBuilder::new()
350 .with_trigger_config(CheckpointTriggerConfig::item_interval(5))
351 .build();
352
353 let effect = get_trigger_config();
355 assert_eq!(
356 effect.run(&env).await.unwrap().agent_completion_interval,
357 Some(5)
358 );
359
360 let new_config = CheckpointTriggerConfig::item_interval(100);
362 let effect = with_trigger_config(new_config, get_trigger_config());
363 assert_eq!(
364 effect.run(&env).await.unwrap().agent_completion_interval,
365 Some(100)
366 );
367
368 let effect = get_trigger_config();
370 assert_eq!(
371 effect.run(&env).await.unwrap().agent_completion_interval,
372 Some(5)
373 );
374 }
375
376 #[test]
377 fn test_checkpoint_env_increment_items() {
378 let env = MockCheckpointEnvBuilder::new().build();
379
380 assert_eq!(env.get_items(), 0);
381 env.increment_items(5);
382 assert_eq!(env.get_items(), 5);
383 env.increment_items(3);
384 assert_eq!(env.get_items(), 8);
385 env.reset_items();
386 assert_eq!(env.get_items(), 0);
387 }
388}