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}
82
83impl RequestReadContext {
84 pub fn deadline(&self) -> Option<RequestReadDeadline> {
86 self.deadline
87 }
88
89 pub fn stop_reason(&self) -> Option<RequestReadStopReason> {
91 if self.cancellations.iter().any(receiver_cancelled) {
92 Some(RequestReadStopReason::Cancelled)
93 } else if self
94 .deadline
95 .is_some_and(|deadline| tokio::time::Instant::now() >= deadline.async_at)
96 {
97 Some(RequestReadStopReason::Deadline)
98 } else {
99 None
100 }
101 }
102
103 pub async fn wait_for_stop(self) -> RequestReadStopReason {
105 let cancellation_receivers = self.cancellations;
106 let request_deadline = self.deadline;
107 let cancellation = async move {
108 match cancellation_receivers.len() {
109 0 => std::future::pending::<()>().await,
110 1 => wait_for_receiver_cancellation(cancellation_receivers[0].clone()).await,
111 2 => {
112 tokio::select! {
113 _ = wait_for_receiver_cancellation(cancellation_receivers[0].clone()) => {},
114 _ = wait_for_receiver_cancellation(cancellation_receivers[1].clone()) => {},
115 }
116 }
117 _ => wait_for_receiver_set(cancellation_receivers).await,
118 }
119 };
120 let deadline = async move {
121 match request_deadline {
122 Some(deadline) => tokio::time::sleep_until(deadline.async_at).await,
123 None => std::future::pending::<()>().await,
124 }
125 };
126 tokio::select! {
127 _ = cancellation => RequestReadStopReason::Cancelled,
128 _ = deadline => RequestReadStopReason::Deadline,
129 }
130 }
131}
132
133tokio::task_local! {
134 static REQUEST_READ_CONTEXT: RequestReadContext;
135}
136
137pub fn capture_request_read_context() -> RequestReadContext {
139 REQUEST_READ_CONTEXT
140 .try_with(Clone::clone)
141 .unwrap_or_default()
142}
143
144pub fn effective_request_read_deadline(candidate: RequestReadDeadline) -> RequestReadDeadline {
147 capture_request_read_context()
148 .deadline
149 .map_or(candidate, |existing| existing.earlier(candidate))
150}
151
152pub async fn scope_request_read_cancellation<F>(
156 cancellation: tokio::sync::watch::Receiver<bool>,
157 future: F,
158) -> F::Output
159where
160 F: Future,
161{
162 let mut context = capture_request_read_context();
163 let mut cancellations = Vec::with_capacity(context.cancellations.len() + 1);
164 cancellations.extend(context.cancellations.iter().cloned());
165 cancellations.push(cancellation);
166 context.cancellations = cancellations.into();
167 REQUEST_READ_CONTEXT.scope(context, future).await
168}
169
170pub async fn scope_request_read_deadline<F>(duration: Duration, future: F) -> F::Output
173where
174 F: Future,
175{
176 scope_request_read_deadline_at(RequestReadDeadline::after(duration), future).await
177}
178
179pub async fn scope_request_read_deadline_at<F>(
181 deadline: RequestReadDeadline,
182 future: F,
183) -> F::Output
184where
185 F: Future,
186{
187 let mut context = capture_request_read_context();
188 context.deadline = Some(match context.deadline {
189 Some(existing) => existing.earlier(deadline),
190 None => deadline,
191 });
192 REQUEST_READ_CONTEXT.scope(context, future).await
193}
194
195pub fn inherit_request_read_context<F>(future: F) -> impl Future<Output = F::Output> + Send
197where
198 F: Future + Send,
199{
200 let inherited = REQUEST_READ_CONTEXT.try_with(Clone::clone).ok();
201 async move {
202 match inherited {
203 Some(context) => REQUEST_READ_CONTEXT.scope(context, future).await,
204 None => future.await,
205 }
206 }
207}
208
209pub fn inherit_request_read_cancellation<F>(future: F) -> impl Future<Output = F::Output> + Send
211where
212 F: Future + Send,
213{
214 inherit_request_read_context(future)
215}
216
217pub fn request_read_is_cancelled() -> bool {
219 capture_request_read_context().stop_reason().is_some()
220}
221
222pub fn ensure_request_read_active(operation: &'static str) -> StorageResult<()> {
224 if request_read_is_cancelled() {
225 Err(StorageError::Timeout {
226 operation: operation.into(),
227 })
228 } else {
229 Ok(())
230 }
231}
232
233pub async fn wait_for_request_read_cancellation() {
235 let _ = capture_request_read_context().wait_for_stop().await;
236}
237
238pub async fn await_request_read_phase<F>(
240 operation: &'static str,
241 future: F,
242) -> StorageResult<F::Output>
243where
244 F: Future,
245{
246 ensure_request_read_active(operation)?;
247 tokio::pin!(future);
248 tokio::select! {
249 biased;
250 _ = wait_for_request_read_cancellation() => Err(StorageError::Timeout {
251 operation: operation.into(),
252 }),
253 output = &mut future => {
254 ensure_request_read_active(operation)?;
255 Ok(output)
256 }
257 }
258}
259
260fn receiver_cancelled(receiver: &tokio::sync::watch::Receiver<bool>) -> bool {
261 *receiver.borrow() || receiver.has_changed().is_err()
262}
263
264async fn wait_for_receiver_cancellation(mut receiver: tokio::sync::watch::Receiver<bool>) {
265 loop {
266 if *receiver.borrow_and_update() {
267 return;
268 }
269 if receiver.changed().await.is_err() {
270 return;
271 }
272 }
273}
274
275async fn wait_for_receiver_set(receivers: Arc<[tokio::sync::watch::Receiver<bool>]>) {
276 let mut waits: Vec<Pin<Box<dyn Future<Output = ()> + Send>>> = receivers
277 .iter()
278 .cloned()
279 .map(|receiver| {
280 Box::pin(wait_for_receiver_cancellation(receiver))
281 as Pin<Box<dyn Future<Output = ()> + Send>>
282 })
283 .collect();
284 std::future::poll_fn(move |cx| {
285 if waits
286 .iter_mut()
287 .any(|wait| wait.as_mut().poll(cx).is_ready())
288 {
289 Poll::Ready(())
290 } else {
291 Poll::Pending
292 }
293 })
294 .await;
295}
296
297#[cfg(test)]
298mod tests {
299 use super::*;
300
301 #[tokio::test]
302 async fn nested_scopes_merge_cancellation_sources() {
303 let (outer_tx, outer_rx) = tokio::sync::watch::channel(false);
304 let (_inner_tx, inner_rx) = tokio::sync::watch::channel(false);
305
306 let stopped = scope_request_read_cancellation(
307 outer_rx,
308 scope_request_read_cancellation(inner_rx, async move {
309 outer_tx.send(true).expect("outer scope remains live");
310 wait_for_request_read_cancellation().await;
311 request_read_is_cancelled()
312 }),
313 )
314 .await;
315
316 assert!(stopped);
317 }
318
319 #[tokio::test]
320 async fn sender_loss_is_request_abandonment() {
321 let (tx, rx) = tokio::sync::watch::channel(false);
322 drop(tx);
323
324 let stopped = scope_request_read_cancellation(rx, async {
325 capture_request_read_context().stop_reason()
326 })
327 .await;
328
329 assert_eq!(stopped, Some(RequestReadStopReason::Cancelled));
330 }
331
332 #[tokio::test]
333 async fn spawned_child_inherits_the_same_cancellation() {
334 let (tx, rx) = tokio::sync::watch::channel(false);
335
336 let stopped = scope_request_read_cancellation(rx, async move {
337 let child = tokio::spawn(inherit_request_read_context(async {
338 wait_for_request_read_cancellation().await;
339 request_read_is_cancelled()
340 }));
341 tx.send(true).expect("child receiver remains live");
342 child.await.expect("child task")
343 })
344 .await;
345
346 assert!(stopped);
347 }
348
349 #[tokio::test]
350 async fn nested_deadline_keeps_the_earlier_absolute_instant() {
351 let earlier = RequestReadDeadline::after(Duration::from_secs(10));
352 let later = RequestReadDeadline::after(Duration::from_secs(20));
353
354 let effective = scope_request_read_deadline_at(earlier, async {
355 effective_request_read_deadline(later)
356 })
357 .await;
358
359 assert_eq!(effective.async_at(), earlier.async_at());
360 }
361
362 #[tokio::test]
363 async fn async_phase_cannot_degrade_cancelled_work_to_success() {
364 let (tx, rx) = tokio::sync::watch::channel(false);
365 tx.send(true).expect("receiver remains live");
366
367 let result = scope_request_read_cancellation(rx, async {
368 await_request_read_phase("test.phase", async { 7_u8 }).await
369 })
370 .await;
371
372 assert!(matches!(result, Err(StorageError::Timeout { .. })));
373 }
374}