Skip to main content

prospicio_prob/
serial.rs

1//! Saving and loading distributions: [`Dist::to_json`] and
2//! [`Dist::from_json`], a versioned JSON document per distribution.
3//!
4//! The document holds the family and the parameters the family's
5//! constructor takes, so a loaded distribution is rebuilt by the same
6//! validated constructor and equals the saved one. Numbers round-trip bit
7//! for bit; a non-finite value is written as `"NaN"`, `"inf"` or `"-inf"`.
8//! A mixture is saved with its components, so it must have been built from
9//! native severities ([`Mixture::from_dists`], which the bindings use). A
10//! [`Custom`](crate::Custom) cannot be saved: it is a function in the
11//! caller's language.
12//!
13//! ```json
14//! {"format": "risk_rs.distribution", "format_version": 1,
15//!  "family": "lognormal", "meanlog": 7.0, "sdlog": 0.5}
16//! ```
17
18use 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
31/// Value of the `format` field.
32const FORMAT: &str = "risk_rs.distribution";
33/// The format version this build writes and the newest it reads.
34const FORMAT_VERSION: u64 = 1;
35
36impl Dist {
37    /// The distribution as a JSON document; [`from_json`](Self::from_json)
38    /// reads it back to an equal distribution.
39    ///
40    /// Fails for a [`Custom`](crate::Custom) (a function in the caller's
41    /// language) and for a mixture built from trait objects rather than
42    /// with [`Mixture::from_dists`].
43    ///
44    /// ```
45    /// use prospicio_prob::{Dist, Distribution, Lognormal};
46    ///
47    /// let d = Dist::from(Lognormal::new(7.0, 0.5).unwrap());
48    /// let back = Dist::from_json(&d.to_json().unwrap()).unwrap();
49    /// assert_eq!(back.family(), "lognormal");
50    /// assert_eq!(back.mean(), d.mean());
51    /// ```
52    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    /// Reads a document written by [`to_json`](Self::to_json). Fails on
63    /// malformed JSON, another format, a newer format version, an unknown
64    /// family, or parameters the family's constructor refuses.
65    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
90/// The family and parameters, without the format fields.
91fn 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}