use anyhow::anyhow;
use json_patch::Patch;
use serde::de::DeserializeOwned;
use serde_json::Value;
use std::pin::Pin;
use std::sync::Arc;
use std::time::Instant;
use thiserror::Error;
use tokio::sync::{watch, Notify};
use tokio::task::JoinHandle;
use tokio::{select, sync::RwLock};
use tokio_stream::wrappers::WatchStream;
use tokio_stream::{Stream, StreamExt};
use tracing::field::display;
use tracing::{debug, error, field, info, info_span, span, trace, warn, Instrument, Level, Span};
#[cfg(debug_assertions)]
mod testing;
#[cfg(debug_assertions)]
pub use testing::*;
use crate::errors::{IOError, InternalError, SerializationError};
use crate::planner::{Domain, Error as PlannerError, Planner};
use crate::state::State;
use crate::system::{Resources, System};
use crate::task::{Error as TaskError, Job};
use crate::workflow::{channel, AggregateError, AutoInterrupt, Interrupt, Sender, WorkflowStatus};
#[derive(Debug, Error)]
#[error(transparent)]
pub struct Panicked(#[from] tokio::task::JoinError);
#[derive(Debug, Error)]
pub enum SeekError {
#[error(transparent)]
Serialization(#[from] SerializationError),
#[error(transparent)]
Panic(#[from] Panicked),
#[error(transparent)]
Planning(#[from] TaskError),
#[error(transparent)]
Internal(#[from] InternalError),
}
#[derive(Debug)]
pub enum SeekStatus {
Success,
NotFound,
Interrupted,
Aborted(Vec<IOError>),
}
impl PartialEq for SeekStatus {
fn eq(&self, other: &Self) -> bool {
matches!(
(self, other),
(SeekStatus::Success, SeekStatus::Success)
| (SeekStatus::NotFound, SeekStatus::NotFound)
| (SeekStatus::Interrupted, SeekStatus::Interrupted)
)
}
}
impl Eq for SeekStatus {}
pub trait WorkerState {}
pub struct Uninitialized {
domain: Domain,
resources: Resources,
}
#[derive(Clone)]
pub struct Ready {
domain: Domain,
resources: Resources,
system_rwlock: Arc<RwLock<System>>,
update_event_channel: watch::Sender<()>,
patch_tx: Sender<Patch>,
notify_writer_closed: Arc<Notify>,
}
pub struct Stopped {}
impl WorkerState for Uninitialized {}
impl WorkerState for Ready {}
impl WorkerState for Stopped {}
pub trait WithResources {
fn insert_resource<R: Send + Sync + 'static>(&mut self, resource: R);
}
impl WithResources for Uninitialized {
fn insert_resource<R>(&mut self, resource: R)
where
R: Send + Sync + 'static,
{
self.resources.insert(resource);
}
}
impl WithResources for Ready {
fn insert_resource<R>(&mut self, resource: R)
where
R: Send + Sync + 'static,
{
self.resources.insert(resource);
}
}
pub struct Worker<O, S: WorkerState = Uninitialized> {
inner: S,
_output: std::marker::PhantomData<O>,
}
impl<O, S: WorkerState> Worker<O, S> {
fn from_inner(inner: S) -> Self {
Worker {
inner,
_output: std::marker::PhantomData,
}
}
pub fn stop(self) -> Worker<O, Stopped> {
Worker::from_inner(Stopped {})
}
}
impl<O> Default for Worker<O, Uninitialized> {
fn default() -> Self {
Worker::new()
}
}
impl<O> Worker<O, Uninitialized> {
pub fn new() -> Self {
Worker::from_inner(Uninitialized {
domain: Domain::new(),
resources: Resources::new(),
})
}
}
impl<O, S: WorkerState + WithResources> Worker<O, S> {
pub fn resource<R>(mut self, res: R) -> Self
where
R: Send + Sync + 'static,
{
self.inner.insert_resource(res);
self
}
pub fn use_resource<R>(&mut self, res: R)
where
R: Send + Sync + 'static,
{
self.inner.insert_resource(res);
}
}
impl<O> Worker<O, Uninitialized> {
pub fn job(mut self, route: &'static str, job: Job) -> Self {
self.inner.domain = self.inner.domain.job(route, job);
self
}
pub fn jobs<const N: usize>(mut self, route: &'static str, list: [Job; N]) -> Self {
self.inner.domain = self.inner.domain.jobs(route, list);
self
}
pub fn initial_state(self, state: O) -> Result<Worker<O, Ready>, SerializationError>
where
O: State,
{
let Uninitialized {
domain, resources, ..
} = self.inner;
let system = System::try_from(state)?;
let system_rwlock = Arc::new(RwLock::new(system));
let (patch_tx, mut patch_rx) = channel::<Patch>(100);
let notify_writer_closed = Arc::new(Notify::new());
let (update_event_channel, _) = watch::channel(());
{
let notify_writer_closed = notify_writer_closed.clone();
let system_writer = Arc::clone(&system_rwlock);
let update_event_tx = update_event_channel.clone();
tokio::spawn(
async move {
while let Some(mut msg) = patch_rx.recv().await {
let changes = std::mem::take(&mut msg.data);
trace!(received=%changes);
let mut system = system_writer.write().await;
if let Err(e) = system.patch(changes) {
error!("patch failed: {e}");
notify_writer_closed.notify_one();
break;
}
trace!("patch successful");
let _ = update_event_tx.send(());
msg.ack();
}
}
.instrument(span!(Level::TRACE, "worker_sync",)),
);
}
Ok(Worker::from_inner(Ready {
domain,
resources,
system_rwlock,
update_event_channel,
patch_tx,
notify_writer_closed,
}))
}
}
struct WorkerStream<T> {
inner: Pin<Box<dyn Stream<Item = T> + Send + 'static>>,
}
impl<T> WorkerStream<T> {
fn new<S>(stream: S) -> Self
where
S: Stream<Item = T> + Send + 'static,
{
Self {
inner: Box::pin(stream),
}
}
}
impl<T> Stream for WorkerStream<T> {
type Item = T;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.inner.as_mut().poll_next(cx)
}
}
fn follow_worker<T>(
channel: watch::Sender<()>,
system_rwlock: Arc<RwLock<System>>,
) -> impl Stream<Item = T>
where
T: DeserializeOwned,
{
let rx = channel.subscribe();
WorkerStream::new(
WatchStream::from_changes(rx)
.then(move |_| {
let sys_reader = Arc::clone(&system_rwlock);
async move {
let system = sys_reader.read().await;
system.state::<T>().ok()
}
})
.filter_map(|opt| opt),
)
}
impl<O: State> Worker<O, Ready> {
pub async fn state(&self) -> Result<O, SerializationError>
where
O: DeserializeOwned,
{
let system = self.inner.system_rwlock.read().await;
let state = system.state()?;
Ok(state)
}
pub fn follow(&self) -> impl Stream<Item = O>
where
O: DeserializeOwned,
{
follow_worker(
self.inner.update_event_channel.clone(),
Arc::clone(&self.inner.system_rwlock),
)
}
pub async fn seek_with_interrupt(
&mut self,
tgt: O::Target,
interrupt: Interrupt,
) -> Result<SeekStatus, SeekError> {
let tgt = serde_json::to_value(tgt).map_err(SerializationError::from)?;
let Ready {
resources,
domain,
system_rwlock,
notify_writer_closed,
patch_tx,
..
} = self.inner.clone();
let planner = Planner::new(domain);
{
let mut system = self.inner.system_rwlock.write().await;
system.set_resources(resources);
}
enum InnerSeekResult {
TargetReached,
WorkflowCompleted,
Interrupted,
}
enum InnerSeekError {
Runtime(AggregateError<TaskError>),
Planning(PlannerError),
}
async fn find_and_run_workflow<T: State>(
planner: &Planner,
sys_reader: &Arc<RwLock<System>>,
tgt: &Value,
patch_tx: &Sender<Patch>,
sigint: &Interrupt,
) -> Result<InnerSeekResult, InnerSeekError> {
info!("searching workflow");
let now = Instant::now();
if tracing::enabled!(tracing::Level::DEBUG) {
let system = sys_reader.read().await;
let cur = system
.state::<T::Target>()
.and_then(serde_json::to_value)
.map_err(SerializationError::from)
.map_err(PlannerError::from)
.map_err(InnerSeekError::Planning)?;
let changes = json_patch::diff(&cur, tgt);
if !changes.0.is_empty() {
debug!("pending changes:");
for change in &changes.0 {
debug!("- {}", change);
}
}
}
let workflow = {
let system = sys_reader.read().await;
let res = planner.find_workflow::<T>(&system, tgt);
if let Err(PlannerError::NotFound) = res {
warn!(time = ?now.elapsed(), "workflow not found");
}
res.map_err(InnerSeekError::Planning)?
};
if workflow.is_empty() {
debug!("nothing to do");
return Ok(InnerSeekResult::TargetReached);
}
info!(time = ?now.elapsed(), "workflow found");
if tracing::enabled!(tracing::Level::WARN) {
warn!("the following paths were ignored during planning");
for path in workflow.ignored() {
warn!("{path}");
}
}
if tracing::enabled!(tracing::Level::DEBUG) {
debug!("will execute the following tasks:");
for line in workflow.to_string().lines() {
debug!("{line}");
}
}
let now = Instant::now();
info!("executing workflow");
let status = workflow
.execute(sys_reader, patch_tx.clone(), sigint.clone())
.await
.map_err(InnerSeekError::Runtime)?;
info!(time = ?now.elapsed(), "workflow executed successfully");
if matches!(status, WorkflowStatus::Interrupted) {
return Ok(InnerSeekResult::Interrupted);
}
Ok(InnerSeekResult::WorkflowCompleted)
}
let drop_interrupt = AutoInterrupt::from(interrupt);
let handle: JoinHandle<Result<SeekStatus, SeekError>> = {
let err_rx = notify_writer_closed;
let interrupt = drop_interrupt.clone();
let sys_reader = system_rwlock;
let patch_tx = patch_tx;
tokio::spawn(async move {
let seek_span = Span::current();
info!("applying target state");
loop {
select! {
biased;
_ = err_rx.notified() => {
return Err(InternalError::from(anyhow!("state patch failed, worker state possibly tainted")))?;
}
res = find_and_run_workflow::<O>(&planner, &sys_reader, &tgt, &patch_tx, &interrupt) => {
match res {
Ok(InnerSeekResult::TargetReached) => {
info!("target state applied");
seek_span.record("result", display("success"));
return Ok(SeekStatus::Success);
}
Ok(InnerSeekResult::WorkflowCompleted) => {}
Ok(InnerSeekResult::Interrupted) => {
warn!("target state apply interrupted by user request");
seek_span.record("result", display("interrupted"));
return Ok(SeekStatus::Interrupted);
}
Err(InnerSeekError::Planning(PlannerError::NotFound)) => {
seek_span.record("result", display("workflow_not_found"));
return Ok(SeekStatus::NotFound);
}
Err(InnerSeekError::Planning(PlannerError::Serialization(e))) => return Err(e)?,
Err(InnerSeekError::Planning(PlannerError::Internal(e))) => return Err(e)?,
Err(InnerSeekError::Planning(PlannerError::Task(e))) => return Err(e)?,
Err(InnerSeekError::Runtime(err)) => {
let mut io = Vec::new();
let mut other = Vec::new();
let AggregateError (all) = err;
for e in all.into_iter() {
match e {
TaskError::IO(re) => io.push(re),
TaskError::ConditionFailed => {},
_ => other.push(e)
}
}
if !other.is_empty() {
return Err(InternalError::from(anyhow!(AggregateError::from(other))))?;
}
if !io.is_empty() {
warn!("target state apply interrupted due to error");
seek_span.record("result", display("aborted"));
return Ok(SeekStatus::Aborted(io))
}
continue;
}
}
}
}
}
}.instrument(info_span!("seek_target", result=field::Empty)))
};
let status = match handle.await {
Ok(Ok(res)) => Ok(res),
Ok(Err(e)) => Err(e),
Err(e) => Err(Panicked(e))?,
}?;
Ok(status)
}
pub async fn seek_target(&mut self, tgt: O::Target) -> Result<SeekStatus, SeekError> {
self.seek_with_interrupt(tgt, Interrupt::new()).await
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::time::Duration;
use super::*;
use crate::extract::{Target, View};
use crate::task::*;
use serde::{Deserialize, Serialize};
use tokio::time::{sleep, timeout};
use tracing_subscriber::fmt::format::FmtSpan;
use tracing_subscriber::{prelude::*, EnvFilter};
#[derive(Debug, Serialize, Deserialize, PartialEq)]
struct Counters(HashMap<String, i32>);
impl State for Counters {
type Target = Self;
}
fn plus_one(mut counter: View<i32>, Target(tgt): Target<i32>) -> IO<i32> {
if *counter < tgt {
*counter += 1;
}
with_io(counter, |counter| async {
sleep(Duration::from_millis(10)).await;
Ok(counter)
})
}
fn buggy_plus_one(mut counter: View<i32>, Target(tgt): Target<i32>) -> View<i32> {
if *counter < tgt {
*counter -= 1;
}
counter
}
fn init() {
tracing_subscriber::registry()
.with(
tracing_subscriber::fmt::layer()
.pretty()
.with_target(false)
.with_thread_names(true)
.with_thread_ids(true)
.with_line_number(true)
.with_span_events(FmtSpan::NEW | FmtSpan::CLOSE),
)
.with(EnvFilter::from_default_env())
.try_init()
.unwrap_or(());
}
#[tokio::test]
async fn test_worker_complex_state() {
init();
let mut worker = Worker::new()
.job("/{counter}", update(plus_one))
.initial_state(Counters(HashMap::from([
("one".to_string(), 0),
("two".to_string(), 0),
])))
.unwrap();
let status = worker
.seek_target(Counters(HashMap::from([
("one".to_string(), 2),
("two".to_string(), 0),
])))
.await
.unwrap();
assert_eq!(status, SeekStatus::Success);
let state = worker.state().await.unwrap();
assert_eq!(
state,
Counters(HashMap::from([
("one".to_string(), 2),
("two".to_string(), 0),
]))
);
}
#[tokio::test]
async fn test_worker_bug() {
init();
let mut worker = Worker::new()
.job("", update(buggy_plus_one))
.initial_state(0)
.unwrap();
let status = worker.seek_target(2).await.unwrap();
assert!(matches!(status, SeekStatus::NotFound));
}
#[tokio::test]
async fn test_worker_follow_updates() {
init();
let mut worker = Worker::new()
.job("", update(plus_one))
.initial_state(0)
.unwrap();
let mut updates = worker.follow();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
let mut res = results.write().await;
let first_update = updates.next().await;
res.push(first_update);
let second_update = updates.next().await;
res.push(second_update);
});
}
let status = worker.seek_target(2).await.unwrap();
assert_eq!(status, SeekStatus::Success);
let results = results.read().await;
assert_eq!(*results, vec![Some(1), Some(2)]);
}
#[tokio::test]
async fn test_worker_follow_best_effort_loss() {
init();
let mut worker = Worker::new()
.job("", update(plus_one))
.initial_state(0)
.unwrap();
let mut updates = worker.follow();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
let mut res = results.write().await;
let first = updates.next().await;
res.push(first);
tokio::time::sleep(Duration::from_millis(200)).await;
let maybe_update = updates.next().await;
res.push(maybe_update);
});
}
let status = worker.seek_target(100).await.unwrap();
assert_eq!(status, SeekStatus::Success);
let results = results.read().await;
assert_eq!(
(*results)
.iter()
.map(|r| r.is_some())
.collect::<Vec<bool>>(),
vec![true, true]
)
}
#[tokio::test]
async fn test_worker_interrupt_status() {
init();
fn sleepy_plus_one(mut counter: View<i32>, Target(tgt): Target<i32>) -> IO<i32> {
if *counter < tgt {
*counter += 1;
}
with_io(counter, |counter| async {
sleep(Duration::from_millis(10)).await;
Ok(counter)
})
}
let mut worker = Worker::new()
.job("", update(sleepy_plus_one))
.initial_state(0)
.unwrap();
let mut updates = worker.follow();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
while let Some(s) = updates.next().await {
let mut res = results.write().await;
res.push(s);
}
});
}
let res = timeout(Duration::from_millis(30), worker.seek_target(10)).await;
assert!(res.is_err());
let results = results.read().await;
assert!(results.len() < 3);
}
#[tokio::test]
async fn test_follow_stream_closes_on_worker_end() {
init();
let mut worker = Worker::new()
.job("", update(plus_one))
.initial_state(0)
.unwrap();
let mut updates = worker.follow();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
let mut res = results.write().await;
let first = updates.next().await;
res.push(first);
let end = updates.next().await;
res.push(end);
});
}
let status = worker.seek_target(1).await.unwrap();
assert_eq!(status, SeekStatus::Success);
worker.stop();
let results = results.read().await;
assert_eq!(*results, vec![Some(1), None]);
}
#[tokio::test]
async fn test_multiple_streams_receive_all_changes() {
init();
let mut worker = Worker::new()
.job("", update(plus_one))
.initial_state(0)
.unwrap();
let mut stream1 = worker.follow();
let mut stream2 = worker.follow();
let mut stream3 = worker.follow();
let results1 = Arc::new(tokio::sync::RwLock::new(Vec::new()));
let results2 = Arc::new(tokio::sync::RwLock::new(Vec::new()));
let results3 = Arc::new(tokio::sync::RwLock::new(Vec::new()));
let task1 = {
let results = Arc::clone(&results1);
tokio::spawn(async move {
let mut res = results.write().await;
while let Some(update) = stream1.next().await {
res.push(update);
}
})
};
let task2 = {
let results = Arc::clone(&results2);
tokio::spawn(async move {
let mut res = results.write().await;
while let Some(update) = stream2.next().await {
res.push(update);
}
})
};
let task3 = {
let results = Arc::clone(&results3);
tokio::spawn(async move {
let mut res = results.write().await;
while let Some(update) = stream3.next().await {
res.push(update);
}
})
};
let status = worker.seek_target(3).await.unwrap();
assert_eq!(status, SeekStatus::Success);
worker.stop();
let _ = tokio::join!(task1, task2, task3);
let results1 = results1.read().await;
let results2 = results2.read().await;
let results3 = results3.read().await;
let expected = vec![1, 2, 3];
assert_eq!(*results1, expected);
assert_eq!(*results2, expected);
assert_eq!(*results3, expected);
}
}