use std::collections::VecDeque;
use std::fmt;
use std::sync::{Arc, Mutex, Weak};
use http::Uri;
use http::uri::{Authority, Scheme};
use crate::Error;
use crate::config::Config;
use crate::http;
use crate::proxy::Proxy;
use crate::transport::time::{Duration, Instant};
use crate::transport::{Buffers, ConnectionDetails, Connector, NextTimeout, Transport};
use crate::util::DebugAuthority;
pub(crate) struct ConnectionPool {
connector: Box<dyn Connector<Out = Box<dyn Transport>>>,
pool: Arc<Mutex<Pool>>,
}
impl ConnectionPool {
pub fn new(connector: Box<dyn Connector<Out = Box<dyn Transport>>>, config: &Config) -> Self {
ConnectionPool {
connector,
pool: Arc::new(Mutex::new(Pool::new(config))),
}
}
pub fn connect(
&self,
details: &ConnectionDetails,
max_idle_age: Duration,
use_pool: bool,
) -> Result<Connection, Error> {
let key = details.into();
if use_pool {
let mut pool = self.pool.lock().unwrap();
pool.purge(details.now);
if let Some(conn) = pool.get(&key, max_idle_age, details.now) {
debug!("Use pooled: {:?}", key);
return Ok(conn);
}
}
let transport = self.run_connector(details)?;
let conn = Connection {
transport,
key,
last_use: details.now,
pool: if use_pool {
Arc::downgrade(&self.pool)
} else {
Weak::new()
},
position_per_host: None,
};
Ok(conn)
}
pub fn run_connector(&self, details: &ConnectionDetails) -> Result<Box<dyn Transport>, Error> {
let transport = self
.connector
.connect(details, None)?
.ok_or(Error::ConnectionFailed)?;
Ok(transport)
}
#[cfg(test)]
pub fn pool_count(&self) -> usize {
let lock = self.pool.lock().unwrap();
lock.lru.len()
}
}
pub(crate) struct Connection {
transport: Box<dyn Transport>,
key: PoolKey,
last_use: Instant,
pool: Weak<Mutex<Pool>>,
position_per_host: Option<usize>,
}
impl Connection {
pub fn buffers(&mut self) -> &mut dyn Buffers {
self.transport.buffers()
}
pub fn transmit_output(&mut self, amount: usize, timeout: NextTimeout) -> Result<(), Error> {
if timeout.after.is_zero() {
return Err(Error::Timeout(timeout.reason));
}
self.transport.transmit_output(amount, timeout)
}
pub fn maybe_await_input(&mut self, timeout: NextTimeout) -> Result<bool, Error> {
if timeout.after.is_zero() {
return Err(Error::Timeout(timeout.reason));
}
self.transport.maybe_await_input(timeout)
}
pub fn consume_input(&mut self, amount: usize) {
self.transport.buffers().input_consume(amount)
}
pub fn close(self) {
debug!("Close: {:?}", self.key);
}
pub fn reuse(mut self, now: Instant) {
if !self.transport.buffers().input().is_empty() {
debug!("Unconsumed input. Closing connection");
return;
}
if !self.transport.is_open() {
return;
}
self.last_use = now;
let Some(arc) = self.pool.upgrade() else {
debug!("Pool gone: {:?}", self.key);
return;
};
debug!("Return to pool: {:?}", self.key);
let mut pool = arc.lock().unwrap();
pool.add(self);
pool.purge(now);
}
pub fn is_tls(&self) -> bool {
self.transport.is_tls()
}
fn age(&self, now: Instant) -> Duration {
now.duration_since(self.last_use)
}
fn is_open(&mut self) -> bool {
self.transport.is_open()
}
}
#[derive(Clone, PartialEq, Eq)]
struct PoolKey(Arc<PoolKeyInner>);
impl PoolKey {
fn new(uri: &Uri, proxy: Option<&Proxy>) -> Self {
let inner = PoolKeyInner(
uri.scheme().expect("uri with scheme").clone(),
uri.authority().expect("uri with authority").clone(),
proxy.cloned(),
);
PoolKey(Arc::new(inner))
}
}
#[derive(PartialEq, Eq)]
struct PoolKeyInner(Scheme, Authority, Option<Proxy>);
#[derive(Debug)]
struct Pool {
lru: VecDeque<Connection>,
max_idle_connections: usize,
max_idle_connections_per_host: usize,
max_idle_age: Duration,
}
impl Pool {
fn new(config: &Config) -> Self {
Pool {
lru: VecDeque::new(),
max_idle_connections: config.max_idle_connections(),
max_idle_connections_per_host: config.max_idle_connections_per_host(),
max_idle_age: config.max_idle_age().into(),
}
}
fn purge(&mut self, now: Instant) {
while self.lru.len() > self.max_idle_connections || self.front_is_too_old(now) {
self.lru.pop_front();
}
self.update_position_per_host();
let max = self.max_idle_connections_per_host;
self.lru.retain(|c| c.position_per_host.unwrap() < max);
}
fn front_is_too_old(&self, now: Instant) -> bool {
self.lru.front().map(|c| c.age(now)) > Some(self.max_idle_age)
}
fn update_position_per_host(&mut self) {
for c in &mut self.lru {
c.position_per_host = None;
}
loop {
let maybe_uncounted = self
.lru
.iter()
.rev()
.find(|c| c.position_per_host.is_none());
let Some(uncounted) = maybe_uncounted else {
break; };
let key_to_count = uncounted.key.clone();
for (position, c) in self
.lru
.iter_mut()
.rev()
.filter(|c| c.key == key_to_count)
.enumerate()
{
c.position_per_host = Some(position);
}
}
}
fn add(&mut self, conn: Connection) {
self.lru.push_back(conn)
}
fn get(&mut self, key: &PoolKey, max_idle_age: Duration, now: Instant) -> Option<Connection> {
while let Some(i) = self.lru.iter().position(|c| c.key == *key) {
let mut conn = self.lru.remove(i).unwrap();
if !conn.is_open() {
continue;
}
if conn.age(now) >= max_idle_age {
continue;
}
return Some(conn);
}
None
}
}
impl fmt::Debug for ConnectionPool {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConnectionPool")
.field("connector", &self.connector)
.finish()
}
}
impl fmt::Debug for Connection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Connection")
.field("key", &self.key)
.field("conn", &self.transport)
.finish()
}
}
impl fmt::Debug for PoolKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PoolKey")
.field("scheme", &self.0.0)
.field("authority", &DebugAuthority(&self.0.1))
.field("proxy", &self.0.2)
.finish()
}
}
impl<'a, 'b> From<&'a ConnectionDetails<'b>> for PoolKey {
fn from(details: &'a ConnectionDetails) -> Self {
PoolKey::new(details.uri, details.config.proxy())
}
}
#[cfg(all(test, feature = "_test"))]
mod test {
use super::*;
#[test]
fn poolkey_new() {
PoolKey::new(&Uri::from_static("zzz://example.com"), None);
}
#[test]
fn no_reuse_with_unconsumed_input() {
use crate::test::init_test_log;
use crate::transport::set_handler;
init_test_log();
set_handler("/trailing", 200, &[("content-length", "5")], b"hellojunk");
let agent = crate::Agent::new_with_defaults();
let mut res = agent.get("https://example.test/trailing").call().unwrap();
assert_eq!(res.body_mut().read_to_string().unwrap(), "hello");
assert_eq!(agent.pool_count(), 0);
}
}
#[cfg(test)]
mod config_pooling_tests {
use super::*;
use crate::Agent;
use crate::transport::LazyBuffers;
use crate::unversioned::resolver::DefaultResolver;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug)]
struct CountingConnector(Arc<AtomicUsize>);
impl Connector for CountingConnector {
type Out = TestTransport;
fn connect(
&self,
_: &ConnectionDetails,
_: Option<()>,
) -> Result<Option<Self::Out>, Error> {
let id = self.0.fetch_add(1, Ordering::SeqCst) + 1;
Ok(Some(TestTransport {
id,
buffers: LazyBuffers::new(1024, 1024),
}))
}
}
#[derive(Debug)]
struct TestTransport {
id: usize,
buffers: LazyBuffers,
}
impl Transport for TestTransport {
fn buffers(&mut self) -> &mut dyn Buffers {
&mut self.buffers
}
fn transmit_output(&mut self, _: usize, _: NextTimeout) -> Result<(), Error> {
Ok(())
}
fn await_input(&mut self, _: NextTimeout) -> Result<bool, Error> {
let body = self.id.to_string();
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n{}",
body.len(),
body
);
self.buffers.input_append_buf()[..response.len()].copy_from_slice(response.as_bytes());
self.buffers.input_appended(response.len());
Ok(true)
}
fn is_open(&mut self) -> bool {
true
}
fn is_tls(&self) -> bool {
true
}
}
fn agent(config: Config) -> Agent {
Agent::with_parts(
config,
CountingConnector(Arc::new(AtomicUsize::new(0))),
DefaultResolver::default(),
)
}
fn request(agent: &Agent, config: Config) -> String {
let mut req = http::Request::get("https://127.0.0.1/").body(()).unwrap();
req.extensions_mut()
.insert(crate::config::RequestLevelConfig(config));
agent.run(req).unwrap().body_mut().read_to_string().unwrap()
}
fn check_override(config: Config) {
let base = Agent::config_builder().proxy(None).build();
check_configs(base, config);
}
fn check_configs(base: Config, config: Config) {
let agent = agent(base.clone());
assert_eq!(request(&agent, base.clone()), "1");
assert_eq!(
request(&agent, config.clone()),
"2",
"override must bypass existing connection"
);
assert_eq!(request(&agent, config), "3", "override must not enter pool");
assert_eq!(
request(&agent, base),
"1",
"Agent connection must remain reusable"
);
}
#[test]
fn connection_overrides_bypass_pool() {
check_override(Agent::config_builder().proxy(None).no_delay(false).build());
check_override(
Agent::config_builder()
.proxy(None)
.ip_family(crate::config::IpFamily::Ipv4Only)
.build(),
);
check_override(
Agent::config_builder()
.proxy(None)
.input_buffer_size(4096)
.build(),
);
check_override(
Agent::config_builder()
.proxy(None)
.output_buffer_size(4096)
.build(),
);
check_override(
Agent::config_builder()
.proxy(None)
.user_agent("custom")
.build(),
);
}
#[test]
fn request_settings_preserve_pooling() {
let base = Agent::config_builder().proxy(None).build();
let agent = agent(base.clone());
assert_eq!(request(&agent, base), "1");
let config = Agent::config_builder()
.proxy(None)
.https_only(true)
.http_status_as_error(false)
.timeout_global(Some(std::time::Duration::from_secs(10)))
.max_response_header_size(4096)
.build();
assert_eq!(request(&agent, config), "1");
}
#[test]
#[cfg(feature = "_tls")]
fn client_identity_isolation() {
use crate::tls::{Certificate, ClientCert, PrivateKey, TlsConfig};
let identity = |bytes: &'static [u8]| {
ClientCert::new_with_certs(
&[Certificate::from_der(bytes)],
PrivateKey::from_pem(
b"-----BEGIN PRIVATE KEY-----\nQQ==\n-----END PRIVATE KEY-----\n",
)
.unwrap(),
)
};
let config = |cert| {
Agent::config_builder()
.proxy(None)
.tls_config(TlsConfig::builder().client_cert(cert).build())
.build()
};
let a = config(Some(identity(b"A")));
let b = config(Some(identity(b"B")));
let none = config(None);
check_configs(a.clone(), b);
check_configs(a.clone(), none.clone());
check_configs(none, a.clone());
let agent = agent(a.clone());
assert_eq!(request(&agent, a.clone()), "1");
assert_eq!(
request(&agent, a),
"1",
"cloned credentials can reuse connections"
);
}
#[test]
#[cfg(feature = "_tls")]
fn tls_overrides_bypass_pool() {
use crate::tls::{RootCerts, TlsConfig, TlsProvider};
for tls in [
TlsConfig::builder().disable_verification(true).build(),
TlsConfig::builder().use_sni(false).build(),
TlsConfig::builder()
.provider(TlsProvider::NativeTls)
.build(),
TlsConfig::builder()
.root_certs(RootCerts::PlatformVerifier)
.build(),
] {
check_override(Agent::config_builder().proxy(None).tls_config(tls).build());
}
}
}