@@ -8,13 +8,15 @@ use std::env;
88use std:: io:: Read ;
99use std:: path:: { Path , PathBuf } ;
1010use std:: process:: { Command , Stdio } ;
11+ use std:: sync:: mpsc:: { self , Receiver , RecvTimeoutError } ;
1112use std:: thread;
1213use std:: time:: Duration ;
1314use wait_timeout:: ChildExt ;
1415
1516const DEFAULT_STDOUT_CAP : usize = 2 * 1024 * 1024 ;
1617const DEFAULT_STDERR_CAP : usize = 64 * 1024 ;
1718const DEFAULT_PROVIDER_TIMEOUT : Duration = Duration :: from_secs ( 20 ) ;
19+ const STREAM_COLLECT_TIMEOUT : Duration = Duration :: from_secs ( 1 ) ;
1820
1921#[ derive( Debug , Clone , Serialize , Deserialize ) ]
2022pub struct Provider {
@@ -162,26 +164,25 @@ pub fn run_provider(name: &str, root: &Path) -> Result<Vec<ProviderFinding>> {
162164 . stderr ( Stdio :: piped ( ) )
163165 . spawn ( )
164166 . with_context ( || format ! ( "spawn provider {name}" ) ) ?;
165- let stdout_handle = child
167+ let stdout_reader = child
166168 . stdout
167169 . take ( )
168- . map ( |stream| thread :: spawn ( move || read_stream_capped ( stream, stdout_cap) ) ) ;
169- let stderr_handle = child
170+ . map ( |stream| spawn_stream_reader ( stream, stdout_cap) ) ;
171+ let stderr_reader = child
170172 . stderr
171173 . take ( )
172- . map ( |stream| thread :: spawn ( move || read_stream_capped ( stream, stderr_cap) ) ) ;
174+ . map ( |stream| spawn_stream_reader ( stream, stderr_cap) ) ;
173175 let status = match child. wait_timeout ( timeout) ? {
174176 Some ( status) => status,
175177 None => {
176178 let _ = child. kill ( ) ;
177179 let _ = child. wait ( ) ;
178- let _ = join_stream ( stdout_handle) ;
179- let _ = join_stream ( stderr_handle) ;
180180 return Err ( anyhow ! ( "provider timed out after {:?}" , timeout) ) ;
181181 }
182182 } ;
183- let ( stdout, stdout_truncated) = join_stream ( stdout_handle) ;
184- let ( stderr, _) = join_stream ( stderr_handle) ;
183+ let ( stdout, stdout_truncated) =
184+ collect_stream ( stdout_reader, "stdout" , STREAM_COLLECT_TIMEOUT ) ?;
185+ let ( stderr, _) = collect_stream ( stderr_reader, "stderr" , STREAM_COLLECT_TIMEOUT ) ?;
185186 if stdout_truncated {
186187 return Err ( anyhow ! ( "provider stdout exceeded {stdout_cap} byte cap" ) ) ;
187188 }
@@ -191,6 +192,18 @@ pub fn run_provider(name: &str, root: &Path) -> Result<Vec<ProviderFinding>> {
191192 parse_provider_output ( name, root, & stdout)
192193}
193194
195+ struct StreamReader {
196+ receiver : Receiver < ( String , bool ) > ,
197+ }
198+
199+ fn spawn_stream_reader ( stream : impl Read + Send + ' static , cap : usize ) -> StreamReader {
200+ let ( sender, receiver) = mpsc:: channel ( ) ;
201+ thread:: spawn ( move || {
202+ let _ = sender. send ( read_stream_capped ( stream, cap) ) ;
203+ } ) ;
204+ StreamReader { receiver }
205+ }
206+
194207fn read_stream_capped ( mut stream : impl Read , cap : usize ) -> ( String , bool ) {
195208 let mut out = Vec :: with_capacity ( cap. min ( 64 * 1024 ) ) ;
196209 let mut truncated = false ;
@@ -213,10 +226,23 @@ fn read_stream_capped(mut stream: impl Read, cap: usize) -> (String, bool) {
213226 ( redact_text ( & String :: from_utf8_lossy ( & out) ) , truncated)
214227}
215228
216- fn join_stream ( handle : Option < thread:: JoinHandle < ( String , bool ) > > ) -> ( String , bool ) {
217- handle
218- . and_then ( |handle| handle. join ( ) . ok ( ) )
219- . unwrap_or_default ( )
229+ fn collect_stream (
230+ reader : Option < StreamReader > ,
231+ label : & str ,
232+ timeout : Duration ,
233+ ) -> Result < ( String , bool ) > {
234+ let Some ( reader) = reader else {
235+ return Ok ( ( String :: new ( ) , false ) ) ;
236+ } ;
237+ match reader. receiver . recv_timeout ( timeout) {
238+ Ok ( result) => Ok ( result) ,
239+ Err ( RecvTimeoutError :: Timeout ) => {
240+ Err ( anyhow ! ( "provider {label} did not close after process exit" ) )
241+ }
242+ Err ( RecvTimeoutError :: Disconnected ) => {
243+ Err ( anyhow ! ( "provider {label} reader ended without output" ) )
244+ }
245+ }
220246}
221247
222248pub fn parse_provider_output (
0 commit comments