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
103pub 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
122pub 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 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 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 #[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 #[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 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 #[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 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}