1use std::sync::Arc;
19
20use prospicio_core::{Error, Result};
21use serde_json::{Map, Number, Value, json};
22
23use crate::dist::{Dist, SeverityDist};
24use crate::evt::Gpd;
25use crate::sampled::Empirical;
26use crate::{
27 Gamma, Grid, LogAffinePareto, Loglogistic, Lognormal, Mixture, Pareto, PiecewisePareto,
28 Sampled, Truncation, Tweedie, Weibull,
29};
30
31const FORMAT: &str = "risk_rs.distribution";
33const FORMAT_VERSION: u64 = 1;
35
36impl Dist {
37 pub fn to_json(&self) -> Result<String> {
53 let mut doc = Map::new();
54 doc.insert("format".into(), json!(FORMAT));
55 doc.insert("format_version".into(), json!(FORMAT_VERSION));
56 for (k, v) in body(self)? {
57 doc.insert(k, v);
58 }
59 Ok(serde_json::to_string(&Value::Object(doc)).expect("a JSON value serializes"))
60 }
61
62 pub fn from_json(text: &str) -> Result<Self> {
66 let doc: Value = serde_json::from_str(text)
67 .map_err(|e| Error::Data(format!("not a JSON document: {e}")))?;
68 let top = object(&doc, "document")?;
69 match top.get("format").and_then(Value::as_str) {
70 Some(FORMAT) => {}
71 other => {
72 return Err(Error::Data(format!(
73 "not a distribution document: format {other:?}, expected {FORMAT:?}"
74 )));
75 }
76 }
77 let version = top
78 .get("format_version")
79 .and_then(Value::as_u64)
80 .ok_or_else(|| Error::Data("format_version is missing".into()))?;
81 if version > FORMAT_VERSION {
82 return Err(Error::Data(format!(
83 "format_version {version} is newer than this build reads ({FORMAT_VERSION})"
84 )));
85 }
86 from_body(top)
87 }
88}
89
90fn body(d: &Dist) -> Result<Map<String, Value>> {
92 let mut m = Map::new();
93 m.insert("family".into(), json!(d.family()));
94 let mut put = |k: &str, v: Value| {
95 m.insert(k.into(), v);
96 };
97 match d {
98 Dist::Lognormal(x) => {
99 put("meanlog", num(x.meanlog()));
100 put("sdlog", num(x.sdlog()));
101 }
102 Dist::Pareto(x) => {
103 put("t", num(x.t()));
104 put("alpha", num(x.alpha()));
105 put("truncation", x.truncation().map_or(Value::Null, num));
106 }
107 Dist::PiecewisePareto(x) => {
108 put("t", nums(x.thresholds()));
109 put("alpha", nums(x.alphas()));
110 match x.truncation() {
111 Some((t, kind)) => {
112 put("truncation", num(t));
113 put("truncation_type", json!(truncation_name(kind)));
114 }
115 None => put("truncation", Value::Null),
116 }
117 }
118 Dist::LogAffinePareto(x) => {
119 put("t", num(x.t()));
120 put("alpha0", num(x.alpha0()));
121 put("gamma", num(x.gamma()));
122 }
123 Dist::GeneralizedPareto(x) => {
124 put("xi", num(x.xi()));
125 put("beta", num(x.beta()));
126 put("location", num(x.location()));
127 }
128 Dist::Gamma(x) => {
129 put("shape", num(x.shape()));
130 put("scale", num(x.scale()));
131 }
132 Dist::Tweedie(x) => {
133 put("mean", num(x.mean_param()));
134 put("dispersion", num(x.dispersion()));
135 put("power", num(x.power()));
136 }
137 Dist::Weibull(x) => {
138 put("shape", num(x.shape()));
139 put("scale", num(x.scale()));
140 }
141 Dist::Loglogistic(x) => {
142 put("shape", num(x.shape()));
143 put("scale", num(x.scale()));
144 }
145 Dist::Mixture(x) => {
146 let dists = x.dists().ok_or_else(|| {
147 Error::Data(
148 "this mixture was built from trait objects; build it with \
149 Mixture::from_dists to save it"
150 .into(),
151 )
152 })?;
153 put("weights", nums(x.weights()));
154 let parts = dists
155 .iter()
156 .map(|c| body(c.dist()).map(Value::Object))
157 .collect::<Result<Vec<_>>>()?;
158 put("components", Value::Array(parts));
159 }
160 Dist::Grid(x) => {
161 put("step", num(x.step()));
162 put("probs", nums(x.probs()));
163 }
164 Dist::Sampled(x) => {
165 put("draws", nums(x.draws()));
166 }
167 Dist::Custom(x) => {
168 return Err(Error::Data(format!(
169 "custom distribution {:?} is a function in the caller's language and \
170 cannot be saved",
171 x.name()
172 )));
173 }
174 }
175 Ok(m)
176}
177
178fn from_body(m: &Map<String, Value>) -> Result<Dist> {
179 let family = m
180 .get("family")
181 .and_then(Value::as_str)
182 .ok_or_else(|| Error::Data("family is missing".into()))?;
183 let f = |k: &str| float(m, k);
184 Ok(match family {
185 "lognormal" => Lognormal::new(f("meanlog")?, f("sdlog")?)?.into(),
186 "pareto" => {
187 let p = Pareto::new(f("t")?, f("alpha")?)?;
188 match optional(m, "truncation")? {
189 Some(t) => p.truncated(t)?,
190 None => p,
191 }
192 .into()
193 }
194 "piecewise_pareto" => {
195 let p = PiecewisePareto::new(floats(m, "t")?, floats(m, "alpha")?)?;
196 match optional(m, "truncation")? {
197 Some(t) => {
198 let kind = match m.get("truncation_type").and_then(Value::as_str) {
199 Some("lp") => Truncation::LastPiece,
200 Some("wd") => Truncation::WholeDistribution,
201 other => {
202 return Err(Error::Data(format!(
203 "truncation_type must be \"lp\" or \"wd\", got {other:?}"
204 )));
205 }
206 };
207 p.truncated(t, kind)?
208 }
209 None => p,
210 }
211 .into()
212 }
213 "log_affine_pareto" => LogAffinePareto::new(f("t")?, f("alpha0")?, f("gamma")?)?.into(),
214 "generalized_pareto" => Gpd::new(f("xi")?, f("beta")?)?
215 .shifted(f("location")?)?
216 .into(),
217 "gamma" => Gamma::new(f("shape")?, f("scale")?)?.into(),
218 "tweedie" => Tweedie::new(f("mean")?, f("dispersion")?, f("power")?)?.into(),
219 "weibull" => Weibull::new(f("shape")?, f("scale")?)?.into(),
220 "loglogistic" => Loglogistic::new(f("shape")?, f("scale")?)?.into(),
221 "mixture" => {
222 let weights = floats(m, "weights")?;
223 let comps = m
224 .get("components")
225 .and_then(Value::as_array)
226 .ok_or_else(|| Error::Data("components is missing".into()))?;
227 if comps.len() != weights.len() {
228 return Err(Error::Data("one weight per mixture component".into()));
229 }
230 let parts = weights
231 .into_iter()
232 .zip(comps)
233 .map(|(w, c)| {
234 let d = from_body(object(c, "component")?)?;
235 let s = SeverityDist::try_from(d).map_err(|_| {
236 Error::Data("a mixture component cannot be sampled draws".into())
237 })?;
238 Ok((w, s))
239 })
240 .collect::<Result<Vec<_>>>()?;
241 Dist::Mixture(Arc::new(Mixture::from_dists(parts)?))
242 }
243 "grid" => Grid::new(f("step")?, floats(m, "probs")?)?.into(),
244 "sampled" => Sampled::new(floats(m, "draws")?)?.into(),
245 other => {
246 return Err(Error::Data(format!("unknown family {other:?}")));
247 }
248 })
249}
250
251fn truncation_name(kind: Truncation) -> &'static str {
252 match kind {
253 Truncation::LastPiece => "lp",
254 Truncation::WholeDistribution => "wd",
255 }
256}
257
258fn num(x: f64) -> Value {
259 match Number::from_f64(x) {
260 Some(n) => Value::Number(n),
261 None if x.is_nan() => json!("NaN"),
262 None if x > 0.0 => json!("inf"),
263 None => json!("-inf"),
264 }
265}
266
267fn nums(xs: &[f64]) -> Value {
268 Value::Array(xs.iter().map(|&x| num(x)).collect())
269}
270
271fn to_f64(v: &Value, name: &str) -> Result<f64> {
272 match v {
273 Value::Number(n) => n.as_f64(),
274 Value::String(s) => match s.as_str() {
275 "NaN" => Some(f64::NAN),
276 "inf" => Some(f64::INFINITY),
277 "-inf" => Some(f64::NEG_INFINITY),
278 _ => None,
279 },
280 _ => None,
281 }
282 .ok_or_else(|| Error::Data(format!("{name} must be a number")))
283}
284
285fn float(m: &Map<String, Value>, name: &str) -> Result<f64> {
286 to_f64(
287 m.get(name)
288 .ok_or_else(|| Error::Data(format!("{name} is missing")))?,
289 name,
290 )
291}
292
293fn optional(m: &Map<String, Value>, name: &str) -> Result<Option<f64>> {
294 match m.get(name) {
295 None | Some(Value::Null) => Ok(None),
296 Some(v) => to_f64(v, name).map(Some),
297 }
298}
299
300fn floats(m: &Map<String, Value>, name: &str) -> Result<Vec<f64>> {
301 m.get(name)
302 .and_then(Value::as_array)
303 .ok_or_else(|| Error::Data(format!("{name} must be a list of numbers")))?
304 .iter()
305 .map(|v| to_f64(v, name))
306 .collect()
307}
308
309fn object<'a>(v: &'a Value, what: &str) -> Result<&'a Map<String, Value>> {
310 v.as_object()
311 .ok_or_else(|| Error::Data(format!("the {what} must be a JSON object")))
312}
313
314#[cfg(test)]
315mod tests {
316 use super::*;
317 use crate::{Distribution, Severity};
318
319 fn every() -> Vec<Dist> {
320 let ln = Lognormal::new(7.0, 0.5).unwrap();
321 vec![
322 ln.into(),
323 Pareto::new(1e5, 1.5).unwrap().into(),
324 Pareto::new(1e5, 1.5)
325 .unwrap()
326 .truncated(1e7)
327 .unwrap()
328 .into(),
329 PiecewisePareto::new(vec![1.0, 10.0, 100.0], vec![1.2, 1.8, 2.5])
330 .unwrap()
331 .into(),
332 PiecewisePareto::new(vec![1.0, 10.0], vec![1.2, 1.8])
333 .unwrap()
334 .truncated(1000.0, Truncation::WholeDistribution)
335 .unwrap()
336 .into(),
337 LogAffinePareto::new(100.0, 1.5, 0.3).unwrap().into(),
338 Gpd::new(0.25, 3.0).unwrap().shifted(10.0).unwrap().into(),
339 Gamma::new(2.0, 500.0).unwrap().into(),
340 Tweedie::new(1000.0, 2.0, 1.5).unwrap().into(),
341 Weibull::new(1.5, 1000.0).unwrap().into(),
342 Loglogistic::new(4.0, 900.0).unwrap().into(),
343 Dist::Mixture(Arc::new(
344 Mixture::from_dists(vec![
345 (0.7, SeverityDist::try_from(Dist::from(ln)).unwrap()),
346 (
347 0.3,
348 SeverityDist::try_from(Dist::from(Pareto::new(1e5, 2.0).unwrap())).unwrap(),
349 ),
350 ])
351 .unwrap(),
352 )),
353 Grid::new(0.5, vec![0.1, 0.4, 0.3, 0.2]).unwrap().into(),
354 Sampled::new(vec![3.0, 1.0, 2.0, 0.1 + 0.2]).unwrap().into(),
355 ]
356 }
357
358 #[test]
359 fn every_family_round_trips_exactly() {
360 for d in every() {
361 let text = d.to_json().unwrap();
362 let back = Dist::from_json(&text).unwrap();
363 assert_eq!(back.family(), d.family(), "{text}");
364 assert_eq!(back.to_json().unwrap(), text, "{text}");
365 assert_eq!(back.mean().to_bits(), d.mean().to_bits(), "{text}");
366 for p in [0.1, 0.5, 0.9] {
367 let (a, b) = (back.quantile(p).unwrap(), d.quantile(p).unwrap());
368 assert_eq!(a.to_bits(), b.to_bits(), "{text} at {p}");
369 }
370 if let (Some(a), Some(b)) = (back.as_severity(), d.as_severity()) {
371 assert_eq!(a.lev(1500.0).to_bits(), b.lev(1500.0).to_bits(), "{text}");
372 }
373 }
374 }
375
376 #[test]
377 fn refuses_what_it_cannot_save_or_read() {
378 let custom = crate::Custom::new(
379 "c",
380 Arc::new(|x: f64| Ok((1.0 - (-x).exp()).max(0.0))),
381 None,
382 true,
383 )
384 .unwrap();
385 assert!(Dist::from(custom).to_json().is_err());
386 let boxed = Mixture::new(vec![(
387 1.0,
388 Box::new(Gamma::new(2.0, 1.0).unwrap()) as Box<dyn Severity + Send + Sync>,
389 )])
390 .unwrap();
391 assert!(Dist::from(boxed).to_json().is_err());
392 assert!(Dist::from_json("not json").is_err());
393 assert!(Dist::from_json(r#"{"format": "risk_rs.glm_fit", "format_version": 1}"#).is_err());
394 let newer = r#"{"format": "risk_rs.distribution", "format_version": 2, "family": "gamma", "shape": 2, "scale": 1}"#;
395 assert!(Dist::from_json(newer).is_err());
396 let bad = r#"{"format": "risk_rs.distribution", "format_version": 1, "family": "gamma", "shape": -2, "scale": 1}"#;
397 assert!(Dist::from_json(bad).is_err());
398 let unknown =
399 r#"{"format": "risk_rs.distribution", "format_version": 1, "family": "cauchy"}"#;
400 assert!(Dist::from_json(unknown).is_err());
401 }
402}