1use 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 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}