forked from mincomk/rsh
205 lines
7.8 KiB
Rust
205 lines
7.8 KiB
Rust
use crate::config::Config;
|
|
use anyhow::{anyhow, Context, Result};
|
|
use futures_util::stream::{SplitSink, SplitStream};
|
|
use futures_util::{SinkExt, StreamExt};
|
|
use rsh_types::{BackendOpMsg, OpEvent, OpMsg, OpReq, OpResp};
|
|
use ssh_key::{HashAlg, PrivateKey};
|
|
use std::collections::HashMap;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use std::sync::Arc;
|
|
use tokio::net::TcpStream;
|
|
use tokio::sync::{mpsc, oneshot, Mutex};
|
|
use tokio_tungstenite::tungstenite::Message;
|
|
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
|
|
|
|
type Ws = WebSocketStream<MaybeTlsStream<TcpStream>>;
|
|
type WsSink = SplitSink<Ws, Message>;
|
|
type WsStream = SplitStream<Ws>;
|
|
|
|
enum PendingKind {
|
|
Once(oneshot::Sender<OpResp>),
|
|
Stream(mpsc::Sender<OpResp>),
|
|
}
|
|
|
|
pub struct AuthedClient {
|
|
next_id: AtomicU64,
|
|
pending: Arc<Mutex<HashMap<u64, PendingKind>>>,
|
|
events_tx: Arc<Mutex<Option<mpsc::Sender<OpEvent>>>>,
|
|
out: mpsc::Sender<OpMsg>,
|
|
_reader: tokio::task::JoinHandle<()>,
|
|
_writer: tokio::task::JoinHandle<()>,
|
|
}
|
|
|
|
impl AuthedClient {
|
|
pub async fn connect(cfg: &Config) -> Result<Self> {
|
|
let key_path = cfg.ssh_key_path();
|
|
let raw = std::fs::read(&key_path).with_context(|| format!("read {}", key_path.display()))?;
|
|
let mut priv_key = PrivateKey::from_openssh(&raw).context("parse openssh private key")?;
|
|
if priv_key.is_encrypted() {
|
|
let pw = inquire::Password::new(&format!("Passphrase for {}:", key_path.display()))
|
|
.without_confirmation()
|
|
.prompt()
|
|
.map_err(|e| anyhow!("password prompt: {e}"))?;
|
|
priv_key = priv_key.decrypt(pw.as_bytes()).context("decrypt private key")?;
|
|
}
|
|
let pub_openssh = priv_key.public_key().to_openssh().context("encode pubkey")?;
|
|
|
|
let ws_url = format!(
|
|
"{}/ws/op",
|
|
cfg.backend_url.trim_end_matches('/').trim_end_matches("/ws/op")
|
|
);
|
|
let (ws, _) = tokio_tungstenite::connect_async(&ws_url)
|
|
.await
|
|
.with_context(|| format!("ws connect {ws_url}"))?;
|
|
let (mut sink, mut stream) = ws.split();
|
|
|
|
send_msg(&mut sink, &OpMsg::AuthInit { pubkey_openssh: pub_openssh }).await?;
|
|
let challenge = recv_msg(&mut stream).await?;
|
|
let nonce = match challenge {
|
|
BackendOpMsg::Challenge { nonce } => nonce,
|
|
BackendOpMsg::AuthFail { reason } => return Err(anyhow!("auth fail: {reason}")),
|
|
other => return Err(anyhow!("unexpected: {other:?}")),
|
|
};
|
|
let sig = priv_key
|
|
.sign("rsh-auth", HashAlg::Sha512, &nonce)
|
|
.context("sign challenge")?;
|
|
let alg = sig.algorithm().as_str().to_string();
|
|
let pem = sig.to_pem(ssh_key::LineEnding::LF).context("encode signature")?;
|
|
send_msg(&mut sink, &OpMsg::AuthSign { signature: pem.into_bytes(), alg }).await?;
|
|
match recv_msg(&mut stream).await? {
|
|
BackendOpMsg::AuthOk => {}
|
|
BackendOpMsg::AuthFail { reason } => return Err(anyhow!("auth fail: {reason}")),
|
|
other => return Err(anyhow!("unexpected after AuthSign: {other:?}")),
|
|
}
|
|
|
|
let pending: Arc<Mutex<HashMap<u64, PendingKind>>> = Arc::new(Mutex::new(HashMap::new()));
|
|
let events_tx: Arc<Mutex<Option<mpsc::Sender<OpEvent>>>> = Arc::new(Mutex::new(None));
|
|
let (out_tx, mut out_rx) = mpsc::channel::<OpMsg>(64);
|
|
|
|
let writer = tokio::spawn(async move {
|
|
while let Some(m) = out_rx.recv().await {
|
|
let txt = match serde_json::to_string(&m) {
|
|
Ok(t) => t,
|
|
Err(_) => break,
|
|
};
|
|
if sink.send(Message::Text(txt)).await.is_err() {
|
|
break;
|
|
}
|
|
}
|
|
let _ = sink.close().await;
|
|
});
|
|
|
|
let pending_r = pending.clone();
|
|
let events_tx_r = events_tx.clone();
|
|
let reader = tokio::spawn(async move {
|
|
while let Some(msg) = stream.next().await {
|
|
let Ok(msg) = msg else { break };
|
|
let txt = match msg {
|
|
Message::Text(t) => t,
|
|
Message::Binary(b) => match String::from_utf8(b) { Ok(s) => s, Err(_) => continue },
|
|
Message::Close(_) => break,
|
|
_ => continue,
|
|
};
|
|
let parsed: BackendOpMsg = match serde_json::from_str(&txt) {
|
|
Ok(v) => v,
|
|
Err(_) => continue,
|
|
};
|
|
match parsed {
|
|
BackendOpMsg::Resp { id, body } => {
|
|
let mut p = pending_r.lock().await;
|
|
match p.get(&id) {
|
|
Some(PendingKind::Stream(tx)) => {
|
|
let _ = tx.send(body).await;
|
|
}
|
|
_ => {
|
|
if let Some(PendingKind::Once(tx)) = p.remove(&id) {
|
|
let _ = tx.send(body);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
BackendOpMsg::Event(ev) => {
|
|
let lock = events_tx_r.lock().await;
|
|
if let Some(tx) = lock.as_ref() {
|
|
let _ = tx.send(ev).await;
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
}
|
|
// Signal all pending callers that the connection is gone.
|
|
pending_r.lock().await.clear();
|
|
*events_tx_r.lock().await = None;
|
|
});
|
|
|
|
Ok(Self {
|
|
next_id: AtomicU64::new(1),
|
|
pending,
|
|
events_tx,
|
|
out: out_tx,
|
|
_reader: reader,
|
|
_writer: writer,
|
|
})
|
|
}
|
|
|
|
pub async fn req(&self, body: OpReq) -> Result<OpResp> {
|
|
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
|
|
let (tx, rx) = oneshot::channel();
|
|
self.pending.lock().await.insert(id, PendingKind::Once(tx));
|
|
self.out
|
|
.send(OpMsg::Req { id, body })
|
|
.await
|
|
.map_err(|_| anyhow!("send failed (connection closed)"))?;
|
|
rx.await.map_err(|_| anyhow!("backend disconnected"))
|
|
}
|
|
|
|
pub async fn req_stream(&self, body: OpReq) -> Result<(u64, mpsc::Receiver<OpResp>)> {
|
|
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
|
|
let (tx, rx) = mpsc::channel(64);
|
|
self.pending.lock().await.insert(id, PendingKind::Stream(tx));
|
|
self.out
|
|
.send(OpMsg::Req { id, body })
|
|
.await
|
|
.map_err(|_| anyhow!("send failed"))?;
|
|
Ok((id, rx))
|
|
}
|
|
|
|
pub async fn drop_stream(&self, id: u64) {
|
|
self.pending.lock().await.remove(&id);
|
|
}
|
|
|
|
pub async fn send_attach_io(&self, frame: rsh_types::AttachIOFrame) -> Result<()> {
|
|
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
|
|
self.out
|
|
.send(OpMsg::Req { id, body: OpReq::AttachIO(frame) })
|
|
.await
|
|
.map_err(|_| anyhow!("send failed"))
|
|
}
|
|
|
|
pub async fn take_events(&self) -> mpsc::Receiver<OpEvent> {
|
|
let (tx, rx) = mpsc::channel(64);
|
|
*self.events_tx.lock().await = Some(tx);
|
|
rx
|
|
}
|
|
}
|
|
|
|
async fn send_msg(sink: &mut WsSink, msg: &OpMsg) -> Result<()> {
|
|
sink.send(Message::Text(serde_json::to_string(msg)?))
|
|
.await
|
|
.map_err(|e| anyhow!(e))
|
|
}
|
|
|
|
async fn recv_msg(stream: &mut WsStream) -> Result<BackendOpMsg> {
|
|
loop {
|
|
let Some(msg) = stream.next().await else { return Err(anyhow!("connection closed")) };
|
|
let msg = msg?;
|
|
let txt = match msg {
|
|
Message::Text(t) => t,
|
|
Message::Binary(b) => String::from_utf8(b)?,
|
|
Message::Close(_) => return Err(anyhow!("connection closed")),
|
|
_ => continue,
|
|
};
|
|
return Ok(serde_json::from_str(&txt)?);
|
|
}
|
|
}
|