use core::hash::Hash;
use core::sync::atomic::{AtomicUsize, Ordering};
use core::time::Duration;
#[cfg(feature = "std")]
use std::collections::{HashMap, VecDeque};
#[cfg(feature = "std")]
use std::sync::{Arc, RwLock, RwLockWriteGuard};
#[cfg(feature = "std")]
use std::time::Instant;
use crate::crypto::profiles::CryptoProvider;
use crate::crypto::{key::SigningKeyProvider, x509::CertificateSpec};
use crate::transport::client::GenericClient;
use crate::transport::error::{TransportError, TransportFailure};
use crate::transport::handshake::HandshakeKeyManager;
use crate::transport::protocols::{PersistentConnection, Protocol};
use crate::transport::MessageCollector;
use crate::transport::{TransportResult, X509ClientConfig};
#[cfg(feature = "aes-gcm")]
use crate::crypto::profiles::DefaultCryptoProvider;
#[cfg(not(feature = "x509"))]
use crate::transport::client::ClientBuilder;
#[cfg(feature = "x509")]
mod x509 {
pub use crate::crypto::x509::store::CertificateTrust;
pub use crate::crypto::x509::Certificate;
pub use crate::transport::handshake::HandshakeProtocolKind;
}
#[cfg(feature = "x509")]
use x509::*;
#[cfg(feature = "transport-policy")]
mod policy {
pub use crate::transport::policy::PolicyConf;
pub use crate::transport::MessageEmitter;
}
#[cfg(feature = "transport-policy")]
use policy::*;
pub trait ConnectionBuilder<P: Protocol>: Sized {
type Output;
fn with_timeout(self, timeout: Duration) -> Self;
#[cfg(feature = "x509")]
fn with_trust_store(self, store: Arc<dyn CertificateTrust>) -> Self;
#[cfg(feature = "x509")]
fn with_client_identity(self, cert: CertificateSpec, key: Arc<dyn SigningKeyProvider>) -> TransportResult<Self>;
fn build(self) -> Self::Output;
}
#[derive(Clone, Debug)]
pub struct PoolConfig {
pub idle_timeout: Option<Duration>,
pub max_connections: usize,
}
impl Default for PoolConfig {
fn default() -> Self {
Self { idle_timeout: None, max_connections: 64 }
}
}
#[cfg(feature = "x509")]
#[derive(Clone)]
struct ClientIdentity<C: CryptoProvider = DefaultCryptoProvider> {
certificate: Arc<Certificate>,
key: Arc<HandshakeKeyManager<C>>,
}
#[cfg(feature = "x509")]
#[derive(Clone, Default)]
struct PoolTlsConfig<C: CryptoProvider = DefaultCryptoProvider> {
trust_store: Option<Arc<dyn CertificateTrust>>,
client_identity: Option<ClientIdentity<C>>,
server_certificate_chain: Option<Arc<[Certificate]>>,
handshake_protocol: Option<HandshakeProtocolKind>,
}
#[cfg(feature = "x509")]
impl<C: CryptoProvider> PoolTlsConfig<C> {
fn set_trust_store(&mut self, store: Arc<dyn CertificateTrust>) {
self.trust_store = Some(store);
}
fn set_client_identity(&mut self, cert: Certificate, key: HandshakeKeyManager<C>) {
let certificate = Arc::new(cert);
let key = Arc::new(key);
self.client_identity = Some(ClientIdentity { certificate, key });
}
fn set_server_certificate_chain(&mut self, chain: Arc<[Certificate]>) {
self.server_certificate_chain = Some(chain);
}
fn set_handshake_protocol(&mut self, kind: HandshakeProtocolKind) {
self.handshake_protocol = Some(kind);
}
fn apply<Pro>(&self, transport: Pro::Transport) -> Pro::Transport
where
Pro: Protocol,
Pro::Transport: MessageEmitter + MessageCollector + PolicyConf + X509ClientConfig<CryptoProvider = C>,
{
let mut configured = transport;
if let Some(store) = &self.trust_store {
let store = Arc::clone(store);
configured = configured.with_trust_store(store);
}
if let Some(identity) = &self.client_identity {
let cert = Arc::clone(&identity.certificate);
let key = Arc::clone(&identity.key);
configured = configured.with_client_identity(cert, key);
}
if let Some(chain) = &self.server_certificate_chain {
let chain = Arc::clone(chain);
configured = configured.with_server_certificate_chain(chain);
}
if let Some(kind) = self.handshake_protocol {
configured = configured.with_handshake_protocol(kind);
}
configured
}
}
pub struct ConnectionPoolBuilder<P: Protocol, C: CryptoProvider = DefaultCryptoProvider> {
config: PoolConfig,
timeout: Option<Duration>,
#[cfg(feature = "x509")]
tls: PoolTlsConfig<C>,
_phantom: core::marker::PhantomData<(P, C)>,
}
impl<P: Protocol, C: CryptoProvider> Default for ConnectionPoolBuilder<P, C> {
fn default() -> Self {
Self {
config: PoolConfig::default(),
timeout: None,
#[cfg(feature = "x509")]
tls: PoolTlsConfig::default(),
_phantom: core::marker::PhantomData,
}
}
}
impl<P: Protocol, C: CryptoProvider> ConnectionPoolBuilder<P, C> {
pub fn with_config(mut self, config: PoolConfig) -> Self {
self.config = config;
self
}
#[cfg(feature = "x509")]
pub fn with_server_certificate_chain(mut self, chain: impl Into<Arc<[Certificate]>>) -> Self {
self.tls.set_server_certificate_chain(chain.into());
self
}
#[cfg(feature = "x509")]
pub fn with_handshake_protocol(mut self, kind: HandshakeProtocolKind) -> Self {
self.tls.set_handshake_protocol(kind);
self
}
}
#[cfg(feature = "std")]
impl<P: Protocol, C: CryptoProvider + Send + Sync + 'static> ConnectionBuilder<P> for ConnectionPoolBuilder<P, C> {
type Output = ConnectionPool<P, C>;
fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
#[cfg(feature = "x509")]
fn with_trust_store(mut self, store: Arc<dyn CertificateTrust>) -> Self {
self.tls.set_trust_store(store);
self
}
#[cfg(feature = "x509")]
fn with_client_identity(
mut self,
cert: CertificateSpec,
key: Arc<dyn SigningKeyProvider>,
) -> TransportResult<Self> {
let cert_converted = Certificate::try_from(cert)?;
let key_converted: HandshakeKeyManager<C> = HandshakeKeyManager::new(key);
self.tls.set_client_identity(cert_converted, key_converted);
Ok(self)
}
fn build(self) -> Self::Output {
ConnectionPool {
pools: Arc::new(RwLock::new(HashMap::new())),
config: self.config,
timeout: self.timeout,
total_connections: Arc::new(AtomicUsize::new(0)),
#[cfg(feature = "x509")]
tls: self.tls,
}
}
}
#[cfg(feature = "std")]
struct AvailableEntry<P: Protocol> {
client: GenericClient<P>,
last_used: Instant,
}
#[cfg(feature = "std")]
struct DestinationPool<P: Protocol> {
available: VecDeque<AvailableEntry<P>>,
in_use: usize,
}
#[cfg(feature = "std")]
pub struct ConnectionPool<P: Protocol, C: CryptoProvider = DefaultCryptoProvider> {
pools: Arc<RwLock<HashMap<P::Address, DestinationPool<P>>>>,
config: PoolConfig,
timeout: Option<Duration>,
total_connections: Arc<AtomicUsize>,
#[cfg(feature = "x509")]
tls: PoolTlsConfig<C>,
}
#[cfg(feature = "std")]
impl<P: Protocol, C: CryptoProvider> ConnectionPool<P, C> {
fn release_connection_count(&self) {
let _ = self
.total_connections
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| current.checked_sub(1));
}
}
#[cfg(feature = "std")]
impl<P: Protocol + Send + Sync, C: CryptoProvider + Send + Sync + 'static> ConnectionPool<P, C>
where
P::Address: Hash + Eq + Clone + Send + Sync,
P::Transport: Send + Sync,
{
pub fn builder() -> ConnectionPoolBuilder<P, C> {
ConnectionPoolBuilder::default()
}
fn wrap_client(self: &Arc<Self>, client: GenericClient<P>, addr: P::Address) -> PooledClient<P, C>
where
P: PersistentConnection,
{
PooledClient { client: Some(client), pool: Arc::clone(self), addr }
}
fn write_pools(&self) -> TransportResult<RwLockWriteGuard<'_, HashMap<P::Address, DestinationPool<P>>>> {
self.pools
.write()
.map_err(|_| TransportError::OperationFailed(TransportFailure::Busy))
}
#[cfg(not(feature = "x509"))]
fn apply_timeout_to_builder<B>(&self, builder: B) -> B
where
B: ConnectionBuilder<P>,
{
if let Some(timeout) = self.timeout {
builder.with_timeout(timeout)
} else {
builder
}
}
fn try_take_ready_client(self: &Arc<Self>, addr: &P::Address) -> TransportResult<Option<GenericClient<P>>>
where
P: PersistentConnection,
{
let mut pools = self.write_pools()?;
if let Some(dest_pool) = pools.get_mut(addr) {
self.prune_idle_locked(dest_pool, Instant::now());
while let Some(entry) = dest_pool.available.pop_front() {
if <P as PersistentConnection>::is_connected(entry.client.transport()) {
dest_pool.in_use += 1;
return Ok(Some(entry.client));
}
self.release_connection_count();
}
}
Ok(None)
}
fn reserve_slot(self: &Arc<Self>, addr: &P::Address) -> TransportResult<SlotGuard<P, C>> {
let reserved = self
.total_connections
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
if current >= self.config.max_connections {
None
} else {
Some(current + 1)
}
});
if reserved.is_err() {
return Err(TransportError::OperationFailed(TransportFailure::Busy));
}
let mut pools = self.write_pools()?;
let dest_pool = pools
.entry(addr.clone())
.or_insert_with(|| DestinationPool { available: VecDeque::new(), in_use: 0 });
self.prune_idle_locked(dest_pool, Instant::now());
dest_pool.in_use += 1;
let pool = Arc::clone(self);
let addr = addr.clone();
Ok(SlotGuard::new(pool, addr))
}
fn prune_idle_locked(&self, dest_pool: &mut DestinationPool<P>, now: Instant) {
if let Some(timeout) = self.config.idle_timeout {
while let Some(entry) = dest_pool.available.front() {
if now.duration_since(entry.last_used) >= timeout {
dest_pool.available.pop_front();
self.release_connection_count();
} else {
break;
}
}
}
}
#[cfg(not(feature = "x509"))]
pub async fn connect(self: &Arc<Self>, addr: P::Address) -> TransportResult<PooledClient<P, C>>
where
P: PersistentConnection + Send + Sync,
P::Transport: MessageEmitter + MessageCollector + PolicyConf + Send + Sync,
{
if let Some(client) = self.try_take_ready_client(&addr)? {
return Ok(self.wrap_client(client, addr));
}
let mut reservation = self.reserve_slot(&addr)?;
let builder = self.apply_timeout_to_builder(ClientBuilder::<P, C>::builder());
let builder = ConnectionBuilder::build(builder);
let client = builder.connect(addr.clone()).await?;
reservation.disarm();
Ok(self.wrap_client(client, addr))
}
#[cfg(feature = "x509")]
pub async fn connect(self: &Arc<Self>, addr: P::Address) -> TransportResult<PooledClient<P, C>>
where
P: PersistentConnection + Send + Sync,
P::Transport:
MessageEmitter + MessageCollector + PolicyConf + X509ClientConfig<CryptoProvider = C> + Send + Sync,
{
if let Some(client) = self.try_take_ready_client(&addr)? {
return Ok(self.wrap_client(client, addr));
}
let mut reservation = self.reserve_slot(&addr)?;
let stream = P::connect(addr.clone()).await.map_err(|e| e.into())?;
let mut transport = self.tls.apply::<P>(P::create_transport(stream));
if let Some(timeout) = self.timeout {
transport = transport.with_timeout(timeout);
}
let client = GenericClient::from_transport_with_addr(transport, addr.clone());
reservation.disarm();
Ok(self.wrap_client(client, addr))
}
pub fn try_acquire(self: &Arc<Self>, addr: &P::Address) -> TransportResult<Option<PooledClient<P, C>>>
where
P: PersistentConnection + Send + Sync,
P::Transport: MessageEmitter + MessageCollector + PolicyConf + Send + Sync,
{
let maybe_client = self.try_take_ready_client(addr)?;
Ok(maybe_client.map(|client| self.wrap_client(client, addr.clone())))
}
}
#[cfg(feature = "std")]
#[cfg(not(feature = "x509"))]
impl<P: Protocol + Send + Sync, C: CryptoProvider + Send + Sync + 'static> ConnectionPool<P, C>
where
P::Address: Hash + Eq + Clone + Send + Sync,
P::Transport: Send + Sync,
{
}
#[cfg(feature = "std")]
pub struct PooledClient<P: Protocol + PersistentConnection, C: CryptoProvider = DefaultCryptoProvider>
where
P::Address: Hash + Eq + Send + Sync,
{
client: Option<GenericClient<P>>,
pool: Arc<ConnectionPool<P, C>>,
addr: P::Address,
}
#[cfg(feature = "std")]
impl<P: Protocol + PersistentConnection, C: CryptoProvider> PooledClient<P, C>
where
P::Address: Hash + Eq + Send + Sync,
{
pub fn conn(&mut self) -> TransportResult<&mut GenericClient<P>> {
self.client
.as_mut()
.ok_or(TransportError::OperationFailed(TransportFailure::Busy))
}
}
#[cfg(feature = "std")]
impl<P: Protocol + PersistentConnection, C: CryptoProvider> Drop for PooledClient<P, C>
where
P::Address: Hash + Eq + Send + Sync,
{
fn drop(&mut self) {
let client = match self.client.take() {
Some(client) => client,
None => return,
};
let mut returned_to_pool = false;
let is_healthy = <P as PersistentConnection>::is_connected(client.transport());
if let Ok(mut pools) = self.pool.pools.write() {
if let Some(dest_pool) = pools.get_mut(&self.addr) {
dest_pool.in_use = dest_pool.in_use.saturating_sub(1);
if is_healthy {
dest_pool
.available
.push_back(AvailableEntry { client, last_used: Instant::now() });
returned_to_pool = true;
}
}
}
if !returned_to_pool {
self.pool.release_connection_count();
}
}
}
#[cfg(feature = "std")]
struct SlotGuard<P: Protocol, C: CryptoProvider = DefaultCryptoProvider>
where
P::Address: Hash + Eq + Clone + Send + Sync,
{
pool: Arc<ConnectionPool<P, C>>,
addr: P::Address,
active: bool,
}
#[cfg(feature = "std")]
impl<P: Protocol, C: CryptoProvider> SlotGuard<P, C>
where
P::Address: Hash + Eq + Clone + Send + Sync,
{
fn new(pool: Arc<ConnectionPool<P, C>>, addr: P::Address) -> Self {
Self { pool, addr, active: true }
}
fn disarm(&mut self) {
self.active = false;
}
}
#[cfg(feature = "std")]
impl<P: Protocol, C: CryptoProvider> Drop for SlotGuard<P, C>
where
P::Address: Hash + Eq + Clone + Send + Sync,
{
fn drop(&mut self) {
if !self.active {
return;
}
self.pool.release_connection_count();
let pools = self.pool.pools.write();
if let Ok(mut pools) = pools {
if let Some(dest_pool) = pools.get_mut(&self.addr) {
dest_pool.in_use = dest_pool.in_use.saturating_sub(1);
}
}
}
}
#[cfg(all(test, feature = "std"))]
mod tests {
use super::*;
#[test]
fn test_pool_config_default() {
let config = PoolConfig::default();
assert!(config.idle_timeout.is_none());
assert_eq!(config.max_connections, 64);
}
#[test]
fn test_pool_config_with_timeout() {
let config = PoolConfig { idle_timeout: Some(Duration::from_secs(30)), max_connections: 64 };
assert_eq!(config.idle_timeout, Some(Duration::from_secs(30)));
}
#[test]
fn test_pool_config_with_max_connections() {
let config = PoolConfig { idle_timeout: None, max_connections: 16 };
assert_eq!(config.max_connections, 16);
}
}