1use 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
71pub 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}