karyon_jsonrpc/server/
acceptor.rs1use 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#[async_trait]
45pub(super) trait AsyncAcceptor: Send + Sync {
46 async fn accept(&self) -> Result<Box<dyn ByteStream>>;
48 async fn handle(&self, stream: Box<dyn ByteStream>, server: &Arc<Server>) -> Result<()>;
50 fn local_endpoint(&self) -> Result<Endpoint>;
51}
52
53#[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
82pub(super) struct StreamAcceptor<C> {
84 pub(super) listener: Box<dyn StreamListener>,
85 pub(super) codec: C,
86 #[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 #[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#[cfg(feature = "ws")]
134pub(super) struct WsAcceptor<W> {
135 pub(super) listener: Box<dyn StreamListener>,
136 pub(super) layer: Arc<WsLayer<W>>,
137 #[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 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 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 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}