use std::collections::{HashMap, HashSet};
use std::fmt;
use std::io::{self, IsTerminal};
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering};
use std::sync::mpsc::{self, Receiver, RecvTimeoutError, Sender};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use console::Term;
use indicatif::{MultiProgress, ProgressBar, ProgressDrawTarget, ProgressStyle};
use serde_json::Value;
use super::error::ReportingError;
use crate::configuration::{ProjectConfig, TaskConfig};
use crate::project::ScientificProject;
const REFRESH_INTERVAL: Duration = Duration::from_millis(100);
static TERMINAL_OWNED: AtomicBool = AtomicBool::new(false);
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum TaskStatus {
Pending,
Running,
Completed,
Failed,
}
impl TaskStatus {
fn encode(self) -> u8 {
match self {
Self::Pending => 0,
Self::Running => 1,
Self::Completed => 2,
Self::Failed => 3,
}
}
fn decode(value: u8) -> Self {
match value {
0 => Self::Pending,
1 => Self::Running,
2 => Self::Completed,
3 => Self::Failed,
_ => unreachable!("task status is written only through TaskStatus::encode"),
}
}
fn label(self) -> &'static str {
match self {
Self::Pending => "pending",
Self::Running => "running",
Self::Completed => "completed",
Self::Failed => "failed",
}
}
}
#[derive(Clone, Debug)]
pub struct TaskIdentity {
task: TaskConfig,
keys: Arc<[Box<str>]>,
label: Arc<str>,
}
impl TaskIdentity {
pub fn label(&self) -> &str {
&self.label
}
pub fn len(&self) -> usize {
self.keys.len()
}
pub fn is_empty(&self) -> bool {
self.keys.is_empty()
}
pub fn value(&self, key: &str) -> Option<&Value> {
self.keys.iter().find_map(|name| {
(name.as_ref() == key)
.then(|| self.task.value(key))
.flatten()
})
}
pub fn iter(&self) -> impl ExactSizeIterator<Item = (&str, &Value)> {
self.keys.iter().map(|name| {
(
name.as_ref(),
self.task
.value(name)
.expect("validated identity keys resolve for every task"),
)
})
}
fn matches(&self, task: &TaskConfig) -> bool {
self.keys
.iter()
.all(|key| task.value(key) == self.task.value(key))
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ProgressSummary {
total: u64,
pending: u64,
running: u64,
completed: u64,
failed: u64,
}
impl ProgressSummary {
pub fn total(&self) -> u64 {
self.total
}
pub fn pending(&self) -> u64 {
self.pending
}
pub fn running(&self) -> u64 {
self.running
}
pub fn completed(&self) -> u64 {
self.completed
}
pub fn failed(&self) -> u64 {
self.failed
}
pub fn is_success(&self) -> bool {
self.completed == self.total && self.pending == 0 && self.running == 0 && self.failed == 0
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum OutputMode {
Auto,
Terminal,
Plain,
Hidden,
}
pub struct ProgressReporterBuilder {
configuration: ProjectConfig,
identity_keys: Option<Vec<String>>,
output: OutputMode,
}
impl ProgressReporterBuilder {
pub fn identify_tasks_by<I, S>(mut self, keys: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.identity_keys = Some(keys.into_iter().map(Into::into).collect());
self
}
pub fn terminal(mut self) -> Self {
self.output = OutputMode::Terminal;
self
}
pub fn plain(mut self) -> Self {
self.output = OutputMode::Plain;
self
}
pub fn hidden(mut self) -> Self {
self.output = OutputMode::Hidden;
self
}
pub fn start(self) -> Result<ProgressReporter, ReportingError> {
let identity_keys: Arc<[Box<str>]> =
validate_identity_keys(&self.configuration, self.identity_keys)?
.into_iter()
.map(String::into_boxed_str)
.collect();
let slots = build_slots(&self.configuration, Arc::clone(&identity_keys))?;
let output = match self.output {
OutputMode::Auto if io::stderr().is_terminal() => OutputMode::Terminal,
OutputMode::Auto => OutputMode::Plain,
explicit => explicit,
};
acquire_terminal()?;
let lease = TerminalLease;
let (events, receiver) = mpsc::channel();
let renderer_slots = Arc::clone(&slots);
let renderer = thread::Builder::new()
.name("scientific-workflow-progress".to_owned())
.spawn(move || render(receiver, renderer_slots, output, lease))
.map_err(|source| ReportingError::StartRenderer { source })?;
Ok(ProgressReporter {
inner: Arc::new(ReporterInner {
slots,
identity_keys,
events,
}),
renderer: Some(renderer),
finished: false,
})
}
}
impl fmt::Debug for ProgressReporterBuilder {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ProgressReporterBuilder")
.field("tasks", &self.configuration.task_count())
.field("identity_keys", &self.identity_keys)
.field("output", &self.output)
.finish_non_exhaustive()
}
}
pub struct ProgressReporter {
inner: Arc<ReporterInner>,
renderer: Option<JoinHandle<()>>,
finished: bool,
}
impl ProgressReporter {
pub fn for_project(project: &ScientificProject) -> ProgressReporterBuilder {
Self::for_configuration(project.configuration())
}
pub fn for_configuration(configuration: &ProjectConfig) -> ProgressReporterBuilder {
ProgressReporterBuilder {
configuration: configuration.clone(),
identity_keys: None,
output: OutputMode::Auto,
}
}
pub fn start_task(
&self,
task: &TaskConfig,
initial_iteration: u64,
target_iteration: Option<u64>,
) -> Result<TaskProgress, ReportingError> {
let ordinal = task.task_ordinal();
let index = usize::try_from(ordinal)
.ok()
.filter(|index| *index < self.inner.slots.len())
.ok_or(ReportingError::UnknownTaskOrdinal {
task_ordinal: ordinal,
})?;
let slot = Arc::clone(&self.inner.slots[index]);
if !slot.identity.matches(task) {
return Err(ReportingError::TaskIdentityMismatch {
task_ordinal: ordinal,
});
}
if let Some(target) = target_iteration.filter(|target| initial_iteration > *target) {
return Err(ReportingError::InitialIterationBeyondTarget {
identity: slot.identity.label().to_owned(),
initial: initial_iteration,
target,
});
}
slot.status
.compare_exchange(
TaskStatus::Pending.encode(),
TaskStatus::Running.encode(),
Ordering::AcqRel,
Ordering::Acquire,
)
.map_err(|_| ReportingError::TaskAlreadyStarted {
identity: slot.identity.label().to_owned(),
})?;
slot.current.store(initial_iteration, Ordering::Relaxed);
if let Some(target) = target_iteration {
slot.target.store(target, Ordering::Relaxed);
slot.target_known.store(true, Ordering::Release);
} else {
slot.target_known.store(false, Ordering::Release);
}
*lock(&slot.phase) = "running".into();
Ok(TaskProgress {
slot,
events: self.inner.events.clone(),
active: true,
})
}
pub fn report(&self, message: impl Into<String>) -> Result<(), ReportingError> {
self.inner
.events
.send(RenderEvent::Message(message.into()))
.map_err(|_| ReportingError::RendererUnavailable)
}
pub fn summary(&self) -> ProgressSummary {
summarize(&self.inner.slots)
}
pub fn complete(
mut self,
message: impl Into<String>,
) -> Result<ProgressSummary, ReportingError> {
let summary = self.summary();
if !summary.is_success() {
self.stop(false, "workflow did not complete".to_owned())?;
return Err(ReportingError::IncompleteProgress {
pending: summary.pending,
running: summary.running,
failed: summary.failed,
});
}
self.stop(true, message.into())?;
Ok(summary)
}
pub fn fail(mut self, message: impl Into<String>) -> Result<ProgressSummary, ReportingError> {
let summary = self.summary();
self.stop(false, message.into())?;
Ok(summary)
}
pub fn report_error(message: impl fmt::Display) {
eprintln!("[error] {message}");
}
fn stop(&mut self, success: bool, message: String) -> Result<(), ReportingError> {
self.inner
.events
.send(RenderEvent::Stop { success, message })
.map_err(|_| ReportingError::RendererUnavailable)?;
self.finished = true;
if self
.renderer
.take()
.expect("an unfinished reporter owns one renderer")
.join()
.is_err()
{
return Err(ReportingError::RendererPanicked);
}
Ok(())
}
}
impl fmt::Debug for ProgressReporter {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ProgressReporter")
.field("tasks", &self.inner.slots.len())
.field("identity_keys", &self.inner.identity_keys)
.field("summary", &self.summary())
.finish_non_exhaustive()
}
}
impl Drop for ProgressReporter {
fn drop(&mut self) {
if self.finished {
return;
}
let _ = self.inner.events.send(RenderEvent::Stop {
success: false,
message: "progress reporter dropped before completion".to_owned(),
});
if let Some(renderer) = self.renderer.take() {
let _ = renderer.join();
}
}
}
pub struct TaskProgress {
slot: Arc<ProgressSlot>,
events: Sender<RenderEvent>,
active: bool,
}
impl TaskProgress {
pub fn identity(&self) -> &TaskIdentity {
&self.slot.identity
}
pub fn current_iteration(&self) -> u64 {
self.slot.current.load(Ordering::Relaxed)
}
pub fn target_iteration(&self) -> Option<u64> {
if self.slot.target_known.load(Ordering::Acquire) {
Some(self.slot.target.load(Ordering::Relaxed))
} else {
None
}
}
pub fn status(&self) -> TaskStatus {
TaskStatus::decode(self.slot.status.load(Ordering::Acquire))
}
pub fn set_iteration(&self, iteration: u64) -> Result<(), ReportingError> {
if let Some(target) = self.target_iteration().filter(|target| iteration > *target) {
return Err(ReportingError::IterationBeyondTarget {
identity: self.identity().label().to_owned(),
iteration,
target,
});
}
let previous = self.slot.current.fetch_max(iteration, Ordering::Relaxed);
if iteration < previous {
return Err(ReportingError::IterationRegressed {
identity: self.identity().label().to_owned(),
current: previous,
attempted: iteration,
});
}
Ok(())
}
pub fn set_phase(&self, phase: impl Into<String>) {
*lock(&self.slot.phase) = phase.into().into_boxed_str();
}
pub fn report(&self, message: impl Into<String>) -> Result<(), ReportingError> {
self.events
.send(RenderEvent::TaskMessage {
identity: self.identity().label().to_owned(),
message: message.into(),
})
.map_err(|_| ReportingError::RendererUnavailable)
}
pub fn complete(mut self) -> Result<(), ReportingError> {
if let Some(target) = self.target_iteration() {
let current = self.current_iteration();
if current != target {
*lock(&self.slot.phase) = "target not reached".into();
self.slot
.status
.store(TaskStatus::Failed.encode(), Ordering::Release);
self.active = false;
return Err(ReportingError::TargetIterationNotReached {
identity: self.identity().label().to_owned(),
current,
target,
});
}
}
*lock(&self.slot.phase) = "completed".into();
self.slot
.status
.store(TaskStatus::Completed.encode(), Ordering::Release);
self.active = false;
Ok(())
}
pub fn fail(mut self, reason: impl Into<String>) {
*lock(&self.slot.phase) = reason.into().into_boxed_str();
self.slot
.status
.store(TaskStatus::Failed.encode(), Ordering::Release);
self.active = false;
}
}
impl fmt::Debug for TaskProgress {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("TaskProgress")
.field("identity", &self.identity().label())
.field("current_iteration", &self.current_iteration())
.field("target_iteration", &self.target_iteration())
.field("active", &self.active)
.finish_non_exhaustive()
}
}
impl Drop for TaskProgress {
fn drop(&mut self) {
if self.active {
*lock(&self.slot.phase) = "interrupted".into();
self.slot
.status
.store(TaskStatus::Failed.encode(), Ordering::Release);
}
}
}
struct ReporterInner {
slots: Arc<[Arc<ProgressSlot>]>,
identity_keys: Arc<[Box<str>]>,
events: Sender<RenderEvent>,
}
struct ProgressSlot {
identity: TaskIdentity,
current: AtomicU64,
target: AtomicU64,
target_known: AtomicBool,
status: AtomicU8,
phase: Mutex<Box<str>>,
}
enum RenderEvent {
Message(String),
TaskMessage { identity: String, message: String },
Stop { success: bool, message: String },
}
struct TerminalLease;
impl Drop for TerminalLease {
fn drop(&mut self) {
TERMINAL_OWNED.store(false, Ordering::Release);
}
}
fn acquire_terminal() -> Result<(), ReportingError> {
TERMINAL_OWNED
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.map(|_| ())
.map_err(|_| ReportingError::TerminalAlreadyOwned)
}
fn validate_identity_keys(
configuration: &ProjectConfig,
requested: Option<Vec<String>>,
) -> Result<Vec<String>, ReportingError> {
let keys = requested.unwrap_or_else(|| {
configuration
.parameters()
.sweep_keys()
.map(str::to_owned)
.collect()
});
let mut seen = HashSet::with_capacity(keys.len());
for key in &keys {
if !seen.insert(key.as_str()) {
return Err(ReportingError::DuplicateIdentityParameter { key: key.clone() });
}
if !configuration.parameters().contains_parameter(key) {
return Err(ReportingError::UnknownIdentityParameter { key: key.clone() });
}
}
Ok(keys)
}
fn build_slots(
configuration: &ProjectConfig,
keys: Arc<[Box<str>]>,
) -> Result<Arc<[Arc<ProgressSlot>]>, ReportingError> {
let capacity = usize::try_from(configuration.task_count()).map_err(|_| {
ReportingError::TaskCountTooLarge {
task_count: configuration.task_count(),
}
})?;
let mut slots = Vec::with_capacity(capacity);
let mut identities = HashMap::<String, u64>::with_capacity(capacity);
for task in configuration.task_configs() {
let label = render_identity(&task, &keys);
if let Some(first_ordinal) = identities.insert(label.clone(), task.task_ordinal()) {
return Err(ReportingError::NonUniqueTaskIdentity {
identity: label,
first_ordinal,
second_ordinal: task.task_ordinal(),
});
}
slots.push(Arc::new(ProgressSlot {
identity: TaskIdentity {
task,
keys: Arc::clone(&keys),
label: label.into(),
},
current: AtomicU64::new(0),
target: AtomicU64::new(0),
target_known: AtomicBool::new(false),
status: AtomicU8::new(TaskStatus::Pending.encode()),
phase: Mutex::new("pending".into()),
}));
}
Ok(slots.into())
}
fn render_identity(task: &TaskConfig, keys: &[Box<str>]) -> String {
if keys.is_empty() {
return "task".to_owned();
}
keys.iter()
.map(|key| {
let value = task
.value(key)
.expect("validated identity keys resolve for every task");
let value = serde_json::to_string(value)
.expect("serde_json::Value always serializes to valid JSON");
format!("{key}={value}")
})
.collect::<Vec<_>>()
.join(", ")
}
fn summarize(slots: &[Arc<ProgressSlot>]) -> ProgressSummary {
let mut summary = ProgressSummary {
total: u64::try_from(slots.len()).expect("slot count originated from a u64 task count"),
pending: 0,
running: 0,
completed: 0,
failed: 0,
};
for slot in slots {
match TaskStatus::decode(slot.status.load(Ordering::Acquire)) {
TaskStatus::Pending => summary.pending += 1,
TaskStatus::Running => summary.running += 1,
TaskStatus::Completed => summary.completed += 1,
TaskStatus::Failed => summary.failed += 1,
}
}
summary
}
fn render(
receiver: Receiver<RenderEvent>,
slots: Arc<[Arc<ProgressSlot>]>,
output: OutputMode,
_lease: TerminalLease,
) {
if output == OutputMode::Terminal {
let _ = Term::stderr().clear_screen();
}
let mut terminal = (output == OutputMode::Terminal).then(|| TerminalDisplay::new(&slots));
let mut last_statuses = vec![TaskStatus::Pending; slots.len()];
loop {
if let Some(display) = &mut terminal {
display.refresh(&slots);
}
match receiver.recv_timeout(REFRESH_INTERVAL) {
Ok(RenderEvent::Message(message)) => write_message(output, terminal.as_ref(), &message),
Ok(RenderEvent::TaskMessage { identity, message }) => {
write_message(output, terminal.as_ref(), &format!("{identity}: {message}"));
}
Ok(RenderEvent::Stop { success, message }) => {
if let Some(display) = &mut terminal {
display.refresh(&slots);
display.finish(&slots);
}
if output == OutputMode::Plain {
write_plain_transitions(&slots, &mut last_statuses);
}
write_final(output, &slots, success, &message);
break;
}
Err(RecvTimeoutError::Timeout) => {
if output == OutputMode::Plain {
write_plain_transitions(&slots, &mut last_statuses);
}
}
Err(RecvTimeoutError::Disconnected) => break,
}
}
}
struct TerminalDisplay {
multi: MultiProgress,
bars: Vec<ProgressBar>,
statuses: Vec<TaskStatus>,
known_style: ProgressStyle,
unknown_style: ProgressStyle,
}
impl TerminalDisplay {
fn new(slots: &[Arc<ProgressSlot>]) -> Self {
let multi = MultiProgress::with_draw_target(ProgressDrawTarget::stderr());
let known_style = ProgressStyle::with_template(
"{prefix:.bold} [{msg}] {wide_bar:.cyan/blue} {pos}/{len} elapsed {elapsed_precise} ETA {eta_precise}",
)
.expect("hard-coded progress template is valid");
let unknown_style = ProgressStyle::with_template(
"{prefix:.bold} [{msg}] {spinner:.cyan} iteration {pos} elapsed {elapsed_precise} ETA unknown",
)
.expect("hard-coded spinner template is valid");
let bars: Vec<_> = slots
.iter()
.map(|slot| {
let bar = multi.add(ProgressBar::new_spinner());
bar.set_prefix(slot.identity.label().to_owned());
bar.set_style(unknown_style.clone());
bar.set_message("pending");
bar
})
.collect();
for bar in &bars {
bar.force_draw();
}
Self {
multi,
bars,
statuses: vec![TaskStatus::Pending; slots.len()],
known_style,
unknown_style,
}
}
fn refresh(&mut self, slots: &[Arc<ProgressSlot>]) {
for ((bar, previous_status), slot) in self.bars.iter().zip(&mut self.statuses).zip(slots) {
if !slot.target_known.load(Ordering::Acquire) {
bar.set_style(self.unknown_style.clone());
} else {
let target = slot.target.load(Ordering::Relaxed);
bar.set_style(self.known_style.clone());
bar.set_length(target);
}
bar.set_position(slot.current.load(Ordering::Relaxed));
let status = TaskStatus::decode(slot.status.load(Ordering::Acquire));
if *previous_status == TaskStatus::Pending && status == TaskStatus::Running {
bar.reset_elapsed();
}
let phase = lock(&slot.phase);
if phase.is_empty() || phase.as_ref() == status.label() {
bar.set_message(status.label());
} else {
bar.set_message(format!("{}: {}", status.label(), phase.as_ref()));
}
bar.tick();
*previous_status = status;
}
}
fn finish(&self, slots: &[Arc<ProgressSlot>]) {
for (bar, slot) in self.bars.iter().zip(slots) {
let status = TaskStatus::decode(slot.status.load(Ordering::Acquire));
bar.finish_with_message(status.label());
}
let _ = self.multi.clear();
}
}
fn write_message(output: OutputMode, terminal: Option<&TerminalDisplay>, message: &str) {
match output {
OutputMode::Terminal => {
if let Some(display) = terminal {
let _ = display.multi.println(message);
}
}
OutputMode::Plain => eprintln!("[progress] {message}"),
OutputMode::Hidden | OutputMode::Auto => {}
}
}
fn write_plain_transitions(slots: &[Arc<ProgressSlot>], previous: &mut [TaskStatus]) {
for (slot, old) in slots.iter().zip(previous) {
let status = TaskStatus::decode(slot.status.load(Ordering::Acquire));
if status != *old {
let phase = lock(&slot.phase);
eprintln!(
"[task] identity={} status={} phase={} iteration={} target={}",
slot.identity.label(),
status.label(),
phase.as_ref(),
slot.current.load(Ordering::Relaxed),
format_target(slot)
);
*old = status;
}
}
}
fn write_final(output: OutputMode, slots: &[Arc<ProgressSlot>], success: bool, message: &str) {
if output == OutputMode::Hidden {
return;
}
let summary = summarize(slots);
eprintln!(
"[workflow] status={} tasks={} completed={} failed={} pending={} message={}",
if success { "completed" } else { "failed" },
summary.total,
summary.completed,
summary.failed,
summary.pending,
message
);
}
fn format_target(slot: &ProgressSlot) -> String {
if slot.target_known.load(Ordering::Acquire) {
slot.target.load(Ordering::Relaxed).to_string()
} else {
"unknown".to_owned()
}
}
fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}