Skip to main content

sequel_mcp/sql/
ssh.rs

1//! SSH direct-tunnel transport (russh port of the legacy tunnel runtime).
2//!
3//! Topology per connection: `mysql_async` always connects to a LOCAL
4//! loopback endpoint; every local TCP connection is bridged onto a fresh
5//! `direct_tcpip` channel of ONE multiplexed SSH session to the bastion,
6//! which forwards to the MySQL host/port as seen FROM the bastion.
7//! Host keys are verified against known_hosts through the strict/lenient
8//! policy engine — mismatches and revoked keys are rejected in every
9//! mode; only genuinely unknown hosts may be accepted under lenient.
10//!
11//! Hardening (#4A): concurrent first users of the same tunnel coalesce
12//! on ONE establishment; every tunnel carries a monotonic GENERATION
13//! that also keys the MySQL pool, so a reused loopback port can never
14//! splice an old pool onto a new transport (ABA); retirement follows a
15//! fixed order (drain → evict pools → stop listener → wind down channel
16//! tasks → disconnect the session); the cache is LRU and bounded;
17//! keepalive is tunable for half-open detection; the tunnel cache key
18//! binds the config revision, the SSH credential generation, and the
19//! known_hosts content, so rotating any of them invalidates the cached
20//! session. Test-mode rules apply to the bastion endpoint before any
21//! socket is opened.
22
23use crate::config::{SshAuthMethod, SshHostKeyPolicy, SshTunnel};
24use crate::sql::known_hosts;
25use russh::client::{self, Handle};
26use russh::keys::PublicKeyBase64;
27use sha2::Digest as _;
28use std::collections::HashMap;
29use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
30use std::sync::{Arc, Mutex, OnceLock};
31use thiserror::Error;
32use tokio::net::TcpListener;
33use tokio::sync::mpsc;
34
35/// Deadline for establishing the SSH session (connect + KEX + auth).
36pub const SSH_CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
37
38/// Deadline for opening a single forwarded channel (a stale/half-open
39/// session must fail the local connection promptly, not hang it).
40pub const CHANNEL_OPEN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
41
42/// Maximum distinct tunnels kept in the process-wide cache (LRU).
43pub const MAX_TUNNELS: usize = 8;
44
45/// Keepalive interval. Default 30 s; tunable (1..=300) via
46/// `SEQUEL_MCP_SSH_KEEPALIVE_SECS` — an operational knob that also lets
47/// the half-open tests run with a realistic detection budget.
48pub fn keepalive_interval() -> std::time::Duration {
49    let secs = std::env::var("SEQUEL_MCP_SSH_KEEPALIVE_SECS")
50        .ok()
51        .and_then(|v| v.trim().parse::<u64>().ok())
52        .unwrap_or(30)
53        .clamp(1, 300);
54    std::time::Duration::from_secs(secs)
55}
56
57#[derive(Debug, Error)]
58pub enum SshError {
59    #[error("{0}")]
60    Transport(String),
61    #[error("authentication as {user:?} failed on {host}:{port}")]
62    Auth {
63        host: String,
64        port: u16,
65        user: String,
66    },
67    #[error("host key rejected for {host}:{port}: {reason}")]
68    HostKey {
69        host: String,
70        port: u16,
71        reason: String,
72    },
73    #[error("tunnel setup: {0}")]
74    Setup(String),
75}
76
77/// `russh` client handler whose only job is host-key verification via
78/// the ported known_hosts engine (strict = fail closed; lenient accepts
79/// ONLY genuinely unknown hosts — mismatches and revoked keys are
80/// rejected in every mode).
81struct HostKeyHandler {
82    policy: SshHostKeyPolicy,
83    entries: Vec<known_hosts::KnownHostEntry>,
84    host: String,
85    port: u16,
86    /// Shared with `establish` so the TOFU/mismatch commentary is
87    /// actually EMITTED (stderr — protocol-safe) after connect returns
88    /// instead of silently buffered (review finding).
89    logs: std::sync::Arc<Mutex<Vec<String>>>,
90}
91
92impl client::Handler for HostKeyHandler {
93    type Error = russh::Error;
94
95    async fn check_server_key(
96        &mut self,
97        server_public_key: &russh::keys::PublicKey,
98    ) -> Result<bool, Self::Error> {
99        let raw = server_public_key.public_key_bytes();
100        let mut logs: Vec<String> = Vec::new();
101        let decision = known_hosts::decide_host_key(
102            self.policy,
103            &self.host,
104            self.port,
105            &self.entries,
106            &raw,
107            &mut |line| logs.push(line),
108        );
109        self.logs.lock().unwrap().extend(logs);
110        Ok(decision == known_hosts::HostKeyDecision::Accept)
111    }
112}
113
114/// A lease on a tunnel: the loopback endpoint for `mysql_async` plus the
115/// transport GENERATION. The generation participates in the MySQL pool
116/// key, so a pool can never outlive its tunnel and get spliced onto a
117/// later transport that reuses the same ephemeral port.
118#[derive(Debug, Clone, PartialEq, Eq)]
119pub struct TunnelLease {
120    pub host: String,
121    pub port: u16,
122    pub generation: u64,
123}
124
125struct TunnelEntry {
126    /// Local loopback listener address handed to `mysql_async`.
127    local: std::net::SocketAddr,
128    /// Multiplexed SSH session; channels are opened per TCP connection.
129    handle: Arc<Handle<HostKeyHandler>>,
130    /// Sender whose message ends the accept-loop task.
131    shutdown: mpsc::Sender<()>,
132    generation: u64,
133    /// Set on retirement: the accept loop stops taking new connections.
134    draining: Arc<AtomicBool>,
135    last_used: Mutex<std::time::Instant>,
136    /// Live forwarder tasks (diagnostics + test assertions).
137    live_tasks: Arc<AtomicU64>,
138}
139
140impl TunnelEntry {
141    fn is_dead(&self) -> bool {
142        self.handle.is_closed()
143    }
144    fn touch(&self) {
145        *self.last_used.lock().unwrap() = std::time::Instant::now();
146    }
147}
148
149struct Tunnels {
150    map: HashMap<String, Arc<TunnelEntry>>,
151    next_generation: u64,
152    /// Per-key async guards so concurrent first users coalesce on ONE
153    /// establishment instead of racing N SSH sessions.
154    inflight: HashMap<String, Arc<tokio::sync::Mutex<()>>>,
155}
156
157static TUNNELS: OnceLock<Mutex<Tunnels>> = OnceLock::new();
158
159fn tunnels() -> &'static Mutex<Tunnels> {
160    TUNNELS.get_or_init(|| {
161        Mutex::new(Tunnels {
162            map: HashMap::new(),
163            next_generation: 1,
164            inflight: HashMap::new(),
165        })
166    })
167}
168
169/// Number of live tunnels (diagnostics/tests).
170pub fn tunnel_count() -> usize {
171    tunnels().lock().unwrap().map.len()
172}
173
174/// Connection-name fragments of the live tunnel keys (LRU tests).
175pub fn tunnel_connection_names() -> Vec<String> {
176    tunnels()
177        .lock()
178        .unwrap()
179        .map
180        .keys()
181        .map(|k| k.split('\u{1}').next().unwrap_or("").to_string())
182        .collect()
183}
184
185/// Live forwarder tasks across all tunnels (drain assertions).
186pub fn tunnel_live_tasks() -> u64 {
187    tunnels()
188        .lock()
189        .unwrap()
190        .map
191        .values()
192        .map(|e| e.live_tasks.load(Ordering::Relaxed))
193        .sum()
194}
195
196/// Drop every cached tunnel and retire each one in order (tests).
197pub fn invalidate_all() {
198    let entries: Vec<Arc<TunnelEntry>> = {
199        let mut state = tunnels().lock().unwrap();
200        state.map.drain().map(|(_, e)| e).collect()
201    };
202    for entry in entries {
203        retire(&entry);
204    }
205}
206
207/// Retire one tunnel in the fixed order: mark draining (no new channel
208/// creation) → evict the MySQL pools keyed to this generation → stop
209/// the listener → let channel tasks wind down → disconnect the SSH
210/// session.
211fn retire(entry: &TunnelEntry) {
212    entry.draining.store(true, Ordering::SeqCst);
213    let evicted = super::pool::evict_by_generation(entry.generation);
214    if evicted > 0 {
215        eprintln!(
216            "[sequel-mcp] ssh tunnel gen {}: evicted {evicted} associated MySQL pool(s)",
217            entry.generation
218        );
219    }
220    let _ = entry.shutdown.try_send(());
221    // Disconnect the SSH session; outstanding channel copies error out
222    // and their tasks finish. Fire-and-forget with a bounded handle.
223    let handle = Arc::clone(&entry.handle);
224    tokio::spawn(async move {
225        let _ = handle
226            .disconnect(russh::Disconnect::ByApplication, "tunnel retired", "en")
227            .await;
228    });
229}
230
231impl SshTunnel {
232    fn auth_method_label(&self) -> &'static str {
233        match self.auth_method {
234            SshAuthMethod::Password => "password",
235            SshAuthMethod::Key => "key",
236        }
237    }
238}
239
240/// 16-hex-char content stamp of the known_hosts file (empty string when
241/// no explicit file is configured — the default path's absence is not a
242/// distinguishing identity).
243fn known_hosts_stamp(path: Option<&std::path::Path>) -> String {
244    let Some(path) = path else {
245        return String::new();
246    };
247    let mut h = sha2::Sha256::new();
248    match std::fs::read(path) {
249        Ok(bytes) => h.update(&bytes),
250        // Unreadable files fail closed in the checked loader; the stamp
251        // only needs to be deterministic per file identity.
252        Err(_) => h.update(format!("unreadable:{}", path.display())),
253    }
254    format!(
255        "{:016x}",
256        u64::from_be_bytes(h.finalize()[..8].try_into().expect("8 bytes"))
257    )
258}
259
260/// In-process credential generation for the SSH secret (never the
261/// secret itself): a process-keyed HMAC digest, hex-encoded for keying.
262fn ssh_credential_fragment(ssh_password: Option<&str>) -> String {
263    let cred = super::pool::CredentialGeneration::derive(ssh_password.unwrap_or(""));
264    cred.key_fragment()
265}
266
267#[allow(clippy::too_many_arguments)]
268fn tunnel_key(
269    conn_name: &str,
270    ssh: &SshTunnel,
271    ssh_password: Option<&str>,
272    target: &str,
273    port: u16,
274    policy_revision: u64,
275    kh_stamp: &str,
276) -> String {
277    let target_endpoint = format!("{target}:{port}");
278    let bridge = ssh
279        .docker
280        .as_ref()
281        .map(|d| format!("{}:{}", d.container, d.bridge_tool.as_str()))
282        .unwrap_or_default();
283    format!(
284        "{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}",
285        conn_name,
286        ssh.host,
287        ssh.port,
288        ssh.user,
289        ssh.auth_method_label(),
290        kh_stamp,
291        ssh_credential_fragment(ssh_password),
292        policy_revision,
293        target_endpoint,
294        bridge
295    )
296}
297
298/// Return a lease on a (reused or freshly established) SSH tunnel for
299/// `conn_name`'s SSH config, forwarding to `target_host:target_port` as
300/// reachable from the bastion. Concurrent first callers coalesce on one
301/// establishment. The lease's generation MUST be carried into the MySQL
302/// pool key (`verified_pool(.., tunnel_generation)`).
303pub async fn tunnel_endpoint(
304    conn_name: &str,
305    ssh: &SshTunnel,
306    ssh_password: Option<&str>,
307    target_host: &str,
308    target_port: u16,
309    policy_revision: u64,
310) -> Result<TunnelLease, SshError> {
311    // Fail-closed test-mode gate on the BASTION endpoint before any
312    // socket is opened (docker bastions publish on loopback).
313    crate::app::test_mode::check_mysql_endpoint(&ssh.host, ssh.port).map_err(SshError::Setup)?;
314
315    let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
316    let kh_stamp = known_hosts_stamp(kh_path);
317    let key = tunnel_key(
318        conn_name,
319        ssh,
320        ssh_password,
321        target_host,
322        target_port,
323        policy_revision,
324        &kh_stamp,
325    );
326
327    // Fast path: a live tunnel for exactly this identity.
328    {
329        let state = tunnels().lock().unwrap();
330        if let Some(entry) = state.map.get(&key)
331            && !entry.is_dead()
332            && !entry.draining.load(Ordering::SeqCst)
333        {
334            entry.touch();
335            return Ok(TunnelLease {
336                host: "127.0.0.1".into(),
337                port: entry.local.port(),
338                generation: entry.generation,
339            });
340        }
341    }
342
343    // Coalescing: one establishment per key even under a thundering
344    // herd — everyone serializes on the per-key guard, and waits find
345    // the winner's entry on the double-check.
346    let guard = {
347        let mut state = tunnels().lock().unwrap();
348        Arc::clone(state.inflight.entry(key.clone()).or_default())
349    };
350    let _hold = guard.lock().await;
351
352    {
353        let state = tunnels().lock().unwrap();
354        if let Some(entry) = state.map.get(&key)
355            && !entry.is_dead()
356            && !entry.draining.load(Ordering::SeqCst)
357        {
358            entry.touch();
359            return Ok(TunnelLease {
360                host: "127.0.0.1".into(),
361                port: entry.local.port(),
362                generation: entry.generation,
363            });
364        }
365    }
366
367    let entry = establish(ssh, ssh_password, target_host, target_port).await?;
368
369    let mut state = tunnels().lock().unwrap();
370    state.inflight.remove(&key);
371    // Retire stale tunnels of the SAME connection (identity rotation:
372    // credential, known_hosts content, or policy revision changed).
373    let conn_prefix = format!("{conn_name}\u{1}");
374    let stale: Vec<String> = state
375        .map
376        .keys()
377        .filter(|k| k.as_str() != key.as_str() && k.starts_with(&conn_prefix))
378        .cloned()
379        .collect();
380    for k in stale {
381        if let Some(e) = state.map.remove(&k) {
382            retire(&e);
383        }
384    }
385    // Drop dead entries; enforce the LRU bound by retiring the
386    // least-recently-used victims.
387    let dead: Vec<String> = state
388        .map
389        .iter()
390        .filter(|(_, e)| e.is_dead())
391        .map(|(k, _)| k.clone())
392        .collect();
393    for k in dead {
394        if let Some(e) = state.map.remove(&k) {
395            retire(&e);
396        }
397    }
398    while state.map.len() >= MAX_TUNNELS {
399        let victim = state
400            .map
401            .iter()
402            .min_by_key(|(_, e)| *e.last_used.lock().unwrap())
403            .map(|(k, _)| k.clone());
404        let Some(victim) = victim else { break };
405        if let Some(e) = state.map.remove(&victim) {
406            eprintln!(
407                "[sequel-mcp] ssh tunnel cache full ({}); retiring LRU entry",
408                MAX_TUNNELS
409            );
410            retire(&e);
411        }
412    }
413    let generation = entry.generation;
414    let port = entry.local.port();
415    state.map.insert(key, entry);
416    Ok(TunnelLease {
417        host: "127.0.0.1".into(),
418        port,
419        generation,
420    })
421}
422
423/// When `ssh.docker` is configured, per-connection forwarding runs over
424/// an SSH **exec** channel (`docker exec -i <container> <tool> …`)
425/// instead of direct-tcpip — the bridge that works even where sshd
426/// denies TCP forwarding. All argv components are validated (no shell,
427/// no spaces), so the exec command line is a plain join.
428fn bridge_command(
429    ssh: &SshTunnel,
430    target_host: &str,
431    target_port: u16,
432) -> Result<Option<String>, SshError> {
433    let Some(docker) = &ssh.docker else {
434        return Ok(None);
435    };
436    let argv = super::docker::bridge_argv(
437        &docker.container,
438        docker.bridge_tool,
439        target_host,
440        target_port,
441    )
442    .map_err(|e| SshError::Setup(format!("docker bridge: {e}")))?;
443    Ok(Some(argv.join(" ")))
444}
445
446async fn establish(
447    ssh: &SshTunnel,
448    ssh_password: Option<&str>,
449    target_host: &str,
450    target_port: u16,
451) -> Result<Arc<TunnelEntry>, SshError> {
452    let bridge_cmd = bridge_command(ssh, target_host, target_port)?;
453    let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
454    let entries =
455        known_hosts::load_known_hosts_checked(kh_path).map_err(|e| SshError::HostKey {
456            host: ssh.host.clone(),
457            port: ssh.port,
458            reason: e,
459        })?;
460    let logs = std::sync::Arc::new(Mutex::new(Vec::new()));
461    let handler = HostKeyHandler {
462        policy: ssh.host_key_policy.unwrap_or(SshHostKeyPolicy::Lenient),
463        entries,
464        host: ssh.host.clone(),
465        port: ssh.port,
466        logs: logs.clone(),
467    };
468
469    let keepalive = keepalive_interval();
470    let config = Arc::new(client::Config {
471        keepalive_interval: Some(keepalive),
472        keepalive_max: 3,
473        nodelay: true,
474        ..client::Config::default()
475    });
476
477    let addr = (ssh.host.as_str(), ssh.port);
478    let session = async {
479        let mut handle = client::connect(config, addr, handler)
480            .await
481            .map_err(|e| SshError::Transport(format!("connect to bastion: {e}")))?;
482        // Host-key commentary (TOFU accepts, migration-compat warnings)
483        // must reach the operator, not rot in a buffer.
484        for line in logs.lock().unwrap().drain(..) {
485            eprintln!("[sequel-mcp] SSH {}:{} {line}", ssh.host, ssh.port);
486        }
487        // Authenticate with EXACTLY the configured method — no silent
488        // fallback between password and key. The secret doubles as the
489        // private-key passphrase under key auth.
490        let auth = match ssh.auth_method {
491            SshAuthMethod::Password => {
492                let Some(password) = ssh_password else {
493                    return Err(SshError::Setup(format!(
494                        "no SSH password stored for {:?} (expected under \"<connection>::ssh\")",
495                        ssh.user
496                    )));
497                };
498                handle
499                    .authenticate_password(ssh.user.clone(), password)
500                    .await
501                    .map_err(|e| SshError::Transport(format!("password auth transport: {e}")))?
502            }
503            SshAuthMethod::Key => {
504                let Some(path) = &ssh.private_key_path else {
505                    return Err(SshError::Setup(
506                        "key auth configured without privateKeyPath".into(),
507                    ));
508                };
509                let expanded = crate::app::paths::expand_tilde(path);
510                let key = russh::keys::load_secret_key(&expanded, ssh_password)
511                    .map_err(|e| SshError::Setup(format!("load private key: {e}")))?;
512                use russh::keys::Algorithm;
513                match key.algorithm() {
514                    Algorithm::Ed25519 | Algorithm::Rsa { .. } | Algorithm::Ecdsa { .. } => {}
515                    other => {
516                        return Err(SshError::Setup(format!(
517                            "unsupported private key algorithm {other:?} (supported: ed25519, rsa, ecdsa)"
518                        )));
519                    }
520                }
521                // Let the server's advertised algorithms pick the RSA
522                // signature hash (SHA-512 preferred, SHA-256 next;
523                // legacy ssh-rsa/SHA-1 is never selected).
524                let hash_alg = if key.algorithm().is_rsa() {
525                    match handle.best_supported_rsa_hash().await {
526                        Ok(best) => best.unwrap_or(Some(russh::keys::HashAlg::Sha256)),
527                        Err(_) => Some(russh::keys::HashAlg::Sha256),
528                    }
529                } else {
530                    None
531                };
532                handle
533                    .authenticate_publickey(
534                        ssh.user.clone(),
535                        russh::keys::PrivateKeyWithHashAlg::new(Arc::new(key), hash_alg),
536                    )
537                    .await
538                    .map_err(|e| SshError::Transport(format!("publickey auth transport: {e}")))?
539            }
540        };
541        if !matches!(auth, russh::client::AuthResult::Success) {
542            return Err(SshError::Auth {
543                host: ssh.host.clone(),
544                port: ssh.port,
545                user: ssh.user.clone(),
546            });
547        }
548        Ok(handle)
549    };
550    let handle = tokio::time::timeout(SSH_CONNECT_TIMEOUT, session)
551        .await
552        .map_err(|_| {
553            SshError::Transport(format!(
554                "ssh session timeout after {}s",
555                SSH_CONNECT_TIMEOUT.as_secs()
556            ))
557        })??;
558
559    // Local loopback forwarder: one fresh direct-tcpip channel per
560    // accepted TCP connection, bounded by CHANNEL_OPEN_TIMEOUT so a
561    // stale session fails the connection instead of hanging it.
562    let listener = TcpListener::bind(("127.0.0.1", 0))
563        .await
564        .map_err(|e| SshError::Setup(format!("local bind: {e}")))?;
565    let local = listener
566        .local_addr()
567        .map_err(|e| SshError::Setup(format!("local addr: {e}")))?;
568
569    let generation = {
570        let mut state = tunnels().lock().unwrap();
571        let assigned = state.next_generation;
572        state.next_generation += 1;
573        assigned
574    };
575
576    let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
577    let forward_host = target_host.to_string();
578    let forward_port = u32::from(target_port);
579    let shared = Arc::new(handle);
580    let accept_session = Arc::clone(&shared);
581    let draining = Arc::new(AtomicBool::new(false));
582    let live_tasks: Arc<AtomicU64> = Arc::new(AtomicU64::new(0));
583    let accept_draining = Arc::clone(&draining);
584    let accept_tasks = Arc::clone(&live_tasks);
585    tokio::spawn(async move {
586        loop {
587            tokio::select! {
588                _ = shutdown_rx.recv() => break,
589                accepted = listener.accept() => {
590                    let Ok((socket, _peer)) = accepted else { break };
591                    if accept_draining.load(Ordering::SeqCst) || accept_session.is_closed() {
592                        break;
593                    }
594                    let session = Arc::clone(&accept_session);
595                    let host = forward_host.clone();
596                    let tasks = Arc::clone(&accept_tasks);
597                    accept_tasks.fetch_add(1, Ordering::Relaxed);
598                    let bridge_cmd = bridge_cmd.clone();
599                    tokio::spawn(async move {
600                        let target = format!("{host}:{forward_port}");
601                        // Bridge connections open a session channel and
602                        // exec the (validated, space-free) docker bridge
603                        // argv; direct connections use direct-tcpip.
604                        let channel = tokio::time::timeout(
605                            CHANNEL_OPEN_TIMEOUT,
606                            async {
607                                match &bridge_cmd {
608                                    Some(cmd) => {
609                                        let ch = session.channel_open_session().await?;
610                                        ch.exec(true, cmd.as_str()).await?;
611                                        Ok::<_, russh::Error>(ch)
612                                    }
613                                    None => session
614                                        .channel_open_direct_tcpip(
615                                            host,
616                                            forward_port,
617                                            "127.0.0.1",
618                                            0,
619                                        )
620                                        .await,
621                                }
622                            },
623                        )
624                        .await;
625                        match channel {
626                            Ok(Ok(channel)) => {
627                                // ChannelStream is the full tokio-IO view
628                                // of a channel; split gives owned halves
629                                // for the bidirectional copy.
630                                let (mut chan_r, mut chan_w) =
631                                    tokio::io::split(channel.into_stream());
632                                let (mut sock_r, mut sock_w) = socket.into_split();
633                                let a = tokio::io::copy(&mut sock_r, &mut chan_w);
634                                let b = tokio::io::copy(&mut chan_r, &mut sock_w);
635                                let _ = tokio::join!(a, b);
636                            }
637                            Ok(Err(e)) => {
638                                eprintln!(
639                                    "[sequel-mcp] ssh forwarder: channel open to {target} failed: {e:?}"
640                                );
641                                // Dropping the socket closes it: the MySQL
642                                // handshake fails promptly.
643                            }
644                            Err(_) => {
645                                eprintln!(
646                                    "[sequel-mcp] ssh forwarder: channel open to {target} timed out after {}s (stale session?)",
647                                    CHANNEL_OPEN_TIMEOUT.as_secs()
648                                );
649                            }
650                        }
651                        tasks.fetch_sub(1, Ordering::Relaxed);
652                    });
653                }
654            }
655        }
656    });
657
658    Ok(Arc::new(TunnelEntry {
659        local,
660        handle: shared,
661        shutdown: shutdown_tx,
662        generation,
663        draining,
664        last_used: Mutex::new(std::time::Instant::now()),
665        live_tasks,
666    }))
667}
668
669#[cfg(test)]
670mod tests {
671    use super::*;
672    use crate::config::SshTunnel;
673
674    fn sample_tunnel() -> SshTunnel {
675        SshTunnel {
676            host: "bastion".into(),
677            port: 22,
678            user: "sshuser".into(),
679            auth_method: SshAuthMethod::Password,
680            ..SshTunnel::default()
681        }
682    }
683
684    #[test]
685    fn tunnel_key_separates_identity() {
686        let ssh = sample_tunnel();
687        let a = tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp");
688        let mut ssh2 = ssh.clone();
689        ssh2.user = "other".into();
690        let b = tunnel_key("c1", &ssh2, Some("pw"), "db", 3306, 1, "stamp");
691        assert_ne!(a, b);
692        assert_ne!(
693            a,
694            tunnel_key("c1", &ssh, Some("pw"), "other", 3306, 1, "stamp")
695        );
696        assert_ne!(
697            a,
698            tunnel_key("c1", &ssh, Some("pw"), "db", 3307, 1, "stamp")
699        );
700        assert_ne!(
701            a,
702            tunnel_key("c2", &ssh, Some("pw"), "db", 3306, 1, "stamp")
703        );
704        // Rotation inputs: credential, known_hosts content, revision.
705        assert_ne!(
706            a,
707            tunnel_key("c1", &ssh, Some("different"), "db", 3306, 1, "stamp")
708        );
709        assert_ne!(
710            a,
711            tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 2, "stamp")
712        );
713        assert_ne!(
714            a,
715            tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp2")
716        );
717        // Bridge identity (container + tool) is part of the key.
718        let mut bridged = ssh.clone();
719        bridged.docker = Some(crate::config::SshDocker {
720            container: "db".into(),
721            bridge_tool: crate::config::BridgeTool::Nc,
722        });
723        assert_ne!(
724            a,
725            tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
726        );
727        let mut bridged2 = bridged.clone();
728        bridged2.docker = Some(crate::config::SshDocker {
729            container: "db".into(),
730            bridge_tool: crate::config::BridgeTool::Socat,
731        });
732        assert_ne!(
733            tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp"),
734            tunnel_key("c1", &bridged2, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
735        );
736    }
737
738    #[test]
739    fn bridge_command_forms() {
740        use crate::config::{BridgeTool, SshDocker};
741        let mk = |tool| SshTunnel {
742            host: "bastion".into(),
743            port: 22,
744            user: "u".into(),
745            auth_method: SshAuthMethod::Key,
746            docker: Some(SshDocker {
747                container: "db-1".into(),
748                bridge_tool: tool,
749            }),
750            ..SshTunnel::default()
751        };
752        assert_eq!(
753            bridge_command(&mk(BridgeTool::Nc), "127.0.0.1", 3306)
754                .unwrap()
755                .unwrap(),
756            "docker exec -i db-1 nc 127.0.0.1 3306"
757        );
758        assert_eq!(
759            bridge_command(&mk(BridgeTool::Ncat), "127.0.0.1", 3306)
760                .unwrap()
761                .unwrap(),
762            "docker exec -i db-1 ncat 127.0.0.1 3306"
763        );
764        assert_eq!(
765            bridge_command(&mk(BridgeTool::Socat), "127.0.0.1", 3306)
766                .unwrap()
767                .unwrap(),
768            "docker exec -i db-1 socat - TCP:127.0.0.1:3306"
769        );
770        // No docker config: direct-tcpip path.
771        let plain = SshTunnel::default();
772        assert_eq!(bridge_command(&plain, "db", 3306).unwrap(), None);
773        // Invalid inputs fail typed before any connection.
774        let mut bad = mk(BridgeTool::Nc);
775        bad.docker = Some(SshDocker {
776            container: "bad name!".into(),
777            bridge_tool: BridgeTool::Nc,
778        });
779        assert!(bridge_command(&bad, "127.0.0.1", 3306).is_err());
780    }
781
782    #[test]
783    fn credential_fragment_hides_the_secret() {
784        let f = ssh_credential_fragment(Some("top-secret-password"));
785        assert_eq!(f.len(), 32, "hex of a 16-byte digest");
786        assert!(!f.contains("top-secret"));
787        assert_ne!(f, ssh_credential_fragment(Some("other")));
788        // A missing and an empty secret are the same non-secret here.
789        assert_eq!(
790            ssh_credential_fragment(None),
791            ssh_credential_fragment(Some(""))
792        );
793    }
794
795    #[test]
796    fn known_hosts_stamp_tracks_content() {
797        let dir = tempfile::TempDir::new().unwrap();
798        let f = dir.path().join("kh");
799        std::fs::write(&f, b"content-a").unwrap();
800        let s1 = known_hosts_stamp(Some(&f));
801        let s2 = known_hosts_stamp(Some(&f));
802        assert_eq!(s1, s2, "stable per content");
803        std::fs::write(&f, b"content-b").unwrap();
804        assert_ne!(s1, known_hosts_stamp(Some(&f)));
805        assert_eq!(known_hosts_stamp(None), "");
806    }
807}