1mod 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 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 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 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 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 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 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 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 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 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}