1use std::fs::{self, File, OpenOptions};
14use std::io::{Read, Write};
15use std::path::{Path, PathBuf};
16use std::time::{Duration, SystemTime};
17
18use fs2::FileExt;
19
20use crate::error::{AUTH_FAILURE_MESSAGE, AppError, Result};
21
22pub const DEFAULT_TTL: Duration = Duration::from_secs(60);
24
25pub const MAX_STALE: Duration = Duration::from_secs(7 * 24 * 3600);
28
29#[derive(Debug, Clone)]
33pub struct Cache {
34 dir: PathBuf,
35}
36
37impl Cache {
38 pub fn for_vendor(vendor: &str) -> Result<Self> {
41 let base = xdg_cache_dir()?.join("ai-usagebar").join(vendor);
42 Ok(Self { dir: base })
43 }
44
45 pub fn for_vendor_account(vendor: &str, label: &str) -> Result<Self> {
50 let base = xdg_cache_dir()?
51 .join("ai-usagebar")
52 .join(vendor)
53 .join(label);
54 Ok(Self { dir: base })
55 }
56
57 pub fn at(path: PathBuf) -> Self {
59 Self { dir: path }
60 }
61
62 pub fn ensure_dir(&self) -> Result<()> {
64 fs::create_dir_all(&self.dir).map_err(|e| AppError::io_at(&self.dir, e))
65 }
66
67 pub fn dir(&self) -> &Path {
68 &self.dir
69 }
70
71 pub fn payload_path(&self) -> PathBuf {
72 self.dir.join("usage.json")
73 }
74 pub fn stale_path(&self) -> PathBuf {
75 self.dir.join(".stale")
76 }
77 pub fn last_error_path(&self) -> PathBuf {
78 self.dir.join(".last_error")
79 }
80 pub fn lock_path(&self) -> PathBuf {
81 self.dir.join(".fetch.lock")
82 }
83
84 pub fn payload_age(&self) -> Option<Duration> {
87 let meta = fs::metadata(self.payload_path()).ok()?;
88 let mtime = meta.modified().ok()?;
89 SystemTime::now().duration_since(mtime).ok()
90 }
91
92 pub fn fresh_payload(&self, ttl: Duration) -> Result<Option<Vec<u8>>> {
95 let Some(age) = self.payload_age() else {
96 return Ok(None);
97 };
98 if age < ttl {
99 self.read_payload().map(Some)
100 } else {
101 Ok(None)
102 }
103 }
104
105 pub fn maybe_payload(&self) -> Result<Option<Vec<u8>>> {
111 if !self.payload_path().exists() {
112 return Ok(None);
113 }
114 self.read_payload().map(Some)
115 }
116
117 pub fn fallback_payload(&self, max_stale: Duration) -> Result<Option<Vec<u8>>> {
123 let Some(age) = self.payload_age() else {
124 return Ok(None);
125 };
126 if age > max_stale {
127 return Ok(None);
128 }
129 self.read_payload().map(Some)
130 }
131
132 fn read_payload(&self) -> Result<Vec<u8>> {
133 let p = self.payload_path();
134 let mut f = File::open(&p).map_err(|e| AppError::io_at(&p, e))?;
135 let mut buf = Vec::new();
136 f.read_to_end(&mut buf)
137 .map_err(|e| AppError::io_at(&p, e))?;
138 Ok(buf)
139 }
140
141 pub fn write_payload(&self, bytes: &[u8]) -> Result<()> {
144 self.ensure_dir()?;
145 let mut tmp = tempfile::Builder::new()
146 .prefix(".usage.")
147 .tempfile_in(&self.dir)
148 .map_err(|e| AppError::io_at(&self.dir, e))?;
149 tmp.write_all(bytes)
150 .map_err(|e| AppError::io_at(tmp.path(), e))?;
151 tmp.as_file_mut()
152 .sync_all()
153 .map_err(|e| AppError::io_at(tmp.path(), e))?;
154 tmp.persist(self.payload_path())
155 .map_err(|e| AppError::io_at(self.payload_path(), e.error))?;
156 let _ = fs::remove_file(self.stale_path());
158 let _ = fs::remove_file(self.last_error_path());
159 Ok(())
160 }
161
162 pub fn mark_stale(&self) {
164 let _ = self.ensure_dir();
165 let _ = File::create(self.stale_path());
166 }
167
168 pub fn is_stale(&self) -> bool {
169 self.stale_path().exists()
170 }
171
172 pub fn write_last_error(&self, code: u16, msg: &str) -> (u16, String) {
184 let _ = self.ensure_dir();
185 let path = self.last_error_path();
186 let msg = if matches!(code, 401 | 403) {
190 AUTH_FAILURE_MESSAGE
191 } else {
192 msg
193 };
194 let msg = crate::display::sanitize_untrusted_field(msg);
195 let body = format!("{code}\n{msg}");
196 let _ = atomic_write(&path, body.as_bytes());
197 (code, msg)
198 }
199
200 pub fn clear_last_error(&self) {
202 let _ = fs::remove_file(self.last_error_path());
203 }
204
205 pub fn read_last_error(&self) -> Option<(u16, String)> {
206 let raw = fs::read_to_string(self.last_error_path()).ok()?;
207 let (code, msg) = raw.split_once('\n').unwrap_or((raw.as_str(), ""));
213 Some((code.parse::<u16>().ok()?, msg.to_string()))
214 }
215}
216
217pub async fn acquire_lock_async(path: &Path, timeout: Duration) -> Result<LockGuard> {
230 let path = path.to_path_buf();
231 tokio::task::spawn_blocking(move || acquire_lock(&path, timeout))
232 .await
233 .map_err(|e| AppError::Other(format!("cache lock task failed: {e}")))?
234}
235
236pub fn acquire_lock(path: &Path, timeout: Duration) -> Result<LockGuard> {
237 if let Some(parent) = path.parent() {
238 fs::create_dir_all(parent).map_err(|e| AppError::io_at(parent, e))?;
239 }
240 let f = OpenOptions::new()
241 .create(true)
242 .read(true)
243 .write(true)
244 .truncate(false)
245 .open(path)
246 .map_err(|e| AppError::io_at(path, e))?;
247
248 let deadline = std::time::Instant::now() + timeout;
249 loop {
250 match f.try_lock_exclusive() {
251 Ok(()) => return Ok(LockGuard { file: f }),
252 Err(_) => {
253 if std::time::Instant::now() >= deadline {
254 return Err(AppError::Other(format!(
255 "cache lock timeout after {:?}",
256 timeout
257 )));
258 }
259 std::thread::sleep(Duration::from_millis(50));
260 }
261 }
262 }
263}
264
265pub struct LockGuard {
269 file: File,
270}
271
272impl Drop for LockGuard {
273 fn drop(&mut self) {
274 let _ = FileExt::unlock(&self.file);
275 }
276}
277
278pub fn atomic_write(path: &Path, bytes: &[u8]) -> Result<()> {
281 let dir = path.parent().ok_or_else(|| {
282 AppError::Other(format!(
283 "atomic_write: path has no parent: {}",
284 path.display()
285 ))
286 })?;
287 fs::create_dir_all(dir).map_err(|e| AppError::io_at(dir, e))?;
288 let mut tmp = tempfile::Builder::new()
289 .prefix(".tmp.")
290 .tempfile_in(dir)
291 .map_err(|e| AppError::io_at(dir, e))?;
292 tmp.write_all(bytes)
293 .map_err(|e| AppError::io_at(tmp.path(), e))?;
294 tmp.as_file_mut()
295 .sync_all()
296 .map_err(|e| AppError::io_at(tmp.path(), e))?;
297 tmp.persist(path)
298 .map_err(|e| AppError::io_at(path, e.error))?;
299 Ok(())
300}
301
302fn xdg_cache_dir() -> Result<PathBuf> {
303 directories::BaseDirs::new()
304 .map(|b| b.cache_dir().to_path_buf())
305 .ok_or_else(|| AppError::Other("could not resolve XDG cache dir (no HOME?)".into()))
306}
307
308pub fn home_dir() -> Result<PathBuf> {
316 directories::BaseDirs::new()
317 .map(|b| b.home_dir().to_path_buf())
318 .ok_or_else(|| AppError::Other("could not resolve home directory (no HOME?)".into()))
319}
320
321#[cfg(test)]
328pub(crate) fn closed_temp_file(name: &str, contents: Option<&str>) -> (tempfile::TempDir, PathBuf) {
329 let dir = tempfile::TempDir::new().unwrap();
330 let path = dir.path().join(name);
331 if let Some(c) = contents {
332 std::fs::write(&path, c).unwrap();
333 }
334 (dir, path)
335}
336
337#[cfg(test)]
338mod tests {
339 use super::*;
340 use tempfile::TempDir;
341
342 fn fixture() -> (TempDir, Cache) {
343 let td = TempDir::new().unwrap();
344 let cache = Cache::at(td.path().join("anthropic"));
345 cache.ensure_dir().unwrap();
346 (td, cache)
347 }
348
349 #[test]
350 fn ensure_dir_is_idempotent() {
351 let (_td, cache) = fixture();
352 cache.ensure_dir().unwrap();
353 cache.ensure_dir().unwrap();
354 assert!(cache.dir().is_dir());
355 }
356
357 #[test]
358 fn write_then_read_round_trip() {
359 let (_td, cache) = fixture();
360 cache.write_payload(b"hello world").unwrap();
361 let got = cache.maybe_payload().unwrap();
362 assert_eq!(got.as_deref(), Some(&b"hello world"[..]));
363 }
364
365 #[test]
366 fn maybe_payload_returns_none_when_missing() {
367 let (_td, cache) = fixture();
368 assert!(cache.maybe_payload().unwrap().is_none());
369 }
370
371 #[test]
372 fn fresh_payload_respects_ttl() {
373 let (_td, cache) = fixture();
374 cache.write_payload(b"x").unwrap();
375 assert!(
377 cache
378 .fresh_payload(Duration::from_secs(10))
379 .unwrap()
380 .is_some()
381 );
382 assert!(
384 cache
385 .fresh_payload(Duration::from_secs(0))
386 .unwrap()
387 .is_none()
388 );
389 }
390
391 #[test]
392 fn write_clears_stale_marker_and_last_error() {
393 let (_td, cache) = fixture();
394 cache.mark_stale();
395 cache.write_last_error(429, "rate limited");
396 assert!(cache.is_stale());
397 assert!(cache.read_last_error().is_some());
398
399 cache.write_payload(b"fresh").unwrap();
400 assert!(!cache.is_stale());
401 assert!(cache.read_last_error().is_none());
402 }
403
404 #[test]
405 fn fallback_payload_refuses_a_payload_older_than_the_limit() {
406 let (_td, cache) = fixture();
407 cache.write_payload(b"old").unwrap();
408
409 std::thread::sleep(Duration::from_millis(60));
415
416 assert!(cache.maybe_payload().unwrap().is_some());
419
420 assert!(
424 cache
425 .fallback_payload(Duration::from_millis(5))
426 .unwrap()
427 .is_none()
428 );
429
430 assert_eq!(
433 cache.fallback_payload(MAX_STALE).unwrap().as_deref(),
434 Some(&b"old"[..])
435 );
436 }
437
438 #[test]
439 fn last_error_round_trip() {
440 let (_td, cache) = fixture();
441 cache.write_last_error(503, "service unavailable");
442 let (code, msg) = cache.read_last_error().unwrap();
443 assert_eq!(code, 503);
444 assert_eq!(msg, "service unavailable");
445 }
446
447 #[test]
448 fn last_error_with_empty_message_round_trips() {
449 let (_td, cache) = fixture();
450 cache.write_last_error(429, "");
451 let (code, msg) = cache.read_last_error().unwrap();
452 assert_eq!(code, 429);
453 assert_eq!(msg, "");
454 }
455
456 #[test]
457 fn last_error_replaces_401_body_with_credential_neutral_message() {
458 let (_td, cache) = fixture();
459 cache.write_last_error(401, "PANCEA user@example.test <credential>&token");
460
461 let persisted = fs::read_to_string(cache.last_error_path()).unwrap();
462 assert_eq!(persisted, format!("401\n{AUTH_FAILURE_MESSAGE}"));
463 assert!(!persisted.contains("PANCEA"));
464 assert!(!persisted.contains("<credential>"));
465 }
466
467 #[test]
468 fn last_error_replaces_403_body_with_credential_neutral_message() {
469 let (_td, cache) = fixture();
470 cache.write_last_error(403, "PANCEA account@example.test <credential>&token");
471
472 let persisted = fs::read_to_string(cache.last_error_path()).unwrap();
473 assert_eq!(persisted, format!("403\n{AUTH_FAILURE_MESSAGE}"));
474 assert!(!persisted.contains("PANCEA"));
475 assert!(!persisted.contains("<credential>"));
476 }
477
478 #[test]
484 fn write_last_error_returns_exactly_what_a_later_run_would_read() {
485 for (code, raw) in [
486 (401u16, "PANCEA user@example.test <credential>&token"),
487 (403, "PANCEA account@example.test <credential>&token"),
488 (429, "rate limited, retry in 60s"),
489 (500, "bad\x1b]52;c;Y2FuYXJ5\x07field"),
490 ] {
491 let (_td, cache) = fixture();
492 let returned = cache.write_last_error(code, raw);
493 assert_eq!(
494 returned,
495 cache.read_last_error().unwrap(),
496 "returned pair diverged from the persisted one for {code}"
497 );
498 }
499 }
500
501 #[test]
511 fn no_vendor_builds_a_last_error_pair_from_a_raw_http_body() {
512 let mut sites = Vec::new();
513 for file in crate::guard::rs_files_in("src") {
514 if !file.ends_with("fetch.rs") {
515 continue;
516 }
517 let source = std::fs::read_to_string(&file).expect("readable module");
518 for (n, line) in crate::guard::production_code(&source).lines().enumerate() {
519 if line.contains("(status, body") {
520 sites.push(format!("{}:{}", file.display(), n + 1));
521 }
522 }
523 }
524 assert!(
525 sites.is_empty(),
526 "a last_error pair must be the return of `write_last_error`, which \
527 redacts 401/403 — building one from the raw body puts the response \
528 body in the widget tooltip. Found: {sites:#?}"
529 );
530 }
531
532 #[test]
539 fn no_vendor_replaces_the_original_error_with_a_no_cache_message() {
540 let mut sites = Vec::new();
541 for file in crate::guard::rs_files_in("src") {
542 if !file.ends_with("fetch.rs") {
543 continue;
544 }
545 let source = std::fs::read_to_string(&file).expect("readable module");
546 for (n, line) in crate::guard::production_code(&source).lines().enumerate() {
547 if line.contains("no usable cache") || line.contains("no cache and network") {
548 sites.push(format!("{}:{}", file.display(), n + 1));
549 }
550 }
551 }
552 assert!(
553 sites.is_empty(),
554 "a cold cache must return the error that caused the fetch to fail, not a generic message about the cache. Thread the original `AppError` into the fallback instead. Found: {sites:#?}"
555 );
556 }
557
558 #[test]
563 fn the_returned_pair_carries_the_auth_redaction() {
564 for code in [401u16, 403] {
565 let (_td, cache) = fixture();
566 let (returned_code, msg) =
567 cache.write_last_error(code, "PANCEA user@example.test <credential>&token");
568 assert_eq!(returned_code, code);
569 assert_eq!(msg, AUTH_FAILURE_MESSAGE);
570 assert!(!msg.contains("PANCEA"), "{msg}");
571 assert!(!msg.contains("<credential>"), "{msg}");
572 }
573 }
574
575 #[test]
579 fn last_error_round_trips_a_multi_line_message() {
580 let (_td, cache) = fixture();
581 let body = "{\n \"error\": \"quota exhausted\",\n \"retry_after\": 3600\n}";
582 cache.write_last_error(429, body);
583
584 let (code, msg) = cache.read_last_error().unwrap();
585 assert_eq!(code, 429);
586 assert_eq!(msg, body);
587 assert!(
588 msg.contains("quota exhausted"),
589 "message was truncated to its first line: {msg:?}"
590 );
591 }
592
593 #[test]
594 fn last_error_strips_terminal_controls_before_persisting() {
595 let (_td, cache) = fixture();
596 cache.write_last_error(500, "bad\x1b]52;c;Y2FuYXJ5\x07\nnext\tfield");
597
598 let (code, msg) = cache.read_last_error().unwrap();
599 assert_eq!(code, 500);
600 assert_eq!(msg, "bad]52;c;Y2FuYXJ5\nnext field");
601 assert!(
602 msg.contains("Y2FuYXJ5"),
603 "non-auth diagnostic was not preserved"
604 );
605 assert!(!msg.chars().any(|ch| ch.is_control() && ch != '\n'));
606 }
607
608 #[test]
612 fn last_error_reads_files_written_by_the_previous_version() {
613 let (_td, cache) = fixture();
614
615 fs::write(cache.last_error_path(), "503\nservice unavailable").unwrap();
616 assert_eq!(
617 cache.read_last_error(),
618 Some((503, "service unavailable".into()))
619 );
620
621 fs::write(cache.last_error_path(), "429").unwrap();
622 assert_eq!(cache.read_last_error(), Some((429, String::new())));
623
624 fs::write(cache.last_error_path(), "not-a-code\nboom").unwrap();
626 assert!(cache.read_last_error().is_none());
627 }
628
629 #[test]
630 fn lock_serializes_concurrent_acquirers() {
631 let (_td, cache) = fixture();
634 let lock_path = cache.lock_path();
635 let _guard = acquire_lock(&lock_path, Duration::from_millis(500)).unwrap();
636
637 let res = acquire_lock(&lock_path, Duration::from_millis(100));
638 assert!(matches!(res, Err(AppError::Other(_))));
639 }
640
641 #[tokio::test(flavor = "current_thread")]
647 async fn async_lock_does_not_stall_the_runtime() {
648 let (_td, cache) = fixture();
649 let lock_path = cache.lock_path();
650 let _held = acquire_lock(&lock_path, Duration::from_millis(500)).unwrap();
651
652 let waiter = acquire_lock_async(&lock_path, Duration::from_millis(400));
654
655 let mut ticks = 0usize;
657 let ticker = async {
658 let mut iv = tokio::time::interval(Duration::from_millis(20));
659 iv.tick().await;
660 loop {
661 iv.tick().await;
662 ticks += 1;
663 }
664 };
665
666 tokio::select! {
667 res = waiter => {
668 assert!(matches!(res, Err(AppError::Other(_))));
670 }
671 _ = ticker => unreachable!("the ticker loops forever"),
672 }
673 assert!(
674 ticks > 1,
675 "runtime was starved while the lock was contended ({ticks} ticks)"
676 );
677 }
678
679 #[test]
680 fn atomic_write_creates_parent_dirs() {
681 let td = TempDir::new().unwrap();
682 let nested = td.path().join("a/b/c/file.txt");
683 atomic_write(&nested, b"abc").unwrap();
684 assert_eq!(fs::read(&nested).unwrap(), b"abc");
685 }
686}