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