1use 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 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(®ion_number)
95 .copied()
96 .flatten()
97 .map(|value| PartitionExprVersion { value }),
98 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 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 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 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 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 let region_masks = partition_rule
194 .split_record_batch(&record_batch)
195 .context(error::SplitInsertSnafu)?;
196 partition_timer.observe_duration();
197
198 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 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(®ion_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(®ion_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 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 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}