Skip to main content

risingwave_stream/executor/
merge.rs

1// Copyright 2022 RisingWave Labs
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::VecDeque;
16use std::pin::Pin;
17use std::task::{Context, Poll};
18
19use risingwave_common::array::StreamChunkBuilder;
20use tokio::sync::mpsc;
21
22use super::exchange::input::BoxedActorInput;
23use super::*;
24use crate::executor::prelude::*;
25use crate::task::LocalBarrierManager;
26
27pub type SelectReceivers = DynamicReceivers<ActorId, ()>;
28
29pub type MergeUpstream = BufferChunks<SelectReceivers>;
30pub type SingletonUpstream = BoxedActorInput;
31
32pub(crate) enum MergeExecutorUpstream {
33    Singleton(SingletonUpstream),
34    Merge(MergeUpstream),
35}
36
37pub(crate) struct MergeExecutorInput {
38    upstream: MergeExecutorUpstream,
39    actor_context: ActorContextRef,
40    upstream_fragment_id: UpstreamFragmentId,
41    local_barrier_manager: LocalBarrierManager,
42    executor_stats: Arc<StreamingMetrics>,
43    pub(crate) info: ExecutorInfo,
44}
45
46impl MergeExecutorInput {
47    pub(crate) fn new(
48        upstream: MergeExecutorUpstream,
49        actor_context: ActorContextRef,
50        upstream_fragment_id: UpstreamFragmentId,
51        local_barrier_manager: LocalBarrierManager,
52        executor_stats: Arc<StreamingMetrics>,
53        info: ExecutorInfo,
54    ) -> Self {
55        Self {
56            upstream,
57            actor_context,
58            upstream_fragment_id,
59            local_barrier_manager,
60            executor_stats,
61            info,
62        }
63    }
64
65    pub(crate) fn into_executor(self, barrier_rx: mpsc::UnboundedReceiver<Barrier>) -> Executor {
66        let fragment_id = self.actor_context.fragment_id;
67        let executor = match self.upstream {
68            MergeExecutorUpstream::Singleton(input) => ReceiverExecutor::new(
69                self.actor_context,
70                fragment_id,
71                self.upstream_fragment_id,
72                input,
73                self.local_barrier_manager,
74                self.executor_stats,
75                barrier_rx,
76            )
77            .boxed(),
78            MergeExecutorUpstream::Merge(inputs) => MergeExecutor::new(
79                self.actor_context,
80                fragment_id,
81                self.upstream_fragment_id,
82                inputs,
83                self.local_barrier_manager,
84                self.executor_stats,
85                barrier_rx,
86            )
87            .boxed(),
88        };
89        (self.info, executor).into()
90    }
91}
92
93impl Stream for MergeExecutorInput {
94    type Item = DispatcherMessageStreamItem;
95
96    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
97        match &mut self.get_mut().upstream {
98            MergeExecutorUpstream::Singleton(input) => input.poll_next_unpin(cx),
99            MergeExecutorUpstream::Merge(inputs) => inputs.poll_next_unpin(cx),
100        }
101    }
102}
103
104mod upstream {
105    use super::*;
106
107    /// Trait unifying operations on [`MergeUpstream`] and [`SingletonUpstream`], so that we can
108    /// reuse code between [`MergeExecutor`] and [`ReceiverExecutor`].
109    pub trait Upstream:
110        Stream<Item = StreamExecutorResult<DispatcherMessage>> + Unpin + Send + 'static
111    {
112        fn upstream_input_ids(&self) -> impl Iterator<Item = ActorId> + '_;
113        fn is_empty(&self) -> bool;
114        fn flush_buffered_watermarks(&mut self);
115        fn update(&mut self, to_add: Vec<BoxedActorInput>, to_remove: &HashSet<ActorId>);
116    }
117
118    impl Upstream for MergeUpstream {
119        fn upstream_input_ids(&self) -> impl Iterator<Item = ActorId> + '_ {
120            self.inner.upstream_input_ids()
121        }
122
123        fn is_empty(&self) -> bool {
124            self.inner.is_empty()
125        }
126
127        fn flush_buffered_watermarks(&mut self) {
128            self.inner.flush_buffered_watermarks();
129        }
130
131        fn update(&mut self, to_add: Vec<BoxedActorInput>, to_remove: &HashSet<ActorId>) {
132            if !to_add.is_empty() {
133                self.inner.add_upstreams_from(to_add);
134            }
135            if !to_remove.is_empty() {
136                self.inner.remove_upstreams(to_remove);
137            }
138        }
139    }
140
141    impl Upstream for SingletonUpstream {
142        fn upstream_input_ids(&self) -> impl Iterator<Item = ActorId> + '_ {
143            std::iter::once(self.id())
144        }
145
146        fn is_empty(&self) -> bool {
147            false
148        }
149
150        fn flush_buffered_watermarks(&mut self) {
151            // For single input, we won't buffer watermarks so there's nothing to do.
152        }
153
154        fn update(&mut self, to_add: Vec<BoxedActorInput>, to_remove: &HashSet<ActorId>) {
155            assert!(
156                to_remove.contains(&self.id()),
157                "the removed upstream actors should contain the current input"
158            );
159
160            // Replace the single input.
161            *self = Itertools::exactly_one(to_add.into_iter())
162                .expect("receiver should have exactly one new upstream");
163        }
164    }
165}
166use upstream::Upstream;
167
168/// The core of `MergeExecutor` and `ReceiverExecutor`.
169pub struct MergeExecutorInner<U> {
170    /// The context of the actor.
171    actor_context: ActorContextRef,
172
173    /// Upstream channels.
174    upstream: U,
175
176    /// Belonged fragment id.
177    fragment_id: FragmentId,
178
179    /// Upstream fragment id.
180    upstream_fragment_id: FragmentId,
181
182    local_barrier_manager: LocalBarrierManager,
183
184    /// Streaming metrics.
185    metrics: Arc<StreamingMetrics>,
186
187    barrier_rx: mpsc::UnboundedReceiver<Barrier>,
188}
189
190impl<U> MergeExecutorInner<U> {
191    pub fn new(
192        ctx: ActorContextRef,
193        fragment_id: FragmentId,
194        upstream_fragment_id: FragmentId,
195        upstream: U,
196        local_barrier_manager: LocalBarrierManager,
197        metrics: Arc<StreamingMetrics>,
198        barrier_rx: mpsc::UnboundedReceiver<Barrier>,
199    ) -> Self {
200        Self {
201            actor_context: ctx,
202            upstream,
203            fragment_id,
204            upstream_fragment_id,
205            local_barrier_manager,
206            metrics,
207            barrier_rx,
208        }
209    }
210}
211
212/// `MergeExecutor` merges data from multiple upstream actors and aligns them with barriers.
213pub type MergeExecutor = MergeExecutorInner<MergeUpstream>;
214
215impl MergeExecutor {
216    #[cfg(test)]
217    pub fn for_test(
218        actor_id: impl Into<ActorId>,
219        inputs: Vec<super::exchange::permit::Receiver>,
220        local_barrier_manager: crate::task::LocalBarrierManager,
221        schema: Schema,
222        chunk_size: usize,
223        barrier_rx: Option<mpsc::UnboundedReceiver<Barrier>>,
224    ) -> Self {
225        let actor_id = actor_id.into();
226        use super::exchange::input::LocalInput;
227        use crate::executor::exchange::input::ActorInput;
228
229        let barrier_rx =
230            barrier_rx.unwrap_or_else(|| local_barrier_manager.subscribe_barrier(actor_id));
231
232        let metrics = StreamingMetrics::unused();
233        let actor_ctx = ActorContext::for_test(actor_id);
234        let upstream = Self::new_merge_upstream(
235            inputs
236                .into_iter()
237                .enumerate()
238                .map(|(idx, input)| LocalInput::new(input, ActorId::new(idx as u32)).boxed_input())
239                .collect(),
240            &metrics,
241            &actor_ctx,
242            chunk_size,
243            schema,
244        );
245
246        Self::new(
247            actor_ctx,
248            514.into(),
249            1919.into(),
250            upstream,
251            local_barrier_manager,
252            metrics.into(),
253            barrier_rx,
254        )
255    }
256
257    pub(crate) fn new_merge_upstream(
258        upstreams: Vec<BoxedActorInput>,
259        metrics: &StreamingMetrics,
260        actor_context: &ActorContext,
261        chunk_size: usize,
262        schema: Schema,
263    ) -> MergeUpstream {
264        let merge_barrier_align_duration = Some(
265            metrics
266                .merge_barrier_align_duration
267                .with_guarded_label_values(&[
268                    &actor_context.id.to_string(),
269                    &actor_context.fragment_id.to_string(),
270                ]),
271        );
272
273        BufferChunks::new(
274            // Futures of all active upstreams.
275            SelectReceivers::new(upstreams, None, merge_barrier_align_duration),
276            chunk_size,
277            schema,
278        )
279    }
280}
281
282impl<U> MergeExecutorInner<U>
283where
284    U: Upstream,
285{
286    #[try_stream(ok = Message, error = StreamExecutorError)]
287    pub(crate) async fn execute_inner(mut self: Box<Self>) {
288        let mut upstream = self.upstream;
289        let actor_id = self.actor_context.id;
290
291        let mut metrics = self.metrics.new_actor_input_metrics(
292            actor_id,
293            self.fragment_id,
294            self.upstream_fragment_id,
295        );
296
297        let mut barrier_buffer = DispatchBarrierBuffer::new(
298            self.barrier_rx,
299            actor_id,
300            self.upstream_fragment_id,
301            self.local_barrier_manager,
302            self.metrics.clone(),
303            self.fragment_id,
304            self.actor_context.config.clone(),
305        );
306
307        loop {
308            let upstream_is_empty = upstream.is_empty();
309            let msg = barrier_buffer
310                .await_next_message(&mut upstream, &metrics, upstream_is_empty)
311                .await?;
312            let msg = match msg {
313                DispatcherMessage::Watermark(watermark) => Message::Watermark(watermark),
314                DispatcherMessage::Chunk(chunk) => {
315                    metrics.actor_in_record_cnt.inc_by(chunk.cardinality() as _);
316                    Message::Chunk(chunk)
317                }
318                DispatcherMessage::Barrier(barrier) => {
319                    let (barrier, new_inputs) =
320                        barrier_buffer.pop_barrier_with_inputs(barrier).await?;
321
322                    if let Some(Mutation::Update(UpdateMutation { dispatchers, .. })) =
323                        barrier.mutation.as_deref()
324                        && upstream
325                            .upstream_input_ids()
326                            .any(|actor_id| dispatchers.contains_key(&actor_id))
327                    {
328                        // `Watermark` of upstream may become stale after downstream scaling.
329                        upstream.flush_buffered_watermarks();
330                    }
331
332                    if let Some(update) =
333                        barrier.as_update_merge(self.actor_context.id, self.upstream_fragment_id)
334                    {
335                        let new_upstream_fragment_id = update
336                            .new_upstream_fragment_id
337                            .unwrap_or(self.upstream_fragment_id);
338                        let removed_upstream_actor_id: HashSet<_> =
339                            if update.new_upstream_fragment_id.is_some() {
340                                upstream.upstream_input_ids().collect()
341                            } else {
342                                update.removed_upstream_actor_id.iter().copied().collect()
343                            };
344
345                        // `Watermark` of upstream may become stale after upstream scaling.
346                        upstream.flush_buffered_watermarks();
347
348                        // Add and remove upstreams.
349                        upstream.update(new_inputs.unwrap_or_default(), &removed_upstream_actor_id);
350
351                        self.upstream_fragment_id = new_upstream_fragment_id;
352                        metrics = self.metrics.new_actor_input_metrics(
353                            actor_id,
354                            self.fragment_id,
355                            self.upstream_fragment_id,
356                        );
357                    }
358
359                    let is_stop = barrier.is_stop(actor_id);
360                    let msg = Message::Barrier(barrier);
361                    if is_stop {
362                        yield msg;
363                        break;
364                    }
365
366                    msg
367                }
368            };
369
370            yield msg;
371        }
372    }
373}
374
375impl<U> Execute for MergeExecutorInner<U>
376where
377    U: Upstream,
378{
379    fn execute(self: Box<Self>) -> BoxedMessageStream {
380        self.execute_inner().boxed()
381    }
382}
383
384/// A wrapper that buffers the `StreamChunk`s from upstream until no more ready items are available.
385/// Besides, any message other than `StreamChunk` will trigger the buffered `StreamChunk`s
386/// to be emitted immediately along with the message itself.
387pub struct BufferChunks<S: Stream> {
388    inner: S,
389    chunk_builder: StreamChunkBuilder,
390
391    /// The items to be emitted. Whenever there's something here, we should return a `Poll::Ready` immediately.
392    pending_items: VecDeque<S::Item>,
393}
394
395impl<S: Stream> BufferChunks<S> {
396    pub(super) fn new(inner: S, chunk_size: usize, schema: Schema) -> Self {
397        assert!(chunk_size > 0);
398        let chunk_builder = StreamChunkBuilder::new(chunk_size, schema.data_types());
399        Self {
400            inner,
401            chunk_builder,
402            pending_items: VecDeque::new(),
403        }
404    }
405}
406
407impl<S: Stream> std::ops::Deref for BufferChunks<S> {
408    type Target = S;
409
410    fn deref(&self) -> &Self::Target {
411        &self.inner
412    }
413}
414
415impl<S: Stream> std::ops::DerefMut for BufferChunks<S> {
416    fn deref_mut(&mut self) -> &mut Self::Target {
417        &mut self.inner
418    }
419}
420
421impl<S: Stream> Stream for BufferChunks<S>
422where
423    S: Stream<Item = DispatcherMessageStreamItem> + Unpin,
424{
425    type Item = S::Item;
426
427    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
428        loop {
429            if let Some(item) = self.pending_items.pop_front() {
430                return Poll::Ready(Some(item));
431            }
432
433            match self.inner.poll_next_unpin(cx) {
434                Poll::Pending => {
435                    return if let Some(chunk_out) = self.chunk_builder.take() {
436                        Poll::Ready(Some(Ok(MessageInner::Chunk(chunk_out))))
437                    } else {
438                        Poll::Pending
439                    };
440                }
441
442                Poll::Ready(Some(result)) => {
443                    if let Ok(MessageInner::Chunk(chunk)) = result {
444                        for row in chunk.records() {
445                            if let Some(chunk_out) = self.chunk_builder.append_record(row) {
446                                self.pending_items
447                                    .push_back(Ok(MessageInner::Chunk(chunk_out)));
448                            }
449                        }
450                    } else {
451                        return if let Some(chunk_out) = self.chunk_builder.take() {
452                            self.pending_items.push_back(result);
453                            Poll::Ready(Some(Ok(MessageInner::Chunk(chunk_out))))
454                        } else {
455                            Poll::Ready(Some(result))
456                        };
457                    }
458                }
459
460                Poll::Ready(None) => {
461                    return if let Some(chunk_out) = self.chunk_builder.take() {
462                        Poll::Ready(Some(Ok(MessageInner::Chunk(chunk_out))))
463                    } else {
464                        Poll::Pending
465                    };
466                }
467            }
468        }
469    }
470}
471
472#[cfg(test)]
473mod tests {
474    use std::sync::atomic::{AtomicBool, Ordering};
475    use std::task::Poll;
476    use std::time::Duration;
477
478    use assert_matches::assert_matches;
479    use futures::future::try_join_all;
480    use futures::{FutureExt, poll};
481    use risingwave_common::array::Op;
482    use risingwave_common::util::epoch::test_epoch;
483    use risingwave_pb::task_service::stream_exchange_service_server::{
484        StreamExchangeService, StreamExchangeServiceServer,
485    };
486    use risingwave_pb::task_service::{GetStreamRequest, GetStreamResponse, PbPermits};
487    use tokio::time::sleep;
488    use tokio_stream::wrappers::ReceiverStream;
489    use tonic::{Request, Response, Status, Streaming};
490
491    use super::*;
492    use crate::executor::exchange::input::{ActorInput, LocalInput, RemoteInput};
493    use crate::executor::exchange::permit::channel_for_test;
494    use crate::executor::{BarrierInner as Barrier, MessageInner as Message};
495    use crate::task::barrier_test_utils::LocalBarrierTestEnv;
496    use crate::task::test_utils::helper_make_local_actor;
497    use crate::task::{NewOutputRequest, TEST_PARTIAL_GRAPH_ID};
498
499    fn build_test_chunk(size: u64) -> StreamChunk {
500        let ops = vec![Op::Insert; size as usize];
501        StreamChunk::new(ops, vec![])
502    }
503
504    #[tokio::test]
505    async fn test_buffer_chunks() {
506        let test_env = LocalBarrierTestEnv::for_test().await;
507
508        let (tx, rx) = channel_for_test();
509        let input = LocalInput::new(rx, 1.into()).boxed_input();
510        let mut buffer = BufferChunks::new(input, 100, Schema::new(vec![]));
511
512        // Send a chunk
513        tx.send(Message::Chunk(build_test_chunk(10)).into())
514            .await
515            .unwrap();
516        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
517            assert_eq!(chunk.ops().len() as u64, 10);
518        });
519
520        // Send 2 chunks and expect them to be merged.
521        tx.send(Message::Chunk(build_test_chunk(10)).into())
522            .await
523            .unwrap();
524        tx.send(Message::Chunk(build_test_chunk(10)).into())
525            .await
526            .unwrap();
527        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
528            assert_eq!(chunk.ops().len() as u64, 20);
529        });
530
531        // Send a watermark.
532        tx.send(
533            Message::Watermark(Watermark {
534                col_idx: 0,
535                data_type: DataType::Int64,
536                val: ScalarImpl::Int64(233),
537            })
538            .into(),
539        )
540        .await
541        .unwrap();
542        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Watermark(watermark) => {
543            assert_eq!(watermark.val, ScalarImpl::Int64(233));
544        });
545
546        // Send 2 chunks before a watermark. Expect the 2 chunks to be merged and the watermark to be emitted.
547        tx.send(Message::Chunk(build_test_chunk(10)).into())
548            .await
549            .unwrap();
550        tx.send(Message::Chunk(build_test_chunk(10)).into())
551            .await
552            .unwrap();
553        tx.send(
554            Message::Watermark(Watermark {
555                col_idx: 0,
556                data_type: DataType::Int64,
557                val: ScalarImpl::Int64(233),
558            })
559            .into(),
560        )
561        .await
562        .unwrap();
563        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
564            assert_eq!(chunk.ops().len() as u64, 20);
565        });
566        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Watermark(watermark) => {
567            assert_eq!(watermark.val, ScalarImpl::Int64(233));
568        });
569
570        // Send a barrier.
571        let barrier = Barrier::new_test_barrier(test_epoch(1));
572        test_env.inject_barrier(&barrier, [2.into()]);
573        tx.send(Message::Barrier(barrier.clone().into_dispatcher()).into())
574            .await
575            .unwrap();
576        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Barrier(Barrier { epoch: barrier_epoch, .. }) => {
577            assert_eq!(barrier_epoch.curr, test_epoch(1));
578        });
579
580        // Send 2 chunks before a barrier. Expect the 2 chunks to be merged and the barrier to be emitted.
581        tx.send(Message::Chunk(build_test_chunk(10)).into())
582            .await
583            .unwrap();
584        tx.send(Message::Chunk(build_test_chunk(10)).into())
585            .await
586            .unwrap();
587        let barrier = Barrier::new_test_barrier(test_epoch(2));
588        test_env.inject_barrier(&barrier, [2.into()]);
589        tx.send(Message::Barrier(barrier.clone().into_dispatcher()).into())
590            .await
591            .unwrap();
592        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
593            assert_eq!(chunk.ops().len() as u64, 20);
594        });
595        assert_matches!(buffer.next().await.unwrap().unwrap(), Message::Barrier(Barrier { epoch: barrier_epoch, .. }) => {
596            assert_eq!(barrier_epoch.curr, test_epoch(2));
597        });
598    }
599
600    #[tokio::test]
601    async fn test_merger() {
602        const CHANNEL_NUMBER: usize = 10;
603        let mut txs = Vec::with_capacity(CHANNEL_NUMBER);
604        let mut rxs = Vec::with_capacity(CHANNEL_NUMBER);
605        for _i in 0..CHANNEL_NUMBER {
606            let (tx, rx) = channel_for_test();
607            txs.push(tx);
608            rxs.push(rx);
609        }
610        let barrier_test_env = LocalBarrierTestEnv::for_test().await;
611        let actor_id = 233.into();
612        let mut handles = Vec::with_capacity(CHANNEL_NUMBER);
613
614        let epochs = (10..1000u64)
615            .step_by(10)
616            .map(|idx| (idx, test_epoch(idx)))
617            .collect_vec();
618        let mut prev_epoch = 0;
619        let prev_epoch = &mut prev_epoch;
620        let barriers: HashMap<_, _> = epochs
621            .iter()
622            .map(|(_, epoch)| {
623                let barrier = Barrier::with_prev_epoch_for_test(*epoch, *prev_epoch);
624                *prev_epoch = *epoch;
625                barrier_test_env.inject_barrier(&barrier, [actor_id]);
626                (*epoch, barrier)
627            })
628            .collect();
629        let b2 = Barrier::with_prev_epoch_for_test(test_epoch(1000), *prev_epoch)
630            .with_mutation(Mutation::Stop(StopMutation::default()));
631        barrier_test_env.inject_barrier(&b2, [actor_id]);
632        barrier_test_env.flush_all_events().await;
633
634        for (tx_id, tx) in txs.into_iter().enumerate() {
635            let epochs = epochs.clone();
636            let barriers = barriers.clone();
637            let b2 = b2.clone();
638            let handle = tokio::spawn(async move {
639                for (idx, epoch) in epochs {
640                    if idx % 20 == 0 {
641                        tx.send(Message::Chunk(build_test_chunk(10)).into())
642                            .await
643                            .unwrap();
644                    } else {
645                        tx.send(
646                            Message::Watermark(Watermark {
647                                col_idx: (idx as usize / 20 + tx_id) % CHANNEL_NUMBER,
648                                data_type: DataType::Int64,
649                                val: ScalarImpl::Int64(idx as i64),
650                            })
651                            .into(),
652                        )
653                        .await
654                        .unwrap();
655                    }
656                    tx.send(Message::Barrier(barriers[&epoch].clone().into_dispatcher()).into())
657                        .await
658                        .unwrap();
659                    sleep(Duration::from_millis(1)).await;
660                }
661                tx.send(Message::Barrier(b2.clone().into_dispatcher()).into())
662                    .await
663                    .unwrap();
664            });
665            handles.push(handle);
666        }
667
668        let merger = MergeExecutor::for_test(
669            actor_id,
670            rxs,
671            barrier_test_env.local_barrier_manager.clone(),
672            Schema::new(vec![]),
673            100,
674            None,
675        );
676        let mut merger = merger.boxed().execute();
677        for (idx, epoch) in epochs {
678            if idx % 20 == 0 {
679                // expect 1 or more chunks with 100 rows in total
680                let mut count = 0usize;
681                while count < 100 {
682                    assert_matches!(merger.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
683                        count += chunk.ops().len();
684                    });
685                }
686                assert_eq!(count, 100);
687            } else if idx as usize / 20 >= CHANNEL_NUMBER - 1 {
688                // expect n watermarks
689                for _ in 0..CHANNEL_NUMBER {
690                    assert_matches!(merger.next().await.unwrap().unwrap(), Message::Watermark(watermark) => {
691                        assert_eq!(watermark.val, ScalarImpl::Int64((idx - 20 * (CHANNEL_NUMBER as u64 - 1)) as i64));
692                    });
693                }
694            }
695            // expect a barrier
696            assert_matches!(merger.next().await.unwrap().unwrap(), Message::Barrier(Barrier{epoch:barrier_epoch,..}) => {
697                assert_eq!(barrier_epoch.curr, epoch);
698            });
699        }
700        assert_matches!(
701            merger.next().await.unwrap().unwrap(),
702            Message::Barrier(Barrier {
703                mutation,
704                ..
705            }) if mutation.as_deref().unwrap().is_stop_mutation()
706        );
707
708        for handle in handles {
709            handle.await.unwrap();
710        }
711    }
712
713    #[tokio::test]
714    async fn empty_dynamic_merge_progresses_across_attach_detach_and_reattach() {
715        let actor_id = 999.into();
716        let resolver_1 = 1000.into();
717        let resolver_2 = 1001.into();
718        let empty_fragment = 0.into();
719        let resolver_fragment_1 = 500.into();
720        let resolver_fragment_2 = 501.into();
721        let barrier_test_env = LocalBarrierTestEnv::for_test().await;
722        let e1 = Barrier::new_test_barrier(test_epoch(1));
723        barrier_test_env.inject_barrier(&e1, [actor_id]);
724        barrier_test_env.flush_all_events().await;
725
726        let metrics = Arc::new(StreamingMetrics::unused());
727        let actor_ctx = ActorContext::for_test(actor_id);
728        let barrier_rx = barrier_test_env
729            .local_barrier_manager
730            .subscribe_barrier(actor_id);
731        let upstream = MergeExecutor::new_merge_upstream(
732            vec![],
733            &metrics,
734            &actor_ctx,
735            100,
736            Schema::empty().clone(),
737        );
738        let merge = MergeExecutor::new(
739            actor_ctx,
740            514.into(),
741            empty_fragment,
742            upstream,
743            barrier_test_env.local_barrier_manager.clone(),
744            metrics,
745            barrier_rx,
746        );
747        let mut merge = merge.boxed().execute();
748
749        macro_rules! assert_recv_pending {
750            () => {
751                assert_matches!(
752                    poll!(merge.as_mut().next()),
753                    Poll::Pending,
754                    "an empty dynamic merge must wait for a barrier rather than terminate or busy-loop"
755                );
756            };
757        }
758
759        macro_rules! recv_barrier {
760            ($epoch:expr) => {
761                assert_matches!(
762                    merge.next().await.unwrap().unwrap(),
763                    Message::Barrier(Barrier { epoch, .. }) if epoch.curr == test_epoch($epoch)
764                );
765            };
766        }
767
768        async fn take_upstream_tx(
769            barrier_test_env: &LocalBarrierTestEnv,
770            upstream_actor_id: ActorId,
771            downstream_actor_id: ActorId,
772        ) -> crate::executor::exchange::permit::Sender {
773            let mut requests = barrier_test_env
774                .take_pending_new_output_requests(upstream_actor_id)
775                .await;
776            assert_eq!(requests.len(), 1);
777            let (actor_id, request) = requests.pop().unwrap();
778            assert_eq!(actor_id, downstream_actor_id);
779            let NewOutputRequest::Local(tx) = request else {
780                unreachable!()
781            };
782            tx
783        }
784
785        assert_recv_pending!();
786
787        recv_barrier!(1);
788        assert_recv_pending!();
789
790        let b1 = Barrier::new_test_barrier(test_epoch(2)).with_mutation(Mutation::Update(
791            UpdateMutation {
792                merges: maplit::hashmap! {
793                    (actor_id, empty_fragment) => MergeUpdate {
794                        actor_id,
795                        upstream_fragment_id: empty_fragment,
796                        new_upstream_fragment_id: Some(resolver_fragment_1),
797                        added_upstream_actors: vec![helper_make_local_actor(resolver_1)],
798                        removed_upstream_actor_id: vec![],
799                    },
800                },
801                ..Default::default()
802            },
803        ));
804        barrier_test_env.inject_barrier(&b1, [actor_id]);
805        barrier_test_env.flush_all_events().await;
806        assert_recv_pending!();
807        barrier_test_env.flush_all_events().await;
808        let resolver_1_tx = take_upstream_tx(&barrier_test_env, resolver_1, actor_id).await;
809        resolver_1_tx
810            .send(Message::Barrier(b1.clone().into_dispatcher()).into())
811            .await
812            .unwrap();
813        recv_barrier!(2);
814
815        resolver_1_tx
816            .send(Message::Chunk(build_test_chunk(1)).into())
817            .await
818            .unwrap();
819        assert_matches!(merge.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
820            assert_eq!(chunk.cardinality(), 1);
821        });
822
823        let b2 = Barrier::new_test_barrier(test_epoch(3)).with_mutation(Mutation::Update(
824            UpdateMutation {
825                merges: maplit::hashmap! {
826                    (actor_id, resolver_fragment_1) => MergeUpdate {
827                        actor_id,
828                        upstream_fragment_id: resolver_fragment_1,
829                        new_upstream_fragment_id: Some(empty_fragment),
830                        added_upstream_actors: vec![],
831                        removed_upstream_actor_id: vec![resolver_1],
832                    },
833                },
834                ..Default::default()
835            },
836        ));
837        barrier_test_env.inject_barrier(&b2, [actor_id]);
838        resolver_1_tx
839            .send(Message::Barrier(b2.clone().into_dispatcher()).into())
840            .await
841            .unwrap();
842        recv_barrier!(3);
843        assert_recv_pending!();
844
845        let e3 = Barrier::new_test_barrier(test_epoch(4));
846        barrier_test_env.inject_barrier(&e3, [actor_id]);
847        recv_barrier!(4);
848        assert_recv_pending!();
849
850        let b4 = Barrier::new_test_barrier(test_epoch(5)).with_mutation(Mutation::Update(
851            UpdateMutation {
852                merges: maplit::hashmap! {
853                    (actor_id, empty_fragment) => MergeUpdate {
854                        actor_id,
855                        upstream_fragment_id: empty_fragment,
856                        new_upstream_fragment_id: Some(resolver_fragment_2),
857                        added_upstream_actors: vec![helper_make_local_actor(resolver_2)],
858                        removed_upstream_actor_id: vec![],
859                    },
860                },
861                ..Default::default()
862            },
863        ));
864        barrier_test_env.inject_barrier(&b4, [actor_id]);
865        barrier_test_env.flush_all_events().await;
866        assert_recv_pending!();
867        barrier_test_env.flush_all_events().await;
868        let resolver_2_tx = take_upstream_tx(&barrier_test_env, resolver_2, actor_id).await;
869        resolver_2_tx
870            .send(Message::Barrier(b4.into_dispatcher()).into())
871            .await
872            .unwrap();
873        recv_barrier!(5);
874
875        resolver_2_tx
876            .send(Message::Chunk(build_test_chunk(1)).into())
877            .await
878            .unwrap();
879        assert_matches!(merge.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
880            assert_eq!(chunk.cardinality(), 1);
881        });
882    }
883
884    #[tokio::test]
885    async fn test_configuration_change() {
886        let actor_id = 233.into();
887        let (untouched, old, new) = (234.into(), 235.into(), 238.into()); // upstream actors
888        let barrier_test_env = LocalBarrierTestEnv::for_test().await;
889        let metrics = Arc::new(StreamingMetrics::unused());
890
891        // untouched -> actor_id
892        // old -> actor_id
893        // new -> actor_id
894
895        let (upstream_fragment_id, fragment_id) = (10.into(), 18.into());
896
897        let actor_ctx = ActorContext::for_test(actor_id);
898
899        let inputs: Vec<_> =
900            try_join_all([untouched, old].into_iter().map(async |upstream_actor_id| {
901                new_input(
902                    &barrier_test_env.local_barrier_manager,
903                    metrics.clone(),
904                    actor_id,
905                    fragment_id,
906                    &helper_make_local_actor(upstream_actor_id),
907                    upstream_fragment_id,
908                    actor_ctx.config.clone(),
909                )
910                .await
911            }))
912            .await
913            .unwrap();
914
915        let merge_updates = maplit::hashmap! {
916            (actor_id, upstream_fragment_id) => MergeUpdate {
917                actor_id,
918                upstream_fragment_id,
919                new_upstream_fragment_id: None,
920                added_upstream_actors: vec![helper_make_local_actor(new)],
921                removed_upstream_actor_id: vec![old],
922            }
923        };
924
925        let b1 = Barrier::new_test_barrier(test_epoch(1)).with_mutation(Mutation::Update(
926            UpdateMutation {
927                merges: merge_updates,
928                ..Default::default()
929            },
930        ));
931        barrier_test_env.inject_barrier(&b1, [actor_id]);
932        barrier_test_env.flush_all_events().await;
933
934        let barrier_rx = barrier_test_env
935            .local_barrier_manager
936            .subscribe_barrier(actor_id);
937        let upstream = MergeExecutor::new_merge_upstream(
938            inputs,
939            &metrics,
940            &actor_ctx,
941            100,
942            Schema::empty().clone(),
943        );
944
945        let mut merge = MergeExecutor::new(
946            actor_ctx,
947            fragment_id,
948            upstream_fragment_id,
949            upstream,
950            barrier_test_env.local_barrier_manager.clone(),
951            metrics.clone(),
952            barrier_rx,
953        )
954        .boxed()
955        .execute();
956
957        let mut txs = HashMap::new();
958        macro_rules! send {
959            ($actors:expr, $msg:expr) => {
960                for actor in $actors {
961                    txs.get(&actor).unwrap().send($msg).await.unwrap();
962                }
963            };
964        }
965
966        macro_rules! assert_recv_pending {
967            () => {
968                assert!(
969                    merge
970                        .next()
971                        .now_or_never()
972                        .flatten()
973                        .transpose()
974                        .unwrap()
975                        .is_none()
976                );
977            };
978        }
979        macro_rules! recv {
980            () => {
981                merge.next().await.transpose().unwrap()
982            };
983        }
984
985        macro_rules! collect_upstream_tx {
986            ($actors:expr) => {
987                for upstream_id in $actors {
988                    let mut output_requests = barrier_test_env
989                        .take_pending_new_output_requests(upstream_id)
990                        .await;
991                    assert_eq!(output_requests.len(), 1);
992                    let (downstream_actor_id, request) = output_requests.pop().unwrap();
993                    assert_eq!(downstream_actor_id, actor_id);
994                    let NewOutputRequest::Local(tx) = request else {
995                        unreachable!()
996                    };
997                    txs.insert(upstream_id, tx);
998                }
999            };
1000        }
1001
1002        assert_recv_pending!();
1003        barrier_test_env.flush_all_events().await;
1004
1005        // 2. Take downstream receivers.
1006        collect_upstream_tx!([untouched, old]);
1007
1008        // 3. Send a chunk.
1009        send!([untouched, old], Message::Chunk(build_test_chunk(1)).into());
1010        assert_eq!(2, recv!().unwrap().as_chunk().unwrap().cardinality()); // We should be able to receive the chunk twice.
1011        assert_recv_pending!();
1012
1013        send!(
1014            [untouched, old],
1015            Message::Barrier(b1.clone().into_dispatcher()).into()
1016        );
1017        assert_recv_pending!(); // We should not receive the barrier, since merger is waiting for the new upstream new.
1018
1019        collect_upstream_tx!([new]);
1020
1021        send!([new], Message::Barrier(b1.clone().into_dispatcher()).into());
1022        recv!().unwrap().as_barrier().unwrap(); // We should now receive the barrier.
1023
1024        // 5. Send a chunk.
1025        send!([untouched, new], Message::Chunk(build_test_chunk(1)).into());
1026        assert_eq!(2, recv!().unwrap().as_chunk().unwrap().cardinality()); // We should be able to receive the chunk twice.
1027        assert_recv_pending!();
1028    }
1029
1030    struct FakeExchangeService {
1031        rpc_called: Arc<AtomicBool>,
1032    }
1033
1034    fn exchange_client_test_barrier() -> crate::executor::Barrier {
1035        Barrier::new_test_barrier(test_epoch(1))
1036    }
1037
1038    #[async_trait::async_trait]
1039    impl StreamExchangeService for FakeExchangeService {
1040        type GetStreamStream = ReceiverStream<std::result::Result<GetStreamResponse, Status>>;
1041
1042        async fn get_stream(
1043            &self,
1044            _request: Request<Streaming<GetStreamRequest>>,
1045        ) -> std::result::Result<Response<Self::GetStreamStream>, Status> {
1046            let (tx, rx) = tokio::sync::mpsc::channel(10);
1047            self.rpc_called.store(true, Ordering::SeqCst);
1048            // send stream_chunk
1049            let stream_chunk = StreamChunk::default().to_protobuf();
1050            tx.send(Ok(GetStreamResponse {
1051                message: Some(PbStreamMessageBatch {
1052                    stream_message_batch: Some(
1053                        risingwave_pb::stream_plan::stream_message_batch::StreamMessageBatch::StreamChunk(
1054                            stream_chunk,
1055                        ),
1056                    ),
1057                }),
1058                permits: Some(PbPermits::default()),
1059            }))
1060            .await
1061            .unwrap();
1062            // send barrier
1063            let barrier = exchange_client_test_barrier();
1064            tx.send(Ok(GetStreamResponse {
1065                message: Some(PbStreamMessageBatch {
1066                    stream_message_batch: Some(
1067                        risingwave_pb::stream_plan::stream_message_batch::StreamMessageBatch::BarrierBatch(
1068                            BarrierBatch {
1069                                barriers: vec![barrier.to_protobuf()],
1070                            },
1071                        ),
1072                    ),
1073                }),
1074                permits: Some(PbPermits::default()),
1075            }))
1076            .await
1077            .unwrap();
1078            Ok(Response::new(ReceiverStream::new(rx)))
1079        }
1080    }
1081
1082    #[tokio::test]
1083    async fn test_stream_exchange_client() {
1084        let rpc_called = Arc::new(AtomicBool::new(false));
1085        let server_run = Arc::new(AtomicBool::new(false));
1086        let addr = "127.0.0.1:12348".parse().unwrap();
1087
1088        // Start a server.
1089        let (shutdown_send, shutdown_recv) = tokio::sync::oneshot::channel();
1090        let exchange_svc = StreamExchangeServiceServer::new(FakeExchangeService {
1091            rpc_called: rpc_called.clone(),
1092        });
1093        let cp_server_run = server_run.clone();
1094        let join_handle = tokio::spawn(async move {
1095            cp_server_run.store(true, Ordering::SeqCst);
1096            tonic::transport::Server::builder()
1097                .add_service(exchange_svc)
1098                .serve_with_shutdown(addr, async move {
1099                    shutdown_recv.await.unwrap();
1100                })
1101                .await
1102                .unwrap();
1103        });
1104
1105        sleep(Duration::from_secs(1)).await;
1106        assert!(server_run.load(Ordering::SeqCst));
1107
1108        let test_env = LocalBarrierTestEnv::for_test().await;
1109
1110        let remote_input = {
1111            RemoteInput::new(
1112                &test_env.local_barrier_manager,
1113                addr.into(),
1114                TEST_PARTIAL_GRAPH_ID,
1115                (0.into(), 0.into()),
1116                (0.into(), 0.into()),
1117                Arc::new(StreamingMetrics::unused()),
1118                Arc::new(StreamingConfig::default()),
1119            )
1120            .await
1121            .unwrap()
1122        };
1123
1124        test_env.inject_barrier(&exchange_client_test_barrier(), [remote_input.id()]);
1125
1126        pin_mut!(remote_input);
1127
1128        assert_matches!(remote_input.next().await.unwrap().unwrap(), Message::Chunk(chunk) => {
1129            let (ops, columns, visibility) = chunk.into_inner();
1130            assert!(ops.is_empty());
1131            assert!(columns.is_empty());
1132            assert!(visibility.is_empty());
1133        });
1134        assert_matches!(remote_input.next().await.unwrap().unwrap(), Message::Barrier(Barrier { epoch: barrier_epoch, .. }) => {
1135            assert_eq!(barrier_epoch.curr, test_epoch(1));
1136        });
1137        assert!(rpc_called.load(Ordering::SeqCst));
1138
1139        shutdown_send.send(()).unwrap();
1140        join_handle.await.unwrap();
1141    }
1142}