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 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
412pub(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
432pub(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
437pub(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
625fn 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}