1use chrono::{DateTime, Utc};
8use serde::{Deserialize, Serialize};
9use sha2::{Digest, Sha256};
10
11#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
12pub struct CachedMessageRef {
13 pub graph_id: String,
14 pub account: String,
15 pub cached_at: DateTime<Utc>,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq)]
20pub enum CacheLookup {
21 NotFound,
22 One(String, CachedMessageRef),
23 Ambiguous(Vec<(String, CachedMessageRef)>),
30}
31
32pub fn short_hash(graph_id: &str) -> String {
35 let mut h = Sha256::new();
36 h.update(graph_id.as_bytes());
37 let result = h.finalize();
38 format!(
39 "{:02x}{:02x}{:02x}{:02x}",
40 result[0], result[1], result[2], result[3]
41 )
42}
43
44use std::collections::HashMap;
45use std::path::{Path, PathBuf};
46
47use crate::error::CoreError;
48
49const MAX_ENTRIES: usize = 1000;
50
51#[derive(Debug, Default, Clone, Serialize, Deserialize)]
52pub struct MessageCache {
53 #[serde(default)]
54 pub entries: HashMap<String, CachedMessageRef>,
55}
56
57impl MessageCache {
58 pub fn default_path() -> Result<PathBuf, CoreError> {
60 let dir = dirs::cache_dir()
61 .ok_or(CoreError::NoConfigDir)?
62 .join("pidge");
63 std::fs::create_dir_all(&dir)?;
64 Ok(dir.join("messages.json"))
65 }
66
67 pub fn load() -> Result<Self, CoreError> {
69 let path = Self::default_path()?;
70 Self::load_from(&path)
71 }
72
73 pub fn load_from(path: &Path) -> Result<Self, CoreError> {
74 if !path.exists() {
75 return Ok(Self::default());
76 }
77 let text = std::fs::read_to_string(path)?;
78 let cache: MessageCache = serde_json::from_str(&text)
79 .map_err(|e| CoreError::Io(std::io::Error::new(std::io::ErrorKind::InvalidData, e)))?;
80 Ok(cache)
81 }
82
83 pub fn save(&self) -> Result<(), CoreError> {
84 let path = Self::default_path()?;
85 self.save_to(&path)
86 }
87
88 pub fn save_to(&self, path: &Path) -> Result<(), CoreError> {
89 let text = serde_json::to_string_pretty(self)
90 .map_err(|e| CoreError::Io(std::io::Error::new(std::io::ErrorKind::InvalidData, e)))?;
91 std::fs::write(path, text)?;
92 Ok(())
93 }
94
95 pub fn insert_many(&mut self, msgs: &[(String, String)]) {
98 let now = Utc::now();
99 for (graph_id, account) in msgs {
100 let hash = short_hash(graph_id);
101 self.entries.insert(
102 hash,
103 CachedMessageRef {
104 graph_id: graph_id.clone(),
105 account: account.clone(),
106 cached_at: now,
107 },
108 );
109 }
110 self.evict_oldest_if_needed();
111 }
112
113 fn evict_oldest_if_needed(&mut self) {
114 if self.entries.len() <= MAX_ENTRIES {
115 return;
116 }
117 let excess = self.entries.len() - MAX_ENTRIES;
118 let mut sorted: Vec<(String, DateTime<Utc>)> = self
119 .entries
120 .iter()
121 .map(|(k, v)| (k.clone(), v.cached_at))
122 .collect();
123 sorted.sort_by_key(|(_, t)| *t);
124 for (k, _) in sorted.into_iter().take(excess) {
125 self.entries.remove(&k);
126 }
127 }
128
129 pub fn find_by_fragment(&self, fragment: &str) -> CacheLookup {
133 if fragment.is_empty() {
134 return CacheLookup::NotFound;
135 }
136 let mut matches: Vec<(String, CachedMessageRef)> = self
137 .entries
138 .iter()
139 .filter(|(k, _)| k.contains(fragment))
140 .map(|(k, v)| (k.clone(), v.clone()))
141 .collect();
142 match matches.len() {
143 0 => CacheLookup::NotFound,
144 1 => {
145 let (k, v) = matches.remove(0);
146 CacheLookup::One(k, v)
147 }
148 _ => CacheLookup::Ambiguous(matches.into_iter().take(10).collect()),
149 }
150 }
151}
152
153#[cfg(test)]
154mod tests {
155 use super::*;
156
157 #[test]
158 fn short_hash_is_deterministic() {
159 let a = short_hash("AAMkAGI2TG93AAA=");
160 let b = short_hash("AAMkAGI2TG93AAA=");
161 assert_eq!(a, b);
162 }
163
164 #[test]
165 fn short_hash_is_8_hex_chars() {
166 let h = short_hash("anything");
167 assert_eq!(h.len(), 8);
168 assert!(h.chars().all(|c| c.is_ascii_hexdigit()));
169 }
170
171 #[test]
172 fn short_hash_differs_for_different_inputs() {
173 assert_ne!(short_hash("x"), short_hash("y"));
174 }
175
176 #[test]
177 fn empty_cache_roundtrips_through_file() {
178 let tmp = tempfile::TempDir::new().unwrap();
179 let path = tmp.path().join("messages.json");
180
181 let cache = MessageCache::default();
182 cache.save_to(&path).unwrap();
183 let loaded = MessageCache::load_from(&path).unwrap();
184 assert_eq!(loaded.entries.len(), 0);
185 }
186
187 #[test]
188 fn populated_cache_roundtrips_through_file() {
189 let tmp = tempfile::TempDir::new().unwrap();
190 let path = tmp.path().join("messages.json");
191
192 let mut cache = MessageCache::default();
193 cache.insert_many(&[
194 ("AAA".to_string(), "user@example.com".to_string()),
195 ("BBB".to_string(), "user@example.com".to_string()),
196 ]);
197 cache.save_to(&path).unwrap();
198
199 let loaded = MessageCache::load_from(&path).unwrap();
200 assert_eq!(loaded.entries.len(), 2);
201 let hash_aaa = short_hash("AAA");
202 let entry = loaded.entries.get(&hash_aaa).unwrap();
203 assert_eq!(entry.graph_id, "AAA");
204 assert_eq!(entry.account, "user@example.com");
205 }
206
207 #[test]
208 fn load_from_missing_file_returns_empty_cache() {
209 let tmp = tempfile::TempDir::new().unwrap();
210 let path = tmp.path().join("nonexistent.json");
211 let cache = MessageCache::load_from(&path).unwrap();
212 assert_eq!(cache.entries.len(), 0);
213 }
214
215 #[test]
216 fn insert_many_evicts_oldest_when_over_max() {
217 let mut cache = MessageCache::default();
218
219 let now = Utc::now();
221 let old_time = now - chrono::Duration::seconds(3600);
222 for i in 0..MAX_ENTRIES {
223 let hash = format!("{:08x}", i);
224 cache.entries.insert(
225 hash,
226 CachedMessageRef {
227 graph_id: format!("old-{}", i),
228 account: "user@example.com".into(),
229 cached_at: old_time,
230 },
231 );
232 }
233 assert_eq!(cache.entries.len(), MAX_ENTRIES);
234
235 cache.insert_many(&[("new-graph-id".into(), "user@example.com".into())]);
237 assert_eq!(cache.entries.len(), MAX_ENTRIES);
238
239 let new_hash = short_hash("new-graph-id");
241 assert!(cache.entries.contains_key(&new_hash));
242 }
243
244 fn cache_with(entries: &[(&str, &str)]) -> MessageCache {
245 let mut cache = MessageCache::default();
246 let pairs: Vec<(String, String)> = entries
247 .iter()
248 .map(|(g, a)| (g.to_string(), a.to_string()))
249 .collect();
250 cache.insert_many(&pairs);
251 cache
252 }
253
254 #[test]
255 fn find_by_fragment_returns_not_found_when_no_match() {
256 let cache = cache_with(&[("hello", "u@e.com")]);
257 assert!(matches!(
258 cache.find_by_fragment("zzzzzz"),
259 CacheLookup::NotFound
260 ));
261 }
262
263 #[test]
264 fn find_by_fragment_returns_not_found_for_empty_fragment() {
265 let cache = cache_with(&[("hello", "u@e.com")]);
266 assert!(matches!(cache.find_by_fragment(""), CacheLookup::NotFound));
267 }
268
269 #[test]
270 fn find_by_fragment_returns_one_for_exact_match() {
271 let cache = cache_with(&[("hello", "u@e.com")]);
272 let hash = short_hash("hello");
273 match cache.find_by_fragment(&hash) {
274 CacheLookup::One(h, _) => assert_eq!(h, hash),
275 other => panic!("expected One, got {other:?}"),
276 }
277 }
278
279 #[test]
280 fn find_by_fragment_matches_prefix() {
281 let cache = cache_with(&[("hello", "u@e.com")]);
282 let hash = short_hash("hello");
283 let prefix = &hash[..3];
284 assert!(matches!(
285 cache.find_by_fragment(prefix),
286 CacheLookup::One(_, _)
287 ));
288 }
289
290 #[test]
291 fn find_by_fragment_matches_suffix() {
292 let cache = cache_with(&[("hello", "u@e.com")]);
293 let hash = short_hash("hello");
294 let suffix = &hash[hash.len() - 3..];
295 assert!(matches!(
296 cache.find_by_fragment(suffix),
297 CacheLookup::One(_, _)
298 ));
299 }
300
301 #[test]
302 fn find_by_fragment_matches_middle_substring() {
303 let cache = cache_with(&[("hello", "u@e.com")]);
304 let hash = short_hash("hello");
305 let middle = &hash[2..5];
306 assert!(matches!(
307 cache.find_by_fragment(middle),
308 CacheLookup::One(_, _)
309 ));
310 }
311
312 #[test]
313 fn find_by_fragment_returns_ambiguous_when_two_hashes_share_fragment() {
314 let mut cache = MessageCache::default();
315 let mut entries = Vec::new();
316 for i in 0..200 {
317 entries.push((format!("graph-{i}"), "u@e.com".to_string()));
318 }
319 cache.insert_many(&entries);
320
321 let hashes: Vec<String> = cache.entries.keys().cloned().collect();
322 let mut found_ambiguous = false;
323 'outer: for h in &hashes {
324 for start in 0..h.len() - 1 {
325 let frag = &h[start..start + 2];
326 let count = hashes.iter().filter(|other| other.contains(frag)).count();
327 if count >= 2 {
328 if let CacheLookup::Ambiguous(matches) = cache.find_by_fragment(frag) {
329 assert!(matches.len() >= 2);
330 found_ambiguous = true;
331 break 'outer;
332 }
333 }
334 }
335 }
336 assert!(
337 found_ambiguous,
338 "expected at least one ambiguous fragment in 200-entry cache"
339 );
340 }
341}