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")]
164 pub topology_auto_selector: Option<JsonValue>,
165 #[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}