use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use magnetar_proto::ConnectionConfig;
use moonpool_core::{Providers, TaskProvider, TimeError, TimeProvider};
use parking_lot::Mutex;
use tokio::sync::Notify;
use crate::dns::DnsResolver;
use crate::driver::{DriverHandle, ReconnectContext, spawn_supervised as spawn_supervised_driver};
use crate::transport::Transport;
use crate::{ConnectionShared, EngineError, handshake_plain, make_shared_with_providers};
#[derive(Clone)]
pub(crate) struct ConnectionFactory<P: Providers> {
pub(crate) addr: String,
pub(crate) bootstrap_config: ConnectionConfig,
pub(crate) operation_retry: Arc<Mutex<magnetar_proto::OperationRetryConfig>>,
pub(crate) providers: P,
pub(crate) service_url_provider: Option<Arc<dyn magnetar_proto::ServiceUrlProvider>>,
pub(crate) dns_resolver: Option<Arc<dyn DnsResolver>>,
}
impl<P: Providers> std::fmt::Debug for ConnectionFactory<P> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConnectionFactory")
.field("addr", &self.addr)
.field(
"has_service_url_provider",
&self.service_url_provider.is_some(),
)
.field("has_dns_resolver", &self.dns_resolver.is_some())
.finish_non_exhaustive()
}
}
type PoolKey = (String, String, usize);
type DialOutcome = Result<Arc<ConnectionShared>, EngineError>;
struct PendingDial {
notify: Arc<Notify>,
result: Arc<Mutex<Option<Arc<DialOutcome>>>>,
cancel: Arc<Notify>,
completed: Arc<Notify>,
is_complete: Arc<AtomicBool>,
}
impl PendingDial {
fn new() -> Self {
Self {
notify: Arc::new(Notify::new()),
result: Arc::new(Mutex::new(None)),
cancel: Arc::new(Notify::new()),
completed: Arc::new(Notify::new()),
is_complete: Arc::new(AtomicBool::new(false)),
}
}
fn handles(&self) -> Self {
Self {
notify: self.notify.clone(),
result: self.result.clone(),
cancel: self.cancel.clone(),
completed: self.completed.clone(),
is_complete: self.is_complete.clone(),
}
}
fn mark_complete(&self) {
self.is_complete.store(true, Ordering::Release);
self.completed.notify_waiters();
}
async fn cancel_and_wait(&self) {
let completed = self.completed.notified();
let mut completed = std::pin::pin!(completed);
completed.as_mut().enable();
self.cancel.notify_one();
if !self.is_complete.load(Ordering::Acquire) {
completed.await;
}
}
}
struct PendingCompletion(PendingDial);
impl Drop for PendingCompletion {
fn drop(&mut self) {
{
let mut result = self.0.result.lock();
if result.is_none() {
*result = Some(Arc::new(Err(EngineError::PeerClosed)));
}
}
self.0.notify.notify_waiters();
self.0.mark_complete();
}
}
enum EntryState {
Pending(PendingDial),
Ready {
shared: Arc<ConnectionShared>,
driver: Mutex<Option<DriverHandle>>,
},
}
pub(crate) struct ProxyConnectionPool<P: Providers> {
factory: ConnectionFactory<P>,
closed: AtomicBool,
entries: Mutex<HashMap<PoolKey, Arc<EntryState>>>,
}
impl<P: Providers> ProxyConnectionPool<P> {
pub(crate) fn set_operation_retry_config(&self, config: magnetar_proto::OperationRetryConfig) {
*self.factory.operation_retry.lock() = config;
}
}
impl<P: Providers> std::fmt::Debug for ProxyConnectionPool<P> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let snapshot: Vec<_> = self.entries.lock().keys().cloned().collect();
f.debug_struct("ProxyConnectionPool")
.field("factory", &self.factory)
.field("closed", &self.closed.load(Ordering::Acquire))
.field("entries", &snapshot)
.finish()
}
}
impl<P: Providers> ProxyConnectionPool<P> {
pub(crate) fn new(factory: ConnectionFactory<P>) -> Arc<Self> {
Arc::new(Self {
factory,
closed: AtomicBool::new(false),
entries: Mutex::new(HashMap::new()),
})
}
#[allow(dead_code)] pub(crate) fn bootstrap_addr(&self) -> &str {
&self.factory.addr
}
#[cfg(test)]
#[must_use]
pub(crate) fn len(&self) -> usize {
self.entries.lock().len()
}
}
impl<P: Providers + Send + Sync> ProxyConnectionPool<P> {
pub(crate) async fn close(&self) {
self.closed.store(true, Ordering::Release);
let drained: Vec<Arc<EntryState>> = self.entries.lock().drain().map(|(_, v)| v).collect();
for state in drained {
match &*state {
EntryState::Ready { shared, driver } => {
{
let mut conn = shared.inner.lock();
conn.close();
}
shared.driver_waker.notify_one();
let handle = driver.lock().take();
if let Some(handle) = handle {
let _ = handle.join().await;
}
}
EntryState::Pending(pending) => {
{
let mut result = pending.result.lock();
if result.is_none() {
*result = Some(Arc::new(Err(EngineError::PeerClosed)));
}
}
pending.notify.notify_waiters();
pending.cancel_and_wait().await;
}
}
}
}
}
pub(crate) async fn get_or_open<P>(
pool: Arc<ProxyConnectionPool<P>>,
logical: &str,
physical: &str,
proxy_to_broker_url: Option<String>,
index: usize,
) -> Result<Arc<ConnectionShared>, EngineError>
where
P: Providers + Send + Sync,
{
if pool.closed.load(Ordering::Acquire) {
return Err(pool_closed_error());
}
let key: PoolKey = (logical.to_owned(), physical.to_owned(), index);
let pending = {
let mut entries = pool.entries.lock();
if pool.closed.load(Ordering::Acquire) {
return Err(pool_closed_error());
}
if let Some(state) = entries.get(&key).cloned() {
match &*state {
EntryState::Ready { shared, .. } => return Ok(shared.clone()),
EntryState::Pending(pending) => pending.handles(),
}
} else {
let pending = PendingDial::new();
let handles = pending.handles();
let entry = Arc::new(EntryState::Pending(pending));
let clobbered = entries.insert(key.clone(), entry.clone());
debug_assert!(
clobbered.is_none(),
"pool entry insert clobbered a live entry — double registration for one key"
);
drop(entries);
spawn_dial(
pool.clone(),
physical.to_owned(),
proxy_to_broker_url,
key.clone(),
entry,
handles.handles(),
);
handles
}
};
loop {
let notified = pending.notify.notified();
let mut notified = std::pin::pin!(notified);
notified.as_mut().enable();
if let Some(outcome) = pending.result.lock().as_ref().map(Arc::clone) {
return match &*outcome {
Ok(shared) => Ok(shared.clone()),
Err(err) => Err(clone_engine_error(err)),
};
}
notified.await;
}
}
pub(crate) async fn get_or_open_bootstrap_sibling<P>(
pool: Arc<ProxyConnectionPool<P>>,
index: usize,
) -> Result<Arc<ConnectionShared>, EngineError>
where
P: Providers + Send + Sync,
{
let authority = pool.factory.addr.clone();
let proxy = pool.factory.bootstrap_config.proxy_to_broker_url.clone();
get_or_open(pool, &authority, &authority, proxy, index).await
}
fn spawn_dial<P>(
pool: Arc<ProxyConnectionPool<P>>,
physical: String,
proxy_to_broker_url: Option<String>,
key: PoolKey,
expected_entry: Arc<EntryState>,
pending: PendingDial,
) where
P: Providers + Send + Sync,
{
let factory = pool.factory.clone();
let task = pool.factory.providers.task().clone();
let _detached = task.spawn_task("magnetar-moonpool-pool-dial", async move {
let _completion = PendingCompletion(pending.handles());
let time = factory.providers.time().clone();
let operation_timeout = factory.bootstrap_config.operation_timeout;
let build = time.timeout(
operation_timeout,
build_entry_async::<P>(&factory, &physical, proxy_to_broker_url),
);
let mut build = std::pin::pin!(build);
let cancelled = pending.cancel.notified();
let mut cancelled = std::pin::pin!(cancelled);
cancelled.as_mut().enable();
let outcome = moonpool_core::select! {
biased;
() = &mut cancelled => Err(EngineError::PeerClosed),
timed = &mut build => match timed {
Ok(outcome) => outcome,
Err(TimeError::Elapsed) => Err(EngineError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"pool dial to {physical} exceeded operation_timeout \
({operation_timeout:?})"
),
))),
Err(TimeError::Shutdown) => Err(EngineError::PeerClosed),
},
};
let mut orphaned_success = None;
let published = {
let mut map = pool.entries.lock();
let is_current_generation = map
.get(&key)
.is_some_and(|current| Arc::ptr_eq(current, &expected_entry));
let may_promote = is_current_generation && !pool.closed.load(Ordering::Acquire);
match outcome {
Ok((shared, driver)) if may_promote => {
let waiter_shared = shared.clone();
map.insert(
key,
Arc::new(EntryState::Ready {
shared,
driver: Mutex::new(Some(driver)),
}),
);
Ok(waiter_shared)
}
Ok(pair) => {
if is_current_generation {
map.remove(&key);
}
orphaned_success = Some(pair);
Err(EngineError::PeerClosed)
}
Err(err) if may_promote => {
map.remove(&key);
Err(clone_engine_error(&err))
}
Err(_) => {
if is_current_generation {
map.remove(&key);
}
Err(EngineError::PeerClosed)
}
}
};
{
let mut result = pending.result.lock();
if result.is_none() {
*result = Some(Arc::new(published));
}
}
pending.notify.notify_waiters();
if let Some((shared, driver)) = orphaned_success {
{
let mut conn = shared.inner.lock();
conn.close();
}
shared.driver_waker.notify_one();
let _ = driver.join().await;
}
});
}
async fn build_entry_async<P: Providers>(
factory: &ConnectionFactory<P>,
physical: &str,
proxy_to_broker_url: Option<String>,
) -> Result<(Arc<ConnectionShared>, DriverHandle), EngineError> {
let mut cfg = factory.bootstrap_config.clone();
cfg.proxy_to_broker_url = proxy_to_broker_url;
let connect_timeout = cfg.connect_timeout;
let operation_timeout = cfg.operation_timeout;
let mut transport = crate::dial_with_retry::<P, _, _>(
factory.providers.time(),
cfg.connect_max_retries,
operation_timeout,
|| {
Transport::<P>::connect_with_resolver(
factory.providers.network(),
physical,
factory.dns_resolver.as_deref(),
factory.providers.time(),
connect_timeout,
)
},
)
.await?;
let shared = make_shared_with_providers::<P>(&factory.providers, cfg);
shared
.inner
.lock()
.set_operation_retry_config(factory.operation_retry.lock().clone());
handshake_plain::<P>(
&shared,
&mut transport,
factory.providers.time(),
None,
physical,
false,
)
.await?;
let ctx = ReconnectContext {
host_port: physical.to_owned(),
service_url_provider: factory.service_url_provider.clone(),
dns_resolver: factory.dns_resolver.clone(),
};
let driver =
spawn_supervised_driver::<P>(shared.clone(), transport, ctx, factory.providers.clone());
Ok((shared, driver))
}
fn pool_closed_error() -> EngineError {
EngineError::Config("connection pool is closed".to_owned())
}
fn clone_engine_error(err: &EngineError) -> EngineError {
match err {
EngineError::Io(io) => EngineError::Io(std::io::Error::new(io.kind(), io.to_string())),
EngineError::PeerClosed => EngineError::PeerClosed,
EngineError::Config(msg) => EngineError::Config(msg.clone()),
EngineError::Protocol(p) => EngineError::Config(format!("protocol error: {p}")),
EngineError::HandshakeFailed(msg) => EngineError::HandshakeFailed(msg.clone()),
EngineError::Tls(t) => EngineError::Config(format!("tls error: {t}")),
EngineError::MemoryLimitExceeded {
current,
limit,
requested,
} => EngineError::MemoryLimitExceeded {
current: *current,
limit: *limit,
requested: *requested,
},
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use moonpool_core::TokioProviders;
use super::*;
fn dummy_factory() -> ConnectionFactory<TokioProviders> {
ConnectionFactory {
addr: "broker.example.com:6650".to_owned(),
bootstrap_config: ConnectionConfig {
operation_timeout: Duration::from_secs(30),
..ConnectionConfig::default()
},
operation_retry: Arc::new(Mutex::new(magnetar_proto::OperationRetryConfig::default())),
providers: TokioProviders::new(),
service_url_provider: None,
dns_resolver: None,
}
}
#[tokio::test(flavor = "current_thread")]
async fn fresh_pool_is_empty() {
let pool = ProxyConnectionPool::new(dummy_factory());
assert_eq!(pool.len(), 0);
let pending = PendingDial::new();
let result = pending.result.clone();
let worker = pending.handles();
let worker_task = tokio::spawn(async move {
worker.cancel.notified().await;
worker.mark_complete();
});
pool.entries.lock().insert(
("logical".to_owned(), "physical".to_owned(), 0),
Arc::new(EntryState::Pending(pending)),
);
pool.close().await;
worker_task.await.expect("pending worker exits on cancel");
let outcome = result
.lock()
.as_ref()
.cloned()
.expect("pool close must resolve pending dials");
assert!(matches!(&*outcome, Err(EngineError::PeerClosed)));
}
#[test]
fn debug_includes_pool_state() {
let pool = ProxyConnectionPool::new(dummy_factory());
let s = format!("{pool:?}");
assert!(s.contains("ProxyConnectionPool"));
assert!(s.contains("entries"));
}
}