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