1use std::collections::HashMap;
6use std::path::Path;
7
8use crate::{DrivenError, Result};
9
10use super::SharedRules;
11
12#[derive(Debug, Clone)]
14pub struct RuleSnapshot {
15 pub id: u64,
17 pub timestamp: u64,
19 pub version: u64,
21 pub rules: HashMap<u32, Vec<u8>>,
23 pub metadata: HashMap<String, String>,
25}
26
27impl RuleSnapshot {
28 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 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 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 pub fn size(&self) -> usize {
66 self.rules.values().map(|v| v.len()).sum()
67 }
68
69 pub fn rule_count(&self) -> usize {
71 self.rules.len()
72 }
73
74 pub fn to_bytes(&self) -> Vec<u8> {
76 let mut output = Vec::new();
77
78 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 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 output.extend_from_slice(&(self.metadata.len() as u32).to_le_bytes());
93
94 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 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 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#[derive(Debug)]
224pub struct SnapshotManager {
225 directory: std::path::PathBuf,
227 next_id: u64,
229 max_snapshots: usize,
231}
232
233impl SnapshotManager {
234 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 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 pub fn with_max_snapshots(mut self, max: usize) -> Self {
251 self.max_snapshots = max;
252 self
253 }
254
255 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 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 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 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 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 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 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}