Skip to main content

wp_arrow/contract/
value.rs

1//! 线协议契约的**值层**:`DataRecord` → 列(按 Arrow 列类型派发)。
2//!
3//! 与 [`crate::contract`] 的 schema 表配套:表回答「这一列该是什么 Arrow 类型」,
4//! 本模块回答「值怎么写进那一列」。两者合起来才是完整的线协议契约。
5//!
6//! # 归属(A-2 第 3 步 / 2c)
7//!
8//! 这些函数**逐字移植自** `wp-connector-utils/src/arrow/record.rs`(那个 crate 是
9//! 「面向 sink 的 connector 工具」,契约的语义归属不在这里)。移植时只改了错误类型:
10//! `SinkResult`/`SinkReason`(来自 `wp-connector-api`,本 crate 不能依赖)→
11//! [`WpArrowError::ArrowBuildError`]。行为、错误文案形状与派发顺序保持不变。
12//!
13//! 派发**按 Arrow 列类型**(不是按 wp-model 类型):所以「某类型编成什么列」由
14//! [`crate::contract::wp_type_to_arrow`] 决定,本模块只认列类型 —— 两处不会各有一套口径。
15
16use std::sync::Arc;
17
18use arrow::array::{
19    ArrayRef, BinaryBuilder, BooleanBuilder, Float64Builder, Int32Builder, Int64Builder,
20    StringBuilder, TimestampNanosecondBuilder,
21};
22use arrow::datatypes::{DataType, Field, Schema, TimeUnit};
23use arrow::record_batch::RecordBatch;
24use wp_model_core::model::{DataRecord, Value};
25
26use crate::error::WpArrowError;
27
28/// 单条 `DataRecord` → 一行 `RecordBatch`。
29///
30/// 每个列按名字在记录里查找;缺字段 → null。
31pub fn encode_record(
32    record: &DataRecord,
33    schema: &Arc<Schema>,
34) -> Result<RecordBatch, WpArrowError> {
35    let mut columns: Vec<ArrayRef> = Vec::with_capacity(schema.fields().len());
36    for field in schema.fields() {
37        let records = [Arc::new(record.clone())];
38        columns.push(build_column_from_field(field, &records)?);
39    }
40    RecordBatch::try_new(Arc::clone(schema), columns)
41        .map_err(|e| WpArrowError::ArrowBuildError(e.to_string()))
42}
43
44/// 多条 `DataRecord` → 一个 `RecordBatch`。
45///
46/// 每个列按名字在**每条**记录里查找;缺字段 → null。`records` 为空时按 schema
47/// 产零行(列类型齐全,便于下游直接 append)。
48pub fn encode_records(
49    records: &[Arc<DataRecord>],
50    schema: &Arc<Schema>,
51) -> Result<RecordBatch, WpArrowError> {
52    if records.is_empty() {
53        let empty_columns: Vec<ArrayRef> = schema
54            .fields()
55            .iter()
56            .map(|f| empty_column_for_type(f.data_type()))
57            .collect::<Result<Vec<_>, _>>()?;
58        return RecordBatch::try_new(Arc::clone(schema), empty_columns)
59            .map_err(|e| WpArrowError::ArrowBuildError(e.to_string()));
60    }
61
62    let mut columns: Vec<ArrayRef> = Vec::with_capacity(schema.fields().len());
63    for field in schema.fields() {
64        columns.push(build_column_from_field(field, records)?);
65    }
66    RecordBatch::try_new(Arc::clone(schema), columns)
67        .map_err(|e| WpArrowError::ArrowBuildError(e.to_string()))
68}
69
70// ---------------------------------------------------------------------------
71// 列构造:按 Arrow 列类型派发
72// ---------------------------------------------------------------------------
73
74fn build_column_from_field(
75    field: &Field,
76    records: &[Arc<DataRecord>],
77) -> Result<ArrayRef, WpArrowError> {
78    let field_name = field.name();
79    match field.data_type() {
80        DataType::Boolean => {
81            let mut builder = BooleanBuilder::with_capacity(records.len());
82            for record in records {
83                match record.field(field_name).map(|f| f.get_value()) {
84                    Some(Value::Bool(v)) => builder.append_value(*v),
85                    Some(Value::Chars(s)) => builder.append_value(s.eq_ignore_ascii_case("true")),
86                    _ => builder.append_null(),
87                }
88            }
89            Ok(Arc::new(builder.finish()) as ArrayRef)
90        }
91        DataType::Int64 => {
92            let mut builder = Int64Builder::with_capacity(records.len());
93            for record in records {
94                match record
95                    .field(field_name)
96                    .and_then(|f| parse_digit(f.get_value()))
97                {
98                    Some(v) => builder.append_value(v),
99                    None => builder.append_null(),
100                }
101            }
102            Ok(Arc::new(builder.finish()) as ArrayRef)
103        }
104        DataType::Int32 => {
105            let mut builder = Int32Builder::with_capacity(records.len());
106            for record in records {
107                match record
108                    .field(field_name)
109                    .and_then(|f| parse_digit(f.get_value()))
110                {
111                    Some(v) => builder.append_value(v as i32),
112                    None => builder.append_null(),
113                }
114            }
115            Ok(Arc::new(builder.finish()) as ArrayRef)
116        }
117        DataType::Binary => {
118            let mut builder = BinaryBuilder::with_capacity(records.len(), records.len() * 64);
119            for record in records {
120                match record.field(field_name).map(|f| f.get_value()) {
121                    Some(v) => {
122                        let bytes = to_raw_bytes(v);
123                        builder.append_value(&bytes[..]);
124                    }
125                    None => builder.append_null(),
126                }
127            }
128            Ok(Arc::new(builder.finish()) as ArrayRef)
129        }
130        DataType::Float64 => {
131            let mut builder = Float64Builder::with_capacity(records.len());
132            for record in records {
133                match record
134                    .field(field_name)
135                    .and_then(|f| parse_float(f.get_value()))
136                {
137                    Some(v) => builder.append_value(v),
138                    None => builder.append_null(),
139                }
140            }
141            Ok(Arc::new(builder.finish()) as ArrayRef)
142        }
143        DataType::Timestamp(TimeUnit::Nanosecond, None) => {
144            let mut builder = TimestampNanosecondBuilder::with_capacity(records.len());
145            for record in records {
146                match record
147                    .field(field_name)
148                    .and_then(|f| parse_timestamp_ns(f.get_value()))
149                {
150                    Some(v) => builder.append_value(v),
151                    None => builder.append_null(),
152                }
153            }
154            Ok(Arc::new(builder.finish()) as ArrayRef)
155        }
156        // Utf8 与其余列类型(结构化字段走的就是这里:JSON 文本)
157        _ => {
158            let mut builder = StringBuilder::with_capacity(records.len(), records.len() * 32);
159            for record in records {
160                match record.field(field_name) {
161                    Some(f) => builder.append_value(format_utf8_value(f.get_value())),
162                    None => builder.append_null(),
163                }
164            }
165            Ok(Arc::new(builder.finish()) as ArrayRef)
166        }
167    }
168}
169
170// ---------------------------------------------------------------------------
171// 值格式化
172// ---------------------------------------------------------------------------
173
174/// [`Value`] → Utf8 列文本。
175///
176/// 结构化(`Obj` / `Array`)序列化为 **JSON**;其余走 `Display`。
177fn format_utf8_value(v: &Value) -> String {
178    match v {
179        Value::Obj(_) | Value::Array(_) => {
180            serde_json::to_string(v).unwrap_or_else(|_| format!("{v:?}"))
181        }
182        _ => v.to_string(),
183    }
184}
185
186/// [`Value`] → Binary 列的原始字节。
187///
188/// `Value::Hex` 取其 `u128` 的**最小大端**字节;其余退化为 Utf8 文本的字节。
189///
190/// 注意:`hex` 字段**不再**走 Binary 列(DIV-1 已对齐为 Utf8),所以 `Value::Hex`
191/// 分支只在**显式声明 Binary** 的列上才可达。
192fn to_raw_bytes(v: &Value) -> Vec<u8> {
193    match v {
194        Value::Hex(h) => {
195            if h.0 == 0 {
196                return vec![0];
197            }
198            let be = h.0.to_be_bytes();
199            let start = be.iter().position(|&b| b != 0).unwrap();
200            be[start..].to_vec()
201        }
202        _ => format_utf8_value(v).into_bytes(),
203    }
204}
205
206// ---------------------------------------------------------------------------
207// 取值助手(带 Chars 回退)
208// ---------------------------------------------------------------------------
209
210fn parse_digit(v: &Value) -> Option<i64> {
211    match v {
212        Value::Int(d) => Some(*d),
213        Value::Float(f) => Some(*f as i64),
214        Value::Chars(s) => s.parse().ok(),
215        _ => None,
216    }
217}
218
219fn parse_float(v: &Value) -> Option<f64> {
220    match v {
221        Value::Float(f) => Some(*f),
222        Value::Int(d) => Some(*d as f64),
223        Value::Chars(s) => s.parse().ok(),
224        _ => None,
225    }
226}
227
228fn parse_timestamp_ns(v: &Value) -> Option<i64> {
229    match v {
230        Value::Time(t) => Some(t.and_utc().timestamp_nanos_opt()?),
231        // 整数时间戳按**毫秒**解读(与移植前的 sink 口径一致)
232        Value::Int(d) => d.checked_mul(1_000_000),
233        Value::Chars(s) => chrono::DateTime::parse_from_rfc3339(s)
234            .ok()
235            .or_else(|| {
236                chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S")
237                    .ok()
238                    .map(|dt| dt.and_utc().fixed_offset())
239            })
240            .and_then(|dt| dt.timestamp_nanos_opt()),
241        _ => None,
242    }
243}
244
245fn empty_column_for_type(data_type: &DataType) -> Result<ArrayRef, WpArrowError> {
246    let arr: ArrayRef = match data_type {
247        DataType::Boolean => Arc::new(arrow::array::BooleanArray::from(Vec::<bool>::new())),
248        DataType::Int32 => Arc::new(arrow::array::Int32Array::from(Vec::<i32>::new())),
249        DataType::Int64 => Arc::new(arrow::array::Int64Array::from(Vec::<i64>::new())),
250        DataType::Float64 => Arc::new(arrow::array::Float64Array::from(Vec::<f64>::new())),
251        DataType::Timestamp(TimeUnit::Nanosecond, None) => Arc::new(
252            arrow::array::TimestampNanosecondArray::from(Vec::<i64>::new()),
253        ),
254        DataType::Binary => Arc::new(arrow::array::BinaryArray::from(Vec::<Option<&[u8]>>::new())),
255        _ => Arc::new(arrow::array::StringArray::from(Vec::<Option<&str>>::new())),
256    };
257    Ok(arr)
258}
259
260#[cfg(test)]
261mod tests {
262    use super::*;
263    // `is_null` / `data_type` 来自 `Array` trait(B 侧的测试模块顶部已导入,这里同样需要)
264    use arrow::array::Array as _;
265    use arrow::array::{
266        BinaryArray, BooleanArray, Float64Array, Int32Array, Int64Array, StringArray,
267        TimestampNanosecondArray,
268    };
269    use wp_model_core::model::types::value::{HexT, ObjectValue};
270    use wp_model_core::model::{DataField, Field as ModelField, FieldStorage};
271
272    /// 线协议**值层**的金标准:覆盖全部列类型 + 缺字段 + 类型回退(Chars/Int/Float/Time 互转)。
273    ///
274    /// 本测与 `wp-connector-utils/src/arrow/record.rs` 的同名测试**逐字同构**
275    /// (同一份夹具、同一组期望)。两份同时通过即证明 A-2 2c 的值层搬迁是**等价**改动,
276    /// 而不只是「看起来一样」;搬迁后本测是**实现侧**钉桩,那边那份变成消费侧拼线。
277    #[test]
278    fn wire_value_encoding_is_pinned_by_golden_values() {
279        let epoch =
280            chrono::NaiveDateTime::parse_from_str("2024-01-01 00:00:00", "%Y-%m-%d %H:%M:%S")
281                .unwrap();
282        let ts_val = epoch + chrono::Duration::seconds(5);
283
284        // row0:全字段齐、尽量走「正路」
285        let mut obj = ObjectValue::new();
286        obj.insert("k", DataField::from_chars("k", "v"));
287        let row0 = DataRecord::from(vec![
288            FieldStorage::from(DataField::from_bool("b", true)),
289            FieldStorage::from(DataField::from_int("i64", 42)),
290            FieldStorage::from(DataField::from_int("i32", 70_000)),
291            FieldStorage::from(DataField::from_float("f", 1.5)),
292            FieldStorage::from(DataField::from_time("ts", ts_val)),
293            FieldStorage::from(DataField::from_chars("bin", "hi")),
294            FieldStorage::from(DataField::from_obj("s", obj)),
295        ]);
296
297        // row1:类型回退(Chars 解析 / Float→Int / Int(ms)→时间戳 / Hex→Binary / Array→JSON)
298        let arr = DataField::from_arr(
299            "s",
300            vec![
301                DataField::from_chars("c", "x"),
302                DataField::from_int("i", 22),
303            ],
304        );
305        let row1 = DataRecord::from(vec![
306            FieldStorage::from(DataField::from_chars("b", "TRUE")),
307            FieldStorage::from(DataField::from_chars("i64", "42")),
308            FieldStorage::from(DataField::from_float("i32", 3.9)),
309            FieldStorage::from(DataField::from_chars("f", "2.71")),
310            FieldStorage::from(DataField::from_int("ts", 1_700_000_000_000)),
311            FieldStorage::from(DataField::from_hex("bin", HexT(0x1A2B))),
312            FieldStorage::from(arr),
313        ]);
314
315        // row2:除 `s` 外全缺 → 其余列 null;`s` 走 Hex 的 Utf8 形态
316        let row2 = DataRecord::from(vec![FieldStorage::from(DataField::from_hex(
317            "s",
318            HexT(0x1A2B),
319        ))]);
320
321        let rows = vec![Arc::new(row0), Arc::new(row1), Arc::new(row2)];
322        let schema = Arc::new(Schema::new(vec![
323            Field::new("b", DataType::Boolean, true),
324            Field::new("i64", DataType::Int64, true),
325            Field::new("i32", DataType::Int32, true),
326            Field::new("f", DataType::Float64, true),
327            Field::new("ts", DataType::Timestamp(TimeUnit::Nanosecond, None), true),
328            Field::new("bin", DataType::Binary, true),
329            Field::new("s", DataType::Utf8, true),
330        ]));
331
332        let batch = encode_records(&rows, &schema).unwrap();
333        assert_eq!(batch.num_rows(), 3);
334
335        let b = batch
336            .column(0)
337            .as_any()
338            .downcast_ref::<BooleanArray>()
339            .unwrap();
340        assert_eq!((b.value(0), b.value(1)), (true, true));
341        assert!(b.is_null(2));
342
343        let i64c = batch
344            .column(1)
345            .as_any()
346            .downcast_ref::<Int64Array>()
347            .unwrap();
348        assert_eq!((i64c.value(0), i64c.value(1)), (42, 42));
349        assert!(i64c.is_null(2));
350
351        let i32c = batch
352            .column(2)
353            .as_any()
354            .downcast_ref::<Int32Array>()
355            .unwrap();
356        assert_eq!((i32c.value(0), i32c.value(1)), (70_000, 3));
357        assert!(i32c.is_null(2));
358
359        let fc = batch
360            .column(3)
361            .as_any()
362            .downcast_ref::<Float64Array>()
363            .unwrap();
364        assert_eq!(fc.value(0), 1.5);
365        assert_eq!(fc.value(1), 2.71);
366        assert!(fc.is_null(2));
367
368        let tsc = batch
369            .column(4)
370            .as_any()
371            .downcast_ref::<TimestampNanosecondArray>()
372            .unwrap();
373        assert_eq!(
374            tsc.value(0),
375            ts_val.and_utc().timestamp_nanos_opt().unwrap()
376        );
377        assert_eq!(tsc.value(1), 1_700_000_000_000 * 1_000_000);
378        assert!(tsc.is_null(2));
379
380        let binc = batch
381            .column(5)
382            .as_any()
383            .downcast_ref::<BinaryArray>()
384            .unwrap();
385        assert_eq!(binc.value(0), b"hi");
386        assert_eq!(binc.value(1), &[0x1A, 0x2B]);
387        assert!(binc.is_null(2));
388
389        let sc = batch
390            .column(6)
391            .as_any()
392            .downcast_ref::<StringArray>()
393            .unwrap();
394        assert!(!sc.is_null(2), "row2 的 s 是有的(Hex),不应 null");
395        // 结构化字段走 JSON(`serde_json::to_string` 同一套规则)——不硬编 JSON 形状,
396        // 但把「必须等于源 Value 的 serde_json 渲染」钉住。
397        for (row, field) in [(0usize, "s"), (1usize, "s")] {
398            let src = rows[row].field(field).unwrap().get_value();
399            let rendered = sc.value(row);
400            match serde_json::to_string(src) {
401                Ok(json) => {
402                    assert_eq!(
403                        rendered, json,
404                        "row{row} 结构化字段应等于源 Value 的 JSON 渲染"
405                    )
406                }
407                Err(e) => panic!(
408                    "row{row}: serde_json 渲染失败({e})→ 列里实际是 {rendered:?},src={src:?}"
409                ),
410            }
411            assert!(
412                serde_json::from_str::<serde_json::Value>(rendered).is_ok(),
413                "结构化字段在 Utf8 列里必须是合法 JSON"
414            );
415        }
416        // row2:Hex 走 Utf8 时是 `{:#X}` 形态(与 `Value::Hex` 的 Display 同形)
417        assert_eq!(sc.value(2), "0x1A2B");
418        assert_eq!(sc.value(2), format!("{:#X}", 0x1A2Bu128));
419    }
420
421    /// 单条入口(`encode_record`)与整批入口行为一致(都按名字查字段)。
422    #[test]
423    fn single_record_entry_matches_batch_entry() {
424        let rec = DataRecord::from(vec![FieldStorage::from(ModelField::from_chars("x", "v"))]);
425        let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, true)]));
426        let one = encode_record(&rec, &schema).unwrap();
427        let many = encode_records(&[Arc::new(rec)], &schema).unwrap();
428        assert_eq!(one.num_rows(), 1);
429        assert_eq!(many.num_rows(), 1);
430        let a = one
431            .column(0)
432            .as_any()
433            .downcast_ref::<StringArray>()
434            .unwrap();
435        let b = many
436            .column(0)
437            .as_any()
438            .downcast_ref::<StringArray>()
439            .unwrap();
440        assert_eq!(a.value(0), b.value(0));
441    }
442
443    /// 空批次:按 schema 产零行,列类型齐全(便于下游 append)。
444    #[test]
445    fn empty_records_produce_typed_zero_row_batch() {
446        let schema = Arc::new(Schema::new(vec![
447            Field::new("i", DataType::Int64, true),
448            Field::new("s", DataType::Utf8, true),
449        ]));
450        let batch = encode_records(&[], &schema).unwrap();
451        assert_eq!(batch.num_rows(), 0);
452        assert_eq!(batch.column(0).data_type(), &DataType::Int64);
453        assert_eq!(batch.column(1).data_type(), &DataType::Utf8);
454    }
455}