#![cfg(feature = "std")]
use std::collections::{BTreeMap, BTreeSet};
use std::io::{ErrorKind, Read, Write};
use std::net::{SocketAddr, TcpListener, TcpStream, ToSocketAddrs};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc::{self, Receiver, Sender, TryRecvError};
use std::thread;
use std::time::{Duration, Instant};
use purecrypto::rng::RngCore;
use crate::auth::{Authenticator, ServerAuth, ServerStep};
use crate::channel::{
ChannelEvent, ChannelOpen, ChannelRequest, ConnectionState, SSH_EXTENDED_DATA_STDERR,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED, SSH_OPEN_RESOURCE_SHORTAGE,
};
use crate::driver::{Event, ServerDriver};
use crate::error::{Error, Result};
use crate::format::Writer;
use crate::hostkey::HostKey;
use crate::transport::kex::{defaults, is_strict_kex_marker};
use crate::transport::rekey::RekeyPolicy;
use crate::transport::{ExtInfo, KexAlgorithmsOwned, KexInit};
const MAX_AUTH_STEPS: usize = 64;
const MAX_CONNECTION_STEPS: usize = 10_000_000;
const MAX_DRAIN_STEPS: usize = 1_000_000;
const SUBSYSTEM_EGRESS_BACKLOG: usize = 32;
const SHELL_EGRESS_BACKLOG: usize = 4 * 1024 * 1024;
const MAX_ENV_PER_CHANNEL: usize = 64;
const MAX_ENV_BYTES_PER_CHANNEL: usize = 16 * 1024;
const SSH_DISCONNECT_BY_APPLICATION: u32 = 11;
const SSH_DISCONNECT_HOST_NOT_ALLOWED: u32 = 9;
#[derive(Debug, Clone)]
pub struct ExecResult {
pub stdout: Vec<u8>,
pub stderr: Vec<u8>,
pub exit_status: u32,
}
#[derive(Debug, Default, Clone)]
pub struct SessionEnv {
vars: BTreeMap<String, String>,
}
impl SessionEnv {
pub fn new() -> Self {
Self {
vars: BTreeMap::new(),
}
}
pub fn insert(&mut self, key: impl Into<String>, value: impl Into<String>) -> Option<String> {
self.vars.insert(key.into(), value.into())
}
pub fn get(&self, key: &str) -> Option<&str> {
self.vars.get(key).map(|s| s.as_str())
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &str)> {
self.vars.iter().map(|(k, v)| (k.as_str(), v.as_str()))
}
pub fn len(&self) -> usize {
self.vars.len()
}
pub fn is_empty(&self) -> bool {
self.vars.is_empty()
}
}
pub trait CommandHandler: Send + Sync {
fn handle(&self, user: &str, env: &SessionEnv, command: &str) -> ExecResult;
}
struct ShellRuntime {
pending_pty: Option<PtySpec>,
session: Option<Box<dyn ShellSession>>,
exited: Option<ShellExitStatus>,
exit_sent: bool,
pending_stdout: Vec<u8>,
}
impl ShellRuntime {
fn new() -> Self {
Self {
pending_pty: None,
session: None,
exited: None,
exit_sent: false,
pending_stdout: Vec::new(),
}
}
}
#[derive(Debug, Clone)]
pub struct PtySpec {
pub term: String,
pub cols: u32,
pub rows: u32,
pub px_w: u32,
pub px_h: u32,
pub modes: Vec<u8>,
}
#[derive(Debug, Clone)]
pub enum ShellExitStatus {
Exited(u32),
Signalled {
name: String,
core_dumped: bool,
message: String,
},
}
pub trait ShellHandler: Send + Sync {
fn spawn(
&self,
user: &str,
env: &SessionEnv,
pty: Option<PtySpec>,
) -> Result<Box<dyn ShellSession>>;
}
pub trait ShellSession: Send {
fn read(&mut self, buf: &mut [u8]) -> Result<usize>;
fn write(&mut self, buf: &[u8]) -> Result<usize>;
fn close_stdin(&mut self) -> Result<()>;
fn resize(&mut self, cols: u32, rows: u32, px_w: u32, px_h: u32) -> Result<()>;
fn try_exit(&mut self) -> Option<ShellExitStatus>;
}
pub use crate::stream::{ChannelEgress, ChannelStream};
struct SubsystemRuntime {
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,
}
#[derive(Debug, Clone, Copy)]
pub struct SessionOpenContext<'a> {
pub user: &'a str,
pub chroot_directory: Option<&'a str>,
pub print_motd: bool,
}
pub type SessionOpenCallback = Arc<dyn Fn(&SessionOpenContext<'_>) -> Result<()> + Send + Sync>;
pub trait SubsystemHandler: Send + Sync {
fn handle(&self, user: &str, env: &SessionEnv, name: &str, stream: ChannelStream)
-> Result<()>;
}
pub trait ExecStreamHandler: Send + Sync {
fn claims(&self, command: &str) -> bool;
fn run(&self, user: &str, env: &SessionEnv, command: &str, stream: ChannelStream)
-> Result<()>;
}
#[derive(Debug, Clone, Copy)]
pub struct DirectTcpipRequest<'a> {
pub dest_host: &'a str,
pub dest_port: u32,
pub orig_host: &'a str,
pub orig_port: u32,
}
pub trait DirectTcpipHandler: Send + Sync {
fn handle(
&self,
user: &str,
request: DirectTcpipRequest<'_>,
stream: ChannelStream,
) -> Result<()>;
}
#[derive(Debug, Clone, Copy)]
pub struct DirectStreamlocalRequest<'a> {
pub socket_path: &'a str,
}
pub trait DirectStreamlocalHandler: Send + Sync {
fn handle(
&self,
user: &str,
request: DirectStreamlocalRequest<'_>,
stream: ChannelStream,
) -> Result<()>;
}
#[derive(Clone)]
pub struct ForwardContext {
req_tx: Sender<ForwardOpenRequest>,
}
pub(crate) struct ForwardOpenRequest {
bound_address: String,
bound_port: u32,
orig_address: String,
orig_port: u32,
reply: std::sync::mpsc::SyncSender<Result<ChannelStream>>,
}
impl ForwardContext {
pub(crate) fn new(req_tx: Sender<ForwardOpenRequest>) -> Self {
Self { req_tx }
}
#[doc(hidden)]
pub fn for_test_no_opens() -> Self {
let (tx, _rx) = mpsc::channel();
drop(_rx);
Self { req_tx: tx }
}
pub fn open_forwarded_tcpip(
&self,
bound_address: &str,
bound_port: u16,
orig_address: &str,
orig_port: u16,
) -> Result<ChannelStream> {
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.req_tx
.send(ForwardOpenRequest {
bound_address: bound_address.to_string(),
bound_port: bound_port as u32,
orig_address: orig_address.to_string(),
orig_port: orig_port as u32,
reply: tx,
})
.map_err(|_| Error::Protocol("forwarded-tcpip: connection closed"))?;
rx.recv()
.map_err(|_| Error::Protocol("forwarded-tcpip: reply dropped"))?
}
}
pub trait TcpipForwardHandler: Send + Sync {
fn bind(
&self,
user: &str,
bind_address: &str,
bind_port: u16,
ctx: ForwardContext,
) -> Result<u16>;
fn unbind(&self, user: &str, bind_address: &str, bind_port: u16) -> Result<()>;
}
#[derive(Clone)]
pub struct StreamlocalForwardContext {
req_tx: Sender<StreamlocalOpenRequest>,
}
pub(crate) struct StreamlocalOpenRequest {
socket_path: String,
reply: std::sync::mpsc::SyncSender<Result<ChannelStream>>,
}
impl StreamlocalForwardContext {
pub(crate) fn new(req_tx: Sender<StreamlocalOpenRequest>) -> Self {
Self { req_tx }
}
#[doc(hidden)]
pub fn for_test_no_opens() -> Self {
let (tx, _rx) = mpsc::channel();
drop(_rx);
Self { req_tx: tx }
}
pub fn open_forwarded_streamlocal(&self, socket_path: &str) -> Result<ChannelStream> {
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.req_tx
.send(StreamlocalOpenRequest {
socket_path: socket_path.to_string(),
reply: tx,
})
.map_err(|_| Error::Protocol("forwarded-streamlocal: connection closed"))?;
rx.recv()
.map_err(|_| Error::Protocol("forwarded-streamlocal: reply dropped"))?
}
}
pub trait StreamlocalForwardHandler: Send + Sync {
fn bind(&self, user: &str, socket_path: &str, ctx: StreamlocalForwardContext) -> Result<()>;
fn unbind(&self, user: &str, socket_path: &str) -> Result<()>;
}
#[derive(Clone)]
pub struct AgentForwardContext {
req_tx: Sender<AgentOpenRequest>,
}
pub(crate) struct AgentOpenRequest {
reply: std::sync::mpsc::SyncSender<Result<ChannelStream>>,
}
impl AgentForwardContext {
pub(crate) fn new(req_tx: Sender<AgentOpenRequest>) -> Self {
Self { req_tx }
}
#[doc(hidden)]
pub fn for_test_no_opens() -> Self {
let (tx, _rx) = mpsc::channel();
drop(_rx);
Self { req_tx: tx }
}
pub fn open_auth_agent(&self) -> Result<ChannelStream> {
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.req_tx
.send(AgentOpenRequest { reply: tx })
.map_err(|_| Error::Protocol("auth-agent: connection closed"))?;
rx.recv()
.map_err(|_| Error::Protocol("auth-agent: reply dropped"))?
}
}
pub struct AgentForwardHandle {
pub auth_sock_path: std::path::PathBuf,
pub stopper: Box<dyn core::any::Any + Send + Sync>,
}
pub trait AgentForwardHandler: Send + Sync {
fn setup(&self, user: &str, ctx: AgentForwardContext) -> Result<AgentForwardHandle>;
}
#[derive(Clone)]
pub struct X11ForwardContext {
req_tx: Sender<X11OpenRequest>,
}
pub(crate) struct X11OpenRequest {
pub orig_host: String,
pub orig_port: u32,
pub reply: std::sync::mpsc::SyncSender<Result<ChannelStream>>,
}
impl X11ForwardContext {
pub(crate) fn new(req_tx: Sender<X11OpenRequest>) -> Self {
Self { req_tx }
}
#[doc(hidden)]
pub fn for_test_no_opens() -> Self {
let (tx, _rx) = mpsc::channel();
drop(_rx);
Self { req_tx: tx }
}
pub fn open_x11(&self, orig_host: String, orig_port: u32) -> Result<ChannelStream> {
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.req_tx
.send(X11OpenRequest {
orig_host,
orig_port,
reply: tx,
})
.map_err(|_| Error::Protocol("x11: connection closed"))?;
rx.recv()
.map_err(|_| Error::Protocol("x11: reply dropped"))?
}
}
pub struct X11ForwardHandle {
pub display_env: String,
pub display_number: u16,
pub stopper: Box<dyn core::any::Any + Send + Sync>,
}
pub trait X11ForwardHandler: Send + Sync {
fn setup(
&self,
user: &str,
single_connection: bool,
auth_protocol: &str,
auth_cookie: &str,
screen: u32,
ctx: X11ForwardContext,
) -> Result<X11ForwardHandle>;
}
pub struct Config {
pub host_keys: Vec<Box<dyn HostKey + Send + Sync>>,
pub authenticator: Arc<dyn AuthenticatorFactory>,
pub allowed_auth_methods: Vec<&'static str>,
pub command_handler: Arc<dyn CommandHandler>,
pub exec_stream_handler: Option<Arc<dyn ExecStreamHandler>>,
pub shell_handler: Option<Arc<dyn ShellHandler>>,
pub subsystem_handler: Option<Arc<dyn SubsystemHandler>>,
pub direct_tcpip_handler: Option<Arc<dyn DirectTcpipHandler>>,
pub tcpip_forward_handler: Option<Arc<dyn TcpipForwardHandler>>,
pub direct_streamlocal_handler: Option<Arc<dyn DirectStreamlocalHandler>>,
pub streamlocal_forward_handler: Option<Arc<dyn StreamlocalForwardHandler>>,
pub agent_forward_handler: Option<Arc<dyn AgentForwardHandler>>,
pub x11_forward_handler: Option<Arc<dyn X11ForwardHandler>>,
pub on_session_open: Option<SessionOpenCallback>,
pub rekey_policy: RekeyPolicy,
pub accept_env: Vec<String>,
pub login_grace_time: Duration,
pub max_connections: Option<usize>,
pub ciphers: Option<Vec<String>>,
pub macs: Option<Vec<String>>,
pub kex_algorithms: Option<Vec<String>>,
pub host_key_algorithms: Option<Vec<String>>,
pub policy: Option<Arc<crate::config::SshServerConfig>>,
pub group_resolver: Option<GroupResolver>,
pub compression: Option<crate::config::Compression>,
pub default_auth_methods: Vec<String>,
pub ca_signature_algorithms: Vec<String>,
}
pub type GroupResolver = Arc<dyn Fn(&str) -> Vec<String> + Send + Sync>;
#[derive(Debug, Clone)]
pub struct EffectivePolicy {
pub allow_agent_forwarding: Option<bool>,
pub x11_forwarding: Option<bool>,
pub max_sessions: Option<u32>,
pub allow_tcp_forwarding: Option<crate::config::TcpForwarding>,
pub permit_open: Option<Vec<crate::config::HostPort>>,
pub permit_listen: Option<Vec<crate::config::HostPort>>,
pub gateway_ports: Option<crate::config::ServerGatewayPorts>,
pub force_command: Option<String>,
pub chroot_directory: Option<String>,
pub client_alive_interval: Option<u32>,
pub client_alive_count_max: Option<u32>,
pub print_motd: Option<bool>,
pub cert_caps: Option<crate::auth::AuthCertCaps>,
}
impl EffectivePolicy {
pub fn unrestricted() -> Self {
EffectivePolicy {
allow_agent_forwarding: None,
x11_forwarding: None,
max_sessions: None,
allow_tcp_forwarding: None,
permit_open: None,
permit_listen: None,
gateway_ports: None,
force_command: None,
chroot_directory: None,
client_alive_interval: None,
client_alive_count_max: None,
print_motd: None,
cert_caps: None,
}
}
fn pty_allowed(&self) -> bool {
self.cert_caps.as_ref().is_none_or(|c| c.permit_pty)
}
fn agent_forwarding_allowed(&self) -> bool {
self.allow_agent_forwarding != Some(false)
&& self
.cert_caps
.as_ref()
.is_none_or(|c| c.permit_agent_forwarding)
}
fn x11_forwarding_allowed(&self) -> bool {
self.x11_forwarding != Some(false)
&& self
.cert_caps
.as_ref()
.is_none_or(|c| c.permit_x11_forwarding)
}
fn local_forwarding_allowed(&self) -> bool {
self.allow_tcp_forwarding.is_none_or(|p| p.local_allowed())
&& self
.cert_caps
.as_ref()
.is_none_or(|c| c.permit_port_forwarding)
}
fn remote_forwarding_allowed(&self) -> bool {
self.allow_tcp_forwarding.is_none_or(|p| p.remote_allowed())
&& self
.cert_caps
.as_ref()
.is_none_or(|c| c.permit_port_forwarding)
}
fn permit_open_allows(&self, host: &str, port: u16) -> bool {
match &self.permit_open {
None => true,
Some(list) => list.iter().any(|e| e.matches(host, port)),
}
}
fn permit_listen_allows(&self, host: &str, port: u16) -> bool {
match &self.permit_listen {
None => true,
Some(list) => list.iter().any(|e| e.matches(host, port)),
}
}
}
pub const HARD_BLOCKED_ENV_NAMES: &[&str] = &[
"LD_PRELOAD",
"LD_LIBRARY_PATH",
"LD_AUDIT",
"LD_BIND_NOT",
"DYLD_INSERT_LIBRARIES",
"DYLD_LIBRARY_PATH",
"BASH_ENV",
"ENV",
"IFS",
"PATH",
"SHELL",
"HOME",
"USER",
"LOGNAME",
];
fn env_glob_match(pattern: &str, name: &str) -> bool {
let pb = pattern.as_bytes();
let nb = name.as_bytes();
fn rec(p: &[u8], n: &[u8]) -> bool {
let mut pi = 0usize;
let mut ni = 0usize;
while pi < p.len() {
match p[pi] {
b'*' => {
while pi < p.len() && p[pi] == b'*' {
pi += 1;
}
if pi == p.len() {
return true;
}
while ni <= n.len() {
if rec(&p[pi..], &n[ni..]) {
return true;
}
ni += 1;
}
return false;
}
b'?' => {
if ni >= n.len() {
return false;
}
pi += 1;
ni += 1;
}
c => {
if ni >= n.len() || n[ni] != c {
return false;
}
pi += 1;
ni += 1;
}
}
}
ni == n.len()
}
rec(pb, nb)
}
pub(crate) fn env_name_accepted(name: &str, accept_env: &[String]) -> bool {
if HARD_BLOCKED_ENV_NAMES.contains(&name) {
return false;
}
accept_env.iter().any(|pat| env_glob_match(pat, name))
}
impl Config {
pub fn new(
host_keys: Vec<Box<dyn HostKey + Send + Sync>>,
authenticator: Arc<dyn AuthenticatorFactory>,
allowed_auth_methods: Vec<&'static str>,
command_handler: Arc<dyn CommandHandler>,
) -> Self {
Self {
host_keys,
authenticator,
allowed_auth_methods,
command_handler,
exec_stream_handler: None,
shell_handler: None,
subsystem_handler: None,
direct_tcpip_handler: None,
tcpip_forward_handler: None,
direct_streamlocal_handler: None,
streamlocal_forward_handler: None,
agent_forward_handler: None,
x11_forward_handler: None,
on_session_open: None,
rekey_policy: RekeyPolicy::default(),
accept_env: Vec::new(),
login_grace_time: Duration::from_secs(120),
max_connections: Some(256),
ciphers: None,
macs: None,
kex_algorithms: None,
host_key_algorithms: None,
policy: None,
group_resolver: None,
compression: None,
default_auth_methods: Vec::new(),
ca_signature_algorithms: Vec::new(),
}
}
pub fn with_auth_methods(mut self, methods: Vec<String>) -> Self {
self.default_auth_methods = methods;
self
}
pub fn with_policy(mut self, policy: Arc<crate::config::SshServerConfig>) -> Self {
self.policy = Some(policy);
self
}
pub fn with_group_resolver(mut self, resolver: GroupResolver) -> Self {
self.group_resolver = Some(resolver);
self
}
pub fn with_algorithms(
mut self,
ciphers: Option<Vec<String>>,
macs: Option<Vec<String>>,
kex_algorithms: Option<Vec<String>>,
host_key_algorithms: Option<Vec<String>>,
) -> Self {
self.ciphers = ciphers;
self.macs = macs;
self.kex_algorithms = kex_algorithms;
self.host_key_algorithms = host_key_algorithms;
self
}
pub fn with_accept_env(mut self, patterns: Vec<String>) -> Self {
self.accept_env = patterns;
self
}
pub fn with_login_grace_time(mut self, dur: Duration) -> Self {
self.login_grace_time = dur;
self
}
pub fn with_max_connections(mut self, max: Option<usize>) -> Self {
self.max_connections = max;
self
}
pub fn with_shell(mut self, handler: Arc<dyn ShellHandler>) -> Self {
self.shell_handler = Some(handler);
self
}
pub fn with_subsystem(mut self, handler: Arc<dyn SubsystemHandler>) -> Self {
self.subsystem_handler = Some(handler);
self
}
pub fn with_exec_stream_handler(mut self, handler: Arc<dyn ExecStreamHandler>) -> Self {
self.exec_stream_handler = Some(handler);
self
}
pub fn with_direct_tcpip(mut self, handler: Arc<dyn DirectTcpipHandler>) -> Self {
self.direct_tcpip_handler = Some(handler);
self
}
pub fn with_tcpip_forward(mut self, handler: Arc<dyn TcpipForwardHandler>) -> Self {
self.tcpip_forward_handler = Some(handler);
self
}
pub fn with_direct_streamlocal(mut self, handler: Arc<dyn DirectStreamlocalHandler>) -> Self {
self.direct_streamlocal_handler = Some(handler);
self
}
pub fn with_streamlocal_forward(mut self, handler: Arc<dyn StreamlocalForwardHandler>) -> Self {
self.streamlocal_forward_handler = Some(handler);
self
}
pub fn with_agent_forward(mut self, handler: Arc<dyn AgentForwardHandler>) -> Self {
self.agent_forward_handler = Some(handler);
self
}
pub fn with_x11_forward(mut self, handler: Arc<dyn X11ForwardHandler>) -> Self {
self.x11_forward_handler = Some(handler);
self
}
pub fn on_session_open<F>(mut self, f: F) -> Self
where
F: Fn(&SessionOpenContext<'_>) -> Result<()> + Send + Sync + 'static,
{
self.on_session_open = Some(Arc::new(f));
self
}
}
pub trait AuthenticatorFactory: Send + Sync {
fn build(&self) -> Box<dyn Authenticator>;
fn build_with_peer(&self, peer: Option<&str>) -> Box<dyn Authenticator> {
let _ = peer;
self.build()
}
}
impl<F> AuthenticatorFactory for F
where
F: Fn() -> Box<dyn Authenticator> + Send + Sync,
{
fn build(&self) -> Box<dyn Authenticator> {
(self)()
}
}
pub struct Server {
listener: TcpListener,
cfg: Arc<Config>,
}
impl Server {
pub fn bind<A: ToSocketAddrs>(addr: A, cfg: Config) -> Result<Self> {
if cfg.host_keys.is_empty() {
return Err(Error::Protocol("server: no host keys configured"));
}
let listener = TcpListener::bind(addr)?;
Ok(Self {
listener,
cfg: Arc::new(cfg),
})
}
pub fn local_addr(&self) -> Result<SocketAddr> {
Ok(self.listener.local_addr()?)
}
pub fn accept_one(&mut self) -> Result<()> {
let (stream, peer) = self.listener.accept()?;
handle_session_with_peer(stream, peer, self.cfg.clone())
}
pub fn serve(&mut self) -> Result<()> {
let live = Arc::new(AtomicUsize::new(0));
loop {
let (stream, peer) = self.listener.accept()?;
if let Some(max) = self.cfg.max_connections {
let prev = live.fetch_add(1, Ordering::AcqRel);
if prev >= max {
live.fetch_sub(1, Ordering::AcqRel);
drop(stream);
continue;
}
} else {
live.fetch_add(1, Ordering::AcqRel);
}
let cfg = self.cfg.clone();
let guard = ConnGuard { live: live.clone() };
let spawn = thread::Builder::new()
.name("puressh-conn".into())
.spawn(move || {
let _guard = guard;
let _ = handle_session_with_peer(stream, peer, cfg);
});
if spawn.is_err() {
continue;
}
}
}
}
struct ConnGuard {
live: Arc<AtomicUsize>,
}
impl Drop for ConnGuard {
fn drop(&mut self) {
self.live.fetch_sub(1, Ordering::AcqRel);
}
}
pub fn handle_session(stream: TcpStream, cfg: Arc<Config>) -> Result<()> {
let peer = stream
.peer_addr()
.unwrap_or_else(|_| SocketAddr::from(([0, 0, 0, 0], 0)));
handle_session_with_peer(stream, peer, cfg)
}
pub fn handle_session_with_peer(
stream: TcpStream,
peer: SocketAddr,
cfg: Arc<Config>,
) -> Result<()> {
let local = stream.local_addr().ok();
handle_connection_inner(stream, peer, local, cfg)
}
fn map_preauth_timeout<T>(r: Result<T>) -> Result<T> {
match r {
Err(Error::Io(e)) if matches!(e.kind(), ErrorKind::TimedOut | ErrorKind::WouldBlock) => {
Err(Error::Io(std::io::Error::new(
ErrorKind::TimedOut,
"pre-auth inactivity timeout (LoginGraceTime)",
)))
}
other => other,
}
}
fn handle_connection_inner(
mut stream: TcpStream,
peer: SocketAddr,
local: Option<SocketAddr>,
cfg: Arc<Config>,
) -> Result<()> {
stream.set_nodelay(true)?;
let grace = cfg.login_grace_time;
let deadline = if grace.is_zero() {
None
} else {
Some(Instant::now() + grace)
};
let mut driver = ServerDriver::new(cfg.clone());
driver.start(Instant::now())?;
map_preauth_timeout(srv_drive_handshake(&mut stream, &mut driver, deadline))?;
let session_id = driver.session_id().to_vec();
let peer_ip = ip_string(&peer);
let local_ip = local.and_then(|a| ip_string(&a));
let local_port = local.map(|a| a.port());
let preauth = resolve_preauth_policy(&cfg, peer_ip.as_deref(), local_ip.as_deref(), local_port);
let (user, cert_caps) = map_preauth_timeout(do_server_auth(
&mut stream,
&mut driver,
&cfg,
session_id,
&preauth,
peer_ip.as_deref(),
local_ip.as_deref(),
local_port,
deadline,
))?;
if !grace.is_zero() {
stream.set_read_timeout(None)?;
}
let groups = resolve_user_groups(&cfg, &user);
let mut effective = resolve_effective_policy(
&cfg,
&user,
groups.as_deref(),
peer_ip.as_deref(),
local_ip.as_deref(),
local_port,
);
if let Some(caps) = cert_caps {
if let Some(forced) = caps.force_command.clone() {
effective.force_command = Some(forced);
}
effective.cert_caps = Some(caps);
}
if let Some(hook) = cfg.on_session_open.clone() {
let ctx = SessionOpenContext {
user: &user,
chroot_directory: effective.chroot_directory.as_deref(),
print_motd: effective.print_motd.unwrap_or(false),
};
hook(&ctx)?;
}
driver.notify_auth_success();
driver.set_rekey_policy(cfg.rekey_policy);
let r = do_connection_phase(&mut stream, &mut driver, &cfg, &user, &effective);
let _ = srv_send_disconnect(
&mut stream,
&mut driver,
SSH_DISCONNECT_BY_APPLICATION,
"closing session",
);
r
}
#[allow(clippy::too_many_arguments)]
fn ip_string(addr: &SocketAddr) -> Option<String> {
let ip = addr.ip();
if ip.is_unspecified() {
None
} else {
Some(ip.to_string())
}
}
struct PreAuthPolicy {
methods: Vec<&'static str>,
max_auth_tries: Option<u32>,
banner: Option<String>,
}
fn resolve_preauth_policy(
cfg: &Config,
address: Option<&str>,
local_address: Option<&str>,
local_port: Option<u16>,
) -> PreAuthPolicy {
let Some(policy) = cfg.policy.as_ref() else {
return PreAuthPolicy {
methods: cfg.allowed_auth_methods.clone(),
max_auth_tries: None,
banner: None,
};
};
let ctx = crate::config::MatchContext {
host: "",
address,
local_address,
local_port,
..crate::config::MatchContext::default()
};
let opts = policy.resolve(&ctx, crate::config::match_block::ExecPolicy::Deny);
let methods = resolve_auth_methods(&cfg.allowed_auth_methods, &opts);
let banner = opts
.banner
.as_deref()
.and_then(|p| std::fs::read_to_string(p).ok());
PreAuthPolicy {
methods,
max_auth_tries: opts.max_auth_tries,
banner,
}
}
fn resolve_auth_methods(
base: &[&'static str],
opts: &crate::config::ServerOptions,
) -> Vec<&'static str> {
let mut methods: Vec<&'static str> = base.to_vec();
if opts.pubkey_authentication == Some(false) {
methods.retain(|m| *m != "publickey");
}
if opts.password_authentication == Some(false) {
methods.retain(|m| *m != "password");
}
if opts.kbd_interactive_authentication == Some(false) {
methods.retain(|m| *m != "keyboard-interactive");
}
methods
}
fn resolve_user_groups(cfg: &Config, user: &str) -> Option<Vec<String>> {
cfg.group_resolver.as_ref().map(|r| r(user))
}
fn resolve_effective_policy(
cfg: &Config,
user: &str,
groups: Option<&[String]>,
address: Option<&str>,
local_address: Option<&str>,
local_port: Option<u16>,
) -> EffectivePolicy {
let Some(policy) = cfg.policy.as_ref() else {
return EffectivePolicy::unrestricted();
};
let ctx = crate::config::MatchContext {
host: "",
user: Some(user),
groups,
address,
local_address,
local_port,
..crate::config::MatchContext::default()
};
let opts = policy.resolve(&ctx, crate::config::match_block::ExecPolicy::Deny);
EffectivePolicy {
allow_agent_forwarding: opts.allow_agent_forwarding,
x11_forwarding: opts.x11_forwarding,
max_sessions: opts.max_sessions,
allow_tcp_forwarding: opts.allow_tcp_forwarding,
permit_open: opts.permit_open,
permit_listen: opts.permit_listen,
gateway_ports: opts.gateway_ports,
force_command: opts.force_command,
chroot_directory: opts.chroot_directory,
client_alive_interval: opts.client_alive_interval,
client_alive_count_max: opts.client_alive_count_max,
print_motd: opts.print_motd,
cert_caps: None,
}
}
#[allow(clippy::too_many_arguments)]
fn do_server_auth(
stream: &mut TcpStream,
driver: &mut ServerDriver,
cfg: &Config,
session_id: Vec<u8>,
preauth: &PreAuthPolicy,
peer_ip: Option<&str>,
local_ip: Option<&str>,
local_port: Option<u16>,
deadline: Option<Instant>,
) -> Result<(String, Option<crate::auth::AuthCertCaps>)> {
let methods = preauth.methods.clone();
let auth_impl = cfg.authenticator.build_with_peer(peer_ip);
let mut server_auth = ServerAuth::new(session_id, methods, auth_impl);
server_auth.set_max_auth_tries(preauth.max_auth_tries);
server_auth.set_now(unix_now());
server_auth.set_ca_signature_algorithms(cfg.ca_signature_algorithms.clone());
if let Some(text) = preauth.banner.as_deref() {
let banner = crate::auth::message::UserauthBanner {
message: text.to_string(),
language: String::new(),
};
srv_send(stream, driver, &banner.encode())?;
}
let mut resolved_user: Option<String> = None;
for _ in 0..MAX_AUTH_STEPS {
let payload = srv_read(stream, driver, deadline)?;
if let Some((user, method)) = ServerAuth::peek_request(&payload) {
if resolved_user.is_none() {
let reres = reresolve_user_policy(
cfg,
&preauth.methods,
&user,
peer_ip,
local_ip,
local_port,
);
server_auth.set_accepted_methods(reres.methods);
server_auth.notify_user_resolved(&user, &reres.auth_methods);
if let Some(text) = reres.banner.as_deref() {
let banner = crate::auth::message::UserauthBanner {
message: text.to_string(),
language: String::new(),
};
srv_send(stream, driver, &banner.encode())?;
}
resolved_user = Some(user);
}
if method != "none" && !server_auth.accepted_methods().contains(&method) {
match server_auth.reject_unadvertised()? {
ServerStep::Send(p) => {
srv_send(stream, driver, &p)?;
continue;
}
ServerStep::Disconnect(reason) => {
let _ = srv_send_disconnect(
stream,
driver,
SSH_DISCONNECT_HOST_NOT_ALLOWED,
reason,
);
return Err(Error::AuthFailed);
}
ServerStep::Authenticated { .. } => unreachable!(),
}
}
}
match server_auth.on_packet(&payload)? {
ServerStep::Send(p) => srv_send(stream, driver, &p)?,
ServerStep::Authenticated {
payload,
user,
cert_caps,
} => {
srv_send(stream, driver, &payload)?;
return Ok((user, cert_caps));
}
ServerStep::Disconnect(reason) => {
let _ =
srv_send_disconnect(stream, driver, SSH_DISCONNECT_HOST_NOT_ALLOWED, reason);
return Err(Error::AuthFailed);
}
}
}
Err(Error::Protocol("auth: too many steps"))
}
struct ReResolvedUserPolicy {
methods: Vec<&'static str>,
banner: Option<String>,
auth_methods: Vec<String>,
}
fn reresolve_user_policy(
cfg: &Config,
base_methods: &[&'static str],
user: &str,
address: Option<&str>,
local_address: Option<&str>,
local_port: Option<u16>,
) -> ReResolvedUserPolicy {
let Some(policy) = cfg.policy.as_ref() else {
return ReResolvedUserPolicy {
methods: base_methods.to_vec(),
banner: None,
auth_methods: cfg.default_auth_methods.clone(),
};
};
let groups = resolve_user_groups(cfg, user);
let ctx = crate::config::MatchContext {
host: "",
user: Some(user),
groups: groups.as_deref(),
address,
local_address,
local_port,
..crate::config::MatchContext::default()
};
let opts = policy.resolve(&ctx, crate::config::match_block::ExecPolicy::Deny);
let methods = resolve_auth_methods(&cfg.allowed_auth_methods, &opts);
let banner = opts
.banner
.as_deref()
.and_then(|p| std::fs::read_to_string(p).ok());
let auth_methods = opts
.authentication_methods
.clone()
.unwrap_or_else(|| cfg.default_auth_methods.clone());
ReResolvedUserPolicy {
methods,
banner,
auth_methods,
}
}
struct ForwardConn {
req_tx: Sender<ForwardOpenRequest>,
req_rx: Receiver<ForwardOpenRequest>,
pending_opens: BTreeMap<u32, std::sync::mpsc::SyncSender<Result<ChannelStream>>>,
owned_bindings: Vec<(String, u16)>,
}
impl ForwardConn {
fn new() -> Self {
let (req_tx, req_rx) = std::sync::mpsc::channel();
Self {
req_tx,
req_rx,
pending_opens: BTreeMap::new(),
owned_bindings: Vec::new(),
}
}
fn drain_pending(
&mut self,
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
) -> Result<()> {
loop {
match self.req_rx.try_recv() {
Ok(req) => {
let kind = ChannelOpen::ForwardedTcpip {
dest_host: req.bound_address.clone(),
dest_port: req.bound_port,
orig_host: req.orig_address.clone(),
orig_port: req.orig_port,
};
let (local_id, payload) = conn.open(kind)?;
srv_send(stream, driver, &payload)?;
self.pending_opens.insert(local_id, req.reply);
}
Err(TryRecvError::Empty) => return Ok(()),
Err(TryRecvError::Disconnected) => return Ok(()),
}
}
}
}
struct StreamlocalForwardConn {
req_tx: Sender<StreamlocalOpenRequest>,
req_rx: Receiver<StreamlocalOpenRequest>,
pending_opens: BTreeMap<u32, std::sync::mpsc::SyncSender<Result<ChannelStream>>>,
owned_bindings: Vec<String>,
}
impl StreamlocalForwardConn {
fn new() -> Self {
let (req_tx, req_rx) = std::sync::mpsc::channel();
Self {
req_tx,
req_rx,
pending_opens: BTreeMap::new(),
owned_bindings: Vec::new(),
}
}
fn drain_pending(
&mut self,
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
) -> Result<()> {
loop {
match self.req_rx.try_recv() {
Ok(req) => {
let kind = ChannelOpen::ForwardedStreamlocal {
socket_path: req.socket_path.clone(),
};
let (local_id, payload) = conn.open(kind)?;
srv_send(stream, driver, &payload)?;
self.pending_opens.insert(local_id, req.reply);
}
Err(TryRecvError::Empty) => return Ok(()),
Err(TryRecvError::Disconnected) => return Ok(()),
}
}
}
}
struct AgentForwardConn {
req_tx: Sender<AgentOpenRequest>,
req_rx: Receiver<AgentOpenRequest>,
pending_opens: BTreeMap<u32, std::sync::mpsc::SyncSender<Result<ChannelStream>>>,
active: BTreeMap<u32, AgentForwardHandle>,
}
impl AgentForwardConn {
fn new() -> Self {
let (req_tx, req_rx) = std::sync::mpsc::channel();
Self {
req_tx,
req_rx,
pending_opens: BTreeMap::new(),
active: BTreeMap::new(),
}
}
fn drain_pending(
&mut self,
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
) -> Result<()> {
loop {
match self.req_rx.try_recv() {
Ok(req) => {
let (local_id, payload) = conn.open(ChannelOpen::AuthAgent)?;
srv_send(stream, driver, &payload)?;
self.pending_opens.insert(local_id, req.reply);
}
Err(TryRecvError::Empty) => return Ok(()),
Err(TryRecvError::Disconnected) => return Ok(()),
}
}
}
}
struct X11ForwardConn {
req_tx: Sender<X11OpenRequest>,
req_rx: Receiver<X11OpenRequest>,
pending_opens: BTreeMap<u32, std::sync::mpsc::SyncSender<Result<ChannelStream>>>,
active: BTreeMap<u32, X11ForwardHandle>,
}
impl X11ForwardConn {
fn new() -> Self {
let (req_tx, req_rx) = std::sync::mpsc::channel();
Self {
req_tx,
req_rx,
pending_opens: BTreeMap::new(),
active: BTreeMap::new(),
}
}
fn drain_pending(
&mut self,
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
) -> Result<()> {
loop {
match self.req_rx.try_recv() {
Ok(req) => {
let kind = ChannelOpen::X11 {
orig_host: req.orig_host.clone(),
orig_port: req.orig_port,
};
let (local_id, payload) = conn.open(kind)?;
srv_send(stream, driver, &payload)?;
self.pending_opens.insert(local_id, req.reply);
}
Err(TryRecvError::Empty) => return Ok(()),
Err(TryRecvError::Disconnected) => return Ok(()),
}
}
}
}
#[allow(clippy::too_many_arguments)]
fn do_connection_phase(
stream: &mut TcpStream,
driver: &mut ServerDriver,
cfg: &Config,
user: &str,
effective: &EffectivePolicy,
) -> Result<()> {
let mut conn = ConnectionState::new();
let mut any_channel_opened = false;
let mut steps = 0usize;
let mut shells: BTreeMap<u32, ShellRuntime> = BTreeMap::new();
let mut subsystems: BTreeMap<u32, SubsystemRuntime> = BTreeMap::new();
let mut envs: BTreeMap<u32, SessionEnv> = BTreeMap::new();
let mut forward = ForwardConn::new();
let mut agent_forward = AgentForwardConn::new();
let mut x11_forward = X11ForwardConn::new();
let mut streamlocal_forward = StreamlocalForwardConn::new();
let mut polling_active = false;
let mut extras = ConnExtras::new();
let result = do_connection_loop(
stream,
driver,
cfg,
user,
effective,
&mut conn,
&mut any_channel_opened,
&mut steps,
&mut shells,
&mut subsystems,
&mut envs,
&mut forward,
&mut agent_forward,
&mut x11_forward,
&mut streamlocal_forward,
&mut polling_active,
&mut extras,
);
if let Some(handler) = cfg.tcpip_forward_handler.clone() {
for (addr, port) in forward.owned_bindings.drain(..) {
let _ = handler.unbind(user, &addr, port);
}
}
for (_id, reply) in forward.pending_opens.drain_filter_compat() {
let _ = reply.send(Err(Error::Protocol(
"forwarded-tcpip: connection torn down",
)));
}
agent_forward.active.clear();
for (_id, reply) in agent_forward.pending_opens.drain_filter_compat() {
let _ = reply.send(Err(Error::Protocol("auth-agent: connection torn down")));
}
x11_forward.active.clear();
for (_id, reply) in x11_forward.pending_opens.drain_filter_compat() {
let _ = reply.send(Err(Error::Protocol("x11: connection torn down")));
}
if let Some(handler) = cfg.streamlocal_forward_handler.clone() {
for path in streamlocal_forward.owned_bindings.drain(..) {
let _ = handler.unbind(user, &path);
}
}
for (_id, reply) in streamlocal_forward.pending_opens.drain_filter_compat() {
let _ = reply.send(Err(Error::Protocol(
"forwarded-streamlocal: connection torn down",
)));
}
result
}
struct ConnExtras {
session_channels: BTreeSet<u32>,
last_activity: Instant,
last_keepalive: Instant,
missed_keepalives: u32,
}
impl ConnExtras {
fn new() -> Self {
let now = Instant::now();
ConnExtras {
session_channels: BTreeSet::new(),
last_activity: now,
last_keepalive: now,
missed_keepalives: 0,
}
}
}
trait DrainFilterCompat<K, V> {
fn drain_filter_compat(&mut self) -> alloc::vec::IntoIter<(K, V)>;
}
impl<K: Ord + Clone, V> DrainFilterCompat<K, V> for BTreeMap<K, V> {
fn drain_filter_compat(&mut self) -> alloc::vec::IntoIter<(K, V)> {
let keys: Vec<K> = self.keys().cloned().collect();
let mut out = Vec::with_capacity(keys.len());
for k in keys {
if let Some(v) = self.remove(&k) {
out.push((k, v));
}
}
out.into_iter()
}
}
#[allow(clippy::too_many_arguments)]
fn do_connection_loop(
stream: &mut TcpStream,
driver: &mut ServerDriver,
cfg: &Config,
user: &str,
effective: &EffectivePolicy,
conn: &mut ConnectionState,
any_channel_opened: &mut bool,
steps: &mut usize,
shells: &mut BTreeMap<u32, ShellRuntime>,
subsystems: &mut BTreeMap<u32, SubsystemRuntime>,
envs: &mut BTreeMap<u32, SessionEnv>,
forward: &mut ForwardConn,
agent_forward: &mut AgentForwardConn,
x11_forward: &mut X11ForwardConn,
streamlocal_forward: &mut StreamlocalForwardConn,
polling_active: &mut bool,
extras: &mut ConnExtras,
) -> Result<()> {
loop {
*steps += 1;
if *steps > MAX_CONNECTION_STEPS {
return Err(Error::Protocol("connection: step cap exceeded"));
}
let any_shell_alive = shells.values().any(|rt| rt.session.is_some());
let any_subsystem_alive = !subsystems.is_empty();
let any_forward_alive = !forward.owned_bindings.is_empty();
let any_agent_fwd_alive = !agent_forward.active.is_empty();
let any_x11_fwd_alive = !x11_forward.active.is_empty();
let any_streamlocal_fwd_alive = !streamlocal_forward.owned_bindings.is_empty();
let keepalive_enabled = effective.client_alive_interval.is_some_and(|i| i > 0);
let want_polling = any_shell_alive
|| any_subsystem_alive
|| any_forward_alive
|| any_agent_fwd_alive
|| any_x11_fwd_alive
|| any_streamlocal_fwd_alive
|| keepalive_enabled;
if want_polling && !*polling_active {
let _ = stream.set_read_timeout(Some(Duration::from_millis(50)));
*polling_active = true;
} else if !want_polling && *polling_active {
let _ = stream.set_read_timeout(None);
*polling_active = false;
}
if *polling_active && !driver.is_kexing() {
drain_shells(stream, driver, conn, shells)?;
finalize_exited_shells(stream, driver, conn, shells)?;
drain_subsystems(stream, driver, conn, subsystems)?;
forward.drain_pending(stream, driver, conn)?;
agent_forward.drain_pending(stream, driver, conn)?;
x11_forward.drain_pending(stream, driver, conn)?;
streamlocal_forward.drain_pending(stream, driver, conn)?;
}
if keepalive_enabled && !driver.is_kexing() {
let interval = Duration::from_secs(effective.client_alive_interval.unwrap_or(0) as u64);
let count_max = effective.client_alive_count_max.unwrap_or(3);
let now = Instant::now();
if now.duration_since(extras.last_activity) >= interval
&& now.duration_since(extras.last_keepalive) >= interval
{
if extras.missed_keepalives >= count_max {
return Err(Error::Protocol(
"client keepalive: ClientAliveCountMax exceeded",
));
}
let p = conn.send_global_request(crate::channel::GlobalRequest::Keepalive, true);
srv_send(stream, driver, &p)?;
extras.last_keepalive = now;
extras.missed_keepalives = extras.missed_keepalives.saturating_add(1);
}
}
if *any_channel_opened
&& !conn.channels().any(|c| !c.is_fully_closed())
&& !any_forward_alive
&& !any_agent_fwd_alive
&& !any_x11_fwd_alive
&& !any_streamlocal_fwd_alive
{
return Ok(());
}
let payload = if *polling_active {
match srv_read_maybe_timeout(stream, driver)? {
Some(p) => p,
None => continue, }
} else {
srv_read(stream, driver, None)?
};
extras.last_activity = Instant::now();
extras.missed_keepalives = 0;
dispatch_app_packet(
stream,
driver,
conn,
cfg,
effective,
user,
&payload,
any_channel_opened,
extras,
shells,
subsystems,
envs,
forward,
agent_forward,
x11_forward,
streamlocal_forward,
)?;
}
}
fn drain_shells(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
shells: &mut BTreeMap<u32, ShellRuntime>,
) -> Result<()> {
let mut buf = [0u8; 8 * 1024];
let channels: Vec<u32> = shells.keys().copied().collect();
for ch in channels {
let Some(rt) = shells.get_mut(&ch) else {
continue;
};
if rt.session.is_none() {
continue;
}
if !rt.pending_stdout.is_empty() {
let leftover = core::mem::take(&mut rt.pending_stdout);
emit_channel_data(stream, driver, conn, ch, &leftover, rt)?;
}
let mut pulled = 0usize;
while rt.pending_stdout.len() < SHELL_EGRESS_BACKLOG && pulled < 64 * 1024 {
if let Some(sess) = rt.session.as_mut() {
let n = sess.read(&mut buf)?;
if n == 0 {
break;
}
pulled += n;
let bytes = buf[..n].to_vec();
emit_channel_data(stream, driver, conn, ch, &bytes, rt)?;
} else {
break;
}
}
if rt.exited.is_none()
&& let Some(sess) = rt.session.as_mut()
&& let Some(status) = sess.try_exit()
{
rt.exited = Some(status);
}
}
Ok(())
}
fn emit_channel_data(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
channel: u32,
bytes: &[u8],
rt: &mut ShellRuntime,
) -> Result<()> {
let mut off = 0usize;
while off < bytes.len() {
let (payload, taken) = conn.send_data(channel, &bytes[off..])?;
if taken == 0 {
rt.pending_stdout.extend_from_slice(&bytes[off..]);
return Ok(());
}
srv_send(stream, driver, &payload)?;
off += taken;
}
Ok(())
}
fn finalize_exited_shells(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
shells: &mut BTreeMap<u32, ShellRuntime>,
) -> Result<()> {
let channels: Vec<u32> = shells.keys().copied().collect();
for ch in channels {
let Some(rt) = shells.get_mut(&ch) else {
continue;
};
if rt.exit_sent {
continue;
}
if !rt.pending_stdout.is_empty() {
continue;
}
let Some(status) = rt.exited.take() else {
continue;
};
let req = match status {
ShellExitStatus::Exited(code) => ChannelRequest::ExitStatus { code },
ShellExitStatus::Signalled {
name,
core_dumped,
message,
} => ChannelRequest::ExitSignal {
name,
core_dumped,
message,
language: String::new(),
},
};
let p = conn.send_request(ch, req, false)?;
srv_send(stream, driver, &p)?;
let p = conn.send_eof(ch)?;
srv_send(stream, driver, &p)?;
let p = conn.send_close(ch)?;
srv_send(stream, driver, &p)?;
rt.exit_sent = true;
rt.session = None;
}
Ok(())
}
fn drain_subsystems(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
subsystems: &mut BTreeMap<u32, SubsystemRuntime>,
) -> Result<()> {
let channels: Vec<u32> = subsystems.keys().copied().collect();
for ch in channels {
let Some(rt) = subsystems.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_subsystem_data(stream, driver, conn, 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_subsystem_data(stream, driver, conn, 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 = conn.send_eof(ch)?;
srv_send(stream, driver, &p)?;
rt.eof_sent = true;
}
if rt.pending_close && !rt.close_sent {
if !rt.eof_sent {
let p = conn.send_eof(ch)?;
srv_send(stream, driver, &p)?;
rt.eof_sent = true;
}
let p = conn.send_close(ch)?;
srv_send(stream, driver, &p)?;
rt.close_sent = true;
}
}
}
Ok(())
}
fn emit_subsystem_data(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
channel: u32,
bytes: &[u8],
rt: &mut SubsystemRuntime,
) -> Result<()> {
let mut off = 0usize;
while off < bytes.len() {
let (payload, taken) = conn.send_data(channel, &bytes[off..])?;
if taken == 0 {
rt.pending_data.extend_from_slice(&bytes[off..]);
return Ok(());
}
srv_send(stream, driver, &payload)?;
off += taken;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn dispatch_app_packet(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
cfg: &Config,
effective: &EffectivePolicy,
user: &str,
payload: &[u8],
any_channel_opened: &mut bool,
extras: &mut ConnExtras,
shells: &mut BTreeMap<u32, ShellRuntime>,
subsystems: &mut BTreeMap<u32, SubsystemRuntime>,
envs: &mut BTreeMap<u32, SessionEnv>,
forward: &mut ForwardConn,
agent_forward: &mut AgentForwardConn,
x11_forward: &mut X11ForwardConn,
streamlocal_forward: &mut StreamlocalForwardConn,
) -> Result<()> {
let ev = conn.on_packet(payload)?;
match ev {
ChannelEvent::OpenConfirmed { channel } => {
if let Some(reply) = forward.pending_opens.remove(&channel) {
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
let _ = reply.send(Ok(cs));
} else if let Some(reply) = agent_forward.pending_opens.remove(&channel) {
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
let _ = reply.send(Ok(cs));
} else if let Some(reply) = x11_forward.pending_opens.remove(&channel) {
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
let _ = reply.send(Ok(cs));
} else if let Some(reply) = streamlocal_forward.pending_opens.remove(&channel) {
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
let _ = reply.send(Ok(cs));
}
}
ChannelEvent::OpenFailed {
channel,
reason: _reason,
description: _description,
} => {
if let Some(reply) = forward.pending_opens.remove(&channel) {
let _ = reply.send(Err(Error::Protocol(
"forwarded-tcpip: open rejected by peer",
)));
} else if let Some(reply) = agent_forward.pending_opens.remove(&channel) {
let _ = reply.send(Err(Error::Protocol("auth-agent: open rejected by peer")));
} else if let Some(reply) = x11_forward.pending_opens.remove(&channel) {
let _ = reply.send(Err(Error::Protocol("x11: open rejected by peer")));
} else if let Some(reply) = streamlocal_forward.pending_opens.remove(&channel) {
let _ = reply.send(Err(Error::Protocol(
"forwarded-streamlocal: open rejected by peer",
)));
}
}
ChannelEvent::OpenRejected {
payload: failure_payload,
reason: _reason,
} => {
srv_send(stream, driver, &failure_payload)?;
}
ChannelEvent::OpenRequest { channel, kind } => match kind {
ChannelOpen::Session => {
*any_channel_opened = true;
if let Some(cap) = effective.max_sessions
&& extras.session_channels.len() as u64 >= cap as u64
{
let p = conn.reject_open(
channel,
SSH_OPEN_RESOURCE_SHORTAGE,
"MaxSessions limit reached",
"",
)?;
srv_send(stream, driver, &p)?;
} else {
let p = conn.accept_open(channel)?;
srv_send(stream, driver, &p)?;
extras.session_channels.insert(channel);
envs.insert(channel, SessionEnv::new());
}
}
ChannelOpen::DirectTcpip {
dest_host,
dest_port,
orig_host,
orig_port,
} => {
let forwarding_ok = effective.local_forwarding_allowed();
let dest_ok = u16::try_from(dest_port)
.map(|p| effective.permit_open_allows(&dest_host, p))
.unwrap_or_else(|_| effective.permit_open.is_none());
if !forwarding_ok || !dest_ok {
let reason = if !forwarding_ok {
"local forwarding administratively prohibited"
} else {
"direct-tcpip destination not permitted by PermitOpen"
};
let p = conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
reason,
"",
)?;
srv_send(stream, driver, &p)?;
} else if let Some(handler) = cfg.direct_tcpip_handler.clone() {
let p = conn.accept_open(channel)?;
srv_send(stream, driver, &p)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let user_owned = user.to_string();
thread::spawn(move || {
let req = DirectTcpipRequest {
dest_host: &dest_host,
dest_port,
orig_host: &orig_host,
orig_port,
};
let _ = handler.handle(&user_owned, req, cs);
});
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
} else {
let p = conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"direct-tcpip not enabled",
"",
)?;
srv_send(stream, driver, &p)?;
}
}
ChannelOpen::DirectStreamlocal { socket_path } => {
let forwarding_ok = effective.local_forwarding_allowed();
if !forwarding_ok {
let p = conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"local forwarding administratively prohibited",
"",
)?;
srv_send(stream, driver, &p)?;
} else if let Some(handler) = cfg.direct_streamlocal_handler.clone() {
let p = conn.accept_open(channel)?;
srv_send(stream, driver, &p)?;
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let user_owned = user.to_string();
thread::spawn(move || {
let req = DirectStreamlocalRequest {
socket_path: &socket_path,
};
let _ = handler.handle(&user_owned, req, cs);
});
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
} else {
let p = conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"direct-streamlocal not enabled",
"",
)?;
srv_send(stream, driver, &p)?;
}
}
_ => {
let p = conn.reject_open(
channel,
SSH_OPEN_ADMINISTRATIVELY_PROHIBITED,
"channel type not supported",
"",
)?;
srv_send(stream, driver, &p)?;
}
},
ChannelEvent::Request {
channel,
request,
want_reply,
} => {
handle_channel_request(
stream,
driver,
conn,
cfg,
effective,
user,
channel,
request,
want_reply,
shells,
subsystems,
envs,
agent_forward,
x11_forward,
)?;
}
ChannelEvent::Data { channel, data } => {
if let Some(rt) = shells.get_mut(&channel)
&& let Some(sess) = rt.session.as_mut()
{
let mut off = 0usize;
let mut retries = 0u32;
while off < data.len() {
let n = sess.write(&data[off..])?;
if n == 0 {
retries += 1;
if retries > 4 {
break;
}
continue;
}
off += n;
}
}
if let Some(rt) = subsystems.get_mut(&channel) {
let _ = rt.ingress_tx.send(Some(data.clone()));
}
if let Some(adj) = conn.replenish_window(channel, data.len() as u32)? {
srv_send(stream, driver, &adj)?;
}
}
ChannelEvent::ExtendedData { channel, data, .. } => {
if let Some(adj) = conn.replenish_window(channel, data.len() as u32)? {
srv_send(stream, driver, &adj)?;
}
}
ChannelEvent::Eof { channel } => {
if let Some(rt) = shells.get_mut(&channel)
&& let Some(sess) = rt.session.as_mut()
{
let _ = sess.close_stdin();
}
if let Some(rt) = subsystems.get_mut(&channel) {
let _ = rt.ingress_tx.send(None);
}
}
ChannelEvent::Close { channel } => {
if let Some(ch) = conn.channel(channel)
&& !ch.local_closed
{
let p = conn.send_close(channel)?;
srv_send(stream, driver, &p)?;
}
shells.remove(&channel);
subsystems.remove(&channel);
envs.remove(&channel);
extras.session_channels.remove(&channel);
agent_forward.active.remove(&channel);
x11_forward.active.remove(&channel);
}
ChannelEvent::WindowAdjust { .. } => {}
ChannelEvent::GlobalRequest {
request,
want_reply,
} => {
handle_global_request(
stream,
driver,
conn,
cfg,
effective,
user,
request,
want_reply,
forward,
streamlocal_forward,
)?;
}
_ => {}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn handle_global_request(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
cfg: &Config,
effective: &EffectivePolicy,
user: &str,
request: crate::channel::GlobalRequest,
want_reply: bool,
forward: &mut ForwardConn,
streamlocal_forward: &mut StreamlocalForwardConn,
) -> Result<()> {
use crate::channel::GlobalRequest;
use crate::format::Writer;
match request {
GlobalRequest::TcpipForward {
bind_address,
bind_port,
} => {
let policy_ok = effective.remote_forwarding_allowed()
&& (bind_port > u16::MAX as u32
|| effective.permit_listen_allows(&bind_address, bind_port as u16));
let effective_bind = crate::forwarding::reverse::apply_gateway_ports(
effective.gateway_ports,
&bind_address,
);
let bound = if !policy_ok || bind_port > u16::MAX as u32 {
None
} else if let Some(handler) = cfg.tcpip_forward_handler.clone() {
let ctx = ForwardContext::new(forward.req_tx.clone());
handler
.bind(user, &effective_bind, bind_port as u16, ctx)
.ok()
} else {
None
};
if !want_reply {
if let Some(port) = bound {
forward.owned_bindings.push((effective_bind, port));
}
return Ok(());
}
match bound {
Some(port) => {
forward.owned_bindings.push((effective_bind, port));
let tail = if bind_port == 0 {
let mut w = Writer::new();
w.write_u32(port as u32);
w.into_vec()
} else {
Vec::new()
};
let p = conn.send_global_success(&tail);
srv_send(stream, driver, &p)?;
}
None => {
let p = conn.send_global_failure();
srv_send(stream, driver, &p)?;
}
}
}
GlobalRequest::CancelTcpipForward {
bind_address,
bind_port,
} => {
let effective_bind = crate::forwarding::reverse::apply_gateway_ports(
effective.gateway_ports,
&bind_address,
);
let ok = if bind_port > u16::MAX as u32 {
false
} else if let Some(handler) = cfg.tcpip_forward_handler.clone() {
let r = handler
.unbind(user, &effective_bind, bind_port as u16)
.is_ok();
if r {
forward
.owned_bindings
.retain(|(a, p)| !(a == &effective_bind && *p == bind_port as u16));
}
r
} else {
false
};
if !want_reply {
return Ok(());
}
let p = if ok {
conn.send_global_success(&[])
} else {
conn.send_global_failure()
};
srv_send(stream, driver, &p)?;
}
GlobalRequest::StreamlocalForward { socket_path } => {
let bound = if !effective.remote_forwarding_allowed() {
false
} else if let Some(handler) = cfg.streamlocal_forward_handler.clone() {
let ctx = StreamlocalForwardContext::new(streamlocal_forward.req_tx.clone());
handler.bind(user, &socket_path, ctx).is_ok()
} else {
false
};
if bound {
streamlocal_forward.owned_bindings.push(socket_path.clone());
}
if !want_reply {
return Ok(());
}
let p = if bound {
conn.send_global_success(&[])
} else {
conn.send_global_failure()
};
srv_send(stream, driver, &p)?;
}
GlobalRequest::CancelStreamlocalForward { socket_path } => {
let ok = if let Some(handler) = cfg.streamlocal_forward_handler.clone() {
let r = handler.unbind(user, &socket_path).is_ok();
if r {
streamlocal_forward
.owned_bindings
.retain(|p| p != &socket_path);
}
r
} else {
false
};
if !want_reply {
return Ok(());
}
let p = if ok {
conn.send_global_success(&[])
} else {
conn.send_global_failure()
};
srv_send(stream, driver, &p)?;
}
GlobalRequest::Keepalive | GlobalRequest::Other { .. } => {
if want_reply {
let p = conn.send_global_failure();
srv_send(stream, driver, &p)?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn handle_channel_request(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
cfg: &Config,
effective: &EffectivePolicy,
user: &str,
channel: u32,
request: ChannelRequest,
want_reply: bool,
shells: &mut BTreeMap<u32, ShellRuntime>,
subsystems: &mut BTreeMap<u32, SubsystemRuntime>,
envs: &mut BTreeMap<u32, SessionEnv>,
agent_forward: &mut AgentForwardConn,
x11_forward: &mut X11ForwardConn,
) -> Result<()> {
let empty_env = SessionEnv::new();
match request {
ChannelRequest::Exec { command } => {
let command = if let Some(forced) = effective.force_command.as_deref() {
envs.entry(channel)
.or_default()
.insert("SSH_ORIGINAL_COMMAND".to_string(), command.clone());
if forced.eq_ignore_ascii_case("internal-sftp") {
return route_internal_sftp(
stream, driver, conn, cfg, user, channel, want_reply, subsystems, envs,
);
}
forced.to_string()
} else {
command
};
if let Some(handler) = cfg.exec_stream_handler.clone()
&& handler.claims(&command)
{
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let user_owned = user.to_string();
let command_owned = command.clone();
let env_snapshot = envs.get(&channel).cloned().unwrap_or_default();
let handler_for_thread = handler;
thread::spawn(move || {
let _ = handler_for_thread.run(&user_owned, &env_snapshot, &command_owned, cs);
});
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
return Ok(());
}
let env_ref = envs.get(&channel).unwrap_or(&empty_env);
let result = cfg.command_handler.handle(user, env_ref, &command);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
drain_send(stream, driver, conn, channel, &result.stdout, None)?;
drain_send(
stream,
driver,
conn,
channel,
&result.stderr,
Some(SSH_EXTENDED_DATA_STDERR),
)?;
let p = conn.send_request(
channel,
ChannelRequest::ExitStatus {
code: result.exit_status,
},
false,
)?;
srv_send(stream, driver, &p)?;
let p = conn.send_eof(channel)?;
srv_send(stream, driver, &p)?;
let p = conn.send_close(channel)?;
srv_send(stream, driver, &p)?;
}
ChannelRequest::PtyReq {
term,
cols,
rows,
px_w,
px_h,
modes,
} => {
if cfg.shell_handler.is_some() && effective.pty_allowed() {
let rt = shells.entry(channel).or_insert_with(ShellRuntime::new);
rt.pending_pty = Some(PtySpec {
term,
cols,
rows,
px_w,
px_h,
modes,
});
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
ChannelRequest::Shell => {
if let Some(forced) = effective.force_command.as_deref() {
envs.entry(channel)
.or_default()
.insert("SSH_ORIGINAL_COMMAND".to_string(), String::new());
if forced.eq_ignore_ascii_case("internal-sftp") {
return route_internal_sftp(
stream, driver, conn, cfg, user, channel, want_reply, subsystems, envs,
);
}
let env_ref = envs.get(&channel).unwrap_or(&empty_env);
let result = cfg.command_handler.handle(user, env_ref, forced);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
drain_send(stream, driver, conn, channel, &result.stdout, None)?;
drain_send(
stream,
driver,
conn,
channel,
&result.stderr,
Some(SSH_EXTENDED_DATA_STDERR),
)?;
let p = conn.send_request(
channel,
ChannelRequest::ExitStatus {
code: result.exit_status,
},
false,
)?;
srv_send(stream, driver, &p)?;
let p = conn.send_eof(channel)?;
srv_send(stream, driver, &p)?;
let p = conn.send_close(channel)?;
srv_send(stream, driver, &p)?;
return Ok(());
}
if let Some(handler) = cfg.shell_handler.clone() {
let rt = shells.entry(channel).or_insert_with(ShellRuntime::new);
let pty = rt.pending_pty.take();
let env_ref = envs.get(&channel).unwrap_or(&empty_env);
match handler.spawn(user, env_ref, pty) {
Ok(sess) => {
rt.session = Some(sess);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
}
Err(_) => {
shells.remove(&channel);
if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
ChannelRequest::WindowChange {
cols,
rows,
px_w,
px_h,
} => {
if let Some(rt) = shells.get_mut(&channel)
&& let Some(sess) = rt.session.as_mut()
{
let _ = sess.resize(cols, rows, px_w, px_h);
}
}
ChannelRequest::Env { name, value } => {
if env_name_accepted(&name, &cfg.accept_env) {
let bag = envs.entry(channel).or_default();
let replacing = bag.get(&name).map(|v| v.len());
let new_count = if replacing.is_some() {
bag.len()
} else {
bag.len().saturating_add(1)
};
let current_bytes: usize = bag.iter().map(|(k, v)| k.len() + v.len()).sum();
let new_bytes = if let Some(old_value_len) = replacing {
current_bytes
.saturating_sub(old_value_len)
.saturating_add(value.len())
} else {
current_bytes
.saturating_add(name.len())
.saturating_add(value.len())
};
if new_count > MAX_ENV_PER_CHANNEL || new_bytes > MAX_ENV_BYTES_PER_CHANNEL {
if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
} else {
bag.insert(name, value);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
ChannelRequest::Subsystem { name } => {
if let Some(handler) = cfg.subsystem_handler.clone() {
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) =
mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let user_owned = user.to_string();
let name_owned = name.clone();
let env_snapshot = envs.get(&channel).cloned().unwrap_or_default();
thread::spawn(move || {
let _ = handler.handle(&user_owned, &env_snapshot, &name_owned, cs);
});
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
ChannelRequest::AuthAgentReq => {
if effective.agent_forwarding_allowed()
&& let Some(handler) = cfg.agent_forward_handler.clone()
{
let ctx = AgentForwardContext::new(agent_forward.req_tx.clone());
match handler.setup(user, ctx) {
Ok(handle) => {
let path_str = handle.auth_sock_path.to_string_lossy().into_owned();
envs.entry(channel)
.or_default()
.insert("SSH_AUTH_SOCK".to_string(), path_str);
agent_forward.active.insert(channel, handle);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
}
Err(_) => {
if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
ChannelRequest::X11Req {
single_connection,
auth_protocol,
auth_cookie,
screen,
} => {
if effective.x11_forwarding_allowed()
&& let Some(handler) = cfg.x11_forward_handler.clone()
{
let ctx = X11ForwardContext::new(x11_forward.req_tx.clone());
match handler.setup(
user,
single_connection,
&auth_protocol,
&auth_cookie,
screen,
ctx,
) {
Ok(handle) => {
envs.entry(channel)
.or_default()
.insert("DISPLAY".to_string(), handle.display_env.clone());
x11_forward.active.insert(channel, handle);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
}
Err(_) => {
if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
_ => {
if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn route_internal_sftp(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
cfg: &Config,
user: &str,
channel: u32,
want_reply: bool,
subsystems: &mut BTreeMap<u32, SubsystemRuntime>,
envs: &mut BTreeMap<u32, SessionEnv>,
) -> Result<()> {
if let Some(handler) = cfg.subsystem_handler.clone() {
let (ingress_tx, ingress_rx) = mpsc::channel::<Option<Vec<u8>>>();
let (egress_tx, egress_rx) = mpsc::sync_channel::<ChannelEgress>(SUBSYSTEM_EGRESS_BACKLOG);
let cs = ChannelStream::new(ingress_rx, egress_tx);
let user_owned = user.to_string();
let env_snapshot = envs.get(&channel).cloned().unwrap_or_default();
thread::spawn(move || {
let _ = handler.handle(&user_owned, &env_snapshot, "sftp", cs);
});
subsystems.insert(
channel,
SubsystemRuntime {
ingress_tx,
egress_rx,
pending_data: Vec::new(),
pending_eof: false,
pending_close: false,
eof_sent: false,
close_sent: false,
},
);
if want_reply {
let p = conn.send_request_success(channel)?;
srv_send(stream, driver, &p)?;
}
} else if want_reply {
let p = conn.send_request_failure(channel)?;
srv_send(stream, driver, &p)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn drain_send(
stream: &mut TcpStream,
driver: &mut ServerDriver,
conn: &mut ConnectionState,
channel: u32,
mut data: &[u8],
extended: Option<u32>,
) -> Result<()> {
let mut iter = 0usize;
while !data.is_empty() {
iter += 1;
if iter > MAX_DRAIN_STEPS {
return Err(Error::Protocol("drain_send did not converge"));
}
let (payload, taken) = if let Some(code) = extended {
conn.send_extended_data(channel, code, data)?
} else {
conn.send_data(channel, data)?
};
if taken > 0 {
srv_send(stream, driver, &payload)?;
data = &data[taken..];
continue;
}
let pkt = srv_read(stream, driver, None)?;
let ev = conn.on_packet(&pkt)?;
match ev {
ChannelEvent::WindowAdjust { channel: c, .. } if c == channel => continue,
ChannelEvent::Close { channel: c } if c == channel => {
return Err(Error::BadChannelState);
}
_ => continue,
}
}
Ok(())
}
pub(crate) fn pick_host_key<'a>(
keys: &'a [Box<dyn HostKey + Send + Sync>],
name: &str,
) -> Option<&'a (dyn HostKey + Send + Sync)> {
for k in keys {
if k.algorithm() == name {
return Some(k.as_ref());
}
}
for k in keys {
let a = k.algorithm();
if (a == "ssh-rsa" || a == "rsa-sha2-256" || a == "rsa-sha2-512")
&& (name == "ssh-rsa" || name == "rsa-sha2-256" || name == "rsa-sha2-512")
{
return Some(k.as_ref());
}
}
None
}
pub(crate) fn server_ext_info() -> ExtInfo {
ExtInfo::new().with_server_sig_algs(
"ssh-ed25519-cert-v01@openssh.com,\
ecdsa-sha2-nistp256-cert-v01@openssh.com,ecdsa-sha2-nistp384-cert-v01@openssh.com,\
ecdsa-sha2-nistp521-cert-v01@openssh.com,\
rsa-sha2-512-cert-v01@openssh.com,rsa-sha2-256-cert-v01@openssh.com,\
ssh-ed25519,ecdsa-sha2-nistp256,ecdsa-sha2-nistp384,ecdsa-sha2-nistp521,\
rsa-sha2-512,rsa-sha2-256",
)
}
fn unix_now() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
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_server_kexinit<R: RngCore>(rng: &mut R, cfg: &Config) -> KexInit {
let host_keys = &cfg.host_keys;
let mut have: Vec<&'static str> = Vec::new();
for n in crate::cert::CERT_KEY_NAMES {
if host_keys.iter().any(|k| k.algorithm() == *n) {
have.push(*n);
}
}
for n in defaults::HOST_KEY {
if host_keys.iter().any(|k| k.algorithm() == *n) {
have.push(*n);
continue;
}
if (*n == "rsa-sha2-256" || *n == "rsa-sha2-512")
&& host_keys.iter().any(|k| {
let a = k.algorithm();
a == "ssh-rsa" || a == "rsa-sha2-256" || a == "rsa-sha2-512"
})
{
have.push(*n);
}
}
let host_key: Vec<String> = match &cfg.host_key_algorithms {
Some(pref) => pref
.iter()
.filter(|p| have.contains(&p.as_str()))
.cloned()
.collect(),
None => have.iter().map(|s| s.to_string()).collect(),
};
let ed25519_excluded = cfg
.host_key_algorithms
.as_ref()
.is_some_and(|pref| !pref.iter().any(|p| p == "ssh-ed25519"));
let host_key = if host_key.is_empty() && !ed25519_excluded {
alloc::vec!["ssh-ed25519".to_string()]
} else {
host_key
};
let default_kex: Vec<&str> = defaults::KEX
.iter()
.copied()
.filter(|n| !is_strict_kex_marker(n))
.collect();
let mut kex = owned_or_default(&cfg.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(&cfg.ciphers, defaults::CIPHERS);
let macs = owned_or_default(&cfg.macs, defaults::MACS);
let comp: Vec<String> = defaults::COMP
.iter()
.filter(|name| {
if cfg.compression == Some(crate::config::Compression::No) {
!name.contains("zlib")
} else {
true
}
})
.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)
}
fn arm_preauth_deadline(stream: &TcpStream, deadline: Option<Instant>) -> Result<()> {
if let Some(deadline) = deadline {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return Err(Error::Io(std::io::Error::new(
ErrorKind::TimedOut,
"pre-auth inactivity timeout (LoginGraceTime)",
)));
}
stream
.set_read_timeout(Some(remaining))
.map_err(Error::Io)?;
}
Ok(())
}
fn srv_pump_out(stream: &mut TcpStream, driver: &mut ServerDriver) -> Result<()> {
while let Some(frame) = driver.poll_transmit() {
stream.write_all(&frame)?;
}
Ok(())
}
fn srv_send(stream: &mut TcpStream, driver: &mut ServerDriver, payload: &[u8]) -> Result<()> {
driver.enqueue_payload(payload)?;
srv_pump_out(stream, driver)
}
fn srv_read_into(
stream: &mut TcpStream,
driver: &mut ServerDriver,
deadline: Option<Instant>,
) -> Result<()> {
arm_preauth_deadline(stream, deadline)?;
let mut tmp = [0u8; 16 * 1024];
let n = stream.read(&mut tmp)?;
if n == 0 {
return Err(Error::Protocol("connection closed"));
}
driver.handle_input(&tmp[..n], Instant::now())?;
Ok(())
}
fn srv_read(
stream: &mut TcpStream,
driver: &mut ServerDriver,
deadline: Option<Instant>,
) -> Result<Vec<u8>> {
loop {
driver.handle_timeout(Instant::now())?;
srv_pump_out(stream, driver)?;
while let Some(ev) = driver.poll_event() {
if let Event::AppData(payload) = ev {
srv_pump_out(stream, driver)?;
return Ok(payload);
}
}
srv_read_into(stream, driver, deadline)?;
}
}
fn srv_read_maybe_timeout(
stream: &mut TcpStream,
driver: &mut ServerDriver,
) -> Result<Option<Vec<u8>>> {
match srv_read(stream, driver, None) {
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 srv_send_disconnect(
stream: &mut TcpStream,
driver: &mut ServerDriver,
reason: u32,
description: &str,
) -> Result<()> {
let mut w = Writer::new();
w.write_u8(1);
w.write_u32(reason);
w.write_string(description.as_bytes());
w.write_string(b"");
srv_send(stream, driver, &w.into_vec())
}
fn srv_drive_handshake(
stream: &mut TcpStream,
driver: &mut ServerDriver,
deadline: Option<Instant>,
) -> Result<()> {
loop {
srv_pump_out(stream, driver)?;
while let Some(ev) = driver.poll_event() {
if matches!(ev, Event::HandshakeComplete) {
srv_pump_out(stream, driver)?;
return Ok(());
}
}
srv_read_into(stream, driver, deadline)?;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::auth::{AuthAttempt, AuthDecision, Authenticator};
use crate::client::{Client, Config as ClientConfig, HostKeyPolicy};
use crate::hostkey::Ed25519HostKey;
use crate::transport::kex::KexAlgorithms;
use crate::transport::{KexRunner, PacketCodec, Role};
use purecrypto::rng::OsRng;
use std::sync::Mutex;
use std::time::Duration;
struct OneKeyAuth {
allowed_user: String,
allowed_blob: Vec<u8>,
}
impl Authenticator for OneKeyAuth {
fn evaluate(&mut self, attempt: AuthAttempt) -> AuthDecision {
match attempt {
AuthAttempt::PublicKey {
user,
public_blob,
probe_only,
verified,
..
} => {
if user != self.allowed_user {
return AuthDecision::Reject;
}
if public_blob != self.allowed_blob {
return AuthDecision::Reject;
}
if probe_only {
return AuthDecision::Accept;
}
if !verified {
return AuthDecision::Reject;
}
AuthDecision::Accept
}
_ => AuthDecision::Reject,
}
}
}
#[test]
fn resolve_auth_methods_subtracts_disabled() {
let base: &[&'static str] = &["publickey", "password", "keyboard-interactive"];
let opts = crate::config::ServerOptions::default();
assert_eq!(
resolve_auth_methods(base, &opts),
vec!["publickey", "password", "keyboard-interactive"]
);
let opts = crate::config::ServerOptions {
password_authentication: Some(false),
..Default::default()
};
assert_eq!(
resolve_auth_methods(base, &opts),
vec!["publickey", "keyboard-interactive"]
);
let opts = crate::config::ServerOptions {
kbd_interactive_authentication: Some(false),
..Default::default()
};
assert_eq!(
resolve_auth_methods(base, &opts),
vec!["publickey", "password"]
);
let opts = crate::config::ServerOptions {
pubkey_authentication: Some(false),
..Default::default()
};
assert_eq!(
resolve_auth_methods(base, &opts),
vec!["password", "keyboard-interactive"]
);
}
struct StaticHandler {
out: Vec<u8>,
}
impl CommandHandler for StaticHandler {
fn handle(&self, _user: &str, _env: &SessionEnv, _command: &str) -> ExecResult {
ExecResult {
stdout: self.out.clone(),
stderr: Vec::new(),
exit_status: 0,
}
}
}
fn fresh_seed() -> [u8; 32] {
let mut s = [0u8; 32];
OsRng.fill_bytes(&mut s);
s
}
fn kexinit_test_config(host_keys: Vec<Box<dyn HostKey + Send + Sync>>) -> Config {
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(|| -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: String::new(),
allowed_blob: Vec::new(),
})
});
Config::new(
host_keys,
factory,
vec!["publickey"],
Arc::new(StaticHandler { out: Vec::new() }),
)
}
struct MemoryShellState {
stdout: Vec<u8>,
stdin: Vec<u8>,
closed_stdin: bool,
pty: Option<PtySpec>,
resizes: Vec<(u32, u32, u32, u32)>,
exit_on_stdin_close: Option<ShellExitStatus>,
exit_now: Option<ShellExitStatus>,
user: String,
}
#[derive(Clone)]
struct MemoryShell {
inner: Arc<Mutex<MemoryShellState>>,
}
impl MemoryShell {
fn new() -> Self {
Self {
inner: Arc::new(Mutex::new(MemoryShellState {
stdout: Vec::new(),
stdin: Vec::new(),
closed_stdin: false,
pty: None,
resizes: Vec::new(),
exit_on_stdin_close: None,
exit_now: None,
user: String::new(),
})),
}
}
fn push_stdout(&self, bytes: &[u8]) {
self.inner.lock().unwrap().stdout.extend_from_slice(bytes);
}
fn arm_exit_on_stdin_close(&self, status: ShellExitStatus) {
self.inner.lock().unwrap().exit_on_stdin_close = Some(status);
}
}
struct MemoryShellHandler {
shell: MemoryShell,
}
impl ShellHandler for MemoryShellHandler {
fn spawn(
&self,
user: &str,
_env: &SessionEnv,
pty: Option<PtySpec>,
) -> Result<Box<dyn ShellSession>> {
{
let mut st = self.shell.inner.lock().unwrap();
st.pty = pty;
st.user = user.to_string();
}
Ok(Box::new(MemoryShellSession {
inner: self.shell.inner.clone(),
}))
}
}
struct MemoryShellSession {
inner: Arc<Mutex<MemoryShellState>>,
}
impl ShellSession for MemoryShellSession {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
let mut st = self.inner.lock().unwrap();
if st.stdout.is_empty() {
return Ok(0);
}
let n = core::cmp::min(buf.len(), st.stdout.len());
buf[..n].copy_from_slice(&st.stdout[..n]);
st.stdout.drain(..n);
Ok(n)
}
fn write(&mut self, data: &[u8]) -> Result<usize> {
self.inner.lock().unwrap().stdin.extend_from_slice(data);
Ok(data.len())
}
fn close_stdin(&mut self) -> Result<()> {
self.inner.lock().unwrap().closed_stdin = true;
Ok(())
}
fn resize(&mut self, cols: u32, rows: u32, px_w: u32, px_h: u32) -> Result<()> {
self.inner
.lock()
.unwrap()
.resizes
.push((cols, rows, px_w, px_h));
Ok(())
}
fn try_exit(&mut self) -> Option<ShellExitStatus> {
let mut st = self.inner.lock().unwrap();
if let Some(s) = st.exit_now.take() {
return Some(s);
}
if st.closed_stdin
&& st.stdout.is_empty()
&& let Some(s) = st.exit_on_stdin_close.take()
{
return Some(s);
}
None
}
}
#[test]
fn loopback_shell_with_pty_and_stdin() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "shell-test-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let memshell = MemoryShell::new();
memshell.push_stdout(b"hello from memshell\n");
memshell.arm_exit_on_stdin_close(ShellExitStatus::Exited(0));
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused-exec\n".to_vec(),
}),
)
.with_shell(Arc::new(MemoryShellHandler {
shell: memshell.clone(),
}));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let out = client
.shell_with_stdin("xterm-256color", 132, 43, b"echo back\n")
.expect("shell_with_stdin");
assert_eq!(out.stdout, b"hello from memshell\n");
assert_eq!(out.exit_status, Some(0));
assert_eq!(out.exit_signal, None);
let st = memshell.inner.lock().unwrap();
let pty = st.pty.as_ref().expect("pty-req captured");
assert_eq!(pty.term, "xterm-256color");
assert_eq!(pty.cols, 132);
assert_eq!(pty.rows, 43);
assert_eq!(st.stdin, b"echo back\n");
assert!(st.closed_stdin, "EOF should reach the backend");
assert_eq!(st.user, user);
drop(st);
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_exec_roundtrip() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "ssh-test-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"loopback-test\n".to_vec(),
}),
);
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let out = client.exec("ignored").expect("exec");
assert_eq!(out.stdout, b"loopback-test\n");
assert_eq!(out.exit_status, Some(0));
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_ping_pong_roundtrip() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "ssh-test-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"after-ping\n".to_vec(),
}),
);
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
client
.send_transport_ping(b"obscure-keystroke-chaff")
.expect("send PING");
let out = client.exec("ignored").expect("exec after ping");
assert_eq!(out.stdout, b"after-ping\n");
assert_eq!(out.exit_status, Some(0));
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_match_user_forbids_publickey() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "blocked-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let policy = crate::config::SshServerConfig::parse(
"Match User blocked-user\n PubkeyAuthentication no\n",
)
.expect("parse policy");
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"never\n".to_vec(),
}),
)
.with_policy(Arc::new(policy));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
let res = client.authenticate_publickey(&user, client_hk);
assert!(
res.is_err(),
"publickey must be rejected for a Match User PubkeyAuthentication no block"
);
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_forces_rekeys_with_tiny_policy() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "ssh-test-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let payload: Vec<u8> = (0..16_384).map(|i| (i & 0xff) as u8).collect();
let mut cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: payload.clone(),
}),
);
cfg.rekey_policy = RekeyPolicy {
max_bytes: 1024,
max_duration: Duration::from_secs(60 * 60),
max_seq: 1u32 << 31,
};
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let session_id_before = client.session_id().to_vec();
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let out = client.exec("ignored").expect("exec");
assert_eq!(out.stdout, payload);
assert_eq!(out.exit_status, Some(0));
assert_eq!(client.session_id(), session_id_before.as_slice());
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn server_kexinit_negotiation_uses_role_server() {
let mut rng = OsRng;
let host_keys: Vec<Box<dyn HostKey + Send + Sync>> =
vec![Box::new(Ed25519HostKey::from_seed(fresh_seed()))];
let cfg = kexinit_test_config(host_keys);
let advert = build_server_kexinit(&mut rng, &cfg);
let mut runner = KexRunner::new(Role::Server, advert.clone());
let mut cookie = [0u8; 16];
rng.fill_bytes(&mut cookie);
let client_init = {
let algs = KexAlgorithms {
kex: &["curve25519-sha256"],
server_host_key: &["ssh-ed25519"],
ciphers_c2s: &["chacha20-poly1305@openssh.com"],
ciphers_s2c: &["chacha20-poly1305@openssh.com"],
macs_c2s: &["hmac-sha2-256"],
macs_s2c: &["hmac-sha2-256"],
comp_c2s: &["none"],
comp_s2c: &["none"],
lang_c2s: &[],
lang_s2c: &[],
};
KexInit::from_algorithms(&algs, cookie)
};
let _ = runner.start(&mut rng).expect("server start");
let mut codec = PacketCodec::new();
let adv = runner
.on_packet(
&mut rng,
&mut codec,
&client_init.encode(),
None,
None,
b"SSH-2.0-test-client",
b"SSH-2.0-test-server",
)
.expect("server processes client kexinit");
assert!(!adv.completed);
let neg = runner.negotiated().expect("negotiated");
assert_eq!(neg.kex, "curve25519-sha256");
assert_eq!(neg.host_key, "ssh-ed25519");
}
#[test]
fn server_cipher_override_replaces_and_keeps_kex_markers() {
let mut rng = OsRng;
let host_keys: Vec<Box<dyn HostKey + Send + Sync>> =
vec![Box::new(Ed25519HostKey::from_seed(fresh_seed()))];
let cfg = kexinit_test_config(host_keys).with_algorithms(
Some(vec!["aes256-ctr".to_string()]),
None,
None,
None,
);
let advert = build_server_kexinit(&mut rng, &cfg);
assert_eq!(advert.ciphers_c2s, vec!["aes256-ctr".to_string()]);
let markers = advert
.kex
.iter()
.filter(|k| is_strict_kex_marker(k))
.count();
assert_eq!(markers, 2);
}
#[test]
fn server_hostkey_override_intersects_with_loaded_keys() {
let mut rng = OsRng;
let host_keys: Vec<Box<dyn HostKey + Send + Sync>> =
vec![Box::new(Ed25519HostKey::from_seed(fresh_seed()))];
let cfg = kexinit_test_config(host_keys).with_algorithms(
None,
None,
None,
Some(vec!["rsa-sha2-512".to_string(), "ssh-ed25519".to_string()]),
);
let advert = build_server_kexinit(&mut rng, &cfg);
assert_eq!(advert.server_host_key, vec!["ssh-ed25519".to_string()]);
}
#[test]
fn server_hostkey_override_excluding_ed25519_is_honoured() {
let mut rng = OsRng;
let host_keys: Vec<Box<dyn HostKey + Send + Sync>> =
vec![Box::new(Ed25519HostKey::from_seed(fresh_seed()))];
let cfg = kexinit_test_config(host_keys).with_algorithms(
None,
None,
None,
Some(vec!["rsa-sha2-512".to_string()]),
);
let advert = build_server_kexinit(&mut rng, &cfg);
assert!(
advert.server_host_key.is_empty(),
"explicit ed25519 exclusion must not be overridden by the fallback net"
);
}
#[test]
fn server_no_override_falls_back_to_ed25519_when_no_keys() {
let mut rng = OsRng;
let cfg = kexinit_test_config(Vec::new());
let advert = build_server_kexinit(&mut rng, &cfg);
assert_eq!(advert.server_host_key, vec!["ssh-ed25519".to_string()]);
}
struct EchoUpperSubsystem {
captured_name: Arc<Mutex<Option<String>>>,
captured_user: Arc<Mutex<Option<String>>>,
}
impl SubsystemHandler for EchoUpperSubsystem {
fn handle(
&self,
user: &str,
_env: &SessionEnv,
name: &str,
mut stream: ChannelStream,
) -> Result<()> {
*self.captured_name.lock().unwrap() = Some(name.to_string());
*self.captured_user.lock().unwrap() = Some(user.to_string());
let mut acc = Vec::new();
let mut tmp = [0u8; 256];
loop {
match std::io::Read::read(&mut stream, &mut tmp) {
Ok(0) => break, Ok(n) => acc.extend_from_slice(&tmp[..n]),
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
std::thread::sleep(Duration::from_millis(5));
continue;
}
Err(_) => break,
}
}
for b in acc.iter_mut() {
b.make_ascii_uppercase();
}
std::io::Write::write_all(&mut stream, &acc).ok();
Ok(())
}
}
#[test]
fn loopback_subsystem_roundtrip() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "subsys-test-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let captured_name = Arc::new(Mutex::new(None));
let captured_user = Arc::new(Mutex::new(None));
let sub = EchoUpperSubsystem {
captured_name: captured_name.clone(),
captured_user: captured_user.clone(),
};
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused-exec\n".to_vec(),
}),
)
.with_subsystem(Arc::new(sub));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let body = b"hello, subsystem world".to_vec();
let resp = client
.subsystem_once("echo", &body)
.expect("subsystem_once");
assert_eq!(resp, b"HELLO, SUBSYSTEM WORLD".to_vec());
assert_eq!(
captured_name.lock().unwrap().as_deref(),
Some("echo"),
"subsystem name reached the handler",
);
assert_eq!(
captured_user.lock().unwrap().as_deref(),
Some(user.as_str()),
"authenticated user reached the handler",
);
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_subsystem_unconfigured_refused() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "subsys-reject-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused-exec\n".to_vec(),
}),
);
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let err = client
.subsystem_once("sftp", b"")
.expect_err("expected rejection");
match err {
Error::Protocol(_) => {}
other => panic!("expected Error::Protocol, got {:?}", other),
}
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[cfg(unix)]
struct SftpSubsystem {
cwd: std::path::PathBuf,
root: std::path::PathBuf,
}
#[cfg(unix)]
impl SubsystemHandler for SftpSubsystem {
fn handle(
&self,
_user: &str,
_env: &SessionEnv,
name: &str,
stream: ChannelStream,
) -> Result<()> {
if name != "sftp" {
return Ok(());
}
let opts = crate::sftp::SftpServerOptions::new(self.cwd.clone())
.with_root(self.root.clone())
.hide_jail_in_realpath(false);
let mut sess = crate::sftp::SftpServerSession::new(opts);
let _ = sess.run(stream);
Ok(())
}
}
#[cfg(unix)]
struct SftpTempDir(std::path::PathBuf);
#[cfg(unix)]
impl SftpTempDir {
fn new(tag: &str) -> Self {
let dir = std::env::temp_dir().join(format!(
"puressh-server-sftp-{}-{}-{}",
tag,
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos(),
));
std::fs::create_dir_all(&dir).unwrap();
Self(dir)
}
fn path(&self) -> &std::path::Path {
&self.0
}
}
#[cfg(unix)]
impl Drop for SftpTempDir {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
#[cfg(unix)]
#[test]
fn loopback_sftp_client_roundtrip() {
let tmp = SftpTempDir::new("roundtrip");
let root = tmp.path().to_path_buf();
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "sftp-test-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let sub = SftpSubsystem {
cwd: root.clone(),
root: root.clone(),
};
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused-exec\n".to_vec(),
}),
)
.with_subsystem(Arc::new(sub));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("local_addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
{
#[allow(deprecated)]
let mut sftp = client.sftp().expect("sftp handshake");
assert!(sftp.server_version() >= 3);
let cwd = sftp.realpath(b".").expect("realpath .");
assert_eq!(cwd.as_slice(), root.as_os_str().as_encoded_bytes());
let target = root.join("hello.txt");
let body = b"hello from sftp\n".to_vec();
let handle = sftp
.open(
target.as_os_str().as_encoded_bytes(),
crate::sftp::FXF_WRITE | crate::sftp::FXF_CREAT | crate::sftp::FXF_TRUNC,
crate::sftp::Attrs::default(),
)
.expect("open for write");
sftp.write(&handle, 0, &body).expect("write");
sftp.close(&handle).expect("close write handle");
let handle = sftp
.open(
target.as_os_str().as_encoded_bytes(),
crate::sftp::FXF_READ,
crate::sftp::Attrs::default(),
)
.expect("open for read");
let got = sftp.read(&handle, 0, 1024).expect("read");
assert_eq!(got, body);
sftp.close(&handle).expect("close read handle");
let dh = sftp
.opendir(root.as_os_str().as_encoded_bytes())
.expect("opendir");
let mut all_names = Vec::<Vec<u8>>::new();
while let Some(batch) = sftp.readdir(&dh).expect("readdir") {
for e in batch {
all_names.push(e.filename);
}
}
sftp.close(&dh).expect("close dir");
assert!(
all_names.iter().any(|n| n == b"hello.txt"),
"readdir saw the new file: {:?}",
all_names
.iter()
.map(|n| String::from_utf8_lossy(n).into_owned())
.collect::<Vec<_>>(),
);
sftp.remove(target.as_os_str().as_encoded_bytes())
.expect("remove");
let err = sftp
.stat(target.as_os_str().as_encoded_bytes())
.expect_err("stat after remove");
match err {
crate::sftp::SftpError::Status {
code: crate::sftp::FxpStatus::NoSuchFile,
..
} => {}
other => panic!("expected NoSuchFile, got {:?}", other),
}
}
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_direct_tcpip_round_trip() {
use std::io::{Read as _, Write as _};
use std::net::TcpListener;
let echo_listener = TcpListener::bind("127.0.0.1:0").expect("bind echo");
let echo_addr = echo_listener.local_addr().expect("echo addr");
let echo_thread = thread::spawn(move || {
if let Ok((mut s, _)) = echo_listener.accept() {
let mut buf = [0u8; 1024];
loop {
match s.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => {
if s.write_all(&buf[..n]).is_err() {
break;
}
}
}
}
}
});
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "direct-tcpip-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused-exec\n".to_vec(),
}),
)
.with_direct_tcpip(Arc::new(
crate::forwarding::direct::DefaultDirectTcpipHandler::permit_all(),
));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind ssh");
let ssh_addr = server.local_addr().expect("ssh addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
ssh_addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
{
#[allow(deprecated)]
let mut s = client
.open_direct_tcpip(
&echo_addr.ip().to_string(),
echo_addr.port(),
"127.0.0.1",
0,
)
.expect("open direct-tcpip");
s.write_all(b"ping").expect("write");
let mut got = [0u8; 4];
s.read_exact(&mut got).expect("read echo");
assert_eq!(&got, b"ping");
}
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
let _ = echo_thread.join();
}
#[test]
#[allow(deprecated)]
fn loopback_direct_tcpip_unconfigured_refused() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "direct-tcpip-reject-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused\n".to_vec(),
}),
);
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
match client.open_direct_tcpip("127.0.0.1", 1, "127.0.0.1", 0) {
Ok(_) => panic!("expected direct-tcpip open to be refused"),
Err(Error::Protocol(_)) => {}
Err(other) => panic!("expected Error::Protocol, got {:?}", other),
}
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_tcpip_forward_round_trip() {
use crate::client::{ClientHandlers, ForwardedTcpipOrigin};
use std::io::{Read as _, Write as _};
use std::net::TcpStream;
use std::sync::atomic::Ordering;
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "tcpip-forward-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused-exec\n".to_vec(),
}),
)
.with_tcpip_forward(Arc::new(
crate::forwarding::reverse::DefaultTcpipForwardHandler::permit_all_interfaces(),
));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind ssh");
let ssh_addr = server.local_addr().expect("ssh addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
ssh_addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let bound_port = client
.request_tcpip_forward("127.0.0.1", 0)
.expect("request_tcpip_forward");
assert!(bound_port > 0);
let origin_seen: Arc<Mutex<Option<ForwardedTcpipOrigin>>> = Arc::new(Mutex::new(None));
let origin_clone = origin_seen.clone();
let cb: Arc<crate::client::ForwardedTcpipCallback> =
Arc::new(move |origin: ForwardedTcpipOrigin, mut s: ChannelStream| {
*origin_clone.lock().unwrap() = Some(origin);
let mut acc = Vec::new();
let mut tmp = [0u8; 256];
loop {
match Read::read(&mut s, &mut tmp) {
Ok(0) => break,
Ok(n) => acc.extend_from_slice(&tmp[..n]),
Err(_) => break,
}
}
for b in acc.iter_mut() {
b.make_ascii_uppercase();
}
let _ = Write::write_all(&mut s, &acc);
});
let handlers = ClientHandlers::new().with_forwarded_tcpip(cb);
let stop = handlers.stop.clone();
let serve_thread = thread::spawn(move || -> std::result::Result<Client, Error> {
client.serve(handlers)?;
Ok(client)
});
let mut s = TcpStream::connect(("127.0.0.1", bound_port)).expect("dial forwarded port");
s.write_all(b"hello").expect("write");
s.shutdown(std::net::Shutdown::Write)
.expect("shutdown write");
let mut got = Vec::new();
s.read_to_end(&mut got).expect("read echo");
assert_eq!(got, b"HELLO");
drop(s);
thread::sleep(Duration::from_millis(100));
stop.store(true, Ordering::SeqCst);
let start = std::time::Instant::now();
while !serve_thread.is_finished() {
if start.elapsed() > Duration::from_secs(10) {
panic!("serve loop did not stop in time");
}
thread::sleep(Duration::from_millis(20));
}
let client_back = serve_thread
.join()
.expect("serve join")
.expect("serve result");
let captured = origin_seen.lock().unwrap().clone().expect("origin latched");
assert_eq!(captured.bound_address, "127.0.0.1");
assert_eq!(captured.bound_port, bound_port);
assert!(captured.orig_port > 0);
drop(client_back);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_tcpip_forward_unconfigured_refused() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "tcpip-forward-reject-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused\n".to_vec(),
}),
);
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind ssh");
let ssh_addr = server.local_addr().expect("ssh addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
ssh_addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
match client.request_tcpip_forward("127.0.0.1", 0) {
Ok(_) => panic!("expected tcpip-forward to be refused"),
Err(Error::Protocol(_)) => {}
Err(other) => panic!("expected Error::Protocol, got {:?}", other),
}
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn loopback_serve_context_direct_tcpip_round_trip() {
use crate::client::ClientHandlers;
use std::io::{Read as _, Write as _};
use std::net::TcpListener;
use std::sync::atomic::Ordering;
let echo_listener = TcpListener::bind("127.0.0.1:0").expect("bind echo");
let echo_addr = echo_listener.local_addr().expect("echo addr");
let echo_thread = thread::spawn(move || {
if let Ok((mut s, _)) = echo_listener.accept() {
let mut buf = [0u8; 1024];
loop {
match s.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => {
if s.write_all(&buf[..n]).is_err() {
break;
}
}
}
}
}
});
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "serve-ctx-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler {
out: b"unused\n".to_vec(),
}),
)
.with_direct_tcpip(Arc::new(
crate::forwarding::direct::DefaultDirectTcpipHandler::permit_all(),
));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind ssh");
let ssh_addr = server.local_addr().expect("ssh addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
ssh_addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
let client_hk: Box<dyn HostKey + Send> = Box::new(Ed25519HostKey::from_seed(client_seed));
client
.authenticate_publickey(&user, client_hk)
.expect("authenticate");
let (handlers, ctx) = ClientHandlers::new().with_serve_context();
let stop = handlers.stop.clone();
let serve_thread = thread::spawn(move || -> std::result::Result<Client, Error> {
client.serve(handlers)?;
Ok(client)
});
let mut s = ctx
.open_direct_tcpip(
&echo_addr.ip().to_string(),
echo_addr.port(),
"127.0.0.1",
0,
)
.expect("open_direct_tcpip via ServeContext");
s.write_all(b"ping").expect("write ping");
let mut got = [0u8; 4];
s.read_exact(&mut got).expect("read echo");
assert_eq!(&got, b"ping");
drop(s);
thread::sleep(Duration::from_millis(100));
drop(ctx);
stop.store(true, Ordering::SeqCst);
let start = std::time::Instant::now();
while !serve_thread.is_finished() {
if start.elapsed() > Duration::from_secs(10) {
panic!("serve loop did not stop in time");
}
thread::sleep(Duration::from_millis(20));
}
let client_back = serve_thread
.join()
.expect("serve join")
.expect("serve result");
drop(client_back);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
let _ = echo_thread.join();
}
#[test]
fn env_filter_blocks_ld_preload_even_when_glob_matches() {
let allow = vec!["*".to_string()];
assert!(
!env_name_accepted("LD_PRELOAD", &allow),
"LD_PRELOAD must NEVER be accepted, even with `AcceptEnv *`"
);
assert!(!env_name_accepted("LD_LIBRARY_PATH", &allow));
assert!(!env_name_accepted("LD_AUDIT", &allow));
assert!(!env_name_accepted("LD_BIND_NOT", &allow));
assert!(!env_name_accepted("DYLD_INSERT_LIBRARIES", &allow));
assert!(!env_name_accepted("DYLD_LIBRARY_PATH", &allow));
assert!(!env_name_accepted("BASH_ENV", &allow));
assert!(!env_name_accepted("ENV", &allow));
assert!(!env_name_accepted("IFS", &allow));
assert!(!env_name_accepted("PATH", &allow));
assert!(!env_name_accepted("SHELL", &allow));
assert!(!env_name_accepted("HOME", &allow));
assert!(!env_name_accepted("USER", &allow));
assert!(!env_name_accepted("LOGNAME", &allow));
assert!(env_name_accepted("LANG", &allow));
assert!(env_name_accepted("LC_ALL", &allow));
assert!(env_name_accepted("TERM", &allow));
}
#[test]
fn env_filter_explicit_listing_of_blocked_name_still_blocks() {
let allow = vec!["LD_PRELOAD".to_string()];
assert!(!env_name_accepted("LD_PRELOAD", &allow));
}
#[test]
fn env_filter_empty_allowlist_drops_everything() {
let allow: Vec<String> = Vec::new();
assert!(!env_name_accepted("LANG", &allow));
assert!(!env_name_accepted("TERM", &allow));
assert!(!env_name_accepted("FOO", &allow));
assert!(!env_name_accepted("LD_PRELOAD", &allow));
}
#[test]
fn env_filter_glob_matching() {
let allow = vec!["LC_*".to_string(), "LANG".to_string(), "X???".to_string()];
assert!(env_name_accepted("LANG", &allow));
assert!(env_name_accepted("LC_ALL", &allow));
assert!(env_name_accepted("LC_TIME", &allow));
assert!(env_name_accepted("LC_CTYPE", &allow));
assert!(env_name_accepted("XABC", &allow));
assert!(!env_name_accepted("LANGUAGE", &allow)); assert!(!env_name_accepted("XABCD", &allow)); assert!(!env_name_accepted("XAB", &allow));
assert!(!env_name_accepted("MC_ALL", &allow));
assert!(!env_name_accepted("", &allow));
}
fn policy_cfg(src: &str) -> Config {
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(fresh_seed()));
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(|| -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: "x".into(),
allowed_blob: Vec::new(),
})
});
let policy = crate::config::SshServerConfig::parse(src).expect("parse policy");
Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler { out: Vec::new() }),
)
.with_policy(Arc::new(policy))
}
#[test]
fn pubkey_auth_no_locks_out_method_set() {
let cfg = policy_cfg("PubkeyAuthentication no\n");
let pre = resolve_preauth_policy(&cfg, Some("203.0.113.5"), None, None);
assert!(
pre.methods.is_empty(),
"expected lockout, got {:?}",
pre.methods
);
let cfg2 = policy_cfg("PubkeyAuthentication yes\n");
let pre2 = resolve_preauth_policy(&cfg2, Some("203.0.113.5"), None, None);
assert_eq!(pre2.methods, vec!["publickey"]);
}
#[test]
fn match_address_gates_pubkey_via_block() {
let cfg = policy_cfg("Match Address 192.0.2.0/24\n PubkeyAuthentication no\n");
let inside = resolve_preauth_policy(&cfg, Some("192.0.2.9"), None, None);
assert!(inside.methods.is_empty());
let outside = resolve_preauth_policy(&cfg, Some("198.51.100.1"), None, None);
assert_eq!(outside.methods, vec!["publickey"]);
}
#[test]
fn match_localport_in_effective_policy() {
let cfg =
policy_cfg("Match LocalPort 2222\n X11Forwarding no\n AllowAgentForwarding no\n");
let eff = resolve_effective_policy(&cfg, "alice", None, None, None, Some(2222));
assert_eq!(eff.x11_forwarding, Some(false));
assert!(!eff.x11_forwarding_allowed());
assert!(!eff.agent_forwarding_allowed());
let eff2 = resolve_effective_policy(&cfg, "alice", None, None, None, Some(22));
assert_eq!(eff2.x11_forwarding, None);
assert!(eff2.x11_forwarding_allowed());
assert!(eff2.agent_forwarding_allowed());
}
#[test]
fn match_group_in_effective_policy() {
let cfg = policy_cfg("Match Group dev\n X11Forwarding no\n");
let groups = vec!["dev".to_string()];
let eff = resolve_effective_policy(&cfg, "alice", Some(&groups), None, None, None);
assert_eq!(eff.x11_forwarding, Some(false));
let eff2 = resolve_effective_policy(&cfg, "alice", None, None, None, None);
assert_eq!(eff2.x11_forwarding, None);
}
#[test]
fn max_auth_tries_flows_into_preauth() {
let cfg = policy_cfg("MaxAuthTries 3\n");
let pre = resolve_preauth_policy(&cfg, Some("203.0.113.5"), None, None);
assert_eq!(pre.max_auth_tries, Some(3));
}
#[test]
fn reresolve_match_user_drops_publickey() {
let cfg = policy_cfg("Match User alice\n PubkeyAuthentication no\n");
let pre = resolve_preauth_policy(&cfg, Some("203.0.113.5"), None, None);
assert_eq!(pre.methods, vec!["publickey"]);
let alice = reresolve_user_policy(&cfg, &pre.methods, "alice", None, None, None);
assert!(
alice.methods.is_empty(),
"alice should lose publickey, got {:?}",
alice.methods
);
let bob = reresolve_user_policy(&cfg, &pre.methods, "bob", None, None, None);
assert_eq!(bob.methods, vec!["publickey"]);
}
#[test]
fn reresolve_match_group_uses_group_resolver() {
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(fresh_seed()));
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(|| -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: "x".into(),
allowed_blob: Vec::new(),
})
});
let policy =
crate::config::SshServerConfig::parse("Match Group dev\n PubkeyAuthentication no\n")
.expect("parse");
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler { out: Vec::new() }),
)
.with_policy(Arc::new(policy))
.with_group_resolver(Arc::new(|user: &str| {
if user == "alice" {
vec!["dev".to_string()]
} else {
vec!["users".to_string()]
}
}));
let base = vec!["publickey"];
let alice = reresolve_user_policy(&cfg, &base, "alice", None, None, None);
assert!(alice.methods.is_empty());
let bob = reresolve_user_policy(&cfg, &base, "bob", None, None, None);
assert_eq!(bob.methods, vec!["publickey"]);
}
#[test]
fn reresolve_match_user_banner() {
use std::io::Write;
let dir = std::env::temp_dir();
let path = dir.join(format!("puressh-banner-{}.txt", std::process::id()));
{
let mut f = std::fs::File::create(&path).expect("create banner");
f.write_all(b"hello alice\n").expect("write banner");
}
let src = format!(
"Match User alice\n Banner {}\n",
path.to_str().expect("utf8 path")
);
let cfg = policy_cfg(&src);
let base = vec!["publickey"];
let alice = reresolve_user_policy(&cfg, &base, "alice", None, None, None);
assert_eq!(alice.banner.as_deref(), Some("hello alice\n"));
let bob = reresolve_user_policy(&cfg, &base, "bob", None, None, None);
assert!(bob.banner.is_none());
let _ = std::fs::remove_file(&path);
}
#[test]
fn reresolve_unreadable_banner_is_skipped() {
let cfg = policy_cfg("Match User alice\n Banner /nonexistent/puressh/banner\n");
let base = vec!["publickey"];
let alice = reresolve_user_policy(&cfg, &base, "alice", None, None, None);
assert!(alice.banner.is_none());
assert_eq!(alice.methods, vec!["publickey"]);
}
#[test]
fn reresolve_no_policy_passes_through() {
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(fresh_seed()));
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(|| -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: "x".into(),
allowed_blob: Vec::new(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler { out: Vec::new() }),
);
let base = vec!["publickey"];
let r = reresolve_user_policy(&cfg, &base, "anyone", None, None, None);
assert_eq!(r.methods, vec!["publickey"]);
assert!(r.banner.is_none());
}
#[test]
fn no_policy_is_unrestricted() {
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(fresh_seed()));
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(|| -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: "x".into(),
allowed_blob: Vec::new(),
})
});
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(StaticHandler { out: Vec::new() }),
);
let pre = resolve_preauth_policy(&cfg, Some("1.2.3.4"), None, None);
assert_eq!(pre.methods, vec!["publickey"]);
assert!(pre.banner.is_none());
let eff = resolve_effective_policy(&cfg, "u", None, None, None, None);
assert!(eff.agent_forwarding_allowed());
assert!(eff.x11_forwarding_allowed());
}
#[test]
fn env_glob_match_grammar_basics() {
assert!(env_glob_match("*", "anything"));
assert!(env_glob_match("*", ""));
assert!(env_glob_match("a*b", "ab"));
assert!(env_glob_match("a*b", "aXYZb"));
assert!(!env_glob_match("a*b", "aXYZ"));
assert!(env_glob_match("?", "a"));
assert!(!env_glob_match("?", "ab"));
assert!(!env_glob_match("?", ""));
assert!(env_glob_match("**", "anything"));
assert!(env_glob_match("FOO", "FOO"));
assert!(!env_glob_match("FOO", "foo"));
}
#[test]
fn w7_fields_resolve_into_effective_policy() {
let cfg = policy_cfg(
"MaxSessions 2\n\
AllowTcpForwarding local\n\
PermitOpen 127.0.0.1:80\n\
PermitListen 127.0.0.1:8080\n\
GatewayPorts yes\n\
ForceCommand /bin/true\n\
ClientAliveInterval 15\n\
ClientAliveCountMax 2\n\
PrintMotd yes\n",
);
let eff = resolve_effective_policy(&cfg, "alice", None, None, None, None);
assert_eq!(eff.max_sessions, Some(2));
assert!(eff.local_forwarding_allowed());
assert!(!eff.remote_forwarding_allowed());
assert!(eff.permit_open_allows("127.0.0.1", 80));
assert!(!eff.permit_open_allows("127.0.0.1", 81));
assert!(eff.permit_listen_allows("127.0.0.1", 8080));
assert!(!eff.permit_listen_allows("10.0.0.1", 8080));
assert_eq!(
eff.gateway_ports,
Some(crate::config::ServerGatewayPorts::Yes)
);
assert_eq!(eff.force_command.as_deref(), Some("/bin/true"));
assert_eq!(eff.client_alive_interval, Some(15));
assert_eq!(eff.client_alive_count_max, Some(2));
assert_eq!(eff.print_motd, Some(true));
}
#[test]
fn w7_unset_policy_is_permissive() {
let cfg = policy_cfg("Port 22\n");
let eff = resolve_effective_policy(&cfg, "alice", None, None, None, None);
assert!(eff.local_forwarding_allowed());
assert!(eff.remote_forwarding_allowed());
assert!(eff.permit_open_allows("anything", 1));
assert!(eff.permit_listen_allows("anything", 1));
assert_eq!(eff.max_sessions, None);
assert!(eff.force_command.is_none());
}
#[test]
fn w7_permit_open_none_denies_all() {
let cfg = policy_cfg("PermitOpen none\n");
let eff = resolve_effective_policy(&cfg, "alice", None, None, None, None);
assert!(!eff.permit_open_allows("127.0.0.1", 22));
assert!(!eff.permit_open_allows("anything", 1));
}
#[test]
fn w7_match_overrides_forwarding_for_user() {
let cfg = policy_cfg(
"Match User alice\n AllowTcpForwarding no\n MaxSessions 1\n ForceCommand internal-sftp\n",
);
let eff = resolve_effective_policy(&cfg, "alice", None, None, None, None);
assert!(!eff.local_forwarding_allowed());
assert!(!eff.remote_forwarding_allowed());
assert_eq!(eff.max_sessions, Some(1));
assert_eq!(eff.force_command.as_deref(), Some("internal-sftp"));
let eff2 = resolve_effective_policy(&cfg, "bob", None, None, None, None);
assert!(eff2.local_forwarding_allowed());
assert_eq!(eff2.max_sessions, None);
assert!(eff2.force_command.is_none());
}
fn caps_all() -> crate::auth::AuthCertCaps {
crate::auth::AuthCertCaps {
permit_pty: true,
permit_port_forwarding: true,
permit_agent_forwarding: true,
permit_x11_forwarding: true,
force_command: None,
}
}
fn caps_none() -> crate::auth::AuthCertCaps {
crate::auth::AuthCertCaps {
permit_pty: false,
permit_port_forwarding: false,
permit_agent_forwarding: false,
permit_x11_forwarding: false,
force_command: None,
}
}
#[test]
fn r1_plain_key_auth_allows_all_capabilities() {
let eff = EffectivePolicy::unrestricted();
assert!(eff.pty_allowed());
assert!(eff.agent_forwarding_allowed());
assert!(eff.x11_forwarding_allowed());
assert!(eff.local_forwarding_allowed());
assert!(eff.remote_forwarding_allowed());
}
#[test]
fn r1_cert_with_all_extensions_allows_all_capabilities() {
let mut eff = EffectivePolicy::unrestricted();
eff.cert_caps = Some(caps_all());
assert!(eff.pty_allowed());
assert!(eff.agent_forwarding_allowed());
assert!(eff.x11_forwarding_allowed());
assert!(eff.local_forwarding_allowed());
assert!(eff.remote_forwarding_allowed());
}
#[test]
fn r1_cert_without_extensions_denies_each_capability() {
let mut eff = EffectivePolicy::unrestricted();
eff.cert_caps = Some(caps_none());
assert!(!eff.pty_allowed());
assert!(!eff.agent_forwarding_allowed());
assert!(!eff.x11_forwarding_allowed());
assert!(!eff.local_forwarding_allowed());
assert!(!eff.remote_forwarding_allowed());
}
#[test]
fn r1_cert_extensions_gate_independently() {
let mut eff = EffectivePolicy::unrestricted();
eff.cert_caps = Some(crate::auth::AuthCertCaps {
permit_pty: true,
permit_port_forwarding: false,
permit_agent_forwarding: false,
permit_x11_forwarding: false,
force_command: None,
});
assert!(eff.pty_allowed());
assert!(!eff.agent_forwarding_allowed());
assert!(!eff.x11_forwarding_allowed());
assert!(!eff.local_forwarding_allowed());
assert!(!eff.remote_forwarding_allowed());
}
#[test]
fn r1_cert_and_config_gates_compose_with_and() {
let cfg = policy_cfg("AllowTcpForwarding no\n");
let mut eff = resolve_effective_policy(&cfg, "alice", None, None, None, None);
eff.cert_caps = Some(caps_all());
assert!(!eff.local_forwarding_allowed());
assert!(!eff.remote_forwarding_allowed());
let cfg2 = policy_cfg("Port 22\n");
let mut eff2 = resolve_effective_policy(&cfg2, "alice", None, None, None, None);
eff2.cert_caps = Some(caps_none());
assert!(!eff2.local_forwarding_allowed());
assert!(!eff2.pty_allowed());
}
#[test]
fn r1_auth_cert_caps_reads_extensions_from_cert() {
let blob = {
let path = format!(
"{}/tests/fixtures/cert/u_ed25519-cert.pub",
env!("CARGO_MANIFEST_DIR")
);
let text = std::fs::read_to_string(&path).expect("read fixture");
let b64 = text.split_whitespace().nth(1).expect("base64");
crate::key::base64::decode(b64.as_bytes()).expect("decode")
};
let cert = crate::cert::Certificate::parse(&blob).expect("parse cert");
let ci = crate::auth::CertInfo::from_certificate(&cert).expect("certinfo");
let caps = crate::auth::AuthCertCaps::from_cert_info(&ci);
assert!(caps.permit_pty);
assert!(caps.permit_port_forwarding);
assert!(caps.permit_agent_forwarding);
assert!(caps.permit_x11_forwarding);
assert!(caps.force_command.is_none());
let mut stripped = ci.clone();
stripped.extensions.clear();
let caps2 = crate::auth::AuthCertCaps::from_cert_info(&stripped);
assert!(!caps2.permit_pty);
assert!(!caps2.permit_port_forwarding);
assert!(!caps2.permit_agent_forwarding);
assert!(!caps2.permit_x11_forwarding);
}
#[test]
fn w7_gateway_ports_rewrite() {
use crate::config::ServerGatewayPorts as GP;
use crate::forwarding::reverse::apply_gateway_ports;
assert_eq!(apply_gateway_ports(None, "0.0.0.0"), "127.0.0.1");
assert_eq!(apply_gateway_ports(Some(GP::No), ""), "127.0.0.1");
assert_eq!(apply_gateway_ports(Some(GP::No), "::"), "::1");
assert_eq!(apply_gateway_ports(Some(GP::Yes), ""), "0.0.0.0");
assert_eq!(apply_gateway_ports(Some(GP::Yes), "127.0.0.1"), "0.0.0.0");
assert_eq!(
apply_gateway_ports(Some(GP::ClientSpecified), "0.0.0.0"),
"0.0.0.0"
);
assert_eq!(
apply_gateway_ports(Some(GP::ClientSpecified), "192.0.2.1"),
"192.0.2.1"
);
}
struct EchoCommandHandler;
impl CommandHandler for EchoCommandHandler {
fn handle(&self, _user: &str, env: &SessionEnv, command: &str) -> ExecResult {
let orig = env.get("SSH_ORIGINAL_COMMAND").unwrap_or("").to_string();
ExecResult {
stdout: format!("CMD={command}\nORIG={orig}\n").into_bytes(),
stderr: Vec::new(),
exit_status: 0,
}
}
}
#[allow(clippy::type_complexity)]
fn w7_server_and_client(
policy_src: &str,
with_direct: bool,
) -> (Client, thread::JoinHandle<Result<()>>, Arc<Mutex<bool>>) {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let client_hk_for_auth = Ed25519HostKey::from_seed(client_seed);
let allowed_blob = client_hk_for_auth.public_blob();
let user = "w7-user".to_string();
let allowed_user_for_factory = user.clone();
let allowed_blob_clone = allowed_blob.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: allowed_user_for_factory.clone(),
allowed_blob: allowed_blob_clone.clone(),
})
});
let policy = crate::config::SshServerConfig::parse(policy_src).expect("parse policy");
let mut cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(EchoCommandHandler),
)
.with_policy(Arc::new(policy));
if with_direct {
cfg = cfg.with_direct_tcpip(Arc::new(
crate::forwarding::direct::DefaultDirectTcpipHandler::permit_all(),
));
}
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind ssh");
let ssh_addr = server.local_addr().expect("ssh addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let server_thread = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
ssh_addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("client connect");
client
.authenticate_publickey(&user, Box::new(Ed25519HostKey::from_seed(client_seed)))
.expect("authenticate");
(client, server_thread, server_done)
}
fn w7_finish(
client: Client,
server_thread: thread::JoinHandle<Result<()>>,
server_done: Arc<Mutex<bool>>,
) {
drop(client);
let start = std::time::Instant::now();
while !*server_done.lock().unwrap() {
if start.elapsed() > Duration::from_secs(10) {
panic!("server thread did not finish in time");
}
thread::sleep(Duration::from_millis(20));
}
let _ = server_thread.join();
}
#[test]
fn w7_force_command_overrides_exec() {
let (mut client, st, sd) = w7_server_and_client("ForceCommand /forced/cmd\n", false);
let out = client.exec("client-asked-this").expect("exec");
let stdout = String::from_utf8_lossy(&out.stdout);
assert!(stdout.contains("CMD=/forced/cmd"), "stdout was {stdout:?}");
assert!(
stdout.contains("ORIG=client-asked-this"),
"stdout was {stdout:?}"
);
w7_finish(client, st, sd);
}
#[test]
fn w7_force_command_match_conditional() {
let (mut client, st, sd) =
w7_server_and_client("Match User w7-user\n ForceCommand /only/forced\n", false);
let out = client.exec("orig").expect("exec");
let stdout = String::from_utf8_lossy(&out.stdout);
assert!(stdout.contains("CMD=/only/forced"), "stdout was {stdout:?}");
w7_finish(client, st, sd);
}
struct RecordingScpStream {
seen: Arc<Mutex<Option<(String, String)>>>,
}
impl ExecStreamHandler for RecordingScpStream {
fn claims(&self, command: &str) -> bool {
let t = command.trim_start();
t.starts_with("scp ") || t == "scp"
}
fn run(
&self,
_user: &str,
env: &SessionEnv,
command: &str,
mut stream: ChannelStream,
) -> Result<()> {
let orig = env.get("SSH_ORIGINAL_COMMAND").unwrap_or("").to_string();
*self.seen.lock().unwrap() = Some((command.to_string(), orig));
let _ = stream.write_all(b"\0");
Ok(())
}
}
#[test]
fn r2_force_command_routes_to_exec_stream_scp() {
let host_seed = fresh_seed();
let client_seed = fresh_seed();
let host_key: Box<dyn HostKey + Send + Sync> =
Box::new(Ed25519HostKey::from_seed(host_seed));
let allowed_blob = Ed25519HostKey::from_seed(client_seed).public_blob();
let user = "scp-user".to_string();
let user_f = user.clone();
let factory: Arc<dyn AuthenticatorFactory> = Arc::new(move || -> Box<dyn Authenticator> {
Box::new(OneKeyAuth {
allowed_user: user_f.clone(),
allowed_blob: allowed_blob.clone(),
})
});
let policy =
crate::config::SshServerConfig::parse("ForceCommand scp -t /tmp/dst\n").expect("parse");
let seen = Arc::new(Mutex::new(None));
let cfg = Config::new(
vec![host_key],
factory,
vec!["publickey"],
Arc::new(EchoCommandHandler),
)
.with_policy(Arc::new(policy))
.with_exec_stream_handler(Arc::new(RecordingScpStream { seen: seen.clone() }));
let mut server = Server::bind("127.0.0.1:0", cfg).expect("bind");
let addr = server.local_addr().expect("addr");
let server_done = Arc::new(Mutex::new(false));
let sd = server_done.clone();
let st = thread::spawn(move || {
let r = server.accept_one();
*sd.lock().unwrap() = true;
r
});
let mut client = Client::connect(
addr,
ClientConfig {
host_key_policy: HostKeyPolicy::AcceptAny,
timeout: Some(Duration::from_secs(10)),
algorithms: Default::default(),
},
)
.expect("connect");
client
.authenticate_publickey(&user, Box::new(Ed25519HostKey::from_seed(client_seed)))
.expect("auth");
let _ = client.exec("the-client-command");
w7_finish(client, st, server_done);
let got = seen.lock().unwrap().clone().expect("exec-stream claimed");
assert_eq!(
got.0, "scp -t /tmp/dst",
"forced scp command reached overlay"
);
assert_eq!(
got.1, "the-client-command",
"original command exposed as SSH_ORIGINAL_COMMAND"
);
}
#[test]
fn w7_max_sessions_zero_rejects_session_open() {
let (mut client, st, sd) = w7_server_and_client("MaxSessions 0\n", false);
match client.exec("anything") {
Ok(_) => panic!("expected session open to be refused by MaxSessions 0"),
Err(Error::Protocol(_)) => {}
Err(other) => panic!("expected Error::Protocol, got {other:?}"),
}
w7_finish(client, st, sd);
}
#[test]
#[allow(deprecated)] fn w7_allow_tcp_forwarding_no_denies_direct() {
let (mut client, st, sd) = w7_server_and_client("AllowTcpForwarding no\n", true);
match client.open_direct_tcpip("127.0.0.1", 80, "127.0.0.1", 0) {
Ok(_) => panic!("expected direct-tcpip to be denied by AllowTcpForwarding no"),
Err(Error::Protocol(_)) => {}
Err(other) => panic!("expected Error::Protocol, got {other:?}"),
}
w7_finish(client, st, sd);
}
#[test]
#[allow(deprecated)] fn w7_permit_open_blocks_disallowed_dest() {
let (mut client, st, sd) = w7_server_and_client("PermitOpen 127.0.0.1:80\n", true);
match client.open_direct_tcpip("127.0.0.1", 81, "127.0.0.1", 0) {
Ok(_) => panic!("expected :81 to be blocked by PermitOpen"),
Err(Error::Protocol(_)) => {}
Err(other) => panic!("expected Error::Protocol, got {other:?}"),
}
let permitted = client.open_direct_tcpip("127.0.0.1", 80, "127.0.0.1", 0);
assert!(
permitted.is_ok(),
"expected :80 open to be accepted by PermitOpen"
);
drop(permitted);
w7_finish(client, st, sd);
}
#[test]
fn w7_compression_no_strips_zlib_from_advert() {
let host_keys: Vec<Box<dyn HostKey + Send + Sync>> =
vec![Box::new(Ed25519HostKey::from_seed(fresh_seed()))];
let mut cfg = kexinit_test_config(host_keys);
cfg.compression = Some(crate::config::Compression::No);
let mut rng = OsRng;
let advert = build_server_kexinit(&mut rng, &cfg);
assert!(advert.comp_s2c.iter().any(|c| c == "none"));
assert!(!advert.comp_s2c.iter().any(|c| c.contains("zlib")));
}
}