Skip to main content

runmat_runtime/builtins/strings/core/
strcmpi.rs

1//! MATLAB-compatible `strcmpi` builtin for RunMat (case-insensitive string comparison).
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::Value;
10
11use crate::builtins::common::broadcast::{broadcast_index, broadcast_shapes, compute_strides};
12use crate::builtins::common::map_control_flow_with_builtin;
13use crate::builtins::common::spec::{
14    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
15    ReductionNaN, ResidencyPolicy, ShapeRequirements,
16};
17use crate::builtins::common::tensor;
18use crate::builtins::strings::common::contains_numeric_or_resident_text_input;
19use crate::builtins::strings::search::text_utils::{logical_result, TextCollection, TextElement};
20use crate::builtins::strings::type_resolvers::logical_text_match_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::strcmpi")]
24pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
25    name: "strcmpi",
26    op_kind: GpuOpKind::Custom("string-compare"),
27    supported_precisions: &[],
28    broadcast: BroadcastSemantics::Matlab,
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: "Runs entirely on the CPU; GPU operands are gathered before comparison.",
37};
38
39#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::strings::core::strcmpi")]
40pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
41    name: "strcmpi",
42    shape: ShapeRequirements::Any,
43    constant_strategy: ConstantStrategy::InlineLiteral,
44    elementwise: None,
45    reduction: None,
46    emits_nan: false,
47    notes: "Produces logical host results; not eligible for GPU fusion.",
48};
49
50const BUILTIN_NAME: &str = "strcmpi";
51
52const STRCMPI_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
53    name: "tf",
54    ty: BuiltinParamType::LogicalArray,
55    arity: BuiltinParamArity::Required,
56    default: None,
57    description: "Logical comparison result.",
58}];
59
60const STRCMPI_INPUTS: [BuiltinParamDescriptor; 2] = [
61    BuiltinParamDescriptor {
62        name: "A",
63        ty: BuiltinParamType::Any,
64        arity: BuiltinParamArity::Required,
65        default: None,
66        description: "First text input (string/char/cell/string array).",
67    },
68    BuiltinParamDescriptor {
69        name: "B",
70        ty: BuiltinParamType::Any,
71        arity: BuiltinParamArity::Required,
72        default: None,
73        description: "Second text input (string/char/cell/string array).",
74    },
75];
76
77const STRCMPI_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
78    label: "tf = strcmpi(A, B)",
79    inputs: &STRCMPI_INPUTS,
80    outputs: &STRCMPI_OUTPUT,
81}];
82
83const STRCMPI_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
84    code: "RM.STRCMPI.INVALID_INPUT",
85    identifier: Some("RunMat:strcmpi:InvalidInput"),
86    when: "At least one input is not a supported text container.",
87    message: "strcmpi: text inputs must be string/char/cell/string-array values",
88};
89
90const STRCMPI_ERROR_SHAPE_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
91    code: "RM.STRCMPI.SHAPE_MISMATCH",
92    identifier: Some("RunMat:strcmpi:ShapeMismatch"),
93    when: "Inputs are not broadcast-compatible for elementwise comparison.",
94    message: "strcmpi: input sizes are not broadcast-compatible",
95};
96
97const STRCMPI_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
98    code: "RM.STRCMPI.INTERNAL",
99    identifier: Some("RunMat:strcmpi:InternalError"),
100    when: "Internal logical result assembly failed.",
101    message: "strcmpi: internal error",
102};
103
104const STRCMPI_ERRORS: [BuiltinErrorDescriptor; 3] = [
105    STRCMPI_ERROR_INVALID_INPUT,
106    STRCMPI_ERROR_SHAPE_MISMATCH,
107    STRCMPI_ERROR_INTERNAL,
108];
109
110pub const STRCMPI_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
111    signatures: &STRCMPI_SIGNATURES,
112    output_mode: BuiltinOutputMode::Fixed,
113    completion_policy: BuiltinCompletionPolicy::Public,
114    errors: &STRCMPI_ERRORS,
115};
116
117pub const STRCMPI_INTEGER_AUDIT: BuiltinIntegerAuditDescriptor = BuiltinIntegerAuditDescriptor {
118    kind: BuiltinIntegerAuditKind::NotApplicable,
119    canonical_builtin: None,
120    notes: "strcmpi compares text containers. An unsupported integer or other numeric input returns scalar logical false as documented, without reading numeric payloads or accessing a provider.",
121};
122
123fn strcmpi_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
124    strcmpi_error_with_message(error.message, error)
125}
126
127fn strcmpi_error_with_message(
128    message: impl Into<String>,
129    error: &'static BuiltinErrorDescriptor,
130) -> RuntimeError {
131    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
132    if let Some(identifier) = error.identifier {
133        builder = builder.with_identifier(identifier);
134    }
135    builder.build()
136}
137
138fn remap_strcmpi_flow(err: RuntimeError) -> RuntimeError {
139    map_control_flow_with_builtin(err, BUILTIN_NAME)
140}
141
142#[runtime_builtin(
143    name = "strcmpi",
144    category = "strings/core",
145    summary = "Compare text inputs for case-insensitive equality.",
146    keywords = "strcmpi,string compare,text equality",
147    accel = "sink",
148    type_resolver(logical_text_match_type),
149    descriptor(crate::builtins::strings::core::strcmpi::STRCMPI_DESCRIPTOR),
150    integer_audit(crate::builtins::strings::core::strcmpi::STRCMPI_INTEGER_AUDIT),
151    builtin_path = "crate::builtins::strings::core::strcmpi"
152)]
153async fn strcmpi_builtin(a: Value, b: Value) -> crate::BuiltinResult<Value> {
154    if contains_numeric_or_resident_text_input(&a) || contains_numeric_or_resident_text_input(&b) {
155        return Ok(Value::Bool(false));
156    }
157    let a = gather_if_needed_async(&a)
158        .await
159        .map_err(remap_strcmpi_flow)?;
160    let b = gather_if_needed_async(&b)
161        .await
162        .map_err(remap_strcmpi_flow)?;
163    let left = TextCollection::from_argument(BUILTIN_NAME, a, "first argument")
164        .map_err(|_| strcmpi_error(&STRCMPI_ERROR_INVALID_INPUT))?;
165    let right = TextCollection::from_argument(BUILTIN_NAME, b, "second argument")
166        .map_err(|_| strcmpi_error(&STRCMPI_ERROR_INVALID_INPUT))?;
167    evaluate_strcmpi(&left, &right)
168}
169
170fn evaluate_strcmpi(left: &TextCollection, right: &TextCollection) -> BuiltinResult<Value> {
171    let shape = broadcast_shapes(BUILTIN_NAME, &left.shape, &right.shape)
172        .map_err(|_| strcmpi_error(&STRCMPI_ERROR_SHAPE_MISMATCH))?;
173    let total = tensor::element_count(&shape);
174    if total == 0 {
175        return logical_result(BUILTIN_NAME, Vec::new(), shape)
176            .map_err(|_| strcmpi_error(&STRCMPI_ERROR_INTERNAL));
177    }
178    let left_strides = compute_strides(&left.shape);
179    let right_strides = compute_strides(&right.shape);
180    let left_lower = left.lowercased();
181    let right_lower = right.lowercased();
182    let mut data = Vec::with_capacity(total);
183    for linear in 0..total {
184        let li = broadcast_index(linear, &shape, &left.shape, &left_strides);
185        let ri = broadcast_index(linear, &shape, &right.shape, &right_strides);
186        let equal = match (&left.elements[li], &right.elements[ri]) {
187            (TextElement::Missing, _) => false,
188            (_, TextElement::Missing) => false,
189            (TextElement::Text(_), TextElement::Text(_)) => {
190                match (&left_lower[li], &right_lower[ri]) {
191                    (Some(lhs), Some(rhs)) => lhs == rhs,
192                    _ => false,
193                }
194            }
195        };
196        data.push(if equal { 1 } else { 0 });
197    }
198    logical_result(BUILTIN_NAME, data, shape).map_err(|_| strcmpi_error(&STRCMPI_ERROR_INTERNAL))
199}
200
201#[cfg(test)]
202pub(crate) mod tests {
203    use super::*;
204    use crate::RuntimeError;
205    use runmat_builtins::{ResolveContext, Type};
206    use runmat_value::{CellArray, CharArray, LogicalArray, StringArray};
207
208    fn strcmpi_builtin(a: Value, b: Value) -> BuiltinResult<Value> {
209        futures::executor::block_on(super::strcmpi_builtin(a, b))
210    }
211
212    fn error_message(err: RuntimeError) -> String {
213        err.to_string()
214    }
215
216    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
217    #[test]
218    fn strcmpi_string_scalar_true_ignores_case() {
219        let result = strcmpi_builtin(
220            Value::String("RunMat".into()),
221            Value::String("runmat".into()),
222        )
223        .expect("strcmpi");
224        assert_eq!(result, Value::Bool(true));
225    }
226
227    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
228    #[test]
229    fn strcmpi_string_scalar_false_when_text_differs() {
230        let result = strcmpi_builtin(
231            Value::String("RunMat".into()),
232            Value::String("runtime".into()),
233        )
234        .expect("strcmpi");
235        assert_eq!(result, Value::Bool(false));
236    }
237
238    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
239    #[test]
240    fn strcmpi_string_array_broadcast_scalar_case_insensitive() {
241        let array = StringArray::new(
242            vec!["red".into(), "green".into(), "blue".into()],
243            vec![1, 3],
244        )
245        .unwrap();
246        let result = strcmpi_builtin(Value::StringArray(array), Value::String("GREEN".into()))
247            .expect("strcmpi");
248        let expected = LogicalArray::new(vec![0, 1, 0], vec![1, 3]).unwrap();
249        assert_eq!(result, Value::LogicalArray(expected));
250    }
251
252    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
253    #[test]
254    fn strcmpi_char_array_row_compare_casefold() {
255        let chars = CharArray::new(vec!['c', 'a', 't', 'D', 'O', 'G'], 2, 3).unwrap();
256        let result =
257            strcmpi_builtin(Value::CharArray(chars), Value::String("CaT".into())).expect("cmp");
258        let expected = LogicalArray::new(vec![1, 0], vec![2, 1]).unwrap();
259        assert_eq!(result, Value::LogicalArray(expected));
260    }
261
262    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
263    #[test]
264    fn strcmpi_char_array_to_char_array_casefold() {
265        let left = CharArray::new(vec!['A', 'b', 'C', 'd'], 2, 2).unwrap();
266        let right = CharArray::new(vec!['a', 'B', 'x', 'Y'], 2, 2).unwrap();
267        let result =
268            strcmpi_builtin(Value::CharArray(left), Value::CharArray(right)).expect("strcmpi");
269        let expected = LogicalArray::new(vec![1, 0], vec![2, 1]).unwrap();
270        assert_eq!(result, Value::LogicalArray(expected));
271    }
272
273    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
274    #[test]
275    fn strcmpi_cell_array_scalar_casefold() {
276        let cell = CellArray::new(
277            vec![
278                Value::from("North"),
279                Value::from("east"),
280                Value::from("South"),
281            ],
282            1,
283            3,
284        )
285        .unwrap();
286        let result =
287            strcmpi_builtin(Value::Cell(cell), Value::String("EAST".into())).expect("strcmpi");
288        let expected = LogicalArray::new(vec![0, 1, 0], vec![1, 3]).unwrap();
289        assert_eq!(result, Value::LogicalArray(expected));
290    }
291
292    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
293    #[test]
294    fn strcmpi_cell_array_vs_cell_array_broadcast() {
295        let left = CellArray::new(vec![Value::from("North"), Value::from("East")], 1, 2).unwrap();
296        let right = CellArray::new(vec![Value::from("north")], 1, 1).unwrap();
297        let result = strcmpi_builtin(Value::Cell(left), Value::Cell(right)).expect("strcmpi");
298        let expected = LogicalArray::new(vec![1, 0], vec![1, 2]).unwrap();
299        assert_eq!(result, Value::LogicalArray(expected));
300    }
301
302    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
303    #[test]
304    fn strcmpi_string_array_multi_dimensional_broadcast() {
305        let left = StringArray::new(vec!["north".into(), "south".into()], vec![2, 1]).unwrap();
306        let right = StringArray::new(
307            vec!["NORTH".into(), "EAST".into(), "SOUTH".into()],
308            vec![1, 3],
309        )
310        .unwrap();
311        let result =
312            strcmpi_builtin(Value::StringArray(left), Value::StringArray(right)).expect("strcmpi");
313        let expected = LogicalArray::new(vec![1, 0, 0, 0, 0, 1], vec![2, 3]).unwrap();
314        assert_eq!(result, Value::LogicalArray(expected));
315    }
316
317    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
318    #[test]
319    fn strcmpi_missing_strings_compare_false() {
320        let strings = StringArray::new(vec!["<missing>".into()], vec![1, 1]).unwrap();
321        let result = strcmpi_builtin(
322            Value::StringArray(strings.clone()),
323            Value::StringArray(strings),
324        )
325        .expect("strcmpi");
326        assert_eq!(result, Value::Bool(false));
327    }
328
329    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
330    #[test]
331    fn strcmpi_char_array_trailing_space_not_equal() {
332        let chars = CharArray::new(vec!['c', 'a', 't', ' '], 1, 4).unwrap();
333        let result =
334            strcmpi_builtin(Value::CharArray(chars), Value::String("cat".into())).expect("strcmpi");
335        assert_eq!(result, Value::Bool(false));
336    }
337
338    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
339    #[test]
340    fn strcmpi_size_mismatch_error() {
341        let left = StringArray::new(vec!["a".into(), "b".into()], vec![2, 1]).unwrap();
342        let right = StringArray::new(vec!["a".into(), "b".into(), "c".into()], vec![3, 1]).unwrap();
343        let err = error_message(
344            strcmpi_builtin(Value::StringArray(left), Value::StringArray(right))
345                .expect_err("size mismatch"),
346        );
347        assert!(err.contains(STRCMPI_ERROR_SHAPE_MISMATCH.message));
348    }
349
350    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
351    #[test]
352    fn strcmpi_unsupported_numeric_argument_returns_false() {
353        let result =
354            strcmpi_builtin(Value::Num(1.0), Value::String("a".into())).expect("comparison");
355        assert_eq!(result, Value::Bool(false));
356    }
357
358    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
359    #[test]
360    fn strcmpi_cell_array_invalid_element_errors() {
361        let cell = CellArray::new(vec![Value::Num(42.0)], 1, 1).unwrap();
362        let err = error_message(
363            strcmpi_builtin(Value::Cell(cell), Value::String("test".into()))
364                .expect_err("cell element type"),
365        );
366        assert!(err.contains(STRCMPI_ERROR_INVALID_INPUT.message));
367    }
368
369    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
370    #[test]
371    fn strcmpi_empty_char_array_returns_empty() {
372        let chars = CharArray::new(Vec::<char>::new(), 0, 3).unwrap();
373        let result = strcmpi_builtin(Value::CharArray(chars), Value::String("anything".into()))
374            .expect("cmp");
375        let expected = LogicalArray::new(Vec::<u8>::new(), vec![0, 1]).unwrap();
376        assert_eq!(result, Value::LogicalArray(expected));
377    }
378
379    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
380    #[test]
381    #[cfg(feature = "wgpu")]
382    fn strcmpi_with_wgpu_provider_matches_expected() {
383        let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
384            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
385        );
386        let names = StringArray::new(vec!["North".into(), "south".into()], vec![2, 1]).unwrap();
387        let comparison = StringArray::new(vec!["north".into()], vec![1, 1]).unwrap();
388        let result = strcmpi_builtin(Value::StringArray(names), Value::StringArray(comparison))
389            .expect("strcmpi");
390        let expected = LogicalArray::new(vec![1, 0], vec![2, 1]).unwrap();
391        assert_eq!(result, Value::LogicalArray(expected));
392    }
393
394    #[test]
395    fn strcmpi_type_is_logical_match() {
396        assert_eq!(
397            logical_text_match_type(
398                &[Type::String, Type::String],
399                &ResolveContext::new(Vec::new()),
400            ),
401            Type::Bool
402        );
403    }
404}