use std::collections::{BTreeMap, HashMap};
use std::future::Future;
use std::sync::Arc;
use std::time::Duration;
use taquba::{
Clock, ExpiryIndex, JobRecord, LeaseHandle, PermanentFailure, Queue, SettlementEffects, Worker,
WorkerError, WorkerHandle,
};
use taquba_cron::PREVIOUS_FIRE_MS_HEADER;
use taquba_workflow::{RunId, RunOptions, RunSpec, RunState};
use tokio_util::sync::CancellationToken;
use crate::definition_store::{DefinitionError, DefinitionStore};
use crate::graph::{Graph, Node};
use crate::hook::{EVENTS_QUEUE, Event};
use crate::input::TaskInput;
use crate::partition::Partition;
use crate::pools::Pools;
use crate::readiness::{
NodeState, current_records, is_ready, node_states, rerun_scope, settled_state,
};
use crate::records::JsonBytes;
use crate::records::{
self, EXPIRY_PREFIX, Entry, Expiring, GraphRecord, GraphRunRecord, GraphRunState, NodeRecord,
ReadError, RecordError, RecordStatus,
};
use crate::task::{self, HEADER_GRAPH, TaskIdentity};
pub const TRIGGERS_QUEUE: &str = "swale-triggers";
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error(transparent)]
Queue(#[from] taquba::Error),
#[error(transparent)]
Workflow(#[from] taquba_workflow::Error),
#[error(transparent)]
ObjectStore(#[from] taquba::object_store::Error),
#[error(transparent)]
Record(#[from] RecordError),
#[error(transparent)]
Definition(#[from] DefinitionError),
#[error("definition `{0}` is not in the definition store")]
UnknownDefinition(String),
#[error("graph `{0}` does not have an adopted definition")]
UnknownGraph(String),
#[error("the trigger of graph `{0}` does not determine a partition")]
NoPartition(String),
#[error("node `{node}`: pool `{pool}` does not have a runtime")]
UnknownPool {
node: String,
pool: String,
},
#[error("graph `{graph}` does not have a run for partition `{partition}`")]
UnknownGraphRun {
graph: String,
partition: Partition,
},
#[error("graph `{graph}` does not have a node `{node}`")]
UnknownNode {
graph: String,
node: String,
},
#[error("the run of graph `{graph}` for partition `{partition}` changed during the transition")]
Contended {
graph: String,
partition: Partition,
},
}
impl From<ReadError> for Error {
fn from(e: ReadError) -> Self {
match e {
ReadError::Queue(e) => Error::Queue(e),
ReadError::Record(e) => Error::Record(e),
}
}
}
impl Error {
pub fn is_permanent(&self) -> bool {
match self {
Error::Queue(e) => e.is_permanent(),
Error::Workflow(e) => e.is_permanent(),
Error::Record(_)
| Error::UnknownGraph(_)
| Error::UnknownGraphRun { .. }
| Error::UnknownNode { .. }
| Error::NoPartition(_) => true,
Error::ObjectStore(_)
| Error::Definition(_)
| Error::UnknownDefinition(_)
| Error::UnknownPool { .. }
| Error::Contended { .. } => false,
}
}
}
#[derive(Debug, Clone)]
pub struct SchedulerOptions {
pub concurrency: usize,
pub poll_interval: Duration,
pub reconcile_interval: Duration,
}
impl Default for SchedulerOptions {
fn default() -> Self {
SchedulerOptions {
concurrency: 4,
poll_interval: Duration::from_millis(250),
reconcile_interval: Duration::from_secs(60),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StartOutcome {
pub started: bool,
pub submitted: Vec<RunId>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RerunOutcome {
Submitted(RunId),
Active(RunId),
NoRecord,
NotReady,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ReconcileReport {
pub active_runs: usize,
pub submitted: usize,
pub settled: usize,
pub cancelled: usize,
}
pub struct Scheduler {
pub(crate) queue: Arc<Queue>,
definitions: Arc<DefinitionStore>,
pools: Arc<Pools>,
pub(crate) clock: Arc<dyn Clock>,
pub(crate) expiry: ExpiryIndex,
}
impl Scheduler {
pub fn new(queue: Arc<Queue>, definitions: Arc<DefinitionStore>, pools: Arc<Pools>) -> Self {
let clock = queue.clock();
Scheduler {
queue,
definitions,
pools,
clock,
expiry: ExpiryIndex::new(EXPIRY_PREFIX),
}
}
pub fn definitions(&self) -> &Arc<DefinitionStore> {
&self.definitions
}
pub async fn graph_record(&self, graph: &str) -> Result<Option<GraphRecord>, Error> {
Ok(records::read(self.queue.view(), &records::graph_key(graph)).await?)
}
pub async fn start_run(
&self,
hash: &str,
partition: &Partition,
) -> Result<StartOutcome, Error> {
let graph = self.graph(hash).await?;
let record = GraphRunRecord {
definition: hash.to_string(),
requested_at_ms: self.clock.now_ms(),
state: GraphRunState::Active,
settled_at_ms: None,
expected_reruns: BTreeMap::new(),
};
let key = records::graph_run_key(graph.name(), partition);
if !self
.queue
.kv_compare_put(&key, None, &record.to_bytes())
.await?
{
return Ok(StartOutcome {
started: false,
submitted: Vec::new(),
});
}
let mut submitted = Vec::new();
for node in graph.roots() {
let (run_id, _) = self
.submit_node(
node,
identity(&graph, hash, partition, node, 0),
&BTreeMap::new(),
|_| SettlementEffects::default(),
)
.await?;
submitted.push(run_id);
}
Ok(StartOutcome {
started: true,
submitted,
})
}
pub async fn rerun(
&self,
graph_name: &str,
partition: &Partition,
node: &str,
) -> Result<RerunOutcome, Error> {
self.rerun_with(graph_name, partition, node, |_| {
SettlementEffects::default()
})
.await
}
pub(crate) async fn rerun_with(
&self,
graph_name: &str,
partition: &Partition,
node: &str,
effects: impl FnOnce(&RunId) -> SettlementEffects,
) -> Result<RerunOutcome, Error> {
let key = records::graph_run_key(graph_name, partition);
let Some((run, bytes)) = self.graph_run(&key).await? else {
return Err(Error::UnknownGraphRun {
graph: graph_name.to_string(),
partition: partition.clone(),
});
};
let graph = self.graph(&run.definition).await?;
let node = graph.node(node).ok_or_else(|| Error::UnknownNode {
graph: graph_name.to_string(),
node: node.to_string(),
})?;
let records = self.node_records(&graph, partition).await?;
let Some(record) = records.get(node.name()) else {
return Ok(RerunOutcome::NoRecord);
};
let mut active = GraphRunRecord {
state: GraphRunState::Active,
settled_at_ms: None,
..run.clone()
};
if record.status == RecordStatus::Succeeded && run.is_current(node.name(), record) {
for scoped in rerun_scope(&graph, node) {
if let Some(scoped_record) = records.get(scoped.name()) {
active
.expected_reruns
.insert(scoped.name().to_string(), scoped_record.rerun + 1);
}
}
}
let current = current_records(&active, &records);
if !is_ready(node, ¤t) {
return Ok(RerunOutcome::NotReady);
}
if active != run
&& !self
.queue
.kv_compare_put(&key, Some(&bytes), &active.to_bytes())
.await?
{
return Err(Error::Contended {
graph: graph_name.to_string(),
partition: partition.clone(),
});
}
let identity = identity(
&graph,
&record.definition,
partition,
node,
record.rerun + 1,
);
let (run_id, new) = self
.submit_node(node, identity, &upstream_records(node, ¤t), effects)
.await?;
Ok(if new {
RerunOutcome::Submitted(run_id)
} else {
RerunOutcome::Active(run_id)
})
}
pub async fn cancel_run(&self, graph_name: &str, partition: &Partition) -> Result<bool, Error> {
let key = records::graph_run_key(graph_name, partition);
let Some((run, bytes)) = self.graph_run(&key).await? else {
return Ok(false);
};
if run.state != GraphRunState::Active {
return Ok(false);
}
let cancelled = GraphRunRecord {
state: GraphRunState::Cancelled,
..run.clone()
};
if !self
.commit_settled(graph_name, partition, &bytes, cancelled)
.await?
{
return Ok(false);
}
let graph = self.graph(&run.definition).await?;
self.cancel_active_runs(&graph, partition, &run).await?;
Ok(true)
}
pub async fn handle_trigger(
&self,
graph_name: &str,
interval_start_ms: Option<u64>,
) -> Result<Option<Partition>, Error> {
let record = self.adopted(graph_name).await?;
let graph = self.graph(&record.definition).await?;
let partition = interval_start_ms
.and_then(|ms| Partition::of_time(graph.partitioning(), ms))
.ok_or_else(|| Error::NoPartition(graph_name.to_string()))?;
let started = self
.start_adopted(graph_name, &record, std::slice::from_ref(&partition))
.await?;
Ok(started.into_iter().next())
}
pub async fn start_runs(
&self,
graph_name: &str,
partitions: &[Partition],
) -> Result<Vec<Partition>, Error> {
let record = self.adopted(graph_name).await?;
self.start_adopted(graph_name, &record, partitions).await
}
async fn adopted(&self, graph_name: &str) -> Result<GraphRecord, Error> {
self.graph_record(graph_name)
.await?
.ok_or_else(|| Error::UnknownGraph(graph_name.to_string()))
}
async fn start_adopted(
&self,
graph_name: &str,
record: &GraphRecord,
partitions: &[Partition],
) -> Result<Vec<Partition>, Error> {
let mut started = Vec::new();
for partition in partitions {
if self.start_run(&record.definition, partition).await?.started {
tracing::info!(graph = %graph_name, %partition, "graph run started");
started.push(partition.clone());
}
}
Ok(started)
}
pub async fn handle_event(&self, event: &Event) -> Result<(), Error> {
let key = records::graph_run_key(&event.graph, &event.partition);
let Some((run, bytes)) = self.graph_run(&key).await? else {
return Ok(());
};
if run.state != GraphRunState::Active {
return Ok(());
}
let graph = self.graph(&run.definition).await?;
let Some(node) = graph.node(&event.node) else {
return Ok(());
};
let downstreams: Vec<&Node> = node
.downstreams()
.iter()
.map(|name| {
graph
.node(name)
.expect("a downstream name is a node of the graph")
})
.collect();
self.advance(&graph, &event.partition, &run, &bytes, downstreams)
.await?;
Ok(())
}
pub async fn reconcile(&self) -> Result<ReconcileReport, Error> {
let mut report = ReconcileReport::default();
let runs: Vec<Entry<GraphRunRecord>> =
records::scan(self.queue.view(), records::GRAPH_RUNS_PREFIX.as_bytes()).await?;
for Entry {
key,
bytes,
record: run,
} in runs
{
let Some((graph_name, partition)) = records::parse_graph_run_key(&key) else {
continue;
};
let graph = match self.definitions.get(&run.definition).await {
Ok(Some(graph)) => graph,
Ok(None) => {
tracing::warn!(graph = %graph_name, %partition, definition = %run.definition, "graph run records an unknown definition");
continue;
}
Err(e) => {
tracing::warn!(graph = %graph_name, %partition, definition = %run.definition, error = %e, "the definition of a graph run does not load");
continue;
}
};
match run.state {
GraphRunState::Active => {
report.active_runs += 1;
let (submitted, settled) = self
.advance(&graph, &partition, &run, &bytes, graph.nodes())
.await?;
report.submitted += submitted;
if settled {
report.settled += 1;
}
}
GraphRunState::Cancelled => {
report.cancelled += self.cancel_active_runs(&graph, &partition, &run).await?;
}
GraphRunState::Complete | GraphRunState::Failed => {}
}
}
Ok(report)
}
pub async fn run<F: Future<Output = ()>>(
self: Arc<Self>,
options: SchedulerOptions,
shutdown: F,
) -> Result<(), Error> {
let stop = CancellationToken::new();
let worker = taquba::run_worker_concurrent(
&self.queue,
EVENTS_QUEUE,
self.clone(),
options.concurrency,
options.poll_interval,
stop.clone().cancelled_owned(),
);
let triggers = taquba::run_worker_concurrent(
&self.queue,
TRIGGERS_QUEUE,
Arc::new(TriggerWorker::new(self.clone())),
options.concurrency,
options.poll_interval,
stop.clone().cancelled_owned(),
);
let reconciler = async {
let mut interval = tokio::time::interval(options.reconcile_interval);
interval.tick().await;
loop {
tokio::select! {
_ = interval.tick() => {
if let Err(e) = self.reconcile().await {
tracing::warn!(error = %e, "reconciler pass failed");
}
}
() = stop.cancelled() => return,
}
}
};
let mut all = std::pin::pin!(async {
let (worker, triggers, ()) = tokio::join!(worker, triggers, reconciler);
worker.and(triggers)
});
tokio::select! {
result = &mut all => Ok(result?),
() = shutdown => {
stop.cancel();
Ok(all.await?)
}
}
}
pub fn spawn<F>(
self: Arc<Self>,
options: SchedulerOptions,
shutdown: F,
) -> WorkerHandle<Result<(), Error>>
where
F: Future<Output = ()> + Send + 'static,
{
WorkerHandle::spawn(shutdown, move |stop| async move {
self.run(options, stop.cancelled_owned()).await
})
}
async fn graph(&self, hash: &str) -> Result<Arc<Graph>, Error> {
self.definitions
.get(hash)
.await?
.ok_or_else(|| Error::UnknownDefinition(hash.to_string()))
}
async fn graph_run(&self, key: &[u8]) -> Result<Option<(GraphRunRecord, Vec<u8>)>, Error> {
let Some(bytes) = self.queue.view().kv_get(key).await? else {
return Ok(None);
};
let record = records::parse::<GraphRunRecord>(key, &bytes)?;
Ok(Some((record, bytes.to_vec())))
}
async fn node_record(
&self,
graph: &Graph,
partition: &Partition,
node: &Node,
) -> Result<Option<NodeRecord>, Error> {
let key = records::node_record_key(graph.name(), partition, node);
Ok(records::read(self.queue.view(), &key).await?)
}
async fn node_records(
&self,
graph: &Graph,
partition: &Partition,
) -> Result<BTreeMap<String, NodeRecord>, Error> {
let mut records = BTreeMap::new();
for node in graph.nodes() {
if let Some(record) = self.node_record(graph, partition, node).await? {
records.insert(node.name().to_string(), record);
}
}
Ok(records)
}
async fn advance<'a>(
&self,
graph: &Graph,
partition: &Partition,
run: &GraphRunRecord,
bytes: &[u8],
candidates: impl IntoIterator<Item = &'a Node>,
) -> Result<(usize, bool), Error> {
let records = self.node_records(graph, partition).await?;
let current = current_records(run, &records);
let states = node_states(graph, ¤t);
let mut submitted = 0;
for node in candidates {
if states[node.name()] != NodeState::Ready {
continue;
}
let rerun = records.get(node.name()).map_or(0, |r| r.rerun + 1);
let (_, new) = self
.submit_node(
node,
identity(graph, &run.definition, partition, node, rerun),
&upstream_records(node, ¤t),
|_| SettlementEffects::default(),
)
.await?;
if new {
submitted += 1;
}
}
let settled = self
.settle(graph, partition, run, bytes, &records, &states)
.await?;
Ok((submitted, settled))
}
async fn submit_node(
&self,
node: &Node,
identity: TaskIdentity,
upstreams: &BTreeMap<String, NodeRecord>,
effects: impl FnOnce(&RunId) -> SettlementEffects,
) -> Result<(RunId, bool), Error> {
let runtime = self
.pools
.runtime(node.pool())
.ok_or_else(|| Error::UnknownPool {
node: node.name().to_string(),
pool: node.pool().to_string(),
})?;
let run_id = identity.run_id();
let outcome = runtime
.submit(RunSpec {
run_id: Some(run_id.clone()),
input: TaskInput::new(node, upstreams).to_bytes(),
options: RunOptions {
headers: identity.headers(),
max_attempts_per_step: Some(node.retries() + 1),
..RunOptions::default()
},
effects: effects(&run_id),
})
.await?;
if outcome.newly_submitted {
tracing::info!(run_id = %run_id, pool = node.pool(), "task instance submitted");
}
Ok((run_id, outcome.newly_submitted))
}
async fn settle(
&self,
graph: &Graph,
partition: &Partition,
run: &GraphRunRecord,
bytes: &[u8],
records: &BTreeMap<String, NodeRecord>,
states: &BTreeMap<String, NodeState>,
) -> Result<bool, Error> {
let Some(state) = settled_state(states) else {
return Ok(false);
};
for node in graph.nodes() {
if let Some(run_id) = unrecorded_run_id(graph, partition, run, node, records)
&& self.run_is_active(node, &run_id).await?
{
return Ok(false);
}
}
let mut settled = GraphRunRecord {
state,
..run.clone()
};
settled.expected_reruns.retain(|name, expected| {
records
.get(name)
.is_none_or(|record| record.rerun < *expected)
});
let written = self
.commit_settled(graph.name(), partition, bytes, settled)
.await?;
if written {
tracing::info!(graph = graph.name(), %partition, state = ?state, "graph run settled");
}
Ok(written)
}
async fn commit_settled(
&self,
graph_name: &str,
partition: &Partition,
expected: &[u8],
settled: GraphRunRecord,
) -> Result<bool, Error> {
let key = records::graph_run_key(graph_name, partition);
let settled_at_ms = self.clock.now_ms();
let settled = GraphRunRecord {
settled_at_ms: Some(settled_at_ms),
..settled
};
let expiring = Expiring::Run {
graph: graph_name.to_string(),
partition: partition.clone(),
};
let effects = SettlementEffects::default()
.kv_put(key.clone(), settled.to_bytes())
.expiry_entry(&self.expiry, settled_at_ms, &expiring.suffix());
Ok(self
.queue
.kv_compare_commit(&key, Some(expected), effects)
.await?
.is_some())
}
async fn run_is_active(&self, node: &Node, run_id: &RunId) -> Result<bool, Error> {
let Some(runtime) = self.pools.runtime(node.pool()) else {
return Ok(false);
};
Ok(matches!(
runtime.status(run_id).await?,
Some(status) if !matches!(status.state, RunState::Terminated(_))
))
}
async fn cancel_active_runs(
&self,
graph: &Graph,
partition: &Partition,
run: &GraphRunRecord,
) -> Result<usize, Error> {
let records = self.node_records(graph, partition).await?;
let mut cancelled = 0;
for node in graph.nodes() {
let Some(run_id) = unrecorded_run_id(graph, partition, run, node, &records) else {
continue;
};
let Some(runtime) = self.pools.runtime(node.pool()) else {
continue;
};
if runtime.cancel(&run_id).await? {
cancelled += 1;
}
}
Ok(cancelled)
}
}
fn upstream_records(
node: &Node,
records: &BTreeMap<String, NodeRecord>,
) -> BTreeMap<String, NodeRecord> {
node.upstreams()
.iter()
.filter_map(|name| Some((name.clone(), records.get(name)?.clone())))
.collect()
}
fn unrecorded_run_id(
graph: &Graph,
partition: &Partition,
run: &GraphRunRecord,
node: &Node,
records: &BTreeMap<String, NodeRecord>,
) -> Option<RunId> {
let rerun = match records.get(node.name()) {
Some(record)
if record.status == RecordStatus::Succeeded && run.is_current(node.name(), record) =>
{
return None;
}
Some(record) => record.rerun + 1,
None => 0,
};
Some(task::run_id(graph.name(), partition, node.name(), rerun))
}
fn identity(
graph: &Graph,
hash: &str,
partition: &Partition,
node: &Node,
rerun: u32,
) -> TaskIdentity {
TaskIdentity {
graph: graph.name().to_string(),
partition: partition.clone(),
node: node.name().to_string(),
asset: node.asset().map(str::to_string),
definition: hash.to_string(),
rerun,
}
}
fn worker_error(error: Error) -> WorkerError {
if error.is_permanent() {
PermanentFailure::new(error.to_string()).into()
} else {
Box::new(error)
}
}
pub fn firing_headers(graph: &str) -> HashMap<String, String> {
HashMap::from([(HEADER_GRAPH.to_string(), graph.to_string())])
}
pub struct TriggerWorker {
scheduler: Arc<Scheduler>,
}
impl TriggerWorker {
pub fn new(scheduler: Arc<Scheduler>) -> Self {
TriggerWorker { scheduler }
}
}
impl Worker for TriggerWorker {
async fn process(&self, job: &JobRecord, _lease: &LeaseHandle) -> Result<(), WorkerError> {
let graph = job.headers.get(HEADER_GRAPH).ok_or_else(|| {
PermanentFailure::new(format!("the job does not have the `{HEADER_GRAPH}` header"))
})?;
let interval_start_ms = job
.headers
.get(PREVIOUS_FIRE_MS_HEADER)
.and_then(|value| value.parse().ok());
self.scheduler
.handle_trigger(graph, interval_start_ms)
.await
.map(|_| ())
.map_err(worker_error)
}
}
impl Worker for Scheduler {
async fn process(&self, job: &JobRecord, _lease: &LeaseHandle) -> Result<(), WorkerError> {
let event = Event::from_bytes(&job.payload)
.map_err(|e| PermanentFailure::new(format!("the payload is not an event: {e}")))?;
self.handle_event(&event).await.map_err(worker_error)
}
}