Skip to main content

runmat_runtime/builtins/strings/core/
strcmp.rs

1//! MATLAB-compatible `strcmp` 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::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::strcmp")]
24pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
25    name: "strcmp",
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: "Performs host-side text comparisons; GPU operands are gathered automatically before evaluation.",
37};
38
39#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::strings::core::strcmp")]
40pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
41    name: "strcmp",
42    shape: ShapeRequirements::Any,
43    constant_strategy: ConstantStrategy::InlineLiteral,
44    elementwise: None,
45    reduction: None,
46    emits_nan: false,
47    notes: "Produces logical results on the host; not eligible for GPU fusion.",
48};
49
50const BUILTIN_NAME: &str = "strcmp";
51
52const STRCMP_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 STRCMP_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 STRCMP_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
78    label: "tf = strcmp(A, B)",
79    inputs: &STRCMP_INPUTS,
80    outputs: &STRCMP_OUTPUT,
81}];
82
83const STRCMP_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
84    code: "RM.STRCMP.INVALID_INPUT",
85    identifier: Some("RunMat:strcmp:InvalidInput"),
86    when: "At least one input is not a supported text container.",
87    message: "strcmp: text inputs must be string/char/cell/string-array values",
88};
89
90const STRCMP_ERROR_SHAPE_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
91    code: "RM.STRCMP.SHAPE_MISMATCH",
92    identifier: Some("RunMat:strcmp:ShapeMismatch"),
93    when: "Inputs are not broadcast-compatible for elementwise comparison.",
94    message: "strcmp: input sizes are not broadcast-compatible",
95};
96
97const STRCMP_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
98    code: "RM.STRCMP.INTERNAL",
99    identifier: Some("RunMat:strcmp:InternalError"),
100    when: "Internal logical result assembly failed.",
101    message: "strcmp: internal error",
102};
103
104const STRCMP_ERRORS: [BuiltinErrorDescriptor; 3] = [
105    STRCMP_ERROR_INVALID_INPUT,
106    STRCMP_ERROR_SHAPE_MISMATCH,
107    STRCMP_ERROR_INTERNAL,
108];
109
110pub const STRCMP_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
111    signatures: &STRCMP_SIGNATURES,
112    output_mode: BuiltinOutputMode::Fixed,
113    completion_policy: BuiltinCompletionPolicy::Public,
114    errors: &STRCMP_ERRORS,
115};
116
117pub const STRCMP_INTEGER_AUDIT: BuiltinIntegerAuditDescriptor = BuiltinIntegerAuditDescriptor {
118    kind: BuiltinIntegerAuditKind::NotApplicable,
119    canonical_builtin: None,
120    notes: "strcmp 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 strcmp_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
124    strcmp_error_with_message(error.message, error)
125}
126
127fn strcmp_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_strcmp_flow(err: RuntimeError) -> RuntimeError {
139    map_control_flow_with_builtin(err, BUILTIN_NAME)
140}
141
142#[runtime_builtin(
143    name = "strcmp",
144    category = "strings/core",
145    summary = "Compare text inputs for exact case-sensitive equality.",
146    keywords = "strcmp,string compare,text equality",
147    accel = "sink",
148    type_resolver(logical_text_match_type),
149    descriptor(crate::builtins::strings::core::strcmp::STRCMP_DESCRIPTOR),
150    integer_audit(crate::builtins::strings::core::strcmp::STRCMP_INTEGER_AUDIT),
151    builtin_path = "crate::builtins::strings::core::strcmp"
152)]
153async fn strcmp_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_strcmp_flow)?;
160    let b = gather_if_needed_async(&b)
161        .await
162        .map_err(remap_strcmp_flow)?;
163    let left = TextCollection::from_argument(BUILTIN_NAME, a, "first argument")
164        .map_err(|_| strcmp_error(&STRCMP_ERROR_INVALID_INPUT))?;
165    let right = TextCollection::from_argument(BUILTIN_NAME, b, "second argument")
166        .map_err(|_| strcmp_error(&STRCMP_ERROR_INVALID_INPUT))?;
167    evaluate_strcmp(&left, &right)
168}
169
170fn evaluate_strcmp(left: &TextCollection, right: &TextCollection) -> BuiltinResult<Value> {
171    let shape = broadcast_shapes(BUILTIN_NAME, &left.shape, &right.shape)
172        .map_err(|_| strcmp_error(&STRCMP_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(|_| strcmp_error(&STRCMP_ERROR_INTERNAL));
177    }
178    let left_strides = compute_strides(&left.shape);
179    let right_strides = compute_strides(&right.shape);
180    let mut data = Vec::with_capacity(total);
181    for linear in 0..total {
182        let li = broadcast_index(linear, &shape, &left.shape, &left_strides);
183        let ri = broadcast_index(linear, &shape, &right.shape, &right_strides);
184        let equal = match (&left.elements[li], &right.elements[ri]) {
185            (TextElement::Missing, _) => false,
186            (_, TextElement::Missing) => false,
187            (TextElement::Text(lhs), TextElement::Text(rhs)) => lhs == rhs,
188        };
189        data.push(if equal { 1 } else { 0 });
190    }
191    logical_result(BUILTIN_NAME, data, shape).map_err(|_| strcmp_error(&STRCMP_ERROR_INTERNAL))
192}
193
194#[cfg(test)]
195pub(crate) mod tests {
196    use super::*;
197    use crate::RuntimeError;
198    use runmat_builtins::{ResolveContext, Type};
199    use runmat_value::{CellArray, CharArray, LogicalArray, StringArray};
200
201    fn strcmp_builtin(a: Value, b: Value) -> BuiltinResult<Value> {
202        futures::executor::block_on(super::strcmp_builtin(a, b))
203    }
204
205    fn error_message(err: RuntimeError) -> String {
206        err.to_string()
207    }
208
209    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
210    #[test]
211    fn strcmp_string_scalar_true() {
212        let result = strcmp_builtin(
213            Value::String("RunMat".into()),
214            Value::String("RunMat".into()),
215        )
216        .expect("strcmp");
217        assert_eq!(result, Value::Bool(true));
218    }
219
220    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
221    #[test]
222    fn strcmp_string_scalar_false() {
223        let result = strcmp_builtin(
224            Value::String("RunMat".into()),
225            Value::String("runmat".into()),
226        )
227        .expect("strcmp");
228        assert_eq!(result, Value::Bool(false));
229    }
230
231    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
232    #[test]
233    fn strcmp_string_array_broadcast_scalar() {
234        let array = StringArray::new(
235            vec!["red".into(), "green".into(), "blue".into()],
236            vec![1, 3],
237        )
238        .unwrap();
239        let result =
240            strcmp_builtin(Value::StringArray(array), Value::String("green".into())).expect("cmp");
241        let expected = LogicalArray::new(vec![0, 1, 0], vec![1, 3]).unwrap();
242        assert_eq!(result, Value::LogicalArray(expected));
243    }
244
245    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
246    #[test]
247    fn strcmp_char_array_row_compare() {
248        let chars = CharArray::new(vec!['c', 'a', 't', 'd', 'o', 'g'], 2, 3).unwrap();
249        let result =
250            strcmp_builtin(Value::CharArray(chars), Value::String("cat".into())).expect("cmp");
251        let expected = LogicalArray::new(vec![1, 0], vec![2, 1]).unwrap();
252        assert_eq!(result, Value::LogicalArray(expected));
253    }
254
255    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
256    #[test]
257    fn strcmp_char_array_to_char_array() {
258        let left = CharArray::new(vec!['a', 'b', 'c', 'd'], 2, 2).unwrap();
259        let right = CharArray::new(vec!['a', 'b', 'x', 'y'], 2, 2).unwrap();
260        let result =
261            strcmp_builtin(Value::CharArray(left), Value::CharArray(right)).expect("strcmp");
262        let expected = LogicalArray::new(vec![1, 0], vec![2, 1]).unwrap();
263        assert_eq!(result, Value::LogicalArray(expected));
264    }
265
266    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
267    #[test]
268    fn strcmp_cell_array_scalar() {
269        let cell = CellArray::new(
270            vec![
271                Value::from("apple"),
272                Value::from("pear"),
273                Value::from("grape"),
274            ],
275            1,
276            3,
277        )
278        .unwrap();
279        let result =
280            strcmp_builtin(Value::Cell(cell), Value::String("grape".into())).expect("strcmp");
281        let expected = LogicalArray::new(vec![0, 0, 1], vec![1, 3]).unwrap();
282        assert_eq!(result, Value::LogicalArray(expected));
283    }
284
285    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
286    #[test]
287    fn strcmp_cell_array_to_cell_array_broadcasts() {
288        let left = CellArray::new(vec![Value::from("red"), Value::from("blue")], 2, 1).unwrap();
289        let right = CellArray::new(vec![Value::from("red")], 1, 1).unwrap();
290        let result = strcmp_builtin(Value::Cell(left), Value::Cell(right)).expect("strcmp");
291        let expected = LogicalArray::new(vec![1, 0], vec![2, 1]).unwrap();
292        assert_eq!(result, Value::LogicalArray(expected));
293    }
294
295    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
296    #[test]
297    fn strcmp_string_array_multi_dimensional_broadcast() {
298        let left = StringArray::new(vec!["north".into(), "south".into()], vec![2, 1]).unwrap();
299        let right = StringArray::new(
300            vec!["north".into(), "east".into(), "south".into()],
301            vec![1, 3],
302        )
303        .unwrap();
304        let result =
305            strcmp_builtin(Value::StringArray(left), Value::StringArray(right)).expect("strcmp");
306        let expected = LogicalArray::new(vec![1, 0, 0, 0, 0, 1], vec![2, 3]).unwrap();
307        assert_eq!(result, Value::LogicalArray(expected));
308    }
309
310    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
311    #[test]
312    fn strcmp_char_array_trailing_space_is_not_equal() {
313        let chars = CharArray::new(vec!['c', 'a', 't', ' '], 1, 4).unwrap();
314        let result =
315            strcmp_builtin(Value::CharArray(chars), Value::String("cat".into())).expect("strcmp");
316        assert_eq!(result, Value::Bool(false));
317    }
318
319    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
320    #[test]
321    fn strcmp_char_array_empty_rows_returns_empty() {
322        let chars = CharArray::new(Vec::new(), 0, 0).unwrap();
323        let result = strcmp_builtin(Value::CharArray(chars), Value::String("anything".into()))
324            .expect("strcmp");
325        match result {
326            Value::LogicalArray(array) => {
327                assert_eq!(array.shape, vec![0, 1]);
328                assert!(array.data.is_empty());
329            }
330            other => panic!("expected empty logical array, got {other:?}"),
331        }
332    }
333
334    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
335    #[test]
336    fn strcmp_missing_strings_compare_false() {
337        let strings = StringArray::new(vec!["<missing>".into()], vec![1, 1]).unwrap();
338        let result = strcmp_builtin(
339            Value::StringArray(strings.clone()),
340            Value::StringArray(strings),
341        )
342        .expect("strcmp");
343        assert_eq!(result, Value::Bool(false));
344    }
345
346    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
347    #[test]
348    fn strcmp_missing_string_false() {
349        let array = StringArray::new(vec!["alpha".into(), "<missing>".into()], vec![1, 2]).unwrap();
350        let result =
351            strcmp_builtin(Value::StringArray(array), Value::String("alpha".into())).expect("cmp");
352        let expected = LogicalArray::new(vec![1, 0], vec![1, 2]).unwrap();
353        assert_eq!(result, Value::LogicalArray(expected));
354    }
355
356    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
357    #[test]
358    fn strcmp_size_mismatch_error() {
359        let left = StringArray::new(vec!["a".into(), "b".into()], vec![2, 1]).unwrap();
360        let right = StringArray::new(vec!["a".into(), "b".into(), "c".into()], vec![3, 1]).unwrap();
361        let err = error_message(
362            strcmp_builtin(Value::StringArray(left), Value::StringArray(right))
363                .expect_err("size mismatch"),
364        );
365        assert!(err.contains(STRCMP_ERROR_SHAPE_MISMATCH.message));
366    }
367
368    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
369    #[test]
370    fn strcmp_unsupported_numeric_argument_returns_false() {
371        let result =
372            strcmp_builtin(Value::Num(1.0), Value::String("a".into())).expect("comparison");
373        assert_eq!(result, Value::Bool(false));
374    }
375
376    #[test]
377    fn strcmp_type_is_logical_match() {
378        assert_eq!(
379            logical_text_match_type(
380                &[Type::String, Type::String],
381                &ResolveContext::new(Vec::new()),
382            ),
383            Type::Bool
384        );
385    }
386}