1use 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
35pub const SSH_CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
37
38pub const CHANNEL_OPEN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
41
42pub const MAX_TUNNELS: usize = 8;
44
45pub fn keepalive_interval() -> std::time::Duration {
49 let secs = std::env::var("SEQUEL_MCP_SSH_KEEPALIVE_SECS")
50 .ok()
51 .and_then(|v| v.trim().parse::<u64>().ok())
52 .unwrap_or(30)
53 .clamp(1, 300);
54 std::time::Duration::from_secs(secs)
55}
56
57#[derive(Debug, Error)]
58pub enum SshError {
59 #[error("{0}")]
60 Transport(String),
61 #[error("authentication as {user:?} failed on {host}:{port}")]
62 Auth {
63 host: String,
64 port: u16,
65 user: String,
66 },
67 #[error("host key rejected for {host}:{port}: {reason}")]
68 HostKey {
69 host: String,
70 port: u16,
71 reason: String,
72 },
73 #[error("tunnel setup: {0}")]
74 Setup(String),
75}
76
77struct HostKeyHandler {
82 policy: SshHostKeyPolicy,
83 entries: Vec<known_hosts::KnownHostEntry>,
84 host: String,
85 port: u16,
86 logs: std::sync::Arc<Mutex<Vec<String>>>,
90}
91
92impl client::Handler for HostKeyHandler {
93 type Error = russh::Error;
94
95 async fn check_server_key(
96 &mut self,
97 server_public_key: &russh::keys::PublicKey,
98 ) -> Result<bool, Self::Error> {
99 let raw = server_public_key.public_key_bytes();
100 let mut logs: Vec<String> = Vec::new();
101 let decision = known_hosts::decide_host_key(
102 self.policy,
103 &self.host,
104 self.port,
105 &self.entries,
106 &raw,
107 &mut |line| logs.push(line),
108 );
109 self.logs.lock().unwrap().extend(logs);
110 Ok(decision == known_hosts::HostKeyDecision::Accept)
111 }
112}
113
114#[derive(Debug, Clone, PartialEq, Eq)]
119pub struct TunnelLease {
120 pub host: String,
121 pub port: u16,
122 pub generation: u64,
123}
124
125struct TunnelEntry {
126 local: std::net::SocketAddr,
128 handle: Arc<Handle<HostKeyHandler>>,
130 shutdown: mpsc::Sender<()>,
132 generation: u64,
133 draining: Arc<AtomicBool>,
135 last_used: Mutex<std::time::Instant>,
136 live_tasks: Arc<AtomicU64>,
138}
139
140impl TunnelEntry {
141 fn is_dead(&self) -> bool {
142 self.handle.is_closed()
143 }
144 fn touch(&self) {
145 *self.last_used.lock().unwrap() = std::time::Instant::now();
146 }
147}
148
149struct Tunnels {
150 map: HashMap<String, Arc<TunnelEntry>>,
151 next_generation: u64,
152 inflight: HashMap<String, Arc<tokio::sync::Mutex<()>>>,
155}
156
157static TUNNELS: OnceLock<Mutex<Tunnels>> = OnceLock::new();
158
159fn tunnels() -> &'static Mutex<Tunnels> {
160 TUNNELS.get_or_init(|| {
161 Mutex::new(Tunnels {
162 map: HashMap::new(),
163 next_generation: 1,
164 inflight: HashMap::new(),
165 })
166 })
167}
168
169pub fn tunnel_count() -> usize {
171 tunnels().lock().unwrap().map.len()
172}
173
174pub fn tunnel_connection_names() -> Vec<String> {
176 tunnels()
177 .lock()
178 .unwrap()
179 .map
180 .keys()
181 .map(|k| k.split('\u{1}').next().unwrap_or("").to_string())
182 .collect()
183}
184
185pub fn tunnel_live_tasks() -> u64 {
187 tunnels()
188 .lock()
189 .unwrap()
190 .map
191 .values()
192 .map(|e| e.live_tasks.load(Ordering::Relaxed))
193 .sum()
194}
195
196pub fn invalidate_all() {
198 let entries: Vec<Arc<TunnelEntry>> = {
199 let mut state = tunnels().lock().unwrap();
200 state.map.drain().map(|(_, e)| e).collect()
201 };
202 for entry in entries {
203 retire(&entry);
204 }
205}
206
207fn retire(entry: &TunnelEntry) {
212 entry.draining.store(true, Ordering::SeqCst);
213 let evicted = super::pool::evict_by_generation(entry.generation);
214 if evicted > 0 {
215 eprintln!(
216 "[sequel-mcp] ssh tunnel gen {}: evicted {evicted} associated MySQL pool(s)",
217 entry.generation
218 );
219 }
220 let _ = entry.shutdown.try_send(());
221 let handle = Arc::clone(&entry.handle);
224 tokio::spawn(async move {
225 let _ = handle
226 .disconnect(russh::Disconnect::ByApplication, "tunnel retired", "en")
227 .await;
228 });
229}
230
231impl SshTunnel {
232 fn auth_method_label(&self) -> &'static str {
233 match self.auth_method {
234 SshAuthMethod::Password => "password",
235 SshAuthMethod::Key => "key",
236 }
237 }
238}
239
240fn known_hosts_stamp(path: Option<&std::path::Path>) -> String {
244 let Some(path) = path else {
245 return String::new();
246 };
247 let mut h = sha2::Sha256::new();
248 match std::fs::read(path) {
249 Ok(bytes) => h.update(&bytes),
250 Err(_) => h.update(format!("unreadable:{}", path.display())),
253 }
254 format!(
255 "{:016x}",
256 u64::from_be_bytes(h.finalize()[..8].try_into().expect("8 bytes"))
257 )
258}
259
260fn ssh_credential_fragment(ssh_password: Option<&str>) -> String {
263 let cred = super::pool::CredentialGeneration::derive(ssh_password.unwrap_or(""));
264 cred.key_fragment()
265}
266
267#[allow(clippy::too_many_arguments)]
268fn tunnel_key(
269 conn_name: &str,
270 ssh: &SshTunnel,
271 ssh_password: Option<&str>,
272 target: &str,
273 port: u16,
274 policy_revision: u64,
275 kh_stamp: &str,
276) -> String {
277 let target_endpoint = format!("{target}:{port}");
278 let bridge = ssh
279 .docker
280 .as_ref()
281 .map(|d| format!("{}:{}", d.container, d.bridge_tool.as_str()))
282 .unwrap_or_default();
283 format!(
284 "{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}\u{1}{}",
285 conn_name,
286 ssh.host,
287 ssh.port,
288 ssh.user,
289 ssh.auth_method_label(),
290 kh_stamp,
291 ssh_credential_fragment(ssh_password),
292 policy_revision,
293 target_endpoint,
294 bridge
295 )
296}
297
298pub async fn tunnel_endpoint(
304 conn_name: &str,
305 ssh: &SshTunnel,
306 ssh_password: Option<&str>,
307 target_host: &str,
308 target_port: u16,
309 policy_revision: u64,
310) -> Result<TunnelLease, SshError> {
311 crate::app::test_mode::check_mysql_endpoint(&ssh.host, ssh.port).map_err(SshError::Setup)?;
314
315 let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
316 let kh_stamp = known_hosts_stamp(kh_path);
317 let key = tunnel_key(
318 conn_name,
319 ssh,
320 ssh_password,
321 target_host,
322 target_port,
323 policy_revision,
324 &kh_stamp,
325 );
326
327 {
329 let state = tunnels().lock().unwrap();
330 if let Some(entry) = state.map.get(&key)
331 && !entry.is_dead()
332 && !entry.draining.load(Ordering::SeqCst)
333 {
334 entry.touch();
335 return Ok(TunnelLease {
336 host: "127.0.0.1".into(),
337 port: entry.local.port(),
338 generation: entry.generation,
339 });
340 }
341 }
342
343 let guard = {
347 let mut state = tunnels().lock().unwrap();
348 Arc::clone(state.inflight.entry(key.clone()).or_default())
349 };
350 let _hold = guard.lock().await;
351
352 {
353 let state = tunnels().lock().unwrap();
354 if let Some(entry) = state.map.get(&key)
355 && !entry.is_dead()
356 && !entry.draining.load(Ordering::SeqCst)
357 {
358 entry.touch();
359 return Ok(TunnelLease {
360 host: "127.0.0.1".into(),
361 port: entry.local.port(),
362 generation: entry.generation,
363 });
364 }
365 }
366
367 let entry = establish(ssh, ssh_password, target_host, target_port).await?;
368
369 let mut state = tunnels().lock().unwrap();
370 state.inflight.remove(&key);
371 let conn_prefix = format!("{conn_name}\u{1}");
374 let stale: Vec<String> = state
375 .map
376 .keys()
377 .filter(|k| k.as_str() != key.as_str() && k.starts_with(&conn_prefix))
378 .cloned()
379 .collect();
380 for k in stale {
381 if let Some(e) = state.map.remove(&k) {
382 retire(&e);
383 }
384 }
385 let dead: Vec<String> = state
388 .map
389 .iter()
390 .filter(|(_, e)| e.is_dead())
391 .map(|(k, _)| k.clone())
392 .collect();
393 for k in dead {
394 if let Some(e) = state.map.remove(&k) {
395 retire(&e);
396 }
397 }
398 while state.map.len() >= MAX_TUNNELS {
399 let victim = state
400 .map
401 .iter()
402 .min_by_key(|(_, e)| *e.last_used.lock().unwrap())
403 .map(|(k, _)| k.clone());
404 let Some(victim) = victim else { break };
405 if let Some(e) = state.map.remove(&victim) {
406 eprintln!(
407 "[sequel-mcp] ssh tunnel cache full ({}); retiring LRU entry",
408 MAX_TUNNELS
409 );
410 retire(&e);
411 }
412 }
413 let generation = entry.generation;
414 let port = entry.local.port();
415 state.map.insert(key, entry);
416 Ok(TunnelLease {
417 host: "127.0.0.1".into(),
418 port,
419 generation,
420 })
421}
422
423fn bridge_command(
429 ssh: &SshTunnel,
430 target_host: &str,
431 target_port: u16,
432) -> Result<Option<String>, SshError> {
433 let Some(docker) = &ssh.docker else {
434 return Ok(None);
435 };
436 let argv = super::docker::bridge_argv(
437 &docker.container,
438 docker.bridge_tool,
439 target_host,
440 target_port,
441 )
442 .map_err(|e| SshError::Setup(format!("docker bridge: {e}")))?;
443 Ok(Some(argv.join(" ")))
444}
445
446async fn establish(
447 ssh: &SshTunnel,
448 ssh_password: Option<&str>,
449 target_host: &str,
450 target_port: u16,
451) -> Result<Arc<TunnelEntry>, SshError> {
452 let bridge_cmd = bridge_command(ssh, target_host, target_port)?;
453 let kh_path = ssh.known_hosts_path.as_deref().map(std::path::Path::new);
454 let entries =
455 known_hosts::load_known_hosts_checked(kh_path).map_err(|e| SshError::HostKey {
456 host: ssh.host.clone(),
457 port: ssh.port,
458 reason: e,
459 })?;
460 let logs = std::sync::Arc::new(Mutex::new(Vec::new()));
461 let handler = HostKeyHandler {
462 policy: ssh.host_key_policy.unwrap_or(SshHostKeyPolicy::Lenient),
463 entries,
464 host: ssh.host.clone(),
465 port: ssh.port,
466 logs: logs.clone(),
467 };
468
469 let keepalive = keepalive_interval();
470 let config = Arc::new(client::Config {
471 keepalive_interval: Some(keepalive),
472 keepalive_max: 3,
473 nodelay: true,
474 ..client::Config::default()
475 });
476
477 let addr = (ssh.host.as_str(), ssh.port);
478 let session = async {
479 let mut handle = client::connect(config, addr, handler)
480 .await
481 .map_err(|e| SshError::Transport(format!("connect to bastion: {e}")))?;
482 for line in logs.lock().unwrap().drain(..) {
485 eprintln!("[sequel-mcp] SSH {}:{} {line}", ssh.host, ssh.port);
486 }
487 let auth = match ssh.auth_method {
491 SshAuthMethod::Password => {
492 let Some(password) = ssh_password else {
493 return Err(SshError::Setup(format!(
494 "no SSH password stored for {:?} (expected under \"<connection>::ssh\")",
495 ssh.user
496 )));
497 };
498 handle
499 .authenticate_password(ssh.user.clone(), password)
500 .await
501 .map_err(|e| SshError::Transport(format!("password auth transport: {e}")))?
502 }
503 SshAuthMethod::Key => {
504 let Some(path) = &ssh.private_key_path else {
505 return Err(SshError::Setup(
506 "key auth configured without privateKeyPath".into(),
507 ));
508 };
509 let expanded = crate::app::paths::expand_tilde(path);
510 let key = russh::keys::load_secret_key(&expanded, ssh_password)
511 .map_err(|e| SshError::Setup(format!("load private key: {e}")))?;
512 use russh::keys::Algorithm;
513 match key.algorithm() {
514 Algorithm::Ed25519 | Algorithm::Rsa { .. } | Algorithm::Ecdsa { .. } => {}
515 other => {
516 return Err(SshError::Setup(format!(
517 "unsupported private key algorithm {other:?} (supported: ed25519, rsa, ecdsa)"
518 )));
519 }
520 }
521 let hash_alg = if key.algorithm().is_rsa() {
525 match handle.best_supported_rsa_hash().await {
526 Ok(best) => best.unwrap_or(Some(russh::keys::HashAlg::Sha256)),
527 Err(_) => Some(russh::keys::HashAlg::Sha256),
528 }
529 } else {
530 None
531 };
532 handle
533 .authenticate_publickey(
534 ssh.user.clone(),
535 russh::keys::PrivateKeyWithHashAlg::new(Arc::new(key), hash_alg),
536 )
537 .await
538 .map_err(|e| SshError::Transport(format!("publickey auth transport: {e}")))?
539 }
540 };
541 if !matches!(auth, russh::client::AuthResult::Success) {
542 return Err(SshError::Auth {
543 host: ssh.host.clone(),
544 port: ssh.port,
545 user: ssh.user.clone(),
546 });
547 }
548 Ok(handle)
549 };
550 let handle = tokio::time::timeout(SSH_CONNECT_TIMEOUT, session)
551 .await
552 .map_err(|_| {
553 SshError::Transport(format!(
554 "ssh session timeout after {}s",
555 SSH_CONNECT_TIMEOUT.as_secs()
556 ))
557 })??;
558
559 let listener = TcpListener::bind(("127.0.0.1", 0))
563 .await
564 .map_err(|e| SshError::Setup(format!("local bind: {e}")))?;
565 let local = listener
566 .local_addr()
567 .map_err(|e| SshError::Setup(format!("local addr: {e}")))?;
568
569 let generation = {
570 let mut state = tunnels().lock().unwrap();
571 let assigned = state.next_generation;
572 state.next_generation += 1;
573 assigned
574 };
575
576 let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
577 let forward_host = target_host.to_string();
578 let forward_port = u32::from(target_port);
579 let shared = Arc::new(handle);
580 let accept_session = Arc::clone(&shared);
581 let draining = Arc::new(AtomicBool::new(false));
582 let live_tasks: Arc<AtomicU64> = Arc::new(AtomicU64::new(0));
583 let accept_draining = Arc::clone(&draining);
584 let accept_tasks = Arc::clone(&live_tasks);
585 tokio::spawn(async move {
586 loop {
587 tokio::select! {
588 _ = shutdown_rx.recv() => break,
589 accepted = listener.accept() => {
590 let Ok((socket, _peer)) = accepted else { break };
591 if accept_draining.load(Ordering::SeqCst) || accept_session.is_closed() {
592 break;
593 }
594 let session = Arc::clone(&accept_session);
595 let host = forward_host.clone();
596 let tasks = Arc::clone(&accept_tasks);
597 accept_tasks.fetch_add(1, Ordering::Relaxed);
598 let bridge_cmd = bridge_cmd.clone();
599 tokio::spawn(async move {
600 let target = format!("{host}:{forward_port}");
601 let channel = tokio::time::timeout(
605 CHANNEL_OPEN_TIMEOUT,
606 async {
607 match &bridge_cmd {
608 Some(cmd) => {
609 let ch = session.channel_open_session().await?;
610 ch.exec(true, cmd.as_str()).await?;
611 Ok::<_, russh::Error>(ch)
612 }
613 None => session
614 .channel_open_direct_tcpip(
615 host,
616 forward_port,
617 "127.0.0.1",
618 0,
619 )
620 .await,
621 }
622 },
623 )
624 .await;
625 match channel {
626 Ok(Ok(channel)) => {
627 let (mut chan_r, mut chan_w) =
631 tokio::io::split(channel.into_stream());
632 let (mut sock_r, mut sock_w) = socket.into_split();
633 let a = tokio::io::copy(&mut sock_r, &mut chan_w);
634 let b = tokio::io::copy(&mut chan_r, &mut sock_w);
635 let _ = tokio::join!(a, b);
636 }
637 Ok(Err(e)) => {
638 eprintln!(
639 "[sequel-mcp] ssh forwarder: channel open to {target} failed: {e:?}"
640 );
641 }
644 Err(_) => {
645 eprintln!(
646 "[sequel-mcp] ssh forwarder: channel open to {target} timed out after {}s (stale session?)",
647 CHANNEL_OPEN_TIMEOUT.as_secs()
648 );
649 }
650 }
651 tasks.fetch_sub(1, Ordering::Relaxed);
652 });
653 }
654 }
655 }
656 });
657
658 Ok(Arc::new(TunnelEntry {
659 local,
660 handle: shared,
661 shutdown: shutdown_tx,
662 generation,
663 draining,
664 last_used: Mutex::new(std::time::Instant::now()),
665 live_tasks,
666 }))
667}
668
669#[cfg(test)]
670mod tests {
671 use super::*;
672 use crate::config::SshTunnel;
673
674 fn sample_tunnel() -> SshTunnel {
675 SshTunnel {
676 host: "bastion".into(),
677 port: 22,
678 user: "sshuser".into(),
679 auth_method: SshAuthMethod::Password,
680 ..SshTunnel::default()
681 }
682 }
683
684 #[test]
685 fn tunnel_key_separates_identity() {
686 let ssh = sample_tunnel();
687 let a = tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp");
688 let mut ssh2 = ssh.clone();
689 ssh2.user = "other".into();
690 let b = tunnel_key("c1", &ssh2, Some("pw"), "db", 3306, 1, "stamp");
691 assert_ne!(a, b);
692 assert_ne!(
693 a,
694 tunnel_key("c1", &ssh, Some("pw"), "other", 3306, 1, "stamp")
695 );
696 assert_ne!(
697 a,
698 tunnel_key("c1", &ssh, Some("pw"), "db", 3307, 1, "stamp")
699 );
700 assert_ne!(
701 a,
702 tunnel_key("c2", &ssh, Some("pw"), "db", 3306, 1, "stamp")
703 );
704 assert_ne!(
706 a,
707 tunnel_key("c1", &ssh, Some("different"), "db", 3306, 1, "stamp")
708 );
709 assert_ne!(
710 a,
711 tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 2, "stamp")
712 );
713 assert_ne!(
714 a,
715 tunnel_key("c1", &ssh, Some("pw"), "db", 3306, 1, "stamp2")
716 );
717 let mut bridged = ssh.clone();
719 bridged.docker = Some(crate::config::SshDocker {
720 container: "db".into(),
721 bridge_tool: crate::config::BridgeTool::Nc,
722 });
723 assert_ne!(
724 a,
725 tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
726 );
727 let mut bridged2 = bridged.clone();
728 bridged2.docker = Some(crate::config::SshDocker {
729 container: "db".into(),
730 bridge_tool: crate::config::BridgeTool::Socat,
731 });
732 assert_ne!(
733 tunnel_key("c1", &bridged, Some("pw"), "127.0.0.1", 3306, 1, "stamp"),
734 tunnel_key("c1", &bridged2, Some("pw"), "127.0.0.1", 3306, 1, "stamp")
735 );
736 }
737
738 #[test]
739 fn bridge_command_forms() {
740 use crate::config::{BridgeTool, SshDocker};
741 let mk = |tool| SshTunnel {
742 host: "bastion".into(),
743 port: 22,
744 user: "u".into(),
745 auth_method: SshAuthMethod::Key,
746 docker: Some(SshDocker {
747 container: "db-1".into(),
748 bridge_tool: tool,
749 }),
750 ..SshTunnel::default()
751 };
752 assert_eq!(
753 bridge_command(&mk(BridgeTool::Nc), "127.0.0.1", 3306)
754 .unwrap()
755 .unwrap(),
756 "docker exec -i db-1 nc 127.0.0.1 3306"
757 );
758 assert_eq!(
759 bridge_command(&mk(BridgeTool::Ncat), "127.0.0.1", 3306)
760 .unwrap()
761 .unwrap(),
762 "docker exec -i db-1 ncat 127.0.0.1 3306"
763 );
764 assert_eq!(
765 bridge_command(&mk(BridgeTool::Socat), "127.0.0.1", 3306)
766 .unwrap()
767 .unwrap(),
768 "docker exec -i db-1 socat - TCP:127.0.0.1:3306"
769 );
770 let plain = SshTunnel::default();
772 assert_eq!(bridge_command(&plain, "db", 3306).unwrap(), None);
773 let mut bad = mk(BridgeTool::Nc);
775 bad.docker = Some(SshDocker {
776 container: "bad name!".into(),
777 bridge_tool: BridgeTool::Nc,
778 });
779 assert!(bridge_command(&bad, "127.0.0.1", 3306).is_err());
780 }
781
782 #[test]
783 fn credential_fragment_hides_the_secret() {
784 let f = ssh_credential_fragment(Some("top-secret-password"));
785 assert_eq!(f.len(), 32, "hex of a 16-byte digest");
786 assert!(!f.contains("top-secret"));
787 assert_ne!(f, ssh_credential_fragment(Some("other")));
788 assert_eq!(
790 ssh_credential_fragment(None),
791 ssh_credential_fragment(Some(""))
792 );
793 }
794
795 #[test]
796 fn known_hosts_stamp_tracks_content() {
797 let dir = tempfile::TempDir::new().unwrap();
798 let f = dir.path().join("kh");
799 std::fs::write(&f, b"content-a").unwrap();
800 let s1 = known_hosts_stamp(Some(&f));
801 let s2 = known_hosts_stamp(Some(&f));
802 assert_eq!(s1, s2, "stable per content");
803 std::fs::write(&f, b"content-b").unwrap();
804 assert_ne!(s1, known_hosts_stamp(Some(&f)));
805 assert_eq!(known_hosts_stamp(None), "");
806 }
807}