1mod batch;
18mod flow_notifier;
19mod metrics;
20mod pending_batch;
21mod pending_worker;
22
23use std::num::NonZeroUsize;
24use std::sync::Arc;
25use std::time::Duration;
26
27use arrow::record_batch::RecordBatch;
28use async_trait::async_trait;
29use common_batcher::flush_limiter::FlushLimiter;
30use common_batcher::flush_policy::timing::TimingFlushPolicy;
31use common_batcher::request_limiter::RequestLimiter;
32use common_batcher::worker_registry::WorkerRegistry;
33use operator::batcher::PendingRowsBatcher;
34use operator::error::{BatchFlushSnafu, Result, UnexpectedSnafu};
35use operator::insert::Inserter;
36use session::context::QueryContextRef;
37use snafu::ResultExt;
38use table::metadata::TableInfoRef;
39use tokio::sync::{OwnedSemaphorePermit, Semaphore, broadcast, oneshot};
40
41use crate::batcher::pending_rows_batch_sync_enabled;
42use crate::batcher::table::flow_notifier::FlowNotifier;
43use crate::batcher::table::metrics::PENDING_WORKERS;
44use crate::batcher::table::pending_batch::PendingBatch;
45use crate::batcher::table::pending_worker::{PendingWorker, WorkerCommand, start_worker};
46
47#[derive(Clone, Debug, Hash, PartialEq, Eq)]
48struct BatchKey {
49 catalog: String,
50 schema: String,
51 table_name: String,
52 skip_wal: bool,
53}
54
55fn batch_key_from_ctx(table_name: &str, ctx: &QueryContextRef) -> BatchKey {
57 BatchKey {
58 catalog: ctx.current_catalog().to_string(),
59 schema: ctx.current_schema().clone(),
60 table_name: table_name.to_string(),
61 skip_wal: ctx.skip_wal(),
62 }
63}
64
65type PendingWorkers = WorkerRegistry<BatchKey, WorkerCommand>;
66
67pub struct TablePendingRowsBatcher {
69 workers: Arc<PendingWorkers>,
70 flush_policy: TimingFlushPolicy,
71 flush_limiter: FlushLimiter,
72 request_limiter: RequestLimiter,
73 worker_channel_capacity: usize,
74 pending_rows_batch_sync: bool,
75 worker_idle_timeout: Duration,
76 inserter: Arc<Inserter>,
77 flow_notifier: FlowNotifier,
78 shutdown: broadcast::Sender<()>,
79}
80
81impl TablePendingRowsBatcher {
82 pub fn try_new(
84 flush_interval: Duration,
85 max_batch_rows: usize,
86 max_concurrent_flushes: usize,
87 worker_channel_capacity: usize,
88 max_inflight_requests: usize,
89 flow_notification_queue_capacity: NonZeroUsize,
90 inserter: Arc<Inserter>,
91 ) -> Option<Arc<Self>> {
92 if worker_channel_capacity == 0
93 || worker_channel_capacity > Semaphore::MAX_PERMITS
94 || max_inflight_requests == 0
95 || max_inflight_requests > Semaphore::MAX_PERMITS
96 {
97 return None;
98 }
99 let flush_policy = TimingFlushPolicy::try_new(flush_interval, max_batch_rows)?;
100 let flush_limiter = FlushLimiter::try_new(max_concurrent_flushes)?;
101 let request_limiter = RequestLimiter::try_new(max_inflight_requests)?;
102 let flow_notifier = FlowNotifier::new(
103 inserter.table_flownode_set_cache().clone(),
104 inserter.node_manager().clone(),
105 flow_notification_queue_capacity,
106 )?;
107 let (shutdown, _) = broadcast::channel(1);
108 PENDING_WORKERS.set(0);
109 Some(Arc::new(Self {
110 workers: Arc::new(WorkerRegistry::new()),
111 flush_policy,
112 flush_limiter,
113 request_limiter,
114 worker_channel_capacity,
115 pending_rows_batch_sync: pending_rows_batch_sync_enabled(),
116 worker_idle_timeout: flush_interval.checked_mul(3).unwrap_or(flush_interval),
117 inserter,
118 flow_notifier,
119 shutdown,
120 }))
121 }
122
123 async fn worker(&self, key: &BatchKey) -> PendingWorker {
124 let (tx, receiver) = self
125 .workers
126 .get_or_create(key.clone(), self.worker_channel_capacity)
127 .await;
128 if let Some(rx) = receiver {
129 start_worker(
130 key.clone(),
131 tx.clone(),
132 self.workers.clone(),
133 rx,
134 self.shutdown.clone(),
135 self.flush_policy,
136 self.flush_limiter.clone(),
137 self.inserter.clone(),
138 self.flow_notifier.clone(),
139 self.worker_idle_timeout,
140 );
141 PENDING_WORKERS.set(self.workers.len().await as i64);
142 }
143 PendingWorker { tx }
144 }
145}
146
147impl Drop for TablePendingRowsBatcher {
148 fn drop(&mut self) {
149 let _ = self.shutdown.send(());
150 }
151}
152
153#[async_trait]
154impl PendingRowsBatcher for TablePendingRowsBatcher {
155 async fn acquire(&self) -> Result<Arc<OwnedSemaphorePermit>> {
158 self.request_limiter.acquire().await.map_err(|_| {
159 UnexpectedSnafu {
160 violated: "batcher admission closed".to_string(),
161 }
162 .build()
163 })
164 }
165
166 async fn submit(
169 &self,
170 table_info: TableInfoRef,
171 batch: RecordBatch,
172 ctx: QueryContextRef,
173 permit: Arc<OwnedSemaphorePermit>,
174 ) -> Result<usize> {
175 let total_rows = batch.num_rows();
176 if total_rows == 0 {
177 return Ok(0);
178 }
179 if table_info.catalog_name != ctx.current_catalog()
180 || table_info.schema_name != ctx.current_schema()
181 {
182 return UnexpectedSnafu {
183 violated: "batch table and request database differ".to_string(),
184 }
185 .fail();
186 }
187 let key = batch_key_from_ctx(&table_info.name, &ctx);
188 let (response_tx, response_rx) = oneshot::channel();
189 let pending = PendingBatch {
190 table_info,
191 batch,
192 ctx,
193 response_tx,
194 _permit: permit,
195 };
196 let mut command = Some(WorkerCommand::Submit(pending));
197 for _ in 0..2 {
198 let worker = self.worker(&key).await;
199 let Some(pending) = command.take() else { break };
200 match worker.tx.send(pending).await {
201 Ok(()) => break,
202 Err(error) => {
203 command = Some(error.0);
204 if self.workers.remove_if_same(&key, &worker.tx).await {
205 PENDING_WORKERS.set(self.workers.len().await as i64);
206 }
207 }
208 }
209 }
210 if command.is_some() {
211 return UnexpectedSnafu {
212 violated: "batch worker channel closed".to_string(),
213 }
214 .fail();
215 }
216 if !self.pending_rows_batch_sync {
217 return Ok(total_rows);
218 }
219 response_rx
220 .await
221 .map_err(|_| {
222 UnexpectedSnafu {
223 violated: "batch worker stopped before reporting its write result".to_string(),
224 }
225 .build()
226 })?
227 .context(BatchFlushSnafu)?;
228 Ok(total_rows)
229 }
230}
231
232#[cfg(test)]
233mod tests {
234 use std::collections::HashMap;
235 use std::num::NonZeroUsize;
236 use std::sync::atomic::{AtomicUsize, Ordering};
237 use std::sync::{Arc, Mutex};
238 use std::time::Duration;
239
240 use api::region::RegionResponse;
241 use api::v1::helper::{tag_column_schema, time_index_column_schema};
242 use api::v1::region::{RegionRequest, bulk_insert_request, region_request};
243 use api::v1::value::ValueData;
244 use api::v1::{ColumnDataType, Row, Rows, Value};
245 use arrow::array::{ArrayRef, Int32Array, TimestampMillisecondArray};
246 use arrow::datatypes::Schema as ArrowSchema;
247 use arrow::record_batch::RecordBatch;
248 use catalog::memory::MemoryCatalogManager;
249 use common_batcher::flush_limiter::FlushLimiter;
250 use common_batcher::flush_policy::timing::TimingFlushPolicy;
251 use common_batcher::request_limiter::RequestLimiter;
252 use common_batcher::worker_registry::WorkerRegistry;
253 use common_catalog::consts::default_engine;
254 use common_grpc::flight::FlightDecoder;
255 use common_meta::error::Result as MetaResult;
256 use common_meta::peer::Peer;
257 use common_meta::test_util::{MockDatanodeHandler, MockDatanodeManager};
258 use common_query::OutputData;
259 use common_query::request::QueryRequest;
260 use common_recordbatch::SendableRecordBatchStream;
261 use common_telemetry::info;
262 use datatypes::schema::{ColumnDefaultConstraint, SchemaBuilder};
263 use datatypes::value::Value as DtValue;
264 use datatypes::vectors::{Int32Vector, TimestampMillisecondVector, VectorRef};
265 use operator::batcher::PendingRowsBatcher;
266 use operator::error::Error;
267 use operator::insert::Inserter;
268 use operator::metrics::DIST_INGEST_ROW_COUNT;
269 use operator::req_convert::insert::rows_to_record_batch;
270 use operator::test_util::{
271 create_partition_rule_manager, new_test_table_info, prepare_mocked_backend,
272 };
273 use session::context::{Channel, QueryContext};
274 use store_api::storage::RegionId;
275 use table::dist_table::DistTable;
276 use table::metadata::TableInfoRef;
277 use table::requests::InsertRequest;
278 use tokio::sync::{broadcast, mpsc, oneshot};
279 use tokio::time::timeout;
280
281 use crate::batcher::table::flow_notifier::FlowNotifier;
282 use crate::batcher::table::pending_batch::PendingBatch;
283 use crate::batcher::table::pending_worker::{WorkerCommand, start_worker};
284 use crate::batcher::table::{PendingWorkers, TablePendingRowsBatcher, batch_key_from_ctx};
285 use crate::batcher::test_util::mock_table_flownode_cache;
286
287 #[derive(Clone)]
288 struct BulkHandler {
289 requests: Arc<Mutex<Vec<RecordBatch>>>,
290 report_missing_row: bool,
291 expected_skip_wal: bool,
292 expected_schema: String,
293 }
294
295 #[async_trait::async_trait]
296 impl MockDatanodeHandler for BulkHandler {
297 async fn handle(&self, peer: &Peer, request: RegionRequest) -> MetaResult<RegionResponse> {
298 assert_eq!(3, peer.id);
299 assert_eq!(self.expected_schema, request.header.unwrap().dbname);
300 let Some(region_request::Body::BulkInsert(request)) = request.body else {
301 panic!("expected bulk insert")
302 };
303 assert_eq!(RegionId::new(1, 3).as_u64(), request.region_id);
304 assert_eq!(self.expected_skip_wal, request.skip_wal);
305 assert!(request.partition_expr_version.is_some());
306 let Some(bulk_insert_request::Body::ArrowIpc(ipc)) = request.body else {
307 panic!("expected Arrow IPC")
308 };
309 let batch = FlightDecoder::try_from_schema_bytes(&ipc.schema)
310 .unwrap()
311 .try_decode_record_batch(&ipc.data_header, &ipc.payload)
312 .unwrap();
313 let rows = batch.num_rows();
314 self.requests.lock().unwrap().push(batch);
315 Ok(RegionResponse::new(
316 rows - usize::from(self.report_missing_row),
317 ))
318 }
319
320 async fn handle_query(
321 &self,
322 _: &Peer,
323 _: QueryRequest,
324 ) -> MetaResult<SendableRecordBatchStream> {
325 panic!("batching must not query the datanode")
326 }
327 }
328
329 #[tokio::test]
330 async fn test_acknowledgement_preserves_admission() {
331 use operator::error::UnexpectedSnafu;
332
333 use crate::batcher::pending_rows_batch_sync_enabled;
334 use crate::batcher::table::batch_key_from_ctx;
335 use crate::batcher::table::pending_batch::notify_batches;
336
337 for sync in [false, true] {
338 for fail in [false, true] {
339 let backend = prepare_mocked_backend().await;
340 let nodes = Arc::new(MockDatanodeManager::new(BulkHandler {
341 requests: Arc::new(Mutex::new(Vec::new())),
342 report_missing_row: false,
343 expected_skip_wal: false,
344 expected_schema: "public".to_string(),
345 }));
346 let inserter = Arc::new(Inserter::new(
347 MemoryCatalogManager::new(),
348 create_partition_rule_manager(backend).await,
349 nodes,
350 mock_table_flownode_cache(1, vec![]).await,
351 true,
352 ));
353 let mut batcher = TablePendingRowsBatcher::try_new(
354 Duration::from_secs(3600),
355 1,
356 1,
357 1,
358 1,
359 NonZeroUsize::new(1).unwrap(),
360 inserter,
361 )
362 .unwrap();
363 assert_eq!(
364 batcher.pending_rows_batch_sync,
365 pending_rows_batch_sync_enabled()
366 );
367 Arc::get_mut(&mut batcher).unwrap().pending_rows_batch_sync = sync;
368 let table = Arc::new(new_test_table_info(1, "ack", [0].into_iter()));
369 let ctx = QueryContext::arc();
370 let key = batch_key_from_ctx(&table.name, &ctx);
371 let (_, receiver) = batcher.workers.get_or_create(key, 1).await;
373 let mut receiver = receiver.unwrap();
374 let batch = RecordBatch::try_from_iter(vec![(
375 "a",
376 Arc::new(Int32Array::from(vec![1])) as ArrayRef,
377 )])
378 .unwrap();
379 let permit = batcher.acquire().await.unwrap();
380 let submitter = batcher.clone();
381 let submitted =
382 tokio::spawn(async move { submitter.submit(table, batch, ctx, permit).await });
383 let WorkerCommand::Submit(pending) =
384 timeout(Duration::from_secs(5), receiver.recv())
385 .await
386 .unwrap()
387 .unwrap();
388 let mut submitted = Some(submitted);
389 if sync {
390 assert!(!submitted.as_ref().unwrap().is_finished());
391 } else {
392 assert_eq!(
393 timeout(Duration::from_secs(5), submitted.take().unwrap())
394 .await
395 .unwrap()
396 .unwrap()
397 .unwrap(),
398 1
399 );
400 }
401 let acquire = batcher.acquire();
403 tokio::pin!(acquire);
404 assert!(futures::poll!(acquire.as_mut()).is_pending());
405 let result = if fail {
406 Err(Arc::new(
407 UnexpectedSnafu {
408 violated: "flush failed".to_string(),
409 }
410 .build(),
411 ))
412 } else {
413 Ok(())
414 };
415 notify_batches(vec![pending], result);
416 if let Some(submitted) = submitted {
417 let result = timeout(Duration::from_secs(5), submitted)
418 .await
419 .unwrap()
420 .unwrap();
421 assert_eq!(result.is_err(), fail);
422 if !fail {
423 assert_eq!(result.unwrap(), 1);
424 }
425 }
426 timeout(Duration::from_secs(5), acquire)
427 .await
428 .unwrap()
429 .unwrap();
430 }
431 }
432 }
433
434 fn rows(value: i32) -> Rows {
435 Rows {
436 schema: vec![
437 tag_column_schema("a", ColumnDataType::Int32),
438 time_index_column_schema("ts", ColumnDataType::TimestampMillisecond),
439 ],
440 rows: vec![Row {
441 values: vec![
442 Value {
443 value_data: Some(ValueData::I32Value(value)),
444 },
445 Value {
446 value_data: Some(ValueData::TimestampMillisecondValue(
447 1000 + i64::from(value),
448 )),
449 },
450 ],
451 }],
452 }
453 }
454
455 async fn run_bulk_case(
456 max_batch_rows: usize,
457 report_missing_row: bool,
458 shared_request: bool,
459 ) -> usize {
460 run_bulk_case_with_wal(max_batch_rows, report_missing_row, shared_request, false).await
461 }
462
463 async fn run_bulk_case_with_wal(
464 max_batch_rows: usize,
465 report_missing_row: bool,
466 shared_request: bool,
467 skip_wal: bool,
468 ) -> usize {
469 let backend = prepare_mocked_backend().await;
470 let partitions = create_partition_rule_manager(backend.clone()).await;
471 let captured = Arc::new(Mutex::new(Vec::new()));
472 static NEXT_DATABASE: AtomicUsize = AtomicUsize::new(0);
474 let schema_name = format!(
475 "ingest_count_{}",
476 NEXT_DATABASE.fetch_add(1, Ordering::Relaxed)
477 );
478 let nodes = Arc::new(MockDatanodeManager::new(BulkHandler {
479 requests: captured.clone(),
480 report_missing_row,
481 expected_skip_wal: skip_wal,
482 expected_schema: schema_name.clone(),
483 }));
484 let inserter = Arc::new(Inserter::new(
485 MemoryCatalogManager::new(),
486 partitions.clone(),
487 nodes.clone(),
488 mock_table_flownode_cache(1, vec![]).await,
489 true,
490 ));
491 let mut table = new_test_table_info(1, "table_1", [1, 2, 3].into_iter());
492 table.schema_name = schema_name;
493 table.meta.engine = default_engine().to_string();
494 let mut columns = table.meta.schema.column_schemas().to_vec();
495 columns[2] = columns[2]
496 .clone()
497 .with_default_constraint(Some(ColumnDefaultConstraint::Value(DtValue::Int32(7))))
498 .unwrap();
499 table.meta.schema = Arc::new(
500 SchemaBuilder::try_from(columns)
501 .unwrap()
502 .version(123)
503 .build()
504 .unwrap(),
505 );
506 let table = Arc::new(table);
507 let first = rows_to_record_batch(&rows(1), &table).unwrap();
508 let second = rows_to_record_batch(&rows(2), &table).unwrap();
509 let batcher: Arc<dyn PendingRowsBatcher> = TablePendingRowsBatcher::try_new(
512 Duration::from_secs(3600),
513 max_batch_rows,
514 1,
515 4,
516 if shared_request { 1 } else { 4 },
517 NonZeroUsize::new(16).unwrap(),
518 inserter,
519 )
520 .unwrap();
521 let context = |channel| {
522 let mut ctx = QueryContext::with_channel("greptime", &table.schema_name, channel);
523 ctx.set_batching_enabled(true);
524 Arc::new(ctx)
525 };
526 let influx_ctx = context(Channel::Influx);
527 let opentsdb_ctx = context(Channel::Opentsdb);
528 let ingest_count =
529 DIST_INGEST_ROW_COUNT.with_label_values(&[influx_ctx.get_db_string().as_str()]);
530 assert_eq!(0, ingest_count.get());
531 influx_ctx.set_skip_wal(skip_wal);
532 opentsdb_ctx.set_skip_wal(skip_wal);
533 let (first_result, second_result) = timeout(Duration::from_secs(5), async {
534 if shared_request {
535 let permit = batcher.acquire().await.unwrap();
536 tokio::join!(
537 batcher.submit(table.clone(), first, influx_ctx, permit.clone()),
538 batcher.submit(table.clone(), second, opentsdb_ctx, permit),
539 )
540 } else {
541 let inserter = Inserter::new(
542 MemoryCatalogManager::new_with_table(DistTable::table(table.clone())),
543 partitions,
544 nodes,
545 mock_table_flownode_cache(1, vec![]).await,
546 true,
547 )
548 .with_pending_rows_batcher(Some(batcher));
549 let request = |value: i32| InsertRequest {
550 catalog_name: table.catalog_name.clone(),
551 schema_name: table.schema_name.clone(),
552 table_name: table.name.clone(),
553 columns_values: HashMap::from([
554 (
555 "a".to_string(),
556 Arc::new(Int32Vector::from_slice([value])) as VectorRef,
557 ),
558 (
559 "ts".to_string(),
560 Arc::new(TimestampMillisecondVector::from_slice([
561 1000 + i64::from(value)
562 ])) as VectorRef,
563 ),
564 ]),
565 skip_wal,
566 };
567 let (first, second) = tokio::join!(
568 inserter.handle_table_insert(request(1), influx_ctx),
569 inserter.handle_table_insert(request(2), opentsdb_ctx),
570 );
571 (
572 first.map(|output| match output.data {
573 OutputData::AffectedRows(rows) => rows,
574 _ => panic!("expected affected rows"),
575 }),
576 second.map(|output| match output.data {
577 OutputData::AffectedRows(rows) => rows,
578 _ => panic!("expected affected rows"),
579 }),
580 )
581 }
582 })
583 .await
584 .expect("the row threshold did not dispatch the combined bulk insert");
585 if report_missing_row {
586 assert!(first_result.is_err());
587 assert!(second_result.is_err());
588 } else {
589 assert_eq!(1, first_result.unwrap());
590 assert_eq!(1, second_result.unwrap());
591 }
592 assert_eq!(if report_missing_row { 0 } else { 2 }, ingest_count.get());
593 let requests = captured.lock().unwrap();
594 let mut actual_rows = Vec::new();
595 for batch in requests.iter() {
596 assert_eq!(
597 vec!["a", "ts", "b"],
598 batch
599 .schema()
600 .fields()
601 .iter()
602 .map(|field| field.name().as_str())
603 .collect::<Vec<_>>()
604 );
605 let a = batch
606 .column(0)
607 .as_any()
608 .downcast_ref::<Int32Array>()
609 .unwrap();
610 let ts = batch
611 .column(1)
612 .as_any()
613 .downcast_ref::<TimestampMillisecondArray>()
614 .unwrap();
615 let b = batch
616 .column(2)
617 .as_any()
618 .downcast_ref::<Int32Array>()
619 .unwrap();
620 actual_rows.extend(
621 (0..batch.num_rows())
622 .map(|index| (a.value(index), ts.value(index), b.value(index))),
623 );
624 }
625 actual_rows.sort_unstable();
626 assert_eq!(vec![(1, 1001, 7), (2, 1002, 7)], actual_rows);
627 requests.len()
628 }
629
630 #[tokio::test]
631 async fn test_inserter_counts_batched_rows_once() {
632 for report_missing_row in [false, true] {
633 assert_eq!(1, run_bulk_case(2, report_missing_row, false).await);
634 }
635 }
636
637 #[tokio::test]
638 async fn test_row_threshold_ablation_reduces_bulk_requests() {
639 for repetition in 0..2 {
640 let direct_count = run_bulk_case(1, false, false).await;
641 let combined_count = run_bulk_case(2, false, false).await;
642 assert_eq!(2, direct_count);
643 assert_eq!(1, combined_count);
644 info!(
645 "repetition={repetition} rows=2 threshold=1 bulk_requests={direct_count}; threshold=2 bulk_requests={combined_count}"
646 );
647 }
648 }
649
650 #[tokio::test]
651 async fn test_original_request_shares_one_admission_slot() {
652 assert_eq!(1, run_bulk_case(2, false, true).await);
653 }
654
655 #[tokio::test]
656 async fn test_bulk_preserves_skip_wal() {
657 for skip_wal in [false, true] {
658 assert_eq!(1, run_bulk_case_with_wal(2, false, false, skip_wal).await);
659 }
660 }
661
662 #[test]
663 fn test_batch_key_uses_name_and_wal_policy() {
664 let ctx = QueryContext::arc();
665 let table = new_test_table_info(1, "table_1", [1, 2, 3].into_iter());
666 let mut replacement = table.clone();
667 replacement.ident.table_id = 2;
668 replacement.ident.version += 1;
669 let key = batch_key_from_ctx(&table.name, &ctx);
670 assert_eq!(key, batch_key_from_ctx(&replacement.name, &ctx));
671 assert_ne!(key, batch_key_from_ctx("table_2", &ctx));
672 for (catalog, schema) in [("other", "public"), ("greptime", "other")] {
673 let other = Arc::new(QueryContext::with_channel(catalog, schema, Channel::Influx));
674 assert_ne!(key, batch_key_from_ctx(&table.name, &other));
675 }
676 ctx.set_skip_wal(true);
677 assert_ne!(key, batch_key_from_ctx(&table.name, &ctx));
678 }
679
680 const WORKER_TIMEOUT: Duration = Duration::from_secs(5);
681
682 #[derive(Clone)]
683 struct GatedHandler {
684 inner: BulkHandler,
685 entered: mpsc::UnboundedSender<oneshot::Sender<()>>,
686 }
687
688 #[async_trait::async_trait]
689 impl MockDatanodeHandler for GatedHandler {
690 async fn handle(&self, peer: &Peer, request: RegionRequest) -> MetaResult<RegionResponse> {
691 let response = self.inner.handle(peer, request).await?;
692 let (release_tx, release_rx) = oneshot::channel();
693 self.entered.send(release_tx).unwrap();
694 release_rx.await.unwrap();
695 Ok(response)
696 }
697
698 async fn handle_query(
699 &self,
700 _: &Peer,
701 _: QueryRequest,
702 ) -> MetaResult<SendableRecordBatchStream> {
703 panic!("worker must not query the datanode")
704 }
705 }
706
707 struct WorkerTest {
708 tx: mpsc::Sender<WorkerCommand>,
709 shutdown: broadcast::Sender<()>,
710 workers: Arc<PendingWorkers>,
711 table: TableInfoRef,
712 entered: mpsc::UnboundedReceiver<oneshot::Sender<()>>,
713 }
714
715 impl Drop for WorkerTest {
716 fn drop(&mut self) {
717 let _ = self.shutdown.send(());
718 }
719 }
720
721 impl WorkerTest {
722 async fn new(max_rows: usize, max_flushes: usize) -> Self {
723 let backend = prepare_mocked_backend().await;
724 let partitions = create_partition_rule_manager(backend.clone()).await;
725 let (entered_tx, entered) = mpsc::unbounded_channel();
726 let nodes = Arc::new(MockDatanodeManager::new(GatedHandler {
727 inner: BulkHandler {
728 requests: Arc::new(Mutex::new(Vec::new())),
729 report_missing_row: false,
730 expected_skip_wal: false,
731 expected_schema: "public".to_string(),
732 },
733 entered: entered_tx,
734 }));
735 let cache = mock_table_flownode_cache(1, vec![]).await;
736 let notifier =
737 FlowNotifier::new(cache.clone(), nodes.clone(), NonZeroUsize::new(16).unwrap())
738 .unwrap();
739 let inserter = Arc::new(Inserter::new(
740 MemoryCatalogManager::new(),
741 partitions,
742 nodes,
743 cache,
744 true,
745 ));
746 let table = Arc::new(new_test_table_info(1, "table_1", [1, 2, 3].into_iter()));
747 let key = batch_key_from_ctx(&table.name, &QueryContext::arc());
748 let workers = Arc::new(WorkerRegistry::new());
749 let (tx, rx) = mpsc::channel(1);
750 workers.get_or_insert_with(key.clone(), || tx.clone()).await;
751 let (shutdown, _) = broadcast::channel(1);
752 start_worker(
753 key,
754 tx.clone(),
755 workers.clone(),
756 rx,
757 shutdown.clone(),
758 TimingFlushPolicy::try_new(Duration::from_secs(3600), max_rows).unwrap(),
759 FlushLimiter::try_new(max_flushes).unwrap(),
760 inserter,
761 notifier,
762 Duration::from_secs(10800),
763 );
764 Self {
765 tx,
766 shutdown,
767 workers,
768 table,
769 entered,
770 }
771 }
772
773 async fn submit(
774 &self,
775 count: usize,
776 changed_schema: bool,
777 ) -> oneshot::Receiver<Result<(), Arc<Error>>> {
778 self.submit_with_table(count, changed_schema, self.table.clone())
779 .await
780 }
781
782 async fn submit_with_table(
783 &self,
784 count: usize,
785 changed_schema: bool,
786 table_info: TableInfoRef,
787 ) -> oneshot::Receiver<Result<(), Arc<Error>>> {
788 let mut input = rows(1);
789 input.rows = vec![input.rows[0].clone(); count];
790 let mut batch = rows_to_record_batch(&input, &self.table).unwrap();
791 if changed_schema {
792 let schema = ArrowSchema::new_with_metadata(
793 batch.schema().fields().clone(),
794 [("test_version".to_string(), "2".to_string())]
795 .into_iter()
796 .collect(),
797 );
798 batch = RecordBatch::try_new(Arc::new(schema), batch.columns().to_vec()).unwrap();
799 }
800 let limiter = RequestLimiter::try_new(1).unwrap();
801 let (response_tx, response_rx) = oneshot::channel();
802 self.tx
803 .send(WorkerCommand::Submit(PendingBatch {
804 table_info,
805 batch,
806 ctx: QueryContext::arc(),
807 response_tx,
808 _permit: limiter.acquire().await.unwrap(),
809 }))
810 .await
811 .unwrap();
812 drop(
814 timeout(WORKER_TIMEOUT, self.tx.reserve())
815 .await
816 .unwrap()
817 .unwrap(),
818 );
819 response_rx
820 }
821
822 async fn entered(&mut self) -> oneshot::Sender<()> {
823 timeout(WORKER_TIMEOUT, self.entered.recv())
824 .await
825 .expect("flush did not reach datanode")
826 .unwrap()
827 }
828
829 async fn stopped(&self) {
830 timeout(WORKER_TIMEOUT, async {
831 while !self.workers.is_empty().await {
832 tokio::task::yield_now().await;
833 }
834 })
835 .await
836 .expect("worker did not remove its registry entry");
837 }
838 }
839
840 #[tokio::test]
841 async fn test_same_schema_flushes_overlap() {
842 let mut worker = WorkerTest::new(1, 2).await;
843 let first = worker.submit(1, false).await;
844 let first_release = worker.entered().await;
845 let second = worker.submit(1, false).await;
846 let second_release = worker.entered().await;
847 first_release.send(()).unwrap();
849 second_release.send(()).unwrap();
850 timeout(WORKER_TIMEOUT, first)
851 .await
852 .unwrap()
853 .unwrap()
854 .unwrap();
855 timeout(WORKER_TIMEOUT, second)
856 .await
857 .unwrap()
858 .unwrap()
859 .unwrap();
860 worker.shutdown.send(()).unwrap();
861 worker.stopped().await;
862 }
863
864 #[tokio::test]
865 async fn test_schema_change_waits_for_prior_flush() {
866 let mut worker = WorkerTest::new(1, 2).await;
867 let first = worker.submit(1, false).await;
868 let first_release = worker.entered().await;
869 let second = worker.submit(1, true).await;
870 assert!(
872 timeout(Duration::from_millis(100), worker.entered.recv())
873 .await
874 .is_err()
875 );
876 first_release.send(()).unwrap();
877 let second_release = worker.entered().await;
878 second_release.send(()).unwrap();
879 timeout(WORKER_TIMEOUT, first)
880 .await
881 .unwrap()
882 .unwrap()
883 .unwrap();
884 timeout(WORKER_TIMEOUT, second)
885 .await
886 .unwrap()
887 .unwrap()
888 .unwrap();
889 worker.shutdown.send(()).unwrap();
890 worker.stopped().await;
891 }
892
893 #[tokio::test]
894 async fn test_table_identity_change_flushes_pending_rows_before_switching() {
895 for recreated in [false, true] {
896 let mut worker = WorkerTest::new(2, 2).await;
897 let first = worker.submit(1, false).await;
898 let mut changed = (*worker.table).clone();
899 if recreated {
900 changed.ident.table_id = 2;
901 } else {
902 changed.ident.version += 1;
903 }
904 let mut second = worker.submit_with_table(2, false, Arc::new(changed)).await;
905 let first_release = worker.entered().await;
907 assert!(
908 timeout(Duration::from_millis(100), worker.entered.recv())
909 .await
910 .is_err()
911 );
912 assert!(
913 timeout(Duration::from_millis(100), &mut second)
914 .await
915 .is_err()
916 );
917 first_release.send(()).unwrap();
918 timeout(WORKER_TIMEOUT, first)
919 .await
920 .unwrap()
921 .unwrap()
922 .unwrap();
923 if recreated {
924 assert!(
926 timeout(WORKER_TIMEOUT, second)
927 .await
928 .unwrap()
929 .unwrap()
930 .is_err()
931 );
932 } else {
933 let second_release = worker.entered().await;
934 second_release.send(()).unwrap();
935 timeout(WORKER_TIMEOUT, second)
936 .await
937 .unwrap()
938 .unwrap()
939 .unwrap();
940 }
941 worker.shutdown.send(()).unwrap();
942 worker.stopped().await;
943 }
944 }
945
946 #[tokio::test]
947 async fn test_shutdown_flushes_inline_without_joining_inflight() {
948 let mut worker = WorkerTest::new(2, 1).await;
949 let mut first = worker.submit(2, false).await;
950 let first_release = worker.entered().await;
951 let pending = worker.submit(1, false).await;
952 worker.shutdown.send(()).unwrap();
953 let pending_release = worker.entered().await;
955 pending_release.send(()).unwrap();
956 timeout(WORKER_TIMEOUT, pending)
957 .await
958 .unwrap()
959 .unwrap()
960 .unwrap();
961 worker.stopped().await;
962 assert!(matches!(
963 first.try_recv(),
964 Err(oneshot::error::TryRecvError::Empty)
965 ));
966 first_release.send(()).unwrap();
967 timeout(WORKER_TIMEOUT, first)
968 .await
969 .unwrap()
970 .unwrap()
971 .unwrap();
972 }
973}