1use runmat_builtins::{
4 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
6 CellArray, CharArray, ResolveContext, StringArray, StructValue, Tensor, Type, Value,
7};
8use runmat_macros::runtime_builtin;
9
10use crate::builtins::common::random_args::keyword_of;
11use crate::builtins::common::spec::{
12 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
13 ReductionNaN, ResidencyPolicy, ShapeRequirements,
14};
15use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
16
17const STATSET: &str = "statset";
18const STATGET: &str = "statget";
19const MAX_OPTION_INTEGER: usize = 1_000_000_000;
20
21const OPTION_FIELDS: [&str; 20] = [
22 "Display",
23 "MaxFunEvals",
24 "MaxIter",
25 "TolBnd",
26 "TolFun",
27 "TolTypeFun",
28 "TolX",
29 "TolTypeX",
30 "GradObj",
31 "Jacobian",
32 "DerivStep",
33 "FunValCheck",
34 "Robust",
35 "RobustWgtFun",
36 "WgtFun",
37 "Tune",
38 "UseParallel",
39 "UseSubstreams",
40 "Streams",
41 "OutputFcn",
42];
43
44const COMMON_STATFUNS: [&str; 48] = [
45 "bootci",
46 "bootstrp",
47 "crossval",
48 "factoran",
49 "fitglm",
50 "fitlm",
51 "fitlme",
52 "fitnlm",
53 "fitrgp",
54 "gamfit",
55 "gevfit",
56 "glmfit",
57 "gmdistribution",
58 "gpfit",
59 "kmeans",
60 "kmedoids",
61 "lasso",
62 "lassoglm",
63 "lognfit",
64 "mlecustom",
65 "mlecov",
66 "mvncdf",
67 "mvtcdf",
68 "nbinfit",
69 "nlinfit",
70 "nnmf",
71 "normfit",
72 "parallel",
73 "pca",
74 "plsregress",
75 "ppca",
76 "rocmetrics",
77 "sequentialfs",
78 "tsne",
79 "wblfit",
80 "copulafit",
81 "coxphfit",
82 "evfit",
83 "fitcox",
84 "fitglme",
85 "fitlmematrix",
86 "mdscale",
87 "nlmefitsa",
88 "treebagger",
89 "candexch",
90 "cordexch",
91 "daugment",
92 "dcovary",
93];
94
95const OUTPUT_OPTIONS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
96 name: "options",
97 ty: BuiltinParamType::Any,
98 arity: BuiltinParamArity::Required,
99 default: None,
100 description: "Statistics options structure.",
101}];
102
103const INPUT_STATFUN: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
104 name: "statfun",
105 ty: BuiltinParamType::StringScalar,
106 arity: BuiltinParamArity::Required,
107 default: None,
108 description: "Statistics function name.",
109}];
110
111const INPUT_PAIRS: [BuiltinParamDescriptor; 2] = [
112 BuiltinParamDescriptor {
113 name: "name",
114 ty: BuiltinParamType::StringScalar,
115 arity: BuiltinParamArity::Required,
116 default: None,
117 description: "Option field name.",
118 },
119 BuiltinParamDescriptor {
120 name: "value",
121 ty: BuiltinParamType::Any,
122 arity: BuiltinParamArity::Variadic,
123 default: None,
124 description: "Option value and additional name-value pairs.",
125 },
126];
127
128const INPUT_STRUCT_PAIRS: [BuiltinParamDescriptor; 3] = [
129 BuiltinParamDescriptor {
130 name: "oldopts",
131 ty: BuiltinParamType::Any,
132 arity: BuiltinParamArity::Required,
133 default: None,
134 description: "Existing statistics options structure.",
135 },
136 BuiltinParamDescriptor {
137 name: "name",
138 ty: BuiltinParamType::StringScalar,
139 arity: BuiltinParamArity::Optional,
140 default: None,
141 description: "Option field name or replacement options structure.",
142 },
143 BuiltinParamDescriptor {
144 name: "value",
145 ty: BuiltinParamType::Any,
146 arity: BuiltinParamArity::Variadic,
147 default: None,
148 description: "Option value and additional name-value pairs.",
149 },
150];
151
152const INPUT_OLD_NEW_OPTIONS: [BuiltinParamDescriptor; 2] = [
153 BuiltinParamDescriptor {
154 name: "oldopts",
155 ty: BuiltinParamType::Any,
156 arity: BuiltinParamArity::Required,
157 default: None,
158 description: "Existing statistics options structure.",
159 },
160 BuiltinParamDescriptor {
161 name: "newopts",
162 ty: BuiltinParamType::Any,
163 arity: BuiltinParamArity::Required,
164 default: None,
165 description: "Replacement statistics options structure. Nonempty fields override oldopts.",
166 },
167];
168
169const STATSET_SIGNATURES: [BuiltinSignatureDescriptor; 5] = [
170 BuiltinSignatureDescriptor {
171 label: "options = statset()",
172 inputs: &[],
173 outputs: &OUTPUT_OPTIONS,
174 },
175 BuiltinSignatureDescriptor {
176 label: "options = statset(statfun)",
177 inputs: &INPUT_STATFUN,
178 outputs: &OUTPUT_OPTIONS,
179 },
180 BuiltinSignatureDescriptor {
181 label: "options = statset(name, value, ...)",
182 inputs: &INPUT_PAIRS,
183 outputs: &OUTPUT_OPTIONS,
184 },
185 BuiltinSignatureDescriptor {
186 label: "options = statset(oldopts, newopts)",
187 inputs: &INPUT_OLD_NEW_OPTIONS,
188 outputs: &OUTPUT_OPTIONS,
189 },
190 BuiltinSignatureDescriptor {
191 label: "options = statset(oldopts, name, value, ...)",
192 inputs: &INPUT_STRUCT_PAIRS,
193 outputs: &OUTPUT_OPTIONS,
194 },
195];
196
197const STATGET_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
198 name: "val",
199 ty: BuiltinParamType::Any,
200 arity: BuiltinParamArity::Required,
201 default: None,
202 description: "Option field value.",
203}];
204
205const STATGET_INPUT_OPTIONS: BuiltinParamDescriptor = BuiltinParamDescriptor {
206 name: "options",
207 ty: BuiltinParamType::Any,
208 arity: BuiltinParamArity::Required,
209 default: None,
210 description: "Statistics options structure.",
211};
212const STATGET_INPUT_FIELD: BuiltinParamDescriptor = BuiltinParamDescriptor {
213 name: "field",
214 ty: BuiltinParamType::StringScalar,
215 arity: BuiltinParamArity::Required,
216 default: None,
217 description: "Option field name or unique leading prefix.",
218};
219const STATGET_INPUT_DEFAULT: BuiltinParamDescriptor = BuiltinParamDescriptor {
220 name: "defaultData",
221 ty: BuiltinParamType::Any,
222 arity: BuiltinParamArity::Optional,
223 default: None,
224 description: "Value returned when the selected option is empty.",
225};
226const STATGET_INPUTS_REQUIRED: [BuiltinParamDescriptor; 2] =
227 [STATGET_INPUT_OPTIONS, STATGET_INPUT_FIELD];
228const STATGET_INPUTS: [BuiltinParamDescriptor; 3] = [
229 STATGET_INPUT_OPTIONS,
230 STATGET_INPUT_FIELD,
231 STATGET_INPUT_DEFAULT,
232];
233
234const STATGET_SIGNATURES: [BuiltinSignatureDescriptor; 2] = [
235 BuiltinSignatureDescriptor {
236 label: "val = statget(options, field)",
237 inputs: &STATGET_INPUTS_REQUIRED,
238 outputs: &STATGET_OUTPUT,
239 },
240 BuiltinSignatureDescriptor {
241 label: "val = statget(options, field, defaultData)",
242 inputs: &STATGET_INPUTS,
243 outputs: &STATGET_OUTPUT,
244 },
245];
246
247const STATSET_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
248 code: "RM.STATSET.INVALID_ARGUMENT",
249 identifier: Some("RunMat:statset:InvalidArgument"),
250 when: "Argument grammar does not match supported statset forms.",
251 message: "statset: invalid argument",
252};
253const STATSET_ERROR_INVALID_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
254 code: "RM.STATSET.INVALID_OPTION",
255 identifier: Some("RunMat:statset:InvalidOption"),
256 when: "An option name or value is malformed.",
257 message: "statset: invalid option",
258};
259const STATSET_ERROR_INVALID_STATFUN: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
260 code: "RM.STATSET.INVALID_STATFUN",
261 identifier: Some("RunMat:statset:InvalidStatfun"),
262 when: "The statfun argument is not a supported statistics function name.",
263 message: "statset: invalid statistics function",
264};
265
266const STATSET_ERRORS: [BuiltinErrorDescriptor; 3] = [
267 STATSET_ERROR_INVALID_ARGUMENT,
268 STATSET_ERROR_INVALID_OPTION,
269 STATSET_ERROR_INVALID_STATFUN,
270];
271
272const STATGET_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
273 code: "RM.STATGET.INVALID_ARGUMENT",
274 identifier: Some("RunMat:statget:InvalidArgument"),
275 when: "Argument grammar does not match statget forms.",
276 message: "statget: invalid argument",
277};
278const STATGET_ERROR_INVALID_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
279 code: "RM.STATGET.INVALID_OPTION",
280 identifier: Some("RunMat:statget:InvalidOption"),
281 when: "The options argument is not a struct or the field argument is not text.",
282 message: "statget: invalid option",
283};
284
285const STATGET_ERRORS: [BuiltinErrorDescriptor; 2] =
286 [STATGET_ERROR_INVALID_ARGUMENT, STATGET_ERROR_INVALID_OPTION];
287
288pub const STATSET_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
289 signatures: &STATSET_SIGNATURES,
290 output_mode: BuiltinOutputMode::Fixed,
291 completion_policy: BuiltinCompletionPolicy::Public,
292 errors: &STATSET_ERRORS,
293};
294
295pub const STATGET_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
296 signatures: &STATGET_SIGNATURES,
297 output_mode: BuiltinOutputMode::Fixed,
298 completion_policy: BuiltinCompletionPolicy::Public,
299 errors: &STATGET_ERRORS,
300};
301
302#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::stats::options")]
303pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
304 name: "statset/statget",
305 op_kind: GpuOpKind::Custom("statistics-options"),
306 supported_precisions: &[],
307 broadcast: BroadcastSemantics::None,
308 provider_hooks: &[],
309 constant_strategy: ConstantStrategy::InlineLiteral,
310 residency: ResidencyPolicy::GatherImmediately,
311 nan_mode: ReductionNaN::Include,
312 two_pass_threshold: None,
313 workgroup_size: None,
314 accepts_nan_mode: false,
315 notes: "Host metadata construction and lookup. gpuArray option values are gathered before use.",
316};
317
318#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::stats::options")]
319pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
320 name: "statset/statget",
321 shape: ShapeRequirements::Any,
322 constant_strategy: ConstantStrategy::InlineLiteral,
323 elementwise: None,
324 reduction: None,
325 emits_nan: false,
326 notes: "Option struct construction and lookup are host metadata work and do not fuse.",
327};
328
329fn stat_options_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
330 Type::Struct {
331 known_fields: Some(
332 OPTION_FIELDS
333 .iter()
334 .map(|field| (*field).to_string())
335 .collect(),
336 ),
337 }
338}
339
340fn statget_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
341 Type::Unknown
342}
343
344#[runtime_builtin(
345 name = "statset",
346 category = "stats/options",
347 summary = "Create or update statistics options structures.",
348 keywords = "statset,statistics options,MaxIter,TolFun,TolX,Display,UseParallel",
349 accel = "cpu",
350 type_resolver(stat_options_type),
351 descriptor(crate::builtins::stats::options::STATSET_DESCRIPTOR),
352 builtin_path = "crate::builtins::stats::options"
353)]
354async fn statset_builtin(rest: Vec<Value>) -> BuiltinResult<Value> {
355 let args = gather_all(rest).await?;
356 Ok(Value::Struct(parse_statset(args)?))
357}
358
359#[runtime_builtin(
360 name = "statget",
361 category = "stats/options",
362 summary = "Access field values in statistics options structures.",
363 keywords = "statget,statset,statistics options",
364 accel = "cpu",
365 type_resolver(statget_type),
366 descriptor(crate::builtins::stats::options::STATGET_DESCRIPTOR),
367 builtin_path = "crate::builtins::stats::options"
368)]
369async fn statget_builtin(options: Value, field: Value, rest: Vec<Value>) -> BuiltinResult<Value> {
370 if rest.len() > 1 {
371 return Err(statget_error(
372 &STATGET_ERROR_INVALID_ARGUMENT,
373 "statget: expected at most one default value",
374 ));
375 }
376 let options = gather_if_needed_async(&options)
377 .await
378 .map_err(|err| statget_error(&STATGET_ERROR_INVALID_ARGUMENT, err.message()))?;
379 let field = gather_if_needed_async(&field)
380 .await
381 .map_err(|err| statget_error(&STATGET_ERROR_INVALID_ARGUMENT, err.message()))?;
382 let default_data = rest.into_iter().next();
383 let Value::Struct(options) = options else {
384 return Err(statget_error(
385 &STATGET_ERROR_INVALID_OPTION,
386 "statget: options must be a struct",
387 ));
388 };
389 let field = text_scalar(&field).map_err(|err| {
390 statget_error(
391 &STATGET_ERROR_INVALID_OPTION,
392 format!("statget: {}", err.message()),
393 )
394 })?;
395 let Some(canonical) = unique_option_match(&field) else {
396 return Ok(empty_numeric());
397 };
398 let value = lookup_struct_field(&options, canonical)
399 .cloned()
400 .unwrap_or_else(empty_numeric);
401 if is_empty_value(&value) {
402 if let Some(default_data) = default_data {
403 gather_if_needed_async(&default_data)
404 .await
405 .map_err(|err| statget_error(&STATGET_ERROR_INVALID_ARGUMENT, err.message()))
406 } else {
407 Ok(value)
408 }
409 } else {
410 Ok(value)
411 }
412}
413
414async fn gather_all(values: Vec<Value>) -> BuiltinResult<Vec<Value>> {
415 let mut out = Vec::with_capacity(values.len());
416 for value in values {
417 out.push(
418 gather_if_needed_async(&value)
419 .await
420 .map_err(|err| statset_error(&STATSET_ERROR_INVALID_ARGUMENT, err.message()))?,
421 );
422 }
423 Ok(out)
424}
425
426fn parse_statset(args: Vec<Value>) -> BuiltinResult<StructValue> {
427 if args.is_empty() {
428 return Ok(empty_options());
429 }
430 let mut index = 0usize;
431 let mut options;
432 match &args[0] {
433 Value::Struct(existing) => {
434 options = canonicalize_options(existing)?;
435 index = 1;
436 if index < args.len() {
437 if let Value::Struct(newopts) = &args[index] {
438 merge_old_into_new(&mut options, newopts)?;
439 index += 1;
440 }
441 }
442 }
443 first if looks_like_option_name(first) => {
444 options = empty_options();
445 }
446 first => {
447 let statfun = text_scalar(first).map_err(|err| {
448 statset_error(
449 &STATSET_ERROR_INVALID_ARGUMENT,
450 format!("statset: {}", err.message()),
451 )
452 })?;
453 options = defaults_for_statfun(&statfun)?;
454 index = 1;
455 }
456 }
457 let remaining = &args[index..];
458 if !remaining.is_empty() {
459 if !remaining.len().is_multiple_of(2) {
460 return Err(statset_error(
461 &STATSET_ERROR_INVALID_ARGUMENT,
462 "statset: expected option name-value pairs",
463 ));
464 }
465 for pair in remaining.chunks(2) {
466 let name = text_scalar(&pair[0]).map_err(|err| {
467 statset_error(
468 &STATSET_ERROR_INVALID_OPTION,
469 format!("statset: {}", err.message()),
470 )
471 })?;
472 let canonical = canonical_option_name(&name)?;
473 options.insert(canonical, validate_option_value(&name, &pair[1])?);
474 }
475 }
476 Ok(options)
477}
478
479fn merge_old_into_new(oldopts: &mut StructValue, newopts: &StructValue) -> BuiltinResult<()> {
480 let canonical_new = canonicalize_options(newopts)?;
481 for field in OPTION_FIELDS {
482 if let Some(new_value) = lookup_struct_field(&canonical_new, field) {
483 if !is_empty_value(new_value) {
484 oldopts.insert(field, new_value.clone());
485 }
486 }
487 }
488 Ok(())
489}
490
491fn canonicalize_options(options: &StructValue) -> BuiltinResult<StructValue> {
492 let mut out = empty_options();
493 for (name, value) in &options.fields {
494 let canonical = canonical_option_name(name)?;
495 out.insert(canonical, validate_option_value(name, value)?);
496 }
497 Ok(out)
498}
499
500fn defaults_for_statfun(statfun: &str) -> BuiltinResult<StructValue> {
501 let key = statfun.to_ascii_lowercase();
502 if !COMMON_STATFUNS.contains(&key.as_str()) {
503 return Err(statset_error(
504 &STATSET_ERROR_INVALID_STATFUN,
505 format!("statset: unsupported statistics function '{statfun}'"),
506 ));
507 }
508 let mut options = empty_options();
509 match key.as_str() {
510 "fitglm" | "fitlm" | "fitnlm" | "glmfit" | "lasso" | "lassoglm" | "nlinfit" | "normfit"
511 | "wblfit" => {
512 options.insert("Display", Value::from("off"));
513 options.insert("MaxIter", Value::Num(100.0));
514 options.insert("TolX", Value::Num(1.0e-6));
515 }
516 "factoran" => {
517 options.insert("Display", Value::from("off"));
518 options.insert("MaxIter", Value::Num(100.0));
519 options.insert("TolX", Value::Num(1.0e-8));
520 }
521 "nbinfit" => {
522 options.insert("Display", Value::from("off"));
523 options.insert("MaxFunEvals", Value::Num(400.0));
524 options.insert("MaxIter", Value::Num(200.0));
525 options.insert("TolBnd", Value::Num(1.0e-6));
526 options.insert("TolFun", Value::Num(1.0e-6));
527 options.insert("TolX", Value::Num(1.0e-6));
528 }
529 "kmeans" | "tsne" | "pca" | "ppca" | "gmdistribution" | "kmedoids" => {
530 options.insert("Display", Value::from("off"));
531 options.insert("MaxIter", Value::Num(100.0));
532 }
533 "bootci" | "bootstrp" | "crossval" | "parallel" | "sequentialfs" => {
534 options.insert("UseParallel", Value::Bool(false));
535 options.insert("UseSubstreams", Value::Bool(false));
536 options.insert("Streams", empty_cell());
537 }
538 _ => {}
539 }
540 Ok(options)
541}
542
543fn empty_options() -> StructValue {
544 let mut out = StructValue::new();
545 for field in OPTION_FIELDS {
546 out.insert(
547 field,
548 if field == "Streams" {
549 empty_cell()
550 } else {
551 empty_numeric()
552 },
553 );
554 }
555 out
556}
557
558fn canonical_option_name(name: &str) -> BuiltinResult<&'static str> {
559 let Some(canonical) = unique_option_match(name) else {
560 return Err(statset_error(
561 &STATSET_ERROR_INVALID_OPTION,
562 format!("statset: unknown option '{name}'"),
563 ));
564 };
565 Ok(canonical)
566}
567
568fn unique_option_match(name: &str) -> Option<&'static str> {
569 let needle = name.to_ascii_lowercase();
570 for field in OPTION_FIELDS {
571 if field.eq_ignore_ascii_case(name) {
572 return Some(field);
573 }
574 }
575 let mut found = None;
576 for field in OPTION_FIELDS {
577 if field.to_ascii_lowercase().starts_with(&needle) {
578 if found.is_some() {
579 return None;
580 }
581 found = Some(field);
582 }
583 }
584 found
585}
586
587fn validate_option_value(name: &str, value: &Value) -> BuiltinResult<Value> {
588 let canonical = canonical_option_name(name)?;
589 match canonical {
590 "Display" => one_of_text(canonical, value, &["off", "final", "iter"]),
591 "FunValCheck" | "GradObj" | "Jacobian" => one_of_text(canonical, value, &["off", "on"]),
592 "TolTypeFun" | "TolTypeX" => one_of_text(canonical, value, &["abs", "rel"]),
593 "RobustWgtFun" => validate_robust_weight(value),
594 "MaxFunEvals" | "MaxIter" => positive_integer_value(canonical, value),
595 "TolBnd" | "TolFun" | "TolX" | "Tune" => positive_scalar_value(canonical, value),
596 "DerivStep" => positive_numeric_value(canonical, value),
597 "UseParallel" | "UseSubstreams" => bool_or_on_off_value(canonical, value),
598 "Streams" | "OutputFcn" | "Robust" | "WgtFun" => Ok(value.clone()),
599 other => Err(statset_error(
600 &STATSET_ERROR_INVALID_OPTION,
601 format!("statset: unsupported option '{other}'"),
602 )),
603 }
604}
605
606fn validate_robust_weight(value: &Value) -> BuiltinResult<Value> {
607 if is_empty_value(value)
608 || matches!(
609 value,
610 Value::FunctionHandle(_)
611 | Value::ExternalFunctionHandle(_)
612 | Value::MethodFunctionHandle(_)
613 | Value::BoundFunctionHandle { .. }
614 )
615 {
616 return Ok(value.clone());
617 }
618 one_of_text(
619 "RobustWgtFun",
620 value,
621 &[
622 "andrews", "bisquare", "cauchy", "fair", "huber", "logistic", "talwar", "welsch",
623 ],
624 )
625}
626
627fn one_of_text(field: &str, value: &Value, allowed: &[&str]) -> BuiltinResult<Value> {
628 if is_empty_value(value) {
629 return Ok(value.clone());
630 }
631 let text = text_scalar(value)?;
632 let lower = text.to_ascii_lowercase();
633 if allowed.contains(&lower.as_str()) {
634 Ok(Value::from(lower))
635 } else {
636 Err(statset_error(
637 &STATSET_ERROR_INVALID_OPTION,
638 format!("statset: {field} must be one of {}", allowed.join(", ")),
639 ))
640 }
641}
642
643fn positive_integer_value(field: &str, value: &Value) -> BuiltinResult<Value> {
644 if is_empty_value(value) {
645 return Ok(value.clone());
646 }
647 let scalar = numeric_scalar(field, value)?;
648 if scalar < 1.0 || scalar.fract() != 0.0 || scalar > MAX_OPTION_INTEGER as f64 {
649 return Err(statset_error(
650 &STATSET_ERROR_INVALID_OPTION,
651 format!("statset: {field} must be a positive integer"),
652 ));
653 }
654 Ok(Value::Num(scalar))
655}
656
657fn positive_scalar_value(field: &str, value: &Value) -> BuiltinResult<Value> {
658 if is_empty_value(value) {
659 return Ok(value.clone());
660 }
661 let scalar = numeric_scalar(field, value)?;
662 if scalar <= 0.0 {
663 return Err(statset_error(
664 &STATSET_ERROR_INVALID_OPTION,
665 format!("statset: {field} must be a positive scalar"),
666 ));
667 }
668 Ok(Value::Num(scalar))
669}
670
671fn positive_numeric_value(field: &str, value: &Value) -> BuiltinResult<Value> {
672 if is_empty_value(value) {
673 return Ok(value.clone());
674 }
675 match value {
676 Value::Num(_) | Value::Int(_) | Value::Bool(_) => positive_scalar_value(field, value),
677 Value::Tensor(tensor) => {
678 if tensor
679 .data
680 .iter()
681 .all(|entry| entry.is_finite() && *entry > 0.0)
682 {
683 Ok(value.clone())
684 } else {
685 Err(statset_error(
686 &STATSET_ERROR_INVALID_OPTION,
687 format!("statset: {field} must contain positive finite values"),
688 ))
689 }
690 }
691 other => Err(statset_error(
692 &STATSET_ERROR_INVALID_OPTION,
693 format!("statset: {field} must be numeric, got {other:?}"),
694 )),
695 }
696}
697
698fn bool_or_on_off_value(field: &str, value: &Value) -> BuiltinResult<Value> {
699 if is_empty_value(value) {
700 return Ok(value.clone());
701 }
702 match value {
703 Value::Bool(flag) => Ok(Value::Bool(*flag)),
704 Value::Num(n) if *n == 0.0 || *n == 1.0 => Ok(Value::Bool(*n != 0.0)),
705 Value::Int(i) if i.to_f64() == 0.0 || i.to_f64() == 1.0 => {
706 Ok(Value::Bool(i.to_f64() != 0.0))
707 }
708 Value::Tensor(tensor)
709 if tensor.data.len() == 1 && (tensor.data[0] == 0.0 || tensor.data[0] == 1.0) =>
710 {
711 Ok(Value::Bool(tensor.data[0] != 0.0))
712 }
713 _ => {
714 let text = text_scalar(value)?;
715 match text.to_ascii_lowercase().as_str() {
716 "on" | "true" => Ok(Value::Bool(true)),
717 "off" | "false" => Ok(Value::Bool(false)),
718 _ => Err(statset_error(
719 &STATSET_ERROR_INVALID_OPTION,
720 format!("statset: {field} must be logical or 'on'/'off'"),
721 )),
722 }
723 }
724 }
725}
726
727fn numeric_scalar(field: &str, value: &Value) -> BuiltinResult<f64> {
728 let scalar = match value {
729 Value::Num(n) => *n,
730 Value::Int(i) => i.to_f64(),
731 Value::Bool(flag) => {
732 if *flag {
733 1.0
734 } else {
735 0.0
736 }
737 }
738 Value::Tensor(tensor) if tensor.data.len() == 1 => tensor.data[0],
739 other => {
740 return Err(statset_error(
741 &STATSET_ERROR_INVALID_OPTION,
742 format!("statset: {field} must be a numeric scalar, got {other:?}"),
743 ))
744 }
745 };
746 if !scalar.is_finite() {
747 return Err(statset_error(
748 &STATSET_ERROR_INVALID_OPTION,
749 format!("statset: {field} must be finite"),
750 ));
751 }
752 Ok(scalar)
753}
754
755fn text_scalar(value: &Value) -> BuiltinResult<String> {
756 if let Some(text) = keyword_of(value) {
757 return Ok(text);
758 }
759 match value {
760 Value::CharArray(CharArray { data, rows: 1, .. }) => Ok(data.iter().collect()),
761 Value::StringArray(StringArray { data, .. }) if data.len() == 1 => Ok(data[0].clone()),
762 other => Err(statset_error(
763 &STATSET_ERROR_INVALID_OPTION,
764 format!("option names must be text scalars, got {other:?}"),
765 )),
766 }
767}
768
769fn looks_like_option_name(value: &Value) -> bool {
770 text_scalar(value)
771 .ok()
772 .and_then(|text| unique_option_match(&text))
773 .is_some()
774}
775
776fn lookup_struct_field<'a>(options: &'a StructValue, name: &str) -> Option<&'a Value> {
777 options
778 .fields
779 .iter()
780 .find(|(field, _)| field.eq_ignore_ascii_case(name))
781 .map(|(_, value)| value)
782}
783
784fn is_empty_value(value: &Value) -> bool {
785 match value {
786 Value::Tensor(tensor) => tensor.data.is_empty(),
787 Value::LogicalArray(array) => array.data.is_empty(),
788 Value::Cell(cell) => cell.data.is_empty(),
789 Value::StringArray(array) => array.data.is_empty(),
790 Value::CharArray(array) => array.data.is_empty(),
791 _ => false,
792 }
793}
794
795fn empty_numeric() -> Value {
796 Value::Tensor(Tensor::new(Vec::new(), vec![0, 0]).expect("empty tensor"))
797}
798
799fn empty_cell() -> Value {
800 Value::Cell(CellArray::new(Vec::new(), 0, 0).expect("empty cell"))
801}
802
803fn statset_error(error: &'static BuiltinErrorDescriptor, detail: impl AsRef<str>) -> RuntimeError {
804 let detail = detail.as_ref();
805 let message = if detail.starts_with("statset:") {
806 detail.to_string()
807 } else {
808 format!("{}: {detail}", error.message)
809 };
810 let mut builder = build_runtime_error(message).with_builtin(STATSET);
811 if let Some(identifier) = error.identifier {
812 builder = builder.with_identifier(identifier);
813 }
814 builder.build()
815}
816
817fn statget_error(error: &'static BuiltinErrorDescriptor, detail: impl AsRef<str>) -> RuntimeError {
818 let detail = detail.as_ref();
819 let message = if detail.starts_with("statget:") {
820 detail.to_string()
821 } else {
822 format!("{}: {detail}", error.message)
823 };
824 let mut builder = build_runtime_error(message).with_builtin(STATGET);
825 if let Some(identifier) = error.identifier {
826 builder = builder.with_identifier(identifier);
827 }
828 builder.build()
829}
830
831#[cfg(test)]
832mod tests {
833 use super::*;
834 use futures::executor::block_on;
835
836 fn struct_value(value: Value) -> StructValue {
837 let Value::Struct(st) = value else {
838 panic!("expected struct, got {value:?}");
839 };
840 st
841 }
842
843 fn num_field(options: &StructValue, name: &str) -> f64 {
844 match options.fields.get(name).unwrap() {
845 Value::Num(value) => *value,
846 other => panic!("expected numeric field {name}, got {other:?}"),
847 }
848 }
849
850 #[test]
851 fn statset_builds_custom_options() {
852 let options = struct_value(
853 block_on(statset_builtin(vec![
854 Value::from("FunValCheck"),
855 Value::from("on"),
856 Value::from("TolX"),
857 Value::Num(1.0e-8),
858 Value::from("UseParallel"),
859 Value::from("off"),
860 ]))
861 .unwrap(),
862 );
863 assert!(matches!(options.fields.get("FunValCheck"), Some(Value::String(s)) if s == "on"));
864 assert_eq!(num_field(&options, "TolX"), 1.0e-8);
865 assert!(matches!(
866 options.fields.get("UseParallel"),
867 Some(Value::Bool(false))
868 ));
869 assert!(
870 matches!(options.fields.get("Streams"), Some(Value::Cell(cell)) if cell.data.is_empty())
871 );
872 }
873
874 #[test]
875 fn statset_applies_function_defaults_and_updates() {
876 let base = block_on(statset_builtin(vec![Value::from("nbinfit")])).unwrap();
877 let options = struct_value(
878 block_on(statset_builtin(vec![
879 base,
880 Value::from("TolX"),
881 Value::Num(1.0e-8),
882 ]))
883 .unwrap(),
884 );
885 assert_eq!(num_field(&options, "MaxIter"), 200.0);
886 assert_eq!(num_field(&options, "TolX"), 1.0e-8);
887 assert_eq!(num_field(&options, "TolFun"), 1.0e-6);
888 }
889
890 #[test]
891 fn statset_old_new_merge_prefers_nonempty_new_fields() {
892 let oldopts = struct_value(
893 block_on(statset_builtin(vec![
894 Value::from("TolX"),
895 Value::Num(1.0e-6),
896 Value::from("MaxIter"),
897 Value::Num(20.0),
898 ]))
899 .unwrap(),
900 );
901 let newopts = struct_value(
902 block_on(statset_builtin(vec![
903 Value::from("TolX"),
904 Value::Num(1.0e-9),
905 ]))
906 .unwrap(),
907 );
908 let merged = struct_value(
909 block_on(statset_builtin(vec![
910 Value::Struct(oldopts),
911 Value::Struct(newopts),
912 ]))
913 .unwrap(),
914 );
915 assert_eq!(num_field(&merged, "TolX"), 1.0e-9);
916 assert_eq!(num_field(&merged, "MaxIter"), 20.0);
917 }
918
919 #[test]
920 fn statget_supports_unique_prefix_and_default_for_empty() {
921 let options = block_on(statset_builtin(vec![
922 Value::from("TolX"),
923 Value::Num(1.0e-8),
924 Value::from("MaxIter"),
925 Value::Num(15.0),
926 ]))
927 .unwrap();
928 let value = block_on(statget_builtin(
929 options.clone(),
930 Value::from("TolX"),
931 Vec::new(),
932 ))
933 .unwrap();
934 assert_eq!(value, Value::Num(1.0e-8));
935
936 let value = block_on(statget_builtin(
937 options.clone(),
938 Value::from("MaxI"),
939 Vec::new(),
940 ))
941 .unwrap();
942 assert_eq!(value, Value::Num(15.0));
943
944 let value = block_on(statget_builtin(
945 options,
946 Value::from("TolFun"),
947 vec![Value::Num(3.0)],
948 ))
949 .unwrap();
950 assert_eq!(value, Value::Num(3.0));
951 }
952
953 #[test]
954 fn statset_rejects_invalid_values() {
955 let err = block_on(statset_builtin(vec![
956 Value::from("MaxIter"),
957 Value::Num(2.5),
958 ]))
959 .unwrap_err();
960 assert_eq!(err.identifier(), Some("RunMat:statset:InvalidOption"));
961
962 let err = block_on(statset_builtin(vec![
963 Value::from("Display"),
964 Value::from("verbose"),
965 ]))
966 .unwrap_err();
967 assert!(err.message.contains("Display"));
968 }
969
970 #[test]
971 fn descriptors_cover_public_forms() {
972 let statset_labels = STATSET_DESCRIPTOR
973 .signatures
974 .iter()
975 .map(|sig| sig.label)
976 .collect::<Vec<_>>();
977 assert!(statset_labels.contains(&"options = statset(statfun)"));
978 assert!(statset_labels.contains(&"options = statset(oldopts, newopts)"));
979
980 let statget_labels = STATGET_DESCRIPTOR
981 .signatures
982 .iter()
983 .map(|sig| sig.label)
984 .collect::<Vec<_>>();
985 assert_eq!(
986 statget_labels,
987 vec![
988 "val = statget(options, field)",
989 "val = statget(options, field, defaultData)",
990 ]
991 );
992 }
993}