Skip to content
Open
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
70 changes: 70 additions & 0 deletions tls/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -606,6 +606,76 @@ mod tests {
handle.await.unwrap();
}

// A version message whose body is shorter than the 4-byte version, sent
// by a TLS-authenticated client, is rejected as Error::ProtocolVersion —
// the server task must error, not panic on a short slice.
#[tokio::test]
async fn short_version_message_rejected() {
let log = logger();
let mock_datadir = mock_datadir();
let addr: SocketAddrV6 = SocketAddrV6::from_str("[::1]:46467").unwrap();

let server_config =
local_config(1, MeasurementConnectionPolicy::Enforced);
let corpus = vec![
mock_datadir.join("corim-rot.cbor"),
mock_datadir.join("corim-sp.cbor"),
];

let log2 = log.clone();
let handle = tokio::spawn(async move {
let server = Server::new(server_config, addr, log2.clone())
.await
.unwrap();

let result = server
.accept(corpus.clone())
.await
.unwrap()
.handshake()
.await;
match result {
Err(Error::ProtocolVersion) => {}
Err(other) => {
panic!("expected ProtocolVersion, got {other:?}")
}
Ok(_) => {
panic!("a malformed version message must not complete")
}
}
});

let client_config = client::Client::new_tls_local_client_config(
mock_datadir.join("test-sprockets-auth-2.key.pem"),
mock_datadir.join("test-sprockets-auth-2.certlist.pem"),
vec![mock_datadir.join("test-root-a.cert.pem")],
log,
)
.unwrap();

let dnsname =
rustls::pki_types::ServerName::try_from("unknown.com").unwrap();
let connector = tokio_rustls::TlsConnector::from(std::sync::Arc::new(
client_config,
));
let stream = loop {
if let Ok(s) = tokio::net::TcpStream::connect(addr).await {
break s;
};
sleep(Duration::from_millis(1)).await;
};
let mut stream = connector.connect(dnsname, stream).await.unwrap();

// A valid length prefix (2) followed by a 2-byte body: the server's
// recv_msg succeeds, but the body is shorter than the 4-byte version
// it must contain.
stream.write_all(&2u32.to_le_bytes()).await.unwrap();
stream.write_all(b"xy").await.unwrap();
stream.shutdown().await.unwrap();

handle.await.unwrap();
}

#[tokio::test]
async fn spawn_accept() {
let log = logger();
Expand Down
10 changes: 8 additions & 2 deletions tls/src/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -138,8 +138,14 @@ impl SprocketsAcceptor {

// get version from the client
let version_bytes = recv_msg(&mut stream).await?;
let version =
u32::from_le_bytes(version_bytes[..4].try_into().unwrap());
// Anything but exactly the 4-byte little-endian version is protocol
// garbage from the peer; reject it rather than index past the end of a
// short message.
let version_bytes: [u8; 4] = version_bytes
.as_slice()
.try_into()
.map_err(|_| Error::ProtocolVersion)?;
let version = u32::from_le_bytes(version_bytes);

if version == CURRENT_PROTOCOL_VERSION {
// we're good to go
Expand Down
Loading