Skip to main content

iscp/transport/
websocket.rs

1//! WebSocketトランスポートの実装
2
3use 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        // Connect
60        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}