use crate::api::{Client, WsConfig};
use crate::error::{Error, Result};
use futures::{SinkExt, StreamExt};
use serde_json::{Value, json};
use std::time::Duration;
use tokio_tungstenite::tungstenite::Message;
const STDIN: u8 = 0x01;
const STDOUT: u8 = 0x02;
pub const PING_INTERVAL: Duration = Duration::from_secs(25);
#[derive(Debug, Clone)]
pub struct ShellOptions {
pub cols: u16,
pub rows: u16,
pub cwd: Option<String>,
pub wake: bool,
pub sandbox_id: Option<String>,
}
impl Default for ShellOptions {
fn default() -> Self {
Self {
cols: 80,
rows: 24,
cwd: None,
wake: true,
sandbox_id: None,
}
}
}
impl ShellOptions {
pub fn size(mut self, cols: u16, rows: u16) -> Self {
self.cols = cols;
self.rows = rows;
self
}
pub fn cwd(mut self, cwd: impl Into<String>) -> Self {
self.cwd = Some(cwd.into());
self
}
pub fn no_wake(mut self) -> Self {
self.wake = false;
self
}
pub fn vm(mut self, sandbox_id: impl Into<String>) -> Self {
self.sandbox_id = Some(sandbox_id.into());
self
}
fn query(&self) -> String {
let mut q = format!(
"cols={}&rows={}&wake={}",
self.cols,
self.rows,
if self.wake { "true" } else { "false" }
);
if let Some(cwd) = &self.cwd {
q.push_str("&cwd=");
q.push_str(&urlencode(cwd));
}
if let Some(id) = &self.sandbox_id {
q.push_str("&sandbox_id=");
q.push_str(&urlencode(id));
}
q
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum ShellEvent {
Output(Vec<u8>),
Error(String),
}
#[derive(Debug, Clone, PartialEq)]
pub struct ShellExit {
pub code: i32,
pub error: Option<String>,
}
impl ShellExit {
pub fn is_clean(&self) -> bool {
self.code == 0 && self.error.is_none()
}
}
#[derive(Debug, Clone, PartialEq)]
enum Incoming {
Ready(String),
Output(Vec<u8>),
Error(String),
Exit(i32),
Ignored,
Closed,
}
fn frame_stdin(bytes: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(bytes.len() + 1);
out.push(STDIN);
out.extend_from_slice(bytes);
out
}
fn parse_incoming(msg: &Message) -> Incoming {
match msg {
Message::Binary(b) => match b.split_first() {
Some((&STDOUT, rest)) => Incoming::Output(rest.to_vec()),
_ => Incoming::Ignored,
},
Message::Text(t) => {
let Ok(v) = serde_json::from_str::<Value>(t) else {
return Incoming::Ignored;
};
match v.get("type").and_then(Value::as_str) {
Some("ready") => Incoming::Ready(
v.get("sandbox_id")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
),
Some("exit") => Incoming::Exit(
v.get("code").and_then(Value::as_i64).unwrap_or(0) as i32,
),
Some("error") => Incoming::Error(
v.get("message")
.and_then(Value::as_str)
.unwrap_or("the server reported an error with no message")
.to_string(),
),
_ => Incoming::Ignored,
}
}
Message::Close(_) => Incoming::Closed,
_ => Incoming::Ignored,
}
}
fn urlencode(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for b in s.bytes() {
match b {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' | b'/' => {
out.push(b as char)
}
_ => out.push_str(&format!("%{b:02X}")),
}
}
out
}
fn ws_url(base: &str, id: &str, query: &str) -> String {
let scheme_swapped = if let Some(rest) = base.strip_prefix("https://") {
format!("wss://{rest}")
} else if let Some(rest) = base.strip_prefix("http://") {
format!("ws://{rest}")
} else {
format!("ws://{base}")
};
format!("{scheme_swapped}/deployments/{}/shell?{query}", urlencode(id))
}
type Socket = tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>;
pub struct Shell {
socket: Socket,
sandbox_id: String,
exit: Option<ShellExit>,
last_error: Option<String>,
ping_at: std::time::Instant,
}
impl std::fmt::Debug for Shell {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Shell")
.field("sandbox_id", &self.sandbox_id)
.field("exit", &self.exit)
.finish()
}
}
impl Client {
pub async fn shell(&self, id: &str, opts: &ShellOptions) -> Result<Shell> {
let cfg = self.ws().ok_or_else(|| {
Error::Invalid(
"this client was built on a custom transport, which has no socket to \
open a shell on"
.into(),
)
})?;
Shell::connect(cfg, id, opts).await
}
}
impl Shell {
async fn connect(cfg: &WsConfig, id: &str, opts: &ShellOptions) -> Result<Self> {
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
let url = ws_url(&cfg.base, id, &opts.query());
if cfg.insecure && url.starts_with("wss://") {
return Err(Error::Invalid(
"insecure TLS is not supported for shell sessions — reach the admin \
listener over an SSH tunnel, or terminate TLS with a trusted certificate"
.into(),
));
}
let mut req = url
.as_str()
.into_client_request()
.map_err(|e| Error::Shell(format!("{url} is not a usable WebSocket URL: {e}")))?;
if let Some(h) = cfg.auth.header() {
req.headers_mut().insert(
"authorization",
h.parse()
.map_err(|_| Error::Invalid("the credential is not a valid header".into()))?,
);
}
let (mut socket, _) = tokio_tungstenite::connect_async(req)
.await
.map_err(|e| match &e {
tokio_tungstenite::tungstenite::Error::Http(resp) => {
let body = resp
.body()
.as_ref()
.map(|b| String::from_utf8_lossy(b).to_string())
.unwrap_or_default();
Error::from_response(
resp.status().as_u16(),
&body,
"deployment",
id,
cfg.auth.credential(),
)
}
_ => Error::Shell(e.to_string()),
})?;
let sandbox_id = loop {
match socket.next().await {
Some(Ok(msg)) => match parse_incoming(&msg) {
Incoming::Ready(id) => break id,
Incoming::Error(m) => return Err(Error::Shell(m)),
Incoming::Closed | Incoming::Exit(_) => {
return Err(Error::Shell(
"the shell closed before it was ready".into(),
));
}
_ => continue,
},
Some(Err(e)) => return Err(Error::Shell(e.to_string())),
None => {
return Err(Error::Shell(
"the shell closed before it was ready".into(),
));
}
}
};
Ok(Self {
socket,
sandbox_id,
exit: None,
last_error: None,
ping_at: std::time::Instant::now(),
})
}
pub fn sandbox_id(&self) -> &str {
&self.sandbox_id
}
pub async fn write(&mut self, bytes: &[u8]) -> Result<()> {
self.socket
.send(Message::Binary(frame_stdin(bytes)))
.await
.map_err(|e| Error::Shell(e.to_string()))
}
pub async fn resize(&mut self, cols: u16, rows: u16) -> Result<()> {
self.socket
.send(Message::Text(
json!({ "type": "resize", "cols": cols, "rows": rows }).to_string(),
))
.await
.map_err(|e| Error::Shell(e.to_string()))
}
pub async fn next(&mut self) -> Option<ShellEvent> {
loop {
if self.exit.is_some() {
return None;
}
if self.ping_at.elapsed() >= PING_INTERVAL {
self.ping_at = std::time::Instant::now();
if self.socket.send(Message::Ping(Vec::new())).await.is_err() {
self.finish(0);
return None;
}
}
let msg = match tokio::time::timeout(PING_INTERVAL, self.socket.next()).await {
Err(_) => continue,
Ok(None) => {
self.finish(0);
return None;
}
Ok(Some(Err(e))) => {
self.last_error = Some(e.to_string());
self.finish(0);
return None;
}
Ok(Some(Ok(m))) => m,
};
match parse_incoming(&msg) {
Incoming::Output(b) => return Some(ShellEvent::Output(b)),
Incoming::Error(m) => {
self.last_error = Some(m.clone());
return Some(ShellEvent::Error(m));
}
Incoming::Exit(code) => {
self.finish(code);
return None;
}
Incoming::Closed => {
self.finish(0);
return None;
}
Incoming::Ready(_) | Incoming::Ignored => continue,
}
}
}
fn finish(&mut self, code: i32) {
self.exit = Some(ShellExit {
code,
error: self.last_error.take(),
});
}
pub fn exit(&self) -> Option<&ShellExit> {
self.exit.as_ref()
}
pub async fn close(&mut self) -> Result<()> {
let _ = self.socket.close(None).await;
if self.exit.is_none() {
self.finish(0);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stdin_is_always_prefixed() {
assert_eq!(frame_stdin(b"ls\n"), vec![0x01, b'l', b's', b'\n']);
assert_eq!(frame_stdin(b""), vec![0x01], "even an empty write is framed");
assert_eq!(frame_stdin(&[0x02])[0], 0x01, "the payload is not the channel");
}
#[test]
fn output_is_unwrapped_and_other_channels_are_not() {
assert_eq!(
parse_incoming(&Message::Binary(vec![0x02, b'h', b'i'])),
Incoming::Output(b"hi".to_vec())
);
assert_eq!(
parse_incoming(&Message::Binary(vec![0x02])),
Incoming::Output(vec![]),
"an empty payload is still output"
);
assert_eq!(parse_incoming(&Message::Binary(vec![0x01, b'x'])), Incoming::Ignored);
assert_eq!(parse_incoming(&Message::Binary(vec![])), Incoming::Ignored);
}
#[test]
fn control_frames_are_understood() {
let t = |s: &str| parse_incoming(&Message::Text(s.to_string()));
assert_eq!(
t(r#"{"type":"ready","sandbox_id":"sb-1"}"#),
Incoming::Ready("sb-1".into())
);
assert_eq!(t(r#"{"type":"exit","code":3}"#), Incoming::Exit(3));
assert_eq!(t(r#"{"type":"exit","code":0}"#), Incoming::Exit(0));
assert_eq!(
t(r#"{"type":"error","message":"boom"}"#),
Incoming::Error("boom".into())
);
}
#[test]
fn unknown_frames_are_ignored_rather_than_fatal() {
let t = |s: &str| parse_incoming(&Message::Text(s.to_string()));
assert_eq!(t(r#"{"type":"something-new","x":1}"#), Incoming::Ignored);
assert_eq!(t("not json at all"), Incoming::Ignored);
assert_eq!(t("{}"), Incoming::Ignored);
assert_eq!(parse_incoming(&Message::Pong(vec![])), Incoming::Ignored);
}
#[test]
fn missing_fields_degrade_predictably() {
let t = |s: &str| parse_incoming(&Message::Text(s.to_string()));
assert_eq!(t(r#"{"type":"exit"}"#), Incoming::Exit(0));
assert_eq!(t(r#"{"type":"ready"}"#), Incoming::Ready(String::new()));
match t(r#"{"type":"error"}"#) {
Incoming::Error(m) => assert!(!m.is_empty(), "an error must always say something"),
other => panic!("{other:?}"),
}
}
#[test]
fn an_exit_after_an_error_is_not_clean() {
let died = ShellExit {
code: 0,
error: Some("gave up after 5 reconnect attempts".into()),
};
assert!(!died.is_clean(), "code 0 after an error is a crash, not a logout");
let logged_out = ShellExit { code: 0, error: None };
assert!(logged_out.is_clean());
let failed = ShellExit { code: 1, error: None };
assert!(!failed.is_clean());
}
#[test]
fn the_url_swaps_scheme_and_carries_the_options() {
let o = ShellOptions::default().size(120, 40);
assert_eq!(
ws_url("http://127.0.0.1:9090", "sb-1", &o.query()),
"ws://127.0.0.1:9090/deployments/sb-1/shell?cols=120&rows=40&wake=true"
);
assert!(ws_url("https://lb.example.com", "sb-1", "").starts_with("wss://"));
assert!(ShellOptions::default().no_wake().query().contains("wake=false"));
assert!(ShellOptions::default().vm("sb-1").query().ends_with("&sandbox_id=sb-1"));
assert!(!ShellOptions::default().query().contains("sandbox_id"));
let with_cwd = ShellOptions::default().cwd("/work space");
assert!(with_cwd.query().contains("cwd=/work%20space"), "{}", with_cwd.query());
}
#[test]
fn a_deployment_id_is_escaped_into_the_url() {
assert!(
ws_url("http://x:1", "a?b=c", "q=1").starts_with("ws://x:1/deployments/a%3Fb%3Dc/shell?"),
"{}",
ws_url("http://x:1", "a?b=c", "q=1")
);
}
}