1use 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}