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 static ACTIVE_RANGE_REFRESH: RefCell<Vec<usize>> = const { RefCell::new(Vec::new()) };
18 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#[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 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 pub fn with_url_refresher(mut self, refresher: Arc<dyn PresignedDownloadUrlRefresher>) -> Self {
79 self.url_refresher = Some(refresher);
80 self
81 }
82
83 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 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}