Skip to main content

runmat_runtime/object/
indexing.rs

1use crate::call::identity::strict_callable_display_name;
2use crate::indexing::EndExpr;
3use crate::runtime_error::semantic_error;
4use crate::RuntimeError;
5use runmat_value::Value;
6
7pub const OBJECT_PROTOCOL_SUBSREF: &str = crate::OBJECT_SUBSREF_METHOD;
8pub const OBJECT_PROTOCOL_SUBSASGN: &str = crate::OBJECT_SUBSASGN_METHOD;
9pub const OBJECT_PROTOCOL_KIND_PAREN: &str = crate::OBJECT_INDEX_PAREN;
10pub const OBJECT_PROTOCOL_KIND_BRACE: &str = crate::OBJECT_INDEX_BRACE;
11pub const OBJECT_PROTOCOL_KIND_MEMBER: &str = crate::OBJECT_INDEX_MEMBER;
12pub const OBJECT_SELECTOR_COLON: &str = ":";
13pub const OBJECT_SELECTOR_END: &str = "end";
14pub const OBJECT_END_RANGE_TAG: &str = "end_expr";
15
16#[derive(Clone, Copy)]
17pub enum ObjectIndexOp {
18    Subsref,
19    Subsasgn,
20}
21
22impl ObjectIndexOp {
23    pub fn protocol_name(self) -> &'static str {
24        match self {
25            Self::Subsref => OBJECT_PROTOCOL_SUBSREF,
26            Self::Subsasgn => OBJECT_PROTOCOL_SUBSASGN,
27        }
28    }
29}
30
31#[derive(Clone, Copy)]
32pub enum ObjectIndexKind {
33    Paren,
34    Brace,
35    Member,
36}
37
38impl ObjectIndexKind {
39    pub fn protocol_name(self) -> &'static str {
40        match self {
41            Self::Paren => OBJECT_PROTOCOL_KIND_PAREN,
42            Self::Brace => OBJECT_PROTOCOL_KIND_BRACE,
43            Self::Member => OBJECT_PROTOCOL_KIND_MEMBER,
44        }
45    }
46}
47
48#[derive(Clone)]
49pub enum ObjectIndexSelector {
50    ScalarIndices { indices: Vec<usize> },
51    IndexValues { values: Vec<Value> },
52    Member(String),
53}
54
55#[derive(Clone)]
56pub struct ObjectIndexDescriptor {
57    base: Value,
58    op: ObjectIndexOp,
59    kind: ObjectIndexKind,
60    selector: ObjectIndexSelector,
61    rhs: Option<Value>,
62}
63
64#[derive(Debug, Clone, Copy)]
65pub struct ObjectParenExprSelectorSpec<'a> {
66    pub dims: usize,
67    pub colon_mask: u32,
68    pub end_mask: u32,
69    pub range_dims: &'a [usize],
70    pub range_params: &'a [(f64, f64)],
71    pub range_start_exprs: &'a [Option<EndExpr>],
72    pub range_step_exprs: &'a [Option<EndExpr>],
73    pub range_end_exprs: &'a [EndExpr],
74    pub end_numeric_exprs: &'a [(usize, EndExpr)],
75    pub numeric: &'a [Value],
76}
77
78impl ObjectIndexDescriptor {
79    pub fn subsref_paren(base: Value, selector: ObjectIndexSelector) -> Self {
80        Self {
81            base,
82            op: ObjectIndexOp::Subsref,
83            kind: ObjectIndexKind::Paren,
84            selector,
85            rhs: None,
86        }
87    }
88
89    pub fn subsref_brace(base: Value, selector: ObjectIndexSelector) -> Self {
90        Self {
91            base,
92            op: ObjectIndexOp::Subsref,
93            kind: ObjectIndexKind::Brace,
94            selector,
95            rhs: None,
96        }
97    }
98
99    pub fn subsasgn_paren(base: Value, selector: ObjectIndexSelector, rhs: Value) -> Self {
100        Self {
101            base,
102            op: ObjectIndexOp::Subsasgn,
103            kind: ObjectIndexKind::Paren,
104            selector,
105            rhs: Some(rhs),
106        }
107    }
108
109    pub fn subsasgn_brace(base: Value, selector: ObjectIndexSelector, rhs: Value) -> Self {
110        Self {
111            base,
112            op: ObjectIndexOp::Subsasgn,
113            kind: ObjectIndexKind::Brace,
114            selector,
115            rhs: Some(rhs),
116        }
117    }
118
119    pub fn subsref_paren_from_slice(
120        base: Value,
121        dims: usize,
122        colon_mask: u32,
123        end_mask: u32,
124        numeric: &[Value],
125    ) -> Result<Self, RuntimeError> {
126        let values = build_object_paren_selector_values(dims, colon_mask, end_mask, numeric)?;
127        Ok(Self::subsref_paren(
128            base,
129            ObjectIndexSelector::IndexValues { values },
130        ))
131    }
132
133    pub fn subsasgn_paren_from_slice(
134        base: Value,
135        dims: usize,
136        colon_mask: u32,
137        end_mask: u32,
138        numeric: &[Value],
139        rhs: Value,
140    ) -> Result<Self, RuntimeError> {
141        let values = build_object_paren_selector_values(dims, colon_mask, end_mask, numeric)?;
142        Ok(Self::subsasgn_paren(
143            base,
144            ObjectIndexSelector::IndexValues { values },
145            rhs,
146        ))
147    }
148
149    pub fn subsasgn_paren_from_expr_slice(
150        base: Value,
151        spec: ObjectParenExprSelectorSpec<'_>,
152        rhs: Value,
153    ) -> Result<Self, RuntimeError> {
154        let values = build_object_paren_expr_selector_values(spec)?;
155        Ok(Self::subsasgn_paren(
156            base,
157            ObjectIndexSelector::IndexValues { values },
158            rhs,
159        ))
160    }
161
162    pub fn subsref_paren_from_expr_slice(
163        base: Value,
164        spec: ObjectParenExprSelectorSpec<'_>,
165    ) -> Result<Self, RuntimeError> {
166        let values = build_object_paren_expr_selector_values(spec)?;
167        Ok(Self::subsref_paren(
168            base,
169            ObjectIndexSelector::IndexValues { values },
170        ))
171    }
172
173    pub fn member(base: Value, op: ObjectIndexOp, field: String, rhs: Option<Value>) -> Self {
174        Self {
175            base,
176            op,
177            kind: ObjectIndexKind::Member,
178            selector: ObjectIndexSelector::Member(field),
179            rhs,
180        }
181    }
182
183    pub fn base(&self) -> &Value {
184        &self.base
185    }
186
187    pub fn operation(&self) -> ObjectIndexOp {
188        self.op
189    }
190
191    pub fn rhs(&self) -> Option<&Value> {
192        self.rhs.as_ref()
193    }
194
195    pub fn into_method_invocation(self) -> Result<(Value, String, Vec<Value>), RuntimeError> {
196        let selector = match self.selector {
197            ObjectIndexSelector::ScalarIndices { indices } => {
198                let values = indices
199                    .into_iter()
200                    .map(|index| Value::Num(index as f64))
201                    .collect();
202                build_protocol_index_cell(values)?
203            }
204            ObjectIndexSelector::IndexValues { values } => build_protocol_index_cell(values)?,
205            ObjectIndexSelector::Member(field) => Value::String(field),
206        };
207        let mut args = vec![
208            Value::String(self.kind.protocol_name().to_string()),
209            selector,
210        ];
211        if let Some(rhs) = self.rhs {
212            args.push(rhs);
213        }
214        Ok((self.base, self.op.protocol_name().to_string(), args))
215    }
216}
217
218fn build_protocol_index_cell(values: Vec<Value>) -> Result<Value, RuntimeError> {
219    let cols = values.len();
220    let cell = build_cell_array_with_shape(values, 1, cols, "object index descriptor build")?;
221    Ok(Value::Cell(cell))
222}
223
224fn matlab_index_type(kind: ObjectIndexKind) -> &'static str {
225    match kind {
226        ObjectIndexKind::Paren => "()",
227        ObjectIndexKind::Brace => "{}",
228        ObjectIndexKind::Member => ".",
229    }
230}
231
232pub fn class_name_from_base(base: &Value) -> Option<&str> {
233    match base {
234        Value::Object(obj) => Some(obj.class_name.as_str()),
235        Value::HandleObject(handle) => Some(handle.class_name.as_str()),
236        _ => None,
237    }
238}
239
240pub fn build_matlab_substruct_arg(
241    descriptor: &ObjectIndexDescriptor,
242) -> Result<Value, RuntimeError> {
243    let subs_value = match &descriptor.selector {
244        ObjectIndexSelector::ScalarIndices { indices } => {
245            let values = indices
246                .iter()
247                .map(|index| Value::Num(*index as f64))
248                .collect();
249            build_protocol_index_cell(values)?
250        }
251        ObjectIndexSelector::IndexValues { values } => build_protocol_index_cell(values.clone())?,
252        ObjectIndexSelector::Member(field) => Value::String(field.clone()),
253    };
254    let mut value = runmat_value::StructValue::new();
255    value.fields.insert(
256        "type".to_string(),
257        Value::String(matlab_index_type(descriptor.kind).to_string()),
258    );
259    value.fields.insert("subs".to_string(), subs_value);
260    Ok(Value::Struct(value))
261}
262
263fn encode_end_expr_value(expr: &EndExpr) -> Result<Value, RuntimeError> {
264    fn mk_cell(items: Vec<Value>) -> Result<Value, RuntimeError> {
265        let cols = items.len();
266        let cell = build_cell_array_with_shape(items, 1, cols, "end expression encoding")?;
267        Ok(Value::Cell(cell))
268    }
269
270    match expr {
271        EndExpr::End => Ok(Value::String("end".to_string())),
272        EndExpr::Const(v) => Ok(Value::Num(*v)),
273        EndExpr::Var(i) => Ok(Value::String(format!("var:{i}"))),
274        EndExpr::ResolvedCall { identity, args, .. } => {
275            let name = strict_callable_display_name(identity).ok_or_else(|| {
276                semantic_error(
277                    "UndefinedFunction",
278                    "end expression call missing callable name",
279                )
280            })?;
281            let mut items = vec![Value::String("call".to_string()), Value::String(name)];
282            for a in args {
283                items.push(encode_end_expr_value(a)?);
284            }
285            mk_cell(items)
286        }
287        EndExpr::Add(a, b) => mk_cell(vec![
288            Value::String("+".to_string()),
289            encode_end_expr_value(a)?,
290            encode_end_expr_value(b)?,
291        ]),
292        EndExpr::Sub(a, b) => mk_cell(vec![
293            Value::String("-".to_string()),
294            encode_end_expr_value(a)?,
295            encode_end_expr_value(b)?,
296        ]),
297        EndExpr::Mul(a, b) => mk_cell(vec![
298            Value::String("*".to_string()),
299            encode_end_expr_value(a)?,
300            encode_end_expr_value(b)?,
301        ]),
302        EndExpr::Div(a, b) => mk_cell(vec![
303            Value::String("/".to_string()),
304            encode_end_expr_value(a)?,
305            encode_end_expr_value(b)?,
306        ]),
307        EndExpr::LeftDiv(a, b) => mk_cell(vec![
308            Value::String("\\".to_string()),
309            encode_end_expr_value(a)?,
310            encode_end_expr_value(b)?,
311        ]),
312        EndExpr::Pow(a, b) => mk_cell(vec![
313            Value::String("^".to_string()),
314            encode_end_expr_value(a)?,
315            encode_end_expr_value(b)?,
316        ]),
317        EndExpr::Neg(a) => mk_cell(vec![
318            Value::String("neg".to_string()),
319            encode_end_expr_value(a)?,
320        ]),
321        EndExpr::Pos(a) => mk_cell(vec![
322            Value::String("pos".to_string()),
323            encode_end_expr_value(a)?,
324        ]),
325        EndExpr::Floor(a) => mk_cell(vec![
326            Value::String("floor".to_string()),
327            encode_end_expr_value(a)?,
328        ]),
329        EndExpr::Ceil(a) => mk_cell(vec![
330            Value::String("ceil".to_string()),
331            encode_end_expr_value(a)?,
332        ]),
333        EndExpr::Round(a) => mk_cell(vec![
334            Value::String("round".to_string()),
335            encode_end_expr_value(a)?,
336        ]),
337        EndExpr::Fix(a) => mk_cell(vec![
338            Value::String("fix".to_string()),
339            encode_end_expr_value(a)?,
340        ]),
341    }
342}
343
344fn build_end_range_descriptor(
345    start: Value,
346    step: Value,
347    end_expr: &EndExpr,
348) -> Result<Value, RuntimeError> {
349    let encoded_end = encode_end_expr_value(end_expr)?;
350    let cell = build_cell_array_with_shape(
351        vec![
352            start,
353            step,
354            Value::String(OBJECT_END_RANGE_TAG.to_string()),
355            encoded_end,
356        ],
357        1,
358        4,
359        "obj range",
360    )?;
361    Ok(Value::Cell(cell))
362}
363
364fn normalize_object_numeric_selector(selector: &Value) -> Result<Value, RuntimeError> {
365    match selector {
366        Value::Num(n) => Ok(Value::Num(*n)),
367        Value::Int(i) => Ok(Value::Int(i.clone())),
368        Value::Tensor(t) => Ok(Value::Tensor(t.clone())),
369        Value::Bool(value) => Ok(Value::Bool(*value)),
370        Value::LogicalArray(array) => Ok(Value::LogicalArray(array.clone())),
371        Value::String(value) => Ok(Value::String(value.clone())),
372        Value::StringArray(array) => Ok(Value::StringArray(array.clone())),
373        Value::CharArray(array) => Ok(Value::CharArray(array.clone())),
374        Value::Cell(cell) => Ok(Value::Cell(cell.clone())),
375        _ => Err(semantic_error(
376            "ObjectSelectorTypeUnsupported",
377            "unsupported index type for object selector",
378        )),
379    }
380}
381
382fn validate_object_range_selector_plan(
383    dims: usize,
384    range_dims: &[usize],
385    range_params: &[(f64, f64)],
386    range_start_exprs: &[Option<EndExpr>],
387    range_step_exprs: &[Option<EndExpr>],
388    range_end_exprs: &[EndExpr],
389) -> Result<Vec<Option<usize>>, RuntimeError> {
390    let count = range_dims.len();
391    if range_params.len() != count
392        || range_start_exprs.len() != count
393        || range_step_exprs.len() != count
394        || range_end_exprs.len() != count
395    {
396        return Err(semantic_error(
397            "InvalidRangeSelectorPlan",
398            "inconsistent object range selector metadata",
399        ));
400    }
401
402    let mut range_pos_by_dim = vec![None; dims];
403    for (pos, &dim) in range_dims.iter().enumerate() {
404        if dim >= dims {
405            return Err(semantic_error(
406                "InvalidRangeSelectorDim",
407                "object range selector dimension is out of bounds",
408            ));
409        }
410        if range_pos_by_dim[dim].replace(pos).is_some() {
411            return Err(semantic_error(
412                "InvalidRangeSelectorPlan",
413                "object range selector dimension appears more than once",
414            ));
415        }
416    }
417    Ok(range_pos_by_dim)
418}
419
420fn validate_object_end_numeric_selector_plan(
421    slot_count: usize,
422    end_numeric_exprs: &[(usize, EndExpr)],
423) -> Result<Vec<Option<&EndExpr>>, RuntimeError> {
424    let mut end_expr_by_slot = vec![None; slot_count];
425    for (position, expr) in end_numeric_exprs {
426        if *position >= slot_count {
427            return Err(semantic_error(
428                "InvalidEndSelectorPlan",
429                "object end-selector position is out of bounds",
430            ));
431        }
432        if end_expr_by_slot[*position].is_some() {
433            return Err(semantic_error(
434                "InvalidEndSelectorPlan",
435                "object end-selector position appears more than once",
436            ));
437        }
438        end_expr_by_slot[*position] = Some(expr);
439    }
440    Ok(end_expr_by_slot)
441}
442
443fn validate_object_selector_masks(
444    dims: usize,
445    colon_mask: u32,
446    end_mask: u32,
447) -> Result<(), RuntimeError> {
448    if (colon_mask & end_mask) != 0 {
449        return Err(semantic_error(
450            "InvalidSelectorMaskPlan",
451            "object selector masks overlap on the same dimension",
452        ));
453    }
454
455    if dims < u32::BITS as usize {
456        let allowed_mask = if dims == 0 { 0 } else { (1u32 << dims) - 1 };
457        if ((colon_mask | end_mask) & !allowed_mask) != 0 {
458            return Err(semantic_error(
459                "InvalidSelectorMaskPlan",
460                "object selector mask dimension is out of bounds",
461            ));
462        }
463    }
464
465    Ok(())
466}
467
468fn object_selector_mask_has_dim(mask: u32, dim: usize) -> bool {
469    dim < u32::BITS as usize && (mask & (1u32 << dim)) != 0
470}
471
472pub fn build_object_paren_selector_values(
473    dims: usize,
474    colon_mask: u32,
475    end_mask: u32,
476    numeric: &[Value],
477) -> Result<Vec<Value>, RuntimeError> {
478    validate_object_selector_masks(dims, colon_mask, end_mask)?;
479    let mut values = Vec::with_capacity(dims);
480    let mut numeric_iter = 0usize;
481    for d in 0..dims {
482        let is_colon = object_selector_mask_has_dim(colon_mask, d);
483        let is_end = object_selector_mask_has_dim(end_mask, d);
484        if is_colon {
485            values.push(Value::String(OBJECT_SELECTOR_COLON.to_string()));
486            continue;
487        }
488        if is_end {
489            values.push(Value::String(OBJECT_SELECTOR_END.to_string()));
490            continue;
491        }
492        let selector = numeric.get(numeric_iter).ok_or(semantic_error(
493            "MissingNumericIndex",
494            "missing numeric index",
495        ))?;
496        values.push(normalize_object_numeric_selector(selector)?);
497        numeric_iter += 1;
498    }
499    if numeric_iter != numeric.len() {
500        return Err(semantic_error(
501            "UnexpectedNumericIndex",
502            "unexpected extra numeric index values",
503        ));
504    }
505    Ok(values)
506}
507
508pub fn build_object_paren_expr_selector_values(
509    spec: ObjectParenExprSelectorSpec<'_>,
510) -> Result<Vec<Value>, RuntimeError> {
511    validate_object_selector_masks(spec.dims, spec.colon_mask, spec.end_mask)?;
512    let range_pos_by_dim = validate_object_range_selector_plan(
513        spec.dims,
514        spec.range_dims,
515        spec.range_params,
516        spec.range_start_exprs,
517        spec.range_step_exprs,
518        spec.range_end_exprs,
519    )?;
520    for (d, range_pos) in range_pos_by_dim.iter().enumerate().take(spec.dims) {
521        if range_pos.is_some() {
522            let is_colon = object_selector_mask_has_dim(spec.colon_mask, d);
523            let is_end = object_selector_mask_has_dim(spec.end_mask, d);
524            if is_colon || is_end {
525                return Err(semantic_error(
526                    "InvalidRangeSelectorPlan",
527                    "object range selector conflicts with colon/end selector masks",
528                ));
529            }
530        }
531    }
532    let slot_count = (0..spec.dims)
533        .filter(|&d| {
534            let is_colon = object_selector_mask_has_dim(spec.colon_mask, d);
535            let is_end = object_selector_mask_has_dim(spec.end_mask, d);
536            !is_colon && !is_end && range_pos_by_dim[d].is_none()
537        })
538        .count();
539    let end_expr_by_slot =
540        validate_object_end_numeric_selector_plan(slot_count, spec.end_numeric_exprs)?;
541    let mut values = Vec::with_capacity(spec.dims);
542    let mut num_iter = 0usize;
543    for (d, range_pos) in range_pos_by_dim.iter().enumerate().take(spec.dims) {
544        let is_colon = object_selector_mask_has_dim(spec.colon_mask, d);
545        let is_end = object_selector_mask_has_dim(spec.end_mask, d);
546        if is_colon {
547            values.push(Value::String(OBJECT_SELECTOR_COLON.to_string()));
548            continue;
549        }
550        if is_end {
551            values.push(Value::String(OBJECT_SELECTOR_END.to_string()));
552            continue;
553        }
554        if let Some(pos) = *range_pos {
555            let (raw_st, raw_sp) = spec.range_params[pos];
556            let st = if let Some(expr) = &spec.range_start_exprs[pos] {
557                encode_end_expr_value(expr)?
558            } else {
559                Value::Num(raw_st)
560            };
561            let sp = if let Some(expr) = &spec.range_step_exprs[pos] {
562                encode_end_expr_value(expr)?
563            } else {
564                Value::Num(raw_sp)
565            };
566            let off = &spec.range_end_exprs[pos];
567            values.push(build_end_range_descriptor(st, sp, off)?);
568            continue;
569        }
570        if let Some(expr) = end_expr_by_slot[num_iter] {
571            values.push(encode_end_expr_value(expr)?);
572            num_iter += 1;
573            continue;
574        }
575        let selector = spec.numeric.get(num_iter).ok_or(semantic_error(
576            "MissingNumericIndex",
577            "missing numeric index",
578        ))?;
579        num_iter += 1;
580        values.push(normalize_object_numeric_selector(selector)?);
581    }
582    if num_iter != spec.numeric.len() {
583        return Err(semantic_error(
584            "UnexpectedNumericIndex",
585            "unexpected extra numeric index values",
586        ));
587    }
588    Ok(values)
589}
590
591fn build_cell_array_with_shape(
592    values: Vec<Value>,
593    rows: usize,
594    cols: usize,
595    context: &str,
596) -> Result<runmat_value::CellArray, RuntimeError> {
597    runmat_value::CellArray::new(values, rows, cols)
598        .map_err(|error| semantic_error("ShapeMismatch", format!("{context}: {error}")))
599}