1use std::{
2 fs::{self, File},
3 io::{Read, Write},
4 path::{Path, PathBuf},
5 time::{Duration, SystemTime, UNIX_EPOCH},
6};
7
8use serde::{Deserialize, Serialize};
9
10use crate::{Error, ProxyNode, Result, subscription::SubscriptionValidators};
11
12const SCHEMA_VERSION: u32 = 1;
13const MAX_CACHE_BYTES: u64 = 8 * 1024 * 1024;
14
15#[derive(Clone, Debug)]
24pub struct CachePolicy {
25 pub directory: PathBuf,
27 pub ttl: Duration,
29 pub max_stale: Duration,
32}
33
34impl CachePolicy {
35 pub fn new(directory: impl Into<PathBuf>) -> Self {
38 Self {
39 directory: directory.into(),
40 ttl: Duration::from_secs(3 * 24 * 60 * 60),
41 max_stale: Duration::from_secs(7 * 24 * 60 * 60),
42 }
43 }
44
45 pub fn validate(&self) -> Result<()> {
47 if self.ttl.is_zero() {
48 return Err(Error::Config("cache TTL must be nonzero"));
49 }
50 if self.ttl.checked_add(self.max_stale).is_none() {
51 return Err(Error::Config("cache lifetime is too large"));
52 }
53 Ok(())
54 }
55}
56
57#[derive(Clone)]
58pub(crate) struct CacheStore {
59 policy: CachePolicy,
60}
61
62pub(crate) struct CachedNodes {
63 pub nodes: Vec<ProxyNode>,
64 pub fresh: bool,
65 pub validators: SubscriptionValidators,
66}
67
68#[derive(Serialize, Deserialize)]
69#[serde(deny_unknown_fields)]
70struct CacheEntry {
71 schema: u32,
72 source_key: String,
73 saved_at: SystemTime,
74 nodes: Vec<ProxyNode>,
75 #[serde(default)]
77 validators: SubscriptionValidators,
78}
79
80impl CacheStore {
81 pub(crate) fn new(policy: CachePolicy) -> Self {
82 Self { policy }
83 }
84
85 pub(crate) async fn load(&self, source_key: &str) -> Result<Option<CachedNodes>> {
86 self.policy.validate()?;
87 validate_source_key(source_key)?;
88 let policy = self.policy.clone();
89 let source_key = source_key.to_owned();
90 tokio::task::spawn_blocking(move || load(&policy, &source_key, SystemTime::now()))
91 .await
92 .map_err(|_| Error::Cache("cache read task failed"))?
93 }
94
95 #[cfg(test)]
96 pub(crate) async fn save(&self, source_key: &str, nodes: &[ProxyNode]) -> Result<()> {
97 self.save_with_validators(source_key, nodes, &SubscriptionValidators::default())
98 .await
99 }
100
101 pub(crate) async fn save_with_validators(
104 &self,
105 source_key: &str,
106 nodes: &[ProxyNode],
107 validators: &SubscriptionValidators,
108 ) -> Result<()> {
109 self.policy.validate()?;
110 validate_source_key(source_key)?;
111 validate_nodes(nodes)?;
112 let directory = self.policy.directory.clone();
113 let entry = CacheEntry {
114 schema: SCHEMA_VERSION,
115 source_key: source_key.to_owned(),
116 saved_at: SystemTime::now(),
117 nodes: nodes.to_vec(),
118 validators: validators.clone(),
119 };
120 tokio::task::spawn_blocking(move || save(&directory, &entry))
121 .await
122 .map_err(|_| Error::Cache("cache write task failed"))?
123 }
124}
125
126fn validate_source_key(source_key: &str) -> Result<()> {
127 if source_key.len() != 64
128 || !source_key
129 .bytes()
130 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
131 {
132 return Err(Error::Cache(
133 "cache source key must be a lowercase SHA-256 digest",
134 ));
135 }
136 Ok(())
137}
138
139fn validate_nodes(nodes: &[ProxyNode]) -> Result<()> {
140 if nodes.is_empty() || nodes.iter().any(|node| node.validate().is_err()) {
141 return Err(Error::Cache("cache must contain valid proxy nodes"));
142 }
143 Ok(())
144}
145
146fn cache_path(directory: &Path, source_key: &str) -> PathBuf {
147 directory.join(format!("{source_key}.json"))
148}
149
150fn load(policy: &CachePolicy, source_key: &str, now: SystemTime) -> Result<Option<CachedNodes>> {
151 let path = cache_path(&policy.directory, source_key);
152 let metadata = match fs::symlink_metadata(&path) {
153 Ok(metadata) => metadata,
154 Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
155 Err(error) => return Err(cache_io_error(error)),
156 };
157 if !metadata.is_file() {
159 return Ok(None);
160 }
161 let file = match File::open(path) {
162 Ok(file) => file,
163 Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
164 Err(error) => return Err(cache_io_error(error)),
165 };
166 let metadata = file.metadata().map_err(cache_io_error)?;
167 if !metadata.is_file() || metadata.len() > MAX_CACHE_BYTES {
168 return Ok(None);
169 }
170 let mut bytes = Vec::new();
171 file.take(MAX_CACHE_BYTES + 1)
172 .read_to_end(&mut bytes)
173 .map_err(cache_io_error)?;
174 if bytes.len() as u64 > MAX_CACHE_BYTES {
175 return Ok(None);
176 }
177 let entry: CacheEntry = match serde_json::from_slice(&bytes) {
178 Ok(entry) => entry,
179 Err(_) => return Ok(None),
180 };
181 if entry.schema != SCHEMA_VERSION
182 || entry.source_key != source_key
183 || entry.saved_at.duration_since(UNIX_EPOCH).is_err()
184 || validate_nodes(&entry.nodes).is_err()
185 || !entry.validators.is_valid()
186 {
187 return Ok(None);
188 }
189 let age = match now.duration_since(entry.saved_at) {
190 Ok(age) => age,
191 Err(_) => return Ok(None),
192 };
193 let fresh = age < policy.ttl;
194 let stale_deadline = policy
195 .ttl
196 .checked_add(policy.max_stale)
197 .ok_or(Error::Config("cache lifetime is too large"))?;
198 if !fresh && (policy.max_stale.is_zero() || age >= stale_deadline) {
199 return Ok(None);
200 }
201 Ok(Some(CachedNodes {
202 nodes: entry.nodes,
203 fresh,
204 validators: entry.validators,
205 }))
206}
207
208fn save(directory: &Path, entry: &CacheEntry) -> Result<()> {
209 let bytes =
210 serde_json::to_vec(entry).map_err(|_| Error::Cache("cache serialization failed"))?;
211 if bytes.len() as u64 > MAX_CACHE_BYTES {
212 return Err(Error::Cache("cache exceeds maximum file size"));
213 }
214 let mut directories = fs::DirBuilder::new();
215 directories.recursive(true);
216 #[cfg(unix)]
217 {
218 use std::os::unix::fs::DirBuilderExt;
219 directories.mode(0o700);
220 }
221 directories.create(directory).map_err(cache_io_error)?;
222 let mut temporary = tempfile::Builder::new()
223 .prefix(".subscription-proxy-")
224 .suffix(".tmp")
225 .tempfile_in(directory)
226 .map_err(cache_io_error)?;
227 #[cfg(unix)]
228 {
229 use std::os::unix::fs::PermissionsExt;
230 temporary
231 .as_file()
232 .set_permissions(fs::Permissions::from_mode(0o600))
233 .map_err(cache_io_error)?;
234 }
235 temporary.write_all(&bytes).map_err(cache_io_error)?;
236 temporary.as_file().sync_all().map_err(cache_io_error)?;
237 temporary
238 .persist(cache_path(directory, &entry.source_key))
239 .map_err(|error| cache_io_error(error.error))?;
240 Ok(())
241}
242
243fn cache_io_error(error: std::io::Error) -> Error {
244 Error::Io(match error.raw_os_error() {
247 Some(code) => std::io::Error::from_raw_os_error(code),
248 None => std::io::Error::from(error.kind()),
249 })
250}
251
252#[cfg(test)]
253mod tests {
254 use super::*;
255
256 const KEY: &str = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef";
257 const OTHER_KEY: &str = "abcdef0123456789abcdef0123456789abcdef0123456789abcdef0123456789";
258
259 fn node(port: u16) -> ProxyNode {
260 ProxyNode::from_url(&format!("http://user:password@localhost:{port}")).unwrap()
261 }
262
263 fn entry(saved_at: SystemTime) -> CacheEntry {
264 CacheEntry {
265 schema: SCHEMA_VERSION,
266 source_key: KEY.to_owned(),
267 saved_at,
268 nodes: vec![node(8080)],
269 validators: SubscriptionValidators::default(),
270 }
271 }
272
273 fn put_entry(directory: &Path, entry: &CacheEntry) {
274 fs::write(
275 cache_path(directory, KEY),
276 serde_json::to_vec(entry).unwrap(),
277 )
278 .unwrap();
279 }
280
281 #[test]
282 fn policy_checks_deadlines() {
283 let mut policy = CachePolicy::new("cache");
284 assert_eq!(policy.ttl, Duration::from_secs(3 * 24 * 60 * 60));
285 assert_eq!(policy.max_stale, Duration::from_secs(7 * 24 * 60 * 60));
286 assert!(policy.validate().is_ok());
287 policy.ttl = Duration::ZERO;
288 assert!(policy.validate().is_err());
289 policy.ttl = Duration::MAX;
290 assert!(policy.validate().is_err());
291 policy.max_stale = Duration::ZERO;
292 assert!(policy.validate().is_ok());
293 }
294
295 #[tokio::test]
296 async fn atomic_replacement_preserves_nodes_and_restricts_file_permissions() {
297 let directory = tempfile::tempdir().unwrap();
298 let store = CacheStore::new(CachePolicy::new(directory.path()));
299 assert!(store.load(KEY).await.unwrap().is_none());
300 store.save(KEY, &[node(8080)]).await.unwrap();
301 let original_file = File::open(cache_path(directory.path(), KEY)).unwrap();
302 store.save(KEY, &[node(9090)]).await.unwrap();
303 let loaded = store.load(KEY).await.unwrap().unwrap();
304 assert!(loaded.fresh);
305 assert_eq!(loaded.nodes, vec![node(9090)]);
306 #[cfg(unix)]
308 {
309 use std::os::unix::fs::PermissionsExt;
310 let old: CacheEntry = serde_json::from_reader(original_file).unwrap();
311 assert_eq!(old.nodes, vec![node(8080)]);
312 let mode = fs::metadata(cache_path(directory.path(), KEY))
313 .unwrap()
314 .permissions()
315 .mode();
316 assert_eq!(mode & 0o777, 0o600);
317 }
318 #[cfg(not(unix))]
319 drop(original_file);
320 assert_eq!(fs::read_dir(directory.path()).unwrap().count(), 1);
321 }
322
323 #[tokio::test]
324 async fn validators_round_trip_and_revalidation_renews_freshness() {
325 let directory = tempfile::tempdir().unwrap();
326 let policy = CachePolicy::new(directory.path());
327 let mut old = entry(SystemTime::now() - policy.ttl - Duration::from_secs(1));
328 old.validators = SubscriptionValidators {
329 etag: Some("W/\"revision-1\"".to_owned()),
330 last_modified: Some("Wed, 21 Oct 2015 07:28:00 GMT".to_owned()),
331 };
332 put_entry(directory.path(), &old);
333 let store = CacheStore::new(policy);
334 let cached = store.load(KEY).await.unwrap().unwrap();
335 assert!(!cached.fresh);
336 assert_eq!(cached.validators.etag, old.validators.etag);
337 assert_eq!(
338 cached.validators.last_modified,
339 old.validators.last_modified
340 );
341
342 store
344 .save_with_validators(KEY, &cached.nodes, &cached.validators)
345 .await
346 .unwrap();
347 let cached = store.load(KEY).await.unwrap().unwrap();
348 assert!(cached.fresh);
349 assert_eq!(cached.nodes, old.nodes);
350 assert_eq!(cached.validators.etag, old.validators.etag);
351 assert_eq!(
352 cached.validators.last_modified,
353 old.validators.last_modified
354 );
355 }
356
357 #[tokio::test]
358 async fn original_schema_one_files_without_validators_remain_readable() {
359 let directory = tempfile::tempdir().unwrap();
360 let mut legacy = serde_json::to_value(entry(SystemTime::now())).unwrap();
361 legacy.as_object_mut().unwrap().remove("validators");
362 fs::write(
363 cache_path(directory.path(), KEY),
364 serde_json::to_vec(&legacy).unwrap(),
365 )
366 .unwrap();
367 let store = CacheStore::new(CachePolicy::new(directory.path()));
368 let cached = store.load(KEY).await.unwrap().unwrap();
369 assert!(cached.fresh);
370 assert_eq!(cached.nodes, vec![node(8080)]);
371 assert!(cached.validators.etag.is_none());
372 assert!(cached.validators.last_modified.is_none());
373 }
374
375 #[cfg(unix)]
376 #[tokio::test]
377 async fn new_directories_are_private_without_changing_existing_permissions() {
378 use std::os::unix::fs::PermissionsExt;
379
380 let directory = tempfile::tempdir().unwrap();
381 fs::set_permissions(directory.path(), fs::Permissions::from_mode(0o750)).unwrap();
382 let parent = directory.path().join("new-parent");
383 let child = parent.join("cache");
384 let store = CacheStore::new(CachePolicy::new(&child));
385 store.save(KEY, &[node(8080)]).await.unwrap();
386 for path in [&parent, &child] {
387 assert_eq!(
388 fs::metadata(path).unwrap().permissions().mode() & 0o777,
389 0o700
390 );
391 }
392 assert_eq!(
393 fs::metadata(directory.path()).unwrap().permissions().mode() & 0o777,
394 0o750
395 );
396 }
397
398 #[cfg(unix)]
399 #[tokio::test]
400 async fn symbolic_link_entries_are_misses_and_writes_replace_only_the_link() {
401 use std::os::unix::fs::symlink;
402
403 let directory = tempfile::tempdir().unwrap();
404 let original = directory.path().join("original-secret");
405 let original_bytes = serde_json::to_vec(&entry(SystemTime::now())).unwrap();
406 fs::write(&original, &original_bytes).unwrap();
407 let cached = cache_path(directory.path(), KEY);
408 symlink(&original, &cached).unwrap();
409 let store = CacheStore::new(CachePolicy::new(directory.path()));
410 assert!(store.load(KEY).await.unwrap().is_none());
411 store.save(KEY, &[node(9090)]).await.unwrap();
412 assert_eq!(fs::read(&original).unwrap(), original_bytes);
413 assert!(fs::symlink_metadata(cached).unwrap().is_file());
414 assert_eq!(
415 store.load(KEY).await.unwrap().unwrap().nodes,
416 vec![node(9090)]
417 );
418 }
419
420 #[tokio::test]
421 async fn concurrent_replacements_keep_validators_with_their_nodes() {
422 let directory = tempfile::tempdir().unwrap();
423 let store = CacheStore::new(CachePolicy::new(directory.path()));
424 let mut writers = tokio::task::JoinSet::new();
425 for port in 8080..8096 {
426 let store = store.clone();
427 writers.spawn(async move {
428 store
429 .save_with_validators(
430 KEY,
431 &[node(port)],
432 &SubscriptionValidators {
433 etag: Some(format!("\"{port}\"")),
434 last_modified: None,
435 },
436 )
437 .await
438 .unwrap();
439 let cached = store.load(KEY).await.unwrap().unwrap();
440 let etag = cached.validators.etag.unwrap();
441 let port = etag.trim_matches('"').parse().unwrap();
442 assert_eq!(cached.nodes, vec![node(port)]);
443 });
444 }
445 while let Some(result) = writers.join_next().await {
446 result.unwrap();
447 }
448 assert_eq!(fs::read_dir(directory.path()).unwrap().count(), 1);
449 }
450
451 #[test]
452 fn contextual_io_errors_do_not_reveal_private_paths() {
453 let error = cache_io_error(std::io::Error::new(
454 std::io::ErrorKind::PermissionDenied,
455 "could not create /private/token=secret/cache.json",
456 ));
457 assert!(!format!("{error:?} {error}").contains("secret"));
458 assert!(
459 matches!(error, Error::Io(error) if error.kind() == std::io::ErrorKind::PermissionDenied)
460 );
461 }
462
463 #[test]
464 fn lifetime_and_stale_fallback_have_fixed_boundaries() {
465 let directory = tempfile::tempdir().unwrap();
466 let mut policy = CachePolicy::new(directory.path());
467 policy.ttl = Duration::from_secs(10);
468 policy.max_stale = Duration::from_secs(20);
469 let saved_at = UNIX_EPOCH + Duration::from_secs(100);
470 put_entry(directory.path(), &entry(saved_at));
471
472 assert!(
473 load(&policy, KEY, saved_at + Duration::from_secs(9))
474 .unwrap()
475 .unwrap()
476 .fresh
477 );
478 for age in [10, 29] {
479 assert!(
480 !load(&policy, KEY, saved_at + Duration::from_secs(age))
481 .unwrap()
482 .unwrap()
483 .fresh
484 );
485 }
486 for age in [30, 31] {
487 assert!(
488 load(&policy, KEY, saved_at + Duration::from_secs(age))
489 .unwrap()
490 .is_none()
491 );
492 }
493 policy.max_stale = Duration::ZERO;
494 assert!(
495 load(&policy, KEY, saved_at + Duration::from_secs(10))
496 .unwrap()
497 .is_none()
498 );
499 }
500
501 #[test]
502 fn corrupt_schema_source_time_and_nodes_are_cache_misses() {
503 let directory = tempfile::tempdir().unwrap();
504 let policy = CachePolicy::new(directory.path());
505 let now = SystemTime::now();
506 fs::write(cache_path(directory.path(), KEY), b"{invalid").unwrap();
507 assert!(load(&policy, KEY, now).unwrap().is_none());
508
509 let mut bad_schema = entry(now);
510 bad_schema.schema += 1;
511 let mut wrong_source = entry(now);
512 wrong_source.source_key = OTHER_KEY.to_owned();
513 let future = entry(now + Duration::from_nanos(1));
514 let mut empty = entry(now);
515 empty.nodes.clear();
516 let mut invalid = entry(now);
517 invalid.nodes = vec![
518 serde_json::from_value(serde_json::json!({
519 "name": "secret", "endpoint": "http://localhost:0"
520 }))
521 .unwrap(),
522 ];
523 for bad in [bad_schema, wrong_source, future, empty, invalid] {
524 put_entry(directory.path(), &bad);
525 assert!(load(&policy, KEY, now).unwrap().is_none());
526 }
527 }
528
529 #[tokio::test]
530 async fn validates_source_keys_and_empty_lists_before_writing() {
531 let directory = tempfile::tempdir().unwrap();
532 let store = CacheStore::new(CachePolicy::new(directory.path()));
533 for key in [
534 "../secret",
535 "/tmp/secret",
536 "https://subscription/?token=secret",
537 "",
538 "ABCDEF",
539 ] {
540 assert!(store.load(key).await.is_err());
541 assert!(store.save(key, &[node(8080)]).await.is_err());
542 }
543 assert!(store.save(KEY, &[]).await.is_err());
544 assert_eq!(fs::read_dir(directory.path()).unwrap().count(), 0);
545 }
546
547 #[tokio::test]
548 async fn separate_sources_never_share_nodes() {
549 let directory = tempfile::tempdir().unwrap();
550 let store = CacheStore::new(CachePolicy::new(directory.path()));
551 store.save(KEY, &[node(8080)]).await.unwrap();
552 assert!(store.load(OTHER_KEY).await.unwrap().is_none());
553 store.save(OTHER_KEY, &[node(9090)]).await.unwrap();
554 assert_eq!(
555 store.load(KEY).await.unwrap().unwrap().nodes,
556 vec![node(8080)]
557 );
558 assert_eq!(
559 store.load(OTHER_KEY).await.unwrap().unwrap().nodes,
560 vec![node(9090)]
561 );
562 }
563
564 #[tokio::test]
565 async fn oversized_files_are_misses_and_oversized_writes_preserve_old_cache() {
566 let directory = tempfile::tempdir().unwrap();
567 let store = CacheStore::new(CachePolicy::new(directory.path()));
568 let path = cache_path(directory.path(), KEY);
569 File::create(&path)
570 .unwrap()
571 .set_len(MAX_CACHE_BYTES + 1)
572 .unwrap();
573 assert!(store.load(KEY).await.unwrap().is_none());
574 store.save(KEY, &[node(8080)]).await.unwrap();
575 let large_node = node(9090).with_name("x".repeat(MAX_CACHE_BYTES as usize));
576 assert!(store.save(KEY, &[large_node]).await.is_err());
577 assert_eq!(
578 store.load(KEY).await.unwrap().unwrap().nodes,
579 vec![node(8080)]
580 );
581 }
582
583 #[tokio::test]
584 async fn filesystem_errors_are_reported() {
585 let directory = tempfile::tempdir().unwrap();
586 let path = directory.path().join("file");
587 fs::write(&path, b"not a directory").unwrap();
588 let store = CacheStore::new(CachePolicy::new(path));
589 assert!(matches!(store.load(KEY).await, Err(Error::Io(_))));
590 assert!(matches!(
591 store.save(KEY, &[node(8080)]).await,
592 Err(Error::Io(_))
593 ));
594 }
595}