Skip to main content

nu_command/matrix/
value.rs

1use ndarray::ArrayD;
2use nu_protocol::{
3    CellPathMutation, CustomValue, ShellError, Span, Type, Value,
4    ast::{Comparison, Math, Operator, PathMember},
5    casing::Casing,
6};
7use serde::{Deserialize, Serialize};
8use std::any::Any;
9use std::cmp::Ordering;
10
11#[derive(Debug, Clone, Serialize, Deserialize)]
12pub struct MatrixValue {
13    pub array: ArrayD<f64>,
14}
15
16#[typetag::serde]
17impl CustomValue for MatrixValue {
18    fn clone_value(&self, span: Span) -> Value {
19        Value::custom(Box::new(self.clone()), span)
20    }
21
22    fn type_name(&self) -> String {
23        "matrix".to_string()
24    }
25
26    fn to_base_value(&self, span: Span) -> Result<Value, ShellError> {
27        Ok(ndarray_to_value(&self.array, span))
28    }
29
30    fn as_any(&self) -> &dyn Any {
31        self
32    }
33
34    fn as_mut_any(&mut self) -> &mut dyn Any {
35        self
36    }
37
38    fn partial_cmp(&self, other: &Value) -> Option<Ordering> {
39        match other {
40            Value::Custom { val, .. } if val.type_name() == self.type_name() => {
41                let other_matrix = val.as_any().downcast_ref::<MatrixValue>()?;
42                if self.array.shape() != other_matrix.array.shape() {
43                    return None;
44                }
45                if ndarray::Zip::from(&self.array)
46                    .and(&other_matrix.array)
47                    .all(|a, b| a == b)
48                {
49                    Some(Ordering::Equal)
50                } else if ndarray::Zip::from(&self.array)
51                    .and(&other_matrix.array)
52                    .all(|a, b| *a <= *b)
53                {
54                    Some(Ordering::Less)
55                } else if ndarray::Zip::from(&self.array)
56                    .and(&other_matrix.array)
57                    .all(|a, b| *a >= *b)
58                {
59                    Some(Ordering::Greater)
60                } else {
61                    None
62                }
63            }
64            _ => None,
65        }
66    }
67
68    fn follow_path_string(
69        &self,
70        self_span: Span,
71        column_name: String,
72        path_span: Span,
73        _optional: bool,
74        casing: Casing,
75    ) -> Result<Value, ShellError> {
76        let col = match casing {
77            Casing::Sensitive => column_name,
78            Casing::Insensitive => column_name.to_lowercase(),
79        };
80
81        match col.as_str() {
82            "shape" => Ok(Value::list(
83                self.array
84                    .shape()
85                    .iter()
86                    .map(|d| Value::int(*d as i64, path_span))
87                    .collect(),
88                path_span,
89            )),
90            "ndim" => Ok(Value::int(self.array.ndim() as i64, path_span)),
91            "size" => Ok(Value::int(self.array.len() as i64, path_span)),
92            _ => Err(ShellError::CantFindColumn {
93                col_name: col,
94                span: Some(path_span),
95                src_span: self_span,
96            }),
97        }
98    }
99
100    fn follow_path_int(
101        &self,
102        self_span: Span,
103        index: usize,
104        path_span: Span,
105        _optional: bool,
106    ) -> Result<Value, ShellError> {
107        if self.array.ndim() == 0 {
108            return Err(ShellError::IncompatiblePathAccess {
109                type_name: self.type_name(),
110                span: path_span,
111            });
112        }
113        if index >= self.array.shape()[0] {
114            return Err(ShellError::AccessBeyondEnd {
115                max_idx: self.array.shape()[0] - 1,
116                span: self_span,
117            });
118        }
119        let subview = self.array.index_axis(ndarray::Axis(0), index);
120        if subview.ndim() == 0 {
121            Ok(Value::float(*subview.first().unwrap_or(&0.0), path_span))
122        } else {
123            Ok(ndarray_to_value(&subview.to_owned(), path_span))
124        }
125    }
126
127    fn update_data_at_cell_path(
128        &self,
129        cell_path: &[PathMember],
130        new_val: Value,
131        action: &CellPathMutation,
132        head: Span,
133    ) -> Result<Value, ShellError> {
134        let mut base = self.to_base_value(head)?;
135        base.mutate_data_at_cell_path(cell_path, new_val, action)?;
136        match base {
137            Value::List { vals, .. } => {
138                MatrixValue::from_list_of_lists(&vals, head).map(|m| m.into_value(head))
139            }
140            other => Ok(other),
141        }
142    }
143
144    fn operation(
145        &self,
146        lhs_span: Span,
147        operator: Operator,
148        op: Span,
149        right: &Value,
150    ) -> Result<Value, ShellError> {
151        match operator {
152            Operator::Math(Math::Add) => {
153                matrix_math_op(self, right, operator, op, lhs_span, |a, b| a + b)
154            }
155            Operator::Math(Math::Subtract) => {
156                matrix_math_op(self, right, operator, op, lhs_span, |a, b| a - b)
157            }
158            Operator::Math(Math::Multiply) => {
159                matrix_math_op(self, right, operator, op, lhs_span, |a, b| a * b)
160            }
161            Operator::Math(Math::Divide) => {
162                matrix_math_op(self, right, operator, op, lhs_span, |a, b| a / b)
163            }
164            Operator::Comparison(comparison @ Comparison::Equal)
165            | Operator::Comparison(comparison @ Comparison::NotEqual)
166            | Operator::Comparison(comparison @ Comparison::LessThan)
167            | Operator::Comparison(comparison @ Comparison::GreaterThan)
168            | Operator::Comparison(comparison @ Comparison::LessThanOrEqual)
169            | Operator::Comparison(comparison @ Comparison::GreaterThanOrEqual) => {
170                compare_matrix(self, right, op, lhs_span, comparison)
171            }
172            _ => Err(ShellError::OperatorUnsupportedType {
173                op: operator,
174                unsupported: Type::Custom(self.type_name().into()),
175                op_span: op,
176                unsupported_span: lhs_span,
177                help: None,
178            }),
179        }
180    }
181}
182
183impl MatrixValue {
184    pub fn new(array: ArrayD<f64>) -> Self {
185        Self { array }
186    }
187
188    pub fn into_value(self, span: Span) -> Value {
189        Value::custom(Box::new(self), span)
190    }
191
192    pub fn from_value(value: &Value) -> Result<Self, ShellError> {
193        let span = value.span();
194        match value {
195            Value::Custom { val, .. } => {
196                val.as_any().downcast_ref::<Self>().cloned().ok_or_else(|| {
197                    ShellError::CantConvert {
198                        to_type: "matrix".into(),
199                        from_type: val.type_name(),
200                        span,
201                        help: Some("expected a matrix value".into()),
202                    }
203                })
204            }
205            x => Err(ShellError::CantConvert {
206                to_type: "matrix".into(),
207                from_type: x.get_type().to_string(),
208                span,
209                help: None,
210            }),
211        }
212    }
213
214    pub fn from_shape_vec(
215        shape: Vec<usize>,
216        data: Vec<f64>,
217        span: Span,
218    ) -> Result<Self, ShellError> {
219        ArrayD::from_shape_vec(shape, data)
220            .map(Self::new)
221            .map_err(|e| {
222                ShellError::Generic(nu_protocol::shell_error::generic::GenericError::new(
223                    "Matrix shape error",
224                    e.to_string(),
225                    span,
226                ))
227            })
228    }
229
230    pub fn from_list_of_lists(values: &[Value], span: Span) -> Result<Self, ShellError> {
231        let mut rows: Vec<Vec<f64>> = Vec::new();
232        let mut ncols: Option<usize> = None;
233
234        for (i, value) in values.iter().enumerate() {
235            match value {
236                Value::List { vals, .. } => {
237                    let row: Result<Vec<f64>, ShellError> =
238                        vals.iter().map(|v| value_to_f64(v, span)).collect();
239                    let row = row?;
240                    if let Some(expected) = ncols {
241                        if row.len() != expected {
242                            return Err(ShellError::Generic(
243                                nu_protocol::shell_error::generic::GenericError::new(
244                                    "Inconsistent row lengths",
245                                    format!(
246                                        "row {} has {} elements, expected {}",
247                                        i,
248                                        row.len(),
249                                        expected
250                                    ),
251                                    span,
252                                ),
253                            ));
254                        }
255                    } else {
256                        ncols = Some(row.len());
257                    }
258                    rows.push(row);
259                }
260                _ => {
261                    return Err(ShellError::Generic(
262                        nu_protocol::shell_error::generic::GenericError::new(
263                            "Invalid matrix input",
264                            format!("row {} is not a list", i),
265                            span,
266                        ),
267                    ));
268                }
269            }
270        }
271
272        let nrows = rows.len();
273        let ncols = ncols.unwrap_or(0);
274        let flat: Vec<f64> = rows.into_iter().flatten().collect();
275
276        Self::from_shape_vec(vec![nrows, ncols], flat, span)
277    }
278
279    pub fn from_list_of_records(values: &[Value], span: Span) -> Result<Self, ShellError> {
280        if values.is_empty() {
281            let array = ArrayD::from_shape_vec(vec![0, 0], vec![]).map_err(|e| {
282                ShellError::Generic(nu_protocol::shell_error::generic::GenericError::new(
283                    "Matrix shape error",
284                    e.to_string(),
285                    span,
286                ))
287            })?;
288            return Ok(Self::new(array));
289        }
290
291        let first_record = match &values[0] {
292            Value::Record { val, .. } => val,
293            _ => {
294                return Err(ShellError::Generic(
295                    nu_protocol::shell_error::generic::GenericError::new(
296                        "Invalid matrix input",
297                        "expected a list of records",
298                        span,
299                    ),
300                ));
301            }
302        };
303
304        let cols: Vec<String> = first_record.columns().cloned().collect();
305        let ncols = cols.len();
306        let nrows = values.len();
307
308        let mut data = Vec::with_capacity(nrows * ncols);
309
310        for (i, value) in values.iter().enumerate() {
311            match value {
312                Value::Record { val, .. } => {
313                    for col in &cols {
314                        let element = val.get(col).ok_or_else(|| {
315                            ShellError::Generic(
316                                nu_protocol::shell_error::generic::GenericError::new(
317                                    "Missing column",
318                                    format!("row {} is missing column '{}'", i, col),
319                                    span,
320                                ),
321                            )
322                        })?;
323                        data.push(value_to_f64(element, span)?);
324                    }
325                }
326                _ => {
327                    return Err(ShellError::Generic(
328                        nu_protocol::shell_error::generic::GenericError::new(
329                            "Invalid matrix input",
330                            format!("row {} is not a record", i),
331                            span,
332                        ),
333                    ));
334                }
335            }
336        }
337
338        Self::from_shape_vec(vec![nrows, ncols], data, span)
339    }
340
341    /// Apply an element-wise binary operation with another matrix or scalar.
342    ///
343    /// When `broadcast` is true, `other` may be broadcast to this matrix's shape.
344    pub fn elementwise_binary<F, G>(
345        self,
346        other: Value,
347        broadcast: bool,
348        head: Span,
349        f_matrix: F,
350        f_scalar: G,
351    ) -> Result<ArrayD<f64>, ShellError>
352    where
353        F: FnOnce(ArrayD<f64>, ArrayD<f64>) -> ArrayD<f64>,
354        G: FnOnce(ArrayD<f64>, f64) -> ArrayD<f64>,
355    {
356        match other {
357            Value::Int { val, .. } => Ok(f_scalar(self.array, val as f64)),
358            Value::Float { val, .. } => Ok(f_scalar(self.array, val)),
359            Value::Custom { .. } => {
360                let other_matrix = MatrixValue::from_value(&other)?;
361                if broadcast {
362                    let target_shape = self.array.shape().to_vec();
363                    let other_view = other_matrix
364                        .array
365                        .broadcast(target_shape.as_slice())
366                        .ok_or_else(|| {
367                            ShellError::Generic(
368                                nu_protocol::shell_error::generic::GenericError::new(
369                                    "Broadcast error",
370                                    "shapes are not compatible for broadcasting",
371                                    head,
372                                ),
373                            )
374                        })?;
375                    Ok(f_matrix(self.array, other_view.to_owned().into_dyn()))
376                } else if self.array.shape() == other_matrix.array.shape() {
377                    Ok(f_matrix(self.array, other_matrix.array))
378                } else {
379                    Err(ShellError::Generic(
380                        nu_protocol::shell_error::generic::GenericError::new(
381                            "Shape mismatch",
382                            format!(
383                                "shapes do not match: {:?} vs {:?}. Use --broadcast to enable broadcasting.",
384                                self.array.shape(),
385                                other_matrix.array.shape()
386                            ),
387                            head,
388                        ),
389                    ))
390                }
391            }
392            _ => Err(ShellError::Generic(
393                nu_protocol::shell_error::generic::GenericError::new(
394                    "Invalid argument",
395                    "expected a matrix, int, or float",
396                    head,
397                ),
398            )),
399        }
400    }
401
402    pub fn test_value(rows: &[&[f64]]) -> Value {
403        let nrows = rows.len();
404        let ncols = if nrows > 0 { rows[0].len() } else { 0 };
405        let flat: Vec<f64> = rows.iter().flat_map(|r| r.iter()).copied().collect();
406        let array = ArrayD::from_shape_vec(vec![nrows, ncols], flat)
407            .expect("test value shape must be valid");
408        Value::test_custom_value(Box::new(Self { array }))
409    }
410}
411
412/// Convert a numeric Value to f64.
413pub(crate) fn value_to_f64(value: &Value, span: Span) -> Result<f64, ShellError> {
414    match value {
415        Value::Int { val, .. } => Ok(*val as f64),
416        Value::Float { val, .. } => Ok(*val),
417        Value::String { val, .. } => val.parse::<f64>().map_err(|_| ShellError::CantConvert {
418            to_type: "float".into(),
419            from_type: "string".into(),
420            span,
421            help: None,
422        }),
423        _ => Err(ShellError::CantConvert {
424            to_type: "float".into(),
425            from_type: value.get_type().to_string(),
426            span,
427            help: None,
428        }),
429    }
430}
431
432/// Convert a list of numeric Values to `Vec<f64>`.
433pub(crate) fn values_to_f64s(vals: &[Value], span: Span) -> Result<Vec<f64>, ShellError> {
434    vals.iter().map(|v| value_to_f64(v, span)).collect()
435}
436
437/// Parse positive dimensions from i64 values, rejecting zero and negatives.
438pub(crate) fn positive_dim(dim: i64, span: Span) -> Result<usize, ShellError> {
439    if dim > 0 {
440        Ok(dim as usize)
441    } else {
442        Err(ShellError::Generic(
443            nu_protocol::shell_error::generic::GenericError::new(
444                "Invalid dimensions",
445                "dimensions must be positive integers",
446                span,
447            ),
448        ))
449    }
450}
451
452fn ndarray_to_value(array: &ArrayD<f64>, span: Span) -> Value {
453    if array.ndim() == 0 {
454        return Value::float(array.first().copied().unwrap_or(0.0), span);
455    }
456
457    if array.ndim() == 1 {
458        let list: Vec<Value> = array.iter().map(|v| Value::float(*v, span)).collect();
459        return Value::list(list, span);
460    }
461
462    if array.ndim() == 2 {
463        let rows: Vec<Value> = array
464            .axis_iter(ndarray::Axis(0))
465            .map(|row| {
466                let vals: Vec<Value> = row.iter().map(|v| Value::float(*v, span)).collect();
467                Value::list(vals, span)
468            })
469            .collect();
470        return Value::list(rows, span);
471    }
472
473    let sub_results: Vec<Value> = array
474        .axis_iter(ndarray::Axis(0))
475        .map(|sub| ndarray_to_value(&sub.to_owned(), span))
476        .collect();
477    Value::list(sub_results, span)
478}
479
480fn matrix_math_op<F>(
481    left: &MatrixValue,
482    right: &Value,
483    operator: Operator,
484    op_span: Span,
485    lhs_span: Span,
486    f: F,
487) -> Result<Value, ShellError>
488where
489    F: Fn(f64, f64) -> f64,
490{
491    match right {
492        Value::Int { val, .. } => {
493            let result = left.array.map(|v| f(*v, *val as f64));
494            Ok(MatrixValue::new(result).into_value(op_span))
495        }
496        Value::Float { val, .. } => {
497            let result = left.array.map(|v| f(*v, *val));
498            Ok(MatrixValue::new(result).into_value(op_span))
499        }
500        Value::Custom { val, .. } => {
501            let other = val.as_any().downcast_ref::<MatrixValue>().ok_or_else(|| {
502                ShellError::OperatorIncompatibleTypes {
503                    op: operator,
504                    lhs: Type::Custom("matrix".into()),
505                    rhs: Type::Custom(val.type_name().into()),
506                    op_span,
507                    lhs_span,
508                    rhs_span: right.span(),
509                    help: None,
510                }
511            })?;
512            if left.array.shape() != other.array.shape() {
513                return Err(ShellError::OperatorIncompatibleTypes {
514                    op: operator,
515                    lhs: Type::Custom("matrix".into()),
516                    rhs: Type::Custom("matrix".into()),
517                    op_span,
518                    lhs_span,
519                    rhs_span: right.span(),
520                    help: Some("shapes do not match"),
521                });
522            }
523            let shape: Vec<usize> = left.array.shape().to_vec();
524            let mut result = ArrayD::zeros(shape);
525            ndarray::Zip::from(&mut result)
526                .and(&left.array)
527                .and(&other.array)
528                .for_each(|r, &a, &b| *r = f(a, b));
529            Ok(MatrixValue::new(result).into_value(op_span))
530        }
531        _ => Err(ShellError::OperatorIncompatibleTypes {
532            op: operator,
533            lhs: Type::Custom("matrix".into()),
534            rhs: right.get_type(),
535            op_span,
536            lhs_span,
537            rhs_span: right.span(),
538            help: Some("expected a matrix or scalar"),
539        }),
540    }
541}
542
543fn compare_matrix(
544    left: &MatrixValue,
545    right: &Value,
546    op_span: Span,
547    lhs_span: Span,
548    comparison: Comparison,
549) -> Result<Value, ShellError> {
550    let op = Operator::Comparison(comparison);
551
552    match right {
553        Value::Custom { val, .. } if val.type_name() == "matrix" => {
554            let other = val.as_any().downcast_ref::<MatrixValue>().ok_or_else(|| {
555                ShellError::OperatorIncompatibleTypes {
556                    op,
557                    lhs: Type::Custom("matrix".into()),
558                    rhs: Type::Custom(val.type_name().into()),
559                    op_span,
560                    lhs_span,
561                    rhs_span: right.span(),
562                    help: None,
563                }
564            })?;
565
566            if left.array.shape() != other.array.shape() {
567                return shape_mismatch_comparison(comparison, op, op_span, lhs_span, right.span());
568            }
569
570            let all_match = match comparison {
571                Comparison::Equal => ndarray::Zip::from(&left.array)
572                    .and(&other.array)
573                    .all(|a, b| a == b),
574                Comparison::NotEqual => !ndarray::Zip::from(&left.array)
575                    .and(&other.array)
576                    .all(|a, b| a == b),
577                Comparison::LessThan => ndarray::Zip::from(&left.array)
578                    .and(&other.array)
579                    .all(|a, b| a < b),
580                Comparison::GreaterThan => ndarray::Zip::from(&left.array)
581                    .and(&other.array)
582                    .all(|a, b| a > b),
583                Comparison::LessThanOrEqual => ndarray::Zip::from(&left.array)
584                    .and(&other.array)
585                    .all(|a, b| a <= b),
586                Comparison::GreaterThanOrEqual => ndarray::Zip::from(&left.array)
587                    .and(&other.array)
588                    .all(|a, b| a >= b),
589                _ => {
590                    return Err(ShellError::OperatorUnsupportedType {
591                        op,
592                        unsupported: Type::Custom("matrix".into()),
593                        op_span,
594                        unsupported_span: lhs_span,
595                        help: None,
596                    });
597                }
598            };
599
600            Ok(Value::bool(all_match, op_span))
601        }
602        Value::Int { val, .. } => {
603            let s = *val as f64;
604            Ok(Value::bool(
605                compare_all_to_scalar(&left.array, s, comparison, op, op_span, lhs_span)?,
606                op_span,
607            ))
608        }
609        Value::Float { val, .. } => Ok(Value::bool(
610            compare_all_to_scalar(&left.array, *val, comparison, op, op_span, lhs_span)?,
611            op_span,
612        )),
613        _ => Err(ShellError::OperatorIncompatibleTypes {
614            op,
615            lhs: Type::Custom("matrix".into()),
616            rhs: right.get_type(),
617            op_span,
618            lhs_span,
619            rhs_span: right.span(),
620            help: Some("expected a matrix or numeric scalar"),
621        }),
622    }
623}
624
625/// Equality on mismatched shapes is false / not-equal is true; ordering requires matching shapes.
626fn shape_mismatch_comparison(
627    comparison: Comparison,
628    op: Operator,
629    op_span: Span,
630    lhs_span: Span,
631    rhs_span: Span,
632) -> Result<Value, ShellError> {
633    match comparison {
634        Comparison::Equal => Ok(Value::bool(false, op_span)),
635        Comparison::NotEqual => Ok(Value::bool(true, op_span)),
636        Comparison::LessThan
637        | Comparison::GreaterThan
638        | Comparison::LessThanOrEqual
639        | Comparison::GreaterThanOrEqual => Err(ShellError::OperatorIncompatibleTypes {
640            op,
641            lhs: Type::Custom("matrix".into()),
642            rhs: Type::Custom("matrix".into()),
643            op_span,
644            lhs_span,
645            rhs_span,
646            help: Some("cannot compare matrices with different shapes"),
647        }),
648        _ => Err(ShellError::OperatorUnsupportedType {
649            op,
650            unsupported: Type::Custom("matrix".into()),
651            op_span,
652            unsupported_span: lhs_span,
653            help: None,
654        }),
655    }
656}
657
658fn compare_all_to_scalar(
659    array: &ArrayD<f64>,
660    scalar: f64,
661    comparison: Comparison,
662    op: Operator,
663    op_span: Span,
664    lhs_span: Span,
665) -> Result<bool, ShellError> {
666    match comparison {
667        Comparison::Equal => Ok(array.iter().all(|v| (*v - scalar).abs() < f64::EPSILON)),
668        Comparison::NotEqual => Ok(array.iter().any(|v| (*v - scalar).abs() >= f64::EPSILON)),
669        Comparison::LessThan => Ok(array.iter().all(|&v| v < scalar)),
670        Comparison::GreaterThan => Ok(array.iter().all(|&v| v > scalar)),
671        Comparison::LessThanOrEqual => Ok(array.iter().all(|&v| v <= scalar)),
672        Comparison::GreaterThanOrEqual => Ok(array.iter().all(|&v| v >= scalar)),
673        _ => Err(ShellError::OperatorUnsupportedType {
674            op,
675            unsupported: Type::Custom("matrix".into()),
676            op_span,
677            unsupported_span: lhs_span,
678            help: None,
679        }),
680    }
681}