Skip to main content

rusty_cat/presigned/
range_download.rs

1use std::cell::RefCell;
2use std::panic::{catch_unwind, AssertUnwindSafe};
3use std::sync::{Arc, Mutex as StdMutex, RwLock, Weak};
4
5use reqwest::header::{HeaderMap, HeaderValue, ACCEPT, RANGE};
6
7use crate::download_trait::{BreakpointDownload, DownloadHeadCtx, DownloadRangeGetCtx};
8use crate::error::{InnerErrorCode, MeowError};
9use crate::TransferTask;
10
11use super::time::now_unix_secs;
12use super::{PresignedDownloadUrlRefresher, PresignedRangeDownloadPlan};
13
14thread_local! {
15    // A stack, rather than one current key, also detects indirect recursion
16    // such as A refresher -> B refresher -> A refresher.
17    static ACTIVE_RANGE_REFRESH: RefCell<Vec<usize>> = const { RefCell::new(Vec::new()) };
18    // `BreakpointDownload` intentionally keeps its existing public API, where
19    // range headers and URL are obtained by two consecutive calls. The default
20    // backend performs both calls in one blocking closure. Remember the URL
21    // paired with the just-built headers on that thread so a concurrent plan
22    // refresh cannot mix two credential generations.
23    static PENDING_RANGE_URL: RefCell<Option<(Weak<StdMutex<()>>, String)>> = const { RefCell::new(None) };
24}
25
26#[derive(Debug)]
27struct ActiveRefreshGuard {
28    key: usize,
29}
30
31impl ActiveRefreshGuard {
32    fn enter(key: usize) -> Result<Self, MeowError> {
33        ACTIVE_RANGE_REFRESH.with(|active| {
34            let mut active = active.borrow_mut();
35            if active.contains(&key) {
36                return Err(MeowError::from_code_str(
37                    InnerErrorCode::InvalidTaskState,
38                    "presigned range URL refresher re-entered the same download",
39                ));
40            }
41            active.push(key);
42            Ok(Self { key })
43        })
44    }
45}
46
47impl Drop for ActiveRefreshGuard {
48    fn drop(&mut self) {
49        ACTIVE_RANGE_REFRESH.with(|active| {
50            let mut active = active.borrow_mut();
51            if let Some(index) = active.iter().rposition(|key| *key == self.key) {
52                active.remove(index);
53            }
54        });
55    }
56}
57
58/// Provider-neutral presigned range-download implementation.
59#[derive(Clone)]
60pub struct PresignedRangeDownload {
61    plan: Arc<RwLock<Arc<PresignedRangeDownloadPlan>>>,
62    refresh_gate: Arc<StdMutex<()>>,
63    url_refresher: Option<Arc<dyn PresignedDownloadUrlRefresher>>,
64}
65
66impl PresignedRangeDownload {
67    /// Creates a download protocol from a plan.
68    pub fn new(plan: PresignedRangeDownloadPlan) -> Self {
69        Self {
70            plan: Arc::new(RwLock::new(Arc::new(plan))),
71            refresh_gate: Arc::new(StdMutex::new(())),
72            url_refresher: None,
73        }
74    }
75
76    /// Adds a synchronous URL refresher used when range URL is expired or close
77    /// to expiry.
78    pub fn with_url_refresher(mut self, refresher: Arc<dyn PresignedDownloadUrlRefresher>) -> Self {
79        self.url_refresher = Some(refresher);
80        self
81    }
82
83    /// Returns a snapshot of the current download plan.
84    pub fn plan(&self) -> Result<PresignedRangeDownloadPlan, MeowError> {
85        self.plan.read().map(|g| g.as_ref().clone()).map_err(|_| {
86            MeowError::from_code_str(
87                InnerErrorCode::InvalidTaskState,
88                "presigned range download plan lock poisoned",
89            )
90        })
91    }
92
93    fn merge_headers(target: &mut HeaderMap, extra: &HeaderMap) {
94        for (k, v) in extra {
95            target.insert(k.clone(), v.clone());
96        }
97    }
98
99    fn should_refresh_plan(plan: &PresignedRangeDownloadPlan) -> Result<bool, MeowError> {
100        let Some(expires_at) = plan.range_expires_at_unix_secs else {
101            return Ok(false);
102        };
103        Ok(now_unix_secs()?.saturating_add(plan.refresh_before_secs) >= expires_at)
104    }
105
106    fn is_plan_expired(plan: &PresignedRangeDownloadPlan) -> Result<bool, MeowError> {
107        let Some(expires_at) = plan.range_expires_at_unix_secs else {
108            return Ok(false);
109        };
110        Ok(now_unix_secs()? >= expires_at)
111    }
112
113    fn plan_snapshot(&self) -> Result<Arc<PresignedRangeDownloadPlan>, MeowError> {
114        self.plan.read().map(|plan| Arc::clone(&plan)).map_err(|_| {
115            MeowError::from_code_str(
116                InnerErrorCode::InvalidTaskState,
117                "presigned range download plan lock poisoned",
118            )
119        })
120    }
121
122    fn ensure_fresh_snapshot(&self) -> Result<Arc<PresignedRangeDownloadPlan>, MeowError> {
123        let plan = self.plan_snapshot()?;
124        if !Self::should_refresh_plan(&plan)? {
125            return Ok(plan);
126        }
127
128        let Some(refresher) = &self.url_refresher else {
129            if Self::is_plan_expired(&plan)? {
130                crate::log::emit_lazy(|| {
131                    crate::log::Log::error(
132                        "range_get",
133                        "presigned range URL expired and no refresher is configured",
134                    )
135                    .with_url(plan.range_url.as_str())
136                });
137                return Err(MeowError::from_code_str(
138                    InnerErrorCode::InvalidTaskState,
139                    "presigned range URL expired and no refresher is configured",
140                ));
141            }
142            return Ok(plan);
143        };
144
145        let refresh_key = Arc::as_ptr(&self.refresh_gate) as usize;
146        if ACTIVE_RANGE_REFRESH.with(|active| active.borrow().contains(&refresh_key)) {
147            return Err(MeowError::from_code_str(
148                InnerErrorCode::InvalidTaskState,
149                "presigned range URL refresher re-entered the same download",
150            ));
151        }
152
153        // Only one caller refreshes a stale generation. Re-check after taking
154        // the gate because another range part may have already published a new
155        // immutable snapshot while this caller waited.
156        let _refresh = self.refresh_gate.lock().map_err(|_| {
157            MeowError::from_code_str(
158                InnerErrorCode::InvalidTaskState,
159                "presigned range refresh lock poisoned",
160            )
161        })?;
162        let plan = self.plan_snapshot()?;
163        if !Self::should_refresh_plan(&plan)? {
164            return Ok(plan);
165        }
166
167        let _active_refresh = ActiveRefreshGuard::enter(refresh_key)?;
168        let refresh_result =
169            catch_unwind(AssertUnwindSafe(|| refresher.refresh_range_download(&plan))).map_err(
170                |_| {
171                    MeowError::from_code_str(
172                        InnerErrorCode::InvalidTaskState,
173                        "presigned range URL refresher panicked",
174                    )
175                },
176            )?;
177        let mut refreshed = refresh_result.inspect_err(|e| {
178            crate::log::emit_lazy(|| {
179                crate::log::Log::error(
180                    "range_get",
181                    format!(
182                        "presigned range URL refresh/re-sign failed: {}",
183                        crate::log::redact_secrets(&e.to_string())
184                    ),
185                )
186                .with_url(plan.range_url.as_str())
187            });
188        })?;
189        if let (Some(old), Some(new)) = (plan.total_size, refreshed.total_size) {
190            if old != new {
191                crate::log::emit_lazy(|| {
192                    crate::log::Log::error(
193                        "range_get",
194                        format!("refreshed range total_size mismatch: old={old} new={new}"),
195                    )
196                    .with_url(plan.range_url.as_str())
197                });
198                return Err(MeowError::from_code(
199                    InnerErrorCode::InvalidTaskState,
200                    format!("refreshed range total_size mismatch: old={old} new={new}"),
201                ));
202            }
203        }
204        if refreshed.total_size.is_none() {
205            refreshed.total_size = plan.total_size;
206        }
207        let refreshed = Arc::new(refreshed);
208        let mut guard = self.plan.write().map_err(|_| {
209            MeowError::from_code_str(
210                InnerErrorCode::InvalidTaskState,
211                "presigned range download plan lock poisoned",
212            )
213        })?;
214        *guard = Arc::clone(&refreshed);
215        Ok(refreshed)
216    }
217
218    fn request_from_fresh_snapshot(
219        &self,
220        range_value: &str,
221        mut base: HeaderMap,
222    ) -> Result<(String, HeaderMap), MeowError> {
223        let plan = self.ensure_fresh_snapshot()?;
224        base.insert(
225            RANGE,
226            HeaderValue::from_str(range_value).map_err(|e| {
227                let detail = format!("invalid range header value '{range_value}': {e}");
228                crate::log::emit_lazy({
229                    let detail = detail.clone();
230                    move || crate::log::Log::warn("range_get", detail)
231                });
232                MeowError::from_code(InnerErrorCode::ParameterEmpty, detail)
233            })?,
234        );
235        if !base.contains_key(ACCEPT) {
236            base.insert(
237                ACCEPT,
238                HeaderValue::from_static(crate::http_breakpoint::DEFAULT_RANGE_ACCEPT),
239            );
240        }
241        Self::merge_headers(&mut base, &plan.range_headers);
242        Ok((plan.range_url.clone(), base))
243    }
244
245    fn remember_range_url(&self, range_url: String) {
246        let key = Arc::downgrade(&self.refresh_gate);
247        PENDING_RANGE_URL.with(|pending| {
248            *pending.borrow_mut() = Some((key, range_url));
249        });
250    }
251
252    fn take_remembered_range_url(&self) -> Option<String> {
253        let key = Arc::downgrade(&self.refresh_gate);
254        PENDING_RANGE_URL.with(|pending| {
255            let mut pending = pending.borrow_mut();
256            match pending.as_ref() {
257                Some((pending_key, _)) if Weak::ptr_eq(pending_key, &key) => {
258                    pending.take().map(|(_, url)| url)
259                }
260                _ => None,
261            }
262        })
263    }
264
265    #[cfg(test)]
266    pub(crate) fn ensure_fresh_plan(&self) -> Result<PresignedRangeDownloadPlan, MeowError> {
267        self.ensure_fresh_snapshot()
268            .map(|plan| plan.as_ref().clone())
269    }
270}
271
272impl BreakpointDownload for PresignedRangeDownload {
273    fn total_size_hint(&self, _task: &TransferTask) -> Option<u64> {
274        self.plan_snapshot().ok().and_then(|plan| plan.total_size)
275    }
276
277    fn head_url(&self, task: &TransferTask) -> String {
278        self.plan_snapshot()
279            .ok()
280            .and_then(|plan| plan.head_url.clone())
281            .unwrap_or_else(|| task.url().to_string())
282    }
283
284    fn range_url(&self, _task: &TransferTask) -> String {
285        self.take_remembered_range_url().unwrap_or_else(|| {
286            self.plan_snapshot()
287                .map(|plan| plan.range_url.clone())
288                .unwrap_or_default()
289        })
290    }
291
292    fn merge_head_headers(&self, ctx: DownloadHeadCtx<'_>) -> Result<(), MeowError> {
293        let plan = self.plan()?;
294        Self::merge_headers(ctx.base, &plan.head_headers);
295        Ok(())
296    }
297
298    fn merge_range_get_headers(&self, ctx: DownloadRangeGetCtx<'_>) -> Result<(), MeowError> {
299        let (range_url, headers) =
300            self.request_from_fresh_snapshot(ctx.range_value, ctx.base.clone())?;
301        *ctx.base = headers;
302        self.remember_range_url(range_url);
303        Ok(())
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use std::sync::atomic::{AtomicUsize, Ordering};
310    use std::sync::{Arc, Barrier, Mutex};
311
312    use super::*;
313
314    struct CountingRefresher {
315        calls: AtomicUsize,
316    }
317
318    impl PresignedDownloadUrlRefresher for CountingRefresher {
319        fn refresh_range_download(
320            &self,
321            plan: &PresignedRangeDownloadPlan,
322        ) -> Result<PresignedRangeDownloadPlan, MeowError> {
323            self.calls.fetch_add(1, Ordering::SeqCst);
324            std::thread::sleep(std::time::Duration::from_millis(20));
325            Ok(plan
326                .clone()
327                .with_range_expires_at_unix_secs(now_unix_secs()? + 3600))
328        }
329    }
330
331    #[test]
332    fn concurrent_expired_readers_refresh_only_once_and_see_complete_snapshot() {
333        let refresher = Arc::new(CountingRefresher {
334            calls: AtomicUsize::new(0),
335        });
336        let download = Arc::new(
337            PresignedRangeDownload::new(
338                PresignedRangeDownloadPlan::new("https://example.com/object")
339                    .with_total_size(42)
340                    .with_range_expires_at_unix_secs(1)
341                    .with_refresh_before_secs(60),
342            )
343            .with_url_refresher(refresher.clone()),
344        );
345        let barrier = Arc::new(Barrier::new(16));
346        let mut threads = Vec::new();
347        for _ in 0..16 {
348            let download = Arc::clone(&download);
349            let barrier = Arc::clone(&barrier);
350            threads.push(std::thread::spawn(move || {
351                barrier.wait();
352                let plan = download.ensure_fresh_plan().expect("fresh plan");
353                assert_eq!(plan.total_size, Some(42));
354                assert_eq!(plan.range_url, "https://example.com/object");
355                assert!(plan.range_expires_at_unix_secs.unwrap_or(0) > 1);
356            }));
357        }
358        for thread in threads {
359            thread.join().expect("reader");
360        }
361        assert_eq!(refresher.calls.load(Ordering::SeqCst), 1);
362    }
363
364    struct SnapshotRefresher;
365
366    impl PresignedDownloadUrlRefresher for SnapshotRefresher {
367        fn refresh_range_download(
368            &self,
369            plan: &PresignedRangeDownloadPlan,
370        ) -> Result<PresignedRangeDownloadPlan, MeowError> {
371            let mut refreshed = plan.clone();
372            refreshed.range_url = "https://new.example.com/object".to_owned();
373            refreshed
374                .range_headers
375                .insert("x-plan-generation", HeaderValue::from_static("new"));
376            refreshed.range_expires_at_unix_secs = Some(now_unix_secs()? + 3600);
377            Ok(refreshed)
378        }
379    }
380
381    #[test]
382    fn range_url_and_headers_are_built_from_one_refreshed_snapshot() {
383        let download = PresignedRangeDownload::new(
384            PresignedRangeDownloadPlan::new("https://old.example.com/object")
385                .with_range_expires_at_unix_secs(1),
386        )
387        .with_url_refresher(Arc::new(SnapshotRefresher));
388
389        let (url, headers) = download
390            .request_from_fresh_snapshot("bytes=0-9", HeaderMap::new())
391            .expect("range request");
392        assert_eq!(url, "https://new.example.com/object");
393        assert_eq!(headers.get("x-plan-generation").unwrap(), "new");
394        assert_eq!(headers.get(RANGE).unwrap(), "bytes=0-9");
395
396        download.remember_range_url(url);
397        *download.plan.write().expect("plan write") = Arc::new(
398            PresignedRangeDownloadPlan::new("https://later.example.com/object")
399                .with_range_expires_at_unix_secs(now_unix_secs().unwrap() + 7200),
400        );
401        assert_eq!(
402            download.take_remembered_range_url().as_deref(),
403            Some("https://new.example.com/object"),
404            "the URL paired with the headers must survive a concurrent plan publication"
405        );
406        assert!(download.take_remembered_range_url().is_none());
407    }
408
409    struct ReentrantRefresher {
410        download: Mutex<Option<PresignedRangeDownload>>,
411    }
412
413    impl PresignedDownloadUrlRefresher for ReentrantRefresher {
414        fn refresh_range_download(
415            &self,
416            _plan: &PresignedRangeDownloadPlan,
417        ) -> Result<PresignedRangeDownloadPlan, MeowError> {
418            self.download
419                .lock()
420                .map_err(|_| {
421                    MeowError::from_code_str(
422                        InnerErrorCode::InvalidTaskState,
423                        "reentrant test lock poisoned",
424                    )
425                })?
426                .as_ref()
427                .ok_or_else(|| {
428                    MeowError::from_code_str(
429                        InnerErrorCode::InvalidTaskState,
430                        "reentrant test download missing",
431                    )
432                })?
433                .ensure_fresh_plan()
434        }
435    }
436
437    #[test]
438    fn refresher_reentry_returns_error_instead_of_deadlocking() {
439        let refresher = Arc::new(ReentrantRefresher {
440            download: Mutex::new(None),
441        });
442        let download = PresignedRangeDownload::new(
443            PresignedRangeDownloadPlan::new("https://example.com/object")
444                .with_range_expires_at_unix_secs(1),
445        )
446        .with_url_refresher(refresher.clone());
447        *refresher.download.lock().expect("refresher download") = Some(download.clone());
448
449        let (sender, receiver) = std::sync::mpsc::channel();
450        let thread = std::thread::spawn(move || {
451            let _ = sender.send(download.ensure_fresh_plan());
452        });
453        let error = receiver
454            .recv_timeout(std::time::Duration::from_secs(1))
455            .expect("reentrant refresh must return promptly")
456            .expect_err("reentrant refresh must fail");
457        assert!(error.to_string().contains("re-entered"));
458        thread.join().expect("refresh thread");
459    }
460
461    #[test]
462    fn indirect_same_thread_refresh_reentry_is_rejected() {
463        let first = ActiveRefreshGuard::enter(11).expect("enter first refresher");
464        let second = ActiveRefreshGuard::enter(22).expect("enter nested refresher");
465        let error = ActiveRefreshGuard::enter(11).expect_err("A -> B -> A must be rejected");
466        assert!(error.to_string().contains("re-entered"));
467        drop(second);
468        drop(first);
469        ActiveRefreshGuard::enter(11).expect("guards must clean up their stack on drop");
470    }
471}