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
99pub 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
118pub 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 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 #[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 #[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 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 #[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 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}