Skip to main content

servers/batcher/
table.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Table-scoped RecordBatch accumulation before bulk routing and encoding.
16
17mod 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
55// Match Prom batching: group by table name and keep WAL policies separate.
56fn 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
67/// Admission and worker lookup for table-scoped bulk writes.
68pub 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    /// Creates a shared timing batcher, rejecting disabled or unsupported limits.
83    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    /// Acquires once per original request, before submitting any of its tables.
156    /// Retain a clone in every submission until its flush has completed.
157    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    /// Acknowledges according to the global batching policy. Cancellation does
167    /// not retract an admitted submission. Flush failures may be partial writes.
168    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                // Hold the worker command to control completion independently of scheduling.
372                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                // Early acknowledgement must not release capacity before completion.
402                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        // Isolate the database counter across concurrent test cases.
473        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        // Keep timer/concurrency/workload fixed; only the supplied row threshold
510        // determines whether these two submissions share a bulk request.
511        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            // With capacity one, this proves the submitted command was dequeued.
813            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        // Both RPCs reached the datanode before either was allowed to complete.
848        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        // A second flush permit is free, but the schema barrier must withhold the RPC.
871        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            // The first batch is below max_rows but must be flushed on identity change.
906            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                // The recreated table has no route in this fixture; its error belongs only to it.
925                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        // The normal flush owns the sole permit. Shutdown must still start pending inline.
954        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}