risingwave_compute/rpc/service/
batch_exchange_service.rs1use 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}