1use crate::analytic::{self, Analytic};
10use crate::builtins;
11use crate::chain::{Chain, Solution};
12use crate::conjugate::{self, Seen};
13use crate::continuous::{Family, Rng};
14use crate::dist::{Budget, Counts, Dist};
15use crate::error::{Fault, OpError, OpResult, Result, RuntimeError};
16use crate::failure::Failures;
17use crate::ops::{self, Truth};
18use crate::report::Sink;
19use crate::value::{Closure, Delayed, Value, fmt_prob};
20use crate::weight::Weight;
21use crate::world::Returned;
22use crate::world::{Flow, World, clear, clear_dead, live_slots, merge, merge_values, state_hash, total_weight};
23use probl_sema::builtins::Lifting;
24use probl_sema::conjugate::{Conjugacy, Likelihood, Update};
25use probl_sema::effects::EvidenceOrder;
26use probl_sema::ir::*;
27use probl_sema::{Builtin, Liveness};
28use probl_syntax::Span;
29use probl_syntax::ast::BinOp;
30use rustc_hash::{FxHashMap, FxHashSet};
31use std::collections::BTreeMap;
32use std::sync::Arc;
33use std::sync::atomic::{AtomicBool, Ordering};
34
35#[derive(Clone, Debug, Default)]
36pub struct Stats {
37 pub peak_worlds: usize,
39 pub world_steps: u64,
41 pub calls: u64,
42 pub memo_hits: u64,
43 pub solved_loops: u64,
45 pub chain_states: u64,
46 pub solved_calls: u64,
49 pub call_rounds: u64,
50 pub updates: Vec<Updates>,
53}
54
55#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
58pub struct Updates {
59 pub delayed: u64,
61 pub exact: u64,
63 pub drawn: u64,
65}
66
67impl Stats {
68 pub fn absorb(&mut self, other: &Stats) {
70 self.peak_worlds = self.peak_worlds.max(other.peak_worlds);
71 self.world_steps += other.world_steps;
72 self.calls += other.calls;
73 self.memo_hits += other.memo_hits;
74 self.solved_loops += other.solved_loops;
75 self.chain_states += other.chain_states;
76 self.solved_calls += other.solved_calls;
77 self.call_rounds += other.call_rounds;
78 if self.updates.len() < other.updates.len() {
79 self.updates.resize(other.updates.len(), Updates::default());
80 }
81 for (mine, theirs) in self.updates.iter_mut().zip(&other.updates) {
82 mine.delayed += theirs.delayed;
83 mine.exact += theirs.exact;
84 mine.drawn += theirs.drawn;
85 }
86 }
87}
88
89#[derive(Clone, Debug)]
91pub struct Config {
92 pub today: Option<i32>,
93 pub epsilon: f64,
94 pub merging: bool,
95 pub memoizing: bool,
96 pub max_worlds: usize,
97 pub max_iterations: u64,
98 pub max_call_depth: usize,
99 pub max_cached_calls: usize,
100 pub max_output: usize,
101 pub budget: Budget,
102 pub cancel: Option<Arc<AtomicBool>>,
103 pub sample_seed: Option<u64>,
106 pub conjugate: bool,
109 pub solving: bool,
111 pub max_chain_states: usize,
113 pub partial: bool,
116}
117
118const MAX_ELIMINATION: u64 = 200_000_000;
121
122pub const BATCH: u64 = 1000;
126
127#[derive(Debug)]
130pub struct Batch {
131 pub sinks: Vec<Sink>,
132 pub totals: SampleTotals,
133 pub observed: bool,
134 pub densities: bool,
135 pub unresolved: Weight,
136 pub last_ruling_out: Option<Span>,
137 pub stats: Stats,
138 pub failures: Failures,
139}
140
141pub type Printed = Vec<(Span, String)>;
143
144#[derive(Clone, Copy, Debug)]
146pub struct SampleTotals {
147 pub weight: Weight,
148 pub squares: Weight,
150}
151
152#[derive(Clone, Debug)]
154pub struct CallResult {
155 pub outcomes: Vec<Returned>,
156 pub unresolved: Weight,
158 pub lost: Weight,
160 pub observed: bool,
162 pub pending: Vec<(usize, Weight)>,
166 pub failures: Failures,
168}
169
170impl CallResult {
171 fn waiting(depth: usize) -> CallResult {
174 CallResult {
175 outcomes: Vec::new(),
176 unresolved: Weight::ONE,
177 lost: Weight::ZERO,
178 observed: false,
179 pending: vec![(depth, Weight::ONE)],
180 failures: Failures::default(),
181 }
182 }
183
184 fn take_pending(&mut self, depth: usize) -> Weight {
186 match self.pending.iter().position(|&(d, _)| d == depth) {
187 Some(i) => self.pending.remove(i).1,
188 None => Weight::ZERO,
189 }
190 }
191}
192
193struct Approx {
196 result: Arc<CallResult>,
197 head: usize,
201 frame: u64,
202 round: u64,
203}
204
205fn add_pending(into: &mut Vec<(usize, Weight)>, from: &[(usize, Weight)], weight: Weight) {
207 for &(depth, w) in from {
208 match into.iter_mut().find(|(d, _)| *d == depth) {
209 Some((_, total)) => *total += weight * w,
210 None => into.push((depth, weight * w)),
211 }
212 }
213}
214
215type CallKey = (FnId, Vec<Value>);
216
217pub struct Engine<'p> {
218 prog: &'p Program,
219 live: &'p Liveness,
220 conj: &'p Conjugacy,
222 config: Config,
223 budget: Budget,
224 memo: FxHashMap<CallKey, Arc<CallResult>>,
225 active: FxHashMap<CallKey, usize>,
227 approx: FxHashMap<CallKey, Approx>,
230 frames_at: Vec<u64>,
233 rounds_at: Vec<u64>,
234 last_number: u64,
235 pending: Vec<(usize, Weight)>,
237 depth: usize,
238 pub unresolved: Weight,
240 lost: Weight,
243 solvable: Vec<bool>,
246 too_large: FxHashSet<StmtId>,
249 pub observed: bool,
251 pub densities: bool,
254 pub last_ruling_out: Option<Span>,
256 pub sinks: Vec<Sink>,
257 dice: FxHashMap<(u32, u32), Value>,
258 pools: FxHashMap<(u32, Value), Value>,
259 record_names: Vec<Arc<str>>,
260 enum_values: Vec<Vec<Value>>,
261 print: &'p mut (dyn FnMut(&str) + Send),
262 inputs: &'p [Value],
264 printed: usize,
266 lines: Option<Printed>,
269 pub stats: Stats,
270 sampler: Option<Rng>,
273 nested: usize,
275 next_latent: u64,
276 callback: Option<&'static str>,
279 pub failures: Failures,
281 handlers: Vec<(usize, &'p [Catch])>,
284 faulted: Vec<(World, RuntimeError)>,
287 evidence: EvidenceOrder,
290}
291
292#[derive(Clone, Copy)]
294struct At {
295 f: FnId,
296 stmt: StmtId,
297 weight: Weight,
298 run: u32,
299}
300
301impl At {
302 fn new(f: FnId, stmt: &Stmt, w: &World) -> At {
303 At {
304 f,
305 stmt: stmt.id,
306 weight: w.weight,
307 run: w.run,
308 }
309 }
310}
311
312macro_rules! each {
318 ($self:ident, $at:expr, world $w:expr, $e:expr $(, $undo:block)?) => {
319 match $e {
320 Ok(v) => v,
321 Err(e) => {
322 let saved = $self.snapshot(&$w);
323 $self.fail(e, $at, saved)?;
324 $($undo)?
325 continue;
326 }
327 }
328 };
329 ($self:ident, $at:expr, saved $saved:ident, $e:expr $(, $undo:block)?) => {
330 match $e {
331 Ok(v) => v,
332 Err(e) => {
333 $self.fail(e, $at, $saved.take())?;
334 $($undo)?
335 continue;
336 }
337 }
338 };
339}
340
341impl<'p> Engine<'p> {
342 pub fn new(
343 prog: &'p Program,
344 live: &'p Liveness,
345 conj: &'p Conjugacy,
346 config: Config,
347 inputs: &'p [Value],
348 print: &'p mut (dyn FnMut(&str) + Send),
349 ) -> Engine<'p> {
350 let sampler = config.sample_seed.map(Rng::new);
351 Engine {
352 prog,
353 live,
354 conj,
355 budget: config.budget.clone(),
356 config,
357 memo: FxHashMap::default(),
358 active: FxHashMap::default(),
359 approx: FxHashMap::default(),
360 frames_at: Vec::new(),
361 rounds_at: Vec::new(),
362 last_number: 0,
363 pending: Vec::new(),
364 depth: 0,
365 unresolved: Weight::ZERO,
366 lost: Weight::ZERO,
367 solvable: probl_sema::effects::solvable_loops(prog),
368 too_large: FxHashSet::default(),
369 observed: false,
370 densities: false,
371 last_ruling_out: None,
372 sinks: vec![Sink::default(); prog.reports.len()],
373 dice: FxHashMap::default(),
374 pools: FxHashMap::default(),
375 record_names: prog.records.iter().map(|r| Arc::from(r.name.as_str())).collect(),
376 enum_values: prog
377 .enums
378 .iter()
379 .enumerate()
380 .map(|(t, e)| {
381 e.variants
382 .iter()
383 .enumerate()
384 .map(|(v, name)| ops::enum_value(t as u32, v as u32, name))
385 .collect()
386 })
387 .collect(),
388 print,
389 inputs,
390 printed: 0,
391 lines: None,
392 stats: Stats::default(),
393 sampler,
394 nested: 0,
395 next_latent: 0,
396 callback: None,
397 failures: Failures::default(),
398 evidence: probl_sema::effects::evidence_order(prog),
399 handlers: Vec::new(),
400 faulted: Vec::new(),
401 }
402 }
403
404 pub fn run_main(&mut self) -> Result<Weight> {
406 let prog = self.prog;
407 let main = prog.main();
408 let world = World {
409 slots: vec![Value::Dead; main.n_slots()],
410 constraints: Default::default(),
411 inherited: Default::default(),
412 weight: Weight::ONE,
413 run: 0,
414 };
415 let flow = self.exec_block(MAIN, &main.body, vec![world])?;
416 uncaught(&flow)?;
417 Ok(total_weight(&flow.next))
418 }
419
420 pub fn run_batch(&mut self, rng: Rng, first: u64, n: u64, printed: usize) -> (Result<Batch>, Printed) {
426 self.sampler = Some(rng);
427 self.sinks = vec![Sink::default(); self.prog.reports.len()];
428 self.observed = false;
429 self.densities = false;
430 self.unresolved = Weight::ZERO;
431 self.last_ruling_out = None;
432 self.stats = Stats::default();
433 self.failures = Failures::default();
434 self.printed = printed;
435 self.lines = Some(Vec::new());
436 self.budget = self.config.budget.clone();
437 let result = self.sample_runs(first, n);
438 self.budget.give_back();
439 (result, self.lines.take().unwrap_or_default())
440 }
441
442 fn sample_runs(&mut self, first: u64, n: u64) -> Result<Batch> {
443 let main = self.prog.main();
444 let worlds = (first..first + n)
445 .map(|run| World {
446 slots: vec![Value::Dead; main.n_slots()],
447 constraints: Default::default(),
448 inherited: Default::default(),
449 weight: Weight::ONE,
450 run: run as u32,
451 })
452 .collect();
453 let flow = self.exec_block(MAIN, &main.body, worlds)?;
454 uncaught(&flow)?;
455 let mut totals = SampleTotals {
456 weight: Weight::ZERO,
457 squares: Weight::ZERO,
458 };
459 for w in &flow.next {
460 totals.weight += w.weight;
461 totals.squares += w.weight * w.weight;
462 }
463 for sink in &mut self.sinks {
464 sink.end_batch();
465 }
466 Ok(Batch {
467 sinks: std::mem::take(&mut self.sinks),
468 totals,
469 observed: self.observed,
470 densities: self.densities,
471 unresolved: self.unresolved,
472 last_ruling_out: self.last_ruling_out,
473 stats: std::mem::take(&mut self.stats),
474 failures: std::mem::take(&mut self.failures),
475 })
476 }
477
478 fn sample(&mut self, v: &Value) -> Value {
482 let rng = self.sampler.as_mut().expect("only called when sampling");
483 match v {
484 Value::Dist(d) if !d.outcomes.is_empty() => {
485 let i = rng.choose(d.outcomes.iter().map(|(_, p)| *p)).unwrap_or(0);
486 match &d.outcomes[i].0 {
487 Value::Continuous(f) => Value::Float(f.sample(rng)),
488 x => x.clone(),
489 }
490 }
491 Value::Continuous(f) => Value::Float(f.sample(rng)),
492 other => other.clone(),
493 }
494 }
495
496 fn sample_continuous(&mut self, v: Value) -> Value {
500 match &v {
501 Value::Continuous(_) => self.sample(&v),
502 Value::Dist(d) if d.outcomes.iter().any(|(x, _)| matches!(x, Value::Continuous(_))) => {
503 let pairs = d.outcomes.iter().map(|(x, p)| (self.sample(x), *p)).collect();
504 Dist::from_pairs(pairs, d.missing).into_value()
505 }
506 _ => v,
507 }
508 }
509
510 fn direct_counts(&mut self, f: FnId, e: &'p Expr, w: &World) -> Result<Option<Counts>> {
514 let ExprKind::Builtin {
515 func: b @ (Builtin::Binomial | Builtin::Poisson | Builtin::Geometric),
516 args,
517 ..
518 } = &e.kind
519 else {
520 return Ok(None);
521 };
522 if self.sampler.is_none() {
523 return Ok(None);
524 }
525 let mut values = Vec::with_capacity(args.len());
526 for a in args {
527 values.push(self.eval(f, a, w)?);
528 }
529 let counts = builtins::counts(*b, &values, &mut self.budget).map_err(|err| err.at(e.span))?;
530 Ok(counts.filter(Counts::direct))
531 }
532
533 fn continuous_draw(&self, span: Span) -> RuntimeError {
534 analytic::unsupported("a continuous draw inside `simulate`")
535 .at(span)
536 .with_note("`simulate` is computed by enumeration, even in sample mode")
537 .with_help("draw the value outside `simulate`; analytic joint distribution recipes aren't supported yet")
538 }
539
540 fn delaying(&self) -> bool {
543 self.config.conjugate && self.sampler.is_some()
544 }
545
546 fn updates(&mut self, variable: u32) -> &mut Updates {
547 let i = variable as usize;
548 if self.stats.updates.len() <= i {
549 let n = self.conj.variables.len().max(i + 1);
550 self.stats.updates.resize(n, Updates::default());
551 }
552 &mut self.stats.updates[i]
553 }
554
555 fn draw_delayed(&mut self, w: &mut World, slot: SlotId) {
558 if let Value::Delayed(d) = &w.slots[slot as usize] {
559 let d = **d;
560 let x = d
561 .family
562 .sample(self.sampler.as_mut().expect("delayed only when sampling"));
563 w.slots[slot as usize] = Value::Float(x);
564 self.updates(d.variable).drawn += 1;
565 }
566 }
567
568 fn observe_exactly(&mut self, f: FnId, u: &Update<'p>, w: &mut World, from: Span) -> Result<Option<f64>> {
577 let Value::Delayed(d) = &w.slots[u.slot as usize] else {
578 return Ok(None);
579 };
580 let d = **d;
581 let pair = matches!(
582 (d.family, u.likelihood),
583 (
584 Family::Beta { .. },
585 Likelihood::Binomial { .. } | Likelihood::Bernoulli { .. }
586 ) | (Family::Gamma { .. }, Likelihood::Poisson { .. })
587 | (Family::Normal { .. }, Likelihood::Normal { .. })
588 );
589 let seen = if pair { self.seen(f, u, w, from)? } else { None };
590 let Some(seen) = seen else {
591 self.draw_delayed(w, u.slot);
592 return Ok(None);
593 };
594 let (ln, posterior) = conjugate::update(&d.family, seen).expect("a conjugate pair");
595 w.slots[u.slot as usize] = Value::Delayed(Arc::new(Delayed {
596 family: posterior,
597 variable: d.variable,
598 }));
599 self.updates(d.variable).exact += 1;
600 self.densities |= matches!(seen, Seen::Normal { .. });
601 Ok(Some(ln))
602 }
603
604 fn seen(&mut self, f: FnId, u: &Update<'p>, w: &World, from: Span) -> Result<Option<Seen>> {
607 let count = |v: &Value| match v {
610 Value::Bool(_) => f64::NAN,
611 _ => v.as_f64().unwrap_or(f64::NAN),
612 };
613 Ok(match u.likelihood {
614 Likelihood::Binomial { value, trials } => {
615 let v = self.eval(f, value, w)?;
616 let n = self.eval(f, trials, w)?;
617 match builtins::counts(Builtin::Binomial, &[n, Value::Prob(0.5)], &mut self.budget)
618 .map_err(|e| e.at(from))?
619 {
620 Some(Counts::Binomial { n, .. }) => Some(Seen::Binomial {
621 trials: n,
622 k: count(&v),
623 }),
624 _ => None,
625 }
626 }
627 Likelihood::Bernoulli { value } => match self.eval(f, value, w)? {
628 Value::Bool(b) => Some(Seen::Bernoulli(b)),
629 _ => None,
630 },
631 Likelihood::Poisson { value } => {
632 let v = self.eval(f, value, w)?;
633 Some(Seen::Poisson { k: count(&v) })
634 }
635 Likelihood::Normal { value, sd } => {
636 let v = self.eval(f, value, w)?;
637 let sd = self.eval(f, sd, w)?;
638 if sd.is_uncertain() {
639 return Ok(None);
640 }
641 let checked = builtins::call_plain(Builtin::Normal, &[Value::Float(0.0), sd], &mut self.budget)
642 .map_err(|e| e.at(from))?;
643 let Value::Continuous(family) = checked else {
644 unreachable!("`normal` gives a continuous distribution")
645 };
646 let Family::Normal { sd, .. } = *family else {
647 unreachable!("`normal` gives a normal distribution")
648 };
649 match v {
650 Value::Int(_) | Value::Float(_) | Value::Prob(_) if v.as_f64().is_some() => Some(Seen::Normal {
651 y: v.as_f64().expect("a number"),
652 sd,
653 }),
654 _ => None,
655 }
656 }
657 })
658 }
659
660 fn merge(&self, worlds: Vec<World>, stmt: StmtId) -> Vec<World> {
661 merge(worlds, &self.live.after[stmt as usize], self.merging())
662 }
663
664 fn merging(&self) -> bool {
667 self.config.merging && self.sampler.is_none()
668 }
669
670 fn partial(&self) -> bool {
674 self.config.partial && self.nested == 0 && self.callback.is_none()
675 }
676
677 fn catches_here(&self, fault: Option<Fault>) -> bool {
679 self.handlers
680 .iter()
681 .any(|(depth, catches)| *depth == self.depth && catches.iter().any(|c| c.catches(fault)))
682 }
683
684 fn catches_above(&self, fault: Option<Fault>) -> bool {
686 self.handlers
687 .iter()
688 .any(|(depth, catches)| *depth < self.depth && catches.iter().any(|c| c.catches(fault)))
689 }
690
691 #[inline]
694 fn snapshot(&self, w: &World) -> Option<World> {
695 if self.handlers.is_empty() {
696 return None;
697 }
698 self.handlers
699 .iter()
700 .any(|(depth, _)| *depth == self.depth)
701 .then(|| w.clone())
702 }
703
704 fn fail(&mut self, e: RuntimeError, at: At, saved: Option<World>) -> Result<()> {
711 if e.fault.is_none() {
712 return Err(e);
713 }
714 if let Some(w) = saved.filter(|_| self.catches_here(e.fault)) {
715 self.faulted.push((w, e));
716 return Ok(());
717 }
718 if !self.partial() && !self.catches_above(e.fault) {
719 return Err(e);
720 }
721 let fun = &self.prog.functions[at.f as usize];
722 let e = match fun.kind {
723 FnKind::Named => e.with_note(format!("in a call to `{}`", fun.name)),
724 FnKind::Lambda => e.with_note("inside a lambda"),
725 FnKind::Simulate => e.with_note("inside a `simulate` block"),
726 FnKind::Main => e,
727 };
728 let stmt = at.stmt as usize;
729 let before_evidence = self.evidence.at[stmt] || self.evidence.after[stmt];
730 let run = self.sampler.is_some().then_some(at.run);
731 self.failures.record(e, at.weight, run, before_evidence);
732 Ok(())
733 }
734
735 fn check_worlds(&self, n: usize, span: Span) -> Result<()> {
738 if n > self.config.max_worlds && self.sampler.is_none() {
739 return Err(RuntimeError::limit(
740 span,
741 format!("more than {} worlds are too many to follow", self.config.max_worlds),
742 )
743 .with_help("simplify the model, raise the limit, or sample it with `@mode sample(runs: 10_000)`"));
744 }
745 Ok(())
746 }
747
748 fn spend(&mut self, n: u64, span: Span) -> Result<()> {
749 self.budget.work(n).map_err(|e| e.at(span))?;
750 if let Some(cancel) = &self.config.cancel {
751 if cancel.load(Ordering::Relaxed) {
752 return Err(RuntimeError::limit(span, "the run was cancelled"));
753 }
754 }
755 Ok(())
756 }
757
758 fn exec_block(&mut self, f: FnId, block: &'p Block, worlds: Vec<World>) -> Result<Flow> {
761 let mut flow = Flow::next(worlds);
762 for stmt in &block.stmts {
763 if flow.next.is_empty() {
764 break;
765 }
766 let input = std::mem::take(&mut flow.next);
767 let out = self.exec_stmt(f, stmt, input)?;
768 flow.next = out.next;
769 clear(&mut flow.next, &self.live.dies[stmt.id as usize]);
770 flow.broke.extend(out.broke);
771 flow.continued.extend(out.continued);
772 flow.returned.extend(out.returned);
773 flow.faulted.extend(out.faulted);
774 }
775 Ok(flow)
776 }
777
778 fn restrict_event(&mut self, event: &analytic::Event, yes: bool, w: &mut World, span: Span) -> Result<f64> {
779 self.budget
780 .collection((w.constraints.len() + 1) as u128)
781 .map_err(|e| e.at(span))?;
782 self.budget
783 .work((w.constraints.len() + event.yes.0.len() + event.draw.domain.0.len()) as u64)
784 .map_err(|e| e.at(span))?;
785 Ok(event.restrict(yes, &mut w.constraints))
786 }
787
788 fn exec_stmt(&mut self, f: FnId, stmt: &'p Stmt, worlds: Vec<World>) -> Result<Flow> {
789 if self.handlers.is_empty() {
791 return self.exec_stmt_kind(f, stmt, worlds);
792 }
793 let outer = std::mem::take(&mut self.faulted);
796 let flow = self.exec_stmt_kind(f, stmt, worlds);
797 let faulted = std::mem::replace(&mut self.faulted, outer);
798 let mut flow = flow?;
799 flow.faulted.extend(faulted);
800 Ok(flow)
801 }
802
803 fn exec_stmt_kind(&mut self, f: FnId, stmt: &'p Stmt, worlds: Vec<World>) -> Result<Flow> {
804 let span = stmt.span;
805 if matches!(
806 stmt.kind,
807 StmtKind::Draw { .. } | StmtKind::Take { .. } | StmtKind::Observe { .. } | StmtKind::Chance { .. }
808 ) {
809 self.check_callback_effect(span)?;
810 }
811 let n = worlds.len();
812 self.stats.world_steps += n as u64;
813 self.stats.peak_worlds = self.stats.peak_worlds.max(n);
814 self.spend(n as u64, span)?;
815 self.check_worlds(n, span)?;
816 let mut worlds = worlds;
817 if self.delaying() {
818 let conj = self.conj;
820 let first = &conj.draws_first[stmt.id as usize];
821 if !first.is_empty() {
822 for w in &mut worlds {
823 for &slot in first {
824 self.draw_delayed(w, slot);
825 }
826 }
827 }
828 }
829 match &stmt.kind {
830 StmtKind::Set { place, value } => {
831 let mut worlds = worlds;
832 let mut failed = Vec::new();
833 for (i, w) in worlds.iter_mut().enumerate() {
834 let mut saved = self.snapshot(w);
835 let v = match self.record_place_type(place, w) {
836 Some(ty) => self.eval_expected(f, value, w, &ty),
837 None => self.eval(f, value, w),
838 };
839 let v = each!(self, At::new(f, stmt, w), saved saved, v, { failed.push(i) });
840 each!(self, At::new(f, stmt, w), saved saved, self.assign(f, place, v, w, span), {
841 failed.push(i)
842 });
843 }
844 drop_failed(&mut worlds, &failed);
845 if self.live.narrows[stmt.id as usize] && self.merging() {
850 return Ok(Flow::next(self.merge(worlds, stmt.id)));
851 }
852 Ok(Flow::next(worlds))
853 }
854 StmtKind::Draw { place, dist } => {
855 let delay = if self.delaying() {
857 self.conj.delays[stmt.id as usize]
858 } else {
859 None
860 };
861 let mut out = Vec::with_capacity(worlds.len());
862 for w in worlds {
863 let at = At::new(f, stmt, &w);
864 let mut saved = self.snapshot(&w);
865 if self.sampler.is_some() {
866 if let Some(counts) = each!(self, at, saved saved, self.direct_counts(f, dist, &w)) {
867 let k = counts.sample(self.sampler.as_mut().expect("sampling"));
868 let mut w = w;
869 each!(self, at, saved saved, self.assign(f, place, Value::Int(k.into()), &mut w, dist.span));
870 out.push(w);
871 continue;
872 }
873 }
874 let d = each!(self, at, saved saved, self.eval(f, dist, &w));
875 if let (Some(variable), Value::Continuous(family)) = (delay, &d) {
876 if conjugate::is_prior(family) {
877 let mut w = w;
878 w.slots[place.slot as usize] = Value::Delayed(Arc::new(Delayed {
879 family: **family,
880 variable,
881 }));
882 self.updates(variable).delayed += 1;
883 out.push(w);
884 continue;
885 }
886 }
887 let mark = out.len();
888 each!(self, at, saved saved, self.split_by(f, place, d, w, dist.span, &mut out), {
889 out.truncate(mark)
890 });
891 self.check_worlds(out.len(), span)?;
892 }
893 Ok(Flow::next(self.merge(out, stmt.id)))
894 }
895 StmtKind::Take { place, bag } => {
896 let mut out = Vec::new();
897 for w in worlds {
898 let at = At::new(f, stmt, &w);
899 let current = each!(self, at, world w, self.read_place(f, bag, &w, span));
900 let Value::Bag(cards) = ¤t else {
901 return Err(RuntimeError::new(
902 span,
903 format!("`take` needs a bag, found {}", ops::article(¤t.kind())),
904 )
905 .with_help("make one with `bag([card: count, …])`"));
906 };
907 let total: u128 = cards.values().map(|n| *n as u128).sum();
908 if total == 0 {
909 let empty = RuntimeError::new(span, "can't take a card from an empty bag")
910 .as_fault(Fault::EmptyCollection);
911 each!(self, at, world w, Err::<(), _>(empty));
912 }
913 let chosen = match &mut self.sampler {
914 Some(rng) => rng.choose(cards.values().map(|n| *n as f64)),
915 None => None,
916 };
917 for (i, (card, count)) in cards.iter().enumerate() {
918 if self.sampler.is_some() && chosen != Some(i) {
919 continue;
920 }
921 let rest = cards.without_nth(i);
922 let mut nw = if self.sampler.is_some() {
923 w.clone()
924 } else {
925 w.clone().scaled(*count as f64 / total as f64)
926 };
927 self.assign(f, bag, Value::multiset(rest), &mut nw, span)?;
928 self.assign(f, place, card.clone(), &mut nw, span)?;
929 out.push(nw);
930 }
931 self.check_worlds(out.len(), span)?;
932 }
933 Ok(Flow::next(self.merge(out, stmt.id)))
934 }
935 StmtKind::Call { dest, callee, args } => {
936 let mut out = Vec::with_capacity(worlds.len());
937 for w in worlds {
938 let at = At::new(f, stmt, &w);
939 let mut key = Vec::with_capacity(args.len() + 4);
940 let evaluated: Result<()> = args.iter().try_for_each(|a| {
941 key.push(self.eval(f, a, &w)?);
942 Ok(())
943 });
944 each!(self, at, world w, evaluated);
945 let func = match callee {
946 Callee::Fn { func, capture_args } => {
947 for &s in capture_args {
948 key.push(self.slot(f, s, &w, span)?);
949 }
950 *func
951 }
952 Callee::Value(e, named) => match each!(self, at, world w, self.eval(f, e, &w)) {
953 Value::Builtin(b) => {
954 let value =
955 each!(self, at, world w, self.builtin_values(b, &key, named, w.weight, span));
956 let mut saved = self.snapshot(&w);
957 let mut nw = w;
958 each!(self, at, saved saved, self.assign(f, dest, value, &mut nw, span));
959 out.push(nw);
960 continue;
961 }
962 Value::Closure(c) => {
963 if !named.is_empty() {
964 return Err(RuntimeError::new(
965 span,
966 "only builtin minimum/maximum accept named defaults",
967 ));
968 }
969 self.check_arity(&c, args.len(), span)?;
970 key.extend(c.captured.iter().cloned());
971 c.func
972 }
973 other => {
974 return Err(RuntimeError::new(
975 e.span,
976 format!("can't call {}", ops::article(&other.kind())),
977 ));
978 }
979 },
980 };
981 let result = each!(self, at, world w, self.call(func, key, span));
982 self.unresolved += w.weight * result.unresolved;
983 self.lost += w.weight * result.lost;
984 let run = self.sampler.is_some().then_some(w.run);
989 let after = self.evidence.after[stmt.id as usize];
990 for g in &result.failures.groups {
991 if self.catches_here(g.error.fault) {
992 let mut caught = w.clone();
993 caught.weight = caught.weight * g.weight;
994 self.faulted.push((caught, g.error.clone()));
995 } else if self.partial() || self.catches_above(g.error.fault) {
996 self.failures.absorb_group(g, w.weight, run, after);
997 } else {
998 return Err(g.error.clone());
999 }
1000 }
1001 add_pending(&mut self.pending, &result.pending, w.weight);
1002 let last = result.outcomes.len().saturating_sub(1);
1003 let mut w = Some(w);
1004 for (i, (v, p, restrictions)) in result.outcomes.iter().enumerate() {
1005 let mut nw = if i == last {
1006 w.take().unwrap()
1007 } else {
1008 w.clone().unwrap()
1009 };
1010 nw.weight = nw.weight * *p;
1011 if !restrictions.is_empty() {
1014 Arc::make_mut(&mut nw.constraints)
1015 .extend(restrictions.iter().map(|(id, d)| (*id, d.clone())));
1016 }
1017 self.assign(f, dest, v.clone(), &mut nw, span)?;
1018 out.push(nw);
1019 }
1020 self.check_worlds(out.len(), span)?;
1021 }
1022 Ok(Flow::next(self.merge(out, stmt.id)))
1023 }
1024 StmtKind::If { cond, then, otherwise } => {
1025 let (mut yes, mut no) = (Vec::new(), Vec::new());
1026 for w in worlds {
1027 let condition = each!(self, At::new(f, stmt, &w), world w, self.eval(f, cond, &w));
1028 if let Value::Event(event) = condition {
1029 self.check_callback_effect(cond.span)?;
1030 let mut y = w.clone();
1031 let mut n = w;
1032 let p = self.restrict_event(&event, true, &mut y, cond.span)?;
1033 let q = self.restrict_event(&event, false, &mut n, cond.span)?;
1034 if p > 0.0 {
1035 yes.push(y.scaled(p));
1036 }
1037 if q > 0.0 {
1038 no.push(n.scaled(q));
1039 }
1040 continue;
1041 }
1042 let c = ops::condition(&condition).map_err(|e| e.at(cond.span))?;
1043 if (c.yes > 0.0 && c.no > 0.0) || c.missing > 0.0 {
1044 self.check_callback_effect(cond.span)?;
1045 }
1046 if let Some(rng) = &mut self.sampler {
1047 match rng.choose([c.yes, c.no].into_iter()) {
1049 Some(0) => yes.push(w),
1050 Some(_) => no.push(w),
1051 None => self.unresolved += w.weight,
1052 }
1053 continue;
1054 }
1055 let (p, q) = (c.yes, c.no);
1056 if c.missing > 0.0 {
1057 self.unresolved += w.weight.scale(c.missing);
1058 }
1059 if q <= 0.0 {
1060 yes.push(w.scaled(p));
1061 } else if p <= 0.0 {
1062 no.push(w.scaled(q));
1063 } else {
1064 no.push(w.clone().scaled(q));
1065 yes.push(w.scaled(p));
1066 }
1067 }
1068 self.check_worlds(yes.len() + no.len(), span)?;
1069 let mut flow = Flow::default();
1070 if !yes.is_empty() {
1071 flow.join(self.exec_block(f, then, yes)?);
1072 }
1073 if !no.is_empty() {
1074 flow.join(self.exec_block(f, otherwise, no)?);
1075 }
1076 flow.next = self.merge(flow.next, stmt.id);
1077 Ok(flow)
1078 }
1079 StmtKind::Chance {
1080 arms,
1081 otherwise,
1082 exhaustive,
1083 } => {
1084 let mut buckets: Vec<Vec<World>> = vec![Vec::new(); arms.len()];
1085 let mut rest = Vec::new();
1086 for w in worlds {
1087 let at = At::new(f, stmt, &w);
1088 let chances: Result<Vec<f64>> = arms
1090 .iter()
1091 .map(|(weight, _)| {
1092 let value = self.eval(f, weight, &w)?;
1093 ops::to_prob(&value).map_err(|e| e.at(weight.span))
1094 })
1095 .collect();
1096 let chances = each!(self, at, world w, chances);
1097 let sum: f64 = chances.iter().sum();
1098 if sum > 1.0 + 1e-9 {
1099 let over = RuntimeError::new(
1100 span,
1101 format!(
1102 "the chances in this `chance` add up to {}, more than 100%",
1103 fmt_prob(sum)
1104 ),
1105 );
1106 each!(self, at, world w, Err::<(), _>(over.as_fault(Fault::DomainError)));
1107 }
1108 let remainder = (1.0 - sum).max(0.0);
1109 if remainder > 1e-9 && *exhaustive && otherwise.is_none() {
1110 let short = RuntimeError::new(
1111 span,
1112 format!("the chances add up to {} and there's no `else` arm", fmt_prob(sum)),
1113 )
1114 .with_help("a `chance` used as a value needs its chances to add up to 100%, or an `else`");
1115 each!(self, at, world w, Err::<(), _>(short.as_fault(Fault::DomainError)));
1116 }
1117 if let Some(rng) = &mut self.sampler {
1118 let options = chances.iter().copied().chain(std::iter::once(remainder));
1120 match rng.choose(options) {
1121 Some(i) if i < arms.len() => buckets[i].push(w),
1122 Some(_) => rest.push(w),
1123 None => self.unresolved += w.weight,
1124 }
1125 continue;
1126 }
1127 for (bucket, p) in buckets.iter_mut().zip(&chances) {
1128 if *p > 0.0 {
1129 bucket.push(w.clone().scaled(*p));
1130 }
1131 }
1132 if remainder > 1e-12 {
1133 rest.push(w.scaled(remainder));
1134 }
1135 }
1136 self.check_worlds(buckets.iter().map(Vec::len).sum::<usize>() + rest.len(), span)?;
1137 let mut flow = Flow::default();
1138 for ((_, body), bucket) in arms.iter().zip(buckets) {
1139 if !bucket.is_empty() {
1140 flow.join(self.exec_block(f, body, bucket)?);
1141 }
1142 }
1143 if !rest.is_empty() {
1144 match otherwise {
1145 Some(body) => flow.join(self.exec_block(f, body, rest)?),
1146 None if *exhaustive => {}
1147 None => flow.next.extend(rest),
1148 }
1149 }
1150 flow.next = self.merge(flow.next, stmt.id);
1151 Ok(flow)
1152 }
1153 StmtKind::Loop { body, bounded } => self.exec_loop(f, stmt, body, *bounded, worlds),
1154 StmtKind::Break => Ok(Flow {
1155 broke: worlds,
1156 ..Flow::default()
1157 }),
1158 StmtKind::Continue => Ok(Flow {
1159 continued: worlds,
1160 ..Flow::default()
1161 }),
1162 StmtKind::Return(value) => {
1163 let mut flow = Flow::default();
1164 for w in worlds {
1165 let v = each!(self, At::new(f, stmt, &w), world w, self.eval(f, value, &w));
1166 let mut constraints = w.constraints;
1167 if constraints.keys().any(|id| !w.inherited.contains(id)) {
1168 Arc::make_mut(&mut constraints).retain(|id, _| w.inherited.contains(id));
1169 }
1170 flow.returned.push((v, w.weight, constraints));
1171 }
1172 Ok(flow)
1173 }
1174 StmtKind::Observe { value, from } => {
1175 self.observed = true;
1176 let exact = if self.delaying() {
1177 probl_sema::conjugate::update(value, from.as_ref())
1178 } else {
1179 None
1180 };
1181 let mut out = Vec::with_capacity(worlds.len());
1182 for mut w in worlds {
1183 let here = At::new(f, stmt, &w);
1184 let mut saved = self.snapshot(&w);
1185 if let Some(u) = &exact {
1186 let at = from.as_ref().map_or(span, |d| d.span);
1187 if let Some(ln) = each!(self, here, saved saved, self.observe_exactly(f, u, &mut w, at)) {
1188 w.weight = w.weight * Weight::from_ln(ln);
1189 if w.weight.is_zero() {
1190 self.last_ruling_out = Some(span);
1191 } else {
1192 out.push(w);
1193 }
1194 continue;
1195 }
1196 }
1197 let (factor, missing, ruled_out) = match from {
1198 None => {
1199 let v = each!(self, here, saved saved, self.eval(f, value, &w));
1200 if let Value::Event(event) = v {
1201 let p = self.restrict_event(&event, true, &mut w, value.span)?;
1202 (p, 0.0, 1.0 - p)
1203 } else {
1204 let b = ops::fact(&v, "observe").map_err(|e| e.at(value.span))?;
1205 if b { (1.0, 0.0, 0.0) } else { (0.0, 0.0, 1.0) }
1206 }
1207 }
1208 Some(d) => {
1209 let v = each!(self, here, saved saved, self.eval(f, value, &w));
1210 match each!(self, here, saved saved, self.direct_counts(f, d, &w)) {
1211 Some(counts) if self.sampler.is_some() => {
1212 let x = match v {
1213 Value::Bool(_) => None,
1214 _ => v.as_f64(),
1215 };
1216 (x.map_or(0.0, |x| counts.pmf(x)), 0.0, 0.0)
1217 }
1218 _ => {
1219 let dist = each!(self, here, saved saved, self.eval(f, d, &w));
1220 if let (Value::Bool(b), Value::Event(event)) = (&v, &dist) {
1221 let p = self.restrict_event(event, *b, &mut w, span)?;
1222 (p, 0.0, 1.0 - p)
1223 } else {
1224 if analytic::contains(&v) || analytic::contains(&dist) {
1225 return Err(analytic::unsupported("this likelihood observation").at(span));
1226 }
1227 self.densities |= is_density(&dist);
1228 each!(
1229 self,
1230 here,
1231 saved saved,
1232 likelihood(&dist, &v, self.sampler.is_some()).map_err(|e| e.at(span))
1233 )
1234 }
1235 }
1236 }
1237 }
1238 };
1239 if missing > 0.0 && self.sampler.is_none() {
1240 self.unresolved += w.weight.scale(missing);
1241 }
1242 if ruled_out > 0.0 && self.sampler.is_none() {
1243 self.lost += w.weight.scale(ruled_out);
1244 }
1245 if factor > 0.0 {
1246 w.weight = w.weight.scale(factor);
1247 out.push(w);
1248 } else {
1249 self.last_ruling_out = Some(span);
1250 }
1251 }
1252 Ok(Flow::next(out))
1253 }
1254 StmtKind::Report { site, value, key } => {
1255 let mut failed = Vec::new();
1256 for (i, w) in worlds.iter().enumerate() {
1257 let v = each!(self, At::new(f, stmt, w), world w, self.eval(f, value, w), { failed.push(i) });
1258 let k = match key {
1259 Some(k) => {
1260 let k_value =
1261 each!(self, At::new(f, stmt, w), world w, self.eval(f, k, w), { failed.push(i) });
1262 if analytic::contains(&k_value) {
1263 return Err(analytic::unsupported("grouping by a continuous outcome").at(k.span));
1264 }
1265 if k_value.is_uncertain() {
1266 return Err(RuntimeError::new(
1267 k.span,
1268 "a `by` key must be a plain value, not a distribution",
1269 )
1270 .with_help("draw a value first with `~`"));
1271 }
1272 k_value
1273 }
1274 None => Value::Unit,
1275 };
1276 let (v, run) = if self.sampler.is_some() {
1277 (self.sample_continuous(v), Some(w.run))
1278 } else {
1279 if analytic::contains(&v) && !matches!(v, Value::Analytic(_) | Value::Event(_)) {
1280 return Err(analytic::unsupported(
1281 "reporting an aggregate containing analytic outcomes; report its fields individually",
1282 )
1283 .at(value.span));
1284 }
1285 (v, None)
1286 };
1287 self.sinks[*site as usize]
1288 .validate_analytic(&k, &v)
1289 .map_err(|e| e.at(value.span))?;
1290 self.sinks[*site as usize].add(k, &v, w.weight, run);
1291 }
1292 let mut worlds = worlds;
1293 drop_failed(&mut worlds, &failed);
1294 Ok(Flow::next(worlds))
1295 }
1296 StmtKind::Fail { message } => Err(RuntimeError::new(span, message.clone())),
1297 StmtKind::Try { body, catches } => {
1298 self.handlers.push((self.depth, catches.as_slice()));
1299 let flow = self.exec_block(f, body, worlds);
1300 self.handlers.pop();
1301 let mut flow = flow?;
1302 let mut caught: Vec<Vec<World>> = vec![Vec::new(); catches.len()];
1305 for (w, e) in std::mem::take(&mut flow.faulted) {
1306 match catches.iter().position(|c| c.catches(e.fault)) {
1307 Some(i) => caught[i].push(w),
1308 None => flow.faulted.push((w, e)),
1309 }
1310 }
1311 for (c, worlds) in catches.iter().zip(caught) {
1312 if !worlds.is_empty() {
1313 flow.join(self.exec_block(f, &c.body, worlds)?);
1314 }
1315 }
1316 flow.next = self.merge(flow.next, stmt.id);
1317 Ok(flow)
1318 }
1319 StmtKind::Check { slot, ty } => {
1320 for w in &mut worlds {
1321 if let Value::Delayed(_) = &w.slots[*slot as usize] {
1324 let always = *ty == TypeSpec::Float;
1325 if always {
1326 continue;
1327 }
1328 self.draw_delayed(w, *slot);
1329 }
1330 w.slots[*slot as usize] = self.coerce(w.slots[*slot as usize].clone(), ty, span)?;
1331 let v = &w.slots[*slot as usize];
1332 if !self.conforms(v, ty) {
1333 let name = &self.prog.functions[f as usize].slots[*slot as usize].name;
1334 let what = if name == TEMP {
1335 "the result".to_string()
1336 } else {
1337 format!("`{name}`")
1338 };
1339 return Err(RuntimeError::new(
1340 span,
1341 format!(
1342 "{what} should be {}, but it's {}",
1343 ops::article(&ty.describe(self.prog)),
1344 ops::article(&v.kind())
1345 ),
1346 ));
1347 }
1348 }
1349 Ok(Flow::next(worlds))
1350 }
1351 }
1352 }
1353
1354 fn exec_loop(
1355 &mut self,
1356 f: FnId,
1357 stmt: &'p Stmt,
1358 body: &'p Block,
1359 bounded: bool,
1360 worlds: Vec<World>,
1361 ) -> Result<Flow> {
1362 let live = self.live;
1363 let span = stmt.span;
1364 let entered = total_weight(&worlds);
1365 let cutoff = entered.scale(self.config.epsilon);
1366 let mut inside = worlds;
1367 let mut out = Flow::default();
1368 let mut iterations: u64 = 0;
1369 let mut met = self.may_solve(stmt, bounded).then(FxHashSet::<u64>::default);
1373 while !inside.is_empty() {
1374 let cycles = match &mut met {
1375 Some(met) => {
1376 let head = live_slots(&live.loop_head[stmt.id as usize], inside[0].slots.len());
1377 let hashes: Vec<u64> = inside.iter().map(|w| state_hash(w, &head)).collect();
1378 let again = hashes.iter().any(|h| met.contains(h));
1379 met.extend(hashes);
1380 again
1381 }
1382 None => false,
1383 };
1384 if cycles {
1385 met = None;
1386 if let Some(solved) = self.solve_loop(f, stmt, body, &inside)? {
1387 out.join(solved);
1388 break;
1389 }
1390 }
1391 let mass = total_weight(&inside);
1392 if !bounded && mass < cutoff && self.sampler.is_none() {
1393 self.unresolved += mass;
1394 break;
1395 }
1396 iterations += 1;
1397 if iterations > self.config.max_iterations {
1398 return Err(RuntimeError::limit(
1399 span,
1400 format!("this loop ran {} times without finishing", self.config.max_iterations),
1401 )
1402 .with_note(format!(
1403 "worlds still inside the loop weigh {} of what entered it",
1404 fmt_prob(mass.ratio(entered))
1405 ))
1406 .with_help("check that every world can leave the loop"));
1407 }
1408 let mut flow = self.exec_block(f, body, inside)?;
1409 clear_dead(&mut flow.broke, &live.after[stmt.id as usize]);
1410 clear_dead(&mut flow.continued, &live.loop_head[stmt.id as usize]);
1411 out.next.extend(flow.broke);
1412 out.returned.extend(flow.returned);
1413 out.faulted.extend(flow.faulted);
1414 let mut again = flow.next;
1415 again.extend(flow.continued);
1416 inside = merge(again, &live.loop_head[stmt.id as usize], self.merging());
1417 }
1418 out.next = self.merge(out.next, stmt.id);
1419 Ok(out)
1420 }
1421
1422 fn may_solve(&self, stmt: &Stmt, bounded: bool) -> bool {
1426 !bounded
1427 && self.config.solving
1428 && self.sampler.is_none()
1429 && self.solvable[stmt.id as usize]
1430 && !self.too_large.contains(&stmt.id)
1431 }
1432
1433 fn solve_loop(&mut self, f: FnId, stmt: &'p Stmt, body: &'p Block, inside: &[World]) -> Result<Option<Flow>> {
1441 if inside
1442 .iter()
1443 .any(|w| !w.constraints.is_empty() || w.slots.iter().any(analytic::contains))
1444 {
1445 self.too_large.insert(stmt.id);
1446 return Ok(None);
1447 }
1448 let head_set = &self.live.loop_head[stmt.id as usize];
1449 let after = &self.live.after[stmt.id as usize];
1450 let n = inside[0].slots.len();
1451 let head = live_slots(head_set, n);
1452 let dead: Vec<SlotId> = head_set.iter_missing(n).collect();
1453 let mut states: Vec<Vec<Value>> = Vec::new();
1456 let mut index: FxHashMap<Vec<Value>, usize> = FxHashMap::default();
1457 let mut intern = |mut slots: Vec<Value>, states: &mut Vec<Vec<Value>>| -> usize {
1458 for &d in &dead {
1459 slots[d as usize] = Value::Dead;
1460 }
1461 let key: Vec<Value> = head.iter().map(|&i| slots[i].clone()).collect();
1462 *index.entry(key).or_insert_with(|| {
1463 states.push(slots);
1464 states.len() - 1
1465 })
1466 };
1467 let total = total_weight(inside);
1468 let mut start: Vec<f64> = Vec::new();
1469 for w in inside {
1470 let i = intern(w.slots.clone(), &mut states);
1471 if start.len() <= i {
1472 start.resize(i + 1, 0.0);
1473 }
1474 start[i] += w.weight.ratio(total);
1475 }
1476 let mut chain = Chain::default();
1477 let mut exits: Vec<Vec<World>> = Vec::new();
1478 let mut returns: Vec<Vec<Returned>> = Vec::new();
1479 let mut unresolved: Vec<Weight> = Vec::new();
1480 let mut failed: Vec<Failures> = Vec::new();
1481 let mut faulted: Vec<Vec<(World, RuntimeError)>> = Vec::new();
1482 let saved_pending = std::mem::take(&mut self.pending);
1485 let mut k = 0;
1486 while k < states.len() {
1487 if !self.pending.is_empty() {
1488 self.pending = saved_pending;
1489 return Ok(None);
1490 }
1491 if states.len() > self.config.max_chain_states {
1492 self.pending = saved_pending;
1493 self.too_large.insert(stmt.id);
1494 return Ok(None);
1495 }
1496 let world = World {
1497 slots: states[k].clone(),
1498 constraints: Default::default(),
1499 inherited: Default::default(),
1500 weight: Weight::ONE,
1501 run: 0,
1502 };
1503 let saved_unresolved = std::mem::replace(&mut self.unresolved, Weight::ZERO);
1504 let saved_lost = std::mem::replace(&mut self.lost, Weight::ZERO);
1505 let saved_failures = std::mem::take(&mut self.failures);
1506 let flow = self.exec_block(f, body, vec![world]);
1507 let left = std::mem::replace(&mut self.unresolved, saved_unresolved);
1508 let lost = std::mem::replace(&mut self.lost, saved_lost);
1509 let failures = std::mem::replace(&mut self.failures, saved_failures);
1510 let mut flow = flow?;
1511 if flow
1512 .next
1513 .iter()
1514 .chain(&flow.continued)
1515 .chain(&flow.broke)
1516 .any(|w| !w.constraints.is_empty() || w.slots.iter().any(analytic::contains))
1517 || flow
1518 .returned
1519 .iter()
1520 .any(|(v, _, c)| !c.is_empty() || analytic::contains(v))
1521 {
1522 self.pending = saved_pending;
1523 self.too_large.insert(stmt.id);
1524 return Ok(None);
1525 }
1526 clear_dead(&mut flow.broke, after);
1527 let mut leave = left + lost + failures.weight;
1530 let faults = std::mem::take(&mut flow.faulted);
1531 for (w, _) in &faults {
1532 leave += w.weight;
1533 }
1534 for w in &flow.broke {
1535 leave += w.weight;
1536 }
1537 for (_, w, _) in &flow.returned {
1538 leave += *w;
1539 }
1540 let mut next: FxHashMap<usize, f64> = FxHashMap::default();
1541 for w in flow.next.into_iter().chain(flow.continued) {
1542 let p = w.weight.to_f64();
1543 if p == 0.0 {
1544 self.too_large.insert(stmt.id);
1546 self.pending = saved_pending;
1547 return Ok(None);
1548 }
1549 let j = intern(w.slots, &mut states);
1550 *next.entry(j).or_insert(0.0) += p;
1551 }
1552 let mut next: Vec<(usize, f64)> = next.into_iter().collect();
1553 next.sort_by_key(|&(j, _)| j);
1554 chain.next.push(next);
1555 chain.leave.push(leave.to_f64());
1556 exits.push(flow.broke);
1557 returns.push(flow.returned);
1558 unresolved.push(left);
1559 failed.push(failures);
1560 faulted.push(faults);
1561 k += 1;
1562 }
1563 let waiting = !self.pending.is_empty();
1564 self.pending = saved_pending;
1565 if waiting {
1566 return Ok(None);
1567 }
1568 start.resize(states.len(), 0.0);
1569 let mut steps = MAX_ELIMINATION;
1570 let solution = chain.visits(&start, &mut steps);
1571 self.spend(MAX_ELIMINATION - steps, stmt.span)?;
1572 let visits = match solution {
1573 Solution::Visits(v) => v,
1574 Solution::TooBig => {
1575 self.too_large.insert(stmt.id);
1576 return Ok(None);
1577 }
1578 Solution::Stuck(s) => return Err(self.never_leaves(f, stmt, &states[s], &head)),
1579 };
1580 self.stats.solved_loops += 1;
1581 self.stats.chain_states += states.len() as u64;
1582 let mut flow = Flow::default();
1583 let mut left = Weight::ZERO;
1584 for (k, v) in visits.into_iter().enumerate() {
1585 if v == 0.0 {
1586 continue;
1587 }
1588 let times = total.scale(v);
1589 for mut w in std::mem::take(&mut exits[k]) {
1590 w.weight = w.weight * times;
1591 flow.next.push(w);
1592 }
1593 for (value, w, c) in std::mem::take(&mut returns[k]) {
1594 flow.returned.push((value, w * times, c));
1595 }
1596 left += unresolved[k] * times;
1597 self.failures.absorb(&failed[k], times, None, false);
1598 for (mut w, e) in std::mem::take(&mut faulted[k]) {
1599 w.weight = w.weight * times;
1600 flow.faulted.push((w, e));
1601 }
1602 }
1603 self.unresolved += left;
1604 Ok(Some(flow))
1605 }
1606
1607 fn never_leaves(&self, f: FnId, stmt: &Stmt, slots: &[Value], head: &[usize]) -> RuntimeError {
1610 let names = &self.prog.functions[f as usize].slots;
1611 let parts: Vec<String> = head
1612 .iter()
1613 .filter(|&&i| names[i].name != TEMP)
1614 .map(|&i| format!("`{}` is {:?}", names[i].name, slots[i]))
1615 .collect();
1616 let err = RuntimeError::new(stmt.span, "some worlds can never leave this loop")
1617 .with_help("check that every world can leave the loop");
1618 if parts.is_empty() {
1619 err
1620 } else {
1621 err.with_note(format!("for example, the worlds where {}", parts.join(" and ")))
1622 }
1623 }
1624
1625 fn split_by(&mut self, f: FnId, place: &Place, d: Value, w: World, span: Span, out: &mut Vec<World>) -> Result<()> {
1627 match d {
1628 Value::Dist(_) | Value::Continuous(_) if self.sampler.is_some() => {
1629 let v = self.sample(&d);
1630 let mut w = w;
1631 self.assign(f, place, v, &mut w, span)?;
1632 out.push(w);
1633 }
1634 Value::Continuous(family) => {
1635 if self.nested > 0 {
1636 return Err(self.continuous_draw(span));
1637 }
1638 self.next_latent += 1;
1639 let v = Analytic::new(self.next_latent, *family)
1640 .value()
1641 .map_err(|e| e.at(span))?;
1642 let mut w = w;
1643 self.assign(f, place, v, &mut w, span)?;
1644 out.push(w);
1645 }
1646 Value::Dist(dist) => {
1647 self.unresolved += w.weight.scale(dist.missing);
1648 let last = dist.outcomes.len().saturating_sub(1);
1649 let mut w = Some(w);
1650 for (i, (v, p)) in dist.outcomes.iter().enumerate() {
1651 let base = if i == last {
1652 w.take().unwrap()
1653 } else {
1654 w.clone().unwrap()
1655 };
1656 let mut nw = base.scaled(*p);
1657 if matches!(v, Value::Continuous(_)) {
1658 self.split_by(f, place, v.clone(), nw, span, out)?;
1659 } else {
1660 self.assign(f, place, v.clone(), &mut nw, span)?;
1661 out.push(nw);
1662 }
1663 }
1664 }
1665 Value::Prob(p) => {
1666 return self.split_by(f, place, Dist::bernoulli(p).into_value(), w, span, out);
1667 }
1668 other => {
1669 let mut w = w;
1670 self.assign(f, place, other, &mut w, span)?;
1671 out.push(w);
1672 }
1673 }
1674 Ok(())
1675 }
1676
1677 fn check_arity(&self, c: &Closure, given: usize, span: Span) -> Result<()> {
1678 let expected = self.prog.functions[c.func as usize].n_params as usize;
1679 if expected != given {
1680 return Err(RuntimeError::new(
1681 span,
1682 format!("this function takes {expected} argument(s), but {given} were given"),
1683 ));
1684 }
1685 Ok(())
1686 }
1687
1688 fn coerce(&mut self, v: Value, ty: &TypeSpec, span: Span) -> Result<Value> {
1691 self.coerce_at(v, ty, span, 0)
1692 }
1693
1694 fn coerce_at(&mut self, v: Value, ty: &TypeSpec, span: Span, depth: usize) -> Result<Value> {
1695 self.budget.work(1).map_err(|e| e.at(span))?;
1696 if depth > 64 {
1697 return Err(RuntimeError::new(
1698 span,
1699 "type conversion nesting exceeds the limit of 64",
1700 ));
1701 }
1702 Ok(match (v, ty) {
1703 (Value::Analytic(_) | Value::Event(_), TypeSpec::Prob) => {
1704 return Err(analytic::unsupported("converting this outcome to prob").at(span));
1705 }
1706 (v @ Value::Float(_), TypeSpec::Int) => Value::Int(
1707 ops::integer(&v, "int conversion", &mut self.budget)
1708 .map_err(|e| e.at(span))?
1709 .into_owned(),
1710 ),
1711 (Value::Analytic(_), TypeSpec::Int) => {
1712 return Err(analytic::unsupported("converting this outcome to int").at(span));
1713 }
1714 (Value::Prob(p), TypeSpec::Float) => Value::Float(p),
1715 (v @ (Value::Int(_) | Value::Float(_)), TypeSpec::Prob) => {
1718 ops::make_prob(&v).map_err(|e| OpError { fault: None, ..e }.at(span))?
1719 }
1720 (Value::List(xs), TypeSpec::List(t)) => {
1721 self.budget.collection(xs.len() as u128).map_err(|e| e.at(span))?;
1722 let mut out = Vec::with_capacity(xs.len());
1723 for x in xs.iter() {
1724 out.push(self.coerce_at(x.clone(), t, span, depth + 1)?);
1725 }
1726 Value::list(out)
1727 }
1728 (Value::Map(xs), TypeSpec::Map(kt, vt)) => {
1729 let mut out = BTreeMap::new();
1730 for (k, v) in xs.iter() {
1731 let key = self.coerce_at(k.clone(), kt, span, depth + 1)?;
1732 let value = self.coerce_at(v.clone(), vt, span, depth + 1)?;
1733 if out.insert(key, value).is_some() {
1734 return Err(RuntimeError::new(span, "map key collision during type conversion"));
1735 }
1736 }
1737 Value::map(out)
1738 }
1739 (Value::Bag(xs), TypeSpec::Bag(t)) => {
1740 let mut out = BTreeMap::<Value, u64>::new();
1741 for (x, n) in xs.iter() {
1742 let key = self.coerce_at(x.clone(), t, span, depth + 1)?;
1743 let count = out.entry(key).or_default();
1744 *count = count
1745 .checked_add(*n)
1746 .ok_or_else(|| RuntimeError::new(span, "bag count overflow during type conversion"))?;
1747 }
1748 Value::multiset(crate::value::Multiset::new(out))
1749 }
1750 (Value::Dist(d), TypeSpec::Dist(t)) => {
1751 let mut out = Vec::with_capacity(d.outcomes.len());
1752 for (x, p) in &d.outcomes {
1753 out.push((self.coerce_at(x.clone(), t, span, depth + 1)?, *p));
1754 }
1755 ops::combine(out, d.missing, &mut self.budget).map_err(|e| e.at(span))?
1756 }
1757 (Value::Record(r), TypeSpec::Record(t)) => {
1758 let types: Vec<_> = self.prog.records[*t as usize]
1759 .fields
1760 .iter()
1761 .map(|f| (f.name.clone(), f.ty.clone()))
1762 .collect();
1763 let mut fields = Vec::with_capacity(r.fields.len());
1764 for (name, value) in &r.fields {
1765 let value = match types.iter().find(|(n, _)| n == &**name) {
1766 Some((_, t)) => self.coerce_at(value.clone(), t, span, depth + 1)?,
1767 None => value.clone(),
1768 };
1769 fields.push((name.clone(), value));
1770 }
1771 ops::make_record(r.ty.clone(), fields)
1772 }
1773 (Value::Record(r), TypeSpec::AnonRecord(types)) => {
1774 let mut fields = Vec::with_capacity(r.fields.len());
1775 for (name, value) in &r.fields {
1776 let value = match types.iter().find(|(n, _)| n == &**name) {
1777 Some((_, t)) => self.coerce_at(value.clone(), t, span, depth + 1)?,
1778 None => value.clone(),
1779 };
1780 fields.push((name.clone(), value));
1781 }
1782 ops::make_record(r.ty.clone(), fields)
1783 }
1784 (v, _) => v,
1785 })
1786 }
1787
1788 fn conforms(&self, v: &Value, ty: &TypeSpec) -> bool {
1790 match (ty, v) {
1791 (TypeSpec::Int, Value::Int(_)) => true,
1792 (TypeSpec::Float, Value::Float(_) | Value::Int(_) | Value::Analytic(_)) => true,
1793 (TypeSpec::Complex, Value::Complex(_)) => true,
1794 (TypeSpec::Prob, Value::Prob(_)) => true,
1795 (TypeSpec::Bool, Value::Bool(_) | Value::Event(_)) => true,
1796 (TypeSpec::Str, Value::Str(_)) => true,
1797 (TypeSpec::Date, Value::Date(_)) => true,
1798 (TypeSpec::Unit, Value::Unit) => true,
1799 (TypeSpec::Function, Value::Closure(_) | Value::Builtin(_)) => true,
1800 (TypeSpec::List(t), Value::List(items)) => items.iter().all(|x| self.conforms(x, t)),
1801 (TypeSpec::List(t), Value::Range(..)) => **t == TypeSpec::Int,
1802 (TypeSpec::Map(k, t), Value::Map(m)) => m.iter().all(|(a, b)| self.conforms(a, k) && self.conforms(b, t)),
1803 (TypeSpec::Bag(t), Value::Bag(b)) => b.keys().all(|x| self.conforms(x, t)),
1804 (TypeSpec::Dist(t), Value::Dist(d)) => d.outcomes.iter().all(|(x, _)| match x {
1805 Value::Continuous(_) => **t == TypeSpec::Float,
1806 _ => self.conforms(x, t),
1807 }),
1808 (TypeSpec::Dist(t), Value::Continuous(_)) => **t == TypeSpec::Float,
1809 (TypeSpec::Record(r), Value::Record(rec)) => {
1810 let declared = &self.prog.records[*r as usize];
1811 rec.ty.as_deref() == Some(declared.name.as_str())
1812 && rec.fields.len() == declared.fields.len()
1813 && declared
1814 .fields
1815 .iter()
1816 .all(|f| rec.get(&f.name).is_some_and(|x| self.conforms(x, &f.ty)))
1817 }
1818 (TypeSpec::Enum(e), Value::Enum(x)) => x.ty == *e,
1819 (TypeSpec::AnonRecord(fields), Value::Record(rec)) => {
1820 rec.fields.len() == fields.len()
1821 && rec
1822 .fields
1823 .iter()
1824 .zip(fields)
1825 .all(|((n, x), (m, t))| &**n == m.as_str() && self.conforms(x, t))
1826 }
1827 _ => false,
1828 }
1829 }
1830
1831 fn call(&mut self, func: FnId, key: Vec<Value>, span: Span) -> Result<Arc<CallResult>> {
1845 self.stats.calls += 1;
1846 let prog = self.prog;
1847 let fun = &prog.functions[func as usize];
1848 let sampling = self.sampler.is_some();
1850 let memoizable = self.config.memoizing && !sampling && !fun.effects.prints && self.callback.is_none();
1853 let call_key = (func, key);
1854 if memoizable {
1855 if let Some(r) = self.memo.get(&call_key) {
1856 self.stats.memo_hits += 1;
1857 self.observed |= r.observed;
1858 return Ok(r.clone());
1859 }
1860 }
1861 if !sampling {
1862 if let Some(&depth) = self.active.get(&call_key) {
1863 if fun.effects.prints {
1864 return Err(RuntimeError::new(
1865 span,
1866 format!("`{}` comes back to itself, and prints", describe_call(fun, &call_key.1)),
1867 )
1868 .with_note("a call that comes back to itself runs several times, until its result settles, and would print each time")
1869 .with_help("remove the `print`, or print the result where the call is made"));
1870 }
1871 let so_far = match self.approx.get(&call_key) {
1872 Some(a) if self.frames_at.get(a.head) == Some(&a.frame) => a.result.clone(),
1873 _ => {
1874 let result = Arc::new(CallResult::waiting(depth));
1875 let approx = Approx {
1876 result: result.clone(),
1877 head: depth,
1878 frame: self.frames_at[depth],
1879 round: self.rounds_at[depth],
1880 };
1881 self.approx.insert(call_key, approx);
1882 result
1883 }
1884 };
1885 return Ok(so_far);
1886 }
1887 if let Some(a) = self.approx.get(&call_key) {
1889 if self.frames_at.get(a.head) == Some(&a.frame) && self.rounds_at.get(a.head) == Some(&a.round) {
1890 self.stats.memo_hits += 1;
1891 return Ok(a.result.clone());
1892 }
1893 }
1894 }
1895 if self.depth >= self.config.max_call_depth {
1896 return Err(RuntimeError::limit(
1897 span,
1898 format!("calls are nested more than {} deep", self.config.max_call_depth),
1899 )
1900 .with_help("check for recursion that doesn't stop, or write it as a loop"));
1901 }
1902 let mut slots = vec![Value::Dead; fun.n_slots()];
1903 let n = fun.n_params as usize;
1904 for (i, v) in call_key.1[..n].iter().enumerate() {
1905 slots[i] = v.clone();
1906 }
1907 for (cap, v) in fun.captures.iter().zip(&call_key.1[n..]) {
1908 slots[cap.slot as usize] = v.clone();
1909 }
1910 let depth = self.depth + 1;
1911 if !sampling {
1912 self.active.insert(call_key.clone(), depth);
1913 if self.frames_at.len() <= depth {
1914 self.frames_at.resize(depth + 1, 0);
1915 self.rounds_at.resize(depth + 1, 0);
1916 }
1917 self.last_number += 1;
1918 self.frames_at[depth] = self.last_number;
1919 }
1920 let result = self.run_call(func, &call_key, slots, depth, span);
1921 if !sampling {
1922 self.active.remove(&call_key);
1923 self.frames_at[depth] = 0;
1925 }
1926 let result = result?;
1927 if memoizable
1929 && result.pending.is_empty()
1930 && result
1931 .outcomes
1932 .iter()
1933 .all(|(v, _, c)| c.is_empty() && !analytic::contains(v))
1934 {
1935 if self.memo.len() >= self.config.max_cached_calls {
1936 self.memo.clear();
1937 }
1938 self.memo.insert(call_key, result.clone());
1939 }
1940 Ok(result)
1941 }
1942
1943 fn run_call(
1953 &mut self,
1954 func: FnId,
1955 call_key: &CallKey,
1956 slots: Vec<Value>,
1957 depth: usize,
1958 span: Span,
1959 ) -> Result<Arc<CallResult>> {
1960 let fun = &self.prog.functions[func as usize];
1961 let sampling = self.sampler.is_some();
1962 let mut inherited = std::collections::BTreeSet::new();
1963 for v in &slots {
1964 analytic::collect_ids(v, &mut inherited);
1965 }
1966 let inherited = Arc::new(inherited);
1967 let mut rounds: u64 = 0;
1968 let mut before = f64::INFINITY;
1969 loop {
1970 if !sampling {
1971 self.last_number += 1;
1972 self.rounds_at[depth] = self.last_number;
1973 }
1974 let saved_unresolved = std::mem::replace(&mut self.unresolved, Weight::ZERO);
1975 let saved_lost = std::mem::replace(&mut self.lost, Weight::ZERO);
1976 let saved_observed = std::mem::replace(&mut self.observed, false);
1977 let saved_pending = std::mem::take(&mut self.pending);
1978 let saved_failures = std::mem::take(&mut self.failures);
1979 self.depth += 1;
1980 let flow = self.exec_block(
1981 func,
1982 &fun.body,
1983 vec![World {
1984 slots: slots.clone(),
1985 constraints: Default::default(),
1986 inherited: inherited.clone(),
1987 weight: Weight::ONE,
1988 run: 0,
1989 }],
1990 );
1991 self.depth -= 1;
1992 let unresolved = std::mem::replace(&mut self.unresolved, saved_unresolved);
1993 let lost = std::mem::replace(&mut self.lost, saved_lost);
1994 let observed = std::mem::replace(&mut self.observed, saved_observed);
1995 let pending = std::mem::replace(&mut self.pending, saved_pending);
1996 let failures = std::mem::replace(&mut self.failures, saved_failures);
1997 self.observed |= observed;
1998 let flow = flow.map_err(|e| match fun.kind {
1999 FnKind::Named => e.with_note(format!("in a call to `{}`", fun.name)),
2000 FnKind::Simulate => e.with_note("inside a `simulate` block"),
2001 FnKind::Lambda => e.with_note("inside a lambda"),
2002 FnKind::Main => e,
2003 })?;
2004 uncaught(&flow)?;
2005 let mut result = CallResult {
2006 outcomes: merge_values(flow.returned),
2007 unresolved,
2008 lost,
2009 observed,
2010 pending,
2011 failures,
2012 };
2013 if let Some(head) = result.pending.iter().map(|&(d, _)| d).filter(|&d| d < depth).min() {
2014 let waiting = Weight::sum(result.pending.iter().map(|&(_, w)| w));
2016 result.pending = vec![(head, waiting)];
2017 let result = Arc::new(result);
2018 let approx = Approx {
2019 result: result.clone(),
2020 head,
2021 frame: self.frames_at[head],
2022 round: self.rounds_at[head],
2023 };
2024 self.approx.insert(call_key.clone(), approx);
2025 return Ok(result);
2026 }
2027 let waiting = result.take_pending(depth);
2028 if waiting.is_zero() {
2029 return Ok(Arc::new(result));
2030 }
2031 if rounds == 0 {
2033 self.stats.solved_calls += 1;
2034 }
2035 let w = waiting.to_f64();
2036 if w <= self.config.epsilon {
2037 self.settle(depth);
2038 return Ok(Arc::new(result));
2039 }
2040 rounds += 1;
2041 self.stats.call_rounds += 1;
2042 if w >= before {
2043 return Err(RuntimeError::new(
2044 span,
2045 format!(
2046 "`{}` never returns for some of its worlds",
2047 describe_call(fun, &call_key.1)
2048 ),
2049 )
2050 .with_note(format!("{} of its weight keeps coming back to it", fmt_prob(w)))
2051 .with_help("check that the recursion can end"));
2052 }
2053 if rounds >= self.config.max_iterations {
2054 return Err(RuntimeError::limit(
2055 span,
2056 format!(
2057 "`{}` came back to itself {} times without settling",
2058 describe_call(fun, &call_key.1),
2059 self.config.max_iterations
2060 ),
2061 )
2062 .with_note(format!("{} of its weight is still waiting on it", fmt_prob(w)))
2063 .with_help("check that the recursion can end, or raise `@epsilon`"));
2064 }
2065 before = w;
2066 result.pending.push((depth, waiting));
2067 let approx = Approx {
2068 result: Arc::new(result),
2069 head: depth,
2070 frame: self.frames_at[depth],
2071 round: self.rounds_at[depth],
2072 };
2073 self.approx.insert(call_key.clone(), approx);
2074 }
2075 }
2076
2077 fn settle(&mut self, depth: usize) {
2081 let frame = self.frames_at[depth];
2082 let epsilon = self.config.epsilon;
2083 let settled: Vec<(CallKey, CallResult)> = self
2084 .approx
2085 .iter()
2086 .filter(|(_, a)| a.head == depth && a.frame == frame)
2087 .filter(|(_, a)| a.result.pending.iter().all(|&(_, w)| w.to_f64() <= epsilon))
2088 .map(|(key, a)| {
2089 let mut result = (*a.result).clone();
2090 result.pending.clear();
2092 (key.clone(), result)
2093 })
2094 .collect();
2095 self.approx.retain(|_, a| a.head != depth);
2096 if !self.config.memoizing {
2097 return;
2098 }
2099 for (key, result) in settled {
2100 if !self.active.contains_key(&key)
2101 && !self.prog.functions[key.0 as usize].effects.prints
2102 && result
2103 .outcomes
2104 .iter()
2105 .all(|(v, _, c)| c.is_empty() && !analytic::contains(v))
2106 {
2107 self.memo.insert(key, Arc::new(result));
2108 }
2109 }
2110 }
2111
2112 fn simulate(&mut self, func: FnId, key: Vec<Value>, span: Span) -> Result<Value> {
2115 if key.iter().any(analytic::contains) {
2116 return Err(analytic::unsupported("capturing an analytic draw inside `simulate`").at(span));
2117 }
2118 let saved_observed = self.observed;
2119 let saved_densities = self.densities;
2120 let sampler = self.sampler.take();
2122 let callback = self.callback.take();
2123 self.nested += 1;
2124 let result = self.call(func, key, span);
2125 self.nested -= 1;
2126 self.sampler = sampler;
2127 self.callback = callback;
2128 self.observed = saved_observed;
2130 self.densities = saved_densities;
2131 let result = result?;
2132 if !result.pending.is_empty() {
2133 return Err(OpError::unsupported(
2134 "a `simulate` block that comes back to a call still running isn't supported yet",
2135 )
2136 .at(span));
2137 }
2138 if let Some(g) = result.failures.groups.first() {
2141 return Err(g.error.clone());
2142 }
2143 let resolved = Weight::sum(result.outcomes.iter().map(|(_, w, _)| *w));
2144 if resolved.is_zero() {
2145 let message = if result.unresolved.is_zero() {
2146 "every world in this `simulate` was ruled out by `observe`"
2147 } else {
2148 "this `simulate` left all of its weight unresolved"
2149 };
2150 return Err(RuntimeError::new(span, message));
2151 }
2152 let denom = resolved + result.unresolved;
2153 let pairs = result
2154 .outcomes
2155 .iter()
2156 .map(|(v, w, _)| (v.clone(), w.ratio(denom)))
2157 .collect();
2158 ops::combine(pairs, result.unresolved.ratio(denom), &mut self.budget).map_err(|e| e.at(span))
2159 }
2160
2161 fn call_pure(&mut self, closure: &Value, args: Vec<Value>, what: &'static str, span: Span) -> Result<Value> {
2163 if let Value::Builtin(b) = closure {
2164 return self.builtin_values(*b, &args, &[], Weight::ONE, span);
2165 }
2166 let Value::Closure(c) = closure else {
2167 return Err(RuntimeError::new(
2168 span,
2169 format!(
2170 "`{what}` needs a function, like `x -> x * 2`, found {}",
2171 ops::article(&closure.kind())
2172 ),
2173 ));
2174 };
2175 self.check_arity(c, args.len(), span)?;
2176 let mut key = args;
2177 key.extend(c.captured.iter().cloned());
2178 let previous = self.callback.replace(what);
2179 let result = self.call(c.func, key, span);
2180 self.callback = previous;
2181 let result = result?;
2182 if !result.pending.is_empty() {
2183 return Err(OpError::unsupported(format!(
2184 "a function given to `{what}` that comes back to a call still running isn't supported yet"
2185 ))
2186 .at(span));
2187 }
2188 if let Some(g) = result.failures.groups.first() {
2191 return Err(g.error.clone());
2192 }
2193 match result.outcomes.as_slice() {
2194 [(v, p, c)] if c.is_empty() && (p.to_f64() - 1.0).abs() < 1e-12 && result.unresolved.is_zero() => {
2195 Ok(v.clone())
2196 }
2197 _ => Err(RuntimeError::new(
2198 span,
2199 format!("the function given to `{what}` can't branch on chances, draw values or observe"),
2200 )),
2201 }
2202 }
2203
2204 fn check_callback_effect(&self, span: Span) -> Result<()> {
2207 match self.callback {
2208 Some(what) => Err(RuntimeError::new(
2209 span,
2210 format!("the function given to `{what}` can't branch on chances, draw values or observe"),
2211 )
2212 .with_help("use a loop for probabilistic traversal, or `simulate` for a local distribution")),
2213 None => Ok(()),
2214 }
2215 }
2216
2217 fn record_place_type(&self, place: &Place, w: &World) -> Option<TypeSpec> {
2218 if place.path.is_empty() {
2219 return None;
2220 }
2221 let Value::Record(r) = &w.slots[place.slot as usize] else {
2222 return None;
2223 };
2224 let id = self
2225 .prog
2226 .records
2227 .iter()
2228 .position(|t| Some(t.name.as_str()) == r.ty.as_deref())?;
2229 let mut ty = TypeSpec::Record(id as u32);
2230 for element in &place.path {
2231 ty = match (element, ty) {
2232 (PathElem::Field(name), TypeSpec::Record(id)) => self.prog.records[id as usize]
2233 .fields
2234 .iter()
2235 .find(|f| f.name == *name)?
2236 .ty
2237 .clone(),
2238 (PathElem::Field(name), TypeSpec::AnonRecord(fields)) => fields.into_iter().find(|(n, _)| n == name)?.1,
2239 (PathElem::Index(_), TypeSpec::List(t) | TypeSpec::Map(_, t)) => *t,
2240 _ => return None,
2241 };
2242 }
2243 Some(ty)
2244 }
2245
2246 fn assign(&mut self, f: FnId, place: &Place, v: Value, w: &mut World, span: Span) -> Result<()> {
2247 if place.path.is_empty() {
2248 w.slots[place.slot as usize] = v;
2249 return Ok(());
2250 }
2251 let v = match self.record_place_type(place, w) {
2252 Some(ty) => {
2253 let v = self.coerce(v, &ty, span)?;
2254 if !self.conforms(&v, &ty) {
2255 return Err(RuntimeError::new(
2256 span,
2257 format!(
2258 "assigned field should be {}, found {}",
2259 ty.describe(self.prog),
2260 v.kind()
2261 ),
2262 ));
2263 }
2264 v
2265 }
2266 None => v,
2267 };
2268 let mut keys = Vec::with_capacity(place.path.len());
2269 for elem in &place.path {
2270 keys.push(match elem {
2271 PathElem::Field(name) => PathKey::Field(name.as_str()),
2272 PathElem::Index(e) => {
2273 let key = self.eval(f, e, w)?;
2274 if analytic::contains(&key) {
2275 return Err(analytic::unsupported("an analytic assignment index").at(e.span));
2276 }
2277 PathKey::Index(key)
2278 }
2279 });
2280 }
2281 let target = &mut w.slots[place.slot as usize];
2282 if matches!(target, Value::Dead) {
2283 return Err(self.no_value(f, place.slot, span));
2284 }
2285 update(target, &keys, v, &mut self.budget).map_err(|e| e.at(span))
2286 }
2287
2288 fn read_place(&mut self, f: FnId, place: &Place, w: &World, span: Span) -> Result<Value> {
2289 let mut v = self.slot(f, place.slot, w, span)?;
2290 for elem in &place.path {
2291 v = match elem {
2292 PathElem::Field(name) => ops::field(&v, name, &mut self.budget),
2293 PathElem::Index(e) => {
2294 let i = self.eval(f, e, w)?;
2295 ops::index(&v, &i, &mut self.budget)
2296 }
2297 }
2298 .map_err(|e| e.at(span))?;
2299 }
2300 Ok(v)
2301 }
2302
2303 fn slot(&mut self, f: FnId, s: SlotId, w: &World, span: Span) -> Result<Value> {
2304 match &w.slots[s as usize] {
2305 Value::Dead => Err(self.no_value(f, s, span)),
2306 Value::Delayed(_) => {
2307 let name = &self.prog.functions[f as usize].slots[s as usize].name;
2308 Err(
2309 OpError::internal("internal error: a variable was read before it was drawn")
2310 .at(span)
2311 .with_note(format!(
2312 "`{name}`'s draw was delayed for an exact update (docs/semantics.md, section 14)"
2313 )),
2314 )
2315 }
2316 v => analytic::resolve(v, &w.constraints, &mut self.budget).map_err(|e| e.at(span)),
2317 }
2318 }
2319
2320 fn no_value(&self, f: FnId, s: SlotId, span: Span) -> RuntimeError {
2321 let name = &self.prog.functions[f as usize].slots[s as usize].name;
2322 RuntimeError::new(span, format!("`{name}` has no value here")).with_help(
2323 "it's used before it's given a value; if a function reads it, call the function after the variable is set",
2324 )
2325 }
2326
2327 fn eval_expected(&mut self, f: FnId, e: &Expr, w: &World, ty: &TypeSpec) -> Result<Value> {
2330 match (&e.kind, ty) {
2331 (ExprKind::List(xs), TypeSpec::List(t)) => {
2332 self.budget.collection(xs.len() as u128).map_err(|err| err.at(e.span))?;
2333 let mut out = Vec::with_capacity(xs.len());
2334 for x in xs {
2335 out.push(self.eval_expected(f, x, w, t)?);
2336 }
2337 Ok(Value::list(out))
2338 }
2339 (ExprKind::Map(xs), TypeSpec::Map(kt, vt)) => {
2340 let mut out = BTreeMap::new();
2341 for (k, v) in xs {
2342 let key = self.eval_expected(f, k, w, kt)?;
2343 if analytic::contains(&key) {
2344 return Err(analytic::unsupported("analytic map keys").at(k.span));
2345 }
2346 if key.is_uncertain() {
2347 return Err(RuntimeError::new(k.span, "map keys can't be distributions"));
2348 }
2349 out.insert(key, self.eval_expected(f, v, w, vt)?);
2350 }
2351 Ok(Value::map(out))
2352 }
2353 (ExprKind::Record { ty: None, fields }, TypeSpec::AnonRecord(types)) => {
2354 let mut out = Vec::with_capacity(fields.len());
2355 for (n, e) in fields {
2356 let v = match types.iter().find(|(name, _)| name == n) {
2357 Some((_, ty)) => self.eval_expected(f, e, w, ty)?,
2358 None => self.eval(f, e, w)?,
2359 };
2360 out.push((Arc::from(n.as_str()), v));
2361 }
2362 Ok(ops::make_record(None, out))
2363 }
2364 _ => {
2365 let v = self.eval(f, e, w)?;
2366 self.coerce(v, ty, e.span)
2367 }
2368 }
2369 }
2370
2371 fn record_field_type(&self, value: &Value, name: &str) -> Option<TypeSpec> {
2375 match value {
2376 Value::Record(r) => self
2377 .prog
2378 .records
2379 .iter()
2380 .find(|t| Some(t.name.as_str()) == r.ty.as_deref())?
2381 .fields
2382 .iter()
2383 .find(|f| f.name == name)
2384 .map(|f| f.ty.clone()),
2385 Value::Dist(d) => {
2386 let first = self.record_field_type(&d.outcomes.first()?.0, name)?;
2387 d.outcomes
2388 .iter()
2389 .all(|(v, _)| self.record_field_type(v, name).as_ref() == Some(&first))
2390 .then_some(first)
2391 }
2392 _ => None,
2393 }
2394 }
2395
2396 fn check_updated_records(&mut self, value: Value, span: Span) -> Result<Value> {
2397 match value {
2398 Value::Record(ref r) => {
2399 let Some(id) = self
2400 .prog
2401 .records
2402 .iter()
2403 .position(|t| Some(t.name.as_str()) == r.ty.as_deref())
2404 else {
2405 return Ok(value);
2406 };
2407 let ty = TypeSpec::Record(id as u32);
2408 let value = self.coerce(value, &ty, span)?;
2409 if !self.conforms(&value, &ty) {
2410 return Err(RuntimeError::new(
2411 span,
2412 format!("updated record doesn't conform to {}", ty.describe(self.prog)),
2413 ));
2414 }
2415 Ok(value)
2416 }
2417 Value::Dist(d) => {
2418 let mut outcomes = Vec::with_capacity(d.outcomes.len());
2419 for (v, p) in &d.outcomes {
2420 outcomes.push((self.check_updated_records(v.clone(), span)?, *p));
2421 }
2422 ops::combine(outcomes, d.missing, &mut self.budget).map_err(|e| e.at(span))
2423 }
2424 v => Ok(v),
2425 }
2426 }
2427
2428 fn eval(&mut self, f: FnId, e: &Expr, w: &World) -> Result<Value> {
2429 let span = e.span;
2430 let at = |err: OpError| err.at(span);
2431 match &e.kind {
2432 ExprKind::Lit(l) => self.literal(l).map_err(at),
2433 ExprKind::Slot(s) => self.slot(f, *s, w, span),
2434 ExprKind::Unary(op, x) => {
2435 let v = self.eval(f, x, w)?;
2436 ops::unary(*op, &v, &mut self.budget).map_err(at)
2437 }
2438 ExprKind::Binary(op @ (BinOp::And | BinOp::Or), a, b) => {
2439 let and = *op == BinOp::And;
2440 let word = if and { "and" } else { "or" };
2441 let va = self.eval(f, a, w)?;
2442 let ta = ops::truth(&va, word).map_err(|err| err.at(a.span))?;
2443 if let Truth::Fact(x) = ta {
2444 if x != and {
2445 return Ok(Value::Bool(x));
2447 }
2448 }
2449 let vb = self.eval(f, b, w)?;
2450 let tb = ops::truth(&vb, word).map_err(|err| err.at(b.span))?;
2451 ops::logic(and, ta, tb, &mut self.budget).map_err(at)
2452 }
2453 ExprKind::Binary(op, a, b) => {
2454 let va = self.eval(f, a, w)?;
2455 let vb = self.eval(f, b, w)?;
2456 ops::binary(*op, &va, &vb, &mut self.budget).map_err(at)
2457 }
2458 ExprKind::List(items) => {
2459 self.budget.collection(items.len() as u128).map_err(at)?;
2460 let mut values = Vec::with_capacity(items.len());
2461 for item in items {
2462 values.push(self.eval(f, item, w)?);
2463 }
2464 Ok(Value::list(values))
2465 }
2466 ExprKind::Map(entries) => {
2467 let mut map = BTreeMap::new();
2468 for (k, v) in entries {
2469 let key = self.eval(f, k, w)?;
2470 if analytic::contains(&key) {
2471 return Err(analytic::unsupported("analytic map keys").at(k.span));
2472 }
2473 if key.is_uncertain() {
2474 return Err(RuntimeError::new(k.span, "map keys can't be distributions"));
2475 }
2476 let value = self.eval(f, v, w)?;
2477 map.insert(key, value);
2478 }
2479 Ok(Value::map(map))
2480 }
2481 ExprKind::Record { ty, fields } => {
2482 let mut values = Vec::with_capacity(fields.len());
2483 for (name, v) in fields {
2484 let field_ty = ty
2485 .and_then(|t| self.prog.records[t as usize].fields.iter().find(|d| d.name == *name))
2486 .map(|d| d.ty.clone());
2487 let value = match field_ty {
2488 Some(t) => {
2489 let value = self.eval(f, v, w)?;
2490 let value = self.coerce(value, &t, v.span)?;
2491 if !self.conforms(&value, &t) {
2492 return Err(RuntimeError::new(
2493 v.span,
2494 format!(
2495 "field `{name}` should be {}, found {}",
2496 t.describe(self.prog),
2497 value.kind()
2498 ),
2499 ));
2500 }
2501 value
2502 }
2503 None => self.eval(f, v, w)?,
2504 };
2505 values.push((Arc::from(name.as_str()), value));
2506 }
2507 Ok(ops::make_record(
2508 ty.map(|t| self.record_names[t as usize].clone()),
2509 values,
2510 ))
2511 }
2512 ExprKind::Field(x, name) => {
2513 let v = self.eval(f, x, w)?;
2514 ops::field(&v, name, &mut self.budget).map_err(at)
2515 }
2516 ExprKind::Index(x, i) => {
2517 let v = self.eval(f, x, w)?;
2518 let i = self.eval(f, i, w)?;
2519 ops::index(&v, &i, &mut self.budget).map_err(at)
2520 }
2521 ExprKind::With(x, fields) => {
2522 let base = self.eval(f, x, w)?;
2523 let mut updates = Vec::with_capacity(fields.len());
2524 for (name, v) in fields {
2525 let field_ty = self.record_field_type(&base, name);
2526 let value = match field_ty {
2527 Some(ty) => {
2528 let value = self.eval_expected(f, v, w, &ty)?;
2529 if !self.conforms(&value, &ty) {
2530 return Err(RuntimeError::new(
2531 v.span,
2532 format!(
2533 "field `{name}` should be {}, found {}",
2534 ty.describe(self.prog),
2535 value.kind()
2536 ),
2537 ));
2538 }
2539 value
2540 }
2541 None => self.eval(f, v, w)?,
2542 };
2543 updates.push((Arc::from(name.as_str()), value));
2544 }
2545 let value = ops::lift1(&base, &mut self.budget, |b, _| ops::with_fields(b, &updates)).map_err(at)?;
2546 self.check_updated_records(value, span)
2547 }
2548 ExprKind::Builtin { func, args, named } => self.builtin(f, *func, args, named, w, span),
2549 ExprKind::Closure { func, capture_args } => Ok(Value::Closure(Arc::new(Closure {
2550 func: *func,
2551 captured: capture_args
2552 .iter()
2553 .map(|&s| self.slot(f, s, w, span))
2554 .collect::<Result<_>>()?,
2555 }))),
2556 ExprKind::Simulate { func, capture_args } => {
2557 let key = capture_args
2558 .iter()
2559 .map(|&s| self.slot(f, s, w, span))
2560 .collect::<Result<_>>()?;
2561 self.simulate(*func, key, span)
2562 }
2563 ExprKind::Interp(parts) => {
2564 let mut text = String::new();
2565 for part in parts {
2566 match part {
2567 InterpPart::Lit(s) => crate::text::push(&mut text, s, &mut self.budget).map_err(at)?,
2568 InterpPart::Expr(e) => {
2569 let value = self.eval(f, e, w)?;
2570 crate::text::push_value(&mut text, &value, &mut self.budget).map_err(at)?;
2571 }
2572 }
2573 }
2574 Ok(Value::str(&text))
2575 }
2576 ExprKind::Input(i) => Ok(self.inputs[*i as usize].clone()),
2577 }
2578 }
2579
2580 fn literal(&mut self, l: &Lit) -> OpResult<Value> {
2581 Ok(match l {
2582 Lit::Builtin(b) => Value::Builtin(*b),
2583 Lit::Unit => Value::Unit,
2584 Lit::Bool(b) => Value::Bool(*b),
2585 Lit::Int(i) => {
2586 self.budget.integer_bits(i.bits())?;
2587 self.budget.work(i.bits().div_ceil(64).max(1))?;
2588 Value::Int(i.clone())
2589 }
2590 Lit::Float(x) | Lit::FloatConstant(x) => Value::Float(*x),
2591 Lit::Prob(p) => Value::Prob(*p),
2592 Lit::Str(s) => crate::text::value(s, &mut self.budget)?,
2593 Lit::Dice { count, sides } => {
2594 if let Some(d) = self.dice.get(&(*count, *sides)) {
2595 return Ok(d.clone());
2596 }
2597 let d = Dist::dice(*count, *sides, &mut self.budget)?.into_value();
2598 self.dice.insert((*count, *sides), d.clone());
2599 d
2600 }
2601 Lit::Enum { ty, variant } => self.enum_values[*ty as usize][*variant as usize].clone(),
2602 })
2603 }
2604
2605 fn builtin(
2606 &mut self,
2607 f: FnId,
2608 b: Builtin,
2609 args: &[Expr],
2610 named: &[(String, Expr)],
2611 w: &World,
2612 span: Span,
2613 ) -> Result<Value> {
2614 if b == Builtin::Typeof {
2615 let v = match args[0].kind {
2618 ExprKind::Slot(s) if matches!(w.slots[s as usize], Value::Delayed(_)) => w.slots[s as usize].clone(),
2619 _ => self.eval(f, &args[0], w)?,
2620 };
2621 return crate::type_name::of(&v, &self.prog.enums, &mut self.budget).map_err(|e| e.at(span));
2622 }
2623 let mut values = Vec::with_capacity(args.len());
2624 for a in args {
2625 values.push(self.eval(f, a, w)?);
2626 }
2627 let mut names = Vec::with_capacity(named.len());
2628 for (name, value) in named {
2629 names.push(name.clone());
2630 values.push(self.eval(f, value, w)?);
2631 }
2632 self.builtin_values(b, &values, &names, w.weight, span)
2633 }
2634
2635 fn check_function_arity(&self, value: &Value, given: usize, span: Span) -> Result<()> {
2636 match value {
2637 Value::Closure(c) => self.check_arity(c, given, span),
2638 Value::Builtin(b) => {
2639 let (min, max) = b.arity();
2640 if given < min || given > max || (*b == Builtin::Date && given == 2) {
2641 return Err(RuntimeError::new(
2642 span,
2643 format!(
2644 "`{}` takes {}, but {given} were given",
2645 b.name(),
2646 if *b == Builtin::Date {
2647 "one ISO string or three integers".into()
2648 } else if min == max {
2649 format!("{min} argument(s)")
2650 } else if max == usize::MAX {
2651 format!("at least {min} arguments")
2652 } else {
2653 format!("{min} to {max} arguments")
2654 }
2655 ),
2656 ));
2657 }
2658 Ok(())
2659 }
2660 _ => Err(RuntimeError::new(
2661 span,
2662 "expected a comparator function, like `(a, b) -> a - b`",
2663 )),
2664 }
2665 }
2666
2667 fn builtin_values(
2669 &mut self,
2670 b: Builtin,
2671 values: &[Value],
2672 named: &[String],
2673 weight: Weight,
2674 span: Span,
2675 ) -> Result<Value> {
2676 let extrema = matches!(b, Builtin::Minimum | Builtin::Maximum);
2677 if !named.is_empty() && (!extrema || named != ["default"]) {
2678 return Err(RuntimeError::new(
2679 span,
2680 format!(
2681 "`{}` does not accept these named arguments; only minimum/maximum support `default`",
2682 b.name()
2683 ),
2684 ));
2685 }
2686 let positional = values.len() - named.len();
2687 if !named.is_empty() && positional == 0 {
2688 return Err(RuntimeError::new(
2689 span,
2690 "a collection argument is required before `default`",
2691 ));
2692 }
2693 if !named.is_empty() && positional >= 3 {
2694 return Err(RuntimeError::new(
2695 span,
2696 "the default was supplied both positionally and by name",
2697 ));
2698 }
2699 self.check_function_arity(&Value::Builtin(b), values.len(), span)?;
2700 let (values, default) = if extrema && (!named.is_empty() || values.len() == 3) {
2701 (&values[..values.len() - 1], values.last())
2702 } else {
2703 (values, None)
2704 };
2705 let at = |err: OpError| err.at(span);
2706 builtins::check_query_input(b, values).map_err(at)?;
2707 if values.iter().any(analytic::contains)
2708 && !matches!(
2709 b,
2710 Builtin::Map
2711 | Builtin::Filter
2712 | Builtin::Reduce
2713 | Builtin::Len
2714 | Builtin::IterItems
2715 | Builtin::Settled
2716 | Builtin::BooleanLaw
2717 )
2718 {
2719 return Err(analytic::unsupported(&format!("`{}` on this outcome", b.name())).at(span));
2720 }
2721 match b {
2722 Builtin::Typeof => crate::type_name::of(&values[0], &self.prog.enums, &mut self.budget).map_err(at),
2723 Builtin::RunDate => self.config.today.map(Value::Date).ok_or_else(|| {
2724 RuntimeError::new(span, "the host did not supply an execution date for `today`")
2725 .with_help("set Options.today, or supply today in the WASM request")
2726 }),
2727 Builtin::Print => {
2728 let mut text = String::new();
2729 for (i, value) in values.iter().enumerate() {
2730 if i != 0 {
2731 crate::text::push(&mut text, " ", &mut self.budget).map_err(at)?;
2732 }
2733 crate::text::push_value(&mut text, value, &mut self.budget).map_err(at)?;
2734 }
2735 let line = if weight == Weight::ONE {
2736 text
2737 } else {
2738 format!("[{}] {text}", fmt_weight(weight))
2739 };
2740 self.printed += line.len() + 1;
2741 if self.printed > self.config.max_output {
2742 return Err(too_much_output(span));
2743 }
2744 match &mut self.lines {
2745 Some(lines) => lines.push((span, line)),
2746 None => (self.print)(&line),
2747 }
2748 Ok(Value::Unit)
2749 }
2750 Builtin::Map | Builtin::Filter | Builtin::Reduce => self.higher_order(b, values, span),
2751 Builtin::Count if values.len() == 2 => self.higher_order(b, values, span),
2752 Builtin::Minimum | Builtin::Maximum => {
2753 if values.len() == 2 {
2754 self.collection_extreme(b, values, default, span)
2755 } else {
2756 builtins::population_extreme(&values[0], b == Builtin::Maximum, default, &mut self.budget)
2757 .map_err(at)
2758 }
2759 }
2760 Builtin::Sort | Builtin::SortDesc if values.len() == 2 => self.higher_order(b, values, span),
2761 Builtin::Highest | Builtin::Lowest if values.len() == 3 => self.higher_order(b, values, span),
2762 Builtin::Min | Builtin::Max => builtins::call_plain(b, values, &mut self.budget).map_err(at),
2763 Builtin::Roll => self.roll(values).map_err(at),
2764 Builtin::Take => Err(RuntimeError::new(
2765 span,
2766 "use `deck.take()` to draw and remove an item from a mutable bag",
2767 )),
2768 _ if b.lifting() == Lifting::Raw => builtins::call_raw(b, values, &mut self.budget).map_err(at),
2769 _ if values.iter().any(|v| matches!(v, Value::Continuous(_))) => {
2770 let kind = values
2771 .iter()
2772 .find(|v| matches!(v, Value::Continuous(_)))
2773 .unwrap()
2774 .kind();
2775 Err(
2776 RuntimeError::new(span, format!("`{}` needs a value, not a {kind}", b.name()))
2777 .with_help("draw a value first, like `let x ~ normal(0, 1)`"),
2778 )
2779 }
2780 _ => ops::lift_n(values, &mut self.budget, &|a, budget| {
2781 builtins::call_plain(b, a, budget)
2782 })
2783 .map_err(at),
2784 }
2785 }
2786
2787 fn roll(&mut self, values: &[Value]) -> OpResult<Value> {
2788 if values[0].is_uncertain() {
2789 return Err(OpError::new("the number of dice to roll must be a plain number")
2790 .help("draw it first, like `let n ~ d4`"));
2791 }
2792 let count = ops::integer(&values[0], "roll count", &mut self.budget)?;
2793 let count = count
2794 .to_u64()
2795 .filter(|n| *n <= 1000)
2796 .ok_or_else(|| OpError::new("roll needs between 0 and 1000 dice"))? as u32;
2797 let die = match &values[1] {
2798 v @ (Value::Int(_) | Value::Float(_)) => {
2799 let sides = ops::integer(v, "roll sides", &mut self.budget)?;
2800 let sides = sides
2801 .to_u64()
2802 .filter(|n| *n >= 1 && *n <= u32::MAX as u64)
2803 .ok_or_else(|| OpError::new("roll needs between 1 and 4294967295 sides"))?;
2804 Dist::dice(1, sides as u32, &mut self.budget)?
2805 }
2806 Value::Dist(d) => (**d).clone(),
2807 v => {
2808 return Err(OpError::new(format!(
2809 "roll needs a die, like `d6`, found {}",
2810 ops::article(&v.kind())
2811 )));
2812 }
2813 };
2814 let key = (count, die.clone().into_value());
2815 if let Some(pool) = self.pools.get(&key) {
2816 return Ok(pool.clone());
2817 }
2818 let pool = Dist::pool(count, &die, &mut self.budget)?.into_value();
2819 self.pools.insert(key, pool.clone());
2820 Ok(pool)
2821 }
2822
2823 fn compare_callback(
2824 &mut self,
2825 comparator: &Value,
2826 a: &Value,
2827 other: &Value,
2828 b: Builtin,
2829 span: Span,
2830 ) -> Result<std::cmp::Ordering> {
2831 self.budget.work(1).map_err(|e| e.at(span))?;
2832 match self.call_pure(comparator, vec![a.clone(), other.clone()], b.name(), span)? {
2833 Value::Int(n) => Ok(n.cmp(&probl_number::Integer::ZERO)),
2834 Value::Float(x) if x.is_finite() => Ok(x.partial_cmp(&0.0).expect("finite comparator result")),
2835 other => Err(RuntimeError::new(
2836 span,
2837 format!(
2838 "the comparator given to `{}` must return a finite int or float (negative, zero, or positive), found {}",
2839 b.name(),
2840 ops::article(&other.kind())
2841 ),
2842 )),
2843 }
2844 }
2845
2846 fn collection_extreme(
2847 &mut self,
2848 b: Builtin,
2849 values: &[Value],
2850 default: Option<&Value>,
2851 span: Span,
2852 ) -> Result<Value> {
2853 if matches!(values[0], Value::Dist(_) | Value::Continuous(_)) {
2854 return Err(RuntimeError::new(
2855 span,
2856 "a comparator is supported only for collection extrema, not distribution support bounds",
2857 ));
2858 }
2859 let comparator = &values[1];
2860 self.check_function_arity(comparator, 2, span)?;
2861 let items = builtins::items(&values[0], b.name(), &mut self.budget).map_err(|e| e.at(span))?;
2862 let mut iter = items.into_iter();
2863 let Some(mut best) = iter.next() else {
2864 return builtins::empty_extreme(b.name(), default).map_err(|e| e.at(span));
2865 };
2866 for x in iter {
2867 let order = self.compare_callback(comparator, &x, &best, b, span)?;
2868 if (b == Builtin::Maximum && order.is_gt()) || (b == Builtin::Minimum && order.is_lt()) {
2869 best = x;
2870 }
2871 }
2872 Ok(best)
2873 }
2874
2875 fn higher_order(&mut self, b: Builtin, values: &[Value], span: Span) -> Result<Value> {
2877 if let Value::Dist(d) = &values[0] {
2878 let mut results = Vec::with_capacity(d.outcomes.len());
2879 for (coll, p) in &d.outcomes {
2880 let mut args = values.to_vec();
2881 args[0] = coll.clone();
2882 results.push((self.higher_order(b, &args, span)?, *p));
2883 }
2884 return ops::combine(results, d.missing, &mut self.budget).map_err(|e| e.at(span));
2885 }
2886 if matches!(b, Builtin::Highest | Builtin::Lowest) {
2887 if let Value::Dist(d) = &values[1] {
2888 let mut results = Vec::with_capacity(d.outcomes.len());
2889 for (count, p) in &d.outcomes {
2890 let mut args = values.to_vec();
2891 args[1] = count.clone();
2892 results.push((self.higher_order(b, &args, span)?, *p));
2893 }
2894 return ops::combine(results, d.missing, &mut self.budget).map_err(|e| e.at(span));
2895 }
2896 }
2897 let items = builtins::items(&values[0], b.name(), &mut self.budget).map_err(|e| e.at(span))?;
2898 match b {
2899 Builtin::Sort | Builtin::SortDesc | Builtin::Highest | Builtin::Lowest => {
2900 let comparator = values.last().unwrap();
2901 self.check_function_arity(comparator, 2, span)?;
2903 let count = if matches!(b, Builtin::Highest | Builtin::Lowest) {
2904 Some(builtins::extreme_count(&values[1], &mut self.budget).map_err(|e| e.at(span))?)
2905 } else {
2906 None
2907 };
2908 crate::ordering::reserve_sort(items.len(), &mut self.budget).map_err(|e| e.at(span))?;
2909 let mut items = items;
2910 crate::ordering::try_sort_by(&mut items, |a, other| {
2911 let order = self.compare_callback(comparator, a, other, b, span)?;
2912 Ok(if matches!(b, Builtin::SortDesc | Builtin::Highest) {
2913 order.reverse()
2914 } else {
2915 order
2916 })
2917 })?;
2918 if let Some(n) = count {
2919 items.truncate(n);
2920 }
2921 Ok(Value::list(items))
2922 }
2923 Builtin::Map => {
2924 let mut out = Vec::with_capacity(items.len());
2925 for x in items {
2926 out.push(self.call_pure(&values[1], vec![x], "map", span)?);
2927 }
2928 Ok(Value::list(out))
2929 }
2930 Builtin::Filter | Builtin::Count => {
2931 let mut kept = Vec::new();
2932 for x in items {
2933 let test = self.call_pure(&values[1], vec![x.clone()], b.name(), span)?;
2934 match test {
2935 Value::Bool(true) => kept.push(x),
2936 Value::Bool(false) => {}
2937 other => {
2938 return Err(RuntimeError::new(
2939 span,
2940 format!("the test given to `{}` must give a fact (true or false)", b.name()),
2941 )
2942 .with_help(format!("it gave {}", ops::article(&other.kind()))));
2943 }
2944 }
2945 }
2946 Ok(if b == Builtin::Count {
2947 Value::Int((kept.len() as i64).into())
2948 } else {
2949 Value::list(kept)
2950 })
2951 }
2952 Builtin::Reduce => {
2953 let callback = &values[1];
2954 if !matches!(callback, Value::Closure(_) | Value::Builtin(_)) {
2955 return Err(
2956 RuntimeError::new(span, "`reduce` needs a function as its second argument")
2957 .with_help("write `reduce(xs, f)` or `reduce(xs, f, initial)`"),
2958 );
2959 }
2960 self.check_function_arity(callback, 2, span)?;
2961 let mut items = items.into_iter();
2962 let mut acc = match values.get(2) {
2963 Some(initial) => initial.clone(),
2964 None => items.next().ok_or_else(|| {
2965 RuntimeError::new(span, "`reduce` of an empty collection needs an initial value")
2966 .with_help("supply an initial value: `reduce(xs, f, initial)`")
2967 })?,
2968 };
2969 for x in items {
2970 acc = self.call_pure(callback, vec![acc, x], "reduce", span)?;
2971 }
2972 Ok(acc)
2973 }
2974 _ => unreachable!(),
2975 }
2976 }
2977}
2978
2979fn describe_call(fun: &Function, key: &[Value]) -> String {
2981 let args: Vec<String> = key[..fun.n_params as usize].iter().map(|v| format!("{v:?}")).collect();
2982 format!("{}({})", fun.name, args.join(", "))
2983}
2984
2985fn uncaught(flow: &Flow) -> Result<()> {
2989 match flow.faulted.first() {
2990 Some((_, e)) => Err(OpError::internal("a fault went past the `try` meant to catch it").at(e.span)),
2991 None => Ok(()),
2992 }
2993}
2994
2995fn drop_failed(worlds: &mut Vec<World>, failed: &[usize]) {
2997 if failed.is_empty() {
2998 return;
2999 }
3000 let mut failed = failed.iter().copied().peekable();
3001 let mut i = 0;
3002 worlds.retain(|_| {
3003 let drop = failed.next_if_eq(&i).is_some();
3004 i += 1;
3005 !drop
3006 });
3007}
3008
3009pub fn too_much_output(span: Span) -> RuntimeError {
3010 RuntimeError::limit(span, "the program printed more than the output limit")
3011}
3012
3013pub fn fmt_weight(w: Weight) -> String {
3015 let x = w.to_f64();
3016 if x == 0.0 && !w.is_zero() {
3017 let l = w.log10() + 2.0;
3018 return format!("{:.1}e{}%", libm::pow(10.0, l - l.floor()), l.floor());
3019 }
3020 fmt_prob(x)
3021}
3022
3023enum PathKey<'a> {
3024 Field(&'a str),
3025 Index(Value),
3026}
3027
3028fn update(target: &mut Value, keys: &[PathKey], v: Value, budget: &mut Budget) -> OpResult<()> {
3029 let Some((first, rest)) = keys.split_first() else {
3030 *target = v;
3031 return Ok(());
3032 };
3033 match (first, target) {
3034 (PathKey::Field(name), Value::Record(r)) => {
3035 let r = Arc::make_mut(r);
3036 let slot = r
3037 .get_mut(name)
3038 .ok_or_else(|| OpError::new(format!("this record has no field `{name}`")))?;
3039 update(slot, rest, v, budget)
3040 }
3041 (PathKey::Index(i), Value::List(items)) => {
3042 let items = Arc::make_mut(items);
3043 let k = ops::as_index(i, items.len() as u128, budget)? as usize;
3044 update(&mut items[k], rest, v, budget)
3045 }
3046 (PathKey::Index(k), Value::Map(m)) => {
3047 let m = Arc::make_mut(m);
3048 if rest.is_empty() {
3049 m.insert(k.clone(), v);
3050 return Ok(());
3051 }
3052 let slot = m
3053 .get_mut(k)
3054 .ok_or_else(|| OpError::new(format!("the key {k:?} isn't in the map")))?;
3055 update(slot, rest, v, budget)
3056 }
3057 (PathKey::Field(name), other) => Err(OpError::new(format!(
3058 "can't set the field `{name}` of {}",
3059 ops::article(&other.kind())
3060 ))),
3061 (PathKey::Index(_), other) => Err(OpError::new(format!(
3062 "can't set an element of {}",
3063 ops::article(&other.kind())
3064 ))),
3065 }
3066}
3067
3068fn is_density(d: &Value) -> bool {
3070 match d {
3071 Value::Continuous(_) => true,
3072 Value::Dist(dist) => dist.outcomes.iter().any(|(x, _)| matches!(x, Value::Continuous(_))),
3073 _ => false,
3074 }
3075}
3076
3077fn likelihood(d: &Value, v: &Value, sampling: bool) -> OpResult<(f64, f64, f64)> {
3082 let continuous = match d {
3083 Value::Continuous(f) => Some(vec![(**f, 1.0)]),
3084 Value::Dist(dist) if dist.outcomes.iter().any(|(x, _)| matches!(x, Value::Continuous(_))) => {
3085 let parts: Option<Vec<_>> = dist
3086 .outcomes
3087 .iter()
3088 .map(|(x, p)| match x {
3089 Value::Continuous(f) => Some((**f, *p)),
3090 _ => None,
3091 })
3092 .collect();
3093 let parts = parts.ok_or_else(|| {
3094 OpError::new("`observe … from` can't mix a density with the probabilities of single values")
3095 })?;
3096 Some(parts)
3097 }
3098 _ => None,
3099 };
3100 if let Some(parts) = continuous {
3101 if !sampling {
3102 return Err(OpError::new("observing a value from a continuous distribution needs sample mode")
3103 .help("its density isn't a probability, so enumeration can't use it; sample with `@mode sample(runs: 10_000)`"));
3104 }
3105 let x = match v {
3106 Value::Bool(_) => None,
3107 _ => v.as_f64(),
3108 }
3109 .ok_or_else(|| {
3110 OpError::new(format!(
3111 "a continuous distribution can't produce {}",
3112 ops::article(&v.kind())
3113 ))
3114 })?;
3115 return Ok((parts.iter().map(|(f, p)| p * f.pdf(x)).sum(), 0.0, 0.0));
3116 }
3117 match d {
3118 Value::Dist(dist) => {
3119 let (mut seen, mut other) = (0.0, 0.0);
3120 for (x, p) in &dist.outcomes {
3121 if ops::equals(x, v) {
3122 seen += p;
3123 } else {
3124 other += p;
3125 }
3126 }
3127 Ok((seen, dist.missing, other))
3128 }
3129 Value::Prob(_) => Err(OpError::new("`observe … from` needs a distribution, not a probability")
3130 .help("to observe that a fact with probability p is true, write `observe true from bernoulli(p)`")),
3131 other => Ok(if ops::equals(other, v) {
3132 (1.0, 0.0, 0.0)
3133 } else {
3134 (0.0, 0.0, 1.0)
3135 }),
3136 }
3137}