Skip to main content

mito2/wal/
entry_distributor.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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
34/// [WalEntryDistributor] distributes Wal entries to specific [WalEntryReceiver]s based on [RegionId].
35pub(crate) struct WalEntryDistributor {
36    raw_wal_reader: Arc<dyn RawEntryReader>,
37    provider: Provider,
38    /// Sends [Entry] to receivers based on [RegionId]
39    senders: HashMap<RegionId, Sender<Entry>>,
40    /// Waits for the arg from the [WalEntryReader].
41    arg_receivers: Vec<(RegionId, oneshot::Receiver<EntryId>)>,
42}
43
44impl WalEntryDistributor {
45    /// Distributes entries to specific [WalEntryReceiver]s based on [RegionId].
46    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        // No subscribers
58        if args.is_empty() {
59            return Ok(());
60        }
61        // Safety: must exist
62        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[&region_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(&region_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/// Receives the Wal entries from [WalEntryDistributor].
98#[derive(Debug)]
99pub(crate) struct WalEntryReceiver {
100    /// Receives the [Entry] from the [WalEntryDistributor].
101    entry_receiver: Option<Receiver<Entry>>,
102    /// Sends the `start_id` to the [WalEntryDistributor].
103    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        // Safety: check via arg_sender
124        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
154/// The default buffer size of the [Entry] receiver.
155pub const DEFAULT_ENTRY_RECEIVER_BUFFER_SIZE: usize = 2048;
156
157/// Returns [WalEntryDistributor] and batch [WalEntryReceiver]s.
158///
159/// ### Note:
160/// Ensures `receiver.read` is called before the `distributor.distribute` in the same thread.
161///
162/// ```text
163/// let (distributor, receivers) = build_wal_entry_distributor_and_receivers(..);
164///  Thread 1                        |
165///                                  |
166/// // may deadlock                  |
167/// distributor.distribute().await;  |
168///                                  |  
169///                                  |
170/// receivers[0].read().await        |
171/// ```
172///
173pub 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 &region_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        // Drops all receivers
281        drop(receivers);
282        // Returns immediately
283        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        // Builds distributor and receivers
339        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        // Should be okay if one of receiver is dropped.
352        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            &region1_expected_wal_entry,
443            3,
444        ));
445        entries.extend(generate_tail_corrupted_stream(
446            provider.clone(),
447            region2,
448            &region2_expected_wal_entry,
449            2,
450        ));
451        entries.extend(generate_tail_corrupted_stream(
452            provider.clone(),
453            region3,
454            &region3_expected_wal_entry,
455            4,
456        ));
457
458        let corrupted_stream = MockRawEntryReader { entries };
459        // Builds distributor and receivers
460        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            &region1_expected_wal_entry,
524            3,
525        ));
526        entries.extend(vec![
527            // The incomplete entry.
528            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            // The incomplete entry.
536            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        // Builds distributor and receivers
547        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        // Builds distributor and receivers
647        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}