Skip to main content

gam_config/
lib.rs

1use gam_inference::formula_dsl::parse_link_choice;
2use gam_inference::model::GroupMetadata;
3use gam_models::fit_orchestration::descriptors::build_analytic_penalty_registry_from_descriptors;
4use gam_models::fit_orchestration::{CtnStage1Recipe, FitConfig};
5use gam_models::survival::location_scale::residual_distribution_inverse_link;
6use gam_models::survival::lognormal_kernel::{FrailtySpec, HazardLoading};
7use gam_models::survival::parse_survival_distribution;
8use gam_models::survival::parse_survival_likelihood_mode;
9use gam_models::transformation_normal::TransformationNormalConfig;
10use gam_problem::types::{
11    InverseLink, LinkComponent, LinkFunction, MixtureLinkSpec, SasLinkSpec, StandardLink,
12};
13use gam_solve::mixture_link::{state_from_beta_logisticspec, state_from_sasspec, state_fromspec};
14use ndarray::Array1;
15
16mod fit_request_document;
17
18pub use fit_request_document::{
19    AnalyticPenaltiesDocument, CtnStage1ConfigDocument, CtnStage1Document, FIT_REQUEST_SCHEMA,
20    FIT_REQUEST_SCHEMA_VERSION, FitRequestConfigDocument, FitRequestDocument,
21    LatentCoordinateDocument, LatentCoordinatesDocument, PrecisionHyperpriorDocument,
22    SmoothDescriptorsDocument,
23};
24
25#[derive(Clone, Copy, Debug, Eq, PartialEq)]
26pub enum CliFrailtyKind {
27    GaussianShift,
28    HazardMultiplier,
29}
30
31#[derive(Clone, Copy, Debug, Eq, PartialEq)]
32pub enum CliHazardLoading {
33    Full,
34    LoadedVsUnloaded,
35}
36
37impl CtnStage1Document {
38    fn into_recipe(self) -> Result<CtnStage1Recipe, String> {
39        let mut config = TransformationNormalConfig::default();
40        if let Some(overrides) = self.config {
41            if let Some(value) = overrides.response_degree {
42                config.response_degree = value;
43            }
44            if let Some(value) = overrides.response_num_internal_knots {
45                config.response_num_internal_knots = value;
46            }
47            if let Some(value) = overrides.response_penalty_order {
48                config.response_penalty_order = value;
49            }
50            if let Some(value) = overrides.response_extra_penalty_orders {
51                config.response_extra_penalty_orders = value;
52            }
53            if let Some(value) = overrides.double_penalty {
54                config.double_penalty = value;
55            }
56        }
57        if config.response_degree == 0 {
58            return Err("ctn_stage1.config.response_degree must be >= 1".to_string());
59        }
60        if config.response_num_internal_knots < 2 {
61            return Err("ctn_stage1.config.response_num_internal_knots must be >= 2".to_string());
62        }
63        if config.response_penalty_order == 0
64            || config
65                .response_extra_penalty_orders
66                .iter()
67                .any(|order| *order == 0)
68        {
69            return Err("ctn_stage1 response penalty orders must be >= 1".to_string());
70        }
71        CtnStage1Recipe::new(
72            &self.response_column,
73            &self.covariate_formula_rhs,
74            config,
75            self.weight_column.as_deref(),
76            self.offset_column.as_deref(),
77        )
78    }
79}
80
81#[derive(Clone, Debug)]
82pub struct ResolvedFitRequest {
83    pub formula: String,
84    pub fit_config: FitConfig,
85}
86
87pub struct SurvivalInverseLinkInput<'a> {
88    pub link: Option<&'a str>,
89    pub mixture_rho: Option<&'a str>,
90    pub sas_init: Option<&'a str>,
91    pub beta_logistic_init: Option<&'a str>,
92    pub survival_distribution: &'a str,
93}
94
95pub fn parse_fit_request_json(request_json: &str) -> Result<ResolvedFitRequest, String> {
96    resolve_fit_request_document(FitRequestDocument::from_json(request_json)?)
97}
98
99/// Validate a fit request through the production resolver and serialize it in
100/// the one canonical byte representation used by every frontend.
101pub fn canonicalize_fit_request_json(request_json: &str) -> Result<String, String> {
102    let request = FitRequestDocument::from_json(request_json)?;
103    resolve_fit_request_document(request.clone())?;
104    request.to_canonical_json()
105}
106
107pub fn resolve_fit_request_document(
108    request: FitRequestDocument,
109) -> Result<ResolvedFitRequest, String> {
110    let formula = request.formula;
111    let fit_config = resolve_fit_request_config(request.config)?;
112    Ok(ResolvedFitRequest {
113        formula,
114        fit_config,
115    })
116}
117
118/// Parse the canonical config object used by non-formula helper APIs.
119/// Formula fit entry points must use [`FitRequestDocument`] instead.
120pub fn parse_fit_config_json(config_json: Option<&str>) -> Result<FitConfig, String> {
121    let config = match config_json {
122        Some(raw) if !raw.trim().is_empty() => {
123            serde_json::from_str::<FitRequestConfigDocument>(raw)
124                .map_err(|error| format!("invalid fit config object: {error}"))?
125        }
126        _ => FitRequestConfigDocument::default(),
127    };
128    resolve_fit_request_config(config)
129}
130
131pub fn resolve_fit_request_config(
132    json_config: FitRequestConfigDocument,
133) -> Result<FitConfig, String> {
134    let mut fit_config = FitConfig::default();
135    fit_config.group_metadata = json_config.group_metadata.and_then(nonempty_group_metadata);
136    if let Some(training_table_kind) = json_config.training_table_kind {
137        fit_config.training_table_kind = training_table_kind;
138    }
139    fit_config.penalty_block_gamma_priors =
140        parse_precision_hyperpriors(json_config.precision_hyperpriors)?;
141    let latent_coordinates = json_config
142        .latent_coordinates
143        .as_ref()
144        .map(|coordinates| coordinates.to_json_value())
145        .transpose()?;
146    let analytic_penalties = json_config
147        .analytic_penalties
148        .as_ref()
149        .map(|penalties| penalties.to_json_value())
150        .transpose()?;
151    build_analytic_penalty_registry_from_descriptors(
152        latent_coordinates.as_ref(),
153        analytic_penalties.as_ref(),
154    )?;
155    fit_config.latents = latent_coordinates;
156    fit_config.analytic_penalties = analytic_penalties;
157    fit_config.smooth_overrides = json_config
158        .smooth_descriptors
159        .as_ref()
160        .map(|descriptors| descriptors.to_json_value())
161        .transpose()?;
162    fit_config.topology_auto_selector = json_config
163        .topology_auto_selector
164        .as_ref()
165        .map(gam_solve::topology_selector::TopologyAutoSelector::from_json)
166        .transpose()?;
167    fit_config.family = json_config.family;
168    fit_config.negative_binomial_theta = json_config.negative_binomial_theta;
169    fit_config.expectile_tau = json_config.expectile_tau;
170    fit_config.offset_column = json_config.offset;
171    fit_config.weight_column = json_config.weights;
172    if let Some(ridge_lambda) = json_config.ridge_lambda {
173        fit_config.ridge_lambda = ridge_lambda;
174    }
175    if let Some(flag) = json_config.transformation_normal {
176        fit_config.transformation_normal = flag;
177    }
178    if let Some(mode) = json_config.survival_likelihood {
179        fit_config.survival_likelihood = mode;
180    }
181    if let Some(distribution) = json_config.survival_distribution {
182        fit_config.survival_distribution = distribution;
183    }
184    if let Some(target) = json_config.baseline_target {
185        fit_config.baseline_target = target;
186    }
187    if let Some(value) = json_config.baseline_scale {
188        fit_config.baseline_scale = Some(value);
189    }
190    if let Some(value) = json_config.baseline_shape {
191        fit_config.baseline_shape = Some(value);
192    }
193    if let Some(value) = json_config.baseline_rate {
194        fit_config.baseline_rate = Some(value);
195    }
196    if let Some(value) = json_config.baseline_makeham {
197        fit_config.baseline_makeham = Some(value);
198    }
199    if let Some(value) = json_config.time_basis {
200        fit_config.time_basis = value;
201    }
202    if let Some(value) = json_config.time_degree {
203        fit_config.time_degree = value;
204    }
205    if let Some(value) = json_config.time_num_internal_knots {
206        fit_config.time_num_internal_knots = value;
207    }
208    if let Some(value) = json_config.time_smooth_lambda {
209        fit_config.time_smooth_lambda = value;
210    }
211    fit_config.threshold_time_k = json_config.threshold_time_k;
212    if let Some(value) = json_config.threshold_time_degree {
213        fit_config.threshold_time_degree = value;
214    }
215    fit_config.sigma_time_k = json_config.sigma_time_k;
216    if let Some(value) = json_config.sigma_time_degree {
217        fit_config.sigma_time_degree = value;
218    }
219    fit_config.z_column = json_config.z_column;
220    if let Some(formula) = json_config.logslope_formula {
221        fit_config.logslope_formula = Some(formula);
222    }
223    if let Some(stage1) = json_config.ctn_stage1 {
224        fit_config.ctn_stage1 = Some(stage1.into_recipe()?);
225    }
226    fit_config.link = json_config.link;
227    if let Some(flag) = json_config.flexible_link {
228        fit_config.flexible_link = flag;
229    }
230    if let Some(flag) = json_config.scale_dimensions {
231        fit_config.scale_dimensions = flag;
232    }
233    if let Some(value) = json_config.pilot_subsample_threshold {
234        fit_config.spatial_optimization.pilot_subsample_threshold = value;
235    }
236    if let Some(flag) = json_config.adaptive_regularization {
237        fit_config.adaptive_regularization = Some(flag);
238    }
239    if let Some(formula) = json_config.noise_formula {
240        fit_config.noise_formula = Some(formula);
241    }
242    if let Some(column) = json_config.noise_offset {
243        fit_config.noise_offset_column = Some(column);
244    }
245    if let Some(flag) = json_config.firth {
246        fit_config.firth = flag;
247    }
248    if let Some(value) = json_config.outer_max_iter {
249        fit_config.outer_max_iter = Some(value);
250    }
251    if let Some(raw_gpu) = json_config.gpu {
252        fit_config.gpu_policy = parse_gpu_policy(&raw_gpu)?;
253    }
254    fit_config.frailty = parse_json_frailty_spec(
255        json_config.frailty_kind,
256        json_config.frailty_sd,
257        json_config.hazard_loading,
258    )?;
259    fit_config = fit_config.resolve()?;
260    Ok(fit_config)
261}
262
263pub fn resolve_cli_frailty_spec(
264    frailty_kind: Option<CliFrailtyKind>,
265    frailty_sd: Option<f64>,
266    hazard_loading: Option<CliHazardLoading>,
267    context: &str,
268) -> Result<FrailtySpec, String> {
269    let validate_sigma = || -> Result<Option<f64>, String> {
270        match frailty_sd {
271            None => Ok(None),
272            Some(sigma) => {
273                if !sigma.is_finite() || sigma < 0.0 {
274                    return Err(format!(
275                        "{context} requires a finite --frailty-sd >= 0, got {sigma}"
276                    ));
277                }
278                Ok(Some(sigma))
279            }
280        }
281    };
282
283    match frailty_kind {
284        None => {
285            if frailty_sd.is_some() || hazard_loading.is_some() {
286                return Err(format!(
287                    "{context} requires --frailty-kind when --frailty-sd or --hazard-loading is provided"
288                ));
289            }
290            Ok(FrailtySpec::None)
291        }
292        Some(CliFrailtyKind::GaussianShift) => {
293            if hazard_loading.is_some() {
294                return Err(format!(
295                    "{context} does not accept --hazard-loading with --frailty-kind gaussian-shift"
296                ));
297            }
298            Ok(FrailtySpec::GaussianShift {
299                sigma_fixed: validate_sigma()?,
300            })
301        }
302        Some(CliFrailtyKind::HazardMultiplier) => Ok(FrailtySpec::HazardMultiplier {
303            sigma_fixed: validate_sigma()?,
304            loading: hazard_loading.map(cli_hazard_loading).ok_or_else(|| {
305                format!("{context} requires --hazard-loading with --frailty-kind hazard-multiplier")
306            })?,
307        }),
308    }
309}
310
311pub fn parse_survival_likelihood_cli(raw: &str) -> Result<String, String> {
312    let normalized = raw.trim().to_ascii_lowercase();
313    parse_survival_likelihood_mode(&normalized)?;
314    Ok(normalized)
315}
316
317pub fn parse_baseline_target_cli(raw: &str) -> Result<String, String> {
318    let normalized = raw.trim().to_ascii_lowercase();
319    match normalized.as_str() {
320        "linear" | "weibull" | "gompertz" | "gompertz-makeham" => Ok(normalized),
321        other => Err(format!(
322            "unsupported --baseline-target '{other}'; use linear|weibull|gompertz|gompertz-makeham"
323        )),
324    }
325}
326
327pub fn parse_comma_f64(v: &str, label: &str) -> Result<Vec<f64>, String> {
328    let mut out = Vec::new();
329    for part in v.split(',') {
330        let t = part.trim();
331        if t.is_empty() {
332            continue;
333        }
334        let parsed = t
335            .parse::<f64>()
336            .map_err(|err| format!("{label} contains non-numeric value '{t}': {err}"))?;
337        if !parsed.is_finite() {
338            return Err(format!("{label} contains non-finite value '{t}'"));
339        }
340        out.push(parsed);
341    }
342    Ok(out)
343}
344
345pub fn effective_link_to_standard(
346    link: LinkFunction,
347    context: &str,
348) -> Result<StandardLink, String> {
349    StandardLink::try_from(link).map_err(|_| {
350        format!(
351            "{context}: state-bearing link `{}` must be routed through `InverseLink::Sas` / `InverseLink::BetaLogistic`, not `Standard(_)`",
352            link.name()
353        )
354    })
355}
356
357pub fn parse_survival_inverse_link(
358    input: SurvivalInverseLinkInput<'_>,
359) -> Result<InverseLink, String> {
360    if let Some(raw) = input.link {
361        let name = raw.trim().to_ascii_lowercase();
362        if name == "loglog" || name == "cauchit" {
363            // `loglog` and `cauchit` have no scalar `LinkFunction`/`StandardLink`
364            // representative, but the blended-link kernels implement their inverse
365            // link and derivative jets exactly (`LinkComponent::LogLog` /
366            // `LinkComponent::Cauchit`). Represent a survival `--link loglog` /
367            // `--link cauchit` as a single-component mixture: it carries weight 1.0
368            // with no free mixing logits, so it evaluates as exactly that link and
369            // flows end-to-end through the fully-wired `InverseLink::Mixture` survival
370            // path (prepare/construct/row-kernel/predict).
371            if input.sas_init.is_some() {
372                return Err("--sas-init requires --link sas".to_string());
373            }
374            if input.beta_logistic_init.is_some() {
375                return Err("--beta-logistic-init requires --link beta-logistic".to_string());
376            }
377            if input.mixture_rho.is_some() {
378                return Err(
379                    "--mixture-rho requires survival --link blended(...)/mixture(...)".to_string(),
380                );
381            }
382            let component = if name == "loglog" {
383                LinkComponent::LogLog
384            } else {
385                LinkComponent::Cauchit
386            };
387            return state_fromspec(&MixtureLinkSpec {
388                components: vec![component],
389                initial_rho: Array1::zeros(0),
390            })
391            .map(InverseLink::Mixture)
392            .map_err(|e| format!("invalid survival {name} link state: {e}"));
393        }
394    }
395    let choice = parse_link_choice(input.link, false).map_err(|err| {
396        let err = err.to_string();
397        if let Some(raw) = input.link {
398            let name = raw.trim().to_ascii_lowercase();
399            if err.starts_with("unsupported --link ") || err.starts_with("unsupported link type ") {
400                return format!(
401                    "unsupported survival --link '{name}'; {}",
402                    survival_link_usage()
403                );
404            }
405        }
406        err
407    })?;
408    if let Some(choice) = choice {
409        if let Some(components) = choice.mixture_components {
410            if input.sas_init.is_some() || input.beta_logistic_init.is_some() {
411                return Err(
412                    "survival blended(...) link does not accept --sas-init/--beta-logistic-init"
413                        .to_string(),
414                );
415            }
416            let expected = components.len().saturating_sub(1);
417            let initial_rho = if let Some(raw) = input.mixture_rho {
418                let vals = parse_comma_f64(raw, "--mixture-rho")?;
419                if vals.len() != expected {
420                    return Err(format!(
421                        "--mixture-rho expects {expected} values for blended({})",
422                        components
423                            .iter()
424                            .map(|component| component.name())
425                            .collect::<Vec<_>>()
426                            .join(",")
427                    ));
428                }
429                Array1::from_vec(vals)
430            } else {
431                Array1::zeros(expected)
432            };
433            return state_fromspec(&MixtureLinkSpec {
434                components,
435                initial_rho,
436            })
437            .map(InverseLink::Mixture)
438            .map_err(|e| format!("invalid survival blended link state: {e}"));
439        }
440
441        if input.mixture_rho.is_some() {
442            return Err(
443                "--mixture-rho requires survival --link blended(...)/mixture(...)".to_string(),
444            );
445        }
446        match choice.link {
447            LinkFunction::Sas => {
448                if input.beta_logistic_init.is_some() {
449                    return Err("--beta-logistic-init requires --link beta-logistic".to_string());
450                }
451                let (epsilon, log_delta) = if let Some(raw) = input.sas_init {
452                    let vals = parse_comma_f64(raw, "--sas-init")?;
453                    if vals.len() != 2 {
454                        return Err(format!(
455                            "--sas-init expects two values: epsilon,log_delta (got {})",
456                            vals.len()
457                        ));
458                    }
459                    (vals[0], vals[1])
460                } else {
461                    (0.0, 0.0)
462                };
463                state_from_sasspec(SasLinkSpec {
464                    initial_epsilon: epsilon,
465                    initial_log_delta: log_delta,
466                })
467                .map(InverseLink::Sas)
468                .map_err(|e| format!("invalid survival SAS link state: {e}"))
469            }
470            LinkFunction::BetaLogistic => {
471                if input.sas_init.is_some() {
472                    return Err("--sas-init requires --link sas".to_string());
473                }
474                let (epsilon, delta) = if let Some(raw) = input.beta_logistic_init {
475                    let vals = parse_comma_f64(raw, "--beta-logistic-init")?;
476                    if vals.len() != 2 {
477                        return Err(format!(
478                            "--beta-logistic-init expects two values: epsilon,delta (got {})",
479                            vals.len()
480                        ));
481                    }
482                    (vals[0], vals[1])
483                } else {
484                    (0.0, 0.0)
485                };
486                state_from_beta_logisticspec(SasLinkSpec {
487                    initial_epsilon: epsilon,
488                    initial_log_delta: delta,
489                })
490                .map(InverseLink::BetaLogistic)
491                .map_err(|e| format!("invalid survival Beta-Logistic link state: {e}"))
492            }
493            LinkFunction::Log => Err(format!(
494                "unsupported survival --link 'log'; {}",
495                survival_link_usage()
496            )),
497            other => {
498                if input.sas_init.is_some() {
499                    return Err("--sas-init requires --link sas".to_string());
500                }
501                if input.beta_logistic_init.is_some() {
502                    return Err("--beta-logistic-init requires --link beta-logistic".to_string());
503                }
504                Ok(InverseLink::Standard(effective_link_to_standard(
505                    other,
506                    "survival inverse link",
507                )?))
508            }
509        }
510    } else {
511        if input.mixture_rho.is_some() {
512            return Err("--mixture-rho requires --link blended(...)/mixture(...)".to_string());
513        }
514        if input.sas_init.is_some() {
515            return Err("--sas-init requires --link sas".to_string());
516        }
517        if input.beta_logistic_init.is_some() {
518            return Err("--beta-logistic-init requires --link beta-logistic".to_string());
519        }
520        let dist = parse_survival_distribution(input.survival_distribution)?;
521        Ok(residual_distribution_inverse_link(dist))
522    }
523}
524
525fn parse_json_frailty_spec(
526    frailty_kind: Option<String>,
527    frailty_sd: Option<f64>,
528    hazard_loading: Option<String>,
529) -> Result<FrailtySpec, String> {
530    if let Some(kind) = frailty_kind {
531        let trimmed = kind.trim().to_ascii_lowercase();
532        let sigma = frailty_sd;
533        let hazard_loading = hazard_loading
534            .as_ref()
535            .map(|raw| raw.trim().to_ascii_lowercase());
536        let frailty = match trimmed.as_str() {
537            "none" | "" => {
538                if sigma.is_some() || hazard_loading.is_some() {
539                    return Err(
540                        "frailty_kind='none' does not accept frailty_sd or hazard_loading"
541                            .to_string(),
542                    );
543                }
544                FrailtySpec::None
545            }
546            "hazard-multiplier" => {
547                let loading = match hazard_loading.as_deref() {
548                    Some("full") | None => HazardLoading::Full,
549                    Some("loaded-vs-unloaded") => HazardLoading::LoadedVsUnloaded,
550                    Some(other) => {
551                        return Err(format!(
552                            "unknown hazard_loading '{other}'; supported: 'full', 'loaded-vs-unloaded'"
553                        ));
554                    }
555                };
556                FrailtySpec::HazardMultiplier {
557                    sigma_fixed: sigma,
558                    loading,
559                }
560            }
561            "gaussian-shift" => {
562                if hazard_loading.is_some() {
563                    return Err(
564                        "hazard_loading is valid only with frailty_kind='hazard-multiplier'"
565                            .to_string(),
566                    );
567                }
568                FrailtySpec::GaussianShift { sigma_fixed: sigma }
569            }
570            other => {
571                return Err(format!(
572                    "unknown frailty_kind '{other}'; supported: 'none', 'hazard-multiplier', 'gaussian-shift'"
573                ));
574            }
575        };
576        Ok(frailty)
577    } else if frailty_sd.is_some() || hazard_loading.is_some() {
578        Err("frailty_kind is required when frailty_sd or hazard_loading is provided".to_string())
579    } else {
580        Ok(FrailtySpec::None)
581    }
582}
583
584fn cli_hazard_loading(loading: CliHazardLoading) -> HazardLoading {
585    match loading {
586        CliHazardLoading::Full => HazardLoading::Full,
587        CliHazardLoading::LoadedVsUnloaded => HazardLoading::LoadedVsUnloaded,
588    }
589}
590
591fn parse_precision_hyperpriors(
592    precision_hyperpriors: Option<std::collections::BTreeMap<String, PrecisionHyperpriorDocument>>,
593) -> Result<Vec<(String, f64, f64)>, String> {
594    let mut out = Vec::with_capacity(precision_hyperpriors.as_ref().map_or(0, |map| map.len()));
595    for (label, prior) in precision_hyperpriors.unwrap_or_default() {
596        if label.trim().is_empty() {
597            return Err("precision_hyperpriors keys must be non-empty".to_string());
598        }
599        if !prior.shape.is_finite() || prior.shape <= 0.0 {
600            return Err(format!(
601                "precision_hyperpriors['{label}'].shape must be finite and > 0"
602            ));
603        }
604        if !prior.rate.is_finite() || prior.rate < 0.0 {
605            return Err(format!(
606                "precision_hyperpriors['{label}'].rate must be finite and >= 0"
607            ));
608        }
609        out.push((label, prior.shape, prior.rate));
610    }
611    Ok(out)
612}
613
614fn nonempty_group_metadata(metadata: GroupMetadata) -> Option<GroupMetadata> {
615    if metadata.is_empty() {
616        None
617    } else {
618        Some(metadata)
619    }
620}
621
622fn parse_gpu_policy(raw_gpu: &str) -> Result<gam_gpu::GpuPolicy, String> {
623    gam_gpu::GpuPolicy::parse(raw_gpu).ok_or_else(|| {
624        format!(
625            "invalid gpu policy '{}'; supported values are auto, off, required",
626            raw_gpu
627        )
628    })
629}
630
631fn survival_link_usage() -> &'static str {
632    "use identity|logit|probit|cloglog|loglog|cauchit|sas|beta-logistic|blended(...)/mixture(...) or flexible(...)"
633}
634
635#[cfg(test)]
636mod tests {
637    use super::*;
638    use gam_models::survival::lognormal_kernel::FrailtySpec;
639    use serde_json::{Value, json};
640
641    struct ParityCase {
642        name: &'static str,
643        cli: FitConfig,
644        json: Value,
645    }
646
647    fn base_cli() -> FitConfig {
648        FitConfig::default()
649    }
650
651    fn resolved_cli(input: FitConfig) -> Result<FitConfig, String> {
652        input.resolve()
653    }
654
655    fn resolved_json(config: Value) -> Result<FitConfig, String> {
656        let config = serde_json::from_value::<FitRequestConfigDocument>(config)
657            .map_err(|error| format!("invalid test fit config: {error}"))?;
658        let request = FitRequestDocument::new("y ~ x", config)?;
659        resolve_fit_request_document(request).map(|resolved| {
660            assert_eq!(resolved.formula, "y ~ x");
661            resolved.fit_config
662        })
663    }
664
665    fn canonical_fit_config(config: FitConfig) -> String {
666        format!("{config:#?}")
667    }
668
669    #[test]
670    fn rich_request_document_resolves_every_frontend_parity_field() {
671        let request = FitRequestDocument::new(
672            "y ~ duchon(z)",
673            FitRequestConfigDocument {
674                ctn_stage1: Some(CtnStage1Document {
675                    response_column: "dose".to_string(),
676                    covariate_formula_rhs: "s(age)".to_string(),
677                    config: Some(CtnStage1ConfigDocument {
678                        response_degree: Some(4),
679                        response_num_internal_knots: Some(9),
680                        response_penalty_order: Some(2),
681                        response_extra_penalty_orders: Some(vec![1, 3]),
682                        double_penalty: Some(false),
683                    }),
684                    weight_column: Some("case_weight".to_string()),
685                    offset_column: Some("stage1_offset".to_string()),
686                }),
687                precision_hyperpriors: Some(std::collections::BTreeMap::from([(
688                    "duchon(z):roughness".to_string(),
689                    PrecisionHyperpriorDocument {
690                        shape: 2.5,
691                        rate: 0.75,
692                    },
693                )])),
694                latent_coordinates: Some(
695                    serde_json::from_value(json!({
696                        "z": {"n": 20, "d": 2, "name": "z", "init": "pca"}
697                    }))
698                    .unwrap(),
699                ),
700                analytic_penalties: Some(AnalyticPenaltiesDocument(vec![json!({
701                    "kind": "orthogonality",
702                    "target": "z",
703                    "weight": 1.25
704                })])),
705                smooth_descriptors: Some(
706                    serde_json::from_value(json!({
707                        "z": {"kind": "duchon", "vars": ["z"], "centers": 8}
708                    }))
709                    .unwrap(),
710                ),
711                ..FitRequestConfigDocument::default()
712            },
713        )
714        .unwrap();
715
716        let canonical = request.to_canonical_json().unwrap();
717        assert_eq!(
718            canonicalize_fit_request_json(&canonical).unwrap(),
719            canonical
720        );
721        let resolved = parse_fit_request_json(&canonical).unwrap();
722        assert_eq!(resolved.formula, "y ~ duchon(z)");
723        let config = resolved.fit_config;
724        let stage1 = config.ctn_stage1.unwrap();
725        assert_eq!(stage1.response_column, "dose");
726        assert_eq!(stage1.covariate_formula_rhs, "s(age)");
727        assert_eq!(stage1.config.response_degree, 4);
728        assert_eq!(stage1.config.response_num_internal_knots, 9);
729        assert_eq!(stage1.config.response_extra_penalty_orders, vec![1, 3]);
730        assert_eq!(
731            config.penalty_block_gamma_priors,
732            vec![("duchon(z):roughness".to_string(), 2.5, 0.75)]
733        );
734        assert_eq!(config.latents.unwrap()["z"]["d"], json!(2));
735        assert_eq!(config.analytic_penalties.unwrap()[0]["target"], "z");
736        assert_eq!(config.smooth_overrides.unwrap()["z"]["kind"], "duchon");
737    }
738
739    #[test]
740    fn rich_request_rejects_invalid_prior_and_order_dependent_penalty_target() {
741        let invalid_prior = FitRequestDocument::new(
742            "y ~ x",
743            FitRequestConfigDocument {
744                precision_hyperpriors: Some(std::collections::BTreeMap::from([(
745                    "x".to_string(),
746                    PrecisionHyperpriorDocument {
747                        shape: 0.0,
748                        rate: 1.0,
749                    },
750                )])),
751                ..FitRequestConfigDocument::default()
752            },
753        )
754        .unwrap();
755        assert!(
756            resolve_fit_request_document(invalid_prior)
757                .unwrap_err()
758                .contains("shape must be finite and > 0")
759        );
760
761        let numeric_target = FitRequestDocument::new(
762            "y ~ s(z)",
763            FitRequestConfigDocument {
764                latent_coordinates: Some(
765                    serde_json::from_value(json!({"z": {"n": 4, "d": 1}})).unwrap(),
766                ),
767                analytic_penalties: Some(AnalyticPenaltiesDocument(vec![json!({
768                    "kind": "orthogonality",
769                    "target": 0
770                })])),
771                ..FitRequestConfigDocument::default()
772            },
773        )
774        .unwrap();
775        assert!(
776            resolve_fit_request_document(numeric_target)
777                .unwrap_err()
778                .contains("target must be a latent-coordinate name")
779        );
780    }
781
782    #[test]
783    fn cli_shaped_and_json_wire_config_resolution_match() {
784        let cases = vec![
785            ParityCase {
786                name: "family and link selection",
787                cli: {
788                    let mut input = base_cli();
789                    input.family = Some("binomial".to_string());
790                    input.link = Some("probit".to_string());
791                    input.flexible_link = true;
792                    input
793                },
794                json: json!({
795                    "family": "binomial",
796                    "link": "probit",
797                    "flexible_link": true
798                }),
799            },
800            ParityCase {
801                name: "offset weights ridge and noise offset columns",
802                cli: {
803                    let mut input = base_cli();
804                    input.offset_column = Some("eta_offset".to_string());
805                    input.weight_column = Some("case_weight".to_string());
806                    input.noise_offset_column = Some("sigma_offset".to_string());
807                    input.ridge_lambda = 0.125;
808                    input
809                },
810                json: json!({
811                    "offset": "eta_offset",
812                    "weights": "case_weight",
813                    "noise_offset": "sigma_offset",
814                    "ridge_lambda": 0.125
815                }),
816            },
817            ParityCase {
818                name: "weibull survival likelihood and baseline scale shape",
819                cli: {
820                    let mut input = base_cli();
821                    input.survival_likelihood = "weibull".to_string();
822                    input.baseline_target = "weibull".to_string();
823                    input.baseline_scale = Some(2.5);
824                    input.baseline_shape = Some(1.75);
825                    input
826                },
827                json: json!({
828                    "survival_likelihood": "weibull",
829                    "baseline_target": "weibull",
830                    "baseline_scale": 2.5,
831                    "baseline_shape": 1.75
832                }),
833            },
834            ParityCase {
835                name: "transformation survival gompertz makeham baseline",
836                cli: {
837                    let mut input = base_cli();
838                    input.survival_likelihood = "transformation".to_string();
839                    input.baseline_target = "gompertz-makeham".to_string();
840                    input.baseline_shape = Some(1.2);
841                    input.baseline_rate = Some(0.04);
842                    input.baseline_makeham = Some(0.01);
843                    input
844                },
845                json: json!({
846                    "survival_likelihood": "transformation",
847                    "baseline_target": "gompertz-makeham",
848                    "baseline_shape": 1.2,
849                    "baseline_rate": 0.04,
850                    "baseline_makeham": 0.01
851                }),
852            },
853            ParityCase {
854                name: "survival likelihood values are canonicalized",
855                cli: {
856                    let mut input = base_cli();
857                    input.survival_likelihood = "TRANSFORMATION".to_string();
858                    input
859                },
860                json: json!({
861                    "survival_likelihood": "Transformation"
862                }),
863            },
864            ParityCase {
865                name: "noise formula logslope z column and scale dimensions",
866                cli: {
867                    let mut input = base_cli();
868                    input.noise_formula = Some("~ s(age) + treatment".to_string());
869                    input.logslope_formula = Some("~ s(dose)".to_string());
870                    input.z_column = Some("dose".to_string());
871                    input.scale_dimensions = true;
872                    input
873                },
874                json: json!({
875                    "noise_formula": "~ s(age) + treatment",
876                    "logslope_formula": "~ s(dose)",
877                    "z_column": "dose",
878                    "scale_dimensions": true
879                }),
880            },
881            ParityCase {
882                name: "firth transformation normal outer iterations and adaptive regularization",
883                cli: {
884                    let mut input = base_cli();
885                    input.firth = true;
886                    input.transformation_normal = true;
887                    input.outer_max_iter = Some(7);
888                    input.adaptive_regularization = Some(true);
889                    input
890                },
891                json: json!({
892                    "firth": true,
893                    "transformation_normal": true,
894                    "outer_max_iter": 7,
895                    "adaptive_regularization": true
896                }),
897            },
898            ParityCase {
899                name: "gpu policy toggle",
900                cli: {
901                    let mut input = base_cli();
902                    input.gpu_policy = gam_gpu::GpuPolicy::Off;
903                    input
904                },
905                json: json!({
906                    "gpu": "off"
907                }),
908            },
909            ParityCase {
910                name: "hazard multiplier frailty fields",
911                cli: {
912                    let mut input = base_cli();
913                    input.frailty = FrailtySpec::HazardMultiplier {
914                        sigma_fixed: Some(0.35),
915                        loading: HazardLoading::LoadedVsUnloaded,
916                    };
917                    input
918                },
919                json: json!({
920                    "frailty_kind": "hazard-multiplier",
921                    "frailty_sd": 0.35,
922                    "hazard_loading": "loaded-vs-unloaded"
923                }),
924            },
925            ParityCase {
926                name: "gaussian shift frailty fields",
927                cli: {
928                    let mut input = base_cli();
929                    input.frailty = FrailtySpec::GaussianShift {
930                        sigma_fixed: Some(0.2),
931                    };
932                    input
933                },
934                json: json!({
935                    "frailty_kind": "gaussian-shift",
936                    "frailty_sd": 0.2
937                }),
938            },
939        ];
940
941        for case in cases {
942            let cli = resolved_cli(case.cli)
943                .unwrap_or_else(|err| panic!("{}: CLI-shaped config failed: {err}", case.name));
944            let json = resolved_json(case.json)
945                .unwrap_or_else(|err| panic!("{}: JSON wire config failed: {err}", case.name));
946            assert_eq!(
947                canonical_fit_config(cli),
948                canonical_fit_config(json),
949                "{}",
950                case.name
951            );
952        }
953    }
954
955    #[test]
956    fn cli_shaped_and_json_wire_config_resolution_rejections_match() {
957        let cases = vec![
958            ParityCase {
959                name: "negative ridge lambda",
960                cli: {
961                    let mut input = base_cli();
962                    input.ridge_lambda = -1.0;
963                    input
964                },
965                json: json!({
966                    "ridge_lambda": -1.0
967                }),
968            },
969            ParityCase {
970                name: "linear baseline rejects shape",
971                cli: {
972                    let mut input = base_cli();
973                    input.baseline_shape = Some(1.1);
974                    input
975                },
976                json: json!({
977                    "baseline_shape": 1.1
978                }),
979            },
980            ParityCase {
981                name: "weibull likelihood rejects gompertz target",
982                cli: {
983                    let mut input = base_cli();
984                    input.survival_likelihood = "weibull".to_string();
985                    input.baseline_target = "gompertz".to_string();
986                    input
987                },
988                json: json!({
989                    "survival_likelihood": "weibull",
990                    "baseline_target": "gompertz"
991                }),
992            },
993        ];
994
995        for case in cases {
996            let cli = resolved_cli(case.cli).expect_err(case.name);
997            let json = resolved_json(case.json).expect_err(case.name);
998            assert_eq!(cli, json, "{}", case.name);
999        }
1000    }
1001
1002    // ── parse_comma_f64 ───────────────────────────────────────────────────
1003
1004    #[test]
1005    fn parse_comma_f64_empty_string_returns_empty_vec() {
1006        assert_eq!(parse_comma_f64("", "x").unwrap(), Vec::<f64>::new());
1007        assert_eq!(parse_comma_f64("   ", "x").unwrap(), Vec::<f64>::new());
1008    }
1009
1010    #[test]
1011    fn parse_comma_f64_single_value() {
1012        assert_eq!(parse_comma_f64("3.14", "x").unwrap(), vec![3.14]);
1013    }
1014
1015    #[test]
1016    fn parse_comma_f64_multiple_values_with_spaces() {
1017        let result = parse_comma_f64("1.0, 2.5, -3.0", "x").unwrap();
1018        assert_eq!(result, vec![1.0, 2.5, -3.0]);
1019    }
1020
1021    #[test]
1022    fn parse_comma_f64_non_numeric_returns_error() {
1023        let err = parse_comma_f64("1.0, bad, 3.0", "--vals").unwrap_err();
1024        assert!(err.contains("--vals"), "error should name the label: {err}");
1025        assert!(
1026            err.contains("bad"),
1027            "error should name the bad token: {err}"
1028        );
1029    }
1030
1031    #[test]
1032    fn parse_comma_f64_infinity_returns_error() {
1033        let err = parse_comma_f64("inf", "--vals").unwrap_err();
1034        assert!(
1035            err.contains("non-finite"),
1036            "error should say non-finite: {err}"
1037        );
1038    }
1039
1040    #[test]
1041    fn parse_comma_f64_nan_returns_error() {
1042        let err = parse_comma_f64("nan", "--vals").unwrap_err();
1043        assert!(
1044            err.contains("non-finite"),
1045            "error should say non-finite: {err}"
1046        );
1047    }
1048
1049    // ── parse_survival_likelihood_cli ─────────────────────────────────────
1050
1051    #[test]
1052    fn parse_survival_likelihood_cli_valid_values() {
1053        assert_eq!(
1054            parse_survival_likelihood_cli("transformation").unwrap(),
1055            "transformation"
1056        );
1057        assert_eq!(parse_survival_likelihood_cli("weibull").unwrap(), "weibull");
1058        // case-insensitive
1059        assert_eq!(parse_survival_likelihood_cli("WEIBULL").unwrap(), "weibull");
1060        assert_eq!(
1061            parse_survival_likelihood_cli("Transformation").unwrap(),
1062            "transformation"
1063        );
1064    }
1065
1066    #[test]
1067    fn parse_survival_likelihood_cli_invalid_returns_error() {
1068        assert!(parse_survival_likelihood_cli("lognormal").is_err());
1069        assert!(parse_survival_likelihood_cli("").is_err());
1070    }
1071
1072    // ── parse_baseline_target_cli ─────────────────────────────────────────
1073
1074    #[test]
1075    fn parse_baseline_target_cli_valid_values() {
1076        for target in &["linear", "weibull", "gompertz", "gompertz-makeham"] {
1077            assert_eq!(
1078                parse_baseline_target_cli(target).unwrap(),
1079                *target,
1080                "should accept '{target}'"
1081            );
1082        }
1083        // trimmed and lowercased
1084        assert_eq!(parse_baseline_target_cli("  Weibull  ").unwrap(), "weibull");
1085    }
1086
1087    #[test]
1088    fn parse_baseline_target_cli_invalid_returns_error() {
1089        let err = parse_baseline_target_cli("cox").unwrap_err();
1090        assert!(
1091            err.contains("cox"),
1092            "error should name the bad value: {err}"
1093        );
1094    }
1095}