Skip to main content

driven/state/
snapshot.rs

1//! Rule Snapshots
2//!
3//! Point-in-time captures of rule state for resumability.
4
5use std::collections::HashMap;
6use std::path::Path;
7
8use crate::{DrivenError, Result};
9
10use super::SharedRules;
11
12/// Snapshot of rule state
13#[derive(Debug, Clone)]
14pub struct RuleSnapshot {
15    /// Snapshot ID
16    pub id: u64,
17    /// Timestamp
18    pub timestamp: u64,
19    /// Version at snapshot time
20    pub version: u64,
21    /// Rule data
22    pub rules: HashMap<u32, Vec<u8>>,
23    /// Metadata
24    pub metadata: HashMap<String, String>,
25}
26
27impl RuleSnapshot {
28    /// Create a new snapshot
29    pub fn new(id: u64) -> Self {
30        Self {
31            id,
32            timestamp: std::time::SystemTime::now()
33                .duration_since(std::time::UNIX_EPOCH)
34                .unwrap_or_default()
35                .as_secs(),
36            version: 0,
37            rules: HashMap::new(),
38            metadata: HashMap::new(),
39        }
40    }
41
42    /// Capture from shared rules
43    pub fn capture(shared: &SharedRules, id: u64) -> Self {
44        let mut snapshot = Self::new(id);
45        snapshot.version = shared.version();
46
47        for rule_id in shared.ids() {
48            if let Some(rule) = shared.get(rule_id) {
49                snapshot.rules.insert(rule_id, rule.as_bytes().to_vec());
50            }
51        }
52
53        snapshot
54    }
55
56    /// Restore to shared rules
57    pub fn restore(&self, shared: &SharedRules) {
58        shared.clear();
59        for (_id, content) in &self.rules {
60            shared.insert(content.clone());
61        }
62    }
63
64    /// Get size in bytes
65    pub fn size(&self) -> usize {
66        self.rules.values().map(|v| v.len()).sum()
67    }
68
69    /// Get rule count
70    pub fn rule_count(&self) -> usize {
71        self.rules.len()
72    }
73
74    /// Serialize to bytes
75    pub fn to_bytes(&self) -> Vec<u8> {
76        let mut output = Vec::new();
77
78        // Header
79        output.extend_from_slice(&self.id.to_le_bytes());
80        output.extend_from_slice(&self.timestamp.to_le_bytes());
81        output.extend_from_slice(&self.version.to_le_bytes());
82        output.extend_from_slice(&(self.rules.len() as u32).to_le_bytes());
83
84        // Rules
85        for (id, content) in &self.rules {
86            output.extend_from_slice(&id.to_le_bytes());
87            output.extend_from_slice(&(content.len() as u32).to_le_bytes());
88            output.extend_from_slice(content);
89        }
90
91        // Metadata count
92        output.extend_from_slice(&(self.metadata.len() as u32).to_le_bytes());
93
94        // Metadata
95        for (key, value) in &self.metadata {
96            output.extend_from_slice(&(key.len() as u16).to_le_bytes());
97            output.extend_from_slice(key.as_bytes());
98            output.extend_from_slice(&(value.len() as u16).to_le_bytes());
99            output.extend_from_slice(value.as_bytes());
100        }
101
102        output
103    }
104
105    /// Deserialize from bytes
106    pub fn from_bytes(data: &[u8]) -> Result<Self> {
107        if data.len() < 28 {
108            return Err(DrivenError::InvalidBinary("Snapshot too small".into()));
109        }
110
111        let mut pos = 0;
112
113        let id = u64::from_le_bytes([
114            data[pos],
115            data[pos + 1],
116            data[pos + 2],
117            data[pos + 3],
118            data[pos + 4],
119            data[pos + 5],
120            data[pos + 6],
121            data[pos + 7],
122        ]);
123        pos += 8;
124
125        let timestamp = u64::from_le_bytes([
126            data[pos],
127            data[pos + 1],
128            data[pos + 2],
129            data[pos + 3],
130            data[pos + 4],
131            data[pos + 5],
132            data[pos + 6],
133            data[pos + 7],
134        ]);
135        pos += 8;
136
137        let version = u64::from_le_bytes([
138            data[pos],
139            data[pos + 1],
140            data[pos + 2],
141            data[pos + 3],
142            data[pos + 4],
143            data[pos + 5],
144            data[pos + 6],
145            data[pos + 7],
146        ]);
147        pos += 8;
148
149        let rule_count =
150            u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]) as usize;
151        pos += 4;
152
153        let mut rules = HashMap::with_capacity(rule_count);
154
155        for _ in 0..rule_count {
156            if pos + 8 > data.len() {
157                return Err(DrivenError::InvalidBinary("Truncated snapshot".into()));
158            }
159
160            let id = u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]);
161            pos += 4;
162
163            let len = u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]])
164                as usize;
165            pos += 4;
166
167            if pos + len > data.len() {
168                return Err(DrivenError::InvalidBinary("Truncated rule data".into()));
169            }
170
171            rules.insert(id, data[pos..pos + len].to_vec());
172            pos += len;
173        }
174
175        // Read metadata
176        let mut metadata = HashMap::new();
177        if pos + 4 <= data.len() {
178            let meta_count =
179                u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]])
180                    as usize;
181            pos += 4;
182
183            for _ in 0..meta_count {
184                if pos + 2 > data.len() {
185                    break;
186                }
187                let key_len = u16::from_le_bytes([data[pos], data[pos + 1]]) as usize;
188                pos += 2;
189
190                if pos + key_len > data.len() {
191                    break;
192                }
193                let key = String::from_utf8_lossy(&data[pos..pos + key_len]).to_string();
194                pos += key_len;
195
196                if pos + 2 > data.len() {
197                    break;
198                }
199                let val_len = u16::from_le_bytes([data[pos], data[pos + 1]]) as usize;
200                pos += 2;
201
202                if pos + val_len > data.len() {
203                    break;
204                }
205                let value = String::from_utf8_lossy(&data[pos..pos + val_len]).to_string();
206                pos += val_len;
207
208                metadata.insert(key, value);
209            }
210        }
211
212        Ok(Self {
213            id,
214            timestamp,
215            version,
216            rules,
217            metadata,
218        })
219    }
220}
221
222/// Snapshot manager for storing and retrieving snapshots
223#[derive(Debug)]
224pub struct SnapshotManager {
225    /// Storage directory
226    directory: std::path::PathBuf,
227    /// Next snapshot ID
228    next_id: u64,
229    /// Maximum snapshots to keep
230    max_snapshots: usize,
231}
232
233impl SnapshotManager {
234    /// Create a new snapshot manager
235    pub fn new(directory: impl AsRef<Path>) -> Result<Self> {
236        let directory = directory.as_ref().to_path_buf();
237        std::fs::create_dir_all(&directory)?;
238
239        // Find highest existing ID
240        let next_id = Self::find_highest_id(&directory)? + 1;
241
242        Ok(Self {
243            directory,
244            next_id,
245            max_snapshots: 10,
246        })
247    }
248
249    /// Set maximum snapshots to keep
250    pub fn with_max_snapshots(mut self, max: usize) -> Self {
251        self.max_snapshots = max;
252        self
253    }
254
255    /// Create a snapshot
256    pub fn create(&mut self, shared: &SharedRules) -> Result<RuleSnapshot> {
257        let snapshot = RuleSnapshot::capture(shared, self.next_id);
258        self.next_id += 1;
259
260        self.save(&snapshot)?;
261        self.cleanup()?;
262
263        Ok(snapshot)
264    }
265
266    /// Save a snapshot to disk
267    pub fn save(&self, snapshot: &RuleSnapshot) -> Result<()> {
268        let path = self.snapshot_path(snapshot.id);
269        let data = snapshot.to_bytes();
270        std::fs::write(&path, &data)?;
271        Ok(())
272    }
273
274    /// Load a snapshot from disk
275    pub fn load(&self, id: u64) -> Result<RuleSnapshot> {
276        let path = self.snapshot_path(id);
277        let data = std::fs::read(&path)?;
278        RuleSnapshot::from_bytes(&data)
279    }
280
281    /// List all snapshots
282    pub fn list(&self) -> Result<Vec<u64>> {
283        let mut ids = Vec::new();
284
285        for entry in std::fs::read_dir(&self.directory)? {
286            let entry = entry?;
287            if let Some(name) = entry.file_name().to_str() {
288                if let Some(id_str) = name
289                    .strip_prefix("snapshot_")
290                    .and_then(|s| s.strip_suffix(".drv"))
291                {
292                    if let Ok(id) = id_str.parse() {
293                        ids.push(id);
294                    }
295                }
296            }
297        }
298
299        ids.sort();
300        Ok(ids)
301    }
302
303    /// Load the latest snapshot
304    pub fn load_latest(&self) -> Result<Option<RuleSnapshot>> {
305        let ids = self.list()?;
306        match ids.last() {
307            Some(&id) => Ok(Some(self.load(id)?)),
308            None => Ok(None),
309        }
310    }
311
312    /// Delete a snapshot
313    pub fn delete(&self, id: u64) -> Result<()> {
314        let path = self.snapshot_path(id);
315        std::fs::remove_file(&path)?;
316        Ok(())
317    }
318
319    /// Clean up old snapshots
320    fn cleanup(&self) -> Result<()> {
321        let mut ids = self.list()?;
322        while ids.len() > self.max_snapshots {
323            if let Some(oldest) = ids.first() {
324                self.delete(*oldest)?;
325                ids.remove(0);
326            }
327        }
328        Ok(())
329    }
330
331    fn snapshot_path(&self, id: u64) -> std::path::PathBuf {
332        self.directory.join(format!("snapshot_{:08}.drv", id))
333    }
334
335    fn find_highest_id(directory: &Path) -> Result<u64> {
336        let mut highest = 0u64;
337
338        if directory.exists() {
339            for entry in std::fs::read_dir(directory)? {
340                let entry = entry?;
341                if let Some(name) = entry.file_name().to_str() {
342                    if let Some(id_str) = name
343                        .strip_prefix("snapshot_")
344                        .and_then(|s| s.strip_suffix(".drv"))
345                    {
346                        if let Ok(id) = id_str.parse::<u64>() {
347                            highest = highest.max(id);
348                        }
349                    }
350                }
351            }
352        }
353
354        Ok(highest)
355    }
356}
357
358#[cfg(test)]
359mod tests {
360    use super::*;
361
362    #[test]
363    fn test_snapshot_roundtrip() {
364        let mut snapshot = RuleSnapshot::new(1);
365        snapshot.rules.insert(1, b"rule one".to_vec());
366        snapshot.rules.insert(2, b"rule two".to_vec());
367        snapshot
368            .metadata
369            .insert("key".to_string(), "value".to_string());
370
371        let bytes = snapshot.to_bytes();
372        let restored = RuleSnapshot::from_bytes(&bytes).unwrap();
373
374        assert_eq!(restored.id, 1);
375        assert_eq!(restored.rules.len(), 2);
376        assert_eq!(restored.rules.get(&1).unwrap(), b"rule one");
377    }
378
379    #[test]
380    fn test_capture_restore() {
381        let shared = SharedRules::new();
382        shared.insert(b"test rule".to_vec());
383
384        let snapshot = RuleSnapshot::capture(&shared, 1);
385        assert_eq!(snapshot.rule_count(), 1);
386
387        let new_shared = SharedRules::new();
388        snapshot.restore(&new_shared);
389        assert_eq!(new_shared.len(), 1);
390    }
391}