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    #[serde(skip_serializing_if = "Option::is_none")]
164    pub topology_auto_selector: Option<JsonValue>,
165    /// Container type of the caller's training table (`"pandas"`, `"polars"`,
166    /// `"pyarrow"`, `"numpy"`, ...), passed through opaquely into the saved
167    /// model payload for the predict-time output-container fallback (#394).
168    #[serde(skip_serializing_if = "Option::is_none")]
169    pub training_table_kind: Option<String>,
170    #[serde(skip_serializing_if = "Option::is_none")]
171    pub transformation_normal: Option<bool>,
172    #[serde(skip_serializing_if = "Option::is_none")]
173    pub weights: Option<String>,
174    #[serde(skip_serializing_if = "Option::is_none")]
175    pub z_column: Option<String>,
176}
177
178#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
179#[serde(deny_unknown_fields)]
180pub struct PrecisionHyperpriorDocument {
181    pub shape: f64,
182    pub rate: f64,
183}
184
185#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
186#[serde(transparent)]
187pub struct LatentCoordinatesDocument(pub BTreeMap<String, LatentCoordinateDocument>);
188
189impl LatentCoordinatesDocument {
190    pub fn to_json_value(&self) -> Result<JsonValue, String> {
191        for (symbol, coordinate) in &self.0 {
192            if symbol.trim().is_empty() {
193                return Err("latent_coordinates keys must be non-empty symbols".to_string());
194            }
195            if coordinate.n == 0 || coordinate.d == 0 {
196                return Err(format!(
197                    "latent_coordinates['{symbol}'] requires positive n and d"
198                ));
199            }
200            if coordinate
201                .name
202                .as_deref()
203                .is_some_and(|name| name.trim().is_empty())
204            {
205                return Err(format!(
206                    "latent_coordinates['{symbol}'].name must be non-empty"
207                ));
208            }
209        }
210        serde_json::to_value(self)
211            .map_err(|error| format!("failed to serialize latent coordinates: {error}"))
212    }
213}
214
215#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
216#[serde(deny_unknown_fields)]
217pub struct LatentCoordinateDocument {
218    pub n: usize,
219    pub d: usize,
220    #[serde(default, skip_serializing_if = "Option::is_none")]
221    pub name: Option<String>,
222    #[serde(default, skip_serializing_if = "Option::is_none")]
223    pub init: Option<JsonValue>,
224    #[serde(default, skip_serializing_if = "Option::is_none")]
225    pub manifold: Option<JsonValue>,
226    #[serde(default, skip_serializing_if = "Option::is_none")]
227    pub retraction: Option<JsonValue>,
228    #[serde(default, skip_serializing_if = "Option::is_none")]
229    pub aux_prior: Option<JsonValue>,
230    #[serde(default, skip_serializing_if = "Option::is_none")]
231    pub dim_selection: Option<JsonValue>,
232    #[serde(default, skip_serializing_if = "Option::is_none")]
233    pub aux_outcome: Option<JsonValue>,
234    #[serde(default, skip_serializing_if = "Option::is_none")]
235    pub id_mode: Option<String>,
236}
237
238#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
239#[serde(transparent)]
240pub struct AnalyticPenaltiesDocument(pub Vec<JsonValue>);
241
242impl AnalyticPenaltiesDocument {
243    pub fn to_json_value(&self) -> Result<JsonValue, String> {
244        for (index, descriptor) in self.0.iter().enumerate() {
245            let descriptor = descriptor
246                .as_object()
247                .ok_or_else(|| format!("analytic_penalties[{index}] must be an object"))?;
248            if !descriptor.get("target").is_some_and(JsonValue::is_string) {
249                return Err(format!(
250                    "analytic_penalties[{index}].target must be a latent-coordinate name"
251                ));
252            }
253        }
254        serde_json::to_value(self)
255            .map_err(|error| format!("failed to serialize analytic penalties: {error}"))
256    }
257}
258
259#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
260#[serde(transparent)]
261pub struct SmoothDescriptorsDocument(pub BTreeMap<String, JsonValue>);
262
263impl SmoothDescriptorsDocument {
264    pub fn to_json_value(&self) -> Result<JsonValue, String> {
265        for (symbol, descriptor) in &self.0 {
266            if symbol.trim().is_empty() {
267                return Err("smooth_descriptors keys must be non-empty symbols".to_string());
268            }
269            if !descriptor.is_object() {
270                return Err(format!("smooth_descriptors['{symbol}'] must be an object"));
271            }
272        }
273        serde_json::to_value(self)
274            .map_err(|error| format!("failed to serialize smooth descriptors: {error}"))
275    }
276}
277
278#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
279#[serde(deny_unknown_fields)]
280pub struct CtnStage1Document {
281    pub response_column: String,
282    pub covariate_formula_rhs: String,
283    #[serde(default, skip_serializing_if = "Option::is_none")]
284    pub config: Option<CtnStage1ConfigDocument>,
285    #[serde(default, skip_serializing_if = "Option::is_none")]
286    pub weight_column: Option<String>,
287    #[serde(default, skip_serializing_if = "Option::is_none")]
288    pub offset_column: Option<String>,
289}
290
291#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
292#[serde(deny_unknown_fields)]
293pub struct CtnStage1ConfigDocument {
294    #[serde(default, skip_serializing_if = "Option::is_none")]
295    pub response_degree: Option<usize>,
296    #[serde(default, skip_serializing_if = "Option::is_none")]
297    pub response_num_internal_knots: Option<usize>,
298    #[serde(default, skip_serializing_if = "Option::is_none")]
299    pub response_penalty_order: Option<usize>,
300    #[serde(default, skip_serializing_if = "Option::is_none")]
301    pub response_extra_penalty_orders: Option<Vec<usize>>,
302    #[serde(default, skip_serializing_if = "Option::is_none")]
303    pub double_penalty: Option<bool>,
304}
305
306#[cfg(test)]
307mod tests {
308    use super::*;
309    use serde_json::json;
310
311    #[test]
312    fn canonical_document_round_trips_identically() {
313        let document = FitRequestDocument::new(
314            "y ~ duchon(x)",
315            FitRequestConfigDocument {
316                ctn_stage1: Some(CtnStage1Document {
317                    response_column: "dose".to_string(),
318                    covariate_formula_rhs: "s(age)".to_string(),
319                    config: Some(CtnStage1ConfigDocument {
320                        response_degree: Some(4),
321                        response_penalty_order: Some(2),
322                        ..CtnStage1ConfigDocument::default()
323                    }),
324                    weight_column: Some("case_weight".to_string()),
325                    offset_column: None,
326                }),
327                latent_coordinates: Some(
328                    serde_json::from_value(json!({
329                        "x": {"d": 2, "init": "pca", "n": 12, "name": "x"}
330                    }))
331                    .unwrap(),
332                ),
333                analytic_penalties: Some(AnalyticPenaltiesDocument(vec![json!(
334                    {"kind": "orthogonality", "target": "x", "weight": 1.0}
335                )])),
336                precision_hyperpriors: Some(BTreeMap::from([(
337                    "s(x):roughness".to_string(),
338                    PrecisionHyperpriorDocument {
339                        shape: 2.0,
340                        rate: 0.5,
341                    },
342                )])),
343                smooth_descriptors: Some(
344                    serde_json::from_value(json!({
345                        "x": {"centers": 8, "kind": "duchon", "vars": ["x"]}
346                    }))
347                    .unwrap(),
348                ),
349                ..FitRequestConfigDocument::default()
350            },
351        )
352        .unwrap();
353
354        let encoded = document.to_canonical_json().unwrap();
355        let decoded = FitRequestDocument::from_json(&encoded).unwrap();
356        assert_eq!(decoded, document);
357        assert_eq!(decoded.to_canonical_json().unwrap(), encoded);
358    }
359
360    #[test]
361    fn parser_rejects_another_schema_or_version() {
362        let wrong_schema = r#"{"schema":"other","schema_version":1,"formula":"y ~ x","config":{}}"#;
363        assert!(
364            FitRequestDocument::from_json(wrong_schema)
365                .unwrap_err()
366                .contains("schema must be")
367        );
368
369        let wrong_version =
370            r#"{"schema":"gam.fit-request","schema_version":2,"formula":"y ~ x","config":{}}"#;
371        assert!(
372            FitRequestDocument::from_json(wrong_version)
373                .unwrap_err()
374                .contains("unsupported fit request schema_version")
375        );
376    }
377}