1use serde::{Deserialize, Serialize};
2use serde_json::Value as JsonValue;
3use std::collections::BTreeMap;
4
5pub const FIT_REQUEST_SCHEMA: &str = "gam.fit-request";
7
8pub const FIT_REQUEST_SCHEMA_VERSION: u32 = 1;
10
11#[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#[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")]
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}