1use mysql_async::Pool;
11use rand::RngCore;
12use std::collections::HashMap;
13use std::fmt;
14use std::sync::Mutex;
15use thiserror::Error;
16use zeroize::Zeroizing;
17
18use crate::config::MySqlConnection;
19
20pub const MAX_POOLS: usize = 16;
22
23#[derive(Clone)]
27pub struct CredentialGeneration([u8; 16]);
28
29impl fmt::Debug for CredentialGeneration {
30 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
31 f.write_str("CredentialGeneration(<redacted>)")
32 }
33}
34
35impl fmt::Display for CredentialGeneration {
36 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
37 f.write_str("<redacted>")
38 }
39}
40
41impl CredentialGeneration {
42 pub fn derive(password: &str) -> Self {
46 use hmac::{KeyInit, Mac};
48 let key = process_key();
49 let mut mac = <hmac::Hmac<sha2::Sha256> as KeyInit>::new_from_slice(key)
50 .expect("HMAC accepts any key length");
51 mac.update(b"sequel-mcp/credential-generation/v1\n");
52 mac.update(password.as_bytes());
53 let tag = mac.finalize().into_bytes();
54 let mut out = [0u8; 16];
55 out.copy_from_slice(&tag[..16]);
56 CredentialGeneration(out)
57 }
58
59 pub fn eq_material(&self, other: &CredentialGeneration) -> bool {
60 let mut diff = 0u8;
61 for (a, b) in self.0.iter().zip(other.0.iter()) {
62 diff |= a ^ b;
63 }
64 diff == 0
65 }
66
67 pub fn key_fragment(&self) -> String {
71 self.0.iter().map(|b| format!("{b:02x}")).collect()
72 }
73}
74
75fn process_key() -> &'static [u8; 32] {
76 use std::sync::OnceLock;
77 static KEY: OnceLock<[u8; 32]> = OnceLock::new();
78 KEY.get_or_init(|| {
79 let mut k = [0u8; 32];
80 rand::rng().fill_bytes(&mut k);
81 k
82 })
83}
84
85#[derive(Debug, Error)]
86pub enum PoolManagerError {
87 #[error("mysql pool initialization failed: {0}")]
88 Init(String),
89 #[error("pool cache is full ({0} live pools); refusing to open another")]
90 CacheFull(usize),
91}
92
93struct PoolEntry {
94 pool: Pool,
95 generation: CredentialGeneration,
96 tunnel_generation: Option<u64>,
100}
101
102pub struct PoolManager {
105 pools: Mutex<HashMap<String, PoolEntry>>,
106}
107
108impl PoolManager {
109 pub fn new() -> Self {
110 Self {
111 pools: Mutex::new(HashMap::new()),
112 }
113 }
114
115 fn config_key(
119 conn: &MySqlConnection,
120 database: Option<&str>,
121 revision: u64,
122 host_override: Option<&str>,
123 port_override: Option<u16>,
124 tunnel_generation: Option<u64>,
125 ) -> String {
126 format!(
127 "{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}",
128 conn.name,
129 host_override.unwrap_or(&conn.host),
130 port_override.unwrap_or(conn.port),
131 conn.user,
132 database.or(conn.database.as_deref()).unwrap_or(""),
133 conn.ssl,
134 conn.ssl_server_name.as_deref().unwrap_or(""),
135 revision,
136 tunnel_generation.map(|g| g.to_string()).unwrap_or_default(),
137 )
138 }
139
140 #[allow(clippy::too_many_arguments)]
145 pub async fn verified_pool(
146 &self,
147 conn: &MySqlConnection,
148 password: &Zeroizing<String>,
149 database: Option<&str>,
150 revision: u64,
151 host_override: Option<&str>,
152 port_override: Option<u16>,
153 tunnel_generation: Option<u64>,
154 ) -> Result<Pool, PoolManagerError> {
155 crate::app::test_mode::check_mysql_endpoint(
160 host_override.unwrap_or(&conn.host),
161 port_override.unwrap_or(conn.port),
162 )
163 .map_err(PoolManagerError::Init)?;
164 let generation = CredentialGeneration::derive(password);
165 let key = Self::config_key(
166 conn,
167 database,
168 revision,
169 host_override,
170 port_override,
171 tunnel_generation,
172 );
173
174 {
176 let pools = self.pools.lock().unwrap();
177 if let Some(entry) = pools.get(&key)
178 && entry.generation.eq_material(&generation)
179 {
180 return Ok(entry.pool.clone());
181 }
182 }
183
184 let opts: mysql_async::Opts =
189 super::mysql::build_opts(conn, password, database, host_override, port_override).into();
190 let candidate = Pool::new(opts);
191 {
192 use mysql_async::prelude::Queryable;
193 let probe = async {
197 let mut conn = candidate
198 .get_conn()
199 .await
200 .map_err(|e| PoolManagerError::Init(e.to_string()))?;
201 conn.query_drop("SELECT 1")
202 .await
203 .map_err(|e| PoolManagerError::Init(format!("health query failed: {e}")))?;
204 Ok::<(), PoolManagerError>(())
205 };
206 tokio::time::timeout(super::mysql::CONNECT_TIMEOUT, probe)
207 .await
208 .map_err(|_| {
209 PoolManagerError::Init(format!(
210 "connect timeout after {}s",
211 super::mysql::CONNECT_TIMEOUT.as_secs()
212 ))
213 })??;
214 }
215
216 let mut pools = self.pools.lock().unwrap();
217 if let Some(entry) = pools.get(&key) {
218 if entry.generation.eq_material(&generation) {
219 return Ok(entry.pool.clone());
221 }
222 let old = pools.remove(&key).expect("entry present");
224 let _close = tokio::task::spawn(old.pool.disconnect());
225 }
226 if pools.len() >= MAX_POOLS {
227 return Err(PoolManagerError::CacheFull(pools.len()));
228 }
229 pools.insert(
230 key,
231 PoolEntry {
232 pool: candidate.clone(),
233 generation,
234 tunnel_generation,
235 },
236 );
237 Ok(candidate)
238 }
239
240 pub fn pool_count(&self) -> usize {
241 self.pools.lock().unwrap().len()
242 }
243
244 pub fn invalidate_all(&self) {
245 let mut pools = self.pools.lock().unwrap();
246 for (_, entry) in pools.drain() {
247 let _close = tokio::task::spawn(entry.pool.disconnect());
248 }
249 }
250}
251
252pub fn evict_by_generation(generation: u64) -> usize {
255 let mgr = super::mysql::pool_manager();
256 let mut pools = mgr.pools.lock().unwrap();
257 let victims: Vec<String> = pools
258 .iter()
259 .filter(|(_, e)| e.tunnel_generation == Some(generation))
260 .map(|(k, _)| k.clone())
261 .collect();
262 let n = victims.len();
263 for k in victims {
264 if let Some(entry) = pools.remove(&k) {
265 let _close = tokio::task::spawn(entry.pool.disconnect());
266 }
267 }
268 n
269}
270
271impl Default for PoolManager {
272 fn default() -> Self {
273 Self::new()
274 }
275}
276
277#[cfg(test)]
278mod tests {
279 use super::*;
280
281 fn conn(host: &str) -> MySqlConnection {
282 MySqlConnection {
283 name: "t".into(),
284 host: host.into(),
285 port: 3306,
286 user: "u".into(),
287 ..MySqlConnection::default()
288 }
289 }
290
291 #[test]
292 fn generation_is_deterministic_per_process() {
293 let a = CredentialGeneration::derive("pw");
294 let b = CredentialGeneration::derive("pw");
295 assert!(a.eq_material(&b));
296 assert!(!a.eq_material(&CredentialGeneration::derive("other")));
297 }
298
299 #[test]
300 fn generation_debug_and_display_are_redacted() {
301 let g = CredentialGeneration::derive("secret-value");
302 assert_eq!(format!("{g:?}"), "CredentialGeneration(<redacted>)");
303 assert_eq!(format!("{g}"), "<redacted>");
304 let s = format!("{g:?}{g}");
306 assert!(!s.contains("secret-value"));
307 }
308
309 #[tokio::test]
310 async fn failed_handshake_never_populates_cache() {
311 let mgr = PoolManager::new();
312 let c = conn("127.0.0.1");
314 let pw = Zeroizing::new("x".to_string());
315 let err = mgr
316 .verified_pool(&c, &pw, None, 1, None, Some(1), None)
317 .await
318 .unwrap_err();
319 assert!(matches!(err, PoolManagerError::Init(_)), "{err:?}");
320 assert_eq!(mgr.pool_count(), 0);
321 }
322
323 #[tokio::test]
324 async fn concurrent_first_access_initializes_once() {
325 use std::sync::Arc;
326 let mgr = Arc::new(PoolManager::new());
330 let c = conn("127.0.0.1");
331 let pw = Zeroizing::new("x".to_string());
332 let mut handles = Vec::new();
333 for _ in 0..4 {
334 let m = mgr.clone();
335 let cc = c.clone();
336 let p = pw.clone();
337 handles.push(tokio::task::spawn(async move {
338 m.verified_pool(&cc, &p, None, 1, None, Some(2), None).await
339 }));
340 }
341 for h in handles {
342 assert!(h.await.unwrap().is_err());
343 }
344 assert_eq!(mgr.pool_count(), 0, "failures must not be cached");
345 }
346
347 #[test]
348 fn config_key_separates_transport_settings() {
349 let c1 = conn("db1.example.invalid");
350 let mut c2 = conn("db2.example.invalid");
351 c2.ssl = true;
352 let k1 = PoolManager::config_key(&c1, None, 1, None, None, None);
353 let k2 = PoolManager::config_key(&c2, None, 1, None, None, None);
354 let k3 = PoolManager::config_key(&c1, None, 2, None, None, None);
355 assert_ne!(k1, k2);
356 assert_ne!(k1, k3);
357 let k4 = PoolManager::config_key(&c1, None, 1, None, None, Some(7));
361 let k5 = PoolManager::config_key(&c1, None, 1, None, None, Some(8));
362 assert_ne!(k1, k4);
363 assert_ne!(k4, k5);
364 }
365}