use crate::lex::error::LexError;
const MAX_RECURSION_DEPTH: usize = 100;
const MAX_TENSOR_ELEMENTS: usize = 10_000_000;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum Tensor {
Scalar(f64),
Array(Vec<Tensor>),
}
impl std::fmt::Display for Tensor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Tensor::Scalar(n) => {
if n.fract() == 0.0 && n.is_finite() {
if *n >= i64::MIN as f64 && *n <= i64::MAX as f64 {
write!(f, "{}", *n as i64)
} else {
write!(f, "{}", n)
}
} else {
write!(f, "{}", n)
}
}
Tensor::Array(items) => {
write!(f, "[")?;
for (i, item) in items.iter().enumerate() {
if i > 0 {
write!(f, ", ")?;
}
write!(f, "{}", item)?;
}
write!(f, "]")
}
}
}
}
impl Tensor {
pub fn is_integer(&self) -> bool {
match self {
Tensor::Scalar(n) => n.fract() == 0.0,
Tensor::Array(items) => items.iter().all(|t| t.is_integer()),
}
}
pub fn shape(&self) -> Vec<usize> {
match self {
Tensor::Scalar(_) => vec![],
Tensor::Array(items) => {
if items.is_empty() {
vec![0]
} else {
let mut shape = vec![items.len()];
shape.extend(items[0].shape());
shape
}
}
}
}
pub fn flatten(&self) -> Vec<f64> {
let capacity = self.count_elements();
let mut result = Vec::with_capacity(capacity);
self.flatten_into(&mut result);
result
}
fn count_elements(&self) -> usize {
match self {
Tensor::Scalar(_) => 1,
Tensor::Array(items) => items.iter().map(|t| t.count_elements()).sum(),
}
}
fn flatten_into(&self, result: &mut Vec<f64>) {
match self {
Tensor::Scalar(n) => result.push(*n),
Tensor::Array(items) => {
for item in items {
item.flatten_into(result);
}
}
}
}
#[inline]
pub fn is_scalar(&self) -> bool {
matches!(self, Tensor::Scalar(_))
}
#[inline]
pub fn is_array(&self) -> bool {
matches!(self, Tensor::Array(_))
}
#[inline]
pub fn ndim(&self) -> usize {
self.shape().len()
}
#[inline]
pub fn len(&self) -> usize {
self.count_elements()
}
#[inline]
pub fn is_empty(&self) -> bool {
match self {
Tensor::Scalar(_) => false,
Tensor::Array(items) => items.is_empty(),
}
}
}
#[inline]
pub fn is_tensor_literal(s: &str) -> bool {
let s = s.trim();
let bytes = s.as_bytes();
if bytes.first() != Some(&b'[') || bytes.last() != Some(&b']') {
return false;
}
let mut depth: i32 = 0;
for &b in bytes {
match b {
b'[' => depth += 1,
b']' => {
depth -= 1;
if depth < 0 {
return false;
}
}
b'0'..=b'9' | b'.' | b'-' | b',' | b' ' | b'\t' => {}
_ => return false,
}
}
depth == 0
}
pub fn parse_tensor(s: &str) -> Result<Tensor, LexError> {
let s = s.trim();
if !s.starts_with('[') {
return Err(LexError::UnexpectedChar(s.chars().next().unwrap_or(' ')));
}
let (tensor, remaining) = parse_tensor_inner(s, 0)?;
if !remaining.trim().is_empty() {
return Err(LexError::UnexpectedChar(
remaining.trim().chars().next().unwrap_or('?'),
));
}
let element_count = tensor.count_elements();
if element_count > MAX_TENSOR_ELEMENTS {
return Err(LexError::InvalidStructure(format!(
"Tensor element count exceeds maximum: {} (max: {})",
element_count, MAX_TENSOR_ELEMENTS
)));
}
Ok(tensor)
}
fn estimate_array_size(s: &str) -> usize {
let mut depth = 0;
let mut comma_count = 0;
for ch in s.chars() {
match ch {
'[' => depth += 1,
']' => {
if depth == 0 {
break;
}
depth -= 1;
}
',' if depth == 0 => comma_count += 1,
_ => {}
}
}
if comma_count > 0 {
comma_count + 1
} else {
2
}
}
fn parse_tensor_inner(s: &str, depth: usize) -> Result<(Tensor, &str), LexError> {
if depth > MAX_RECURSION_DEPTH {
return Err(LexError::InvalidStructure(format!(
"Recursion depth exceeded (max: {})",
MAX_RECURSION_DEPTH
)));
}
let s = s.trim();
if let Some(remaining_str) = s.strip_prefix('[') {
let mut remaining = remaining_str;
let estimated_capacity = estimate_array_size(remaining_str);
let mut items = Vec::with_capacity(estimated_capacity);
loop {
remaining = remaining.trim_start();
if remaining.is_empty() {
return Err(LexError::UnbalancedBrackets);
}
if remaining.starts_with(']') {
remaining = &remaining[1..];
break;
}
if !items.is_empty() {
if !remaining.starts_with(',') {
return Err(LexError::UnexpectedChar(
remaining.chars().next().unwrap_or('?'),
));
}
remaining = remaining[1..].trim_start();
}
if remaining.starts_with(']') {
remaining = &remaining[1..];
break;
}
if remaining.starts_with('[') {
let (tensor, rest) = parse_tensor_inner(remaining, depth + 1)?;
items.push(tensor);
remaining = rest;
} else {
let (num, rest) = parse_number(remaining)?;
items.push(Tensor::Scalar(num));
remaining = rest;
}
}
if items.is_empty() {
return Err(LexError::EmptyTensor);
}
if items.len() > 1 {
let first_shape = items[0].shape();
for item in &items[1..] {
if item.shape() != first_shape {
return Err(LexError::InconsistentDimensions);
}
}
}
Ok((Tensor::Array(items), remaining))
} else {
let (num, rest) = parse_number(s)?;
Ok((Tensor::Scalar(num), rest))
}
}
fn parse_number(s: &str) -> Result<(f64, &str), LexError> {
let s = s.trim_start();
let bytes = s.as_bytes();
let mut end = 0;
let mut has_dot = false;
if bytes.first() == Some(&b'-') {
end = 1;
}
while end < bytes.len() {
match bytes[end] {
b'0'..=b'9' => end += 1,
b'.' if !has_dot => {
has_dot = true;
end += 1;
}
_ => break,
}
}
if end == 0 || (end == 1 && bytes[0] == b'-') {
return Err(LexError::InvalidNumber(
s.chars().take(10).collect::<String>(),
));
}
let num_str = &s[..end];
let num: f64 = num_str.parse().map_err(|_| {
let context = if num_str.len() > 80 {
format!("{}...", &num_str[..80])
} else {
num_str.to_string()
};
LexError::InvalidNumber(context)
})?;
if !num.is_finite() {
return Err(LexError::InvalidNumber(format!(
"{} (non-finite values not allowed)",
num_str
)));
}
Ok((num, &s[end..]))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_tensor_literal_valid() {
assert!(is_tensor_literal("[1, 2, 3]"));
assert!(is_tensor_literal("[[1, 2], [3, 4]]"));
assert!(is_tensor_literal("[1.5, 2.5]"));
assert!(is_tensor_literal("[-1, -2]"));
assert!(is_tensor_literal(" [1, 2, 3] "));
}
#[test]
fn test_is_tensor_literal_invalid() {
assert!(!is_tensor_literal("hello"));
assert!(!is_tensor_literal("@reference"));
assert!(!is_tensor_literal("123"));
assert!(!is_tensor_literal(""));
assert!(!is_tensor_literal("[1, 2"));
assert!(!is_tensor_literal("[a, b]"));
}
#[test]
fn test_parse_1d() {
let t = parse_tensor("[1, 2, 3]").unwrap();
assert_eq!(t.shape(), vec![3]);
assert_eq!(t.flatten(), vec![1.0, 2.0, 3.0]);
}
#[test]
fn test_parse_2d() {
let t = parse_tensor("[[1, 2], [3, 4]]").unwrap();
assert_eq!(t.shape(), vec![2, 2]);
assert_eq!(t.flatten(), vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn test_parse_floats() {
let t = parse_tensor("[1.5, 2.5, 3.5]").unwrap();
assert_eq!(t.flatten(), vec![1.5, 2.5, 3.5]);
assert!(!t.is_integer());
}
#[test]
fn test_parse_negatives() {
let t = parse_tensor("[-1, -2, -3]").unwrap();
assert_eq!(t.flatten(), vec![-1.0, -2.0, -3.0]);
}
#[test]
fn test_parse_trailing_comma() {
let t = parse_tensor("[1, 2, 3,]").unwrap();
assert_eq!(t.flatten(), vec![1.0, 2.0, 3.0]);
}
#[test]
fn test_empty_tensor_error() {
assert!(matches!(parse_tensor("[]"), Err(LexError::EmptyTensor)));
}
#[test]
fn test_unbalanced_brackets_error() {
assert!(matches!(
parse_tensor("[1, 2"),
Err(LexError::UnbalancedBrackets)
));
}
#[test]
fn test_inconsistent_dimensions_error() {
assert!(matches!(
parse_tensor("[[1, 2], [3]]"),
Err(LexError::InconsistentDimensions)
));
}
#[test]
fn test_invalid_number_error() {
assert!(matches!(
parse_tensor("[abc]"),
Err(LexError::InvalidNumber(_))
));
}
#[test]
fn test_tensor_display() {
let t = parse_tensor("[1, 2, 3]").unwrap();
assert_eq!(format!("{}", t), "[1, 2, 3]");
let t = parse_tensor("[[1, 2], [3, 4]]").unwrap();
assert_eq!(format!("{}", t), "[[1, 2], [3, 4]]");
}
#[test]
fn test_tensor_methods() {
let scalar = Tensor::Scalar(42.0);
assert!(scalar.is_scalar());
assert!(!scalar.is_array());
assert_eq!(scalar.ndim(), 0);
assert_eq!(scalar.len(), 1);
assert!(!scalar.is_empty());
let array = parse_tensor("[1, 2, 3]").unwrap();
assert!(!array.is_scalar());
assert!(array.is_array());
assert_eq!(array.ndim(), 1);
assert_eq!(array.len(), 3);
assert!(!array.is_empty());
}
#[test]
fn test_tensor_equality() {
let t1 = parse_tensor("[1, 2, 3]").unwrap();
let t2 = parse_tensor("[1, 2, 3]").unwrap();
assert_eq!(t1, t2);
let t3 = parse_tensor("[1, 2, 4]").unwrap();
assert_ne!(t1, t3);
}
#[test]
fn test_tensor_clone() {
let t1 = parse_tensor("[[1, 2], [3, 4]]").unwrap();
let t2 = t1.clone();
assert_eq!(t1, t2);
}
#[test]
fn test_display_roundtrip() {
let original = parse_tensor("[1, 2, 3]").unwrap();
let serialized = original.to_string();
let parsed = parse_tensor(&serialized).unwrap();
assert_eq!(original, parsed);
let original = parse_tensor("[[1.5, 2.5], [3.5, 4.5]]").unwrap();
let serialized = original.to_string();
let parsed = parse_tensor(&serialized).unwrap();
assert_eq!(original, parsed);
}
}