1use sim_transport_ports::{
2 Datagram, DnsPort, Half, IpcAddress, IpcListener, IpcPort, Listener, Result, SocketAddress,
3 SocketPort, Stream, TransportError, TransportErrorKind,
4};
5use std::{
6 io::{Read, Write},
7 net::{Shutdown, TcpListener, TcpStream, ToSocketAddrs, UdpSocket},
8 time::Duration,
9};
10
11fn native(error: std::io::Error) -> TransportError {
12 let kind = match error.kind() {
13 std::io::ErrorKind::AddrInUse => TransportErrorKind::AddressInUse,
14 std::io::ErrorKind::ConnectionRefused => TransportErrorKind::ConnectionRefused,
15 std::io::ErrorKind::TimedOut => TransportErrorKind::TimedOut,
16 std::io::ErrorKind::WouldBlock => TransportErrorKind::WouldBlock,
17 std::io::ErrorKind::NotFound => TransportErrorKind::NotFound,
18 std::io::ErrorKind::Interrupted => TransportErrorKind::Cancelled,
19 _ => TransportErrorKind::ProviderFault,
20 };
21 let message = error.to_string();
22 drop(error);
23 TransportError::new(kind, message)
24}
25fn socket(address: &SocketAddress) -> std::net::SocketAddr {
26 match address {
27 SocketAddress::Ip { address, port } => (*address, *port).into(),
28 }
29}
30fn wrapped(address: std::net::SocketAddr) -> SocketAddress {
31 SocketAddress::Ip {
32 address: address.ip(),
33 port: address.port(),
34 }
35}
36
37pub struct LinuxSocketPort;
38impl SocketPort for LinuxSocketPort {
39 fn listen_tcp(&self, address: &SocketAddress) -> Result<Box<dyn Listener>> {
40 let listener = TcpListener::bind(socket(address)).map_err(native)?;
41 listener.set_nonblocking(true).map_err(native)?;
42 Ok(Box::new(NativeListener(listener)))
43 }
44 fn connect_tcp(&self, address: &SocketAddress) -> Result<Box<dyn Stream>> {
45 let stream = TcpStream::connect(socket(address)).map_err(native)?;
46 stream.set_nodelay(true).map_err(native)?;
47 Ok(Box::new(NativeStream(stream)))
48 }
49 fn bind_udp(&self, address: &SocketAddress) -> Result<Box<dyn Datagram>> {
50 let socket = UdpSocket::bind(socket(address)).map_err(native)?;
51 socket.set_nonblocking(true).map_err(native)?;
52 Ok(Box::new(NativeDatagram(socket)))
53 }
54}
55struct NativeListener(TcpListener);
56impl Listener for NativeListener {
57 fn local_address(&self) -> Result<SocketAddress> {
58 self.0.local_addr().map(wrapped).map_err(native)
59 }
60 fn accept(&self) -> Result<Option<Box<dyn Stream>>> {
61 match self.0.accept() {
62 Ok((s, _)) => {
63 s.set_nodelay(true).map_err(native)?;
64 Ok(Some(Box::new(NativeStream(s))))
65 }
66 Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(None),
67 Err(e) => Err(native(e)),
68 }
69 }
70 fn close(&self) -> Result<()> {
71 Ok(())
72 }
73}
74struct NativeStream(TcpStream);
75impl Read for NativeStream {
76 fn read(&mut self, b: &mut [u8]) -> std::io::Result<usize> {
77 self.0.read(b)
78 }
79}
80impl Write for NativeStream {
81 fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
82 self.0.write(b)
83 }
84 fn flush(&mut self) -> std::io::Result<()> {
85 self.0.flush()
86 }
87}
88impl Stream for NativeStream {
89 fn set_read_timeout(&self, t: Option<Duration>) -> Result<()> {
90 self.0.set_read_timeout(t).map_err(native)
91 }
92 fn shutdown(&self, h: Half) -> Result<()> {
93 self.0
94 .shutdown(match h {
95 Half::Read => Shutdown::Read,
96 Half::Write => Shutdown::Write,
97 Half::Both => Shutdown::Both,
98 })
99 .map_err(native)
100 }
101}
102struct NativeDatagram(UdpSocket);
103impl Datagram for NativeDatagram {
104 fn local_address(&self) -> Result<SocketAddress> {
105 self.0.local_addr().map(wrapped).map_err(native)
106 }
107 fn send_to(&mut self, b: &[u8], a: &SocketAddress) -> Result<usize> {
108 self.0.send_to(b, socket(a)).map_err(native)
109 }
110 fn recv_from(&mut self, b: &mut [u8]) -> Result<Option<(usize, SocketAddress)>> {
111 match self.0.recv_from(b) {
112 Ok((n, a)) => Ok(Some((n, wrapped(a)))),
113 Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(None),
114 Err(e) => Err(native(e)),
115 }
116 }
117 fn close(&self) -> Result<()> {
118 Ok(())
119 }
120}
121pub struct LinuxDnsPort;
122impl DnsPort for LinuxDnsPort {
123 fn resolve(&self, host: &str, port: u16) -> Result<Vec<SocketAddress>> {
124 (host, port)
125 .to_socket_addrs()
126 .map(|v| v.map(wrapped).collect())
127 .map_err(|e| TransportError::new(TransportErrorKind::DnsFailure, e.to_string()))
128 }
129}
130
131pub struct LinuxIpcPort;
133
134pub fn bind_transport_services() -> Result<()> {
141 sim_transport_ports::bind_services(sim_transport_ports::TransportServices {
142 sockets: std::sync::Arc::new(LinuxSocketPort),
143 dns: std::sync::Arc::new(LinuxDnsPort),
144 ipc: Some(std::sync::Arc::new(LinuxIpcPort)),
145 })
146}
147#[cfg(unix)]
148mod unix {
149 use super::{Duration, Half, IpcListener, Read, Result, Shutdown, Stream, Write, native};
150 use std::os::unix::net::{UnixListener, UnixStream};
151 pub(super) struct UListener(pub UnixListener);
152 pub(super) struct UStream(pub UnixStream);
153 impl IpcListener for UListener {
154 fn accept(&self) -> Result<Option<Box<dyn Stream>>> {
155 match self.0.accept() {
156 Ok((s, _)) => Ok(Some(Box::new(UStream(s)))),
157 Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(None),
158 Err(e) => Err(native(e)),
159 }
160 }
161 fn close(&self) -> Result<()> {
162 Ok(())
163 }
164 }
165 impl Read for UStream {
166 fn read(&mut self, b: &mut [u8]) -> std::io::Result<usize> {
167 self.0.read(b)
168 }
169 }
170 impl Write for UStream {
171 fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
172 self.0.write(b)
173 }
174 fn flush(&mut self) -> std::io::Result<()> {
175 self.0.flush()
176 }
177 }
178 impl Stream for UStream {
179 fn set_read_timeout(&self, t: Option<Duration>) -> Result<()> {
180 self.0.set_read_timeout(t).map_err(native)
181 }
182 fn shutdown(&self, h: Half) -> Result<()> {
183 self.0
184 .shutdown(match h {
185 Half::Read => Shutdown::Read,
186 Half::Write => Shutdown::Write,
187 Half::Both => Shutdown::Both,
188 })
189 .map_err(native)
190 }
191 }
192}
193impl IpcPort for LinuxIpcPort {
194 fn listen(&self, address: &IpcAddress) -> Result<Box<dyn IpcListener>> {
195 match address {
196 #[cfg(unix)]
197 IpcAddress::UnixPath(path) => {
198 let l = std::os::unix::net::UnixListener::bind(path).map_err(native)?;
199 l.set_nonblocking(true).map_err(native)?;
200 Ok(Box::new(unix::UListener(l)))
201 }
202 _ => Err(TransportError::new(
203 TransportErrorKind::Unsupported,
204 "Linux IPC requires a UnixPath",
205 )),
206 }
207 }
208 fn connect(&self, address: &IpcAddress) -> Result<Box<dyn Stream>> {
209 match address {
210 #[cfg(unix)]
211 IpcAddress::UnixPath(path) => std::os::unix::net::UnixStream::connect(path)
212 .map(|s| Box::new(unix::UStream(s)) as Box<dyn Stream>)
213 .map_err(native),
214 _ => Err(TransportError::new(
215 TransportErrorKind::Unsupported,
216 "Linux IPC requires a UnixPath",
217 )),
218 }
219 }
220}