use crate::effect::SuspendReason;
use crate::filter::{Distribution, FilterKind, FilterMeta};
use crate::graph::NodeId;
use crate::schema::Schema;
use crate::step::StepMeta;
use crate::value::Value;
use serde::{Deserialize, Serialize};
#[derive(Debug)]
pub enum NodeOutcome {
Produced(Value),
HandOff {
target: NodeId,
carry: Value,
},
Paused {
turn: usize,
reason: SuspendReason,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NodeMeta {
pub name: String,
pub effectful: bool,
pub kind: FilterKind,
pub cacheable: bool,
pub deterministic: bool,
pub differentiable: bool,
pub distribution: Distribution,
pub input_schema: Option<Schema>,
pub output_schema: Option<Schema>,
}
impl NodeMeta {
pub fn trainable(&self) -> bool {
!self.effectful && self.kind == FilterKind::Trainable
}
}
impl From<FilterMeta> for NodeMeta {
fn from(m: FilterMeta) -> Self {
Self {
name: m.name,
effectful: false,
kind: m.kind,
cacheable: m.cacheable,
deterministic: m.deterministic,
differentiable: m.differentiable,
distribution: m.distribution,
input_schema: m.input_schema,
output_schema: m.output_schema,
}
}
}
impl From<StepMeta> for NodeMeta {
fn from(m: StepMeta) -> Self {
Self {
name: m.name,
effectful: true,
kind: FilterKind::Opaque,
cacheable: false,
deterministic: false,
differentiable: false,
distribution: m.distribution,
input_schema: m.input_schema,
output_schema: m.output_schema,
}
}
}
impl NodeMeta {
pub fn as_filter_meta(&self) -> FilterMeta {
FilterMeta {
name: self.name.clone(),
kind: self.kind,
cacheable: self.cacheable,
differentiable: self.differentiable,
deterministic: self.deterministic,
stream_mode: crate::filter::StreamMode::FixedState,
distribution: self.distribution.clone(),
input_schema: self.input_schema.clone(),
output_schema: self.output_schema.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn filter_meta() -> FilterMeta {
FilterMeta {
name: "Scaler".into(),
kind: FilterKind::Trainable,
cacheable: true,
differentiable: false,
deterministic: true,
stream_mode: crate::filter::StreamMode::FixedState,
distribution: Distribution::Local,
input_schema: None,
output_schema: None,
}
}
#[test]
fn a_filter_keeps_its_caching_contract() {
let meta = NodeMeta::from(filter_meta());
assert!(!meta.effectful);
assert!(meta.cacheable);
assert!(meta.deterministic);
assert!(meta.trainable());
}
#[test]
fn a_step_is_not_output_cacheable() {
let meta = NodeMeta::from(StepMeta::new("ReactStep"));
assert!(meta.effectful);
assert!(!meta.cacheable);
assert!(!meta.deterministic);
assert!(!meta.differentiable);
assert!(!meta.trainable());
}
#[test]
fn schemas_survive_both_directions() {
let mut sm = StepMeta::new("Judge");
sm.input_schema = Some(Schema::text());
sm.output_schema = Some(Schema::messages());
let meta = NodeMeta::from(sm);
assert_eq!(meta.input_schema, Some(Schema::text()));
assert_eq!(meta.output_schema, Some(Schema::messages()));
}
}