1pub mod do_put;
16
17use std::collections::HashMap;
18use std::sync::Arc;
19
20use api::v1::{AffectedRows, FlightMetadata, Metrics};
21use arrow_flight::FlightData;
22use arrow_flight::utils::flight_data_to_arrow_batch;
23use common_base::bytes::Bytes;
24use common_recordbatch::DfRecordBatch;
25use datatypes::arrow;
26use datatypes::arrow::array::ArrayRef;
27use datatypes::arrow::buffer::Buffer;
28use datatypes::arrow::datatypes::{Schema as ArrowSchema, SchemaRef};
29use datatypes::arrow::error::ArrowError;
30use datatypes::arrow::ipc::{MessageHeader, convert, reader, root_as_message, writer};
31use flatbuffers::FlatBufferBuilder;
32use prost::Message;
33use prost::bytes::Bytes as ProstBytes;
34use snafu::{OptionExt, ResultExt};
35use vec1::{Vec1, vec1};
36
37use crate::error;
38use crate::error::{DecodeFlightDataSnafu, InvalidFlightDataSnafu, Result};
39
40pub const FLOW_EXTENSIONS_METADATA_KEY: &str = "x-greptime-flow-extensions";
42pub const SNAPSHOT_SEQS_METADATA_KEY: &str = "x-greptime-snapshot-seqs";
44
45#[derive(Debug, Clone)]
46pub enum FlightMessage {
47 Schema(SchemaRef),
48 RecordBatch(DfRecordBatch),
49 AffectedRows {
50 rows: usize,
51 metrics: Option<String>,
52 },
53 Metrics(String),
54}
55
56pub struct FlightEncoder {
57 write_options: writer::IpcWriteOptions,
58 data_gen: writer::IpcDataGenerator,
59 dictionary_tracker: writer::DictionaryTracker,
60}
61
62impl Default for FlightEncoder {
63 fn default() -> Self {
64 let write_options = writer::IpcWriteOptions::default()
65 .try_with_compression(Some(arrow::ipc::CompressionType::LZ4_FRAME))
66 .unwrap();
67
68 Self {
69 write_options,
70 data_gen: writer::IpcDataGenerator::default(),
71 dictionary_tracker: writer::DictionaryTracker::new(false),
72 }
73 }
74}
75
76impl FlightEncoder {
77 pub fn with_compression_disabled() -> Self {
79 let write_options = writer::IpcWriteOptions::default()
80 .try_with_compression(None)
81 .unwrap();
82
83 Self {
84 write_options,
85 data_gen: writer::IpcDataGenerator::default(),
86 dictionary_tracker: writer::DictionaryTracker::new(false),
87 }
88 }
89
90 pub fn encode_schema(&mut self, schema: &ArrowSchema) -> FlightData {
92 self.data_gen
93 .schema_to_bytes_with_dictionary_tracker(
94 schema,
95 &mut self.dictionary_tracker,
96 &self.write_options,
97 )
98 .into()
99 }
100
101 pub fn encode(&mut self, flight_message: FlightMessage) -> Vec1<FlightData> {
107 match flight_message {
108 FlightMessage::Schema(schema) => vec1![self.encode_schema(schema.as_ref())],
109 FlightMessage::RecordBatch(record_batch) => {
110 let (encoded_dictionaries, encoded_batch) = self
111 .data_gen
112 .encode(
113 &record_batch,
114 &mut self.dictionary_tracker,
115 &self.write_options,
116 &mut Default::default(),
117 )
118 .expect("DictionaryTracker configured above to not fail on replacement");
119
120 Vec1::from_vec_push(
121 encoded_dictionaries.into_iter().map(Into::into).collect(),
122 encoded_batch.into(),
123 )
124 }
125 FlightMessage::AffectedRows { rows, metrics } => {
126 let metadata = FlightMetadata {
127 affected_rows: Some(AffectedRows { value: rows as _ }),
128 metrics: metrics.map(|s| Metrics {
129 metrics: s.into_bytes(),
130 }),
131 }
132 .encode_to_vec();
133 vec1![FlightData {
134 flight_descriptor: None,
135 data_header: build_none_flight_msg().into(),
136 app_metadata: metadata.into(),
137 data_body: ProstBytes::default(),
138 }]
139 }
140 FlightMessage::Metrics(s) => {
141 let metadata = FlightMetadata {
142 affected_rows: None,
143 metrics: Some(Metrics {
144 metrics: s.as_bytes().to_vec(),
145 }),
146 }
147 .encode_to_vec();
148 vec1![FlightData {
149 flight_descriptor: None,
150 data_header: build_none_flight_msg().into(),
151 app_metadata: metadata.into(),
152 data_body: ProstBytes::default(),
153 }]
154 }
155 }
156 }
157}
158
159#[derive(Default)]
160pub struct FlightDecoder {
161 schema: Option<SchemaRef>,
162 schema_bytes: Option<bytes::Bytes>,
163 dictionaries_by_id: HashMap<i64, ArrayRef>,
164}
165
166impl FlightDecoder {
167 pub fn try_from_schema_bytes(schema_bytes: &bytes::Bytes) -> Result<Self> {
169 let arrow_schema = convert::try_schema_from_flatbuffer_bytes(&schema_bytes[..])
170 .context(error::ArrowSnafu)?;
171 Ok(Self {
172 schema: Some(Arc::new(arrow_schema)),
173 schema_bytes: Some(schema_bytes.clone()),
174 dictionaries_by_id: HashMap::new(),
175 })
176 }
177
178 pub fn try_decode_record_batch(
179 &mut self,
180 data_header: &bytes::Bytes,
181 data_body: &bytes::Bytes,
182 ) -> Result<DfRecordBatch> {
183 let schema = self
184 .schema
185 .as_ref()
186 .context(InvalidFlightDataSnafu {
187 reason: "Should have decoded schema first!",
188 })?
189 .clone();
190 let message = root_as_message(&data_header[..])
191 .map_err(|err| {
192 ArrowError::ParseError(format!("Unable to get root as message: {err:?}"))
193 })
194 .context(error::ArrowSnafu)?;
195 let result = message
196 .header_as_record_batch()
197 .ok_or_else(|| {
198 ArrowError::ParseError(
199 "Unable to convert flight data header to a record batch".to_string(),
200 )
201 })
202 .and_then(|batch| {
203 reader::read_record_batch(
204 &Buffer::from(data_body.as_ref()),
205 batch,
206 schema,
207 &HashMap::new(),
208 None,
209 &message.version(),
210 )
211 })
212 .context(error::ArrowSnafu)?;
213 Ok(result)
214 }
215
216 pub fn try_decode(&mut self, flight_data: &FlightData) -> Result<Option<FlightMessage>> {
223 let message = root_as_message(&flight_data.data_header).map_err(|e| {
224 InvalidFlightDataSnafu {
225 reason: e.to_string(),
226 }
227 .build()
228 })?;
229 match message.header_type() {
230 MessageHeader::NONE => {
231 let metadata = FlightMetadata::decode(flight_data.app_metadata.clone())
232 .context(DecodeFlightDataSnafu)?;
233 if let Some(AffectedRows { value }) = metadata.affected_rows {
234 return Ok(Some(FlightMessage::AffectedRows {
235 rows: value as _,
236 metrics: metadata
237 .metrics
238 .map(|m| String::from_utf8_lossy(&m.metrics).to_string()),
239 }));
240 }
241 if let Some(Metrics { metrics }) = metadata.metrics {
242 return Ok(Some(FlightMessage::Metrics(
243 String::from_utf8_lossy(&metrics).to_string(),
244 )));
245 }
246 InvalidFlightDataSnafu {
247 reason: "Expecting FlightMetadata have some meaningful content.",
248 }
249 .fail()
250 }
251 MessageHeader::Schema => {
252 let arrow_schema = Arc::new(ArrowSchema::try_from(flight_data).map_err(|e| {
253 InvalidFlightDataSnafu {
254 reason: e.to_string(),
255 }
256 .build()
257 })?);
258 self.schema = Some(arrow_schema.clone());
259 self.schema_bytes = Some(flight_data.data_header.clone());
260 Ok(Some(FlightMessage::Schema(arrow_schema)))
261 }
262 MessageHeader::RecordBatch => {
263 let schema = self.schema.clone().context(InvalidFlightDataSnafu {
264 reason: "Should have decoded schema first!",
265 })?;
266 let arrow_batch = flight_data_to_arrow_batch(
267 flight_data,
268 schema.clone(),
269 &self.dictionaries_by_id,
270 )
271 .map_err(|e| {
272 InvalidFlightDataSnafu {
273 reason: e.to_string(),
274 }
275 .build()
276 })?;
277 Ok(Some(FlightMessage::RecordBatch(arrow_batch)))
278 }
279 MessageHeader::DictionaryBatch => {
280 let dictionary_batch =
281 message
282 .header_as_dictionary_batch()
283 .context(InvalidFlightDataSnafu {
284 reason: "could not get dictionary batch from DictionaryBatch message",
285 })?;
286
287 let schema = self.schema.as_ref().context(InvalidFlightDataSnafu {
288 reason: "schema message is not present previously",
289 })?;
290
291 reader::read_dictionary(
292 &flight_data.data_body.clone().into(),
293 dictionary_batch,
294 schema,
295 &mut self.dictionaries_by_id,
296 &message.version(),
297 )
298 .context(error::ArrowSnafu)?;
299 Ok(None)
300 }
301 other => {
302 let name = other.variant_name().unwrap_or("UNKNOWN");
303 InvalidFlightDataSnafu {
304 reason: format!("Unsupported FlightData type: {name}"),
305 }
306 .fail()
307 }
308 }
309 }
310
311 pub fn schema(&self) -> Option<&SchemaRef> {
312 self.schema.as_ref()
313 }
314
315 pub fn schema_bytes(&self) -> Option<bytes::Bytes> {
316 self.schema_bytes.clone()
317 }
318}
319
320pub fn flight_messages_to_recordbatches(
321 messages: Vec<FlightMessage>,
322) -> Result<Vec<DfRecordBatch>> {
323 if messages.is_empty() {
324 Ok(vec![])
325 } else {
326 let mut recordbatches = Vec::with_capacity(messages.len() - 1);
327
328 match &messages[0] {
329 FlightMessage::Schema(_schema) => {}
330 _ => {
331 return InvalidFlightDataSnafu {
332 reason: "First Flight Message must be schema!",
333 }
334 .fail();
335 }
336 };
337
338 for message in messages.into_iter().skip(1) {
339 match message {
340 FlightMessage::RecordBatch(recordbatch) => recordbatches.push(recordbatch),
341 _ => {
342 return InvalidFlightDataSnafu {
343 reason: "Expect the following Flight Messages are all Recordbatches!",
344 }
345 .fail();
346 }
347 }
348 }
349
350 Ok(recordbatches)
351 }
352}
353
354fn build_none_flight_msg() -> Bytes {
355 let mut builder = FlatBufferBuilder::new();
356
357 let mut message = arrow::ipc::MessageBuilder::new(&mut builder);
358 message.add_version(arrow::ipc::MetadataVersion::V5);
359 message.add_header_type(MessageHeader::NONE);
360 message.add_bodyLength(0);
361
362 let data = message.finish();
363 builder.finish(data, None);
364
365 builder.finished_data().into()
366}
367
368#[cfg(test)]
369mod test {
370 use arrow_flight::utils::batches_to_flight_data;
371 use datatypes::arrow::array::{
372 DictionaryArray, Int32Array, ListArray, StringArray, UInt8Array, UInt32Array,
373 };
374 use datatypes::arrow::buffer::OffsetBuffer;
375 use datatypes::arrow::datatypes::{DataType, Field, Schema};
376
377 use super::*;
378 use crate::Error;
379
380 #[test]
381 fn test_try_decode() -> Result<()> {
382 let schema = Arc::new(ArrowSchema::new(vec![Field::new(
383 "n",
384 DataType::Int32,
385 true,
386 )]));
387
388 let batch1 = DfRecordBatch::try_new(
389 schema.clone(),
390 vec![Arc::new(Int32Array::from(vec![Some(1), None, Some(3)])) as _],
391 )
392 .unwrap();
393 let batch2 = DfRecordBatch::try_new(
394 schema.clone(),
395 vec![Arc::new(Int32Array::from(vec![None, Some(5)])) as _],
396 )
397 .unwrap();
398
399 let flight_data =
400 batches_to_flight_data(&schema, vec![batch1.clone(), batch2.clone()]).unwrap();
401 assert_eq!(flight_data.len(), 3);
402 let [d1, d2, d3] = flight_data.as_slice() else {
403 unreachable!()
404 };
405
406 let decoder = &mut FlightDecoder::default();
407 assert!(decoder.schema.is_none());
408
409 let result = decoder.try_decode(d2);
410 assert!(matches!(result, Err(Error::InvalidFlightData { .. })));
411 assert!(
412 result
413 .unwrap_err()
414 .to_string()
415 .contains("Should have decoded schema first!")
416 );
417
418 let message = decoder.try_decode(d1)?.unwrap();
419 assert!(matches!(message, FlightMessage::Schema(_)));
420 let FlightMessage::Schema(decoded_schema) = message else {
421 unreachable!()
422 };
423 assert_eq!(decoded_schema, schema);
424
425 let _ = decoder.schema.as_ref().unwrap();
426
427 let message = decoder.try_decode(d2)?.unwrap();
428 assert!(matches!(message, FlightMessage::RecordBatch(_)));
429 let FlightMessage::RecordBatch(actual_batch) = message else {
430 unreachable!()
431 };
432 assert_eq!(actual_batch, batch1);
433
434 let message = decoder.try_decode(d3)?.unwrap();
435 assert!(matches!(message, FlightMessage::RecordBatch(_)));
436 let FlightMessage::RecordBatch(actual_batch) = message else {
437 unreachable!()
438 };
439 assert_eq!(actual_batch, batch2);
440 Ok(())
441 }
442
443 #[test]
444 fn test_affected_rows_metrics_encode_decode() -> Result<()> {
445 let metrics = r#"{"region_watermarks":[{"region_id":42,"watermark":7}]}"#;
446 let mut encoder = FlightEncoder::default();
447 let encoded = encoder.encode(FlightMessage::AffectedRows {
448 rows: 3,
449 metrics: Some(metrics.to_string()),
450 });
451
452 assert_eq!(encoded.len(), 1);
453
454 let mut decoder = FlightDecoder::default();
455 let decoded = decoder.try_decode(encoded.first())?.unwrap();
456 let FlightMessage::AffectedRows {
457 rows,
458 metrics: decoded_metrics,
459 } = decoded
460 else {
461 unreachable!()
462 };
463 assert_eq!(rows, 3);
464 assert_eq!(decoded_metrics.as_deref(), Some(metrics));
465
466 let encoded = encoder.encode(FlightMessage::AffectedRows {
467 rows: 5,
468 metrics: None,
469 });
470 let decoded = decoder.try_decode(encoded.first())?.unwrap();
471 let FlightMessage::AffectedRows {
472 rows,
473 metrics: decoded_metrics,
474 } = decoded
475 else {
476 unreachable!()
477 };
478 assert_eq!(rows, 5);
479 assert!(decoded_metrics.is_none());
480
481 Ok(())
482 }
483
484 #[test]
485 fn test_flight_messages_to_recordbatches() {
486 let schema = Arc::new(Schema::new(vec![Field::new("m", DataType::Int32, true)]));
487 let batch1 = DfRecordBatch::try_new(
488 schema.clone(),
489 vec![Arc::new(Int32Array::from(vec![Some(2), None, Some(4)])) as _],
490 )
491 .unwrap();
492 let batch2 = DfRecordBatch::try_new(
493 schema.clone(),
494 vec![Arc::new(Int32Array::from(vec![None, Some(6)])) as _],
495 )
496 .unwrap();
497 let recordbatches = vec![batch1.clone(), batch2.clone()];
498
499 let m1 = FlightMessage::Schema(schema);
500 let m2 = FlightMessage::RecordBatch(batch1);
501 let m3 = FlightMessage::RecordBatch(batch2);
502
503 let result = flight_messages_to_recordbatches(vec![m2.clone(), m1.clone(), m3.clone()]);
504 assert!(matches!(result, Err(Error::InvalidFlightData { .. })));
505 assert!(
506 result
507 .unwrap_err()
508 .to_string()
509 .contains("First Flight Message must be schema!")
510 );
511
512 let result = flight_messages_to_recordbatches(vec![m1.clone(), m2.clone(), m1.clone()]);
513 assert!(matches!(result, Err(Error::InvalidFlightData { .. })));
514 assert!(
515 result
516 .unwrap_err()
517 .to_string()
518 .contains("Expect the following Flight Messages are all Recordbatches!")
519 );
520
521 let actual = flight_messages_to_recordbatches(vec![m1, m2, m3]).unwrap();
522 assert_eq!(actual, recordbatches);
523 }
524
525 #[test]
526 fn test_flight_encode_decode_with_dictionary_array() -> Result<()> {
527 let schema = Arc::new(Schema::new(vec![
528 Field::new("i", DataType::UInt8, true),
529 Field::new_dictionary("s", DataType::UInt32, DataType::Utf8, true),
530 ]));
531 let batch1 = DfRecordBatch::try_new(
532 schema.clone(),
533 vec![
534 Arc::new(UInt8Array::from_iter_values(vec![1, 2, 3])) as _,
535 Arc::new(DictionaryArray::new(
536 UInt32Array::from_value(0, 3),
537 Arc::new(StringArray::from_iter_values(["x"])),
538 )) as _,
539 ],
540 )
541 .unwrap();
542 let batch2 = DfRecordBatch::try_new(
543 schema.clone(),
544 vec![
545 Arc::new(UInt8Array::from_iter_values(vec![4, 5, 6, 7, 8])) as _,
546 Arc::new(DictionaryArray::new(
547 UInt32Array::from_iter_values([0, 1, 2, 2, 3]),
548 Arc::new(StringArray::from_iter_values(["h", "e", "l", "o"])),
549 )) as _,
550 ],
551 )
552 .unwrap();
553
554 let message_1 = FlightMessage::Schema(schema.clone());
555 let message_2 = FlightMessage::RecordBatch(batch1);
556 let message_3 = FlightMessage::RecordBatch(batch2);
557
558 let mut encoder = FlightEncoder::default();
559 let encoded_1 = encoder.encode(message_1);
560 let encoded_2 = encoder.encode(message_2);
561 let encoded_3 = encoder.encode(message_3);
562 assert_eq!(encoded_1.len(), 1);
564 assert_eq!(encoded_2.len(), 2);
567 assert_eq!(encoded_3.len(), 2);
568
569 let mut decoder = FlightDecoder::default();
570 let decoded_1 = decoder.try_decode(encoded_1.first())?;
571 let Some(FlightMessage::Schema(actual_schema)) = decoded_1 else {
572 unreachable!()
573 };
574 assert_eq!(actual_schema, schema);
575 let decoded_2 = decoder.try_decode(&encoded_2[0])?;
576 assert!(decoded_2.is_none());
578 let Some(FlightMessage::RecordBatch(decoded_2)) = decoder.try_decode(&encoded_2[1])? else {
579 unreachable!()
580 };
581 let decoded_3 = decoder.try_decode(&encoded_3[0])?;
582 assert!(decoded_3.is_none());
584 let Some(FlightMessage::RecordBatch(decoded_3)) = decoder.try_decode(&encoded_3[1])? else {
585 unreachable!()
586 };
587 let actual = arrow::util::pretty::pretty_format_batches(&[decoded_2, decoded_3])
588 .unwrap()
589 .to_string();
590 let expected = r"
591+---+---+
592| i | s |
593+---+---+
594| 1 | x |
595| 2 | x |
596| 3 | x |
597| 4 | h |
598| 5 | e |
599| 6 | l |
600| 7 | l |
601| 8 | o |
602+---+---+";
603 assert_eq!(actual, expected.trim());
604 Ok(())
605 }
606
607 #[test]
608 fn test_encode_schema_with_nested_dictionary_array() -> Result<()> {
609 let item = Arc::new(Field::new_dictionary(
610 "item",
611 DataType::UInt32,
612 DataType::Utf8,
613 true,
614 ));
615 let schema = Arc::new(Schema::new(vec![Field::new(
616 "tags",
617 DataType::List(item.clone()),
618 true,
619 )]));
620 let values = DictionaryArray::new(
621 UInt32Array::from_iter_values([0, 1, 0]),
622 Arc::new(StringArray::from_iter_values(["host-a", "host-b"])),
623 );
624 let list = ListArray::new(
625 item,
626 OffsetBuffer::from_lengths([2, 1]),
627 Arc::new(values),
628 None,
629 );
630 let batch = DfRecordBatch::try_new(schema.clone(), vec![Arc::new(list)]).unwrap();
631
632 let mut encoder = FlightEncoder::default();
633 let encoded_schema = encoder.encode_schema(schema.as_ref());
634 let encoded_batch = encoder.encode(FlightMessage::RecordBatch(batch.clone()));
635
636 let mut decoder = FlightDecoder::default();
637 assert!(matches!(
638 decoder.try_decode(&encoded_schema)?,
639 Some(FlightMessage::Schema(actual)) if actual == schema
640 ));
641 for data in encoded_batch.iter().take(encoded_batch.len() - 1) {
642 assert!(decoder.try_decode(data)?.is_none());
643 }
644 assert!(matches!(
645 decoder.try_decode(encoded_batch.last())?,
646 Some(FlightMessage::RecordBatch(actual)) if actual == batch
647 ));
648 Ok(())
649 }
650
651 #[test]
652 fn test_affected_rows_roundtrip_through_flight_codec() {
653 let mut encoder = FlightEncoder::default();
657 let mut decoder = FlightDecoder::default();
658
659 let encoded = encoder.encode(FlightMessage::AffectedRows {
661 rows: 7,
662 metrics: None,
663 });
664 let decoded = decoder.try_decode(encoded.first()).unwrap().unwrap();
665 assert!(matches!(
666 decoded,
667 FlightMessage::AffectedRows {
668 rows: 7,
669 metrics: None,
670 }
671 ));
672
673 let json = r#"{"region_watermarks":[{"region_id":1,"watermark":99}]}"#;
675 let encoded = encoder.encode(FlightMessage::AffectedRows {
676 rows: 42,
677 metrics: Some(json.to_string()),
678 });
679 let decoded = decoder.try_decode(encoded.first()).unwrap().unwrap();
680 assert!(matches!(
681 decoded,
682 FlightMessage::AffectedRows {
683 rows: 42,
684 metrics: Some(_),
685 }
686 ));
687 }
688
689 #[test]
692 fn test_old_affected_rows_format_decoded_by_new_code() {
693 use arrow_flight::FlightData;
694 use prost::bytes::Bytes as ProstBytes;
695
696 let old_wire_bytes = FlightData {
701 flight_descriptor: None,
702 data_header: build_none_flight_msg().into(),
703 app_metadata: FlightMetadata {
704 affected_rows: Some(AffectedRows { value: 99 }),
705 metrics: None, }
707 .encode_to_vec()
708 .into(),
709 data_body: ProstBytes::default(),
710 };
711
712 let mut decoder = FlightDecoder::default();
713 let decoded = decoder.try_decode(&old_wire_bytes).unwrap().unwrap();
714 assert!(matches!(
715 decoded,
716 FlightMessage::AffectedRows {
717 rows: 99,
718 metrics: None,
719 }
720 ));
721 }
722}