use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum WireArgType {
Primitive {
arrow: String,
},
CypherValue,
Vector {
len: usize,
element: String,
},
Variadic {
inner: Box<WireArgType>,
},
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn vector_manifest_parses() {
let json = r#"{"kind":"vector","len":128,"element":"float32"}"#;
let parsed: WireArgType = serde_json::from_str(json).expect("vector must parse");
assert_eq!(
parsed,
WireArgType::Vector {
len: 128,
element: "float32".to_owned(),
}
);
}
#[test]
fn variadic_manifest_parses() {
let json = r#"{"kind":"variadic","inner":{"kind":"cypher_value"}}"#;
let parsed: WireArgType = serde_json::from_str(json).expect("variadic must parse");
assert_eq!(
parsed,
WireArgType::Variadic {
inner: Box::new(WireArgType::CypherValue),
}
);
}
#[test]
fn round_trips_through_json() {
for value in [
WireArgType::Primitive {
arrow: "int64".to_owned(),
},
WireArgType::CypherValue,
WireArgType::Vector {
len: 4,
element: "float64".to_owned(),
},
WireArgType::Variadic {
inner: Box::new(WireArgType::Primitive {
arrow: "utf8".to_owned(),
}),
},
] {
let json = serde_json::to_string(&value).expect("serialize");
let back: WireArgType = serde_json::from_str(&json).expect("deserialize");
assert_eq!(back, value, "round-trip must be lossless for {value:?}");
}
}
#[test]
fn unknown_kind_is_rejected() {
let json = r#"{"kind":"quaternion"}"#;
assert!(serde_json::from_str::<WireArgType>(json).is_err());
}
#[test]
fn unknown_field_on_struct_variant_is_rejected() {
let json = r#"{"kind":"primitive","arrow":"int64","bogus":1}"#;
assert!(serde_json::from_str::<WireArgType>(json).is_err());
}
#[test]
fn unknown_field_on_unit_variant_is_tolerated() {
let json = r#"{"kind":"cypher_value","bogus":1}"#;
assert_eq!(
serde_json::from_str::<WireArgType>(json).expect("unit variant ignores extra keys"),
WireArgType::CypherValue
);
}
}