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