Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
141 changes: 80 additions & 61 deletions src/commands/serve.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Expand All @@ -798,7 +807,7 @@ impl ServeCommand {
}
}

Ok(())
result
}

async fn serve_on_listener(
Expand Down Expand Up @@ -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);
});
Expand Down Expand Up @@ -1251,6 +1261,7 @@ type Request = hyper::Request<hyper::body::Incoming>;
async fn handle_client(
client: impl AsyncRead + AsyncWrite + Send + Sync + Unpin + 'static,
handler: &ProxyHandler<HostHandlerState>,
shutdown: &GracefulShutdown,
debuggee_store: Option<&mut Store<Host>>,
) {
// Hyper's `service_fn` takes an `Fn` closure, so to bridge the need to
Expand All @@ -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::<ErrorResponse>()
.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::<ErrorResponse>()
.map(|e| e.status())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let error_html = format!(
"\
<!doctype html>
<html>
<head>
Expand All @@ -1292,22 +1301,32 @@ async fn handle_client(
</center>
</body>
</html>"
);
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:?}");
}
}
Expand Down
Loading