use std::time::Duration;
use rama_core::Layer;
use rama_core::error::{BoxError, BoxErrorExt as _};
use rama_core::extensions::ExtensionsRef;
use rama_http_types::Request;
use rama_net::address::{HostWithOptPort, ProxyAddress};
use rama_net::client::pool::{ConnID, MultiplexPool, MuxSelection, PooledConnector, ReqToConnID};
use rama_net::client::{ConnectorService, ConnectorTarget};
use rama_net::{AuthorityInputExt, Protocol, ProtocolInputExt};
use super::{BindBodyToConnLayer, BindBodyToConnector};
#[derive(Clone, Debug, Default)]
#[non_exhaustive]
pub struct BasicHttpConnIdentifier;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct BasicHttpConId {
pub protocol: Option<Protocol>,
pub authority: HostWithOptPort,
pub proxy_address: Option<ProxyAddress>,
pub connector_target: Option<ConnectorTarget>,
}
impl ConnID for BasicHttpConId {
#[cfg(feature = "opentelemetry")]
fn attributes(&self) -> impl Iterator<Item = rama_core::telemetry::opentelemetry::KeyValue> {
self.protocol
.as_ref()
.map(|protocol| {
rama_core::telemetry::opentelemetry::KeyValue::new("protocol", protocol.to_string())
})
.into_iter()
.chain([rama_core::telemetry::opentelemetry::KeyValue::new(
"authority",
self.authority.to_string(),
)])
}
}
impl<Body> ReqToConnID<Request<Body>> for BasicHttpConnIdentifier {
type ID = BasicHttpConId;
fn id(&self, req: &Request<Body>) -> Result<Self::ID, BoxError> {
let authority = req
.authority()
.ok_or_else(|| BoxError::from_static_str("no authority found in http request"))?;
let protocol = req.protocol().cloned();
Ok(BasicHttpConId {
protocol,
authority,
proxy_address: req.extensions().get_ref().cloned(),
connector_target: req.extensions().get_ref().cloned(),
})
}
}
#[derive(Debug, Clone)]
pub struct HttpPooledConnectorConfig {
pub max_total: usize,
pub max_concurrent_streams: usize,
pub selection: MuxSelection,
pub idle_timeout: Option<Duration>,
pub wait_for_pool_timeout: Option<Duration>,
}
impl Default for HttpPooledConnectorConfig {
fn default() -> Self {
Self {
max_total: 50,
max_concurrent_streams: 100,
selection: MuxSelection::default(),
idle_timeout: Some(Duration::from_secs(300)),
wait_for_pool_timeout: Some(Duration::from_secs(120)),
}
}
}
impl HttpPooledConnectorConfig {
pub fn build_connector<S>(
self,
inner: S,
) -> Result<
BindBodyToConnector<
PooledConnector<
S,
MultiplexPool<S::Connection, BasicHttpConId>,
BasicHttpConnIdentifier,
>,
>,
BoxError,
>
where
S: ConnectorService<Request>,
{
let pool = MultiplexPool::try_new(self.max_concurrent_streams, self.max_total)?
.with_selection(self.selection)
.maybe_with_idle_timeout(self.idle_timeout);
let connector = PooledConnector::new(inner, pool, BasicHttpConnIdentifier)
.maybe_with_wait_for_pool_timeout(self.wait_for_pool_timeout);
Ok(BindBodyToConnLayer::new().into_layer(connector))
}
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use rama_core::error::BoxError;
use rama_core::extensions::ExtensionsRef;
use rama_core::rt::Executor;
use rama_core::service::service_fn;
use rama_core::{Layer, Service};
use rama_http_types::body::util::BodyExt as _;
use rama_http_types::{Body, HeaderValue, Request, Response, Version};
use rama_net::client::ConnectorService;
use rama_net::test_utils::client::MockConnectorService;
use rama_utils::octets::kib;
use tokio::time::sleep;
use super::HttpPooledConnectorConfig;
use crate::client::HttpConnectorLayer;
use crate::server::HttpServer;
fn create_test_request(version: Version) -> Request {
Request::builder()
.uri("https://www.example.com")
.version(version)
.body(Body::from("a random request body"))
.unwrap()
}
fn tagging_mock_connector() -> impl ConnectorService<
Request,
Connection: Service<Request, Output = Response, Error = BoxError> + ExtensionsRef,
> {
let conns = Arc::new(AtomicUsize::new(0));
HttpConnectorLayer::default().into_layer(MockConnectorService::new(move || {
let conn_id = conns.fetch_add(1, Ordering::Relaxed);
let resps = Arc::new(AtomicUsize::new(0));
HttpServer::auto(Executor::default()).service(service_fn(move |_req: Request| {
let resps = resps.clone();
async move {
let resp_id = resps.fetch_add(1, Ordering::Relaxed);
let mut resp = Response::new(Body::from("ok"));
let headers = resp.headers_mut();
headers.insert("x-conn-id", HeaderValue::from(conn_id as u64));
headers.insert("x-resp-id", HeaderValue::from(resp_id as u64));
Ok::<_, Infallible>(resp)
}
}))
}))
}
fn conn_id(resp: &Response) -> u64 {
resp.headers()
.get("x-conn-id")
.unwrap()
.to_str()
.unwrap()
.parse()
.unwrap()
}
#[tokio::test]
async fn pool_keeps_h2_connection_in_use_until_response_body_consumed() {
let connector = HttpPooledConnectorConfig {
max_concurrent_streams: 1,
max_total: 4,
..Default::default()
}
.build_connector(tagging_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_2);
let est1 = connector.serve(req()).await.unwrap();
let resp1 = est1.conn.serve(req()).await.unwrap();
drop(est1);
let est2 = connector.serve(req()).await.unwrap();
let resp2 = est2.conn.serve(req()).await.unwrap();
assert_eq!(conn_id(&resp1), 0);
assert_eq!(
conn_id(&resp2),
1,
"second request must not reuse a connection whose response body is still in flight"
);
}
#[tokio::test]
async fn pool_keeps_h1_connection_in_use_until_response_body_consumed() {
let connector = HttpPooledConnectorConfig {
max_concurrent_streams: 1,
max_total: 4,
..Default::default()
}
.build_connector(tagging_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_11);
let est1 = connector.serve(req()).await.unwrap();
let resp1 = est1.conn.serve(req()).await.unwrap();
drop(est1);
let est2 = connector.serve(req()).await.unwrap();
let resp2 = est2.conn.serve(req()).await.unwrap();
assert_eq!(conn_id(&resp1), 0);
assert_eq!(
conn_id(&resp2),
1,
"h1: second request must not reuse a connection whose response body is still in flight"
);
}
#[tokio::test(start_paused = true)]
async fn pool_does_not_reuse_h1_connection_after_server_close() {
let conns = Arc::new(AtomicUsize::new(0));
let inner =
HttpConnectorLayer::default().into_layer(MockConnectorService::new(move || {
let conn_id = conns.fetch_add(1, Ordering::Relaxed);
HttpServer::auto(Executor::default()).service(service_fn(
move |_req: Request| async move {
let mut resp = Response::new(Body::from("ok"));
let headers = resp.headers_mut();
headers.insert("x-conn-id", HeaderValue::from(conn_id as u64));
headers.insert("connection", HeaderValue::from_static("close"));
Ok::<_, Infallible>(resp)
},
))
}));
let connector = HttpPooledConnectorConfig {
max_total: 4,
..Default::default()
}
.build_connector(inner)
.unwrap();
let req = || create_test_request(Version::HTTP_11);
let est1 = connector.serve(req()).await.unwrap();
let resp1 = est1.conn.serve(req()).await.unwrap();
drop(est1);
let id1 = conn_id(&resp1);
resp1.into_body().collect().await.unwrap();
sleep(Duration::from_millis(50)).await;
let est2 = connector.serve(req()).await.unwrap();
let resp2 = est2.conn.serve(req()).await.unwrap();
assert_eq!(id1, 0);
assert_ne!(
conn_id(&resp2),
id1,
"must not reuse an h1 connection the server closed"
);
}
#[tokio::test(start_paused = true)]
async fn pool_does_not_reuse_h1_connection_after_body_dropped_early() {
let conns = Arc::new(AtomicUsize::new(0));
let inner =
HttpConnectorLayer::default().into_layer(MockConnectorService::new(move || {
let conn_id = conns.fetch_add(1, Ordering::Relaxed);
HttpServer::auto(Executor::default()).service(service_fn(
move |_req: Request| async move {
let mut resp = Response::new(Body::from(vec![0u8; kib(1024)]));
resp.headers_mut()
.insert("x-conn-id", HeaderValue::from(conn_id as u64));
Ok::<_, Infallible>(resp)
},
))
}));
let connector = HttpPooledConnectorConfig {
max_total: 4,
..Default::default()
}
.build_connector(inner)
.unwrap();
let req = || create_test_request(Version::HTTP_11);
let est1 = connector.serve(req()).await.unwrap();
let resp1 = est1.conn.serve(req()).await.unwrap();
drop(est1);
let id1 = conn_id(&resp1);
drop(resp1);
sleep(Duration::from_millis(50)).await;
let est2 = connector.serve(req()).await.unwrap();
let resp2 = est2.conn.serve(req()).await.unwrap();
assert_eq!(id1, 0);
assert_ne!(
conn_id(&resp2),
id1,
"must not reuse an h1 connection whose response body was abandoned mid-stream"
);
}
#[tokio::test]
async fn pool_reuses_connection_after_body_consumed() {
let connector = HttpPooledConnectorConfig {
max_concurrent_streams: 1,
max_total: 4,
..Default::default()
}
.build_connector(tagging_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_2);
let resp1 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(conn_id(&resp1), 0);
resp1.into_body().collect().await.unwrap();
let resp2 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(
conn_id(&resp2),
0,
"connection must be reused once its response body is consumed"
);
}
#[tokio::test]
async fn pool_multiplexes_on_h2() {
let connector = HttpPooledConnectorConfig::default()
.build_connector(tagging_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_2);
let resp1 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
let resp2 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(conn_id(&resp1), 0);
assert_eq!(
conn_id(&resp2),
0,
"h2: a second in-flight request multiplexes onto the same connection"
);
}
#[tokio::test]
async fn pool_does_not_multiplex_on_h1() {
let connector = HttpPooledConnectorConfig::default()
.build_connector(tagging_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_11);
let resp1 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
let resp2 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(conn_id(&resp1), 0);
assert_eq!(
conn_id(&resp2),
1,
"h1 does not multiplex: a second in-flight request needs a new connection"
);
}
#[tokio::test]
async fn pool_respects_max_concurrent_streams() {
let connector = HttpPooledConnectorConfig {
max_concurrent_streams: 2,
max_total: 4,
..Default::default()
}
.build_connector(tagging_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_2);
let resp1 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
let resp2 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
let resp3 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(conn_id(&resp1), 0);
assert_eq!(conn_id(&resp2), 0, "second request fits on connection 0");
assert_eq!(
conn_id(&resp3),
1,
"third request exceeds the per-connection limit, needs a new connection"
);
}
fn large_body_mock_connector() -> impl ConnectorService<
Request,
Connection: Service<Request, Output = Response, Error = BoxError> + ExtensionsRef,
> {
let conns = Arc::new(AtomicUsize::new(0));
HttpConnectorLayer::default().into_layer(MockConnectorService::new(move || {
let conn_id = conns.fetch_add(1, Ordering::Relaxed);
HttpServer::auto(Executor::default()).service(service_fn(
move |_req: Request| async move {
let mut resp = Response::new(Body::from(vec![0u8; kib(1024)]));
resp.headers_mut()
.insert("x-conn-id", HeaderValue::from(conn_id as u64));
Ok::<_, Infallible>(resp)
},
))
}))
}
#[tokio::test]
async fn pool_binds_connection_across_streaming_body() {
let connector = HttpPooledConnectorConfig {
max_concurrent_streams: 1,
max_total: 4,
..Default::default()
}
.build_connector(large_body_mock_connector())
.unwrap();
let req = || create_test_request(Version::HTTP_2);
let resp1 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
let id1 = conn_id(&resp1);
let mut body1 = resp1.into_body();
assert!(
body1.frame().await.is_some(),
"streaming body should yield at least one frame"
);
let resp2 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(id1, 0);
assert_ne!(
conn_id(&resp2),
id1,
"a connection still streaming its response body must not be reused"
);
while let Some(frame) = body1.frame().await {
frame.unwrap();
}
let resp3 = connector
.serve(req())
.await
.unwrap()
.conn
.serve(req())
.await
.unwrap();
assert_eq!(
conn_id(&resp3),
id1,
"a connection is reused once its streaming body reaches end-of-stream"
);
}
}