Skip to main content

runmat_runtime/builtins/strings/core/
strlength.rs

1//! MATLAB-compatible `strlength` builtin for RunMat.
2
3use runmat_builtins::{
4    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
6};
7use runmat_builtins::{BuiltinIntegerAuditDescriptor, BuiltinIntegerAuditKind};
8use runmat_macros::runtime_builtin;
9use runmat_value::{CellArray, CharArray, StringArray, Tensor, Value};
10
11use crate::builtins::common::map_control_flow_with_builtin;
12use crate::builtins::common::spec::{
13    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
14    ReductionNaN, ResidencyPolicy, ShapeRequirements,
15};
16use crate::builtins::common::tensor;
17use crate::builtins::strings::common::{
18    contains_numeric_or_resident_text_input, is_missing_string,
19};
20use crate::builtins::strings::type_resolvers::numeric_text_scalar_or_tensor_type;
21use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
22
23#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::strings::core::strlength")]
24pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
25    name: "strlength",
26    op_kind: GpuOpKind::Custom("string-metadata"),
27    supported_precisions: &[],
28    broadcast: BroadcastSemantics::None,
29    provider_hooks: &[],
30    constant_strategy: ConstantStrategy::InlineLiteral,
31    residency: ResidencyPolicy::GatherImmediately,
32    nan_mode: ReductionNaN::Include,
33    two_pass_threshold: None,
34    workgroup_size: None,
35    accepts_nan_mode: false,
36    notes: "Measures string lengths on the CPU; any GPU-resident inputs are gathered before evaluation.",
37};
38
39#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::strings::core::strlength")]
40pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
41    name: "strlength",
42    shape: ShapeRequirements::Any,
43    constant_strategy: ConstantStrategy::InlineLiteral,
44    elementwise: None,
45    reduction: None,
46    emits_nan: true,
47    notes: "Metadata-only builtin; not eligible for fusion and never emits GPU kernels.",
48};
49
50const BUILTIN_NAME: &str = "strlength";
51
52const STRLENGTH_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
53    name: "L",
54    ty: BuiltinParamType::NumericArray,
55    arity: BuiltinParamArity::Required,
56    default: None,
57    description: "Character counts for each text element.",
58}];
59
60const STRLENGTH_INPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
61    name: "str",
62    ty: BuiltinParamType::Any,
63    arity: BuiltinParamArity::Required,
64    default: None,
65    description: "String array, character array, or cell array of text scalars.",
66}];
67
68const STRLENGTH_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
69    label: "L = strlength(str)",
70    inputs: &STRLENGTH_INPUT,
71    outputs: &STRLENGTH_OUTPUT,
72}];
73
74const STRLENGTH_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
75    code: "RM.STRLENGTH.INVALID_INPUT",
76    identifier: Some("RunMat:strlength:InvalidInput"),
77    when: "Input is not a string array, character array, or cell array of text scalars.",
78    message: "strlength: first argument must be a string array, character array, or cell array of character vectors",
79};
80
81const STRLENGTH_ERROR_INVALID_CELL_ELEMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
82    code: "RM.STRLENGTH.INVALID_CELL_ELEMENT",
83    identifier: Some("RunMat:strlength:InvalidCellElement"),
84    when: "A cell-array element is not a character row vector or scalar string.",
85    message: "strlength: cell array elements must be character vectors or string scalars",
86};
87
88const STRLENGTH_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
89    code: "RM.STRLENGTH.INTERNAL",
90    identifier: Some("RunMat:strlength:InternalError"),
91    when: "Internal tensor construction failed while building length results.",
92    message: "strlength: internal error",
93};
94
95const STRLENGTH_ERRORS: [BuiltinErrorDescriptor; 3] = [
96    STRLENGTH_ERROR_INVALID_INPUT,
97    STRLENGTH_ERROR_INVALID_CELL_ELEMENT,
98    STRLENGTH_ERROR_INTERNAL,
99];
100
101pub const STRLENGTH_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
102    signatures: &STRLENGTH_SIGNATURES,
103    output_mode: BuiltinOutputMode::Fixed,
104    completion_policy: BuiltinCompletionPolicy::Public,
105    errors: &STRLENGTH_ERRORS,
106};
107
108pub const STRLENGTH_INTEGER_AUDIT: BuiltinIntegerAuditDescriptor =
109    BuiltinIntegerAuditDescriptor {
110        kind: BuiltinIntegerAuditKind::NotApplicable,
111        canonical_builtin: None,
112        notes: "strlength measures string, character, and cellstr input and returns double character counts. Integer and resident numeric inputs reject before provider access and are never interpreted as character codes.",
113    };
114
115fn strlength_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
116    strlength_error_with_message(error.message, error)
117}
118
119fn strlength_error_with_message(
120    message: impl Into<String>,
121    error: &'static BuiltinErrorDescriptor,
122) -> RuntimeError {
123    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
124    if let Some(identifier) = error.identifier {
125        builder = builder.with_identifier(identifier);
126    }
127    builder.build()
128}
129
130fn remap_strlength_flow(err: RuntimeError) -> RuntimeError {
131    map_control_flow_with_builtin(err, BUILTIN_NAME)
132}
133
134#[runtime_builtin(
135    name = "strlength",
136    category = "strings/core",
137    summary = "Count characters in each element of text inputs.",
138    keywords = "strlength,string length,text,count,characters",
139    accel = "sink",
140    type_resolver(numeric_text_scalar_or_tensor_type),
141    descriptor(crate::builtins::strings::core::strlength::STRLENGTH_DESCRIPTOR),
142    integer_audit(crate::builtins::strings::core::strlength::STRLENGTH_INTEGER_AUDIT),
143    builtin_path = "crate::builtins::strings::core::strlength"
144)]
145async fn strlength_builtin(value: Value) -> crate::BuiltinResult<Value> {
146    if contains_numeric_or_resident_text_input(&value) {
147        return Err(strlength_error(&STRLENGTH_ERROR_INVALID_INPUT));
148    }
149    let gathered = gather_if_needed_async(&value)
150        .await
151        .map_err(remap_strlength_flow)?;
152    match gathered {
153        Value::StringArray(array) => strlength_string_array(array),
154        Value::String(text) => Ok(Value::Num(string_scalar_length(&text))),
155        Value::CharArray(array) => strlength_char_array(array),
156        Value::Cell(cell) => strlength_cell_array(cell),
157        _ => Err(strlength_error(&STRLENGTH_ERROR_INVALID_INPUT)),
158    }
159}
160
161fn strlength_string_array(array: StringArray) -> BuiltinResult<Value> {
162    let StringArray { data, shape, .. } = array;
163    let mut lengths = Vec::with_capacity(data.len());
164    for text in &data {
165        lengths.push(string_scalar_length(text));
166    }
167    let tensor =
168        Tensor::new(lengths, shape).map_err(|_| strlength_error(&STRLENGTH_ERROR_INTERNAL))?;
169    Ok(tensor::tensor_into_value(tensor))
170}
171
172fn strlength_char_array(array: CharArray) -> BuiltinResult<Value> {
173    let rows = array.rows;
174    let mut lengths = Vec::with_capacity(rows);
175    for row in 0..rows {
176        let length = if array.rows <= 1 {
177            array.cols
178        } else {
179            trimmed_row_length(&array, row)
180        } as f64;
181        lengths.push(length);
182    }
183    let tensor = Tensor::new(lengths, vec![rows, 1])
184        .map_err(|_| strlength_error(&STRLENGTH_ERROR_INTERNAL))?;
185    Ok(tensor::tensor_into_value(tensor))
186}
187
188fn strlength_cell_array(cell: CellArray) -> BuiltinResult<Value> {
189    let CellArray {
190        data, rows, cols, ..
191    } = cell;
192    let mut lengths = Vec::with_capacity(rows * cols);
193    for col in 0..cols {
194        for row in 0..rows {
195            let idx = row * cols + col;
196            let value: &Value = &data[idx];
197            let length = match value {
198                Value::String(text) => string_scalar_length(text),
199                Value::StringArray(sa) if sa.data.len() == 1 => string_scalar_length(&sa.data[0]),
200                Value::CharArray(char_vec) if char_vec.rows == 1 => char_vec.cols as f64,
201                Value::CharArray(_) => {
202                    return Err(strlength_error(&STRLENGTH_ERROR_INVALID_CELL_ELEMENT));
203                }
204                _ => return Err(strlength_error(&STRLENGTH_ERROR_INVALID_CELL_ELEMENT)),
205            };
206            lengths.push(length);
207        }
208    }
209    let tensor = Tensor::new(lengths, vec![rows, cols])
210        .map_err(|_| strlength_error(&STRLENGTH_ERROR_INTERNAL))?;
211    Ok(tensor::tensor_into_value(tensor))
212}
213
214fn string_scalar_length(text: &str) -> f64 {
215    if is_missing_string(text) {
216        f64::NAN
217    } else {
218        text.chars().count() as f64
219    }
220}
221
222fn trimmed_row_length(array: &CharArray, row: usize) -> usize {
223    let cols = array.cols;
224    let mut end = cols;
225    while end > 0 {
226        let ch = array.data[row * cols + end - 1];
227        if ch == ' ' {
228            end -= 1;
229        } else {
230            break;
231        }
232    }
233    end
234}
235
236#[cfg(test)]
237pub(crate) mod tests {
238    use super::*;
239    use runmat_builtins::{ResolveContext, Type};
240
241    fn strlength_builtin(value: Value) -> BuiltinResult<Value> {
242        futures::executor::block_on(super::strlength_builtin(value))
243    }
244
245    fn error_message(err: crate::RuntimeError) -> String {
246        err.message().to_string()
247    }
248
249    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
250    #[test]
251    fn strlength_string_scalar() {
252        let result = strlength_builtin(Value::String("RunMat".into())).expect("strlength");
253        assert_eq!(result, Value::Num(6.0));
254    }
255
256    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
257    #[test]
258    fn strlength_string_array_with_missing() {
259        let array = StringArray::new(vec!["alpha".into(), "<missing>".into()], vec![2, 1]).unwrap();
260        let result = strlength_builtin(Value::StringArray(array)).expect("strlength");
261        match result {
262            Value::Tensor(tensor) => {
263                assert_eq!(tensor.shape, vec![2, 1]);
264                assert_eq!(tensor.materialize_f64().len(), 2);
265                assert_eq!(tensor.materialize_f64()[0], 5.0);
266                assert!(tensor.materialize_f64()[1].is_nan());
267            }
268            other => panic!("expected tensor result, got {other:?}"),
269        }
270    }
271
272    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
273    #[test]
274    fn strlength_char_array_multiple_rows() {
275        let data: Vec<char> = vec!['c', 'a', 't', ' ', ' ', 'h', 'o', 'r', 's', 'e'];
276        let array = CharArray::new(data, 2, 5).unwrap();
277        let result = strlength_builtin(Value::CharArray(array)).expect("strlength");
278        match result {
279            Value::Tensor(tensor) => {
280                assert_eq!(tensor.shape, vec![2, 1]);
281                assert_eq!(tensor.materialize_f64(), vec![3.0, 5.0]);
282            }
283            other => panic!("expected tensor result, got {other:?}"),
284        }
285    }
286
287    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
288    #[test]
289    fn strlength_char_vector_retains_explicit_spaces() {
290        let data: Vec<char> = "hi   ".chars().collect();
291        let array = CharArray::new(data, 1, 5).unwrap();
292        let result = strlength_builtin(Value::CharArray(array)).expect("strlength");
293        assert_eq!(result, Value::Num(5.0));
294    }
295
296    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
297    #[test]
298    fn strlength_cell_array_of_char_vectors() {
299        let cell = CellArray::new(
300            vec![
301                Value::CharArray(CharArray::new_row("red")),
302                Value::CharArray(CharArray::new_row("green")),
303            ],
304            1,
305            2,
306        )
307        .unwrap();
308        let result = strlength_builtin(Value::Cell(cell)).expect("strlength");
309        match result {
310            Value::Tensor(tensor) => {
311                assert_eq!(tensor.shape, vec![1, 2]);
312                assert_eq!(tensor.materialize_f64(), vec![3.0, 5.0]);
313            }
314            other => panic!("expected tensor result, got {other:?}"),
315        }
316    }
317
318    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
319    #[test]
320    fn strlength_cell_array_with_string_scalars() {
321        let cell = CellArray::new(
322            vec![
323                Value::String("alpha".into()),
324                Value::String("beta".into()),
325                Value::String("<missing>".into()),
326            ],
327            1,
328            3,
329        )
330        .unwrap();
331        let result = strlength_builtin(Value::Cell(cell)).expect("strlength");
332        match result {
333            Value::Tensor(tensor) => {
334                assert_eq!(tensor.shape, vec![1, 3]);
335                assert_eq!(tensor.materialize_f64().len(), 3);
336                assert_eq!(tensor.materialize_f64()[0], 5.0);
337                assert_eq!(tensor.materialize_f64()[1], 4.0);
338                assert!(tensor.materialize_f64()[2].is_nan());
339            }
340            other => panic!("expected tensor result, got {other:?}"),
341        }
342    }
343
344    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
345    #[test]
346    fn strlength_string_array_preserves_shape() {
347        let array = StringArray::new(
348            vec!["ab".into(), "c".into(), "def".into(), "".into()],
349            vec![2, 2],
350        )
351        .unwrap();
352        let result = strlength_builtin(Value::StringArray(array)).expect("strlength");
353        match result {
354            Value::Tensor(tensor) => {
355                assert_eq!(tensor.shape, vec![2, 2]);
356                assert_eq!(tensor.materialize_f64(), vec![2.0, 1.0, 3.0, 0.0]);
357            }
358            other => panic!("expected tensor result, got {other:?}"),
359        }
360    }
361
362    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
363    #[test]
364    fn strlength_char_array_trims_padding() {
365        let data: Vec<char> = vec!['d', 'o', 'g', ' ', ' ', 'h', 'o', 'r', 's', 'e'];
366        let array = CharArray::new(data, 2, 5).unwrap();
367        let result = strlength_builtin(Value::CharArray(array)).expect("strlength");
368        match result {
369            Value::Tensor(tensor) => {
370                assert_eq!(tensor.shape, vec![2, 1]);
371                assert_eq!(tensor.materialize_f64(), vec![3.0, 5.0]);
372            }
373            other => panic!("expected tensor result, got {other:?}"),
374        }
375    }
376
377    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
378    #[test]
379    fn strlength_errors_on_invalid_input() {
380        let err = error_message(strlength_builtin(Value::Num(1.0)).unwrap_err());
381        assert_eq!(err, STRLENGTH_ERROR_INVALID_INPUT.message);
382    }
383
384    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
385    #[test]
386    fn strlength_rejects_cell_with_invalid_element() {
387        let cell = CellArray::new(
388            vec![Value::CharArray(CharArray::new_row("ok")), Value::Num(5.0)],
389            1,
390            2,
391        )
392        .unwrap();
393        let err = error_message(strlength_builtin(Value::Cell(cell)).unwrap_err());
394        assert_eq!(err, STRLENGTH_ERROR_INVALID_CELL_ELEMENT.message);
395    }
396
397    #[test]
398    fn strlength_type_is_numeric_text_scalar_or_tensor() {
399        assert_eq!(
400            numeric_text_scalar_or_tensor_type(&[Type::String], &ResolveContext::new(Vec::new())),
401            Type::Num
402        );
403    }
404}