1use 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
27const MAX_FRAME_LEN: u32 = 8 * 1024 * 1024;
31
32const MAX_HANDSHAKE_LEN: u32 = 65535;
35
36const MAX_CHUNK: usize = 60_000;
42
43pub(crate) struct NoiseParams {
45 pub addr: SocketAddr,
46 pub local_key: Option<NoiseIdentity>,
48 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
128fn 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
171fn 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
184fn 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
243pub struct NoiseTransport {
249 state: TransportState,
250 stream: std::net::TcpStream,
251 remote: [u8; 32],
252}
253
254impl NoiseTransport {
255 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 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
373pub 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 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}