use core::fmt::Debug;
use std::time::Duration;
use super::conn::{ConnectorService, EstablishedClientConnection};
use rama_core::error::{BoxError, ErrorContext};
use rama_core::extensions::{Extension, ExtensionsRef};
use rama_core::telemetry::tracing::trace;
use rama_core::{Layer, Service};
use rama_utils::macros::generate_set_and_with;
use tokio::sync::OwnedSemaphorePermit;
use tokio::time::timeout;
#[cfg(feature = "opentelemetry")]
#[cfg_attr(docsrs, doc(cfg(feature = "opentelemetry")))]
pub mod metrics;
mod exclusive;
#[doc(inline)]
pub use exclusive::{LeasedConnection, LruDropPool, ReuseStrategy};
pub mod multiplex;
#[doc(inline)]
pub use multiplex::{MultiplexPool, MultiplexedConnection, MuxSelection};
pub trait Pool<C, ID>: Send + Sync + 'static {
type Connection: Send + ExtensionsRef;
type CreatePermit: Send;
fn get_conn(
&self,
id: &ID,
) -> impl Future<
Output = Result<ConnectionResult<Self::Connection, Self::CreatePermit>, BoxError>,
> + Send;
fn create(
&self,
id: ID,
conn: C,
create_permit: Self::CreatePermit,
) -> impl Future<Output = Self::Connection> + Send;
}
pub enum ConnectionResult<C, P> {
Connection(C),
CreatePermit(P),
}
impl<C: Debug, P: Debug> Debug for ConnectionResult<C, P> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Connection(arg0) => f.debug_tuple("Connection").field(arg0).finish(),
Self::CreatePermit(arg0) => f.debug_tuple("CreatePermit").field(arg0).finish(),
}
}
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct NoPool;
impl<C, ID> Pool<C, ID> for NoPool
where
C: Send + ExtensionsRef + 'static,
ID: Clone + Send + Sync + PartialEq + 'static,
{
type Connection = C;
type CreatePermit = ();
async fn get_conn(
&self,
_id: &ID,
) -> Result<ConnectionResult<Self::Connection, Self::CreatePermit>, BoxError> {
Ok(ConnectionResult::CreatePermit(()))
}
async fn create(&self, _id: ID, conn: C, _permit: Self::CreatePermit) -> Self::Connection {
conn
}
}
#[expect(dead_code)]
#[derive(Debug)]
pub struct ActiveSlot(OwnedSemaphorePermit);
#[expect(dead_code)]
#[derive(Debug)]
pub struct PoolSlot(OwnedSemaphorePermit);
pub trait ReqToConnID<Input: ExtensionsRef>: Sized + Clone + Send + Sync + 'static {
type ID: ConnID;
fn id(&self, input: &Input) -> Result<Self::ID, BoxError>;
}
pub trait ConnID: Send + Sync + PartialEq + Clone + Debug + 'static {
#[cfg(feature = "opentelemetry")]
fn attributes(&self) -> impl Iterator<Item = rama_core::telemetry::opentelemetry::KeyValue> {
core::iter::empty()
}
}
impl<Input, ID, F> ReqToConnID<Input> for F
where
F: Fn(&Input) -> Result<ID, BoxError> + Clone + Send + Sync + 'static,
ID: ConnID,
Input: ExtensionsRef,
{
type ID = ID;
fn id(&self, request: &Input) -> Result<Self::ID, BoxError> {
self(request)
}
}
pub struct PooledConnector<S, P, R> {
inner: S,
pool: P,
req_to_conn_id: R,
wait_for_pool_timeout: Option<Duration>,
}
impl<S, P, R> PooledConnector<S, P, R> {
pub fn new(inner: S, pool: P, req_to_conn_id: R) -> Self {
Self {
inner,
pool,
req_to_conn_id,
wait_for_pool_timeout: None,
}
}
generate_set_and_with!(
pub fn wait_for_pool_timeout(mut self, timeout: Option<Duration>) -> Self {
self.wait_for_pool_timeout = timeout;
self
}
);
}
impl<Input, S, P, R> Service<Input> for PooledConnector<S, P, R>
where
S: ConnectorService<Input>,
Input: Send + ExtensionsRef + 'static,
P: Pool<S::Connection, R::ID> + Extension,
R: ReqToConnID<Input>,
{
type Output = EstablishedClientConnection<P::Connection, Input>;
type Error = BoxError;
async fn serve(&self, input: Input) -> Result<Self::Output, Self::Error> {
let conn_id = self.req_to_conn_id.id(&input)?;
let pool = if let Some(pool) = input.extensions().get_ref::<P>() {
trace!("pooled connector: using pool from ctx");
pool
} else {
trace!("pooled connector: using pool from connector");
&self.pool
};
let pool_result = if let Some(duration) = self.wait_for_pool_timeout {
timeout(duration, pool.get_conn(&conn_id))
.await
.inspect_err(|err|{
trace!(%err, "pooled connector: timeout triggered while waiting for a connection (/w conn id: {conn_id:?}) from pool");
})?
} else {
pool.get_conn(&conn_id).await
};
match pool_result? {
ConnectionResult::Connection(conn) => {
trace!(
"pooled connector: got connection (w/ conn id: {conn_id:?}) from pool (running health checks now)"
);
Ok(EstablishedClientConnection { conn, input })
}
ConnectionResult::CreatePermit(permit) => {
trace!(
"pooled connector: no connection (w/ conn id: {conn_id:?}) found, received permit to create a new one"
);
let EstablishedClientConnection { input, conn } =
self.inner.connect(input).await.into_box_error()?;
trace!(
"pooled connector: returning new pooled connection (w/ conn id: {conn_id:?}"
);
let pool = input.extensions().get_ref::<P>().unwrap_or(&self.pool);
let conn = pool.create(conn_id, conn, permit).await;
Ok(EstablishedClientConnection { input, conn })
}
}
}
}
pub struct PooledConnectorLayer<P, R> {
pool: P,
req_to_conn_id: R,
wait_for_pool_timeout: Option<Duration>,
}
impl<P, R> PooledConnectorLayer<P, R> {
pub fn new(pool: P, req_to_conn_id: R) -> Self {
Self {
pool,
req_to_conn_id,
wait_for_pool_timeout: None,
}
}
generate_set_and_with!(
pub fn wait_for_pool_timeout(mut self, timeout: Option<Duration>) -> Self {
self.wait_for_pool_timeout = timeout;
self
}
);
}
impl<S, P: Clone, R: Clone> Layer<S> for PooledConnectorLayer<P, R> {
type Service = PooledConnector<S, P, R>;
fn layer(&self, inner: S) -> Self::Service {
PooledConnector::new(inner, self.pool.clone(), self.req_to_conn_id.clone())
.maybe_with_wait_for_pool_timeout(self.wait_for_pool_timeout)
}
fn into_layer(self, inner: S) -> Self::Service {
PooledConnector::new(inner, self.pool, self.req_to_conn_id)
.maybe_with_wait_for_pool_timeout(self.wait_for_pool_timeout)
}
}