Skip to main content

risingwave_compute/rpc/service/
batch_exchange_service.rs

1// Copyright 2025 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::pin::Pin;
16use std::sync::Arc;
17use std::task::{Context, Poll};
18
19use risingwave_batch::task::BatchManager;
20use risingwave_common_service::GrpcCall;
21use risingwave_pb::task_service::batch_exchange_service_server::BatchExchangeService;
22use risingwave_pb::task_service::{GetDataRequest, GetDataResponse};
23use thiserror_ext::AsReport;
24use tokio_stream::wrappers::ReceiverStream;
25use tonic::{Request, Response, Status};
26
27type BatchData = std::result::Result<GetDataResponse, Status>;
28
29pub struct BatchDataStream {
30    inner: ReceiverStream<BatchData>,
31    await_tree_root: Option<await_tree::TreeRoot>,
32}
33
34impl BatchDataStream {
35    fn new(
36        receiver: tokio::sync::mpsc::Receiver<BatchData>,
37        await_tree_root: Option<await_tree::TreeRoot>,
38    ) -> Self {
39        Self {
40            inner: ReceiverStream::new(receiver),
41            await_tree_root,
42        }
43    }
44}
45
46impl tokio_stream::Stream for BatchDataStream {
47    type Item = BatchData;
48
49    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
50        let this = self.get_mut();
51        let result = Pin::new(&mut this.inner).poll_next(cx);
52        if matches!(&result, Poll::Ready(None)) {
53            this.await_tree_root.take();
54        }
55        result
56    }
57}
58
59#[derive(Clone)]
60pub struct BatchExchangeServiceImpl {
61    batch_mgr: Arc<BatchManager>,
62}
63
64impl BatchExchangeServiceImpl {
65    pub fn new(batch_mgr: Arc<BatchManager>) -> Self {
66        Self { batch_mgr }
67    }
68}
69
70#[async_trait::async_trait]
71impl BatchExchangeService for BatchExchangeServiceImpl {
72    type GetDataStream = BatchDataStream;
73
74    async fn get_data(
75        &self,
76        request: Request<GetDataRequest>,
77    ) -> std::result::Result<Response<Self::GetDataStream>, Status> {
78        let peer_addr = request
79            .remote_addr()
80            .ok_or_else(|| Status::unavailable("connection unestablished"))?;
81        let pb_task_output_id = request
82            .into_inner()
83            .task_output_id
84            .expect("Failed to get task output id.");
85        let (tx, rx) =
86            tokio::sync::mpsc::channel(self.batch_mgr.config().developer.receiver_channel_size);
87        if let Err(e) = self.batch_mgr.get_data(tx, peer_addr, &pb_task_output_id) {
88            error!(
89                %peer_addr,
90                error = %e.as_report(),
91                "Failed to serve exchange RPC"
92            );
93            return Err(e.into());
94        }
95
96        let await_tree_root = self.batch_mgr.await_tree_reg().map(|registry| {
97            let key = GrpcCall::new(format!(
98                "{peer_addr} - /task_service.BatchExchangeService/GetData - \
99                 {pb_task_output_id:?}"
100            ));
101            registry.register(key, "/task_service.BatchExchangeService/GetData")
102        });
103        Ok(Response::new(BatchDataStream::new(rx, await_tree_root)))
104    }
105}