Skip to main content

runmat_runtime/builtins/strings/core/
strncmp.rs

1//! MATLAB-compatible `strncmp` builtin for RunMat.
2
3use runmat_builtins::{
4    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor, Value,
6};
7use runmat_macros::runtime_builtin;
8
9use crate::builtins::common::broadcast::{broadcast_index, broadcast_shapes, compute_strides};
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::common::tensor;
16use crate::builtins::strings::search::text_utils::{logical_result, TextCollection, TextElement};
17use crate::builtins::strings::type_resolvers::logical_text_match_type;
18use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
19
20const FN_NAME: &str = "strncmp";
21
22const STRNCMP_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
23    name: "tf",
24    ty: BuiltinParamType::LogicalArray,
25    arity: BuiltinParamArity::Required,
26    default: None,
27    description: "Logical prefix-comparison result.",
28}];
29
30const STRNCMP_INPUTS: [BuiltinParamDescriptor; 3] = [
31    BuiltinParamDescriptor {
32        name: "A",
33        ty: BuiltinParamType::Any,
34        arity: BuiltinParamArity::Required,
35        default: None,
36        description: "First text input (string/char/cell/string array).",
37    },
38    BuiltinParamDescriptor {
39        name: "B",
40        ty: BuiltinParamType::Any,
41        arity: BuiltinParamArity::Required,
42        default: None,
43        description: "Second text input (string/char/cell/string array).",
44    },
45    BuiltinParamDescriptor {
46        name: "N",
47        ty: BuiltinParamType::IntegerScalar,
48        arity: BuiltinParamArity::Required,
49        default: None,
50        description: "Prefix length to compare.",
51    },
52];
53
54const STRNCMP_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
55    label: "tf = strncmp(A, B, N)",
56    inputs: &STRNCMP_INPUTS,
57    outputs: &STRNCMP_OUTPUT,
58}];
59
60const STRNCMP_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
61    code: "RM.STRNCMP.INVALID_INPUT",
62    identifier: Some("RunMat:strncmp:InvalidInput"),
63    when: "At least one text input is not a supported text container.",
64    message: "strncmp: text inputs must be string/char/cell/string-array values",
65};
66
67const STRNCMP_ERROR_SHAPE_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
68    code: "RM.STRNCMP.SHAPE_MISMATCH",
69    identifier: Some("RunMat:strncmp:ShapeMismatch"),
70    when: "Text inputs are not broadcast-compatible.",
71    message: "strncmp: input sizes are not broadcast-compatible",
72};
73
74const STRNCMP_ERROR_INVALID_PREFIX_LENGTH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
75    code: "RM.STRNCMP.INVALID_PREFIX_LENGTH",
76    identifier: Some("RunMat:strncmp:InvalidPrefixLength"),
77    when: "Prefix length argument is not a finite nonnegative integer scalar.",
78    message: "strncmp: prefix length must be a finite nonnegative integer scalar",
79};
80
81const STRNCMP_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
82    code: "RM.STRNCMP.INTERNAL",
83    identifier: Some("RunMat:strncmp:InternalError"),
84    when: "Internal logical result assembly failed.",
85    message: "strncmp: internal error",
86};
87
88const STRNCMP_ERRORS: [BuiltinErrorDescriptor; 4] = [
89    STRNCMP_ERROR_INVALID_INPUT,
90    STRNCMP_ERROR_SHAPE_MISMATCH,
91    STRNCMP_ERROR_INVALID_PREFIX_LENGTH,
92    STRNCMP_ERROR_INTERNAL,
93];
94
95pub const STRNCMP_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
96    signatures: &STRNCMP_SIGNATURES,
97    output_mode: BuiltinOutputMode::Fixed,
98    completion_policy: BuiltinCompletionPolicy::Public,
99    errors: &STRNCMP_ERRORS,
100};
101
102#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::strings::core::strncmp")]
103pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
104    name: "strncmp",
105    op_kind: GpuOpKind::Custom("string-prefix-compare"),
106    supported_precisions: &[],
107    broadcast: BroadcastSemantics::Matlab,
108    provider_hooks: &[],
109    constant_strategy: ConstantStrategy::InlineLiteral,
110    residency: ResidencyPolicy::GatherImmediately,
111    nan_mode: ReductionNaN::Include,
112    two_pass_threshold: None,
113    workgroup_size: None,
114    accepts_nan_mode: false,
115    notes: "Performs host-side prefix comparisons; GPU inputs are gathered before evaluation.",
116};
117
118#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::strings::core::strncmp")]
119pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
120    name: "strncmp",
121    shape: ShapeRequirements::Any,
122    constant_strategy: ConstantStrategy::InlineLiteral,
123    elementwise: None,
124    reduction: None,
125    emits_nan: false,
126    notes: "Produces logical host results and is not eligible for GPU fusion.",
127};
128
129fn strncmp_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
130    strncmp_error_with_message(error.message, error)
131}
132
133fn strncmp_error_with_message(
134    message: impl Into<String>,
135    error: &'static BuiltinErrorDescriptor,
136) -> RuntimeError {
137    let mut builder = build_runtime_error(message).with_builtin(FN_NAME);
138    if let Some(identifier) = error.identifier {
139        builder = builder.with_identifier(identifier);
140    }
141    builder.build()
142}
143
144fn remap_strncmp_flow(err: RuntimeError) -> RuntimeError {
145    map_control_flow_with_builtin(err, FN_NAME)
146}
147
148#[runtime_builtin(
149    name = "strncmp",
150    category = "strings/core",
151    summary = "Compare text inputs case-sensitively up to N leading characters.",
152    keywords = "strncmp,string compare,prefix,text equality",
153    accel = "sink",
154    type_resolver(logical_text_match_type),
155    descriptor(crate::builtins::strings::core::strncmp::STRNCMP_DESCRIPTOR),
156    builtin_path = "crate::builtins::strings::core::strncmp"
157)]
158async fn strncmp_builtin(a: Value, b: Value, n: Value) -> crate::BuiltinResult<Value> {
159    let a = gather_if_needed_async(&a)
160        .await
161        .map_err(remap_strncmp_flow)?;
162    let b = gather_if_needed_async(&b)
163        .await
164        .map_err(remap_strncmp_flow)?;
165    let n = gather_if_needed_async(&n)
166        .await
167        .map_err(remap_strncmp_flow)?;
168
169    let limit = parse_prefix_length(n)?;
170    let left = TextCollection::from_argument(FN_NAME, a, "first argument")
171        .map_err(|_| strncmp_error(&STRNCMP_ERROR_INVALID_INPUT))?;
172    let right = TextCollection::from_argument(FN_NAME, b, "second argument")
173        .map_err(|_| strncmp_error(&STRNCMP_ERROR_INVALID_INPUT))?;
174    evaluate_strncmp(&left, &right, limit)
175}
176
177fn evaluate_strncmp(
178    left: &TextCollection,
179    right: &TextCollection,
180    limit: usize,
181) -> BuiltinResult<Value> {
182    let shape = broadcast_shapes(FN_NAME, &left.shape, &right.shape)
183        .map_err(|_| strncmp_error(&STRNCMP_ERROR_SHAPE_MISMATCH))?;
184    let total = tensor::element_count(&shape);
185    if total == 0 {
186        return logical_result(FN_NAME, Vec::new(), shape)
187            .map_err(|_| strncmp_error(&STRNCMP_ERROR_INTERNAL));
188    }
189
190    let left_strides = compute_strides(&left.shape);
191    let right_strides = compute_strides(&right.shape);
192    let mut data = Vec::with_capacity(total);
193
194    for linear in 0..total {
195        let li = broadcast_index(linear, &shape, &left.shape, &left_strides);
196        let ri = broadcast_index(linear, &shape, &right.shape, &right_strides);
197        let equal = if limit == 0 {
198            true
199        } else {
200            match (&left.elements[li], &right.elements[ri]) {
201                (TextElement::Missing, _) | (_, TextElement::Missing) => false,
202                (TextElement::Text(lhs), TextElement::Text(rhs)) => prefix_equal(lhs, rhs, limit),
203            }
204        };
205        data.push(if equal { 1 } else { 0 });
206    }
207
208    logical_result(FN_NAME, data, shape).map_err(|_| strncmp_error(&STRNCMP_ERROR_INTERNAL))
209}
210
211fn prefix_equal(lhs: &str, rhs: &str, limit: usize) -> bool {
212    if limit == 0 {
213        return true;
214    }
215    let mut lhs_iter = lhs.chars();
216    let mut rhs_iter = rhs.chars();
217    let mut compared = 0usize;
218
219    while compared < limit {
220        let left_char = lhs_iter.next();
221        let right_char = rhs_iter.next();
222        match (left_char, right_char) {
223            (Some(lc), Some(rc)) => {
224                if lc != rc {
225                    return false;
226                }
227            }
228            (None, Some(_)) | (Some(_), None) => {
229                return false;
230            }
231            (None, None) => {
232                return true;
233            }
234        }
235        compared += 1;
236    }
237
238    true
239}
240
241fn parse_prefix_length(value: Value) -> BuiltinResult<usize> {
242    match value {
243        Value::Int(i) => {
244            let raw = i.to_i64();
245            if raw < 0 {
246                return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
247            }
248            Ok(raw as usize)
249        }
250        Value::Num(n) => parse_prefix_length_from_float(n),
251        Value::Bool(b) => Ok(if b { 1 } else { 0 }),
252        Value::Tensor(tensor) => {
253            if tensor.data.len() != 1 {
254                return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
255            }
256            parse_prefix_length_from_float(tensor.data[0])
257        }
258        Value::LogicalArray(array) => {
259            if array.data.len() != 1 {
260                return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
261            }
262            Ok(if array.data[0] != 0 { 1 } else { 0 })
263        }
264        _ => Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH)),
265    }
266}
267
268fn parse_prefix_length_from_float(value: f64) -> BuiltinResult<usize> {
269    if !value.is_finite() {
270        return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
271    }
272    if value < 0.0 {
273        return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
274    }
275    let rounded = value.round();
276    if (rounded - value).abs() > f64::EPSILON {
277        return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
278    }
279    if rounded > (usize::MAX as f64) {
280        return Err(strncmp_error(&STRNCMP_ERROR_INVALID_PREFIX_LENGTH));
281    }
282    Ok(rounded as usize)
283}
284
285#[cfg(test)]
286pub(crate) mod tests {
287    use super::*;
288    #[cfg(feature = "wgpu")]
289    use runmat_accelerate_api::AccelProvider;
290    use runmat_builtins::{
291        CellArray, CharArray, IntValue, LogicalArray, ResolveContext, StringArray, Tensor, Type,
292    };
293
294    fn strncmp_builtin(a: Value, b: Value, n: Value) -> BuiltinResult<Value> {
295        futures::executor::block_on(super::strncmp_builtin(a, b, n))
296    }
297
298    fn error_message(err: crate::RuntimeError) -> String {
299        err.to_string()
300    }
301
302    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
303    #[test]
304    fn strncmp_exact_prefix_true() {
305        let result = strncmp_builtin(
306            Value::String("RunMat".into()),
307            Value::String("Runway".into()),
308            Value::Int(IntValue::I32(3)),
309        )
310        .expect("strncmp");
311        assert_eq!(result, Value::Bool(true));
312    }
313
314    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
315    #[test]
316    fn strncmp_mismatch_within_prefix_false() {
317        let result = strncmp_builtin(
318            Value::String("RunMat".into()),
319            Value::String("Runway".into()),
320            Value::Int(IntValue::I32(4)),
321        )
322        .expect("strncmp");
323        assert_eq!(result, Value::Bool(false));
324    }
325
326    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
327    #[test]
328    fn strncmp_longer_string_after_prefix_false() {
329        let result = strncmp_builtin(
330            Value::String("cat".into()),
331            Value::String("cater".into()),
332            Value::Int(IntValue::I32(4)),
333        )
334        .expect("strncmp");
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 strncmp_zero_length_always_true() {
341        let result = strncmp_builtin(
342            Value::String("alpha".into()),
343            Value::String("omega".into()),
344            Value::Num(0.0),
345        )
346        .expect("strncmp");
347        assert_eq!(result, Value::Bool(true));
348    }
349
350    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
351    #[test]
352    fn strncmp_prefix_length_bool_true_compares_first_character() {
353        let result = strncmp_builtin(
354            Value::String("alpha".into()),
355            Value::String("array".into()),
356            Value::Bool(true),
357        )
358        .expect("strncmp");
359        assert_eq!(result, Value::Bool(true));
360    }
361
362    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
363    #[test]
364    fn strncmp_prefix_length_bool_false_treated_as_zero() {
365        let result = strncmp_builtin(
366            Value::String("alpha".into()),
367            Value::String("omega".into()),
368            Value::Bool(false),
369        )
370        .expect("strncmp");
371        assert_eq!(result, Value::Bool(true));
372    }
373
374    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
375    #[test]
376    fn strncmp_prefix_length_logical_array_scalar() {
377        let logical = LogicalArray::new(vec![1], vec![1]).unwrap();
378        let result = strncmp_builtin(
379            Value::String("beta".into()),
380            Value::String("theta".into()),
381            Value::LogicalArray(logical),
382        )
383        .expect("strncmp");
384        assert_eq!(result, Value::Bool(false));
385    }
386
387    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
388    #[test]
389    fn strncmp_prefix_length_tensor_scalar_double() {
390        let limit = Tensor::new(vec![2.0], vec![1, 1]).unwrap();
391        let result = strncmp_builtin(
392            Value::String("gamma".into()),
393            Value::String("gamut".into()),
394            Value::Tensor(limit),
395        )
396        .expect("strncmp");
397        assert_eq!(result, Value::Bool(true));
398    }
399
400    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
401    #[test]
402    fn strncmp_char_array_rows() {
403        let chars = CharArray::new(
404            vec![
405                'c', 'a', 't', ' ', ' ', 'c', 'a', 'm', 'e', 'l', 'c', 'o', 'w', ' ', ' ',
406            ],
407            3,
408            5,
409        )
410        .unwrap();
411        let result = strncmp_builtin(
412            Value::CharArray(chars),
413            Value::String("ca".into()),
414            Value::Int(IntValue::I32(2)),
415        )
416        .expect("strncmp");
417        let expected = LogicalArray::new(vec![1, 1, 0], vec![3, 1]).unwrap();
418        assert_eq!(result, Value::LogicalArray(expected));
419    }
420
421    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
422    #[test]
423    fn strncmp_cell_arrays_broadcast() {
424        let left = CellArray::new(
425            vec![
426                Value::from("red"),
427                Value::from("green"),
428                Value::from("blue"),
429            ],
430            1,
431            3,
432        )
433        .unwrap();
434        let right = CellArray::new(
435            vec![
436                Value::from("rose"),
437                Value::from("gray"),
438                Value::from("black"),
439            ],
440            1,
441            3,
442        )
443        .unwrap();
444        let result = strncmp_builtin(
445            Value::Cell(left),
446            Value::Cell(right),
447            Value::Int(IntValue::I32(2)),
448        )
449        .expect("strncmp");
450        let expected = LogicalArray::new(vec![0, 1, 1], vec![1, 3]).unwrap();
451        assert_eq!(result, Value::LogicalArray(expected));
452    }
453
454    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
455    #[test]
456    fn strncmp_string_array_broadcast_scalar() {
457        let strings = StringArray::new(
458            vec!["north".into(), "south".into(), "east".into()],
459            vec![1, 3],
460        )
461        .unwrap();
462        let result = strncmp_builtin(
463            Value::StringArray(strings),
464            Value::String("no".into()),
465            Value::Int(IntValue::I32(2)),
466        )
467        .expect("strncmp");
468        let expected = LogicalArray::new(vec![1, 0, 0], vec![1, 3]).unwrap();
469        assert_eq!(result, Value::LogicalArray(expected));
470    }
471
472    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
473    #[test]
474    fn strncmp_missing_string_false_when_prefix_positive() {
475        let strings =
476            StringArray::new(vec!["<missing>".into(), "value".into()], vec![1, 2]).unwrap();
477        let result = strncmp_builtin(
478            Value::StringArray(strings),
479            Value::String("val".into()),
480            Value::Int(IntValue::I32(3)),
481        )
482        .expect("strncmp");
483        let expected = LogicalArray::new(vec![0, 1], vec![1, 2]).unwrap();
484        assert_eq!(result, Value::LogicalArray(expected));
485    }
486
487    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
488    #[test]
489    fn strncmp_missing_zero_length_true() {
490        let strings = StringArray::new(vec!["<missing>".into()], vec![1, 1]).unwrap();
491        let result = strncmp_builtin(
492            Value::StringArray(strings),
493            Value::String("anything".into()),
494            Value::Int(IntValue::I32(0)),
495        )
496        .expect("strncmp");
497        assert_eq!(result, Value::Bool(true));
498    }
499
500    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
501    #[test]
502    fn strncmp_size_mismatch_error() {
503        let left = StringArray::new(vec!["a".into(), "b".into()], vec![2, 1]).unwrap();
504        let right = StringArray::new(vec!["a".into(), "b".into(), "c".into()], vec![3, 1]).unwrap();
505        let err = error_message(
506            strncmp_builtin(
507                Value::StringArray(left),
508                Value::StringArray(right),
509                Value::Int(IntValue::I32(1)),
510            )
511            .expect_err("size mismatch"),
512        );
513        assert!(err.contains(STRNCMP_ERROR_SHAPE_MISMATCH.message));
514    }
515
516    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
517    #[test]
518    fn strncmp_invalid_length_type_errors() {
519        let err = error_message(
520            strncmp_builtin(
521                Value::String("abc".into()),
522                Value::String("abc".into()),
523                Value::String("3".into()),
524            )
525            .expect_err("invalid prefix length"),
526        );
527        assert!(err.contains("prefix length"));
528    }
529
530    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
531    #[test]
532    fn strncmp_negative_length_errors() {
533        let err = error_message(
534            strncmp_builtin(
535                Value::String("abc".into()),
536                Value::String("abc".into()),
537                Value::Num(-1.0),
538            )
539            .expect_err("negative length"),
540        );
541        assert!(err.to_ascii_lowercase().contains("nonnegative"));
542    }
543
544    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
545    #[test]
546    #[cfg(feature = "wgpu")]
547    fn strncmp_prefix_length_from_gpu_tensor() {
548        use runmat_accelerate::backend::wgpu::provider::{
549            register_wgpu_provider, WgpuProviderOptions,
550        };
551        use runmat_accelerate_api::HostTensorView;
552
553        let provider = match register_wgpu_provider(WgpuProviderOptions::default()) {
554            Ok(provider) => provider,
555            Err(_) => return,
556        };
557        let tensor = Tensor::new(vec![3.0], vec![1, 1]).unwrap();
558        let view = HostTensorView {
559            data: &tensor.data,
560            shape: &tensor.shape,
561        };
562        let handle = provider.upload(&view).expect("upload prefix length to GPU");
563        let result = strncmp_builtin(
564            Value::String("delta".into()),
565            Value::String("deluge".into()),
566            Value::GpuTensor(handle.clone()),
567        )
568        .expect("strncmp");
569        assert_eq!(result, Value::Bool(true));
570        let _ = provider.free(&handle);
571    }
572
573    #[test]
574    fn strncmp_type_is_logical_match() {
575        assert_eq!(
576            logical_text_match_type(
577                &[Type::String, Type::String],
578                &ResolveContext::new(Vec::new()),
579            ),
580            Type::Bool
581        );
582    }
583}