dragonfly_client_rs/
durable_cache.rs

1//! Batched, optional database cache transport. File bytes never leave the worker.
2use color_eyre::{eyre::ensure, Result};
3use reqwest::blocking::Client;
4use serde::{de::DeserializeOwned, Deserialize, Serialize};
5use sha2::{Digest, Sha256};
6use std::{
7    collections::HashMap,
8    fs::File,
9    io::Read,
10    path::Path,
11    time::{Duration, Instant},
12};
13
14pub(crate) const BATCH_SIZE: usize = 128;
15const MAX_RESULT_BYTES: usize = 16_384;
16const MAX_BATCH_BYTES: usize = 512 * 1024;
17const MAX_RESPONSE_BYTES: u64 = 3 * 1024 * 1024;
18const MAX_READ_BYTES: usize = 16 * 1024 * 1024;
19const NETWORK_BUDGET: Duration = Duration::from_secs(2);
20
21#[derive(Clone, Serialize)]
22pub(crate) struct Context {
23    scanner: &'static str,
24    rules_commit: String,
25    rules_digest: String,
26    engine_digest: String,
27}
28
29#[derive(Clone, Serialize)]
30pub(crate) struct Lease {
31    pub name: String,
32    pub version: String,
33    pub assignment_id: String,
34    pub attempt: u64,
35}
36
37#[derive(Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
38pub(crate) struct Key {
39    pub file_digest: String,
40    pub language: String,
41}
42
43#[derive(Clone, Serialize, Deserialize)]
44struct Value {
45    #[serde(flatten)]
46    key: Key,
47    result: String,
48}
49
50#[derive(Deserialize)]
51struct Reply {
52    revoked: bool,
53    entries: Vec<Value>,
54}
55
56#[derive(Deserialize)]
57struct QuarantineResult {
58    quarantined: usize,
59}
60
61pub(crate) struct DurableCache {
62    client: Client,
63    url: String,
64    context: Context,
65    lease: Option<Lease>,
66    keys: HashMap<String, Key>,
67    hits: HashMap<String, String>,
68    pending: Vec<Value>,
69    quarantines: Vec<Key>,
70    pending_bytes: usize,
71    read_bytes: usize,
72    spent: Duration,
73    unavailable: bool,
74    pub revoked: bool,
75}
76
77impl DurableCache {
78    pub fn new(
79        client: Client,
80        url: &str,
81        scanner: &'static str,
82        commit: &str,
83        rules: &HashMap<String, String>,
84        engine: Option<&Path>,
85    ) -> Result<Self> {
86        let mut digest = Sha256::new();
87        digest.update(hash_path(&std::env::current_exe()?)?);
88        if let Some(engine) = engine {
89            digest.update(hash_path(engine)?);
90        }
91        digest.update(b"durable-file-results-v1");
92        Ok(Self {
93            client,
94            url: format!("{}/scan-cache", url.trim_end_matches('/')),
95            context: Context {
96                scanner,
97                rules_commit: commit.to_owned(),
98                rules_digest: rules_digest(rules),
99                engine_digest: format!("{:x}", digest.finalize()),
100            },
101            lease: None,
102            keys: HashMap::new(),
103            hits: HashMap::new(),
104            pending: Vec::new(),
105            quarantines: Vec::new(),
106            pending_bytes: 0,
107            read_bytes: 0,
108            spent: Duration::ZERO,
109            unavailable: false,
110            revoked: false,
111        })
112    }
113
114    pub fn begin_job(&mut self, lease: Lease) {
115        self.lease = Some(lease);
116        self.keys.clear();
117        self.hits.clear();
118        self.pending.clear();
119        self.pending_bytes = 0;
120        self.read_bytes = 0;
121        self.spent = Duration::ZERO;
122        self.unavailable = false;
123    }
124
125    fn available(&self) -> bool {
126        !self.unavailable && !self.revoked && self.lease.is_some() && self.spent < NETWORK_BUDGET
127    }
128
129    fn post<T: DeserializeOwned>(&mut self, route: &str, body: &impl Serialize) -> Result<T> {
130        let started = Instant::now();
131        let result = (|| {
132            let response = self
133                .client
134                .post(format!("{}/{route}", self.url))
135                .timeout(Duration::from_millis(750))
136                .json(body)
137                .send()?
138                .error_for_status()?;
139            let mut bytes = Vec::new();
140            response
141                .take(MAX_RESPONSE_BYTES + 1)
142                .read_to_end(&mut bytes)?;
143            ensure!(
144                bytes.len() as u64 <= MAX_RESPONSE_BYTES,
145                "Cache response exceeds limit"
146            );
147            Ok(serde_json::from_slice(&bytes)?)
148        })();
149        self.spent += started.elapsed();
150        if result.is_err() {
151            self.unavailable = true;
152        }
153        result
154    }
155
156    pub fn prefetch(&mut self, keys: &[(String, Key)]) -> Result<()> {
157        self.flush_quarantines()?;
158        if !self.available() {
159            return Ok(());
160        }
161        for chunk in keys.chunks(BATCH_SIZE) {
162            if !self.available() || self.keys.len() + chunk.len() > 65_536 {
163                break;
164            }
165            let by_key: HashMap<Key, String> = chunk
166                .iter()
167                .map(|(local, key)| (key.clone(), local.clone()))
168                .collect();
169            let body = serde_json::json!({"context":self.context,"keys":by_key.keys().collect::<Vec<_>>()});
170            let reply: Reply = self.post("lookup", &body)?;
171            if reply.revoked {
172                self.revoked = true;
173                self.hits.clear();
174                break;
175            }
176            for (local, key) in chunk {
177                self.keys.insert(local.clone(), key.clone());
178            }
179            for value in reply.entries {
180                if value.result.len() > MAX_RESULT_BYTES {
181                    self.unavailable = true;
182                    self.hits.clear();
183                    color_eyre::eyre::bail!("Cache result exceeds limit");
184                }
185                if self.read_bytes + value.result.len() > MAX_READ_BYTES {
186                    continue;
187                }
188                if let Some(local) = by_key.get(&value.key) {
189                    self.read_bytes += value.result.len();
190                    self.hits.insert(local.clone(), value.result);
191                }
192            }
193        }
194        Ok(())
195    }
196
197    pub fn lookup<T: DeserializeOwned>(&self, key: &str) -> Result<Option<T>> {
198        if self.unavailable || self.revoked {
199            return Ok(None);
200        }
201        self.hits
202            .get(key)
203            .map(|value| serde_json::from_str(value).map_err(Into::into))
204            .transpose()
205    }
206
207    pub fn insert(&mut self, local: &str, result: &impl Serialize) -> Result<()> {
208        if !self.available() || self.hits.contains_key(local) {
209            return Ok(());
210        }
211        let Some(key) = self.keys.get(local).cloned() else {
212            return Ok(());
213        };
214        let result = serde_json::to_string(result)?;
215        if result.len() > MAX_RESULT_BYTES {
216            return Ok(());
217        }
218        let value = Value { key, result };
219        let encoded_bytes = serde_json::to_vec(&value)?.len();
220        let envelope_bytes = serde_json::to_vec(&serde_json::json!({
221            "context": self.context, "lease": self.lease, "entries": [], "revoke": false
222        }))?
223        .len();
224        if self.pending.len() >= BATCH_SIZE
225            || envelope_bytes + self.pending_bytes + encoded_bytes + self.pending.len()
226                > MAX_BATCH_BYTES
227        {
228            self.flush()?;
229        }
230        if !self.available() {
231            return Ok(());
232        }
233        if envelope_bytes + encoded_bytes > MAX_BATCH_BYTES {
234            return Ok(());
235        }
236        self.pending_bytes += encoded_bytes;
237        self.pending.push(value);
238        Ok(())
239    }
240
241    pub fn flush(&mut self) -> Result<()> {
242        if self.pending.is_empty() {
243            return Ok(());
244        }
245        let entries = std::mem::take(&mut self.pending);
246        self.pending_bytes = 0;
247        if !self.available() {
248            return Ok(());
249        }
250        let body = serde_json::json!({"context":self.context,"lease":self.lease,"entries":entries,"revoke":false});
251        let _: serde_json::Value = self.post("write", &body)?;
252        Ok(())
253    }
254
255    pub fn quarantine(&mut self, local: &str) -> Result<()> {
256        let key = self
257            .keys
258            .get(local)
259            .cloned()
260            .ok_or_else(|| color_eyre::eyre::eyre!("Missing cache key for quarantine"))?;
261        self.hits.remove(local);
262        self.pending.retain(|value| value.key != key);
263        self.pending_bytes = self
264            .pending
265            .iter()
266            .map(|value| serde_json::to_vec(value).map(|bytes| bytes.len()))
267            .sum::<std::result::Result<usize, _>>()?;
268        if !self.quarantines.contains(&key) {
269            ensure!(
270                self.quarantines.len() < 4096,
271                "Pending quarantine capacity reached"
272            );
273            self.quarantines.push(key);
274        }
275        self.flush_quarantines()
276    }
277
278    fn flush_quarantines(&mut self) -> Result<()> {
279        while !self.quarantines.is_empty() && self.available() {
280            let count = self.quarantines.len().min(BATCH_SIZE);
281            let body = serde_json::json!({"context":self.context,"lease":self.lease,"entries":[],"quarantine":self.quarantines[..count]});
282            let reply: QuarantineResult = self.post("write", &body)?;
283            ensure!(
284                reply.quarantined <= count,
285                "Invalid quarantine acknowledgement"
286            );
287            // HTTP success commits the entire Mainframe transaction. This count is
288            // newly changed rows, not per-key acknowledgements: already quarantined
289            // or expired/absent entries contribute zero and require no retry.
290            self.quarantines.drain(..count);
291        }
292        Ok(())
293    }
294
295    #[cfg(test)]
296    pub fn revoke(&mut self) -> Result<()> {
297        self.revoked = true;
298        self.hits.clear();
299        self.pending.clear();
300        self.pending_bytes = 0;
301        let body = serde_json::json!({"context":self.context,"lease":self.lease,"entries":[],"revoke":true});
302        let _: serde_json::Value = self.post("write", &body)?;
303        Ok(())
304    }
305}
306
307pub(crate) fn hash_path(path: &Path) -> Result<[u8; 32]> {
308    let mut file = File::open(path)?;
309    let mut digest = Sha256::new();
310    let mut buffer = [0; 8192];
311    loop {
312        let size = file.read(&mut buffer)?;
313        if size == 0 {
314            break;
315        }
316        digest.update(&buffer[..size]);
317    }
318    Ok(digest.finalize().into())
319}
320
321fn rules_digest(rules: &HashMap<String, String>) -> String {
322    let mut pairs = rules.iter().collect::<Vec<_>>();
323    pairs.sort_unstable_by_key(|(name, _)| *name);
324    let mut digest = Sha256::new();
325    for (name, contents) in pairs {
326        for value in [name, contents] {
327            digest.update((value.len() as u64).to_be_bytes());
328            digest.update(value.as_bytes());
329        }
330    }
331    format!("{:x}", digest.finalize())
332}
333
334#[cfg(test)]
335mod tests {
336    use super::*;
337    use std::{io::Write, net::TcpListener, sync::mpsc, thread};
338
339    fn server(
340        responses: Vec<(u16, String)>,
341    ) -> (
342        String,
343        mpsc::Receiver<serde_json::Value>,
344        thread::JoinHandle<()>,
345    ) {
346        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
347        let url = format!("http://{}", listener.local_addr().unwrap());
348        let (tx, rx) = mpsc::channel();
349        let handle = thread::spawn(move || {
350            for (status, body) in responses {
351                let (mut stream, _) = listener.accept().unwrap();
352                stream
353                    .set_read_timeout(Some(Duration::from_secs(5)))
354                    .unwrap();
355                let mut bytes = Vec::new();
356                let mut buffer = [0; 4096];
357                let (offset, length) = loop {
358                    let n = stream.read(&mut buffer).unwrap();
359                    assert!(n > 0);
360                    bytes.extend_from_slice(&buffer[..n]);
361                    if let Some(offset) = bytes.windows(4).position(|v| v == b"\r\n\r\n") {
362                        let headers = String::from_utf8_lossy(&bytes[..offset]);
363                        let length = headers
364                            .lines()
365                            .find_map(|line| {
366                                line.to_lowercase()
367                                    .strip_prefix("content-length: ")
368                                    .map(|v| v.parse::<usize>().unwrap())
369                            })
370                            .unwrap();
371                        if bytes.len() >= offset + 4 + length {
372                            break (offset + 4, length);
373                        }
374                    }
375                };
376                tx.send(serde_json::from_slice(&bytes[offset..offset + length]).unwrap())
377                    .unwrap();
378                write!(stream,"HTTP/1.1 {status} Response\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",body.len()).unwrap();
379            }
380        });
381        (url, rx, handle)
382    }
383
384    fn client(url: &str) -> DurableCache {
385        let mut cache =
386            DurableCache::new(Client::new(), url, "yara", "rules", &HashMap::new(), None).unwrap();
387        cache.begin_job(Lease {
388            name: "test".into(),
389            version: "1".into(),
390            assignment_id: "lease".into(),
391            attempt: 1,
392        });
393        cache
394    }
395
396    fn key() -> (String, Key) {
397        (
398            "local".into(),
399            Key {
400                file_digest: "a".repeat(64),
401                language: "py".into(),
402            },
403        )
404    }
405
406    #[test]
407    fn fresh_worker_reads_durable_result_and_persistent_revocation() {
408        let (_, key) = key();
409        let hit=serde_json::json!({"revoked":false,"entries":[{"file_digest":key.file_digest,"language":"py","result":"[1,2]"}]}).to_string();
410        let (url, rx, server) = server(vec![
411            (200, "{\"revoked\":false,\"entries\":[]}".into()),
412            (200, "{\"inserted\":1,\"skipped\":0}".into()),
413            (200, hit),
414            (200, "{}".into()),
415            (200, "{\"revoked\":true,\"entries\":[]}".into()),
416        ]);
417        let mut first = client(&url);
418        first.prefetch(&[self::key()]).unwrap();
419        assert!(first.lookup::<Vec<u64>>("local").unwrap().is_none());
420        first.insert("local", &vec![1_u64, 2]).unwrap();
421        first.flush().unwrap();
422        drop(first);
423        let mut second = client(&url);
424        second.prefetch(&[self::key()]).unwrap();
425        assert_eq!(
426            second.lookup::<Vec<u64>>("local").unwrap().unwrap(),
427            vec![1, 2]
428        );
429        second.revoke().unwrap();
430        assert!(second.lookup::<Vec<u64>>("local").unwrap().is_none());
431        let mut third = client(&url);
432        third.prefetch(&[self::key()]).unwrap();
433        assert!(third.revoked);
434        server.join().unwrap();
435        let bodies = rx.try_iter().collect::<Vec<_>>();
436        assert_eq!(bodies[0]["context"], bodies[2]["context"]);
437        assert_eq!(bodies[1]["entries"][0]["result"], "[1,2]");
438        assert_eq!(bodies[3]["revoke"], true);
439    }
440
441    #[test]
442    fn unavailable_cache_stops_requests_for_this_job() {
443        let (url, rx, server) = server(vec![(503, "{}".into())]);
444        let mut cache = client(&url);
445        assert!(cache.prefetch(&[key()]).is_err());
446        cache.prefetch(&[key()]).unwrap();
447        cache.insert("local", &Vec::<u8>::new()).unwrap();
448        cache.flush().unwrap();
449        assert!(cache.lookup::<Vec<u8>>("local").unwrap().is_none());
450        server.join().unwrap();
451        assert_eq!(rx.try_iter().count(), 1);
452    }
453
454    #[test]
455    fn quarantine_keeps_unrelated_writes_and_accepts_idempotent_zero_changes() {
456        let (url, rx, server) = server(vec![
457            (200, "{\"revoked\":false,\"entries\":[]}".into()),
458            (200, "{\"quarantined\":0}".into()),
459            (200, "{\"inserted\":1}".into()),
460        ]);
461        let mut cache = client(&url);
462        let good = (
463            "good".to_owned(),
464            Key {
465                file_digest: "b".repeat(64),
466                language: "py".into(),
467            },
468        );
469        cache.prefetch(&[key(), good]).unwrap();
470        cache.insert("local", &vec![1_u8]).unwrap();
471        cache.insert("good", &vec![7_u8]).unwrap();
472        cache.quarantine("local").unwrap();
473        assert!(cache.quarantines.is_empty());
474        assert!(!cache.revoked);
475        cache.flush().unwrap();
476        server.join().unwrap();
477        let bodies = rx.try_iter().collect::<Vec<_>>();
478        assert_eq!(bodies[2]["entries"].as_array().unwrap().len(), 1);
479        assert_eq!(bodies[2]["entries"][0]["file_digest"], "b".repeat(64));
480        assert_eq!(bodies[2]["entries"][0]["result"], "[7]");
481    }
482
483    #[test]
484    fn quarantine_retries_after_failure_without_revoking_namespace() {
485        let hit = serde_json::json!({"revoked":false,"entries":[
486            {"file_digest":"a".repeat(64),"language":"py","result":"[1]"}
487        ]})
488        .to_string();
489        let (url, rx, server) = server(vec![
490            (200, hit),
491            (503, "{}".into()),
492            (200, "{\"quarantined\":1}".into()),
493            (200, "{\"revoked\":false,\"entries\":[]}".into()),
494            (200, "{\"revoked\":false,\"entries\":[]}".into()),
495        ]);
496        let mut cache = client(&url);
497        cache.prefetch(&[key()]).unwrap();
498        assert!(cache.quarantine("local").is_err());
499        assert!(!cache.revoked);
500        assert!(cache.lookup::<Vec<u8>>("local").unwrap().is_none());
501        cache.begin_job(Lease {
502            name: "next".into(),
503            version: "1".into(),
504            assignment_id: "lease".into(),
505            attempt: 1,
506        });
507        cache.prefetch(&[key()]).unwrap();
508        assert!(cache.quarantines.is_empty());
509        assert!(cache.lookup::<Vec<u8>>("local").unwrap().is_none());
510        let mut restarted = client(&url);
511        restarted.prefetch(&[key()]).unwrap();
512        assert!(restarted.lookup::<Vec<u8>>("local").unwrap().is_none());
513        assert!(!restarted.revoked);
514        server.join().unwrap();
515        let bodies = rx.try_iter().collect::<Vec<_>>();
516        assert_eq!(bodies.len(), 5);
517        assert_eq!(bodies[1]["quarantine"], bodies[2]["quarantine"]);
518        assert!(bodies[1].get("revoke").is_none());
519    }
520
521    #[test]
522    fn oversized_response_disables_previously_loaded_hits() {
523        let (_, key) = key();
524        let hit = serde_json::json!({"revoked":false,"entries":[
525            {"file_digest":key.file_digest,"language":"py","result":"[]"}
526        ]})
527        .to_string();
528        let bad = serde_json::json!({"revoked":false,"entries":[
529            {"file_digest":key.file_digest,"language":"py","result":"x".repeat(MAX_RESULT_BYTES + 1)}
530        ]}).to_string();
531        let (url, rx, server) = server(vec![(200, hit), (200, bad)]);
532        let mut cache = client(&url);
533        cache.prefetch(&[self::key()]).unwrap();
534        assert!(cache.lookup::<Vec<u8>>("local").unwrap().is_some());
535        assert!(cache.prefetch(&[self::key()]).is_err());
536        assert!(cache.lookup::<Vec<u8>>("local").unwrap().is_none());
537        cache.prefetch(&[self::key()]).unwrap();
538        server.join().unwrap();
539        assert_eq!(rx.try_iter().count(), 2);
540    }
541
542    #[test]
543    fn writes_bound_the_entire_encoded_request() {
544        let (url, rx, server) = server(vec![(200, "{}".into()), (200, "{}".into())]);
545        let mut cache = client(&url);
546        for n in 0..32 {
547            let local = n.to_string();
548            cache.keys.insert(
549                local.clone(),
550                Key {
551                    file_digest: format!("{n:064x}"),
552                    language: "py".into(),
553                },
554            );
555            cache.insert(&local, &vec!["\"".repeat(7_000)]).unwrap();
556        }
557        cache.flush().unwrap();
558        server.join().unwrap();
559        let requests = rx.try_iter().collect::<Vec<_>>();
560        assert_eq!(requests.len(), 2);
561        assert_eq!(
562            requests
563                .iter()
564                .map(|v| v["entries"].as_array().unwrap().len())
565                .sum::<usize>(),
566            32
567        );
568        for request in requests {
569            assert!(serde_json::to_vec(&request).unwrap().len() <= MAX_BATCH_BYTES);
570        }
571    }
572
573    #[test]
574    fn hash_and_corpus_fingerprints_are_stable_and_content_sensitive() {
575        let file = tempfile::NamedTempFile::new().unwrap();
576        std::fs::write(file.path(), b"abc").unwrap();
577        let digest = hash_path(file.path()).unwrap();
578        assert_eq!(digest.to_vec(), Sha256::digest(b"abc").to_vec());
579        let a = HashMap::from([("a".into(), "bc".into())]);
580        let b = HashMap::from([("ab".into(), "c".into())]);
581        assert_ne!(rules_digest(&a), rules_digest(&b));
582    }
583}