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