#![cfg(feature = "std")]
use std::collections::{BTreeMap, VecDeque};
use std::io::{ErrorKind, Read, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc::{self, Receiver, Sender, TryRecvError};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use purecrypto::hash::{Digest, Sha256};
use purecrypto::rng::RngCore;
use crate::auth::{ClientAuth, ClientCredential, ClientStep};
use crate::channel::{
ChannelEvent, ChannelOpen, ChannelRequest, ConnectionState, SSH_EXTENDED_DATA_STDERR,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
};
use crate::driver::{ClientDriver, Event};
use crate::error::{Error, Result};
use crate::hostkey::{HostKey, HostKeyVerify, host_key_verify_by_name};
use crate::known_hosts::{KnownHosts, LookupResult};
use crate::sftp::SftpClient;
pub use crate::stream::{ChannelEgress, ChannelStream};
use crate::transport::kex::{defaults, is_strict_kex_marker};
use crate::transport::ping::encode_ping;
use crate::transport::{KexAlgorithmsOwned, KexInit, KexRunner};
pub trait Transport: Read + Write + Send {
fn set_read_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()>;
fn set_write_timeout(&mut self, _t: Option<Duration>) -> std::io::Result<()> {
Ok(())
}
}
impl Transport for TcpStream {
fn set_read_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()> {
TcpStream::set_read_timeout(self, t)
}
fn set_write_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()> {
TcpStream::set_write_timeout(self, t)
}
}
const MAX_BANNER_LINES: usize = 32;
const MAX_EXEC_OUTPUT: usize = 64 * 1024 * 1024;
const MAX_KEX_STEPS: usize = 32;
const MAX_AUTH_STEPS: usize = 64;
const MAX_EXEC_ITER: usize = 1_000_000;
const SERVE_EGRESS_BACKLOG: usize = 32;
const MAX_SERVE_STEPS: usize = 100_000_000;
const SERVE_POLL_INTERVAL: Duration = Duration::from_millis(50);
pub enum HostKeyPolicy {
AcceptAny,
AcceptFingerprint([u8; 32]),
KnownHosts(KnownHostsPolicy),
}
pub struct KnownHostsPolicy {
pub store: Arc<Mutex<KnownHosts>>,
pub save_path: Option<PathBuf>,
pub hash_new: bool,
pub on_unknown: TofuAction,
pub on_mismatch: TofuAction,
}
impl KnownHostsPolicy {
pub fn strict(store: Arc<Mutex<KnownHosts>>) -> Self {
Self {
store,
save_path: None,
hash_new: false,
on_unknown: TofuAction::Reject,
on_mismatch: TofuAction::Reject,
}
}
}
pub type TofuPromptFn = dyn Fn(&str, u16, &str, &[u8]) -> bool + Send + Sync;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HostKeyChange {
Unknown,
Changed,
}
#[non_exhaustive]
pub struct HostKeyPrompt<'a> {
pub change: HostKeyChange,
pub host: &'a str,
pub port: u16,
pub key_type: &'a str,
pub key_blob: &'a [u8],
pub fingerprint: String,
pub existing: &'a [(String, Vec<u8>)],
}
impl HostKeyPrompt<'_> {
pub fn existing_fingerprints(&self) -> Vec<String> {
self.existing
.iter()
.map(|(_, blob)| host_key_fingerprint(blob))
.collect()
}
}
pub type HostKeyPromptFn = dyn Fn(&HostKeyPrompt<'_>) -> bool + Send + Sync;
#[non_exhaustive]
pub enum TofuAction {
Reject,
Accept,
Prompt(Arc<TofuPromptFn>),
AcceptWithWarning,
PromptDetailed(Arc<HostKeyPromptFn>),
}
pub struct Config {
pub host_key_policy: HostKeyPolicy,
pub timeout: Option<Duration>,
pub algorithms: AlgoOverrides,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AlgoOverrides {
pub ciphers: Option<Vec<String>>,
pub macs: Option<Vec<String>>,
pub kex_algorithms: Option<Vec<String>>,
pub host_key_algorithms: Option<Vec<String>>,
pub pubkey_accepted_algorithms: Option<Vec<String>>,
pub ca_signature_algorithms: Option<Vec<String>>,
pub compression: Option<bool>,
}
impl Config {
pub fn insecure() -> Self {
Self {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: None,
algorithms: AlgoOverrides::default(),
}
}
pub fn with_known_hosts(store: Arc<Mutex<KnownHosts>>) -> Self {
Self {
host_key_policy: HostKeyPolicy::KnownHosts(KnownHostsPolicy::strict(store)),
timeout: None,
algorithms: AlgoOverrides::default(),
}
}
}
pub struct ExecOutput {
pub stdout: Vec<u8>,
pub stderr: Vec<u8>,
pub exit_status: Option<u32>,
pub exit_signal: Option<String>,
}
#[derive(Debug, Clone)]
pub struct ForwardedTcpipOrigin {
pub bound_address: String,
pub bound_port: u16,
pub orig_address: String,
pub orig_port: u16,
}
pub type ForwardedTcpipCallback =
dyn Fn(ForwardedTcpipOrigin, ChannelStream) + Send + Sync + 'static;
#[derive(Debug, Clone)]
pub struct ForwardedStreamlocalOrigin {
pub socket_path: String,
}
pub type ForwardedStreamlocalCallback =
dyn Fn(ForwardedStreamlocalOrigin, ChannelStream) + Send + Sync + 'static;
pub type AuthAgentCallback = dyn Fn(ChannelStream) + Send + Sync + 'static;
pub type X11Callback = dyn Fn(ChannelStream) + Send + Sync + 'static;
pub struct ClientHandlers {
pub on_forwarded_tcpip: Option<Arc<ForwardedTcpipCallback>>,
pub on_forwarded_streamlocal: Option<Arc<ForwardedStreamlocalCallback>>,
pub on_auth_agent: Option<Arc<AuthAgentCallback>>,
pub on_x11: Option<Arc<X11Callback>>,
pub stop: Arc<AtomicBool>,
cmd_rx: Option<Receiver<ServeCommand>>,
}
impl Default for ClientHandlers {
fn default() -> Self {
Self::new()
}
}
impl ClientHandlers {
pub fn new() -> Self {
Self {
on_forwarded_tcpip: None,
on_forwarded_streamlocal: None,
on_auth_agent: None,
on_x11: None,
stop: Arc::new(AtomicBool::new(false)),
cmd_rx: None,
}
}
pub fn with_forwarded_tcpip(mut self, cb: Arc<ForwardedTcpipCallback>) -> Self {
self.on_forwarded_tcpip = Some(cb);
self
}
pub fn with_forwarded_streamlocal(mut self, cb: Arc<ForwardedStreamlocalCallback>) -> Self {
self.on_forwarded_streamlocal = Some(cb);
self
}
pub fn with_auth_agent(mut self, cb: Arc<AuthAgentCallback>) -> Self {
self.on_auth_agent = Some(cb);
self
}
pub fn with_x11(mut self, cb: Arc<X11Callback>) -> Self {
self.on_x11 = Some(cb);
self
}
pub fn with_serve_context(mut self) -> (Self, ServeContext) {
let (tx, rx) = mpsc::channel();
self.cmd_rx = Some(rx);
(self, ServeContext { cmd_tx: tx })
}
}
pub enum ServeCommand {
OpenDirectTcpip {
dest_host: String,
dest_port: u16,
orig_host: String,
orig_port: u16,
reply: mpsc::SyncSender<Result<ChannelStream>>,
},
OpenDirectStreamlocal {
socket_path: String,
reply: mpsc::SyncSender<Result<ChannelStream>>,
},
}
#[derive(Clone)]
pub struct ServeContext {
cmd_tx: Sender<ServeCommand>,
}
impl ServeContext {
pub fn open_direct_tcpip(
&self,
dest_host: &str,
dest_port: u16,
orig_host: &str,
orig_port: u16,
) -> Result<ChannelStream> {
let (reply_tx, reply_rx) = mpsc::sync_channel::<Result<ChannelStream>>(1);
self.cmd_tx
.send(ServeCommand::OpenDirectTcpip {
dest_host: dest_host.to_string(),
dest_port,
orig_host: orig_host.to_string(),
orig_port,
reply: reply_tx,
})
.map_err(|_| Error::Protocol("serve loop terminated"))?;
reply_rx
.recv()
.map_err(|_| Error::Protocol("serve loop terminated"))?
}
pub fn open_direct_streamlocal(&self, socket_path: &str) -> Result<ChannelStream> {
let (reply_tx, reply_rx) = mpsc::sync_channel::<Result<ChannelStream>>(1);
self.cmd_tx
.send(ServeCommand::OpenDirectStreamlocal {
socket_path: socket_path.to_string(),
reply: reply_tx,
})
.map_err(|_| Error::Protocol("serve loop terminated"))?;
reply_rx
.recv()
.map_err(|_| Error::Protocol("serve loop terminated"))?
}
}
struct PendingOutboundOpen {
stream: Option<ChannelStream>,
ingress_tx: Sender<Option<Vec<u8>>>,
egress_rx: Option<Receiver<ChannelEgress>>,
reply: mpsc::SyncSender<Result<ChannelStream>>,
}
struct ServeRuntime {
ingress_tx: Sender<Option<Vec<u8>>>,
egress_rx: Receiver<ChannelEgress>,
pending_data: Vec<u8>,
pending_eof: bool,
pending_close: bool,
eof_sent: bool,
close_sent: bool,
}
pub struct Client {
stream: Box<dyn Transport>,
out_buf: VecDeque<u8>,
pub(crate) conn: ConnectionState,
driver: ClientDriver,
algo_overrides: AlgoOverrides,
request_auth_agent: bool,
request_x11: Option<X11ReqArgs>,
session_env: Vec<(String, String)>,
keepalive: Option<(Duration, u32)>,
request_pty: Option<PtyReqArgs>,
tcpip_forward_grants: Vec<(String, u16)>,
streamlocal_forward_grants: Vec<String>,
}
#[derive(Clone)]
struct X11ReqArgs {
single_connection: bool,
auth_protocol: String,
auth_cookie: String,
screen: u32,
}
#[derive(Clone)]
struct PtyReqArgs {
term: String,
cols: u32,
rows: u32,
px_w: u32,
px_h: u32,
modes: Vec<u8>,
}
impl Client {
pub fn connect<A: ToSocketAddrs>(addr: A, cfg: Config) -> Result<Self> {
let stream = TcpStream::connect(addr)?;
if let Some(t) = cfg.timeout {
stream.set_read_timeout(Some(t))?;
stream.set_write_timeout(Some(t))?;
}
stream.set_nodelay(true)?;
Self::from_transport(Box::new(stream), "", 0, cfg)
}
pub fn connect_to_host(host: &str, port: u16, cfg: Config) -> Result<Self> {
let stream = TcpStream::connect((host, port))?;
if let Some(t) = cfg.timeout {
stream.set_read_timeout(Some(t))?;
stream.set_write_timeout(Some(t))?;
}
stream.set_nodelay(true)?;
Self::from_transport(Box::new(stream), host, port, cfg)
}
pub fn connect_via(
stream: Box<dyn Transport>,
host: &str,
port: u16,
cfg: Config,
) -> Result<Self> {
Self::from_transport(stream, host, port, cfg)
}
fn from_transport(
stream: Box<dyn Transport>,
host: &str,
port: u16,
cfg: Config,
) -> Result<Self> {
let host_key_policy = cfg.host_key_policy;
let target_host = host.to_string();
let ca_sig_algos = cfg.algorithms.ca_signature_algorithms.clone();
let verifier_factory: crate::driver::client::VerifierFactory =
Box::new(move |reply: &[u8], runner: &KexRunner| {
build_verifier(
reply,
&host_key_policy,
runner,
&target_host,
port,
ca_sig_algos.as_deref(),
unix_now(),
)
});
let driver = ClientDriver::new(cfg.algorithms.clone(), verifier_factory);
let mut me = Self {
stream,
out_buf: VecDeque::new(),
conn: ConnectionState::new(),
driver,
algo_overrides: cfg.algorithms,
request_auth_agent: false,
request_x11: None,
session_env: Vec::new(),
keepalive: None,
request_pty: None,
tcpip_forward_grants: Vec::new(),
streamlocal_forward_grants: Vec::new(),
};
me.driver.start(Instant::now())?;
me.drive_handshake()?;
Ok(me)
}
pub fn set_rekey_policy(&mut self, policy: crate::transport::RekeyPolicy) {
self.driver.set_rekey_policy(policy);
}
pub fn authenticate(&mut self, user: &str, credentials: Vec<ClientCredential>) -> Result<()> {
let mut auth = ClientAuth::new(user, self.driver.session_id().to_vec());
if let Some(accepted) = self.algo_overrides.pubkey_accepted_algorithms.clone() {
auth.set_pubkey_accepted(accepted);
}
if let Some(ext) = self.driver.peer_ext_info()
&& let Some(algs) = ext.server_sig_algs.as_deref()
{
auth.set_server_sig_algs(algs);
}
for c in credentials {
auth.add_credential(c);
}
self.run_auth(auth)
}
pub fn run_auth(&mut self, mut auth: ClientAuth) -> Result<()> {
let first = auth.start();
self.write_payload(&first)?;
for _ in 0..MAX_AUTH_STEPS {
let payload = self.read_one_packet()?;
match auth.on_packet(&payload)? {
ClientStep::Send(p) => self.write_payload(&p)?,
ClientStep::Success => {
self.driver.notify_auth_success();
return Ok(());
}
ClientStep::Failed { .. } => return Err(Error::AuthFailed),
ClientStep::Banner { .. } => {}
ClientStep::Idle => {}
}
}
Err(Error::Protocol("auth: too many steps without termination"))
}
pub fn new_auth_driver(&self, user: &str) -> ClientAuth {
let mut auth = ClientAuth::new(user, self.driver.session_id().to_vec());
if let Some(accepted) = self.algo_overrides.pubkey_accepted_algorithms.clone() {
auth.set_pubkey_accepted(accepted);
}
if let Some(ext) = self.driver.peer_ext_info()
&& let Some(algs) = ext.server_sig_algs.as_deref()
{
auth.set_server_sig_algs(algs);
}
auth
}
pub fn authenticate_password(&mut self, user: &str, password: &str) -> Result<()> {
self.authenticate(user, vec![ClientCredential::Password(password.into())])
}
pub fn session_id(&self) -> &[u8] {
self.driver.session_id()
}
pub fn peer_ext_info(&self) -> Option<&crate::transport::ExtInfo> {
self.driver.peer_ext_info()
}
pub fn negotiated_compression(&self) -> Option<(String, String)> {
self.driver.negotiated_compression()
}
pub fn authenticate_publickey(
&mut self,
user: &str,
key: Box<dyn HostKey + Send>,
) -> Result<()> {
self.authenticate(user, vec![ClientCredential::PublicKey(key)])
}
pub fn open_session_for_agent_forward(&mut self) -> Result<u32> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut iter_guard = 0usize;
loop {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("agent-forward: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => break,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("agent-forward: channel open failed"));
}
_ => {}
}
}
let p = self
.conn
.send_request(local_id, ChannelRequest::AuthAgentReq, false)?;
self.write_payload(&p)?;
Ok(local_id)
}
pub fn open_session_for_x11_forward(
&mut self,
single_connection: bool,
auth_protocol: &str,
auth_cookie: &str,
screen: u32,
) -> Result<u32> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut iter_guard = 0usize;
loop {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("x11-forward: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => break,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("x11-forward: channel open failed"));
}
_ => {}
}
}
let p = self.conn.send_request(
local_id,
ChannelRequest::X11Req {
single_connection,
auth_protocol: auth_protocol.to_string(),
auth_cookie: auth_cookie.to_string(),
screen,
},
false,
)?;
self.write_payload(&p)?;
Ok(local_id)
}
pub fn set_read_timeout(&mut self, t: Option<core::time::Duration>) -> std::io::Result<()> {
self.stream.set_read_timeout(t)?;
self.stream.set_write_timeout(t)
}
pub fn close_session(&mut self, channel: u32) -> Result<()> {
let payload = match self.conn.send_close(channel) {
Ok(p) => p,
Err(_) => return Ok(()),
};
self.write_payload(&payload)?;
Ok(())
}
pub fn set_request_auth_agent_forwarding(&mut self, on: bool) {
self.request_auth_agent = on;
}
pub(crate) fn maybe_send_auth_agent_req(&mut self, channel: u32) -> Result<()> {
if self.request_auth_agent {
let p = self
.conn
.send_request(channel, ChannelRequest::AuthAgentReq, false)?;
self.write_payload(&p)?;
}
Ok(())
}
pub fn set_request_x11_forwarding(&mut self, args: Option<(bool, String, String, u32)>) {
self.request_x11 =
args.map(
|(single_connection, auth_protocol, auth_cookie, screen)| X11ReqArgs {
single_connection,
auth_protocol,
auth_cookie,
screen,
},
);
}
pub(crate) fn maybe_send_x11_req(&mut self, channel: u32) -> Result<()> {
if let Some(args) = self.request_x11.clone() {
let p = self.conn.send_request(
channel,
ChannelRequest::X11Req {
single_connection: args.single_connection,
auth_protocol: args.auth_protocol,
auth_cookie: args.auth_cookie,
screen: args.screen,
},
false,
)?;
self.write_payload(&p)?;
}
Ok(())
}
pub fn set_session_env(&mut self, env: Vec<(String, String)>) {
self.session_env = env;
}
pub(crate) fn maybe_send_env(&mut self, channel: u32) -> Result<()> {
for (name, value) in self.session_env.clone() {
let p = self
.conn
.send_request(channel, ChannelRequest::Env { name, value }, false)?;
self.write_payload(&p)?;
}
Ok(())
}
pub fn set_keepalive(&mut self, interval_secs: u32, count_max: u32) {
self.keepalive = if interval_secs == 0 {
None
} else {
Some((Duration::from_secs(interval_secs as u64), count_max.max(1)))
};
}
pub fn set_request_pty(&mut self, args: Option<(String, u32, u32, u32, u32, Vec<u8>)>) {
self.request_pty = args.map(|(term, cols, rows, px_w, px_h, modes)| PtyReqArgs {
term,
cols,
rows,
px_w,
px_h,
modes,
});
}
pub(crate) fn maybe_send_pty_req(&mut self, channel: u32) -> Result<()> {
if let Some(args) = self.request_pty.clone() {
let p = self.conn.send_request(
channel,
ChannelRequest::PtyReq {
term: args.term,
cols: args.cols,
rows: args.rows,
px_w: args.px_w,
px_h: args.px_h,
modes: args.modes,
},
false,
)?;
self.write_payload(&p)?;
}
Ok(())
}
pub fn exec(&mut self, command: &str) -> Result<ExecOutput> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut opened = false;
let mut iter_guard = 0usize;
while !opened {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("exec: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => {
opened = true;
}
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("channel open failed"));
}
_ => {}
}
}
self.maybe_send_auth_agent_req(local_id)?;
self.maybe_send_x11_req(local_id)?;
self.maybe_send_env(local_id)?;
self.maybe_send_pty_req(local_id)?;
let exec_req = self.conn.send_request(
local_id,
ChannelRequest::Exec {
command: command.into(),
},
true,
)?;
self.write_payload(&exec_req)?;
let mut exec_accepted = false;
iter_guard = 0;
while !exec_accepted {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("exec: request loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::Success { channel } if channel == local_id => exec_accepted = true,
ChannelEvent::Failure { channel } if channel == local_id => {
return Err(Error::Protocol("exec request denied"));
}
_ => {}
}
}
let mut out = ExecOutput {
stdout: Vec::new(),
stderr: Vec::new(),
exit_status: None,
exit_signal: None,
};
let mut local_eof_sent = false;
let mut local_close_sent = false;
let mut remote_close_seen = false;
for _ in 0..MAX_EXEC_ITER {
if remote_close_seen && local_close_sent {
break;
}
let payload = self.read_one_packet()?;
let ev = self.conn.on_packet(&payload)?;
match ev {
ChannelEvent::Data { channel, data } if channel == local_id => {
if out.stdout.len() + out.stderr.len() + data.len() > MAX_EXEC_OUTPUT {
return Err(Error::Protocol("exec output too large"));
}
let n = data.len() as u32;
out.stdout.extend_from_slice(&data);
if let Some(adj) = self.conn.replenish_window(local_id, n)? {
self.write_payload(&adj)?;
}
}
ChannelEvent::ExtendedData {
channel,
code,
data,
} if channel == local_id => {
if out.stdout.len() + out.stderr.len() + data.len() > MAX_EXEC_OUTPUT {
return Err(Error::Protocol("exec output too large"));
}
let n = data.len() as u32;
if code == SSH_EXTENDED_DATA_STDERR {
out.stderr.extend_from_slice(&data);
} else {
out.stdout.extend_from_slice(&data);
}
if let Some(adj) = self.conn.replenish_window(local_id, n)? {
self.write_payload(&adj)?;
}
}
ChannelEvent::Request {
channel,
request,
want_reply,
} if channel == local_id => {
match request {
ChannelRequest::ExitStatus { code } => out.exit_status = Some(code),
ChannelRequest::ExitSignal { name, .. } => out.exit_signal = Some(name),
_ => {}
}
if want_reply {
let p = self.conn.send_request_failure(local_id)?;
self.write_payload(&p)?;
}
}
ChannelEvent::Eof { channel } if channel == local_id && !local_eof_sent => {
let p = self.conn.send_eof(local_id)?;
self.write_payload(&p)?;
local_eof_sent = true;
}
ChannelEvent::Close { channel } if channel == local_id => {
remote_close_seen = true;
if !local_close_sent {
let p = self.conn.send_close(local_id)?;
self.write_payload(&p)?;
local_close_sent = true;
}
}
ChannelEvent::WindowAdjust { .. } => {}
_ => {}
}
}
if !(remote_close_seen && local_close_sent) {
return Err(Error::Protocol("exec: drain loop exceeded iteration cap"));
}
Ok(out)
}
pub fn shell_with_stdin(
&mut self,
term: &str,
cols: u32,
rows: u32,
stdin: &[u8],
) -> Result<ExecOutput> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut opened = false;
let mut iter_guard = 0usize;
while !opened {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("shell: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => {
opened = true;
}
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("channel open failed"));
}
_ => {}
}
}
self.maybe_send_auth_agent_req(local_id)?;
self.maybe_send_x11_req(local_id)?;
let pty_req = self.conn.send_request(
local_id,
ChannelRequest::PtyReq {
term: term.into(),
cols,
rows,
px_w: 0,
px_h: 0,
modes: Vec::new(),
},
true,
)?;
self.write_payload(&pty_req)?;
self.await_request_reply(local_id, "pty-req")?;
let shell_req = self
.conn
.send_request(local_id, ChannelRequest::Shell, true)?;
self.write_payload(&shell_req)?;
self.await_request_reply(local_id, "shell")?;
if !stdin.is_empty() {
let mut off = 0usize;
iter_guard = 0;
while off < stdin.len() {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("shell: stdin drain loop did not converge"));
}
let (payload, taken) = self.conn.send_data(local_id, &stdin[off..])?;
if taken == 0 {
let pkt = self.read_one_packet()?;
match self.conn.on_packet(&pkt)? {
ChannelEvent::WindowAdjust { channel, .. } if channel == local_id => {}
ChannelEvent::Close { channel } if channel == local_id => {
return Err(Error::Protocol(
"shell: peer closed channel before stdin drain",
));
}
_ => {}
}
continue;
}
self.write_payload(&payload)?;
off += taken;
}
}
let eof = self.conn.send_eof(local_id)?;
self.write_payload(&eof)?;
let mut out = ExecOutput {
stdout: Vec::new(),
stderr: Vec::new(),
exit_status: None,
exit_signal: None,
};
let mut local_close_sent = false;
let mut remote_close_seen = false;
for _ in 0..MAX_EXEC_ITER {
if remote_close_seen && local_close_sent {
break;
}
let payload = self.read_one_packet()?;
let ev = self.conn.on_packet(&payload)?;
match ev {
ChannelEvent::Data { channel, data } if channel == local_id => {
if out.stdout.len() + out.stderr.len() + data.len() > MAX_EXEC_OUTPUT {
return Err(Error::Protocol("shell output too large"));
}
let n = data.len() as u32;
out.stdout.extend_from_slice(&data);
if let Some(adj) = self.conn.replenish_window(local_id, n)? {
self.write_payload(&adj)?;
}
}
ChannelEvent::ExtendedData {
channel,
code,
data,
} if channel == local_id => {
if out.stdout.len() + out.stderr.len() + data.len() > MAX_EXEC_OUTPUT {
return Err(Error::Protocol("shell output too large"));
}
let n = data.len() as u32;
if code == SSH_EXTENDED_DATA_STDERR {
out.stderr.extend_from_slice(&data);
} else {
out.stdout.extend_from_slice(&data);
}
if let Some(adj) = self.conn.replenish_window(local_id, n)? {
self.write_payload(&adj)?;
}
}
ChannelEvent::Request {
channel,
request,
want_reply,
} if channel == local_id => {
match request {
ChannelRequest::ExitStatus { code } => out.exit_status = Some(code),
ChannelRequest::ExitSignal { name, .. } => out.exit_signal = Some(name),
_ => {}
}
if want_reply {
let p = self.conn.send_request_failure(local_id)?;
self.write_payload(&p)?;
}
}
ChannelEvent::Eof { channel } if channel == local_id => {}
ChannelEvent::Close { channel } if channel == local_id => {
remote_close_seen = true;
if !local_close_sent {
let p = self.conn.send_close(local_id)?;
self.write_payload(&p)?;
local_close_sent = true;
}
}
ChannelEvent::WindowAdjust { .. } => {}
_ => {}
}
}
if !(remote_close_seen && local_close_sent) {
return Err(Error::Protocol("shell: drain loop exceeded iteration cap"));
}
Ok(out)
}
pub fn subsystem_once(&mut self, name: &str, stdin: &[u8]) -> Result<Vec<u8>> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut opened = false;
let mut iter_guard = 0usize;
while !opened {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("subsystem: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => opened = true,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("channel open failed"));
}
_ => {}
}
}
let sub_req = self.conn.send_request(
local_id,
ChannelRequest::Subsystem { name: name.into() },
true,
)?;
self.write_payload(&sub_req)?;
self.await_request_reply(local_id, "subsystem")?;
if !stdin.is_empty() {
let mut off = 0usize;
iter_guard = 0;
while off < stdin.len() {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol(
"subsystem: stdin drain loop did not converge",
));
}
let (payload, taken) = self.conn.send_data(local_id, &stdin[off..])?;
if taken == 0 {
let pkt = self.read_one_packet()?;
match self.conn.on_packet(&pkt)? {
ChannelEvent::WindowAdjust { channel, .. } if channel == local_id => {}
ChannelEvent::Close { channel } if channel == local_id => {
return Err(Error::Protocol(
"subsystem: peer closed channel before stdin drain",
));
}
_ => {}
}
continue;
}
self.write_payload(&payload)?;
off += taken;
}
}
let eof = self.conn.send_eof(local_id)?;
self.write_payload(&eof)?;
let mut out = Vec::<u8>::new();
let mut local_close_sent = false;
let mut remote_close_seen = false;
for _ in 0..MAX_EXEC_ITER {
if remote_close_seen && local_close_sent {
break;
}
let payload = self.read_one_packet()?;
let ev = self.conn.on_packet(&payload)?;
match ev {
ChannelEvent::Data { channel, data } if channel == local_id => {
if out.len() + data.len() > MAX_EXEC_OUTPUT {
return Err(Error::Protocol("subsystem output too large"));
}
let n = data.len() as u32;
out.extend_from_slice(&data);
if let Some(adj) = self.conn.replenish_window(local_id, n)? {
self.write_payload(&adj)?;
}
}
ChannelEvent::ExtendedData {
channel,
code: _,
data,
} if channel == local_id => {
let n = data.len() as u32;
if let Some(adj) = self.conn.replenish_window(local_id, n)? {
self.write_payload(&adj)?;
}
}
ChannelEvent::Eof { channel } if channel == local_id => {}
ChannelEvent::Close { channel } if channel == local_id => {
remote_close_seen = true;
if !local_close_sent {
let p = self.conn.send_close(local_id)?;
self.write_payload(&p)?;
local_close_sent = true;
}
}
ChannelEvent::WindowAdjust { .. } => {}
_ => {}
}
}
if !(remote_close_seen && local_close_sent) {
return Err(Error::Protocol(
"subsystem: drain loop exceeded iteration cap",
));
}
Ok(out)
}
#[cfg_attr(
feature = "multichannel",
deprecated(
since = "0.0.2",
note = "Use SharedClient::sftp instead; the borrow-based API \
prevents multiple concurrent channels on one connection."
)
)]
pub fn sftp(&mut self) -> Result<SftpClient<ClientChannelStream<'_>>> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut opened = false;
let mut iter_guard = 0usize;
while !opened {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("sftp: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => opened = true,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("channel open failed"));
}
_ => {}
}
}
self.maybe_send_auth_agent_req(local_id)?;
self.maybe_send_x11_req(local_id)?;
let sub_req = self.conn.send_request(
local_id,
ChannelRequest::Subsystem {
name: "sftp".into(),
},
true,
)?;
self.write_payload(&sub_req)?;
self.await_request_reply(local_id, "subsystem")?;
let stream = ClientChannelStream {
client: self,
channel: local_id,
read_buf: Vec::new(),
stderr_buf: Vec::new(),
remote_eof: false,
local_close_sent: false,
};
match SftpClient::new(stream) {
Ok(c) => Ok(c),
Err(e) => Err(Error::Protocol(match e {
crate::sftp::SftpError::Protocol(s) => s,
_ => "sftp: handshake failed",
})),
}
}
#[cfg_attr(
feature = "multichannel",
deprecated(
since = "0.0.2",
note = "Use SharedClient::exec_stream for multi-channel support."
)
)]
pub fn exec_stream(&mut self, command: &str) -> Result<ClientChannelStream<'_>> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::Session)?;
self.write_payload(&open_payload)?;
let mut opened = false;
let mut iter_guard = 0usize;
while !opened {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("exec_stream: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => opened = true,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("exec_stream: channel open failed"));
}
_ => {}
}
}
self.maybe_send_auth_agent_req(local_id)?;
self.maybe_send_x11_req(local_id)?;
self.maybe_send_env(local_id)?;
self.maybe_send_pty_req(local_id)?;
let exec_req = self.conn.send_request(
local_id,
ChannelRequest::Exec {
command: command.into(),
},
true,
)?;
self.write_payload(&exec_req)?;
self.await_request_reply(local_id, "exec")?;
Ok(ClientChannelStream {
client: self,
channel: local_id,
read_buf: Vec::new(),
stderr_buf: Vec::new(),
remote_eof: false,
local_close_sent: false,
})
}
#[cfg_attr(
feature = "multichannel",
deprecated(
since = "0.0.2",
note = "Use SharedClient::open_direct_tcpip for multi-channel support."
)
)]
pub fn open_direct_tcpip(
&mut self,
dest_host: &str,
dest_port: u16,
orig_host: &str,
orig_port: u16,
) -> Result<ClientChannelStream<'_>> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::DirectTcpip {
dest_host: dest_host.to_string(),
dest_port: dest_port as u32,
orig_host: orig_host.to_string(),
orig_port: orig_port as u32,
})?;
self.write_payload(&open_payload)?;
let mut iter_guard = 0usize;
loop {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol("direct-tcpip: open loop did not converge"));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => break,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("direct-tcpip: open failed"));
}
_ => {}
}
}
Ok(ClientChannelStream {
client: self,
channel: local_id,
read_buf: Vec::new(),
stderr_buf: Vec::new(),
remote_eof: false,
local_close_sent: false,
})
}
#[cfg_attr(
feature = "multichannel",
deprecated(
since = "0.0.7",
note = "Use ServeContext::open_direct_streamlocal for multi-channel support."
)
)]
pub fn open_direct_streamlocal(
&mut self,
socket_path: &str,
) -> Result<ClientChannelStream<'_>> {
let (local_id, open_payload) = self.conn.open(ChannelOpen::DirectStreamlocal {
socket_path: socket_path.to_string(),
})?;
self.write_payload(&open_payload)?;
let mut iter_guard = 0usize;
loop {
iter_guard += 1;
if iter_guard > MAX_EXEC_ITER {
return Err(Error::Protocol(
"direct-streamlocal: open loop did not converge",
));
}
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::OpenConfirmed { channel } if channel == local_id => break,
ChannelEvent::OpenFailed { channel, .. } if channel == local_id => {
return Err(Error::Protocol("direct-streamlocal: open failed"));
}
_ => {}
}
}
Ok(ClientChannelStream {
client: self,
channel: local_id,
read_buf: Vec::new(),
stderr_buf: Vec::new(),
remote_eof: false,
local_close_sent: false,
})
}
pub fn scp_send_to(
&mut self,
sources: &[&std::path::Path],
remote_dest: &str,
opts: crate::scp::ScpSendOptions,
) -> Result<()> {
let cmd = build_scp_to_cmd(remote_dest, &opts)?;
#[allow(deprecated)]
let mut stream = self.exec_stream(&cmd)?;
let result = (|| -> Result<()> {
let mut sender = crate::scp::Sender::new(&mut stream)
.map_err(|e| scp_proto(e, "scp_send_to: handshake"))?;
for src in sources {
sender
.send_path(src, &opts)
.map_err(|e| scp_proto(e, "scp_send_to: send_path"))?;
}
Ok(())
})();
let stderr = stream.take_stderr();
match result {
Ok(()) => Ok(()),
Err(e) => {
if !stderr.is_empty() {
let msg = String::from_utf8_lossy(&stderr).trim().to_string();
eprintln!("scp_send_to: remote stderr: {}", msg);
}
Err(e)
}
}
}
pub fn scp_recv_from(
&mut self,
remote_source: &str,
local_dest: &std::path::Path,
mut opts: crate::scp::ScpRecvOptions,
) -> Result<()> {
let cmd = build_scp_from_cmd(remote_source, &opts)?;
if !opts.target_is_file && !opts.recursive {
if let Ok(md) = std::fs::metadata(local_dest) {
if !md.is_dir() {
opts.target_is_file = true;
}
} else {
opts.target_is_file = true;
}
}
let expected_name: Option<String> = if opts.recursive {
None
} else {
let base = remote_source.rsplit('/').next().unwrap_or(remote_source);
let is_glob = base.contains(['*', '?', '[']);
if base.is_empty() || is_glob {
None
} else {
Some(base.to_string())
}
};
#[allow(deprecated)]
let mut stream = self.exec_stream(&cmd)?;
let result = (|| -> Result<()> {
let mut recv = crate::scp::Receiver::new(&mut stream, local_dest, opts)
.map_err(|e| scp_proto(e, "scp_recv_from: handshake"))?
.with_expected_name(expected_name.as_deref());
recv.run().map_err(|e| scp_proto(e, "scp_recv_from: run"))?;
Ok(())
})();
let stderr = stream.take_stderr();
match result {
Ok(()) => Ok(()),
Err(e) => {
if !stderr.is_empty() {
let msg = String::from_utf8_lossy(&stderr).trim().to_string();
eprintln!("scp_recv_from: remote stderr: {}", msg);
}
Err(e)
}
}
}
pub fn request_tcpip_forward(&mut self, bind_address: &str, bind_port: u16) -> Result<u16> {
use crate::channel::GlobalRequest;
let payload = self.conn.send_global_request(
GlobalRequest::TcpipForward {
bind_address: bind_address.to_string(),
bind_port: bind_port as u32,
},
true,
);
self.write_payload(&payload)?;
let data = self.await_global_reply("tcpip-forward")?;
let granted_port = if bind_port == 0 {
let mut r = crate::format::Reader::new(&data);
let p = r
.read_u32()
.map_err(|_| Error::Protocol("tcpip-forward: server omitted assigned-port tail"))?;
if p > u16::MAX as u32 {
return Err(Error::Protocol(
"tcpip-forward: server returned out-of-range port",
));
}
p as u16
} else {
bind_port
};
self.tcpip_forward_grants
.push((bind_address.to_string(), granted_port));
Ok(granted_port)
}
pub fn cancel_tcpip_forward(&mut self, bind_address: &str, bind_port: u16) -> Result<()> {
use crate::channel::GlobalRequest;
let payload = self.conn.send_global_request(
GlobalRequest::CancelTcpipForward {
bind_address: bind_address.to_string(),
bind_port: bind_port as u32,
},
true,
);
self.write_payload(&payload)?;
let _ = self.await_global_reply("cancel-tcpip-forward")?;
if let Some(idx) = self
.tcpip_forward_grants
.iter()
.position(|(a, p)| a == bind_address && *p == bind_port)
.or_else(|| {
self.tcpip_forward_grants
.iter()
.position(|(_, p)| *p == bind_port)
})
{
self.tcpip_forward_grants.swap_remove(idx);
}
Ok(())
}
pub fn request_streamlocal_forward(&mut self, socket_path: &str) -> Result<()> {
use crate::channel::GlobalRequest;
let payload = self.conn.send_global_request(
GlobalRequest::StreamlocalForward {
socket_path: socket_path.to_string(),
},
true,
);
self.write_payload(&payload)?;
let _ = self.await_global_reply("streamlocal-forward@openssh.com")?;
self.streamlocal_forward_grants
.push(socket_path.to_string());
Ok(())
}
pub fn cancel_streamlocal_forward(&mut self, socket_path: &str) -> Result<()> {
use crate::channel::GlobalRequest;
let payload = self.conn.send_global_request(
GlobalRequest::CancelStreamlocalForward {
socket_path: socket_path.to_string(),
},
true,
);
self.write_payload(&payload)?;
let _ = self.await_global_reply("cancel-streamlocal-forward@openssh.com")?;
if let Some(idx) = self
.streamlocal_forward_grants
.iter()
.position(|p| p == socket_path)
{
self.streamlocal_forward_grants.swap_remove(idx);
}
Ok(())
}
pub fn serve(&mut self, handlers: ClientHandlers) -> Result<()> {
let mut runtimes: BTreeMap<u32, ServeRuntime> = BTreeMap::new();
let mut pending_opens: BTreeMap<u32, PendingOutboundOpen> = BTreeMap::new();
let _ = self.stream.set_read_timeout(Some(SERVE_POLL_INTERVAL));
let mut steps = 0usize;
let mut last_activity = Instant::now();
let mut probes_pending: u32 = 0;
let result = loop {
steps += 1;
if steps > MAX_SERVE_STEPS {
break Err(Error::Protocol("serve: step cap exceeded"));
}
if let Some((interval, count_max)) = self.keepalive
&& !self.driver.is_kexing()
&& last_activity.elapsed() >= interval
{
if probes_pending >= count_max {
break Err(Error::Protocol(
"serve: server keepalive timed out (ServerAliveCountMax exceeded)",
));
}
let probe = self
.conn
.send_global_request(crate::channel::GlobalRequest::Keepalive, true);
if let Err(e) = self.write_payload(&probe) {
break Err(e);
}
probes_pending += 1;
last_activity = Instant::now();
}
if !self.driver.is_kexing()
&& let Some(rx) = handlers.cmd_rx.as_ref()
&& let Err(e) = serve_drain_commands(self, rx, &mut pending_opens)
{
break Err(e);
}
if !self.driver.is_kexing() {
if let Err(e) = serve_drain_runtimes(self, &mut runtimes) {
break Err(e);
}
runtimes.retain(|_, rt| !rt.close_sent);
}
if handlers.stop.load(Ordering::SeqCst)
&& runtimes.is_empty()
&& pending_opens.is_empty()
{
break Ok(());
}
let payload = match self.read_one_packet_maybe_timeout() {
Ok(Some(p)) => p,
Ok(None) => continue, Err(e) => break Err(e),
};
if self.keepalive.is_some() {
probes_pending = 0;
last_activity = Instant::now();
}
if let Err(e) =
serve_dispatch_packet(self, &handlers, &mut runtimes, &mut pending_opens, &payload)
{
break Err(e);
}
};
let stale_opens = core::mem::take(&mut pending_opens);
for (_ch, po) in stale_opens {
let _ = po.reply.send(Err(Error::Protocol("serve loop terminated")));
}
runtimes.clear();
let _ = self.stream.set_read_timeout(None);
result
}
pub(crate) fn read_one_packet_maybe_timeout(&mut self) -> Result<Option<Vec<u8>>> {
match self.read_one_packet() {
Ok(p) => Ok(Some(p)),
Err(Error::Io(e))
if e.kind() == ErrorKind::WouldBlock || e.kind() == ErrorKind::TimedOut =>
{
Ok(None)
}
Err(e) => Err(e),
}
}
fn await_global_reply(&mut self, what: &'static str) -> Result<Vec<u8>> {
for _ in 0..MAX_EXEC_ITER {
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::GlobalSuccess { data } => return Ok(data),
ChannelEvent::GlobalFailure => {
let _ = what; return Err(Error::Protocol("global request denied"));
}
_ => {}
}
}
Err(Error::Protocol(
"global request: reply loop did not converge",
))
}
pub(crate) fn await_request_reply(&mut self, channel: u32, what: &'static str) -> Result<()> {
for _ in 0..MAX_EXEC_ITER {
let payload = self.read_one_packet()?;
match self.conn.on_packet(&payload)? {
ChannelEvent::Success { channel: c } if c == channel => return Ok(()),
ChannelEvent::Failure { channel: c } if c == channel => {
let _ = what; return Err(Error::Protocol("shell: channel request denied"));
}
_ => {}
}
}
Err(Error::Protocol(
"shell: request-reply loop did not converge",
))
}
pub(crate) fn pump_out(&mut self) -> Result<()> {
while let Some(frame) = self.driver.poll_transmit() {
self.out_buf.extend(frame);
}
self.flush_out_buf()
}
fn flush_out_buf(&mut self) -> Result<()> {
while !self.out_buf.is_empty() {
let front_len = self.out_buf.as_slices().0.len();
let n = match self.stream.write(self.out_buf.as_slices().0) {
Ok(0) => return Err(Error::Protocol("connection closed")),
Ok(n) => n,
Err(ref e) if e.kind() == ErrorKind::Interrupted => continue,
Err(ref e)
if e.kind() == ErrorKind::WouldBlock || e.kind() == ErrorKind::TimedOut =>
{
return Ok(());
}
Err(e) => return Err(Error::Io(e)),
};
debug_assert!(n <= front_len);
self.out_buf.drain(..n);
}
Ok(())
}
fn read_into_driver(&mut self) -> Result<()> {
let mut tmp = [0u8; 16 * 1024];
let n = self.stream.read(&mut tmp)?;
if n == 0 {
return Err(Error::Protocol("connection closed"));
}
self.driver.handle_input(&tmp[..n], Instant::now())?;
Ok(())
}
fn drive_handshake(&mut self) -> Result<()> {
for _ in 0..MAX_KEX_STEPS.saturating_mul(MAX_BANNER_LINES + 4) {
self.pump_out()?;
while let Some(ev) = self.driver.poll_event() {
if matches!(ev, Event::HandshakeComplete) {
self.pump_out()?;
return Ok(());
}
}
self.read_into_driver()?;
}
Err(Error::Protocol("kex: too many steps"))
}
pub(crate) fn read_one_packet(&mut self) -> Result<Vec<u8>> {
loop {
self.driver.handle_timeout(Instant::now())?;
self.pump_out()?;
while let Some(ev) = self.driver.poll_event() {
if let Event::AppData(payload) = ev {
self.pump_out()?;
return Ok(payload);
}
}
self.read_into_driver()?;
}
}
pub(crate) fn write_payload(&mut self, payload: &[u8]) -> Result<()> {
self.driver.enqueue_payload(payload)?;
self.pump_out()
}
pub(crate) fn send_transport_ping(&mut self, data: &[u8]) -> Result<()> {
let ping = encode_ping(data);
self.write_payload(&ping)
}
}
pub struct ClientChannelStream<'a> {
client: &'a mut Client,
channel: u32,
read_buf: Vec<u8>,
stderr_buf: Vec<u8>,
remote_eof: bool,
local_close_sent: bool,
}
impl ClientChannelStream<'_> {
pub fn take_stderr(&mut self) -> Vec<u8> {
core::mem::take(&mut self.stderr_buf)
}
fn pump_one(&mut self) -> std::io::Result<()> {
let payload = self.client.read_one_packet().map_err(io_err)?;
let ev = self.client.conn.on_packet(&payload).map_err(io_err)?;
match ev {
ChannelEvent::Data { channel, data } if channel == self.channel => {
let n = data.len() as u32;
self.read_buf.extend_from_slice(&data);
if let Some(adj) = self
.client
.conn
.replenish_window(self.channel, n)
.map_err(io_err)?
{
self.client.write_payload(&adj).map_err(io_err)?;
}
}
ChannelEvent::ExtendedData {
channel,
code: _,
data,
} if channel == self.channel => {
let n = data.len() as u32;
self.stderr_buf.extend_from_slice(&data);
if let Some(adj) = self
.client
.conn
.replenish_window(self.channel, n)
.map_err(io_err)?
{
self.client.write_payload(&adj).map_err(io_err)?;
}
}
ChannelEvent::Eof { channel } if channel == self.channel => {
self.remote_eof = true;
}
ChannelEvent::Close { channel } if channel == self.channel => {
self.remote_eof = true;
if !self.local_close_sent {
let p = self.client.conn.send_close(self.channel).map_err(io_err)?;
self.client.write_payload(&p).map_err(io_err)?;
self.local_close_sent = true;
}
}
_ => {}
}
Ok(())
}
}
impl Read for ClientChannelStream<'_> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
while self.read_buf.is_empty() && !self.remote_eof {
self.pump_one()?;
}
if self.read_buf.is_empty() {
return Ok(0);
}
let n = core::cmp::min(buf.len(), self.read_buf.len());
buf[..n].copy_from_slice(&self.read_buf[..n]);
self.read_buf.drain(..n);
Ok(n)
}
}
impl Write for ClientChannelStream<'_> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
loop {
let (payload, taken) = self
.client
.conn
.send_data(self.channel, buf)
.map_err(io_err)?;
if taken > 0 {
self.client.write_payload(&payload).map_err(io_err)?;
return Ok(taken);
}
if self.remote_eof {
return Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"channel closed by peer mid-write",
));
}
self.pump_one()?;
}
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl Drop for ClientChannelStream<'_> {
fn drop(&mut self) {
if !self.local_close_sent {
if let Ok(p) = self.client.conn.send_eof(self.channel) {
let _ = self.client.write_payload(&p);
}
if let Ok(p) = self.client.conn.send_close(self.channel) {
let _ = self.client.write_payload(&p);
}
self.local_close_sent = true;
}
const MAX_DRAIN: usize = 128;
for _ in 0..MAX_DRAIN {
if self.remote_eof {
break;
}
if self.pump_one().is_err() {
break;
}
}
}
}
pub(crate) fn io_err(e: Error) -> std::io::Error {
match e {
Error::Io(io) => io,
other => std::io::Error::other(format!("{:?}", other)),
}
}
#[cfg(all(feature = "std", unix))]
pub fn encode_termios_modes(t: &libc::termios) -> Vec<u8> {
const TTY_OP_END: u8 = 0;
const VINTR: u8 = 1;
const VQUIT: u8 = 2;
const VERASE: u8 = 3;
const VKILL: u8 = 4;
const VEOF: u8 = 5;
const VEOL: u8 = 6;
const VEOL2: u8 = 7;
const VSTART: u8 = 8;
const VSTOP: u8 = 9;
const VSUSP: u8 = 10;
const VREPRINT: u8 = 12;
const VWERASE: u8 = 13;
const VLNEXT: u8 = 14;
const IGNPAR: u8 = 30;
const PARMRK: u8 = 31;
const INPCK: u8 = 32;
const ISTRIP: u8 = 33;
const INLCR: u8 = 34;
const IGNCR: u8 = 35;
const ICRNL: u8 = 36;
const IXON: u8 = 39;
const IXANY: u8 = 40;
const IXOFF: u8 = 41;
const IMAXBEL: u8 = 42;
const ISIG: u8 = 50;
const ICANON: u8 = 51;
const ECHO: u8 = 53;
const ECHOE: u8 = 54;
const ECHOK: u8 = 55;
const ECHONL: u8 = 56;
const NOFLSH: u8 = 57;
const TOSTOP: u8 = 58;
const IEXTEN: u8 = 59;
const ECHOCTL: u8 = 60;
const ECHOKE: u8 = 61;
const OPOST: u8 = 70;
const ONLCR: u8 = 72;
const OCRNL: u8 = 73;
const ONOCR: u8 = 74;
const ONLRET: u8 = 75;
const CS7: u8 = 90;
const CS8: u8 = 91;
const PARENB: u8 = 92;
const PARODD: u8 = 93;
const TTY_OP_ISPEED: u8 = 128;
const TTY_OP_OSPEED: u8 = 129;
let mut out = Vec::with_capacity(128);
let mut push = |op: u8, val: u32| {
out.push(op);
out.extend_from_slice(&val.to_be_bytes());
};
let cc = |i: usize| -> u32 { t.c_cc[i] as u32 };
push(VINTR, cc(libc::VINTR));
push(VQUIT, cc(libc::VQUIT));
push(VERASE, cc(libc::VERASE));
push(VKILL, cc(libc::VKILL));
push(VEOF, cc(libc::VEOF));
push(VEOL, cc(libc::VEOL));
push(VEOL2, cc(libc::VEOL2));
push(VSTART, cc(libc::VSTART));
push(VSTOP, cc(libc::VSTOP));
push(VSUSP, cc(libc::VSUSP));
push(VREPRINT, cc(libc::VREPRINT));
push(VWERASE, cc(libc::VWERASE));
push(VLNEXT, cc(libc::VLNEXT));
let iflag = t.c_iflag;
let bit_i = |mask: libc::tcflag_t| -> u32 { if iflag & mask != 0 { 1 } else { 0 } };
push(IGNPAR, bit_i(libc::IGNPAR));
push(PARMRK, bit_i(libc::PARMRK));
push(INPCK, bit_i(libc::INPCK));
push(ISTRIP, bit_i(libc::ISTRIP));
push(INLCR, bit_i(libc::INLCR));
push(IGNCR, bit_i(libc::IGNCR));
push(ICRNL, bit_i(libc::ICRNL));
push(IXON, bit_i(libc::IXON));
push(IXANY, bit_i(libc::IXANY));
push(IXOFF, bit_i(libc::IXOFF));
push(IMAXBEL, bit_i(libc::IMAXBEL));
let lflag = t.c_lflag;
let bit_l = |mask: libc::tcflag_t| -> u32 { if lflag & mask != 0 { 1 } else { 0 } };
push(ISIG, bit_l(libc::ISIG));
push(ICANON, bit_l(libc::ICANON));
push(ECHO, bit_l(libc::ECHO));
push(ECHOE, bit_l(libc::ECHOE));
push(ECHOK, bit_l(libc::ECHOK));
push(ECHONL, bit_l(libc::ECHONL));
push(NOFLSH, bit_l(libc::NOFLSH));
push(TOSTOP, bit_l(libc::TOSTOP));
push(IEXTEN, bit_l(libc::IEXTEN));
push(ECHOCTL, bit_l(libc::ECHOCTL));
push(ECHOKE, bit_l(libc::ECHOKE));
let oflag = t.c_oflag;
let bit_o = |mask: libc::tcflag_t| -> u32 { if oflag & mask != 0 { 1 } else { 0 } };
push(OPOST, bit_o(libc::OPOST));
push(ONLCR, bit_o(libc::ONLCR));
push(OCRNL, bit_o(libc::OCRNL));
push(ONOCR, bit_o(libc::ONOCR));
push(ONLRET, bit_o(libc::ONLRET));
let cflag = t.c_cflag;
let cs = cflag & libc::CSIZE;
push(CS7, if cs == libc::CS7 { 1 } else { 0 });
push(CS8, if cs == libc::CS8 { 1 } else { 0 });
push(PARENB, if cflag & libc::PARENB != 0 { 1 } else { 0 });
push(PARODD, if cflag & libc::PARODD != 0 { 1 } else { 0 });
let _ = t; push(TTY_OP_ISPEED, 38_400);
push(TTY_OP_OSPEED, 38_400);
out.push(TTY_OP_END);
out
}
fn build_scp_to_cmd(remote_dest: &str, opts: &crate::scp::ScpSendOptions) -> Result<String> {
let quoted = single_quote_for_remote(remote_dest)?;
let mut s = String::from("scp -t");
if opts.recursive {
s.push_str(" -r");
}
if opts.preserve_times {
s.push_str(" -p");
}
s.push_str(" -- ");
s.push_str("ed);
Ok(s)
}
fn build_scp_from_cmd(remote_source: &str, opts: &crate::scp::ScpRecvOptions) -> Result<String> {
let quoted = single_quote_for_remote(remote_source)?;
let mut s = String::from("scp -f");
if opts.recursive {
s.push_str(" -r");
}
if opts.preserve_times {
s.push_str(" -p");
}
s.push_str(" -- ");
s.push_str("ed);
Ok(s)
}
fn single_quote_for_remote(p: &str) -> Result<String> {
if p.contains('\'') {
return Err(Error::Protocol("scp: remote path contains single quote"));
}
if p.contains('\n') {
return Err(Error::Protocol("scp: remote path contains newline"));
}
if p.contains('\0') {
return Err(Error::Protocol("scp: remote path contains NUL"));
}
if p.starts_with('-') {
return Err(Error::Protocol("scp: remote path starts with '-'"));
}
let mut q = String::with_capacity(p.len() + 2);
q.push('\'');
q.push_str(p);
q.push('\'');
Ok(q)
}
fn scp_proto(e: crate::scp::ScpError, _stage: &'static str) -> Error {
match e {
crate::scp::ScpError::Io(io) => Error::Io(io),
crate::scp::ScpError::Remote(_) => Error::Protocol("scp: remote fatal frame"),
crate::scp::ScpError::Warning(_) => Error::Protocol("scp: remote warning frame"),
crate::scp::ScpError::BadHeader(_) => Error::Protocol("scp: malformed header"),
crate::scp::ScpError::BadName(_) => Error::Protocol("scp: invalid name"),
crate::scp::ScpError::PathEscape => Error::Protocol("scp: path escapes base"),
crate::scp::ScpError::Unexpected(_) => Error::Protocol("scp: unexpected protocol state"),
}
}
pub fn host_key_fingerprint(blob: &[u8]) -> String {
const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
let d = Sha256::digest(blob);
let bytes: &[u8] = d.as_ref();
let mut out = String::with_capacity(bytes.len().div_ceil(3) * 4 + 7);
out.push_str("SHA256:");
let mut i = 0;
while i + 3 <= bytes.len() {
let b = ((bytes[i] as u32) << 16) | ((bytes[i + 1] as u32) << 8) | (bytes[i + 2] as u32);
out.push(ALPHABET[((b >> 18) & 0x3F) as usize] as char);
out.push(ALPHABET[((b >> 12) & 0x3F) as usize] as char);
out.push(ALPHABET[((b >> 6) & 0x3F) as usize] as char);
out.push(ALPHABET[(b & 0x3F) as usize] as char);
i += 3;
}
let rem = bytes.len() - i;
if rem == 1 {
let b = (bytes[i] as u32) << 16;
out.push(ALPHABET[((b >> 18) & 0x3F) as usize] as char);
out.push(ALPHABET[((b >> 12) & 0x3F) as usize] as char);
} else if rem == 2 {
let b = ((bytes[i] as u32) << 16) | ((bytes[i + 1] as u32) << 8);
out.push(ALPHABET[((b >> 18) & 0x3F) as usize] as char);
out.push(ALPHABET[((b >> 12) & 0x3F) as usize] as char);
out.push(ALPHABET[((b >> 6) & 0x3F) as usize] as char);
}
out
}
fn print_mismatch_banner(
host: &str,
port: u16,
expected: &[(String, Vec<u8>)],
new_key_type: &str,
new_key_blob: &[u8],
) {
let target = if port == 22 {
host.to_string()
} else {
format!("[{host}]:{port}")
};
eprintln!("@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@");
eprintln!("@ WARNING: REMOTE HOST IDENTIFICATION HAS CHANGED! @");
eprintln!("@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@");
eprintln!("IT IS POSSIBLE THAT SOMEONE IS DOING SOMETHING NASTY!");
eprintln!("Someone could be eavesdropping on you right now (man-in-the-middle attack)!");
eprintln!("It is also possible that a host key has just been changed.");
eprintln!(
"The host key for {target} has changed; the known_hosts entry does not match \
what the server presented."
);
if expected.is_empty() {
eprintln!("Old fingerprint: <none on file>");
} else {
for (kt, blob) in expected {
eprintln!("Old fingerprint: {} ({kt})", host_key_fingerprint(blob));
}
}
eprintln!(
"New fingerprint: {} ({new_key_type})",
host_key_fingerprint(new_key_blob),
);
}
fn owned_or_default(over: &Option<Vec<String>>, default: &[&str]) -> Vec<String> {
match over {
Some(v) => v.clone(),
None => default.iter().map(|s| s.to_string()).collect(),
}
}
pub(crate) fn build_default_kexinit<R: RngCore>(rng: &mut R, over: &AlgoOverrides) -> KexInit {
let default_kex: Vec<&str> = defaults::KEX
.iter()
.copied()
.filter(|n| !is_strict_kex_marker(n))
.collect();
let mut kex = owned_or_default(&over.kex_algorithms, &default_kex);
for marker in defaults::KEX.iter().filter(|n| is_strict_kex_marker(n)) {
if !kex.iter().any(|k| k == marker) {
kex.push((*marker).to_string());
}
}
let ciphers = owned_or_default(&over.ciphers, defaults::CIPHERS);
let macs = owned_or_default(&over.macs, defaults::MACS);
let host_key = match &over.host_key_algorithms {
Some(v) => v.clone(),
None => {
let mut v: Vec<String> = crate::cert::CERT_KEY_NAMES
.iter()
.map(|s| s.to_string())
.collect();
v.extend(defaults::HOST_KEY.iter().map(|s| s.to_string()));
v
}
};
let comp: Vec<String> = if cfg!(feature = "compress") && over.compression == Some(true) {
vec!["zlib@openssh.com".to_string(), "none".to_string()]
} else {
defaults::COMP.iter().map(|s| s.to_string()).collect()
};
let algs = KexAlgorithmsOwned {
kex,
server_host_key: host_key,
ciphers_c2s: ciphers.clone(),
ciphers_s2c: ciphers,
macs_c2s: macs.clone(),
macs_s2c: macs,
comp_c2s: comp.clone(),
comp_s2c: comp,
lang_c2s: Vec::new(),
lang_s2c: Vec::new(),
};
let mut cookie = [0u8; 16];
rng.fill_bytes(&mut cookie);
KexInit::from_algorithms_owned(algs, cookie)
}
pub(crate) fn unix_now() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
pub(crate) fn build_verifier(
reply_payload: &[u8],
policy: &HostKeyPolicy,
runner: &KexRunner,
target_host: &str,
target_port: u16,
ca_signature_algorithms: Option<&[String]>,
now: u64,
) -> Result<Box<dyn HostKeyVerify>> {
if reply_payload.len() < 5 {
return Err(Error::Format("kex-ecdh-reply too short"));
}
let k_s_len = u32::from_be_bytes([
reply_payload[1],
reply_payload[2],
reply_payload[3],
reply_payload[4],
]) as usize;
if reply_payload.len() < 5 + k_s_len {
return Err(Error::Format("kex-ecdh-reply truncated"));
}
let k_s = &reply_payload[5..5 + k_s_len];
if matches!(policy, HostKeyPolicy::KnownHosts(_))
&& (target_host.is_empty() || target_port == 0)
{
return Err(Error::Config(
"HostKeyPolicy::KnownHosts requires Client::connect_to_host",
));
}
let neg = runner
.negotiated()
.ok_or(Error::Protocol("kex: no negotiated algorithms"))?;
if crate::cert::is_cert_name(&neg.host_key) {
let cert = crate::cert::Certificate::parse(k_s)?;
let ca_algos: Vec<&str> = match ca_signature_algorithms {
Some(list) => list.iter().map(|s| s.as_str()).collect(),
None => crate::config::algos::CA_SIGNATURE_DEFAULTS.to_vec(),
};
match policy {
HostKeyPolicy::AcceptAny => {
cert.check_type(crate::cert::CertType::Host)?;
cert.check_validity(now)?;
cert.require_known_critical_options()?;
}
HostKeyPolicy::AcceptFingerprint(fp) => {
let digest = Sha256::digest(k_s);
if digest.as_ref() != fp {
return Err(Error::HostKeyRejected);
}
cert.check_type(crate::cert::CertType::Host)?;
cert.check_validity(now)?;
cert.require_known_critical_options()?;
}
HostKeyPolicy::KnownHosts(kh) => {
let store = kh.store.lock().map_err(|_| Error::HostKeyRejected)?;
store.verify_host_cert(target_host, target_port, &cert, &ca_algos, now)?;
}
}
return host_key_verify_by_name(&neg.host_key, k_s);
}
match policy {
HostKeyPolicy::AcceptAny => {}
HostKeyPolicy::AcceptFingerprint(fp) => {
let digest = Sha256::digest(k_s);
if digest.as_ref() != fp {
return Err(Error::HostKeyRejected);
}
}
HostKeyPolicy::KnownHosts(kh) => {
let mut store = kh.store.lock().map_err(|_| Error::HostKeyRejected)?;
let lookup = store.lookup(target_host, target_port, &neg.host_key, k_s);
match lookup {
LookupResult::Match => {}
LookupResult::Mismatch { expected } => {
if !matches!(&kh.on_mismatch, TofuAction::PromptDetailed(_)) {
print_mismatch_banner(
target_host,
target_port,
&expected,
&neg.host_key,
k_s,
);
}
let accept = match &kh.on_mismatch {
TofuAction::Reject => false,
TofuAction::Accept => true,
TofuAction::AcceptWithWarning => {
eprintln!(
"Connecting anyway because StrictHostKeyChecking is set to no; \
the trusted entry in known_hosts is NOT being updated."
);
true
}
TofuAction::Prompt(cb) => {
drop(store);
let ok = cb(target_host, target_port, &neg.host_key, k_s);
store = kh.store.lock().map_err(|_| Error::HostKeyRejected)?;
ok
}
TofuAction::PromptDetailed(cb) => {
drop(store);
let prompt = HostKeyPrompt {
change: HostKeyChange::Changed,
host: target_host,
port: target_port,
key_type: &neg.host_key,
key_blob: k_s,
fingerprint: host_key_fingerprint(k_s),
existing: &expected,
};
let ok = cb(&prompt);
store = kh.store.lock().map_err(|_| Error::HostKeyRejected)?;
ok
}
};
if !accept {
return Err(Error::HostKeyRejected);
}
if matches!(
&kh.on_mismatch,
TofuAction::Accept | TofuAction::Prompt(_) | TofuAction::PromptDetailed(_)
) {
let _ = store.remove(target_host, target_port);
store.add(target_host, target_port, &neg.host_key, k_s, kh.hash_new);
if let Some(path) = &kh.save_path {
store.save(path).map_err(Error::from)?;
}
}
}
LookupResult::Unknown => {
let accept = match &kh.on_unknown {
TofuAction::Reject => false,
TofuAction::Accept | TofuAction::AcceptWithWarning => true,
TofuAction::Prompt(cb) => {
drop(store);
let ok = cb(target_host, target_port, &neg.host_key, k_s);
store = kh.store.lock().map_err(|_| Error::HostKeyRejected)?;
ok
}
TofuAction::PromptDetailed(cb) => {
drop(store);
let prompt = HostKeyPrompt {
change: HostKeyChange::Unknown,
host: target_host,
port: target_port,
key_type: &neg.host_key,
key_blob: k_s,
fingerprint: host_key_fingerprint(k_s),
existing: &[],
};
let ok = cb(&prompt);
store = kh.store.lock().map_err(|_| Error::HostKeyRejected)?;
ok
}
};
if !accept {
return Err(Error::HostKeyRejected);
}
store.add(target_host, target_port, &neg.host_key, k_s, kh.hash_new);
if let Some(path) = &kh.save_path {
store.save(path).map_err(Error::from)?;
}
}
}
}
}
host_key_verify_by_name(&neg.host_key, k_s)
}
fn serve_drain_commands(
client: &mut Client,
cmd_rx: &Receiver<ServeCommand>,
pending_opens: &mut BTreeMap<u32, PendingOutboundOpen>,
) -> Result<()> {
loop {
match cmd_rx.try_recv() {
Ok(ServeCommand::OpenDirectTcpip {
dest_host,
dest_port,
orig_host,
orig_port,
reply,
}) => {
let (local_id, open_payload) = client.conn.open(ChannelOpen::DirectTcpip {
dest_host,
dest_port: dest_port as u32,
orig_host,
orig_port: orig_port as u32,
})?;
client.write_payload(&open_payload)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SERVE_EGRESS_BACKLOG);
let stream = ChannelStream::new(ingress_rx, egress_tx);
pending_opens.insert(
local_id,
PendingOutboundOpen {
stream: Some(stream),
ingress_tx,
egress_rx: Some(egress_rx),
reply,
},
);
}
Ok(ServeCommand::OpenDirectStreamlocal { socket_path, reply }) => {
let (local_id, open_payload) = client
.conn
.open(ChannelOpen::DirectStreamlocal { socket_path })?;
client.write_payload(&open_payload)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SERVE_EGRESS_BACKLOG);
let stream = ChannelStream::new(ingress_rx, egress_tx);
pending_opens.insert(
local_id,
PendingOutboundOpen {
stream: Some(stream),
ingress_tx,
egress_rx: Some(egress_rx),
reply,
},
);
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => break,
}
}
Ok(())
}
fn serve_drain_runtimes(
client: &mut Client,
runtimes: &mut BTreeMap<u32, ServeRuntime>,
) -> Result<()> {
let channels: Vec<u32> = runtimes.keys().copied().collect();
for ch in channels {
let Some(rt) = runtimes.get_mut(&ch) else {
continue;
};
if rt.close_sent {
continue;
}
if !rt.pending_data.is_empty() {
let leftover = core::mem::take(&mut rt.pending_data);
emit_serve_data(client, ch, &leftover, rt)?;
if !rt.pending_data.is_empty() {
continue;
}
}
loop {
if !rt.pending_data.is_empty() {
break;
}
match rt.egress_rx.try_recv() {
Ok(ChannelEgress::Data(bytes)) => {
emit_serve_data(client, ch, &bytes, rt)?;
}
Ok(ChannelEgress::Eof) => {
rt.pending_eof = true;
break;
}
Ok(ChannelEgress::Close) => {
rt.pending_close = true;
break;
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => {
rt.pending_close = true;
break;
}
}
}
if rt.pending_data.is_empty() {
if rt.pending_eof && !rt.eof_sent {
let p = client.conn.send_eof(ch)?;
client.write_payload(&p)?;
rt.eof_sent = true;
}
if rt.pending_close && !rt.close_sent {
if !rt.eof_sent {
let p = client.conn.send_eof(ch)?;
client.write_payload(&p)?;
rt.eof_sent = true;
}
let p = client.conn.send_close(ch)?;
client.write_payload(&p)?;
rt.close_sent = true;
}
}
}
Ok(())
}
fn emit_serve_data(
client: &mut Client,
channel: u32,
bytes: &[u8],
rt: &mut ServeRuntime,
) -> Result<()> {
let mut off = 0usize;
while off < bytes.len() {
let (payload, taken) = client.conn.send_data(channel, &bytes[off..])?;
if taken == 0 {
rt.pending_data.extend_from_slice(&bytes[off..]);
return Ok(());
}
client.write_payload(&payload)?;
off += taken;
}
Ok(())
}
fn serve_dispatch_packet(
client: &mut Client,
handlers: &ClientHandlers,
runtimes: &mut BTreeMap<u32, ServeRuntime>,
pending_opens: &mut BTreeMap<u32, PendingOutboundOpen>,
payload: &[u8],
) -> Result<()> {
let ev = client.conn.on_packet(payload)?;
match ev {
ChannelEvent::OpenConfirmed { channel } => {
if let Some(mut po) = pending_opens.remove(&channel) {
let stream = po
.stream
.take()
.ok_or(Error::Protocol("pending open: stream taken twice"))?;
let egress_rx = po
.egress_rx
.take()
.ok_or(Error::Protocol("pending open: egress taken twice"))?;
if po.reply.send(Ok(stream)).is_err() {
let p = client.conn.send_close(channel)?;
client.write_payload(&p)?;
return Ok(());
}
runtimes.insert(
channel,
ServeRuntime {
ingress_tx: po.ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
}
}
ChannelEvent::OpenFailed { channel, .. } => {
if let Some(po) = pending_opens.remove(&channel) {
let _ = po
.reply
.send(Err(Error::Protocol("direct-tcpip: open failed")));
}
}
ChannelEvent::OpenRequest { channel, kind } => match kind {
ChannelOpen::ForwardedTcpip {
dest_host,
dest_port,
orig_host,
orig_port,
} => {
if client.tcpip_forward_grants.is_empty() {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"no tcpip-forward requested",
"",
)?;
client.write_payload(&p)?;
return Ok(());
}
if let Some(cb) = handlers.on_forwarded_tcpip.clone() {
let p = client.conn.accept_open(channel)?;
client.write_payload(&p)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SERVE_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let origin = ForwardedTcpipOrigin {
bound_address: dest_host,
bound_port: clamp_u16(dest_port),
orig_address: orig_host,
orig_port: clamp_u16(orig_port),
};
thread::spawn(move || {
cb(origin, cs);
});
runtimes.insert(
channel,
ServeRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
} else {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"forwarded-tcpip not enabled",
"",
)?;
client.write_payload(&p)?;
}
}
ChannelOpen::ForwardedStreamlocal { socket_path } => {
if client.streamlocal_forward_grants.is_empty() {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"no streamlocal-forward requested",
"",
)?;
client.write_payload(&p)?;
return Ok(());
}
if let Some(cb) = handlers.on_forwarded_streamlocal.clone() {
let p = client.conn.accept_open(channel)?;
client.write_payload(&p)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SERVE_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let origin = ForwardedStreamlocalOrigin { socket_path };
thread::spawn(move || {
cb(origin, cs);
});
runtimes.insert(
channel,
ServeRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
} else {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"forwarded-streamlocal not enabled",
"",
)?;
client.write_payload(&p)?;
}
}
ChannelOpen::AuthAgent => {
if let Some(cb) = handlers.on_auth_agent.clone() {
let p = client.conn.accept_open(channel)?;
client.write_payload(&p)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SERVE_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
thread::spawn(move || {
cb(cs);
});
runtimes.insert(
channel,
ServeRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
} else {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"auth-agent not enabled",
"",
)?;
client.write_payload(&p)?;
}
}
ChannelOpen::X11 {
orig_host: _,
orig_port: _,
} => {
if let Some(cb) = handlers.on_x11.clone() {
let p = client.conn.accept_open(channel)?;
client.write_payload(&p)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SERVE_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
thread::spawn(move || {
cb(cs);
});
runtimes.insert(
channel,
ServeRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
} else {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"x11 not enabled",
"",
)?;
client.write_payload(&p)?;
}
}
_ => {
let p = client.conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"channel type not supported",
"",
)?;
client.write_payload(&p)?;
}
},
ChannelEvent::Data { channel, data } => {
if let Some(rt) = runtimes.get_mut(&channel) {
let _ = rt.ingress_tx.send(Some(data.clone()));
}
if let Some(adj) = client.conn.replenish_window(channel, data.len() as u32)? {
client.write_payload(&adj)?;
}
}
ChannelEvent::ExtendedData { channel, data, .. } => {
if let Some(adj) = client.conn.replenish_window(channel, data.len() as u32)? {
client.write_payload(&adj)?;
}
}
ChannelEvent::Eof { channel } => {
if let Some(rt) = runtimes.get_mut(&channel) {
let _ = rt.ingress_tx.send(None);
}
}
ChannelEvent::Close { channel } => {
if let Some(ch) = client.conn.channel(channel)
&& !ch.local_closed
{
let p = client.conn.send_close(channel)?;
client.write_payload(&p)?;
}
runtimes.remove(&channel);
}
ChannelEvent::GlobalRequest {
want_reply: true, ..
} => {
let p = client.conn.send_global_failure();
client.write_payload(&p)?;
}
_ => {}
}
Ok(())
}
fn clamp_u16(v: u32) -> u16 {
if v > u16::MAX as u32 {
u16::MAX
} else {
v as u16
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hostkey::Ed25519HostKey;
use crate::transport::version::LOCAL_VERSION;
use crate::transport::{PacketCodec, Role, VersionExchange};
use purecrypto::rng::OsRng;
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
const SSH_MSG_KEX_ECDH_REPLY: u8 = 31;
const MAX_INBOX_BYTES: usize = 8 * 1024 * 1024;
fn read_line<S: Read>(stream: &mut S, buf: &mut Vec<u8>, max_len: usize) -> Result<()> {
let mut byte = [0u8; 1];
loop {
let n = stream.read(&mut byte)?;
if n == 0 {
return Err(Error::Protocol("connection closed before newline"));
}
buf.push(byte[0]);
if byte[0] == b'\n' {
return Ok(());
}
if buf.len() >= max_len {
return Err(Error::Protocol("banner line too long"));
}
}
}
#[test]
fn config_insecure_constructor_is_accept_any() {
let cfg = Config::insecure();
assert!(matches!(cfg.host_key_policy, HostKeyPolicy::AcceptAny));
assert!(cfg.timeout.is_none());
}
#[test]
fn known_hosts_strict_constructor_defaults_reject_reject() {
let store = Arc::new(Mutex::new(KnownHosts::new()));
let p = KnownHostsPolicy::strict(store);
assert!(matches!(p.on_unknown, TofuAction::Reject));
assert!(matches!(p.on_mismatch, TofuAction::Reject));
assert!(!p.hash_new);
assert!(p.save_path.is_none());
}
#[test]
fn build_verifier_fails_hard_on_empty_host() {
use crate::transport::kex::KexAlgorithms;
use crate::transport::{KexInit, KexRunner};
let store = Arc::new(Mutex::new(KnownHosts::new()));
let policy = HostKeyPolicy::KnownHosts(KnownHostsPolicy::strict(store));
let runner = KexRunner::new(
Role::Client,
KexInit::from_algorithms(
&KexAlgorithms {
kex: defaults::KEX,
server_host_key: defaults::HOST_KEY,
ciphers_c2s: defaults::CIPHERS,
ciphers_s2c: defaults::CIPHERS,
macs_c2s: defaults::MACS,
macs_s2c: defaults::MACS,
comp_c2s: defaults::COMP,
comp_s2c: defaults::COMP,
lang_c2s: &[],
lang_s2c: &[],
},
[0u8; 16],
),
);
let mut reply = vec![SSH_MSG_KEX_ECDH_REPLY];
reply.extend_from_slice(&0u32.to_be_bytes());
let err = build_verifier(&reply, &policy, &runner, "", 22, None, 0);
assert!(matches!(err, Err(Error::Config(_))));
let err = build_verifier(&reply, &policy, &runner, "host", 0, None, 0);
assert!(matches!(err, Err(Error::Config(_))));
}
#[test]
fn exec_output_constructible() {
let _ = ExecOutput {
stdout: Vec::new(),
stderr: Vec::new(),
exit_status: Some(0),
exit_signal: None,
};
}
fn run_server(
listener: TcpListener,
host_key_seed: [u8; 32],
) -> thread::JoinHandle<std::result::Result<Vec<u8>, String>> {
thread::spawn(move || -> std::result::Result<Vec<u8>, String> {
let (mut s, _) = listener.accept().map_err(|e| e.to_string())?;
let server_hk = Ed25519HostKey::from_seed(host_key_seed);
s.write_all(&VersionExchange::outgoing_bytes())
.map_err(|e| e.to_string())?;
let mut line = Vec::new();
let v_c: Vec<u8> = {
read_line(&mut s, &mut line, 1024).map_err(|e| format!("{e:?}"))?;
if !line.starts_with(b"SSH-") {
return Err("client did not send SSH banner".into());
}
let parsed = VersionExchange::parse_remote(&line).map_err(|e| format!("{e:?}"))?;
parsed.into_bytes()
};
let v_s = LOCAL_VERSION.as_bytes().to_vec();
let mut codec = PacketCodec::new();
let server_over = AlgoOverrides {
host_key_algorithms: Some(vec!["ssh-ed25519".to_string()]),
..Default::default()
};
let advert = build_default_kexinit(&mut OsRng, &server_over);
let mut runner = KexRunner::new(Role::Server, advert);
let mut inbox: Vec<u8> = Vec::new();
let mut rng = OsRng;
let initial = runner.start(&mut rng).map_err(|e| format!("{e:?}"))?;
for p in initial.outbound {
let frame = codec.encode(&p, &mut rng).map_err(|e| format!("{e:?}"))?;
s.write_all(&frame).map_err(|e| e.to_string())?;
}
let mut steps = 0;
loop {
steps += 1;
if steps > MAX_KEX_STEPS {
return Err("server kex did not converge".into());
}
let payload = read_one_packet_local(&mut s, &mut codec, &mut inbox)
.map_err(|e| format!("{e:?}"))?;
let adv = runner
.on_packet(
&mut rng,
&mut codec,
&payload,
Some(&server_hk),
None,
&v_c,
&v_s,
)
.map_err(|e| format!("{e:?}"))?;
for p in adv.outbound {
let frame = codec.encode(&p, &mut rng).map_err(|e| format!("{e:?}"))?;
s.write_all(&frame).map_err(|e| e.to_string())?;
}
if adv.completed {
break;
}
}
let sid = runner.session_id().unwrap().to_vec();
Ok(sid)
})
}
fn read_one_packet_local(
s: &mut TcpStream,
codec: &mut PacketCodec,
inbox: &mut Vec<u8>,
) -> Result<Vec<u8>> {
loop {
if let Some((payload, consumed)) = codec.decode(inbox)? {
inbox.drain(..consumed);
return Ok(payload);
}
let mut tmp = [0u8; 4096];
let n = s.read(&mut tmp)?;
if n == 0 {
return Err(Error::Protocol("connection closed"));
}
inbox.extend_from_slice(&tmp[..n]);
if inbox.len() > MAX_INBOX_BYTES {
return Err(Error::Protocol("inbound buffer too large"));
}
}
}
#[test]
fn handshake_over_real_loopback_socket() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut seed = [0u8; 32];
OsRng.fill_bytes(&mut seed);
let server = run_server(listener, seed);
let client = Client::connect(addr, Config::insecure()).expect("client connect");
let server_sid = server.join().unwrap().expect("server handshake");
assert_eq!(client.session_id(), server_sid.as_slice());
assert!(!client.session_id().is_empty());
}
#[test]
fn client_answers_ping_with_pong() {
use crate::transport::ping::{SSH_MSG_PONG, encode_ping, encode_pong};
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut seed = [0u8; 32];
OsRng.fill_bytes(&mut seed);
let server = thread::spawn(move || -> std::result::Result<Vec<u8>, String> {
let (mut s, _) = listener.accept().map_err(|e| e.to_string())?;
let server_hk = Ed25519HostKey::from_seed(seed);
s.write_all(&VersionExchange::outgoing_bytes())
.map_err(|e| e.to_string())?;
let mut line = Vec::new();
let v_c: Vec<u8> = {
read_line(&mut s, &mut line, 1024).map_err(|e| format!("{e:?}"))?;
let parsed = VersionExchange::parse_remote(&line).map_err(|e| format!("{e:?}"))?;
parsed.into_bytes()
};
let v_s = LOCAL_VERSION.as_bytes().to_vec();
let mut codec = PacketCodec::new();
let server_over = AlgoOverrides {
host_key_algorithms: Some(vec!["ssh-ed25519".to_string()]),
..Default::default()
};
let advert = build_default_kexinit(&mut OsRng, &server_over);
let mut runner = KexRunner::new(Role::Server, advert);
let mut inbox: Vec<u8> = Vec::new();
let mut rng = OsRng;
let initial = runner.start(&mut rng).map_err(|e| format!("{e:?}"))?;
for p in initial.outbound {
let frame = codec.encode(&p, &mut rng).map_err(|e| format!("{e:?}"))?;
s.write_all(&frame).map_err(|e| e.to_string())?;
}
let mut steps = 0;
loop {
steps += 1;
if steps > MAX_KEX_STEPS {
return Err("server kex did not converge".into());
}
let payload = read_one_packet_local(&mut s, &mut codec, &mut inbox)
.map_err(|e| format!("{e:?}"))?;
let adv = runner
.on_packet(
&mut rng,
&mut codec,
&payload,
Some(&server_hk),
None,
&v_c,
&v_s,
)
.map_err(|e| format!("{e:?}"))?;
for p in adv.outbound {
let frame = codec.encode(&p, &mut rng).map_err(|e| format!("{e:?}"))?;
s.write_all(&frame).map_err(|e| e.to_string())?;
}
if adv.completed {
break;
}
}
let stray = encode_pong(b"unsolicited");
let frame = codec
.encode(&stray, &mut rng)
.map_err(|e| format!("{e:?}"))?;
s.write_all(&frame).map_err(|e| e.to_string())?;
let ping = encode_ping(b"chaff-1234");
let frame = codec
.encode(&ping, &mut rng)
.map_err(|e| format!("{e:?}"))?;
s.write_all(&frame).map_err(|e| e.to_string())?;
let reply = read_one_packet_local(&mut s, &mut codec, &mut inbox)
.map_err(|e| format!("{e:?}"))?;
if reply.first().copied() != Some(SSH_MSG_PONG) {
return Err(format!("expected PONG, got msg {:?}", reply.first()));
}
let mut r = crate::format::Reader::new(&reply);
r.read_u8().map_err(|e| format!("{e:?}"))?;
let data = r.read_string().map_err(|e| format!("{e:?}"))?;
if data != b"chaff-1234" {
return Err(format!("PONG echoed wrong data: {data:?}"));
}
Ok(b"ok".to_vec())
});
let mut client = Client::connect(addr, Config::insecure()).expect("client connect");
let _ = client.read_one_packet();
let res = server.join().unwrap();
assert_eq!(res.expect("server PONG assertions"), b"ok");
}
#[test]
fn fingerprint_mismatch_rejected() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut seed = [0u8; 32];
OsRng.fill_bytes(&mut seed);
let server = run_server(listener, seed);
let cfg = Config {
host_key_policy: HostKeyPolicy::AcceptFingerprint([0xffu8; 32]),
timeout: None,
algorithms: Default::default(),
};
let err = Client::connect(addr, cfg).err().expect("must fail");
assert!(matches!(err, Error::HostKeyRejected));
let _ = server.join();
}
#[test]
fn known_hosts_promptdetailed_unknown_adds_entry() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut seed = [0u8; 32];
OsRng.fill_bytes(&mut seed);
let server = run_server(listener, seed);
type SeenUnknown = Arc<Mutex<Option<(HostKeyChange, String, usize)>>>;
let store = Arc::new(Mutex::new(KnownHosts::new()));
let seen: SeenUnknown = Arc::new(Mutex::new(None));
let seen_cb = seen.clone();
let cb: Arc<HostKeyPromptFn> = Arc::new(move |p: &HostKeyPrompt<'_>| {
*seen_cb.lock().unwrap() = Some((p.change, p.fingerprint.clone(), p.existing.len()));
true
});
let policy = KnownHostsPolicy {
store: store.clone(),
save_path: None,
hash_new: false,
on_unknown: TofuAction::PromptDetailed(cb),
on_mismatch: TofuAction::Reject,
};
let cfg = Config {
host_key_policy: HostKeyPolicy::KnownHosts(policy),
timeout: None,
algorithms: Default::default(),
};
let _client =
Client::connect_to_host("127.0.0.1", addr.port(), cfg).expect("connect should succeed");
let _ = server.join();
let (change, fp, existing_len) = seen.lock().unwrap().take().expect("callback fired");
assert_eq!(change, HostKeyChange::Unknown);
assert!(fp.starts_with("SHA256:"), "fingerprint was {fp}");
assert_eq!(existing_len, 0);
let real_blob = Ed25519HostKey::from_seed(seed).public_blob();
assert!(matches!(
store
.lock()
.unwrap()
.lookup("127.0.0.1", addr.port(), "ssh-ed25519", &real_blob),
LookupResult::Match
));
}
#[test]
fn known_hosts_promptdetailed_unknown_reject() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut seed = [0u8; 32];
OsRng.fill_bytes(&mut seed);
let server = run_server(listener, seed);
let store = Arc::new(Mutex::new(KnownHosts::new()));
let cb: Arc<HostKeyPromptFn> = Arc::new(|_p: &HostKeyPrompt<'_>| false);
let policy = KnownHostsPolicy {
store: store.clone(),
save_path: None,
hash_new: false,
on_unknown: TofuAction::PromptDetailed(cb),
on_mismatch: TofuAction::Reject,
};
let cfg = Config {
host_key_policy: HostKeyPolicy::KnownHosts(policy),
timeout: None,
algorithms: Default::default(),
};
let err = Client::connect_to_host("127.0.0.1", addr.port(), cfg)
.err()
.expect("connect must fail");
assert!(matches!(err, Error::HostKeyRejected));
let _ = server.join();
let real_blob = Ed25519HostKey::from_seed(seed).public_blob();
assert!(matches!(
store
.lock()
.unwrap()
.lookup("127.0.0.1", addr.port(), "ssh-ed25519", &real_blob),
LookupResult::Unknown
));
}
#[test]
fn known_hosts_promptdetailed_mismatch_rotates() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut seed = [0u8; 32];
OsRng.fill_bytes(&mut seed);
let server = run_server(listener, seed);
let bogus_blob = Ed25519HostKey::from_seed([0x11u8; 32]).public_blob();
let store = Arc::new(Mutex::new(KnownHosts::new()));
store
.lock()
.unwrap()
.add("127.0.0.1", addr.port(), "ssh-ed25519", &bogus_blob, false);
type SeenChanged = Arc<Mutex<Option<(HostKeyChange, Vec<String>)>>>;
let seen: SeenChanged = Arc::new(Mutex::new(None));
let seen_cb = seen.clone();
let cb: Arc<HostKeyPromptFn> = Arc::new(move |p: &HostKeyPrompt<'_>| {
*seen_cb.lock().unwrap() = Some((p.change, p.existing_fingerprints()));
true
});
let policy = KnownHostsPolicy {
store: store.clone(),
save_path: None,
hash_new: false,
on_unknown: TofuAction::Reject,
on_mismatch: TofuAction::PromptDetailed(cb),
};
let cfg = Config {
host_key_policy: HostKeyPolicy::KnownHosts(policy),
timeout: None,
algorithms: Default::default(),
};
let _client =
Client::connect_to_host("127.0.0.1", addr.port(), cfg).expect("connect should succeed");
let _ = server.join();
let (change, old_fps) = seen.lock().unwrap().take().expect("callback fired");
assert_eq!(change, HostKeyChange::Changed);
assert_eq!(old_fps, vec![host_key_fingerprint(&bogus_blob)]);
let real_blob = Ed25519HostKey::from_seed(seed).public_blob();
let guard = store.lock().unwrap();
assert!(matches!(
guard.lookup("127.0.0.1", addr.port(), "ssh-ed25519", &real_blob),
LookupResult::Match
));
assert!(matches!(
guard.lookup("127.0.0.1", addr.port(), "ssh-ed25519", &bogus_blob),
LookupResult::Mismatch { .. }
));
}
#[test]
fn client_advert_no_overrides_matches_defaults() {
let advert = build_default_kexinit(&mut OsRng, &AlgoOverrides::default());
let want_ciphers: Vec<String> = defaults::CIPHERS.iter().map(|s| s.to_string()).collect();
assert_eq!(advert.ciphers_c2s, want_ciphers);
assert_eq!(advert.ciphers_s2c, want_ciphers);
let n = advert.kex.len();
assert!(crate::transport::kex::is_strict_kex_marker(
&advert.kex[n - 1]
));
assert!(crate::transport::kex::is_strict_kex_marker(
&advert.kex[n - 2]
));
}
#[test]
fn client_cipher_override_replaces_and_keeps_markers() {
let over = AlgoOverrides {
ciphers: Some(vec!["aes128-ctr".to_string()]),
..Default::default()
};
let advert = build_default_kexinit(&mut OsRng, &over);
assert_eq!(advert.ciphers_c2s, vec!["aes128-ctr".to_string()]);
assert_eq!(advert.ciphers_s2c, vec!["aes128-ctr".to_string()]);
assert!(
advert
.kex
.iter()
.any(|k| crate::transport::kex::is_strict_kex_marker(k))
);
}
#[test]
fn client_kex_override_reappends_markers() {
let over = AlgoOverrides {
kex_algorithms: Some(vec!["curve25519-sha256".to_string()]),
..Default::default()
};
let advert = build_default_kexinit(&mut OsRng, &over);
assert_eq!(advert.kex[0], "curve25519-sha256");
let markers = advert
.kex
.iter()
.filter(|k| crate::transport::kex::is_strict_kex_marker(k))
.count();
assert_eq!(markers, 2, "both strict-kex markers must be re-appended");
}
#[test]
fn restricted_client_ciphers_negotiate_against_default_server() {
use crate::transport::kexinit::negotiate;
let client_over = AlgoOverrides {
ciphers: Some(vec!["aes128-ctr".to_string()]),
..Default::default()
};
let client =
build_default_kexinit(&mut OsRng, &client_over).with_ext_info_marker(Role::Client);
let server = build_default_kexinit(&mut OsRng, &AlgoOverrides::default());
let neg = negotiate(&client, &server).expect("should negotiate");
assert_eq!(neg.cipher_c2s, "aes128-ctr");
assert_eq!(neg.cipher_s2c, "aes128-ctr");
assert!(neg.strict_kex_enabled);
}
#[test]
fn disjoint_ciphers_fail_with_no_common_algorithm() {
use crate::transport::kexinit::negotiate;
let client_over = AlgoOverrides {
ciphers: Some(vec!["aes128-ctr".to_string()]),
..Default::default()
};
let server_over = AlgoOverrides {
ciphers: Some(vec!["aes256-ctr".to_string()]),
..Default::default()
};
let client = build_default_kexinit(&mut OsRng, &client_over);
let server = build_default_kexinit(&mut OsRng, &server_over);
assert!(matches!(
negotiate(&client, &server),
Err(Error::NoCommonAlgorithm(_))
));
}
}