use mysql_async::Pool;
use rand::RngCore;
use std::collections::HashMap;
use std::fmt;
use std::sync::Mutex;
use thiserror::Error;
use zeroize::Zeroizing;
use crate::config::MySqlConnection;
pub const MAX_POOLS: usize = 16;
#[derive(Clone)]
pub struct CredentialGeneration([u8; 16]);
impl fmt::Debug for CredentialGeneration {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("CredentialGeneration(<redacted>)")
}
}
impl fmt::Display for CredentialGeneration {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("<redacted>")
}
}
impl CredentialGeneration {
pub fn derive(password: &str) -> Self {
use hmac::{KeyInit, Mac};
let key = process_key();
let mut mac = <hmac::Hmac<sha2::Sha256> as KeyInit>::new_from_slice(key)
.expect("HMAC accepts any key length");
mac.update(b"sequel-mcp/credential-generation/v1\n");
mac.update(password.as_bytes());
let tag = mac.finalize().into_bytes();
let mut out = [0u8; 16];
out.copy_from_slice(&tag[..16]);
CredentialGeneration(out)
}
pub fn eq_material(&self, other: &CredentialGeneration) -> bool {
let mut diff = 0u8;
for (a, b) in self.0.iter().zip(other.0.iter()) {
diff |= a ^ b;
}
diff == 0
}
pub fn key_fragment(&self) -> String {
self.0.iter().map(|b| format!("{b:02x}")).collect()
}
}
fn process_key() -> &'static [u8; 32] {
use std::sync::OnceLock;
static KEY: OnceLock<[u8; 32]> = OnceLock::new();
KEY.get_or_init(|| {
let mut k = [0u8; 32];
rand::rng().fill_bytes(&mut k);
k
})
}
#[derive(Debug, Error)]
pub enum PoolManagerError {
#[error("mysql pool initialization failed: {0}")]
Init(String),
#[error("pool cache is full ({0} live pools); refusing to open another")]
CacheFull(usize),
}
struct PoolEntry {
pool: Pool,
generation: CredentialGeneration,
tunnel_generation: Option<u64>,
}
pub struct PoolManager {
pools: Mutex<HashMap<String, PoolEntry>>,
}
impl PoolManager {
pub fn new() -> Self {
Self {
pools: Mutex::new(HashMap::new()),
}
}
fn config_key(
conn: &MySqlConnection,
database: Option<&str>,
revision: u64,
host_override: Option<&str>,
port_override: Option<u16>,
tunnel_generation: Option<u64>,
) -> String {
format!(
"{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}",
conn.name,
host_override.unwrap_or(&conn.host),
port_override.unwrap_or(conn.port),
conn.user,
database.or(conn.database.as_deref()).unwrap_or(""),
conn.ssl,
conn.ssl_server_name.as_deref().unwrap_or(""),
revision,
tunnel_generation.map(|g| g.to_string()).unwrap_or_default(),
)
}
#[allow(clippy::too_many_arguments)]
pub async fn verified_pool(
&self,
conn: &MySqlConnection,
password: &Zeroizing<String>,
database: Option<&str>,
revision: u64,
host_override: Option<&str>,
port_override: Option<u16>,
tunnel_generation: Option<u64>,
) -> Result<Pool, PoolManagerError> {
crate::app::test_mode::check_mysql_endpoint(
host_override.unwrap_or(&conn.host),
port_override.unwrap_or(conn.port),
)
.map_err(PoolManagerError::Init)?;
let generation = CredentialGeneration::derive(password);
let key = Self::config_key(
conn,
database,
revision,
host_override,
port_override,
tunnel_generation,
);
{
let pools = self.pools.lock().unwrap();
if let Some(entry) = pools.get(&key)
&& entry.generation.eq_material(&generation)
{
return Ok(entry.pool.clone());
}
}
let opts: mysql_async::Opts =
super::mysql::build_opts(conn, password, database, host_override, port_override).into();
let candidate = Pool::new(opts);
{
use mysql_async::prelude::Queryable;
let probe = async {
let mut conn = candidate
.get_conn()
.await
.map_err(|e| PoolManagerError::Init(e.to_string()))?;
conn.query_drop("SELECT 1")
.await
.map_err(|e| PoolManagerError::Init(format!("health query failed: {e}")))?;
Ok::<(), PoolManagerError>(())
};
tokio::time::timeout(super::mysql::CONNECT_TIMEOUT, probe)
.await
.map_err(|_| {
PoolManagerError::Init(format!(
"connect timeout after {}s",
super::mysql::CONNECT_TIMEOUT.as_secs()
))
})??;
}
let mut pools = self.pools.lock().unwrap();
if let Some(entry) = pools.get(&key) {
if entry.generation.eq_material(&generation) {
return Ok(entry.pool.clone());
}
let old = pools.remove(&key).expect("entry present");
let _close = tokio::task::spawn(old.pool.disconnect());
}
if pools.len() >= MAX_POOLS {
return Err(PoolManagerError::CacheFull(pools.len()));
}
pools.insert(
key,
PoolEntry {
pool: candidate.clone(),
generation,
tunnel_generation,
},
);
Ok(candidate)
}
pub fn pool_count(&self) -> usize {
self.pools.lock().unwrap().len()
}
pub fn invalidate_all(&self) {
let mut pools = self.pools.lock().unwrap();
for (_, entry) in pools.drain() {
let _close = tokio::task::spawn(entry.pool.disconnect());
}
}
}
pub fn evict_by_generation(generation: u64) -> usize {
let mgr = super::mysql::pool_manager();
let mut pools = mgr.pools.lock().unwrap();
let victims: Vec<String> = pools
.iter()
.filter(|(_, e)| e.tunnel_generation == Some(generation))
.map(|(k, _)| k.clone())
.collect();
let n = victims.len();
for k in victims {
if let Some(entry) = pools.remove(&k) {
let _close = tokio::task::spawn(entry.pool.disconnect());
}
}
n
}
impl Default for PoolManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn conn(host: &str) -> MySqlConnection {
MySqlConnection {
name: "t".into(),
host: host.into(),
port: 3306,
user: "u".into(),
..MySqlConnection::default()
}
}
#[test]
fn generation_is_deterministic_per_process() {
let a = CredentialGeneration::derive("pw");
let b = CredentialGeneration::derive("pw");
assert!(a.eq_material(&b));
assert!(!a.eq_material(&CredentialGeneration::derive("other")));
}
#[test]
fn generation_debug_and_display_are_redacted() {
let g = CredentialGeneration::derive("secret-value");
assert_eq!(format!("{g:?}"), "CredentialGeneration(<redacted>)");
assert_eq!(format!("{g}"), "<redacted>");
let s = format!("{g:?}{g}");
assert!(!s.contains("secret-value"));
}
#[tokio::test]
async fn failed_handshake_never_populates_cache() {
let mgr = PoolManager::new();
let c = conn("127.0.0.1");
let pw = Zeroizing::new("x".to_string());
let err = mgr
.verified_pool(&c, &pw, None, 1, None, Some(1), None)
.await
.unwrap_err();
assert!(matches!(err, PoolManagerError::Init(_)), "{err:?}");
assert_eq!(mgr.pool_count(), 0);
}
#[tokio::test]
async fn concurrent_first_access_initializes_once() {
use std::sync::Arc;
let mgr = Arc::new(PoolManager::new());
let c = conn("127.0.0.1");
let pw = Zeroizing::new("x".to_string());
let mut handles = Vec::new();
for _ in 0..4 {
let m = mgr.clone();
let cc = c.clone();
let p = pw.clone();
handles.push(tokio::task::spawn(async move {
m.verified_pool(&cc, &p, None, 1, None, Some(2), None).await
}));
}
for h in handles {
assert!(h.await.unwrap().is_err());
}
assert_eq!(mgr.pool_count(), 0, "failures must not be cached");
}
#[test]
fn config_key_separates_transport_settings() {
let c1 = conn("db1.example.invalid");
let mut c2 = conn("db2.example.invalid");
c2.ssl = true;
let k1 = PoolManager::config_key(&c1, None, 1, None, None, None);
let k2 = PoolManager::config_key(&c2, None, 1, None, None, None);
let k3 = PoolManager::config_key(&c1, None, 2, None, None, None);
assert_ne!(k1, k2);
assert_ne!(k1, k3);
let k4 = PoolManager::config_key(&c1, None, 1, None, None, Some(7));
let k5 = PoolManager::config_key(&c1, None, 1, None, None, Some(8));
assert_ne!(k1, k4);
assert_ne!(k4, k5);
}
}