Skip to main content

gam_config/
fit_request_document.rs

1use serde::{Deserialize, Serialize};
2use serde_json::Value as JsonValue;
3use std::collections::BTreeMap;
4
5/// Stable identity of the serialized fit-request document.
6pub const FIT_REQUEST_SCHEMA: &str = "gam.fit-request";
7
8/// Current fit-request schema version.
9pub const FIT_REQUEST_SCHEMA_VERSION: u32 = 1;
10
11/// A complete, frontend-neutral formula fit request.
12///
13/// Training data is intentionally not embedded: Rust callers supply a
14/// dataset, Python supplies an in-memory table/array, and the CLI
15/// supplies a dataset path. Everything that changes the fitted model belongs in
16/// this document.
17#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
18#[serde(deny_unknown_fields)]
19pub struct FitRequestDocument {
20    pub schema: String,
21    pub schema_version: u32,
22    pub formula: String,
23    #[serde(default)]
24    pub config: FitRequestConfigDocument,
25}
26
27impl FitRequestDocument {
28    pub fn new(
29        formula: impl Into<String>,
30        config: FitRequestConfigDocument,
31    ) -> Result<Self, String> {
32        let document = Self {
33            schema: FIT_REQUEST_SCHEMA.to_string(),
34            schema_version: FIT_REQUEST_SCHEMA_VERSION,
35            formula: formula.into(),
36            config,
37        };
38        document.validate()?;
39        Ok(document)
40    }
41
42    pub fn from_json(raw: &str) -> Result<Self, String> {
43        let document = serde_json::from_str::<Self>(raw)
44            .map_err(|error| format!("invalid fit request document: {error}"))?;
45        document.validate()?;
46        Ok(document)
47    }
48
49    pub fn to_canonical_json(&self) -> Result<String, String> {
50        self.validate()?;
51        serde_json::to_string(self)
52            .map_err(|error| format!("failed to serialize fit request document: {error}"))
53    }
54
55    fn validate(&self) -> Result<(), String> {
56        if self.schema != FIT_REQUEST_SCHEMA {
57            return Err(format!(
58                "fit request schema must be '{FIT_REQUEST_SCHEMA}', got {:?}",
59                self.schema
60            ));
61        }
62        if self.schema_version != FIT_REQUEST_SCHEMA_VERSION {
63            return Err(format!(
64                "unsupported fit request schema_version {}; expected {}",
65                self.schema_version, FIT_REQUEST_SCHEMA_VERSION
66            ));
67        }
68        if self.formula.trim().is_empty() {
69            return Err("fit request formula must be non-empty".to_string());
70        }
71        Ok(())
72    }
73}
74
75/// Serializable model configuration shared by Rust, Python, and the CLI.
76///
77/// Optional fields mean "use the core [`gam_models::fit_orchestration::FitConfig`]
78/// default". The document deliberately has one spelling for each concept; the
79/// parser does not carry aliases or legacy wire formats.
80#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
81#[serde(deny_unknown_fields)]
82pub struct FitRequestConfigDocument {
83    #[serde(skip_serializing_if = "Option::is_none")]
84    pub adaptive_regularization: Option<bool>,
85    #[serde(skip_serializing_if = "Option::is_none")]
86    pub baseline_makeham: Option<f64>,
87    #[serde(skip_serializing_if = "Option::is_none")]
88    pub baseline_rate: Option<f64>,
89    #[serde(skip_serializing_if = "Option::is_none")]
90    pub baseline_scale: Option<f64>,
91    #[serde(skip_serializing_if = "Option::is_none")]
92    pub baseline_shape: Option<f64>,
93    #[serde(skip_serializing_if = "Option::is_none")]
94    pub baseline_target: Option<String>,
95    #[serde(skip_serializing_if = "Option::is_none")]
96    pub ctn_stage1: Option<CtnStage1Document>,
97    #[serde(skip_serializing_if = "Option::is_none")]
98    pub expectile_tau: Option<f64>,
99    #[serde(skip_serializing_if = "Option::is_none")]
100    pub family: Option<String>,
101    #[serde(skip_serializing_if = "Option::is_none")]
102    pub firth: Option<bool>,
103    #[serde(skip_serializing_if = "Option::is_none")]
104    pub flexible_link: Option<bool>,
105    #[serde(skip_serializing_if = "Option::is_none")]
106    pub frailty_kind: Option<String>,
107    #[serde(skip_serializing_if = "Option::is_none")]
108    pub frailty_sd: Option<f64>,
109    #[serde(skip_serializing_if = "Option::is_none")]
110    pub gpu: Option<String>,
111    #[serde(skip_serializing_if = "Option::is_none")]
112    pub group_metadata: Option<BTreeMap<String, JsonValue>>,
113    #[serde(skip_serializing_if = "Option::is_none")]
114    pub hazard_loading: Option<String>,
115    #[serde(skip_serializing_if = "Option::is_none")]
116    pub latent_coordinates: Option<LatentCoordinatesDocument>,
117    #[serde(skip_serializing_if = "Option::is_none")]
118    pub link: Option<String>,
119    #[serde(skip_serializing_if = "Option::is_none")]
120    pub logslope_formula: Option<String>,
121    #[serde(skip_serializing_if = "Option::is_none")]
122    pub negative_binomial_theta: Option<f64>,
123    #[serde(skip_serializing_if = "Option::is_none")]
124    pub noise_formula: Option<String>,
125    #[serde(skip_serializing_if = "Option::is_none")]
126    pub noise_offset: Option<String>,
127    #[serde(skip_serializing_if = "Option::is_none")]
128    pub offset: Option<String>,
129    #[serde(skip_serializing_if = "Option::is_none")]
130    pub outer_max_iter: Option<usize>,
131    #[serde(skip_serializing_if = "Option::is_none")]
132    pub analytic_penalties: Option<AnalyticPenaltiesDocument>,
133    #[serde(skip_serializing_if = "Option::is_none")]
134    pub pilot_subsample_threshold: Option<usize>,
135    #[serde(skip_serializing_if = "Option::is_none")]
136    pub precision_hyperpriors: Option<BTreeMap<String, PrecisionHyperpriorDocument>>,
137    #[serde(skip_serializing_if = "Option::is_none")]
138    pub ridge_lambda: Option<f64>,
139    #[serde(skip_serializing_if = "Option::is_none")]
140    pub scale_dimensions: Option<bool>,
141    #[serde(skip_serializing_if = "Option::is_none")]
142    pub sigma_time_degree: Option<usize>,
143    #[serde(skip_serializing_if = "Option::is_none")]
144    pub sigma_time_k: Option<usize>,
145    #[serde(skip_serializing_if = "Option::is_none")]
146    pub smooth_descriptors: Option<SmoothDescriptorsDocument>,
147    #[serde(skip_serializing_if = "Option::is_none")]
148    pub survival_distribution: Option<String>,
149    #[serde(skip_serializing_if = "Option::is_none")]
150    pub survival_likelihood: Option<String>,
151    #[serde(skip_serializing_if = "Option::is_none")]
152    pub threshold_time_degree: Option<usize>,
153    #[serde(skip_serializing_if = "Option::is_none")]
154    pub threshold_time_k: Option<usize>,
155    #[serde(skip_serializing_if = "Option::is_none")]
156    pub time_basis: Option<String>,
157    #[serde(skip_serializing_if = "Option::is_none")]
158    pub time_degree: Option<usize>,
159    #[serde(skip_serializing_if = "Option::is_none")]
160    pub time_num_internal_knots: Option<usize>,
161    #[serde(skip_serializing_if = "Option::is_none")]
162    pub time_smooth_lambda: Option<f64>,
163    /// Container type of the caller's training table (`"pandas"`, `"polars"`,
164    /// `"pyarrow"`, `"numpy"`, ...), passed through opaquely into the saved
165    /// model payload for the predict-time output-container fallback (#394).
166    #[serde(skip_serializing_if = "Option::is_none")]
167    pub training_table_kind: Option<String>,
168    #[serde(skip_serializing_if = "Option::is_none")]
169    pub transformation_normal: Option<bool>,
170    #[serde(skip_serializing_if = "Option::is_none")]
171    pub weights: Option<String>,
172    #[serde(skip_serializing_if = "Option::is_none")]
173    pub z_column: Option<String>,
174}
175
176#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
177#[serde(deny_unknown_fields)]
178pub struct PrecisionHyperpriorDocument {
179    pub shape: f64,
180    pub rate: f64,
181}
182
183#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
184#[serde(transparent)]
185pub struct LatentCoordinatesDocument(pub BTreeMap<String, LatentCoordinateDocument>);
186
187impl LatentCoordinatesDocument {
188    pub fn to_json_value(&self) -> Result<JsonValue, String> {
189        for (symbol, coordinate) in &self.0 {
190            if symbol.trim().is_empty() {
191                return Err("latent_coordinates keys must be non-empty symbols".to_string());
192            }
193            if coordinate.n == 0 || coordinate.d == 0 {
194                return Err(format!(
195                    "latent_coordinates['{symbol}'] requires positive n and d"
196                ));
197            }
198            if coordinate
199                .name
200                .as_deref()
201                .is_some_and(|name| name.trim().is_empty())
202            {
203                return Err(format!(
204                    "latent_coordinates['{symbol}'].name must be non-empty"
205                ));
206            }
207        }
208        serde_json::to_value(self)
209            .map_err(|error| format!("failed to serialize latent coordinates: {error}"))
210    }
211}
212
213#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
214#[serde(deny_unknown_fields)]
215pub struct LatentCoordinateDocument {
216    pub n: usize,
217    pub d: usize,
218    #[serde(default, skip_serializing_if = "Option::is_none")]
219    pub name: Option<String>,
220    #[serde(default, skip_serializing_if = "Option::is_none")]
221    pub init: Option<JsonValue>,
222    #[serde(default, skip_serializing_if = "Option::is_none")]
223    pub manifold: Option<JsonValue>,
224    #[serde(default, skip_serializing_if = "Option::is_none")]
225    pub retraction: Option<JsonValue>,
226    #[serde(default, skip_serializing_if = "Option::is_none")]
227    pub aux_prior: Option<JsonValue>,
228    #[serde(default, skip_serializing_if = "Option::is_none")]
229    pub dim_selection: Option<JsonValue>,
230    #[serde(default, skip_serializing_if = "Option::is_none")]
231    pub aux_outcome: Option<JsonValue>,
232    #[serde(default, skip_serializing_if = "Option::is_none")]
233    pub id_mode: Option<String>,
234}
235
236#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
237#[serde(transparent)]
238pub struct AnalyticPenaltiesDocument(pub Vec<JsonValue>);
239
240impl AnalyticPenaltiesDocument {
241    pub fn to_json_value(&self) -> Result<JsonValue, String> {
242        for (index, descriptor) in self.0.iter().enumerate() {
243            let descriptor = descriptor
244                .as_object()
245                .ok_or_else(|| format!("analytic_penalties[{index}] must be an object"))?;
246            if !descriptor.get("target").is_some_and(JsonValue::is_string) {
247                return Err(format!(
248                    "analytic_penalties[{index}].target must be a latent-coordinate name"
249                ));
250            }
251        }
252        serde_json::to_value(self)
253            .map_err(|error| format!("failed to serialize analytic penalties: {error}"))
254    }
255}
256
257#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
258#[serde(transparent)]
259pub struct SmoothDescriptorsDocument(pub BTreeMap<String, JsonValue>);
260
261impl SmoothDescriptorsDocument {
262    pub fn to_json_value(&self) -> Result<JsonValue, String> {
263        for (symbol, descriptor) in &self.0 {
264            if symbol.trim().is_empty() {
265                return Err("smooth_descriptors keys must be non-empty symbols".to_string());
266            }
267            if !descriptor.is_object() {
268                return Err(format!("smooth_descriptors['{symbol}'] must be an object"));
269            }
270        }
271        serde_json::to_value(self)
272            .map_err(|error| format!("failed to serialize smooth descriptors: {error}"))
273    }
274}
275
276#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
277#[serde(deny_unknown_fields)]
278pub struct CtnStage1Document {
279    pub response_column: String,
280    pub covariate_formula_rhs: String,
281    #[serde(default, skip_serializing_if = "Option::is_none")]
282    pub config: Option<CtnStage1ConfigDocument>,
283    #[serde(default, skip_serializing_if = "Option::is_none")]
284    pub weight_column: Option<String>,
285    #[serde(default, skip_serializing_if = "Option::is_none")]
286    pub offset_column: Option<String>,
287}
288
289#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
290#[serde(deny_unknown_fields)]
291pub struct CtnStage1ConfigDocument {
292    #[serde(default, skip_serializing_if = "Option::is_none")]
293    pub response_degree: Option<usize>,
294    #[serde(default, skip_serializing_if = "Option::is_none")]
295    pub response_num_internal_knots: Option<usize>,
296    #[serde(default, skip_serializing_if = "Option::is_none")]
297    pub response_penalty_order: Option<usize>,
298    #[serde(default, skip_serializing_if = "Option::is_none")]
299    pub response_extra_penalty_orders: Option<Vec<usize>>,
300    #[serde(default, skip_serializing_if = "Option::is_none")]
301    pub double_penalty: Option<bool>,
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307    use serde_json::json;
308
309    #[test]
310    fn canonical_document_round_trips_identically() {
311        let document = FitRequestDocument::new(
312            "y ~ duchon(x)",
313            FitRequestConfigDocument {
314                ctn_stage1: Some(CtnStage1Document {
315                    response_column: "dose".to_string(),
316                    covariate_formula_rhs: "s(age)".to_string(),
317                    config: Some(CtnStage1ConfigDocument {
318                        response_degree: Some(4),
319                        response_penalty_order: Some(2),
320                        ..CtnStage1ConfigDocument::default()
321                    }),
322                    weight_column: Some("case_weight".to_string()),
323                    offset_column: None,
324                }),
325                latent_coordinates: Some(
326                    serde_json::from_value(json!({
327                        "x": {"d": 2, "init": "pca", "n": 12, "name": "x"}
328                    }))
329                    .unwrap(),
330                ),
331                analytic_penalties: Some(AnalyticPenaltiesDocument(vec![json!(
332                    {"kind": "orthogonality", "target": "x", "weight": 1.0}
333                )])),
334                precision_hyperpriors: Some(BTreeMap::from([(
335                    "s(x):roughness".to_string(),
336                    PrecisionHyperpriorDocument {
337                        shape: 2.0,
338                        rate: 0.5,
339                    },
340                )])),
341                smooth_descriptors: Some(
342                    serde_json::from_value(json!({
343                        "x": {"centers": 8, "kind": "duchon", "vars": ["x"]}
344                    }))
345                    .unwrap(),
346                ),
347                ..FitRequestConfigDocument::default()
348            },
349        )
350        .unwrap();
351
352        let encoded = document.to_canonical_json().unwrap();
353        let decoded = FitRequestDocument::from_json(&encoded).unwrap();
354        assert_eq!(decoded, document);
355        assert_eq!(decoded.to_canonical_json().unwrap(), encoded);
356    }
357
358    #[test]
359    fn parser_rejects_another_schema_or_version() {
360        let wrong_schema = r#"{"schema":"other","schema_version":1,"formula":"y ~ x","config":{}}"#;
361        assert!(
362            FitRequestDocument::from_json(wrong_schema)
363                .unwrap_err()
364                .contains("schema must be")
365        );
366
367        let wrong_version =
368            r#"{"schema":"gam.fit-request","schema_version":2,"formula":"y ~ x","config":{}}"#;
369        assert!(
370            FitRequestDocument::from_json(wrong_version)
371                .unwrap_err()
372                .contains("unsupported fit request schema_version")
373        );
374    }
375
376    #[test]
377    fn parser_rejects_removed_topology_selector_descriptor() {
378        let legacy = r#"{
379            "schema": "gam.fit-request",
380            "schema_version": 1,
381            "formula": "y ~ x",
382            "config": {"topology_auto_selector": {"candidates": ["circle"]}}
383        }"#;
384        let error = FitRequestDocument::from_json(legacy)
385            .expect_err("removed no-op selector field must not be silently ignored");
386        assert!(
387            error.contains("unknown field `topology_auto_selector`"),
388            "{error}"
389        );
390    }
391}