Skip to main content

iscp/wire/
mod.rs

1//! Wire層の定義
2
3mod read_loop;
4
5use std::{
6    sync::{Arc, Mutex},
7    time::Duration,
8};
9
10use bytes::BytesMut;
11use crossbeam::atomic::AtomicCell;
12use tokio::sync::{broadcast, mpsc, oneshot};
13use tokio_util::sync::CancellationToken;
14
15use crate::{
16    encoding::{Encoder, Encoding},
17    error::Error,
18    internal::{Waiter, timeout_with_ct},
19    message::{
20        DownstreamCall, HasRequestId, HasResultCode, Message, RequestMessage, UpstreamCallAck,
21    },
22    transport::{
23        Compression, Compressor, Connector, Transport, TransportCloser, TransportReader,
24        TransportWriter,
25    },
26};
27
28pub(crate) use read_loop::*;
29
30#[derive(Clone)]
31pub struct Conn {
32    inner: Arc<ConnInner>,
33}
34
35impl std::fmt::Debug for Conn {
36    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37        f.debug_struct("Conn").finish()
38    }
39}
40
41struct ConnInner {
42    tx_write_message: mpsc::Sender<Message>,
43    tx_unreliable_write_message: Option<mpsc::Sender<Message>>,
44    tx_read_loop_command: mpsc::UnboundedSender<read_loop::ReadLoopCommand>,
45    tx_unreliable_read_loop_command: Option<mpsc::UnboundedSender<read_loop::ReadLoopCommand>>,
46    rx_downstream_call: broadcast::Receiver<DownstreamCall>,
47    response_senders: Mutex<fnv::FnvHashMap<u32, oneshot::Sender<Message>>>,
48    request_id_counter: RequestIdCounter,
49    ct: CancellationToken,
50    waiter: Waiter,
51    waiter_rw: Waiter,
52    response_message_timeout: Duration,
53    downstream_stream_id_alias_counter: AtomicCell<u32>,
54    channel_size: usize,
55}
56
57impl Conn {
58    pub(crate) fn new<C: Connector>(
59        transport: Transport<C>,
60        encoding: Encoding,
61        compression: Compression,
62        channel_size: usize,
63        timeout: Duration,
64        ping_interval: Duration,
65        ping_timeout: Duration,
66    ) -> Self {
67        let Transport {
68            mut reader,
69            mut writer,
70            closer,
71            unreliable_reader,
72            unreliable_writer,
73        } = transport;
74
75        let (tx_write_message, rx_write_message) = mpsc::channel(channel_size);
76        let (tx_pong, rx_pong) = mpsc::channel(channel_size);
77        // `channel_size` (default 1024), not 1: at capacity 1 a slow receiver loses
78        // every DownstreamCall but the newest, dropping reply calls.
79        let (tx_downstream_call, rx_downstream_call) = broadcast::channel(channel_size);
80        let (tx_read_loop_command, rx_read_loop_command) = mpsc::unbounded_channel();
81        let (waiter, wg) = Waiter::new();
82        let (waiter_rw, wg_rw) = Waiter::new();
83        let Encoding {
84            encoder, decoder, ..
85        } = encoding;
86        let (compressor, extractor) = compression.converters();
87        let (compressor_for_unreliable, extractor_for_unreliable) = compression.converters();
88
89        let (tx_unreliable_write_message, rx_unreliable_write_message) =
90            if unreliable_writer.is_some() {
91                let (tx, rx) = mpsc::channel(channel_size);
92                (Some(tx), Some(rx))
93            } else {
94                (None, None)
95            };
96        let (tx_unreliable_read_loop_command, rx_unreliable_read_loop_command) =
97            if unreliable_reader.is_some() {
98                let (tx, rx) = mpsc::unbounded_channel();
99                (Some(tx), Some(rx))
100            } else {
101                (None, None)
102            };
103
104        let inner = Arc::new(ConnInner {
105            tx_write_message,
106            tx_unreliable_write_message,
107            tx_read_loop_command,
108            tx_unreliable_read_loop_command,
109            rx_downstream_call,
110            response_senders: Mutex::new(fnv::FnvHashMap::default()),
111            request_id_counter: RequestIdCounter::new(),
112            ct: CancellationToken::new(),
113            waiter,
114            waiter_rw,
115            response_message_timeout: timeout,
116            downstream_stream_id_alias_counter: AtomicCell::new(1),
117            channel_size,
118        });
119
120        // Spawn unreliable reader task
121        if let Some(mut unreliable_reader) = unreliable_reader {
122            let d = decoder.clone();
123            let inner_clone = inner.clone();
124            let wg_rw_clone = wg_rw.clone();
125            let tx_pong_clone = tx_pong.clone();
126            let tx_downstream_call_clone = tx_downstream_call.clone();
127            tokio::spawn(async move {
128                read_loop::read_loop(
129                    inner_clone,
130                    &mut unreliable_reader,
131                    rx_unreliable_read_loop_command.unwrap(),
132                    tx_pong_clone,
133                    tx_downstream_call_clone,
134                    d,
135                    extractor_for_unreliable,
136                )
137                .await;
138                if let Err(e) = unreliable_reader.close().await {
139                    log::error!("transport reader close error {e}");
140                }
141                log::debug!("exit read loop");
142                std::mem::drop(wg_rw_clone);
143            });
144        }
145
146        // Spawn reader task
147        let inner_clone = inner.clone();
148        let wg_rw_clone = wg_rw.clone();
149        let d = decoder.clone();
150        tokio::spawn(async move {
151            read_loop::read_loop(
152                inner_clone,
153                &mut reader,
154                rx_read_loop_command,
155                tx_pong,
156                tx_downstream_call,
157                d,
158                extractor,
159            )
160            .await;
161            if let Err(e) = reader.close().await {
162                log::error!("transport reader close error {e}");
163            }
164            log::debug!("exit read loop");
165            std::mem::drop(wg_rw_clone);
166        });
167
168        // Spawn writer task
169        let inner_clone = inner.clone();
170        let e = encoder.clone();
171        tokio::spawn(async move {
172            write_loop(inner_clone, &mut writer, rx_write_message, e, compressor).await;
173            if let Err(e) = writer.close().await {
174                log::error!("transport writer close error {e}");
175            }
176            log::debug!("exit write loop");
177            std::mem::drop(wg_rw);
178        });
179
180        // Spawn closer task
181        let inner_clone = inner.clone();
182        let wg_clone = wg.clone();
183        tokio::spawn(async move {
184            close_task::<C>(inner_clone, closer).await;
185            log::debug!("exit close task");
186            std::mem::drop(wg_clone);
187        });
188
189        // Spawn unreliable writer task
190        let e = encoder.clone();
191        if let Some(mut unreliable_writer) = unreliable_writer {
192            let inner_clone = inner.clone();
193            let wg_clone = wg.clone();
194            let rx_write_message = rx_unreliable_write_message.unwrap();
195            tokio::spawn(async move {
196                write_loop(
197                    inner_clone,
198                    &mut unreliable_writer,
199                    rx_write_message,
200                    e,
201                    compressor_for_unreliable,
202                )
203                .await;
204                log::debug!("unreliable write loop");
205                std::mem::drop(wg_clone);
206            });
207        }
208
209        // Spawn ping pong task
210        let inner_clone = inner.clone();
211        tokio::spawn(async move {
212            ping_pong_loop(inner_clone, rx_pong, ping_interval, ping_timeout).await;
213            log::debug!("exit ping pong loop");
214            std::mem::drop(wg);
215        });
216
217        Conn { inner }
218    }
219
220    pub async fn close(&self) -> Result<(), Error> {
221        if self.inner.ct.is_cancelled() {
222            return Ok(());
223        }
224        self.inner.ct.cancel();
225        self.inner.waiter.wait().await;
226        Ok(())
227    }
228
229    pub(crate) fn cancel(&self) {
230        self.inner.ct.cancel();
231    }
232
233    pub(crate) async fn send_message<T: Into<Message>>(&self, msg: T) -> Result<(), Error> {
234        let result = timeout_with_ct(
235            &self.inner.ct,
236            self.inner.response_message_timeout,
237            self.inner.tx_write_message.send(msg.into()),
238        )
239        .await?;
240        if result.is_err() {
241            return Err(Error::unexpected("write loop closed"));
242        }
243        Ok(())
244    }
245
246    pub(crate) async fn request_message_need_response<T: RequestMessage>(
247        &self,
248        mut request: T,
249    ) -> Result<T::Response, Error> {
250        let request_id = self.inner.request_id_counter.get();
251        let (tx, rx) = oneshot::channel();
252        self.inner.add_response_sender(request_id, tx);
253        request.set_request_id(request_id);
254        self.send_message(request.into()).await?;
255        let result =
256            timeout_with_ct(&self.inner.ct, self.inner.response_message_timeout, rx).await?;
257        let Ok(response) = result else {
258            return Err(Error::unexpected("internal channel closed"));
259        };
260        let Ok(response) = TryInto::<T::Response>::try_into(response) else {
261            return Err(Error::invalid_value("unexpected message returned"));
262        };
263
264        let Some(result_code) = response.result_code() else {
265            return Err(Error::invalid_value("unknown result code"));
266        };
267        if result_code != crate::message::ResultCode::Succeeded {
268            return Err(Error::FailedMessage {
269                result_code,
270                detail: response.result_string().to_owned(),
271            });
272        }
273        Ok(response)
274    }
275
276    pub(crate) async fn send_message_unreliable<T: Into<Message>>(
277        &self,
278        msg: T,
279    ) -> Result<(), Error> {
280        let msg = msg.into();
281        if let Some(tx) = &self.inner.tx_unreliable_write_message {
282            let result = timeout_with_ct(
283                &self.inner.ct,
284                self.inner.response_message_timeout,
285                tx.send(msg),
286            )
287            .await?;
288            if result.is_err() {
289                return Err(Error::unexpected("unreliable write loop closed"));
290            }
291            Ok(())
292        } else {
293            self.send_message(msg).await
294        }
295    }
296
297    pub(crate) async fn send_message_with_qos<T: Into<Message>>(
298        &self,
299        msg: T,
300        qos: crate::message::QoS,
301    ) -> Result<(), Error> {
302        if qos == crate::message::QoS::Unreliable {
303            self.send_message_unreliable(msg).await
304        } else {
305            self.send_message(msg).await
306        }
307    }
308
309    pub(crate) async fn cancelled(&self) {
310        self.inner.ct.cancelled().await
311    }
312
313    pub(crate) async fn add_upstream(
314        &self,
315        stream_id_alias: u32,
316    ) -> Result<
317        (
318            mpsc::Receiver<crate::message::UpstreamChunkAck>,
319            SendCommandGuard,
320        ),
321        Error,
322    > {
323        log::debug!("add upstream stream_id_alias = {stream_id_alias}");
324
325        let (tx, rx) = mpsc::channel(self.inner.channel_size);
326        let command = ReadLoopCommand::AddUpstream {
327            stream_id_alias,
328            tx_upstream_chunk_ack: tx,
329        };
330
331        self.inner
332            .tx_read_loop_command
333            .send(command)
334            .map_err(|_| Error::ConnectionClosed)?;
335
336        let guard = SendCommandGuard::new(
337            &self.inner.tx_read_loop_command,
338            ReadLoopCommand::RemoveUpstream { stream_id_alias },
339        );
340
341        Ok((rx, guard))
342    }
343
344    pub(crate) async fn add_downstream(
345        &self,
346        stream_id_alias: u32,
347        unreliable: bool,
348    ) -> Result<
349        (
350            mpsc::Receiver<ReceivableDownstreamMsg>,
351            (SendCommandGuard, Option<SendCommandGuard>),
352        ),
353        Error,
354    > {
355        log::debug!("add downstream stream_id_alias = {stream_id_alias}");
356
357        let (tx, rx) = mpsc::channel(self.inner.channel_size);
358
359        let unreliable_guard = if unreliable {
360            if let Some(tx_command) = &self.inner.tx_unreliable_read_loop_command {
361                let command = ReadLoopCommand::AddDownstream {
362                    stream_id_alias,
363                    tx_downstream_msg: tx.clone(),
364                };
365                tx_command
366                    .send(command)
367                    .map_err(|_| Error::ConnectionClosed)?;
368                Some(SendCommandGuard::new(
369                    tx_command,
370                    ReadLoopCommand::RemoveDownstream { stream_id_alias },
371                ))
372            } else {
373                None
374            }
375        } else {
376            None
377        };
378
379        let command = ReadLoopCommand::AddDownstream {
380            stream_id_alias,
381            tx_downstream_msg: tx,
382        };
383        self.inner
384            .tx_read_loop_command
385            .send(command)
386            .map_err(|_| Error::ConnectionClosed)?;
387
388        let guard = SendCommandGuard::new(
389            &self.inner.tx_read_loop_command,
390            ReadLoopCommand::RemoveDownstream { stream_id_alias },
391        );
392
393        Ok((rx, (guard, unreliable_guard)))
394    }
395
396    pub(crate) fn downstream_stream_id_alias(&self) -> Result<u32, Error> {
397        let stream_id_alias = self.inner.downstream_stream_id_alias_counter.fetch_add(1);
398        if stream_id_alias == 0 {
399            return Err(Error::unexpected("stream id alias max reached"));
400        }
401        Ok(stream_id_alias)
402    }
403
404    pub(crate) async fn subscribe_call_ack(
405        &self,
406        call_id: String,
407    ) -> Result<(oneshot::Receiver<UpstreamCallAck>, SendCommandGuard), Error> {
408        let (tx, rx) = oneshot::channel();
409        let command = ReadLoopCommand::SubscribeCallAck {
410            call_id: call_id.clone(),
411            tx,
412        };
413
414        self.inner
415            .tx_read_loop_command
416            .send(command)
417            .map_err(|_| Error::ConnectionClosed)?;
418
419        let guard = SendCommandGuard::new(
420            &self.inner.tx_read_loop_command,
421            ReadLoopCommand::RemoveCallAck { call_id },
422        );
423
424        Ok((rx, guard))
425    }
426
427    pub(crate) fn subscribe_downstream_call(&self) -> DownstreamCallReceiver {
428        DownstreamCallReceiver(self.inner.rx_downstream_call.resubscribe())
429    }
430
431    pub fn is_connected(&self) -> bool {
432        !self.inner.ct.is_cancelled()
433    }
434}
435
436async fn write_loop<T: TransportWriter>(
437    inner: Arc<ConnInner>,
438    writer: &mut T,
439    mut rx_write_message: mpsc::Receiver<Message>,
440    mut encoder: Encoder,
441    mut compressor: Compressor,
442) {
443    let _ct_guard = inner.ct.clone().drop_guard();
444    let mut buf = BytesMut::new();
445
446    loop {
447        buf.clear();
448        let msg = tokio::select! {
449            msg = rx_write_message.recv() => {
450                if let Some(msg) = msg {
451                    msg
452                } else {
453                    return;
454                }
455            }
456            _ = inner.ct.cancelled() => {
457                break;
458            }
459        };
460        log::trace!("write message: {msg:?}");
461        if let Err(e) = encoder.encode_to(&mut buf, &msg) {
462            log::warn!("cannot encode message: {e}");
463            continue;
464        }
465        if let Err(e) = compressor.compress(&mut buf) {
466            log::error!("message compression error: {e}");
467            return;
468        }
469        if let Err(e) = writer.write(&buf).await {
470            log::error!("transport write error: {e}");
471            return;
472        }
473    }
474
475    // Write all remaining messages
476    while let Ok(msg) = rx_write_message.try_recv() {
477        buf.clear();
478        log::trace!("write message: {msg:?}");
479        if let Err(e) = encoder.encode_to(&mut buf, &msg) {
480            log::warn!("cannot encode message: {e}");
481            continue;
482        }
483        if let Err(e) = compressor.compress(&mut buf) {
484            log::error!("message compression error: {e}");
485            return;
486        }
487        if let Err(e) = writer.write(&buf).await {
488            log::error!("transport write error: {e}");
489            return;
490        }
491    }
492}
493
494async fn close_task<T: Connector>(inner: Arc<ConnInner>, mut closer: T::Closer) {
495    let _ = inner.ct.cancelled().await;
496    // closer.close() is called after reader and writer close
497    inner.waiter_rw.wait().await;
498    if let Err(e) = closer.close().await {
499        log::error!("transport close error: {e}");
500    }
501}
502
503async fn ping_pong_loop(
504    inner: Arc<ConnInner>,
505    mut rx_pong: mpsc::Receiver<crate::message::Pong>,
506    ping_interval: Duration,
507    ping_timeout: Duration,
508) {
509    let _ct_guard = inner.ct.clone().drop_guard();
510    loop {
511        cancelled_return!(inner.ct, tokio::time::sleep(ping_interval));
512
513        let request_id = inner.request_id_counter.get();
514        let ping = crate::message::Ping {
515            request_id,
516            ..Default::default()
517        };
518        if cancelled_return!(
519            inner.ct,
520            inner
521                .tx_write_message
522                .send_timeout(ping.into(), ping_timeout)
523        )
524        .is_err()
525        {
526            log::error!("ping sending timeout reached");
527            return;
528        }
529        loop {
530            let timeout = tokio::time::Instant::now() + ping_timeout;
531            let pong = tokio::select! {
532                pong = rx_pong.recv() => pong,
533                _ = tokio::time::sleep_until(timeout) => {
534                    log::error!("ping timeout reached");
535                    return;
536                }
537                _ = inner.ct.cancelled() => {
538                    return;
539                }
540            };
541            let Some(pong) = pong else {
542                return;
543            };
544            if pong.request_id >= request_id {
545                break;
546            }
547        }
548    }
549}
550
551impl ConnInner {
552    fn add_response_sender(&self, request_id: u32, sender: oneshot::Sender<Message>) {
553        let old = self
554            .response_senders
555            .lock()
556            .expect("add_response_sender")
557            .insert(request_id, sender);
558        if old.is_some() {
559            log::warn!("request_id duplication detected");
560        }
561    }
562
563    fn remove_response_sender(&self, request_id: u32) -> Option<oneshot::Sender<Message>> {
564        self.response_senders
565            .lock()
566            .expect("remove_response_sender")
567            .remove(&request_id)
568    }
569}
570
571struct RequestIdCounter(AtomicCell<u32>);
572
573impl RequestIdCounter {
574    fn new() -> Self {
575        Self(AtomicCell::new(0))
576    }
577
578    fn get(&self) -> u32 {
579        self.0.fetch_add(2)
580    }
581}