Skip to main content

trace_task/
trace_task.rs

1// Copyright 2025 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use crate::triggers::{Trigger, TriggerAction, TriggersWatcher};
6use crate::{TracingError, trace_shutdown};
7use async_lock::Mutex;
8use flex_client::{AsyncSocket, ProxyHasDomain, socket_to_async};
9use flex_fuchsia_tracing_controller::{self as trace, StopResult, TraceConfig};
10use fuchsia_async::Task;
11use futures::io::AsyncWrite;
12use futures::prelude::*;
13use futures::task::{Context as FutContext, Poll};
14use std::pin::Pin;
15use std::sync::Arc;
16use std::sync::atomic::{AtomicBool, AtomicU64};
17use std::time::{Duration, Instant};
18use zstd::stream::raw::Operation;
19
20static SERIAL: AtomicU64 = AtomicU64::new(100);
21
22pub struct TraceTask {
23    /// Domain for FDomain client handles. It is here to prevent the domain
24    /// being dropped while it could still be used.
25    _domain: flex_client::ClientArg,
26    /// Unique identifier for this task. The value of this id monotonicallly increases.
27    task_id: u64,
28    /// Tag used to identify this task in the log.
29    debug_tag: String,
30    /// Trace configuration.
31    config: trace::TraceConfig,
32    /// Requested categories. These are unexpanded from the user.
33    requested_categories: Vec<String>,
34    /// Duration to capture trace. None indicates capture until canceled.
35    duration: Option<Duration>,
36    /// Triggers for terminating the trace.
37    triggers: Vec<Trigger>,
38    /// True when the task is cleaning up.
39    terminating: Arc<AtomicBool>,
40    /// Start time of the task.
41    start_time: Instant,
42    /// Channel used to shutdown this task.
43    shutdown_sender: async_channel::Sender<()>,
44    /// The task.
45    task: Task<Option<trace::StopResult>>,
46    /// The socket to read the trace data from when tracing is completed.
47    read_socket: Option<AsyncSocket>,
48    /// The compression algorithm to use.
49    compression: trace::CompressionType,
50    /// True when the task was cancelled (aborted).
51    cancelled: Arc<AtomicBool>,
52}
53
54/// Implement Debug explicitly since `flex_client::ClientArg` does not implement Debug,
55/// so we skip it in this implementation.
56impl std::fmt::Debug for TraceTask {
57    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
58        f.debug_struct("TraceTask")
59            .field("task_id", &self.task_id)
60            .field("debug_tag", &self.debug_tag)
61            .field("config", &self.config)
62            .field("requested_categories", &self.requested_categories)
63            .field("duration", &self.duration)
64            .field("triggers", &self.triggers)
65            .field("terminating", &self.terminating)
66            .field("start_time", &self.start_time)
67            .field("shutdown_sender", &self.shutdown_sender)
68            .field("read_socket", &self.read_socket)
69            .field("compression", &self.compression)
70            .field("cancelled", &self.cancelled)
71            .finish()
72    }
73}
74
75// This is just implemented for convenience so the wrapper is await-able.
76impl Future for TraceTask {
77    type Output = Option<trace::StopResult>;
78
79    fn poll(mut self: Pin<&mut Self>, cx: &mut FutContext<'_>) -> Poll<Self::Output> {
80        Pin::new(&mut self.task).poll(cx)
81    }
82}
83
84impl TraceTask {
85    pub async fn new(
86        debug_tag: String,
87        config: trace::TraceConfig,
88        duration: Option<Duration>,
89        triggers: Vec<Trigger>,
90        requested_categories: Option<Vec<String>>,
91        compression: trace::CompressionType,
92        provisioner: trace::ProvisionerProxy,
93    ) -> Result<Self, TracingError> {
94        // Start the tracing session immediately. Maybe we should consider separating the creating
95        // of the session and the actual starting of it. This seems like a side-effect.
96        log::info!("TraceTask::new called with compression: {:?}", compression);
97        let task_id = SERIAL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
98        let domain = provisioner.domain();
99        let (client, server) = domain.create_stream_socket();
100        let (client_end, server_end) = domain.create_proxy::<trace::SessionMarker>();
101        provisioner.initialize_tracing(server_end, &config, server)?;
102
103        client_end
104            .start_tracing(&trace::StartOptions::default())
105            .await?
106            .map_err(Into::<TracingError>::into)?;
107
108        let logging_prefix_og = format!("Task {task_id} ({debug_tag})");
109        let terminate_result = Arc::new(Mutex::new(None));
110        let (shutdown_sender, shutdown_receiver) = async_channel::bounded::<()>(1);
111
112        let controller = client_end.clone();
113        let shutdown_controller = client_end.clone();
114        let triggers_watcher =
115            TriggersWatcher::new(controller, triggers.clone(), shutdown_receiver);
116        let terminating = Arc::new(AtomicBool::new(false));
117        let cancelled = Arc::new(AtomicBool::new(false));
118        let terminating_clone = terminating.clone();
119        let cancelled_clone = cancelled.clone();
120        let terminate_result_clone = terminate_result.clone();
121
122        let shutdown_fut = {
123            let logging_prefix = logging_prefix_og.clone();
124            async move {
125                if terminating_clone
126                    .compare_exchange(
127                        false,
128                        true,
129                        std::sync::atomic::Ordering::SeqCst,
130                        std::sync::atomic::Ordering::Relaxed,
131                    )
132                    .is_ok()
133                {
134                    log::info!("{logging_prefix} Running shutdown future.");
135                    let is_cancelled = cancelled_clone.load(std::sync::atomic::Ordering::Relaxed);
136                    let result = trace_shutdown(&shutdown_controller, is_cancelled).await;
137
138                    let mut done = terminate_result_clone.lock().await;
139                    if done.is_none() {
140                        match result {
141                            Ok(stop) => {
142                                log::info!("{logging_prefix} call to trace_shutdown successful.");
143                                *done = Some(stop)
144                            }
145                            Err(e) => {
146                                log::error!(
147                                    "{logging_prefix} call to trace_shutdown failed: {e:?}"
148                                );
149                            }
150                        }
151                    }
152                } else {
153                    log::debug!("Shutdown already triggered");
154                }
155                "shutdown future completed"
156            }
157        };
158
159        let task = Self::make_task(
160            task_id,
161            debug_tag,
162            duration,
163            shutdown_fut,
164            triggers_watcher,
165            terminate_result,
166        );
167        Ok(Self {
168            _domain: domain,
169            task_id,
170            debug_tag: logging_prefix_og,
171            config,
172            duration,
173            triggers: triggers.clone(),
174            terminating,
175            requested_categories: requested_categories.unwrap_or_default(),
176            start_time: Instant::now(),
177            shutdown_sender,
178            read_socket: Some(socket_to_async(client)),
179            compression,
180            task,
181            cancelled,
182        })
183    }
184
185    /// Shutdown the tracing task.
186    async fn shutdown(self) -> Result<trace::StopResult, TracingError> {
187        if !self.terminating.load(std::sync::atomic::Ordering::SeqCst) {
188            log::info!("{} Sending shutdown message.", self.debug_tag);
189            if self.shutdown_sender.send(()).await.is_err() {
190                log::warn!(
191                    "{} Shutdown channel was closed. Task may have already completed.",
192                    self.debug_tag
193                );
194            }
195        } else {
196            log::debug!("{} Shutdown already in progress.", self.debug_tag);
197        }
198
199        self.await
200            .map(|r| Ok(r))
201            .unwrap_or_else(|| Err(TracingError::RecordingStop("Error awaiting".into())))
202    }
203
204    /// Abort the tracing task without writing results.
205    pub async fn abort(mut self) -> Result<trace::StopResult, TracingError> {
206        if !self.terminating.load(std::sync::atomic::Ordering::SeqCst) {
207            log::info!("{} Sending cancel message for task", self.debug_tag);
208            self.cancelled.store(true, std::sync::atomic::Ordering::SeqCst);
209            if self.shutdown_sender.send(()).await.is_err() {
210                log::warn!(
211                    "{} Shutdown channel was closed. Task may have already completed.",
212                    self.debug_tag
213                );
214            }
215        } else {
216            log::debug!("{} Shutdown already in progress.", self.debug_tag);
217        }
218
219        // Close the socket without reading.
220        if let Some(mut socket) = self.read_socket.take() {
221            let res = socket.close().await;
222            if res.is_err() {
223                log::warn!("{} Failed to close socket: {:?}", self.debug_tag, res);
224            }
225            drop(socket);
226        }
227        self.shutdown().await
228    }
229
230    fn make_task(
231        task_id: u64,
232        debug_tag: String,
233        duration: Option<Duration>,
234        shutdown_fut: impl Future<Output = &'static str> + 'static + std::marker::Send,
235        trigger_watcher: TriggersWatcher<'static>,
236        terminate_result: Arc<Mutex<Option<StopResult>>>,
237    ) -> Task<Option<trace::StopResult>> {
238        Task::local(async move {
239            let mut timeout_fut = Box::pin(async move {
240                if let Some(duration) = duration {
241                    fuchsia_async::Timer::new(duration).await;
242                } else {
243                    std::future::pending::<()>().await;
244                }
245            })
246            .fuse();
247            let mut trigger_fut = trigger_watcher.fuse();
248
249            futures::select! {
250                // Timeout, clean up and wait for copying to finish.
251                _ = timeout_fut => {
252                    log::info!("Trace {task_id} (debug_tag): timeout of {} successfully completed. Stopping and cleaning up.",
253                     duration.map(|d| format!("{} secs", d.as_secs())).unwrap_or_else(|| "infinite?".into()));
254
255                    shutdown_fut.await;
256                     log::debug!("done with timeout!");
257
258                }
259
260                // Trigger hit, shutdown and copy the trace.
261                action = trigger_fut => {
262                    if let Some(action) = action {
263                        match action {
264                            TriggerAction::Terminate => {
265                                log::info!("Task {task_id} ({debug_tag}): received terminate trigger");
266                            }
267                        }
268                    } else {
269                        // This usually means the proxy was closed.
270                        log::debug!("Task {task_id} ({debug_tag}): Trigger future completed without an action!");
271                    }
272                    shutdown_fut.await;
273                     log::debug!("done with trigger future!");
274                }
275            };
276            log::debug!("end of task waiting for terminate_result lock");
277            let res = terminate_result.lock().await.clone();
278            log::debug!("got res in task is some: {}", res.is_some());
279            res
280        })
281    }
282
283    pub fn triggers(&self) -> Vec<Trigger> {
284        self.triggers.clone()
285    }
286    pub fn config(&self) -> TraceConfig {
287        self.config.clone()
288    }
289
290    pub fn start_time(&self) -> Instant {
291        self.start_time
292    }
293
294    pub fn duration(&self) -> Option<Duration> {
295        self.duration.clone()
296    }
297
298    pub fn requested_categories(&self) -> Vec<String> {
299        self.requested_categories.clone()
300    }
301
302    pub fn task_id(&self) -> u64 {
303        self.task_id
304    }
305
306    /// Signals the trace session to stop, copies all trace data to the
307    /// provided writer, and awaits task completion.
308    pub async fn stop_and_receive_data<W>(
309        mut self,
310        mut writer: W,
311    ) -> Result<trace::StopResult, TracingError>
312    where
313        W: AsyncWrite + Unpin + Send + 'static,
314    {
315        if !self.terminating.load(std::sync::atomic::Ordering::SeqCst) {
316            log::info!("{} Sending shutdown message for task", self.debug_tag);
317            if self.shutdown_sender.send(()).await.is_err() {
318                log::warn!(
319                    "{} Shutdown channel was closed. Task may have already completed.",
320                    self.debug_tag
321                );
322            }
323        } else {
324            log::debug!("{} Shutdown already in progress.", self.debug_tag);
325        }
326
327        let mut read_socket = self.read_socket.take().unwrap();
328        let res = match self.compression {
329            trace::CompressionType::Zstd => compress_zstd(&mut read_socket, &mut writer).await,
330            _ => futures::io::copy(&mut read_socket, &mut writer)
331                .await
332                .map(|_| ())
333                .map_err(|e| TracingError::GeneralError(format!("{e:?}"))),
334        };
335
336        if res.is_ok() { self.shutdown().await } else { Err(res.err().unwrap()) }
337    }
338
339    /// Waits for the tracing task to complete and copies the trace data to the writer.
340    /// If the tracing should be stopped vs. waiting, call |stop_and_receive_data|.
341    pub async fn await_completion_and_receive_data<W>(
342        mut self,
343        mut writer: W,
344    ) -> Result<StopResult, TracingError>
345    where
346        W: AsyncWrite + Unpin + Send + 'static,
347    {
348        let mut read_socket = self.read_socket.take().unwrap();
349        let res = match self.compression {
350            trace::CompressionType::Zstd => compress_zstd(&mut read_socket, &mut writer).await,
351            _ => futures::io::copy(&mut read_socket, &mut writer)
352                .await
353                .map(|_| ())
354                .map_err(|e| TracingError::RecordingStop(e.to_string())),
355        };
356
357        match res {
358            Ok(_) => match self.await {
359                Some(r) => Ok(r),
360                None => Err(TracingError::RecordingStop("could not await task".into())),
361            },
362            Err(e) => Err(e),
363        }
364    }
365}
366
367async fn compress_zstd<R, W>(mut reader: R, mut writer: W) -> Result<(), TracingError>
368where
369    R: AsyncRead + Unpin,
370    W: AsyncWrite + Unpin,
371{
372    let mut encoder = zstd::stream::raw::Encoder::new(0)
373        .map_err(|e| TracingError::GeneralError(format!("zstd init: {e:?}")))?;
374    // 128KB is the recommended size for Zstd (ZSTD_CStreamInSize/ZSTD_CStreamOutSize)
375    let mut input_buf = vec![0u8; 128 * 1024];
376    let mut output_buf = vec![0u8; 128 * 1024];
377
378    // Read the stream until EOF, compressing and writing out fully compressed buffers as we go.
379    while let n = reader
380        .read(&mut input_buf)
381        .await
382        .map_err(|e| TracingError::GeneralError(format!("read: {e:?}")))?
383        && n > 0
384    {
385        let mut read_offset = 0;
386        while read_offset < n {
387            let status = encoder
388                .run_on_buffers(&input_buf[read_offset..n], &mut output_buf)
389                .map_err(|e| TracingError::GeneralError(format!("zstd run: {e:?}")))?;
390            read_offset += status.bytes_read;
391            if status.bytes_written > 0 {
392                writer
393                    .write_all(&output_buf[..status.bytes_written])
394                    .await
395                    .map_err(|e| TracingError::GeneralError(format!("write: {e:?}")))?;
396            }
397        }
398    }
399
400    // Flush remaining of the last compressed buffer.
401    loop {
402        let mut out_wrapper = zstd::stream::raw::OutBuffer::around(&mut output_buf);
403        let remaining = encoder
404            .finish(&mut out_wrapper, true)
405            .map_err(|e| TracingError::GeneralError(format!("zstd finish: {e:?}")))?;
406        let bytes = out_wrapper.as_slice();
407        if !bytes.is_empty() {
408            writer
409                .write_all(bytes)
410                .await
411                .map_err(|e| TracingError::GeneralError(format!("write: {e:?}")))?;
412        }
413        if remaining == 0 {
414            break;
415        }
416    }
417    Ok(())
418}
419
420#[cfg(test)]
421mod tests {
422    use super::*;
423    use flex_client::fidl::Responder;
424    use flex_fuchsia_tracing_controller::StartError;
425
426    const FAKE_CONTROLLER_TRACE_OUTPUT: &'static str = "HOWDY HOWDY HOWDY";
427    fn setup_fake_provisioner_proxy_with_payload(
428        start_error: Option<StartError>,
429        trigger_name: Option<&'static str>,
430        expected_write_results: bool,
431        payload_bytes: impl AsRef<[u8]>,
432    ) -> trace::ProvisionerProxy {
433        let payload_bytes = payload_bytes.as_ref().to_vec();
434        let client = fdomain_local::local_client_empty();
435        let (proxy, mut stream) = client.create_proxy_and_stream::<trace::ProvisionerMarker>();
436        fuchsia_async::Task::local(async move {
437            let _client = client;
438            while let Ok(Some(req)) = stream.try_next().await {
439                match req {
440                    trace::ProvisionerRequest::InitializeTracing { controller, output, .. } => {
441                        let mut stream = controller.into_stream();
442                        let mut async_output = socket_to_async(output);
443                        while let Ok(Some(req)) = stream.try_next().await {
444                            match req {
445                                trace::SessionRequest::StartTracing { responder, .. } => {
446                                    let response = match start_error {
447                                        Some(e) => Err(e),
448                                        None => Ok(()),
449                                    };
450                                    responder.send(response).expect("Failed to start")
451                                }
452                                trace::SessionRequest::StopTracing { responder, payload } => {
453                                    if start_error.is_some() {
454                                        responder
455                                            .send(Err(trace::StopError::NotStarted))
456                                            .expect("Failed to stop");
457                                    } else {
458                                        assert_eq!(
459                                            payload.write_results.unwrap(),
460                                            expected_write_results
461                                        );
462                                        if expected_write_results && !payload_bytes.is_empty() {
463                                            let _ = async_output.write_all(&payload_bytes).await;
464                                        }
465                                        let _ = async_output.close().await;
466                                        let stop_result = trace::StopResult {
467                                            provider_stats: Some(vec![]),
468                                            ..Default::default()
469                                        };
470                                        responder.send(Ok(&stop_result)).expect("Failed to stop");
471                                    }
472                                }
473                                trace::SessionRequest::WatchAlert { responder } => {
474                                    if let Some(trigger) = trigger_name {
475                                        responder.send(trigger).expect("Unable to send alert");
476                                    } else {
477                                        responder.drop_without_shutdown();
478                                    }
479                                }
480                                r => panic!("unexpected request: {:#?}", r),
481                            }
482                        }
483                    }
484                    r => panic!("unexpected request: {:#?}", r),
485                }
486            }
487        })
488        .detach();
489        proxy
490    }
491
492    #[fuchsia::test]
493    async fn test_trace_task_start_stop_write_check_with_vec() {
494        let provisioner = setup_fake_provisioner_proxy_with_payload(
495            None,
496            None,
497            true,
498            FAKE_CONTROLLER_TRACE_OUTPUT,
499        );
500
501        let trace_task = TraceTask::new(
502            "test_trace_start_stop_write_check".into(),
503            trace::TraceConfig::default(),
504            None,
505            vec![],
506            None,
507            trace::CompressionType::None,
508            provisioner,
509        )
510        .await
511        .expect("tracing task started");
512
513        let shutdown_result = trace_task.shutdown().await.expect("tracing shutdown");
514        assert_eq!(
515            shutdown_result,
516            trace::StopResult { provider_stats: Some(vec![]), ..Default::default() }.into()
517        );
518    }
519
520    #[cfg(not(target_os = "fuchsia"))]
521    #[fuchsia::test]
522    async fn test_trace_task_start_stop_write_check_with_file() {
523        let temp_dir = tempfile::TempDir::new().unwrap();
524        let output = temp_dir.path().join("trace-test.fxt");
525
526        let provisioner = setup_fake_provisioner_proxy_with_payload(
527            None,
528            None,
529            true,
530            FAKE_CONTROLLER_TRACE_OUTPUT,
531        );
532        let writer = async_fs::File::create(&output).await.unwrap();
533
534        let trace_task = TraceTask::new(
535            "test_trace_start_stop_write_check".into(),
536            trace::TraceConfig::default(),
537            None,
538            vec![],
539            None,
540            trace::CompressionType::None,
541            provisioner,
542        )
543        .await
544        .expect("tracing task started");
545
546        let shutdown_result =
547            trace_task.stop_and_receive_data(writer).await.expect("tracing shutdown");
548
549        let res = async_fs::read_to_string(&output).await.unwrap();
550        assert_eq!(res, FAKE_CONTROLLER_TRACE_OUTPUT.to_string());
551        let expected = trace::StopResult { provider_stats: Some(vec![]), ..Default::default() };
552        assert_eq!(shutdown_result, expected);
553    }
554
555    #[fuchsia::test]
556    async fn test_trace_error_handling_already_started() {
557        let provisioner = setup_fake_provisioner_proxy_with_payload(
558            Some(StartError::AlreadyStarted),
559            None,
560            true,
561            FAKE_CONTROLLER_TRACE_OUTPUT,
562        );
563
564        let trace_task_result = TraceTask::new(
565            "test_trace_error_handling_already_started".into(),
566            trace::TraceConfig::default(),
567            None,
568            vec![],
569            None,
570            trace::CompressionType::None,
571            provisioner,
572        )
573        .await
574        .err();
575
576        assert_eq!(trace_task_result, Some(TracingError::RecordingAlreadyStarted));
577    }
578
579    #[cfg(not(target_os = "fuchsia"))]
580    #[fuchsia::test]
581    async fn test_trace_task_start_with_duration() {
582        let temp_dir = tempfile::TempDir::new().unwrap();
583        let output = temp_dir.path().join("trace-test.fxt");
584
585        let provisioner = setup_fake_provisioner_proxy_with_payload(
586            None,
587            None,
588            true,
589            FAKE_CONTROLLER_TRACE_OUTPUT,
590        );
591        let writer = async_fs::File::create(&output).await.unwrap();
592
593        let trace_task = TraceTask::new(
594            "test_trace_task_start_with_duration".into(),
595            trace::TraceConfig::default(),
596            Some(Duration::from_millis(100)),
597            vec![],
598            None,
599            trace::CompressionType::None,
600            provisioner,
601        )
602        .await
603        .expect("tracing task started");
604
605        let res = trace_task.await_completion_and_receive_data(writer).await;
606        if let Some(ref stop_result) = res.as_ref().ok() {
607            assert!(stop_result.provider_stats.is_some());
608        } else {
609            panic!("Expected stop result from trace_task.await: {res:?}");
610        }
611
612        let mut f = async_fs::File::open(std::path::PathBuf::from(output)).await.unwrap();
613        let mut res = String::new();
614        f.read_to_string(&mut res).await.unwrap();
615        assert_eq!(res, FAKE_CONTROLLER_TRACE_OUTPUT.to_string());
616    }
617
618    #[cfg(not(target_os = "fuchsia"))]
619    #[fuchsia::test]
620    async fn test_triggers_valid() {
621        let temp_dir = tempfile::TempDir::new().unwrap();
622        let output = temp_dir.path().join("trace-test.fxt");
623        let alert_name = "some_alert";
624        let provisioner = setup_fake_provisioner_proxy_with_payload(
625            None,
626            Some(alert_name.into()),
627            true,
628            FAKE_CONTROLLER_TRACE_OUTPUT,
629        );
630        let writer = async_fs::File::create(output.clone()).await.unwrap();
631
632        let trace_task = TraceTask::new(
633            "test_triggers_valid".into(),
634            trace::TraceConfig::default(),
635            None,
636            vec![Trigger {
637                alert: Some(alert_name.into()),
638                action: Some(TriggerAction::Terminate),
639            }],
640            None,
641            trace::CompressionType::None,
642            provisioner,
643        )
644        .await
645        .expect("tracing task started");
646
647        trace_task.await_completion_and_receive_data(writer).await.unwrap();
648        let res = async_fs::read_to_string(&output).await.unwrap();
649        assert_eq!(res, FAKE_CONTROLLER_TRACE_OUTPUT.to_string());
650    }
651
652    #[fuchsia::test]
653    async fn test_trace_task_abort() {
654        let provisioner = setup_fake_provisioner_proxy_with_payload(
655            None,
656            None,
657            false,
658            FAKE_CONTROLLER_TRACE_OUTPUT,
659        );
660
661        let trace_task = TraceTask::new(
662            "test_trace_task_abort".into(),
663            trace::TraceConfig::default(),
664            None,
665            vec![],
666            None,
667            trace::CompressionType::None,
668            provisioner,
669        )
670        .await
671        .expect("tracing task started");
672
673        let shutdown_result = trace_task.abort().await.expect("tracing abort");
674        assert_eq!(
675            shutdown_result,
676            trace::StopResult { provider_stats: Some(vec![]), ..Default::default() }
677        );
678    }
679
680    struct FailingAsyncWriter {
681        fail_after_bytes: usize,
682        written: usize,
683    }
684
685    impl AsyncWrite for FailingAsyncWriter {
686        fn poll_write(
687            mut self: Pin<&mut Self>,
688            _cx: &mut FutContext<'_>,
689            buf: &[u8],
690        ) -> Poll<std::io::Result<usize>> {
691            if self.written >= self.fail_after_bytes {
692                return Poll::Ready(Err(std::io::Error::new(
693                    std::io::ErrorKind::BrokenPipe,
694                    "simulated broken pipe in stress test",
695                )));
696            }
697            let allowed = std::cmp::min(buf.len(), self.fail_after_bytes - self.written);
698            self.written += allowed;
699            Poll::Ready(Ok(allowed))
700        }
701
702        fn poll_flush(self: Pin<&mut Self>, _cx: &mut FutContext<'_>) -> Poll<std::io::Result<()>> {
703            Poll::Ready(Ok(()))
704        }
705
706        fn poll_close(self: Pin<&mut Self>, _cx: &mut FutContext<'_>) -> Poll<std::io::Result<()>> {
707            Poll::Ready(Ok(()))
708        }
709    }
710
711    #[fuchsia::test]
712    async fn test_stress_faulty_writer_error_handling() {
713        let payload = b"PAYLOAD_FOR_FAULTY_WRITER_TEST";
714
715        // Failure mode 1: Immediate failure on uncompressed stop_and_receive_data
716        {
717            let provisioner =
718                setup_fake_provisioner_proxy_with_payload(None, None, true, payload.to_vec());
719            let trace_task = TraceTask::new(
720                "faulty_writer_1".into(),
721                trace::TraceConfig::default(),
722                None,
723                vec![],
724                None,
725                trace::CompressionType::None,
726                provisioner,
727            )
728            .await
729            .unwrap();
730
731            let res = trace_task
732                .stop_and_receive_data(FailingAsyncWriter { fail_after_bytes: 0, written: 0 })
733                .await;
734            assert!(res.is_err(), "Expected error when writer fails immediately");
735        }
736
737        // Failure mode 2: Immediate failure on uncompressed await_completion_and_receive_data
738        {
739            let provisioner =
740                setup_fake_provisioner_proxy_with_payload(None, None, true, payload.to_vec());
741            let trace_task = TraceTask::new(
742                "faulty_writer_2".into(),
743                trace::TraceConfig::default(),
744                Some(Duration::from_millis(10)),
745                vec![],
746                None,
747                trace::CompressionType::None,
748                provisioner,
749            )
750            .await
751            .unwrap();
752
753            let res = trace_task
754                .await_completion_and_receive_data(FailingAsyncWriter {
755                    fail_after_bytes: 0,
756                    written: 0,
757                })
758                .await;
759            assert!(res.is_err());
760        }
761
762        // Failure mode 3: Immediate failure on Zstd stop_and_receive_data
763        {
764            let provisioner =
765                setup_fake_provisioner_proxy_with_payload(None, None, true, payload.to_vec());
766            let trace_task = TraceTask::new(
767                "faulty_writer_3".into(),
768                trace::TraceConfig::default(),
769                None,
770                vec![],
771                None,
772                trace::CompressionType::Zstd,
773                provisioner,
774            )
775            .await
776            .unwrap();
777
778            let res = trace_task
779                .stop_and_receive_data(FailingAsyncWriter { fail_after_bytes: 0, written: 0 })
780                .await;
781            assert!(res.is_err());
782        }
783
784        // Failure mode 4: Failure mid-stream (after 10 bytes) on uncompressed
785        {
786            let provisioner =
787                setup_fake_provisioner_proxy_with_payload(None, None, true, payload.to_vec());
788            let trace_task = TraceTask::new(
789                "faulty_writer_4".into(),
790                trace::TraceConfig::default(),
791                None,
792                vec![],
793                None,
794                trace::CompressionType::None,
795                provisioner,
796            )
797            .await
798            .unwrap();
799
800            let res = trace_task
801                .stop_and_receive_data(FailingAsyncWriter { fail_after_bytes: 10, written: 0 })
802                .await;
803            assert!(res.is_err());
804        }
805
806        // Failure mode 5: Failure mid-stream (after 10 bytes) on Zstd
807        {
808            let provisioner =
809                setup_fake_provisioner_proxy_with_payload(None, None, true, payload.to_vec());
810            let trace_task = TraceTask::new(
811                "faulty_writer_5".into(),
812                trace::TraceConfig::default(),
813                None,
814                vec![],
815                None,
816                trace::CompressionType::Zstd,
817                provisioner,
818            )
819            .await
820            .unwrap();
821
822            let res = trace_task
823                .stop_and_receive_data(FailingAsyncWriter { fail_after_bytes: 10, written: 0 })
824                .await;
825            assert!(res.is_err());
826        }
827    }
828}