1use std::net::IpAddr;
13use std::str::FromStr;
14use std::sync::Arc;
15
16use arrow::array::{
17 Array, BooleanArray, Decimal256Array, Float64Array, Int64Array, ListArray, StringArray,
18 TimestampNanosecondArray,
19};
20use arrow::array::{
21 ArrayRef, BooleanBuilder, Decimal256Builder, Float64Builder, Int64Builder, ListBuilder,
22 RecordBatch, StringBuilder, TimestampNanosecondBuilder,
23};
24use arrow::datatypes::i256;
25use chrono::DateTime;
26use num_bigint::BigUint;
27
28use wp_model_core::model::{
29 DataRecord, DataType, FValueStr, Field, FieldStorage, HexT, IpNetValue, Value,
30};
31
32use crate::error::WpArrowError;
33use crate::schema::{BIGINT_DECIMAL_PRECISION, FieldDef, WpDataType, to_arrow_schema};
34
35pub fn records_to_batch(
41 records: &[DataRecord],
42 field_defs: &[FieldDef],
43) -> Result<RecordBatch, WpArrowError> {
44 let schema = to_arrow_schema(field_defs)?;
45 let columns: Vec<ArrayRef> = field_defs
46 .iter()
47 .map(|fd| build_column(fd, records))
48 .collect::<Result<_, _>>()?;
49 RecordBatch::try_new(Arc::new(schema), columns)
50 .map_err(|e| WpArrowError::ArrowBuildError(e.to_string()))
51}
52
53pub fn batch_to_records(
58 batch: &RecordBatch,
59 field_defs: &[FieldDef],
60) -> Result<Vec<DataRecord>, WpArrowError> {
61 if field_defs.len() != batch.num_columns() {
62 return Err(WpArrowError::SchemaMismatch {
63 expected: field_defs.len(),
64 actual: batch.num_columns(),
65 });
66 }
67
68 let num_rows = batch.num_rows();
69 let mut records = Vec::with_capacity(num_rows);
70
71 for row_idx in 0..num_rows {
72 let mut items = Vec::with_capacity(field_defs.len());
73 for (col_idx, fd) in field_defs.iter().enumerate() {
74 let col = batch.column(col_idx);
75 if col.is_null(row_idx) {
76 continue;
77 }
78 let value = extract_value(col, row_idx, &fd.data_type, &fd.name)?;
79 let meta = wp_type_to_model_meta(&fd.data_type);
80 let field = Field::new(meta, fd.name.as_str(), value);
81 items.push(FieldStorage::from_owned(field));
82 }
83 let mut record = DataRecord::from(items);
84 record.id = row_idx as u64;
85 records.push(record);
86 }
87
88 Ok(records)
89}
90
91fn build_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
96 match &fd.data_type {
97 WpDataType::Chars | WpDataType::Ip | WpDataType::Hex => build_string_column(fd, records),
98 WpDataType::Digit => build_digit_column(fd, records),
99 WpDataType::BigInt => build_bigint_column(fd, records),
101 WpDataType::Float => build_float_column(fd, records),
102 WpDataType::Bool => build_bool_column(fd, records),
103 WpDataType::Time => build_time_column(fd, records),
104 WpDataType::Array(inner) => build_list_column(fd, records, inner),
105 }
106}
107
108fn build_string_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
109 let mut builder = StringBuilder::with_capacity(records.len(), records.len() * 32);
110 for rec in records {
111 match rec.get_value(&fd.name) {
112 Some(Value::Null) | None => {
113 handle_null(&mut builder, fd, |b| b.append_null())?;
114 }
115 Some(val) => {
116 let s = value_to_string(val, &fd.data_type, &fd.name)?;
117 builder.append_value(&s);
118 }
119 }
120 }
121 Ok(Arc::new(builder.finish()))
122}
123
124fn build_bigint_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
126 let mut builder = Decimal256Builder::with_capacity(records.len()).with_data_type(
127 arrow::datatypes::DataType::Decimal256(BIGINT_DECIMAL_PRECISION, 0),
128 );
129 for rec in records {
130 match rec.get_value(&fd.name) {
131 Some(Value::Null) | None => {
132 handle_null(&mut builder, fd, |b| b.append_null())?;
133 }
134 Some(Value::BigUint(v)) => {
135 let dec = i256::from_str(&v.to_string()).map_err(|err| {
136 WpArrowError::ValueConversionError {
137 field_name: fd.name.clone(),
138 expected: "BigInt(decimal)".to_string(),
139 actual: format!("{} (i256 parse: {err})", v),
140 }
141 })?;
142 builder.append_value(dec);
143 }
144 Some(other) => {
145 return Err(WpArrowError::ValueConversionError {
146 field_name: fd.name.clone(),
147 expected: "BigUint".to_string(),
148 actual: other.tag().to_string(),
149 });
150 }
151 }
152 }
153 Ok(Arc::new(builder.finish()))
154}
155
156fn build_digit_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
157 let mut builder = Int64Builder::with_capacity(records.len());
158 for rec in records {
159 match rec.get_value(&fd.name) {
160 Some(Value::Null) | None => {
161 handle_null(&mut builder, fd, |b| b.append_null())?;
162 }
163 Some(Value::Int(v)) => builder.append_value(*v),
164 Some(other) => {
165 return Err(WpArrowError::ValueConversionError {
166 field_name: fd.name.clone(),
167 expected: "Digit".to_string(),
168 actual: other.tag().to_string(),
169 });
170 }
171 }
172 }
173 Ok(Arc::new(builder.finish()))
174}
175
176fn build_float_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
177 let mut builder = Float64Builder::with_capacity(records.len());
178 for rec in records {
179 match rec.get_value(&fd.name) {
180 Some(Value::Null) | None => {
181 handle_null(&mut builder, fd, |b| b.append_null())?;
182 }
183 Some(Value::Float(v)) => builder.append_value(*v),
184 Some(other) => {
185 return Err(WpArrowError::ValueConversionError {
186 field_name: fd.name.clone(),
187 expected: "Float".to_string(),
188 actual: other.tag().to_string(),
189 });
190 }
191 }
192 }
193 Ok(Arc::new(builder.finish()))
194}
195
196fn build_bool_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
197 let mut builder = BooleanBuilder::with_capacity(records.len());
198 for rec in records {
199 match rec.get_value(&fd.name) {
200 Some(Value::Null) | None => {
201 handle_null(&mut builder, fd, |b| b.append_null())?;
202 }
203 Some(Value::Bool(v)) => builder.append_value(*v),
204 Some(other) => {
205 return Err(WpArrowError::ValueConversionError {
206 field_name: fd.name.clone(),
207 expected: "Bool".to_string(),
208 actual: other.tag().to_string(),
209 });
210 }
211 }
212 }
213 Ok(Arc::new(builder.finish()))
214}
215
216fn build_time_column(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
217 let mut builder = TimestampNanosecondBuilder::with_capacity(records.len());
218 for rec in records {
219 match rec.get_value(&fd.name) {
220 Some(Value::Null) | None => {
221 handle_null(&mut builder, fd, |b| b.append_null())?;
222 }
223 Some(Value::Time(ndt)) => {
224 let nanos = ndt.and_utc().timestamp_nanos_opt().ok_or_else(|| {
225 WpArrowError::TimestampOverflow {
226 field_name: fd.name.clone(),
227 }
228 })?;
229 builder.append_value(nanos);
230 }
231 Some(other) => {
232 return Err(WpArrowError::ValueConversionError {
233 field_name: fd.name.clone(),
234 expected: "Time".to_string(),
235 actual: other.tag().to_string(),
236 });
237 }
238 }
239 }
240 Ok(Arc::new(builder.finish()))
241}
242
243fn build_list_column(
244 fd: &FieldDef,
245 records: &[DataRecord],
246 inner_type: &WpDataType,
247) -> Result<ArrayRef, WpArrowError> {
248 match inner_type {
249 WpDataType::Chars | WpDataType::Ip | WpDataType::Hex => {
250 build_list_string(fd, records, inner_type)
251 }
252 WpDataType::Digit => build_list_digit(fd, records),
253 WpDataType::BigInt => build_list_bigint(fd, records),
254 WpDataType::Float => build_list_float(fd, records),
255 WpDataType::Bool => build_list_bool(fd, records),
256 WpDataType::Time => build_list_time(fd, records),
257 WpDataType::Array(_) => Err(WpArrowError::UnsupportedDataType(
258 "nested array<array<...>> not supported".to_string(),
259 )),
260 }
261}
262
263fn build_list_string(
264 fd: &FieldDef,
265 records: &[DataRecord],
266 inner_type: &WpDataType,
267) -> Result<ArrayRef, WpArrowError> {
268 let mut builder = ListBuilder::new(StringBuilder::new());
269 for rec in records {
270 match rec.get_value(&fd.name) {
271 Some(Value::Null) | None => {
272 handle_null(&mut builder, fd, |b| b.append_null())?;
273 }
274 Some(Value::Array(items)) => {
275 for item in items {
276 let val = item.get_value();
277 if matches!(val, Value::Null) {
278 builder.values().append_null();
279 } else {
280 let s = value_to_string(val, inner_type, &fd.name)?;
281 builder.values().append_value(&s);
282 }
283 }
284 builder.append(true);
285 }
286 Some(other) => {
287 return Err(WpArrowError::ValueConversionError {
288 field_name: fd.name.clone(),
289 expected: "Array".to_string(),
290 actual: other.tag().to_string(),
291 });
292 }
293 }
294 }
295 Ok(Arc::new(builder.finish()))
296}
297
298fn build_list_bigint(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
300 let mut builder = ListBuilder::new(Decimal256Builder::with_capacity(0).with_data_type(
301 arrow::datatypes::DataType::Decimal256(BIGINT_DECIMAL_PRECISION, 0),
302 ));
303 for rec in records {
304 match rec.get_value(&fd.name) {
305 Some(Value::Null) | None => {
306 handle_null(&mut builder, fd, |b| b.append_null())?;
307 }
308 Some(Value::Array(items)) => {
309 for item in items {
310 match item.get_value() {
311 Value::BigUint(v) => {
312 let dec = i256::from_str(&v.to_string()).map_err(|err| {
313 WpArrowError::ValueConversionError {
314 field_name: fd.name.clone(),
315 expected: "BigInt(decimal)".to_string(),
316 actual: format!("{} (i256 parse: {err})", v),
317 }
318 })?;
319 builder.values().append_value(dec);
320 }
321 Value::Null => builder.values().append_null(),
322 other => {
323 return Err(WpArrowError::ValueConversionError {
324 field_name: fd.name.clone(),
325 expected: "BigUint".to_string(),
326 actual: other.tag().to_string(),
327 });
328 }
329 }
330 }
331 builder.append(true);
332 }
333 Some(other) => {
334 return Err(WpArrowError::ValueConversionError {
335 field_name: fd.name.clone(),
336 expected: "Array".to_string(),
337 actual: other.tag().to_string(),
338 });
339 }
340 }
341 }
342 Ok(Arc::new(builder.finish()))
343}
344
345fn build_list_digit(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
346 let mut builder = ListBuilder::new(Int64Builder::new());
347 for rec in records {
348 match rec.get_value(&fd.name) {
349 Some(Value::Null) | None => {
350 handle_null(&mut builder, fd, |b| b.append_null())?;
351 }
352 Some(Value::Array(items)) => {
353 for item in items {
354 match item.get_value() {
355 Value::Int(v) => builder.values().append_value(*v),
356 Value::Null => builder.values().append_null(),
357 other => {
358 return Err(WpArrowError::ValueConversionError {
359 field_name: fd.name.clone(),
360 expected: "Digit".to_string(),
361 actual: other.tag().to_string(),
362 });
363 }
364 }
365 }
366 builder.append(true);
367 }
368 Some(other) => {
369 return Err(WpArrowError::ValueConversionError {
370 field_name: fd.name.clone(),
371 expected: "Array".to_string(),
372 actual: other.tag().to_string(),
373 });
374 }
375 }
376 }
377 Ok(Arc::new(builder.finish()))
378}
379
380fn build_list_float(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
381 let mut builder = ListBuilder::new(Float64Builder::new());
382 for rec in records {
383 match rec.get_value(&fd.name) {
384 Some(Value::Null) | None => {
385 handle_null(&mut builder, fd, |b| b.append_null())?;
386 }
387 Some(Value::Array(items)) => {
388 for item in items {
389 match item.get_value() {
390 Value::Float(v) => builder.values().append_value(*v),
391 Value::Null => builder.values().append_null(),
392 other => {
393 return Err(WpArrowError::ValueConversionError {
394 field_name: fd.name.clone(),
395 expected: "Float".to_string(),
396 actual: other.tag().to_string(),
397 });
398 }
399 }
400 }
401 builder.append(true);
402 }
403 Some(other) => {
404 return Err(WpArrowError::ValueConversionError {
405 field_name: fd.name.clone(),
406 expected: "Array".to_string(),
407 actual: other.tag().to_string(),
408 });
409 }
410 }
411 }
412 Ok(Arc::new(builder.finish()))
413}
414
415fn build_list_bool(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
416 let mut builder = ListBuilder::new(BooleanBuilder::new());
417 for rec in records {
418 match rec.get_value(&fd.name) {
419 Some(Value::Null) | None => {
420 handle_null(&mut builder, fd, |b| b.append_null())?;
421 }
422 Some(Value::Array(items)) => {
423 for item in items {
424 match item.get_value() {
425 Value::Bool(v) => builder.values().append_value(*v),
426 Value::Null => builder.values().append_null(),
427 other => {
428 return Err(WpArrowError::ValueConversionError {
429 field_name: fd.name.clone(),
430 expected: "Bool".to_string(),
431 actual: other.tag().to_string(),
432 });
433 }
434 }
435 }
436 builder.append(true);
437 }
438 Some(other) => {
439 return Err(WpArrowError::ValueConversionError {
440 field_name: fd.name.clone(),
441 expected: "Array".to_string(),
442 actual: other.tag().to_string(),
443 });
444 }
445 }
446 }
447 Ok(Arc::new(builder.finish()))
448}
449
450fn build_list_time(fd: &FieldDef, records: &[DataRecord]) -> Result<ArrayRef, WpArrowError> {
451 let mut builder = ListBuilder::new(TimestampNanosecondBuilder::new());
452 for rec in records {
453 match rec.get_value(&fd.name) {
454 Some(Value::Null) | None => {
455 handle_null(&mut builder, fd, |b| b.append_null())?;
456 }
457 Some(Value::Array(items)) => {
458 for item in items {
459 match item.get_value() {
460 Value::Time(ndt) => {
461 let nanos = ndt.and_utc().timestamp_nanos_opt().ok_or_else(|| {
462 WpArrowError::TimestampOverflow {
463 field_name: fd.name.clone(),
464 }
465 })?;
466 builder.values().append_value(nanos);
467 }
468 Value::Null => builder.values().append_null(),
469 other => {
470 return Err(WpArrowError::ValueConversionError {
471 field_name: fd.name.clone(),
472 expected: "Time".to_string(),
473 actual: other.tag().to_string(),
474 });
475 }
476 }
477 }
478 builder.append(true);
479 }
480 Some(other) => {
481 return Err(WpArrowError::ValueConversionError {
482 field_name: fd.name.clone(),
483 expected: "Array".to_string(),
484 actual: other.tag().to_string(),
485 });
486 }
487 }
488 }
489 Ok(Arc::new(builder.finish()))
490}
491
492fn value_to_string(
494 val: &Value,
495 wp_type: &WpDataType,
496 field_name: &str,
497) -> Result<String, WpArrowError> {
498 match (wp_type, val) {
499 (WpDataType::Chars, Value::Chars(s)) => Ok(s.to_string()),
501 (WpDataType::Chars, Value::Domain(d)) => Ok(d.to_string()),
502 (WpDataType::Chars, Value::Url(u)) => Ok(u.to_string()),
503 (WpDataType::Chars, Value::Email(e)) => Ok(e.to_string()),
504 (WpDataType::Ip, Value::IpAddr(ip)) => Ok(ip.to_string()),
506 (WpDataType::Ip, Value::IpNet(net)) => Ok(net.to_string()),
507 (WpDataType::Ip, Value::Chars(s)) => Ok(s.to_string()),
508 (WpDataType::Hex, Value::Hex(h)) => Ok(format!("{:#X}", h.0)),
510 _ => Err(WpArrowError::ValueConversionError {
511 field_name: field_name.to_string(),
512 expected: format!("{:?}", wp_type),
513 actual: val.tag().to_string(),
514 }),
515 }
516}
517
518fn handle_null<B, F>(builder: &mut B, fd: &FieldDef, append_null: F) -> Result<(), WpArrowError>
520where
521 F: FnOnce(&mut B),
522{
523 if fd.nullable {
524 append_null(builder);
525 Ok(())
526 } else {
527 Err(WpArrowError::MissingRequiredField {
528 field_name: fd.name.clone(),
529 })
530 }
531}
532
533fn extract_value(
539 col: &ArrayRef,
540 row_idx: usize,
541 wp_type: &WpDataType,
542 field_name: &str,
543) -> Result<Value, WpArrowError> {
544 match wp_type {
545 WpDataType::Chars => {
546 let arr = col
547 .as_any()
548 .downcast_ref::<StringArray>()
549 .ok_or_else(|| WpArrowError::ArrowBuildError("expected StringArray".to_string()))?;
550 Ok(Value::Chars(FValueStr::from(arr.value(row_idx))))
551 }
552 WpDataType::Digit => {
553 let arr = col
554 .as_any()
555 .downcast_ref::<Int64Array>()
556 .ok_or_else(|| WpArrowError::ArrowBuildError("expected Int64Array".to_string()))?;
557 Ok(Value::Int(arr.value(row_idx)))
558 }
559 WpDataType::BigInt => {
560 let arr = col
561 .as_any()
562 .downcast_ref::<Decimal256Array>()
563 .ok_or_else(|| {
564 WpArrowError::ArrowBuildError("expected Decimal256Array".to_string())
565 })?;
566 let v = arr.value(row_idx);
567 match BigUint::from_str(&v.to_string()) {
569 Ok(v) => Ok(Value::BigUint(v)),
570 Err(err) => Err(WpArrowError::ValueConversionError {
571 field_name: field_name.to_string(),
572 expected: "BigInt(decimal)".to_string(),
573 actual: format!("{v} (parse: {err})"),
574 }),
575 }
576 }
577 WpDataType::Float => {
578 let arr = col.as_any().downcast_ref::<Float64Array>().ok_or_else(|| {
579 WpArrowError::ArrowBuildError("expected Float64Array".to_string())
580 })?;
581 Ok(Value::Float(arr.value(row_idx)))
582 }
583 WpDataType::Bool => {
584 let arr = col.as_any().downcast_ref::<BooleanArray>().ok_or_else(|| {
585 WpArrowError::ArrowBuildError("expected BooleanArray".to_string())
586 })?;
587 Ok(Value::Bool(arr.value(row_idx)))
588 }
589 WpDataType::Time => {
590 let arr = col
591 .as_any()
592 .downcast_ref::<TimestampNanosecondArray>()
593 .ok_or_else(|| {
594 WpArrowError::ArrowBuildError("expected TimestampNanosecondArray".to_string())
595 })?;
596 let nanos = arr.value(row_idx);
597 let ndt = DateTime::from_timestamp_nanos(nanos).naive_utc();
598 Ok(Value::Time(ndt))
599 }
600 WpDataType::Ip => {
601 let arr = col
602 .as_any()
603 .downcast_ref::<StringArray>()
604 .ok_or_else(|| WpArrowError::ArrowBuildError("expected StringArray".to_string()))?;
605 let s = arr.value(row_idx);
606 Ok(parse_ip_value(s, field_name)?)
607 }
608 WpDataType::Hex => {
609 let arr = col
610 .as_any()
611 .downcast_ref::<StringArray>()
612 .ok_or_else(|| WpArrowError::ArrowBuildError("expected StringArray".to_string()))?;
613 let s = arr.value(row_idx);
614 Ok(parse_hex_value(s, field_name)?)
615 }
616 WpDataType::Array(inner) => {
617 let arr = col
618 .as_any()
619 .downcast_ref::<ListArray>()
620 .ok_or_else(|| WpArrowError::ArrowBuildError("expected ListArray".to_string()))?;
621 let inner_arr = arr.value(row_idx);
622 let inner_meta = wp_type_to_model_meta(inner);
623 let mut items = Vec::new();
624 for i in 0..inner_arr.len() {
625 if inner_arr.is_null(i) {
626 items.push(FieldStorage::from_owned(Field::new(
627 inner_meta.clone(),
628 "item",
629 Value::Null,
630 )));
631 } else {
632 let val = extract_value(&inner_arr, i, inner, field_name)?;
633 items.push(FieldStorage::from_owned(Field::new(
634 inner_meta.clone(),
635 "item",
636 val,
637 )));
638 }
639 }
640 Ok(Value::Array(items))
641 }
642 }
643}
644
645fn parse_ip_value(s: &str, field_name: &str) -> Result<Value, WpArrowError> {
647 if s.contains('/') {
648 let parts: Vec<&str> = s.splitn(2, '/').collect();
650 let addr: IpAddr = parts[0].parse().map_err(|e| WpArrowError::ParseError {
651 field_name: field_name.to_string(),
652 detail: format!("invalid IP address: {e}"),
653 })?;
654 let prefix: u8 = parts[1].parse().map_err(|e| WpArrowError::ParseError {
655 field_name: field_name.to_string(),
656 detail: format!("invalid prefix length: {e}"),
657 })?;
658 let net = IpNetValue::new(addr, prefix).ok_or_else(|| WpArrowError::ParseError {
659 field_name: field_name.to_string(),
660 detail: format!("invalid prefix length {prefix} for {addr}"),
661 })?;
662 Ok(Value::IpNet(net))
663 } else {
664 let addr: IpAddr = s.parse().map_err(|e| WpArrowError::ParseError {
666 field_name: field_name.to_string(),
667 detail: format!("invalid IP address: {e}"),
668 })?;
669 Ok(Value::IpAddr(addr))
670 }
671}
672
673fn parse_hex_value(s: &str, field_name: &str) -> Result<Value, WpArrowError> {
675 let hex_str = s
676 .strip_prefix("0x")
677 .or_else(|| s.strip_prefix("0X"))
678 .unwrap_or(s);
679 let v = u128::from_str_radix(hex_str, 16).map_err(|e| WpArrowError::ParseError {
680 field_name: field_name.to_string(),
681 detail: format!("invalid hex: {e}"),
682 })?;
683 Ok(Value::Hex(HexT(v)))
684}
685
686fn wp_type_to_model_meta(wp_type: &WpDataType) -> DataType {
688 match wp_type {
689 WpDataType::Chars => DataType::Chars,
690 WpDataType::Digit => DataType::Int,
691 WpDataType::BigInt => DataType::BigInt,
692 WpDataType::Float => DataType::Float,
693 WpDataType::Bool => DataType::Bool,
694 WpDataType::Time => DataType::Time,
695 WpDataType::Ip => DataType::IP,
696 WpDataType::Hex => DataType::Hex,
697 WpDataType::Array(inner) => {
698 let inner_name = match inner.as_ref() {
699 WpDataType::Chars => "chars",
700 WpDataType::Digit => "digit",
701 WpDataType::BigInt => "bigint",
702 WpDataType::Float => "float",
703 WpDataType::Bool => "bool",
704 WpDataType::Time => "time",
705 WpDataType::Ip => "ip",
706 WpDataType::Hex => "hex",
707 WpDataType::Array(_) => "array",
708 };
709 DataType::Array(inner_name.into())
710 }
711 }
712}
713
714#[cfg(test)]
715mod tests {
716 use super::*;
717 use crate::schema::{FieldDef, WpDataType};
718 use arrow::array::AsArray;
719 use chrono::NaiveDateTime;
720 use std::net::{IpAddr, Ipv4Addr};
721 use wp_model_core::model::{DataField, DataRecord, Field, Value};
722
723 fn make_record(fields: Vec<DataField>) -> DataRecord {
725 DataRecord::from(fields)
726 }
727
728 #[test]
733 fn r2b_basic_types() {
734 let fds = vec![
735 FieldDef::new("name", WpDataType::Chars),
736 FieldDef::new("count", WpDataType::Digit),
737 FieldDef::new("ratio", WpDataType::Float),
738 FieldDef::new("active", WpDataType::Bool),
739 ];
740 let records = vec![
741 make_record(vec![
742 Field::from_chars("name", "Alice"),
743 Field::from_int("count", 10),
744 Field::from_float("ratio", 1.5),
745 Field::from_bool("active", true),
746 ]),
747 make_record(vec![
748 Field::from_chars("name", "Bob"),
749 Field::from_int("count", 20),
750 Field::from_float("ratio", 2.5),
751 Field::from_bool("active", false),
752 ]),
753 ];
754
755 let batch = records_to_batch(&records, &fds).unwrap();
756 assert_eq!(batch.num_columns(), 4);
757 assert_eq!(batch.num_rows(), 2);
758
759 let names = batch.column(0).as_string::<i32>();
760 assert_eq!(names.value(0), "Alice");
761 assert_eq!(names.value(1), "Bob");
762
763 let counts = batch
764 .column(1)
765 .as_primitive::<arrow::datatypes::Int64Type>();
766 assert_eq!(counts.value(0), 10);
767 assert_eq!(counts.value(1), 20);
768
769 let ratios = batch
770 .column(2)
771 .as_primitive::<arrow::datatypes::Float64Type>();
772 assert!((ratios.value(0) - 1.5).abs() < f64::EPSILON);
773
774 let actives = batch.column(3).as_boolean();
775 assert!(actives.value(0));
776 assert!(!actives.value(1));
777 }
778
779 #[test]
780 fn r2b_time_field() {
781 let fds = vec![FieldDef::new("ts", WpDataType::Time)];
782 let ndt =
783 NaiveDateTime::parse_from_str("2024-06-15 12:30:00", "%Y-%m-%d %H:%M:%S").unwrap();
784 let records = vec![make_record(vec![Field::from_time("ts", ndt)])];
785
786 let batch = records_to_batch(&records, &fds).unwrap();
787 let arr = batch
788 .column(0)
789 .as_any()
790 .downcast_ref::<TimestampNanosecondArray>()
791 .unwrap();
792 let expected_nanos = ndt.and_utc().timestamp_nanos_opt().unwrap();
793 assert_eq!(arr.value(0), expected_nanos);
794 }
795
796 #[test]
797 fn r2b_ip_field() {
798 let fds = vec![FieldDef::new("addr", WpDataType::Ip)];
799 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
800 let net = IpNetValue::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 0)), 8).unwrap();
801 let records = vec![
802 make_record(vec![Field::from_ip("addr", ip)]),
803 make_record(vec![Field::new(DataType::IP, "addr", Value::IpNet(net))]),
804 ];
805
806 let batch = records_to_batch(&records, &fds).unwrap();
807 let arr = batch.column(0).as_string::<i32>();
808 assert_eq!(arr.value(0), "192.168.1.1");
809 assert_eq!(arr.value(1), "10.0.0.0/8");
810 }
811
812 #[test]
813 fn r2b_hex_field() {
814 let fds = vec![FieldDef::new("color", WpDataType::Hex)];
815 let records = vec![make_record(vec![Field::from_hex("color", HexT(255))])];
816
817 let batch = records_to_batch(&records, &fds).unwrap();
818 let arr = batch.column(0).as_string::<i32>();
819 assert_eq!(arr.value(0), "0xFF");
820 }
821
822 #[test]
823 fn r2b_nullable_missing() {
824 let fds = vec![
825 FieldDef::new("name", WpDataType::Chars),
826 FieldDef::new("opt", WpDataType::Digit), ];
828 let records = vec![
829 make_record(vec![Field::from_chars("name", "Alice")]),
830 ];
832
833 let batch = records_to_batch(&records, &fds).unwrap();
834 assert!(batch.column(1).is_null(0));
835 }
836
837 #[test]
838 fn r2b_required_missing() {
839 let fds = vec![FieldDef::new("required", WpDataType::Digit).with_nullable(false)];
840 let records = vec![make_record(vec![Field::from_chars("other", "x")])];
841
842 let err = records_to_batch(&records, &fds).unwrap_err();
843 assert!(matches!(err, WpArrowError::MissingRequiredField { .. }));
844 }
845
846 #[test]
847 fn r2b_null_value_nullable() {
848 let fds = vec![FieldDef::new("val", WpDataType::Chars)];
849 let records = vec![make_record(vec![Field::new(
850 DataType::Chars,
851 "val",
852 Value::Null,
853 )])];
854
855 let batch = records_to_batch(&records, &fds).unwrap();
856 assert!(batch.column(0).is_null(0));
857 }
858
859 #[test]
860 fn r2b_empty_records() {
861 let fds = vec![FieldDef::new("x", WpDataType::Digit)];
862 let records: Vec<DataRecord> = vec![];
863
864 let batch = records_to_batch(&records, &fds).unwrap();
865 assert_eq!(batch.num_rows(), 0);
866 assert_eq!(batch.num_columns(), 1);
867 }
868
869 #[test]
870 fn r2b_extra_fields_ignored() {
871 let fds = vec![FieldDef::new("a", WpDataType::Digit)];
872 let records = vec![make_record(vec![
873 Field::from_int("a", 1),
874 Field::from_chars("extra", "ignored"),
875 ])];
876
877 let batch = records_to_batch(&records, &fds).unwrap();
878 assert_eq!(batch.num_columns(), 1);
879 let arr = batch
880 .column(0)
881 .as_primitive::<arrow::datatypes::Int64Type>();
882 assert_eq!(arr.value(0), 1);
883 }
884
885 #[test]
886 fn r2b_array_field() {
887 let fds = vec![FieldDef::new(
888 "tags",
889 WpDataType::Array(Box::new(WpDataType::Digit)),
890 )];
891 let items: Vec<DataField> = vec![Field::from_int("item", 10), Field::from_int("item", 20)];
892 let records = vec![make_record(vec![Field::from_arr("tags", items)])];
893
894 let batch = records_to_batch(&records, &fds).unwrap();
895 let arr = batch
896 .column(0)
897 .as_any()
898 .downcast_ref::<ListArray>()
899 .unwrap();
900 assert_eq!(arr.len(), 1);
901 let inner = arr.value(0);
902 let inner_vals = inner.as_any().downcast_ref::<Int64Array>().unwrap();
903 assert_eq!(inner_vals.value(0), 10);
904 assert_eq!(inner_vals.value(1), 20);
905 }
906
907 #[test]
908 fn r2b_type_mismatch() {
909 let fds = vec![FieldDef::new("num", WpDataType::Digit)];
910 let records = vec![make_record(vec![Field::from_chars("num", "not_a_number")])];
911
912 let err = records_to_batch(&records, &fds).unwrap_err();
913 assert!(matches!(err, WpArrowError::ValueConversionError { .. }));
914 }
915
916 #[test]
917 fn r2b_large_batch() {
918 let fds = vec![
919 FieldDef::new("id", WpDataType::Digit),
920 FieldDef::new("name", WpDataType::Chars),
921 ];
922 let records: Vec<DataRecord> = (0..10000)
923 .map(|i| {
924 make_record(vec![
925 Field::from_int("id", i),
926 Field::from_chars("name", format!("row_{i}")),
927 ])
928 })
929 .collect();
930
931 let batch = records_to_batch(&records, &fds).unwrap();
932 assert_eq!(batch.num_rows(), 10000);
933
934 let ids = batch
935 .column(0)
936 .as_primitive::<arrow::datatypes::Int64Type>();
937 assert_eq!(ids.value(0), 0);
938 assert_eq!(ids.value(9999), 9999);
939 }
940
941 #[test]
946 fn b2r_basic_types() {
947 let fds = vec![
948 FieldDef::new("name", WpDataType::Chars),
949 FieldDef::new("count", WpDataType::Digit),
950 FieldDef::new("ratio", WpDataType::Float),
951 FieldDef::new("active", WpDataType::Bool),
952 ];
953 let records_in = vec![make_record(vec![
955 Field::from_chars("name", "Alice"),
956 Field::from_int("count", 42),
957 Field::from_float("ratio", 1.23),
958 Field::from_bool("active", true),
959 ])];
960 let batch = records_to_batch(&records_in, &fds).unwrap();
961 let records_out = batch_to_records(&batch, &fds).unwrap();
962
963 assert_eq!(records_out.len(), 1);
964 let rec = &records_out[0];
965 assert_eq!(
966 rec.get_value("name"),
967 Some(&Value::Chars(FValueStr::from("Alice")))
968 );
969 assert_eq!(rec.get_value("count"), Some(&Value::Int(42)));
970 assert_eq!(rec.get_value("ratio"), Some(&Value::Float(1.23)));
971 assert_eq!(rec.get_value("active"), Some(&Value::Bool(true)));
972 }
973
974 #[test]
975 fn b2r_timestamp() {
976 let fds = vec![FieldDef::new("ts", WpDataType::Time)];
977 let ndt =
978 NaiveDateTime::parse_from_str("2024-06-15 12:30:00", "%Y-%m-%d %H:%M:%S").unwrap();
979 let records_in = vec![make_record(vec![Field::from_time("ts", ndt)])];
980 let batch = records_to_batch(&records_in, &fds).unwrap();
981 let records_out = batch_to_records(&batch, &fds).unwrap();
982
983 assert_eq!(records_out[0].get_value("ts"), Some(&Value::Time(ndt)));
984 }
985
986 #[test]
987 fn b2r_ip_parsing() {
988 let fds = vec![FieldDef::new("addr", WpDataType::Ip)];
989 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
990 let net = IpNetValue::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 0)), 8).unwrap();
991 let records_in = vec![
992 make_record(vec![Field::from_ip("addr", ip)]),
993 make_record(vec![Field::new(
994 DataType::IP,
995 "addr",
996 Value::IpNet(net.clone()),
997 )]),
998 ];
999 let batch = records_to_batch(&records_in, &fds).unwrap();
1000 let records_out = batch_to_records(&batch, &fds).unwrap();
1001
1002 assert_eq!(records_out[0].get_value("addr"), Some(&Value::IpAddr(ip)));
1003 assert_eq!(records_out[1].get_value("addr"), Some(&Value::IpNet(net)));
1004 }
1005
1006 #[test]
1007 fn b2r_hex_parsing() {
1008 let fds = vec![FieldDef::new("color", WpDataType::Hex)];
1009 let records_in = vec![make_record(vec![Field::from_hex("color", HexT(255))])];
1010 let batch = records_to_batch(&records_in, &fds).unwrap();
1011 let records_out = batch_to_records(&batch, &fds).unwrap();
1012
1013 assert_eq!(
1014 records_out[0].get_value("color"),
1015 Some(&Value::Hex(HexT(255)))
1016 );
1017 }
1018
1019 #[test]
1020 fn b2r_sequential_ids() {
1021 let fds = vec![FieldDef::new("x", WpDataType::Digit)];
1022 let records_in = vec![
1023 make_record(vec![Field::from_int("x", 1)]),
1024 make_record(vec![Field::from_int("x", 2)]),
1025 make_record(vec![Field::from_int("x", 3)]),
1026 ];
1027 let batch = records_to_batch(&records_in, &fds).unwrap();
1028 let records_out = batch_to_records(&batch, &fds).unwrap();
1029
1030 assert_eq!(records_out[0].id, 0);
1031 assert_eq!(records_out[1].id, 1);
1032 assert_eq!(records_out[2].id, 2);
1033 }
1034
1035 #[test]
1036 fn b2r_schema_mismatch() {
1037 let fds_2 = vec![
1038 FieldDef::new("a", WpDataType::Digit),
1039 FieldDef::new("b", WpDataType::Digit),
1040 ];
1041 let fds_1 = vec![FieldDef::new("a", WpDataType::Digit)];
1042 let records = vec![make_record(vec![Field::from_int("a", 1)])];
1043 let batch = records_to_batch(&records, &fds_1).unwrap();
1044
1045 let err = batch_to_records(&batch, &fds_2).unwrap_err();
1046 assert!(matches!(
1047 err,
1048 WpArrowError::SchemaMismatch {
1049 expected: 2,
1050 actual: 1
1051 }
1052 ));
1053 }
1054
1055 #[test]
1060 fn roundtrip_all_types() {
1061 let ip = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
1062 let net = IpNetValue::new(IpAddr::V4(Ipv4Addr::new(172, 16, 0, 0)), 12).unwrap();
1063 let ndt =
1064 NaiveDateTime::parse_from_str("2025-01-01 00:00:00", "%Y-%m-%d %H:%M:%S").unwrap();
1065
1066 let fds = vec![
1067 FieldDef::new("chars", WpDataType::Chars),
1068 FieldDef::new("digit", WpDataType::Digit),
1069 FieldDef::new("float", WpDataType::Float),
1070 FieldDef::new("bool", WpDataType::Bool),
1071 FieldDef::new("time", WpDataType::Time),
1072 FieldDef::new("ip", WpDataType::Ip),
1073 FieldDef::new("hex", WpDataType::Hex),
1074 FieldDef::new("nums", WpDataType::Array(Box::new(WpDataType::Digit))),
1075 ];
1076
1077 let arr_items: Vec<DataField> =
1078 vec![Field::from_int("item", 100), Field::from_int("item", 200)];
1079
1080 let records_in = vec![
1081 make_record(vec![
1082 Field::from_chars("chars", "hello"),
1083 Field::from_int("digit", 42),
1084 Field::from_float("float", 9.876),
1085 Field::from_bool("bool", true),
1086 Field::from_time("time", ndt),
1087 Field::from_ip("ip", ip),
1088 Field::from_hex("hex", HexT(0xDEAD)),
1089 Field::from_arr("nums", arr_items),
1090 ]),
1091 make_record(vec![
1092 Field::from_chars("chars", "world"),
1093 Field::from_int("digit", -1),
1094 Field::from_float("float", 0.0),
1095 Field::from_bool("bool", false),
1096 Field::from_time("time", ndt),
1097 Field::new(DataType::IP, "ip", Value::IpNet(net.clone())),
1098 Field::from_hex("hex", HexT(0)),
1099 Field::from_arr("nums", vec![Field::from_int("item", 300)]),
1100 ]),
1101 ];
1102
1103 let batch = records_to_batch(&records_in, &fds).unwrap();
1104 let records_out = batch_to_records(&batch, &fds).unwrap();
1105
1106 assert_eq!(records_out.len(), 2);
1107
1108 assert_eq!(
1110 records_out[0].get_value("chars"),
1111 Some(&Value::Chars(FValueStr::from("hello")))
1112 );
1113 assert_eq!(records_out[0].get_value("digit"), Some(&Value::Int(42)));
1114 assert_eq!(
1115 records_out[0].get_value("float"),
1116 Some(&Value::Float(9.876))
1117 );
1118 assert_eq!(records_out[0].get_value("bool"), Some(&Value::Bool(true)));
1119 assert_eq!(records_out[0].get_value("time"), Some(&Value::Time(ndt)));
1120 assert_eq!(records_out[0].get_value("ip"), Some(&Value::IpAddr(ip)));
1121 assert_eq!(
1122 records_out[0].get_value("hex"),
1123 Some(&Value::Hex(HexT(0xDEAD)))
1124 );
1125
1126 if let Some(Value::Array(items)) = records_out[0].get_value("nums") {
1128 assert_eq!(items.len(), 2);
1129 assert_eq!(items[0].get_value(), &Value::Int(100));
1130 assert_eq!(items[1].get_value(), &Value::Int(200));
1131 } else {
1132 panic!("expected Array value for 'nums'");
1133 }
1134
1135 assert_eq!(records_out[1].get_value("ip"), Some(&Value::IpNet(net)));
1137 assert_eq!(records_out[1].get_value("hex"), Some(&Value::Hex(HexT(0))));
1138 }
1139
1140 #[test]
1141 fn roundtrip_with_nulls() {
1142 let fds = vec![
1143 FieldDef::new("name", WpDataType::Chars),
1144 FieldDef::new("opt_digit", WpDataType::Digit),
1145 ];
1146
1147 let records_in = vec![
1148 make_record(vec![
1149 Field::from_chars("name", "row1"),
1150 Field::from_int("opt_digit", 100),
1151 ]),
1152 make_record(vec![
1153 Field::from_chars("name", "row2"),
1154 ]),
1156 ];
1157
1158 let batch = records_to_batch(&records_in, &fds).unwrap();
1159 let records_out = batch_to_records(&batch, &fds).unwrap();
1160
1161 assert_eq!(records_out.len(), 2);
1162 assert_eq!(
1163 records_out[0].get_value("opt_digit"),
1164 Some(&Value::Int(100))
1165 );
1166 assert_eq!(records_out[1].get_value("opt_digit"), None);
1168 }
1169
1170 #[test]
1171 fn roundtrip_bigint_ipv6_key() {
1172 let fds = vec![FieldDef::new("ip_num", WpDataType::BigInt)];
1174
1175 let v4 = BigUint::from_str("134744072").unwrap();
1176 let v6 = BigUint::from_str("382824323044708348099391746388336347272").unwrap();
1177
1178 let records_in = vec![
1179 make_record(vec![Field::new(
1180 DataType::BigInt,
1181 "ip_num",
1182 Value::BigUint(v4.clone()),
1183 )]),
1184 make_record(vec![Field::new(
1185 DataType::BigInt,
1186 "ip_num",
1187 Value::BigUint(v6.clone()),
1188 )]),
1189 ];
1190
1191 let batch = records_to_batch(&records_in, &fds).unwrap();
1192 let records_out = batch_to_records(&batch, &fds).unwrap();
1193
1194 assert_eq!(records_out.len(), 2);
1195 assert_eq!(
1196 records_out[0].get_value("ip_num"),
1197 Some(&Value::BigUint(v4))
1198 );
1199 assert_eq!(
1200 records_out[1].get_value("ip_num"),
1201 Some(&Value::BigUint(v6))
1202 );
1203 assert_eq!(
1205 records_out[1].field("ip_num").map(|f| f.get_meta()),
1206 Some(&DataType::BigInt)
1207 );
1208 }
1209
1210 #[test]
1211 fn roundtrip_bigint_list() {
1212 let fds = vec![FieldDef::new(
1214 "nums",
1215 WpDataType::Array(Box::new(WpDataType::BigInt)),
1216 )];
1217
1218 let a = BigUint::from(1u32);
1219 let b = BigUint::from_str("340282366920938463463374607431768211456").unwrap();
1220
1221 let records_in = vec![make_record(vec![Field::from_arr(
1222 "nums",
1223 vec![
1224 Field::new(DataType::BigInt, "item", Value::BigUint(a.clone())),
1225 Field::new(DataType::BigInt, "item", Value::BigUint(b.clone())),
1226 ],
1227 )])];
1228
1229 let batch = records_to_batch(&records_in, &fds).unwrap();
1230 let records_out = batch_to_records(&batch, &fds).unwrap();
1231
1232 if let Some(Value::Array(items)) = records_out[0].get_value("nums") {
1233 assert_eq!(items.len(), 2);
1234 assert_eq!(items[0].get_value(), &Value::BigUint(a));
1235 assert_eq!(items[1].get_value(), &Value::BigUint(b));
1236 } else {
1237 panic!("expected array value");
1238 }
1239 }
1240}