Skip to main content

runmat_runtime/builtins/strings/transform/
erase.rs

1//! MATLAB-compatible `erase` builtin with GPU-aware semantics for RunMat.
2use regex::Regex;
3use runmat_builtins::{
4    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
6    CellArray, CharArray, StringArray, Value,
7};
8use runmat_macros::runtime_builtin;
9
10use crate::builtins::common::map_control_flow_with_builtin;
11use crate::builtins::common::spec::{
12    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
13    ReductionNaN, ResidencyPolicy, ShapeRequirements,
14};
15use crate::builtins::strings::common::{char_row_to_string_slice, is_missing_string};
16use crate::builtins::strings::core::compat::pattern_regex;
17use crate::builtins::strings::type_resolvers::text_preserve_type;
18use crate::{
19    build_runtime_error, gather_if_needed_async, make_cell_with_shape, BuiltinResult, RuntimeError,
20};
21
22#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::strings::transform::erase")]
23pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
24    name: "erase",
25    op_kind: GpuOpKind::Custom("string-transform"),
26    supported_precisions: &[],
27    broadcast: BroadcastSemantics::None,
28    provider_hooks: &[],
29    constant_strategy: ConstantStrategy::InlineLiteral,
30    residency: ResidencyPolicy::GatherImmediately,
31    nan_mode: ReductionNaN::Include,
32    two_pass_threshold: None,
33    workgroup_size: None,
34    accepts_nan_mode: false,
35    notes:
36        "Executes on the CPU; GPU-resident inputs are gathered to host memory before substrings are removed.",
37};
38
39#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::strings::transform::erase")]
40pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
41    name: "erase",
42    shape: ShapeRequirements::Any,
43    constant_strategy: ConstantStrategy::InlineLiteral,
44    elementwise: None,
45    reduction: None,
46    emits_nan: false,
47    notes:
48        "String manipulation builtin; not eligible for fusion plans and always gathers GPU inputs before execution.",
49};
50
51const BUILTIN_NAME: &str = "erase";
52
53const ERASE_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
54    name: "newStr",
55    ty: BuiltinParamType::Any,
56    arity: BuiltinParamArity::Required,
57    default: None,
58    description: "Text with substring occurrences removed, preserving input container kind.",
59}];
60
61const ERASE_INPUTS: [BuiltinParamDescriptor; 2] = [
62    BuiltinParamDescriptor {
63        name: "str",
64        ty: BuiltinParamType::Any,
65        arity: BuiltinParamArity::Required,
66        default: None,
67        description: "Input text (string/char/cell).",
68    },
69    BuiltinParamDescriptor {
70        name: "pattern",
71        ty: BuiltinParamType::Any,
72        arity: BuiltinParamArity::Required,
73        default: None,
74        description: "Pattern text list (scalar or array/cell).",
75    },
76];
77
78const ERASE_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
79    label: "newStr = erase(str, pattern)",
80    inputs: &ERASE_INPUTS,
81    outputs: &ERASE_OUTPUT,
82}];
83
84const ERASE_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
85    code: "RM.ERASE.INVALID_INPUT",
86    identifier: Some("RunMat:erase:InvalidInput"),
87    when: "First argument is not a string array, char array, or cell array of text scalars.",
88    message:
89        "erase: first argument must be a string array, character array, or cell array of character vectors",
90};
91
92const ERASE_ERROR_PATTERN_TYPE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
93    code: "RM.ERASE.PATTERN_TYPE",
94    identifier: Some("RunMat:erase:PatternType"),
95    when: "Second argument is not a text scalar/array/cell of text scalars.",
96    message:
97        "erase: second argument must be a string array, character array, or cell array of character vectors",
98};
99
100const ERASE_ERROR_CELL_ELEMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
101    code: "RM.ERASE.CELL_ELEMENT",
102    identifier: Some("RunMat:erase:CellElement"),
103    when: "Cell arrays contain non-text elements or non-row char arrays.",
104    message: "erase: cell array elements must be string scalars or character vectors",
105};
106
107const ERASE_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
108    code: "RM.ERASE.INTERNAL",
109    identifier: Some("RunMat:erase:InternalError"),
110    when: "Internal output container construction failed.",
111    message: "erase: internal error",
112};
113
114const ERASE_ERRORS: [BuiltinErrorDescriptor; 4] = [
115    ERASE_ERROR_INVALID_INPUT,
116    ERASE_ERROR_PATTERN_TYPE,
117    ERASE_ERROR_CELL_ELEMENT,
118    ERASE_ERROR_INTERNAL,
119];
120
121pub const ERASE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
122    signatures: &ERASE_SIGNATURES,
123    output_mode: BuiltinOutputMode::Fixed,
124    completion_policy: BuiltinCompletionPolicy::Public,
125    errors: &ERASE_ERRORS,
126};
127
128fn map_flow(err: RuntimeError) -> RuntimeError {
129    map_control_flow_with_builtin(err, BUILTIN_NAME)
130}
131
132fn erase_error_with_message(
133    message: impl Into<String>,
134    error: &'static BuiltinErrorDescriptor,
135) -> RuntimeError {
136    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
137    if let Some(identifier) = error.identifier {
138        builder = builder.with_identifier(identifier);
139    }
140    builder.build()
141}
142
143fn erase_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
144    erase_error_with_message(error.message, error)
145}
146
147#[runtime_builtin(
148    name = "erase",
149    category = "strings/transform",
150    summary = "Remove substring occurrences from text inputs.",
151    keywords = "erase,remove substring,strings,character array,text",
152    accel = "sink",
153    type_resolver(text_preserve_type),
154    descriptor(crate::builtins::strings::transform::erase::ERASE_DESCRIPTOR),
155    builtin_path = "crate::builtins::strings::transform::erase"
156)]
157async fn erase_builtin(text: Value, pattern: Value) -> BuiltinResult<Value> {
158    let text = gather_if_needed_async(&text).await.map_err(map_flow)?;
159    let pattern = gather_if_needed_async(&pattern).await.map_err(map_flow)?;
160
161    let patterns = PatternList::from_value(&pattern)?;
162
163    match text {
164        Value::String(s) => Ok(Value::String(erase_string_scalar(s, &patterns))),
165        Value::StringArray(sa) => erase_string_array(sa, &patterns),
166        Value::CharArray(ca) => erase_char_array(ca, &patterns),
167        Value::Cell(cell) => erase_cell_array(cell, &patterns),
168        _ => Err(erase_error(&ERASE_ERROR_INVALID_INPUT)),
169    }
170}
171
172struct PatternList {
173    entries: Vec<PatternEntry>,
174}
175
176enum PatternEntry {
177    Literal(String),
178    Regex(Regex),
179}
180
181impl PatternList {
182    fn from_value(value: &Value) -> BuiltinResult<Self> {
183        let entries = match value {
184            Value::Object(_) => vec![PatternEntry::Regex(
185                Regex::new(&pattern_regex(value, BUILTIN_NAME).map_err(|err| {
186                    erase_error_with_message(err.message().to_string(), &ERASE_ERROR_PATTERN_TYPE)
187                })?)
188                .map_err(|err| {
189                    erase_error_with_message(err.to_string(), &ERASE_ERROR_PATTERN_TYPE)
190                })?,
191            )],
192            Value::String(text) => vec![PatternEntry::Literal(text.clone())],
193            Value::StringArray(array) => array
194                .data
195                .iter()
196                .cloned()
197                .map(PatternEntry::Literal)
198                .collect(),
199            Value::CharArray(array) => {
200                if array.rows == 0 {
201                    Vec::new()
202                } else {
203                    let mut list = Vec::with_capacity(array.rows);
204                    for row in 0..array.rows {
205                        list.push(PatternEntry::Literal(char_row_to_string_slice(
206                            &array.data,
207                            array.cols,
208                            row,
209                        )));
210                    }
211                    list
212                }
213            }
214            Value::Cell(cell) => {
215                let mut list = Vec::with_capacity(cell.data.len());
216                for handle in &cell.data {
217                    match &handle {
218                        Value::String(text) => list.push(PatternEntry::Literal(text.clone())),
219                        Value::StringArray(sa) if sa.data.len() == 1 => {
220                            list.push(PatternEntry::Literal(sa.data[0].clone()));
221                        }
222                        Value::CharArray(ca) if ca.rows == 0 => {
223                            list.push(PatternEntry::Literal(String::new()));
224                        }
225                        Value::CharArray(ca) if ca.rows == 1 => {
226                            list.push(PatternEntry::Literal(char_row_to_string_slice(
227                                &ca.data, ca.cols, 0,
228                            )));
229                        }
230                        Value::CharArray(_) => return Err(erase_error(&ERASE_ERROR_CELL_ELEMENT)),
231                        _ => return Err(erase_error(&ERASE_ERROR_CELL_ELEMENT)),
232                    }
233                }
234                list
235            }
236            _ => return Err(erase_error(&ERASE_ERROR_PATTERN_TYPE)),
237        };
238        Ok(Self { entries })
239    }
240
241    fn apply(&self, input: &str) -> String {
242        if self.entries.is_empty() {
243            return input.to_string();
244        }
245        let mut current = input.to_string();
246        for pattern in &self.entries {
247            match pattern {
248                PatternEntry::Literal(pattern) => {
249                    if pattern.is_empty() {
250                        continue;
251                    }
252                    current = current.replace(pattern, "");
253                }
254                PatternEntry::Regex(pattern) => {
255                    current = pattern.replace_all(&current, "").to_string();
256                }
257            }
258            if current.is_empty() {
259                break;
260            }
261        }
262        current
263    }
264}
265
266fn erase_string_scalar(text: String, patterns: &PatternList) -> String {
267    if is_missing_string(&text) {
268        text
269    } else {
270        patterns.apply(&text)
271    }
272}
273
274fn erase_string_array(array: StringArray, patterns: &PatternList) -> BuiltinResult<Value> {
275    let StringArray { data, shape, .. } = array;
276    let mut erased = Vec::with_capacity(data.len());
277    for entry in data {
278        if is_missing_string(&entry) {
279            erased.push(entry);
280        } else {
281            erased.push(patterns.apply(&entry));
282        }
283    }
284    StringArray::new(erased, shape)
285        .map(Value::StringArray)
286        .map_err(|e| {
287            erase_error_with_message(format!("{BUILTIN_NAME}: {e}"), &ERASE_ERROR_INTERNAL)
288        })
289}
290
291fn erase_char_array(array: CharArray, patterns: &PatternList) -> BuiltinResult<Value> {
292    let CharArray { data, rows, cols } = array;
293    if rows == 0 {
294        return Ok(Value::CharArray(CharArray { data, rows, cols }));
295    }
296
297    let mut processed: Vec<String> = Vec::with_capacity(rows);
298    let mut target_cols = 0usize;
299    for row in 0..rows {
300        let slice = char_row_to_string_slice(&data, cols, row);
301        let erased = patterns.apply(&slice);
302        let len = erased.chars().count();
303        if len > target_cols {
304            target_cols = len;
305        }
306        processed.push(erased);
307    }
308
309    let mut flattened: Vec<char> = Vec::with_capacity(rows * target_cols);
310    for row_text in processed {
311        let mut chars: Vec<char> = row_text.chars().collect();
312        if chars.len() < target_cols {
313            chars.resize(target_cols, ' ');
314        }
315        flattened.extend(chars);
316    }
317
318    CharArray::new(flattened, rows, target_cols)
319        .map(Value::CharArray)
320        .map_err(|e| {
321            erase_error_with_message(format!("{BUILTIN_NAME}: {e}"), &ERASE_ERROR_INTERNAL)
322        })
323}
324
325fn erase_cell_array(cell: CellArray, patterns: &PatternList) -> BuiltinResult<Value> {
326    let shape = cell.shape.clone();
327    let mut values = Vec::with_capacity(cell.data.len());
328    for handle in &cell.data {
329        values.push(erase_cell_element(handle, patterns)?);
330    }
331    make_cell_with_shape(values, shape).map_err(|e| {
332        erase_error_with_message(format!("{BUILTIN_NAME}: {e}"), &ERASE_ERROR_INTERNAL)
333    })
334}
335
336fn erase_cell_element(value: &Value, patterns: &PatternList) -> BuiltinResult<Value> {
337    match value {
338        Value::String(text) => Ok(Value::String(erase_string_scalar(text.clone(), patterns))),
339        Value::StringArray(sa) if sa.data.len() == 1 => Ok(Value::String(erase_string_scalar(
340            sa.data[0].clone(),
341            patterns,
342        ))),
343        Value::CharArray(ca) if ca.rows == 0 => Ok(Value::CharArray(ca.clone())),
344        Value::CharArray(ca) if ca.rows == 1 => {
345            let slice = char_row_to_string_slice(&ca.data, ca.cols, 0);
346            let erased = patterns.apply(&slice);
347            Ok(Value::CharArray(CharArray::new_row(&erased)))
348        }
349        Value::CharArray(_) => Err(erase_error(&ERASE_ERROR_CELL_ELEMENT)),
350        _ => Err(erase_error(&ERASE_ERROR_CELL_ELEMENT)),
351    }
352}
353
354#[cfg(test)]
355pub(crate) mod tests {
356    use super::*;
357    use runmat_builtins::{ResolveContext, Type};
358
359    fn erase_builtin(text: Value, pattern: Value) -> BuiltinResult<Value> {
360        futures::executor::block_on(super::erase_builtin(text, pattern))
361    }
362
363    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
364    #[test]
365    fn erase_string_scalar_single_pattern() {
366        let result = erase_builtin(
367            Value::String("RunMat runtime".into()),
368            Value::String(" runtime".into()),
369        )
370        .expect("erase");
371        assert_eq!(result, Value::String("RunMat".into()));
372    }
373
374    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
375    #[test]
376    fn erase_string_array_multiple_patterns() {
377        let strings = StringArray::new(
378            vec!["gpu".into(), "cpu".into(), "<missing>".into()],
379            vec![3, 1],
380        )
381        .unwrap();
382        let result = erase_builtin(
383            Value::StringArray(strings),
384            Value::StringArray(StringArray::new(vec!["g".into(), "c".into()], vec![2, 1]).unwrap()),
385        )
386        .expect("erase");
387        match result {
388            Value::StringArray(sa) => {
389                assert_eq!(sa.shape, vec![3, 1]);
390                assert_eq!(
391                    sa.data,
392                    vec![
393                        String::from("pu"),
394                        String::from("pu"),
395                        String::from("<missing>")
396                    ]
397                );
398            }
399            other => panic!("expected string array, got {other:?}"),
400        }
401    }
402
403    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
404    #[test]
405    fn erase_string_array_shape_mismatch_applies_all_patterns() {
406        let strings =
407            StringArray::new(vec!["GPU kernel".into(), "CPU kernel".into()], vec![2, 1]).unwrap();
408        let patterns = StringArray::new(vec!["GPU ".into(), "CPU ".into()], vec![1, 2]).unwrap();
409        let result = erase_builtin(Value::StringArray(strings), Value::StringArray(patterns))
410            .expect("erase");
411        match result {
412            Value::StringArray(sa) => {
413                assert_eq!(sa.shape, vec![2, 1]);
414                assert_eq!(
415                    sa.data,
416                    vec![String::from("kernel"), String::from("kernel")]
417                );
418            }
419            other => panic!("expected string array, got {other:?}"),
420        }
421    }
422
423    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
424    #[test]
425    fn erase_char_array_adjusts_width() {
426        let chars = CharArray::new("matrix".chars().collect(), 1, 6).unwrap();
427        let result =
428            erase_builtin(Value::CharArray(chars), Value::String("tr".into())).expect("erase");
429        match result {
430            Value::CharArray(out) => {
431                assert_eq!(out.rows, 1);
432                assert_eq!(out.cols, 4);
433                let expected: Vec<char> = "maix".chars().collect();
434                assert_eq!(out.data, expected);
435            }
436            other => panic!("expected char array, got {other:?}"),
437        }
438    }
439
440    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
441    #[test]
442    fn erase_char_array_handles_full_removal() {
443        let chars = CharArray::new_row("abc");
444        let result = erase_builtin(Value::CharArray(chars.clone()), Value::String("abc".into()))
445            .expect("erase");
446        match result {
447            Value::CharArray(out) => {
448                assert_eq!(out.rows, 1);
449                assert_eq!(out.cols, 0);
450                assert!(out.data.is_empty());
451            }
452            other => panic!("expected empty char array, got {other:?}"),
453        }
454    }
455
456    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
457    #[test]
458    fn erase_char_array_multiple_rows_sequential_patterns() {
459        let chars = CharArray::new(
460            vec![
461                'G', 'P', 'U', ' ', 'p', 'i', 'p', 'e', 'l', 'i', 'n', 'e', 'C', 'P', 'U', ' ',
462                'p', 'i', 'p', 'e', 'l', 'i', 'n', 'e',
463            ],
464            2,
465            12,
466        )
467        .unwrap();
468        let patterns = CharArray::new_row("GPU ");
469        let result =
470            erase_builtin(Value::CharArray(chars), Value::CharArray(patterns)).expect("erase");
471        match result {
472            Value::CharArray(out) => {
473                assert_eq!(out.rows, 2);
474                assert_eq!(out.cols, 12);
475                let first = char_row_to_string_slice(&out.data, out.cols, 0);
476                let second = char_row_to_string_slice(&out.data, out.cols, 1);
477                assert_eq!(first.trim_end(), "pipeline");
478                assert_eq!(second.trim_end(), "CPU pipeline");
479            }
480            other => panic!("expected char array, got {other:?}"),
481        }
482    }
483
484    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
485    #[test]
486    fn erase_cell_array_mixed_content() {
487        let cell = CellArray::new(
488            vec![
489                Value::CharArray(CharArray::new_row("Kernel Planner")),
490                Value::String("GPU Fusion".into()),
491            ],
492            1,
493            2,
494        )
495        .unwrap();
496        let result = erase_builtin(
497            Value::Cell(cell),
498            Value::Cell(
499                CellArray::new(
500                    vec![
501                        Value::String("Kernel ".into()),
502                        Value::String("GPU ".into()),
503                    ],
504                    1,
505                    2,
506                )
507                .unwrap(),
508            ),
509        )
510        .expect("erase");
511        match result {
512            Value::Cell(out) => {
513                let first = out.get(0, 0).unwrap();
514                let second = out.get(0, 1).unwrap();
515                assert_eq!(first, Value::CharArray(CharArray::new_row("Planner")));
516                assert_eq!(second, Value::String("Fusion".into()));
517            }
518            other => panic!("expected cell array, got {other:?}"),
519        }
520    }
521
522    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
523    #[test]
524    fn erase_cell_array_preserves_shape() {
525        let cell = CellArray::new(
526            vec![
527                Value::String("alpha".into()),
528                Value::String("beta".into()),
529                Value::String("gamma".into()),
530                Value::String("delta".into()),
531            ],
532            2,
533            2,
534        )
535        .unwrap();
536        let patterns = StringArray::new(vec!["a".into()], vec![1, 1]).unwrap();
537        let result = erase_builtin(Value::Cell(cell), Value::StringArray(patterns)).expect("erase");
538        match result {
539            Value::Cell(out) => {
540                assert_eq!(out.rows, 2);
541                assert_eq!(out.cols, 2);
542                assert_eq!(out.get(0, 0).unwrap(), Value::String("lph".into()));
543                assert_eq!(out.get(1, 1).unwrap(), Value::String("delt".into()));
544            }
545            other => panic!("expected cell array, got {other:?}"),
546        }
547    }
548
549    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
550    #[test]
551    fn erase_preserves_missing_string() {
552        let result = erase_builtin(
553            Value::String("<missing>".into()),
554            Value::String("missing".into()),
555        )
556        .expect("erase");
557        assert_eq!(result, Value::String("<missing>".into()));
558    }
559
560    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
561    #[test]
562    fn erase_allows_empty_pattern_list() {
563        let strings = StringArray::new(vec!["alpha".into(), "beta".into()], vec![2, 1]).unwrap();
564        let pattern = StringArray::new(Vec::<String>::new(), vec![0, 0]).unwrap();
565        let result = erase_builtin(
566            Value::StringArray(strings.clone()),
567            Value::StringArray(pattern),
568        )
569        .expect("erase");
570        assert_eq!(result, Value::StringArray(strings));
571    }
572
573    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
574    #[test]
575    fn erase_errors_on_invalid_first_argument() {
576        let err = erase_builtin(Value::Num(1.0), Value::String("a".into())).unwrap_err();
577        assert_eq!(err.to_string(), ERASE_ERROR_INVALID_INPUT.message);
578    }
579
580    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
581    #[test]
582    fn erase_errors_on_invalid_pattern_type() {
583        let err = erase_builtin(Value::String("abc".into()), Value::Num(1.0)).unwrap_err();
584        assert_eq!(err.to_string(), ERASE_ERROR_PATTERN_TYPE.message);
585    }
586
587    #[test]
588    fn erase_accepts_pattern_object() {
589        let pattern = crate::builtins::strings::core::compat::pattern_object(r"\d+");
590        let result = erase_builtin(Value::String("run42mat".into()), pattern).expect("erase");
591        assert_eq!(result, Value::String("runmat".into()));
592    }
593
594    #[test]
595    fn erase_type_preserves_text() {
596        assert_eq!(
597            text_preserve_type(&[Type::String], &ResolveContext::new(Vec::new())),
598            Type::String
599        );
600    }
601}