1use std::fmt;
21use std::sync::{Arc, Mutex};
22
23use prospicio_core::{Error, Result};
24use prospicio_math::integrate::gauss_legendre;
25use prospicio_math::roots::bisect_log;
26
27use crate::distribution::{Distribution, check_probability};
28use crate::severity::Severity;
29
30pub type Callback = Arc<dyn Fn(f64) -> std::result::Result<f64, String> + Send + Sync>;
33
34const PANEL_PROBS: [f64; 16] = [
36 0.01, 0.1, 0.25, 0.5, 0.75, 0.9, 0.95, 0.99, 0.999, 1e-4, 1e-5, 1e-6, 1e-7, 1e-8, 1e-10, 1e-12,
37];
38
39const PIECES: usize = 8;
41
42#[derive(Clone)]
58pub struct Custom {
59 name: String,
60 cdf: Callback,
61 quantile: Option<Callback>,
62 parallel_safe: bool,
63 breaks: Vec<f64>,
65 mean: f64,
66 second_moment: f64,
67 error: Arc<Mutex<Option<String>>>,
69}
70
71impl fmt::Debug for Custom {
72 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73 f.debug_struct("Custom")
74 .field("name", &self.name)
75 .field("has_quantile", &self.quantile.is_some())
76 .field("parallel_safe", &self.parallel_safe)
77 .field("mean", &self.mean)
78 .field("upper", &self.upper())
79 .finish()
80 }
81}
82
83fn callback_error(name: &str, what: &str, x: f64, msg: &str) -> Error {
84 Error::Data(format!(
85 "custom distribution {name:?}: {what}({x}) failed: {msg}"
86 ))
87}
88
89impl Custom {
90 pub fn new(
98 name: impl Into<String>,
99 cdf: Callback,
100 quantile: Option<Callback>,
101 parallel_safe: bool,
102 ) -> Result<Self> {
103 let mut d = Self {
104 name: name.into(),
105 cdf,
106 quantile,
107 parallel_safe,
108 breaks: vec![0.0],
109 mean: f64::NAN,
110 second_moment: f64::NAN,
111 error: Arc::new(Mutex::new(None)),
112 };
113 d.try_cdf(0.0)?;
114 let mut breaks = vec![0.0];
115 for p in PANEL_PROBS {
116 let p = if p < 0.01 { 1.0 - p } else { p };
117 let q = d.try_quantile(p)?;
118 if !q.is_finite() {
119 return Err(Error::Data(format!(
120 "custom distribution {:?}: its cdf does not reach {p} at any finite loss",
121 d.name
122 )));
123 }
124 if q > *breaks.last().unwrap() {
125 breaks.push(q);
126 } else if q < *breaks.last().unwrap() {
127 return Err(Error::Data(format!(
128 "custom distribution {:?}: quantiles must increase with p, but q({p}) = {q} \
129 is below an earlier quantile",
130 d.name
131 )));
132 }
133 }
134 d.breaks = breaks;
135 let upper = d.upper();
137 d.mean = d.integrate(0.0, upper, |_, s| s)?;
138 d.second_moment = d.integrate(0.0, upper, |x, s| 2.0 * x * s)?;
139 Ok(d)
140 }
141
142 pub fn name(&self) -> &str {
144 &self.name
145 }
146
147 pub fn has_quantile(&self) -> bool {
150 self.quantile.is_some()
151 }
152
153 pub fn upper(&self) -> f64 {
155 *self.breaks.last().unwrap()
156 }
157
158 pub fn error(&self) -> Option<String> {
161 self.error.lock().map(|e| e.clone()).unwrap_or(None)
162 }
163
164 fn record(&self, e: &Error) {
165 if let Ok(mut slot) = self.error.lock()
166 && slot.is_none()
167 {
168 *slot = Some(e.to_string());
169 }
170 }
171
172 fn try_cdf(&self, x: f64) -> Result<f64> {
173 let v = (self.cdf)(x).map_err(|m| callback_error(&self.name, "cdf", x, &m))?;
174 if !(0.0..=1.0).contains(&v) {
175 return Err(callback_error(
176 &self.name,
177 "cdf",
178 x,
179 &format!("returned {v}, outside [0, 1]"),
180 ));
181 }
182 Ok(v)
183 }
184
185 fn try_quantile(&self, p: f64) -> Result<f64> {
186 check_probability(p)?;
187 if let Some(q) = &self.quantile {
188 let v = q(p).map_err(|m| callback_error(&self.name, "quantile", p, &m))?;
189 if v.is_nan() || v < 0.0 {
190 return Err(callback_error(
191 &self.name,
192 "quantile",
193 p,
194 &format!("returned {v}; losses are non-negative"),
195 ));
196 }
197 return Ok(v);
198 }
199 if self.try_cdf(0.0)? >= p {
200 return Ok(0.0);
201 }
202 let mut hi = 1.0;
204 while self.try_cdf(hi)? < p {
205 hi *= 2.0;
206 if !hi.is_finite() {
207 return Ok(f64::INFINITY);
208 }
209 }
210 let mut lo = hi;
211 loop {
212 lo *= 0.5;
213 if lo == 0.0 {
214 return Ok(0.0);
215 }
216 if self.try_cdf(lo)? < p {
217 break;
218 }
219 }
220 let mut failure = None;
221 let x = bisect_log(lo, hi, |x| match self.try_cdf(x) {
222 Ok(c) => c < p,
223 Err(e) => {
224 failure.get_or_insert(e);
225 false
226 }
227 });
228 match failure {
229 Some(e) => Err(e),
230 None => Ok(x),
231 }
232 }
233
234 fn integrate(&self, lo: f64, hi: f64, g: impl Fn(f64, f64) -> f64) -> Result<f64> {
236 let (lo, hi) = (lo.max(0.0), hi.min(self.upper()));
237 let mut total = 0.0;
238 for w in self.breaks.windows(2) {
239 let (a, b) = (w[0].max(lo), w[1].min(hi));
240 if a >= b {
241 continue;
242 }
243 let step = (b - a) / PIECES as f64;
244 for k in 0..PIECES {
245 let x0 = a + step * k as f64;
246 let x1 = if k + 1 == PIECES { b } else { x0 + step };
247 total += gauss_legendre(|x| self.try_cdf(x).map(|c| g(x, 1.0 - c)), x0, x1)?;
248 }
249 }
250 Ok(total)
251 }
252
253 fn or_nan(&self, r: Result<f64>) -> f64 {
254 r.unwrap_or_else(|e| {
255 self.record(&e);
256 f64::NAN
257 })
258 }
259}
260
261impl Distribution for Custom {
262 fn mean(&self) -> f64 {
263 self.mean
264 }
265
266 fn variance(&self) -> f64 {
267 self.second_moment - self.mean * self.mean
268 }
269
270 fn cdf(&self, x: f64) -> f64 {
271 if x < 0.0 {
272 return 0.0;
273 }
274 self.or_nan(self.try_cdf(x))
275 }
276
277 fn quantile(&self, p: f64) -> Result<f64> {
278 check_probability(p)?;
279 self.try_quantile(p).inspect_err(|e| self.record(e))
280 }
281
282 fn is_parallel_safe(&self) -> bool {
283 self.parallel_safe
284 }
285}
286
287impl Severity for Custom {
288 fn lev(&self, limit: f64) -> f64 {
289 if limit <= 0.0 {
290 return limit;
291 }
292 if limit == f64::INFINITY {
293 return self.mean;
294 }
295 self.or_nan(self.integrate(0.0, limit, |_, s| s))
296 }
297
298 fn stop_loss(&self, retention: f64) -> f64 {
301 if retention <= 0.0 {
302 return self.mean - retention;
303 }
304 self.or_nan(self.integrate(retention, f64::INFINITY, |_, s| s))
305 }
306
307 fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
309 let a = attachment.max(0.0);
310 self.or_nan(self.integrate(a, a + limit, |x, s| 2.0 * (x - a) * s))
311 }
312}
313
314#[cfg(test)]
315mod tests {
316 use super::*;
317 use crate::Lognormal;
318
319 fn from<D: Distribution + Send + Sync + 'static>(d: D, with_quantile: bool) -> Custom {
320 let d = Arc::new(d);
321 let c = d.clone();
322 let cdf: Callback = Arc::new(move |x| Ok(c.cdf(x)));
323 let quantile: Option<Callback> = with_quantile.then(|| {
324 let q = d.clone();
325 Arc::new(move |p| q.quantile(p).map_err(|e| e.to_string())) as Callback
326 });
327 Custom::new("test", cdf, quantile, true).unwrap()
328 }
329
330 #[test]
331 fn moments_and_layers_match_a_closed_form() {
332 let ln = Lognormal::from_mean_cv(1000.0, 1.0).unwrap();
333 for with_quantile in [true, false] {
334 let c = from(ln, with_quantile);
335 assert!((c.mean() / ln.mean() - 1.0).abs() < 1e-7, "{with_quantile}");
336 assert!((c.variance() / ln.variance() - 1.0).abs() < 1e-4);
337 for l in [100.0, 1000.0, 5000.0] {
338 assert!((c.lev(l) / ln.lev(l) - 1.0).abs() < 1e-9, "lev {l}");
339 assert!(
340 (c.stop_loss(l) / ln.stop_loss(l) - 1.0).abs() < 1e-6,
341 "sl {l}"
342 );
343 }
344 let (m, v) = (c.layer(2000.0, 1000.0), c.layer_variance(2000.0, 1000.0));
345 assert!((m / ln.layer(2000.0, 1000.0) - 1.0).abs() < 1e-8);
346 assert!((v / ln.layer_variance(2000.0, 1000.0) - 1.0).abs() < 1e-8);
347 for p in [0.001, 0.5, 0.99] {
348 let q = c.quantile(p).unwrap();
349 assert!((q / ln.quantile(p).unwrap() - 1.0).abs() < 1e-12, "q {p}");
350 }
351 }
352 }
353
354 #[test]
355 fn a_point_mass_at_zero_is_allowed() {
356 let cdf: Callback = Arc::new(|x| {
358 Ok(if x < 0.0 {
359 0.0
360 } else {
361 0.3 + 0.7 * (1.0 - (-x / 10.0).exp())
362 })
363 });
364 let c = Custom::new("zero-inflated", cdf, None, true).unwrap();
365 assert!((c.mean() - 7.0).abs() < 1e-7);
366 assert_eq!(c.quantile(0.2).unwrap(), 0.0);
367 assert_eq!(c.cdf(-1.0), 0.0);
368 }
369
370 #[test]
371 fn bad_callbacks_are_reported() {
372 let failing: Callback = Arc::new(|x| {
373 if x > 50.0 {
374 Err("boom".into())
375 } else {
376 Ok(x / 100.0)
377 }
378 });
379 let e = Custom::new("bad", failing, None, true)
380 .unwrap_err()
381 .to_string();
382 assert!(e.contains("boom"), "{e}");
383
384 let outside: Callback = Arc::new(|_| Ok(1.5));
385 assert!(Custom::new("bad", outside, None, true).is_err());
386
387 let never: Callback = Arc::new(|x| Ok(0.5 * (1.0 - (-x).exp())));
388 let e = Custom::new("bad", never, None, true)
389 .unwrap_err()
390 .to_string();
391 assert!(e.contains("does not reach"), "{e}");
392
393 let late: Callback = Arc::new(|x| {
395 if x == 12345.0 {
396 Err("late".into())
397 } else {
398 Ok(1.0 - (-x).exp())
399 }
400 });
401 let c = Custom::new("late", late, None, false).unwrap();
402 assert!(c.error().is_none());
403 assert!(c.cdf(12345.0).is_nan());
404 assert!(c.error().unwrap().contains("late"));
405 assert!(!c.is_parallel_safe());
406 }
407}