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