use super::Runner;
use crate::EventBus;
use crate::filter_library::FilterLibrary;
use somatize_compiler::ExecutionPlan;
use somatize_core::cache::CacheStore;
use somatize_core::error::Result;
use somatize_core::value::Value;
use std::collections::HashMap;
use std::sync::Arc;
pub trait Transport: Send + Sync {
fn execute(
&self,
plan: &ExecutionPlan,
filters: &FilterLibrary,
input: &Value,
y: Option<&Value>,
fit_mode: bool,
) -> Result<(Value, HashMap<String, Value>)>;
fn get_state(&self, node_ids: &[String]) -> Result<HashMap<String, Value>>;
fn set_state(&self, states: &HashMap<String, Value>) -> Result<()>;
fn get_gradients(&self, node_ids: &[String]) -> Result<HashMap<String, Value>>;
fn apply_gradients(&self, gradients: &HashMap<String, Value>) -> Result<()>;
fn execute_node(&self, node_id: &str, input: Option<&Value>) -> Result<Value> {
let plan = ExecutionPlan::Execute {
node_id: node_id.to_string(),
};
let input_val = input.cloned().unwrap_or(Value::Empty);
let filters = crate::filter_library::FilterLibrary::new();
let (output, _) = self.execute(&plan, &filters, &input_val, None, false)?;
Ok(output)
}
}
pub struct RemoteRunner {
transport: Box<dyn Transport>,
}
impl RemoteRunner {
pub fn new(transport: impl Transport + 'static) -> Self {
Self {
transport: Box::new(transport),
}
}
pub fn transport(&self) -> &dyn Transport {
self.transport.as_ref()
}
}
impl Runner for RemoteRunner {
fn fit(
&self,
plan: &ExecutionPlan,
filters: &FilterLibrary,
_cache: &dyn CacheStore,
_event_bus: &Arc<EventBus>,
input: &Value,
y: Option<&Value>,
) -> Result<(Value, HashMap<String, Value>)> {
self.transport.execute(plan, filters, input, y, true)
}
fn forward(
&self,
plan: &ExecutionPlan,
filters: &FilterLibrary,
_cache: &dyn CacheStore,
_event_bus: &Arc<EventBus>,
input: &Value,
) -> Result<Value> {
let (output, _states) = self.transport.execute(plan, filters, input, None, false)?;
Ok(output)
}
}