Skip to main content

hashtree_network/
mesh_forwarding_route.rs

1//! Hashtree mesh-forwarding ownership for the blob hop budget.
2
3use std::collections::HashMap;
4use std::sync::{Arc, Mutex};
5use std::time::Instant;
6
7use async_trait::async_trait;
8use hashtree_core::{BlobReply, BlobRequest, BlobRoute, BlobRouteContext, Hash, StoreError};
9use tokio::sync::watch;
10
11const MAX_TRACKED_MESH_FORWARDS: usize = 1_024;
12
13#[derive(Clone)]
14enum ForwardResult {
15    Complete(Result<BlobReply, String>),
16    Retry,
17}
18
19struct ActiveForward {
20    htl: u8,
21    attempt_budget: Option<usize>,
22    result: watch::Sender<Option<ForwardResult>>,
23}
24
25enum ForwardOwnership {
26    Forward(Option<ForwardOwnerGuard>),
27    Wait(Arc<ActiveForward>),
28    SuppressCycle,
29}
30
31struct ForwardOwnerGuard {
32    active: Arc<Mutex<HashMap<Hash, Arc<ActiveForward>>>>,
33    hash: Hash,
34    forward: Arc<ActiveForward>,
35    completed: bool,
36    deadline: Option<Instant>,
37}
38
39impl ForwardOwnerGuard {
40    fn complete(mut self, result: &Result<BlobReply, StoreError>) {
41        let shared = match result {
42            Ok(reply) => ForwardResult::Complete(Ok(reply.clone())),
43            Err(_)
44                if self
45                    .deadline
46                    .is_some_and(|deadline| deadline <= Instant::now()) =>
47            {
48                ForwardResult::Retry
49            }
50            Err(error) => ForwardResult::Complete(Err(error.to_string())),
51        };
52        self.remove();
53        self.forward.result.send_replace(Some(shared));
54        self.completed = true;
55    }
56
57    fn remove(&self) {
58        let mut active = self
59            .active
60            .lock()
61            .unwrap_or_else(|poisoned| poisoned.into_inner());
62        if active
63            .get(&self.hash)
64            .is_some_and(|current| Arc::ptr_eq(current, &self.forward))
65        {
66            active.remove(&self.hash);
67        }
68    }
69}
70
71impl Drop for ForwardOwnerGuard {
72    fn drop(&mut self) {
73        if self.completed {
74            return;
75        }
76        self.remove();
77        self.forward.result.send_replace(Some(ForwardResult::Retry));
78    }
79}
80
81/// Marks exactly one Hashtree mesh-forwarding decision.
82///
83/// Transport routes remain opaque and preserve the request they carry. Local
84/// and terminal routes must not use this adapter. An exhausted request is a
85/// route-local miss, so the outer search may still try other authorities.
86pub struct MeshForwardingRoute {
87    inner: Arc<dyn BlobRoute>,
88    active: Arc<Mutex<HashMap<Hash, Arc<ActiveForward>>>>,
89}
90
91impl MeshForwardingRoute {
92    pub fn new(inner: Arc<dyn BlobRoute>) -> Self {
93        Self {
94            inner,
95            active: Arc::new(Mutex::new(HashMap::new())),
96        }
97    }
98
99    fn forwarded_request(request: BlobRequest) -> Option<BlobRequest> {
100        Some(BlobRequest {
101            hash: request.hash,
102            htl: request.htl.checked_sub(1)?,
103        })
104    }
105
106    fn claim_forward(
107        &self,
108        request: BlobRequest,
109        context: Option<BlobRouteContext>,
110    ) -> ForwardOwnership {
111        let attempt_budget = context.map(|context| context.attempt_budget);
112        let mut active = self
113            .active
114            .lock()
115            .unwrap_or_else(|poisoned| poisoned.into_inner());
116        if let Some(current) = active.get(&request.hash) {
117            if request.htl < current.htl {
118                return ForwardOwnership::SuppressCycle;
119            }
120            if request.htl == current.htl && attempt_budget == current.attempt_budget {
121                return ForwardOwnership::Wait(current.clone());
122            }
123        }
124        if active.len() >= MAX_TRACKED_MESH_FORWARDS && !active.contains_key(&request.hash) {
125            return ForwardOwnership::Forward(None);
126        }
127        let (result, _) = watch::channel(None);
128        let forward = Arc::new(ActiveForward {
129            htl: request.htl,
130            attempt_budget,
131            result,
132        });
133        active.insert(request.hash, forward.clone());
134        ForwardOwnership::Forward(Some(ForwardOwnerGuard {
135            active: self.active.clone(),
136            hash: request.hash,
137            forward,
138            completed: false,
139            deadline: context.map(|context| context.deadline),
140        }))
141    }
142
143    async fn wait_for_forward(
144        forward: Arc<ActiveForward>,
145        context: Option<BlobRouteContext>,
146    ) -> Result<Option<BlobReply>, StoreError> {
147        let mut result = forward.result.subscribe();
148        let wait = async {
149            loop {
150                if let Some(result) = result.borrow().clone() {
151                    return match result {
152                        ForwardResult::Retry => Ok(None),
153                        ForwardResult::Complete(result) => {
154                            result.map(Some).map_err(StoreError::Other)
155                        }
156                    };
157                }
158                result.changed().await.map_err(|_| {
159                    StoreError::Other("mesh forwarding owner closed without a result".to_string())
160                })?;
161            }
162        };
163        if let Some(context) = context {
164            tokio::time::timeout_at(tokio::time::Instant::from_std(context.deadline), wait)
165                .await
166                .map_err(|_| {
167                    StoreError::Other(
168                        "coalesced mesh forwarding deadline expired before completion".to_string(),
169                    )
170                })?
171        } else {
172            wait.await
173        }
174    }
175
176    async fn route_inner(
177        &self,
178        request: BlobRequest,
179        context: Option<BlobRouteContext>,
180    ) -> Result<BlobReply, StoreError> {
181        let Some(forwarded) = Self::forwarded_request(request) else {
182            return Ok(BlobReply::NoResult);
183        };
184        loop {
185            if context.is_some_and(|context| context.deadline <= Instant::now()) {
186                return Err(StoreError::Other(
187                    "mesh forwarding deadline expired".to_string(),
188                ));
189            }
190            match self.claim_forward(request, context) {
191                ForwardOwnership::SuppressCycle => return Ok(BlobReply::NoResult),
192                ForwardOwnership::Wait(forward) => {
193                    if let Some(reply) = Self::wait_for_forward(forward, context).await? {
194                        return Ok(reply);
195                    }
196                    // Cancellation or an earlier owner's deadline belongs to that
197                    // reader. Let one remaining reader take over the shared work.
198                }
199                ForwardOwnership::Forward(owner) => {
200                    let result = match context {
201                        Some(context) => self.inner.route_with_context(forwarded, context).await,
202                        None => self.inner.route(forwarded).await,
203                    };
204                    if let Some(owner) = owner {
205                        owner.complete(&result);
206                    }
207                    return result;
208                }
209            }
210        }
211    }
212}
213
214#[async_trait]
215impl BlobRoute for MeshForwardingRoute {
216    async fn route(&self, request: BlobRequest) -> Result<BlobReply, StoreError> {
217        self.route_inner(request, None).await
218    }
219
220    async fn route_with_context(
221        &self,
222        request: BlobRequest,
223        context: BlobRouteContext,
224    ) -> Result<BlobReply, StoreError> {
225        self.route_inner(request, Some(context)).await
226    }
227}
228
229#[cfg(test)]
230mod tests {
231    use std::sync::atomic::{AtomicUsize, Ordering};
232    use std::sync::OnceLock;
233    use std::time::{Duration, Instant};
234
235    use super::*;
236    use hashtree_core::sha256;
237
238    #[derive(Default)]
239    struct RecordingRoute {
240        requests: Mutex<Vec<BlobRequest>>,
241        contexts: Mutex<Vec<BlobRouteContext>>,
242    }
243
244    #[derive(Default)]
245    struct CycleRoute {
246        next: OnceLock<Arc<dyn BlobRoute>>,
247        requests: Mutex<Vec<BlobRequest>>,
248    }
249
250    struct SlowRoute {
251        calls: AtomicUsize,
252    }
253
254    #[async_trait]
255    impl BlobRoute for SlowRoute {
256        async fn route(&self, _request: BlobRequest) -> Result<BlobReply, StoreError> {
257            self.calls.fetch_add(1, Ordering::SeqCst);
258            tokio::time::sleep(Duration::from_millis(20)).await;
259            Ok(BlobReply::NoResult)
260        }
261    }
262
263    #[async_trait]
264    impl BlobRoute for CycleRoute {
265        async fn route(&self, request: BlobRequest) -> Result<BlobReply, StoreError> {
266            self.requests.lock().unwrap().push(request);
267            self.next.get().unwrap().route(request).await
268        }
269    }
270
271    #[async_trait]
272    impl BlobRoute for RecordingRoute {
273        async fn route(&self, request: BlobRequest) -> Result<BlobReply, StoreError> {
274            self.requests.lock().unwrap().push(request);
275            Ok(BlobReply::NoResult)
276        }
277
278        async fn route_with_context(
279            &self,
280            request: BlobRequest,
281            context: BlobRouteContext,
282        ) -> Result<BlobReply, StoreError> {
283            self.requests.lock().unwrap().push(request);
284            self.contexts.lock().unwrap().push(context);
285            Ok(BlobReply::NoResult)
286        }
287    }
288
289    #[tokio::test]
290    async fn one_mesh_decision_consumes_exactly_one_hop() {
291        let inner = Arc::new(RecordingRoute::default());
292        let route = MeshForwardingRoute::new(inner.clone());
293        let hash = sha256(b"one hop");
294
295        assert_eq!(
296            route.route(BlobRequest { hash, htl: 2 }).await.unwrap(),
297            BlobReply::NoResult
298        );
299        assert_eq!(
300            inner.requests.lock().unwrap().as_slice(),
301            &[BlobRequest { hash, htl: 1 }]
302        );
303    }
304
305    #[tokio::test]
306    async fn exhausted_request_is_route_local_no_result() {
307        let inner = Arc::new(RecordingRoute::default());
308        let route = MeshForwardingRoute::new(inner.clone());
309        let hash = sha256(b"exhausted");
310
311        assert_eq!(
312            route.route(BlobRequest { hash, htl: 0 }).await.unwrap(),
313            BlobReply::NoResult
314        );
315        assert!(inner.requests.lock().unwrap().is_empty());
316    }
317
318    #[tokio::test]
319    async fn nested_mesh_decisions_observe_two_one_zero() {
320        let terminal = Arc::new(RecordingRoute::default());
321        let second: Arc<dyn BlobRoute> = Arc::new(MeshForwardingRoute::new(terminal.clone()));
322        let first = MeshForwardingRoute::new(second);
323        let hash = sha256(b"two hops");
324
325        assert_eq!(
326            first.route(BlobRequest { hash, htl: 2 }).await.unwrap(),
327            BlobReply::NoResult
328        );
329        assert_eq!(
330            terminal.requests.lock().unwrap().as_slice(),
331            &[BlobRequest { hash, htl: 0 }]
332        );
333    }
334
335    #[tokio::test]
336    async fn forwarding_preserves_deadline_and_attempt_budget() {
337        let inner = Arc::new(RecordingRoute::default());
338        let route = MeshForwardingRoute::new(inner.clone());
339        let hash = sha256(b"context");
340        let context = BlobRouteContext {
341            deadline: Instant::now() + Duration::from_secs(3),
342            attempt_budget: 2,
343        };
344
345        route
346            .route_with_context(BlobRequest { hash, htl: 1 }, context)
347            .await
348            .unwrap();
349        assert_eq!(
350            inner.requests.lock().unwrap().as_slice(),
351            &[BlobRequest { hash, htl: 0 }]
352        );
353        let observed = inner.contexts.lock().unwrap();
354        assert_eq!(observed.len(), 1);
355        assert_eq!(observed[0].deadline, context.deadline);
356        assert_eq!(observed[0].attempt_budget, context.attempt_budget);
357    }
358
359    #[tokio::test]
360    async fn lower_htl_cycle_reentry_is_suppressed_before_repeating_provider_work() {
361        let first = Arc::new(CycleRoute::default());
362        let second = Arc::new(CycleRoute::default());
363        let first_forwarder: Arc<dyn BlobRoute> = Arc::new(MeshForwardingRoute::new(first.clone()));
364        let second_forwarder: Arc<dyn BlobRoute> =
365            Arc::new(MeshForwardingRoute::new(second.clone()));
366        assert!(first.next.set(second_forwarder.clone()).is_ok());
367        assert!(second.next.set(first_forwarder.clone()).is_ok());
368        let hash = sha256(b"cycle");
369
370        assert_eq!(
371            first_forwarder
372                .route(BlobRequest { hash, htl: 3 })
373                .await
374                .unwrap(),
375            BlobReply::NoResult
376        );
377        assert_eq!(
378            first.requests.lock().unwrap().as_slice(),
379            &[BlobRequest { hash, htl: 2 }]
380        );
381        assert_eq!(
382            second.requests.lock().unwrap().as_slice(),
383            &[BlobRequest { hash, htl: 1 }]
384        );
385    }
386
387    #[tokio::test]
388    async fn equal_in_flight_requests_share_one_mesh_attempt() {
389        let inner = Arc::new(SlowRoute {
390            calls: AtomicUsize::new(0),
391        });
392        let route = Arc::new(MeshForwardingRoute::new(inner.clone()));
393        let request = BlobRequest {
394            hash: sha256(b"duplicate"),
395            htl: 2,
396        };
397
398        let (first, second) = tokio::join!(route.route(request), route.route(request));
399        assert_eq!(first.unwrap(), BlobReply::NoResult);
400        assert_eq!(second.unwrap(), BlobReply::NoResult);
401        assert_eq!(inner.calls.load(Ordering::SeqCst), 1);
402    }
403}