Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions crates/bssh-russh/src/cipher/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -345,10 +345,12 @@ pub(crate) async fn read<R: AsyncRead + Unpin>(
.len()
.checked_sub(padding_length)
.ok_or(Error::IndexOutOfBounds)?;
let payload_len = plaintext_end.saturating_sub(PADDING_LENGTH_LEN);

// Sequence numbers are on 32 bits and wrap.
// https://tools.ietf.org/html/rfc4253#section-6.4
buffer.seqn += Wrapping(1);
buffer.bytes = buffer.bytes.saturating_add(payload_len);
buffer.len = 0;

// Remove the padding
Expand Down
88 changes: 87 additions & 1 deletion crates/bssh-russh/src/client/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,6 @@ use std::convert::TryInto;
use std::num::Wrapping;
use std::pin::Pin;
use std::sync::Arc;
#[cfg(not(target_arch = "wasm32"))]
use std::time::Duration;

use futures::Future;
Expand Down Expand Up @@ -1315,6 +1314,9 @@ impl Session {
// application output in its bounded receivers while a channel is
// window-blocked.
let can_receive_outbound = !self.kex.active() && !self.common.has_any_pending_data();
let rekey_timer =
crate::future_or_pending(self.rekey_time_remaining(), tokio::time::sleep);
pin!(rekey_timer);
tokio::select! {
r = &mut reading => {
let (stream_read, mut buffer, mut opening_cipher) = match r {
Expand All @@ -1336,14 +1338,26 @@ impl Session {
result = self.process_disconnect(&pkt).map_err(H::Error::from);
} else {
self.common.received_data = true;
let kex_was_active = self.kex.active();
reply(self, handler, kex_done_signal, &mut pkt).await?;
buffer.seqn = pkt.seqn; // TODO reply changes seqn internall, find cleaner way

if kex_was_active && !self.kex.active() {
buffer.bytes = 0;
} else if self.read_rekey_limit_reached(buffer.bytes) {
debug!("rekey limit reached after {} inbound bytes", buffer.bytes);
self.initiate_rekey()?;
}
}
}

std::mem::swap(&mut opening_cipher, &mut self.common.remote_to_local);
reading.set(start_reading(stream_read, buffer, opening_cipher));
}
() = &mut rekey_timer => {
debug!("rekey time limit reached");
self.initiate_rekey()?;
}
() = &mut keepalive_timer => {
if let Some(ref mut enc) = self.common.encrypted {
if matches!(enc.state, EncryptedState::Authenticated) {
Expand Down Expand Up @@ -1739,6 +1753,34 @@ impl Session {
Ok(())
}

fn read_rekey_limit_reached(&self, read_bytes: usize) -> bool {
!self.kex.active()
&& self
.common
.encrypted
.as_ref()
.is_some_and(|enc| matches!(enc.state, EncryptedState::Authenticated))
&& read_bytes >= self.common.config.limits.rekey_read_limit
}

fn rekey_time_remaining(&self) -> Option<Duration> {
if self.kex.active() {
return None;
}

let enc = self.common.encrypted.as_ref()?;
if !matches!(enc.state, EncryptedState::Authenticated) {
return None;
}

let limit = self.common.config.limits.rekey_time_limit;
if limit == Duration::MAX {
return None;
}

Some(limit.saturating_sub(Instant::now().duration_since(enc.last_rekey)))
}

/// Flush the temporary cleartext buffer into the encryption
/// buffer. This does *not* flush to the socket.
fn flush(&mut self) -> Result<(), crate::Error> {
Expand Down Expand Up @@ -2022,6 +2064,50 @@ mod tests {
);
}

#[test]
fn automatic_rekey_starts_at_inbound_read_limit() {
let (mut session, sender, _replies) = keyboard_interactive_session();
Arc::get_mut(&mut session.common.config)
.expect("test session owns its config")
.limits = crate::Limits::new(
1 << 30,
16,
std::time::Duration::from_secs(3600),
);
let encrypted = session.common.encrypted.as_mut().unwrap();
encrypted.state = EncryptedState::Authenticated;
encrypted.kex = KEXES.get(&crate::kex::CURVE25519).unwrap().make();
drop(sender);

assert!(!session.read_rekey_limit_reached(15));
assert!(session.read_rekey_limit_reached(16));
session.initiate_rekey().unwrap();

assert!(session.kex.active());
}

#[tokio::test]
async fn automatic_rekey_deadline_wakes_an_idle_authenticated_session() {
let (mut session, sender, _replies) = keyboard_interactive_session();
Arc::get_mut(&mut session.common.config)
.expect("test session owns its config")
.limits = crate::Limits::new(1 << 30, 1 << 30, Duration::ZERO);
let encrypted = session.common.encrypted.as_mut().unwrap();
encrypted.state = EncryptedState::Authenticated;
encrypted.kex = KEXES.get(&crate::kex::CURVE25519).unwrap().make();
drop(sender);

let remaining = session
.rekey_time_remaining()
.expect("authenticated sessions with a finite limit need a deadline");
tokio::time::timeout(Duration::from_millis(50), tokio::time::sleep(remaining))
.await
.expect("an expired rekey deadline must wake without transport activity");
session.initiate_rekey().unwrap();

assert!(session.kex.active());
}

#[cfg(feature = "flate2")]
fn authenticated_session() -> Session {
let config = Arc::new(Config::default());
Expand Down
131 changes: 131 additions & 0 deletions crates/bssh-russh/src/client/test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
mod tests {
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::sync::atomic::{AtomicUsize, Ordering};

use log::debug;
use ssh_key::PrivateKey;
Expand Down Expand Up @@ -83,6 +84,31 @@ mod tests {
}
}

struct CountingClient {
kex_count: Arc<AtomicUsize>,
}

impl Handler for CountingClient {
type Error = Error;

async fn check_server_key(
&mut self,
_: &PublicKeyOrCertificate,
) -> Result<bool, Self::Error> {
Ok(true)
}

async fn kex_done(
&mut self,
_: Option<&[u8]>,
_: &crate::negotiation::Names,
_: &mut crate::client::Session,
) -> Result<(), Self::Error> {
self.kex_count.fetch_add(1, Ordering::Relaxed);
Ok(())
}
}

#[tokio::test]
async fn test_client_connects_to_protocol_1_99() {
let _ = env_logger::try_init();
Expand Down Expand Up @@ -164,4 +190,109 @@ mod tests {
msg => panic!("Unexpected message {msg:?}"),
}
}

#[tokio::test]
async fn automatic_rekey_repeated_bidirectional_transfers_complete() {
let client_key = PrivateKey::random(&mut rng(), ssh_key::Algorithm::Ed25519).unwrap();

let mut server_config = server::Config::default();
server_config.auth_rejection_time = std::time::Duration::from_millis(1);
server_config.inactivity_timeout = None;
server_config.limits = crate::Limits::new(
8 * 1024,
8 * 1024,
std::time::Duration::from_secs(3600),
);
server_config
.keys
.push(PrivateKey::random(&mut rng(), ssh_key::Algorithm::Ed25519).unwrap());
let server_config = Arc::new(server_config);

let socket = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = socket.local_addr().unwrap();
let server = TestServer {
clients: Arc::new(Mutex::new(HashMap::new())),
id: 0,
};
let server_task = tokio::spawn(async move {
let (socket, _) = socket.accept().await.unwrap();
let running = server::run_stream(server_config, socket, server)
.await
.unwrap();
running.await
});

let mut client_config = Config::default();
client_config.limits = crate::Limits::new(
8 * 1024,
8 * 1024,
std::time::Duration::from_secs(3600),
);
let kex_count = Arc::new(AtomicUsize::new(0));

tokio::time::timeout(std::time::Duration::from_secs(20), async {
let mut session = connect(
Arc::new(client_config),
addr,
CountingClient {
kex_count: Arc::clone(&kex_count),
},
)
.await
.unwrap();
let authenticated = session
.authenticate_publickey(
"rekey-user",
PrivateKeyWithHashAlg::new(
Arc::new(client_key),
session.best_supported_rsa_hash().await.unwrap().flatten(),
),
)
.await
.unwrap();
assert!(authenticated.success());
let kex_count_before_transfer = kex_count.load(Ordering::Relaxed);

let mut channel = session.channel_open_session().await.unwrap();
for round in 0_u8..6 {
let payload = vec![round; 32 * 1024];
channel.data(&payload[..]).await.unwrap();
let mut echoed = Vec::with_capacity(payload.len());
while echoed.len() < payload.len() {
match channel.wait().await.unwrap() {
crate::channels::ChannelMsg::Data { data } => {
echoed.extend_from_slice(&data)
}
message => panic!("unexpected message during rekey transfer: {message:?}"),
}
}
assert!(echoed == payload, "echo mismatch in rekey round {round}");
}

tokio::time::timeout(std::time::Duration::from_secs(2), async {
while kex_count.load(Ordering::Relaxed) < kex_count_before_transfer + 2 {
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
}
})
.await
.expect("byte thresholds must complete repeated rekeys during transfer");
assert!(
kex_count.load(Ordering::Relaxed) >= kex_count_before_transfer + 2,
"six threshold-crossing rounds must observe repeated key exchanges"
);

session
.disconnect(crate::Disconnect::ByApplication, "test complete", "")
.await
.unwrap();
})
.await
.expect("repeated bidirectional rekeys must not stall");

tokio::time::timeout(std::time::Duration::from_secs(2), server_task)
.await
.expect("server session must stop after client disconnect")
.unwrap()
.unwrap();
}
}
Loading
Loading