use crate::exec_state::{Internal, RegistrySealed};
use crate::*;
pub(crate) fn register_container_rebuild_from_spec(
eg: &mut EGraph,
sort_name: &str,
spec: &ContainerRebuildSpec,
) {
let Some(container_sort) = eg.get_sort_by_name(sort_name).cloned() else {
return;
};
let mut uf_names = HashMap::default();
collect_element_uf_names(eg, &container_sort, &mut uf_names);
eg.add_read_primitive(
ContainerRebuild {
name: spec.internal_rebuild_prim.clone(),
container_sort: container_sort.clone(),
uf_names: uf_names.clone(),
proof_mode: spec.internal_rebuild_proof_prim.is_some(),
},
None,
);
if let Some(proof_prim) = &spec.internal_rebuild_proof_prim {
let mut cproof_names = HashMap::default();
collect_container_proof_names(eg, &container_sort, &mut cproof_names);
let names = &eg.proof_state.proof_names;
let congr_name = names.congr_constructor.clone();
let trans_name = names.eq_trans_constructor.clone();
let sym_name = names.eq_sym_constructor.clone();
let container_normalize_name = names.container_normalize_constructor.clone();
let proof_sort: ArcSort = std::sync::Arc::new(EqSort {
name: names.proof_datatype.clone(),
});
eg.add_full_primitive(
ContainerRebuildProof {
name: proof_prim.clone(),
container_sort,
proof_sort,
uf_names,
cproof_names,
congr_name,
trans_name,
sym_name,
container_normalize_name,
},
None,
);
}
}
fn collect_element_uf_names(eg: &EGraph, sort: &ArcSort, out: &mut HashMap<String, String>) {
for elem in sort.inner_sorts() {
if elem.is_eq_sort() {
if let Some(uf) = eg.proof_state.uf_function.get(elem.name()) {
out.insert(elem.name().to_string(), uf.clone());
}
} else if elem.is_eq_container_sort() {
collect_element_uf_names(eg, &elem, out);
}
}
}
fn collect_container_proof_names(eg: &EGraph, sort: &ArcSort, out: &mut HashMap<String, String>) {
if let Some(cp) = eg.proof_state.proof_func_parent.get(sort.name()) {
out.insert(sort.name().to_string(), cp.clone());
}
for elem in sort.inner_sorts() {
if elem.is_eq_container_sort() {
collect_container_proof_names(eg, &elem, out);
}
}
}
fn rebuild_with_leaders(
cvs: &ContainerValues,
es: &mut ExecutionState,
sort: &ArcSort,
value: Value,
leaders: &HashMap<Value, Value>,
) -> Value {
let type_id = sort
.value_type()
.expect("container sorts have a value type");
cvs.rebuild_val_with(type_id, value, es, &|v| {
leaders.get(&v).copied().unwrap_or(v)
})
.unwrap_or(value)
}
fn rebuild_container_value_rec(
state: &mut ReadState,
sort: &ArcSort,
value: Value,
uf_names: &HashMap<String, String>,
proof_mode: bool,
) -> Option<Value> {
let elements = {
let cvs = state.container_values();
sort.inner_values(cvs, value)
};
let mut leaders: HashMap<Value, Value> = HashMap::default();
for (esort, eval) in &elements {
let new = if esort.is_eq_sort() {
match self_lookup_leader(state, uf_names, esort, *eval, proof_mode)? {
Some(leader) => leader,
None => *eval,
}
} else if esort.is_eq_container_sort() {
rebuild_container_value_rec(state, esort, *eval, uf_names, proof_mode)?
} else {
*eval
};
if new != *eval {
leaders.insert(*eval, new);
}
}
let cvs = state.container_values();
let es = state.raw_exec_state();
Some(rebuild_with_leaders(cvs, es, sort, value, &leaders))
}
fn self_lookup_leader(
state: &mut ReadState,
uf_names: &HashMap<String, String>,
esort: &ArcSort,
eval: Value,
proof_mode: bool,
) -> Option<Option<Value>> {
let Some(uf_name) = uf_names.get(esort.name()) else {
return Some(None);
};
let Some(looked_up) = state
.lookup(uf_name, eval)
.expect("union-find index lookup failed")
else {
return Some(None);
};
let leader = if proof_mode {
state
.container_values()
.get_val::<crate::sort::PairContainer>(looked_up)?
.first
} else {
looked_up
};
Some(Some(leader))
}
#[derive(Clone)]
struct ContainerRebuild {
name: String,
container_sort: ArcSort,
uf_names: HashMap<String, String>,
proof_mode: bool,
}
impl Primitive for ContainerRebuild {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
&self.name,
vec![self.container_sort.clone(), self.container_sort.clone()],
span.clone(),
)
.into_box()
}
}
impl ReadPrim for ContainerRebuild {
fn apply<'a, 'db>(&self, mut state: ReadState<'a, 'db>, args: &[Value]) -> Option<Value> {
rebuild_container_value_rec(
&mut state,
&self.container_sort,
args[0],
&self.uf_names,
self.proof_mode,
)
}
}
#[derive(Clone)]
struct ContainerRebuildProof {
name: String,
container_sort: ArcSort,
proof_sort: ArcSort,
uf_names: HashMap<String, String>,
cproof_names: HashMap<String, String>,
congr_name: String,
trans_name: String,
sym_name: String,
container_normalize_name: String,
}
impl Primitive for ContainerRebuildProof {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
&self.name,
vec![self.container_sort.clone(), self.proof_sort.clone()],
span.clone(),
)
.into_box()
}
}
impl FullPrim for ContainerRebuildProof {
fn apply<'a, 'db>(&self, mut state: FullState<'a, 'db>, args: &[Value]) -> Option<Value> {
let (_rebuilt, proof) =
rebuild_container_proof_rec(&mut state, self, &self.container_sort, args[0])?;
Some(proof)
}
}
fn rebuild_container_proof_rec(
state: &mut FullState,
prim: &ContainerRebuildProof,
sort: &ArcSort,
value: Value,
) -> Option<(Value, Value)> {
let base = state
.lookup(prim.cproof_names.get(sort.name())?, value)
.expect("container proof lookup failed")?;
let elements = {
let cvs = state.container_values();
sort.inner_values(cvs, value)
};
let mut leaders: HashMap<Value, Value> = HashMap::default();
let mut child_proofs: Vec<(usize, Value)> = vec![];
for (j, (esort, eval)) in elements.iter().enumerate() {
if esort.is_eq_sort() {
if let Some(uf_name) = prim.uf_names.get(esort.name())
&& let Some(pair_val) = state
.lookup(uf_name, *eval)
.expect("union-find index lookup failed")
{
let (leader, proof) = {
let pc = state
.container_values()
.get_val::<crate::sort::PairContainer>(pair_val)?;
(pc.first, pc.second)
};
if leader != *eval {
leaders.insert(*eval, leader);
child_proofs.push((j, proof));
}
}
} else if esort.is_eq_container_sort() {
let (rebuilt_child, child_proof) =
rebuild_container_proof_rec(state, prim, esort, *eval)?;
if rebuilt_child != *eval {
leaders.insert(*eval, rebuilt_child);
child_proofs.push((j, child_proof));
}
}
}
let rebuilt = {
let cvs = state.container_values();
let es = state.raw_exec_state();
rebuild_with_leaders(cvs, es, sort, value, &leaders)
};
let congr_action = state.registry().lookup_table(&prim.congr_name)?.clone();
let mut current = base;
for (j, proof) in child_proofs {
let j_val = state.base_values().get::<i64>(j as i64);
current =
congr_action.lookup_or_insert(state.raw_exec_state(), &[current, j_val, proof])?;
}
let normalize_action = state
.registry()
.lookup_table(&prim.container_normalize_name)?
.clone();
current = normalize_action.lookup_or_insert(state.raw_exec_state(), &[current])?;
if rebuilt != value {
let sym_action = state.registry().lookup_table(&prim.sym_name)?.clone();
let trans_action = state.registry().lookup_table(&prim.trans_name)?.clone();
let cproof_action = state
.registry()
.lookup_table(prim.cproof_names.get(sort.name())?)?
.clone();
let sym_p = sym_action.lookup_or_insert(state.raw_exec_state(), &[current])?;
let refl = trans_action.lookup_or_insert(state.raw_exec_state(), &[sym_p, current])?;
cproof_action.insert(state.raw_exec_state(), [rebuilt, refl].into_iter());
}
Some((rebuilt, current))
}