studiole_command/services/
cli_progress.rs1use 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
8pub 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 #[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 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 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 #[cfg(test)]
87 pub fn hide(&self) {
88 self.bar
89 .set_draw_target(indicatif::ProgressDrawTarget::hidden());
90 }
91
92 #[cfg(test)]
94 pub fn position(&self) -> u64 {
95 self.bar.position()
96 }
97
98 #[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 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 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_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 #[tokio::test]
166 async fn cli_progress_new_receives_all_events() {
167 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 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_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}