1use crate::dist::{Budget, Dist};
4use crate::error::{OpError, OpResult};
5use crate::value::{EnumValue, Record, Value};
6use probl_number::Integer;
7use probl_syntax::ast::{BinOp, UnOp};
8use std::borrow::Cow;
9use std::sync::Arc;
10
11pub fn integer<'a>(v: &'a Value, context: &str, budget: &mut Budget) -> OpResult<Cow<'a, Integer>> {
14 match v {
15 Value::Int(n) => Ok(Cow::Borrowed(n)),
16 Value::Float(f) if f.is_finite() && f.fract() == 0.0 => {
17 let n = Integer::from_f64(*f).expect("finite integral float fits the integer hard limit");
18 budget.integer_allocation(n.bits(), 1)?;
19 budget.work(n.bits().div_ceil(64).max(1))?;
20 Ok(Cow::Owned(n))
21 }
22 _ => {
23 let found = match v {
24 Value::Float(f) => format!("the float {f:?}"),
25 _ => article(&v.kind()),
26 };
27 Err(OpError::new(format!(
28 "`{context}` needs an int (or an exactly integral finite float), found {found}"
29 )))
30 }
31 }
32}
33
34pub fn outcomes(v: &Value) -> std::borrow::Cow<'_, [(Value, f64)]> {
36 match v {
37 Value::Dist(d) => std::borrow::Cow::Borrowed(&d.outcomes),
38 other => std::borrow::Cow::Owned(vec![(other.clone(), 1.0)]),
39 }
40}
41
42fn missing(v: &Value) -> f64 {
43 match v {
44 Value::Dist(d) => d.missing,
45 _ => 0.0,
46 }
47}
48
49pub fn combine(results: Vec<(Value, f64)>, missing: f64, budget: &mut Budget) -> OpResult<Value> {
54 if results.is_empty() {
55 return Err(OpError::new("the result would be an empty distribution"));
56 }
57 let mut flat = Vec::with_capacity(results.len());
58 let mut missing = missing;
59 for (v, w) in results {
60 match v {
61 Value::Dist(d) => {
62 missing += w * d.missing;
63 flat.extend(d.outcomes.iter().map(|(x, p)| (x.clone(), w * p)));
64 }
65 other => flat.push((other, w)),
66 }
67 }
68 budget.outcomes(flat.len() as u128)?;
69 budget.work(flat.len() as u64)?;
70 Ok(Dist::from_pairs(flat, missing.min(1.0)).into_value())
71}
72
73pub fn lift1(a: &Value, budget: &mut Budget, f: impl Fn(&Value, &mut Budget) -> OpResult<Value>) -> OpResult<Value> {
75 if !a.is_dist() {
76 return f(a, budget);
77 }
78 let mut results = Vec::new();
79 for (v, w) in outcomes(a).iter() {
80 results.push((f(v, budget)?, *w));
81 }
82 combine(results, missing(a), budget)
83}
84
85pub fn lift2(
87 a: &Value,
88 b: &Value,
89 budget: &mut Budget,
90 f: impl Fn(&Value, &Value, &mut Budget) -> OpResult<Value>,
91) -> OpResult<Value> {
92 if !a.is_dist() && !b.is_dist() {
93 return f(a, b, budget);
94 }
95 let (oa, ob) = (outcomes(a), outcomes(b));
96 budget.outcomes(oa.len() as u128 * ob.len() as u128)?;
97 budget.work((oa.len() * ob.len()) as u64)?;
98 let mut results = Vec::with_capacity(oa.len() * ob.len());
99 for (x, p) in oa.iter() {
100 for (y, q) in ob.iter() {
101 results.push((f(x, y, budget)?, p * q));
102 }
103 }
104 let m = missing(a) + (1.0 - missing(a)) * missing(b);
105 combine(results, m, budget)
106}
107
108pub type NaryFn<'a> = &'a dyn Fn(&[Value], &mut Budget) -> OpResult<Value>;
111
112pub fn lift_n(args: &[Value], budget: &mut Budget, f: NaryFn) -> OpResult<Value> {
114 if !args.iter().any(Value::is_dist) {
115 return f(args, budget);
116 }
117 let size = args
118 .iter()
119 .fold(1u128, |acc, a| acc.saturating_mul(outcomes(a).len() as u128));
120 budget.outcomes(size)?;
121 budget.work(size as u64)?;
122 let missing = args.iter().fold(0.0, |m, a| m + (1.0 - m) * missing(a));
123 let mut results = Vec::new();
124 let mut current = Vec::with_capacity(args.len());
125 product(args, 0, &mut current, 1.0, f, budget, &mut results)?;
126 combine(results, missing, budget)
127}
128
129fn product(
130 args: &[Value],
131 i: usize,
132 current: &mut Vec<Value>,
133 weight: f64,
134 f: NaryFn,
135 budget: &mut Budget,
136 results: &mut Vec<(Value, f64)>,
137) -> OpResult<()> {
138 if i == args.len() {
139 results.push((f(current, budget)?, weight));
140 return Ok(());
141 }
142 for (v, w) in outcomes(&args[i]).iter() {
143 current.push(v.clone());
144 product(args, i + 1, current, weight * w, f, budget, results)?;
145 current.pop();
146 }
147 Ok(())
148}
149
150pub fn article(kind: &str) -> String {
153 let vowel = kind.starts_with(['a', 'e', 'i', 'o', 'u']);
154 format!("{} {kind}", if vowel { "an" } else { "a" })
155}
156
157pub fn to_prob(v: &Value) -> OpResult<f64> {
160 if crate::analytic::contains(v) {
161 return Err(crate::analytic::unsupported(
162 "using an analytic outcome as a probability parameter",
163 ));
164 }
165 match v {
166 Value::Prob(_) | Value::Float(_) | Value::Int(_) => match make_prob(v)? {
167 Value::Prob(p) => Ok(p),
168 _ => unreachable!("make_prob returns a probability"),
169 },
170 Value::Bool(_) => Err(OpError::new("expected a probability, found a fact (true or false)")
171 .help("convert a boolean explicitly with `prob(fact)`")),
172 Value::Dist(_) => Err(OpError::new(format!("expected a probability, found a {}", v.kind()))
173 .help("`P(…)` gives the probability that a distribution of facts is true")),
174 other => Err(OpError::new(format!(
175 "expected a probability, found {}",
176 article(&other.kind())
177 ))),
178 }
179}
180
181pub fn make_prob(v: &Value) -> OpResult<Value> {
183 let p = match v {
184 Value::Bool(b) => {
185 if *b {
186 1.0
187 } else {
188 0.0
189 }
190 }
191 Value::Prob(p) | Value::Float(p) => *p,
192 Value::Int(n) if *n == 0 => 0.0,
193 Value::Int(n) if *n == 1 => 1.0,
194 Value::Int(_) => return Err(OpError::new("prob needs a finite number between 0 and 1")),
195 _ => {
196 return Err(
197 OpError::new(format!("prob needs a number or bool, found {}", article(&v.kind())))
198 .help("draw distribution outcomes explicitly; use `P(d)` to query a boolean distribution"),
199 );
200 }
201 };
202 if !p.is_finite() || !(0.0..=1.0).contains(&p) {
203 return Err(OpError::new("prob needs a finite number between 0 and 1"));
204 }
205 Ok(Value::Prob(p))
206}
207
208pub(crate) fn computed_prob(p: f64, what: &str) -> OpResult<Value> {
212 const ROUNDING: f64 = 8.0 * f64::EPSILON;
213 if !p.is_finite() || !(-ROUNDING..=1.0 + ROUNDING).contains(&p) {
214 return Err(OpError::new(format!(
215 "`{what}` could not compute a finite probability between 0 and 1"
216 )));
217 }
218 Ok(Value::Prob(if p <= 0.0 { 0.0 } else { p.min(1.0) }))
219}
220
221pub fn fact(v: &Value, context: &str) -> OpResult<bool> {
222 match v {
223 Value::Bool(b) => Ok(*b),
224 _ => Err(
225 OpError::new(format!("`{context}` needs a bool, found {}", article(&v.kind())))
226 .help("draw an outcome first with `~`; for a probability, draw from `bernoulli(p)`"),
227 ),
228 }
229}
230
231#[derive(Clone, Copy, Debug, PartialEq)]
233pub struct Condition {
234 pub yes: f64,
236 pub no: f64,
238 pub missing: f64,
241}
242
243pub enum Truth {
245 Fact(bool),
246 Analytic(Arc<crate::analytic::Event>),
247 Probability(f64),
248 Uncertain(Arc<Dist>),
249}
250
251pub fn truth(v: &Value, op: &str) -> OpResult<Truth> {
253 match v {
254 Value::Bool(b) => Ok(Truth::Fact(*b)),
255 Value::Event(e) => Ok(Truth::Analytic(e.clone())),
256 Value::Dist(d) if d.truth().is_some() => Ok(Truth::Uncertain(d.clone())),
257 Value::Prob(_) => Ok(Truth::Probability(to_prob(v)?)),
258 other => Err(OpError::new(format!(
259 "`{op}` needs bool, prob or dist[bool], found {}",
260 article(&other.kind())
261 ))),
262 }
263}
264
265pub fn not(v: &Value, budget: &mut Budget) -> OpResult<Value> {
266 match truth(v, "not")? {
267 Truth::Fact(b) => Ok(Value::Bool(!b)),
268 Truth::Analytic(e) => Ok(crate::analytic::Event {
269 draw: e.draw.clone(),
270 yes: e.yes.complement(),
271 }
272 .value()),
273 Truth::Probability(p) => computed_prob(1.0 - p, "not"),
274 Truth::Uncertain(d) => lift1(&Value::Dist(d), budget, |x, _| match x {
275 Value::Bool(b) => Ok(Value::Bool(!b)),
276 _ => unreachable!("checked by `truth`"),
277 }),
278 }
279}
280
281pub fn logic(and: bool, a: Truth, b: Truth, budget: &mut Budget) -> OpResult<Value> {
283 if matches!(a, Truth::Analytic(_)) || matches!(b, Truth::Analytic(_)) {
284 let value = |t| match t {
285 Truth::Fact(b) => Value::Bool(b),
286 Truth::Probability(p) => Value::Prob(p),
287 Truth::Uncertain(d) => Value::Dist(d),
288 Truth::Analytic(e) => Value::Event(e),
289 };
290 let (a, b) = (value(a), value(b));
291 for v in [&a, &b] {
292 if let Value::Event(e) = v {
293 budget.collection(e.yes.0.len() as u128 + e.draw.domain.0.len() as u128)?;
294 budget.work((e.yes.0.len() + e.draw.domain.0.len()) as u64)?;
295 }
296 }
297 return crate::analytic::logic(and, &a, &b);
298 }
299 let op = |x: bool, y: bool| if and { x && y } else { x || y };
300 match (a, b) {
301 (Truth::Fact(x), Truth::Fact(y)) => Ok(Value::Bool(op(x, y))),
302 (Truth::Probability(p), Truth::Probability(q)) => computed_prob(
303 if and { p * q } else { p + (1.0 - p) * q },
304 if and { "and" } else { "or" },
305 ),
306 (Truth::Probability(p), Truth::Fact(b)) | (Truth::Fact(b), Truth::Probability(p)) => {
307 Ok(Value::Prob(if b == and {
308 p
309 } else if b {
310 1.0
311 } else {
312 0.0
313 }))
314 }
315 (Truth::Fact(x), Truth::Uncertain(d)) | (Truth::Uncertain(d), Truth::Fact(x)) => {
316 lift1(&Value::Dist(d), budget, |v, _| match v {
317 Value::Bool(y) => Ok(Value::Bool(op(x, *y))),
318 _ => unreachable!("checked by `truth`"),
319 })
320 }
321 (a, b) => {
322 let law = |v| match v {
323 Truth::Uncertain(d) => Value::Dist(d),
324 Truth::Probability(p) => Dist::bernoulli(p).into_value(),
325 Truth::Fact(b) => Value::Bool(b),
326 Truth::Analytic(_) => unreachable!("handled above"),
327 };
328 lift2(&law(a), &law(b), budget, |x, y, _| match (x, y) {
329 (Value::Bool(x), Value::Bool(y)) => Ok(Value::Bool(op(*x, *y))),
330 _ => unreachable!("checked by `truth`"),
331 })
332 }
333 }
334}
335
336pub fn boolean_law(v: &Value) -> OpResult<Value> {
338 match truth(v, "observe ~")? {
339 Truth::Analytic(e) => Ok(Value::Event(e)),
340 Truth::Fact(b) => Ok(Dist::bernoulli(if b { 1.0 } else { 0.0 }).into_value()),
341 Truth::Probability(p) => Ok(Dist::bernoulli(p).into_value()),
342 Truth::Uncertain(d) => Ok(Value::Dist(d)),
343 }
344}
345
346pub fn condition(v: &Value) -> OpResult<Condition> {
347 if matches!(v, Value::Analytic(_)) {
348 return Err(crate::analytic::unsupported(
349 "using a continuous outcome as a probability condition",
350 ));
351 }
352 if matches!(v, Value::Float(_) | Value::Int(_)) {
353 let p = to_prob(v)?;
354 return Ok(Condition {
355 yes: p,
356 no: 1.0 - p,
357 missing: 0.0,
358 });
359 }
360 match truth(v, "condition")? {
361 Truth::Analytic(e) => Ok(Condition {
362 yes: e.probability(),
363 no: 1.0 - e.probability(),
364 missing: 0.0,
365 }),
366 Truth::Fact(b) => Ok(Condition {
367 yes: if b { 1.0 } else { 0.0 },
368 no: if b { 0.0 } else { 1.0 },
369 missing: 0.0,
370 }),
371 Truth::Probability(p) => Ok(Condition {
372 yes: p,
373 no: 1.0 - p,
374 missing: 0.0,
375 }),
376 Truth::Uncertain(d) => {
377 let (yes, no) = d.truth().expect("checked by truth");
378 Ok(Condition {
379 yes,
380 no,
381 missing: d.missing,
382 })
383 }
384 }
385}
386
387pub fn unary(op: UnOp, v: &Value, budget: &mut Budget) -> OpResult<Value> {
390 match op {
391 UnOp::Neg => lift1(v, budget, |x, budget| match x {
392 Value::Int(i) => {
393 budget.integer_work(i, &Integer::ZERO, false)?;
394 {
395 let n = i.negated();
396 budget.integer_allocation(n.bits(), 1)?;
397 Ok(Value::Int(n))
398 }
399 }
400 Value::Float(f) | Value::Prob(f) => Ok(Value::Float(-f)),
401 Value::Complex(z) => Ok(Value::Complex(z.negated())),
402 Value::Analytic(a) => {
403 let mut a = (**a).clone();
404 a.scale = -a.scale;
405 a.offset = -a.offset;
406 a.value()
407 }
408 other => Err(OpError::new(format!("can't negate {}", article(&other.kind())))),
409 }),
410 UnOp::Not => not(v, budget),
411 UnOp::Typeof => unreachable!("typeof is lowered to an interpreter intrinsic"),
412 }
413}
414
415pub fn binary(op: BinOp, a: &Value, b: &Value, budget: &mut Budget) -> OpResult<Value> {
416 match op {
417 BinOp::And | BinOp::Or => unreachable!("`and` and `or` are evaluated lazily by the interpreter"),
418 BinOp::To => unreachable!("`to` is lowered to a built-in"),
419 BinOp::Range | BinOp::RangeExcl => range(op, a, b, budget),
420 BinOp::In => lift2(a, b, budget, |x, coll, budget| {
421 contains(coll, x, budget).map(Value::Bool)
422 }),
423 BinOp::NotIn => lift2(a, b, budget, |x, coll, budget| {
424 contains(coll, x, budget).map(|c| Value::Bool(!c))
425 }),
426 _ => lift2(a, b, budget, |x, y, budget| binary_plain(op, x, y, budget)),
427 }
428}
429
430fn binary_plain(op: BinOp, a: &Value, b: &Value, budget: &mut Budget) -> OpResult<Value> {
431 if matches!(a, Value::Analytic(_)) || matches!(b, Value::Analytic(_)) {
432 return crate::analytic::binary(op, a, b);
433 }
434 if matches!(
435 op,
436 BinOp::Eq | BinOp::Ne | BinOp::Lt | BinOp::Le | BinOp::Gt | BinOp::Ge
437 ) && (crate::analytic::contains(a) || crate::analytic::contains(b))
438 {
439 if matches!(op, BinOp::Eq | BinOp::Ne) {
440 if let (Value::Event(x), Value::Bool(b)) | (Value::Bool(b), Value::Event(x)) = (a, b) {
441 return if *b == (op == BinOp::Eq) {
442 Ok(Value::Event(x.clone()))
443 } else {
444 not(&Value::Event(x.clone()), budget)
445 };
446 }
447 }
448 if let (Value::Event(x), Value::Event(y)) = (a, b) {
449 if x.draw.id == y.draw.id && matches!(op, BinOp::Eq | BinOp::Ne) {
450 let both = x.yes.intersect(&y.yes);
451 let neither = x.yes.complement().intersect(&y.yes.complement());
452 let equal = both.complement().intersect(&neither.complement()).complement();
453 return Ok(crate::analytic::Event {
454 draw: x.draw.clone(),
455 yes: if op == BinOp::Eq { equal } else { equal.complement() },
456 }
457 .value());
458 }
459 }
460 return Err(crate::analytic::unsupported(
461 "comparing aggregate or boolean analytic outcomes",
462 ));
463 }
464 if let Some(v) = continuous_binary(op, a, b)? {
465 return Ok(v);
466 }
467 if matches!(
468 op,
469 BinOp::Eq | BinOp::Ne | BinOp::Lt | BinOp::Le | BinOp::Gt | BinOp::Ge
470 ) {
471 for v in [a, b] {
472 if let Value::Int(n) = v {
473 budget.integer_work(n, &Integer::ONE, false)?;
474 }
475 if let Value::Str(s) = v {
476 budget.string_work(s)?;
477 }
478 }
479 }
480 match op {
481 BinOp::Eq => Ok(Value::Bool(equals(a, b))),
482 BinOp::Ne => Ok(Value::Bool(!equals(a, b))),
483 BinOp::Lt | BinOp::Le | BinOp::Gt | BinOp::Ge => {
484 let ord = compare(a, b)?;
485 Ok(Value::Bool(match op {
486 BinOp::Lt => ord.is_lt(),
487 BinOp::Le => ord.is_le(),
488 BinOp::Gt => ord.is_gt(),
489 _ => ord.is_ge(),
490 }))
491 }
492 BinOp::Add => add(a, b, budget),
493 BinOp::Sub => sub(a, b, budget),
494 _ => arith(op, a, b, budget),
495 }
496}
497
498fn add(a: &Value, b: &Value, budget: &mut Budget) -> OpResult<Value> {
499 match (a, b) {
500 (Value::Str(x), Value::Str(y)) => {
501 let size = x
502 .len()
503 .checked_add(y.len())
504 .ok_or_else(|| OpError::limit("string size overflow"))?;
505 budget.string_size(size)?;
506 let mut out = String::new();
507 crate::text::push(&mut out, x, budget)?;
508 crate::text::push(&mut out, y, budget)?;
509 Ok(Value::str(&out))
510 }
511 (Value::List(x), Value::List(y)) => {
512 budget.collection(x.len() as u128 + y.len() as u128)?;
513 budget.work(x.len() as u64 + y.len() as u64)?;
514 let mut items = x.to_vec();
515 items.extend(y.iter().cloned());
516 Ok(Value::list(items))
517 }
518 (Value::Date(d), n @ (Value::Int(_) | Value::Float(_)))
519 | (n @ (Value::Int(_) | Value::Float(_)), Value::Date(d)) => {
520 let n = integer(n, "date offset", budget)?;
521 date_plus(*d, n.to_i64().ok_or_else(|| OpError::new("date out of range"))?)
522 }
523 (Value::Str(_), _) | (_, Value::Str(_)) => {
524 Err(
525 OpError::new(format!("can't add {} and {}", article(&a.kind()), article(&b.kind())))
526 .help("to build text, use interpolation: \"total: {x}\""),
527 )
528 }
529 _ => arith(BinOp::Add, a, b, budget),
530 }
531}
532
533fn sub(a: &Value, b: &Value, budget: &mut Budget) -> OpResult<Value> {
534 match (a, b) {
535 (Value::Date(x), Value::Date(y)) => Ok(Value::Int((*x as i64 - *y as i64).into())),
536 (Value::Date(d), n @ (Value::Int(_) | Value::Float(_))) => {
537 match integer(n, "date offset", budget)?.negated().to_i64() {
538 Some(m) => date_plus(*d, m),
539 None => Err(OpError::new("date out of range")),
540 }
541 }
542 _ => arith(BinOp::Sub, a, b, budget),
543 }
544}
545
546fn date_plus(d: i32, n: i64) -> OpResult<Value> {
547 crate::dates::add_days(d, n)
548 .map(Value::Date)
549 .ok_or_else(|| OpError::new("date out of range"))
550}
551
552fn number(v: &Value) -> Option<f64> {
554 match v {
555 Value::Bool(_) => None,
556 _ => v.as_f64(),
557 }
558}
559
560fn arith(op: BinOp, a: &Value, b: &Value, budget: &mut Budget) -> OpResult<Value> {
561 let bad = || {
562 OpError::new(format!(
563 "can't use `{}` with {} and {}",
564 op.symbol(),
565 article(&a.kind()),
566 article(&b.kind())
567 ))
568 };
569 if let (BinOp::Pow, Value::Complex(z), Value::Int(n)) = (op, a, b) {
570 budget.integer_work(n, &Integer::ONE, false)?;
571 budget.work(n.bits())?;
572 return Ok(Value::Complex(z.pow_integer(n)?));
573 }
574 if matches!(a, Value::Complex(_)) || matches!(b, Value::Complex(_)) {
575 let (Some(x), Some(y)) = (a.as_complex(), b.as_complex()) else {
576 return Err(bad());
577 };
578 let z = match op {
579 BinOp::Add => x.plus(y),
580 BinOp::Sub => x.minus(y),
581 BinOp::Mul => x.times(y),
582 BinOp::Div => x.divided_by(y),
583 BinOp::Pow => match b {
584 Value::Int(n) => {
585 budget.integer_work(n, &Integer::ONE, false)?;
586 budget.work(n.bits())?;
587 x.pow_integer(n)
588 }
589 _ => Err(OpError::new("a complex power needs an int exponent")),
590 },
591 _ => Err(bad()),
592 }?;
593 return Ok(Value::Complex(z));
594 }
595 if let (Value::Int(x), Value::Int(y)) = (a, b) {
596 budget.integer_work(x, y, matches!(op, BinOp::Mul | BinOp::Div | BinOp::IntDiv | BinOp::Mod))?;
597 let n = match op {
598 BinOp::Add => x.add(y)?,
599 BinOp::Sub => x.sub(y)?,
600 BinOp::Mul => {
601 if !x.is_zero() && !y.is_zero() {
602 budget.integer_bits((x.bits() + y.bits()).saturating_sub(1))?;
603 }
604 x.mul(y)?
605 }
606 BinOp::Div => {
607 if y.is_zero() {
608 return Err(division_by_zero());
609 }
610 return x
611 .ratio(y)
612 .map(Value::Float)
613 .ok_or_else(|| OpError::new("division result is too large for a finite float"));
614 }
615 BinOp::IntDiv => x.div_mod(y)?.0,
616 BinOp::Mod => x.div_mod(y)?.1,
617 BinOp::Pow => return int_power(x, y, budget),
618 _ => return Err(bad()),
619 };
620 budget.integer_allocation(n.bits(), 1)?;
621 return Ok(Value::Int(n));
622 }
623 let (Some(x), Some(y)) = (number(a), number(b)) else {
624 return Err(bad());
625 };
626 let v = match op {
627 BinOp::Add => x + y,
628 BinOp::Sub => x - y,
629 BinOp::Mul => x * y,
630 BinOp::Div => {
631 if y == 0.0 {
632 return Err(division_by_zero());
633 }
634 x / y
635 }
636 BinOp::IntDiv => {
637 if y == 0.0 {
638 return Err(division_by_zero());
639 }
640 let q = (x / y).floor();
641 let n = Integer::from_f64(q).ok_or_else(|| OpError::new("integer division result is not finite"))?;
642 budget.integer_allocation(n.bits(), 1)?;
643 return Ok(Value::Int(n));
644 }
645 BinOp::Mod => {
646 if y == 0.0 {
647 return Err(division_by_zero());
648 }
649 x - (x / y).floor() * y
650 }
651 BinOp::Pow => libm::pow(x, y),
652 _ => return Err(bad()),
653 };
654 if !v.is_finite() {
655 return Err(OpError::new(format!(
656 "`{}` gave a result that isn't a finite number",
657 op.symbol()
658 )));
659 }
660 Ok(Value::Float(v))
661}
662
663fn continuous_binary(op: BinOp, a: &Value, b: &Value) -> OpResult<Option<Value>> {
667 let (family, other, flipped) = match (a, b) {
668 (Value::Continuous(f), other) => (f, other, false),
669 (other, Value::Continuous(f)) => (f, other, true),
670 _ => return Ok(None),
671 };
672 let needs_value = || {
673 OpError::new(format!(
674 "`{}` needs a value, not a {} distribution",
675 op.symbol(),
676 family.name()
677 ))
678 .help("draw a value first, like `let x ~ normal(0, 1)`; comparing a distribution with a number works too")
679 };
680 let x = match number(other) {
681 Some(x) if !x.is_nan() => x,
682 _ => return Err(needs_value()),
683 };
684 let below = family.cdf(x);
686 let yes = match (op, flipped) {
687 (BinOp::Lt | BinOp::Le, false) | (BinOp::Gt | BinOp::Ge, true) => below,
688 (BinOp::Gt | BinOp::Ge, false) | (BinOp::Lt | BinOp::Le, true) => 1.0 - below,
689 (BinOp::Eq, _) => 0.0,
690 (BinOp::Ne, _) => 1.0,
691 _ => return Err(needs_value()),
692 };
693 Ok(Some(Dist::bernoulli(yes).into_value()))
694}
695
696fn division_by_zero() -> OpError {
697 OpError::new("division by zero")
698}
699
700fn int_power(base: &Integer, exponent: &Integer, budget: &mut Budget) -> OpResult<Value> {
701 if exponent.is_zero() {
702 return Ok(Value::Int(Integer::ONE));
703 }
704 if base.is_zero() {
705 return if exponent.is_negative() {
706 Err(division_by_zero())
707 } else {
708 Ok(Value::Int(Integer::ZERO))
709 };
710 }
711 if exponent.is_negative() {
712 let reciprocal = Integer::ONE.ratio(&base.abs()).unwrap_or(0.0);
713 let power = exponent.abs().to_f64().unwrap_or(f64::INFINITY);
714 let result = libm::pow(reciprocal, power);
715 return Ok(Value::Float(if base.is_negative() && exponent.is_odd() {
716 -result
717 } else {
718 result
719 }));
720 }
721 if *base == 1 || *base == -1 {
722 return Ok(Value::Int(if *base == -1 && exponent.is_odd() {
723 (-1).into()
724 } else {
725 Integer::ONE
726 }));
727 }
728 let mut n = exponent
729 .to_u64()
730 .ok_or_else(|| OpError::limit("integer power exceeds the integer size limit"))?;
731 budget.integer_bits(base.bits().saturating_sub(1).saturating_mul(n).saturating_add(1))?;
732 let mut result = Integer::ONE;
733 let mut b = base.clone();
734 while n != 0 {
735 if n & 1 != 0 {
736 budget.integer_work(&result, &b, true)?;
737 result = result.mul(&b)?;
738 budget.integer_bits(result.bits())?;
739 }
740 n >>= 1;
741 if n != 0 {
742 budget.integer_work(&b, &b, true)?;
743 b = b.mul(&b)?;
744 budget.integer_bits(b.bits())?;
745 }
746 }
747 budget.integer_allocation(result.bits(), 1)?;
748 Ok(Value::Int(result))
749}
750
751fn numeric_compare(a: &Value, b: &Value) -> Option<std::cmp::Ordering> {
752 match (a, b) {
753 (Value::Int(a), Value::Int(b)) => Some(a.cmp(b)),
754 (Value::Int(a), Value::Float(b) | Value::Prob(b)) => a.cmp_f64(*b),
755 (Value::Float(a) | Value::Prob(a), Value::Int(b)) => b.cmp_f64(*a).map(std::cmp::Ordering::reverse),
756 (Value::Float(a) | Value::Prob(a), Value::Float(b) | Value::Prob(b)) => a.partial_cmp(b),
757 _ => None,
758 }
759}
760
761pub fn equals(a: &Value, b: &Value) -> bool {
763 if let (Value::Complex(x), Value::Complex(y)) = (a, b) {
764 return x == y;
765 }
766 if let (Value::Complex(z), real) | (real, Value::Complex(z)) = (a, b) {
767 if z.im() != 0.0 {
768 return false;
769 }
770 return match real {
771 Value::Int(n) => {
773 let x = z.re();
774 n.cmp_f64(x).is_some_and(|c| c.is_eq())
775 }
776 Value::Float(x) | Value::Prob(x) => z.re() == *x,
777 _ => false,
778 };
779 }
780 if let Some(c) = numeric_compare(a, b) {
781 return c.is_eq();
782 }
783 if matches!(a, Value::Float(x) | Value::Prob(x) if x.is_nan())
784 || matches!(b, Value::Float(x) | Value::Prob(x) if x.is_nan())
785 {
786 return false;
787 }
788 a == b
789}
790
791pub fn compare(a: &Value, b: &Value) -> OpResult<std::cmp::Ordering> {
793 if matches!(a, Value::Complex(_)) || matches!(b, Value::Complex(_)) {
794 return Err(OpError::new("complex values have no ordering").help("compare `abs(z)`, `real(z)` or `imag(z)`"));
795 }
796 if let Some(c) = numeric_compare(a, b) {
797 return Ok(c);
798 }
799 if matches!(a, Value::Float(x) | Value::Prob(x) if x.is_nan())
800 || matches!(b, Value::Float(x) | Value::Prob(x) if x.is_nan())
801 {
802 return Err(OpError::new("can't order NaN"));
803 }
804 match (a, b) {
805 (Value::Str(x), Value::Str(y)) => Ok(x.cmp(y)),
806 (Value::Date(x), Value::Date(y)) => Ok(x.cmp(y)),
807 (Value::Enum(x), Value::Enum(y)) if x.ty == y.ty => Ok(x.variant.cmp(&y.variant)),
808 (Value::List(x), Value::List(y)) => {
809 for (p, q) in x.iter().zip(y.iter()) {
810 let c = compare(p, q)?;
811 if c.is_ne() {
812 return Ok(c);
813 }
814 }
815 Ok(x.len().cmp(&y.len()))
816 }
817 _ => Err(OpError::new(format!(
818 "can't compare {} with {}",
819 article(&a.kind()),
820 article(&b.kind())
821 ))),
822 }
823}
824
825pub fn contains(coll: &Value, item: &Value, budget: &mut Budget) -> OpResult<bool> {
827 if crate::analytic::contains(coll) || crate::analytic::contains(item) {
828 return Err(crate::analytic::unsupported("membership involving analytic outcomes"));
829 }
830 if let Value::Str(s) = coll {
831 budget.string_work(s)?;
832 if let Value::Str(s) = item {
833 budget.string_work(s)?;
834 }
835 }
836 Ok(match coll {
837 Value::List(items) => items.iter().any(|x| equals(x, item)),
838 Value::Map(m) => m.contains_key(item),
839 Value::Bag(b) => b.get(item).is_some_and(|n| *n > 0),
840 Value::Range(lo, hi) => match item {
841 Value::Int(n) => n >= lo && n <= hi,
842 Value::Float(x) | Value::Prob(x) if x.is_finite() && x.fract() == 0.0 => {
843 lo.cmp_f64(*x).is_some_and(|c| !c.is_gt()) && hi.cmp_f64(*x).is_some_and(|c| !c.is_lt())
844 }
845 _ => false,
846 },
847 Value::Str(s) => match item {
848 Value::Str(sub) => s.contains(&**sub),
849 _ => {
850 return Err(OpError::new(format!(
851 "can't look for {} in a string",
852 article(&item.kind())
853 )));
854 }
855 },
856 other => return Err(OpError::new(format!("can't look inside {}", article(&other.kind())))),
857 })
858}
859
860fn range(op: BinOp, a: &Value, b: &Value, budget: &mut Budget) -> OpResult<Value> {
861 if a.is_uncertain() || b.is_uncertain() {
862 return Err(OpError::new("a range needs plain whole numbers, not distributions")
863 .help("draw a value first, like `let n ~ d6`"));
864 }
865 let lo = integer(a, "range bound", budget)?;
866 let hi = integer(b, "range bound", budget)?;
867 let hi = if op == BinOp::RangeExcl {
868 hi.sub(&Integer::ONE)?
869 } else {
870 hi.into_owned()
871 };
872 budget.integer_work(&lo, &hi, false)?;
873 if op == BinOp::RangeExcl {
874 budget.integer_allocation(hi.bits(), 1)?;
875 }
876 Ok(Value::Range(lo.into_owned(), hi))
877}
878
879pub fn range_len(lo: &Integer, hi: &Integer) -> OpResult<Integer> {
881 if hi < lo {
882 Ok(Integer::ZERO)
883 } else {
884 Ok(hi.sub(lo)?.add(&Integer::ONE)?)
885 }
886}
887
888pub fn range_count(lo: &Integer, hi: &Integer) -> OpResult<u128> {
890 range_len(lo, hi)?
891 .to_u128()
892 .ok_or_else(|| OpError::limit("range has too many elements to materialize"))
893}
894
895pub fn field(v: &Value, name: &str, budget: &mut Budget) -> OpResult<Value> {
898 lift1(v, budget, |x, _| match x {
899 Value::Record(r) => r.get(name).cloned().ok_or_else(|| {
900 let known: Vec<String> = r.fields.iter().map(|(n, _)| format!("`{n}`")).collect();
901 OpError::new(format!("{} has no field `{name}`", article(&x.kind())))
902 .help(format!("its fields are {}", known.join(", ")))
903 }),
904 other => Err(OpError::new(format!(
905 "can't read the field `{name}` of {}",
906 article(&other.kind())
907 ))),
908 })
909}
910
911pub fn index(coll: &Value, i: &Value, budget: &mut Budget) -> OpResult<Value> {
912 lift2(coll, i, budget, |c, i, budget| {
913 let v = index_plain(c, i, budget)?;
914 if let (Value::Range(lo, hi), Value::Int(n)) = (c, &v) {
915 budget.integer_work(lo, hi, false)?;
916 budget.integer_allocation(n.bits(), 1)?;
917 }
918 Ok(v)
919 })
920}
921
922pub fn index_plain(coll: &Value, i: &Value, budget: &mut Budget) -> OpResult<Value> {
923 match coll {
924 Value::List(items) => {
925 let k = as_index(i, items.len() as u128, budget)?;
926 Ok(items[k as usize].clone())
927 }
928 Value::Range(lo, hi) => {
929 let k = integer(i, "index", budget)?;
930 if k.is_negative() {
931 return Err(OpError::new("index out of range"));
932 }
933 let n = lo.add(&k)?;
934 if &n > hi {
935 return Err(OpError::new("index out of range"));
936 }
937 Ok(Value::Int(n))
938 }
939 Value::Str(s) => {
940 budget.string_work(s)?;
941 let k = as_index(i, s.chars().count() as u128, budget)?;
942 let c = s.chars().nth(k as usize).expect("checked scalar index");
943 crate::text::value(c.encode_utf8(&mut [0; 4]), budget)
944 }
945 Value::Map(m) => m.get(i).cloned().ok_or_else(|| {
946 OpError::new(format!("the key {i:?} isn't in the map"))
947 .help("use `get(key, default)` for keys that may be missing")
948 }),
949 other => Err(OpError::new(format!("can't index {}", article(&other.kind())))),
950 }
951}
952
953pub fn as_index(i: &Value, len: u128, budget: &mut Budget) -> OpResult<u128> {
954 let k = integer(i, "index", budget)?.to_u128();
955 k.filter(|k| *k < len).ok_or_else(|| {
956 OpError::new(format!("index {i} is out of range for a length of {len}")).help("indices start at 0")
957 })
958}
959
960pub fn make_record(ty: Option<Arc<str>>, mut fields: Vec<(Arc<str>, Value)>) -> Value {
961 fields.sort_by(|a, b| a.0.cmp(&b.0));
962 Value::record(Record { ty, fields })
963}
964
965pub fn with_fields(base: &Value, updates: &[(Arc<str>, Value)]) -> OpResult<Value> {
966 let Value::Record(r) = base else {
967 return Err(OpError::new(format!(
968 "`with` needs a record, found {}",
969 article(&base.kind())
970 )));
971 };
972 let mut r = Record::clone(r);
973 for (name, v) in updates {
974 match r.get_mut(name) {
975 Some(slot) => *slot = v.clone(),
976 None => {
977 return Err(OpError::new(format!("{} has no field `{name}`", article(&base.kind()))));
978 }
979 }
980 }
981 Ok(Value::record(r))
982}
983
984pub fn enum_value(ty: u32, variant: u32, name: &str) -> Value {
985 Value::Enum(Arc::new(EnumValue {
986 ty,
987 variant,
988 name: Arc::from(name),
989 }))
990}
991
992pub fn is_certain(v: &Value, truth: bool) -> bool {
994 matches!(v, Value::Bool(b) if *b == truth)
995}
996
997#[cfg(test)]
998mod probability_tests {
999 use super::*;
1000
1001 #[test]
1002 fn computed_probabilities_only_correct_boundary_roundoff() {
1003 assert_eq!(computed_prob(1.0 + f64::EPSILON, "test").unwrap(), Value::Prob(1.0));
1004 assert_eq!(computed_prob(-f64::EPSILON, "test").unwrap(), Value::Prob(0.0));
1005 assert_eq!(computed_prob(1e-300, "test").unwrap(), Value::Prob(1e-300));
1006 for p in [1.001, -0.001, f64::INFINITY, f64::NEG_INFINITY, f64::NAN] {
1007 assert!(computed_prob(p, "test").is_err());
1008 }
1009 }
1010}