Skip to main content

studiole_command/services/
cli_progress.rs

1use crate::prelude::*;
2use indicatif::ProgressBar;
3use std::sync::atomic::{AtomicBool, Ordering};
4use tokio::spawn;
5use tokio::sync::broadcast::error::RecvError;
6use tracing::{error, warn};
7
8/// Display command progress as a terminal progress bar.
9pub struct CliProgress<T: ICommandInfo> {
10    mediator: Arc<CommandMediator<T>>,
11    bar: Arc<ProgressBar>,
12    handle: Mutex<Option<JoinHandle<()>>>,
13    finished: Arc<AtomicBool>,
14}
15
16impl<T: ICommandInfo + 'static> CliProgress<T> {
17    /// Create a new [`CliProgress`] whose progress bar is attached to a shared
18    /// [`MultiProgress`].
19    ///
20    /// - Bars added via `multi.add(...)` redraw cooperatively with each other
21    /// - Pair with [`ProgressWriterFactory`] passed to `LoggerBuilder::with_writer`
22    ///   to prevent collisions with `tracing` log output
23    #[must_use]
24    pub fn new(mediator: Arc<CommandMediator<T>>, multi: MultiProgress) -> Self {
25        Self {
26            mediator,
27            bar: Arc::new(multi.add(ProgressBar::new(0))),
28            handle: Mutex::default(),
29            finished: Arc::new(AtomicBool::new(false)),
30        }
31    }
32
33    /// Start listening for events and updating the progress bar.
34    pub async fn start(&self) {
35        let mut handle_guard = self.handle.lock().await;
36        if handle_guard.is_some() {
37            return;
38        }
39        let mediator = self.mediator.clone();
40        let mut receiver = mediator.subscribe();
41        let bar = self.bar.clone();
42        let finished = self.finished.clone();
43        let mut total: u64 = 0;
44        let handle = spawn(async move {
45            while !finished.load(Ordering::Acquire) {
46                match receiver.recv().await {
47                    Ok(event) => Self::handle_event(&bar, &mut total, event),
48                    Err(RecvError::Lagged(count)) => {
49                        warn!("CLI Progress missed {count} events due to lagging");
50                    }
51                    Err(RecvError::Closed) => {
52                        error!("Event pipe was closed. CLI Progress can't proceed.");
53                        break;
54                    }
55                }
56            }
57        });
58        *handle_guard = Some(handle);
59    }
60
61    fn handle_event(bar: &ProgressBar, total: &mut u64, event: T::Event) {
62        match event.get_kind() {
63            EventKind::Queued => {
64                *total += 1;
65                bar.set_length(*total);
66            }
67            EventKind::Executing => {}
68            EventKind::Succeeded | EventKind::Failed => {
69                bar.inc(1);
70            }
71        }
72    }
73
74    /// Signal completion and abort the listener task.
75    pub async fn finish(&self) {
76        self.finished.store(true, Ordering::Release);
77        let mut handle_guard = self.handle.lock().await;
78        if let Some(handle) = handle_guard.take() {
79            handle.abort();
80        }
81        drop(handle_guard);
82        self.bar.finish();
83    }
84
85    /// Hide the progress bar output.
86    #[cfg(test)]
87    pub fn hide(&self) {
88        self.bar
89            .set_draw_target(indicatif::ProgressDrawTarget::hidden());
90    }
91
92    /// Progress bar position (completed items).
93    #[cfg(test)]
94    pub fn position(&self) -> u64 {
95        self.bar.position()
96    }
97
98    /// Progress bar total length (queued items).
99    #[cfg(test)]
100    pub fn length(&self) -> Option<u64> {
101        self.bar.length()
102    }
103}
104
105impl<T: ICommandInfo + 'static> FromServices for CliProgress<T> {
106    type Error = ResolveError;
107
108    fn from_services(services: &ServiceProvider) -> Result<Self, Report<Self::Error>> {
109        let mediator = services.get::<CommandMediator<T>>()?;
110        let factory = services.get::<ProgressWriterFactory>()?;
111        Ok(Self::new(mediator, factory.multi()))
112    }
113}
114
115#[cfg(all(test, feature = "server"))]
116mod tests {
117    #![expect(
118        clippy::as_conversions,
119        reason = "usize to u64 cast in test assertions"
120    )]
121    use super::*;
122
123    const COMMAND_COUNT: usize = CHANNEL_CAPACITY * 2;
124    const WORKER_COUNT: usize = 4;
125    const DELAY_MS: u64 = 1;
126
127    #[tokio::test]
128    async fn cli_progress_receives_all_events() {
129        // Arrange
130        let services = ServiceBuilder::new()
131            .with_test_services()
132            .build()
133            .expect_init();
134        let runner = services.expect_async::<CommandRunner<CommandInfo>>().await;
135        let progress = services.expect::<CliProgress<CommandInfo>>();
136        progress.hide();
137
138        // Act
139        progress.start().await;
140        runner.start(WORKER_COUNT).await;
141        for i in 1..=COMMAND_COUNT {
142            let request = DelayRequest::new(format!("P{i}"), DELAY_MS);
143            runner
144                .queue_request(request)
145                .await
146                .expect("should be able to queue request");
147        }
148        runner.drain().await;
149        progress.finish().await;
150
151        // Assert
152        assert_eq!(
153            progress.length(),
154            Some(COMMAND_COUNT as u64),
155            "progress bar total should match queued commands"
156        );
157        assert_eq!(
158            progress.position(),
159            COMMAND_COUNT as u64,
160            "progress bar position should match completed commands"
161        );
162    }
163
164    /// Direct construction attaches the bar to a shared [`MultiProgress`] and still receives events.
165    #[tokio::test]
166    async fn cli_progress_new_receives_all_events() {
167        // Arrange
168        use indicatif::ProgressDrawTarget;
169        let multi = MultiProgress::with_draw_target(ProgressDrawTarget::hidden());
170        let services = ServiceBuilder::new()
171            .with_test_services()
172            .build()
173            .expect_init();
174        let mediator = services.expect::<CommandMediator<CommandInfo>>();
175        let runner = services.expect_async::<CommandRunner<CommandInfo>>().await;
176        let progress = CliProgress::new(mediator, multi);
177
178        // Act
179        progress.start().await;
180        runner.start(WORKER_COUNT).await;
181        for i in 1..=COMMAND_COUNT {
182            let request = DelayRequest::new(format!("M{i}"), DELAY_MS);
183            runner
184                .queue_request(request)
185                .await
186                .expect("should be able to queue request");
187        }
188        runner.drain().await;
189        progress.finish().await;
190
191        // Assert
192        assert_eq!(
193            progress.length(),
194            Some(COMMAND_COUNT as u64),
195            "progress bar total should match queued commands"
196        );
197        assert_eq!(
198            progress.position(),
199            COMMAND_COUNT as u64,
200            "progress bar position should match completed commands"
201        );
202    }
203}