1use std::collections::HashMap;
5use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
6use std::sync::{Arc, Mutex, Weak};
7use std::time::Duration;
8
9use async_trait::async_trait;
10use camel_api::exchange::Exchange;
11use camel_api::resequencer::BatchCompletion;
12use camel_api::value::cmp_values;
13use camel_language_api::Expression;
14use tokio::sync::mpsc;
15use tokio_util::sync::CancellationToken;
16
17use super::ResequencePolicy;
18
19pub const DEFAULT_MAX_BUCKETS: usize = 10_000;
23
24pub const DEFAULT_MAX_BUCKET_SIZE: usize = 1_000;
29
30pub const DEFAULT_MAX_TIMEOUT_TASKS: usize = 1024;
35
36#[derive(Default)]
38struct Bucket {
39 exchanges: Vec<Exchange>,
40}
41
42struct TimeoutEntry {
46 generation: u64,
47 cancel: CancellationToken,
48}
49
50pub struct BatchPolicy {
57 correlation_expr: Arc<dyn Expression>,
58 sort_expr: Arc<dyn Expression>,
59 completion: BatchCompletion,
60
61 weak_self: Weak<Self>,
63
64 buckets: Mutex<HashMap<String, Bucket>>,
66
67 timeout_tasks: Mutex<HashMap<String, TimeoutEntry>>,
72
73 timeout_generation: AtomicU64,
77
78 #[cfg(test)]
85 interleave_hook: Mutex<Option<Arc<dyn Fn() + Send + Sync>>>,
86
87 driver_tx: Mutex<Option<mpsc::Sender<Exchange>>>,
90
91 shutdown_started: AtomicBool,
94
95 max_buckets: usize,
97
98 max_bucket_size: usize,
100
101 max_timeout_tasks: usize,
103}
104
105impl BatchPolicy {
106 pub fn new_cyclic(
109 correlation_expr: Arc<dyn Expression>,
110 sort_expr: Arc<dyn Expression>,
111 completion: BatchCompletion,
112 ) -> Arc<Self> {
113 Self::with_limits(
114 correlation_expr,
115 sort_expr,
116 completion,
117 DEFAULT_MAX_BUCKETS,
118 DEFAULT_MAX_BUCKET_SIZE,
119 DEFAULT_MAX_TIMEOUT_TASKS,
120 )
121 }
122
123 pub fn with_limits(
126 correlation_expr: Arc<dyn Expression>,
127 sort_expr: Arc<dyn Expression>,
128 completion: BatchCompletion,
129 max_buckets: usize,
130 max_bucket_size: usize,
131 max_timeout_tasks: usize,
132 ) -> Arc<Self> {
133 Arc::new_cyclic(|weak| Self {
134 correlation_expr,
135 sort_expr,
136 completion,
137 weak_self: weak.clone(),
138 buckets: Mutex::new(HashMap::new()),
139 timeout_tasks: Mutex::new(HashMap::new()),
140 timeout_generation: AtomicU64::new(0),
141 #[cfg(test)]
142 interleave_hook: Mutex::new(None),
143 driver_tx: Mutex::new(None),
144 shutdown_started: AtomicBool::new(false),
145 max_buckets,
146 max_bucket_size,
147 max_timeout_tasks,
148 })
149 }
150
151 fn set_driver_tx(&self, tx: mpsc::Sender<Exchange>) {
154 let mut guard = self.driver_tx.lock().unwrap_or_else(|e| e.into_inner());
155 *guard = Some(tx);
156 }
157
158 async fn eval_key(&self, exchange: &Exchange) -> Result<String, String> {
160 self.correlation_expr
161 .evaluate(exchange)
162 .await
163 .map(|v| match v {
166 serde_json::Value::String(s) => s,
167 other => other.to_string(),
168 })
169 .map_err(|e| format!("correlation expression evaluation failed: {e}"))
170 }
171
172 async fn drain_and_sort(&self, mut bucket: Bucket) -> Vec<Exchange> {
174 let mut indexed: Vec<(serde_json::Value, Exchange)> = Vec::new();
175 for ex in bucket.exchanges.drain(..) {
176 let val = self
177 .sort_expr
178 .evaluate(&ex)
179 .await
180 .unwrap_or(serde_json::Value::Null);
181 indexed.push((val, ex));
182 }
183 indexed.sort_by(|a, b| cmp_values(&a.0, &b.0));
184 indexed.into_iter().map(|(_, ex)| ex).collect()
185 }
186
187 fn is_complete_by_size(&self, count: usize) -> bool {
189 match self.completion {
190 BatchCompletion::Size(s) => count >= s,
191 BatchCompletion::SizeOrTimeout(s, _) => count >= s,
192 _ => false,
194 }
195 }
196
197 fn needs_timeout(&self) -> bool {
199 matches!(
200 self.completion,
201 BatchCompletion::Timeout(_) | BatchCompletion::SizeOrTimeout(..)
202 )
203 }
204
205 fn take_bucket(&self, key: &str) -> Option<Bucket> {
207 let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
208 buckets.remove(key)
209 }
210
211 fn cancel_timeout(&self, key: &str) {
215 if let Some(entry) = self
216 .timeout_tasks
217 .lock()
218 .unwrap_or_else(|e| e.into_inner())
219 .remove(key)
220 {
221 entry.cancel.cancel();
222 }
223 }
224
225 fn take_bucket_if_current_timeout_task(&self, key: &str, generation: u64) -> Option<Bucket> {
233 let tasks = self.timeout_tasks.lock().unwrap_or_else(|e| e.into_inner());
234 if !tasks
235 .get(key)
236 .is_some_and(|entry| entry.generation == generation)
237 {
238 return None;
239 }
240 #[cfg(test)]
246 if let Some(hook) = self
247 .interleave_hook
248 .lock()
249 .unwrap_or_else(|e| e.into_inner())
250 .clone()
251 {
252 hook();
253 }
254 let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
255 buckets.remove(key)
256 }
257
258 fn remove_timeout_task_if_current(&self, key: &str, generation: u64) {
260 let mut tasks = self.timeout_tasks.lock().unwrap_or_else(|e| e.into_inner());
261 if tasks
262 .get(key)
263 .is_some_and(|entry| entry.generation == generation)
264 {
265 tasks.remove(key);
266 }
267 }
268
269 fn spawn_timeout_task(&self, key: String, timeout_ms: u64) {
281 let generation = self.timeout_generation.fetch_add(1, Ordering::SeqCst) + 1;
282 let cancel = CancellationToken::new();
283 let cancel_clone = cancel.clone();
284
285 {
287 let mut tasks = self.timeout_tasks.lock().unwrap_or_else(|e| e.into_inner());
288 tasks.insert(key.clone(), TimeoutEntry { generation, cancel });
289 }
290
291 let weak = self.weak_self.clone();
292 let key_clone = key.clone();
293 let driver_tx_opt = {
294 let guard = self.driver_tx.lock().unwrap_or_else(|e| e.into_inner());
295 guard.clone()
296 };
297
298 tokio::spawn(async move {
299 let timeout = Duration::from_millis(timeout_ms);
300
301 tokio::select! {
302 _ = tokio::time::sleep(timeout) => {
303 if cancel_clone.is_cancelled() {
304 return;
305 }
306 }
307 _ = cancel_clone.cancelled() => {
308 return;
309 }
310 }
311
312 let Some(policy) = weak.upgrade() else {
314 return;
315 };
316
317 if policy.shutdown_started.load(Ordering::SeqCst) {
319 return;
320 }
321
322 let bucket = policy.take_bucket_if_current_timeout_task(&key_clone, generation);
327 let Some(bucket) = bucket else {
328 policy.remove_timeout_task_if_current(&key_clone, generation);
334 return;
335 };
336
337 let sorted = policy.drain_and_sort(bucket).await;
338
339 if let Some(tx) = driver_tx_opt {
341 for ex in sorted {
342 if tx.send(ex).await.is_err() {
343 tracing::debug!(
344 key = %key_clone,
345 "BatchPolicy timeout: driver channel closed during emission"
346 );
347 break;
348 }
349 }
350 }
351
352 policy.remove_timeout_task_if_current(&key_clone, generation);
354 });
355 }
356}
357
358#[async_trait]
359impl ResequencePolicy for BatchPolicy {
360 async fn accept(&self, input: Exchange) -> Vec<Exchange> {
361 let correlation_id = input.correlation_id().to_owned();
362 let key = match self.eval_key(&input).await {
363 Ok(k) => k,
364 Err(e) => {
365 tracing::warn!(
367 error = %e,
368 correlation_id = %correlation_id,
369 "BatchPolicy: correlation expression failed, dropping exchange"
370 );
371 return vec![];
372 }
373 };
374
375 let bucket_count = {
376 let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
377 if !buckets.contains_key(&key) && buckets.len() >= self.max_buckets {
380 tracing::warn!(
382 correlation_id = %correlation_id,
383 max_buckets = self.max_buckets,
384 "BatchPolicy: bucket cap reached, dropping exchange"
385 );
386 return vec![];
387 }
388 let bucket = buckets.entry(key.clone()).or_default();
389 if bucket.exchanges.len() >= self.max_bucket_size {
391 tracing::warn!(
393 correlation_id = %correlation_id,
394 max_bucket_size = self.max_bucket_size,
395 "BatchPolicy: per-bucket cap reached, dropping exchange"
396 );
397 return vec![];
398 }
399 bucket.exchanges.push(input);
400 bucket.exchanges.len()
401 };
402
403 if bucket_count == 1 && self.needs_timeout() {
408 let live_tasks = {
409 let tasks = self.timeout_tasks.lock().unwrap_or_else(|e| e.into_inner());
410 tasks.len()
411 };
412 if live_tasks < self.max_timeout_tasks {
413 let timeout_ms = match self.completion {
414 BatchCompletion::Timeout(t) | BatchCompletion::SizeOrTimeout(_, t) => t,
415 _ => unreachable!(),
416 };
417 self.spawn_timeout_task(key.clone(), timeout_ms);
418 } else {
419 tracing::warn!(
421 correlation_id = %correlation_id,
422 max_timeout_tasks = self.max_timeout_tasks,
423 "BatchPolicy: timeout-task cap reached; bucket relies on size/flush completion"
424 );
425 }
426 }
427
428 if self.is_complete_by_size(bucket_count) {
430 self.cancel_timeout(&key);
431 if let Some(bucket) = self.take_bucket(&key) {
432 return self.drain_and_sort(bucket).await;
433 }
434 }
435
436 vec![]
437 }
438
439 async fn flush(&self) -> Vec<Exchange> {
440 self.shutdown_started.store(true, Ordering::SeqCst);
442
443 let all_keys: Vec<String> = {
444 let buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
445 buckets.keys().cloned().collect()
446 };
447
448 let mut all_sorted = Vec::new();
449 for key in &all_keys {
450 self.cancel_timeout(key);
451 if let Some(bucket) = self.take_bucket(key) {
452 let sorted = self.drain_and_sort(bucket).await;
453 all_sorted.extend(sorted);
454 }
455 }
456
457 {
460 let tasks: HashMap<String, TimeoutEntry> = {
461 let mut guard = self.timeout_tasks.lock().unwrap_or_else(|e| e.into_inner());
462 std::mem::take(&mut *guard)
463 };
464 for (_, entry) in tasks {
465 entry.cancel.cancel();
466 }
467 }
468
469 all_sorted
470 }
471
472 fn name(&self) -> &'static str {
473 "batch-resequencer"
474 }
475
476 fn buffered(&self) -> usize {
477 let buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
478 buckets.values().map(|b| b.exchanges.len()).sum()
479 }
480
481 fn set_timeout_tx(&self, tx: tokio::sync::mpsc::Sender<Exchange>) {
482 self.set_driver_tx(tx);
483 }
484}
485
486#[cfg(test)]
489mod tests {
490 use super::*;
491 use camel_api::exchange::ExchangePattern;
492 use camel_api::message::Message;
493
494 struct PropExpr(String);
496
497 #[async_trait::async_trait]
498 impl Expression for PropExpr {
499 async fn evaluate(
500 &self,
501 exchange: &Exchange,
502 ) -> Result<serde_json::Value, camel_language_api::LanguageError> {
503 Ok(exchange
504 .property(&self.0)
505 .cloned()
506 .unwrap_or(serde_json::Value::Null))
507 }
508 }
509
510 struct ConstExpr(String);
512
513 #[async_trait::async_trait]
514 impl Expression for ConstExpr {
515 async fn evaluate(
516 &self,
517 _exchange: &Exchange,
518 ) -> Result<serde_json::Value, camel_language_api::LanguageError> {
519 Ok(serde_json::Value::String(self.0.clone()))
520 }
521 }
522
523 struct FailingExpr;
525
526 #[async_trait::async_trait]
527 impl Expression for FailingExpr {
528 async fn evaluate(
529 &self,
530 _exchange: &Exchange,
531 ) -> Result<serde_json::Value, camel_language_api::LanguageError> {
532 Err(camel_language_api::LanguageError::EvalError(
533 "mock eval failure".into(),
534 ))
535 }
536 }
537
538 fn mk_exchange(seq: i64) -> Exchange {
539 let mut ex = Exchange::new(Message::new(camel_api::body::Body::Text(format!(
540 "msg-{seq}"
541 ))));
542 ex.set_property("seq", serde_json::json!(seq));
543 ex.pattern = ExchangePattern::InOnly;
544 ex
545 }
546
547 fn mk_exchange_with_key(seq: i64, key_prop: &str, key_val: &str) -> Exchange {
548 let mut ex = Exchange::new(Message::new(camel_api::body::Body::Text(format!(
549 "msg-{seq}"
550 ))));
551 ex.set_property("seq", serde_json::json!(seq));
552 ex.set_property(key_prop, serde_json::Value::String(key_val.to_string()));
553 ex.pattern = ExchangePattern::InOnly;
554 ex
555 }
556
557 #[tokio::test]
560 async fn batch_size_completion_emits_sorted_burst() {
561 let policy = BatchPolicy::new_cyclic(
562 Arc::new(ConstExpr("same".into())),
563 Arc::new(PropExpr("seq".into())),
564 BatchCompletion::Size(3),
565 );
566
567 assert!(policy.accept(mk_exchange(3)).await.is_empty());
568 assert!(policy.accept(mk_exchange(1)).await.is_empty());
569
570 let emitted = policy.accept(mk_exchange(2)).await;
571 assert_eq!(emitted.len(), 3, "should emit all 3 on completion");
572 let seqs: Vec<i64> = emitted
573 .iter()
574 .map(|ex| ex.property("seq").and_then(|v| v.as_i64()).unwrap_or(-1))
575 .collect();
576 assert_eq!(seqs, vec![1, 2, 3], "should be sorted ascending");
577 }
578
579 #[tokio::test]
582 async fn batch_timeout_completion_emits_after_timeout() {
583 let policy = BatchPolicy::new_cyclic(
584 Arc::new(ConstExpr("same".into())),
585 Arc::new(PropExpr("seq".into())),
586 BatchCompletion::Timeout(50),
587 );
588
589 let (tx, mut rx) = mpsc::channel::<Exchange>(16);
590 policy.set_driver_tx(tx);
591
592 assert!(policy.accept(mk_exchange(3)).await.is_empty());
593 assert!(policy.accept(mk_exchange(1)).await.is_empty());
594
595 let emitted: Vec<Exchange> = tokio::time::timeout(Duration::from_millis(500), async {
596 let mut out = Vec::new();
597 out.push(rx.recv().await.unwrap());
598 out.push(rx.recv().await.unwrap());
599 out
600 })
601 .await
602 .expect("timeout should fire within 500ms");
603
604 assert_eq!(emitted.len(), 2);
605 let seqs: Vec<i64> = emitted
606 .iter()
607 .map(|ex| ex.property("seq").and_then(|v| v.as_i64()).unwrap_or(-1))
608 .collect();
609 assert_eq!(seqs, vec![1, 3], "should be sorted ascending");
610 }
611
612 #[tokio::test]
614 async fn batch_size_or_timeout_size_wins() {
615 let policy = BatchPolicy::new_cyclic(
616 Arc::new(ConstExpr("same".into())),
617 Arc::new(PropExpr("seq".into())),
618 BatchCompletion::SizeOrTimeout(3, 5_000),
619 );
620
621 assert!(policy.accept(mk_exchange(2)).await.is_empty());
622 assert!(policy.accept(mk_exchange(1)).await.is_empty());
623
624 let emitted = policy.accept(mk_exchange(3)).await;
625 assert_eq!(emitted.len(), 3);
626 let seqs: Vec<i64> = emitted
627 .iter()
628 .map(|ex| ex.property("seq").and_then(|v| v.as_i64()).unwrap_or(-1))
629 .collect();
630 assert_eq!(seqs, vec![1, 2, 3]);
631 }
632
633 #[tokio::test]
635 async fn batch_multi_key_independence() {
636 let policy = BatchPolicy::new_cyclic(
637 Arc::new(PropExpr("region".into())),
638 Arc::new(PropExpr("seq".into())),
639 BatchCompletion::Size(2),
640 );
641
642 let _ = policy
643 .accept(mk_exchange_with_key(2, "region", "east"))
644 .await;
645 let east_emit = policy
646 .accept(mk_exchange_with_key(1, "region", "east"))
647 .await;
648 assert_eq!(east_emit.len(), 2, "east bucket should complete at size 2");
649
650 let west_result = policy
651 .accept(mk_exchange_with_key(3, "region", "west"))
652 .await;
653 assert!(
654 west_result.is_empty(),
655 "west bucket should NOT complete yet"
656 );
657 }
658
659 #[tokio::test]
662 async fn batch_flush_emits_remaining_sorted() {
663 let policy = BatchPolicy::new_cyclic(
664 Arc::new(ConstExpr("same".into())),
665 Arc::new(PropExpr("seq".into())),
666 BatchCompletion::Size(10),
667 );
668
669 assert!(policy.accept(mk_exchange(5)).await.is_empty());
670 assert!(policy.accept(mk_exchange(3)).await.is_empty());
671 assert!(policy.accept(mk_exchange(1)).await.is_empty());
672
673 let flushed = policy.flush().await;
674 assert_eq!(flushed.len(), 3);
675 let seqs: Vec<i64> = flushed
676 .iter()
677 .map(|ex| ex.property("seq").and_then(|v| v.as_i64()).unwrap_or(-1))
678 .collect();
679 assert_eq!(seqs, vec![1, 3, 5]);
680 }
681
682 #[tokio::test]
685 async fn batch_correlation_eval_failure_returns_empty() {
686 let policy = BatchPolicy::new_cyclic(
687 Arc::new(FailingExpr),
688 Arc::new(PropExpr("seq".into())),
689 BatchCompletion::Size(2),
690 );
691
692 let result = policy.accept(mk_exchange(1)).await;
693 assert!(
694 result.is_empty(),
695 "failed correlation should return empty vec, not crash"
696 );
697 }
698
699 #[tokio::test]
701 async fn batch_pure_size_no_timeout_needed() {
702 let policy = BatchPolicy::new_cyclic(
703 Arc::new(ConstExpr("same".into())),
704 Arc::new(PropExpr("seq".into())),
705 BatchCompletion::Size(2),
706 );
707
708 assert!(!policy.needs_timeout());
709 }
710
711 #[tokio::test]
717 async fn batch_bucket_count_cap_drops_new_keys() {
718 let policy = BatchPolicy::with_limits(
719 Arc::new(PropExpr("key".into())),
720 Arc::new(PropExpr("seq".into())),
721 BatchCompletion::Size(100), 4, 1000, 16, );
726
727 for i in 0..4 {
728 let mut ex = mk_exchange(i);
729 ex.set_property("key", serde_json::json!(format!("k{i}")));
730 assert!(policy.accept(ex).await.is_empty(), "buffered, not emitted");
731 }
732 assert_eq!(policy.buckets.lock().unwrap().len(), 4);
733
734 let mut ex = mk_exchange(99);
736 ex.set_property("key", serde_json::json!("k-overflow"));
737 assert!(policy.accept(ex).await.is_empty());
738 assert_eq!(
739 policy.buckets.lock().unwrap().len(),
740 4,
741 "no new bucket past the cap"
742 );
743
744 let mut ex = mk_exchange(100);
746 ex.set_property("key", serde_json::json!("k0"));
747 assert!(policy.accept(ex).await.is_empty());
748 assert_eq!(
749 policy.buckets.lock().unwrap()["k0"].exchanges.len(),
750 2,
751 "existing bucket keeps accepting"
752 );
753 }
754
755 #[tokio::test]
757 async fn batch_per_bucket_cap_drops_overflow() {
758 let policy = BatchPolicy::with_limits(
759 Arc::new(ConstExpr("hot".into())),
760 Arc::new(PropExpr("seq".into())),
761 BatchCompletion::Size(1_000_000), 10, 3, 16, );
766
767 for i in 0..3 {
768 assert!(policy.accept(mk_exchange(i)).await.is_empty());
769 }
770 assert_eq!(policy.buffered(), 3);
771
772 assert!(policy.accept(mk_exchange(3)).await.is_empty());
774 assert!(policy.accept(mk_exchange(4)).await.is_empty());
775 assert_eq!(
776 policy.buffered(),
777 3,
778 "bucket must not grow past max_bucket_size"
779 );
780 }
781
782 #[tokio::test]
785 async fn batch_timeout_task_cap_stops_spawning() {
786 let policy = BatchPolicy::with_limits(
787 Arc::new(PropExpr("key".into())),
788 Arc::new(PropExpr("seq".into())),
789 BatchCompletion::Timeout(60_000), 16, 100, 2, );
794
795 for i in 0..4 {
796 let mut ex = mk_exchange(i);
797 ex.set_property("key", serde_json::json!(format!("k{i}")));
798 assert!(policy.accept(ex).await.is_empty());
799 }
800
801 assert_eq!(
802 policy.timeout_tasks.lock().unwrap().len(),
803 2,
804 "no more than max_timeout_tasks tasks spawned"
805 );
806 assert_eq!(
807 policy.buckets.lock().unwrap().len(),
808 4,
809 "buckets still buffered even without their own timer"
810 );
811
812 let flushed = policy.flush().await;
814 assert_eq!(flushed.len(), 4);
815 }
816
817 #[tokio::test]
821 async fn batch_sequential_unique_keys_do_not_leak_timeout_entries() {
822 let policy = BatchPolicy::with_limits(
823 Arc::new(PropExpr("key".into())),
824 Arc::new(PropExpr("seq".into())),
825 BatchCompletion::Timeout(30),
826 100, 100, 16, );
830
831 for i in 0..8 {
832 let mut ex = mk_exchange(i);
833 ex.set_property("key", serde_json::json!(format!("k{i}")));
834 assert!(policy.accept(ex).await.is_empty());
835 tokio::time::sleep(Duration::from_millis(80)).await;
837 assert_eq!(
838 policy.timeout_tasks.lock().unwrap().len(),
839 0,
840 "timeout entry must be removed on natural completion (k{i})"
841 );
842 }
843 assert_eq!(policy.buckets.lock().unwrap().len(), 0);
844 }
845
846 #[tokio::test]
852 async fn batch_timeout_key_reuse_stale_take_leaves_newer_bucket() {
853 let policy = BatchPolicy::with_limits(
854 Arc::new(ConstExpr("hot".into())),
855 Arc::new(PropExpr("seq".into())),
856 BatchCompletion::Timeout(60_000), 16,
858 100,
859 16,
860 );
861
862 assert!(policy.accept(mk_exchange(1)).await.is_empty());
864 let first_gen = {
865 let tasks = policy.timeout_tasks.lock().unwrap();
866 tasks.get("hot").map(|e| e.generation).unwrap()
867 };
868
869 policy.cancel_timeout("hot");
873 assert!(policy.take_bucket("hot").is_some());
874 assert!(policy.timeout_tasks.lock().unwrap().is_empty());
875 assert!(policy.accept(mk_exchange(2)).await.is_empty());
876 let second_gen = {
877 let tasks = policy.timeout_tasks.lock().unwrap();
878 tasks.get("hot").map(|e| e.generation).unwrap()
879 };
880 assert_ne!(first_gen, second_gen);
881
882 policy.remove_timeout_task_if_current("hot", first_gen);
885 assert_eq!(
886 policy.timeout_tasks.lock().unwrap().len(),
887 1,
888 "stale generation cleanup must not remove the newer entry"
889 );
890
891 let stolen = policy.take_bucket_if_current_timeout_task("hot", first_gen);
894 assert!(
895 stolen.is_none(),
896 "stale generation must not take the bucket"
897 );
898 assert_eq!(
899 policy.buffered(),
900 1,
901 "newer bucket must remain after the stale combined take"
902 );
903
904 let bucket = policy.take_bucket_if_current_timeout_task("hot", second_gen);
906 assert_eq!(
907 bucket
908 .expect("current generation drains its bucket")
909 .exchanges
910 .len(),
911 1
912 );
913 policy.remove_timeout_task_if_current("hot", second_gen);
914 assert!(policy.timeout_tasks.lock().unwrap().is_empty());
915
916 let flushed = policy.flush().await;
918 assert!(flushed.is_empty());
919 }
920
921 #[tokio::test]
932 async fn batch_timeout_combined_take_atomic_under_interleaved_supersede() {
933 let policy = BatchPolicy::with_limits(
934 Arc::new(ConstExpr("hot".into())),
935 Arc::new(PropExpr("seq".into())),
936 BatchCompletion::Timeout(60_000), 16,
938 100,
939 16,
940 );
941
942 assert!(policy.accept(mk_exchange(1)).await.is_empty());
944 let first_gen = {
945 let tasks = policy.timeout_tasks.lock().unwrap();
946 tasks.get("hot").map(|e| e.generation).unwrap()
947 };
948
949 let (start_tx, start_rx) = std::sync::mpsc::channel::<()>();
950 let (supersede_done_tx, supersede_done_rx) = std::sync::mpsc::channel::<()>();
951 let supersede_done_rx = std::sync::Arc::new(std::sync::Mutex::new(supersede_done_rx));
952 *policy.interleave_hook.lock().unwrap() = Some(Arc::new(move || {
953 start_tx.send(()).expect("hook start signal");
961 let _ = supersede_done_rx
962 .lock()
963 .unwrap()
964 .recv_timeout(std::time::Duration::from_secs(2));
965 }));
966
967 let bg_policy = Arc::clone(&policy);
968 let bg = std::thread::spawn(move || {
969 start_rx.recv().expect("bg start");
970 bg_policy.cancel_timeout("hot");
972 {
973 let mut buckets = bg_policy.buckets.lock().unwrap();
974 buckets.insert(
975 "hot".to_string(),
976 Bucket {
977 exchanges: vec![mk_exchange(2)],
978 },
979 );
980 }
981 let newer_gen = bg_policy.timeout_generation.fetch_add(1, Ordering::SeqCst) + 1;
982 bg_policy.timeout_tasks.lock().unwrap().insert(
983 "hot".to_string(),
984 TimeoutEntry {
985 generation: newer_gen,
986 cancel: CancellationToken::new(),
987 },
988 );
989 supersede_done_tx.send(()).expect("supersede done signal");
990 });
991
992 let taken = policy.take_bucket_if_current_timeout_task("hot", first_gen);
996 bg.join().expect("background thread clean");
997
998 let taken = taken.expect("generation 1 still owned the take at check time");
999 let seq_of = |ex: &Exchange| ex.property("seq").cloned().unwrap_or_default();
1001 assert_eq!(
1002 seq_of(&taken.exchanges[0]),
1003 serde_json::json!(1),
1004 "combined take must retrieve the original bucket, not steal the newer one"
1005 );
1006 assert_eq!(
1008 policy.buffered(),
1009 1,
1010 "newer bucket must remain after the interleaved supersede"
1011 );
1012 let newer_bucket = policy.take_bucket("hot").expect("newer bucket present");
1013 assert_eq!(seq_of(&newer_bucket.exchanges[0]), serde_json::json!(2));
1014
1015 policy.flush().await;
1016 }
1017}