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("authentication aborted as {user:?} on {host}:{port}: {reason}")]
68    AuthAborted {
69        host: String,
70        port: u16,
71        user: String,
72        reason: String,
73    },
74    #[error("host key rejected for {host}:{port}: {reason}")]
75    HostKey {
76        host: String,
77        port: u16,
78        reason: String,
79    },
80    #[error("tunnel setup: {0}")]
81    Setup(String),
82}
83
84/// Why an auth exchange can end WITHOUT a server verdict. russh
85/// (0.62 and 0.63 alike) collapses a dead session into
86/// `AuthResult::Failure` with an empty remaining-method set (the reply
87/// channel closes and `wait_recv_reply` maps `None` to `Failure`), so
88/// the 0.10.2 incident — RSA signing compiled out, ssh-key rejecting
89/// the signature as `AlgorithmUnsupported { algorithm: Rsa { hash:
90/// None } }` AFTER the server had already answered USERAUTH_PK_OK —
91/// surfaced as "authentication failed", indistinguishable from a
92/// rejected credential. This reason must make the difference visible.
93/// (russh 0.63.2 also logs-and-fails this case inside the session task
94/// — Eugeny/russh#758 — but the caller-visible shape is unchanged, so
95/// this discrimination remains ours.)
96fn auth_abort_reason(key_is_rsa: bool) -> String {
97    if key_is_rsa && cfg!(not(feature = "rsa")) {
98        "no server verdict was delivered — the SSH session ended mid-auth. \
99         This build cannot SIGN RSA keys: russh's `rsa` feature is missing \
100         (enabled by sequel-mcp's default `rsa` feature). The key itself is \
101         fine and the server had not rejected it"
102            .into()
103    } else if key_is_rsa {
104        "no server verdict was delivered — the SSH session ended mid-auth, a \
105         signing/transport failure and NOT a rejected credential. For RSA keys \
106         the classic cause is a build without russh's `rsa` feature; \
107         RUST_LOG=russh=debug shows the underlying error"
108            .into()
109    } else {
110        "no server verdict was delivered — the SSH session ended mid-auth, a \
111         transport failure and NOT a rejected credential; RUST_LOG=russh=debug \
112         shows the underlying error"
113            .into()
114    }
115}
116
117/// True when the error means "we could not produce the signature the
118/// protocol asked for" rather than "the network/auth exchange broke".
119/// ssh-key reports exactly this shape when RSA signing support is
120/// compiled out (`Rsa { hash: None }` = the negotiated hash never
121/// reached the signer).
122fn is_signature_algorithm_error(e: &russh::Error) -> bool {
123    matches!(
124        e,
125        russh::Error::SshKey(russh::keys::ssh_key::Error::AlgorithmUnsupported { .. })
126    )
127}
128
129/// `russh` client handler whose only job is host-key verification via
130/// the ported known_hosts engine (strict = fail closed; lenient accepts
131/// ONLY genuinely unknown hosts — mismatches and revoked keys are
132/// rejected in every mode).
133struct HostKeyHandler {
134    policy: SshHostKeyPolicy,
135    entries: Vec<known_hosts::KnownHostEntry>,
136    host: String,
137    port: u16,
138    /// Shared with `establish` so the TOFU/mismatch commentary is
139    /// actually EMITTED (stderr — protocol-safe) after connect returns
140    /// instead of silently buffered (review finding).
141    logs: std::sync::Arc<Mutex<Vec<String>>>,
142}
143
144impl client::Handler for HostKeyHandler {
145    type Error = russh::Error;
146
147    async fn check_server_key(
148        &mut self,
149        server_public_key: &russh::keys::PublicKeyOrCertificate,
150    ) -> Result<bool, Self::Error> {
151        // The known_hosts engine compares raw key bytes. russh 0.63 can
152        // also surface host CERTIFICATES here; no stored entry can ever
153        // match one, so they fail closed rather than being reduced to
154        // their signing key (which would smuggle an unverified
155        // certificate's key past the pin).
156        let raw = match server_public_key {
157            russh::keys::PublicKeyOrCertificate::PublicKey { key, .. } => key.public_key_bytes(),
158            russh::keys::PublicKeyOrCertificate::Certificate(_) => return Ok(false),
159        };
160        let mut logs: Vec<String> = Vec::new();
161        let decision = known_hosts::decide_host_key(
162            self.policy,
163            &self.host,
164            self.port,
165            &self.entries,
166            &raw,
167            &mut |line| logs.push(line),
168        );
169        self.logs.lock().unwrap().extend(logs);
170        Ok(decision == known_hosts::HostKeyDecision::Accept)
171    }
172}
173
174/// A lease on a tunnel: the loopback endpoint for `mysql_async` plus the
175/// transport GENERATION. The generation participates in the MySQL pool
176/// key, so a pool can never outlive its tunnel and get spliced onto a
177/// later transport that reuses the same ephemeral port.
178#[derive(Debug, Clone, PartialEq, Eq)]
179pub struct TunnelLease {
180    pub host: String,
181    pub port: u16,
182    pub generation: u64,
183}
184
185struct TunnelEntry {
186    /// Local loopback listener address handed to `mysql_async`.
187    local: std::net::SocketAddr,
188    /// Multiplexed SSH session; channels are opened per TCP connection.
189    handle: Arc<Handle<HostKeyHandler>>,
190    /// Sender whose message ends the accept-loop task.
191    shutdown: mpsc::Sender<()>,
192    generation: u64,
193    /// Set on retirement: the accept loop stops taking new connections.
194    draining: Arc<AtomicBool>,
195    last_used: Mutex<std::time::Instant>,
196    /// Live forwarder tasks (diagnostics + test assertions).
197    live_tasks: Arc<AtomicU64>,
198}
199
200impl TunnelEntry {
201    fn is_dead(&self) -> bool {
202        self.handle.is_closed()
203    }
204    fn touch(&self) {
205        *self.last_used.lock().unwrap() = std::time::Instant::now();
206    }
207}
208
209struct Tunnels {
210    map: HashMap<String, Arc<TunnelEntry>>,
211    next_generation: u64,
212    /// Per-key async guards so concurrent first users coalesce on ONE
213    /// establishment instead of racing N SSH sessions.
214    inflight: HashMap<String, Arc<tokio::sync::Mutex<()>>>,
215}
216
217static TUNNELS: OnceLock<Mutex<Tunnels>> = OnceLock::new();
218
219fn tunnels() -> &'static Mutex<Tunnels> {
220    TUNNELS.get_or_init(|| {
221        Mutex::new(Tunnels {
222            map: HashMap::new(),
223            next_generation: 1,
224            inflight: HashMap::new(),
225        })
226    })
227}
228
229/// Number of live tunnels (diagnostics/tests).
230pub fn tunnel_count() -> usize {
231    tunnels().lock().unwrap().map.len()
232}
233
234/// Connection-name fragments of the live tunnel keys (LRU tests).
235pub fn tunnel_connection_names() -> Vec<String> {
236    tunnels()
237        .lock()
238        .unwrap()
239        .map
240        .keys()
241        .map(|k| k.split('\u{1}').next().unwrap_or("").to_string())
242        .collect()
243}
244
245/// Live forwarder tasks across all tunnels (drain assertions).
246pub fn tunnel_live_tasks() -> u64 {
247    tunnels()
248        .lock()
249        .unwrap()
250        .map
251        .values()
252        .map(|e| e.live_tasks.load(Ordering::Relaxed))
253        .sum()
254}
255
256/// Drop every cached tunnel and retire each one in order (tests).
257pub fn invalidate_all() {
258    let entries: Vec<Arc<TunnelEntry>> = {
259        let mut state = tunnels().lock().unwrap();
260        state.map.drain().map(|(_, e)| e).collect()
261    };
262    for entry in entries {
263        retire(&entry);
264    }
265}
266
267/// Retire one tunnel in the fixed order: mark draining (no new channel
268/// creation) → evict the MySQL pools keyed to this generation → stop
269/// the listener → let channel tasks wind down → disconnect the SSH
270/// session.
271fn retire(entry: &TunnelEntry) {
272    entry.draining.store(true, Ordering::SeqCst);
273    let evicted = super::pool::evict_by_generation(entry.generation);
274    if evicted > 0 {
275        eprintln!(
276            "[sequel-mcp] ssh tunnel gen {}: evicted {evicted} associated MySQL pool(s)",
277            entry.generation
278        );
279    }
280    let _ = entry.shutdown.try_send(());
281    // Disconnect the SSH session; outstanding channel copies error out
282    // and their tasks finish. Fire-and-forget with a bounded handle.
283    let handle = Arc::clone(&entry.handle);
284    tokio::spawn(async move {
285        let _ = handle
286            .disconnect(russh::Disconnect::ByApplication, "tunnel retired", "en")
287            .await;
288    });
289}
290
291impl SshTunnel {
292    fn auth_method_label(&self) -> &'static str {
293        match self.auth_method {
294            SshAuthMethod::Password => "password",
295            SshAuthMethod::Key => "key",
296        }
297    }
298}
299
300/// 16-hex-char content stamp of the known_hosts file (empty string when
301/// no explicit file is configured — the default path's absence is not a
302/// distinguishing identity).
303fn known_hosts_stamp(path: Option<&std::path::Path>) -> String {
304    let Some(path) = path else {
305        return String::new();
306    };
307    let mut h = sha2::Sha256::new();
308    match std::fs::read(path) {
309        Ok(bytes) => h.update(&bytes),
310        // Unreadable files fail closed in the checked loader; the stamp
311        // only needs to be deterministic per file identity.
312        Err(_) => h.update(format!("unreadable:{}", path.display())),
313    }
314    format!(
315        "{:016x}",
316        u64::from_be_bytes(h.finalize()[..8].try_into().expect("8 bytes"))
317    )
318}
319
320/// In-process credential generation for the SSH secret (never the
321/// secret itself): a process-keyed HMAC digest, hex-encoded for keying.
322fn ssh_credential_fragment(ssh_password: Option<&str>) -> String {
323    let cred = super::pool::CredentialGeneration::derive(ssh_password.unwrap_or(""));
324    cred.key_fragment()
325}
326
327#[allow(clippy::too_many_arguments)]
328fn tunnel_key(
329    conn_name: &str,
330    ssh: &SshTunnel,
331    ssh_password: Option<&str>,
332    target: &str,
333    port: u16,
334    policy_revision: u64,
335    kh_stamp: &str,
336) -> String {
337    let target_endpoint = format!("{target}:{port}");
338    let bridge = ssh
339        .docker
340        .as_ref()
341        .map(|d| format!("{}:{}", d.container, d.bridge_tool.as_str()))
342        .unwrap_or_default();
343    format!(
344        "{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}",
345        conn_name,
346        ssh.host,
347        ssh.port,
348        ssh.user,
349        ssh.auth_method_label(),
350        kh_stamp,
351        ssh_credential_fragment(ssh_password),
352        policy_revision,
353        target_endpoint,
354        bridge
355    )
356}
357
358/// Return a lease on a (reused or freshly established) SSH tunnel for
359/// `conn_name`'s SSH config, forwarding to `target_host:target_port` as
360/// reachable from the bastion. Concurrent first callers coalesce on one
361/// establishment. The lease's generation MUST be carried into the MySQL
362/// pool key (`verified_pool(.., tunnel_generation)`).
363pub async fn tunnel_endpoint(
364    conn_name: &str,
365    ssh: &SshTunnel,
366    ssh_password: Option<&str>,
367    target_host: &str,
368    target_port: u16,
369    policy_revision: u64,
370) -> Result<TunnelLease, SshError> {
371    // Fail-closed test-mode gate on the BASTION endpoint before any
372    // socket is opened (docker bastions publish on loopback).
373    crate::app::test_mode::check_mysql_endpoint(&ssh.host, ssh.port).map_err(SshError::Setup)?;
374
375    let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
376    let kh_stamp = known_hosts_stamp(kh_path);
377    let key = tunnel_key(
378        conn_name,
379        ssh,
380        ssh_password,
381        target_host,
382        target_port,
383        policy_revision,
384        &kh_stamp,
385    );
386
387    // Fast path: a live tunnel for exactly this identity.
388    {
389        let state = tunnels().lock().unwrap();
390        if let Some(entry) = state.map.get(&key)
391            && !entry.is_dead()
392            && !entry.draining.load(Ordering::SeqCst)
393        {
394            entry.touch();
395            return Ok(TunnelLease {
396                host: "127.0.0.1".into(),
397                port: entry.local.port(),
398                generation: entry.generation,
399            });
400        }
401    }
402
403    // Coalescing: one establishment per key even under a thundering
404    // herd — everyone serializes on the per-key guard, and waits find
405    // the winner's entry on the double-check.
406    let guard = {
407        let mut state = tunnels().lock().unwrap();
408        Arc::clone(state.inflight.entry(key.clone()).or_default())
409    };
410    let _hold = guard.lock().await;
411
412    {
413        let state = tunnels().lock().unwrap();
414        if let Some(entry) = state.map.get(&key)
415            && !entry.is_dead()
416            && !entry.draining.load(Ordering::SeqCst)
417        {
418            entry.touch();
419            return Ok(TunnelLease {
420                host: "127.0.0.1".into(),
421                port: entry.local.port(),
422                generation: entry.generation,
423            });
424        }
425    }
426
427    let entry = establish(ssh, ssh_password, target_host, target_port).await?;
428
429    let mut state = tunnels().lock().unwrap();
430    state.inflight.remove(&key);
431    // Retire stale tunnels of the SAME connection (identity rotation:
432    // credential, known_hosts content, or policy revision changed).
433    let conn_prefix = format!("{conn_name}\u{1}");
434    let stale: Vec<String> = state
435        .map
436        .keys()
437        .filter(|k| k.as_str() != key.as_str() && k.starts_with(&conn_prefix))
438        .cloned()
439        .collect();
440    for k in stale {
441        if let Some(e) = state.map.remove(&k) {
442            retire(&e);
443        }
444    }
445    // Drop dead entries; enforce the LRU bound by retiring the
446    // least-recently-used victims.
447    let dead: Vec<String> = state
448        .map
449        .iter()
450        .filter(|(_, e)| e.is_dead())
451        .map(|(k, _)| k.clone())
452        .collect();
453    for k in dead {
454        if let Some(e) = state.map.remove(&k) {
455            retire(&e);
456        }
457    }
458    while state.map.len() >= MAX_TUNNELS {
459        let victim = state
460            .map
461            .iter()
462            .min_by_key(|(_, e)| *e.last_used.lock().unwrap())
463            .map(|(k, _)| k.clone());
464        let Some(victim) = victim else { break };
465        if let Some(e) = state.map.remove(&victim) {
466            eprintln!(
467                "[sequel-mcp] ssh tunnel cache full ({}); retiring LRU entry",
468                MAX_TUNNELS
469            );
470            retire(&e);
471        }
472    }
473    let generation = entry.generation;
474    let port = entry.local.port();
475    state.map.insert(key, entry);
476    Ok(TunnelLease {
477        host: "127.0.0.1".into(),
478        port,
479        generation,
480    })
481}
482
483/// When `ssh.docker` is configured, per-connection forwarding runs over
484/// an SSH **exec** channel (`docker exec -i <container> <tool> …`)
485/// instead of direct-tcpip — the bridge that works even where sshd
486/// denies TCP forwarding. All argv components are validated (no shell,
487/// no spaces), so the exec command line is a plain join.
488fn bridge_command(
489    ssh: &SshTunnel,
490    target_host: &str,
491    target_port: u16,
492) -> Result<Option<String>, SshError> {
493    let Some(docker) = &ssh.docker else {
494        return Ok(None);
495    };
496    let argv = super::docker::bridge_argv(
497        &docker.container,
498        docker.bridge_tool,
499        target_host,
500        target_port,
501    )
502    .map_err(|e| SshError::Setup(format!("docker bridge: {e}")))?;
503    Ok(Some(argv.join(" ")))
504}
505
506async fn establish(
507    ssh: &SshTunnel,
508    ssh_password: Option<&str>,
509    target_host: &str,
510    target_port: u16,
511) -> Result<Arc<TunnelEntry>, SshError> {
512    let bridge_cmd = bridge_command(ssh, target_host, target_port)?;
513    let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
514    let entries =
515        known_hosts::load_known_hosts_checked(kh_path).map_err(|e| SshError::HostKey {
516            host: ssh.host.clone(),
517            port: ssh.port,
518            reason: e,
519        })?;
520    let logs = std::sync::Arc::new(Mutex::new(Vec::new()));
521    let handler = HostKeyHandler {
522        policy: ssh.host_key_policy.unwrap_or(SshHostKeyPolicy::Lenient),
523        entries,
524        host: ssh.host.clone(),
525        port: ssh.port,
526        logs: logs.clone(),
527    };
528
529    let keepalive = keepalive_interval();
530    let config = Arc::new(client::Config {
531        keepalive_interval: Some(keepalive),
532        keepalive_max: 3,
533        nodelay: true,
534        ..client::Config::default()
535    });
536
537    let addr = (ssh.host.as_str(), ssh.port);
538    let session = async {
539        let mut handle = client::connect(config, addr, handler)
540            .await
541            .map_err(|e| SshError::Transport(format!("connect to bastion: {e}")))?;
542        // Host-key commentary (TOFU accepts, migration-compat warnings)
543        // must reach the operator, not rot in a buffer.
544        for line in logs.lock().unwrap().drain(..) {
545            eprintln!("[sequel-mcp] SSH {}:{} {line}", ssh.host, ssh.port);
546        }
547        // Authenticate with EXACTLY the configured method — no silent
548        // fallback between password and key. The secret doubles as the
549        // private-key passphrase under key auth.
550        let mut key_algorithm_was_rsa = false;
551        let auth = match ssh.auth_method {
552            SshAuthMethod::Password => {
553                let Some(password) = ssh_password else {
554                    return Err(SshError::Setup(format!(
555                        "no SSH password stored for {:?} (expected under \"<connection>::ssh\")",
556                        ssh.user
557                    )));
558                };
559                handle
560                    .authenticate_password(ssh.user.clone(), password)
561                    .await
562                    .map_err(|e| SshError::Transport(format!("password auth transport: {e}")))?
563            }
564            SshAuthMethod::Key => {
565                let Some(path) = &ssh.private_key_path else {
566                    return Err(SshError::Setup(
567                        "key auth configured without privateKeyPath".into(),
568                    ));
569                };
570                let expanded = crate::app::paths::expand_tilde(path);
571                let key = russh::keys::load_secret_key(&expanded, ssh_password)
572                    .map_err(|e| SshError::Setup(format!("load private key: {e}")))?;
573                use russh::keys::Algorithm;
574                match key.algorithm() {
575                    Algorithm::Ed25519 | Algorithm::Rsa { .. } | Algorithm::Ecdsa { .. } => {}
576                    other => {
577                        return Err(SshError::Setup(format!(
578                            "unsupported private key algorithm {other:?} (supported: ed25519, rsa, ecdsa)"
579                        )));
580                    }
581                }
582                // Let the server's advertised algorithms pick the RSA
583                // signature hash (SHA-512 preferred, SHA-256 next;
584                // legacy ssh-rsa/SHA-1 is never selected).
585                key_algorithm_was_rsa = key.algorithm().is_rsa();
586                let hash_alg = if key_algorithm_was_rsa {
587                    match handle.best_supported_rsa_hash().await {
588                        Ok(best) => best.unwrap_or(Some(russh::keys::HashAlg::Sha256)),
589                        Err(_) => Some(russh::keys::HashAlg::Sha256),
590                    }
591                } else {
592                    None
593                };
594                handle
595                    .authenticate_publickey(
596                        ssh.user.clone(),
597                        russh::keys::PrivateKeyWithHashAlg::new(Arc::new(key), hash_alg),
598                    )
599                    .await
600                    .map_err(|e| {
601                        // A signature we cannot produce is NOT a
602                        // transport problem: name it as an aborted
603                        // auth so it cannot masquerade as a rejected
604                        // credential (the 0.10.2 incident).
605                        if is_signature_algorithm_error(&e) {
606                            SshError::AuthAborted {
607                                host: ssh.host.clone(),
608                                port: ssh.port,
609                                user: ssh.user.clone(),
610                                reason: auth_abort_reason(key_algorithm_was_rsa),
611                            }
612                        } else {
613                            SshError::Transport(format!("publickey auth transport: {e}"))
614                        }
615                    })?
616            }
617        };
618        match auth {
619            russh::client::AuthResult::Success => {}
620            russh::client::AuthResult::Failure {
621                remaining_methods, ..
622            } => {
623                // A genuine USERAUTH_FAILURE carries the server's
624                // remaining-method list. russh 0.62 ALSO maps "the
625                // session died during auth" — reply channel closed,
626                // e.g. the signature could not be produced — to
627                // Failure, but with an EMPTY method set. Collapsing
628                // both into Auth is what made the 0.10.2 RSA signing
629                // failure read as a rejected credential.
630                if remaining_methods.is_empty() {
631                    return Err(SshError::AuthAborted {
632                        host: ssh.host.clone(),
633                        port: ssh.port,
634                        user: ssh.user.clone(),
635                        reason: auth_abort_reason(key_algorithm_was_rsa),
636                    });
637                }
638                return Err(SshError::Auth {
639                    host: ssh.host.clone(),
640                    port: ssh.port,
641                    user: ssh.user.clone(),
642                });
643            }
644        }
645        Ok(handle)
646    };
647    let handle = tokio::time::timeout(SSH_CONNECT_TIMEOUT, session)
648        .await
649        .map_err(|_| {
650            SshError::Transport(format!(
651                "ssh session timeout after {}s",
652                SSH_CONNECT_TIMEOUT.as_secs()
653            ))
654        })??;
655
656    // Local loopback forwarder: one fresh direct-tcpip channel per
657    // accepted TCP connection, bounded by CHANNEL_OPEN_TIMEOUT so a
658    // stale session fails the connection instead of hanging it.
659    let listener = TcpListener::bind(("127.0.0.1", 0))
660        .await
661        .map_err(|e| SshError::Setup(format!("local bind: {e}")))?;
662    let local = listener
663        .local_addr()
664        .map_err(|e| SshError::Setup(format!("local addr: {e}")))?;
665
666    let generation = {
667        let mut state = tunnels().lock().unwrap();
668        let assigned = state.next_generation;
669        state.next_generation += 1;
670        assigned
671    };
672
673    let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
674    let forward_host = target_host.to_string();
675    let forward_port = u32::from(target_port);
676    let shared = Arc::new(handle);
677    let accept_session = Arc::clone(&shared);
678    let draining = Arc::new(AtomicBool::new(false));
679    let live_tasks: Arc<AtomicU64> = Arc::new(AtomicU64::new(0));
680    let accept_draining = Arc::clone(&draining);
681    let accept_tasks = Arc::clone(&live_tasks);
682    tokio::spawn(async move {
683        loop {
684            tokio::select! {
685                _ = shutdown_rx.recv() => break,
686                accepted = listener.accept() => {
687                    let Ok((socket, _peer)) = accepted else { break };
688                    if accept_draining.load(Ordering::SeqCst) || accept_session.is_closed() {
689                        break;
690                    }
691                    let session = Arc::clone(&accept_session);
692                    let host = forward_host.clone();
693                    let tasks = Arc::clone(&accept_tasks);
694                    accept_tasks.fetch_add(1, Ordering::Relaxed);
695                    let bridge_cmd = bridge_cmd.clone();
696                    tokio::spawn(async move {
697                        let target = format!("{host}:{forward_port}");
698                        // Bridge connections open a session channel and
699                        // exec the (validated, space-free) docker bridge
700                        // argv; direct connections use direct-tcpip.
701                        let channel = tokio::time::timeout(
702                            CHANNEL_OPEN_TIMEOUT,
703                            async {
704                                match &bridge_cmd {
705                                    Some(cmd) => {
706                                        let ch = session.channel_open_session().await?;
707                                        ch.exec(true, cmd.as_str()).await?;
708                                        Ok::<_, russh::Error>(ch)
709                                    }
710                                    None => session
711                                        .channel_open_direct_tcpip(
712                                            host,
713                                            forward_port,
714                                            "127.0.0.1",
715                                            0,
716                                        )
717                                        .await,
718                                }
719                            },
720                        )
721                        .await;
722                        match channel {
723                            Ok(Ok(channel)) => {
724                                // ChannelStream is the full tokio-IO view
725                                // of a channel; split gives owned halves
726                                // for the bidirectional copy.
727                                let (mut chan_r, mut chan_w) =
728                                    tokio::io::split(channel.into_stream());
729                                let (mut sock_r, mut sock_w) = socket.into_split();
730                                let a = tokio::io::copy(&mut sock_r, &mut chan_w);
731                                let b = tokio::io::copy(&mut chan_r, &mut sock_w);
732                                let _ = tokio::join!(a, b);
733                            }
734                            Ok(Err(e)) => {
735                                eprintln!(
736                                    "[sequel-mcp] ssh forwarder: channel open to {target} failed: {e:?}"
737                                );
738                                // Dropping the socket closes it: the MySQL
739                                // handshake fails promptly.
740                            }
741                            Err(_) => {
742                                eprintln!(
743                                    "[sequel-mcp] ssh forwarder: channel open to {target} timed out after {}s (stale session?)",
744                                    CHANNEL_OPEN_TIMEOUT.as_secs()
745                                );
746                            }
747                        }
748                        tasks.fetch_sub(1, Ordering::Relaxed);
749                    });
750                }
751            }
752        }
753    });
754
755    Ok(Arc::new(TunnelEntry {
756        local,
757        handle: shared,
758        shutdown: shutdown_tx,
759        generation,
760        draining,
761        last_used: Mutex::new(std::time::Instant::now()),
762        live_tasks,
763    }))
764}
765
766#[cfg(test)]
767mod tests {
768    use super::*;
769    use crate::config::SshTunnel;
770
771    fn sample_tunnel() -> SshTunnel {
772        SshTunnel {
773            host: "bastion".into(),
774            port: 22,
775            user: "sshuser".into(),
776            auth_method: SshAuthMethod::Password,
777            ..SshTunnel::default()
778        }
779    }
780
781    #[test]
782    fn tunnel_key_separates_identity() {
783        let ssh = sample_tunnel();
784        let a = tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp");
785        let mut ssh2 = ssh.clone();
786        ssh2.user = "other".into();
787        let b = tunnel_key("c1", &ssh2, Some("pw"), "db", 3306, 1, "stamp");
788        assert_ne!(a, b);
789        assert_ne!(
790            a,
791            tunnel_key("c1", &ssh, Some("pw"), "other", 3306, 1, "stamp")
792        );
793        assert_ne!(
794            a,
795            tunnel_key("c1", &ssh, Some("pw"), "db", 3307, 1, "stamp")
796        );
797        assert_ne!(
798            a,
799            tunnel_key("c2", &ssh, Some("pw"), "db", 3306, 1, "stamp")
800        );
801        // Rotation inputs: credential, known_hosts content, revision.
802        assert_ne!(
803            a,
804            tunnel_key("c1", &ssh, Some("different"), "db", 3306, 1, "stamp")
805        );
806        assert_ne!(
807            a,
808            tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 2, "stamp")
809        );
810        assert_ne!(
811            a,
812            tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp2")
813        );
814        // Bridge identity (container + tool) is part of the key.
815        let mut bridged = ssh.clone();
816        bridged.docker = Some(crate::config::SshDocker {
817            container: "db".into(),
818            bridge_tool: crate::config::BridgeTool::Nc,
819        });
820        assert_ne!(
821            a,
822            tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
823        );
824        let mut bridged2 = bridged.clone();
825        bridged2.docker = Some(crate::config::SshDocker {
826            container: "db".into(),
827            bridge_tool: crate::config::BridgeTool::Socat,
828        });
829        assert_ne!(
830            tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp"),
831            tunnel_key("c1", &bridged2, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
832        );
833    }
834
835    #[test]
836    fn bridge_command_forms() {
837        use crate::config::{BridgeTool, SshDocker};
838        let mk = |tool| SshTunnel {
839            host: "bastion".into(),
840            port: 22,
841            user: "u".into(),
842            auth_method: SshAuthMethod::Key,
843            docker: Some(SshDocker {
844                container: "db-1".into(),
845                bridge_tool: tool,
846            }),
847            ..SshTunnel::default()
848        };
849        assert_eq!(
850            bridge_command(&mk(BridgeTool::Nc), "127.0.0.1", 3306)
851                .unwrap()
852                .unwrap(),
853            "docker exec -i db-1 nc 127.0.0.1 3306"
854        );
855        assert_eq!(
856            bridge_command(&mk(BridgeTool::Ncat), "127.0.0.1", 3306)
857                .unwrap()
858                .unwrap(),
859            "docker exec -i db-1 ncat 127.0.0.1 3306"
860        );
861        assert_eq!(
862            bridge_command(&mk(BridgeTool::Socat), "127.0.0.1", 3306)
863                .unwrap()
864                .unwrap(),
865            "docker exec -i db-1 socat - TCP:127.0.0.1:3306"
866        );
867        // No docker config: direct-tcpip path.
868        let plain = SshTunnel::default();
869        assert_eq!(bridge_command(&plain, "db", 3306).unwrap(), None);
870        // Invalid inputs fail typed before any connection.
871        let mut bad = mk(BridgeTool::Nc);
872        bad.docker = Some(SshDocker {
873            container: "bad name!".into(),
874            bridge_tool: BridgeTool::Nc,
875        });
876        assert!(bridge_command(&bad, "127.0.0.1", 3306).is_err());
877    }
878
879    #[test]
880    fn credential_fragment_hides_the_secret() {
881        let f = ssh_credential_fragment(Some("top-secret-password"));
882        assert_eq!(f.len(), 32, "hex of a 16-byte digest");
883        assert!(!f.contains("top-secret"));
884        assert_ne!(f, ssh_credential_fragment(Some("other")));
885        // A missing and an empty secret are the same non-secret here.
886        assert_eq!(
887            ssh_credential_fragment(None),
888            ssh_credential_fragment(Some(""))
889        );
890    }
891
892    #[test]
893    fn known_hosts_stamp_tracks_content() {
894        let dir = tempfile::TempDir::new().unwrap();
895        let f = dir.path().join("kh");
896        std::fs::write(&f, b"content-a").unwrap();
897        let s1 = known_hosts_stamp(Some(&f));
898        let s2 = known_hosts_stamp(Some(&f));
899        assert_eq!(s1, s2, "stable per content");
900        std::fs::write(&f, b"content-b").unwrap();
901        assert_ne!(s1, known_hosts_stamp(Some(&f)));
902        assert_eq!(known_hosts_stamp(None), "");
903    }
904
905    #[test]
906    fn auth_abort_reason_never_claims_rejection() {
907        // RSA: the reason must point at signing support (the feature
908        // gap), never at the credential. Wording differs between a
909        // feature-less build ("cannot SIGN") and a full build ("NOT a
910        // rejected credential"), so accept either shape.
911        let rsa = auth_abort_reason(true);
912        assert!(rsa.to_lowercase().contains("rsa"), "{rsa}");
913        assert!(
914            rsa.to_lowercase().contains("not rejected")
915                || rsa.to_lowercase().contains("not a rejected credential")
916                || rsa.to_lowercase().contains("cannot sign"),
917            "must not read as a rejected credential: {rsa}"
918        );
919        // Non-RSA: generic transport wording, no RSA red herring.
920        let other = auth_abort_reason(false);
921        assert!(
922            other.to_lowercase().contains("not a rejected credential"),
923            "{other}"
924        );
925        assert!(!other.to_lowercase().contains("rsa"), "{other}");
926    }
927
928    #[test]
929    fn signature_algorithm_error_discriminates() {
930        // The exact ssh-key shape from the 0.10.2 incident: an RSA
931        // signature whose negotiated hash never reached the signer.
932        let unsupported = russh::Error::SshKey(russh::keys::ssh_key::Error::AlgorithmUnsupported {
933            algorithm: russh::keys::Algorithm::Rsa { hash: None },
934        });
935        assert!(is_signature_algorithm_error(&unsupported));
936        // Other SshKey errors stay transport-classified.
937        let other = russh::Error::SshKey(russh::keys::ssh_key::Error::AlgorithmUnknown);
938        assert!(!is_signature_algorithm_error(&other));
939    }
940}