use std::convert::Infallible;
use std::sync::Arc;
use hyper::body::Incoming;
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Request, Response};
use hyper_util::rt::{TokioIo, TokioTimer};
use tokio::sync::broadcast;
use crate::primitives::request_body_policy::RequestBodyPolicy;
use crate::response::BoxBodyInner;
use crate::server::config::RuntimeConfig;
use crate::server::service::{Service, ServiceError};
use crate::server::RuntimeState;
pub async fn serve_connection<I, S>(
io: TokioIo<I>,
service: S,
config: &RuntimeConfig,
shutdown_rx: &mut broadcast::Receiver<()>,
conn_id: u64,
) where
I: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
S: hyper::service::Service<
Request<Incoming>,
Response = Response<BoxBodyInner>,
Error = Infallible,
> + 'static,
{
let conn = http1::Builder::new()
.timer(TokioTimer::new())
.header_read_timeout(config.header_read_timeout)
.serve_connection(io, service)
.with_upgrades();
let mut conn = std::pin::pin!(conn);
tokio::select! {
result = tokio::time::timeout(config.connection_total_timeout, &mut conn) => {
match result {
Ok(Ok(())) => {
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::KeepAliveClosed,
"connection closed",
)
.connection_id(conn_id),
);
}
Ok(Err(e)) => {
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::ClientDisconnect,
format!("connection error: {}", e),
)
.connection_id(conn_id),
);
}
Err(_elapsed) => {
crate::ops::global_counters().connection_total_timeouts.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Warn,
crate::ops::EventKind::ConnectionTotalTimeout,
"connection total timeout",
)
.connection_id(conn_id),
);
conn.as_mut().graceful_shutdown();
let _ = conn.await;
}
}
}
_ = shutdown_rx.recv() => {
conn.as_mut().graceful_shutdown();
let _ = conn.await;
}
}
}
#[allow(clippy::too_many_arguments)]
pub async fn serve_connection_with_runtime_state<I, S>(
io: TokioIo<I>,
service: S,
config: &RuntimeConfig,
runtime_state: Arc<RuntimeState>,
shutdown_rx: &mut broadcast::Receiver<()>,
conn_id: u64,
local_addr: std::net::SocketAddr,
remote_addr: std::net::SocketAddr,
tls: bool,
tls_info: Option<crate::primitives::connection_info::TlsInfo>,
) where
I: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
S: Service,
{
let service = std::sync::Arc::new(service);
let config = Arc::new(config.clone());
let handler_timeout = config.handler_timeout;
let body_read_timeout = config.body_read_timeout;
let max_body_bytes = config.max_request_body_bytes;
let tls_info = std::sync::Arc::new(tls_info);
let file_stream_semaphore = runtime_state.file_stream_semaphore().clone();
let response_config = config.clone();
let hyper_service = service_fn(move |req: Request<Incoming>| {
let service = service.clone();
let tls_info = tls_info.clone();
let file_stream_semaphore = file_stream_semaphore.clone();
let config = response_config.clone();
async move {
let head = match convert_request_head(&req) {
Ok(h) => h,
Err(e) => {
return Ok::<_, Infallible>(finalize_runtime_response(
e.to_response(),
&config,
));
}
};
if head.method().as_str() == "TRACE"
&& (req
.headers()
.get(hyper::header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok())
.is_some_and(|length| length > 0)
|| req.headers().contains_key(hyper::header::TRANSFER_ENCODING))
{
let mut response = crate::response::bad_request(false);
response.headers_mut().insert(
hyper::header::CONNECTION,
hyper::header::HeaderValue::from_static("close"),
);
return Ok::<_, Infallible>(finalize_runtime_response(response, &config));
}
let service_policy = service.request_body_policy(&head);
let effective_policy = select_body_policy(service_policy, max_body_bytes);
let (parts, body) = req.into_parts();
if let Err(e) = validate_body_framing(&parts.headers) {
crate::ops::global_counters()
.parser_rejects
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::ParserRejection,
format!("parser rejection: {}", e),
)
.connection_id(conn_id),
);
return Ok::<_, Infallible>(finalize_runtime_response(e.to_response(), &config));
}
let declared_length = parts
.headers
.get(hyper::header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok());
if let Some(len) = declared_length {
if let Some(limit) = effective_policy.max_bytes() {
if len > limit {
crate::ops::global_counters()
.body_rejections
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::BodyPolicyRejection,
"body too large",
)
.connection_id(conn_id)
.field(crate::ops::Field::U64("declared_bytes".into(), len))
.field(crate::ops::Field::U64("limit_bytes".into(), limit)),
);
let err = crate::primitives::request_body_error::RequestBodyError::DeclaredLengthTooLarge {
declared: len,
limit,
};
return Ok::<_, Infallible>(finalize_runtime_response(
body_error_to_response(err, &head),
&config,
));
}
}
}
if effective_policy.is_reject() {
if let Some(expect) = parts.headers.get(hyper::header::EXPECT) {
if expect == "100-continue" {
crate::ops::global_counters()
.body_rejections
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::BodyPolicyRejection,
"100-continue rejected by body policy",
)
.connection_id(conn_id),
);
let mut response = crate::response::payload_too_large(false);
response.headers_mut().insert(
hyper::header::CONNECTION,
hyper::header::HeaderValue::from_static("close"),
);
return Ok::<_, Infallible>(finalize_runtime_response(response, &config));
}
}
}
let has_body = declared_length.is_some_and(|len| len > 0)
|| parts.headers.contains_key(hyper::header::TRANSFER_ENCODING);
if effective_policy.is_reject() && has_body {
crate::ops::global_counters()
.body_rejections
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::BodyPolicyRejection,
"request body rejected by policy",
)
.connection_id(conn_id),
);
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::ServiceInvocationSuppressed,
"service invocation suppressed: body rejected by policy",
)
.connection_id(conn_id),
);
let mut response = crate::response::payload_too_large(false);
response.headers_mut().insert(
hyper::header::CONNECTION,
hyper::header::HeaderValue::from_static("close"),
);
return Ok::<_, Infallible>(finalize_runtime_response(response, &config));
}
let body_limit = effective_policy.max_bytes().unwrap_or(u64::MAX);
let request_body = match &effective_policy {
RequestBodyPolicy::Reject => crate::primitives::request_body::RequestBody::empty(),
_ => crate::primitives::request_body::RequestBody::from_incoming(
wrap_incoming_body(body),
declared_length,
body_limit,
),
};
let consumed_flag = request_body.consumed_flag();
match &effective_policy {
RequestBodyPolicy::Reject => {
let connection =
build_connection_info(local_addr, remote_addr, tls, (*tls_info).clone());
let request =
crate::primitives::request::Request::new(head, request_body, connection);
let result = tokio::time::timeout(handler_timeout, service.call(request)).await;
let response = match result {
Ok(Ok(canonical)) => {
match crate::primitives::canonical::to_hyper_response_with_file_stream_semaphore(canonical, &file_stream_semaphore) {
Ok(r) => r,
Err(crate::primitives::canonical::ResponseConstructionError::FileStreamLimit) => crate::response::service_unavailable(),
Err(_) => crate::response::internal_error(),
}
}
Ok(Err(service_err)) => {
let severity = if service_err.is_panic() || !service_err.is_timeout() {
crate::ops::Severity::Error
} else {
crate::ops::Severity::Warn
};
crate::ops::Logger::global().emit(
crate::ops::Event::new(
severity,
crate::ops::EventKind::ServiceError,
service_err.to_string(),
)
.connection_id(conn_id),
);
service_err.to_response()
}
Err(_elapsed) => {
crate::ops::Logger::global().emit(crate::ops::Event::new(
crate::ops::Severity::Warn,
crate::ops::EventKind::ServiceTimeout,
"handler timed out",
));
ServiceError::timeout("handler timed out".to_string()).to_response()
}
};
Ok::<_, Infallible>(finalize_runtime_response(response, &config))
}
RequestBodyPolicy::Buffer { .. } => {
let request_body = match tokio::time::timeout(
body_read_timeout,
request_body.read_all(),
)
.await
{
Ok(Ok(bytes)) => crate::primitives::request_body::RequestBody::from_bytes(
bytes, body_limit,
),
Ok(Err(err)) => {
return Ok::<_, Infallible>(finalize_runtime_response(
body_error_to_response(err, &head),
&config,
));
}
Err(_elapsed) => {
crate::ops::global_counters()
.body_read_timeouts
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
crate::ops::Logger::global().emit(crate::ops::Event::new(
crate::ops::Severity::Warn,
crate::ops::EventKind::BodyReadTimeout,
"body read timeout",
));
let err = crate::primitives::request_body_error::RequestBodyError::ReadTimeout;
return Ok::<_, Infallible>(finalize_runtime_response(
body_error_to_response(err, &head),
&config,
));
}
};
let connection =
build_connection_info(local_addr, remote_addr, tls, (*tls_info).clone());
let request =
crate::primitives::request::Request::new(head, request_body, connection);
let result = tokio::time::timeout(handler_timeout, service.call(request)).await;
let response = match result {
Ok(Ok(canonical)) => {
match crate::primitives::canonical::to_hyper_response_with_file_stream_semaphore(canonical, &file_stream_semaphore) {
Ok(r) => r,
Err(crate::primitives::canonical::ResponseConstructionError::FileStreamLimit) => crate::response::service_unavailable(),
Err(_) => crate::response::internal_error(),
}
}
Ok(Err(service_err)) => {
let severity = if service_err.is_panic() || !service_err.is_timeout() {
crate::ops::Severity::Error
} else {
crate::ops::Severity::Warn
};
crate::ops::Logger::global().emit(
crate::ops::Event::new(
severity,
crate::ops::EventKind::ServiceError,
service_err.to_string(),
)
.connection_id(conn_id),
);
service_err.to_response()
}
Err(_elapsed) => {
crate::ops::Logger::global().emit(crate::ops::Event::new(
crate::ops::Severity::Warn,
crate::ops::EventKind::ServiceTimeout,
"handler timed out",
));
ServiceError::timeout("handler timed out".to_string()).to_response()
}
};
Ok::<_, Infallible>(finalize_runtime_response(response, &config))
}
RequestBodyPolicy::Stream { .. } => {
let effective_timeout = body_read_timeout.min(handler_timeout);
let connection =
build_connection_info(local_addr, remote_addr, tls, (*tls_info).clone());
let request =
crate::primitives::request::Request::new(head, request_body, connection);
let result =
tokio::time::timeout(effective_timeout, service.call(request)).await;
let response = match result {
Ok(Ok(canonical)) => {
match crate::primitives::canonical::to_hyper_response_with_file_stream_semaphore(canonical, &file_stream_semaphore) {
Ok(r) => r,
Err(crate::primitives::canonical::ResponseConstructionError::FileStreamLimit) => crate::response::service_unavailable(),
Err(_) => crate::response::internal_error(),
}
}
Ok(Err(service_err)) => {
let severity = if service_err.is_panic() || !service_err.is_timeout() {
crate::ops::Severity::Error
} else {
crate::ops::Severity::Warn
};
crate::ops::Logger::global().emit(
crate::ops::Event::new(
severity,
crate::ops::EventKind::ServiceError,
service_err.to_string(),
)
.connection_id(conn_id),
);
service_err.to_response()
}
Err(_elapsed) => {
crate::ops::Logger::global().emit(crate::ops::Event::new(
crate::ops::Severity::Warn,
crate::ops::EventKind::ServiceTimeout,
"handler timed out",
));
ServiceError::timeout("handler timed out".to_string()).to_response()
}
};
let incomplete = !consumed_flag.load(std::sync::atomic::Ordering::Acquire);
if incomplete {
crate::ops::Logger::global().emit(
crate::ops::Event::new(
crate::ops::Severity::Debug,
crate::ops::EventKind::IncompleteBodyClose,
"service returned with unconsumed body; connection will close",
)
.connection_id(conn_id),
);
}
let mut response = finalize_runtime_response(response, &config);
if incomplete {
response.headers_mut().insert(
hyper::header::CONNECTION,
hyper::header::HeaderValue::from_static("close"),
);
}
Ok::<_, Infallible>(response)
}
}
}
});
serve_connection(io, hyper_service, &config, shutdown_rx, conn_id).await;
}
fn select_body_policy(service_policy: RequestBodyPolicy, max_body_bytes: u64) -> RequestBodyPolicy {
match service_policy {
RequestBodyPolicy::Reject => RequestBodyPolicy::Reject,
RequestBodyPolicy::Buffer { max_bytes } => {
let effective = max_bytes.min(max_body_bytes);
if effective == 0 {
RequestBodyPolicy::Reject
} else {
RequestBodyPolicy::Buffer {
max_bytes: effective,
}
}
}
RequestBodyPolicy::Stream { max_bytes } => {
let effective = max_bytes.min(max_body_bytes);
if effective == 0 {
RequestBodyPolicy::Reject
} else {
RequestBodyPolicy::Stream {
max_bytes: effective,
}
}
}
}
}
fn body_error_to_response(
err: crate::primitives::request_body_error::RequestBodyError,
_head: &crate::primitives::request_head::RequestHead,
) -> hyper::Response<BoxBodyInner> {
let status = err.to_status_code();
let status =
hyper::StatusCode::from_u16(status).unwrap_or(hyper::StatusCode::INTERNAL_SERVER_ERROR);
let should_close = matches!(
status,
hyper::StatusCode::BAD_REQUEST
| hyper::StatusCode::REQUEST_TIMEOUT
| hyper::StatusCode::PAYLOAD_TOO_LARGE
| hyper::StatusCode::HTTP_VERSION_NOT_SUPPORTED
);
let body_text = match status.as_u16() {
400 => "400 Bad Request\n",
408 => "408 Request Timeout\n",
413 => "413 Payload Too Large\n",
501 => "501 Not Implemented\n",
_ => "500 Internal Server Error\n",
};
let is_head = _head.method().is_head();
let mut resp = crate::response::canonical_error(status, body_text, is_head);
if should_close {
resp.headers_mut().insert(
hyper::header::CONNECTION,
hyper::header::HeaderValue::from_static("close"),
);
}
resp
}
fn build_connection_info(
local_addr: std::net::SocketAddr,
remote_addr: std::net::SocketAddr,
tls: bool,
tls_info: Option<crate::primitives::connection_info::TlsInfo>,
) -> crate::primitives::connection_info::ConnectionInfo {
crate::primitives::connection_info::ConnectionInfo {
local_addr,
remote_addr,
scheme: if tls {
crate::primitives::connection_info::Scheme::Https
} else {
crate::primitives::connection_info::Scheme::Http
},
tls: tls_info,
}
}
fn finalize_runtime_response(
mut response: hyper::Response<BoxBodyInner>,
config: &RuntimeConfig,
) -> hyper::Response<BoxBodyInner> {
response.headers_mut().remove(hyper::header::SERVER);
if let Some(value) = &config.server_header {
if let Ok(value) = hyper::header::HeaderValue::from_str(value) {
response.headers_mut().insert(hyper::header::SERVER, value);
}
}
response
}
fn validate_body_framing(headers: &hyper::HeaderMap) -> Result<(), ServiceError> {
let has_te = headers.contains_key(hyper::header::TRANSFER_ENCODING);
let cl_values: Vec<_> = headers
.get_all(hyper::header::CONTENT_LENGTH)
.iter()
.collect();
let has_cl = !cl_values.is_empty();
let duplicate_cl = cl_values.len() > 1;
if has_te && has_cl {
return Err(ServiceError::rejected(
400,
"conflicting Transfer-Encoding and Content-Length",
));
}
if duplicate_cl {
return Err(ServiceError::rejected(
400,
"duplicate Content-Length headers",
));
}
Ok(())
}
fn wrap_incoming_body(
body: Incoming,
) -> impl futures_util::Stream<
Item = Result<bytes::Bytes, crate::primitives::request_body::IncomingError>,
> + Send
+ 'static {
use futures_util::StreamExt;
http_body_util::BodyStream::new(body).filter_map(|result| async {
match result {
Ok(frame) => frame.into_data().ok().map(Ok),
Err(e) => Some(Err(crate::primitives::request_body::IncomingError(
e.to_string(),
))),
}
})
}
fn convert_request_head(
req: &Request<Incoming>,
) -> Result<crate::primitives::request_head::RequestHead, ServiceError> {
use crate::primitives::header_block::HeaderBlock;
use crate::primitives::method::Method;
use crate::primitives::request_target::RequestTarget;
use crate::primitives::version::HttpVersion;
let method = match req.method().as_str() {
"GET" => Method::get(),
"HEAD" => Method::head(),
"POST" => Method::post(),
"PUT" => Method::put(),
"DELETE" => Method::delete(),
"PATCH" => Method::patch(),
"OPTIONS" => Method::options(),
"TRACE" => Method::trace(),
other => Method::new(other)
.map_err(|_| ServiceError::rejected(400, format!("invalid method: {}", other)))?,
};
let version = match req.version() {
hyper::Version::HTTP_10 => HttpVersion::Http10,
hyper::Version::HTTP_11 => HttpVersion::Http11,
other => {
return Err(ServiceError::rejected(
505,
format!("unsupported HTTP version: {:?}", other),
))
}
};
let raw_target = req
.uri()
.path_and_query()
.map(|pq| pq.as_str())
.unwrap_or("/");
if req.uri().scheme_str().is_some() {
return Err(ServiceError::rejected(
400,
"absolute-form request target not allowed",
));
}
if raw_target == "*" {
return Err(ServiceError::rejected(
405,
format!("method not allowed: {}", method.as_str()),
));
}
let target = RequestTarget::parse(raw_target)
.map_err(|e| ServiceError::rejected(400, format!("invalid request target: {}", e)))?;
let mut headers = HeaderBlock::new();
for (name, value) in req.headers().iter() {
let header_name = crate::primitives::header_block::HeaderName::new(name.as_str())
.map_err(|_| ServiceError::rejected(400, format!("invalid header name: {}", name)))?;
let header_value = match value.to_str() {
Ok(v) => crate::primitives::header_block::HeaderValue::new(v).map_err(|_| {
ServiceError::rejected(400, format!("invalid header value for {}", name))
})?,
Err(_) => {
return Err(ServiceError::rejected(
400,
format!("non-UTF-8 header value for {}", name),
))
}
};
headers.push(header_name, header_value);
}
Ok(crate::primitives::request_head::RequestHead::new(
method, target, version, headers,
))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{ServeConfig, ServeState};
use crate::server::static_service::StaticService;
use std::sync::Arc;
use tempfile::TempDir;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
fn build_state(tmp: &TempDir) -> Arc<ServeState> {
let config = Arc::new(ServeConfig {
root: tmp.path().to_path_buf(),
..ServeConfig::default()
});
Arc::new(ServeState::new(config).unwrap())
}
#[tokio::test]
async fn serve_connection_handles_get() {
let tmp = TempDir::new().unwrap();
std::fs::write(tmp.path().join("hello.txt"), "hello").unwrap();
let state = build_state(&tmp);
let config = RuntimeConfig::default();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (tx, _rx) = broadcast::channel::<()>(1);
let state_clone = state.clone();
let server = tokio::spawn(async move {
let (stream, remote_addr) = listener.accept().await.unwrap();
let mut shutdown_rx = tx.subscribe();
let runtime_state = Arc::new(RuntimeState::new(&config));
serve_connection_with_runtime_state(
TokioIo::new(stream),
StaticService::from_state(state_clone),
&config,
runtime_state,
&mut shutdown_rx,
1,
addr,
remote_addr,
false,
None,
)
.await;
});
let mut client = tokio::net::TcpStream::connect(addr).await.unwrap();
client
.write_all(b"GET /hello.txt HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut buf = Vec::new();
client.read_to_end(&mut buf).await.unwrap();
let _ = server.await;
let response = String::from_utf8_lossy(&buf);
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"unexpected response: {}",
response
);
}
#[test]
fn runtime_server_header_replaces_service_value() {
let config = RuntimeConfig::builder()
.server_header("eggserve-test".into())
.build()
.unwrap();
let mut response = crate::response::not_found(false);
response.headers_mut().insert(
hyper::header::SERVER,
hyper::header::HeaderValue::from_static("spoofed"),
);
let response = finalize_runtime_response(response, &config);
assert_eq!(
response.headers().get(hyper::header::SERVER).unwrap(),
"eggserve-test"
);
assert_eq!(
response
.headers()
.get_all(hyper::header::SERVER)
.iter()
.count(),
1
);
}
}