use anyhow::Result;
use tracing::{debug, info, warn};
use crate::config::Server;
use crate::pool::DeadpoolConnectionProvider;
async fn prewarm_single_pool(
provider: DeadpoolConnectionProvider,
server_name: String,
max_connections: usize,
) -> Result<usize> {
info!(
"Prewarming pool for '{}' with {} connections",
server_name, max_connections
);
let tasks: Vec<_> = (0..max_connections)
.map(|i| {
let provider = provider.clone();
let server_name = server_name.clone();
tokio::spawn(async move {
provider
.get_pooled_connection()
.await
.inspect(|_conn| {
debug!(
"Created connection {}/{} for '{}'",
i + 1,
max_connections,
server_name
);
})
.ok()
})
})
.collect();
let mut connections = Vec::with_capacity(max_connections);
for task in tasks {
if let Ok(Some(conn)) = task.await {
connections.push(conn);
}
}
let created = connections.len();
drop(connections);
info!(
"Pool '{}' ready: {}/{} connections created",
server_name, created, max_connections
);
Ok(created)
}
pub async fn prewarm_pools(
providers: &[DeadpoolConnectionProvider],
servers: &[Server],
) -> Result<()> {
info!("Prewarming all connection pools...");
let tasks: Vec<_> = servers
.iter()
.enumerate()
.map(|(i, server)| {
let provider = providers[i].clone();
let server_name = server.name.to_string();
let max_connections = server.max_connections.get();
tokio::spawn(prewarm_single_pool(provider, server_name, max_connections))
})
.collect();
let mut total_created = 0;
let mut total_expected = 0;
for (task, server) in tasks.into_iter().zip(servers.iter()) {
total_expected += server.max_connections.get();
match task.await {
Ok(Ok(created)) => total_created += created,
Ok(Err(e)) => warn!(
"Failed to prewarm pool for '{}': {}",
server.name.as_str(),
e
),
Err(e) => warn!(
"Prewarming task panicked for '{}': {}",
server.name.as_str(),
e
),
}
}
info!(
"Prewarming complete: {}/{} connections ready across all pools",
total_created, total_expected
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Server;
use tokio::net::TcpListener;
async fn spawn_mock_server_on_random_port() -> (u16, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let handle = tokio::spawn(async move {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
while let Ok((mut stream, _)) = listener.accept().await {
tokio::spawn(async move {
let _ = stream.write_all(b"200 Mock Server Ready\r\n").await;
let mut buffer = [0; 1024];
while let Ok(n) = stream.read(&mut buffer).await {
if n == 0 {
break;
}
if buffer[..n].starts_with(b"QUIT") {
let _ = stream.write_all(b"205 Goodbye\r\n").await;
break;
}
let _ = stream.write_all(b"200 OK\r\n").await;
}
});
}
});
(port, handle)
}
#[tokio::test]
async fn test_prewarm_pools_basic() {
let (port, _server) = spawn_mock_server_on_random_port().await;
let servers = vec![
Server::builder("127.0.0.1", crate::types::Port::try_new(port).unwrap())
.name("TestServer1")
.max_connections(crate::types::MaxConnections::try_new(2).unwrap())
.build()
.unwrap(),
];
let providers = servers
.iter()
.map(|s| {
crate::pool::DeadpoolConnectionProvider::new(
s.host.to_string(),
s.port.get(),
s.name.to_string(),
s.max_connections.get(),
s.username.clone(),
s.password.clone(),
)
})
.collect::<Vec<_>>();
let result = prewarm_pools(&providers, &servers).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_prewarm_pools_multiple_servers() {
let (port1, _server1) = spawn_mock_server_on_random_port().await;
let (port2, _server2) = spawn_mock_server_on_random_port().await;
let servers = vec![
Server::builder("127.0.0.1", crate::types::Port::try_new(port1).unwrap())
.name("Server1")
.max_connections(crate::types::MaxConnections::try_new(2).unwrap())
.build()
.unwrap(),
Server::builder("127.0.0.1", crate::types::Port::try_new(port2).unwrap())
.name("Server2")
.max_connections(crate::types::MaxConnections::try_new(1).unwrap())
.build()
.unwrap(),
];
let providers = servers
.iter()
.map(|s| {
crate::pool::DeadpoolConnectionProvider::new(
s.host.to_string(),
s.port.get(),
s.name.to_string(),
s.max_connections.get(),
s.username.clone(),
s.password.clone(),
)
})
.collect::<Vec<_>>();
let result = prewarm_pools(&providers, &servers).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_prewarm_pools_with_unreachable_server() {
let bad_port = 65500;
let servers = vec![
Server::builder("127.0.0.1", crate::types::Port::try_new(bad_port).unwrap())
.name("UnreachableServer")
.max_connections(crate::types::MaxConnections::try_new(1).unwrap())
.build()
.unwrap(),
];
let providers = servers
.iter()
.map(|s| {
crate::pool::DeadpoolConnectionProvider::new(
s.host.to_string(),
s.port.get(),
s.name.to_string(),
s.max_connections.get(),
s.username.clone(),
s.password.clone(),
)
})
.collect::<Vec<_>>();
let result = prewarm_pools(&providers, &servers).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_prewarm_pools_empty() {
let servers = vec![];
let providers = vec![];
let result = prewarm_pools(&providers, &servers).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_prewarm_single_pool_success() {
let (port, _server) = spawn_mock_server_on_random_port().await;
let provider = crate::pool::DeadpoolConnectionProvider::new(
"127.0.0.1".to_string(),
port,
"TestServer".to_string(),
3,
None,
None,
);
let result = prewarm_single_pool(provider, "TestServer".to_string(), 3).await;
assert!(result.is_ok());
let created = result.unwrap();
assert_eq!(created, 3);
}
#[tokio::test]
async fn test_prewarm_single_pool_single_connection() {
let (port, _server) = spawn_mock_server_on_random_port().await;
let provider = crate::pool::DeadpoolConnectionProvider::new(
"127.0.0.1".to_string(),
port,
"TestServer".to_string(),
1,
None,
None,
);
let result = prewarm_single_pool(provider, "TestServer".to_string(), 1).await;
assert!(result.is_ok());
let created = result.unwrap();
assert_eq!(created, 1);
}
#[tokio::test]
async fn test_prewarm_single_pool_zero_connections() {
let provider = crate::pool::DeadpoolConnectionProvider::new(
"127.0.0.1".to_string(),
1,
"TestServer".to_string(),
0, None,
None,
);
let result = prewarm_single_pool(provider, "TestServer".to_string(), 0).await;
assert!(result.is_ok());
let created = result.unwrap();
assert_eq!(created, 0);
}
#[tokio::test]
async fn test_prewarm_pools_mixed_success_failure() {
let (good_port, _server) = spawn_mock_server_on_random_port().await;
let bad_port = 65499;
let servers = vec![
Server::builder("127.0.0.1", crate::types::Port::try_new(good_port).unwrap())
.name("GoodServer")
.max_connections(crate::types::MaxConnections::try_new(1).unwrap())
.build()
.unwrap(),
Server::builder("127.0.0.1", crate::types::Port::try_new(bad_port).unwrap())
.name("BadServer")
.max_connections(crate::types::MaxConnections::try_new(1).unwrap())
.build()
.unwrap(),
];
let providers = servers
.iter()
.map(|s| {
crate::pool::DeadpoolConnectionProvider::new(
s.host.to_string(),
s.port.get(),
s.name.to_string(),
s.max_connections.get(),
s.username.clone(),
s.password.clone(),
)
})
.collect::<Vec<_>>();
let result = prewarm_pools(&providers, &servers).await;
assert!(result.is_ok());
}
}