sim_lib_server/transport/
socket.rs1use 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
15pub struct TcpServerTransport {
17 address: ServerAddress,
18 listener: Box<dyn Listener>,
19}
20
21impl TcpServerTransport {
22 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 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}