Skip to main content

subscription_proxy_pool/
cache.rs

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/// Optional persistent subscription cache.
16///
17/// Files contain proxy credentials, but never the subscription URL. Keep the
18/// directory private and outside source control. Each write atomically replaces
19/// the previous file; on Unix, new directories and files have owner-only
20/// permissions (`0700` and `0600`). Existing directory permissions are preserved.
21/// Missing, invalid, and expired files are cache misses; other filesystem errors
22/// are returned to the caller instead of silently hiding configuration problems.
23#[derive(Clone, Debug)]
24pub struct CachePolicy {
25    /// Directory in which hashed subscription keys identify cache files.
26    pub directory: PathBuf,
27    /// How long a successful download or HTTP revalidation remains fresh.
28    pub ttl: Duration,
29    /// Additional time after `ttl` during which stale fallback is permitted.
30    /// Zero disables stale fallback.
31    pub max_stale: Duration,
32}
33
34impl CachePolicy {
35    /// Create a policy with a three-day lifetime and seven additional days of
36    /// stale fallback. A failed download never renews either deadline.
37    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    /// Check that cache deadlines can be represented and the lifetime is nonzero.
46    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    // Version-one cache files written before conditional requests remain usable.
76    #[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    /// Save a successfully downloaded or revalidated representation. Callers
102    /// must not use this to renew stale data after an unsuccessful download.
103    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    // Cache entries are regular files. Do not follow links to unrelated files.
158    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    // Some filesystem helpers include paths in their contextual error strings.
245    // Preserve OS codes or the error kind without retaining caller-supplied paths.
246    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        // An already-open handle keeps the old inode when replacement is atomic.
307        #[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        // A 304 confirms these nodes are still current and renews the deadlines.
343        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}