1use serde::{Deserialize, Serialize};
2use serde_json::Value as JsonValue;
3use std::collections::BTreeMap;
4use std::path::PathBuf;
5
6pub const FIT_REQUEST_SCHEMA: &str = "gam.fit-request";
8
9pub const FIT_REQUEST_SCHEMA_VERSION: u32 = 1;
11
12#[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#[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 pub precompute_conformal: Option<bool>,
150 #[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 #[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 #[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 #[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}