use std::{fmt::Debug, future::Future, marker::PhantomData, sync::Arc, time::Duration};
use kube::{runtime::controller::Action, Resource};
use serde::de::DeserializeOwned;
use tokio::task::JoinSet;
use tracing::{debug, error, instrument};
use crate::{record_resource_metadata, Reconciler};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TaskStatus {
Completed,
InProgress,
}
pub trait Task<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
>: Send + Sync
{
fn name(&self) -> &str;
fn run(
&self,
res: Arc<RESOURCE>,
ctx: Arc<CONTEXT>,
) -> impl Future<Output = Result<TaskStatus, ERROR>> + Send;
}
pub struct Dag<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
> {
action: Action,
requeue_delay: Duration,
tasks: Vec<Vec<Arc<TASK>>>,
_ctx: PhantomData<CONTEXT>,
_err: PhantomData<ERROR>,
_res: PhantomData<RESOURCE>,
}
impl<
CONTEXT: Send + Sync + 'static,
ERROR: std::error::Error + Send + Sync + 'static,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync + 'static,
TASK: Task<CONTEXT, ERROR, RESOURCE> + 'static,
> Dag<CONTEXT, ERROR, RESOURCE, TASK>
{
#[instrument(
fields(
resource.api_version = %RESOURCE::api_version(&()),
resource.name,
resource.namespace,
),
skip(self, res, ctx)
)]
pub async fn run(&self, res: Arc<RESOURCE>, ctx: Arc<CONTEXT>) -> Result<Action, ERROR> {
record_resource_metadata!(res.meta());
let mut idx = 0;
while idx < self.tasks.len() {
let mut completed = 0;
let tasks = &self.tasks[idx];
let mut handles = JoinSet::new();
for task in tasks {
let task = task.clone();
let res = res.clone();
let ctx = ctx.clone();
let name = task.name().to_string();
debug!("starting task `{name}`");
handles.spawn(async move { task.run(res, ctx).await.map(|status| (name, status)) });
}
while let Some(res) = handles.join_next().await {
match res {
Ok(Ok((name, TaskStatus::Completed))) => {
debug!("task `{name}` completed");
completed += 1;
}
Ok(Ok((task, TaskStatus::InProgress))) => {
debug!("task `{task}` still in progress");
}
Ok(Err(err)) => return Err(err),
Err(err) => {
error!("failed to wait for task run: {err}");
}
}
}
if completed == tasks.len() {
idx += 1;
} else {
break;
}
}
if idx == self.tasks.len() {
Ok(self.action.clone())
} else {
Ok(Action::requeue(self.requeue_delay))
}
}
}
impl<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
> Default for Dag<CONTEXT, ERROR, RESOURCE, TASK>
{
fn default() -> Self {
Self {
action: Action::await_change(),
requeue_delay: Duration::from_secs(15),
tasks: vec![],
_ctx: PhantomData,
_err: PhantomData,
_res: PhantomData,
}
}
}
#[derive(Default)]
pub struct DagBuilder<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
>(Dag<CONTEXT, ERROR, RESOURCE, TASK>);
impl<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
> DagBuilder<CONTEXT, ERROR, RESOURCE, TASK>
{
pub fn new() -> Self {
Self(Default::default())
}
pub fn action(mut self, action: Action) -> Self {
self.0.action = action;
self
}
pub fn requeue_delay(mut self, delay: Duration) -> Self {
self.0.requeue_delay = delay;
self
}
pub fn start_with<TASKS: IntoIterator<Item = TASK>>(
mut self,
tasks: TASKS,
) -> DagBuilderThen<CONTEXT, ERROR, RESOURCE, TASK> {
self.0.tasks = vec![tasks.into_iter().map(Arc::new).collect()];
DagBuilderThen(self.0)
}
}
pub struct DagBuilderThen<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
>(Dag<CONTEXT, ERROR, RESOURCE, TASK>);
impl<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
> DagBuilderThen<CONTEXT, ERROR, RESOURCE, TASK>
{
pub fn action(mut self, action: Action) -> Self {
self.0.action = action;
self
}
pub fn build(self) -> Dag<CONTEXT, ERROR, RESOURCE, TASK> {
self.0
}
pub fn requeue_delay(mut self, delay: Duration) -> Self {
self.0.requeue_delay = delay;
self
}
pub fn then<TASKS: IntoIterator<Item = TASK>>(mut self, tasks: TASKS) -> Self {
self.0.tasks.push(tasks.into_iter().map(Arc::new).collect());
self
}
}
#[derive(Default)]
pub struct DagReconciler<
CONTEXT: Send + Sync,
ERROR: std::error::Error + Send + Sync,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync,
TASK: Task<CONTEXT, ERROR, RESOURCE>,
> {
on_create_or_update: Option<Dag<CONTEXT, ERROR, RESOURCE, TASK>>,
on_delete: Option<Dag<CONTEXT, ERROR, RESOURCE, TASK>>,
}
impl<
CONTEXT: Send + Sync + 'static,
ERROR: std::error::Error + Send + Sync + 'static,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync + 'static,
TASK: Task<CONTEXT, ERROR, RESOURCE> + 'static,
> DagReconciler<CONTEXT, ERROR, RESOURCE, TASK>
{
pub fn new() -> Self {
Self {
on_create_or_update: None,
on_delete: None,
}
}
pub fn on_create_or_update(mut self, dag: Dag<CONTEXT, ERROR, RESOURCE, TASK>) -> Self {
self.on_create_or_update = Some(dag);
self
}
pub fn on_delete(mut self, dag: Dag<CONTEXT, ERROR, RESOURCE, TASK>) -> Self {
self.on_delete = Some(dag);
self
}
}
impl<
CONTEXT: Send + Sync + 'static,
ERROR: std::error::Error + Send + Sync + 'static,
RESOURCE: Clone + Debug + DeserializeOwned + Resource<DynamicType = ()> + Send + Sync + 'static,
TASK: Task<CONTEXT, ERROR, RESOURCE> + 'static,
> Reconciler<CONTEXT, ERROR, RESOURCE> for DagReconciler<CONTEXT, ERROR, RESOURCE, TASK>
{
#[instrument(
fields(
resource.api_version = %RESOURCE::api_version(&()),
resource.name,
resource.namespace,
),
skip(self, res, ctx)
)]
async fn reconcile_creation_or_update(
&self,
res: Arc<RESOURCE>,
ctx: Arc<CONTEXT>,
) -> Result<Action, ERROR> {
record_resource_metadata!(res.meta());
if let Some(dag) = &self.on_create_or_update {
dag.run(res, ctx).await
} else {
Ok(Action::await_change())
}
}
#[instrument(
fields(
resource.api_version = %RESOURCE::api_version(&()),
resource.name,
resource.namespace,
),
skip(self, res, ctx)
)]
async fn reconcile_deletion(
&self,
res: Arc<RESOURCE>,
ctx: Arc<CONTEXT>,
) -> Result<Action, ERROR> {
record_resource_metadata!(res.meta());
if let Some(dag) = &self.on_delete {
dag.run(res, ctx).await
} else {
Ok(Action::await_change())
}
}
}
#[cfg(test)]
mod test {
use std::convert::Infallible;
use k8s_openapi::api::core::v1::Namespace;
use super::*;
struct DummyTask(TaskStatus);
impl Task<(), Infallible, Namespace> for DummyTask {
fn name(&self) -> &str {
"dummy"
}
async fn run(&self, _res: Arc<Namespace>, _ctx: Arc<()>) -> Result<TaskStatus, Infallible> {
Ok(self.0)
}
}
mod dag {
use super::*;
mod run {
use super::*;
#[tokio::test]
async fn when_first_group_in_progress() {
let delay = Duration::from_secs(10);
let dag = DagBuilder::new()
.requeue_delay(delay)
.start_with([
DummyTask(TaskStatus::InProgress),
DummyTask(TaskStatus::Completed),
])
.then([DummyTask(TaskStatus::InProgress)])
.build();
let action = dag
.run(Arc::new(Namespace::default()), Arc::new(()))
.await
.unwrap();
assert_eq!(action, Action::requeue(delay));
}
#[tokio::test]
async fn when_second_group_in_progress() {
let delay = Duration::from_secs(10);
let dag = DagBuilder::new()
.requeue_delay(delay)
.start_with([
DummyTask(TaskStatus::Completed),
DummyTask(TaskStatus::Completed),
])
.then([DummyTask(TaskStatus::InProgress)])
.build();
let action = dag
.run(Arc::new(Namespace::default()), Arc::new(()))
.await
.unwrap();
assert_eq!(action, Action::requeue(delay));
}
#[tokio::test]
async fn when_completed() {
let exepcted = Action::await_change();
let dag = DagBuilder::new()
.action(exepcted.clone())
.start_with([
DummyTask(TaskStatus::Completed),
DummyTask(TaskStatus::Completed),
])
.then([DummyTask(TaskStatus::Completed)])
.build();
let action = dag
.run(Arc::new(Namespace::default()), Arc::new(()))
.await
.unwrap();
assert_eq!(action, exepcted);
}
}
}
}