From 12558f4ce95eebabcf12d54dfe9f8d53bf29d9e6 Mon Sep 17 00:00:00 2001 From: WenyXu Date: Fri, 11 Sep 2026 07:56:09 +0000 Subject: [PATCH] refactor: validate batcher limits inside fallible constructors Signed-off-by: WenyXu --- src/common/batcher/src/flush_limiter.rs | 18 +++++---- src/common/batcher/src/flush_policy/timing.rs | 24 +++++------- src/common/batcher/src/notifier.rs | 16 ++++---- src/common/batcher/src/pending_worker.rs | 6 +-- src/common/batcher/src/request_limiter.rs | 22 +++++------ src/servers/src/pending_rows_batcher.rs | 38 ++++++++----------- 6 files changed, 55 insertions(+), 69 deletions(-) diff --git a/src/common/batcher/src/flush_limiter.rs b/src/common/batcher/src/flush_limiter.rs index 1e50baf55c..98b6b2a2f0 100644 --- a/src/common/batcher/src/flush_limiter.rs +++ b/src/common/batcher/src/flush_limiter.rs @@ -12,7 +12,6 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::num::NonZeroUsize; use std::sync::Arc; use tokio::sync::{AcquireError, OwnedSemaphorePermit, Semaphore}; @@ -24,10 +23,10 @@ pub struct FlushLimiter { } impl FlushLimiter { - /// Returns `None` if the limit exceeds Tokio's supported semaphore capacity. - pub fn try_new(max_concurrent_flushes: NonZeroUsize) -> Option { - let permits = max_concurrent_flushes.get(); - (permits <= Semaphore::MAX_PERMITS).then(|| Self { + /// Returns `None` if the limit is zero or exceeds Tokio's supported semaphore capacity. + pub fn try_new(max_concurrent_flushes: usize) -> Option { + let permits = max_concurrent_flushes; + ((1..=Semaphore::MAX_PERMITS).contains(&permits)).then(|| Self { semaphore: Arc::new(Semaphore::new(permits)), }) } @@ -46,12 +45,17 @@ mod tests { #[test] fn test_invalid_capacity() { - assert!(FlushLimiter::try_new(NonZeroUsize::new(usize::MAX).unwrap()).is_none()); + for capacity in [0, Semaphore::MAX_PERMITS + 1, usize::MAX] { + assert!(FlushLimiter::try_new(capacity).is_none()); + } + for capacity in [1, Semaphore::MAX_PERMITS] { + assert!(FlushLimiter::try_new(capacity).is_some()); + } } #[tokio::test] async fn test_clones_share_flush_budget() { - let limiter = FlushLimiter::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let limiter = FlushLimiter::try_new(1).unwrap(); let other = limiter.clone(); let permit = limiter.acquire().await.unwrap(); let waiting = other.acquire(); diff --git a/src/common/batcher/src/flush_policy/timing.rs b/src/common/batcher/src/flush_policy/timing.rs index f9d431a6df..56a9172eaa 100644 --- a/src/common/batcher/src/flush_policy/timing.rs +++ b/src/common/batcher/src/flush_policy/timing.rs @@ -12,7 +12,6 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::num::NonZeroUsize; use std::time::Duration; use tokio::time::Instant; @@ -27,13 +26,13 @@ use crate::pending_batch::PendingBatch; #[derive(Debug, Clone, Copy)] pub struct TimingFlushPolicy { flush_interval: Duration, - max_batch_rows: NonZeroUsize, + max_batch_rows: usize, } impl TimingFlushPolicy { - /// Creates a policy, rejecting zero intervals and unrepresentable deadlines. - pub fn try_new(flush_interval: Duration, max_batch_rows: NonZeroUsize) -> Option { - if !Self::validate(flush_interval) { + /// Creates a policy, rejecting zero limits, zero intervals, and unrepresentable deadlines. + pub fn try_new(flush_interval: Duration, max_batch_rows: usize) -> Option { + if max_batch_rows == 0 || !Self::validate(flush_interval) { return None; } Some(Self { @@ -49,7 +48,7 @@ impl TimingFlushPolicy { /// Checks only the row trigger when a caller drives deadline events separately. pub fn reached_row_threshold(&self, batch: &PendingBatch) -> bool { - !batch.is_empty() && batch.total_rows() >= self.max_batch_rows.get() + !batch.is_empty() && batch.total_rows() >= self.max_batch_rows } } @@ -87,16 +86,14 @@ mod tests { for interval in [Duration::ZERO, Duration::from_secs(1), Duration::MAX] { assert_eq!( TimingFlushPolicy::validate(interval), - TimingFlushPolicy::try_new(interval, NonZeroUsize::MIN).is_some() + TimingFlushPolicy::try_new(interval, 1).is_some() ); } } #[test] fn test_first_submission_deadline() { - let policy = - TimingFlushPolicy::try_new(Duration::from_millis(10), NonZeroUsize::new(100).unwrap()) - .unwrap(); + let policy = TimingFlushPolicy::try_new(Duration::from_millis(10), 100).unwrap(); let mut batch = PendingBatch::new(); let first = Instant::now(); assert_eq!(policy.deadline(&batch), None); @@ -136,9 +133,7 @@ mod tests { #[test] fn test_row_threshold_preserves_complete_submissions() { - let policy = - TimingFlushPolicy::try_new(Duration::from_secs(1), NonZeroUsize::new(3).unwrap()) - .unwrap(); + let policy = TimingFlushPolicy::try_new(Duration::from_secs(1), 3).unwrap(); let now = Instant::now(); for rows in [2, 3, 4] { let mut batch = PendingBatch::new(); @@ -154,7 +149,8 @@ mod tests { #[test] fn test_invalid_interval() { - let max_batch_rows = NonZeroUsize::new(1).unwrap(); + assert!(TimingFlushPolicy::try_new(Duration::from_secs(1), 0).is_none()); + let max_batch_rows = 1; assert!(TimingFlushPolicy::try_new(Duration::ZERO, max_batch_rows).is_none()); assert!(TimingFlushPolicy::try_new(Duration::MAX, max_batch_rows).is_none()); } diff --git a/src/common/batcher/src/notifier.rs b/src/common/batcher/src/notifier.rs index 7f956e4f1d..efc795b938 100644 --- a/src/common/batcher/src/notifier.rs +++ b/src/common/batcher/src/notifier.rs @@ -36,12 +36,12 @@ impl Clone for Notifier { } impl Notifier { - /// Returns `None` if the queue exceeds Tokio's supported capacity. - pub fn try_new(capacity: NonZeroUsize) -> Option<(Self, Receiver)> { - if capacity.get() > Semaphore::MAX_PERMITS { + /// Returns `None` if the queue capacity is zero or exceeds Tokio's supported capacity. + pub fn try_new(capacity: usize) -> Option<(Self, Receiver)> { + if capacity == 0 || capacity > Semaphore::MAX_PERMITS { return None; } - let (sender, receiver) = mpsc::channel(capacity.get()); + let (sender, receiver) = mpsc::channel(capacity); Some((Self { sender }, receiver)) } @@ -92,8 +92,9 @@ mod tests { #[test] fn test_admission() { - assert!(Notifier::::try_new(NonZeroUsize::new(usize::MAX).unwrap()).is_none()); - let (notifier, receiver) = Notifier::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + assert!(Notifier::::try_new(0).is_none()); + assert!(Notifier::::try_new(usize::MAX).is_none()); + let (notifier, receiver) = Notifier::try_new(1).unwrap(); let other = notifier.clone(); assert_eq!(other.max_capacity(), 1); assert_eq!(notifier.try_notify(1), Ok(())); @@ -107,8 +108,7 @@ mod tests { // One-factor comparison: only delivery concurrency changes. for _ in 0..2 { for concurrency in [1, 2] { - let (notifier, receiver) = - Notifier::try_new(NonZeroUsize::new(3).unwrap()).unwrap(); + let (notifier, receiver) = Notifier::try_new(3).unwrap(); for item in 0..3 { notifier.try_notify(item).unwrap(); } diff --git a/src/common/batcher/src/pending_worker.rs b/src/common/batcher/src/pending_worker.rs index c70731fd5a..2adda17a5e 100644 --- a/src/common/batcher/src/pending_worker.rs +++ b/src/common/batcher/src/pending_worker.rs @@ -117,7 +117,6 @@ impl PendingWorker { #[cfg(test)] mod tests { use std::future::Future; - use std::num::NonZeroUsize; use std::task::Poll; use std::time::Duration; @@ -125,10 +124,7 @@ mod tests { use crate::flush_policy::timing::TimingFlushPolicy; fn worker(rows: usize) -> PendingWorker { - PendingWorker::new( - TimingFlushPolicy::try_new(Duration::from_millis(10), NonZeroUsize::new(rows).unwrap()) - .unwrap(), - ) + PendingWorker::new(TimingFlushPolicy::try_new(Duration::from_millis(10), rows).unwrap()) } async fn assert_wait_pending(worker: &mut PendingWorker) { diff --git a/src/common/batcher/src/request_limiter.rs b/src/common/batcher/src/request_limiter.rs index 46b8a96f3f..c488bec02d 100644 --- a/src/common/batcher/src/request_limiter.rs +++ b/src/common/batcher/src/request_limiter.rs @@ -12,7 +12,6 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::num::NonZeroUsize; use std::sync::Arc; use tokio::sync::{AcquireError, OwnedSemaphorePermit, Semaphore}; @@ -27,10 +26,10 @@ pub struct RequestLimiter { } impl RequestLimiter { - /// Returns `None` if the limit exceeds Tokio's supported semaphore capacity. - pub fn try_new(max_inflight_requests: NonZeroUsize) -> Option { - let permits = max_inflight_requests.get(); - (permits <= Semaphore::MAX_PERMITS).then(|| Self { + /// Returns `None` if the limit is zero or exceeds Tokio's supported semaphore capacity. + pub fn try_new(max_inflight_requests: usize) -> Option { + let permits = max_inflight_requests; + ((1..=Semaphore::MAX_PERMITS).contains(&permits)).then(|| Self { semaphore: Arc::new(Semaphore::new(permits)), }) } @@ -48,7 +47,6 @@ impl RequestLimiter { #[cfg(test)] mod tests { use std::future::{Future, poll_fn}; - use std::num::NonZeroUsize; use std::pin::Pin; use std::task::Poll; @@ -63,16 +61,16 @@ mod tests { #[test] fn test_capacity_boundaries() { for capacity in [1, Semaphore::MAX_PERMITS] { - assert!(RequestLimiter::try_new(NonZeroUsize::new(capacity).unwrap()).is_some()); + assert!(RequestLimiter::try_new(capacity).is_some()); } - for capacity in [Semaphore::MAX_PERMITS + 1, usize::MAX] { - assert!(RequestLimiter::try_new(NonZeroUsize::new(capacity).unwrap()).is_none()); + for capacity in [0, Semaphore::MAX_PERMITS + 1, usize::MAX] { + assert!(RequestLimiter::try_new(capacity).is_none()); } } #[tokio::test] async fn test_last_submission_releases_request_slot() { - let limiter = RequestLimiter::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let limiter = RequestLimiter::try_new(1).unwrap(); let other = limiter.clone(); let request = limiter.acquire().await.unwrap(); let first_submission = request.clone(); @@ -90,7 +88,7 @@ mod tests { #[tokio::test] async fn test_cancelled_acquisition_does_not_leak_capacity() { - let limiter = RequestLimiter::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let limiter = RequestLimiter::try_new(1).unwrap(); let permit = limiter.acquire().await.unwrap(); let mut cancelled = Box::pin(limiter.acquire()); assert!(is_pending(cancelled.as_mut()).await); @@ -105,7 +103,7 @@ mod tests { #[tokio::test] async fn test_independent_requests_consume_separate_slots() { - let limiter = RequestLimiter::try_new(NonZeroUsize::new(2).unwrap()).unwrap(); + let limiter = RequestLimiter::try_new(2).unwrap(); let first = limiter.acquire().await.unwrap(); let second = limiter.acquire().await.unwrap(); let mut third = Box::pin(limiter.acquire()); diff --git a/src/servers/src/pending_rows_batcher.rs b/src/servers/src/pending_rows_batcher.rs index 442e8d6ef0..f945641f2e 100644 --- a/src/servers/src/pending_rows_batcher.rs +++ b/src/servers/src/pending_rows_batcher.rs @@ -382,13 +382,12 @@ impl PendingRowsBatcher { return None; } - let flush_policy = - TimingFlushPolicy::try_new(flush_interval, NonZeroUsize::new(max_batch_rows)?)?; - let flush_limiter = FlushLimiter::try_new(NonZeroUsize::new(max_concurrent_flushes)?)?; + let flush_policy = TimingFlushPolicy::try_new(flush_interval, max_batch_rows)?; + let flush_limiter = FlushLimiter::try_new(max_concurrent_flushes)?; - let request_limiter = RequestLimiter::try_new(NonZeroUsize::new(max_inflight_requests)?)?; + let request_limiter = RequestLimiter::try_new(max_inflight_requests)?; let (flow_notification_tx, flow_notification_rx) = - Notifier::try_new(flow_notification_queue_capacity)?; + Notifier::try_new(flow_notification_queue_capacity.get())?; let (shutdown, _) = broadcast::channel(1); let pending_rows_batch_sync = pending_rows_batch_sync_enabled(); @@ -1507,8 +1506,7 @@ fn notify_flow_dirty_windows_after_flush( table_flownode_set_cache: TableFlownodeSetCacheRef, node_manager: NodeManagerRef, ) { - let (tx, rx) = - Notifier::try_new(NonZeroUsize::new(table_batches.len().max(1)).unwrap()).unwrap(); + let (tx, rx) = Notifier::try_new(table_batches.len().max(1)).unwrap(); start_flow_notification_worker(rx, table_flownode_set_cache, node_manager); enqueue_flow_notifications(table_batches, &tx); } @@ -1834,7 +1832,6 @@ mod tests { use std::any::Any; use std::collections::{HashMap, HashSet}; use std::future::{Future, poll_fn}; - use std::num::NonZeroUsize; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::task::Poll; @@ -2126,7 +2123,7 @@ mod tests { #[test] fn test_flow_notification_queue_drops_when_full() { - let (tx, mut rx) = Notifier::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let (tx, mut rx) = Notifier::try_new(1).unwrap(); let notification = |table_id| crate::pending_rows_batcher::FlowNotification { table_id, timestamps: vec![table_id as i64], @@ -2347,10 +2344,8 @@ mod tests { let ctx = session::context::QueryContext::arc(); let (response_tx, _response_rx) = oneshot::channel(); let permit = Arc::new(Semaphore::new(1)).try_acquire_owned().unwrap(); - let mut pending_flush = PendingCore::new( - TimingFlushPolicy::try_new(Duration::from_secs(10), NonZeroUsize::new(1).unwrap()) - .unwrap(), - ); + let mut pending_flush = + PendingCore::new(TimingFlushPolicy::try_new(Duration::from_secs(10), 1).unwrap()); pending_flush.submit( FlushWaiter { response_tx, @@ -2392,10 +2387,8 @@ mod tests { #[tokio::test] async fn test_drain_batch_preserves_unready_state_and_clears_zero_rows() { for total_rows in [0, 1] { - let mut pending_flush = PendingCore::new( - TimingFlushPolicy::try_new(Duration::from_secs(10), NonZeroUsize::new(2).unwrap()) - .unwrap(), - ); + let mut pending_flush = + PendingCore::new(TimingFlushPolicy::try_new(Duration::from_secs(10), 2).unwrap()); let mut batch = Some(PendingBatch::new(session::context::QueryContext::arc())); let semaphore = Arc::new(Semaphore::new(1)); let (response_tx, mut response_rx) = oneshot::channel(); @@ -2606,7 +2599,7 @@ mod tests { cache: TableFlownodeSetCacheRef, node_manager: NodeManagerRef, ) -> Notifier { - let (tx, rx) = Notifier::try_new(NonZeroUsize::new(16).unwrap()).unwrap(); + let (tx, rx) = Notifier::try_new(16).unwrap(); start_flow_notification_worker(rx, cache, node_manager); tx } @@ -2940,7 +2933,7 @@ mod tests { #[tokio::test] async fn test_cancelled_waiter_retains_request_slot_until_notification() { - let limiter = RequestLimiter::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let limiter = RequestLimiter::try_new(1).unwrap(); let (response_tx, response_rx) = oneshot::channel(); let waiter = FlushWaiter { response_tx, @@ -3092,11 +3085,10 @@ mod tests { datanodes: Arc::new(HashMap::new()), }); let catalog_manager = MemoryCatalogManager::with_default_setup(); - let (flow_notification_tx, _flow_notification_rx) = - Notifier::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let (flow_notification_tx, _flow_notification_rx) = Notifier::try_new(1).unwrap(); let (shutdown, _) = broadcast::channel(1); - let flush_limiter = FlushLimiter::try_new(NonZeroUsize::new(1).unwrap()).unwrap(); + let flush_limiter = FlushLimiter::try_new(1).unwrap(); start_worker( key.clone(), worker_tx.clone(), @@ -3108,7 +3100,7 @@ mod tests { catalog_manager, flow_notification_tx, worker_idle_timeout, - TimingFlushPolicy::try_new(flush_interval, NonZeroUsize::new(3).unwrap()).unwrap(), + TimingFlushPolicy::try_new(flush_interval, 3).unwrap(), flush_limiter.clone(), );