use gantz_core::compile::entrypoint::{self, EvalKind, EvalSource};
use gantz_core::node;
use serde::{Deserialize, Serialize};
use steel::SteelVal;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum Action {
SetState {
path: Vec<node::Id>,
values: Vec<Value>,
eval: Option<Source>,
},
Eval { sources: Vec<Source> },
Custom { tag: String, data: Vec<u8> },
}
#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
pub struct Source {
pub path: Vec<node::Id>,
pub kind: Kind,
pub conns: node::Conns,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
pub enum Kind {
Push,
Pull,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum Value {
Unit,
Bool(bool),
Int(i64),
Num(f64),
Char(char),
Str(String),
List(Vec<Value>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct StateWrite {
pub path: Vec<node::Id>,
pub value: Value,
}
#[derive(Clone, Debug)]
pub struct StateWritten(pub StateWrite);
pub(crate) fn state_written(
writes: &mut Vec<StateWrite>,
) -> impl Iterator<Item = crate::DynResponse> + '_ {
writes
.drain(..)
.map(|w| crate::DynResponse::new(StateWritten(w)))
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct UnsupportedValue;
impl std::fmt::Display for UnsupportedValue {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "steel value has no wire-encodable representation")
}
}
impl std::error::Error for UnsupportedValue {}
impl TryFrom<&SteelVal> for Value {
type Error = UnsupportedValue;
fn try_from(val: &SteelVal) -> Result<Self, Self::Error> {
Ok(match val {
SteelVal::Void => Value::Unit,
SteelVal::BoolV(b) => Value::Bool(*b),
SteelVal::IntV(i) => Value::Int(*i as i64),
SteelVal::NumV(n) => Value::Num(*n),
SteelVal::CharV(c) => Value::Char(*c),
SteelVal::StringV(s) => Value::Str(s.to_string()),
SteelVal::ListV(l) => {
Value::List(l.iter().map(Value::try_from).collect::<Result<_, _>>()?)
}
_ => return Err(UnsupportedValue),
})
}
}
impl From<Value> for SteelVal {
fn from(v: Value) -> Self {
match v {
Value::Unit => SteelVal::Void,
Value::Bool(b) => SteelVal::BoolV(b),
Value::Int(i) => SteelVal::IntV(i as isize),
Value::Num(n) => SteelVal::NumV(n),
Value::Char(c) => SteelVal::CharV(c),
Value::Str(s) => SteelVal::StringV(s.into()),
Value::List(l) => SteelVal::ListV(l.into_iter().map(SteelVal::from).collect()),
}
}
}
impl From<EvalKind> for Kind {
fn from(kind: EvalKind) -> Self {
match kind {
EvalKind::Push => Kind::Push,
EvalKind::Pull => Kind::Pull,
}
}
}
impl From<Kind> for EvalKind {
fn from(kind: Kind) -> Self {
match kind {
Kind::Push => EvalKind::Push,
Kind::Pull => EvalKind::Pull,
}
}
}
impl From<EvalSource> for Source {
fn from(src: EvalSource) -> Self {
Self {
path: src.path,
kind: src.kind.into(),
conns: src.conns,
}
}
}
impl From<Source> for EvalSource {
fn from(src: Source) -> Self {
Self {
path: src.path,
kind: src.kind.into(),
conns: src.conns,
}
}
}
pub fn entrypoint(sources: impl IntoIterator<Item = Source>) -> entrypoint::Entrypoint {
entrypoint::from_sources(sources.into_iter().map(EvalSource::from))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn value_round_trips_through_steel() {
let values = [
Value::Unit,
Value::Bool(true),
Value::Int(-42),
Value::Num(1.5),
Value::Char('g'),
Value::Str("gantz".to_string()),
Value::List(vec![Value::Int(1), Value::List(vec![Value::Num(2.0)])]),
];
for value in values {
let steel = SteelVal::from(value.clone());
assert_eq!(Value::try_from(&steel).unwrap(), value);
}
}
#[test]
fn value_preserves_int_vs_num() {
assert!(matches!(SteelVal::from(Value::Int(1)), SteelVal::IntV(1)));
assert!(matches!(
SteelVal::from(Value::Num(1.0)),
SteelVal::NumV(n) if n == 1.0
));
}
#[test]
fn unsupported_values_are_rejected() {
let steel = SteelVal::SymbolV("nope".into());
assert_eq!(Value::try_from(&steel), Err(UnsupportedValue));
let steel = SteelVal::ListV(
[SteelVal::IntV(1), SteelVal::SymbolV("nope".into())]
.into_iter()
.collect(),
);
assert_eq!(Value::try_from(&steel), Err(UnsupportedValue));
}
#[test]
fn eval_sources_rebuild_the_identical_entrypoint() {
let ep = entrypoint::push(vec![3, 1], 2);
let sources: Vec<Source> = ep.0.iter().cloned().map(Source::from).collect();
let rebuilt = entrypoint(sources);
assert_eq!(rebuilt, ep);
assert_eq!(rebuilt.id(), ep.id());
}
}