1use std::sync::Arc;
10
11use prospicio_core::{Result, StreamRng};
12
13use crate::distribution::Distribution;
14use crate::evt::Gpd;
15use crate::severity::Severity;
16use crate::{
17 Custom, Gamma, Grid, LogAffinePareto, Loglogistic, Lognormal, Mixture, Pareto, PiecewisePareto,
18 Sampled, Tweedie, Weibull,
19};
20
21#[derive(Debug, Clone)]
43pub enum Dist {
44 Lognormal(Lognormal),
45 Pareto(Pareto),
46 PiecewisePareto(PiecewisePareto),
47 LogAffinePareto(LogAffinePareto),
48 GeneralizedPareto(Gpd),
49 Gamma(Gamma),
50 Tweedie(Tweedie),
51 Weibull(Weibull),
52 Loglogistic(Loglogistic),
53 Mixture(Arc<Mixture>),
54 Grid(Grid),
55 Sampled(Sampled),
56 Custom(Custom),
58}
59
60macro_rules! each {
62 ($self:ident, $d:ident => $call:expr) => {
63 match $self {
64 Dist::Lognormal($d) => $call,
65 Dist::Pareto($d) => $call,
66 Dist::PiecewisePareto($d) => $call,
67 Dist::LogAffinePareto($d) => $call,
68 Dist::GeneralizedPareto($d) => $call,
69 Dist::Gamma($d) => $call,
70 Dist::Tweedie($d) => $call,
71 Dist::Weibull($d) => $call,
72 Dist::Loglogistic($d) => $call,
73 Dist::Mixture($d) => $call,
74 Dist::Grid($d) => $call,
75 Dist::Sampled($d) => $call,
76 Dist::Custom($d) => $call,
77 }
78 };
79}
80
81impl Dist {
82 pub fn family(&self) -> &'static str {
87 match self {
88 Self::Lognormal(_) => "lognormal",
89 Self::Pareto(_) => "pareto",
90 Self::PiecewisePareto(_) => "piecewise_pareto",
91 Self::LogAffinePareto(_) => "log_affine_pareto",
92 Self::GeneralizedPareto(_) => "generalized_pareto",
93 Self::Gamma(_) => "gamma",
94 Self::Tweedie(_) => "tweedie",
95 Self::Weibull(_) => "weibull",
96 Self::Loglogistic(_) => "loglogistic",
97 Self::Mixture(_) => "mixture",
98 Self::Grid(_) => "grid",
99 Self::Sampled(_) => "sampled",
100 Self::Custom(_) => "custom",
101 }
102 }
103
104 pub fn as_severity(&self) -> Option<&(dyn Severity + Send + Sync)> {
108 Some(match self {
109 Self::Lognormal(d) => d,
110 Self::Pareto(d) => d,
111 Self::PiecewisePareto(d) => d,
112 Self::LogAffinePareto(d) => d,
113 Self::GeneralizedPareto(d) => d,
114 Self::Gamma(d) => d,
115 Self::Tweedie(d) => d,
116 Self::Weibull(d) => d,
117 Self::Loglogistic(d) => d,
118 Self::Mixture(d) => d.as_ref(),
119 Self::Grid(d) => d,
120 Self::Custom(d) => d,
121 Self::Sampled(_) => return None,
122 })
123 }
124}
125
126impl Distribution for Dist {
127 fn mean(&self) -> f64 {
128 each!(self, d => d.mean())
129 }
130
131 fn variance(&self) -> f64 {
132 each!(self, d => d.variance())
133 }
134
135 fn cdf(&self, x: f64) -> f64 {
136 each!(self, d => d.cdf(x))
137 }
138
139 fn survival(&self, x: f64) -> f64 {
140 each!(self, d => d.survival(x))
141 }
142
143 fn quantile(&self, p: f64) -> Result<f64> {
144 each!(self, d => d.quantile(p))
145 }
146
147 fn sample(&self, rng: &mut StreamRng, n: usize) -> Vec<f64> {
148 each!(self, d => d.sample(rng, n))
149 }
150
151 fn is_parallel_safe(&self) -> bool {
152 each!(self, d => d.is_parallel_safe())
153 }
154}
155
156macro_rules! from_family {
157 ($($variant:ident($ty:ty)),* $(,)?) => {
158 $(
159 impl From<$ty> for Dist {
160 fn from(d: $ty) -> Self {
161 Self::$variant(d)
162 }
163 }
164 )*
165 };
166}
167
168from_family!(
169 Lognormal(Lognormal),
170 Pareto(Pareto),
171 PiecewisePareto(PiecewisePareto),
172 LogAffinePareto(LogAffinePareto),
173 GeneralizedPareto(Gpd),
174 Gamma(Gamma),
175 Tweedie(Tweedie),
176 Weibull(Weibull),
177 Loglogistic(Loglogistic),
178 Grid(Grid),
179 Sampled(Sampled),
180 Custom(Custom),
181);
182
183#[derive(Debug, Clone)]
201pub struct SeverityDist(Dist);
202
203impl SeverityDist {
204 pub fn dist(&self) -> &Dist {
206 &self.0
207 }
208
209 pub fn into_dist(self) -> Dist {
211 self.0
212 }
213}
214
215impl TryFrom<Dist> for SeverityDist {
216 type Error = Dist;
218
219 fn try_from(d: Dist) -> std::result::Result<Self, Dist> {
220 match d {
221 Dist::Sampled(_) => Err(d),
222 d => Ok(Self(d)),
223 }
224 }
225}
226
227impl From<SeverityDist> for Dist {
228 fn from(s: SeverityDist) -> Self {
229 s.0
230 }
231}
232
233macro_rules! each_severity {
235 ($self:ident, $d:ident => $call:expr) => {
236 match &$self.0 {
237 Dist::Lognormal($d) => $call,
238 Dist::Pareto($d) => $call,
239 Dist::PiecewisePareto($d) => $call,
240 Dist::LogAffinePareto($d) => $call,
241 Dist::GeneralizedPareto($d) => $call,
242 Dist::Gamma($d) => $call,
243 Dist::Tweedie($d) => $call,
244 Dist::Weibull($d) => $call,
245 Dist::Loglogistic($d) => $call,
246 Dist::Mixture($d) => $call,
247 Dist::Grid($d) => $call,
248 Dist::Custom($d) => $call,
249 Dist::Sampled(_) => unreachable!("SeverityDist never holds Sampled"),
250 }
251 };
252}
253
254impl Distribution for SeverityDist {
255 fn mean(&self) -> f64 {
256 self.0.mean()
257 }
258
259 fn variance(&self) -> f64 {
260 self.0.variance()
261 }
262
263 fn cdf(&self, x: f64) -> f64 {
264 self.0.cdf(x)
265 }
266
267 fn survival(&self, x: f64) -> f64 {
268 self.0.survival(x)
269 }
270
271 fn quantile(&self, p: f64) -> Result<f64> {
272 self.0.quantile(p)
273 }
274
275 fn sample(&self, rng: &mut StreamRng, n: usize) -> Vec<f64> {
276 self.0.sample(rng, n)
277 }
278
279 fn is_parallel_safe(&self) -> bool {
280 self.0.is_parallel_safe()
281 }
282}
283
284impl Severity for SeverityDist {
285 fn lev(&self, limit: f64) -> f64 {
286 each_severity!(self, d => d.lev(limit))
287 }
288
289 fn stop_loss(&self, retention: f64) -> f64 {
290 each_severity!(self, d => d.stop_loss(retention))
291 }
292
293 fn layer(&self, limit: f64, attachment: f64) -> f64 {
294 each_severity!(self, d => d.layer(limit, attachment))
295 }
296
297 fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
298 each_severity!(self, d => d.layer_second_moment(limit, attachment))
299 }
300
301 fn layer_variance(&self, limit: f64, attachment: f64) -> f64 {
302 each_severity!(self, d => d.layer_variance(limit, attachment))
303 }
304}
305
306impl From<Mixture> for Dist {
307 fn from(d: Mixture) -> Self {
308 Self::Mixture(Arc::new(d))
309 }
310}
311
312impl From<Arc<Mixture>> for Dist {
313 fn from(d: Arc<Mixture>) -> Self {
314 Self::Mixture(d)
315 }
316}
317
318#[cfg(test)]
319mod tests {
320 use super::*;
321
322 fn all() -> Vec<Dist> {
323 let ln = Lognormal::from_mean_cv(1000.0, 0.8).unwrap();
324 vec![
325 ln.into(),
326 Pareto::new(500.0, 2.5).unwrap().into(),
327 Gamma::new(2.0, 500.0).unwrap().into(),
328 Weibull::new(1.5, 1000.0).unwrap().into(),
329 Loglogistic::new(4.0, 900.0).unwrap().into(),
330 Tweedie::new(1000.0, 2.0, 1.5).unwrap().into(),
331 Gpd::new(0.2, 300.0).unwrap().into(),
332 Mixture::new(vec![
333 (0.5, Box::new(ln) as Box<dyn Severity + Send + Sync>),
334 (0.5, Box::new(Gamma::new(2.0, 500.0).unwrap())),
335 ])
336 .unwrap()
337 .into(),
338 Sampled::new(vec![1.0, 2.0, 3.0, 10.0]).unwrap().into(),
339 Custom::new(
340 "exponential",
341 Arc::new(|x: f64| Ok((1.0 - (-x / 500.0).exp()).max(0.0))),
342 None,
343 false,
344 )
345 .unwrap()
346 .into(),
347 ]
348 }
349
350 #[test]
351 fn dispatch_matches_the_family() {
352 let ln = Lognormal::from_mean_cv(1000.0, 0.8).unwrap();
353 let d = Dist::from(ln);
354 assert_eq!(d.family(), "lognormal");
355 assert_eq!(d.mean(), ln.mean());
356 assert_eq!(d.variance(), ln.variance());
357 assert_eq!(d.cdf(700.0), ln.cdf(700.0));
358 assert_eq!(d.quantile(0.9).unwrap(), ln.quantile(0.9).unwrap());
359 let s = d.as_severity().unwrap();
360 assert_eq!(s.lev(1500.0), ln.lev(1500.0));
361 assert_eq!(s.layer(500.0, 1000.0), ln.layer(500.0, 1000.0));
362 let (mut a, mut b) = (StreamRng::new(1, 2), StreamRng::new(1, 2));
363 assert_eq!(d.sample(&mut a, 5), ln.sample(&mut b, 5));
364 }
365
366 #[test]
367 fn every_variant_is_a_distribution_and_severities_have_layers() {
368 for d in all() {
369 assert!(d.mean().is_finite(), "{}", d.family());
370 assert_eq!(d.is_parallel_safe(), d.family() != "custom");
371 let q = d.quantile(0.5).unwrap();
372 assert!(d.cdf(q) >= 0.5 - 1e-9, "{}", d.family());
373 match d.as_severity() {
374 Some(s) => {
375 let m = s.lev(f64::INFINITY);
376 assert!((m / d.mean() - 1.0).abs() < 1e-6, "{}", d.family());
377 }
378 None => assert_eq!(d.family(), "sampled"),
379 }
380 let c = d.clone();
382 assert_eq!(c.mean(), d.mean());
383 }
384 }
385
386 #[test]
387 fn severity_dist_matches_as_severity_and_rejects_sampled() {
388 for d in all() {
389 let family = d.family();
390 match SeverityDist::try_from(d.clone()) {
391 Ok(s) => {
392 let r = d.as_severity().unwrap();
393 assert_eq!(s.lev(700.0), r.lev(700.0), "{family}");
394 assert_eq!(s.stop_loss(700.0), r.stop_loss(700.0), "{family}");
395 assert_eq!(s.layer(500.0, 200.0), r.layer(500.0, 200.0), "{family}");
396 assert_eq!(
397 s.layer_second_moment(500.0, 200.0),
398 r.layer_second_moment(500.0, 200.0),
399 "{family}"
400 );
401 assert_eq!(s.mean(), d.mean(), "{family}");
402 assert_eq!(s.dist().family(), family);
403 }
404 Err(back) => {
405 assert_eq!(family, "sampled");
406 assert_eq!(back.family(), "sampled");
407 }
408 }
409 }
410 }
411}