1use runmat_builtins::{BuiltinIntegerAuditDescriptor, BuiltinIntegerAuditKind};
4use std::borrow::Cow;
5
6use runmat_builtins::{
7 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
8 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
9};
10use runmat_macros::runtime_builtin;
11use runmat_value::{CellArray, CharArray, StringArray, Tensor, Value};
12
13use crate::builtins::common::map_control_flow_with_builtin;
14use crate::builtins::common::spec::{
15 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
16 ReductionNaN, ResidencyPolicy, ShapeRequirements,
17};
18use crate::builtins::common::tensor;
19use crate::builtins::strings::common::contains_numeric_or_resident_text_input;
20use crate::builtins::strings::type_resolvers::numeric_text_scalar_or_tensor_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::str2double")]
24pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
25 name: "str2double",
26 op_kind: GpuOpKind::Custom("conversion"),
27 supported_precisions: &[],
28 broadcast: BroadcastSemantics::None,
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: "Parses text on the CPU; GPU-resident inputs are gathered before conversion.",
37};
38
39#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::strings::core::str2double")]
40pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
41 name: "str2double",
42 shape: ShapeRequirements::Any,
43 constant_strategy: ConstantStrategy::InlineLiteral,
44 elementwise: None,
45 reduction: None,
46 emits_nan: true,
47 notes: "Conversion builtin; not eligible for fusion and materialises host-side doubles.",
48};
49
50const BUILTIN_NAME: &str = "str2double";
51
52const STR2DOUBLE_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
53 name: "X",
54 ty: BuiltinParamType::NumericArray,
55 arity: BuiltinParamArity::Required,
56 default: None,
57 description: "Parsed double values; invalid parses become NaN.",
58}];
59
60const STR2DOUBLE_INPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
61 name: "str",
62 ty: BuiltinParamType::Any,
63 arity: BuiltinParamArity::Required,
64 default: None,
65 description: "String, character, or cell-array text input to parse.",
66}];
67
68const STR2DOUBLE_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
69 label: "X = str2double(str)",
70 inputs: &STR2DOUBLE_INPUT,
71 outputs: &STR2DOUBLE_OUTPUT,
72}];
73
74const STR2DOUBLE_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
75 code: "RM.STR2DOUBLE.INVALID_INPUT",
76 identifier: Some("RunMat:str2double:InvalidInput"),
77 when: "Input is not a supported text container.",
78 message: "str2double: input must be a string array, character array, or cell array of character vectors",
79};
80
81const STR2DOUBLE_ERROR_INVALID_CELL_ELEMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
82 code: "RM.STR2DOUBLE.INVALID_CELL_ELEMENT",
83 identifier: Some("RunMat:str2double:InvalidCellElement"),
84 when: "Cell array contains non-text or non-scalar text entries.",
85 message: "str2double: cell array elements must be character vectors or string scalars",
86};
87
88const STR2DOUBLE_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
89 code: "RM.STR2DOUBLE.INTERNAL",
90 identifier: Some("RunMat:str2double:InternalError"),
91 when: "Internal tensor assembly failed while building parsed output.",
92 message: "str2double: internal error",
93};
94
95const STR2DOUBLE_ERRORS: [BuiltinErrorDescriptor; 3] = [
96 STR2DOUBLE_ERROR_INVALID_INPUT,
97 STR2DOUBLE_ERROR_INVALID_CELL_ELEMENT,
98 STR2DOUBLE_ERROR_INTERNAL,
99];
100
101pub const STR2DOUBLE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
102 signatures: &STR2DOUBLE_SIGNATURES,
103 output_mode: BuiltinOutputMode::Fixed,
104 completion_policy: BuiltinCompletionPolicy::Public,
105 errors: &STR2DOUBLE_ERRORS,
106};
107
108pub const STR2DOUBLE_INTEGER_AUDIT: BuiltinIntegerAuditDescriptor =
109 BuiltinIntegerAuditDescriptor {
110 kind: BuiltinIntegerAuditKind::NotApplicable,
111 canonical_builtin: None,
112 notes: "str2double parses string, character, or cellstr input and returns double. Integer, numeric, and provider-resident numeric inputs reject before provider access rather than being implicitly converted to text.",
113 };
114
115fn str2double_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
116 str2double_error_with_message(error.message, error)
117}
118
119fn str2double_error_with_message(
120 message: impl Into<String>,
121 error: &'static BuiltinErrorDescriptor,
122) -> RuntimeError {
123 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
124 if let Some(identifier) = error.identifier {
125 builder = builder.with_identifier(identifier);
126 }
127 builder.build()
128}
129
130fn remap_str2double_flow(err: RuntimeError) -> RuntimeError {
131 map_control_flow_with_builtin(err, BUILTIN_NAME)
132}
133
134#[runtime_builtin(
135 name = "str2double",
136 category = "strings/core",
137 summary = "Convert text representations of numbers into double-precision values.",
138 keywords = "str2double,string to double,text conversion,gpu",
139 accel = "sink",
140 type_resolver(numeric_text_scalar_or_tensor_type),
141 descriptor(crate::builtins::strings::core::str2double::STR2DOUBLE_DESCRIPTOR),
142 integer_audit(crate::builtins::strings::core::str2double::STR2DOUBLE_INTEGER_AUDIT),
143 builtin_path = "crate::builtins::strings::core::str2double"
144)]
145async fn str2double_builtin(value: Value) -> crate::BuiltinResult<Value> {
146 if contains_numeric_or_resident_text_input(&value) {
147 return Err(str2double_error(&STR2DOUBLE_ERROR_INVALID_INPUT));
148 }
149 let gathered = gather_if_needed_async(&value)
150 .await
151 .map_err(remap_str2double_flow)?;
152 match gathered {
153 Value::String(text) => Ok(Value::Num(parse_numeric_scalar(&text))),
154 Value::StringArray(array) => str2double_string_array(array),
155 Value::CharArray(array) => str2double_char_array(array),
156 Value::Cell(cell) => str2double_cell_array(cell),
157 _ => Err(str2double_error(&STR2DOUBLE_ERROR_INVALID_INPUT)),
158 }
159}
160
161fn str2double_string_array(array: StringArray) -> BuiltinResult<Value> {
162 let StringArray { data, shape, .. } = array;
163 let mut values = Vec::with_capacity(data.len());
164 for text in &data {
165 values.push(parse_numeric_scalar(text));
166 }
167 let tensor =
168 Tensor::new(values, shape).map_err(|_| str2double_error(&STR2DOUBLE_ERROR_INTERNAL))?;
169 Ok(tensor::tensor_into_value(tensor))
170}
171
172fn str2double_char_array(array: CharArray) -> BuiltinResult<Value> {
173 let rows = array.rows;
174 let cols = array.cols;
175 let mut values = Vec::with_capacity(rows);
176 for row in 0..rows {
177 let start = row * cols;
178 let end = start + cols;
179 let row_text: String = array.data[start..end].iter().collect();
180 values.push(parse_numeric_scalar(&row_text));
181 }
182 let tensor = Tensor::new(values, vec![rows, 1])
183 .map_err(|_| str2double_error(&STR2DOUBLE_ERROR_INTERNAL))?;
184 Ok(tensor::tensor_into_value(tensor))
185}
186
187fn str2double_cell_array(cell: CellArray) -> BuiltinResult<Value> {
188 let CellArray {
189 data, rows, cols, ..
190 } = cell;
191 let mut values = Vec::with_capacity(rows * cols);
192 for col in 0..cols {
193 for row in 0..rows {
194 let idx = row * cols + col;
195 let element: &Value = &data[idx];
196 let numeric = match element {
197 Value::String(text) => parse_numeric_scalar(text),
198 Value::StringArray(sa) if sa.data.len() == 1 => parse_numeric_scalar(&sa.data[0]),
199 Value::CharArray(char_vec) if char_vec.rows == 1 => {
200 let row_text: String = char_vec.data.iter().collect();
201 parse_numeric_scalar(&row_text)
202 }
203 Value::CharArray(_) => {
204 return Err(str2double_error(&STR2DOUBLE_ERROR_INVALID_CELL_ELEMENT));
205 }
206 _ => return Err(str2double_error(&STR2DOUBLE_ERROR_INVALID_CELL_ELEMENT)),
207 };
208 values.push(numeric);
209 }
210 }
211 let tensor = Tensor::new(values, vec![rows, cols])
212 .map_err(|_| str2double_error(&STR2DOUBLE_ERROR_INTERNAL))?;
213 Ok(tensor::tensor_into_value(tensor))
214}
215
216fn parse_numeric_scalar(text: &str) -> f64 {
217 let trimmed = text.trim();
218 if trimmed.is_empty() {
219 return f64::NAN;
220 }
221
222 let lowered = trimmed.to_ascii_lowercase();
223 match lowered.as_str() {
224 "nan" => return f64::NAN,
225 "inf" | "+inf" | "infinity" | "+infinity" => return f64::INFINITY,
226 "-inf" | "-infinity" => return f64::NEG_INFINITY,
227 _ => {}
228 }
229
230 let normalized: Cow<'_, str> = if trimmed.chars().any(|c| c == 'd' || c == 'D') {
231 Cow::Owned(
232 trimmed
233 .chars()
234 .map(|c| if c == 'd' || c == 'D' { 'e' } else { c })
235 .collect(),
236 )
237 } else {
238 Cow::Borrowed(trimmed)
239 };
240
241 normalized.parse::<f64>().unwrap_or(f64::NAN)
242}
243
244#[cfg(test)]
245pub(crate) mod tests {
246 use super::*;
247 use runmat_builtins::{ResolveContext, Type};
248
249 fn str2double_builtin(value: Value) -> BuiltinResult<Value> {
250 futures::executor::block_on(super::str2double_builtin(value))
251 }
252
253 fn error_message(err: crate::RuntimeError) -> String {
254 err.message().to_string()
255 }
256
257 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
258 #[test]
259 fn str2double_string_scalar() {
260 let result = str2double_builtin(Value::String("42.5".into())).expect("str2double");
261 assert_eq!(result, Value::Num(42.5));
262 }
263
264 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
265 #[test]
266 fn str2double_string_scalar_invalid_returns_nan() {
267 let result = str2double_builtin(Value::String("abc".into())).expect("str2double");
268 match result {
269 Value::Num(v) => assert!(v.is_nan()),
270 other => panic!("expected scalar result, got {other:?}"),
271 }
272 }
273
274 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
275 #[test]
276 fn str2double_string_array_preserves_shape() {
277 let array =
278 StringArray::new(vec!["1".into(), " 2.5 ".into(), "foo".into()], vec![3, 1]).unwrap();
279 let result = str2double_builtin(Value::StringArray(array)).expect("str2double");
280 match result {
281 Value::Tensor(tensor) => {
282 assert_eq!(tensor.shape, vec![3, 1]);
283 assert_eq!(tensor.materialize_f64()[0], 1.0);
284 assert_eq!(tensor.materialize_f64()[1], 2.5);
285 assert!(tensor.materialize_f64()[2].is_nan());
286 }
287 Value::Num(_) => panic!("expected tensor"),
288 other => panic!("unexpected result {other:?}"),
289 }
290 }
291
292 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
293 #[test]
294 fn str2double_char_array_multiple_rows() {
295 let data: Vec<char> = vec!['4', '2', ' ', ' ', '1', '0', '0', ' '];
296 let array = CharArray::new(data, 2, 4).unwrap();
297 let result = str2double_builtin(Value::CharArray(array)).expect("str2double");
298 match result {
299 Value::Tensor(tensor) => {
300 assert_eq!(tensor.shape, vec![2, 1]);
301 assert_eq!(tensor.materialize_f64()[0], 42.0);
302 assert_eq!(tensor.materialize_f64()[1], 100.0);
303 }
304 other => panic!("expected tensor result, got {other:?}"),
305 }
306 }
307
308 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
309 #[test]
310 fn str2double_char_array_empty_rows() {
311 let array = CharArray::new(Vec::new(), 0, 0).unwrap();
312 let result = str2double_builtin(Value::CharArray(array)).expect("str2double");
313 match result {
314 Value::Tensor(tensor) => {
315 assert_eq!(tensor.shape, vec![0, 1]);
316 assert_eq!(tensor.materialize_f64().len(), 0);
317 }
318 other => panic!("expected empty tensor, got {other:?}"),
319 }
320 }
321
322 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
323 #[test]
324 #[allow(
325 clippy::approx_constant,
326 reason = "Test ensures literal 3.14 text stays 3.14, not π"
327 )]
328 fn str2double_cell_array_of_text() {
329 let cell = CellArray::new(
330 vec![
331 Value::String("3.14".into()),
332 Value::CharArray(CharArray::new_row("NaN")),
333 Value::String("-Inf".into()),
334 ],
335 1,
336 3,
337 )
338 .unwrap();
339 let result = str2double_builtin(Value::Cell(cell)).expect("str2double");
340 match result {
341 Value::Tensor(tensor) => {
342 assert_eq!(tensor.shape, vec![1, 3]);
343 assert_eq!(tensor.materialize_f64()[0], 3.14);
344 assert!(tensor.materialize_f64()[1].is_nan());
345 assert_eq!(tensor.materialize_f64()[2], f64::NEG_INFINITY);
346 }
347 other => panic!("expected tensor result, got {other:?}"),
348 }
349 }
350
351 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
352 #[test]
353 fn str2double_cell_array_invalid_element_errors() {
354 let cell = CellArray::new(vec![Value::Num(5.0)], 1, 1).unwrap();
355 let err = error_message(str2double_builtin(Value::Cell(cell)).unwrap_err());
356 assert!(
357 err.contains("str2double"),
358 "unexpected error message: {err}"
359 );
360 }
361
362 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
363 #[test]
364 fn str2double_supports_d_exponent() {
365 let result = str2double_builtin(Value::String("1.5D3".into())).expect("str2double");
366 match result {
367 Value::Num(v) => assert_eq!(v, 1500.0),
368 other => panic!("expected scalar result, got {other:?}"),
369 }
370 }
371
372 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
373 #[test]
374 fn str2double_recognises_infinity_forms() {
375 let array = StringArray::new(
376 vec!["Inf".into(), "-Infinity".into(), "+inf".into()],
377 vec![3, 1],
378 )
379 .unwrap();
380 let result = str2double_builtin(Value::StringArray(array)).expect("str2double");
381 match result {
382 Value::Tensor(tensor) => {
383 assert_eq!(tensor.materialize_f64()[0], f64::INFINITY);
384 assert_eq!(tensor.materialize_f64()[1], f64::NEG_INFINITY);
385 assert_eq!(tensor.materialize_f64()[2], f64::INFINITY);
386 }
387 other => panic!("expected tensor result, got {other:?}"),
388 }
389 }
390
391 #[test]
392 fn str2double_type_is_numeric_text_scalar_or_tensor() {
393 assert_eq!(
394 numeric_text_scalar_or_tensor_type(&[Type::String], &ResolveContext::new(Vec::new())),
395 Type::Num
396 );
397 }
398}