Skip to main content

runmat_runtime/builtins/stats/
options.rs

1//! Statistics options structure helpers (`statset` / `statget`).
2
3use 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}