use core::fmt;
use core::ops::{Deref, DerefMut};
#[allow(unused_imports)]
use {
crate::error::{Error, Result},
log::{debug, error, info, log, trace, warn},
};
use zeroize::Zeroize;
#[cfg(feature = "alloc")]
use alloc::boxed::Box;
use heapless::Deque;
use crate::encrypt::{KeyState, KeysRecv, KeysSend, SSH_PAYLOAD_START};
use crate::ident::RemoteVersion;
use crate::*;
use crate::{
channel::{ChanData, ChanNum},
packets::Packet,
};
const DEFER_COUNT: usize = 10;
#[cfg_attr(not(fuzzing), derive(zeroize::ZeroizeOnDrop))]
enum SliceOrVec<'a> {
Borrowed(&'a mut [u8]),
#[cfg(feature = "alloc")]
Owned(Box<[u8; config::SSH_MAX_PACKET]>),
}
impl Deref for SliceOrVec<'_> {
type Target = [u8];
fn deref(&self) -> &Self::Target {
match self {
Self::Borrowed(r) => r,
#[cfg(feature = "alloc")]
Self::Owned(r) => r.as_ref(),
}
}
}
impl DerefMut for SliceOrVec<'_> {
fn deref_mut(&mut self) -> &mut Self::Target {
match self {
Self::Borrowed(r) => r,
#[cfg(feature = "alloc")]
Self::Owned(r) => r.as_mut(),
}
}
}
impl Zeroize for SliceOrVec<'_> {
fn zeroize(&mut self) {
self.deref_mut().zeroize();
}
}
pub struct TrafIn<'a> {
buf: SliceOrVec<'a>,
state: RxState,
}
#[derive(Debug)]
enum RxState {
Idle,
ReadInitial { idx: usize },
Read { idx: usize, expect: usize },
ReadComplete { len: usize },
InPayload { len: usize, seq: u32 },
InChannelData {
chan: ChanNum,
dt: ChanData,
idx: usize,
len: usize,
},
}
impl core::fmt::Debug for TrafIn<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TrafIn").field("state", &self.state).finish_non_exhaustive()
}
}
impl<'a> TrafIn<'a> {
pub fn new(buf: &'a mut [u8]) -> Self {
Self { buf: SliceOrVec::Borrowed(buf), state: RxState::Idle }
}
pub fn is_input_ready(&self) -> bool {
match self.state {
RxState::Idle | RxState::ReadInitial { .. } | RxState::Read { .. } => {
true
}
RxState::ReadComplete { .. }
| RxState::InPayload { .. }
| RxState::InChannelData { .. } => false,
}
}
pub fn input(
&mut self,
keys: &mut KeyState,
remote_version: &mut RemoteVersion,
buf: &[u8],
) -> Result<usize, Error> {
let mut inlen = 0;
debug_assert!(self.is_input_ready());
if remote_version.version().is_none() && matches!(self.state, RxState::Idle)
{
inlen += remote_version.consume(buf)?;
}
let buf = &buf[inlen..];
inlen += self.fill_input(keys, buf)?;
Ok(inlen)
}
pub(crate) fn done_payload(&mut self) {
if let RxState::InPayload { .. } = self.state {
self.state = RxState::Idle
}
}
pub(crate) fn zeroize_payload(&mut self) {
if let RxState::InPayload { len, .. } = self.state {
self.buf[SSH_PAYLOAD_START..SSH_PAYLOAD_START + len].zeroize();
self.done_payload()
}
}
pub(crate) fn payload(&self) -> Option<(&[u8], u32)> {
match self.state {
RxState::InPayload { len, seq } => {
let payload = &self.buf[SSH_PAYLOAD_START..SSH_PAYLOAD_START + len];
Some((payload, seq))
}
_ => None,
}
}
fn fill_input(
&mut self,
keys: &mut KeyState,
buf: &[u8],
) -> Result<usize, Error> {
let size_block = keys.size_block_dec();
let mut r = buf;
trace!("fill_input {:?}", self.state);
if let Some(idx) = match self.state {
RxState::Idle if !r.is_empty() => Some(0),
RxState::ReadInitial { idx } => Some(idx),
_ => None,
} {
trace!("fill_input idle idx {idx}");
let need = (size_block - idx).clamp(0, r.len());
let x;
(x, r) = r.split_at(need);
let w = &mut self.buf[idx..idx + need];
w.copy_from_slice(x);
self.state = RxState::ReadInitial { idx: idx + need }
}
if let RxState::ReadInitial { idx } = self.state {
trace!("fill_input readinit {idx}");
if idx >= size_block {
let w = &mut self.buf[..size_block];
let total_len = keys.decrypt_first_block(w)?;
if total_len > self.buf.len() {
return Err(Error::BigPacket { size: total_len });
}
if total_len < size_block {
return Err(Error::BadDecrypt);
}
trace!("fill_input set read {idx} ex {total_len}");
self.state = RxState::Read { idx, expect: total_len }
}
}
if let RxState::Read { ref mut idx, expect } = self.state {
trace!("expect {expect} idx {idx}");
let need = (expect - *idx).min(r.len());
let x;
(x, r) = r.split_at(need);
let w = &mut self.buf[*idx..*idx + need];
w.copy_from_slice(x);
*idx += need;
if *idx == expect {
self.state = RxState::ReadComplete { len: expect }
}
}
if let RxState::ReadComplete { len } = self.state {
let w = &mut self.buf[..len];
let seq = keys.recv_seq();
let payload_len = keys.decrypt(w)?;
self.state = RxState::InPayload { len: payload_len, seq }
}
trace!("out");
Ok(buf.len() - r.len())
}
pub fn read_channel_ready(&self) -> Option<(ChanNum, ChanData, usize)> {
match self.state {
RxState::InChannelData { chan, dt, idx, len } => {
debug_assert!(len > idx);
let rem = len - idx;
Some((chan, dt, rem))
}
_ => None,
}
}
pub fn set_read_channel_data(
&mut self,
di: channel::DataIn,
) -> Result<(ChanNum, ChanData)> {
match self.state {
RxState::InPayload { .. } => {
let idx = SSH_PAYLOAD_START + di.dt.packet_offset();
self.state = RxState::InChannelData {
chan: di.num,
dt: di.dt,
idx,
len: idx + di.len.get(),
};
Ok((di.num, di.dt))
}
_ => Error::bug(),
}
}
pub fn read_channel(
&mut self,
chan: ChanNum,
dt: ChanData,
buf: &mut [u8],
) -> (usize, Option<usize>) {
match self.state {
RxState::InChannelData { chan: c, dt: e, ref mut idx, len }
if (c, e) == (chan, dt) =>
{
debug_assert!(len > *idx);
let wlen = (len - *idx).min(buf.len());
buf[..wlen].copy_from_slice(&self.buf[*idx..*idx + wlen]);
*idx += wlen;
if *idx == len {
self.state = RxState::Idle;
(wlen, Some(len))
} else {
(wlen, None)
}
}
_ => (0, None),
}
}
pub fn read_channel_either(
&mut self,
chan: ChanNum,
buf: &mut [u8],
) -> (usize, Option<usize>, ChanData) {
match self.state {
RxState::InChannelData { chan: c, dt, ref mut idx, len }
if c == chan =>
{
debug_assert!(len > *idx);
let wlen = (len - *idx).min(buf.len());
buf[..wlen].copy_from_slice(&self.buf[*idx..*idx + wlen]);
*idx += wlen;
if *idx == len {
self.state = RxState::Idle;
(wlen, Some(len), dt)
} else {
(wlen, None, dt)
}
}
_ => (0, None, ChanData::Normal),
}
}
pub fn discard_read_channel(&mut self, chan: ChanNum) -> usize {
match self.state {
RxState::InChannelData { chan: c, len, .. } if c == chan => {
self.state = RxState::Idle;
len
}
_ => 0,
}
}
}
pub(crate) struct TrafOut<'a> {
buf: SliceOrVec<'a>,
state: TxState,
drain: bool,
sending_kex: bool,
deferred_packets: Deque<DeferredPacket, DEFER_COUNT>,
}
#[derive(Debug)]
enum TxState {
Idle,
Write {
idx: usize,
len: usize,
},
Closed,
}
impl core::fmt::Debug for TrafOut<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TrafOut").field("state", &self.state).finish_non_exhaustive()
}
}
#[cfg(feature = "alloc")]
impl TrafIn<'static> {
pub fn new_owned() -> Self {
let mut s = Self::new(&mut []);
s.buf = SliceOrVec::Owned(Box::new([0; _]));
s
}
}
impl<'a> TrafOut<'a> {
pub fn new(buf: &'a mut [u8]) -> Self {
Self {
buf: SliceOrVec::Borrowed(buf),
state: TxState::Idle,
drain: false,
sending_kex: false,
deferred_packets: Deque::new(),
}
}
pub(crate) fn send_packet(
&mut self,
p: packets::Packet,
keys: &mut KeyState,
) -> Result<()> {
let is_kex = matches!(p.category(), packets::Category::Kex);
if is_kex || (self.deferred_packets.is_empty() && !self.sending_kex) {
match self.send_packet_inner(&p, keys) {
Err(Error::NoRoom { .. }) => {
debug_assert!(!is_kex, "KEX packets should have room");
}
res => return res,
}
}
let pnum = p.message_num();
trace!("Delay packet type {pnum:?}");
let Ok(dp) = DeferredPacket::try_from(p) else {
trace!("NoRoom packet type {pnum:?}");
return error::BusySend { packet: pnum, unsupported: true }.fail();
};
self.deferred_packets.push_front(dp).map_err(|_| {
error!("No space to queue packet");
trace!("NoRoom packet type {pnum:?}");
error::BusySend { packet: pnum, unsupported: false }.build()
})
}
fn track_send_packet(
&mut self,
p: &packets::Packet,
keys: &mut KeyState,
) -> Result<()> {
match p.category() {
packets::Category::All | packets::Category::Kex => (),
_ => {
if keys.is_send_cleartext() {
return Error::bug_msg("send cleartext");
}
}
}
if self.sending_kex {
debug_assert!(matches!(
p.category(),
packets::Category::All | packets::Category::Kex
));
}
match p {
Packet::KexInit(_) => {
debug_assert!(!self.sending_kex);
self.sending_kex = true;
}
Packet::NewKeys(_) => {
debug_assert!(self.sending_kex);
self.sending_kex = false;
}
_ => (),
}
Ok(())
}
pub fn send_packet_inner(
&mut self,
p: &packets::Packet,
keys: &mut KeyState,
) -> Result<()> {
self.track_send_packet(p, keys)?;
let (idx, len) = match self.state {
TxState::Idle => (0, 0),
TxState::Write { idx, len } => (idx, len),
TxState::Closed => {
trace!("Dropped output after close {p:?}");
return Ok(());
}
};
let wbuf = &mut self.buf[len..];
if wbuf.len() < SSH_PAYLOAD_START {
return error::NoRoom.fail();
}
let plen = sshwire::write_ssh(&mut wbuf[SSH_PAYLOAD_START..], &p)?;
trace!("Sending {p:?}");
let elen = keys.encrypt(plen, wbuf)?;
self.state = TxState::Write { idx, len: len + elen };
Ok(())
}
pub fn send_deferred_packets(&mut self, keys: &mut KeyState) -> Result<()> {
while let Some(d) = self.deferred_packets.back() {
let p = Packet::from(d);
match self.send_packet_inner(&p, keys) {
Ok(()) => {
self.deferred_packets.pop_back();
}
Err(Error::NoRoom { .. }) => {
break;
}
Err(e) => return Err(e),
}
}
Ok(())
}
pub fn have_deferred_packets(&self) -> bool {
!self.deferred_packets.is_empty()
}
pub fn is_output_pending(&self) -> bool {
trace!("is_output_pending st {:?}", self.state);
matches!(self.state, TxState::Write { .. })
}
pub fn send_allowed(&self, keys: &KeyState) -> usize {
if !self.deferred_packets.is_empty() {
return 0;
}
match self.state {
TxState::Write { len, .. } => keys.max_enc_payload(self.buf.len() - len),
TxState::Idle => keys.max_enc_payload(self.buf.len()),
TxState::Closed => self.buf.len(),
}
}
pub fn close(&mut self) {
self.state = TxState::Closed
}
pub fn closed(&self) -> bool {
matches!(self.state, TxState::Closed)
}
pub fn send_version(&mut self) -> Result<(), Error> {
if !matches!(self.state, TxState::Idle) {
return Error::bug();
}
let len = ident::write_version(&mut self.buf)?;
self.state = TxState::Write { idx: 0, len };
Ok(())
}
pub fn output_buf(&mut self) -> &[u8] {
match self.state {
TxState::Write { ref mut idx, len } => {
let wlen = len - *idx;
&self.buf[*idx..*idx + wlen]
}
_ => &[],
}
}
pub fn consume_output(&mut self, l: usize) {
if let TxState::Write { ref mut idx, len } = self.state {
let wlen = (len - *idx).min(l);
*idx += wlen;
if *idx == len {
self.state = TxState::Idle
}
}
}
pub fn sender<'s>(&'s mut self, keys: &'s mut KeyState) -> TrafSend<'s, 'a> {
TrafSend::new(self, keys)
}
pub fn is_draining(&self) -> bool {
self.drain
}
}
#[cfg(feature = "alloc")]
impl TrafOut<'static> {
pub fn new_owned() -> Self {
let mut s = Self::new(&mut []);
s.buf = SliceOrVec::Owned(Box::new([0; _]));
s
}
}
pub(crate) struct TrafSend<'s, 'a> {
out: &'s mut TrafOut<'a>,
keys: &'s mut KeyState,
}
impl<'s, 'a> TrafSend<'s, 'a> {
fn new(out: &'s mut TrafOut<'a>, keys: &'s mut KeyState) -> Self {
Self { out, keys }
}
pub fn send<'p, P: Into<packets::Packet<'p>>>(&mut self, p: P) -> Result<()> {
self.out.send_packet(p.into(), self.keys)
}
pub fn rekey_send(&mut self, keys: KeysSend) {
self.keys.rekey_send(keys);
}
pub fn rekey_recv(&mut self, keys: KeysRecv) {
self.keys.rekey_recv(keys)
}
pub fn send_version(&mut self) -> Result<(), Error> {
self.out.send_version()
}
pub fn recv_seq(&self) -> u32 {
self.keys.seq_decrypt.0
}
pub fn enable_strict_kex(&mut self) {
self.keys.enable_strict_kex();
}
pub fn is_rekey_needed(&self) -> bool {
self.keys.is_rekey_needed()
}
pub fn set_drain_output(&mut self, drain: bool) {
debug_assert!(drain != self.out.drain, "set_drain_output() dupe");
self.out.drain = drain;
}
pub fn is_drained(&self) -> bool {
debug_assert!(self.out.drain);
matches!(self.out.state, TxState::Idle)
&& self.out.deferred_packets.is_empty()
}
}
#[derive(Debug)]
pub enum DeferredPacket {
ChannelSuccess(packets::ChannelSuccess),
ChannelFailure(packets::ChannelFailure),
ChannelOpenFailure(packets::ChannelOpenFailure<'static>),
ChannelOpenConfirmation(packets::ChannelOpenConfirmation),
ChannelEof(packets::ChannelEof),
ChannelClose(packets::ChannelClose),
Unimplemented(packets::Unimplemented),
RequestFailure(packets::RequestFailure),
RequestSuccess(packets::RequestSuccess),
}
impl DeferredPacket {}
impl From<&DeferredPacket> for Packet<'static> {
fn from(d: &DeferredPacket) -> Self {
match d {
DeferredPacket::ChannelSuccess(p) => (p.clone()).into(),
DeferredPacket::ChannelFailure(p) => (p.clone()).into(),
DeferredPacket::ChannelOpenFailure(p) => (p.clone()).into(),
DeferredPacket::ChannelOpenConfirmation(p) => (p.clone()).into(),
DeferredPacket::ChannelEof(p) => (p.clone()).into(),
DeferredPacket::ChannelClose(p) => (p.clone()).into(),
DeferredPacket::Unimplemented(p) => (p.clone()).into(),
DeferredPacket::RequestFailure(p) => (p.clone()).into(),
DeferredPacket::RequestSuccess(p) => (p.clone()).into(),
}
}
}
impl<'a> TryFrom<Packet<'a>> for DeferredPacket {
type Error = Error;
fn try_from(packet: Packet<'a>) -> Result<Self> {
Ok(match packet {
Packet::ChannelSuccess(p) => Self::ChannelSuccess(p),
Packet::ChannelFailure(p) => Self::ChannelFailure(p),
Packet::ChannelOpenConfirmation(p) => Self::ChannelOpenConfirmation(p),
Packet::ChannelEof(p) => Self::ChannelEof(p),
Packet::ChannelClose(p) => Self::ChannelClose(p),
Packet::Unimplemented(p) => Self::Unimplemented(p),
Packet::RequestFailure(p) => Self::RequestFailure(p),
Packet::RequestSuccess(p) => Self::RequestSuccess(p),
Packet::ChannelOpenFailure(p) => {
Self::ChannelOpenFailure(packets::ChannelOpenFailure {
desc: TextString::new(),
lang: "",
..p
})
}
_ => return error::SSHProtoUnsupported.fail(),
})
}
}