1use 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}