use crate::config::{SshAuthMethod, SshHostKeyPolicy, SshTunnel};
use crate::sql::known_hosts;
use russh::client::{self, Handle};
use russh::keys::PublicKeyBase64;
use sha2::Digest as _;
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use thiserror::Error;
use tokio::net::TcpListener;
use tokio::sync::mpsc;
pub const SSH_CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
pub const CHANNEL_OPEN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
pub const MAX_TUNNELS: usize = 8;
pub fn keepalive_interval() -> std::time::Duration {
let secs = std::env::var("SEQUEL_MCP_SSH_KEEPALIVE_SECS")
.ok()
.and_then(|v| v.trim().parse::<u64>().ok())
.unwrap_or(30)
.clamp(1, 300);
std::time::Duration::from_secs(secs)
}
#[derive(Debug, Error)]
pub enum SshError {
#[error("{0}")]
Transport(String),
#[error("authentication as {user:?} failed on {host}:{port}")]
Auth {
host: String,
port: u16,
user: String,
},
#[error("authentication aborted as {user:?} on {host}:{port}: {reason}")]
AuthAborted {
host: String,
port: u16,
user: String,
reason: String,
},
#[error("host key rejected for {host}:{port}: {reason}")]
HostKey {
host: String,
port: u16,
reason: String,
},
#[error("tunnel setup: {0}")]
Setup(String),
}
fn auth_abort_reason(key_is_rsa: bool) -> String {
if key_is_rsa && cfg!(not(feature = "rsa")) {
"no server verdict was delivered — the SSH session ended mid-auth. \
This build cannot SIGN RSA keys: russh's `rsa` feature is missing \
(enabled by sequel-mcp's default `rsa` feature). The key itself is \
fine and the server had not rejected it"
.into()
} else if key_is_rsa {
"no server verdict was delivered — the SSH session ended mid-auth, a \
signing/transport failure and NOT a rejected credential. For RSA keys \
the classic cause is a build without russh's `rsa` feature; \
RUST_LOG=russh=debug shows the underlying error"
.into()
} else {
"no server verdict was delivered — the SSH session ended mid-auth, a \
transport failure and NOT a rejected credential; RUST_LOG=russh=debug \
shows the underlying error"
.into()
}
}
fn is_signature_algorithm_error(e: &russh::Error) -> bool {
matches!(
e,
russh::Error::SshKey(russh::keys::ssh_key::Error::AlgorithmUnsupported { .. })
)
}
struct HostKeyHandler {
policy: SshHostKeyPolicy,
entries: Vec<known_hosts::KnownHostEntry>,
host: String,
port: u16,
logs: std::sync::Arc<Mutex<Vec<String>>>,
}
impl client::Handler for HostKeyHandler {
type Error = russh::Error;
async fn check_server_key(
&mut self,
server_public_key: &russh::keys::PublicKeyOrCertificate,
) -> Result<bool, Self::Error> {
let raw = match server_public_key {
russh::keys::PublicKeyOrCertificate::PublicKey { key, .. } => key.public_key_bytes(),
russh::keys::PublicKeyOrCertificate::Certificate(_) => return Ok(false),
};
let mut logs: Vec<String> = Vec::new();
let decision = known_hosts::decide_host_key(
self.policy,
&self.host,
self.port,
&self.entries,
&raw,
&mut |line| logs.push(line),
);
self.logs.lock().unwrap().extend(logs);
Ok(decision == known_hosts::HostKeyDecision::Accept)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TunnelLease {
pub host: String,
pub port: u16,
pub generation: u64,
}
struct TunnelEntry {
local: std::net::SocketAddr,
handle: Arc<Handle<HostKeyHandler>>,
shutdown: mpsc::Sender<()>,
generation: u64,
draining: Arc<AtomicBool>,
last_used: Mutex<std::time::Instant>,
live_tasks: Arc<AtomicU64>,
}
impl TunnelEntry {
fn is_dead(&self) -> bool {
self.handle.is_closed()
}
fn touch(&self) {
*self.last_used.lock().unwrap() = std::time::Instant::now();
}
}
struct Tunnels {
map: HashMap<String, Arc<TunnelEntry>>,
next_generation: u64,
inflight: HashMap<String, Arc<tokio::sync::Mutex<()>>>,
}
static TUNNELS: OnceLock<Mutex<Tunnels>> = OnceLock::new();
fn tunnels() -> &'static Mutex<Tunnels> {
TUNNELS.get_or_init(|| {
Mutex::new(Tunnels {
map: HashMap::new(),
next_generation: 1,
inflight: HashMap::new(),
})
})
}
pub fn tunnel_count() -> usize {
tunnels().lock().unwrap().map.len()
}
pub fn tunnel_connection_names() -> Vec<String> {
tunnels()
.lock()
.unwrap()
.map
.keys()
.map(|k| k.split('\u{1}').next().unwrap_or("").to_string())
.collect()
}
pub fn tunnel_live_tasks() -> u64 {
tunnels()
.lock()
.unwrap()
.map
.values()
.map(|e| e.live_tasks.load(Ordering::Relaxed))
.sum()
}
pub fn invalidate_all() {
let entries: Vec<Arc<TunnelEntry>> = {
let mut state = tunnels().lock().unwrap();
state.map.drain().map(|(_, e)| e).collect()
};
for entry in entries {
retire(&entry);
}
}
fn retire(entry: &TunnelEntry) {
entry.draining.store(true, Ordering::SeqCst);
let evicted = super::pool::evict_by_generation(entry.generation);
if evicted > 0 {
eprintln!(
"[sequel-mcp] ssh tunnel gen {}: evicted {evicted} associated MySQL pool(s)",
entry.generation
);
}
let _ = entry.shutdown.try_send(());
let handle = Arc::clone(&entry.handle);
tokio::spawn(async move {
let _ = handle
.disconnect(russh::Disconnect::ByApplication, "tunnel retired", "en")
.await;
});
}
impl SshTunnel {
fn auth_method_label(&self) -> &'static str {
match self.auth_method {
SshAuthMethod::Password => "password",
SshAuthMethod::Key => "key",
}
}
}
fn known_hosts_stamp(path: Option<&std::path::Path>) -> String {
let Some(path) = path else {
return String::new();
};
let mut h = sha2::Sha256::new();
match std::fs::read(path) {
Ok(bytes) => h.update(&bytes),
Err(_) => h.update(format!("unreadable:{}", path.display())),
}
format!(
"{:016x}",
u64::from_be_bytes(h.finalize()[..8].try_into().expect("8 bytes"))
)
}
fn ssh_credential_fragment(ssh_password: Option<&str>) -> String {
let cred = super::pool::CredentialGeneration::derive(ssh_password.unwrap_or(""));
cred.key_fragment()
}
#[allow(clippy::too_many_arguments)]
fn tunnel_key(
conn_name: &str,
ssh: &SshTunnel,
ssh_password: Option<&str>,
target: &str,
port: u16,
policy_revision: u64,
kh_stamp: &str,
) -> String {
let target_endpoint = format!("{target}:{port}");
let bridge = ssh
.docker
.as_ref()
.map(|d| format!("{}:{}", d.container, d.bridge_tool.as_str()))
.unwrap_or_default();
format!(
"{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}",
conn_name,
ssh.host,
ssh.port,
ssh.user,
ssh.auth_method_label(),
kh_stamp,
ssh_credential_fragment(ssh_password),
policy_revision,
target_endpoint,
bridge
)
}
pub async fn tunnel_endpoint(
conn_name: &str,
ssh: &SshTunnel,
ssh_password: Option<&str>,
target_host: &str,
target_port: u16,
policy_revision: u64,
) -> Result<TunnelLease, SshError> {
crate::app::test_mode::check_mysql_endpoint(&ssh.host, ssh.port).map_err(SshError::Setup)?;
let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
let kh_stamp = known_hosts_stamp(kh_path);
let key = tunnel_key(
conn_name,
ssh,
ssh_password,
target_host,
target_port,
policy_revision,
&kh_stamp,
);
{
let state = tunnels().lock().unwrap();
if let Some(entry) = state.map.get(&key)
&& !entry.is_dead()
&& !entry.draining.load(Ordering::SeqCst)
{
entry.touch();
return Ok(TunnelLease {
host: "127.0.0.1".into(),
port: entry.local.port(),
generation: entry.generation,
});
}
}
let guard = {
let mut state = tunnels().lock().unwrap();
Arc::clone(state.inflight.entry(key.clone()).or_default())
};
let _hold = guard.lock().await;
{
let state = tunnels().lock().unwrap();
if let Some(entry) = state.map.get(&key)
&& !entry.is_dead()
&& !entry.draining.load(Ordering::SeqCst)
{
entry.touch();
return Ok(TunnelLease {
host: "127.0.0.1".into(),
port: entry.local.port(),
generation: entry.generation,
});
}
}
let entry = establish(ssh, ssh_password, target_host, target_port).await?;
let mut state = tunnels().lock().unwrap();
state.inflight.remove(&key);
let conn_prefix = format!("{conn_name}\u{1}");
let stale: Vec<String> = state
.map
.keys()
.filter(|k| k.as_str() != key.as_str() && k.starts_with(&conn_prefix))
.cloned()
.collect();
for k in stale {
if let Some(e) = state.map.remove(&k) {
retire(&e);
}
}
let dead: Vec<String> = state
.map
.iter()
.filter(|(_, e)| e.is_dead())
.map(|(k, _)| k.clone())
.collect();
for k in dead {
if let Some(e) = state.map.remove(&k) {
retire(&e);
}
}
while state.map.len() >= MAX_TUNNELS {
let victim = state
.map
.iter()
.min_by_key(|(_, e)| *e.last_used.lock().unwrap())
.map(|(k, _)| k.clone());
let Some(victim) = victim else { break };
if let Some(e) = state.map.remove(&victim) {
eprintln!(
"[sequel-mcp] ssh tunnel cache full ({}); retiring LRU entry",
MAX_TUNNELS
);
retire(&e);
}
}
let generation = entry.generation;
let port = entry.local.port();
state.map.insert(key, entry);
Ok(TunnelLease {
host: "127.0.0.1".into(),
port,
generation,
})
}
fn bridge_command(
ssh: &SshTunnel,
target_host: &str,
target_port: u16,
) -> Result<Option<String>, SshError> {
let Some(docker) = &ssh.docker else {
return Ok(None);
};
let argv = super::docker::bridge_argv(
&docker.container,
docker.bridge_tool,
target_host,
target_port,
)
.map_err(|e| SshError::Setup(format!("docker bridge: {e}")))?;
Ok(Some(argv.join(" ")))
}
async fn establish(
ssh: &SshTunnel,
ssh_password: Option<&str>,
target_host: &str,
target_port: u16,
) -> Result<Arc<TunnelEntry>, SshError> {
let bridge_cmd = bridge_command(ssh, target_host, target_port)?;
let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
let entries =
known_hosts::load_known_hosts_checked(kh_path).map_err(|e| SshError::HostKey {
host: ssh.host.clone(),
port: ssh.port,
reason: e,
})?;
let logs = std::sync::Arc::new(Mutex::new(Vec::new()));
let handler = HostKeyHandler {
policy: ssh.host_key_policy.unwrap_or(SshHostKeyPolicy::Lenient),
entries,
host: ssh.host.clone(),
port: ssh.port,
logs: logs.clone(),
};
let keepalive = keepalive_interval();
let config = Arc::new(client::Config {
keepalive_interval: Some(keepalive),
keepalive_max: 3,
nodelay: true,
..client::Config::default()
});
let addr = (ssh.host.as_str(), ssh.port);
let session = async {
let mut handle = client::connect(config, addr, handler)
.await
.map_err(|e| SshError::Transport(format!("connect to bastion: {e}")))?;
for line in logs.lock().unwrap().drain(..) {
eprintln!("[sequel-mcp] SSH {}:{} {line}", ssh.host, ssh.port);
}
let mut key_algorithm_was_rsa = false;
let auth = match ssh.auth_method {
SshAuthMethod::Password => {
let Some(password) = ssh_password else {
return Err(SshError::Setup(format!(
"no SSH password stored for {:?} (expected under \"<connection>::ssh\")",
ssh.user
)));
};
handle
.authenticate_password(ssh.user.clone(), password)
.await
.map_err(|e| SshError::Transport(format!("password auth transport: {e}")))?
}
SshAuthMethod::Key => {
let Some(path) = &ssh.private_key_path else {
return Err(SshError::Setup(
"key auth configured without privateKeyPath".into(),
));
};
let expanded = crate::app::paths::expand_tilde(path);
let key = russh::keys::load_secret_key(&expanded, ssh_password)
.map_err(|e| SshError::Setup(format!("load private key: {e}")))?;
use russh::keys::Algorithm;
match key.algorithm() {
Algorithm::Ed25519 | Algorithm::Rsa { .. } | Algorithm::Ecdsa { .. } => {}
other => {
return Err(SshError::Setup(format!(
"unsupported private key algorithm {other:?} (supported: ed25519, rsa, ecdsa)"
)));
}
}
key_algorithm_was_rsa = key.algorithm().is_rsa();
let hash_alg = if key_algorithm_was_rsa {
match handle.best_supported_rsa_hash().await {
Ok(best) => best.unwrap_or(Some(russh::keys::HashAlg::Sha256)),
Err(_) => Some(russh::keys::HashAlg::Sha256),
}
} else {
None
};
handle
.authenticate_publickey(
ssh.user.clone(),
russh::keys::PrivateKeyWithHashAlg::new(Arc::new(key), hash_alg),
)
.await
.map_err(|e| {
if is_signature_algorithm_error(&e) {
SshError::AuthAborted {
host: ssh.host.clone(),
port: ssh.port,
user: ssh.user.clone(),
reason: auth_abort_reason(key_algorithm_was_rsa),
}
} else {
SshError::Transport(format!("publickey auth transport: {e}"))
}
})?
}
};
match auth {
russh::client::AuthResult::Success => {}
russh::client::AuthResult::Failure {
remaining_methods, ..
} => {
if remaining_methods.is_empty() {
return Err(SshError::AuthAborted {
host: ssh.host.clone(),
port: ssh.port,
user: ssh.user.clone(),
reason: auth_abort_reason(key_algorithm_was_rsa),
});
}
return Err(SshError::Auth {
host: ssh.host.clone(),
port: ssh.port,
user: ssh.user.clone(),
});
}
}
Ok(handle)
};
let handle = tokio::time::timeout(SSH_CONNECT_TIMEOUT, session)
.await
.map_err(|_| {
SshError::Transport(format!(
"ssh session timeout after {}s",
SSH_CONNECT_TIMEOUT.as_secs()
))
})??;
let listener = TcpListener::bind(("127.0.0.1", 0))
.await
.map_err(|e| SshError::Setup(format!("local bind: {e}")))?;
let local = listener
.local_addr()
.map_err(|e| SshError::Setup(format!("local addr: {e}")))?;
let generation = {
let mut state = tunnels().lock().unwrap();
let assigned = state.next_generation;
state.next_generation += 1;
assigned
};
let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
let forward_host = target_host.to_string();
let forward_port = u32::from(target_port);
let shared = Arc::new(handle);
let accept_session = Arc::clone(&shared);
let draining = Arc::new(AtomicBool::new(false));
let live_tasks: Arc<AtomicU64> = Arc::new(AtomicU64::new(0));
let accept_draining = Arc::clone(&draining);
let accept_tasks = Arc::clone(&live_tasks);
tokio::spawn(async move {
loop {
tokio::select! {
_ = shutdown_rx.recv() => break,
accepted = listener.accept() => {
let Ok((socket, _peer)) = accepted else { break };
if accept_draining.load(Ordering::SeqCst) || accept_session.is_closed() {
break;
}
let session = Arc::clone(&accept_session);
let host = forward_host.clone();
let tasks = Arc::clone(&accept_tasks);
accept_tasks.fetch_add(1, Ordering::Relaxed);
let bridge_cmd = bridge_cmd.clone();
tokio::spawn(async move {
let target = format!("{host}:{forward_port}");
let channel = tokio::time::timeout(
CHANNEL_OPEN_TIMEOUT,
async {
match &bridge_cmd {
Some(cmd) => {
let ch = session.channel_open_session().await?;
ch.exec(true, cmd.as_str()).await?;
Ok::<_, russh::Error>(ch)
}
None => session
.channel_open_direct_tcpip(
host,
forward_port,
"127.0.0.1",
0,
)
.await,
}
},
)
.await;
match channel {
Ok(Ok(channel)) => {
let (mut chan_r, mut chan_w) =
tokio::io::split(channel.into_stream());
let (mut sock_r, mut sock_w) = socket.into_split();
let a = tokio::io::copy(&mut sock_r, &mut chan_w);
let b = tokio::io::copy(&mut chan_r, &mut sock_w);
let _ = tokio::join!(a, b);
}
Ok(Err(e)) => {
eprintln!(
"[sequel-mcp] ssh forwarder: channel open to {target} failed: {e:?}"
);
}
Err(_) => {
eprintln!(
"[sequel-mcp] ssh forwarder: channel open to {target} timed out after {}s (stale session?)",
CHANNEL_OPEN_TIMEOUT.as_secs()
);
}
}
tasks.fetch_sub(1, Ordering::Relaxed);
});
}
}
}
});
Ok(Arc::new(TunnelEntry {
local,
handle: shared,
shutdown: shutdown_tx,
generation,
draining,
last_used: Mutex::new(std::time::Instant::now()),
live_tasks,
}))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::SshTunnel;
fn sample_tunnel() -> SshTunnel {
SshTunnel {
host: "bastion".into(),
port: 22,
user: "sshuser".into(),
auth_method: SshAuthMethod::Password,
..SshTunnel::default()
}
}
#[test]
fn tunnel_key_separates_identity() {
let ssh = sample_tunnel();
let a = tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp");
let mut ssh2 = ssh.clone();
ssh2.user = "other".into();
let b = tunnel_key("c1", &ssh2, Some("pw"), "db", 3306, 1, "stamp");
assert_ne!(a, b);
assert_ne!(
a,
tunnel_key("c1", &ssh, Some("pw"), "other", 3306, 1, "stamp")
);
assert_ne!(
a,
tunnel_key("c1", &ssh, Some("pw"), "db", 3307, 1, "stamp")
);
assert_ne!(
a,
tunnel_key("c2", &ssh, Some("pw"), "db", 3306, 1, "stamp")
);
assert_ne!(
a,
tunnel_key("c1", &ssh, Some("different"), "db", 3306, 1, "stamp")
);
assert_ne!(
a,
tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 2, "stamp")
);
assert_ne!(
a,
tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp2")
);
let mut bridged = ssh.clone();
bridged.docker = Some(crate::config::SshDocker {
container: "db".into(),
bridge_tool: crate::config::BridgeTool::Nc,
});
assert_ne!(
a,
tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
);
let mut bridged2 = bridged.clone();
bridged2.docker = Some(crate::config::SshDocker {
container: "db".into(),
bridge_tool: crate::config::BridgeTool::Socat,
});
assert_ne!(
tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp"),
tunnel_key("c1", &bridged2, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
);
}
#[test]
fn bridge_command_forms() {
use crate::config::{BridgeTool, SshDocker};
let mk = |tool| SshTunnel {
host: "bastion".into(),
port: 22,
user: "u".into(),
auth_method: SshAuthMethod::Key,
docker: Some(SshDocker {
container: "db-1".into(),
bridge_tool: tool,
}),
..SshTunnel::default()
};
assert_eq!(
bridge_command(&mk(BridgeTool::Nc), "127.0.0.1", 3306)
.unwrap()
.unwrap(),
"docker exec -i db-1 nc 127.0.0.1 3306"
);
assert_eq!(
bridge_command(&mk(BridgeTool::Ncat), "127.0.0.1", 3306)
.unwrap()
.unwrap(),
"docker exec -i db-1 ncat 127.0.0.1 3306"
);
assert_eq!(
bridge_command(&mk(BridgeTool::Socat), "127.0.0.1", 3306)
.unwrap()
.unwrap(),
"docker exec -i db-1 socat - TCP:127.0.0.1:3306"
);
let plain = SshTunnel::default();
assert_eq!(bridge_command(&plain, "db", 3306).unwrap(), None);
let mut bad = mk(BridgeTool::Nc);
bad.docker = Some(SshDocker {
container: "bad name!".into(),
bridge_tool: BridgeTool::Nc,
});
assert!(bridge_command(&bad, "127.0.0.1", 3306).is_err());
}
#[test]
fn credential_fragment_hides_the_secret() {
let f = ssh_credential_fragment(Some("top-secret-password"));
assert_eq!(f.len(), 32, "hex of a 16-byte digest");
assert!(!f.contains("top-secret"));
assert_ne!(f, ssh_credential_fragment(Some("other")));
assert_eq!(
ssh_credential_fragment(None),
ssh_credential_fragment(Some(""))
);
}
#[test]
fn known_hosts_stamp_tracks_content() {
let dir = tempfile::TempDir::new().unwrap();
let f = dir.path().join("kh");
std::fs::write(&f, b"content-a").unwrap();
let s1 = known_hosts_stamp(Some(&f));
let s2 = known_hosts_stamp(Some(&f));
assert_eq!(s1, s2, "stable per content");
std::fs::write(&f, b"content-b").unwrap();
assert_ne!(s1, known_hosts_stamp(Some(&f)));
assert_eq!(known_hosts_stamp(None), "");
}
#[test]
fn auth_abort_reason_never_claims_rejection() {
let rsa = auth_abort_reason(true);
assert!(rsa.to_lowercase().contains("rsa"), "{rsa}");
assert!(
rsa.to_lowercase().contains("not rejected")
|| rsa.to_lowercase().contains("not a rejected credential")
|| rsa.to_lowercase().contains("cannot sign"),
"must not read as a rejected credential: {rsa}"
);
let other = auth_abort_reason(false);
assert!(
other.to_lowercase().contains("not a rejected credential"),
"{other}"
);
assert!(!other.to_lowercase().contains("rsa"), "{other}");
}
#[test]
fn signature_algorithm_error_discriminates() {
let unsupported = russh::Error::SshKey(russh::keys::ssh_key::Error::AlgorithmUnsupported {
algorithm: russh::keys::Algorithm::Rsa { hash: None },
});
assert!(is_signature_algorithm_error(&unsupported));
let other = russh::Error::SshKey(russh::keys::ssh_key::Error::AlgorithmUnknown);
assert!(!is_signature_algorithm_error(&other));
}
}