Skip to main content

wp_arrow/
schema.rs

1//! `WpDataType`(本 crate 的 9 变体枚举)→ Arrow 类型映射。
2//!
3//! **这是「类型化前端」,不是线协议契约的入口。** 契约(wparse sink ↔ wfusion 接收)见
4//! [`crate::contract::wp_type_to_arrow`]:它穷尽 `wp_model_core::model::DataType` 的 37 个变体,
5//! 而本模块的 [`WpDataType`] 是 9 变体的强类型前端,它自己的 `Array` → `List(inner)`、
6//! `BigInt` → `Decimal256` **与线协议无关**(契约口径是保守的:结构化与大整数一律 `Utf8`)。
7//!
8//! ⚠️ 所以不要用本模块的口径判断契约是否一致。详见 crate 级文档与
9//! `wp-reactor/docs/design/arrow-type-mapping.md`(§1 A-0 / A-2)。
10//!
11//! 注意 [`WpDataType::Digit`] 的名字**刻意保留**:家族那轮 `Digit → Int` 正名没触及它,
12//! 因为改名会同时改字段元数据字符串与行为。
13
14use std::sync::Arc;
15
16use arrow::datatypes::{DataType as ArrowDataType, Field as ArrowField, Schema, TimeUnit};
17
18use crate::error::WpArrowError;
19
20/// WPL data types that can be mapped to Apache Arrow types.
21#[derive(Debug, Clone, PartialEq, Eq, Hash)]
22pub enum WpDataType {
23    Chars,
24    Digit,
25    /// 任意精度无符号整数(IPv4/IPv6 统一数值键等),以 Decimal256 在 Arrow 中传输
26    BigInt,
27    Float,
28    Bool,
29    Time,
30    Ip,
31    Hex,
32    Array(Box<WpDataType>),
33}
34
35/// BigInt 的 Decimal256 精度:2^129-1(IPv6 统一键上限)为 39 位十进制。
36/// 若未来编码位数超过 39,需同步增大此值。
37pub const BIGINT_DECIMAL_PRECISION: u8 = 39;
38
39/// A named, typed field definition for building Arrow schemas.
40#[derive(Debug, Clone, PartialEq, Eq)]
41pub struct FieldDef {
42    pub name: String,
43    pub data_type: WpDataType,
44    pub nullable: bool,
45}
46
47impl FieldDef {
48    pub fn new(name: impl Into<String>, data_type: WpDataType) -> Self {
49        Self {
50            name: name.into(),
51            data_type,
52            nullable: true,
53        }
54    }
55
56    pub fn with_nullable(mut self, nullable: bool) -> Self {
57        self.nullable = nullable;
58        self
59    }
60}
61
62/// Maps a [`WpDataType`] to the corresponding [`ArrowDataType`].
63pub fn to_arrow_type(wp_type: &WpDataType) -> ArrowDataType {
64    match wp_type {
65        WpDataType::Chars => ArrowDataType::Utf8,
66        WpDataType::Digit => ArrowDataType::Int64,
67        // 任意精度整数:Decimal256(39, 0) 保留数值语义,可无损表示 2^129-1
68        WpDataType::BigInt => ArrowDataType::Decimal256(BIGINT_DECIMAL_PRECISION, 0),
69        WpDataType::Float => ArrowDataType::Float64,
70        WpDataType::Bool => ArrowDataType::Boolean,
71        WpDataType::Time => ArrowDataType::Timestamp(TimeUnit::Nanosecond, None),
72        WpDataType::Ip => ArrowDataType::Utf8,
73        WpDataType::Hex => ArrowDataType::Utf8,
74        WpDataType::Array(inner) => {
75            let inner_arrow = to_arrow_type(inner);
76            ArrowDataType::List(Arc::new(ArrowField::new("item", inner_arrow, true)))
77        }
78    }
79}
80
81/// Converts a [`FieldDef`] into an Arrow [`ArrowField`].
82///
83/// Returns an error if the field name is empty.
84pub fn to_arrow_field(field: &FieldDef) -> Result<ArrowField, WpArrowError> {
85    if field.name.is_empty() {
86        return Err(WpArrowError::EmptyFieldName);
87    }
88    let arrow_type = to_arrow_type(&field.data_type);
89    Ok(ArrowField::new(&field.name, arrow_type, field.nullable))
90}
91
92/// Converts a slice of [`FieldDef`] into an Arrow [`Schema`].
93pub fn to_arrow_schema(fields: &[FieldDef]) -> Result<Schema, WpArrowError> {
94    let arrow_fields: Vec<ArrowField> = fields
95        .iter()
96        .map(to_arrow_field)
97        .collect::<Result<_, _>>()?;
98    Ok(Schema::new(arrow_fields))
99}
100
101/// Parses a WPL type string into a [`WpDataType`].
102///
103/// Supported formats:
104/// - Basic types: `"chars"`, `"digit"`, `"float"`, `"bool"`, `"time"`, `"ip"`, `"hex"`
105/// - Array types: `"array<chars>"`, `"array<array<digit>>"`
106///
107/// Type names are case-insensitive.
108pub fn parse_wp_type(s: &str) -> Result<WpDataType, WpArrowError> {
109    let s = s.trim();
110    let lower = s.to_ascii_lowercase();
111
112    match lower.as_str() {
113        "chars" => Ok(WpDataType::Chars),
114        "digit" => Ok(WpDataType::Digit),
115        "bigint" => Ok(WpDataType::BigInt),
116        "float" => Ok(WpDataType::Float),
117        "bool" => Ok(WpDataType::Bool),
118        "time" => Ok(WpDataType::Time),
119        "ip" => Ok(WpDataType::Ip),
120        "hex" => Ok(WpDataType::Hex),
121        _ if lower.starts_with("array<") && lower.ends_with('>') => {
122            let inner_str = &s[6..s.len() - 1];
123            let inner_trimmed = inner_str.trim();
124            if inner_trimmed.is_empty() {
125                return Err(WpArrowError::InvalidArrayInnerType(String::new()));
126            }
127            let inner = parse_wp_type(inner_trimmed)?;
128            Ok(WpDataType::Array(Box::new(inner)))
129        }
130        _ => Err(WpArrowError::UnsupportedDataType(s.to_string())),
131    }
132}
133
134#[cfg(test)]
135mod tests {
136    use super::*;
137
138    // ---------------------------------------------------------------
139    // to_arrow_type: basic types
140    // ---------------------------------------------------------------
141
142    #[test]
143    fn arrow_type_chars() {
144        assert_eq!(to_arrow_type(&WpDataType::Chars), ArrowDataType::Utf8);
145    }
146
147    #[test]
148    fn arrow_type_digit() {
149        assert_eq!(to_arrow_type(&WpDataType::Digit), ArrowDataType::Int64);
150    }
151
152    #[test]
153    fn arrow_type_bigint() {
154        // 任意精度整数以 Decimal256(39, 0) 传输(可表示 2^129-1),保留数值语义
155        assert_eq!(
156            to_arrow_type(&WpDataType::BigInt),
157            ArrowDataType::Decimal256(BIGINT_DECIMAL_PRECISION, 0)
158        );
159    }
160
161    #[test]
162    fn arrow_type_float() {
163        assert_eq!(to_arrow_type(&WpDataType::Float), ArrowDataType::Float64);
164    }
165
166    #[test]
167    fn arrow_type_bool() {
168        assert_eq!(to_arrow_type(&WpDataType::Bool), ArrowDataType::Boolean);
169    }
170
171    #[test]
172    fn arrow_type_time() {
173        assert_eq!(
174            to_arrow_type(&WpDataType::Time),
175            ArrowDataType::Timestamp(TimeUnit::Nanosecond, None)
176        );
177    }
178
179    #[test]
180    fn arrow_type_ip() {
181        assert_eq!(to_arrow_type(&WpDataType::Ip), ArrowDataType::Utf8);
182    }
183
184    #[test]
185    fn arrow_type_hex() {
186        assert_eq!(to_arrow_type(&WpDataType::Hex), ArrowDataType::Utf8);
187    }
188
189    // ---------------------------------------------------------------
190    // to_arrow_type: array types
191    // ---------------------------------------------------------------
192
193    #[test]
194    fn arrow_type_array_digit() {
195        let wp = WpDataType::Array(Box::new(WpDataType::Digit));
196        let arrow = to_arrow_type(&wp);
197        assert_eq!(
198            arrow,
199            ArrowDataType::List(Arc::new(ArrowField::new(
200                "item",
201                ArrowDataType::Int64,
202                true
203            )))
204        );
205    }
206
207    #[test]
208    fn arrow_type_array_chars() {
209        let wp = WpDataType::Array(Box::new(WpDataType::Chars));
210        let arrow = to_arrow_type(&wp);
211        assert_eq!(
212            arrow,
213            ArrowDataType::List(Arc::new(ArrowField::new("item", ArrowDataType::Utf8, true)))
214        );
215    }
216
217    #[test]
218    fn arrow_type_nested_array() {
219        let wp = WpDataType::Array(Box::new(WpDataType::Array(Box::new(WpDataType::Float))));
220        let inner_list = ArrowDataType::List(Arc::new(ArrowField::new(
221            "item",
222            ArrowDataType::Float64,
223            true,
224        )));
225        let expected = ArrowDataType::List(Arc::new(ArrowField::new("item", inner_list, true)));
226        assert_eq!(to_arrow_type(&wp), expected);
227    }
228
229    // ---------------------------------------------------------------
230    // to_arrow_field
231    // ---------------------------------------------------------------
232
233    #[test]
234    fn arrow_field_basic() {
235        let fd = FieldDef::new("src_ip", WpDataType::Ip);
236        let field = to_arrow_field(&fd).unwrap();
237        assert_eq!(field.name(), "src_ip");
238        assert_eq!(field.data_type(), &ArrowDataType::Utf8);
239        assert!(field.is_nullable());
240    }
241
242    #[test]
243    fn arrow_field_non_nullable() {
244        let fd = FieldDef::new("count", WpDataType::Digit).with_nullable(false);
245        let field = to_arrow_field(&fd).unwrap();
246        assert!(!field.is_nullable());
247    }
248
249    #[test]
250    fn arrow_field_empty_name_errors() {
251        let fd = FieldDef::new("", WpDataType::Chars);
252        assert_eq!(to_arrow_field(&fd), Err(WpArrowError::EmptyFieldName));
253    }
254
255    // ---------------------------------------------------------------
256    // to_arrow_schema
257    // ---------------------------------------------------------------
258
259    #[test]
260    fn arrow_schema_firewall_log() {
261        let fields = vec![
262            FieldDef::new("src_ip", WpDataType::Ip),
263            FieldDef::new("dst_ip", WpDataType::Ip),
264            FieldDef::new("port", WpDataType::Digit),
265            FieldDef::new("protocol", WpDataType::Chars),
266            FieldDef::new("timestamp", WpDataType::Time),
267            FieldDef::new("allowed", WpDataType::Bool),
268        ];
269        let schema = to_arrow_schema(&fields).unwrap();
270        assert_eq!(schema.fields().len(), 6);
271        assert_eq!(schema.field(0).name(), "src_ip");
272        assert_eq!(schema.field(2).data_type(), &ArrowDataType::Int64);
273        assert_eq!(
274            schema.field(4).data_type(),
275            &ArrowDataType::Timestamp(TimeUnit::Nanosecond, None)
276        );
277    }
278
279    #[test]
280    fn arrow_schema_with_array_field() {
281        let fields = vec![
282            FieldDef::new("name", WpDataType::Chars),
283            FieldDef::new("tags", WpDataType::Array(Box::new(WpDataType::Chars))),
284        ];
285        let schema = to_arrow_schema(&fields).unwrap();
286        assert_eq!(schema.fields().len(), 2);
287        assert!(matches!(
288            schema.field(1).data_type(),
289            ArrowDataType::List(_)
290        ));
291    }
292
293    #[test]
294    fn arrow_schema_empty_fields() {
295        let schema = to_arrow_schema(&[]).unwrap();
296        assert_eq!(schema.fields().len(), 0);
297    }
298
299    #[test]
300    fn arrow_schema_error_propagation() {
301        let fields = vec![
302            FieldDef::new("ok", WpDataType::Chars),
303            FieldDef::new("", WpDataType::Digit),
304        ];
305        assert_eq!(to_arrow_schema(&fields), Err(WpArrowError::EmptyFieldName));
306    }
307
308    // ---------------------------------------------------------------
309    // parse_wp_type: basic types
310    // ---------------------------------------------------------------
311
312    #[test]
313    fn parse_chars() {
314        assert_eq!(parse_wp_type("chars"), Ok(WpDataType::Chars));
315    }
316
317    #[test]
318    fn parse_digit() {
319        assert_eq!(parse_wp_type("digit"), Ok(WpDataType::Digit));
320    }
321
322    #[test]
323    fn parse_bigint() {
324        assert_eq!(parse_wp_type("bigint"), Ok(WpDataType::BigInt));
325        assert_eq!(parse_wp_type("BIGINT"), Ok(WpDataType::BigInt));
326    }
327
328    #[test]
329    fn parse_float() {
330        assert_eq!(parse_wp_type("float"), Ok(WpDataType::Float));
331    }
332
333    #[test]
334    fn parse_bool() {
335        assert_eq!(parse_wp_type("bool"), Ok(WpDataType::Bool));
336    }
337
338    #[test]
339    fn parse_time() {
340        assert_eq!(parse_wp_type("time"), Ok(WpDataType::Time));
341    }
342
343    #[test]
344    fn parse_ip() {
345        assert_eq!(parse_wp_type("ip"), Ok(WpDataType::Ip));
346    }
347
348    #[test]
349    fn parse_hex() {
350        assert_eq!(parse_wp_type("hex"), Ok(WpDataType::Hex));
351    }
352
353    // ---------------------------------------------------------------
354    // parse_wp_type: case insensitivity
355    // ---------------------------------------------------------------
356
357    #[test]
358    fn parse_case_insensitive() {
359        assert_eq!(parse_wp_type("CHARS"), Ok(WpDataType::Chars));
360        assert_eq!(parse_wp_type("Digit"), Ok(WpDataType::Digit));
361        assert_eq!(parse_wp_type("BOOL"), Ok(WpDataType::Bool));
362    }
363
364    // ---------------------------------------------------------------
365    // parse_wp_type: array types
366    // ---------------------------------------------------------------
367
368    #[test]
369    fn parse_array_chars() {
370        assert_eq!(
371            parse_wp_type("array<chars>"),
372            Ok(WpDataType::Array(Box::new(WpDataType::Chars)))
373        );
374    }
375
376    #[test]
377    fn parse_array_digit() {
378        assert_eq!(
379            parse_wp_type("array<digit>"),
380            Ok(WpDataType::Array(Box::new(WpDataType::Digit)))
381        );
382    }
383
384    #[test]
385    fn parse_nested_array() {
386        assert_eq!(
387            parse_wp_type("array<array<float>>"),
388            Ok(WpDataType::Array(Box::new(WpDataType::Array(Box::new(
389                WpDataType::Float
390            )))))
391        );
392    }
393
394    #[test]
395    fn parse_array_with_whitespace() {
396        assert_eq!(
397            parse_wp_type("  array< chars >  "),
398            Ok(WpDataType::Array(Box::new(WpDataType::Chars)))
399        );
400    }
401
402    // ---------------------------------------------------------------
403    // parse_wp_type: error cases
404    // ---------------------------------------------------------------
405
406    #[test]
407    fn parse_unsupported_type() {
408        let err = parse_wp_type("unknown").unwrap_err();
409        assert_eq!(
410            err,
411            WpArrowError::UnsupportedDataType("unknown".to_string())
412        );
413    }
414
415    #[test]
416    fn parse_array_empty_inner() {
417        let err = parse_wp_type("array<>").unwrap_err();
418        assert_eq!(err, WpArrowError::InvalidArrayInnerType(String::new()));
419    }
420
421    #[test]
422    fn parse_array_invalid_inner() {
423        let err = parse_wp_type("array<invalid>").unwrap_err();
424        assert_eq!(
425            err,
426            WpArrowError::UnsupportedDataType("invalid".to_string())
427        );
428    }
429
430    // ---------------------------------------------------------------
431    // Property tests: Clone, Eq, Hash, FieldDef defaults
432    // ---------------------------------------------------------------
433
434    #[test]
435    fn wf_data_type_clone_eq() {
436        let a = WpDataType::Array(Box::new(WpDataType::Chars));
437        let b = a.clone();
438        assert_eq!(a, b);
439    }
440
441    #[test]
442    fn wf_data_type_hash_consistent() {
443        use std::collections::HashSet;
444        let mut set = HashSet::new();
445        set.insert(WpDataType::Digit);
446        set.insert(WpDataType::Digit);
447        assert_eq!(set.len(), 1);
448    }
449
450    #[test]
451    fn field_def_default_nullable() {
452        let fd = FieldDef::new("test", WpDataType::Bool);
453        assert!(fd.nullable);
454    }
455}