1use std::sync::Arc;
15
16use arrow::datatypes::{DataType as ArrowDataType, Field as ArrowField, Schema, TimeUnit};
17
18use crate::error::WpArrowError;
19
20#[derive(Debug, Clone, PartialEq, Eq, Hash)]
22pub enum WpDataType {
23 Chars,
24 Digit,
25 BigInt,
27 Float,
28 Bool,
29 Time,
30 Ip,
31 Hex,
32 Array(Box<WpDataType>),
33}
34
35pub const BIGINT_DECIMAL_PRECISION: u8 = 39;
38
39#[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
62pub fn to_arrow_type(wp_type: &WpDataType) -> ArrowDataType {
64 match wp_type {
65 WpDataType::Chars => ArrowDataType::Utf8,
66 WpDataType::Digit => ArrowDataType::Int64,
67 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
81pub 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
92pub 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
101pub 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 #[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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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}