forked from Karylab-cklius/vllm
[Rust Frontend] Refactor TLS serve path with unified MaybeTlsListener (#47101)
Signed-off-by: Bugen Zhao <i@bugenzhao.com>
This commit is contained in:
@@ -61,6 +61,7 @@ expect-test.workspace = true
|
||||
rmp-serde.workspace = true
|
||||
serial_test.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio = { workspace = true, features = ["test-util"] }
|
||||
tower.workspace = true
|
||||
vllm-engine-core-client = { workspace = true, features = ["test-util"] }
|
||||
zeromq.workspace = true
|
||||
|
||||
@@ -4,21 +4,16 @@ mod convert;
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use futures::{Stream, StreamExt as _, stream};
|
||||
use futures::{Stream, StreamExt as _};
|
||||
use thiserror_ext::AsReport as _;
|
||||
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_openssl::SslStream;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use tonic::transport::server::{Connected, TcpConnectInfo};
|
||||
use tonic::{Request, Response, Status};
|
||||
use tracing::info;
|
||||
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
|
||||
|
||||
use self::convert::ResponseOpts;
|
||||
use crate::listener::{Listener, ListenerIo};
|
||||
use crate::state::AppState;
|
||||
|
||||
/// Generated protobuf/gRPC types for the `vllm` package.
|
||||
@@ -31,78 +26,6 @@ pub use pb::generate_server::GenerateServer;
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
/// Newtype over `tokio-openssl`'s `SslStream` so we can implement tonic's
|
||||
/// [`Connected`] on it (the orphan rule blocks doing so on the foreign type).
|
||||
pub(crate) struct GrpcTlsStream {
|
||||
inner: SslStream<ListenerIo>,
|
||||
}
|
||||
|
||||
impl GrpcTlsStream {
|
||||
pub(crate) fn new(inner: SslStream<ListenerIo>) -> Self {
|
||||
Self { inner }
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncRead for GrpcTlsStream {
|
||||
fn poll_read(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &mut ReadBuf<'_>,
|
||||
) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_read(cx, buf)
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for GrpcTlsStream {
|
||||
fn poll_write(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<std::io::Result<usize>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
|
||||
}
|
||||
}
|
||||
|
||||
impl Connected for GrpcTlsStream {
|
||||
type ConnectInfo = TcpConnectInfo;
|
||||
|
||||
fn connect_info(&self) -> TcpConnectInfo {
|
||||
self.inner.get_ref().connect_info()
|
||||
}
|
||||
}
|
||||
|
||||
/// Adapt the shared server listener into tonic's incoming stream shape.
|
||||
pub(crate) fn incoming(listener: Listener) -> impl Stream<Item = std::io::Result<ListenerIo>> {
|
||||
stream::unfold(listener, |mut listener| async move {
|
||||
let (io, _) = axum::serve::Listener::accept(&mut listener).await;
|
||||
Some((Ok(io), listener))
|
||||
})
|
||||
}
|
||||
|
||||
/// Wrap the gRPC listener so each accepted connection completes a TLS handshake
|
||||
/// before tonic serves it.
|
||||
pub(crate) fn tls_incoming(
|
||||
listener: Listener,
|
||||
context: openssl::ssl::SslContext,
|
||||
handshake_timeout: std::time::Duration,
|
||||
) -> impl Stream<Item = std::io::Result<GrpcTlsStream>> {
|
||||
tls_listener::builder(context)
|
||||
.handshake_timeout(handshake_timeout)
|
||||
.listen(listener)
|
||||
.map(|res| {
|
||||
res.map(|(inner, _addr)| GrpcTlsStream::new(inner))
|
||||
.map_err(std::io::Error::other)
|
||||
})
|
||||
}
|
||||
|
||||
/// gRPC Generate service implementation backed by the shared application state.
|
||||
pub struct GenerateServiceImpl {
|
||||
state: Arc<AppState>,
|
||||
|
||||
@@ -30,8 +30,8 @@ use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
use super::pb::generate_client::GenerateClient;
|
||||
use super::{GenerateServer, GenerateServiceImpl, incoming, pb, tls_incoming};
|
||||
use crate::listener::Listener;
|
||||
use super::{GenerateServer, GenerateServiceImpl, pb};
|
||||
use crate::listener::{Listener, MaybeTlsListener};
|
||||
use crate::state::AppState;
|
||||
use crate::tls;
|
||||
use crate::tls_tests::{TestCerts, server_tls};
|
||||
@@ -287,7 +287,7 @@ async fn grpc_test_server(
|
||||
let addr = listener.local_addr().expect("local addr");
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let incoming = incoming(Listener::Tcp(listener));
|
||||
let incoming = MaybeTlsListener::plain(Listener::Tcp(listener));
|
||||
TonicServer::builder()
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
@@ -318,7 +318,7 @@ async fn grpc_tls_test_server(
|
||||
let addr = listener.local_addr().expect("local addr").to_string();
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let incoming = tls_incoming(Listener::Tcp(listener), context, tls::TLS_HANDSHAKE_TIMEOUT);
|
||||
let incoming = MaybeTlsListener::tls(Listener::Tcp(listener), context);
|
||||
TonicServer::builder()
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
@@ -413,7 +413,7 @@ async fn grpc_server_with_keepalive(
|
||||
}
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let incoming = incoming(Listener::Tcp(listener));
|
||||
let incoming = MaybeTlsListener::plain(Listener::Tcp(listener));
|
||||
builder
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
|
||||
+17
-62
@@ -27,7 +27,6 @@ pub use config::{
|
||||
ApiServerOptions, Config, CoordinatorMode, CorsConfig, DEFAULT_KEEP_ALIVE_TIMEOUT,
|
||||
HttpListenerMode, TlsConfig,
|
||||
};
|
||||
use futures::FutureExt as _;
|
||||
use hyper::body::Incoming;
|
||||
use hyper::server::conn::http1;
|
||||
use hyper_util::rt::{TokioIo, TokioTimer};
|
||||
@@ -45,7 +44,7 @@ use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_text::TextLlm;
|
||||
|
||||
use crate::listener::Listener;
|
||||
use crate::listener::{Listener, MaybeTlsListener};
|
||||
use crate::routes::build_router;
|
||||
use crate::server_info::ServerInfoSnapshot;
|
||||
use crate::state::AppState;
|
||||
@@ -179,7 +178,7 @@ where
|
||||
let listener = Listener::bind(&config.listener_mode)
|
||||
.await
|
||||
.context("failed to bind listener for OpenAI server")?;
|
||||
let bind_address = listener.local_addr()?;
|
||||
let bind_address = listener.local_addr_display()?;
|
||||
let model = state.primary_model_name().to_owned();
|
||||
let app = extend_router(build_router(state.clone()));
|
||||
|
||||
@@ -252,7 +251,6 @@ where
|
||||
// silent client cannot hold the connection open.
|
||||
let keep_alive_timeout = config.keep_alive_timeout;
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: tls::TLS_HANDSHAKE_TIMEOUT,
|
||||
header_read: if keep_alive_timeout.is_zero() {
|
||||
DEFAULT_KEEP_ALIVE_TIMEOUT
|
||||
} else {
|
||||
@@ -266,9 +264,15 @@ where
|
||||
let server_shutdown = server_shutdown.clone();
|
||||
let force_shutdown = force_shutdown.clone();
|
||||
async move {
|
||||
let listener = match tls_config {
|
||||
Some(context) => MaybeTlsListener::tls(listener, context),
|
||||
None => MaybeTlsListener::plain(listener),
|
||||
};
|
||||
let server = serve_connections(listener, app, shutdown.cancelled_owned(), timeouts);
|
||||
|
||||
let result = tokio::select! {
|
||||
result = serve_listener(listener, tls_config, app, shutdown.cancelled_owned(), timeouts) => {
|
||||
result
|
||||
result = server => {
|
||||
result.context("HTTP server failed")
|
||||
}
|
||||
_ = force_shutdown.cancelled() => {
|
||||
warn!("HTTP graceful shutdown deadline elapsed; aborting server");
|
||||
@@ -292,18 +296,11 @@ where
|
||||
shutdown.cancelled().await;
|
||||
return Ok(());
|
||||
};
|
||||
// Box to unify the TLS and plaintext arms' different stream types.
|
||||
let server = match grpc_tls {
|
||||
Some(context) => {
|
||||
let incoming =
|
||||
grpc::tls_incoming(grpc_listener, context, tls::TLS_HANDSHAKE_TIMEOUT);
|
||||
svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned()).boxed()
|
||||
}
|
||||
None => {
|
||||
let incoming = grpc::incoming(grpc_listener);
|
||||
svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned()).boxed()
|
||||
}
|
||||
let incoming = match grpc_tls {
|
||||
Some(context) => MaybeTlsListener::tls(grpc_listener, context),
|
||||
None => MaybeTlsListener::plain(grpc_listener),
|
||||
};
|
||||
let server = svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned());
|
||||
|
||||
let result = tokio::select! {
|
||||
result = server => {
|
||||
@@ -333,61 +330,19 @@ where
|
||||
/// Per-connection timeouts applied while serving HTTP/HTTPS.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct ConnectionTimeouts {
|
||||
/// Max time for a client to complete the TLS handshake (TLS path only).
|
||||
pub(crate) handshake: Duration,
|
||||
/// HTTP/1 header-read timeout (bounds idle keep-alive and the head read).
|
||||
pub(crate) header_read: Duration,
|
||||
/// Whether HTTP/1 keep-alive is enabled; `false` closes after each response.
|
||||
pub(crate) keep_alive_enabled: bool,
|
||||
}
|
||||
|
||||
/// Apply optional TLS termination and per-connection HTTP timeouts, then serve
|
||||
/// `app`. Shared by [`serve_with_router_extension`] and the TLS tests.
|
||||
async fn serve_listener(
|
||||
listener: Listener,
|
||||
tls: Option<openssl::ssl::SslContext>,
|
||||
app: Router,
|
||||
shutdown: impl Future<Output = ()> + Send + 'static,
|
||||
timeouts: ConnectionTimeouts,
|
||||
) -> Result<()> {
|
||||
match tls {
|
||||
Some(context) => {
|
||||
// tls-listener terminates TLS (handshake + timeout); serve_connections
|
||||
// owns the HTTP keep-alive/idle bound that axum::serve cannot express.
|
||||
// Failed handshakes (incl. timeouts) log at ERROR via tls-listener.
|
||||
let listener = tls_listener::builder(context)
|
||||
.handshake_timeout(timeouts.handshake)
|
||||
.listen(listener);
|
||||
serve_connections(
|
||||
listener,
|
||||
app,
|
||||
shutdown,
|
||||
timeouts.header_read,
|
||||
timeouts.keep_alive_enabled,
|
||||
)
|
||||
.await
|
||||
.context("HTTPS server failed")
|
||||
}
|
||||
None => serve_connections(
|
||||
listener,
|
||||
app,
|
||||
shutdown,
|
||||
timeouts.header_read,
|
||||
timeouts.keep_alive_enabled,
|
||||
)
|
||||
.await
|
||||
.context("HTTP server failed"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Serve `app` per connection (HTTP/1) with a keep-alive idle timeout and
|
||||
/// graceful drain. Hand-rolled on hyper because [`axum::serve()`] takes no config.
|
||||
async fn serve_connections<L>(
|
||||
mut listener: L,
|
||||
app: Router,
|
||||
shutdown: impl Future<Output = ()> + Send,
|
||||
header_read: Duration,
|
||||
keep_alive_enabled: bool,
|
||||
timeouts: ConnectionTimeouts,
|
||||
) -> Result<()>
|
||||
where
|
||||
L: axum::serve::Listener,
|
||||
@@ -404,8 +359,8 @@ where
|
||||
app.clone().map_request(|req: Request<Incoming>| req.map(Body::new)),
|
||||
);
|
||||
let mut builder = http1::Builder::new();
|
||||
builder.timer(TokioTimer::new()).header_read_timeout(header_read);
|
||||
if !keep_alive_enabled {
|
||||
builder.timer(TokioTimer::new()).header_read_timeout(timeouts.header_read);
|
||||
if !timeouts.keep_alive_enabled {
|
||||
builder.keep_alive(false);
|
||||
}
|
||||
let connection = builder.serve_connection(TokioIo::new(io), service);
|
||||
|
||||
@@ -12,13 +12,14 @@ use std::pin::Pin;
|
||||
use std::task::{Context, Poll, ready};
|
||||
|
||||
use auto_enums::enum_derive;
|
||||
use openssl::ssl::SslContext;
|
||||
use socket2::Socket;
|
||||
use tls_listener::{AsyncAccept, AsyncListener};
|
||||
use tokio::net::{TcpListener, TcpStream, UnixListener, UnixStream};
|
||||
use tonic::transport::server::{Connected, TcpConnectInfo};
|
||||
use tracing::trace;
|
||||
|
||||
use crate::HttpListenerMode;
|
||||
use crate::{HttpListenerMode, tls};
|
||||
|
||||
/// Runtime listener type used by the OpenAI-compatible HTTP or gRPC server,
|
||||
/// which is either a TCP listener or a Unix-domain listener.
|
||||
@@ -61,7 +62,7 @@ impl Listener {
|
||||
|
||||
/// Return a log-friendly local address string for either TCP or Unix
|
||||
/// sockets.
|
||||
pub fn local_addr(&self) -> Result<String> {
|
||||
pub fn local_addr_display(&self) -> Result<String> {
|
||||
match self {
|
||||
Self::Tcp(listener) => Ok(listener.local_addr()?.to_string()),
|
||||
Self::Unix(listener) => Ok(match listener.local_addr()?.as_pathname() {
|
||||
@@ -92,7 +93,7 @@ impl Listener {
|
||||
}
|
||||
}
|
||||
|
||||
fn listener_addr(&self) -> Result<ListenerAddr> {
|
||||
fn local_addr(&self) -> Result<ListenerAddr> {
|
||||
match self {
|
||||
Self::Tcp(listener) => listener.local_addr().map(ListenerAddr::Tcp),
|
||||
Self::Unix(listener) => listener.local_addr().map(ListenerAddr::Unix),
|
||||
@@ -100,6 +101,7 @@ impl Listener {
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the unified listener to plug directly into tonic's gRPC server.
|
||||
impl Connected for ListenerIo {
|
||||
type ConnectInfo = TcpConnectInfo;
|
||||
|
||||
@@ -144,7 +146,7 @@ impl axum::serve::Listener for Listener {
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> Result<Self::Addr> {
|
||||
self.listener_addr()
|
||||
self.local_addr()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -173,10 +175,98 @@ impl AsyncAccept for Listener {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncListener for Listener {
|
||||
fn local_addr(&self) -> Result<Self::Address> {
|
||||
self.listener_addr()
|
||||
self.local_addr()
|
||||
}
|
||||
}
|
||||
|
||||
/// A listener that may be either a plain TCP/UDS listener or a TLS listener over it.
|
||||
pub enum MaybeTlsListener {
|
||||
Plain(Listener),
|
||||
Tls(tls_listener::TlsListener<Listener, SslContext>),
|
||||
}
|
||||
|
||||
impl MaybeTlsListener {
|
||||
/// Create a plain listener without TLS.
|
||||
pub fn plain(listener: Listener) -> Self {
|
||||
Self::Plain(listener)
|
||||
}
|
||||
|
||||
/// Create a TLS listener over the given plain listener.
|
||||
pub fn tls(listener: Listener, context: SslContext) -> Self {
|
||||
Self::Tls(
|
||||
tls_listener::builder(context)
|
||||
.handshake_timeout(tls::TLS_HANDSHAKE_TIMEOUT)
|
||||
.listen(listener),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Listener I/O type that may be either a plain TCP/UDS stream or a TLS stream over it.
|
||||
#[derive(Debug)]
|
||||
#[enum_derive(tokio1::AsyncRead, tokio1::AsyncWrite)]
|
||||
pub enum MaybeTlsStream {
|
||||
Plain(ListenerIo),
|
||||
Tls(tokio_openssl::SslStream<ListenerIo>),
|
||||
}
|
||||
|
||||
/// Allow the maybe-TLS listener to plug directly into `axum::serve(...)`.
|
||||
impl axum::serve::Listener for MaybeTlsListener {
|
||||
type Addr = ListenerAddr;
|
||||
type Io = MaybeTlsStream;
|
||||
|
||||
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
|
||||
match self {
|
||||
Self::Plain(listener) => {
|
||||
let (io, addr) = axum::serve::Listener::accept(listener).await;
|
||||
(MaybeTlsStream::Plain(io), addr)
|
||||
}
|
||||
Self::Tls(tls_listener) => {
|
||||
let (io, addr) = axum::serve::Listener::accept(tls_listener).await;
|
||||
(MaybeTlsStream::Tls(io), addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> tokio::io::Result<Self::Addr> {
|
||||
match self {
|
||||
Self::Plain(listener) => listener.local_addr(),
|
||||
Self::Tls(tls_listener) => tls_listener.local_addr(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the maybe-TLS listener to plug directly into tonic's gRPC server.
|
||||
impl Connected for MaybeTlsStream {
|
||||
type ConnectInfo = TcpConnectInfo;
|
||||
|
||||
fn connect_info(&self) -> TcpConnectInfo {
|
||||
match self {
|
||||
Self::Plain(stream) => stream.connect_info(),
|
||||
Self::Tls(stream) => stream.get_ref().connect_info(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the maybe-TLS listener to be adaptable to tonic's incoming stream shape.
|
||||
impl futures::Stream for MaybeTlsListener {
|
||||
type Item = std::io::Result<MaybeTlsStream>;
|
||||
|
||||
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
match self.get_mut() {
|
||||
Self::Plain(listener) => {
|
||||
let listener = Pin::new(listener);
|
||||
let (io, _) = ready!(listener.poll_accept(cx))?;
|
||||
Poll::Ready(Some(Ok(MaybeTlsStream::Plain(io))))
|
||||
}
|
||||
Self::Tls(tls_listener) => {
|
||||
let tls_listener = Pin::new(tls_listener);
|
||||
let (io, _) =
|
||||
ready!(tls_listener.poll_accept(cx)).map_err(std::io::Error::other)?;
|
||||
Poll::Ready(Some(Ok(MaybeTlsStream::Tls(io))))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
//! TLS tests: `build_server_config` unit checks plus end-to-end OpenSSL handshakes
|
||||
//! through the production `serve_listener` path, with a trivial router since TLS
|
||||
//! through the production listener/connection path, with a trivial router since TLS
|
||||
//! terminates below the app.
|
||||
|
||||
use std::pin::Pin;
|
||||
@@ -23,8 +23,8 @@ use tokio_openssl::SslStream;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::config::{HttpListenerMode, TlsConfig};
|
||||
use crate::listener::Listener;
|
||||
use crate::{ConnectionTimeouts, serve_listener, tls};
|
||||
use crate::listener::{Listener, MaybeTlsListener};
|
||||
use crate::{ConnectionTimeouts, serve_connections, tls};
|
||||
|
||||
// ============================================================================
|
||||
// Test infrastructure
|
||||
@@ -251,7 +251,6 @@ fn build_tls(certs: &TestCerts, cert: &str, key: Option<&str>) -> TlsConfig {
|
||||
|
||||
/// Generous per-connection timeouts that never fire during the fast tests.
|
||||
const TEST_TIMEOUTS: ConnectionTimeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_secs(5),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
@@ -261,9 +260,8 @@ async fn spawn_server(tls_config: Option<TlsConfig>) -> (String, CancellationTok
|
||||
}
|
||||
|
||||
/// Bind an ephemeral listener and serve a trivial router via the production
|
||||
/// `serve_listener`, optionally with TLS. The listener is bound (and thus
|
||||
/// accepting into the backlog) before returning, so a client may connect
|
||||
/// immediately without a sleep.
|
||||
/// listener/connection path. The listener is bound (and thus accepting into the
|
||||
/// backlog) before returning, so a client may connect immediately without a sleep.
|
||||
async fn spawn_server_with_timeouts(
|
||||
tls_config: Option<TlsConfig>,
|
||||
timeouts: ConnectionTimeouts,
|
||||
@@ -274,7 +272,7 @@ async fn spawn_server_with_timeouts(
|
||||
})
|
||||
.await
|
||||
.expect("bind listener");
|
||||
let addr = listener.local_addr().expect("local addr");
|
||||
let addr = listener.local_addr_display().expect("local addr");
|
||||
|
||||
let server_config =
|
||||
tls_config.map(|cfg| tls::build_server_config(&cfg).expect("build server config"));
|
||||
@@ -282,14 +280,11 @@ async fn spawn_server_with_timeouts(
|
||||
let shutdown = CancellationToken::new();
|
||||
let server_shutdown = shutdown.clone();
|
||||
tokio::spawn(async move {
|
||||
let _ = serve_listener(
|
||||
listener,
|
||||
server_config,
|
||||
app,
|
||||
server_shutdown.cancelled_owned(),
|
||||
timeouts,
|
||||
)
|
||||
.await;
|
||||
let listener = match server_config {
|
||||
Some(context) => MaybeTlsListener::tls(listener, context),
|
||||
None => MaybeTlsListener::plain(listener),
|
||||
};
|
||||
let _ = serve_connections(listener, app, server_shutdown.cancelled_owned(), timeouts).await;
|
||||
});
|
||||
(addr, shutdown)
|
||||
}
|
||||
@@ -521,20 +516,19 @@ async fn plain_http_serves_when_tls_is_disabled() {
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn tls_handshake_timeout_drops_silent_client() {
|
||||
// Silent client (no ClientHello) must be dropped at the handshake deadline.
|
||||
let certs = TestCerts::generate();
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_millis(150),
|
||||
header_read: Duration::from_secs(5),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(Some(server_tls(&certs, 0)), timeouts).await;
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 0))).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
tokio::task::yield_now().await;
|
||||
tokio::time::advance(tls::TLS_HANDSHAKE_TIMEOUT + Duration::from_millis(1)).await;
|
||||
tokio::task::yield_now().await;
|
||||
|
||||
let mut buf = [0u8; 1];
|
||||
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
|
||||
let read = tokio::time::timeout(Duration::from_secs(1), tcp.read(&mut buf)).await;
|
||||
assert!(
|
||||
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
|
||||
"server must drop a stalled TLS handshake (expected close, got {read:?})"
|
||||
@@ -546,7 +540,6 @@ async fn tls_handshake_timeout_drops_silent_client() {
|
||||
async fn keep_alive_timeout_closes_idle_connection() {
|
||||
// Idle keep-alive connection must be closed at the deadline.
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
@@ -582,7 +575,6 @@ async fn keep_alive_timeout_closes_idle_tls_connection() {
|
||||
// still fires through tls-listener's post-handshake SslStream, not just plaintext.
|
||||
let certs = TestCerts::generate();
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
@@ -618,7 +610,6 @@ async fn keep_alive_timeout_closes_idle_tls_connection() {
|
||||
async fn idle_timeout_closes_silent_client() {
|
||||
// Silent client closed by the header-read timeout (http1-only arms it from byte 0).
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
@@ -638,7 +629,6 @@ async fn idle_timeout_closes_silent_client() {
|
||||
async fn keep_alive_zero_disables_keep_alive() {
|
||||
// 0 disables keep-alive (serve, then close), like uvicorn's timeout_keep_alive=0.
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_secs(5),
|
||||
keep_alive_enabled: false,
|
||||
};
|
||||
@@ -671,7 +661,6 @@ async fn disabled_keep_alive_still_closes_silent_client() {
|
||||
// Even with keep-alive off, the head read stays bounded, so a silent client
|
||||
// is dropped rather than held open.
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: false,
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user