use std::any::TypeId;
use std::ops::Deref;
use crate::Error;
use crate::api::{ApiError, IntoValue, IntoValues};
use crate::core_relations::{
BaseValue, BaseValues, ContainerValue, ContainerValues, ExecutionState, ExternalFunctionId,
Value,
};
use crate::{
TypeInfo,
ast::{FunctionSubtype, Literal, ResolvedExpr},
core::ResolvedCall,
sort::{F, S},
typechecking::FuncType,
};
use egglog_bridge::{ActionRegistry, TableAction, TableKind};
use smallvec::SmallVec;
type ValueRow = SmallVec<[Value; 8]>;
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, enum_map::Enum)]
pub enum Context {
Pure,
Write,
Read,
Full,
}
impl Context {
pub const ALL: [Context; 4] = [Context::Pure, Context::Write, Context::Read, Context::Full];
}
pub(crate) trait Internal<'a, 'db: 'a>: 'a {
fn es(&self) -> &ExecutionState<'db>;
fn es_mut(&mut self) -> &mut ExecutionState<'db>;
fn ctx(&self) -> Context;
fn call_external_func(&mut self, id: ExternalFunctionId, args: &[Value]) -> Option<Value> {
self.es_mut().call_external_func(id, args)
}
fn raw_exec_state(&mut self) -> &mut ExecutionState<'db> {
self.es_mut()
}
fn apply_resolved_function(&mut self, _func: &FuncType, _args: &[Value]) -> Option<Value> {
None
}
fn apply_table_function(
&mut self,
subtype: FunctionSubtype,
action: &TableAction,
args: &[Value],
) -> Option<Value> {
match (subtype, self.ctx()) {
(FunctionSubtype::Constructor, Context::Write | Context::Full) => {
action.lookup_or_insert(self.es_mut(), args)
}
(FunctionSubtype::Custom, Context::Read | Context::Full) => {
action.lookup(self.es(), args)
}
_ => None,
}
}
}
pub(crate) trait RegistrySealed<'a, 'db: 'a>: Internal<'a, 'db> {
fn registry(&self) -> &ActionRegistry;
fn type_info(&self) -> Option<&'db TypeInfo> {
self.es().external_context()?.downcast_ref()
}
}
#[allow(private_bounds)]
pub trait Core<'a, 'db: 'a>: Internal<'a, 'db> {
fn base_values(&self) -> &'a BaseValues {
self.es().base_values()
}
fn trigger_early_stop(&self) {
self.es().trigger_early_stop()
}
fn should_stop(&self) -> bool {
self.es().should_stop()
}
fn container_values(&self) -> &'a ContainerValues {
self.es().container_values()
}
fn register_container<C: ContainerValue>(&mut self, container: C) -> Value {
let cv = self.container_values();
let es = self.es_mut();
cv.register_val(container, es)
}
fn map_container(
&mut self,
type_id: TypeId,
value: Value,
remap: &(dyn Fn(Value) -> Value + Send + Sync),
) -> Option<Value> {
let cv = self.container_values();
let es = self.es_mut();
cv.rebuild_val_with(type_id, value, es, remap)
}
fn value_to_base<T: BaseValue>(&self, x: Value) -> T {
self.es().base_values().unwrap::<T>(x)
}
fn base_to_value<T: BaseValue>(&self, x: T) -> Value {
self.es().base_values().get::<T>(x)
}
fn value_to_container<T: ContainerValue>(
&self,
x: Value,
) -> Option<impl Deref<Target = T> + 'a> {
self.es().container_values().get_val::<T>(x)
}
fn container_to_value<T: ContainerValue>(&mut self, x: T) -> Value {
self.register_container(x)
}
fn apply_function(
&mut self,
fc: &crate::sort::FunctionContainer,
args: &[Value],
) -> Option<Value> {
let ctx = self.ctx();
let mut pure = PureState::wrap(self.raw_exec_state(), ctx);
fc.apply(&mut pure, args)
}
fn apply_primitive(
&mut self,
primitive: &crate::core::SpecializedPrimitive,
args: &[Value],
) -> Option<Value> {
let id = primitive.external_id(self.ctx());
self.es_mut().call_external_func(id, args)
}
fn eval_resolved_expr(
&mut self,
expr: &ResolvedExpr,
bindings: &[(&str, Value)],
) -> Option<Value> {
match expr {
ResolvedExpr::Lit(_, literal) => Some(match literal {
Literal::Int(x) => self.base_to_value(*x),
Literal::Float(x) => self.base_to_value(F::from(*x)),
Literal::String(x) => self.base_to_value(S::new(x.clone())),
Literal::Bool(x) => self.base_to_value(*x),
Literal::Unit => self.base_to_value(()),
}),
ResolvedExpr::Var(_, resolved_var) => {
assert!(
!resolved_var.is_global_ref,
"global variable {:?} reached direct expression evaluation before remove_globals",
resolved_var.name
);
bindings
.iter()
.find_map(|(name, value)| (*name == resolved_var.name).then_some(*value))
}
ResolvedExpr::Call(_, resolved_call, children) => {
let mut values = Vec::with_capacity(children.len());
for child in children {
values.push(self.eval_resolved_expr(child, bindings)?);
}
match resolved_call {
ResolvedCall::Primitive(primitive) => self.apply_primitive(primitive, &values),
ResolvedCall::Func(func) => self.apply_resolved_function(func, &values),
}
}
}
}
}
#[allow(private_bounds)]
pub trait Read<'a, 'db: 'a>: Core<'a, 'db> + RegistrySealed<'a, 'db> {
fn lookup<K: IntoValues>(&self, name: &str, key: K) -> Result<Option<Value>, Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Function)?;
let key_values: ValueRow = key.into_values(self.base_values()).collect();
check_arity(name, &action, key_values.len())?;
Ok(action.lookup(self.es(), &key_values))
}
fn eclass_of<K: IntoValues>(&self, name: &str, inputs: K) -> Result<Option<Value>, Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Constructor)?;
let key_values: ValueRow = inputs.into_values(self.base_values()).collect();
check_arity(name, &action, key_values.len())?;
Ok(action.lookup(self.es(), &key_values))
}
fn contains<K: IntoValues>(&self, name: &str, key: K) -> Result<bool, Error> {
let action = lookup_action(self.registry(), name)?;
let key_values: ValueRow = key.into_values(self.base_values()).collect();
check_arity(name, &action, key_values.len())?;
Ok(action.lookup(self.es(), &key_values).is_some())
}
fn constructor_schema(&self, name: &str) -> Result<&'db FuncType, Error> {
func_type_of(self.type_info(), name, FunctionSubtype::Constructor)
}
fn function_schema(&self, name: &str) -> Result<&'db FuncType, Error> {
func_type_of(self.type_info(), name, FunctionSubtype::Custom)
}
fn table_subtype(&self, name: &str) -> Option<FunctionSubtype> {
Some(self.type_info()?.get_func_type(name)?.subtype)
}
fn table_size(&self, name: &str) -> Option<usize> {
self.registry()
.lookup_table(name)
.map(|action| action.row_count(self.es()))
}
fn table_sizes(&self) -> Vec<(&str, usize)> {
self.registry().table_sizes(self.es())
}
fn constructor_enodes(&self, name: &str, mut f: impl FnMut(Enode<'_>)) -> Result<(), Error> {
self.constructor_enodes_while(name, |enode| {
f(enode);
true
})
}
fn constructor_enodes_while(
&self,
name: &str,
mut f: impl FnMut(Enode<'_>) -> bool,
) -> Result<(), Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Constructor)?;
action.for_each_while(self.es(), |row| {
let (eclass, children) = row
.vals
.split_last()
.expect("constructor row has at least an eclass column");
f(Enode {
children,
eclass: *eclass,
subsumed: row.subsumed,
})
});
Ok(())
}
fn enodes_for_eclass(
&self,
name: &str,
eclass: Value,
mut f: impl FnMut(Enode<'_>),
) -> Result<(), Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Constructor)?;
action.for_each_output_value(self.es(), eclass, |row| {
let (eclass, children) = row
.vals
.split_last()
.expect("constructor row has at least an eclass column");
f(Enode {
children,
eclass: *eclass,
subsumed: row.subsumed,
});
});
Ok(())
}
fn function_entries(
&self,
name: &str,
mut f: impl FnMut(FunctionEntry<'_>),
) -> Result<(), Error> {
self.function_entries_while(name, |entry| {
f(entry);
true
})
}
fn function_entries_while(
&self,
name: &str,
mut f: impl FnMut(FunctionEntry<'_>) -> bool,
) -> Result<(), Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Function)?;
action.for_each_while(self.es(), |row| {
let (output, inputs) = row
.vals
.split_last()
.expect("function row has at least an output column");
f(FunctionEntry {
inputs,
output: *output,
subsumed: row.subsumed,
})
});
Ok(())
}
}
#[derive(Clone, Copy, Debug)]
pub struct Enode<'a> {
pub children: &'a [Value],
pub eclass: Value,
pub subsumed: bool,
}
#[derive(Clone, Copy, Debug)]
pub struct FunctionEntry<'a> {
pub inputs: &'a [Value],
pub output: Value,
pub subsumed: bool,
}
#[allow(private_bounds)]
pub trait Write<'a, 'db: 'a>: Core<'a, 'db> + RegistrySealed<'a, 'db> {
fn set<K: IntoValues, V: IntoValue>(
&mut self,
name: &str,
key: K,
value: V,
) -> Result<(), Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Function)?;
let bv = self.base_values();
let mut row: ValueRow = key.into_values(bv).collect();
check_arity(name, &action, row.len())?;
row.push(value.into_value(bv));
action.insert(self.es_mut(), row.into_iter());
Ok(())
}
fn add<R: IntoValues>(&mut self, name: &str, inputs: R) -> Result<Value, Error> {
let action = lookup_action(self.registry(), name)?;
check_subtype(name, &action, TableKind::Constructor)?;
let key: ValueRow = inputs.into_values(self.base_values()).collect();
check_arity(name, &action, key.len())?;
let value = action
.lookup_or_insert(self.es_mut(), &key)
.expect("constructor lookup_or_insert returned None");
Ok(value)
}
fn remove<K: IntoValues>(&mut self, name: &str, key: K) -> Result<(), Error> {
let action = lookup_action(self.registry(), name)?;
let key_values: ValueRow = key.into_values(self.base_values()).collect();
check_arity(name, &action, key_values.len())?;
action.remove(self.es_mut(), &key_values);
Ok(())
}
fn subsume<K: IntoValues>(&mut self, name: &str, key: K) -> Result<(), Error> {
let action = lookup_action(self.registry(), name)?;
let key_values: ValueRow = key.into_values(self.base_values()).collect();
check_arity(name, &action, key_values.len())?;
action.subsume(self.es_mut(), key_values.into_iter());
Ok(())
}
fn union(&mut self, x: Value, y: Value) -> Result<(), Error> {
let action = *self.registry().union_action();
action.union(self.es_mut(), x, y);
Ok(())
}
fn panic(&mut self) -> Option<()> {
let panic_id = self.registry().default_panic_id();
self.es_mut().call_external_func(panic_id, &[]);
None
}
}
fn func_type_of<'db>(
type_info: Option<&'db TypeInfo>,
name: &str,
expected: FunctionSubtype,
) -> Result<&'db FuncType, Error> {
let Some(type_info) = type_info else {
return Err(ApiError::SchemasUnavailable {
name: name.to_string(),
}
.into());
};
let Some(func_type) = type_info.get_func_type(name) else {
return Err(ApiError::MissingTable {
name: name.to_string(),
}
.into());
};
if func_type.subtype != expected {
return Err(ApiError::WrongSubtype {
name: name.to_string(),
expected: expected.label(),
actual: func_type.subtype.label(),
}
.into());
}
Ok(func_type)
}
fn lookup_action(registry: &ActionRegistry, name: &str) -> Result<TableAction, Error> {
registry.lookup_table(name).cloned().ok_or_else(|| {
ApiError::MissingTable {
name: name.to_string(),
}
.into()
})
}
fn check_subtype(name: &str, action: &TableAction, expected: TableKind) -> Result<(), Error> {
if action.kind() == expected {
return Ok(());
}
Err(ApiError::WrongSubtype {
name: name.to_string(),
expected: expected.label(),
actual: action.kind().label(),
}
.into())
}
fn check_arity(table: &str, action: &TableAction, got: usize) -> Result<(), Error> {
let expected = action.input_arity();
if got != expected {
return Err(ApiError::WrongArity {
table: table.to_string(),
expected,
got,
}
.into());
}
Ok(())
}
fn apply_registered_function<'a, 'db: 'a>(
state: &mut impl RegistrySealed<'a, 'db>,
func: &FuncType,
args: &[Value],
) -> Option<Value> {
let action = lookup_action(state.registry(), &func.name).ok()?;
state.apply_table_function(func.subtype, &action, args)
}
pub struct PureState<'a, 'db> {
pub(crate) inner: &'a mut ExecutionState<'db>,
pub(crate) ctx: Context,
}
pub struct ReadState<'a, 'db> {
pub(crate) inner: &'a mut ExecutionState<'db>,
pub(crate) registry: &'a ActionRegistry,
pub(crate) ctx: Context,
}
pub struct WriteState<'a, 'db> {
pub(crate) inner: &'a mut ExecutionState<'db>,
pub(crate) registry: &'a ActionRegistry,
pub(crate) ctx: Context,
}
pub struct FullState<'a, 'db> {
pub(crate) inner: &'a mut ExecutionState<'db>,
pub(crate) registry: &'a ActionRegistry,
pub(crate) ctx: Context,
}
impl<'a, 'db: 'a> PureState<'a, 'db> {
pub(crate) fn wrap(es: &'a mut ExecutionState<'db>, ctx: Context) -> Self {
Self { inner: es, ctx }
}
pub const fn valid_contexts() -> &'static [Context] {
&Context::ALL
}
}
impl<'a, 'db: 'a> ReadState<'a, 'db> {
pub(crate) fn wrap(
es: &'a mut ExecutionState<'db>,
registry: &'a ActionRegistry,
ctx: Context,
) -> Self {
Self {
inner: es,
registry,
ctx,
}
}
pub const fn valid_contexts() -> &'static [Context] {
&[Context::Read, Context::Full]
}
}
impl<'a, 'db: 'a> WriteState<'a, 'db> {
pub(crate) fn wrap(
es: &'a mut ExecutionState<'db>,
registry: &'a ActionRegistry,
ctx: Context,
) -> Self {
Self {
inner: es,
registry,
ctx,
}
}
pub const fn valid_contexts() -> &'static [Context] {
&[Context::Write, Context::Full]
}
}
impl<'a, 'db: 'a> FullState<'a, 'db> {
pub(crate) fn wrap(
es: &'a mut ExecutionState<'db>,
registry: &'a ActionRegistry,
ctx: Context,
) -> Self {
Self {
inner: es,
registry,
ctx,
}
}
pub const fn valid_contexts() -> &'static [Context] {
&[Context::Full]
}
}
impl<'a, 'db: 'a> Internal<'a, 'db> for PureState<'a, 'db> {
fn es(&self) -> &ExecutionState<'db> {
self.inner
}
fn es_mut(&mut self) -> &mut ExecutionState<'db> {
self.inner
}
fn ctx(&self) -> Context {
self.ctx
}
}
impl<'a, 'db: 'a> Core<'a, 'db> for PureState<'a, 'db> {}
impl<'a, 'db: 'a> Internal<'a, 'db> for ReadState<'a, 'db> {
fn es(&self) -> &ExecutionState<'db> {
self.inner
}
fn es_mut(&mut self) -> &mut ExecutionState<'db> {
self.inner
}
fn ctx(&self) -> Context {
self.ctx
}
fn apply_resolved_function(&mut self, func: &FuncType, args: &[Value]) -> Option<Value> {
apply_registered_function(self, func, args)
}
}
impl<'a, 'db: 'a> RegistrySealed<'a, 'db> for ReadState<'a, 'db> {
fn registry(&self) -> &ActionRegistry {
self.registry
}
}
impl<'a, 'db: 'a> Core<'a, 'db> for ReadState<'a, 'db> {}
impl<'a, 'db: 'a> Read<'a, 'db> for ReadState<'a, 'db> {}
impl<'a, 'db: 'a> Internal<'a, 'db> for WriteState<'a, 'db> {
fn es(&self) -> &ExecutionState<'db> {
self.inner
}
fn es_mut(&mut self) -> &mut ExecutionState<'db> {
self.inner
}
fn ctx(&self) -> Context {
self.ctx
}
fn apply_resolved_function(&mut self, func: &FuncType, args: &[Value]) -> Option<Value> {
apply_registered_function(self, func, args)
}
}
impl<'a, 'db: 'a> RegistrySealed<'a, 'db> for WriteState<'a, 'db> {
fn registry(&self) -> &ActionRegistry {
self.registry
}
}
impl<'a, 'db: 'a> Core<'a, 'db> for WriteState<'a, 'db> {}
impl<'a, 'db: 'a> Write<'a, 'db> for WriteState<'a, 'db> {}
impl<'a, 'db: 'a> Internal<'a, 'db> for FullState<'a, 'db> {
fn es(&self) -> &ExecutionState<'db> {
self.inner
}
fn es_mut(&mut self) -> &mut ExecutionState<'db> {
self.inner
}
fn ctx(&self) -> Context {
self.ctx
}
fn apply_resolved_function(&mut self, func: &FuncType, args: &[Value]) -> Option<Value> {
apply_registered_function(self, func, args)
}
}
impl<'a, 'db: 'a> RegistrySealed<'a, 'db> for FullState<'a, 'db> {
fn registry(&self) -> &ActionRegistry {
self.registry
}
}
impl<'a, 'db: 'a> Core<'a, 'db> for FullState<'a, 'db> {}
impl<'a, 'db: 'a> Read<'a, 'db> for FullState<'a, 'db> {}
impl<'a, 'db: 'a> Write<'a, 'db> for FullState<'a, 'db> {}