1pub mod analytic;
5pub mod builtins;
6pub mod chain;
7pub mod complex;
8pub mod conjugate;
9pub mod continuous;
10pub mod data;
11pub mod dates;
12pub mod dist;
13pub mod error;
14pub mod interp;
15mod math;
16pub mod ops;
17mod ordering;
18pub mod report;
19mod stats;
20mod text;
21mod type_name;
22pub mod value;
23pub mod weight;
24pub mod world;
25
26pub use error::{ErrorKind, RuntimeError};
27pub use interp::{Stats, Updates};
28pub use weight::Weight;
29
30use continuous::Rng;
31use dist::Budget;
32use error::OpError;
33use probl_sema::Liveness;
34use probl_sema::conjugate::{Conjugacy, Variable};
35use probl_sema::ir::{Mode, Program};
36use probl_syntax::Span;
37use report::{Format, Sink};
38use std::any::Any;
39use std::collections::BTreeMap;
40use std::panic::AssertUnwindSafe;
41use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
42use std::sync::{Arc, Condvar, Mutex, PoisonError, mpsc};
43use value::Value;
44
45#[derive(Clone, Debug)]
48pub struct Limits {
49 pub max_integer_bits: u64,
51 pub max_integer_bytes: u64,
53 pub max_string_bytes: usize,
55 pub max_string_alloc_bytes: u64,
57 pub max_worlds: usize,
59 pub max_outcomes: usize,
61 pub max_collection: usize,
63 pub max_work: u64,
65 pub max_iterations: u64,
67 pub max_call_depth: usize,
69 pub max_cached_calls: usize,
71 pub max_output: usize,
73 pub stack_size: usize,
76 pub max_threads: usize,
79 pub max_chain_states: usize,
82}
83
84impl Default for Limits {
85 fn default() -> Limits {
86 Limits {
87 max_integer_bits: probl_number::MAX_INTEGER_BITS,
88 max_integer_bytes: 256 * 1024 * 1024,
89 max_string_bytes: 16 * 1024 * 1024,
90 max_string_alloc_bytes: 256 * 1024 * 1024,
91 max_worlds: 10_000_000,
92 max_outcomes: 2_000_000,
93 max_collection: 10_000_000,
94 max_work: 20_000_000_000,
95 max_iterations: 10_000_000,
96 max_call_depth: 500,
97 max_cached_calls: 1_000_000,
98 max_output: 64 * 1024 * 1024,
99 stack_size: 64 * 1024 * 1024,
100 max_threads: std::thread::available_parallelism().map_or(1, |n| n.get()),
101 max_chain_states: 50_000,
102 }
103 }
104}
105
106#[derive(Clone, Debug)]
109pub struct Options {
110 pub today: Option<i32>,
114 pub merge: bool,
117 pub memoize: bool,
119 pub epsilon: Option<f64>,
121 pub fractions: bool,
123 pub limits: Limits,
124 pub cancel: Option<Arc<AtomicBool>>,
126 pub mode: Option<Mode>,
128 pub inputs: Option<Arc<data::Inputs>>,
131 pub conjugate: bool,
135 pub solve: bool,
139 pub progress: Option<Progress>,
142}
143
144#[derive(Clone)]
146pub struct Progress(pub Arc<dyn Fn(u64, u64) + Send + Sync>);
147
148impl std::fmt::Debug for Progress {
149 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
150 f.write_str("Progress")
151 }
152}
153
154impl Default for Options {
155 fn default() -> Options {
156 Options {
157 today: None,
158 merge: true,
159 memoize: true,
160 epsilon: None,
161 fractions: false,
162 limits: Limits::default(),
163 cancel: None,
164 mode: None,
165 inputs: None,
166 conjugate: true,
167 solve: true,
168 progress: None,
169 }
170 }
171}
172
173#[derive(Clone, Debug)]
174pub struct Outcome {
175 pub today: Option<i32>,
177 pub output: String,
179 pub stats: Stats,
180 pub unresolved: Weight,
182 pub evidence: Option<Weight>,
184 pub reports: Vec<Sink>,
186 pub results: Vec<report::ReportResult>,
188 pub format: Format,
191 pub sample: Option<Sampled>,
193 pub data: Vec<data::SourceInfo>,
195 pub updates: Vec<(Variable, Updates)>,
198}
199
200#[derive(Clone, Debug)]
202pub struct Sampled {
203 pub runs: u64,
204 pub seed: u64,
205 pub effective: f64,
207 pub weight: Weight,
209 pub squares: Weight,
211 pub evidence_se: f64,
214 pub densities: bool,
217}
218
219pub fn run(
223 program: &Program,
224 options: &Options,
225 print: &mut (dyn FnMut(&str) + Send),
226) -> Result<Outcome, RuntimeError> {
227 std::thread::scope(|scope| {
228 let handle = std::thread::Builder::new()
229 .name("probl-engine".into())
230 .stack_size(options.limits.stack_size)
231 .spawn_scoped(scope, || run_here(program, options, print));
232 let handle = match handle {
233 Ok(h) => h,
234 Err(e) => {
235 return Err(internal(format!("couldn't start the engine: {e}")));
236 }
237 };
238 handle
239 .join()
240 .unwrap_or_else(|panic| Err(internal(panic_detail(&*panic))))
241 })
242}
243
244pub fn run_on_this_thread(
251 program: &Program,
252 options: &Options,
253 print: &mut (dyn FnMut(&str) + Send),
254) -> Result<Outcome, RuntimeError> {
255 run_here(program, options, print)
256}
257
258fn panic_detail(panic: &(dyn Any + Send)) -> String {
260 panic
261 .downcast_ref::<String>()
262 .cloned()
263 .or_else(|| panic.downcast_ref::<&str>().map(|s| s.to_string()))
264 .unwrap_or_else(|| "no details".to_string())
265}
266
267fn internal(detail: String) -> RuntimeError {
268 OpError::internal("internal error: the engine crashed")
269 .at(Span::default())
270 .with_note(detail)
271}
272
273fn run_here(
274 program: &Program,
275 options: &Options,
276 print: &mut (dyn FnMut(&str) + Send),
277) -> Result<Outcome, RuntimeError> {
278 if options.today.is_some_and(|d| !dates::valid(d)) {
279 return Err(RuntimeError::new(
280 Default::default(),
281 "execution date must be within 0001-01-01..9999-12-31",
282 ));
283 }
284 let settings = &program.settings;
285 let mode = options.mode.clone().unwrap_or_else(|| settings.mode.clone());
286 let sample = match mode {
287 Mode::Auto | Mode::Enumerate => None,
288 Mode::Sample { runs, seed } => Some((runs, seed)),
289 ref other => {
290 return Err(
291 OpError::unsupported(format!("`{}` mode isn't implemented yet", other.name()))
292 .help("enumeration and `@mode sample(runs: 10_000)` are; particles and beam search come later")
293 .at(settings.mode_span.unwrap_or_default()),
294 );
295 }
296 };
297 if let Some((runs, _)) = sample {
298 if runs == 0 || runs > u32::MAX as u64 {
299 return Err(OpError::new(format!(
300 "sample mode needs between 1 and {} runs",
301 report::thousands(u32::MAX as i64)
302 ))
303 .at(settings.mode_span.unwrap_or_default()));
304 }
305 }
306 let limits = &options.limits;
307 let config = interp::Config {
308 today: options.today,
309 epsilon: options.epsilon.unwrap_or(settings.epsilon),
310 merging: options.merge,
311 memoizing: options.memoize,
312 max_worlds: limits.max_worlds.min(settings.max_worlds),
313 max_iterations: limits.max_iterations.min(settings.max_iterations),
314 max_call_depth: limits.max_call_depth,
315 max_cached_calls: limits.max_cached_calls,
316 max_output: limits.max_output,
317 budget: Budget {
318 max_integer_bits: limits.max_integer_bits,
319 integer_bytes_left: Arc::new(AtomicU64::new(limits.max_integer_bytes)),
320 max_string_bytes: limits.max_string_bytes,
321 string_bytes_left: Arc::new(AtomicU64::new(limits.max_string_alloc_bytes)),
322 cancel: options.cancel.clone(),
323 max_outcomes: limits.max_outcomes,
324 max_collection: limits.max_collection,
325 work_left: limits.max_work,
326 shared: None,
327 },
328 cancel: options.cancel.clone(),
329 sample_seed: sample.map(|(_, seed)| seed),
330 conjugate: options.conjugate,
331 solving: options.solve,
332 max_chain_states: limits.max_chain_states,
333 };
334 let epsilon = config.epsilon;
335 let inputs = inputs(program, options)?;
336 let live = probl_sema::analyze(program);
337 let conj = probl_sema::conjugate::analyze(program);
338 if let Some((runs, seed)) = sample {
339 return sampled(program, &live, &conj, config, options, print, runs, seed);
340 }
341 let mut engine = interp::Engine::new(program, &live, &conj, config, inputs, print);
342 let finished = engine.run_main()?;
343 let unresolved = engine.unresolved;
344
345 if engine.observed && finished.is_zero() {
346 let span = engine.last_ruling_out.unwrap_or_default();
347 return Err(if unresolved.is_zero() {
348 RuntimeError::new(
349 span,
350 "the evidence is impossible: every world was ruled out by `observe`",
351 )
352 .with_help("check the observations; there's no answer to condition on")
353 } else {
354 RuntimeError::new(span, "every world that fits the evidence was left unresolved")
355 .with_help("lower `@epsilon` so that loops run longer")
356 });
357 }
358
359 let format = Format {
361 fractions: options.fractions,
362 weighted: program.main().effects.observes,
363 unresolved,
364 program_total: finished,
365 run_squares: None,
366 };
367 let plain = Format {
368 fractions: false,
369 ..format
370 };
371 let mut header = String::from("enumerated");
372 let evidence = engine.observed.then_some(finished);
373 if let Some(z) = evidence {
374 let (lo, hi) = (finished.to_f64(), (finished + unresolved).to_f64());
375 if hi - lo >= 0.00005 {
376 header.push_str(&format!(
377 " · evidence {}–{}",
378 report::pct(lo, plain),
379 report::pct(hi, plain)
380 ));
381 } else if z.to_f64() < 0.0001 {
382 header.push_str(&format!(" · evidence {}", scientific(z)));
383 } else {
384 header.push_str(&format!(" · evidence {}", report::pct(z.to_f64(), plain)));
385 }
386 }
387 if !unresolved.is_zero() {
388 if unresolved.to_f64() <= epsilon {
389 header.push_str(&format!(" · unresolved < {epsilon:e}"));
390 } else {
391 header.push_str(&format!(" · unresolved {:.1e}", unresolved.to_f64()));
392 }
393 }
394 let results = report::results(&program.reports, &engine.sinks, format, unresolved);
395 let body = report::render_results(&program.reports, &results, format);
396 let output = if body.is_empty() {
397 header
398 } else {
399 format!("{header}\n\n{body}")
400 };
401 Ok(Outcome {
402 today: options.today,
403 output,
404 stats: engine.stats.clone(),
405 unresolved,
406 evidence,
407 reports: std::mem::take(&mut engine.sinks),
408 results,
409 format,
410 sample: None,
411 data: sources(options),
412 updates: Vec::new(),
413 })
414}
415
416fn scientific(w: Weight) -> String {
418 let l = w.log10();
419 let exponent = l.floor();
420 let mantissa = libm::pow(10.0, l - exponent);
421 let (mantissa, exponent) = if format!("{mantissa:.2}") == "10.00" {
423 (1.0, exponent + 1.0)
424 } else {
425 (mantissa, exponent)
426 };
427 format!("{mantissa:.2}e{exponent}")
428}
429
430fn evidence_estimate(z: Weight, relative_se: f64, densities: bool) -> String {
433 let known = relative_se.is_finite();
434 if densities {
435 let ln = z.log10() * std::f64::consts::LN_10;
436 if !known {
437 return format!("log evidence {ln:.2}");
438 }
439 let decimals = if relative_se > 0.0 {
440 (-libm::log10(relative_se).floor()).clamp(2.0, 6.0) as usize
441 } else {
442 2
443 };
444 return format!("log evidence {ln:.decimals$} ± {relative_se:.decimals$}");
445 }
446 let x = z.to_f64();
447 if x >= 1e-4 {
448 return match known {
449 true => format!("evidence {}", report::estimate(x, x * relative_se)),
450 false => format!("evidence {}", value::fmt_prob(x)),
451 };
452 }
453 if !known {
454 return format!("evidence {}", scientific(z));
455 }
456 let pct = relative_se * 100.0;
457 let decimals = if pct >= 10.0 {
458 0
459 } else if pct >= 1.0 {
460 1
461 } else {
462 2
463 };
464 format!("evidence {} (± {pct:.decimals$}%)", scientific(z))
465}
466
467fn inputs<'a>(program: &Program, options: &'a Options) -> Result<&'a [Value], RuntimeError> {
469 let Some(first) = program.inputs.first() else {
470 return Ok(&[]);
471 };
472 match &options.inputs {
473 Some(inputs) if inputs.fit(program) => {
474 if inputs.max_string_bytes() > options.limits.max_string_bytes {
475 return Err(RuntimeError::limit(
476 first.span,
477 "an input string exceeds the runtime string size limit",
478 ));
479 }
480 if inputs.max_integer_bits() > options.limits.max_integer_bits {
481 return Err(RuntimeError::limit(
482 first.span,
483 "an input integer exceeds the runtime integer size limit",
484 ));
485 }
486 Ok(inputs.values())
487 }
488 Some(_) => Err(internal("the data was loaded for a different program".to_string())),
489 None => Err(
490 RuntimeError::new(first.span, "the program reads data, which wasn't loaded")
491 .with_help("load it with `probl_engine::data::load` before running the program"),
492 ),
493 }
494}
495
496fn sources(options: &Options) -> Vec<data::SourceInfo> {
497 options.inputs.as_ref().map_or_else(Vec::new, |i| i.sources().to_vec())
498}
499
500struct Combined {
502 sinks: Vec<Sink>,
503 totals: interp::SampleTotals,
504 observed: bool,
505 densities: bool,
506 unresolved: Weight,
507 last_ruling_out: Option<Span>,
508 stats: Stats,
509}
510
511impl Combined {
512 fn new(reports: usize) -> Combined {
513 Combined {
514 sinks: vec![Sink::default(); reports],
515 totals: interp::SampleTotals {
516 weight: Weight::ZERO,
517 squares: Weight::ZERO,
518 },
519 observed: false,
520 densities: false,
521 unresolved: Weight::ZERO,
522 last_ruling_out: None,
523 stats: Stats::default(),
524 }
525 }
526
527 fn absorb(&mut self, batch: interp::Batch) {
528 for (sink, theirs) in self.sinks.iter_mut().zip(batch.sinks) {
529 sink.absorb(theirs);
530 }
531 self.totals.weight += batch.totals.weight;
532 self.totals.squares += batch.totals.squares;
533 self.observed |= batch.observed;
534 self.densities |= batch.densities;
535 self.unresolved += batch.unresolved;
536 self.last_ruling_out = batch.last_ruling_out.or(self.last_ruling_out);
537 self.stats.absorb(&batch.stats);
538 }
539}
540
541type Sent = (u64, Result<interp::Batch, RuntimeError>, interp::Printed);
544
545#[allow(clippy::too_many_arguments)]
550fn run_batches(
551 program: &Program,
552 live: &Liveness,
553 conj: &Conjugacy,
554 config: &interp::Config,
555 options: &Options,
556 print: &mut (dyn FnMut(&str) + Send),
557 runs: u64,
558 seed: u64,
559) -> Result<Combined, RuntimeError> {
560 let batches = runs.div_ceil(interp::BATCH);
561 let threads = (options.limits.max_threads.max(1) as u64).min(batches);
562 let inputs = inputs(program, options)?;
563 let ahead = 2 * threads;
566 let mut config = config.clone();
567 config.budget.shared = Some(Arc::new(AtomicU64::new(config.budget.work_left)));
568 config.budget.work_left = 0;
569 let config = &config;
570 if threads == 1 {
571 return run_batches_here(program, live, conj, config, inputs, options, print, runs, seed);
572 }
573
574 let next = AtomicU64::new(0);
575 let failed = AtomicU64::new(u64::MAX);
577 let printed = AtomicUsize::new(0);
579 let combined_upto = Mutex::new(0u64);
581 let caught_up = Condvar::new();
582 let stop = |error: RuntimeError| {
583 let _upto = combined_upto.lock().unwrap_or_else(PoisonError::into_inner);
585 failed.store(0, Ordering::Relaxed);
586 caught_up.notify_all();
587 Err(error)
588 };
589 let (tx, rx) = mpsc::channel::<Sent>();
590 std::thread::scope(|scope| {
591 for t in 0..threads {
592 let tx = tx.clone();
593 let (next, failed, printed) = (&next, &failed, &printed);
594 let (combined_upto, caught_up) = (&combined_upto, &caught_up);
595 let worker = move || {
596 let mut ignore = |_: &str| {};
597 let copied: Vec<Value> = inputs.iter().map(Value::unshared).collect();
600 let mut engine = interp::Engine::new(program, live, conj, config.clone(), &copied, &mut ignore);
601 loop {
602 let index = next.fetch_add(1, Ordering::Relaxed);
603 if index >= batches {
604 break;
605 }
606 let mut upto = combined_upto.lock().unwrap_or_else(PoisonError::into_inner);
607 while index >= *upto + ahead && index <= failed.load(Ordering::Relaxed) {
608 upto = caught_up.wait(upto).unwrap_or_else(PoisonError::into_inner);
609 }
610 drop(upto);
611 if index > failed.load(Ordering::Relaxed) {
612 break;
613 }
614 let first = index * interp::BATCH;
615 let n = (runs - first).min(interp::BATCH);
616 let before = printed.load(Ordering::Relaxed);
617 let (result, lines) = std::panic::catch_unwind(AssertUnwindSafe(|| {
618 engine.run_batch(Rng::stream(seed, index), first, n, before)
619 }))
620 .unwrap_or_else(|panic| (Err(internal(panic_detail(&*panic))), Vec::new()));
621 let stopped = result.is_err();
622 if stopped {
623 failed.fetch_min(index, Ordering::Relaxed);
624 }
625 if tx.send((index, result, lines)).is_err() || stopped {
626 break;
627 }
628 }
629 };
630 let spawned = std::thread::Builder::new()
631 .name("probl-sampler".into())
632 .stack_size(options.limits.stack_size)
633 .spawn_scoped(scope, worker);
634 if let Err(e) = spawned {
635 if t == 0 {
637 return Err(internal(format!("couldn't start a sampling thread: {e}")));
638 }
639 break;
640 }
641 }
642 drop(tx);
643
644 let mut combined = Combined::new(program.reports.len());
645 let mut pending = BTreeMap::new();
646 let mut bytes = 0;
647 for index in 0..batches {
648 let (result, lines) = loop {
649 if let Some(sent) = pending.remove(&index) {
650 break sent;
651 }
652 match rx.recv() {
653 Ok((i, result, lines)) => {
654 pending.insert(i, (result, lines));
655 }
656 Err(_) => return stop(internal("a sampling thread stopped before its batch was done".into())),
657 }
658 };
659 for (span, line) in lines {
660 bytes += line.len() + 1;
661 if bytes > options.limits.max_output {
662 return stop(interp::too_much_output(span));
663 }
664 print(&line);
665 }
666 printed.store(bytes, Ordering::Relaxed);
667 match result {
668 Ok(batch) => combined.absorb(batch),
669 Err(e) => return stop(e),
670 }
671 if let Some(progress) = &options.progress {
672 (progress.0)(((index + 1) * interp::BATCH).min(runs), runs);
673 }
674 *combined_upto.lock().unwrap_or_else(PoisonError::into_inner) = index + 1;
675 caught_up.notify_all();
676 }
677 Ok(combined)
678 })
679}
680
681#[allow(clippy::too_many_arguments)]
684fn run_batches_here(
685 program: &Program,
686 live: &Liveness,
687 conj: &Conjugacy,
688 config: &interp::Config,
689 inputs: &[Value],
690 options: &Options,
691 print: &mut (dyn FnMut(&str) + Send),
692 runs: u64,
693 seed: u64,
694) -> Result<Combined, RuntimeError> {
695 let mut ignore = |_: &str| {};
696 let mut engine = interp::Engine::new(program, live, conj, config.clone(), inputs, &mut ignore);
697 let mut combined = Combined::new(program.reports.len());
698 let mut bytes = 0;
699 for index in 0..runs.div_ceil(interp::BATCH) {
700 let first = index * interp::BATCH;
701 let n = (runs - first).min(interp::BATCH);
702 let (result, lines) = engine.run_batch(Rng::stream(seed, index), first, n, bytes);
703 for (span, line) in lines {
704 bytes += line.len() + 1;
705 if bytes > options.limits.max_output {
706 return Err(interp::too_much_output(span));
707 }
708 print(&line);
709 }
710 combined.absorb(result?);
711 if let Some(progress) = &options.progress {
712 (progress.0)(first + n, runs);
713 }
714 }
715 Ok(combined)
716}
717
718#[allow(clippy::too_many_arguments)]
720fn sampled(
721 program: &Program,
722 live: &Liveness,
723 conj: &Conjugacy,
724 config: interp::Config,
725 options: &Options,
726 print: &mut (dyn FnMut(&str) + Send),
727 runs: u64,
728 seed: u64,
729) -> Result<Outcome, RuntimeError> {
730 let mut engine = run_batches(program, live, conj, &config, options, print, runs, seed)?;
731 let totals = engine.totals;
732 if totals.weight.is_zero() {
733 let span = engine.last_ruling_out.unwrap_or_default();
734 return Err(RuntimeError::new(span, "every run was ruled out by `observe`")
735 .with_note(format!("{} runs were tried", report::thousands(runs as i64)))
736 .with_help(
737 "the evidence may be impossible, or too unlikely for this many runs: enumerate the model, or add runs",
738 ));
739 }
740 let effective = (totals.weight * totals.weight).ratio(totals.squares);
741 let mut header = format!("sample · {} runs · seed {seed}", report::thousands(runs as i64));
742 let evidence = totals.weight.scale(1.0 / runs as f64);
745 let n = runs as f64;
746 let evidence_se = if runs > 1 {
747 ((n / effective - 1.0).max(0.0) / (n - 1.0)).sqrt()
748 } else {
749 f64::NAN
750 };
751 if engine.observed {
752 header.push_str(&format!(
753 " · {}",
754 evidence_estimate(evidence, evidence_se, engine.densities)
755 ));
756 header.push_str(&format!(
757 " · effective sample size {}",
758 report::thousands(effective.round() as i64)
759 ));
760 }
761 let unresolved = engine.unresolved;
762 if !unresolved.is_zero() {
763 header.push_str(&format!(
764 " · unresolved {:.1e}",
765 unresolved.ratio(totals.weight + unresolved)
766 ));
767 }
768 let format = Format {
769 fractions: false,
770 unresolved: Weight::ZERO,
771 program_total: totals.weight,
772 run_squares: Some(totals.squares),
773 weighted: program.main().effects.observes,
774 };
775 let results = report::results(&program.reports, &engine.sinks, format, unresolved);
778 let body = report::render_results(&program.reports, &results, format);
779 let output = if body.is_empty() {
780 header
781 } else {
782 format!("{header}\n\n{body}")
783 };
784 Ok(Outcome {
785 today: options.today,
786 output,
787 stats: engine.stats.clone(),
788 unresolved,
789 evidence: engine.observed.then_some(evidence),
790 reports: std::mem::take(&mut engine.sinks),
791 results,
792 format,
793 sample: Some(Sampled {
794 runs,
795 seed,
796 effective,
797 weight: totals.weight,
798 squares: totals.squares,
799 evidence_se,
800 densities: engine.densities,
801 }),
802 data: sources(options),
803 updates: conj
804 .variables
805 .iter()
806 .enumerate()
807 .map(|(i, v)| (v.clone(), engine.stats.updates.get(i).copied().unwrap_or_default()))
808 .collect(),
809 })
810}