use crate::effects::{indirect_prints, may_print};
use crate::ir::*;
pub fn move_draws(program: &mut Program) {
let callable_prints = indirect_prints(&program.functions);
for i in 0..program.functions.len() {
let body = std::mem::take(&mut program.functions[i].body);
let body = block(body, &program.functions, callable_prints);
program.functions[i].body = body;
}
}
fn block(b: Block, functions: &[Function], callable_prints: bool) -> Block {
let mut out: Vec<Stmt> = Vec::with_capacity(b.stmts.len());
let mut carried: Vec<(Vec<Stmt>, SlotId)> = Vec::new();
let mut just_carried = false;
for mut s in b.stmts {
if let (StmtKind::Check { slot, .. }, Some((draw, last))) = (&s.kind, carried.last_mut()) {
if just_carried && slot == last {
draw.push(s);
continue;
}
}
nested(&mut s, functions, callable_prints);
let (land, pass): (Vec<_>, Vec<_>) = carried
.into_iter()
.partition(|(_, slot)| stops(&s, *slot, functions, callable_prints));
out.extend(land.into_iter().flat_map(|(d, _)| d));
carried = pass;
just_carried = match movable(&s) {
Some(slot) => {
carried.push((vec![s], slot));
true
}
None => {
out.push(s);
false
}
};
}
out.extend(carried.into_iter().flat_map(|(d, _)| d));
Block { stmts: out }
}
fn nested(s: &mut Stmt, functions: &[Function], callable_prints: bool) {
match &mut s.kind {
StmtKind::If { then, otherwise, .. } => {
*then = block(std::mem::take(then), functions, callable_prints);
*otherwise = block(std::mem::take(otherwise), functions, callable_prints);
}
StmtKind::Chance { arms, otherwise, .. } => {
for (_, b) in arms.iter_mut() {
*b = block(std::mem::take(b), functions, callable_prints);
}
if let Some(b) = otherwise {
*b = block(std::mem::take(b), functions, callable_prints);
}
}
StmtKind::Loop { body, .. } => *body = block(std::mem::take(body), functions, callable_prints),
StmtKind::Try { catches, .. } => {
for c in catches.iter_mut() {
c.body = block(std::mem::take(&mut c.body), functions, callable_prints);
}
}
_ => {}
}
}
fn movable(s: &Stmt) -> Option<SlotId> {
let StmtKind::Draw { place, dist } = &s.kind else {
return None;
};
(place.path.is_empty() && written_out(dist)).then_some(place.slot)
}
fn written_out(e: &Expr) -> bool {
let probability = |e: &Expr| matches!(e.kind, ExprKind::Lit(Lit::Prob(p)) if (0.0..=1.0).contains(&p));
let value = |e: &Expr| {
matches!(
e.kind,
ExprKind::Lit(
Lit::Int(_)
| Lit::Float(_)
| Lit::FloatConstant(_)
| Lit::Prob(_)
| Lit::Str(_)
| Lit::Bool(_)
| Lit::Enum { .. }
)
)
};
match &e.kind {
ExprKind::Lit(Lit::Dice { count, sides }) => (*count as u64) * (*sides as u64 - 1) < 10_000,
ExprKind::Builtin { func, args, named } if named.is_empty() && args.len() == 1 => match func {
crate::Builtin::Bernoulli => probability(&args[0]),
crate::Builtin::OneOf => {
matches!(&args[0].kind, ExprKind::List(items) if !items.is_empty() && items.len() <= 10_000 && items.iter().all(value))
}
_ => false,
},
_ => false,
}
}
fn stops(s: &Stmt, slot: SlotId, functions: &[Function], callable_prints: bool) -> bool {
let mut uses = false;
let mut leaves = false;
visit(s, &mut |x| uses |= x == slot, &mut leaves);
uses || leaves || may_print(s, functions, callable_prints)
}
fn visit(s: &Stmt, f: &mut impl FnMut(SlotId), leaves: &mut bool) {
match &s.kind {
StmtKind::Set { place, value } => {
self::place(place, f);
expr(value, f);
}
StmtKind::Draw { place, dist } => {
self::place(place, f);
expr(dist, f);
}
StmtKind::Take { place, bag } => {
self::place(place, f);
self::place(bag, f);
}
StmtKind::Call { dest, callee, args } => {
self::place(dest, f);
match callee {
Callee::Fn { capture_args, .. } => capture_args.iter().for_each(|&s| f(s)),
Callee::Value(e, _) => expr(e, f),
}
args.iter().for_each(|a| expr(a, f));
}
StmtKind::If { cond, then, otherwise } => {
expr(cond, f);
then.stmts
.iter()
.chain(&otherwise.stmts)
.for_each(|s| visit(s, f, leaves));
}
StmtKind::Chance { arms, otherwise, .. } => {
for (w, b) in arms {
expr(w, f);
b.stmts.iter().for_each(|s| visit(s, f, leaves));
}
if let Some(b) = otherwise {
b.stmts.iter().for_each(|s| visit(s, f, leaves));
}
}
StmtKind::Loop { body, .. } => {
let mut inner = false;
body.stmts.iter().for_each(|s| visit(s, f, &mut inner));
*leaves |= returns(body);
}
StmtKind::Return(e) => {
expr(e, f);
*leaves = true;
}
StmtKind::Break | StmtKind::Continue | StmtKind::Fail { .. } => *leaves = true,
StmtKind::Observe { value, from } => {
expr(value, f);
if let Some(d) = from {
expr(d, f);
}
}
StmtKind::Report { value, key, .. } => {
expr(value, f);
if let Some(k) = key {
expr(k, f);
}
}
StmtKind::Check { slot, .. } => f(*slot),
StmtKind::Try { body, catches } => {
body.stmts
.iter()
.chain(catches.iter().flat_map(|c| &c.body.stmts))
.for_each(|s| visit(s, f, leaves));
}
}
}
fn returns(b: &Block) -> bool {
b.stmts.iter().any(|s| match &s.kind {
StmtKind::Return(_) | StmtKind::Fail { .. } => true,
StmtKind::If { then, otherwise, .. } => returns(then) || returns(otherwise),
StmtKind::Chance { arms, otherwise, .. } => {
arms.iter().any(|(_, b)| returns(b)) || otherwise.as_ref().is_some_and(returns)
}
StmtKind::Loop { body, .. } => returns(body),
StmtKind::Try { body, catches } => returns(body) || catches.iter().any(|c| returns(&c.body)),
_ => false,
})
}
fn place(p: &Place, f: &mut impl FnMut(SlotId)) {
f(p.slot);
for elem in &p.path {
if let PathElem::Index(e) = elem {
expr(e, f);
}
}
}
pub(crate) fn expr(e: &Expr, f: &mut impl FnMut(SlotId)) {
expr_reads(e, f, true);
}
pub(crate) fn value_reads(e: &Expr, f: &mut impl FnMut(SlotId)) {
expr_reads(e, f, false);
}
fn expr_reads(e: &Expr, f: &mut impl FnMut(SlotId), include_type_reads: bool) {
let recurse = |e: &Expr, f: &mut _| expr_reads(e, f, include_type_reads);
match &e.kind {
ExprKind::Builtin {
func: crate::Builtin::Typeof,
args,
..
} if !include_type_reads && matches!(args[0].kind, ExprKind::Slot(_)) => {}
ExprKind::Lit(_) | ExprKind::Input(_) => {}
ExprKind::Slot(s) => f(*s),
ExprKind::Unary(_, x) | ExprKind::Field(x, _) => recurse(x, f),
ExprKind::Binary(_, a, b) | ExprKind::Index(a, b) => {
recurse(a, f);
recurse(b, f);
}
ExprKind::List(items) => items.iter().for_each(|x| recurse(x, f)),
ExprKind::Map(entries) => entries.iter().for_each(|(k, v)| {
recurse(k, f);
recurse(v, f);
}),
ExprKind::Record { fields, .. } => fields.iter().for_each(|(_, x)| recurse(x, f)),
ExprKind::With(base, fields) => {
recurse(base, f);
fields.iter().for_each(|(_, x)| recurse(x, f));
}
ExprKind::Builtin { args, named, .. } => {
args.iter().for_each(|x| recurse(x, f));
named.iter().for_each(|(_, x)| recurse(x, f));
}
ExprKind::Closure { capture_args, .. } | ExprKind::Simulate { capture_args, .. } => {
capture_args.iter().for_each(|&s| f(s))
}
ExprKind::Interp(parts) => parts.iter().for_each(|p| {
if let InterpPart::Expr(x) = p {
recurse(x, f)
}
}),
}
}