iscp/transport/
websocket.rs1use bytes::BytesMut;
4use reqwest_websocket::{Message, Upgrade};
5use tokio::sync::mpsc;
6use url::Url;
7
8use super::{
9 Certificate, Connector, NegotiationParams, Transport, TransportCloser, TransportError,
10 TransportReader, TransportWriter, UnreliableNotSupported,
11};
12use futures::{SinkExt, StreamExt, stream::SplitStream};
13
14pub type WebSocketTransport = Transport<WebSocketConnector>;
15
16#[derive(Clone, Debug)]
17pub struct WebSocketConnector {
18 url: String,
19 skip_server_verification: bool,
20 client_auth_cert_and_key: Option<(Certificate, Certificate)>,
21}
22
23impl WebSocketConnector {
24 pub fn new<S: ToString>(url: S) -> Self {
25 Self {
26 url: url.to_string(),
27 skip_server_verification: false,
28 client_auth_cert_and_key: None,
29 }
30 }
31
32 pub fn skip_server_verification(mut self, skip_server_verification: bool) -> Self {
33 self.skip_server_verification = skip_server_verification;
34 self
35 }
36
37 pub fn client_auth_cert_and_key<T: Into<Option<(Certificate, Certificate)>>>(
38 mut self,
39 client_auth_cert_and_key: T,
40 ) -> Self {
41 self.client_auth_cert_and_key = client_auth_cert_and_key.into();
42 self
43 }
44}
45
46impl Connector for WebSocketConnector {
47 type Reader = WebSocketReader;
48 type Writer = WebSocketWriter;
49 type Closer = WebSocketCloser;
50 type UnreliableReader = UnreliableNotSupported;
51 type UnreliableWriter = UnreliableNotSupported;
52
53 async fn connect(
54 &self,
55 negotiation_params: NegotiationParams,
56 ) -> Result<Transport<Self>, TransportError> {
57 let url: Url = self.url.parse().map_err(TransportError::new)?;
58
59 let mut url = url.clone();
61 url.set_query(Some(&negotiation_params.to_uri_query_string()?));
62
63 log::debug!("websocket request to {url}");
64
65 let rustls_config = super::tls::rustls_config_client(
66 self.skip_server_verification,
67 self.client_auth_cert_and_key.clone(),
68 )?;
69 let builder = reqwest::Client::builder();
70 let builder = if let Some(rustls_config) = rustls_config {
71 builder.use_preconfigured_tls(rustls_config)
72 } else {
73 builder
74 };
75
76 let client = builder.http1_only().build().map_err(TransportError::new)?;
77 let response = client
78 .get(url.clone())
79 .upgrade()
80 .send()
81 .await
82 .map_err(TransportError::new)?;
83 let websocket = response
84 .into_websocket()
85 .await
86 .map_err(TransportError::new)?;
87
88 log::info!("reqwest websocket connector connected to {url}");
89
90 let (mut sink, stream) = websocket.split();
91 let (tx_write_message, mut rx_write_message) = mpsc::channel(16);
92
93 tokio::spawn(async move {
94 while let Some(msg) = rx_write_message.recv().await {
95 if let Err(e) = sink.send(msg).await {
96 log::debug!("websocket write error {e:?}");
97 break;
98 }
99 if let Err(e) = sink.flush().await {
100 log::debug!("websocket flush error {e:?}");
101 break;
102 }
103 }
104 log::debug!("websocket write loop closed");
105 });
106
107 let writer = WebSocketWriter {
108 tx_write_message: tx_write_message.clone(),
109 };
110
111 let reader = WebSocketReader {
112 stream,
113 tx_write_message,
114 };
115
116 Ok(Transport {
117 reader,
118 writer,
119 closer: WebSocketCloser {},
120 unreliable_reader: None,
121 unreliable_writer: None,
122 })
123 }
124
125 fn _compat_with_context_takeover() -> bool {
126 true
127 }
128}
129
130pub struct WebSocketReader {
131 stream: SplitStream<reqwest_websocket::WebSocket>,
132 tx_write_message: mpsc::Sender<Message>,
133}
134
135pub struct WebSocketWriter {
136 tx_write_message: mpsc::Sender<Message>,
137}
138
139pub struct WebSocketCloser {}
140
141impl TransportReader for WebSocketReader {
142 async fn read(&mut self, buf: &mut BytesMut) -> Result<(), TransportError> {
143 while let Some(res) = self.stream.next().await {
144 let msg = res.map_err(TransportError::new)?;
145 match msg {
146 Message::Close { code, reason } => {
147 log::debug!("receive websocket close frame, code = {code}, reason = {reason}");
148 break;
149 }
150 Message::Ping(p) => {
151 if self
152 .tx_write_message
153 .send_timeout(Message::Pong(p), std::time::Duration::from_secs(10))
154 .await
155 .is_err()
156 {
157 return Err(TransportError::from_msg(
158 "websocket write task closed or timeout, cannot send pong",
159 ));
160 }
161 continue;
162 }
163 Message::Binary(bin) => {
164 buf.extend_from_slice(bin.as_ref());
165 return Ok(());
166 }
167 _ => continue,
168 };
169 }
170 Err(TransportError::from_msg("web socket stream closed"))
171 }
172
173 async fn close(&mut self) -> Result<(), TransportError> {
174 Ok(())
175 }
176}
177
178impl TransportWriter for WebSocketWriter {
179 async fn write(&mut self, data: &[u8]) -> Result<(), TransportError> {
180 let msg = Message::Binary(bytes::Bytes::copy_from_slice(data));
181 self.tx_write_message
182 .send(msg)
183 .await
184 .map_err(|_| TransportError::from_msg("websocket write task closed"))
185 }
186
187 async fn close(&mut self) -> Result<(), TransportError> {
188 self.tx_write_message
189 .send(Message::Close {
190 code: reqwest_websocket::CloseCode::Normal,
191 reason: "OK".into(),
192 })
193 .await
194 .map_err(|_| TransportError::from_msg("websocket write task closed"))
195 }
196}
197
198impl TransportCloser for WebSocketCloser {
199 async fn close(&mut self) -> Result<(), TransportError> {
200 Ok(())
201 }
202}