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