1use std::collections::hash_map::DefaultHasher;
2use std::hash::{Hash, Hasher};
3use std::num::NonZeroUsize;
4use std::sync::Mutex;
5use std::time::{Duration, Instant};
6
7use crate::types::{InvestigationReport, TargetClass};
8
9#[derive(Debug, Clone)]
11struct TtlEntry<V> {
12 value: V,
13 expires_at: Instant,
14}
15
16impl<V> TtlEntry<V> {
17 fn new(value: V, ttl: Duration) -> Self {
18 Self {
19 value,
20 expires_at: Instant::now() + ttl,
21 }
22 }
23
24 fn is_expired(&self) -> bool {
25 Instant::now() >= self.expires_at
26 }
27}
28
29pub(crate) struct LruTtlStore<V> {
44 ttl: Duration,
45 inner: Mutex<lru::LruCache<String, TtlEntry<V>>>,
46}
47
48impl<V: Clone> LruTtlStore<V> {
49 #[must_use]
51 pub(crate) fn new(capacity: NonZeroUsize, ttl: Duration) -> Self {
52 Self {
53 ttl,
54 inner: Mutex::new(lru::LruCache::new(capacity)),
55 }
56 }
57
58 #[must_use]
60 pub(crate) const fn ttl(&self) -> Duration {
61 self.ttl
62 }
63
64 pub(crate) fn get(&self, key: &str) -> Option<V> {
67 let Ok(mut cache) = self.inner.lock() else {
68 return None;
69 };
70
71 match cache.get(key) {
72 Some(entry) if entry.is_expired() => {
73 cache.pop(key);
74 None
75 }
76 Some(entry) => Some(entry.value.clone()),
77 None => None,
78 }
79 }
80
81 pub(crate) fn mutate<F>(&self, key: String, f: F)
96 where
97 F: FnOnce(Option<V>) -> V,
98 {
99 let Ok(mut cache) = self.inner.lock() else {
100 return;
101 };
102
103 let expired = cache.peek(&key).is_some_and(TtlEntry::is_expired);
107 if expired {
108 cache.pop(&key);
109 }
110 let existing = cache.peek(&key).map(|e| e.value.clone());
111
112 let value = f(existing);
113 cache.put(key, TtlEntry::new(value, self.ttl));
114 }
115
116 pub(crate) fn put(&self, key: String, value: V) {
118 let Ok(mut cache) = self.inner.lock() else {
119 return;
120 };
121
122 cache.put(key, TtlEntry::new(value, self.ttl));
123 }
124
125 pub(crate) fn invalidate(&self, key: &str) {
127 if let Ok(mut cache) = self.inner.lock() {
128 cache.pop(key);
129 }
130 }
131
132 pub(crate) fn clear(&self) {
134 if let Ok(mut cache) = self.inner.lock() {
135 cache.clear();
136 }
137 }
138
139 #[allow(dead_code)]
142 pub(crate) fn len(&self) -> usize {
143 self.inner.lock().map_or(0, |cache| cache.len())
144 }
145
146 #[allow(dead_code)]
148 pub(crate) fn is_empty(&self) -> bool {
149 self.len() == 0
150 }
151}
152
153pub trait InvestigationReportCache: Send + Sync {
158 fn get(&self, key: &str) -> Option<InvestigationReport>;
160
161 fn put(&self, key: String, report: InvestigationReport);
163
164 fn invalidate(&self, key: &str);
166
167 fn clear(&self);
169}
170
171#[must_use]
173pub fn investigation_cache_key(har_json: &str, target_class: TargetClass) -> String {
174 let mut hasher = DefaultHasher::new();
175 har_json.hash(&mut hasher);
176 target_class.hash(&mut hasher);
177 format!("charon:investigation:{:016x}", hasher.finish())
178}
179
180pub struct MemoryInvestigationCache {
182 store: LruTtlStore<InvestigationReport>,
183}
184
185impl MemoryInvestigationCache {
186 #[must_use]
188 pub fn new(capacity: NonZeroUsize, ttl: Duration) -> Self {
189 Self {
190 store: LruTtlStore::new(capacity, ttl),
191 }
192 }
193
194 #[must_use]
196 pub fn len(&self) -> usize {
197 self.store.len()
198 }
199
200 #[must_use]
202 pub fn is_empty(&self) -> bool {
203 self.store.is_empty()
204 }
205}
206
207impl InvestigationReportCache for MemoryInvestigationCache {
208 fn get(&self, key: &str) -> Option<InvestigationReport> {
209 self.store.get(key)
210 }
211
212 fn put(&self, key: String, report: InvestigationReport) {
213 self.store.put(key, report);
214 }
215
216 fn invalidate(&self, key: &str) {
217 self.store.invalidate(key);
218 }
219
220 fn clear(&self) {
221 self.store.clear();
222 }
223}
224
225#[cfg(feature = "redis-cache")]
227pub struct RedisInvestigationCache {
228 client: redis::Client,
229 ttl: Duration,
230 key_prefix: String,
231}
232
233#[cfg(feature = "redis-cache")]
234impl RedisInvestigationCache {
235 pub fn new(redis_url: &str, ttl: Duration) -> redis::RedisResult<Self> {
241 let client = redis::Client::open(redis_url)?;
242 Ok(Self {
243 client,
244 ttl,
245 key_prefix: "charon:investigation".to_string(),
246 })
247 }
248
249 fn prefixed_key(&self, key: &str) -> String {
250 format!("{}:{}", self.key_prefix, key)
251 }
252}
253
254#[cfg(feature = "redis-cache")]
255impl InvestigationReportCache for RedisInvestigationCache {
256 fn get(&self, key: &str) -> Option<InvestigationReport> {
257 let mut connection = self.client.get_connection().ok()?;
258 let payload: Option<String> = redis::cmd("GET")
259 .arg(self.prefixed_key(key))
260 .query(&mut connection)
261 .ok()?;
262 payload.and_then(|value| serde_json::from_str::<InvestigationReport>(&value).ok())
263 }
264
265 fn put(&self, key: String, report: InvestigationReport) {
266 let Ok(payload) = serde_json::to_string(&report) else {
267 return;
268 };
269 let ttl_seconds = self.ttl.as_secs();
270 let Ok(mut connection) = self.client.get_connection() else {
271 return;
272 };
273 let _: redis::RedisResult<()> = redis::cmd("SETEX")
274 .arg(self.prefixed_key(&key))
275 .arg(ttl_seconds)
276 .arg(payload)
277 .query(&mut connection);
278 }
279
280 fn invalidate(&self, key: &str) {
281 let Ok(mut connection) = self.client.get_connection() else {
282 return;
283 };
284 let _: redis::RedisResult<()> = redis::cmd("DEL")
285 .arg(self.prefixed_key(key))
286 .query(&mut connection);
287 }
288
289 fn clear(&self) {
290 let pattern = format!("{}:*", self.key_prefix);
291 let Ok(mut connection) = self.client.get_connection() else {
292 return;
293 };
294 let keys: redis::RedisResult<Vec<String>> =
295 redis::cmd("KEYS").arg(pattern).query(&mut connection);
296 if let Ok(keys) = keys
297 && !keys.is_empty()
298 {
299 let _: redis::RedisResult<()> = redis::cmd("DEL").arg(keys).query(&mut connection);
300 }
301 }
302}
303
304#[cfg(test)]
305#[allow(
306 clippy::unwrap_used,
307 clippy::expect_used,
308 clippy::panic,
309 clippy::indexing_slicing
310)]
311mod tests {
312 use super::*;
313 use crate::types::{AntiBotProvider, Detection, InvestigationReport};
314 use std::collections::BTreeMap;
315
316 fn sample_report() -> InvestigationReport {
317 InvestigationReport {
318 page_title: Some("https://example.com".to_string()),
319 total_requests: 10,
320 blocked_requests: 2,
321 status_histogram: BTreeMap::from([(200, 8), (403, 2)]),
322 resource_type_histogram: BTreeMap::new(),
323 provider_histogram: BTreeMap::new(),
324 marker_histogram: BTreeMap::new(),
325 top_markers: Vec::new(),
326 hosts: Vec::new(),
327 suspicious_requests: Vec::new(),
328 aggregate: Detection {
329 provider: AntiBotProvider::Unknown,
330 confidence: 0.1,
331 markers: Vec::new(),
332 },
333 target_class: Some(TargetClass::Api),
334 }
335 }
336
337 #[test]
338 fn memory_cache_round_trips_report() {
339 let capacity = NonZeroUsize::new(2).unwrap_or(NonZeroUsize::MIN);
340 let cache = MemoryInvestigationCache::new(capacity, Duration::from_mins(1));
341 let key = investigation_cache_key("{\"log\":{}}", TargetClass::Api);
342 let report = sample_report();
343
344 cache.put(key.clone(), report.clone());
345
346 assert_eq!(cache.get(&key), Some(report));
347 }
348
349 #[test]
350 fn memory_cache_expires_entries() {
351 let capacity = NonZeroUsize::new(2).unwrap_or(NonZeroUsize::MIN);
352 let cache = MemoryInvestigationCache::new(capacity, Duration::from_millis(1));
353 let key = investigation_cache_key("{\"log\":{}}", TargetClass::Api);
354 cache.put(key.clone(), sample_report());
355 std::thread::sleep(Duration::from_millis(5));
356 assert!(cache.get(&key).is_none());
357 }
358
359 #[test]
360 fn cache_key_changes_by_target_class() {
361 let har = "{\"log\":{\"entries\":[]}}";
362 let api = investigation_cache_key(har, TargetClass::Api);
363 let high = investigation_cache_key(har, TargetClass::HighSecurity);
364 assert_ne!(api, high);
365 }
366}