diff --git a/rust/src/server/Cargo.toml b/rust/src/server/Cargo.toml index 59c7f8d1744..00183d7f9c2 100644 --- a/rust/src/server/Cargo.toml +++ b/rust/src/server/Cargo.toml @@ -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 diff --git a/rust/src/server/src/grpc/mod.rs b/rust/src/server/src/grpc/mod.rs index 0a71f2edc12..1fcb8674fee 100644 --- a/rust/src/server/src/grpc/mod.rs +++ b/rust/src/server/src/grpc/mod.rs @@ -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, -} - -impl GrpcTlsStream { - pub(crate) fn new(inner: SslStream) -> Self { - Self { inner } - } -} - -impl AsyncRead for GrpcTlsStream { - fn poll_read( - self: Pin<&mut Self>, - cx: &mut Context<'_>, - buf: &mut ReadBuf<'_>, - ) -> Poll> { - 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> { - Pin::new(&mut self.get_mut().inner).poll_write(cx, buf) - } - - fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - Pin::new(&mut self.get_mut().inner).poll_flush(cx) - } - - fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - 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> { - 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> { - 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, diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index 311b819e0c8..65eedbf7e87 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -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) diff --git a/rust/src/server/src/lib.rs b/rust/src/server/src/lib.rs index 55b4f5ffae3..78b18f8f477 100644 --- a/rust/src/server/src/lib.rs +++ b/rust/src/server/src/lib.rs @@ -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, - app: Router, - shutdown: impl Future + 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( mut listener: L, app: Router, shutdown: impl Future + 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| 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); diff --git a/rust/src/server/src/listener.rs b/rust/src/server/src/listener.rs index 75fe25aef51..b1dcc919678 100644 --- a/rust/src/server/src/listener.rs +++ b/rust/src/server/src/listener.rs @@ -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 { + pub fn local_addr_display(&self) -> Result { 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 { + fn local_addr(&self) -> Result { 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.listener_addr() + self.local_addr() } } @@ -173,10 +175,98 @@ impl AsyncAccept for Listener { } } } - impl AsyncListener for Listener { fn local_addr(&self) -> Result { - 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), +} + +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), +} + +/// 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 { + 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; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + 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)))) + } + } } } diff --git a/rust/src/server/src/tls_tests.rs b/rust/src/server/src/tls_tests.rs index 7e0a6dccc30..c9bd063d77d 100644 --- a/rust/src/server/src/tls_tests.rs +++ b/rust/src/server/src/tls_tests.rs @@ -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) -> (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, 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, };