Skip to main content

servers/batcher/table/
pending_batch.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
15use std::sync::Arc;
16
17use arrow::record_batch::RecordBatch;
18use operator::error::Error;
19use session::context::QueryContextRef;
20use table::metadata::TableInfoRef;
21use tokio::sync::{OwnedSemaphorePermit, oneshot};
22
23/// One complete input and its completion/admission ownership.
24pub(in crate::batcher::table) struct PendingBatch {
25    pub table_info: TableInfoRef,
26    pub batch: RecordBatch,
27    pub ctx: QueryContextRef,
28    pub response_tx: oneshot::Sender<Result<(), Arc<Error>>>,
29    pub _permit: Arc<OwnedSemaphorePermit>,
30}
31
32pub(in crate::batcher::table) fn notify_batches(
33    batches: Vec<PendingBatch>,
34    result: Result<(), Arc<Error>>,
35) {
36    for batch in batches {
37        let _ = batch.response_tx.send(result.clone());
38    }
39}
40
41#[cfg(test)]
42mod tests {
43    use arrow::array::Int32Array;
44    use arrow::datatypes::{DataType, Field, Schema};
45    use operator::error::UnexpectedSnafu;
46    use operator::test_util::new_test_table_info;
47    use session::context::QueryContext;
48    use tokio::sync::Semaphore;
49
50    use crate::batcher::table::pending_batch::*;
51
52    #[tokio::test]
53    async fn test_completion_fans_out_and_releases_admission() {
54        for fail in [false, true] {
55            let semaphore = Arc::new(Semaphore::new(3));
56            let mut batches = Vec::new();
57            let mut receivers = Vec::new();
58            for value in 0..3 {
59                let (response_tx, response_rx) = oneshot::channel();
60                let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]));
61                batches.push(PendingBatch {
62                    table_info: Arc::new(new_test_table_info(1, "t", [0].into_iter())),
63                    batch: RecordBatch::try_new(
64                        schema,
65                        vec![Arc::new(Int32Array::from(vec![value]))],
66                    )
67                    .unwrap(),
68                    ctx: QueryContext::arc(),
69                    response_tx,
70                    _permit: Arc::new(semaphore.clone().acquire_owned().await.unwrap()),
71                });
72                receivers.push(response_rx);
73            }
74            assert_eq!(0, semaphore.available_permits());
75            // A cancelled waiter must not prevent delivery to the others.
76            drop(receivers.pop());
77            let error = Arc::new(
78                UnexpectedSnafu {
79                    violated: "test flush failure".to_string(),
80                }
81                .build(),
82            );
83            notify_batches(batches, if fail { Err(error.clone()) } else { Ok(()) });
84            assert_eq!(3, semaphore.available_permits());
85            for receiver in receivers {
86                match receiver.await.unwrap() {
87                    Ok(()) => assert!(!fail),
88                    Err(actual) => {
89                        assert!(fail);
90                        assert!(Arc::ptr_eq(&error, &actual));
91                    }
92                }
93            }
94        }
95    }
96}