use std::collections::{BTreeMap, BTreeSet, HashSet};
use std::fmt::Debug;
use std::hash::{DefaultHasher, Hash, Hasher};
use anyhow::{anyhow, Context as AnyhowCtx};
use json_patch::{Patch, PatchOperation};
use jsonptr::PointerBuf;
use serde_json::Value;
use thiserror::Error;
use tracing::field::display;
use tracing::{error, field, instrument, trace, trace_span, warn, Span};
use crate::errors::{InternalError, MethodError, SerializationError};
use crate::path::Path;
use crate::state::{AsInternal, State};
use crate::system::System;
use crate::task::{self, Context, Operation, Task};
use crate::workflow::{Dag, WorkUnit, Workflow};
mod distance;
mod domain;
use distance::*;
pub use domain::*;
#[derive(Debug, Clone)]
pub struct Planner(Domain);
#[derive(Debug, Error)]
enum SearchFailed {
#[error("method error: {0}")]
BadMethod(#[from] PathSearchError),
#[error("task error: {0:?}")]
BadTask(#[from] task::Error),
#[error("task not applicable")]
EmptyTask,
#[error("internal error: {0:?}")]
Internal(#[from] anyhow::Error),
}
fn select_non_conflicting_prefer_prefixes<'a, I>(paths: I) -> Vec<Path>
where
I: IntoIterator<Item = &'a Path>,
{
let mut result: Vec<Path> = Vec::new();
for p in paths.into_iter() {
if !result.iter().any(|selected| selected.is_prefix_of(p)) {
result.retain(|selected| !p.is_prefix_of(selected));
result.push(p.clone());
}
}
result
}
fn domains_are_conflicting<'a, I>(cumulative_domain: &BTreeSet<Path>, domain: I) -> bool
where
I: Iterator<Item = &'a Path>,
{
for path1 in domain {
for path2 in cumulative_domain.iter() {
if path2.is_prefix_of(path1) || path1.is_prefix_of(path2) {
return true;
}
}
}
false
}
fn longest_common_prefix<'a, I>(paths: I) -> Path
where
I: IntoIterator<Item = &'a Path>,
{
let mut iter = paths.into_iter();
let first = match iter.next() {
Some(path) => path.as_ref().tokens().collect::<Vec<_>>(),
None => return Path::default(),
};
let mut prefix = first;
for path in iter {
let tokens = path.as_ref().tokens().collect::<Vec<_>>();
let mut new_prefix = vec![];
for (a, b) in prefix.iter().zip(tokens.iter()) {
if a == b {
new_prefix.push(a.clone());
} else {
break;
}
}
prefix = new_prefix;
if prefix.is_empty() {
break;
}
}
let buf = PointerBuf::from_tokens(&prefix);
Path::new(&buf)
}
fn hash_state(state: &System) -> u64 {
let value = state.root();
let mut hasher = DefaultHasher::new();
value.hash(&mut hasher);
hasher.finish()
}
#[derive(Clone, PartialEq, Eq)]
struct Candidate {
partial_plan: Dag<WorkUnit>,
changes: Vec<PatchOperation>,
path: Path,
domain: BTreeSet<Path>,
operation: Operation,
priority: u8,
is_method: bool,
}
impl PartialOrd for Candidate {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Candidate {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other
.path
.cmp(&self.path)
.then(self.is_method.cmp(&other.is_method))
.then(self.operation.cmp(&other.operation))
.then(self.priority.cmp(&other.priority))
}
}
#[derive(Debug, Error)]
pub(crate) enum Error {
#[error(transparent)]
Serialization(#[from] SerializationError),
#[error(transparent)]
Task(#[from] task::Error),
#[error("workflow not found")]
NotFound,
#[error(transparent)]
Internal(#[from] InternalError),
}
impl Planner {
pub fn new(domain: Domain) -> Self {
Self(domain)
}
fn try_task(
&self,
task: &Task,
cur_state: &System,
domain: &mut BTreeSet<Path>,
changes: &mut Vec<PatchOperation>,
) -> Result<Dag<WorkUnit>, SearchFailed> {
let span = Span::current();
match task {
Task::Action(action) => {
let work_id = WorkUnit::new_id(action, cur_state.root());
let patch = action.dry_run(cur_state).map_err(SearchFailed::BadTask)?;
if patch.is_empty() {
return Err(SearchFailed::EmptyTask);
}
span.record("selected", display(true));
span.record("changes", display(&patch));
let Patch(ops) = patch;
let new_plan = Dag::from(WorkUnit::new(work_id, action.clone(), ops.clone()));
domain.insert(action.domain());
changes.extend(ops);
Ok(new_plan)
}
Task::Method(method) => {
let tasks = method.expand(cur_state).map_err(SearchFailed::BadTask)?;
let mut extended_tasks = Vec::new();
for mut t in tasks.into_iter() {
let task_id = t.id().to_string();
let Context {
args: method_args, ..
} = method.context();
let Context { args, .. } = t.context_mut();
for (k, v) in method_args.iter() {
if !args.contains_key(k) {
args.insert(k, v);
}
}
let path = self.0.find_path_for_job(&task_id, args)?;
let job = self
.0
.find_job(&path, &task_id)
.ok_or(anyhow!("failed to find job for path {path}"))?;
let task = job.new_task(t.context().to_owned()).with_path(path.clone());
extended_tasks.push(task);
}
let mut plan_branches = Vec::new();
let mut cumulative_domain = BTreeSet::new();
let mut cur_state = cur_state.clone();
for task in extended_tasks {
let mut task_domain = BTreeSet::new();
let mut task_changes = Vec::new();
let partial_plan =
self.try_task(&task, &cur_state, &mut task_domain, &mut task_changes)?;
let partial_plan =
if domains_are_conflicting(&cumulative_domain, task_domain.iter()) {
let dag = Dag::new(plan_branches).prepend(partial_plan);
plan_branches = Vec::new();
dag
} else {
partial_plan
};
cur_state
.patch(Patch(task_changes.to_vec()))
.with_context(|| format!("failed to apply patch {task_changes:?}"))?;
plan_branches.push(partial_plan);
cumulative_domain.extend(task_domain);
changes.extend(task_changes);
}
let new_plan = Dag::new(plan_branches);
domain.extend(cumulative_domain);
let patch = Patch(changes.to_vec());
span.record("selected", display(true));
span.record("changes", display(patch));
Ok(new_plan)
}
}
}
#[instrument(level = "trace", skip_all, err(level = "trace"))]
pub(crate) fn find_workflow<T>(&self, system: &System, tgt: &Value) -> Result<Workflow, Error>
where
T: State,
{
trace!(initial=%system, target=%tgt, "searching for workflow");
let mut stack = vec![(system.clone(), Dag::default(), 0)];
let mut visited_states = HashSet::new();
let find_workflow_span = Span::current();
let initial_state_with_meta = system
.state::<T>()
.and_then(|t| serde_json::to_value(AsInternal(&t)))
.map_err(SerializationError::from)?;
let mut halted_state_paths: Vec<Path> = Vec::new();
while let Some((cur_state, cur_plan, depth)) = stack.pop() {
if depth >= 256 {
warn!(parent: &find_workflow_span, "reached max search depth (256)");
return Err(Error::NotFound)?;
}
let cur = cur_state
.state::<T::Target>()
.and_then(serde_json::to_value)
.map_err(SerializationError::from)?;
visited_states.insert(hash_state(&cur_state));
let distance = Distance::new(&cur, tgt, &halted_state_paths);
if distance.is_empty() {
return Ok(Workflow::new(cur_plan.reverse()).with_ignored(halted_state_paths));
}
let next_span = trace_span!("find_next", distance = %distance, cur_plan=field::Empty);
next_span.in_scope(|| {
next_span.record("cur_plan", field::display(cur_plan.clone().reverse()));
});
let _enter = next_span.enter();
let mut candidates: Vec<Candidate> = Vec::new();
for op in distance.operations() {
let path = Path::new(op.path());
let pointer = path.as_ref();
if halted_state_paths.iter().any(|p| p.is_prefix_of(&path)) {
continue;
}
let state = pointer
.resolve(&initial_state_with_meta)
.unwrap_or(&Value::Null);
if let Some(true) = state
.get("__mahler(halted)")
.and_then(|value| value.as_bool())
{
halted_state_paths.push(path.clone());
continue;
}
let target = pointer.resolve(tgt).unwrap_or(&Value::Null);
if let Some((args, jobs)) = self.0.find_matching_jobs(path.as_str()) {
let context = Context {
path: path.clone(),
args,
target: target.clone(),
};
for job in jobs.filter(|j| j.operation() != &Operation::None) {
if op.matches(job.operation()) || job.operation() == &Operation::Any {
let task = job.new_task(context.clone());
let mut changes = Vec::new();
let mut domain = BTreeSet::new();
match self.try_task(&task, &cur_state, &mut domain, &mut changes) {
Ok(partial_plan) if !changes.is_empty() => {
candidates.push(Candidate {
partial_plan,
changes,
path: task.path().clone(),
domain,
is_method: task.is_method(),
operation: job.operation().clone(),
priority: job.priority(),
});
}
Err(SearchFailed::EmptyTask)
| Err(SearchFailed::BadTask(task::Error::ConditionFailed)) => {}
Err(SearchFailed::Internal(err)) => {
return Err(InternalError::from(err))?;
}
Err(SearchFailed::BadMethod(err)) => {
let err = MethodError::new(err);
if cfg!(debug_assertions) {
return Err(task::Error::from(err))?;
}
warn!(
parent: &find_workflow_span,
"task {} failed: {} ... ignoring",
task.id(),
err
);
}
Err(SearchFailed::BadTask(err)) => {
if cfg!(debug_assertions) {
return Err(err)?;
}
warn!(
parent: &find_workflow_span,
"task {} failed: {} ... ignoring",
task.id(),
err
);
}
_ => {}
}
}
}
}
}
let non_conflicting_paths = select_non_conflicting_prefer_prefixes(
candidates.iter().map(|Candidate { path, .. }| path),
);
let mut concurrent_candidates: BTreeMap<Path, Candidate> = BTreeMap::new();
let mut cumulative_domain = BTreeSet::new();
for candidate in candidates.iter() {
if let Some(prev_candidate) = concurrent_candidates.get(&candidate.path) {
if *prev_candidate >= *candidate {
continue;
}
}
if non_conflicting_paths.iter().any(|p| p == &candidate.path)
&& !domains_are_conflicting(&cumulative_domain, candidate.domain.iter())
{
cumulative_domain.extend(candidate.domain.clone());
concurrent_candidates.insert(candidate.path.clone(), candidate.clone());
}
}
if concurrent_candidates.len() > 1 {
let mut plan_branches = Vec::new();
let mut changes = Vec::new();
let mut domain = BTreeSet::new();
let mut total_priority = 0;
let mut is_method = true;
let path = longest_common_prefix(concurrent_candidates.keys());
for candidate in concurrent_candidates.into_values() {
candidates.retain(|c| *c != candidate);
let Candidate {
partial_plan,
changes: candidate_changes,
domain: candidate_domain,
is_method: candidate_is_method,
priority,
..
} = candidate;
plan_branches.push(partial_plan);
changes.extend(candidate_changes);
domain.extend(candidate_domain);
total_priority += priority;
is_method = is_method && candidate_is_method;
}
candidates.push(Candidate {
partial_plan: Dag::new(plan_branches),
changes,
domain,
path,
is_method,
operation: Operation::Update,
priority: total_priority,
})
}
candidates.sort();
trace!(candidates=%candidates.len());
for Candidate {
partial_plan,
changes,
..
} in candidates.into_iter().rev()
{
let mut new_state = cur_state.clone();
new_state
.patch(Patch(changes))
.with_context(|| "failed to apply patch")
.map_err(InternalError::from)?;
let state_hash = hash_state(&new_state);
if visited_states.contains(&state_hash) {
continue;
}
if cur_plan.any(|unit| partial_plan.any(|u| u.id == unit.id)) {
continue;
}
let new_plan = cur_plan.shallow_clone().prepend(partial_plan);
stack.push((new_state, new_plan, depth + 1));
break;
}
if stack.is_empty() {
trace!(last_evaluated_state=%cur_state, "no plan was found");
}
}
Err(Error::NotFound)?
}
}
#[cfg(test)]
mod tests {
use pretty_assertions::assert_eq;
use serde::{Deserialize, Serialize};
use serde_json::json;
use std::collections::HashMap;
use std::fmt::Display;
use super::*;
use crate::extract::{Args, System, Target, View};
use crate::state::{AsInternal, Map, State};
use crate::{dag, par, seq, task::*, workflow::Dag};
use tracing_subscriber::fmt::format::FmtSpan;
use tracing_subscriber::{prelude::*, EnvFilter};
fn init() {
tracing_subscriber::registry()
.with(
tracing_subscriber::fmt::layer()
.pretty()
.with_target(false)
.with_span_events(FmtSpan::NEW | FmtSpan::CLOSE),
)
.with(EnvFilter::from_default_env())
.try_init()
.unwrap_or(());
}
fn plus_one(mut counter: View<i32>, Target(tgt): Target<i32>) -> View<i32> {
if *counter < tgt {
*counter += 1;
}
counter
}
fn buggy_plus_one(mut counter: View<i32>, Target(tgt): Target<i32>) -> View<i32> {
if *counter < tgt {
*counter -= 1;
}
counter
}
fn plus_two(counter: View<i32>, Target(tgt): Target<i32>) -> Vec<Task> {
if tgt - *counter > 1 {
return vec![plus_one.with_target(tgt), plus_one.with_target(tgt)];
}
vec![]
}
fn plus_three(counter: View<i32>, Target(tgt): Target<i32>) -> Vec<Task> {
if tgt - *counter > 2 {
return vec![plus_two.with_target(tgt), plus_one.with_target(tgt)];
}
vec![]
}
fn minus_one(mut counter: View<i32>, Target(tgt): Target<i32>) -> View<i32> {
if *counter > tgt {
*counter -= 1;
}
counter
}
pub fn find_plan<T>(planner: Planner, cur: T, tgt: T::Target) -> Result<Workflow, super::Error>
where
T: State,
{
let tgt = serde_json::to_value(tgt).expect("failed to serialize target state");
let system =
crate::system::System::try_from(cur).expect("failed to serialize current state");
let res = planner.find_workflow::<T>(&system, &tgt)?;
Ok(res)
}
#[test]
fn it_calculates_a_linear_workflow() {
let domain = Domain::new()
.job("", update(plus_one))
.job("", update(minus_one));
let planner = Planner::new(domain);
let workflow = find_plan(planner, 0, 2).unwrap();
let expected: Dag<&str> = seq!(
"mahler_core::planner::tests::plus_one()",
"mahler_core::planner::tests::plus_one()"
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn it_ignores_none_jobs() {
let domain = Domain::new().job("", none(plus_one));
let planner = Planner::new(domain);
let workflow = find_plan(planner, 0, 2);
assert!(matches!(workflow, Err(super::Error::NotFound)));
}
#[test]
fn it_aborts_search_if_plan_length_grows_too_much() {
let domain = Domain::new()
.job("", update(buggy_plus_one))
.job("", update(minus_one));
let planner = Planner::new(domain);
let workflow = find_plan(planner, 0, 2);
assert!(workflow.is_err());
}
#[test]
fn it_calculates_a_linear_workflow_with_compound_tasks() {
init();
let domain = Domain::new()
.job("", update(plus_two))
.job("", none(plus_one));
let planner = Planner::new(domain);
let workflow = find_plan(planner, 0, 2).unwrap();
let expected: Dag<&str> = seq!(
"mahler_core::planner::tests::plus_one()",
"mahler_core::planner::tests::plus_one()"
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn it_calculates_a_linear_workflow_on_a_complex_state() {
#[derive(Serialize, Deserialize)]
struct MyState {
counters: HashMap<String, i32>,
}
impl State for MyState {
type Target = Self;
}
let initial = MyState {
counters: HashMap::from([("one".to_string(), 0), ("two".to_string(), 0)]),
};
let target = MyState {
counters: HashMap::from([("one".to_string(), 2), ("two".to_string(), 2)]),
};
let domain = Domain::new()
.job("/counters/{counter}", update(minus_one))
.job("/counters/{counter}", update(plus_one));
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = par!(
"mahler_core::planner::tests::plus_one(/counters/one)",
"mahler_core::planner::tests::plus_one(/counters/two)",
) + par!(
"mahler_core::planner::tests::plus_one(/counters/one)",
"mahler_core::planner::tests::plus_one(/counters/two)",
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn it_calculates_a_linear_workflow_on_a_complex_state_with_compound_tasks() {
#[derive(Serialize, Deserialize)]
struct MyState {
counters: HashMap<String, i32>,
}
impl State for MyState {
type Target = Self;
}
let initial = MyState {
counters: HashMap::from([("one".to_string(), 0), ("two".to_string(), 0)]),
};
let target = MyState {
counters: HashMap::from([("one".to_string(), 2), ("two".to_string(), 2)]),
};
let domain = Domain::new()
.job("/counters/{counter}", none(plus_one))
.job("/counters/{counter}", update(plus_two));
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = dag!(
seq!(
"mahler_core::planner::tests::plus_one(/counters/one)",
"mahler_core::planner::tests::plus_one(/counters/one)",
),
seq!(
"mahler_core::planner::tests::plus_one(/counters/two)",
"mahler_core::planner::tests::plus_one(/counters/two)",
)
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn it_calculates_a_linear_workflow_on_a_complex_state_with_deep_compound_tasks() {
#[derive(Serialize, Deserialize)]
struct MyState {
counters: HashMap<String, i32>,
}
impl State for MyState {
type Target = Self;
}
let initial = MyState {
counters: HashMap::from([("one".to_string(), 0), ("two".to_string(), 0)]),
};
let target = MyState {
counters: HashMap::from([("one".to_string(), 3), ("two".to_string(), 0)]),
};
let domain = Domain::new()
.job("/counters/{counter}", none(plus_one))
.job("/counters/{counter}", none(plus_two))
.job("/counters/{counter}", update(plus_three));
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = seq!(
"mahler_core::planner::tests::plus_one(/counters/one)",
"mahler_core::planner::tests::plus_one(/counters/one)",
"mahler_core::planner::tests::plus_one(/counters/one)",
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn it_avoids_conflicts_from_methods() {
init();
let initial = Map::from([("one".to_string(), 0), ("two".to_string(), 0)]);
let target = Map::from([("one".to_string(), 1), ("two".to_string(), 1)]);
fn plus_other(Target(tgt): Target<i32>) -> Vec<Task> {
vec![
plus_one.with_arg("counter", "one").with_target(tgt),
plus_one.with_arg("counter", "two").with_target(tgt),
]
}
let domain = Domain::new()
.job("/{counter}", none(plus_one))
.job("/{counter}", update(plus_other));
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = par!(
"mahler_core::planner::tests::plus_one(/one)",
"mahler_core::planner::tests::plus_one(/two)",
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn it_ignores_halted_sub_states_when_planning() {
init();
#[derive(Serialize, Deserialize)]
struct AppTarget {
running: bool,
}
#[derive(Serialize, Deserialize)]
struct App {
running: bool,
#[serde(default)]
install_failed: bool,
}
impl State for App {
type Target = AppTarget;
fn is_halted(&self) -> bool {
self.install_failed
}
fn as_internal<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeStruct;
let mut state = serializer.serialize_struct("App", 3)?;
state.serialize_field("__mahler(halted)", &self.is_halted())?;
state.serialize_field("running", &self.running)?;
state.serialize_field("install_failed", &self.install_failed)?;
state.end()
}
}
#[derive(Serialize, Deserialize)]
struct Device {
apps: Map<String, App>,
}
#[derive(Serialize, Deserialize)]
struct DeviceTarget {
apps: Map<String, AppTarget>,
}
impl State for Device {
type Target = DeviceTarget;
fn as_internal<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeStruct;
let mut state = serializer.serialize_struct("Device", 2)?;
state.serialize_field("apps", &AsInternal(&self.apps))?;
state.end()
}
}
fn prepare_app(mut app: View<Option<App>>) -> View<Option<App>> {
app.replace(App {
running: false,
install_failed: false,
});
app
}
fn install_app(mut app: View<App>) -> View<App> {
app.running = true;
app
}
let domain = Domain::new().jobs(
"/apps/{app_name}",
[
create(prepare_app).with_description(|Args(app_name): Args<String>| {
format!("prepare app {app_name}")
}),
update(install_app).with_description(|Args(app_name): Args<String>| {
format!("install app {app_name}")
}),
],
);
let initial = serde_json::from_value::<Device>(json!({
"apps": {
"one": {
"running": false,
},
"two": {
"running": false,
"install_failed": true,
}
}
}))
.unwrap();
let target = serde_json::from_value::<DeviceTarget>(json!({
"apps": {
"one": {
"running": true,
},
"two": {
"running": true,
},
"three": {
"running": true,
}
}
}))
.unwrap();
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> =
par!("install app one", "prepare app three") + seq!("install app three");
assert_eq!(expected.to_string(), workflow.to_string());
}
#[test]
fn it_fails_to_find_a_plan_for_a_buggy_task() {
init();
#[derive(Serialize, Deserialize, PartialEq, Eq, Clone)]
struct App {
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
}
impl State for App {
type Target = Self;
}
type Config = Map<String, String>;
#[derive(Serialize, Deserialize, PartialEq, Eq, Clone)]
struct Device {
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[serde(default)]
apps: HashMap<String, App>,
#[serde(default)]
config: Config,
#[serde(default)]
needs_cleanup: bool,
}
impl State for Device {
type Target = Self;
}
fn store_config(
mut config: View<Config>,
Target(tgt_config): Target<Config>,
) -> View<Config> {
*config = tgt_config;
config
}
fn set_device_name(
mut name: View<Option<String>>,
Target(tgt): Target<Option<String>>,
) -> View<Option<String>> {
*name = tgt;
name
}
fn ensure_cleanup(mut device: View<Device>) -> View<Device> {
device.needs_cleanup = true;
device
}
fn complete_cleanup(mut device: View<Device>) -> View<Device> {
device.needs_cleanup = false;
device
}
fn dummy_task() {}
fn do_cleanup(
System(device): System<Device>,
Target(tgt_device): Target<Device>,
) -> Vec<Task> {
let into_tgt = Device {
needs_cleanup: false,
..device.clone()
};
if into_tgt != tgt_device || !device.needs_cleanup {
return vec![];
}
vec![dummy_task.into_task(), complete_cleanup.into_task()]
}
fn prepare_app(
mut app: View<Option<App>>,
Target(tgt_app): Target<App>,
) -> View<Option<App>> {
app.replace(tgt_app);
app
}
let domain = Domain::new()
.job(
"/name",
any(set_device_name).with_description(|| "set device name"),
)
.job(
"/config",
task::update(store_config).with_description(|| "store configuration"),
)
.jobs(
"",
[
update(ensure_cleanup).with_description(|| "ensure cleanup"),
update(do_cleanup),
none(complete_cleanup).with_description(|| "complete cleanup"),
],
)
.job("", none(dummy_task).with_description(|| "dummy task"))
.job(
"/apps/{app_uuid}",
create(prepare_app).with_description(|Args(app_uuid): Args<String>| {
format!("prepare app {app_uuid}")
}),
);
let initial = serde_json::from_value::<Device>(json!({})).unwrap();
let target = serde_json::from_value::<Device>(json!({
"name": "my-device",
"apps": {
"my-app": {"name": "my-app-name"}
},
"config": {
"some-var": "some-value"
}
}))
.unwrap();
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target);
assert!(workflow.is_err(), "unexpected plan:\n{}", workflow.unwrap());
}
#[test]
fn it_avoids_conflict_in_tasks_returned_from_methods() {
init();
#[derive(Serialize, Deserialize)]
struct Service {
image: String,
}
impl State for Service {
type Target = Self;
}
#[derive(Serialize, Deserialize)]
struct Image {}
#[derive(Serialize, Deserialize)]
struct MySys {
services: Map<String, Service>,
images: Map<String, Image>,
}
impl State for MySys {
type Target = MySysTarget;
}
#[derive(Serialize, Deserialize)]
struct MySysTarget {
services: Map<String, Service>,
}
fn create_image(mut view: View<Option<Image>>) -> View<Option<Image>> {
*view = Some(Image {});
view
}
fn create_service_image(
Target(tgt): Target<Service>,
System(state): System<MySys>,
) -> Option<Task> {
if !state.images.contains_key(&tgt.image) {
return Some(create_image.with_arg("image_name", tgt.image));
}
None
}
fn create_service(
mut view: View<Option<Service>>,
Target(tgt): Target<Service>,
System(state): System<MySys>,
) -> View<Option<Service>> {
if state.images.contains_key(&tgt.image) {
*view = Some(tgt);
}
view
}
let domain = Domain::new()
.job(
"/images/{image_name}",
none(create_image).with_description(|Args(image_name): Args<String>| {
format!("create image '{image_name}'")
}),
)
.jobs(
"/services/{service_name}",
[
create(create_service).with_description(|Args(service_name): Args<String>| {
format!("create service '{service_name}'")
}),
create(create_service_image),
],
);
let initial =
serde_json::from_value::<MySys>(json!({ "images": {}, "services": {} })).unwrap();
let target = serde_json::from_value::<MySysTarget>(
json!({ "services": {"one":{"image": "ubuntu"}, "two": {"image": "ubuntu"}} }),
)
.unwrap();
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = seq!(
"create image 'ubuntu'",
"create service 'one'",
"create service 'two'",
);
assert_eq!(expected.to_string(), workflow.to_string());
}
#[test]
fn it_calculates_concurrent_workflows_from_non_conflicting_paths() {
init();
type Config = Map<String, String>;
#[derive(Serialize, Deserialize)]
struct MyState {
config: Config,
counters: Map<String, i32>,
}
impl State for MyState {
type Target = Self;
}
fn new_counter(
mut counter: View<Option<i32>>,
Target(tgt): Target<i32>,
) -> View<Option<i32>> {
counter.replace(tgt);
counter
}
fn update_config(mut config: View<Config>, Target(tgt): Target<Config>) -> View<Config> {
*config = tgt;
config
}
fn new_config(
mut config: View<Option<String>>,
Target(tgt): Target<String>,
) -> View<Option<String>> {
config.replace(tgt);
config
}
let domain = Domain::new()
.job(
"/counters/{counter}",
create(new_counter).with_description(|Args(counter): Args<String>| {
format!("create counter '{counter}'")
}),
)
.job(
"/config/{config}",
create(new_config).with_description(|Args(config): Args<String>| {
format!("create config '{config}'")
}),
)
.job(
"/config",
update(update_config).with_description(|| "update configurations"),
);
let initial =
serde_json::from_value::<MyState>(json!({ "config": {}, "counters": {} })).unwrap();
let target = serde_json::from_value::<MyState>(
json!({ "config": {"some_var":"one", "other_var": "two"}, "counters": {"one": 0} }),
)
.unwrap();
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = par!("update configurations", "create counter 'one'");
assert_eq!(expected.to_string(), workflow.to_string());
}
#[test]
fn it_finds_concurrent_plans_with_nested_forks() {
init();
type Counters = Map<String, i32>;
#[derive(Serialize, Deserialize, Debug)]
struct MyState {
counters: Counters,
}
impl State for MyState {
type Target = Self;
}
fn multi_increment(counters: View<Counters>, target: Target<Counters>) -> Vec<Task> {
counters
.keys()
.filter(|k| {
target.get(k.as_str()).unwrap_or(&0) - counters.get(k.as_str()).unwrap_or(&0)
> 1
})
.map(|k| {
plus_two
.with_arg("counter", k)
.with_target(target.get(k.as_str()))
})
.collect::<Vec<Task>>()
}
fn chunker(counters: View<Counters>, target: Target<Counters>) -> Vec<Task> {
let mut tasks = Vec::new();
for k in counters
.keys()
.filter(|k| {
target.get(k.as_str()).unwrap_or(&0) - counters.get(k.as_str()).unwrap_or(&0)
> 1
})
.take(2)
{
let mut tgt = (*counters).clone();
if target.contains_key(k.as_str()) {
tgt.insert(k.to_string(), *target.get(k.as_str()).unwrap_or(&0));
}
tasks.push(multi_increment.with_target(tgt));
}
tasks
}
let domain = Domain::new()
.job(
"/counters/{counter}",
update(plus_one)
.with_description(|Args(counter): Args<String>| format!("{counter}++")),
)
.job("/counters/{counter}", update(plus_two))
.job("/counters", update(chunker))
.job("/counters", none(multi_increment));
let initial = MyState {
counters: Map::from([
("a".to_string(), 0),
("b".to_string(), 0),
("c".to_string(), 0),
("d".to_string(), 0),
]),
};
let target = MyState {
counters: Map::from([
("a".to_string(), 3),
("b".to_string(), 2),
("c".to_string(), 2),
("d".to_string(), 2),
]),
};
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = dag!(seq!("a++", "a++"), seq!("b++", "b++"))
+ dag!(seq!("c++", "c++"), seq!("d++", "d++"))
+ seq!("a++");
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn test_array_element_conflicts() {
init();
#[derive(Serialize, Deserialize)]
struct MySys {
items: Vec<String>,
configs: HashMap<String, String>,
}
impl State for MySys {
type Target = Self;
}
fn update_item(mut item: View<String>, Target(tgt): Target<String>) -> View<String> {
*item = tgt;
item
}
fn update_config(mut config: View<String>, Target(tgt): Target<String>) -> View<String> {
*config = tgt;
config
}
fn create_item(
mut item: View<Option<String>>,
Target(tgt): Target<String>,
) -> View<Option<String>> {
*item = Some(tgt);
item
}
fn create_config(
mut config: View<Option<String>>,
Target(tgt): Target<String>,
) -> View<Option<String>> {
*config = Some(tgt);
config
}
fn non_conflicting_updates(Target(tgt): Target<MySys>) -> Vec<Task> {
vec![
update_item
.with_arg("index", "0")
.with_target(tgt.items[0].clone()),
update_item
.with_arg("index", "1")
.with_target(tgt.items[1].clone()),
update_config
.with_arg("key", "server")
.with_target(tgt.configs.get("server").unwrap().clone()),
update_config
.with_arg("key", "database")
.with_target(tgt.configs.get("database").unwrap().clone()),
]
}
let domain = Domain::new()
.job("/items/{index}", update(update_item))
.job("/configs/{key}", update(update_config))
.job("/items/{index}", create(create_item))
.job("/configs/{key}", create(create_config))
.job("/", update(non_conflicting_updates));
let initial = MySys {
items: vec!["old1".to_string(), "old2".to_string()],
configs: HashMap::from([
("server".to_string(), "oldserver".to_string()),
("database".to_string(), "olddatabase".to_string()),
]),
};
let target = MySys {
items: vec!["new1".to_string(), "new2".to_string()],
configs: HashMap::from([
("server".to_string(), "newserver".to_string()),
("database".to_string(), "newdatabase".to_string()),
]),
};
let planner = Planner::new(domain);
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = par!("mahler_core::planner::tests::test_array_element_conflicts::update_config(/configs/database)",
"mahler_core::planner::tests::test_array_element_conflicts::update_config(/configs/server)",
"mahler_core::planner::tests::test_array_element_conflicts::update_item(/items/0)",
"mahler_core::planner::tests::test_array_element_conflicts::update_item(/items/1)",
);
assert_eq!(workflow.to_string(), expected.to_string());
}
#[test]
fn test_stacking_problem() {
init();
#[derive(Serialize, Deserialize, PartialEq, Eq, PartialOrd, Ord, Clone, Debug)]
enum Block {
A,
B,
C,
}
impl State for Block {
type Target = Self;
}
impl Display for Block {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{self:?}")
}
}
#[derive(Serialize, Deserialize, PartialEq, Eq, Debug)]
enum Location {
Blk(Block),
Table,
Hand,
}
impl State for Location {
type Target = Self;
}
impl Location {
fn is_block(&self) -> bool {
matches!(self, Location::Blk(_))
}
}
type Blocks = Map<Block, Location>;
#[derive(Serialize, Deserialize, Debug)]
struct World {
blocks: Blocks,
}
impl State for World {
type Target = Self;
}
fn is_clear(blocks: &Blocks, loc: &Location) -> bool {
if loc.is_block() || loc == &Location::Hand {
return blocks.iter().all(|(_, l)| l != loc);
}
true
}
fn is_holding(blocks: &Blocks) -> bool {
!is_clear(blocks, &Location::Hand)
}
fn all_clear(blocks: &Blocks) -> Vec<&Block> {
blocks
.iter()
.filter(|(b, _)| is_clear(blocks, &Location::Blk((*b).clone())))
.map(|(b, _)| b)
.collect()
}
fn pickup(
mut loc: View<Location>,
System(sys): System<World>,
Args(block): Args<Block>,
) -> View<Location> {
if *loc == Location::Table
&& is_clear(&sys.blocks, &Location::Blk(block))
&& !is_holding(&sys.blocks)
{
*loc = Location::Hand;
}
loc
}
fn unstack(
mut loc: View<Location>,
System(sys): System<World>,
Args(block): Args<Block>,
) -> Option<View<Location>> {
if loc.is_block()
&& is_clear(&sys.blocks, &Location::Blk(block))
&& !is_holding(&sys.blocks)
{
*loc = Location::Hand;
return Some(loc);
}
None
}
fn putdown(mut loc: View<Location>) -> View<Location> {
if *loc == Location::Hand {
*loc = Location::Table
}
loc
}
fn stack(
mut loc: View<Location>,
Target(tgt): Target<Location>,
System(sys): System<World>,
) -> View<Location> {
if *loc == Location::Hand && is_clear(&sys.blocks, &tgt) {
*loc = tgt
}
loc
}
fn take(
loc: View<Location>,
System(sys): System<World>,
Args(block): Args<Block>,
) -> Option<Task> {
if is_clear(&sys.blocks, &Location::Blk(block)) {
if *loc == Location::Table {
return Some(pickup.into_task());
} else {
return Some(unstack.into_task());
}
}
None
}
fn put(loc: View<Location>, Target(tgt): Target<Location>) -> Option<Task> {
if *loc == Location::Hand {
if tgt == Location::Table {
return Some(putdown.into_task());
} else {
return Some(stack.with_target(tgt));
}
}
None
}
fn move_blks(blocks: View<Blocks>, Target(target): Target<Blocks>) -> Vec<Task> {
for blk in all_clear(&blocks) {
let tgt_loc = target.get(blk).unwrap();
let cur_loc = blocks.get(blk).unwrap();
if cur_loc != tgt_loc && is_clear(&blocks, tgt_loc) {
return vec![
take.with_arg("block", blk.to_string()),
put.with_arg("block", blk.to_string()).with_target(tgt_loc),
];
}
}
let mut to_table: Vec<Task> = vec![];
for b in all_clear(&blocks) {
to_table.push(take.with_arg("block", b.to_string()));
to_table.push(
put.with_target(Location::Table)
.with_arg("block", b.to_string()),
);
}
to_table
}
let domain = Domain::new()
.jobs(
"/blocks/{block}",
[
update(pickup).with_description(|Args(block): Args<String>| {
format!("pick up block {block}")
}),
update(unstack).with_description(|Args(block): Args<String>| {
format!("unstack block {block}")
}),
update(putdown).with_description(|Args(block): Args<String>| {
format!("put down block {block}")
}),
update(stack).with_description(
|Args(block): Args<String>, Target(tgt): Target<Location>| {
let tgt_block = match tgt {
Location::Blk(block) => format!("{block:?}"),
_ => format!("{tgt:?}"),
};
format!("stack block {block} on top of block {tgt_block}")
},
),
update(take),
update(put),
],
)
.job("/blocks", update(move_blks));
let planner = Planner::new(domain);
let initial = World {
blocks: Map::from([
(Block::A, Location::Table),
(Block::B, Location::Blk(Block::A)),
(Block::C, Location::Blk(Block::B)),
]),
};
let target = World {
blocks: Map::from([
(Block::A, Location::Blk(Block::B)),
(Block::B, Location::Blk(Block::C)),
(Block::C, Location::Table),
]),
};
let workflow = find_plan(planner, initial, target).unwrap();
let expected: Dag<&str> = seq!(
"unstack block C",
"put down block C",
"unstack block B",
"stack block B on top of block C",
"pick up block A",
"stack block A on top of block B",
);
assert_eq!(workflow.to_string(), expected.to_string(),);
}
#[test]
fn test_select_non_conflicting_prefer_prefixes_basic() {
let paths = vec![Path::from_static("/a"), Path::from_static("/b")];
let result = select_non_conflicting_prefer_prefixes(&paths);
assert_eq!(result, paths);
}
#[test]
fn test_select_non_conflicting_prefer_prefixes_with_conflicts() {
let paths = vec![
Path::from_static("/config/other_var"),
Path::from_static("/config/some_var"),
Path::from_static("/counters/one"),
Path::from_static("/config"),
];
let result = select_non_conflicting_prefer_prefixes(&paths);
let expected = vec![
Path::from_static("/counters/one"),
Path::from_static("/config"),
];
assert_eq!(result, expected);
}
#[test]
fn test_select_non_conflicting_prefer_prefixes_your_example() {
let paths = vec![
Path::from_static("/a"),
Path::from_static("/b"),
Path::from_static("/b/c"),
Path::from_static("/b/d"),
];
let result = select_non_conflicting_prefer_prefixes(&paths);
let expected = vec![Path::from_static("/a"), Path::from_static("/b")];
assert_eq!(result, expected);
}
#[test]
fn test_select_non_conflicting_prefer_prefixes_no_later_prefix() {
let paths = vec![
Path::from_static("/config/server/host"),
Path::from_static("/config/server/port"),
Path::from_static("/database/host"),
];
let result = select_non_conflicting_prefer_prefixes(&paths);
assert_eq!(result, paths);
}
#[test]
fn test_select_non_conflicting_prefer_prefixes_prefix_first() {
let paths = vec![
Path::from_static("/config"),
Path::from_static("/config/server"),
Path::from_static("/config/client"),
];
let result = select_non_conflicting_prefer_prefixes(&paths);
let expected = vec![Path::from_static("/config")];
assert_eq!(result, expected);
}
#[test]
fn test_select_non_conflicting_prefer_prefixes_root_path() {
let paths = vec![
Path::from_static(""),
Path::from_static("/config"),
Path::from_static("/counters"),
];
let result = select_non_conflicting_prefer_prefixes(&paths);
let expected = vec![Path::from_static("")];
assert_eq!(result, expected);
}
#[test]
fn test_select_non_conflicting_proper_path_prefix_vs_string_prefix() {
let paths = vec![
Path::from_static("/a"),
Path::from_static("/aa"),
Path::from_static("/a/b"),
];
let result = select_non_conflicting_prefer_prefixes(&paths);
let expected = vec![Path::from_static("/a"), Path::from_static("/aa")];
assert_eq!(result, expected);
}
#[test]
fn test_longest_common_prefix_empty() {
let paths: Vec<Path> = vec![];
let result = longest_common_prefix(&paths);
assert_eq!(result.as_str(), "");
}
#[test]
fn test_longest_common_prefix_single_path() {
let paths = vec![Path::from_static("/config/server")];
let result = longest_common_prefix(&paths);
assert_eq!(result.as_str(), "/config/server");
}
#[test]
fn test_longest_common_prefix_common_prefix() {
let paths = vec![
Path::from_static("/config/server/host"),
Path::from_static("/config/server/port"),
Path::from_static("/config/server/ssl"),
];
let result = longest_common_prefix(&paths);
assert_eq!(result.as_str(), "/config/server");
}
#[test]
fn test_longest_common_prefix_no_common_prefix() {
let paths = vec![
Path::from_static("/config"),
Path::from_static("/counters"),
Path::from_static("/settings"),
];
let result = longest_common_prefix(&paths);
assert_eq!(result.as_str(), "");
}
#[test]
fn test_longest_common_prefix_root_paths() {
let paths = vec![
Path::from_static("/a/b"),
Path::from_static("/a/c"),
Path::from_static("/a/d"),
];
let result = longest_common_prefix(&paths);
assert_eq!(result.as_str(), "/a");
}
}