bunshin/src/channel.rs

123 lines
4.1 KiB
Rust
Raw Normal View History

//! Noise_NX handshake and the framed transport on top of it.
//!
//! Two independent layers: Noise messages under a u16 length prefix, and
//! application frames (u32 length || u8 op || body) split across as many
//! Noise messages as they need and reassembled from them.
use std::io::{self, Read, Write};
use std::net::TcpStream;
use snow::{Builder, TransportState};
use crate::proto::{ProtocolError, MAX_FRAME, NOISE_PARAMS, NOISE_PAYLOAD, PROLOGUE};
/// Runs the Noise_NX responder handshake. The initiator stays anonymous;
/// only we hold a static key. Returns the transport state and the
/// handshake hash (needed later to verify AUTH frames).
pub fn handshake(
stream: &mut TcpStream,
static_key: &[u8],
) -> anyhow::Result<(TransportState, Vec<u8>)> {
let params: snow::params::NoiseParams = NOISE_PARAMS.parse()?;
let mut noise = Builder::new(params)
.local_private_key(static_key)
.prologue(PROLOGUE)
.build_responder()?;
let mut buf = [0u8; 65535];
let mut msg = [0u8; 65535];
let len = read_u16_len(stream)?;
read_exact_into(stream, &mut buf[..len])?;
noise.read_message(&buf[..len], &mut msg)?;
let len = noise.write_message(&[], &mut buf)?;
write_u16_len(stream, &buf[..len])?;
anyhow::ensure!(noise.is_handshake_finished(), "handshake did not complete");
let hash = noise.get_handshake_hash().to_vec();
let transport = noise.into_transport_mode()?;
Ok((transport, hash))
}
fn read_exact_into(stream: &mut TcpStream, buf: &mut [u8]) -> io::Result<()> {
stream.read_exact(buf)
}
fn read_u16_len(stream: &mut TcpStream) -> io::Result<usize> {
let mut len_buf = [0u8; 2];
stream.read_exact(&mut len_buf)?;
Ok(u16::from_be_bytes(len_buf) as usize)
}
fn write_u16_len(stream: &mut TcpStream, packet: &[u8]) -> io::Result<()> {
stream.write_all(&(packet.len() as u16).to_be_bytes())?;
stream.write_all(packet)?;
Ok(())
}
pub struct Channel {
stream: TcpStream,
transport: TransportState,
buf: Vec<u8>,
}
impl Channel {
pub fn new(stream: TcpStream, transport: TransportState) -> Self {
Channel {
stream,
transport,
buf: Vec::new(),
}
}
fn read_noise(&mut self) -> anyhow::Result<Vec<u8>> {
let len = read_u16_len(&mut self.stream)?;
let mut ciphertext = vec![0u8; len];
read_exact_into(&mut self.stream, &mut ciphertext)?;
let mut plaintext = vec![0u8; len];
let n = self.transport.read_message(&ciphertext, &mut plaintext)?;
plaintext.truncate(n);
Ok(plaintext)
}
fn write_noise(&mut self, payload: &[u8]) -> anyhow::Result<()> {
let mut packet = vec![0u8; payload.len() + 16];
let n = self.transport.write_message(payload, &mut packet)?;
packet.truncate(n);
write_u16_len(&mut self.stream, &packet)?;
Ok(())
}
/// Reads one application frame, blocking until a full frame is available.
pub fn read_frame(&mut self) -> anyhow::Result<(u8, Vec<u8>)> {
while self.buf.len() < 5 {
let chunk = self.read_noise()?;
self.buf.extend_from_slice(&chunk);
}
let length = u32::from_be_bytes(self.buf[..4].try_into().unwrap()) as usize;
if length < 1 || length > MAX_FRAME {
return Err(ProtocolError::new(format!("frame length {length} out of range")).into());
}
while self.buf.len() < 4 + length {
let chunk = self.read_noise()?;
self.buf.extend_from_slice(&chunk);
}
let frame: Vec<u8> = self.buf[4..4 + length].to_vec();
self.buf.drain(..4 + length);
Ok((frame[0], frame[1..].to_vec()))
}
pub fn write_frame(&mut self, op: u8, body: &[u8]) -> anyhow::Result<()> {
let mut frame = Vec::with_capacity(5 + body.len());
frame.extend_from_slice(&((1 + body.len()) as u32).to_be_bytes());
frame.push(op);
frame.extend_from_slice(body);
for chunk in frame.chunks(NOISE_PAYLOAD) {
self.write_noise(chunk)?;
}
Ok(())
}
}