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