khive_storage/
request_context.rs1use 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 fn request_read_timeout_from_env() -> Duration {
17 let seconds = std::env::var("KHIVE_REQUEST_READ_TIMEOUT_SECS")
18 .ok()
19 .and_then(|value| value.parse::<u64>().ok())
20 .filter(|seconds| (1..=3_600).contains(seconds))
21 .unwrap_or(DEFAULT_REQUEST_READ_TIMEOUT_SECS);
22 Duration::from_secs(seconds)
23}
24
25#[derive(Clone, Copy, Debug)]
31pub struct RequestReadDeadline {
32 async_at: tokio::time::Instant,
33 blocking_at: WallInstant,
34}
35
36impl RequestReadDeadline {
37 pub fn after(duration: Duration) -> Self {
39 Self {
40 async_at: tokio::time::Instant::now() + duration,
41 blocking_at: WallInstant::now() + duration,
42 }
43 }
44
45 pub fn async_at(self) -> tokio::time::Instant {
47 self.async_at
48 }
49
50 pub fn blocking_at(self) -> WallInstant {
52 self.blocking_at
53 }
54
55 fn earlier(self, other: Self) -> Self {
56 if self.async_at <= other.async_at {
57 self
58 } else {
59 other
60 }
61 }
62}
63
64#[derive(Clone, Copy, Debug, Eq, PartialEq)]
66pub enum RequestReadStopReason {
67 Cancelled,
69 Deadline,
71}
72
73#[derive(Clone, Default)]
78pub struct RequestReadContext {
79 cancellations: Arc<[tokio::sync::watch::Receiver<bool>]>,
80 deadline: Option<RequestReadDeadline>,
81 store_acquisition_operation: Option<&'static str>,
82}
83
84impl RequestReadContext {
85 pub fn scope_store_acquisition<T>(
89 mut self,
90 operation: &'static str,
91 work: impl FnOnce() -> T,
92 ) -> T {
93 self.store_acquisition_operation = Some(operation);
94 REQUEST_READ_CONTEXT.sync_scope(self, work)
95 }
96
97 pub fn store_acquisition_operation(&self) -> Option<&'static str> {
99 self.store_acquisition_operation
100 }
101
102 pub fn blocking_stop_reason(&self) -> Option<RequestReadStopReason> {
104 if self.cancellations.iter().any(receiver_cancelled) {
105 Some(RequestReadStopReason::Cancelled)
106 } else if self
107 .deadline
108 .is_some_and(|deadline| WallInstant::now() >= deadline.blocking_at)
109 {
110 Some(RequestReadStopReason::Deadline)
111 } else {
112 None
113 }
114 }
115
116 pub fn deadline(&self) -> Option<RequestReadDeadline> {
118 self.deadline
119 }
120
121 pub fn stop_reason(&self) -> Option<RequestReadStopReason> {
123 if self.cancellations.iter().any(receiver_cancelled) {
124 Some(RequestReadStopReason::Cancelled)
125 } else if self
126 .deadline
127 .is_some_and(|deadline| tokio::time::Instant::now() >= deadline.async_at)
128 {
129 Some(RequestReadStopReason::Deadline)
130 } else {
131 None
132 }
133 }
134
135 pub async fn wait_for_stop(self) -> RequestReadStopReason {
137 let cancellation_receivers = self.cancellations;
138 let request_deadline = self.deadline;
139 let cancellation = async move {
140 match cancellation_receivers.len() {
141 0 => std::future::pending::<()>().await,
142 1 => wait_for_receiver_cancellation(cancellation_receivers[0].clone()).await,
143 2 => {
144 tokio::select! {
145 _ = wait_for_receiver_cancellation(cancellation_receivers[0].clone()) => {},
146 _ = wait_for_receiver_cancellation(cancellation_receivers[1].clone()) => {},
147 }
148 }
149 _ => wait_for_receiver_set(cancellation_receivers).await,
150 }
151 };
152 let deadline = async move {
153 match request_deadline {
154 Some(deadline) => tokio::time::sleep_until(deadline.async_at).await,
155 None => std::future::pending::<()>().await,
156 }
157 };
158 tokio::select! {
159 _ = cancellation => RequestReadStopReason::Cancelled,
160 _ = deadline => RequestReadStopReason::Deadline,
161 }
162 }
163}
164
165tokio::task_local! {
166 static REQUEST_READ_CONTEXT: RequestReadContext;
167}
168
169pub fn capture_request_read_context() -> RequestReadContext {
171 REQUEST_READ_CONTEXT
172 .try_with(Clone::clone)
173 .unwrap_or_default()
174}
175
176pub fn effective_request_read_deadline(candidate: RequestReadDeadline) -> RequestReadDeadline {
179 capture_request_read_context()
180 .deadline
181 .map_or(candidate, |existing| existing.earlier(candidate))
182}
183
184pub async fn scope_request_read_cancellation<F>(
188 cancellation: tokio::sync::watch::Receiver<bool>,
189 future: F,
190) -> F::Output
191where
192 F: Future,
193{
194 let mut context = capture_request_read_context();
195 let mut cancellations = Vec::with_capacity(context.cancellations.len() + 1);
196 cancellations.extend(context.cancellations.iter().cloned());
197 cancellations.push(cancellation);
198 context.cancellations = cancellations.into();
199 REQUEST_READ_CONTEXT.scope(context, future).await
200}
201
202pub async fn scope_request_read_deadline<F>(duration: Duration, future: F) -> F::Output
205where
206 F: Future,
207{
208 scope_request_read_deadline_at(RequestReadDeadline::after(duration), future).await
209}
210
211pub async fn scope_request_read_deadline_at<F>(
213 deadline: RequestReadDeadline,
214 future: F,
215) -> F::Output
216where
217 F: Future,
218{
219 let mut context = capture_request_read_context();
220 context.deadline = Some(match context.deadline {
221 Some(existing) => existing.earlier(deadline),
222 None => deadline,
223 });
224 REQUEST_READ_CONTEXT.scope(context, future).await
225}
226
227pub fn inherit_request_read_context<F>(future: F) -> impl Future<Output = F::Output> + Send
229where
230 F: Future + Send,
231{
232 let inherited = REQUEST_READ_CONTEXT.try_with(Clone::clone).ok();
233 async move {
234 match inherited {
235 Some(context) => REQUEST_READ_CONTEXT.scope(context, future).await,
236 None => future.await,
237 }
238 }
239}
240
241pub fn inherit_request_read_cancellation<F>(future: F) -> impl Future<Output = F::Output> + Send
243where
244 F: Future + Send,
245{
246 inherit_request_read_context(future)
247}
248
249pub fn request_read_is_cancelled() -> bool {
251 capture_request_read_context().stop_reason().is_some()
252}
253
254pub fn ensure_request_read_active(operation: &'static str) -> StorageResult<()> {
256 if request_read_is_cancelled() {
257 Err(StorageError::Timeout {
258 operation: operation.into(),
259 })
260 } else {
261 Ok(())
262 }
263}
264
265pub async fn wait_for_request_read_cancellation() {
267 let _ = capture_request_read_context().wait_for_stop().await;
268}
269
270pub async fn await_request_read_phase<F>(
272 operation: &'static str,
273 future: F,
274) -> StorageResult<F::Output>
275where
276 F: Future,
277{
278 ensure_request_read_active(operation)?;
279 tokio::pin!(future);
280 tokio::select! {
281 biased;
282 _ = wait_for_request_read_cancellation() => Err(StorageError::Timeout {
283 operation: operation.into(),
284 }),
285 output = &mut future => {
286 ensure_request_read_active(operation)?;
287 Ok(output)
288 }
289 }
290}
291
292fn receiver_cancelled(receiver: &tokio::sync::watch::Receiver<bool>) -> bool {
293 *receiver.borrow() || receiver.has_changed().is_err()
294}
295
296async fn wait_for_receiver_cancellation(mut receiver: tokio::sync::watch::Receiver<bool>) {
297 loop {
298 if *receiver.borrow_and_update() {
299 return;
300 }
301 if receiver.changed().await.is_err() {
302 return;
303 }
304 }
305}
306
307async fn wait_for_receiver_set(receivers: Arc<[tokio::sync::watch::Receiver<bool>]>) {
308 let mut waits: Vec<Pin<Box<dyn Future<Output = ()> + Send>>> = receivers
309 .iter()
310 .cloned()
311 .map(|receiver| {
312 Box::pin(wait_for_receiver_cancellation(receiver))
313 as Pin<Box<dyn Future<Output = ()> + Send>>
314 })
315 .collect();
316 std::future::poll_fn(move |cx| {
317 if waits
318 .iter_mut()
319 .any(|wait| wait.as_mut().poll(cx).is_ready())
320 {
321 Poll::Ready(())
322 } else {
323 Poll::Pending
324 }
325 })
326 .await;
327}
328
329#[cfg(test)]
330mod tests {
331 use super::*;
332
333 #[tokio::test]
334 async fn nested_scopes_merge_cancellation_sources() {
335 let (outer_tx, outer_rx) = tokio::sync::watch::channel(false);
336 let (_inner_tx, inner_rx) = tokio::sync::watch::channel(false);
337
338 let stopped = scope_request_read_cancellation(
339 outer_rx,
340 scope_request_read_cancellation(inner_rx, async move {
341 outer_tx.send(true).expect("outer scope remains live");
342 wait_for_request_read_cancellation().await;
343 request_read_is_cancelled()
344 }),
345 )
346 .await;
347
348 assert!(stopped);
349 }
350
351 #[tokio::test]
352 async fn sender_loss_is_request_abandonment() {
353 let (tx, rx) = tokio::sync::watch::channel(false);
354 drop(tx);
355
356 let stopped = scope_request_read_cancellation(rx, async {
357 capture_request_read_context().stop_reason()
358 })
359 .await;
360
361 assert_eq!(stopped, Some(RequestReadStopReason::Cancelled));
362 }
363
364 #[tokio::test]
365 async fn spawned_child_inherits_the_same_cancellation() {
366 let (tx, rx) = tokio::sync::watch::channel(false);
367
368 let stopped = scope_request_read_cancellation(rx, async move {
369 let child = tokio::spawn(inherit_request_read_context(async {
370 wait_for_request_read_cancellation().await;
371 request_read_is_cancelled()
372 }));
373 tx.send(true).expect("child receiver remains live");
374 child.await.expect("child task")
375 })
376 .await;
377
378 assert!(stopped);
379 }
380
381 #[tokio::test]
382 async fn nested_deadline_keeps_the_earlier_absolute_instant() {
383 let earlier = RequestReadDeadline::after(Duration::from_secs(10));
384 let later = RequestReadDeadline::after(Duration::from_secs(20));
385
386 let effective = scope_request_read_deadline_at(earlier, async {
387 effective_request_read_deadline(later)
388 })
389 .await;
390
391 assert_eq!(effective.async_at(), earlier.async_at());
392 }
393
394 #[tokio::test]
395 async fn async_phase_cannot_degrade_cancelled_work_to_success() {
396 let (tx, rx) = tokio::sync::watch::channel(false);
397 tx.send(true).expect("receiver remains live");
398
399 let result = scope_request_read_cancellation(rx, async {
400 await_request_read_phase("test.phase", async { 7_u8 }).await
401 })
402 .await;
403
404 assert!(matches!(result, Err(StorageError::Timeout { .. })));
405 }
406}