diff --git a/src/commands/serve.rs b/src/commands/serve.rs index d36cc9fe6d81..8f04c697f071 100644 --- a/src/commands/serve.rs +++ b/src/commands/serve.rs @@ -748,40 +748,49 @@ impl ServeCommand { _shutdown_guard: Box::new(shutdown.clone().increment(false)), }); - if debuggee_store.is_some() { - for server in servers { - Self::serve_on_listener( - server, - shutdown.clone(), - sem_connections, - handler.clone(), - debuggee_store, - ) - .await?; - - // There can only be one socket with a debugger attached. - break; - } + // If any listener fails then a shutdown is requested to bring down all + // the other listeners as well, but the drain below still happens to + // gracefully finish in-flight requests. The first error seen is then + // returned at the end. + let result = if debuggee_store.is_some() { + debug_assert_eq!(servers.len(), 1); + let server = servers.into_iter().next().unwrap(); + Self::serve_on_listener( + server, + shutdown.clone(), + sem_connections, + handler.clone(), + debuggee_store, + ) + .await } else { - let mut listener_tasks = vec![]; + let mut listener_tasks = tokio::task::JoinSet::new(); for server in servers { - listener_tasks.push(tokio::task::spawn(Self::serve_on_listener( + listener_tasks.spawn(Self::serve_on_listener( server, shutdown.clone(), sem_connections.clone(), handler.clone(), None, - ))); + )); } - for task in listener_tasks { - task.await??; + let mut result = Ok(()); + while let Some(task_result) = listener_tasks.join_next().await { + let task_result = task_result + .map_err(wasmtime::Error::from) + .and_then(|result| result); + if let Err(e) = task_result { + eprintln!("listener error: {e:?}"); + shutdown.request_shutdown(); + if result.is_ok() { + result = Err(e); + } + } } - } - - // Don't allow any further requests to get picked up. - handler.state().sem_requests.close(); + result + }; drop(handler); @@ -798,7 +807,7 @@ impl ServeCommand { } } - Ok(()) + result } async fn serve_on_listener( @@ -830,11 +839,12 @@ impl ServeCommand { // concurrent requests can't be served. Otherwise though spawn a // task to handle this client. match &mut debuggee_store { - Some(store) => handle_client(stream, &handler, Some(store)).await, + Some(store) => handle_client(stream, &handler, &shutdown, Some(store)).await, None => { let handler = handler.clone(); + let shutdown = shutdown.clone(); tokio::task::spawn(async move { - handle_client(stream, &handler, None).await; + handle_client(stream, &handler, &shutdown, None).await; drop(shutdown_guard); drop(connection_permit); }); @@ -1251,6 +1261,7 @@ type Request = hyper::Request; async fn handle_client( client: impl AsyncRead + AsyncWrite + Send + Sync + Unpin + 'static, handler: &ProxyHandler, + shutdown: &GracefulShutdown, debuggee_store: Option<&mut Store>, ) { // Hyper's `service_fn` takes an `Fn` closure, so to bridge the need to @@ -1259,26 +1270,24 @@ async fn handle_client( // `Send`. let lock = &debuggee_store.map(tokio::sync::Mutex::new); - if let Err(e) = http1::Builder::new() - .keep_alive(true) - .serve_connection( - TokioIo::new(client), - hyper::service::service_fn(move |req| async move { - let mut debuggee_store = match &lock { - Some(store) => Some(store.lock().await), - None => None, - }; - let debuggee_store = debuggee_store.as_mut().map(|s| &mut ***s); - match handle_request(handler, debuggee_store, req).await { - Ok(r) => Ok::<_, Infallible>(r), - Err(e) => { - eprintln!("error: {e:?}"); - let status = e - .downcast_ref::() - .map(|e| e.status()) - .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); - let error_html = format!( - "\ + let conn = http1::Builder::new().keep_alive(true).serve_connection( + TokioIo::new(client), + hyper::service::service_fn(move |req| async move { + let mut debuggee_store = match &lock { + Some(store) => Some(store.lock().await), + None => None, + }; + let debuggee_store = debuggee_store.as_mut().map(|s| &mut ***s); + match handle_request(handler, debuggee_store, req).await { + Ok(r) => Ok::<_, Infallible>(r), + Err(e) => { + eprintln!("error: {e:?}"); + let status = e + .downcast_ref::() + .map(|e| e.status()) + .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); + let error_html = format!( + "\ @@ -1292,22 +1301,32 @@ async fn handle_client( " - ); - Ok(Response::builder() - .status(status) - .header("Content-Type", "text/html; charset=UTF-8") - .body( - Full::new(bytes::Bytes::from(error_html)) - .map_err(|_| unreachable!()) - .boxed_unsync(), - ) - .unwrap()) - } + ); + Ok(Response::builder() + .status(status) + .header("Content-Type", "text/html; charset=UTF-8") + .body( + Full::new(bytes::Bytes::from(error_html)) + .map_err(|_| unreachable!()) + .boxed_unsync(), + ) + .unwrap()) } - }), - ) - .await - { + } + }), + ); + let mut conn = std::pin::pin!(conn); + + // Use hyper's built-in support for graceful shutdown whenever the server + // here gets a shutdown request. + let result = tokio::select! { + result = conn.as_mut() => result, + _ = shutdown.requested() => { + conn.as_mut().graceful_shutdown(); + conn.await + } + }; + if let Err(e) = result { eprintln!("error: {e:?}"); } }