Skip to main content

karyon_jsonrpc/server/
acceptor.rs

1//! Acceptors used by the stream-based and WebSocket backends.
2//! Accepting is split in two phases so the accept loop never blocks on
3//! a handshake: `accept` does the kernel accept, `handle` runs the
4//! TLS/WS handshake and hands the split halves to the server. The loop
5//! runs `handle` off to a separate task.
6
7use std::sync::Arc;
8
9use async_trait::async_trait;
10
11use karyon_net::{framed, ByteStream, Endpoint};
12
13use crate::{
14    codec::JsonRpcCodec,
15    error::{Error, Result},
16    server::Server,
17};
18
19#[cfg(any(feature = "tls", feature = "ws"))]
20use std::time::Duration;
21
22#[cfg(any(feature = "tls", feature = "ws"))]
23use karyon_core::async_util::timeout;
24
25#[cfg(feature = "tls")]
26use karyon_net::tls::TlsLayer;
27
28#[cfg(feature = "ws")]
29use std::net::SocketAddr;
30
31#[cfg(feature = "ws")]
32use karyon_net::{
33    layers::ws::{WsConn, WsLayer},
34    Error as NetError,
35};
36
37#[cfg(any(feature = "tls", feature = "ws"))]
38use karyon_net::ServerLayer;
39
40#[cfg(feature = "ws")]
41use crate::codec::JsonRpcWsCodec;
42
43/// Accepts raw streams and upgrades/handles them.
44#[async_trait]
45pub(super) trait AsyncAcceptor: Send + Sync {
46    /// Kernel accept only, before any handshake.
47    async fn accept(&self) -> Result<Box<dyn ByteStream>>;
48    /// Upgrade the stream, frame it, and hand it to the server.
49    async fn handle(&self, stream: Box<dyn ByteStream>, server: &Arc<Server>) -> Result<()>;
50    fn local_endpoint(&self) -> Result<Endpoint>;
51}
52
53/// A listener that produces byte streams.
54#[async_trait]
55pub(super) trait StreamListener: Send + Sync {
56    async fn accept(&self) -> karyon_net::Result<Box<dyn ByteStream>>;
57    fn local_endpoint(&self) -> karyon_net::Result<Endpoint>;
58}
59
60#[cfg(feature = "tcp")]
61#[async_trait]
62impl StreamListener for karyon_net::tcp::TcpListener {
63    async fn accept(&self) -> karyon_net::Result<Box<dyn ByteStream>> {
64        self.accept().await
65    }
66    fn local_endpoint(&self) -> karyon_net::Result<Endpoint> {
67        self.local_endpoint()
68    }
69}
70
71#[cfg(all(feature = "unix", target_family = "unix"))]
72#[async_trait]
73impl StreamListener for karyon_net::unix::UnixListener {
74    async fn accept(&self) -> karyon_net::Result<Box<dyn ByteStream>> {
75        self.accept().await
76    }
77    fn local_endpoint(&self) -> karyon_net::Result<Endpoint> {
78        self.local_endpoint()
79    }
80}
81
82/// Byte-stream acceptor, with an optional TLS handshake.
83pub(super) struct StreamAcceptor<C> {
84    pub(super) listener: Box<dyn StreamListener>,
85    pub(super) codec: C,
86    /// Set for `tls://` endpoints; applied in `handle`.
87    #[cfg(feature = "tls")]
88    pub(super) tls: Option<TlsLayer>,
89    #[cfg(feature = "tls")]
90    pub(super) handshake_timeout: Duration,
91}
92
93#[async_trait]
94impl<C> AsyncAcceptor for StreamAcceptor<C>
95where
96    C: JsonRpcCodec,
97{
98    async fn accept(&self) -> Result<Box<dyn ByteStream>> {
99        self.listener.accept().await.map_err(Error::from)
100    }
101
102    async fn handle(&self, stream: Box<dyn ByteStream>, server: &Arc<Server>) -> Result<()> {
103        #[cfg(feature = "tls")]
104        let stream = match &self.tls {
105            Some(layer) => {
106                timeout(
107                    self.handshake_timeout,
108                    ServerLayer::handshake(layer, stream),
109                )
110                .await??
111            }
112            None => stream,
113        };
114        let conn = framed(stream, self.codec.clone());
115        let peer = conn.peer_endpoint();
116        let (reader, writer) = conn.split();
117        server.handle_message_conn(reader, writer, peer);
118        Ok(())
119    }
120
121    fn local_endpoint(&self) -> Result<Endpoint> {
122        let ep = self.listener.local_endpoint().map_err(Error::from)?;
123        // The listener is plain TCP; report the TLS scheme.
124        #[cfg(feature = "tls")]
125        if self.tls.is_some() {
126            return Ok(Endpoint::Tls(ep.addr()?, ep.port()?));
127        }
128        Ok(ep)
129    }
130}
131
132/// WebSocket acceptor, with an optional TLS handshake for `wss://`.
133#[cfg(feature = "ws")]
134pub(super) struct WsAcceptor<W> {
135    pub(super) listener: Box<dyn StreamListener>,
136    pub(super) layer: Arc<WsLayer<W>>,
137    /// Set for `wss://` endpoints; applied before the WS handshake.
138    #[cfg(feature = "tls")]
139    pub(super) tls: Option<TlsLayer>,
140    pub(super) handshake_timeout: Duration,
141}
142
143#[cfg(feature = "ws")]
144impl<W> WsAcceptor<W>
145where
146    W: JsonRpcWsCodec,
147{
148    /// Runs the TLS handshake (for `wss://`) then the WS handshake.
149    async fn upgrade(&self, stream: Box<dyn ByteStream>) -> Result<WsConn<W>> {
150        #[cfg(feature = "tls")]
151        let stream = match &self.tls {
152            Some(layer) => ServerLayer::handshake(layer, stream).await?,
153            None => stream,
154        };
155        let conn = ServerLayer::handshake(self.layer.as_ref(), stream).await?;
156        Ok(conn)
157    }
158}
159
160#[cfg(feature = "ws")]
161#[async_trait]
162impl<W> AsyncAcceptor for WsAcceptor<W>
163where
164    W: JsonRpcWsCodec,
165{
166    async fn accept(&self) -> Result<Box<dyn ByteStream>> {
167        self.listener.accept().await.map_err(Error::from)
168    }
169
170    async fn handle(&self, stream: Box<dyn ByteStream>, server: &Arc<Server>) -> Result<()> {
171        // One budget for both handshakes.
172        let conn = timeout(self.handshake_timeout, self.upgrade(stream)).await??;
173        let peer = conn.peer_endpoint();
174        let (reader, writer) = conn.split();
175        server.handle_message_conn(reader, writer, peer);
176        Ok(())
177    }
178
179    fn local_endpoint(&self) -> Result<Endpoint> {
180        // The listener reports `tcp://...`; rewrite to the WS scheme so
181        // a client building from this endpoint runs the WS handshake.
182        let inner = self.listener.local_endpoint().map_err(Error::from)?;
183        let addr = SocketAddr::try_from(inner.clone()).map_err(Error::from)?;
184        #[cfg(feature = "tls")]
185        let scheme = if self.tls.is_some() { "wss" } else { "ws" };
186        #[cfg(not(feature = "tls"))]
187        let scheme = "ws";
188        format!("{scheme}://{addr}/")
189            .parse()
190            .map_err(|e: NetError| Error::from(e))
191    }
192}