1use std::collections::HashMap;
16use std::sync::Arc;
17
18use async_stream::stream;
19use common_telemetry::{debug, error, warn};
20use futures::future::join_all;
21use snafu::OptionExt;
22use store_api::logstore::entry::Entry;
23use store_api::logstore::provider::Provider;
24use store_api::storage::RegionId;
25use tokio::sync::mpsc::{self, Receiver, Sender};
26use tokio::sync::oneshot;
27use tokio_stream::StreamExt;
28
29use crate::error::{self, Result};
30use crate::wal::entry_reader::{WalEntryReader, decode_raw_entry};
31use crate::wal::raw_entry_reader::RawEntryReader;
32use crate::wal::{EntryId, WalEntryStream};
33
34pub(crate) struct WalEntryDistributor {
36 raw_wal_reader: Arc<dyn RawEntryReader>,
37 provider: Provider,
38 senders: HashMap<RegionId, Sender<Entry>>,
40 arg_receivers: Vec<(RegionId, oneshot::Receiver<EntryId>)>,
42}
43
44impl WalEntryDistributor {
45 pub async fn distribute(mut self) -> Result<()> {
47 let arg_futures = self
48 .arg_receivers
49 .iter_mut()
50 .map(|(region_id, receiver)| async { (*region_id, receiver.await.ok()) });
51 let args = join_all(arg_futures)
52 .await
53 .into_iter()
54 .filter_map(|(region_id, start_id)| start_id.map(|start_id| (region_id, start_id)))
55 .collect::<Vec<_>>();
56
57 if args.is_empty() {
59 return Ok(());
60 }
61 let min_start_id = args.iter().map(|(_, start_id)| *start_id).min().unwrap();
63 let receivers: HashMap<_, _> = args
64 .into_iter()
65 .map(|(region_id, start_id)| {
66 (
67 region_id,
68 EntryReceiver {
69 start_id,
70 sender: self.senders[®ion_id].clone(),
71 },
72 )
73 })
74 .collect();
75
76 let mut stream = self.raw_wal_reader.read(&self.provider, min_start_id)?;
77 while let Some(entry) = stream.next().await {
78 let entry = entry?;
79 let entry_id = entry.entry_id();
80 let region_id = entry.region_id();
81
82 if let Some(EntryReceiver { sender, start_id }) = receivers.get(®ion_id) {
83 if entry_id >= *start_id
84 && let Err(err) = sender.send(entry).await
85 {
86 error!(err; "Failed to distribute raw entry, entry_id:{}, region_id: {}", entry_id, region_id);
87 }
88 } else {
89 debug!("Subscriber not found, region_id: {}", region_id);
90 }
91 }
92
93 Ok(())
94 }
95}
96
97#[derive(Debug)]
99pub(crate) struct WalEntryReceiver {
100 entry_receiver: Option<Receiver<Entry>>,
102 arg_sender: Option<oneshot::Sender<EntryId>>,
104}
105
106impl WalEntryReceiver {
107 pub fn new(entry_receiver: Receiver<Entry>, arg_sender: oneshot::Sender<EntryId>) -> Self {
108 Self {
109 entry_receiver: Some(entry_receiver),
110 arg_sender: Some(arg_sender),
111 }
112 }
113}
114
115impl WalEntryReader for WalEntryReceiver {
116 fn read(&mut self, _provider: &Provider, start_id: EntryId) -> Result<WalEntryStream<'static>> {
117 let arg_sender =
118 self.arg_sender
119 .take()
120 .with_context(|| error::InvalidWalReadRequestSnafu {
121 reason: format!("Call WalEntryReceiver multiple time, start_id: {start_id}"),
122 })?;
123 let mut entry_receiver = self.entry_receiver.take().unwrap();
125
126 if arg_sender.send(start_id).is_err() {
127 return error::InvalidWalReadRequestSnafu {
128 reason: format!(
129 "WalEntryDistributor is dropped, failed to send arg, start_id: {start_id}"
130 ),
131 }
132 .fail();
133 }
134
135 let stream = stream! {
136 while let Some(entry) = entry_receiver.recv().await {
137 if entry.is_complete() {
138 yield decode_raw_entry(entry);
139 } else {
140 warn!("Ignoring incomplete entry: {}", entry);
141 }
142 }
143 };
144
145 Ok(Box::pin(stream))
146 }
147}
148
149struct EntryReceiver {
150 start_id: EntryId,
151 sender: Sender<Entry>,
152}
153
154pub const DEFAULT_ENTRY_RECEIVER_BUFFER_SIZE: usize = 2048;
156
157pub fn build_wal_entry_distributor_and_receivers(
174 provider: Provider,
175 raw_wal_reader: Arc<dyn RawEntryReader>,
176 region_ids: &[RegionId],
177 buffer_size: usize,
178) -> (WalEntryDistributor, Vec<WalEntryReceiver>) {
179 let mut senders = HashMap::with_capacity(region_ids.len());
180 let mut readers = Vec::with_capacity(region_ids.len());
181 let mut arg_receivers = Vec::with_capacity(region_ids.len());
182
183 for ®ion_id in region_ids {
184 let (entry_sender, entry_receiver) = mpsc::channel(buffer_size);
185 let (arg_sender, arg_receiver) = oneshot::channel();
186
187 senders.insert(region_id, entry_sender);
188 arg_receivers.push((region_id, arg_receiver));
189 readers.push(WalEntryReceiver::new(entry_receiver, arg_sender));
190 }
191
192 (
193 WalEntryDistributor {
194 provider,
195 raw_wal_reader,
196 senders,
197 arg_receivers,
198 },
199 readers,
200 )
201}
202
203#[cfg(test)]
204mod tests {
205
206 use std::time::Duration;
207
208 use api::v1::{Mutation, OpType, WalEntry};
209 use futures::{StreamExt, TryStreamExt, stream};
210 use prost::Message;
211 use store_api::logstore::entry::{Entry, MultiplePartEntry, MultiplePartHeader, NaiveEntry};
212
213 use super::*;
214 use crate::test_util::wal_util::generate_tail_corrupted_stream;
215 use crate::wal::EntryId;
216 use crate::wal::raw_entry_reader::{EntryStream, RawEntryReader};
217
218 struct MockRawEntryReader {
219 entries: Vec<Entry>,
220 }
221
222 impl MockRawEntryReader {
223 pub fn new(entries: Vec<Entry>) -> MockRawEntryReader {
224 Self { entries }
225 }
226 }
227
228 impl RawEntryReader for MockRawEntryReader {
229 fn read(&self, _provider: &Provider, _start_id: EntryId) -> Result<EntryStream<'static>> {
230 let stream = stream::iter(self.entries.clone().into_iter().map(Ok));
231 Ok(Box::pin(stream))
232 }
233 }
234
235 #[tokio::test]
236 async fn test_delivers_complete_entry_while_channel_is_alive() {
237 let provider = Provider::kafka_provider("my_topic".to_string());
238 let region_id = RegionId::new(1024, 1);
239 let wal_entry = WalEntry::default();
240 let (sender, receiver) = mpsc::channel(1);
241 let (arg_sender, arg_receiver) = oneshot::channel();
242 let mut receiver = WalEntryReceiver::new(receiver, arg_sender);
243 let mut stream = receiver.read(&provider, 0).unwrap();
244 assert_eq!(arg_receiver.await.unwrap(), 0);
245 sender
246 .send(Entry::Naive(NaiveEntry {
247 provider,
248 region_id,
249 entry_id: 1,
250 data: wal_entry.encode_to_vec(),
251 }))
252 .await
253 .unwrap();
254
255 let entry = tokio::time::timeout(Duration::from_secs(1), stream.next())
256 .await
257 .unwrap()
258 .unwrap()
259 .unwrap();
260 assert_eq!(entry, (1, wal_entry));
261 }
262
263 #[tokio::test]
264 async fn test_wal_entry_distributor_without_receivers() {
265 let provider = Provider::kafka_provider("my_topic".to_string());
266 let reader = Arc::new(MockRawEntryReader::new(vec![Entry::Naive(NaiveEntry {
267 region_id: RegionId::new(1024, 1),
268 provider: provider.clone(),
269 entry_id: 1,
270 data: vec![1],
271 })]));
272
273 let (distributor, receivers) = build_wal_entry_distributor_and_receivers(
274 provider,
275 reader,
276 &[RegionId::new(1024, 1), RegionId::new(1025, 1)],
277 128,
278 );
279
280 drop(receivers);
282 distributor.distribute().await.unwrap();
284 }
285
286 #[tokio::test]
287 async fn test_wal_entry_distributor() {
288 common_telemetry::init_default_ut_logging();
289 let provider = Provider::kafka_provider("my_topic".to_string());
290 let reader = Arc::new(MockRawEntryReader::new(vec![
291 Entry::Naive(NaiveEntry {
292 provider: provider.clone(),
293 region_id: RegionId::new(1024, 1),
294 entry_id: 1,
295 data: WalEntry {
296 mutations: vec![Mutation {
297 op_type: OpType::Put as i32,
298 sequence: 1u64,
299 rows: None,
300 write_hint: None,
301 }],
302 bulk_entries: vec![],
303 }
304 .encode_to_vec(),
305 }),
306 Entry::Naive(NaiveEntry {
307 provider: provider.clone(),
308 region_id: RegionId::new(1024, 2),
309 entry_id: 2,
310 data: WalEntry {
311 mutations: vec![Mutation {
312 op_type: OpType::Put as i32,
313 sequence: 2u64,
314 rows: None,
315 write_hint: None,
316 }],
317 bulk_entries: vec![],
318 }
319 .encode_to_vec(),
320 }),
321 Entry::Naive(NaiveEntry {
322 provider: provider.clone(),
323 region_id: RegionId::new(1024, 3),
324 entry_id: 3,
325 data: WalEntry {
326 mutations: vec![Mutation {
327 op_type: OpType::Put as i32,
328 sequence: 3u64,
329 rows: None,
330 write_hint: None,
331 }],
332 bulk_entries: vec![],
333 }
334 .encode_to_vec(),
335 }),
336 ]));
337
338 let (distributor, mut receivers) = build_wal_entry_distributor_and_receivers(
340 provider.clone(),
341 reader,
342 &[
343 RegionId::new(1024, 1),
344 RegionId::new(1024, 2),
345 RegionId::new(1024, 3),
346 ],
347 128,
348 );
349 assert_eq!(receivers.len(), 3);
350
351 let last = receivers.pop().unwrap();
353 drop(last);
354
355 let mut streams = receivers
356 .iter_mut()
357 .map(|receiver| receiver.read(&provider, 0).unwrap())
358 .collect::<Vec<_>>();
359 distributor.distribute().await.unwrap();
360 let entries = streams
361 .get_mut(0)
362 .unwrap()
363 .try_collect::<Vec<_>>()
364 .await
365 .unwrap();
366 assert_eq!(
367 entries,
368 vec![(
369 1,
370 WalEntry {
371 mutations: vec![Mutation {
372 op_type: OpType::Put as i32,
373 sequence: 1u64,
374 rows: None,
375 write_hint: None,
376 }],
377 bulk_entries: vec![],
378 }
379 )]
380 );
381 let entries = streams
382 .get_mut(1)
383 .unwrap()
384 .try_collect::<Vec<_>>()
385 .await
386 .unwrap();
387 assert_eq!(
388 entries,
389 vec![(
390 2,
391 WalEntry {
392 mutations: vec![Mutation {
393 op_type: OpType::Put as i32,
394 sequence: 2u64,
395 rows: None,
396 write_hint: None,
397 }],
398 bulk_entries: vec![],
399 }
400 )]
401 );
402 }
403
404 #[tokio::test]
405 async fn test_tail_corrupted_stream() {
406 common_telemetry::init_default_ut_logging();
407 let mut entries = vec![];
408 let region1 = RegionId::new(1, 1);
409 let region1_expected_wal_entry = WalEntry {
410 mutations: vec![Mutation {
411 op_type: OpType::Put as i32,
412 sequence: 1u64,
413 rows: None,
414 write_hint: None,
415 }],
416 bulk_entries: vec![],
417 };
418 let region2 = RegionId::new(1, 2);
419 let region2_expected_wal_entry = WalEntry {
420 mutations: vec![Mutation {
421 op_type: OpType::Put as i32,
422 sequence: 3u64,
423 rows: None,
424 write_hint: None,
425 }],
426 bulk_entries: vec![],
427 };
428 let region3 = RegionId::new(1, 3);
429 let region3_expected_wal_entry = WalEntry {
430 mutations: vec![Mutation {
431 op_type: OpType::Put as i32,
432 sequence: 3u64,
433 rows: None,
434 write_hint: None,
435 }],
436 bulk_entries: vec![],
437 };
438 let provider = Provider::kafka_provider("my_topic".to_string());
439 entries.extend(generate_tail_corrupted_stream(
440 provider.clone(),
441 region1,
442 ®ion1_expected_wal_entry,
443 3,
444 ));
445 entries.extend(generate_tail_corrupted_stream(
446 provider.clone(),
447 region2,
448 ®ion2_expected_wal_entry,
449 2,
450 ));
451 entries.extend(generate_tail_corrupted_stream(
452 provider.clone(),
453 region3,
454 ®ion3_expected_wal_entry,
455 4,
456 ));
457
458 let corrupted_stream = MockRawEntryReader { entries };
459 let (distributor, mut receivers) = build_wal_entry_distributor_and_receivers(
461 provider.clone(),
462 Arc::new(corrupted_stream),
463 &[region1, region2, region3],
464 128,
465 );
466 assert_eq!(receivers.len(), 3);
467 let mut streams = receivers
468 .iter_mut()
469 .map(|receiver| receiver.read(&provider, 0).unwrap())
470 .collect::<Vec<_>>();
471 distributor.distribute().await.unwrap();
472
473 assert_eq!(
474 streams
475 .get_mut(0)
476 .unwrap()
477 .try_collect::<Vec<_>>()
478 .await
479 .unwrap(),
480 vec![(0, region1_expected_wal_entry)]
481 );
482
483 assert_eq!(
484 streams
485 .get_mut(1)
486 .unwrap()
487 .try_collect::<Vec<_>>()
488 .await
489 .unwrap(),
490 vec![(0, region2_expected_wal_entry)]
491 );
492
493 assert_eq!(
494 streams
495 .get_mut(2)
496 .unwrap()
497 .try_collect::<Vec<_>>()
498 .await
499 .unwrap(),
500 vec![(0, region3_expected_wal_entry)]
501 );
502 }
503
504 #[tokio::test]
505 async fn test_part_corrupted_stream() {
506 common_telemetry::init_default_ut_logging();
507 let mut entries = vec![];
508 let region1 = RegionId::new(1, 1);
509 let region1_expected_wal_entry = WalEntry {
510 mutations: vec![Mutation {
511 op_type: OpType::Put as i32,
512 sequence: 1u64,
513 rows: None,
514 write_hint: None,
515 }],
516 bulk_entries: vec![],
517 };
518 let region2 = RegionId::new(1, 2);
519 let provider = Provider::kafka_provider("my_topic".to_string());
520 entries.extend(generate_tail_corrupted_stream(
521 provider.clone(),
522 region1,
523 ®ion1_expected_wal_entry,
524 3,
525 ));
526 entries.extend(vec![
527 Entry::MultiplePart(MultiplePartEntry {
529 provider: provider.clone(),
530 region_id: region2,
531 entry_id: 0,
532 headers: vec![MultiplePartHeader::First],
533 parts: vec![vec![1; 100]],
534 }),
535 Entry::MultiplePart(MultiplePartEntry {
537 provider: provider.clone(),
538 region_id: region2,
539 entry_id: 0,
540 headers: vec![MultiplePartHeader::First],
541 parts: vec![vec![1; 100]],
542 }),
543 ]);
544
545 let corrupted_stream = MockRawEntryReader { entries };
546 let (distributor, mut receivers) = build_wal_entry_distributor_and_receivers(
548 provider.clone(),
549 Arc::new(corrupted_stream),
550 &[region1, region2],
551 128,
552 );
553 assert_eq!(receivers.len(), 2);
554 let mut streams = receivers
555 .iter_mut()
556 .map(|receiver| receiver.read(&provider, 0).unwrap())
557 .collect::<Vec<_>>();
558 distributor.distribute().await.unwrap();
559 assert_eq!(
560 streams
561 .get_mut(0)
562 .unwrap()
563 .try_collect::<Vec<_>>()
564 .await
565 .unwrap(),
566 vec![(0, region1_expected_wal_entry)]
567 );
568
569 assert_eq!(
570 streams
571 .get_mut(1)
572 .unwrap()
573 .try_collect::<Vec<_>>()
574 .await
575 .unwrap(),
576 vec![]
577 );
578 }
579
580 #[tokio::test]
581 async fn test_wal_entry_receiver_start_id() {
582 let provider = Provider::kafka_provider("my_topic".to_string());
583 let reader = Arc::new(MockRawEntryReader::new(vec![
584 Entry::Naive(NaiveEntry {
585 provider: provider.clone(),
586 region_id: RegionId::new(1024, 1),
587 entry_id: 1,
588 data: WalEntry {
589 mutations: vec![Mutation {
590 op_type: OpType::Put as i32,
591 sequence: 1u64,
592 rows: None,
593 write_hint: None,
594 }],
595 bulk_entries: vec![],
596 }
597 .encode_to_vec(),
598 }),
599 Entry::Naive(NaiveEntry {
600 provider: provider.clone(),
601 region_id: RegionId::new(1024, 2),
602 entry_id: 2,
603 data: WalEntry {
604 mutations: vec![Mutation {
605 op_type: OpType::Put as i32,
606 sequence: 2u64,
607 rows: None,
608 write_hint: None,
609 }],
610 bulk_entries: vec![],
611 }
612 .encode_to_vec(),
613 }),
614 Entry::Naive(NaiveEntry {
615 provider: provider.clone(),
616 region_id: RegionId::new(1024, 1),
617 entry_id: 3,
618 data: WalEntry {
619 mutations: vec![Mutation {
620 op_type: OpType::Put as i32,
621 sequence: 3u64,
622 rows: None,
623 write_hint: None,
624 }],
625 bulk_entries: vec![],
626 }
627 .encode_to_vec(),
628 }),
629 Entry::Naive(NaiveEntry {
630 provider: provider.clone(),
631 region_id: RegionId::new(1024, 2),
632 entry_id: 4,
633 data: WalEntry {
634 mutations: vec![Mutation {
635 op_type: OpType::Put as i32,
636 sequence: 4u64,
637 rows: None,
638 write_hint: None,
639 }],
640 bulk_entries: vec![],
641 }
642 .encode_to_vec(),
643 }),
644 ]));
645
646 let (distributor, mut receivers) = build_wal_entry_distributor_and_receivers(
648 provider.clone(),
649 reader,
650 &[RegionId::new(1024, 1), RegionId::new(1024, 2)],
651 128,
652 );
653 assert_eq!(receivers.len(), 2);
654 let mut streams = receivers
655 .iter_mut()
656 .map(|receiver| receiver.read(&provider, 4).unwrap())
657 .collect::<Vec<_>>();
658 distributor.distribute().await.unwrap();
659
660 assert_eq!(
661 streams
662 .get_mut(1)
663 .unwrap()
664 .try_collect::<Vec<_>>()
665 .await
666 .unwrap(),
667 vec![(
668 4,
669 WalEntry {
670 mutations: vec![Mutation {
671 op_type: OpType::Put as i32,
672 sequence: 4u64,
673 rows: None,
674 write_hint: None,
675 }],
676 bulk_entries: vec![],
677 }
678 )]
679 );
680 }
681}