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::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}