use std::collections::HashMap;
use std::rc::Rc;
use arora_behavior::graph::{Graph, GraphDiff, Io, Link, LinkSource, Node as GraphNode, Port};
use arora_behavior::{
interpreter_module, BehaviorContext, BehaviorError, BehaviorInterpreter, BehaviorStatus,
RunPolicy, TaskHandle, TaskId,
};
use arora_types::call::Call;
use arora_types::data::{DataStore, Key, Slot, StateChange};
use arora_types::value::Value;
use arora_types::value_serde;
use uuid::Uuid;
use crate::arora_generated::behavior_tree::status::Status;
use crate::graph::build_behavior_tree;
use crate::nodes::{
PARALLEL_FUNCTION_ID, RUN_CALL_FUNCTION_ID, RUN_CALL_PARAM_ID, RUN_STATUS_FUNCTION_ID,
RUN_STATUS_LATCH_PARAM_ID, RUN_STATUS_OUT_PARAM_ID,
};
use crate::{lower_behavior_tree, schema_groot, LoweredTree, ModuleFunction};
struct Run {
decorator: Uuid,
call_node: Uuid,
status_key: Key,
}
pub struct BehaviorTreeInterpreter {
graph: Graph,
runner: Uuid,
main: Vec<Uuid>,
main_root: Option<Uuid>,
runs: HashMap<TaskId, Run>,
finished: Vec<(Uuid, Uuid)>,
pending_halts: Vec<TaskId>,
function_index: Rc<HashMap<Uuid, ModuleFunction>>,
lowered: Option<LoweredTree>,
dirty: bool,
}
fn direct_resolver(store: &dyn DataStore) -> impl Fn(&str) -> Option<Box<dyn Slot>> + '_ {
move |name: &str| Some(store.slot(&Key::from(name)))
}
impl BehaviorTreeInterpreter {
pub fn new(function_index: Rc<HashMap<Uuid, ModuleFunction>>) -> Self {
let runner = Uuid::new_v4();
let mut graph = Graph::empty();
graph.nodes.insert(
runner,
GraphNode {
id: runner,
function: PARALLEL_FUNCTION_ID,
children: Some(Vec::new()),
..GraphNode::default()
},
);
graph.root = Some(runner);
Self {
graph,
runner,
main: Vec::new(),
main_root: None,
runs: HashMap::new(),
finished: Vec::new(),
pending_halts: Vec::new(),
function_index,
lowered: None,
dirty: true,
}
}
pub fn load_groot(&mut self, xml: &str) -> Result<(), crate::error::BehaviorTreeError> {
let groot = schema_groot::BehaviorTree::try_from_groot_xml(xml)?;
let graph = groot.into_graph(self.function_index.as_ref())?;
self.load(graph)
.map_err(|e| crate::error::BehaviorTreeError::InconsistentTreeError {
message: e.to_string(),
})
}
pub fn graph(&self) -> &Graph {
&self.graph
}
fn runner_node(&self) -> GraphNode {
let mut children = Vec::with_capacity(1 + self.runs.len() + self.finished.len());
children.extend(self.main_root);
children.extend(self.runs.values().map(|run| run.decorator));
children.extend(self.finished.iter().map(|(decorator, _)| *decorator));
GraphNode {
id: self.runner,
function: PARALLEL_FUNCTION_ID,
children: Some(children),
..GraphNode::default()
}
}
fn edit(&mut self, diff: GraphDiff) -> Result<(), BehaviorError> {
self.graph.apply(diff).map_err(|e| BehaviorError {
message: format!("graph diff: {e}"),
})?;
self.dirty = true;
Ok(())
}
fn prune_finished(&mut self) -> Result<(), BehaviorError> {
if self.finished.is_empty() {
return Ok(());
}
let remove_nodes = self
.finished
.drain(..)
.flat_map(|(decorator, call_node)| [decorator, call_node])
.collect();
let diff = GraphDiff {
remove_nodes,
add_nodes: vec![self.runner_node()],
..GraphDiff::default()
};
self.edit(diff)
}
fn process_halts(&mut self, store: &dyn DataStore) -> Result<(), BehaviorError> {
if self.pending_halts.is_empty() {
return Ok(());
}
let mut remove_nodes = Vec::new();
for task in std::mem::take(&mut self.pending_halts) {
let Some(run) = self.runs.remove(&task) else {
continue;
};
let failure: Value = Status::Failure.into();
store
.write(StateChange::set(run.status_key.clone(), failure))
.map_err(|e| BehaviorError {
message: format!("task run {}: writing halt status: {e}", task.0),
})?;
remove_nodes.push(run.decorator);
remove_nodes.push(run.call_node);
}
if remove_nodes.is_empty() {
return Ok(());
}
let diff = GraphDiff {
remove_nodes,
add_nodes: vec![self.runner_node()],
..GraphDiff::default()
};
self.edit(diff)
}
fn sweep_terminal_runs(&mut self, store: &dyn DataStore) {
let terminal: Vec<TaskId> = self
.runs
.iter()
.filter(|(_, run)| {
let value = store.read(std::slice::from_ref(&run.status_key));
matches!(
value.first().and_then(|v| v.clone()).map(Status::try_from),
Some(Ok(Status::Success)) | Some(Ok(Status::Failure))
)
})
.map(|(task, _)| *task)
.collect();
for task in terminal {
let run = self.runs.remove(&task).expect("the id came from the map");
self.finished.push((run.decorator, run.call_node));
}
}
}
impl BehaviorInterpreter for BehaviorTreeInterpreter {
fn tick(&mut self, ctx: &mut BehaviorContext) -> Result<BehaviorStatus, BehaviorError> {
self.process_halts(ctx.store)?;
if self.dirty {
self.prune_finished()?;
if let Some(lowered) = self.lowered.take() {
lowered.unregister(ctx.call_bridge);
}
let tree =
build_behavior_tree(&self.graph, &direct_resolver(ctx.store)).map_err(|e| {
BehaviorError {
message: format!("lowering the scaffold: {e:?}"),
}
})?;
self.lowered = Some(
lower_behavior_tree(tree, self.function_index.clone(), ctx.call_bridge, false)
.map_err(|e| BehaviorError {
message: format!("registering the scaffold: {e:?}"),
})?,
);
self.dirty = false;
}
if let Some(lowered) = &self.lowered {
lowered.tick(ctx.call_bridge).map_err(|e| BehaviorError {
message: format!("behavior tree: {e:?}"),
})?;
}
self.sweep_terminal_runs(ctx.store);
Ok(BehaviorStatus::Running)
}
fn apply(&mut self, diff: GraphDiff) -> Result<(), BehaviorError> {
let mut diff = diff;
diff.remove_nodes.retain(|id| *id != self.runner);
for id in &diff.remove_nodes {
self.main.retain(|main| main != id);
if self.main_root == Some(*id) {
self.main_root = None;
}
}
for node in &diff.add_nodes {
if node.id != self.runner && !self.main.contains(&node.id) {
self.main.push(node.id);
}
}
if let Some(root) = diff.set_root.take() {
self.main_root = Some(root);
}
diff.add_nodes.push(self.runner_node());
self.edit(diff)
}
fn load(&mut self, graph: Graph) -> Result<(), BehaviorError> {
let mut diff = GraphDiff {
remove_nodes: std::mem::take(&mut self.main),
..GraphDiff::default()
};
self.main = graph.nodes.keys().copied().collect();
self.main_root = graph.root;
let load = GraphDiff::load(graph);
diff.add_nodes = load.add_nodes;
diff.add_links = load.add_links;
diff.variables = load.variables;
diff.add_nodes.push(self.runner_node());
self.edit(diff)
}
fn spawn(&mut self, call: Call, policy: RunPolicy) -> Result<TaskHandle, BehaviorError> {
let _ = policy;
let task = TaskId(Uuid::new_v4());
let module = call
.module_id
.map(|m| m.to_string())
.unwrap_or_else(|| "none".to_string());
let prefix = format!("arora/tasks/{module}/{}/{}", call.id, task.0);
let status_key = Key::from(format!("{prefix}/status"));
let call_value = value_serde::to_value(&call).map_err(|e| BehaviorError {
message: format!("the spawned call does not convert to a value: {e}"),
})?;
let call_node = Uuid::new_v4();
let decorator = Uuid::new_v4();
let mut diff = GraphDiff {
add_nodes: vec![
GraphNode {
id: call_node,
function: RUN_CALL_FUNCTION_ID,
inputs: vec![Io::new(RUN_CALL_PARAM_ID)],
..GraphNode::default()
},
GraphNode {
id: decorator,
function: RUN_STATUS_FUNCTION_ID,
inputs: vec![Io::new(RUN_STATUS_LATCH_PARAM_ID)],
outputs: vec![Io {
predetermined_key: Some(status_key.path.clone()),
..Io::new(RUN_STATUS_OUT_PARAM_ID)
}],
children: Some(vec![call_node]),
},
],
add_links: vec![
Link::new(
Port::new(call_node, RUN_CALL_PARAM_ID),
LinkSource::Literal(call_value),
),
Link::new(
Port::new(decorator, RUN_STATUS_LATCH_PARAM_ID),
LinkSource::Literal(Value::Unit),
),
],
..GraphDiff::default()
};
self.runs.insert(
task,
Run {
decorator,
call_node,
status_key: status_key.clone(),
},
);
diff.add_nodes.push(self.runner_node());
self.edit(diff).inspect_err(|_| {
self.runs.remove(&task);
})?;
Ok(TaskHandle {
id: task,
stop: interpreter_module::encode_halt(task),
status: status_key,
feedback: vec![Key::from(format!("{prefix}/feedback"))],
result: vec![Key::from(format!("{prefix}/result"))],
update: vec![Key::from(format!("{prefix}/update"))],
})
}
fn halt(&mut self, task: TaskId) -> Result<(), BehaviorError> {
self.pending_halts.push(task);
Ok(())
}
}