1use std::future::Future;
4use std::pin::Pin;
5use std::sync::Arc;
6use std::task::Poll;
7use std::time::{Duration, Instant as WallInstant};
8
9use crate::{StorageError, StorageResult};
10
11pub const DEFAULT_REQUEST_READ_TIMEOUT_SECS: u64 = 30;
13
14pub const MAX_REQUEST_READ_TIMEOUT_SECS: u64 = 86_400;
17
18pub fn request_read_timeout_from_env() -> Duration {
26 let raw = std::env::var("KHIVE_REQUEST_READ_TIMEOUT_SECS").ok();
27 let (seconds, correction) = resolve_request_read_timeout_secs(raw.as_deref());
28 if let Some(correction) = correction {
29 static WARNED: std::sync::Once = std::sync::Once::new();
30 WARNED.call_once(|| {
31 tracing::warn!(
32 value = raw.as_deref().unwrap_or_default(),
33 ceiling_secs = seconds,
34 max_secs = MAX_REQUEST_READ_TIMEOUT_SECS,
35 "KHIVE_REQUEST_READ_TIMEOUT_SECS {correction}"
36 );
37 });
38 }
39 Duration::from_secs(seconds)
40}
41
42fn resolve_request_read_timeout_secs(raw: Option<&str>) -> (u64, Option<&'static str>) {
45 let Some(raw) = raw else {
46 return (DEFAULT_REQUEST_READ_TIMEOUT_SECS, None);
47 };
48 match raw.trim().parse::<u64>() {
49 Ok(seconds) if (1..=MAX_REQUEST_READ_TIMEOUT_SECS).contains(&seconds) => (seconds, None),
50 Ok(seconds) if seconds > MAX_REQUEST_READ_TIMEOUT_SECS => (
51 MAX_REQUEST_READ_TIMEOUT_SECS,
52 Some("is above the maximum; clamped to the maximum"),
53 ),
54 _ => (
55 DEFAULT_REQUEST_READ_TIMEOUT_SECS,
56 Some("is not a whole number of seconds from 1 to the maximum; using the default"),
57 ),
58 }
59}
60
61#[derive(Clone, Copy, Debug)]
67pub struct RequestReadDeadline {
68 async_at: tokio::time::Instant,
69 blocking_at: WallInstant,
70}
71
72impl RequestReadDeadline {
73 pub fn after(duration: Duration) -> Self {
75 Self {
76 async_at: tokio::time::Instant::now() + duration,
77 blocking_at: WallInstant::now() + duration,
78 }
79 }
80
81 pub fn async_at(self) -> tokio::time::Instant {
83 self.async_at
84 }
85
86 pub fn blocking_at(self) -> WallInstant {
88 self.blocking_at
89 }
90
91 fn earlier(self, other: Self) -> Self {
92 if self.async_at <= other.async_at {
93 self
94 } else {
95 other
96 }
97 }
98}
99
100#[derive(Clone, Copy, Debug, Eq, PartialEq)]
102pub enum RequestReadStopReason {
103 Cancelled,
105 Deadline,
107}
108
109#[derive(Clone, Default)]
114pub struct RequestReadContext {
115 cancellations: Arc<[tokio::sync::watch::Receiver<bool>]>,
116 deadline: Option<RequestReadDeadline>,
117 store_acquisition_operation: Option<&'static str>,
118}
119
120impl RequestReadContext {
121 pub fn scope_store_acquisition<T>(
125 mut self,
126 operation: &'static str,
127 work: impl FnOnce() -> T,
128 ) -> T {
129 self.store_acquisition_operation = Some(operation);
130 REQUEST_READ_CONTEXT.sync_scope(self, work)
131 }
132
133 pub fn store_acquisition_operation(&self) -> Option<&'static str> {
135 self.store_acquisition_operation
136 }
137
138 pub fn blocking_stop_reason(&self) -> Option<RequestReadStopReason> {
140 if self.cancellations.iter().any(receiver_cancelled) {
141 Some(RequestReadStopReason::Cancelled)
142 } else if self
143 .deadline
144 .is_some_and(|deadline| WallInstant::now() >= deadline.blocking_at)
145 {
146 Some(RequestReadStopReason::Deadline)
147 } else {
148 None
149 }
150 }
151
152 pub fn deadline(&self) -> Option<RequestReadDeadline> {
154 self.deadline
155 }
156
157 pub fn stop_reason(&self) -> Option<RequestReadStopReason> {
159 if self.cancellations.iter().any(receiver_cancelled) {
160 Some(RequestReadStopReason::Cancelled)
161 } else if self
162 .deadline
163 .is_some_and(|deadline| tokio::time::Instant::now() >= deadline.async_at)
164 {
165 Some(RequestReadStopReason::Deadline)
166 } else {
167 None
168 }
169 }
170
171 pub async fn wait_for_stop(self) -> RequestReadStopReason {
173 let cancellation_receivers = self.cancellations;
174 let request_deadline = self.deadline;
175 let cancellation = async move {
176 match cancellation_receivers.len() {
177 0 => std::future::pending::<()>().await,
178 1 => wait_for_receiver_cancellation(cancellation_receivers[0].clone()).await,
179 2 => {
180 tokio::select! {
181 _ = wait_for_receiver_cancellation(cancellation_receivers[0].clone()) => {},
182 _ = wait_for_receiver_cancellation(cancellation_receivers[1].clone()) => {},
183 }
184 }
185 _ => wait_for_receiver_set(cancellation_receivers).await,
186 }
187 };
188 let deadline = async move {
189 match request_deadline {
190 Some(deadline) => tokio::time::sleep_until(deadline.async_at).await,
191 None => std::future::pending::<()>().await,
192 }
193 };
194 tokio::select! {
195 _ = cancellation => RequestReadStopReason::Cancelled,
196 _ = deadline => RequestReadStopReason::Deadline,
197 }
198 }
199}
200
201tokio::task_local! {
202 static REQUEST_READ_CONTEXT: RequestReadContext;
203}
204
205pub fn capture_request_read_context() -> RequestReadContext {
207 REQUEST_READ_CONTEXT
208 .try_with(Clone::clone)
209 .unwrap_or_default()
210}
211
212pub fn effective_request_read_deadline(candidate: RequestReadDeadline) -> RequestReadDeadline {
215 capture_request_read_context()
216 .deadline
217 .map_or(candidate, |existing| existing.earlier(candidate))
218}
219
220pub async fn scope_request_read_cancellation<F>(
224 cancellation: tokio::sync::watch::Receiver<bool>,
225 future: F,
226) -> F::Output
227where
228 F: Future,
229{
230 let mut context = capture_request_read_context();
231 let mut cancellations = Vec::with_capacity(context.cancellations.len() + 1);
232 cancellations.extend(context.cancellations.iter().cloned());
233 cancellations.push(cancellation);
234 context.cancellations = cancellations.into();
235 REQUEST_READ_CONTEXT.scope(context, future).await
236}
237
238pub async fn scope_request_read_deadline<F>(duration: Duration, future: F) -> F::Output
241where
242 F: Future,
243{
244 scope_request_read_deadline_at(RequestReadDeadline::after(duration), future).await
245}
246
247pub async fn scope_request_read_deadline_at<F>(
249 deadline: RequestReadDeadline,
250 future: F,
251) -> F::Output
252where
253 F: Future,
254{
255 let mut context = capture_request_read_context();
256 context.deadline = Some(match context.deadline {
257 Some(existing) => existing.earlier(deadline),
258 None => deadline,
259 });
260 REQUEST_READ_CONTEXT.scope(context, future).await
261}
262
263pub fn inherit_request_read_context<F>(future: F) -> impl Future<Output = F::Output> + Send
265where
266 F: Future + Send,
267{
268 let inherited = REQUEST_READ_CONTEXT.try_with(Clone::clone).ok();
269 async move {
270 match inherited {
271 Some(context) => REQUEST_READ_CONTEXT.scope(context, future).await,
272 None => future.await,
273 }
274 }
275}
276
277pub fn inherit_request_read_cancellation<F>(future: F) -> impl Future<Output = F::Output> + Send
279where
280 F: Future + Send,
281{
282 inherit_request_read_context(future)
283}
284
285pub fn request_read_is_cancelled() -> bool {
287 capture_request_read_context().stop_reason().is_some()
288}
289
290pub fn ensure_request_read_active(operation: &'static str) -> StorageResult<()> {
292 if request_read_is_cancelled() {
293 Err(StorageError::Timeout {
294 operation: operation.into(),
295 })
296 } else {
297 Ok(())
298 }
299}
300
301pub async fn wait_for_request_read_cancellation() {
303 let _ = capture_request_read_context().wait_for_stop().await;
304}
305
306pub async fn await_request_read_phase<F>(
308 operation: &'static str,
309 future: F,
310) -> StorageResult<F::Output>
311where
312 F: Future,
313{
314 ensure_request_read_active(operation)?;
315 tokio::pin!(future);
316 tokio::select! {
317 biased;
318 _ = wait_for_request_read_cancellation() => Err(StorageError::Timeout {
319 operation: operation.into(),
320 }),
321 output = &mut future => {
322 ensure_request_read_active(operation)?;
323 Ok(output)
324 }
325 }
326}
327
328fn receiver_cancelled(receiver: &tokio::sync::watch::Receiver<bool>) -> bool {
329 *receiver.borrow() || receiver.has_changed().is_err()
330}
331
332async fn wait_for_receiver_cancellation(mut receiver: tokio::sync::watch::Receiver<bool>) {
333 loop {
334 if *receiver.borrow_and_update() {
335 return;
336 }
337 if receiver.changed().await.is_err() {
338 return;
339 }
340 }
341}
342
343async fn wait_for_receiver_set(receivers: Arc<[tokio::sync::watch::Receiver<bool>]>) {
344 let mut waits: Vec<Pin<Box<dyn Future<Output = ()> + Send>>> = receivers
345 .iter()
346 .cloned()
347 .map(|receiver| {
348 Box::pin(wait_for_receiver_cancellation(receiver))
349 as Pin<Box<dyn Future<Output = ()> + Send>>
350 })
351 .collect();
352 std::future::poll_fn(move |cx| {
353 if waits
354 .iter_mut()
355 .any(|wait| wait.as_mut().poll(cx).is_ready())
356 {
357 Poll::Ready(())
358 } else {
359 Poll::Pending
360 }
361 })
362 .await;
363}
364
365#[cfg(test)]
366mod tests {
367 use super::*;
368
369 #[test]
370 fn read_ceiling_accepts_up_to_one_day_and_corrects_loudly_past_it() {
371 let cases: [(Option<&str>, u64, bool); 9] = [
372 (None, DEFAULT_REQUEST_READ_TIMEOUT_SECS, false),
373 (Some("1"), 1, false),
374 (Some("3600"), 3_600, false),
375 (Some("7200"), 7_200, false),
376 (Some("86400"), MAX_REQUEST_READ_TIMEOUT_SECS, false),
377 (Some("86401"), MAX_REQUEST_READ_TIMEOUT_SECS, true),
378 (Some("0"), DEFAULT_REQUEST_READ_TIMEOUT_SECS, true),
379 (Some("-5"), DEFAULT_REQUEST_READ_TIMEOUT_SECS, true),
380 (Some("two hours"), DEFAULT_REQUEST_READ_TIMEOUT_SECS, true),
381 ];
382 for (raw, seconds, corrected) in cases {
383 let (got, correction) = resolve_request_read_timeout_secs(raw);
384 assert_eq!(got, seconds, "ceiling for {raw:?}");
385 assert_eq!(correction.is_some(), corrected, "correction for {raw:?}");
386 }
387 }
388
389 #[tokio::test]
390 async fn nested_scopes_merge_cancellation_sources() {
391 let (outer_tx, outer_rx) = tokio::sync::watch::channel(false);
392 let (_inner_tx, inner_rx) = tokio::sync::watch::channel(false);
393
394 let stopped = scope_request_read_cancellation(
395 outer_rx,
396 scope_request_read_cancellation(inner_rx, async move {
397 outer_tx.send(true).expect("outer scope remains live");
398 wait_for_request_read_cancellation().await;
399 request_read_is_cancelled()
400 }),
401 )
402 .await;
403
404 assert!(stopped);
405 }
406
407 #[tokio::test]
408 async fn sender_loss_is_request_abandonment() {
409 let (tx, rx) = tokio::sync::watch::channel(false);
410 drop(tx);
411
412 let stopped = scope_request_read_cancellation(rx, async {
413 capture_request_read_context().stop_reason()
414 })
415 .await;
416
417 assert_eq!(stopped, Some(RequestReadStopReason::Cancelled));
418 }
419
420 #[tokio::test]
421 async fn spawned_child_inherits_the_same_cancellation() {
422 let (tx, rx) = tokio::sync::watch::channel(false);
423
424 let stopped = scope_request_read_cancellation(rx, async move {
425 let child = tokio::spawn(inherit_request_read_context(async {
426 wait_for_request_read_cancellation().await;
427 request_read_is_cancelled()
428 }));
429 tx.send(true).expect("child receiver remains live");
430 child.await.expect("child task")
431 })
432 .await;
433
434 assert!(stopped);
435 }
436
437 #[tokio::test]
438 async fn nested_deadline_keeps_the_earlier_absolute_instant() {
439 let earlier = RequestReadDeadline::after(Duration::from_secs(10));
440 let later = RequestReadDeadline::after(Duration::from_secs(20));
441
442 let effective = scope_request_read_deadline_at(earlier, async {
443 effective_request_read_deadline(later)
444 })
445 .await;
446
447 assert_eq!(effective.async_at(), earlier.async_at());
448 }
449
450 #[tokio::test]
451 async fn async_phase_cannot_degrade_cancelled_work_to_success() {
452 let (tx, rx) = tokio::sync::watch::channel(false);
453 tx.send(true).expect("receiver remains live");
454
455 let result = scope_request_read_cancellation(rx, async {
456 await_request_read_phase("test.phase", async { 7_u8 }).await
457 })
458 .await;
459
460 assert!(matches!(result, Err(StorageError::Timeout { .. })));
461 }
462}