use std::collections::{BTreeMap, BTreeSet};
use super::*;
use crate::hir::analyze::short_circuit::{
DecisionEdge, build_decision_expr, same_value_merge_shape,
};
use crate::structure::DefId;
type StatementValueMergeOutput<'c> = (&'c ShortCircuitCandidate, TempId);
impl<'a, 'b> StructuredBodyLowerer<'a, 'b> {
pub(super) fn try_lower_conditional_reassign_branch(
&mut self,
block: BlockRef,
stop: Option<BlockRef>,
stmts: &mut Vec<HirStmt>,
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> Option<Option<BlockRef>> {
let short = value_merge_candidate_by_header(self.lowering, block)?;
let ShortCircuitExit::ValueMerge(merge) = short.exit else {
return None;
};
if Some(merge) == stop {
return None;
}
if let Some(bvm) = self.branch_value_merges_by_header.get(&block)
&& bvm
.values
.iter()
.any(|v| Some(v.phi_id) != short.result_phi_id)
{
return None;
}
if let Some(stop) = stop
&& stop != merge
&& short.blocks.contains(&stop)
{
return None;
}
if value_merge_defs_are_overridden(self.lowering, short, target_overrides) {
return None;
}
let plan = build_conditional_reassign_plan(self.lowering, block)?;
if merge_has_other_live_phi(self.lowering, plan.merge, plan.phi_id) {
return None;
}
stmts.extend(self.lower_block_prefix(block, true, target_overrides)?);
self.visited.insert(block);
self.visited.extend(value_merge_skipped_blocks(short));
self.overrides.suppress_phi(plan.phi_id);
stmts.push(assign_stmt(
vec![HirLValue::Temp(plan.target_temp)],
vec![plan.init_value],
));
stmts.push(branch_stmt(
plan.cond,
HirBlock {
stmts: vec![assign_stmt(
vec![HirLValue::Temp(plan.target_temp)],
vec![plan.assigned_value],
)],
},
None,
));
Some(Some(plan.merge))
}
pub(super) fn try_lower_statement_value_merge_branch(
&mut self,
block: BlockRef,
stop: Option<BlockRef>,
stmts: &mut Vec<HirStmt>,
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> Option<Option<BlockRef>> {
let short = value_merge_candidate_by_header(self.lowering, block)?;
let ShortCircuitExit::ValueMerge(merge) = short.exit else {
return None;
};
let allowed_blocks = BTreeSet::from([block]);
if recover_short_value_merge_expr_with_allowed_blocks(self.lowering, short, &allowed_blocks)
.is_some()
{
return None;
}
if let Some(stop) = stop
&& stop != merge
&& short.blocks.contains(&stop)
{
return None;
}
let outputs = self.statement_value_merge_outputs(short)?;
let mut short_stmts = self.lower_block_prefix(block, true, target_overrides)?;
short_stmts.extend(
self.lower_value_merge_node(short, short.entry, &outputs, true, target_overrides)?
.stmts,
);
self.visited.insert(block);
self.visited.extend(value_merge_skipped_blocks(short));
for (output_short, _) in &outputs {
self.overrides.suppress_phi(output_short.result_phi_id?);
}
stmts.extend(short_stmts);
if let Some(bvm) = self.branch_value_merges_by_header.get(&block) {
for value in &bvm.values {
if Some(value.phi_id) == short.result_phi_id {
continue;
}
if let Some(decision_expr) =
self.build_secondary_value_merge_decision(short, value.reg)
{
let bvm_temp = self.lowering.bindings.phi_temps[value.phi_id.index()];
let mut stmt =
assign_stmt(vec![HirLValue::Temp(bvm_temp)], vec![decision_expr]);
apply_loop_rewrites(std::slice::from_mut(&mut stmt), target_overrides);
stmts.push(stmt);
self.overrides.suppress_phi(value.phi_id);
}
}
}
Some(Some(merge))
}
fn statement_value_merge_outputs(
&self,
short: &'b ShortCircuitCandidate,
) -> Option<Vec<StatementValueMergeOutput<'b>>> {
let mut outputs = Vec::new();
for candidate in &self.lowering.structure.short_circuit_candidates {
if !same_statement_value_merge_tree(short, candidate) {
continue;
}
let temp = *self
.lowering
.bindings
.phi_temps
.get(candidate.result_phi_id?.index())?;
outputs.push((candidate, temp));
}
(!outputs.is_empty()).then_some(outputs)
}
pub(super) fn try_lower_value_merge_branch(
&mut self,
block: BlockRef,
stop: Option<BlockRef>,
stmts: &mut Vec<HirStmt>,
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> Option<Option<BlockRef>> {
let short = value_merge_candidate_by_header(self.lowering, block)?;
let ShortCircuitExit::ValueMerge(merge) = short.exit else {
return None;
};
if let Some(bvm) = self.branch_value_merges_by_header.get(&block)
&& bvm
.values
.iter()
.any(|v| Some(v.phi_id) != short.result_phi_id)
{
return None;
}
let allowed_blocks = BTreeSet::from([block]);
let recovery = recover_short_value_merge_expr_recovery_with_allowed_blocks(
self.lowering,
short,
&allowed_blocks,
)?;
if let Some(stop) = stop
&& stop != merge
&& short.blocks.contains(&stop)
{
return None;
}
if recovery.consumes_header_subject() {
self.overrides
.suppress_instrs(consumed_value_merge_subject_instrs(self.lowering, block));
}
stmts.extend(self.lower_block_prefix(block, true, target_overrides)?);
self.visited.insert(block);
self.visited.extend(value_merge_skipped_blocks(short));
self.merge_allowed_blocks
.entry(merge)
.or_default()
.insert(block);
Some(Some(merge))
}
fn lower_value_merge_node(
&self,
short: &ShortCircuitCandidate,
node_ref: ShortCircuitNodeRef,
outputs: &[StatementValueMergeOutput<'_>],
prefix_emitted: bool,
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> Option<HirBlock> {
let node = short.nodes.get(node_ref.index())?;
let mut stmts = Vec::new();
if !prefix_emitted {
stmts.extend(self.lower_block_prefix(node.header, true, target_overrides)?);
}
let mut cond = lower_short_circuit_subject(self.lowering, node.header)?;
rewrite_expr_temps(&mut cond, &temp_expr_overrides(target_overrides));
let truthy = self.lower_value_merge_target(
short,
node.header,
&node.truthy,
outputs,
target_overrides,
)?;
let falsy = self.lower_value_merge_target(
short,
node.header,
&node.falsy,
outputs,
target_overrides,
)?;
stmts.push(branch_stmt(cond, truthy, Some(falsy)));
Some(HirBlock { stmts })
}
pub(super) fn branch_entry_target_overrides(
&self,
header: BlockRef,
entry: Option<BlockRef>,
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> BTreeMap<TempId, HirLValue> {
let Some(entry) = entry else {
return target_overrides.clone();
};
let Some(candidate) = self.branch_by_header.get(&header).copied() else {
return target_overrides.clone();
};
if entry == candidate.then_entry {
return self.branch_value_then_target_overrides(header, target_overrides);
}
if Some(entry) == candidate.else_entry {
return self.branch_value_else_target_overrides(header, target_overrides);
}
target_overrides.clone()
}
fn lower_value_merge_target(
&self,
short: &ShortCircuitCandidate,
current_header: BlockRef,
target: &ShortCircuitTarget,
outputs: &[StatementValueMergeOutput<'_>],
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> Option<HirBlock> {
match target {
ShortCircuitTarget::Node(next_ref) => {
self.lower_value_merge_node(short, *next_ref, outputs, false, target_overrides)
}
ShortCircuitTarget::Value(block) => {
self.lower_value_merge_leaf(current_header, *block, outputs, target_overrides)
}
ShortCircuitTarget::TruthyExit | ShortCircuitTarget::FalsyExit => None,
}
}
fn lower_value_merge_leaf(
&self,
current_header: BlockRef,
block: BlockRef,
outputs: &[StatementValueMergeOutput<'_>],
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> Option<HirBlock> {
let mut stmts = if block == current_header {
Vec::new()
} else {
self.lower_block_prefix(block, false, target_overrides)?
};
for (short, target_temp) in outputs {
let value = if block == current_header
&& header_subject_is_value_carrier(self.lowering, current_header, short.result_reg)
{
lower_short_circuit_subject(self.lowering, block)?
} else {
lower_materialized_value_leaf_expr(self.lowering, short, block)?
};
let mut stmt = assign_stmt(vec![HirLValue::Temp(*target_temp)], vec![value]);
apply_loop_rewrites(std::slice::from_mut(&mut stmt), target_overrides);
stmts.push(stmt);
}
Some(HirBlock { stmts })
}
fn build_secondary_value_merge_decision(
&self,
short: &ShortCircuitCandidate,
reg: Reg,
) -> Option<HirExpr> {
let decision = build_decision_expr(
self.lowering,
short,
short.entry,
lower_short_circuit_subject,
|_, target| match target {
ShortCircuitTarget::Node(next_ref) => Some(DecisionEdge::Node(*next_ref)),
ShortCircuitTarget::Value(block) => Some(DecisionEdge::Leaf(
HirDecisionTarget::Expr(expr_for_reg_at_block_exit(self.lowering, *block, reg)),
)),
ShortCircuitTarget::TruthyExit | ShortCircuitTarget::FalsyExit => None,
},
)?;
Some(HirExpr::Decision(Box::new(decision)))
}
pub(super) fn install_stop_boundary_value_merge_override(
&mut self,
header: BlockRef,
branch_stop: Option<BlockRef>,
target_overrides: &BTreeMap<TempId, HirLValue>,
) {
let Some(merge) = branch_stop else {
return;
};
let Some(short) = value_merge_candidate_by_header(self.lowering, header) else {
return;
};
let ShortCircuitExit::ValueMerge(short_merge) = short.exit else {
return;
};
if short_merge != merge {
return;
}
let Some(phi_id) = short.result_phi_id else {
return;
};
let Some(reg) = short.result_reg else {
return;
};
let Some(expr) = shared_target_expr_from_overrides(self.lowering, short, target_overrides)
else {
return;
};
self.replace_phi_with_entry_expr(merge, phi_id, reg, expr);
}
}
fn value_merge_defs_are_overridden(
lowering: &ProtoLowering<'_>,
short: &ShortCircuitCandidate,
target_overrides: &BTreeMap<TempId, HirLValue>,
) -> bool {
if target_overrides.is_empty() {
return false;
}
let is_overridden = |def: &DefId| {
lowering
.bindings
.fixed_temps
.get(def.index())
.is_some_and(|temp| target_overrides.contains_key(temp))
};
short.entry_defs.iter().any(is_overridden)
|| short
.value_incomings
.iter()
.any(|inc| inc.defs.iter().any(is_overridden))
}
fn merge_has_other_live_phi(
lowering: &ProtoLowering<'_>,
merge: BlockRef,
consumed_phi_id: PhiId,
) -> bool {
lowering
.dataflow
.phi_candidates_in_block(merge)
.iter()
.any(|phi| phi.id != consumed_phi_id && !lowering.dead_phis.contains(&phi.id))
}
fn same_statement_value_merge_tree(
base: &ShortCircuitCandidate,
candidate: &ShortCircuitCandidate,
) -> bool {
base.reducible
&& candidate.reducible
&& base.result_phi_id.is_some()
&& candidate.result_phi_id.is_some()
&& base.result_reg.is_some()
&& candidate.result_reg.is_some()
&& same_value_merge_shape(base, candidate)
}