Skip to main content

operator/
bulk_insert.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::collections::HashSet;
16
17use ahash::{HashMap, HashMapExt};
18use api::v1::flow::{DirtyWindowRequest, DirtyWindowRequests};
19use api::v1::region::{
20    BulkInsertRequest, RegionRequest, RegionRequestHeader, bulk_insert_request, region_request,
21};
22use api::v1::{ArrowIpc, PartitionExprVersion};
23use arrow::compute::filter_record_batch;
24use arrow::record_batch::RecordBatch;
25use bytes::Bytes;
26use common_base::AffectedRows;
27use common_error::ext::BoxedError;
28use common_grpc::FlightData;
29use common_grpc::flight::{FlightEncoder, FlightMessage, record_batch_to_ipc};
30use common_telemetry::error;
31use common_telemetry::tracing_context::TracingContext;
32use futures::future::{join_all, try_join_all};
33use meter_core::data::MeterRecord;
34use meter_macros::write_meter;
35use session::context::{Channel, QueryContextRef};
36use snafu::{ResultExt, ensure};
37use store_api::storage::RegionId;
38use table::TableRef;
39use table::metadata::TableInfoRef;
40
41use crate::error::Result;
42use crate::insert::Inserter;
43use crate::req_convert::insert::extract_timestamps;
44use crate::{error, metrics};
45
46impl Inserter {
47    /// Routes and writes a prepared table batch, returning the affected row count.
48    ///
49    /// Callers must exclude instant-TTL tables and handle Flow notifications after
50    /// successful writes. This execution helper does not perform either step.
51    pub async fn flush_bulk_batch(
52        &self,
53        table_info: TableInfoRef,
54        batch: RecordBatch,
55        ctx: QueryContextRef,
56    ) -> Result<AffectedRows> {
57        let (rule, versions) = self
58            .partition_manager
59            .find_table_partition_rule(&table_info)
60            .await
61            .context(error::InvalidPartitionSnafu)?;
62        let masks = rule
63            .split_record_batch(&batch)
64            .context(error::SplitInsertSnafu)?;
65        let mut writes = Vec::with_capacity(masks.len());
66        for (region_number, mask) in masks {
67            if mask.select_none() {
68                continue;
69            }
70            let region_id = RegionId::new(table_info.table_id(), region_number);
71            let selected = if mask.select_all() {
72                batch.clone()
73            } else {
74                filter_record_batch(&batch, mask.array()).context(error::ComputeArrowSnafu)?
75            };
76            let (schema, data_header, payload) = record_batch_to_ipc(selected)
77                .map_err(BoxedError::new)
78                .context(error::ExternalSnafu)?;
79            let peer = self
80                .partition_manager
81                .find_region_leader(region_id)
82                .await
83                .context(error::FindRegionLeaderSnafu)?;
84            let request = RegionRequest {
85                header: Some(RegionRequestHeader {
86                    dbname: ctx.get_db_string(),
87                    tracing_context: TracingContext::from_current_span().to_w3c(),
88                    ..Default::default()
89                }),
90                body: Some(region_request::Body::BulkInsert(BulkInsertRequest {
91                    skip_wal: ctx.skip_wal(),
92                    region_id: region_id.as_u64(),
93                    partition_expr_version: versions
94                        .get(&region_number)
95                        .copied()
96                        .flatten()
97                        .map(|value| PartitionExprVersion { value }),
98                    // Let the datanode revalidate against its current schema.
99                    aligned_schema_version: None,
100                    body: Some(bulk_insert_request::Body::ArrowIpc(ArrowIpc {
101                        schema,
102                        data_header,
103                        payload,
104                    })),
105                })),
106            };
107            writes.push((peer, request));
108        }
109        let results = join_all(writes.into_iter().map(|(peer, request)| async move {
110            self.node_manager
111                .datanode(&peer)
112                .await
113                .handle(request)
114                .await
115                .context(error::RequestInsertsSnafu)
116        }))
117        .await;
118        let affected_rows = results
119            .into_iter()
120            .map(|result| result.map(|response| response.affected_rows))
121            .sum::<Result<usize>>()?;
122        Ok(affected_rows)
123    }
124
125    /// Handle bulk insert request.
126    pub async fn handle_bulk_insert(
127        &self,
128        table: TableRef,
129        raw_flight_data: FlightData,
130        record_batch: RecordBatch,
131        schema_bytes: Bytes,
132        skip_wal: bool,
133        channel: Channel,
134    ) -> Result<AffectedRows> {
135        let table_info = table.table_info();
136        let table_id = table_info.table_id();
137        let db_name = table_info.get_db_string();
138
139        if record_batch.num_rows() == 0 {
140            return Ok(0);
141        }
142
143        // Bulk storage paths may discard unknown columns instead of rejecting them.
144        for field in record_batch.schema_ref().fields() {
145            ensure!(
146                table_info
147                    .meta
148                    .schema
149                    .column_schema_by_name(field.name())
150                    .is_some(),
151                error::InvalidInsertRequestSnafu {
152                    reason: format!(
153                        "Column '{}' not found in table '{}'",
154                        field.name(),
155                        table_info.full_table_name()
156                    ),
157                }
158            );
159        }
160
161        // The zero value is WCU, not bytes. Bulk writes have no WCU accounting;
162        // preserve that behavior while admitting their rows before dispatch.
163        write_meter!(MeterRecord::new(
164            table_info.catalog_name.clone(),
165            table_info.schema_name.clone(),
166            0,
167            record_batch.num_rows() as u64,
168            channel as u8,
169        ))
170        .await
171        .context(error::WriteRejectedSnafu)?;
172
173        let body_size = raw_flight_data.data_body.len();
174        // TODO(yingwen): Fill record batch impure default values. Note that we should override `raw_flight_data` if we have to fill defaults.
175        // notify flownode to update dirty timestamps if flow is configured.
176        self.maybe_update_flow_dirty_window(table_info.clone(), record_batch.clone());
177
178        metrics::BULK_REQUEST_MESSAGE_SIZE.observe(body_size as f64);
179        metrics::BULK_REQUEST_ROWS
180            .with_label_values(&["raw"])
181            .observe(record_batch.num_rows() as f64);
182
183        let partition_timer = metrics::HANDLE_BULK_INSERT_ELAPSED
184            .with_label_values(&["partition"])
185            .start_timer();
186        let (partition_rule, partition_versions) = self
187            .partition_manager
188            .find_table_partition_rule(&table_info)
189            .await
190            .context(error::InvalidPartitionSnafu)?;
191
192        // find partitions for each row in the record batch
193        let region_masks = partition_rule
194            .split_record_batch(&record_batch)
195            .context(error::SplitInsertSnafu)?;
196        partition_timer.observe_duration();
197
198        // fast path: only one region.
199        if region_masks.len() == 1 {
200            metrics::BULK_REQUEST_ROWS
201                .with_label_values(&["rows_per_region"])
202                .observe(record_batch.num_rows() as f64);
203
204            // SAFETY: region masks length checked
205            let (region_number, _) = region_masks.into_iter().next().unwrap();
206            let region_id = RegionId::new(table_id, region_number);
207            let partition_expr_version = partition_versions
208                .get(&region_number)
209                .copied()
210                .unwrap_or_default();
211            let datanode = self
212                .partition_manager
213                .find_region_leader(region_id)
214                .await
215                .context(error::FindRegionLeaderSnafu)?;
216
217            let request = RegionRequest {
218                header: Some(RegionRequestHeader {
219                    tracing_context: TracingContext::from_current_span().to_w3c(),
220                    ..Default::default()
221                }),
222                body: Some(region_request::Body::BulkInsert(BulkInsertRequest {
223                    skip_wal,
224                    region_id: region_id.as_u64(),
225                    partition_expr_version: partition_expr_version
226                        .map(|value| PartitionExprVersion { value }),
227                    aligned_schema_version: None,
228                    body: Some(bulk_insert_request::Body::ArrowIpc(ArrowIpc {
229                        schema: schema_bytes.clone(),
230                        data_header: raw_flight_data.data_header,
231                        payload: raw_flight_data.data_body,
232                    })),
233                })),
234            };
235
236            let _datanode_handle_timer = metrics::HANDLE_BULK_INSERT_ELAPSED
237                .with_label_values(&["datanode_handle"])
238                .start_timer();
239            let datanode = self.node_manager.datanode(&datanode).await;
240            let result = datanode
241                .handle(request)
242                .await
243                .context(error::RequestRegionSnafu)
244                .map(|r| r.affected_rows);
245            if let Ok(rows) = result {
246                crate::metrics::DIST_INGEST_ROW_COUNT
247                    .with_label_values(&[db_name.as_str()])
248                    .inc_by(rows as u64);
249            }
250            return result;
251        }
252
253        let mut mask_per_datanode = HashMap::with_capacity(region_masks.len());
254        for (region_number, mask) in region_masks {
255            let region_id = RegionId::new(table_id, region_number);
256            let datanode = self
257                .partition_manager
258                .find_region_leader(region_id)
259                .await
260                .context(error::FindRegionLeaderSnafu)?;
261            mask_per_datanode
262                .entry(datanode)
263                .or_insert_with(Vec::new)
264                .push((region_id, mask));
265        }
266
267        let wait_all_datanode_timer = metrics::HANDLE_BULK_INSERT_ELAPSED
268            .with_label_values(&["wait_all_datanode"])
269            .start_timer();
270
271        let mut handles = Vec::with_capacity(mask_per_datanode.len());
272
273        for (peer, masks) in mask_per_datanode {
274            for (region_id, mask) in masks {
275                if mask.select_none() {
276                    continue;
277                }
278                let partition_expr_version = partition_versions
279                    .get(&region_id.region_number())
280                    .copied()
281                    .unwrap_or_default();
282                let rb = record_batch.clone();
283                let schema_bytes = schema_bytes.clone();
284                let node_manager = self.node_manager.clone();
285                let peer = peer.clone();
286                let raw_header_and_data = if mask.select_all() {
287                    Some((
288                        raw_flight_data.data_header.clone(),
289                        raw_flight_data.data_body.clone(),
290                    ))
291                } else {
292                    None
293                };
294                let handle: common_runtime::JoinHandle<Result<api::region::RegionResponse>> =
295                    common_runtime::spawn_global(async move {
296                        let (header, payload) = if mask.select_all() {
297                            // SAFETY: raw data must be present, we can avoid re-encoding.
298                            raw_header_and_data.unwrap()
299                        } else {
300                            let filter_timer = metrics::HANDLE_BULK_INSERT_ELAPSED
301                                .with_label_values(&["filter"])
302                                .start_timer();
303                            let batch = filter_record_batch(&rb, mask.array())
304                                .context(error::ComputeArrowSnafu)?;
305                            filter_timer.observe_duration();
306                            metrics::BULK_REQUEST_ROWS
307                                .with_label_values(&["rows_per_region"])
308                                .observe(batch.num_rows() as f64);
309
310                            let encode_timer = metrics::HANDLE_BULK_INSERT_ELAPSED
311                                .with_label_values(&["encode"])
312                                .start_timer();
313                            let mut iter = FlightEncoder::default()
314                                .encode(FlightMessage::RecordBatch(batch))
315                                .into_iter();
316                            let Some(flight_data) = iter.next() else {
317                                // Safety: `iter` on a type of `Vec1`, which is guaranteed to have
318                                // at least one element.
319                                unreachable!()
320                            };
321                            ensure!(
322                                iter.next().is_none(),
323                                error::NotSupportedSnafu {
324                                    feat: "bulk insert RecordBatch with dictionary arrays",
325                                }
326                            );
327                            encode_timer.observe_duration();
328                            (flight_data.data_header, flight_data.data_body)
329                        };
330                        let _datanode_handle_timer = metrics::HANDLE_BULK_INSERT_ELAPSED
331                            .with_label_values(&["datanode_handle"])
332                            .start_timer();
333                        let request = RegionRequest {
334                            header: Some(RegionRequestHeader {
335                                tracing_context: TracingContext::from_current_span().to_w3c(),
336                                ..Default::default()
337                            }),
338                            body: Some(region_request::Body::BulkInsert(BulkInsertRequest {
339                                skip_wal,
340                                region_id: region_id.as_u64(),
341                                partition_expr_version: partition_expr_version
342                                    .map(|value| PartitionExprVersion { value }),
343                                aligned_schema_version: None,
344                                body: Some(bulk_insert_request::Body::ArrowIpc(ArrowIpc {
345                                    schema: schema_bytes,
346                                    data_header: header,
347                                    payload,
348                                })),
349                            })),
350                        };
351
352                        let datanode = node_manager.datanode(&peer).await;
353                        datanode
354                            .handle(request)
355                            .await
356                            .context(error::RequestRegionSnafu)
357                    });
358                handles.push(handle);
359            }
360        }
361
362        let region_responses = try_join_all(handles).await.context(error::JoinTaskSnafu)?;
363        wait_all_datanode_timer.observe_duration();
364        let mut rows_inserted: usize = 0;
365        for res in region_responses {
366            rows_inserted += res?.affected_rows;
367        }
368        crate::metrics::DIST_INGEST_ROW_COUNT
369            .with_label_values(&[db_name.as_str()])
370            .inc_by(rows_inserted as u64);
371        Ok(rows_inserted)
372    }
373
374    fn maybe_update_flow_dirty_window(&self, table_info: TableInfoRef, record_batch: RecordBatch) {
375        let table_id = table_info.table_id();
376        let table_flownode_set_cache = self.table_flownode_set_cache.clone();
377        let node_manager = self.node_manager.clone();
378        common_runtime::spawn_global(async move {
379            let result = table_flownode_set_cache
380                .get(table_id)
381                .await
382                .context(error::RequestInsertsSnafu);
383            let flownodes = match result {
384                Ok(flownodes) => flownodes.unwrap_or_default(),
385                Err(e) => {
386                    error!(e; "Failed to get flownodes for table id: {}", table_id);
387                    return;
388                }
389            };
390
391            let peers: HashSet<_> = flownodes.values().cloned().collect();
392            if peers.is_empty() {
393                return;
394            }
395
396            let Ok(timestamps) = extract_timestamps(
397                &record_batch,
398                &table_info
399                    .meta
400                    .schema
401                    .timestamp_column()
402                    .as_ref()
403                    .unwrap()
404                    .name,
405            )
406            .inspect_err(|e| {
407                error!(e; "Failed to extract timestamps from record batch");
408            }) else {
409                return;
410            };
411
412            for peer in peers {
413                let node_manager = node_manager.clone();
414                let timestamps = timestamps.clone();
415                common_runtime::spawn_global(async move {
416                    if let Err(e) = node_manager
417                        .flownode(&peer)
418                        .await
419                        .handle_mark_window_dirty(DirtyWindowRequests {
420                            requests: vec![DirtyWindowRequest {
421                                table_id,
422                                timestamps,
423                                time_ranges: vec![],
424                            }],
425                        })
426                        .await
427                        .context(error::RequestInsertsSnafu)
428                    {
429                        error!(e; "Failed to mark timestamps as dirty, table: {}", table_id);
430                    }
431                });
432            }
433        });
434    }
435}