1use std::io::{ErrorKind, Read, Write};
23use std::net::{TcpStream, ToSocketAddrs};
24use std::os::unix::net::UnixStream;
25use std::path::PathBuf;
26use std::sync::Arc;
27use std::sync::atomic::AtomicUsize;
28use std::sync::mpsc::{TryRecvError, sync_channel};
29use std::time::Duration;
30
31use serde_json::Value;
32use tungstenite::{Message, WebSocket};
33
34use super::http::{Peer, Request};
35use super::mcp::Caller;
36use super::terminal::{Limits, Pty, plain_name, query_param};
37use crate::error::{Error, Result};
38
39#[derive(Debug, Clone, PartialEq, Eq)]
41pub struct SshRequest {
42 pub instance: String,
44 pub keys_of: Option<String>,
48 pub forwarded_keys: Option<Vec<String>>,
53}
54
55pub const KEYS_HEADER: &str = "X-Isb-Ssh-Keys";
59
60pub fn keys_header(keys: &[String]) -> String {
62 use base64::Engine;
63 base64::engine::general_purpose::URL_SAFE_NO_PAD
64 .encode(serde_json::to_vec(keys).unwrap_or_default())
65}
66
67fn parse_keys_header(v: &str) -> std::result::Result<Vec<String>, String> {
68 use base64::Engine;
69 let b = base64::engine::general_purpose::URL_SAFE_NO_PAD
70 .decode(v.trim())
71 .map_err(|_| "the forwarded SSH keys are not base64url")?;
72 let keys: Vec<String> =
73 serde_json::from_slice(&b).map_err(|_| "the forwarded SSH keys are not a JSON list")?;
74 if keys.is_empty() || keys.len() > 64 || keys.iter().any(|k| k.contains(['\n', '\r', '\0'])) {
76 return Err("the forwarded SSH keys are not a list of keys".into());
77 }
78 Ok(keys)
79}
80
81impl SshRequest {
82 pub fn query(&self) -> String {
83 let mut q = format!("instance={}", self.instance);
84 if let Some(e) = &self.keys_of {
85 q.push_str(&format!("&as={}", encode(e)));
86 }
87 q
88 }
89}
90
91pub type Ssh =
94 Arc<dyn Fn(&Caller, &crate::org::OrgId, &SshRequest) -> Result<Box<dyn Pty>> + Send + Sync>;
95
96static ACTIVE: AtomicUsize = AtomicUsize::new(0);
97
98pub static LIMITS: Limits = Limits {
101 active: &ACTIVE,
102 max_sessions: 64,
103 idle: Duration::from_secs(2 * 3600),
104 max_age: Duration::from_secs(24 * 3600),
105 busy: "too many SSH sessions are open on this server; close one and try again",
106};
107
108pub fn ssh_request(req: &Request) -> std::result::Result<SshRequest, String> {
110 let instance = query_param(req, "instance").ok_or("instance= is required")?;
111 if !plain_name(&instance) {
112 return Err("instance= is not an instance name".into());
113 }
114 let keys_of = match query_param(req, "as") {
115 None => None,
116 Some(e) => {
117 let e = decode(&e).ok_or("as= is not an email")?;
118 if e.is_empty() || e.len() > 254 || !e.contains('@') || e.contains(char::is_control) {
119 return Err("as= is not an email".into());
120 }
121 Some(e)
122 }
123 };
124 let forwarded_keys = match req.header(KEYS_HEADER) {
125 Some(h) => Some(parse_keys_header(h)?),
126 None => None,
127 };
128 Ok(SshRequest {
129 instance,
130 keys_of,
131 forwarded_keys,
132 })
133}
134
135pub fn origin_allowed(req: &Request) -> bool {
139 match req.header("origin") {
140 Some(_) => super::terminal::origin_allowed(req),
141 None => matches!(req.peer, Peer::Unix { .. }) || super::terminal::origin_allowed(req),
142 }
143}
144
145fn encode(s: &str) -> String {
146 s.bytes()
147 .map(|b| match b {
148 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' | b'@' => {
149 (b as char).to_string()
150 }
151 _ => format!("%{b:02X}"),
152 })
153 .collect()
154}
155
156fn decode(s: &str) -> Option<String> {
157 let b = s.as_bytes();
158 let mut out = Vec::with_capacity(b.len());
159 let mut i = 0;
160 while i < b.len() {
161 match b[i] {
162 b'%' => {
163 let h = std::str::from_utf8(b.get(i + 1..i + 3)?).ok()?;
164 out.push(u8::from_str_radix(h, 16).ok()?);
165 i += 3;
166 }
167 b'+' => {
168 out.push(b' ');
169 i += 1;
170 }
171 c => {
172 out.push(c);
173 i += 1;
174 }
175 }
176 }
177 String::from_utf8(out).ok()
178}
179
180#[derive(Clone)]
185pub enum Remote {
186 Socket(PathBuf),
187 Url { base: String, token: Option<String> },
188}
189
190impl std::fmt::Debug for Remote {
191 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
192 match self {
193 Remote::Socket(p) => write!(f, "Socket({})", p.display()),
194 Remote::Url { base, token } => write!(
196 f,
197 "Url({base}, token: {})",
198 if token.is_some() { "set" } else { "none" }
199 ),
200 }
201 }
202}
203
204pub trait Conn: Read + Write + Send {
206 fn set_timeout(&self, t: Option<Duration>) -> std::io::Result<()>;
207}
208
209impl Conn for TcpStream {
210 fn set_timeout(&self, t: Option<Duration>) -> std::io::Result<()> {
211 self.set_read_timeout(t)
212 }
213}
214
215impl Conn for UnixStream {
216 fn set_timeout(&self, t: Option<Duration>) -> std::io::Result<()> {
217 self.set_read_timeout(t)
218 }
219}
220
221impl Conn for rustls::StreamOwned<rustls::ClientConnection, TcpStream> {
222 fn set_timeout(&self, t: Option<Duration>) -> std::io::Result<()> {
223 self.sock.set_read_timeout(t)
224 }
225}
226
227pub fn split_base(base: &str) -> Result<(bool, String, u16, String)> {
229 let (tls, rest) = if let Some(r) = base.strip_prefix("https://") {
230 (true, r)
231 } else if let Some(r) = base.strip_prefix("http://") {
232 (false, r)
233 } else {
234 return Err(Error::invalid(format!(
235 "{base:?}: want an http:// or https:// URL"
236 )));
237 };
238 let (authority, prefix) = match rest.find('/') {
239 Some(i) => (&rest[..i], rest[i..].trim_end_matches('/')),
240 None => (rest, ""),
241 };
242 let (host, port) = match authority.rsplit_once(':') {
243 Some((h, p)) if !h.ends_with(']') || h.starts_with('[') => match p.parse::<u16>() {
244 Ok(p) => (h, p),
245 Err(_) => (authority, if tls { 443 } else { 80 }),
246 },
247 _ => (authority, if tls { 443 } else { 80 }),
248 };
249 if host.is_empty() {
250 return Err(Error::invalid(format!("{base:?}: no host")));
251 }
252 Ok((tls, host.to_string(), port, prefix.to_string()))
253}
254
255impl Remote {
256 pub fn call_tool(&self, org: &str, tool: &str, mut args: Value) -> Result<Value> {
259 match self {
260 Remote::Socket(p) => {
261 args["org"] = Value::String(org.to_string());
262 super::client::call_tool(p, tool, args, Duration::from_secs(60))
263 }
264 Remote::Url { .. } => {
265 let (status, v) = self.http(
266 "POST",
267 &format!("/orgs/{org}/api/v1/tools/{tool}"),
268 Some(&args),
269 )?;
270 if status == 200 {
271 return Ok(v.get("result").cloned().unwrap_or(Value::Null));
272 }
273 Err(answer_error(status, &v))
274 }
275 }
276 }
277
278 pub fn http(&self, method: &str, path: &str, body: Option<&Value>) -> Result<(u16, Value)> {
280 let Remote::Url { base, token } = self else {
281 return Err(Error::invalid(
282 "this needs isb serve's URL: pass --url or set ISB_URL",
283 ));
284 };
285 let url = format!("{}{path}", base.trim_end_matches('/'));
286 let agent: ureq::Agent = ureq::Agent::config_builder()
287 .timeout_global(Some(Duration::from_secs(60)))
288 .http_status_as_error(false)
289 .user_agent(concat!("isb/", env!("CARGO_PKG_VERSION")))
290 .build()
291 .into();
292 let auth = token.as_ref().map(|t| format!("Bearer {}", t.trim()));
293 let payload = match body {
294 Some(b) => serde_json::to_vec(b)?,
295 None => Vec::new(),
296 };
297 macro_rules! go {
298 ($req:expr) => {{
299 let mut r = $req.header("X-Isb-Csrf", "1");
300 if let Some(a) = &auth {
301 r = r.header("Authorization", a);
302 }
303 r
304 }};
305 }
306 let resp = match method {
307 "GET" => go!(agent.get(&url)).call(),
308 "DELETE" => go!(agent.delete(&url)).call(),
309 "POST" => go!(agent.post(&url))
310 .header("Content-Type", "application/json")
311 .send(&payload[..]),
312 m => return Err(Error::invalid(format!("unsupported method {m}"))),
313 };
314 let mut resp = resp.map_err(|e| Error::invalid(format!("{method} {url}: {e}")))?;
315 let status = resp.status().as_u16();
316 let text = resp
317 .body_mut()
318 .with_config()
319 .limit(16 << 20)
320 .read_to_string()
321 .unwrap_or_default();
322 Ok((status, serde_json::from_str(&text).unwrap_or(Value::Null)))
323 }
324
325 pub fn websocket(&self, path: &str) -> Result<WebSocket<Box<dyn Conn>>> {
327 use tungstenite::client::IntoClientRequest;
328 let timeout = Duration::from_secs(15);
329 let (conn, url, token): (Box<dyn Conn>, String, Option<&String>) = match self {
330 Remote::Socket(p) => {
331 let s = UnixStream::connect(p).map_err(|e| {
332 Error::invalid(format!(
333 "cannot connect to isb serve at {}: {e}",
334 p.display()
335 ))
336 })?;
337 (Box::new(s), format!("ws://localhost{path}"), None)
338 }
339 Remote::Url { base, token } => {
340 let (tls, host, port, prefix) = split_base(base)?;
341 let addr = (host.trim_start_matches('[').trim_end_matches(']'), port)
342 .to_socket_addrs()
343 .map_err(|e| Error::invalid(format!("{host}: {e}")))?
344 .next()
345 .ok_or_else(|| Error::invalid(format!("{host} does not resolve")))?;
346 let sock = TcpStream::connect_timeout(&addr, timeout)
347 .map_err(|e| Error::invalid(format!("cannot connect to {base}: {e}")))?;
348 let _ = sock.set_nodelay(true);
349 sock.set_read_timeout(Some(timeout))?;
350 let authority = if (tls && port == 443) || (!tls && port == 80) {
351 host.clone()
352 } else {
353 format!("{host}:{port}")
354 };
355 let scheme = if tls { "wss" } else { "ws" };
356 let url = format!("{scheme}://{authority}{prefix}{path}");
357 let conn: Box<dyn Conn> = if tls {
358 let name = rustls::pki_types::ServerName::try_from(
359 host.trim_start_matches('[')
360 .trim_end_matches(']')
361 .to_string(),
362 )
363 .map_err(|e| Error::invalid(format!("{host}: {e}")))?;
364 let c = rustls::ClientConnection::new(crate::net::default_tls(), name)
365 .map_err(|e| Error::invalid(format!("TLS: {e}")))?;
366 Box::new(rustls::StreamOwned::new(c, sock))
367 } else {
368 Box::new(sock)
369 };
370 (conn, url, token.as_ref())
371 }
372 };
373 conn.set_timeout(Some(timeout))?;
374 let mut req = url
375 .into_client_request()
376 .map_err(|e| Error::WebSocket(e.to_string()))?;
377 if let Some(t) = token {
378 req.headers_mut().insert(
379 "authorization",
380 format!("Bearer {}", t.trim())
381 .parse()
382 .map_err(|_| Error::invalid("the token is not a valid header value"))?,
383 );
384 }
385 let (ws, _) = tungstenite::client(req, conn).map_err(|e| match e {
386 tungstenite::HandshakeError::Failure(tungstenite::Error::Http(r)) => {
387 let status = r.status().as_u16();
388 let body = r
389 .body()
390 .as_ref()
391 .map(|b| String::from_utf8_lossy(b).into_owned())
392 .unwrap_or_default();
393 let v: Value = serde_json::from_str(&body).unwrap_or_else(|_| match body.trim() {
394 "" => serde_json::json!({ "message": match status {
397 401 => "the API token was refused (unknown, expired or revoked)",
398 403 => "refused: SSH needs exec in the org (members and up; not viewers, and not read- or deploy-scoped tokens), from this site",
399 404 => "no SSH here: no such org, or SSH is off on this server (--deny-tools sandbox_exec)",
400 _ => "the server refused the connection",
401 }}),
402 b => serde_json::json!({ "message": b }),
403 });
404 answer_error(status, &v)
405 }
406 e => Error::WebSocket(e.to_string()),
407 })?;
408 Ok(ws)
409 }
410}
411
412pub fn answer_error(status: u16, v: &Value) -> Error {
414 let message = v["message"]
415 .as_str()
416 .map(String::from)
417 .unwrap_or_else(|| format!("HTTP {status}"));
418 match status {
419 401 | 403 => Error::Forbidden(message),
420 404 => Error::NotFound(message),
421 _ => Error::Remote {
422 code: v["error"].as_str().unwrap_or("server_error").into(),
423 message,
424 data: v.get("data").cloned().unwrap_or(Value::Null),
425 },
426 }
427}
428
429fn would_block(e: &tungstenite::Error) -> bool {
430 matches!(e, tungstenite::Error::Io(i) if matches!(i.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut))
431}
432
433pub fn pump<S: Conn + ?Sized>(
437 ws: &mut WebSocket<Box<S>>,
438 input: impl Read + Send + 'static,
439 mut output: impl Write,
440) -> Result<()> {
441 const CHUNK: usize = 32 * 1024;
442 let (tx, rx) = sync_channel::<Option<Vec<u8>>>(64);
443 std::thread::spawn(move || {
444 let mut input = input;
445 let mut buf = vec![0u8; CHUNK];
446 loop {
447 match input.read(&mut buf) {
448 Ok(0) | Err(_) => {
449 let _ = tx.send(None);
450 return;
451 }
452 Ok(n) => {
453 if tx.send(Some(buf[..n].to_vec())).is_err() {
454 return;
455 }
456 }
457 }
458 }
459 });
460 ws.get_ref().set_timeout(Some(Duration::from_millis(10)))?;
461 loop {
462 match ws.read() {
463 Ok(Message::Binary(b)) => {
464 output.write_all(&b)?;
465 output.flush()?;
466 }
467 Ok(Message::Text(t)) => {
468 let v: Value = serde_json::from_str(t.as_str()).unwrap_or(Value::Null);
469 match v["type"].as_str() {
470 Some("error") => {
471 return Err(Error::Remote {
472 code: "ssh".into(),
473 message: v["message"].as_str().unwrap_or("refused").to_string(),
474 data: Value::Null,
475 });
476 }
477 Some("exit") => return Ok(()),
478 _ => {}
479 }
480 }
481 Ok(Message::Close(_)) => return Ok(()),
482 Ok(_) => {}
483 Err(e) if would_block(&e) => {}
484 Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed) => {
485 return Ok(());
486 }
487 Err(e) => return Err(Error::WebSocket(e.to_string())),
488 }
489 loop {
491 match rx.try_recv() {
492 Ok(Some(d)) => match ws.send(Message::binary(d)) {
493 Ok(()) => {}
494 Err(e) if would_block(&e) => {}
495 Err(e) => return Err(Error::WebSocket(e.to_string())),
496 },
497 Ok(None) | Err(TryRecvError::Disconnected) => {
498 let _ = ws.close(None);
500 let _ = ws.flush();
501 return Ok(());
502 }
503 Err(TryRecvError::Empty) => break,
504 }
505 }
506 }
507}
508
509#[cfg(test)]
510mod tests {
511 use super::*;
512
513 fn req(query: &str, peer: Peer, headers: &[(&str, &str)]) -> Request {
514 Request {
515 method: "GET".into(),
516 path: "/orgs/acme/api/v1/ssh".into(),
517 query: Some(query.into()),
518 headers: headers
519 .iter()
520 .map(|(k, v)| (k.to_string(), v.to_string()))
521 .collect(),
522 body: vec![],
523 peer,
524 }
525 }
526
527 fn tcp() -> Peer {
528 Peer::Tcp("127.0.0.1:5000".parse().unwrap())
529 }
530
531 #[test]
532 fn parses_requests() {
533 let r = ssh_request(&req("instance=box", tcp(), &[])).unwrap();
534 assert_eq!(
535 r,
536 SshRequest {
537 instance: "box".into(),
538 keys_of: None,
539 forwarded_keys: None,
540 }
541 );
542 let k = vec!["ssh-ed25519 AAAA".to_string()];
543 let h = keys_header(&k);
544 let r = ssh_request(&req("instance=box", tcp(), &[(KEYS_HEADER, &h)])).unwrap();
545 assert_eq!(r.forwarded_keys, Some(k));
546 let two_lines = keys_header(&["a\nb".to_string()]);
547 for bad in ["!", &two_lines, &keys_header(&[])] {
548 assert!(ssh_request(&req("instance=box", tcp(), &[(KEYS_HEADER, bad)])).is_err());
549 }
550 let r = ssh_request(&req("instance=box&as=a%2Bb%40example.com", tcp(), &[])).unwrap();
551 assert_eq!(r.keys_of.as_deref(), Some("a+b@example.com"));
552 assert_eq!(r.query(), "instance=box&as=a%2Bb@example.com");
553 assert_eq!(ssh_request(&req(&r.query(), tcp(), &[])).unwrap(), r);
554 for bad in [
555 "",
556 "instance=",
557 "instance=Box",
558 "instance=../x",
559 "instance=a&as=nobody",
560 "instance=a&as=%zz",
561 ] {
562 assert!(ssh_request(&req(bad, tcp(), &[])).is_err(), "{bad}");
563 }
564 }
565
566 #[test]
567 fn origin_rules() {
568 let host = ("Host", "isb.example.com");
569 assert!(origin_allowed(&req("", Peer::Unix { uid: None }, &[])));
571 assert!(!origin_allowed(&req(
572 "",
573 Peer::Unix { uid: None },
574 &[host, ("Origin", "https://evil.example")]
575 )));
576 assert!(origin_allowed(&req(
578 "",
579 tcp(),
580 &[host, ("Authorization", "Bearer x")]
581 )));
582 assert!(!origin_allowed(&req(
584 "",
585 tcp(),
586 &[host, ("Cookie", "isb_session=x")]
587 )));
588 assert!(!origin_allowed(&req(
589 "",
590 tcp(),
591 &[
592 host,
593 ("Cookie", "isb_session=x"),
594 ("Origin", "https://evil.example")
595 ]
596 )));
597 assert!(origin_allowed(&req(
598 "",
599 tcp(),
600 &[
601 host,
602 ("Cookie", "isb_session=x"),
603 ("Origin", "https://isb.example.com")
604 ]
605 )));
606 }
607
608 #[test]
609 fn splits_bases() {
610 assert_eq!(
611 split_base("https://isb.example.com").unwrap(),
612 (true, "isb.example.com".into(), 443, "".into())
613 );
614 assert_eq!(
615 split_base("http://127.0.0.1:8092/").unwrap(),
616 (false, "127.0.0.1".into(), 8092, "".into())
617 );
618 assert_eq!(
619 split_base("https://h.example:8443/isb/").unwrap(),
620 (true, "h.example".into(), 8443, "/isb".into())
621 );
622 assert!(split_base("isb.example.com").is_err());
623 assert!(split_base("https://").is_err());
624 let r = Remote::Url {
625 base: "https://x".into(),
626 token: Some("isb_tok_secret".into()),
627 };
628 assert!(!format!("{r:?}").contains("secret"));
629 }
630
631 #[test]
634 fn pumps_bytes_both_ways() {
635 use super::super::http::Duplex;
636 use super::super::terminal::{PtyOutput, bridge};
637 use std::sync::mpsc::{Receiver, Sender, channel};
638
639 struct Upper(Sender<PtyOutput>, Receiver<PtyOutput>);
640 impl Pty for Upper {
641 fn input(&mut self, d: &[u8]) -> crate::Result<()> {
642 if d == b"bye" {
643 self.0
644 .send(PtyOutput::Failed("the key was removed".into()))
645 .unwrap();
646 } else {
647 self.0
648 .send(PtyOutput::Data(d.to_ascii_uppercase()))
649 .unwrap();
650 }
651 Ok(())
652 }
653 fn resize(&mut self, _: u16, _: u16) {}
654 fn output(&mut self, w: Duration) -> PtyOutput {
655 self.1.recv_timeout(w).unwrap_or(PtyOutput::Idle)
656 }
657 fn close(&mut self) {}
658 }
659 let (a, b) = UnixStream::pair().unwrap();
660 let server = std::thread::spawn(move || {
661 let mut s = a;
662 let d: &mut dyn Duplex = &mut s;
663 let mut ws = WebSocket::from_raw_socket(d, tungstenite::protocol::Role::Server, None);
664 let (tx, rx) = channel();
665 bridge(
666 &mut ws,
667 Box::new(Upper(tx, rx)),
668 Duration::from_secs(5),
669 Duration::from_secs(5),
670 );
671 });
672 let conn: Box<UnixStream> = Box::new(b);
673 let mut ws = WebSocket::from_raw_socket(conn, tungstenite::protocol::Role::Client, None);
674 let (mut w, r) = UnixStream::pair().unwrap();
676 let feeder = std::thread::spawn(move || {
677 w.write_all(b"hello").unwrap();
678 std::thread::sleep(Duration::from_millis(200));
679 w.write_all(b"bye").unwrap();
680 std::thread::sleep(Duration::from_millis(2000));
681 });
682 let mut out = Vec::new();
683 let e = pump(&mut ws, r, &mut out).unwrap_err();
684 assert_eq!(out, b"HELLO");
685 assert!(e.to_string().contains("the key was removed"), "{e}");
686 server.join().unwrap();
687 feeder.join().unwrap();
688 }
689}