use std::collections::HashMap;
use std::convert::TryInto;
use std::fmt::Write;
use std::io::{stderr, stdout, Write as WriteIo};
use std::sync::{Arc, Mutex, RwLock};
use std::thread;
use std::time::{Duration, Instant};
use indicatif::{MultiProgress, ProgressBar, ProgressDrawTarget, ProgressStyle};
use itertools::Itertools;
use lazy_static::lazy_static;
use tracing::warn;
use crate::core::formatting::Glyphs;
#[allow(missing_docs)]
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum OperationType {
BuildRebasePlan,
CalculateDiff,
CalculatePatchId,
CheckForCycles,
CheckTouchedPaths,
DetectDuplicateCommits,
FilterCommits,
FindPathToMergeBase,
GetTouchedPaths,
GetMergeBase,
GetUpstreamPatchIds,
InitializeRebase,
MakeGraph,
ProcessEvents,
WalkCommits,
}
impl ToString for OperationType {
fn to_string(&self) -> String {
let s = match self {
OperationType::BuildRebasePlan => "Building rebase plan",
OperationType::CalculateDiff => "Computing diffs",
OperationType::CalculatePatchId => "Hashing commit contents",
OperationType::CheckForCycles => "Checking for cycles",
OperationType::CheckTouchedPaths => "Checking touched paths",
OperationType::DetectDuplicateCommits => "Checking for duplicate commits",
OperationType::FilterCommits => "Filtering commits",
OperationType::FindPathToMergeBase => "Finding path to merge-base",
OperationType::GetMergeBase => "Calculating merge-bases",
OperationType::GetTouchedPaths => "Getting touched paths",
OperationType::GetUpstreamPatchIds => "Enumerating patch IDs",
OperationType::InitializeRebase => "Initializing rebase",
OperationType::MakeGraph => "Examining local history",
OperationType::ProcessEvents => "Processing events",
OperationType::WalkCommits => "Walking commits",
};
s.to_string()
}
}
#[derive(Clone, Debug)]
enum OutputDest {
Stdout,
Suppress,
BufferForTest(Arc<Mutex<Vec<u8>>>),
}
#[derive(Debug)]
struct OperationState {
operation_type: OperationType,
progress_bar: ProgressBar,
has_meter: bool,
start_times: Vec<Instant>,
elapsed_duration: Duration,
}
impl OperationState {
pub fn set_progress(&mut self, current: usize, total: usize) {
self.has_meter = true;
self.progress_bar.set_position(current.try_into().unwrap());
self.progress_bar.set_length(total.try_into().unwrap());
}
pub fn inc_progress(&mut self, increment: usize) {
self.progress_bar.inc(increment.try_into().unwrap());
}
pub fn tick(&self) {
lazy_static! {
static ref CHECKMARK: String = console::style("✓").green().to_string();
static ref IN_PROGRESS_SPINNER_STYLE: ProgressStyle =
ProgressStyle::default_spinner().template("{prefix}{spinner} {msg}");
static ref IN_PROGRESS_BAR_STYLE: ProgressStyle =
ProgressStyle::default_bar().template("{prefix}{spinner} {msg} {bar} {pos}/{len}");
static ref FINISHED_PROGRESS_STYLE: ProgressStyle = IN_PROGRESS_SPINNER_STYLE
.clone()
.tick_strings(&[&CHECKMARK, &CHECKMARK]);
}
let elapsed_duration = match self.start_times.iter().min() {
None => self.elapsed_duration,
Some(start_time) => {
let additional_duration = Instant::now().saturating_duration_since(*start_time);
self.elapsed_duration + additional_duration
}
};
self.progress_bar.set_message(format!(
"{} ({:.1}s)",
self.operation_type.to_string(),
elapsed_duration.as_secs_f64(),
));
self.progress_bar
.set_style(match (self.start_times.as_slice(), self.has_meter) {
([], _) => FINISHED_PROGRESS_STYLE.clone(),
([..], false) => IN_PROGRESS_SPINNER_STYLE.clone(),
([..], true) => IN_PROGRESS_BAR_STYLE.clone(),
});
self.progress_bar.tick();
}
}
#[derive(Clone)]
pub struct Effects {
glyphs: Glyphs,
dest: OutputDest,
multi_progress: Arc<MultiProgress>,
nesting_level: usize,
operation_states: Arc<RwLock<HashMap<OperationType, OperationState>>>,
}
impl std::fmt::Debug for Effects {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"<Output fancy={}>",
self.glyphs.should_write_ansi_escape_codes
)
}
}
fn spawn_progress_updater_thread(
multi_progress: &Arc<MultiProgress>,
operation_states: &Arc<RwLock<HashMap<OperationType, OperationState>>>,
) {
multi_progress.set_draw_target(ProgressDrawTarget::hidden());
let multi_progress = Arc::downgrade(multi_progress);
let operation_states = Arc::downgrade(operation_states);
thread::spawn(move || {
thread::sleep(Duration::from_millis(250));
if let Some(multi_progress) = multi_progress.upgrade() {
multi_progress.set_draw_target(ProgressDrawTarget::stderr());
}
loop {
match operation_states.upgrade() {
None => return,
Some(operation_states) => {
let operation_states = operation_states.read().unwrap();
for operation_state in operation_states.values() {
operation_state.tick();
}
}
}
thread::sleep(Duration::from_millis(100));
}
});
}
impl Effects {
pub fn new(glyphs: Glyphs) -> Self {
let multi_progress = Default::default();
let operation_states = Default::default();
spawn_progress_updater_thread(&multi_progress, &operation_states);
Effects {
glyphs,
dest: OutputDest::Stdout,
multi_progress,
nesting_level: Default::default(),
operation_states,
}
}
pub fn new_suppress_for_test(glyphs: Glyphs) -> Self {
Effects {
glyphs,
dest: OutputDest::Suppress,
multi_progress: Default::default(),
nesting_level: Default::default(),
operation_states: Default::default(),
}
}
pub fn new_from_buffer_for_test(glyphs: Glyphs, buffer: &Arc<Mutex<Vec<u8>>>) -> Self {
Effects {
glyphs,
dest: OutputDest::BufferForTest(Arc::clone(buffer)),
multi_progress: Default::default(),
nesting_level: Default::default(),
operation_states: Default::default(),
}
}
pub fn enable_tui_mode(&self) -> Self {
let multi_progress = Arc::clone(&self.multi_progress);
multi_progress.set_draw_target(ProgressDrawTarget::hidden());
Self {
dest: OutputDest::Suppress,
..self.clone()
}
}
pub fn start_operation(&self, operation_type: OperationType) -> (Effects, ProgressHandle) {
let progress = ProgressHandle {
effects: self,
operation_type,
};
match self.dest {
OutputDest::Stdout => {}
OutputDest::Suppress | OutputDest::BufferForTest(_) => return (self.clone(), progress),
}
let now = Instant::now();
let mut operation_states = self.operation_states.write().unwrap();
let mut nesting_level = self.nesting_level;
let operation_state = operation_states.entry(operation_type).or_insert_with(|| {
let progress_bar = self.multi_progress.add(ProgressBar::new_spinner());
nesting_level += 1;
progress_bar.set_prefix(" ".repeat(nesting_level));
let operation_state = OperationState {
operation_type,
progress_bar,
start_times: Vec::new(),
has_meter: false,
elapsed_duration: Default::default(),
};
operation_state.tick();
operation_state
});
operation_state.start_times.push(now);
let effects = Self {
nesting_level,
..self.clone()
};
(effects, progress)
}
fn on_notify_progress(&self, operation_type: OperationType, current: usize, total: usize) {
let mut operation_states = self.operation_states.write().unwrap();
let operation_state = match operation_states.get_mut(&operation_type) {
Some(operation_state) => operation_state,
None => return,
};
operation_state.set_progress(current, total);
}
fn on_notify_progress_inc(&self, operation_type: OperationType, increment: usize) {
let mut operation_states = self.operation_states.write().unwrap();
let operation_state = match operation_states.get_mut(&operation_type) {
Some(operation_state) => operation_state,
None => return,
};
operation_state.inc_progress(increment);
}
fn on_drop_progress_handle(&self, operation_type: OperationType) {
match self.dest {
OutputDest::Stdout => {}
OutputDest::Suppress | OutputDest::BufferForTest(_) => return,
}
let now = Instant::now();
let mut operation_states = self.operation_states.write().unwrap();
let operation_state = match operation_states.get_mut(&operation_type) {
Some(operation_state) => operation_state,
None => {
warn!("Progress operation not started");
return;
}
};
let previous_start_time = match operation_state
.start_times
.iter()
.position_max_by_key(|x| *x)
{
Some(start_time_index) => operation_state.start_times.remove(start_time_index),
None => {
warn!("Progress operation ended without matching start call");
return;
}
};
operation_state.elapsed_duration += if operation_state.start_times.is_empty() {
now.saturating_duration_since(previous_start_time)
} else {
Duration::ZERO
};
if operation_states
.values()
.all(|operation_state| operation_state.start_times.is_empty())
{
self.multi_progress.clear().unwrap();
operation_states.clear();
}
}
pub fn get_glyphs(&self) -> &Glyphs {
&self.glyphs
}
pub fn get_output_stream(&self) -> OutputStream {
OutputStream {
dest: self.dest.clone(),
}
}
pub fn get_error_stream(&self) -> ErrorStream {
ErrorStream {
dest: self.dest.clone(),
}
}
}
pub struct OutputStream {
dest: OutputDest,
}
impl Write for OutputStream {
fn write_str(&mut self, s: &str) -> std::fmt::Result {
match &self.dest {
OutputDest::Stdout => {
print!("{}", s);
stdout().flush().unwrap();
}
OutputDest::Suppress => {
}
OutputDest::BufferForTest(buffer) => {
let mut buffer = buffer.lock().unwrap();
write!(buffer, "{}", s).unwrap();
}
}
Ok(())
}
}
pub struct ErrorStream {
dest: OutputDest,
}
impl Write for ErrorStream {
fn write_str(&mut self, s: &str) -> std::fmt::Result {
match &self.dest {
OutputDest::Stdout => {
eprint!("{}", s);
stderr().flush().unwrap();
}
OutputDest::Suppress => {
}
OutputDest::BufferForTest(_) => {
}
}
Ok(())
}
}
pub struct ProgressHandle<'a> {
effects: &'a Effects,
operation_type: OperationType,
}
impl Drop for ProgressHandle<'_> {
fn drop(&mut self) {
self.effects.on_drop_progress_handle(self.operation_type)
}
}
impl ProgressHandle<'_> {
pub fn notify_progress(&self, current: usize, total: usize) {
self.effects
.on_notify_progress(self.operation_type, current, total);
}
pub fn notify_progress_inc(&self, increment: usize) {
self.effects
.on_notify_progress_inc(self.operation_type, increment);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_effects_progress() -> eyre::Result<()> {
let effects = Effects::new(Glyphs::text());
let (effects2, progress2) = effects.start_operation(OperationType::GetMergeBase);
{
let operation_states = effects.operation_states.read().unwrap();
let get_merge_base_operation =
operation_states.get(&OperationType::GetMergeBase).unwrap();
assert_eq!(get_merge_base_operation.start_times.len(), 1);
}
std::thread::sleep(Duration::from_millis(1));
let (_effects3, progress3) = effects.start_operation(OperationType::GetMergeBase);
let earlier_start_time = {
let operation_states = effects.operation_states.read().unwrap();
let get_merge_base_operation =
operation_states.get(&OperationType::GetMergeBase).unwrap();
assert_eq!(get_merge_base_operation.start_times.len(), 2);
get_merge_base_operation.start_times[0]
};
drop(progress3);
{
let operation_states = effects.operation_states.read().unwrap();
let get_merge_base_operation =
operation_states.get(&OperationType::GetMergeBase).unwrap();
assert_eq!(
get_merge_base_operation.start_times,
vec![earlier_start_time]
);
}
let (_effects4, progress4) = effects2.start_operation(OperationType::CalculateDiff);
std::thread::sleep(Duration::from_millis(1));
drop(progress4);
{
let operation_states = effects.operation_states.read().unwrap();
let calculate_diff_operation =
operation_states.get(&OperationType::CalculateDiff).unwrap();
assert!(calculate_diff_operation.start_times.is_empty());
assert!(calculate_diff_operation.elapsed_duration >= Duration::from_millis(1));
}
drop(progress2);
{
let operation_states = effects.operation_states.read().unwrap();
assert!(operation_states.is_empty());
}
Ok(())
}
}