1use std::{
2 collections::HashMap,
3 sync::Arc,
4 time::{Duration, Instant, SystemTime},
5};
6
7use crossbeam::atomic::AtomicCell;
8use tokio::sync::{mpsc, oneshot};
9use tokio_util::sync::CancellationToken;
10use uuid::Uuid;
11
12use super::{
13 SharedWireConn,
14 down_order::*,
15 down_state::*,
16 misc::ReconnectWaiter,
17 types::{DataPointGroup, DownstreamChunk, DownstreamMetadata},
18};
19use crate::{
20 error::Error,
21 internal::{WaitGroup, timeout_with_ct},
22 message::{
23 DataId, DownstreamFilter, QoS, ResultCode, data_point_group::DataIdOrAlias,
24 downstream_chunk::UpstreamOrAlias,
25 },
26 wire::Conn as WireConn,
27};
28
29#[derive(Clone, Debug)]
31#[non_exhaustive]
32pub struct DownstreamConfig {
33 pub filters: Vec<DownstreamFilter>,
35 pub expiry_interval: Duration,
37 pub qos: QoS,
39 pub data_ids: Vec<DataId>,
41 pub ack_interval: Duration,
43 pub omit_empty_chunk: bool,
45 pub reordering: DownstreamReordering,
47 pub reordering_chunks: usize,
49 pub close_timeout: Duration,
51}
52
53#[derive(Clone, Copy, PartialEq, Eq, Default, Debug)]
55pub enum DownstreamReordering {
56 #[default]
57 None,
58 BestEffort,
59 Strict,
60}
61
62impl Default for DownstreamConfig {
63 fn default() -> Self {
64 Self {
65 filters: Vec::new(),
66 expiry_interval: Duration::from_secs(60),
67 qos: QoS::Unreliable,
68 data_ids: Vec::new(),
69 ack_interval: Duration::from_millis(100),
70 omit_empty_chunk: false,
71 reordering: DownstreamReordering::default(),
72 reordering_chunks: 0xFF,
73 close_timeout: Duration::from_secs(10),
74 }
75 }
76}
77
78pub struct Downstream {
80 inner: Arc<DownstreamInner>,
81 rx_downstream_chunk: mpsc::Receiver<DownstreamChunk>,
82}
83
84pub struct DownstreamMetadataReader {
86 inner: Arc<DownstreamInner>,
87 rx: mpsc::Receiver<DownstreamMetadata>,
88}
89
90pub struct DownstreamInner {
91 config: Arc<DownstreamConfig>,
92 stream_id: Uuid,
93 stream_id_alias: AtomicCell<u32>,
94 state: State,
95 resume_token: AtomicCell<String>,
96 tx_result: AtomicCell<Option<oneshot::Sender<Result<(), Error>>>>,
97 server_time: SystemTime,
98 close_cause: AtomicCell<Option<Error>>,
99 ct: CancellationToken,
100}
101
102impl Downstream {
103 pub(crate) async fn new(
104 config: Arc<DownstreamConfig>,
105 shared_wire_conn: SharedWireConn,
106 wg: WaitGroup,
107 channel_size: usize,
108 ) -> Result<(Self, DownstreamMetadataReader, CancellationToken), Error> {
109 let wire_conn = shared_wire_conn.get();
110 let (response, stream_id_alias, data_id_aliases) =
111 request_open(&wire_conn, &config).await?;
112
113 let stream_id = super::misc::parse_stream_id(&response.assigned_stream_id)?;
114 let server_time = super::misc::unix_epoch_to_system_time(response.server_time)?;
115
116 let ct = CancellationToken::new();
117
118 let (tx_downstream_chunk, rx_downstream_chunk) = mpsc::channel(channel_size);
119 let (tx_metadata, rx_metadata) = mpsc::channel(channel_size);
120
121 let inner = Arc::new(DownstreamInner {
122 config,
123 stream_id,
124 stream_id_alias: AtomicCell::new(stream_id_alias),
125 state: State::new(data_id_aliases),
126 resume_token: AtomicCell::new(response.resume_token),
127 tx_result: AtomicCell::new(None),
128 server_time,
129 close_cause: AtomicCell::new(None),
130 ct: ct.clone(),
131 });
132
133 let inner_clone = inner.clone();
134 tokio::spawn(async move {
135 if let Err(e) = downstream_loop(
136 wire_conn,
137 shared_wire_conn,
138 inner_clone.clone(),
139 tx_downstream_chunk,
140 tx_metadata,
141 )
142 .await
143 {
144 inner_clone.close_cause.store(Some(e));
145 }
146 log::debug!("exit downstream loop");
147 std::mem::drop(wg);
148 });
149
150 log::info!("opened downstream {stream_id}");
151 Ok((
152 Self {
153 inner: inner.clone(),
154 rx_downstream_chunk,
155 },
156 DownstreamMetadataReader {
157 inner,
158 rx: rx_metadata,
159 },
160 ct,
161 ))
162 }
163
164 pub async fn read_chunk(&mut self) -> Result<DownstreamChunk, Error> {
166 tokio::select! {
167 result = self.rx_downstream_chunk.recv() => {
168 return result.ok_or_else(|| self.inner.close_cause.take().unwrap_or(Error::StreamClosed));
169 }
170 _ = self.inner.ct.cancelled() => (),
171 }
172
173 self.rx_downstream_chunk
174 .try_recv()
175 .map_err(|_| self.inner.close_cause.take().unwrap_or(Error::StreamClosed))
176 }
177
178 pub async fn close(&mut self) -> Result<(), Error> {
180 let (tx, rx) = oneshot::channel();
181 self.inner.tx_result.store(Some(tx));
182 self.inner.ct.cancel();
183 tokio::time::timeout(self.inner.config.close_timeout, rx)
184 .await
185 .map_err(|_| Error::unexpected("close timeout"))?
186 .map_err(|_| Error::unexpected("cannot get close result"))?
187 }
188
189 pub fn stream_id(&self) -> Uuid {
191 self.inner.stream_id
192 }
193
194 pub fn server_time(&self) -> SystemTime {
196 self.inner.server_time
197 }
198
199 pub fn config(&self) -> Arc<DownstreamConfig> {
201 self.inner.config.clone()
202 }
203
204 pub fn state(&self) -> DownstreamState {
206 self.inner.state.state()
207 }
208}
209
210impl DownstreamMetadataReader {
211 pub async fn read(&mut self) -> Result<DownstreamMetadata, Error> {
212 tokio::select! {
213 result = self.rx.recv() => {
214 return result.ok_or_else(|| Error::StreamClosed);
215 }
216 _ = self.inner.ct.cancelled() => (),
217 }
218
219 self.rx.try_recv().map_err(|_| Error::StreamClosed)
220 }
221}
222
223#[allow(clippy::while_let_loop)]
224async fn downstream_loop(
225 mut wire_conn: WireConn,
226 mut shared_wire_conn: SharedWireConn,
227 inner: Arc<DownstreamInner>,
228 tx_downstream_chunk: mpsc::Sender<DownstreamChunk>,
229 tx_metadata: mpsc::Sender<DownstreamMetadata>,
230) -> Result<(), Error> {
231 let _ct_guard = inner.ct.clone().drop_guard();
232 let mut need_resume = false;
233 let mut ack_id_complete = 0;
234 let ct = &inner.ct;
235 let mut reorderer = DownReorderer::new(&inner.config);
236 let omit_empty_chunk = inner.config.omit_empty_chunk && !reorderer.enabled();
237
238 loop {
239 let mut resume_retry_waiter = ReconnectWaiter::new();
240 let mut resume_token = String::new();
241 let result = timeout_with_ct(ct, inner.config.expiry_interval, async {
243 loop {
244 if wire_conn.is_connected() {
245 if need_resume {
246 if inner.config.expiry_interval.is_zero() {
247 break None;
248 }
249 if resume_token.is_empty() {
250 resume_token = inner.resume_token.take();
251 }
252 match request_resume(&wire_conn, inner.stream_id, resume_token.clone())
253 .await
254 {
255 Ok((stream_id_alias, resume_token)) => {
256 inner.stream_id_alias.store(stream_id_alias);
257 inner.resume_token.store(resume_token);
258 log::info!("resume success downstream {}", inner.stream_id);
259 need_resume = false;
260 }
261 Err(e) => {
262 if e.result_code().is_some() {
263 log::warn!("cancel resume by: {e}");
264 break None;
265 } else if e.can_retry_resume() {
266 log::warn!("cannot resume and retry: {e}");
267 resume_retry_waiter.wait().await;
268 } else {
269 log::warn!("cannot resume: {e}");
270 break None;
271 }
272 }
273 }
274 } else if let Ok(result) = wire_conn
275 .add_downstream(
276 inner.stream_id_alias.load(),
277 inner.config.qos == QoS::Unreliable,
278 )
279 .await
280 {
281 break Some(result);
282 }
283 continue;
284 }
285 wire_conn = if let Ok(wire_conn) = shared_wire_conn.get_updated().await {
286 need_resume = true;
287 wire_conn
288 } else {
289 break None;
290 };
291 }
292 })
293 .await;
294
295 let (mut rx_msg, _guard) = match result {
296 Ok(Some(result)) => result,
297 Ok(None) => {
298 return Ok(());
299 }
300 Err(_) => {
301 log::error!("resume timeout in downstream {}", inner.stream_id);
302 return Ok(());
303 }
304 };
305
306 let mut ack_send = Instant::now() + inner.config.ack_interval;
308
309 loop {
310 let msg = tokio::select! {
311 _ = ct.cancelled() => { break; },
312 msg = rx_msg.recv() => {
313 if let Some(msg) = msg {
314 msg
315 } else {
316 break;
318 }
319 }
320 _ = tokio::time::sleep_until(ack_send.into()) => {
321 ack_send = Instant::now() + inner.config.ack_interval;
322 if let Some(ack) = inner.state.take_ack(inner.stream_id_alias.load())
323 && let Err(e) = wire_conn.send_message(ack).await {
324 log::warn!("cannot send downstream ack: {e}");
325 break;
326 }
327 continue;
328 }
329 };
330
331 match msg {
332 crate::wire::ReceivableDownstreamMsg::Chunk(chunk) => {
333 match convert_downstream_chunk(chunk, &inner, omit_empty_chunk) {
334 Ok(Some(chunk)) => {
335 for chunk in reorderer.iter(chunk)? {
336 if tx_downstream_chunk.send(chunk).await.is_err() {
337 ct.cancel(); break;
339 }
340 }
341 }
342 Err(e) => {
343 log::error!("{e}");
344 }
345 _ => (),
346 }
347 }
348 crate::wire::ReceivableDownstreamMsg::Metadata(metadata) => {
349 match convert_downstream_metadata(metadata) {
350 Ok((metadata, ack)) => {
351 if tx_metadata.send(metadata).await.is_err() {
352 ct.cancel(); break;
354 }
355 if let Err(e) = wire_conn.send_message(ack).await {
356 log::warn!("cannot send metadata ack: {e}");
357 break;
358 }
359 }
360 Err(e) => {
361 log::error!("{e}");
362 }
363 }
364 }
365 crate::wire::ReceivableDownstreamMsg::ChunkAckComplete(complete) => {
366 check_chunk_ack_complete(complete, &mut ack_id_complete);
367 }
368 }
369 }
370
371 while let Ok(msg) = rx_msg.try_recv() {
373 match msg {
374 crate::wire::ReceivableDownstreamMsg::Chunk(chunk) => {
375 match convert_downstream_chunk(chunk, &inner, omit_empty_chunk) {
376 Ok(Some(chunk)) => {
377 for chunk in reorderer.iter(chunk)? {
378 let _ = tx_downstream_chunk.try_send(chunk);
379 }
380 }
381 Err(e) => {
382 log::error!("{e}");
383 }
384 _ => (),
385 }
386 }
387 crate::wire::ReceivableDownstreamMsg::Metadata(metadata) => {
388 match convert_downstream_metadata(metadata) {
389 Ok((metadata, ack)) => {
390 let _ = tx_metadata.try_send(metadata);
391 if wire_conn.is_connected()
392 && let Err(e) = wire_conn.send_message(ack).await
393 {
394 log::warn!("cannot send metadata ack: {e}");
395 }
396 }
397 Err(e) => {
398 log::error!("{e}");
399 }
400 }
401 }
402 crate::wire::ReceivableDownstreamMsg::ChunkAckComplete(complete) => {
403 check_chunk_ack_complete(complete, &mut ack_id_complete);
404 }
405 }
406 }
407
408 if wire_conn.is_connected()
410 && let Some(ack) = inner.state.take_ack(inner.stream_id_alias.load())
411 && let Err(e) = wire_conn.send_message(ack).await
412 {
413 log::warn!("cannot send downstream ack: {e}");
414 }
415
416 if inner.config.expiry_interval.is_zero() || ct.is_cancelled() {
417 if wire_conn.is_connected() {
419 let result = tokio::time::timeout(inner.config.close_timeout, async {
421 loop {
422 if let Some(last_issued) = inner.state.last_issued_chunk_ack_id() {
423 if last_issued <= ack_id_complete {
424 break;
425 }
426 } else {
427 break;
428 }
429
430 match rx_msg.recv().await {
431 Some(crate::wire::ReceivableDownstreamMsg::ChunkAckComplete(
432 complete,
433 )) => {
434 check_chunk_ack_complete(complete, &mut ack_id_complete);
435 }
436 None => break,
437 _ => (),
438 }
439 }
440 })
441 .await;
442 if result.is_err() {
443 log::warn!("close timeout at downstream {}", inner.stream_id);
444 }
445
446 let close_msg = crate::message::DownstreamCloseRequest {
448 stream_id: inner.stream_id.as_bytes().to_vec().into(),
449 ..Default::default()
450 };
451 let result =
452 if let Err(e) = wire_conn.request_message_need_response(close_msg).await {
453 log::warn!("cannot send downstream close message: {e}");
454 Err(Error::ConnectionClosed)
455 } else {
456 Ok(())
457 };
458 if let Some(tx_result) = inner.tx_result.take() {
459 let _ = tx_result.send(result);
460 }
461 }
462 return Ok(());
463 }
464
465 log::info!("try to resume downstream: {}", inner.stream_id);
466 }
467}
468
469async fn request_open(
470 wire_conn: &WireConn,
471 config: &DownstreamConfig,
472) -> Result<
473 (
474 crate::message::DownstreamOpenResponse,
475 u32,
476 HashMap<u32, DataId>,
477 ),
478 Error,
479> {
480 let Ok(expiry_interval) = config.expiry_interval.as_secs().try_into() else {
481 return Err(Error::invalid_value("expiry_interval overflow"));
482 };
483 let desired_stream_id_alias = wire_conn.downstream_stream_id_alias()?;
484
485 let data_id_aliases: HashMap<_, _> = config
486 .data_ids
487 .iter()
488 .enumerate()
489 .map(|(i, data_id)| (i as u32 + 1, data_id.clone()))
490 .collect();
491
492 let request = crate::message::DownstreamOpenRequest {
493 desired_stream_id_alias,
494 downstream_filters: config.filters.clone(),
495 expiry_interval,
496 qos: config.qos.into(),
497 omit_empty_chunk: config.omit_empty_chunk,
498 data_id_aliases: data_id_aliases.clone(),
499 ..Default::default()
500 };
501
502 log::debug!("downstream open request: {request:?}");
503
504 let response = wire_conn.request_message_need_response(request).await?;
505 Ok((response, desired_stream_id_alias, data_id_aliases))
506}
507
508async fn request_resume(
509 wire_conn: &WireConn,
510 stream_id: Uuid,
511 resume_token: String,
512) -> Result<(u32, String), Error> {
513 let desired_stream_id_alias = wire_conn.downstream_stream_id_alias()?;
514 let request = crate::message::DownstreamResumeRequest {
515 desired_stream_id_alias,
516 stream_id: stream_id.as_bytes().to_vec().into(),
517 resume_token,
518 ..Default::default()
519 };
520
521 let response = wire_conn.request_message_need_response(request).await?;
522 Ok((desired_stream_id_alias, response.resume_token))
523}
524
525fn convert_downstream_chunk(
526 chunk: crate::message::DownstreamChunk,
527 inner: &DownstreamInner,
528 omit_empty_chunk: bool,
529) -> Result<Option<DownstreamChunk>, Error> {
530 let Some(stream_chunk) = chunk.stream_chunk else {
531 return Ok(None);
532 };
533 if omit_empty_chunk && stream_chunk.data_point_groups.is_empty() {
534 return Ok(None);
535 }
536 let sequence_number = stream_chunk.sequence_number;
537
538 let mut data_point_groups = Vec::new();
539 for dpg in stream_chunk.data_point_groups.into_iter() {
540 let data_id = match dpg.data_id_or_alias {
541 Some(DataIdOrAlias::DataId(data_id)) => data_id,
542 Some(DataIdOrAlias::DataIdAlias(alias)) => {
543 if let Some(data_id) = inner.state.get_data_id(alias) {
544 data_id
545 } else {
546 return Err(Error::invalid_value(format!("unknown data id {alias}")));
547 }
548 }
549 None => {
550 return Err(Error::invalid_value("invalid data_id_or_alias"));
551 }
552 };
553 data_point_groups.push(DataPointGroup {
554 data_id,
555 data_points: dpg.data_points,
556 });
557 }
558
559 let upstream = match chunk.upstream_or_alias {
560 Some(UpstreamOrAlias::UpstreamInfo(info)) => inner.state.add_upstream_info(info)?,
561 Some(UpstreamOrAlias::UpstreamAlias(alias)) => inner.state.get_upstream_info(alias)?,
562 None => {
563 return Err(Error::invalid_value("invalid upstream_or_alias"));
564 }
565 };
566
567 inner
568 .state
569 .add_sequence_number(upstream.stream_id, sequence_number);
570
571 Ok(Some(DownstreamChunk {
572 sequence_number,
573 data_point_groups,
574 upstream,
575 }))
576}
577
578fn convert_downstream_metadata(
579 metadata: crate::message::DownstreamMetadata,
580) -> Result<(DownstreamMetadata, crate::message::DownstreamMetadataAck), Error> {
581 let Some(m) = metadata.metadata else {
582 return Err(Error::invalid_value("invalid metadata"));
583 };
584 let m = super::metadata::ReceivableMetadata::from_prost(m)?;
585 let m = DownstreamMetadata {
586 source_node_id: metadata.source_node_id,
587 metadata: m,
588 };
589 let ack = crate::message::DownstreamMetadataAck {
590 request_id: metadata.request_id,
591 result_code: ResultCode::Succeeded.into(),
592 result_string: "OK".into(),
593 ..Default::default()
594 };
595
596 Ok((m, ack))
597}
598
599fn check_chunk_ack_complete(
600 complete: crate::message::DownstreamChunkAckComplete,
601 ack_id_complete: &mut u32,
602) {
603 if *ack_id_complete < complete.ack_id {
604 *ack_id_complete = complete.ack_id;
605 }
606 if complete.result_code != ResultCode::Succeeded as i32 {
607 log::error!(
608 "failed downstream chunk ack complete: result_code = {}, {}",
609 complete.result_code,
610 complete.result_string
611 );
612 }
613}
614
615impl std::fmt::Debug for Downstream {
616 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
617 f.debug_struct("Downstream")
618 .field("stream_id", &self.inner.stream_id)
619 .finish()
620 }
621}
622
623impl std::fmt::Debug for DownstreamMetadataReader {
624 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
625 f.debug_struct("DownstreamMetadataReader")
626 .field("stream_id", &self.inner.stream_id)
627 .finish()
628 }
629}