use crate::type_tracking::{NativeKind, TypeTracker};
use shape_ast::ast::{Expr, Literal, TypeAnnotation};
pub fn infer_array_element_type(elements: &[Expr], type_tracker: &TypeTracker) -> Option<NativeKind> {
if elements.is_empty() {
return None;
}
if elements.iter().any(|e| matches!(e, Expr::Array(..))) {
return None;
}
if let Some(kind) = infer_from_literals(elements) {
return Some(kind);
}
infer_from_tracked_types(elements, type_tracker)
}
pub fn typed_array_from_annotation(annotation: &TypeAnnotation) -> Option<NativeKind> {
match annotation {
TypeAnnotation::Generic { name, args } if *name == "Array" && args.len() == 1 => {
scalar_annotation_to_slot_kind(&args[0])
}
TypeAnnotation::Array(inner) => scalar_annotation_to_slot_kind(inner),
_ => None,
}
}
fn scalar_annotation_to_slot_kind(annotation: &TypeAnnotation) -> Option<NativeKind> {
match annotation {
TypeAnnotation::Basic(name) => match name.as_str() {
"number" => Some(NativeKind::Float64),
"int" => Some(NativeKind::Int64),
"i8" => Some(NativeKind::Int8),
"u8" => Some(NativeKind::UInt8),
"i16" => Some(NativeKind::Int16),
"u16" => Some(NativeKind::UInt16),
"i32" => Some(NativeKind::Int32),
"u32" => Some(NativeKind::UInt32),
"u64" => Some(NativeKind::UInt64),
"isize" => Some(NativeKind::IntSize),
"usize" => Some(NativeKind::UIntSize),
"bool" => Some(NativeKind::Bool),
"string" => Some(NativeKind::String),
"decimal" => Some(NativeKind::DecimalV2),
_ => None,
},
_ => None,
}
}
fn infer_from_literals(elements: &[Expr]) -> Option<NativeKind> {
let mut kind: Option<NativeKind> = None;
for elem in elements {
let elem_kind = match elem {
Expr::Literal(Literal::Number(_), _) => NativeKind::Float64,
Expr::Literal(Literal::Int(_), _) => NativeKind::Int64,
Expr::Literal(Literal::Bool(_), _) => NativeKind::Bool,
Expr::Literal(Literal::String(_), _) => NativeKind::String,
Expr::Literal(Literal::Decimal(_), _) => NativeKind::DecimalV2,
Expr::Literal(Literal::TypedInt(_, w), _) => typed_int_width_to_slot(*w),
Expr::StructLiteral { .. }
| Expr::Object(..) => {
NativeKind::Ptr(shape_value::HeapKind::TypedObject)
}
_ => return None,
};
match kind {
Some(prev) if prev != elem_kind => return None, Some(_) => {} None => kind = Some(elem_kind),
}
}
kind
}
fn typed_int_width_to_slot(w: shape_ast::IntWidth) -> NativeKind {
use shape_ast::IntWidth;
match w {
IntWidth::I8 => NativeKind::Int8,
IntWidth::U8 => NativeKind::UInt8,
IntWidth::I16 => NativeKind::Int16,
IntWidth::U16 => NativeKind::UInt16,
IntWidth::I32 => NativeKind::Int32,
IntWidth::U32 => NativeKind::UInt32,
IntWidth::U64 => NativeKind::UInt64,
}
}
fn infer_from_tracked_types(elements: &[Expr], type_tracker: &TypeTracker) -> Option<NativeKind> {
let mut kind: Option<NativeKind> = None;
for elem in elements {
let elem_kind = expr_storage_hint(elem, type_tracker)?;
match kind {
Some(prev) if prev != elem_kind => return None,
Some(_) => {}
None => kind = Some(elem_kind),
}
}
kind
}
fn expr_storage_hint(expr: &Expr, type_tracker: &TypeTracker) -> Option<NativeKind> {
let _ = (expr, type_tracker);
None
}
#[cfg(test)]
mod tests {
use super::*;
use shape_ast::ast::Span;
fn span() -> Span {
Span::default()
}
fn num_lit(v: f64) -> Expr {
Expr::Literal(Literal::Number(v), span())
}
fn int_lit(v: i64) -> Expr {
Expr::Literal(Literal::Int(v), span())
}
fn bool_lit(v: bool) -> Expr {
Expr::Literal(Literal::Bool(v), span())
}
fn string_lit(s: &str) -> Expr {
Expr::Literal(Literal::String(s.to_string()), span())
}
fn typed_int_lit(v: i64, w: shape_ast::IntWidth) -> Expr {
Expr::Literal(Literal::TypedInt(v, w), span())
}
fn ident(name: &str) -> Expr {
Expr::Identifier(name.to_string(), span())
}
fn tracker() -> TypeTracker {
TypeTracker::empty()
}
#[test]
fn test_empty_array_returns_none() {
let tt = tracker();
assert_eq!(infer_array_element_type(&[], &tt), None);
}
#[test]
fn test_all_numbers() {
let tt = tracker();
let elems = vec![num_lit(1.0), num_lit(2.5), num_lit(3.14)];
assert_eq!(
infer_array_element_type(&elems, &tt),
Some(NativeKind::Float64)
);
}
#[test]
fn test_all_ints() {
let tt = tracker();
let elems = vec![int_lit(1), int_lit(2), int_lit(3)];
assert_eq!(
infer_array_element_type(&elems, &tt),
Some(NativeKind::Int64)
);
}
#[test]
fn test_all_bools() {
let tt = tracker();
let elems = vec![bool_lit(true), bool_lit(false), bool_lit(true)];
assert_eq!(
infer_array_element_type(&elems, &tt),
Some(NativeKind::Bool)
);
}
#[test]
fn test_all_strings() {
let tt = tracker();
let elems = vec![string_lit("a"), string_lit("b")];
assert_eq!(
infer_array_element_type(&elems, &tt),
Some(NativeKind::String)
);
}
#[test]
fn test_all_typed_i32() {
let tt = tracker();
let elems = vec![
typed_int_lit(1, shape_ast::IntWidth::I32),
typed_int_lit(2, shape_ast::IntWidth::I32),
];
assert_eq!(
infer_array_element_type(&elems, &tt),
Some(NativeKind::Int32)
);
}
#[test]
fn test_mixed_int_and_number_returns_none() {
let tt = tracker();
let elems = vec![int_lit(1), num_lit(2.0)];
assert_eq!(infer_array_element_type(&elems, &tt), None);
}
#[test]
fn test_mixed_literals_and_identifiers_returns_none() {
let tt = tracker();
let elems = vec![int_lit(1), ident("x")];
assert_eq!(infer_array_element_type(&elems, &tt), None);
}
#[test]
fn test_single_element_number() {
let tt = tracker();
let elems = vec![num_lit(42.0)];
assert_eq!(
infer_array_element_type(&elems, &tt),
Some(NativeKind::Float64)
);
}
#[test]
fn test_mixed_typed_int_widths_returns_none() {
let tt = tracker();
let elems = vec![
typed_int_lit(1, shape_ast::IntWidth::I32),
typed_int_lit(2, shape_ast::IntWidth::U8),
];
assert_eq!(infer_array_element_type(&elems, &tt), None);
}
#[test]
fn test_all_identifiers_without_tracking_returns_none() {
let tt = tracker();
let elems = vec![ident("a"), ident("b")];
assert_eq!(infer_array_element_type(&elems, &tt), None);
}
#[test]
fn test_nested_array_literal_element_returns_none() {
let tt = tracker();
let inner = Expr::Array(vec![num_lit(1.0), num_lit(2.0)], span());
let outer_elems = vec![inner.clone(), inner];
assert_eq!(infer_array_element_type(&outer_elems, &tt), None);
}
#[test]
fn test_mixed_scalar_and_nested_array_returns_none() {
let tt = tracker();
let elems = vec![
num_lit(1.0),
Expr::Array(vec![num_lit(2.0), num_lit(3.0)], span()),
];
assert_eq!(infer_array_element_type(&elems, &tt), None);
}
#[test]
fn test_annotation_array_number() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("number".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::Float64));
}
#[test]
fn test_annotation_array_int() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("int".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::Int64));
}
#[test]
fn test_annotation_array_i32() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("i32".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::Int32));
}
#[test]
fn test_annotation_array_bool() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("bool".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::Bool));
}
#[test]
fn test_annotation_array_string() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("string".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::String));
}
#[test]
fn test_annotation_array_u8() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("u8".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::UInt8));
}
#[test]
fn test_annotation_array_sugar_number() {
let ann = TypeAnnotation::Array(Box::new(TypeAnnotation::Basic("number".to_string())));
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::Float64));
}
#[test]
fn test_annotation_array_sugar_int() {
let ann = TypeAnnotation::Array(Box::new(TypeAnnotation::Basic("int".to_string())));
assert_eq!(typed_array_from_annotation(&ann), Some(NativeKind::Int64));
}
#[test]
fn test_annotation_array_custom_type_returns_none() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("Point".to_string())],
};
assert_eq!(typed_array_from_annotation(&ann), None);
}
#[test]
fn test_annotation_non_array_generic_returns_none() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("HashMap"),
args: vec![
TypeAnnotation::Basic("string".to_string()),
TypeAnnotation::Basic("int".to_string()),
],
};
assert_eq!(typed_array_from_annotation(&ann), None);
}
#[test]
fn test_annotation_basic_type_returns_none() {
let ann = TypeAnnotation::Basic("number".to_string());
assert_eq!(typed_array_from_annotation(&ann), None);
}
#[test]
fn test_annotation_array_nested_generic_returns_none() {
use shape_ast::ast::type_path::TypePath;
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic("int".to_string())],
}],
};
assert_eq!(typed_array_from_annotation(&ann), None);
}
#[test]
fn test_all_scalar_widths_via_annotation() {
use shape_ast::ast::type_path::TypePath;
let cases = vec![
("number", NativeKind::Float64),
("int", NativeKind::Int64),
("i8", NativeKind::Int8),
("u8", NativeKind::UInt8),
("i16", NativeKind::Int16),
("u16", NativeKind::UInt16),
("i32", NativeKind::Int32),
("u32", NativeKind::UInt32),
("u64", NativeKind::UInt64),
("isize", NativeKind::IntSize),
("usize", NativeKind::UIntSize),
("bool", NativeKind::Bool),
("string", NativeKind::String),
];
for (type_name, expected_kind) in cases {
let ann = TypeAnnotation::Generic {
name: TypePath::simple("Array"),
args: vec![TypeAnnotation::Basic(type_name.to_string())],
};
assert_eq!(
typed_array_from_annotation(&ann),
Some(expected_kind),
"Array<{type_name}> should map to {expected_kind:?}"
);
}
}
#[test]
fn test_typed_int_width_mapping() {
use shape_ast::IntWidth;
assert_eq!(typed_int_width_to_slot(IntWidth::I8), NativeKind::Int8);
assert_eq!(typed_int_width_to_slot(IntWidth::U8), NativeKind::UInt8);
assert_eq!(typed_int_width_to_slot(IntWidth::I16), NativeKind::Int16);
assert_eq!(typed_int_width_to_slot(IntWidth::U16), NativeKind::UInt16);
assert_eq!(typed_int_width_to_slot(IntWidth::I32), NativeKind::Int32);
assert_eq!(typed_int_width_to_slot(IntWidth::U32), NativeKind::UInt32);
assert_eq!(typed_int_width_to_slot(IntWidth::U64), NativeKind::UInt64);
}
}