Skip to main content

sim_lib_server/transport/
socket.rs

1use std::{sync::Arc, time::Duration};
2
3use sim_transport_ports::{Half, IpcAddress, IpcListener, Listener, SocketAddress, Stream};
4
5use sim_kernel::{Cx, Error, Result, Symbol};
6
7use crate::{EvalSite, FrameKind, ServerAddress, ServerFrame, ServerRuntime};
8
9use super::{
10    ConnectionTransport, SERVER_CONNECTION_IO_TIMEOUT_MS, ServerTransport, answer_or_negotiate,
11    bound_transport_services, error_frame_from_error, is_timeout, read_frame_from,
12    update_negotiated_codec_from_reply, write_frame_to,
13};
14
15/// TCP listener transport for server-frame connections.
16pub struct TcpServerTransport {
17    address: ServerAddress,
18    listener: Box<dyn Listener>,
19}
20
21impl TcpServerTransport {
22    /// Binds a TCP listener to `address`.
23    pub fn bind(address: ServerAddress) -> Result<Self> {
24        let ServerAddress::Tcp { host, port } = &address else {
25            return Err(Error::Eval(
26                "tcp transport requires a tcp address".to_owned(),
27            ));
28        };
29        let ports = bound_transport_services().map_err(port_error)?;
30        let resolved = ports.dns.resolve(host, *port).map_err(port_error)?;
31        let target = resolved
32            .first()
33            .ok_or_else(|| Error::HostError("DNS returned no addresses".to_owned()))?;
34        let listener = ports.sockets.listen_tcp(target).map_err(port_error)?;
35        let local_addr = listener.local_address().map_err(port_error)?;
36        let SocketAddress::Ip {
37            port: local_port, ..
38        } = local_addr;
39        Ok(Self {
40            address: ServerAddress::Tcp {
41                host: host.clone(),
42                port: local_port,
43            },
44            listener,
45        })
46    }
47
48    #[cfg_attr(not(test), allow(dead_code))]
49    /// Returns the bound local port.
50    pub fn local_port(&self) -> Result<u16> {
51        let SocketAddress::Ip { port, .. } = self.listener.local_address().map_err(port_error)?;
52        Ok(port)
53    }
54}
55
56impl ServerTransport for TcpServerTransport {
57    fn address(&self) -> &ServerAddress {
58        &self.address
59    }
60
61    fn accept(&self, cx: &mut Cx) -> Result<Box<dyn ConnectionTransport>> {
62        loop {
63            if let Some(connection) = self.accept_timeout(cx, Duration::from_millis(25))? {
64                return Ok(connection);
65            }
66        }
67    }
68
69    fn shutdown(&self, _cx: &mut Cx) -> Result<()> {
70        Ok(())
71    }
72
73    fn accept_timeout(
74        &self,
75        _cx: &mut Cx,
76        _timeout: Duration,
77    ) -> Result<Option<Box<dyn ConnectionTransport>>> {
78        self.listener
79            .accept()
80            .map(|stream| {
81                stream.map(|stream| {
82                    Box::new(TcpConnectionTransport::server_side(stream))
83                        as Box<dyn ConnectionTransport>
84                })
85            })
86            .map_err(port_error)
87    }
88}
89
90pub struct TcpConnectionTransport {
91    stream: Box<dyn Stream>,
92}
93
94impl TcpConnectionTransport {
95    pub fn connect(address: &ServerAddress) -> Result<Self> {
96        let ServerAddress::Tcp { host, port } = address else {
97            return Err(Error::Eval("tcp connect requires a tcp address".to_owned()));
98        };
99        let ports = bound_transport_services().map_err(port_error)?;
100        let resolved = ports.dns.resolve(host, *port).map_err(port_error)?;
101        let target = resolved
102            .first()
103            .ok_or_else(|| Error::HostError("DNS returned no addresses".to_owned()))?;
104        let stream = ports.sockets.connect_tcp(target).map_err(port_error)?;
105        Ok(Self { stream })
106    }
107
108    fn server_side(stream: Box<dyn Stream>) -> Self {
109        Self { stream }
110    }
111
112    fn serve(&mut self, runtime: &Arc<ServerRuntime>, site: &Arc<dyn EvalSite>) -> Result<()> {
113        let session_id = runtime.open_session(
114            Symbol::qualified("codec", "binary"),
115            runtime.session_isolation().clone(),
116        )?;
117        let mut inflight = 0usize;
118        loop {
119            if runtime.is_stopping() {
120                let _ = runtime.close_session(session_id);
121                return Ok(());
122            }
123
124            let frame = match self.recv_frame_for_serve() {
125                Ok(Some(frame)) => frame,
126                Ok(None) => continue,
127                Err(error) => {
128                    let _ = runtime.close_session(session_id);
129                    return Err(error);
130                }
131            };
132            let Some(frame) = frame else {
133                let _ = runtime.close_session(session_id);
134                return Ok(());
135            };
136            runtime.note_message_received();
137            if runtime.is_stopping() {
138                let _ = runtime.close_session(session_id);
139                return Ok(());
140            }
141            if matches!(frame.kind, FrameKind::Request | FrameKind::Notify)
142                && inflight >= runtime.max_inflight()
143            {
144                let reply = runtime.with_cx(|cx| {
145                    error_frame_from_error(
146                        cx,
147                        &frame,
148                        &Error::Eval(format!(
149                            "connection max-inflight {} exceeded",
150                            runtime.max_inflight()
151                        )),
152                    )
153                })?;
154                write_frame_to(&mut self.stream, &reply)?;
155                runtime.note_message_sent();
156                continue;
157            }
158            if matches!(frame.kind, FrameKind::Request | FrameKind::Notify) {
159                inflight = inflight.saturating_add(1);
160            }
161            let reply = match runtime.with_cx(|cx| answer_or_negotiate(cx, site, frame.clone())) {
162                Ok(reply) => {
163                    update_negotiated_codec_from_reply(runtime, session_id, &frame, &reply)?;
164                    reply
165                }
166                Err(error) => runtime.with_cx(|cx| error_frame_from_error(cx, &frame, &error))?,
167            };
168            if runtime.is_stopping() {
169                let _ = runtime.close_session(session_id);
170                return Ok(());
171            }
172            write_frame_to(&mut self.stream, &reply)?;
173            runtime.note_message_sent();
174            if matches!(frame.kind, FrameKind::Request | FrameKind::Notify) {
175                inflight = inflight.saturating_sub(1);
176            }
177        }
178    }
179
180    fn recv_frame_for_serve(&mut self) -> Result<Option<Option<ServerFrame>>> {
181        self.stream
182            .set_read_timeout(Some(Duration::from_millis(SERVER_CONNECTION_IO_TIMEOUT_MS)))
183            .map_err(port_error)?;
184        match read_frame_from(&mut self.stream) {
185            Ok(frame) => Ok(Some(frame)),
186            Err(error) if is_timeout(&error) => Ok(None),
187            Err(error) => Err(error),
188        }
189    }
190}
191
192impl ConnectionTransport for TcpConnectionTransport {
193    fn send_frame(&mut self, _cx: &mut Cx, frame: ServerFrame) -> Result<()> {
194        write_frame_to(&mut self.stream, &frame)
195    }
196
197    fn recv_frame(
198        &mut self,
199        _cx: &mut Cx,
200        timeout: Option<Duration>,
201    ) -> Result<Option<ServerFrame>> {
202        self.stream.set_read_timeout(timeout).map_err(port_error)?;
203        match read_frame_from(&mut self.stream) {
204            Ok(frame) => Ok(frame),
205            Err(error) if is_timeout(&error) => Ok(None),
206            Err(error) => Err(error),
207        }
208    }
209
210    fn close(&mut self, _cx: &mut Cx) -> Result<()> {
211        let _ = self.stream.shutdown(Half::Both);
212        Ok(())
213    }
214
215    fn as_any(&self) -> &dyn std::any::Any {
216        self
217    }
218
219    fn serve_connection(
220        &mut self,
221        runtime: &Arc<ServerRuntime>,
222        site: &Arc<dyn EvalSite>,
223    ) -> Result<()> {
224        self.serve(runtime, site)
225    }
226}
227
228#[cfg(unix)]
229pub struct UnixServerTransport {
230    address: ServerAddress,
231    listener: Box<dyn IpcListener>,
232}
233
234#[cfg(unix)]
235impl UnixServerTransport {
236    pub fn bind(address: ServerAddress) -> Result<Self> {
237        let ServerAddress::Unix { path } = &address else {
238            return Err(Error::Eval(
239                "unix transport requires a unix address".to_owned(),
240            ));
241        };
242        let listener = bound_transport_services()
243            .map_err(port_error)?
244            .ipc
245            .ok_or_else(|| Error::HostError("local IPC service is unavailable".to_owned()))?
246            .listen(&IpcAddress::UnixPath(path.clone()))
247            .map_err(port_error)?;
248        Ok(Self { address, listener })
249    }
250}
251
252#[cfg(unix)]
253impl ServerTransport for UnixServerTransport {
254    fn address(&self) -> &ServerAddress {
255        &self.address
256    }
257
258    fn accept(&self, cx: &mut Cx) -> Result<Box<dyn ConnectionTransport>> {
259        loop {
260            if let Some(connection) = self.accept_timeout(cx, Duration::from_millis(25))? {
261                return Ok(connection);
262            }
263        }
264    }
265
266    fn shutdown(&self, _cx: &mut Cx) -> Result<()> {
267        self.listener.close().map_err(port_error)
268    }
269
270    fn accept_timeout(
271        &self,
272        _cx: &mut Cx,
273        _timeout: Duration,
274    ) -> Result<Option<Box<dyn ConnectionTransport>>> {
275        self.listener
276            .accept()
277            .map(|stream| {
278                stream.map(|stream| {
279                    Box::new(UnixConnectionTransport::server_side(stream))
280                        as Box<dyn ConnectionTransport>
281                })
282            })
283            .map_err(port_error)
284    }
285}
286
287#[cfg(unix)]
288pub struct UnixConnectionTransport {
289    stream: Box<dyn Stream>,
290}
291
292#[cfg(unix)]
293impl UnixConnectionTransport {
294    pub fn connect(address: &ServerAddress) -> Result<Self> {
295        let ServerAddress::Unix { path } = address else {
296            return Err(Error::Eval(
297                "unix connect requires a unix address".to_owned(),
298            ));
299        };
300        let stream = bound_transport_services()
301            .map_err(port_error)?
302            .ipc
303            .ok_or_else(|| Error::HostError("local IPC service is unavailable".to_owned()))?
304            .connect(&IpcAddress::UnixPath(path.clone()))
305            .map_err(port_error)?;
306        Ok(Self { stream })
307    }
308
309    fn server_side(stream: Box<dyn Stream>) -> Self {
310        Self { stream }
311    }
312
313    fn serve(&mut self, runtime: &Arc<ServerRuntime>, site: &Arc<dyn EvalSite>) -> Result<()> {
314        let session_id = runtime.open_session(
315            Symbol::qualified("codec", "binary"),
316            runtime.session_isolation().clone(),
317        )?;
318        let mut inflight = 0usize;
319        loop {
320            if runtime.is_stopping() {
321                let _ = runtime.close_session(session_id);
322                return Ok(());
323            }
324
325            let frame = match self.recv_frame_for_serve() {
326                Ok(Some(frame)) => frame,
327                Ok(None) => continue,
328                Err(error) => {
329                    let _ = runtime.close_session(session_id);
330                    return Err(error);
331                }
332            };
333            let Some(frame) = frame else {
334                let _ = runtime.close_session(session_id);
335                return Ok(());
336            };
337            runtime.note_message_received();
338            if runtime.is_stopping() {
339                let _ = runtime.close_session(session_id);
340                return Ok(());
341            }
342            if matches!(frame.kind, FrameKind::Request | FrameKind::Notify)
343                && inflight >= runtime.max_inflight()
344            {
345                let reply = runtime.with_cx(|cx| {
346                    error_frame_from_error(
347                        cx,
348                        &frame,
349                        &Error::Eval(format!(
350                            "connection max-inflight {} exceeded",
351                            runtime.max_inflight()
352                        )),
353                    )
354                })?;
355                write_frame_to(&mut self.stream, &reply)?;
356                runtime.note_message_sent();
357                continue;
358            }
359            if matches!(frame.kind, FrameKind::Request | FrameKind::Notify) {
360                inflight = inflight.saturating_add(1);
361            }
362            let reply = match runtime.with_cx(|cx| answer_or_negotiate(cx, site, frame.clone())) {
363                Ok(reply) => {
364                    update_negotiated_codec_from_reply(runtime, session_id, &frame, &reply)?;
365                    reply
366                }
367                Err(error) => runtime.with_cx(|cx| error_frame_from_error(cx, &frame, &error))?,
368            };
369            if runtime.is_stopping() {
370                let _ = runtime.close_session(session_id);
371                return Ok(());
372            }
373            write_frame_to(&mut self.stream, &reply)?;
374            runtime.note_message_sent();
375            if matches!(frame.kind, FrameKind::Request | FrameKind::Notify) {
376                inflight = inflight.saturating_sub(1);
377            }
378        }
379    }
380
381    fn recv_frame_for_serve(&mut self) -> Result<Option<Option<ServerFrame>>> {
382        self.stream
383            .set_read_timeout(Some(Duration::from_millis(SERVER_CONNECTION_IO_TIMEOUT_MS)))
384            .map_err(port_error)?;
385        match read_frame_from(&mut self.stream) {
386            Ok(frame) => Ok(Some(frame)),
387            Err(error) if is_timeout(&error) => Ok(None),
388            Err(error) => Err(error),
389        }
390    }
391}
392
393#[cfg(unix)]
394impl ConnectionTransport for UnixConnectionTransport {
395    fn send_frame(&mut self, _cx: &mut Cx, frame: ServerFrame) -> Result<()> {
396        write_frame_to(&mut self.stream, &frame)
397    }
398
399    fn recv_frame(
400        &mut self,
401        _cx: &mut Cx,
402        timeout: Option<Duration>,
403    ) -> Result<Option<ServerFrame>> {
404        self.stream.set_read_timeout(timeout).map_err(port_error)?;
405        match read_frame_from(&mut self.stream) {
406            Ok(frame) => Ok(frame),
407            Err(error) if is_timeout(&error) => Ok(None),
408            Err(error) => Err(error),
409        }
410    }
411
412    fn close(&mut self, _cx: &mut Cx) -> Result<()> {
413        Ok(())
414    }
415
416    fn as_any(&self) -> &dyn std::any::Any {
417        self
418    }
419
420    fn serve_connection(
421        &mut self,
422        runtime: &Arc<ServerRuntime>,
423        site: &Arc<dyn EvalSite>,
424    ) -> Result<()> {
425        self.serve(runtime, site)
426    }
427}
428
429fn port_error(error: sim_transport_ports::TransportError) -> Error {
430    Error::HostError(error.to_string())
431}