Skip to main content

confium_net_noise/
transport.rs

1//! Noise transport: handshake + session over a framed TCP stream.
2
3use std::io::Read;
4use std::io::Write;
5use std::net::SocketAddr;
6use std::net::ToSocketAddrs;
7
8use snafu::IntoError;
9use snafu::ResultExt;
10use snow::Builder;
11use snow::HandshakeState;
12use snow::TransportState;
13use url::Url;
14
15use confium_net::Listener;
16use confium_net::Result;
17use confium_net::Transport;
18use confium_net::error::ClosedSnafu;
19use confium_net::error::IoSnafu;
20use confium_net::error::MalformedUrlSnafu;
21
22use crate::keys::NoiseIdentity;
23use crate::keys::fingerprint_of;
24use crate::keys::hex;
25use crate::keys::noise_params;
26
27/// Maximum plaintext payload per frame (8 MiB), matching the TCP
28/// transport's guard against hostile length prefixes. Ciphertext
29/// frames carry at most this plus the Noise tag (16 bytes).
30const MAX_FRAME_LEN: u32 = 8 * 1024 * 1024;
31
32/// Handshake frames use a smaller cap: snow handshake messages are
33/// always far below this.
34const MAX_HANDSHAKE_LEN: u32 = 65535;
35
36/// snow caps transport-mode messages at 65535 bytes, so application
37/// payloads larger than this are fragmented. Each encrypted chunk
38/// carries a one-byte prefix — 0x00 = more chunks follow, 0x01 =
39/// final — and the receiver reassembles before returning, preserving
40/// the one-send-one-recv contract.
41const MAX_CHUNK: usize = 60_000;
42
43/// Parsed `noise://` URL parameters.
44pub(crate) struct NoiseParams {
45    pub addr: SocketAddr,
46    /// Provisioned local static key (`key=<hex>`); ephemeral if absent.
47    pub local_key: Option<NoiseIdentity>,
48    /// Pinned remote fingerprint (`pinned=<hex>`).
49    pub pinned: Option<[u8; 32]>,
50}
51
52pub(crate) fn parse_url(url: &Url) -> Result<NoiseParams> {
53    let host = url.host_str().ok_or_else(|| {
54        MalformedUrlSnafu {
55            scheme: "noise",
56            url: url.as_str(),
57            reason: "noise URL requires a host",
58        }
59        .build()
60    })?;
61    let port = url.port().ok_or_else(|| {
62        MalformedUrlSnafu {
63            scheme: "noise",
64            url: url.as_str(),
65            reason: "noise URL requires an explicit port",
66        }
67        .build()
68    })?;
69    let addr = (host, port)
70        .to_socket_addrs()
71        .map_err(|_| {
72            MalformedUrlSnafu {
73                scheme: "noise",
74                url: url.as_str(),
75                reason: "could not resolve host and port",
76            }
77            .build()
78        })?
79        .next()
80        .ok_or_else(|| {
81            MalformedUrlSnafu {
82                scheme: "noise",
83                url: url.as_str(),
84                reason: "no address resolved for host and port",
85            }
86            .build()
87        })?;
88
89    let mut local_key = None;
90    let mut pinned = None;
91    for (k, v) in url.query_pairs() {
92        match k.as_ref() {
93            "key" => {
94                local_key = Some(NoiseIdentity::from_hex(&v).map_err(|_| {
95                    MalformedUrlSnafu {
96                        scheme: "noise",
97                        url: url.as_str(),
98                        reason: "key= must be 32 bytes of hex (private key)",
99                    }
100                    .build()
101                })?);
102            }
103            "pinned" => {
104                pinned = Some(
105                    crate::keys::unhex(&v)
106                        .ok()
107                        .and_then(|b| <[u8; 32]>::try_from(b).ok())
108                        .ok_or_else(|| {
109                            MalformedUrlSnafu {
110                                scheme: "noise",
111                                url: url.as_str(),
112                                reason: "pinned= must be 32 bytes of hex (fingerprint)",
113                            }
114                            .build()
115                        })?,
116                );
117            }
118            _ => {}
119        }
120    }
121    Ok(NoiseParams {
122        addr,
123        local_key,
124        pinned,
125    })
126}
127
128// ---- framing ---------------------------------------------------------
129
130fn write_frame<W: Write>(w: &mut W, data: &[u8]) -> std::io::Result<()> {
131    w.write_all(&(data.len() as u32).to_be_bytes())?;
132    w.write_all(data)?;
133    w.flush()
134}
135
136fn read_frame<R: Read>(r: &mut R, max: u32) -> std::io::Result<Option<Vec<u8>>> {
137    let mut prefix = [0u8; 4];
138    if !fill(r, &mut prefix)? {
139        return Ok(None);
140    }
141    let len = u32::from_be_bytes(prefix);
142    if len > max {
143        return Err(std::io::Error::new(
144            std::io::ErrorKind::InvalidData,
145            format!("frame length {len} exceeds maximum {max}"),
146        ));
147    }
148    let mut buf = vec![0u8; len as usize];
149    if !fill(r, &mut buf)? {
150        return Err(std::io::Error::new(
151            std::io::ErrorKind::UnexpectedEof,
152            "EOF inside frame",
153        ));
154    }
155    Ok(Some(buf))
156}
157
158fn fill<R: Read>(r: &mut R, buf: &mut [u8]) -> std::io::Result<bool> {
159    let mut off = 0;
160    while off < buf.len() {
161        match r.read(&mut buf[off..]) {
162            Ok(0) => return Ok(false),
163            Ok(n) => off += n,
164            Err(ref e) if e.kind() == std::io::ErrorKind::Interrupted => {}
165            Err(e) => return Err(e),
166        }
167    }
168    Ok(true)
169}
170
171// ---- handshake -------------------------------------------------------
172
173fn io_err(msg: String) -> std::io::Error {
174    std::io::Error::new(std::io::ErrorKind::InvalidData, msg)
175}
176
177fn eof() -> std::io::Error {
178    std::io::Error::new(
179        std::io::ErrorKind::UnexpectedEof,
180        "peer closed during noise handshake",
181    )
182}
183
184/// Run the Noise_XX handshake over a connected byte stream. Returns
185/// the established session state and the remote static public key.
186fn handshake<S: Read + Write>(
187    stream: &mut S,
188    identity: &NoiseIdentity,
189    initiator: bool,
190    pinned: Option<[u8; 32]>,
191) -> std::io::Result<(TransportState, [u8; 32])> {
192    let builder = Builder::new(noise_params())
193        .local_private_key(&identity.private)
194        .map_err(|e| io_err(format!("noise local key rejected: {e}")))?;
195    let mut state: HandshakeState = if initiator {
196        builder
197            .build_initiator()
198            .map_err(|e| io_err(format!("noise initiator build: {e}")))?
199    } else {
200        builder
201            .build_responder()
202            .map_err(|e| io_err(format!("noise responder build: {e}")))?
203    };
204
205    let mut buf = vec![0u8; MAX_HANDSHAKE_LEN as usize];
206    while !state.is_handshake_finished() {
207        if state.is_my_turn() {
208            let n = state
209                .write_message(&[], &mut buf)
210                .map_err(|e| io_err(format!("noise handshake write: {e}")))?;
211            write_frame(stream, &buf[..n])?;
212        } else {
213            let frame = read_frame(stream, MAX_HANDSHAKE_LEN)?.ok_or_else(eof)?;
214            state
215                .read_message(&frame, &mut buf)
216                .map_err(|e| io_err(format!("noise handshake read: {e}")))?;
217        }
218    }
219
220    let remote: [u8; 32] = state
221        .get_remote_static()
222        .ok_or_else(|| io_err("noise handshake finished without a remote static key".into()))?
223        .try_into()
224        .expect("noise static keys are 32 bytes");
225
226    if let Some(expected) = pinned {
227        let got = fingerprint_of(&remote);
228        if got != expected {
229            return Err(io_err(format!(
230                "pinned fingerprint mismatch: expected {}, got {}",
231                hex(&expected),
232                hex(&got)
233            )));
234        }
235    }
236
237    state
238        .into_transport_mode()
239        .map(|t| (t, remote))
240        .map_err(|e| io_err(format!("noise transport mode: {e}")))
241}
242
243// ---- transport -------------------------------------------------------
244
245/// An established Noise session over TCP, framed per the
246/// [`confium_net::Transport`] contract: one `send` observed as one
247/// `recv` payload.
248pub struct NoiseTransport {
249    state: TransportState,
250    stream: std::net::TcpStream,
251    remote: [u8; 32],
252}
253
254impl NoiseTransport {
255    /// SHA-256 fingerprint of the authenticated remote static key.
256    pub fn remote_fingerprint(&self) -> [u8; 32] {
257        fingerprint_of(&self.remote)
258    }
259
260    pub(crate) fn connect(params: &NoiseParams) -> Result<Self> {
261        let identity = params
262            .local_key
263            .clone()
264            .unwrap_or_else(NoiseIdentity::generate);
265        let mut stream = std::net::TcpStream::connect(params.addr).context(IoSnafu)?;
266        // A non-noise peer accepts the TCP connection and then never
267        // speaks the handshake; without a deadline the client would
268        // block on the first read forever. 10s bounds a stalled or
269        // mismatched peer; established sessions are not affected.
270        // DeadlineStream polls instead of using SO_RCVTIMEO: under
271        // MRI Ruby on windows-gnu the first recv on a timeout'd
272        // socket fails WSAENOTSOCK (audit ledger).
273        let (state, remote) = {
274            let mut bounded = confium_net::deadline::DeadlineStream::new(
275                &mut stream,
276                std::time::Duration::from_secs(10),
277            )
278            .context(IoSnafu)?;
279            handshake(&mut bounded, &identity, true, params.pinned).context(IoSnafu)?
280        };
281        Ok(Self {
282            state,
283            stream,
284            remote,
285        })
286    }
287
288    pub(crate) fn accept(stream: std::net::TcpStream, params: &NoiseParams) -> Result<Self> {
289        let identity = params
290            .local_key
291            .clone()
292            .unwrap_or_else(NoiseIdentity::generate);
293        let mut stream = stream;
294        let (state, remote) =
295            handshake(&mut stream, &identity, false, params.pinned).context(IoSnafu)?;
296        Ok(Self {
297            state,
298            stream,
299            remote,
300        })
301    }
302}
303
304impl NoiseTransport {
305    fn send_chunk(&mut self, data: &[u8]) -> Result<()> {
306        let mut buf = vec![0u8; data.len() + 128];
307        let n = self
308            .state
309            .write_message(data, &mut buf)
310            .map_err(|e| io_err(format!("noise write: {e}")))
311            .context(IoSnafu)?;
312        write_frame(&mut self.stream, &buf[..n]).context(IoSnafu)
313    }
314
315    fn recv_chunk(&mut self) -> Result<Vec<u8>> {
316        let frame = read_frame(&mut self.stream, MAX_FRAME_LEN + 256)
317            .context(IoSnafu)?
318            .ok_or(ClosedSnafu.build())?;
319        let mut plain = vec![0u8; frame.len()];
320        let n = self
321            .state
322            .read_message(&frame, &mut plain)
323            .map_err(|e| io_err(format!("noise decrypt failed: {e}")))
324            .context(IoSnafu)?;
325        plain.truncate(n);
326        Ok(plain)
327    }
328}
329
330impl Transport for NoiseTransport {
331    fn send(&mut self, data: &[u8]) -> Result<()> {
332        if data.len() > MAX_FRAME_LEN as usize {
333            return Err(IoSnafu.into_error(io_err("payload exceeds frame maximum".into())));
334        }
335        if data.is_empty() {
336            return self.send_chunk(&[0x01]);
337        }
338        let mut chunks = data.chunks(MAX_CHUNK).peekable();
339        while let Some(chunk) = chunks.next() {
340            let final_chunk = chunks.peek().is_none();
341            let mut framed = Vec::with_capacity(chunk.len() + 1);
342            framed.push(if final_chunk { 0x01 } else { 0x00 });
343            framed.extend_from_slice(chunk);
344            self.send_chunk(&framed)?;
345        }
346        Ok(())
347    }
348
349    fn recv(&mut self, buf: &mut [u8]) -> Result<usize> {
350        let mut message = Vec::new();
351        loop {
352            let chunk = self.recv_chunk()?;
353            if chunk.is_empty() {
354                return Err(ClosedSnafu.build());
355            }
356            let (flag, body) = chunk.split_first().expect("chunk has flag byte");
357            message.extend_from_slice(body);
358            if *flag == 0x01 {
359                break;
360            }
361        }
362        let n = message.len().min(buf.len());
363        buf[..n].copy_from_slice(&message[..n]);
364        Ok(n)
365    }
366
367    fn close(&mut self) -> Result<()> {
368        let _ = self.stream.shutdown(std::net::Shutdown::Both);
369        Ok(())
370    }
371}
372
373/// Accepts inbound Noise sessions on a TCP socket.
374pub struct NoiseListener {
375    inner: std::net::TcpListener,
376    params: NoiseParams,
377}
378
379impl NoiseListener {
380    pub(crate) fn new(params: NoiseParams) -> Result<Self> {
381        let inner = std::net::TcpListener::bind(params.addr).context(IoSnafu)?;
382        Ok(Self { inner, params })
383    }
384
385    /// The bound local address (useful when binding port 0 in tests).
386    pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
387        self.inner.local_addr()
388    }
389}
390
391impl Listener for NoiseListener {
392    fn accept(&mut self) -> Result<Box<dyn Transport>> {
393        let (stream, _) = self.inner.accept().context(IoSnafu)?;
394        Ok(Box::new(NoiseTransport::accept(stream, &self.params)?))
395    }
396}