use serde::{Deserialize, Serialize};
#[derive(Debug, Default, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum Type {
Int,
Float,
Bool,
String,
Bytes,
#[default]
Unit,
List { elem: Box<Type> },
Tuple { elems: Vec<Type> },
Record { fields: Vec<RecordField> },
Sum { variants: Vec<Variant> },
Fun { params: Vec<Type>, ret: Box<Type> },
Var { name: String },
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RecordField {
pub name: String,
pub ty: Type,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Variant {
pub name: String,
pub fields: Vec<Type>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn primitive_round_trips_through_json() {
let ty = Type::Int;
let json = serde_json::to_string(&ty).unwrap();
let back: Type = serde_json::from_str(&json).unwrap();
assert_eq!(ty, back);
}
#[test]
fn function_type_round_trips_through_json() {
let ty = Type::Fun {
params: vec![Type::Int, Type::String],
ret: Box::new(Type::Bool),
};
let json = serde_json::to_string(&ty).unwrap();
let back: Type = serde_json::from_str(&json).unwrap();
assert_eq!(ty, back);
}
#[test]
fn sum_type_round_trips_through_json() {
let ty = Type::Sum {
variants: vec![
Variant { name: "Circle".to_string(), fields: vec![Type::Float] },
Variant { name: "Rect".to_string(), fields: vec![Type::Float, Type::Float] },
],
};
let json = serde_json::to_string(&ty).unwrap();
let back: Type = serde_json::from_str(&json).unwrap();
assert_eq!(ty, back);
}
#[test]
fn struct_like_variant_wire_format_is_locked() {
let ty = Type::List { elem: Box::new(Type::Int) };
assert_eq!(
serde_json::to_string(&ty).unwrap(),
r#"{"type":"List","elem":{"type":"Int"}}"#
);
}
#[test]
fn record_wire_format_is_locked() {
let ty = Type::Record {
fields: vec![RecordField { name: "x".to_string(), ty: Type::Int }],
};
assert_eq!(
serde_json::to_string(&ty).unwrap(),
r#"{"type":"Record","fields":[{"name":"x","ty":{"type":"Int"}}]}"#
);
}
#[test]
fn var_round_trips_through_json() {
let ty = Type::Var { name: "a".to_string() };
let json = serde_json::to_string(&ty).unwrap();
let back: Type = serde_json::from_str(&json).unwrap();
assert_eq!(ty, back);
}
#[test]
fn nested_composite_type_round_trips_through_json() {
let ty = Type::List {
elem: Box::new(Type::Record {
fields: vec![RecordField { name: "x".to_string(), ty: Type::Float }],
}),
};
let json = serde_json::to_string(&ty).unwrap();
let back: Type = serde_json::from_str(&json).unwrap();
assert_eq!(ty, back);
}
}