Skip to main content

gam_config/
fit_request_document.rs

1use serde::{Deserialize, Serialize};
2use serde_json::Value as JsonValue;
3use std::collections::BTreeMap;
4use std::path::PathBuf;
5
6/// Stable identity of the serialized fit-request document.
7pub const FIT_REQUEST_SCHEMA: &str = "gam.fit-request";
8
9/// Current fit-request schema version.
10pub const FIT_REQUEST_SCHEMA_VERSION: u32 = 1;
11
12/// A complete, frontend-neutral formula fit request.
13///
14/// Training data is intentionally not embedded: Rust callers supply a
15/// dataset, Python supplies an in-memory table/array, and the CLI
16/// supplies a dataset path. Everything that changes the fitted model belongs in
17/// this document.
18#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
19#[serde(deny_unknown_fields)]
20pub struct FitRequestDocument {
21    pub schema: String,
22    pub schema_version: u32,
23    pub formula: String,
24    #[serde(default)]
25    pub config: FitRequestConfigDocument,
26}
27
28impl FitRequestDocument {
29    pub fn new(
30        formula: impl Into<String>,
31        config: FitRequestConfigDocument,
32    ) -> Result<Self, String> {
33        let document = Self {
34            schema: FIT_REQUEST_SCHEMA.to_string(),
35            schema_version: FIT_REQUEST_SCHEMA_VERSION,
36            formula: formula.into(),
37            config,
38        };
39        document.validate()?;
40        Ok(document)
41    }
42
43    pub fn from_json(raw: &str) -> Result<Self, String> {
44        let document = serde_json::from_str::<Self>(raw)
45            .map_err(|error| format!("invalid fit request document: {error}"))?;
46        document.validate()?;
47        Ok(document)
48    }
49
50    fn validate(&self) -> Result<(), String> {
51        if self.schema != FIT_REQUEST_SCHEMA {
52            return Err(format!(
53                "fit request schema must be '{FIT_REQUEST_SCHEMA}', got {:?}",
54                self.schema
55            ));
56        }
57        if self.schema_version != FIT_REQUEST_SCHEMA_VERSION {
58            return Err(format!(
59                "unsupported fit request schema_version {}; expected {}",
60                self.schema_version, FIT_REQUEST_SCHEMA_VERSION
61            ));
62        }
63        if self.formula.trim().is_empty() {
64            return Err("fit request formula must be non-empty".to_string());
65        }
66        Ok(())
67    }
68}
69
70/// Serializable model configuration shared by Rust, Python, and the CLI.
71///
72/// Optional fields mean "use the core [`gam_models::fit_orchestration::FitConfig`]
73/// default". The document deliberately has one spelling for each concept; the
74/// parser does not carry aliases or legacy wire formats.
75#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
76#[serde(deny_unknown_fields)]
77pub struct FitRequestConfigDocument {
78    #[serde(skip_serializing_if = "Option::is_none")]
79    pub adaptive_regularization: Option<bool>,
80    #[serde(skip_serializing_if = "Option::is_none")]
81    pub baseline_makeham: Option<f64>,
82    #[serde(skip_serializing_if = "Option::is_none")]
83    pub baseline_rate: Option<f64>,
84    #[serde(skip_serializing_if = "Option::is_none")]
85    pub baseline_scale: Option<f64>,
86    #[serde(skip_serializing_if = "Option::is_none")]
87    pub baseline_shape: Option<f64>,
88    #[serde(skip_serializing_if = "Option::is_none")]
89    pub baseline_target: Option<String>,
90    #[serde(skip_serializing_if = "Option::is_none")]
91    pub ctn_stage1: Option<CtnStage1Document>,
92    #[serde(skip_serializing_if = "Option::is_none")]
93    pub expectile_tau: Option<f64>,
94    #[serde(skip_serializing_if = "Option::is_none")]
95    pub family: Option<String>,
96    #[serde(skip_serializing_if = "Option::is_none")]
97    pub firth: Option<bool>,
98    #[serde(skip_serializing_if = "Option::is_none")]
99    pub flexible_link: Option<bool>,
100    #[serde(skip_serializing_if = "Option::is_none")]
101    pub frailty_kind: Option<String>,
102    #[serde(skip_serializing_if = "Option::is_none")]
103    pub frailty_sd: Option<f64>,
104    #[serde(skip_serializing_if = "Option::is_none")]
105    pub gpu: Option<String>,
106    #[serde(skip_serializing_if = "Option::is_none")]
107    pub group_metadata: Option<BTreeMap<String, JsonValue>>,
108    #[serde(skip_serializing_if = "Option::is_none")]
109    pub hazard_loading: Option<String>,
110    #[serde(skip_serializing_if = "Option::is_none")]
111    pub latent_coordinates: Option<LatentCoordinatesDocument>,
112    #[serde(skip_serializing_if = "Option::is_none")]
113    pub link: Option<String>,
114    #[serde(skip_serializing_if = "Option::is_none")]
115    pub slope_formula: Option<String>,
116    #[serde(skip_serializing_if = "Option::is_none")]
117    pub negative_binomial_theta: Option<f64>,
118    #[serde(skip_serializing_if = "Option::is_none")]
119    pub noise_formula: Option<String>,
120    #[serde(skip_serializing_if = "Option::is_none")]
121    pub noise_offset: Option<String>,
122    #[serde(skip_serializing_if = "Option::is_none")]
123    pub offset: Option<String>,
124    #[serde(skip_serializing_if = "Option::is_none")]
125    pub outer_max_iter: Option<usize>,
126    #[serde(skip_serializing_if = "Option::is_none")]
127    pub analytic_penalties: Option<AnalyticPenaltiesDocument>,
128    #[serde(skip_serializing_if = "Option::is_none")]
129    pub pilot_subsample_threshold: Option<usize>,
130    #[serde(skip_serializing_if = "Option::is_none")]
131    pub precision_hyperpriors: Option<BTreeMap<String, PrecisionHyperpriorDocument>>,
132    #[serde(skip_serializing_if = "Option::is_none")]
133    /// Whether to precompute the distribution-free conformal substrates (#942
134    /// jackknife+, #1098 exact full-conformal) at fit time and persist them on
135    /// the saved model. Omit to keep the default of precomputing whenever the
136    /// fit is eligible; `false` skips both.
137    ///
138    /// Measured on `y ~ s(x1,k=6) + s(x2,k=6)` (#2633): the two substrates are
139    /// 94% of a saved Gaussian model at n=20,000 (10.2 MB of 10.85 MB) and grow
140    /// linearly with the training rows. Rebuilding both costs ~5.6 ms, 0.3% of
141    /// the fit, and stays under half a second out to p=253. So turning this off
142    /// yields a ~16x smaller model (10.85 MB -> ~0.65 MB at n=20,000).
143    ///
144    /// It is opt-OUT because rebuilding needs the training design AND response
145    /// back, which a saved model deliberately does not carry: a model shipped to
146    /// a host that never sees the training data must keep them or it cannot
147    /// produce a conformal interval at all. Turn it off when the caller retains
148    /// its training data, fits in batch, or never asks for conformal intervals.
149    pub precompute_conformal: Option<bool>,
150    /// Explicit root for cross-process warm starts. Omit to disable on-disk
151    /// persistence. The path is used exactly as supplied; no temp/cache
152    /// discovery or environment fallback is performed.
153    #[serde(skip_serializing_if = "Option::is_none")]
154    pub persistent_warm_start_root: Option<PathBuf>,
155    #[serde(skip_serializing_if = "Option::is_none")]
156    pub ridge_lambda: Option<f64>,
157    #[serde(skip_serializing_if = "Option::is_none")]
158    pub scale_dimensions: Option<bool>,
159    #[serde(skip_serializing_if = "Option::is_none")]
160    pub slope_time_degree: Option<usize>,
161    /// Number of B-spline basis functions on the `log t` margin of the
162    /// survival marginal-slope slope block (gam#2765, gam#2767). Omitted =
163    /// a slope that does not move along follow-up.
164    #[serde(skip_serializing_if = "Option::is_none")]
165    pub slope_time_k: Option<usize>,
166    #[serde(skip_serializing_if = "Option::is_none")]
167    pub sigma_time_degree: Option<usize>,
168    #[serde(skip_serializing_if = "Option::is_none")]
169    pub sigma_time_k: Option<usize>,
170    #[serde(skip_serializing_if = "Option::is_none")]
171    pub smooth_descriptors: Option<SmoothDescriptorsDocument>,
172    #[serde(skip_serializing_if = "Option::is_none")]
173    pub survival_distribution: Option<String>,
174    #[serde(skip_serializing_if = "Option::is_none")]
175    pub survival_likelihood: Option<String>,
176    /// Explicit centering anchor for the survival baseline time basis, in the
177    /// data's own time units. Omit to let the fit pick it from the likelihood
178    /// mode and the truncation shape of the data — the robust interior median
179    /// exit for marginal-slope and for any genuinely left-truncated dataset
180    /// (#751/#1790), the earliest entry age otherwise.
181    ///
182    /// The CLI's `--survival-time-anchor` declares a conflict with `--request` on
183    /// the premise that this document carries the complete scientific model
184    /// configuration; until #2631 the document had no field for it, so the
185    /// premise was false.
186    #[serde(skip_serializing_if = "Option::is_none")]
187    pub survival_time_anchor: Option<f64>,
188    #[serde(skip_serializing_if = "Option::is_none")]
189    pub threshold_time_degree: Option<usize>,
190    #[serde(skip_serializing_if = "Option::is_none")]
191    pub threshold_time_k: Option<usize>,
192    #[serde(skip_serializing_if = "Option::is_none")]
193    pub time_basis: Option<String>,
194    #[serde(skip_serializing_if = "Option::is_none")]
195    pub time_degree: Option<usize>,
196    #[serde(skip_serializing_if = "Option::is_none")]
197    pub time_num_internal_knots: Option<usize>,
198    #[serde(skip_serializing_if = "Option::is_none")]
199    pub time_smooth_lambda: Option<f64>,
200    /// Container type of the caller's training table (`"pandas"`, `"polars"`,
201    /// `"pyarrow"`, `"numpy"`, ...), passed through opaquely into the saved
202    /// model payload for the predict-time output-container fallback (#394).
203    #[serde(skip_serializing_if = "Option::is_none")]
204    pub training_table_kind: Option<String>,
205    #[serde(skip_serializing_if = "Option::is_none")]
206    pub transformation_normal: Option<bool>,
207    #[serde(skip_serializing_if = "Option::is_none")]
208    pub weights: Option<String>,
209    #[serde(skip_serializing_if = "Option::is_none")]
210    pub z_column: Option<String>,
211}
212
213#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
214#[serde(deny_unknown_fields)]
215pub struct PrecisionHyperpriorDocument {
216    pub shape: f64,
217    pub rate: f64,
218}
219
220#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
221#[serde(transparent)]
222pub struct LatentCoordinatesDocument(pub BTreeMap<String, LatentCoordinateDocument>);
223
224impl LatentCoordinatesDocument {
225    pub fn to_json_value(&self) -> Result<JsonValue, String> {
226        for (symbol, coordinate) in &self.0 {
227            if symbol.trim().is_empty() {
228                return Err("latent_coordinates keys must be non-empty symbols".to_string());
229            }
230            if coordinate.n == 0 || coordinate.d == 0 {
231                return Err(format!(
232                    "latent_coordinates['{symbol}'] requires positive n and d"
233                ));
234            }
235            if coordinate
236                .name
237                .as_deref()
238                .is_some_and(|name| name.trim().is_empty())
239            {
240                return Err(format!(
241                    "latent_coordinates['{symbol}'].name must be non-empty"
242                ));
243            }
244        }
245        serde_json::to_value(self)
246            .map_err(|error| format!("failed to serialize latent coordinates: {error}"))
247    }
248}
249
250#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
251#[serde(deny_unknown_fields)]
252pub struct LatentCoordinateDocument {
253    pub n: usize,
254    pub d: usize,
255    #[serde(default, skip_serializing_if = "Option::is_none")]
256    pub name: Option<String>,
257    #[serde(default, skip_serializing_if = "Option::is_none")]
258    pub init: Option<JsonValue>,
259    #[serde(default, skip_serializing_if = "Option::is_none")]
260    pub manifold: Option<JsonValue>,
261    #[serde(default, skip_serializing_if = "Option::is_none")]
262    pub retraction: Option<JsonValue>,
263    #[serde(default, skip_serializing_if = "Option::is_none")]
264    pub aux_prior: Option<JsonValue>,
265    #[serde(default, skip_serializing_if = "Option::is_none")]
266    pub dim_selection: Option<JsonValue>,
267    #[serde(default, skip_serializing_if = "Option::is_none")]
268    pub aux_outcome: Option<JsonValue>,
269    #[serde(default, skip_serializing_if = "Option::is_none")]
270    pub id_mode: Option<String>,
271}
272
273#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
274#[serde(transparent)]
275pub struct AnalyticPenaltiesDocument(pub Vec<JsonValue>);
276
277impl AnalyticPenaltiesDocument {
278    pub fn to_json_value(&self) -> Result<JsonValue, String> {
279        for (index, descriptor) in self.0.iter().enumerate() {
280            let descriptor = descriptor
281                .as_object()
282                .ok_or_else(|| format!("analytic_penalties[{index}] must be an object"))?;
283            if !descriptor.get("target").is_some_and(JsonValue::is_string) {
284                return Err(format!(
285                    "analytic_penalties[{index}].target must be a latent-coordinate name"
286                ));
287            }
288        }
289        serde_json::to_value(self)
290            .map_err(|error| format!("failed to serialize analytic penalties: {error}"))
291    }
292}
293
294#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
295#[serde(transparent)]
296pub struct SmoothDescriptorsDocument(pub BTreeMap<String, JsonValue>);
297
298impl SmoothDescriptorsDocument {
299    pub fn to_json_value(&self) -> Result<JsonValue, String> {
300        for (symbol, descriptor) in &self.0 {
301            if symbol.trim().is_empty() {
302                return Err("smooth_descriptors keys must be non-empty symbols".to_string());
303            }
304            if !descriptor.is_object() {
305                return Err(format!("smooth_descriptors['{symbol}'] must be an object"));
306            }
307        }
308        serde_json::to_value(self)
309            .map_err(|error| format!("failed to serialize smooth descriptors: {error}"))
310    }
311}
312
313#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
314#[serde(deny_unknown_fields)]
315pub struct CtnStage1Document {
316    pub response_column: String,
317    pub covariate_formula_rhs: String,
318    #[serde(default, skip_serializing_if = "Option::is_none")]
319    pub config: Option<CtnStage1ConfigDocument>,
320    #[serde(default, skip_serializing_if = "Option::is_none")]
321    pub weight_column: Option<String>,
322    #[serde(default, skip_serializing_if = "Option::is_none")]
323    pub offset_column: Option<String>,
324}
325
326#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
327#[serde(deny_unknown_fields)]
328pub struct CtnStage1ConfigDocument {
329    #[serde(default, skip_serializing_if = "Option::is_none")]
330    pub response_degree: Option<usize>,
331    #[serde(default, skip_serializing_if = "Option::is_none")]
332    pub response_num_internal_knots: Option<usize>,
333    #[serde(default, skip_serializing_if = "Option::is_none")]
334    pub response_penalty_order: Option<usize>,
335    #[serde(default, skip_serializing_if = "Option::is_none")]
336    pub response_extra_penalty_orders: Option<Vec<usize>>,
337    #[serde(default, skip_serializing_if = "Option::is_none")]
338    pub double_penalty: Option<bool>,
339}
340
341#[cfg(test)]
342mod tests {
343    use super::*;
344
345    #[test]
346    fn parser_rejects_another_schema_or_version() {
347        let wrong_schema = r#"{"schema":"other","schema_version":1,"formula":"y ~ x","config":{}}"#;
348        assert!(
349            FitRequestDocument::from_json(wrong_schema)
350                .unwrap_err()
351                .contains("schema must be")
352        );
353
354        let wrong_version =
355            r#"{"schema":"gam.fit-request","schema_version":2,"formula":"y ~ x","config":{}}"#;
356        assert!(
357            FitRequestDocument::from_json(wrong_version)
358                .unwrap_err()
359                .contains("unsupported fit request schema_version")
360        );
361    }
362
363    #[test]
364    fn parser_rejects_removed_topology_selector_descriptor() {
365        let legacy = r#"{
366            "schema": "gam.fit-request",
367            "schema_version": 1,
368            "formula": "y ~ x",
369            "config": {"topology_auto_selector": {"candidates": ["circle"]}}
370        }"#;
371        let error = FitRequestDocument::from_json(legacy)
372            .expect_err("removed no-op selector field must not be silently ignored");
373        assert!(
374            error.contains("unknown field `topology_auto_selector`"),
375            "{error}"
376        );
377    }
378}