Skip to main content

confium_net_tcp/
listener.rs

1//! Listening endpoint for `tcp://`, `tcp4://`, `tcp6://` URLs.
2//!
3//! [`TcpListener`] wraps [`std::net::TcpListener`] and implements
4//! [`confium_net::Listener`]. Each [`accept`](confium_net::Listener::accept)
5//! yields a [`crate::TcpTransport`] wrapping the accepted
6//! [`std::net::TcpStream`], so accepted peers speak the same
7//! length-prefixed framing as dial-out peers.
8
9use std::net::IpAddr;
10use std::net::Ipv4Addr;
11use std::net::Ipv6Addr;
12use std::net::SocketAddr;
13use std::net::TcpListener as StdTcpListener;
14
15use confium_net::Listener;
16use confium_net::Result;
17use confium_net::Transport;
18use confium_net::error::ClosedSnafu;
19
20use crate::TcpTransport;
21use crate::transport::address_family;
22
23/// Listening endpoint for inbound TCP connections.
24///
25/// Holding this alive keeps the bound socket open; dropping it closes
26/// the socket (the OS releases the port).
27pub struct TcpListener {
28    inner: Option<StdTcpListener>,
29}
30
31impl TcpListener {
32    /// Bind a new listener at `host:port`, honoring the address-family
33    /// hint `tcp4` / `tcp6` / `tcp` encoded in `scheme`. Passing port
34    /// `0` requests an ephemeral port from the OS; the caller can read
35    /// the assigned port back via [`local_addr`](Self::local_addr).
36    pub fn bind(scheme: &str, host: &str, port: u16) -> std::io::Result<Self> {
37        let listener = match address_family(scheme) {
38            Some(false) => {
39                // tcp4: bind IPv4 only. `host` may be a literal IPv4 or
40                // the wildcard `0.0.0.0`.
41                let ip: Ipv4Addr = parse_ipv4(host)?;
42                StdTcpListener::bind(SocketAddr::new(IpAddr::V4(ip), port))?
43            }
44            Some(true) => {
45                // tcp6: bind IPv6 only.
46                let ip: Ipv6Addr = parse_ipv6(host)?;
47                StdTcpListener::bind(SocketAddr::new(IpAddr::V6(ip), port))?
48            }
49            None => {
50                // tcp: let the standard library resolve `host:port`
51                // (literal IP, `0.0.0.0`, `[::]`, or DNS name). Bind to
52                // the first working address.
53                StdTcpListener::bind((host, port))?
54            }
55        };
56        listener.set_nonblocking(false).ok();
57        Ok(Self {
58            inner: Some(listener),
59        })
60    }
61
62    /// The locally-bound socket address. Useful for reading back the
63    /// OS-assigned ephemeral port after binding with port `0`.
64    pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
65        self.inner
66            .as_ref()
67            .ok_or_else(|| std::io::Error::from(std::io::ErrorKind::NotConnected))?
68            .local_addr()
69    }
70}
71
72impl Listener for TcpListener {
73    fn accept(&mut self) -> Result<Box<dyn Transport>> {
74        let listener = match &self.inner {
75            Some(l) => l,
76            None => return ClosedSnafu.fail(),
77        };
78        let (stream, _peer) = listener.accept().map_err(crate::transport::io_to_closed)?;
79        // Disable Nagle on accepted streams for the same latency
80        // reasons as dial-out peers (see [`TcpTransport::connect`]).
81        stream.set_nodelay(true).ok();
82        Ok(Box::new(TcpTransport::from_stream(stream)))
83    }
84}
85
86impl Drop for TcpListener {
87    fn drop(&mut self) {
88        // Dropping the inner listener closes the socket; take() is for
89        // explicitness so a future `close` method can report errors.
90        self.inner.take();
91    }
92}
93
94fn parse_ipv4(host: &str) -> std::io::Result<Ipv4Addr> {
95    host.parse::<Ipv4Addr>().map_err(|_| {
96        std::io::Error::new(
97            std::io::ErrorKind::InvalidInput,
98            format!("invalid IPv4 bind address '{host}'"),
99        )
100    })
101}
102
103fn parse_ipv6(host: &str) -> std::io::Result<Ipv6Addr> {
104    host.parse::<Ipv6Addr>().map_err(|_| {
105        std::io::Error::new(
106            std::io::ErrorKind::InvalidInput,
107            format!("invalid IPv6 bind address '{host}'"),
108        )
109    })
110}