Skip to main content

risingwave_batch/rpc/service/
task_service.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::future::Future;
16use std::sync::Arc;
17
18use anyhow::Context;
19use futures::stream::{FuturesOrdered, StreamExt};
20use risingwave_common::array::StreamChunk;
21use risingwave_common::util::tracing::TracingContext;
22use risingwave_dml::TableDmlHandleRef;
23use risingwave_dml::dml_manager::DmlManagerRef;
24use risingwave_dml::error::DmlError;
25use risingwave_pb::batch_plan::TaskOutputId;
26use risingwave_pb::task_service::task_service_server::TaskService;
27use risingwave_pb::task_service::{
28    CancelTaskRequest, CancelTaskResponse, CreateTaskRequest, ExecuteRequest, FastInsertRequest,
29    FastInsertResponse, GetDataResponse, IngestDmlAckResponse, IngestDmlInitRequest,
30    IngestDmlInitResponse, IngestDmlPayloadRequest, IngestDmlRequest, IngestDmlResponse,
31    TaskInfoResponse, fast_insert_response, ingest_dml_request, ingest_dml_response,
32};
33use thiserror_ext::AsReport;
34use tokio_stream::wrappers::ReceiverStream;
35use tonic::{Request, Response, Status};
36
37use crate::error::BatchError;
38use crate::executor::{FastInsertExecutor, inject_optional_row_id_column};
39use crate::rpc::service::exchange::GrpcExchangeWriter;
40use crate::task::{
41    BatchEnvironment, BatchManager, BatchTaskExecution, ComputeNodeContext, StateReporter,
42    TASK_STATUS_BUFFER_SIZE,
43};
44
45#[derive(Clone)]
46pub struct BatchServiceImpl {
47    mgr: Arc<BatchManager>,
48    env: BatchEnvironment,
49}
50
51impl BatchServiceImpl {
52    pub fn new(mgr: Arc<BatchManager>, env: BatchEnvironment) -> Self {
53        BatchServiceImpl { mgr, env }
54    }
55}
56
57pub type TaskInfoResponseResult = Result<TaskInfoResponse, Status>;
58pub type GetDataResponseResult = Result<GetDataResponse, Status>;
59pub type IngestDmlResponseResult = Result<IngestDmlResponse, Status>;
60
61#[async_trait::async_trait]
62impl TaskService for BatchServiceImpl {
63    type CreateTaskStream = ReceiverStream<TaskInfoResponseResult>;
64    type ExecuteStream = ReceiverStream<GetDataResponseResult>;
65    type IngestDmlStream = ReceiverStream<IngestDmlResponseResult>;
66
67    async fn create_task(
68        &self,
69        request: Request<CreateTaskRequest>,
70    ) -> Result<Response<Self::CreateTaskStream>, Status> {
71        #[cfg(madsim)]
72        crate::rpc::service::madsim_test_utils::record_create_task();
73
74        let CreateTaskRequest {
75            task_id,
76            plan,
77            tracing_context,
78            expr_context,
79        } = request.into_inner();
80
81        let (state_tx, state_rx) = tokio::sync::mpsc::channel(TASK_STATUS_BUFFER_SIZE);
82        let state_reporter = StateReporter::new_with_dist_sender(state_tx);
83        let res = self
84            .mgr
85            .fire_task(
86                task_id.as_ref().expect("no task id found"),
87                plan.expect("no plan found").clone(),
88                ComputeNodeContext::create(self.env.clone()),
89                state_reporter,
90                TracingContext::from_protobuf(&tracing_context),
91                expr_context.expect("no expression context found"),
92            )
93            .await;
94        match res {
95            Ok(_) => Ok(Response::new(ReceiverStream::new(
96                // Create receiver stream from state receiver.
97                // The state receiver is init in `.async_execute()`.
98                // Will be used for receive task status update.
99                // Note: we introduce this hack cuz `.execute()` do not produce a status stream,
100                // but still share `.async_execute()` and `.try_execute()`.
101                state_rx,
102            ))),
103            Err(e) => {
104                error!(error = %e.as_report(), "failed to fire task");
105                Err(e.into())
106            }
107        }
108    }
109
110    async fn cancel_task(
111        &self,
112        req: Request<CancelTaskRequest>,
113    ) -> Result<Response<CancelTaskResponse>, Status> {
114        let req = req.into_inner();
115        tracing::trace!("Aborting task: {:?}", req.get_task_id().unwrap());
116        self.mgr
117            .cancel_task(req.get_task_id().expect("no task id found"));
118        Ok(Response::new(CancelTaskResponse { status: None }))
119    }
120
121    async fn execute(
122        &self,
123        req: Request<ExecuteRequest>,
124    ) -> Result<Response<Self::ExecuteStream>, Status> {
125        let req = req.into_inner();
126        let env = self.env.clone();
127        let mgr = self.mgr.clone();
128        BatchServiceImpl::get_execute_stream(env, mgr, req).await
129    }
130
131    async fn fast_insert(
132        &self,
133        request: Request<FastInsertRequest>,
134    ) -> Result<Response<FastInsertResponse>, Status> {
135        let req = request.into_inner();
136        let res = self.do_fast_insert(req).await;
137        match res {
138            Ok(_) => Ok(Response::new(FastInsertResponse {
139                status: fast_insert_response::Status::Succeeded.into(),
140                error_message: "".to_owned(),
141            })),
142            Err(e) => match e {
143                BatchError::Dml(e) => Ok(Response::new(FastInsertResponse {
144                    status: fast_insert_response::Status::DmlFailed.into(),
145                    error_message: format!("{}", e.as_report()),
146                })),
147                _ => {
148                    error!(error = %e.as_report(), "failed to fast insert");
149                    Err(e.into())
150                }
151            },
152        }
153    }
154
155    async fn ingest_dml(
156        &self,
157        request: Request<tonic::Streaming<IngestDmlRequest>>,
158    ) -> Result<Response<Self::IngestDmlStream>, Status> {
159        let mut req_stream = request.into_inner();
160        let init = match req_stream.message().await? {
161            Some(req) => match req.request {
162                Some(ingest_dml_request::Request::Init(init)) => init,
163                Some(ingest_dml_request::Request::Payload(_)) => {
164                    return Err(Status::invalid_argument(
165                        "first ingest dml message must be init",
166                    ));
167                }
168                None => return Err(Status::invalid_argument("empty ingest dml request")),
169            },
170            None => return Err(Status::invalid_argument("empty ingest dml stream")),
171        };
172
173        let (tx, rx) = tokio::sync::mpsc::channel(64);
174
175        let (table_dml_handle, request_id, row_id_index) = self.init_ingest_dml(&init)?;
176        let _ = tx.send(Ok(Self::ingest_dml_init_response())).await;
177
178        let dml_manager = self.env.dml_manager_ref();
179        tokio::spawn(async move {
180            let result: Result<(), String> = async {
181                let mut pending_acks = FuturesOrdered::new();
182
183                loop {
184                    tokio::select! {
185                        req = req_stream.message() => {
186                            let req = req
187                                .map_err(|err| format!("ingest dml stream read failed: {}", err.as_report()))?
188                                .ok_or_else(|| "ingest dml stream closed unexpectedly".to_owned())?;
189                            let payload = match req.request {
190                                Some(ingest_dml_request::Request::Payload(payload)) => payload,
191                                Some(ingest_dml_request::Request::Init(_)) | None => {
192                                    Err("unexpected non-payload request in ingest dml stream".to_owned())?
193                                }
194                            };
195
196                            let dml_batch_id = payload.dml_batch_id;
197                            let wait_fut = Self::do_ingest_dml_payload(
198                                table_dml_handle.clone(),
199                                dml_manager.clone(),
200                                request_id,
201                                row_id_index,
202                                payload,
203                            )
204                            .await
205                            .map_err(|err| format!("ingest dml batch {} failed: {}", dml_batch_id, err.as_report()))?;
206
207                            pending_acks.push_back(async move { wait_fut.await.map(|()| dml_batch_id) });
208                        }
209                        ack = pending_acks.next(), if !pending_acks.is_empty() => {
210                            let ack_dml_batch_id = ack
211                                .expect("branch guarded by non-empty pending_acks")
212                                .map_err(|err: DmlError| format!("ingest dml persistence failed: {}", err.as_report()))?;
213
214                            if tx
215                                .send(Ok(BatchServiceImpl::ingest_dml_ack_response(ack_dml_batch_id)))
216                                .await
217                                .is_err()
218                            {
219                                return Ok(());
220                            }
221                        }
222                    }
223                }
224            }
225            .await;
226
227            if let Err(err) = result {
228                let _ = tx.send(Err(Status::internal(err))).await;
229            }
230        });
231
232        Ok(Response::new(ReceiverStream::new(rx)))
233    }
234}
235
236impl BatchServiceImpl {
237    async fn get_execute_stream(
238        env: BatchEnvironment,
239        mgr: Arc<BatchManager>,
240        req: ExecuteRequest,
241    ) -> Result<Response<ReceiverStream<GetDataResponseResult>>, Status> {
242        let ExecuteRequest {
243            task_id,
244            plan,
245            tracing_context,
246            expr_context,
247        } = req;
248
249        let task_id = task_id.expect("no task id found");
250        let plan = plan.expect("no plan found").clone();
251        let tracing_context = TracingContext::from_protobuf(&tracing_context);
252        let expr_context = expr_context.expect("no expression context found");
253
254        let context = ComputeNodeContext::create(env.clone());
255        trace!(
256            "local execute request: plan:{:?} with task id:{:?}",
257            plan, task_id
258        );
259        let task = BatchTaskExecution::new(
260            &task_id,
261            plan,
262            context,
263            mgr.runtime(),
264            mgr.await_tree_reg().cloned(),
265        )?;
266        let task = Arc::new(task);
267        let (tx, rx) = tokio::sync::mpsc::channel(mgr.config().developer.local_execute_buffer_size);
268        if let Err(e) = task
269            .clone()
270            .async_execute(None, tracing_context, expr_context)
271            .await
272        {
273            error!(
274                error = %e.as_report(),
275                ?task_id,
276                "failed to build executors and trigger execution"
277            );
278            return Err(e.into());
279        }
280
281        let pb_task_output_id = TaskOutputId {
282            task_id: Some(task_id.clone()),
283            // Since this is local execution path, the exchange would follow single distribution,
284            // therefore we would only have one data output.
285            output_id: 0,
286        };
287        let mut output = task.get_task_output(&pb_task_output_id).inspect_err(|e| {
288            error!(
289                error = %e.as_report(),
290                ?task_id,
291                "failed to get task output in local execution mode",
292            );
293        })?;
294        let mut writer = GrpcExchangeWriter::new(tx.clone());
295        // Always spawn a task and do not block current function.
296        mgr.runtime().spawn(async move {
297            match output.take_data(&mut writer).await {
298                Ok(_) => Ok(()),
299                Err(e) => tx.send(Err(e.into())).await,
300            }
301        });
302        Ok(Response::new(ReceiverStream::new(rx)))
303    }
304
305    async fn do_fast_insert(&self, insert_req: FastInsertRequest) -> Result<(), BatchError> {
306        let wait_for_persistence = insert_req.wait_for_persistence;
307        let (executor, data_chunk) =
308            FastInsertExecutor::build(self.env.dml_manager_ref(), insert_req)?;
309        executor
310            .do_execute(data_chunk, wait_for_persistence)
311            .await?;
312        Ok(())
313    }
314
315    fn init_ingest_dml(
316        &self,
317        init: &IngestDmlInitRequest,
318    ) -> Result<(TableDmlHandleRef, u32, Option<u32>), Status> {
319        let table_id = init.table_id;
320        let table_version_id = init.table_version_id;
321        let table_dml_handle = self
322            .env
323            .dml_manager_ref()
324            .table_dml_handle(table_id, table_version_id)
325            .map_err(|err| Status::internal(format!("{}", err.as_report())))?;
326        Ok((table_dml_handle, init.request_id, init.row_id_index))
327    }
328
329    async fn do_ingest_dml_payload(
330        table_dml_handle: TableDmlHandleRef,
331        dml_manager: DmlManagerRef,
332        request_id: u32,
333        row_id_index: Option<u32>,
334        payload: IngestDmlPayloadRequest,
335    ) -> Result<impl Future<Output = risingwave_dml::error::Result<()>> + Send + 'static, BatchError>
336    {
337        let pb_chunk = payload.chunk.ok_or_else(|| {
338            BatchError::Internal(anyhow::anyhow!("no chunk in IngestDmlPayloadRequest"))
339        })?;
340        let mut chunk = StreamChunk::from_protobuf(&pb_chunk)
341            .context("failed to decode chunk")
342            .map_err(BatchError::Internal)?;
343        chunk = inject_optional_row_id_column(chunk, row_id_index.map(|index| index as usize));
344        let txn_id = dml_manager.gen_txn_id();
345        let mut write_handle = table_dml_handle
346            .write_handle(request_id, txn_id)
347            .map_err(BatchError::Dml)?;
348
349        write_handle.begin().map_err(BatchError::Dml)?;
350        write_handle
351            .write_chunk(chunk)
352            .await
353            .map_err(BatchError::Dml)?;
354        let persistence_future = write_handle
355            .end_wait_persistence()
356            .map_err(BatchError::Dml)?;
357        Ok(persistence_future)
358    }
359
360    fn ingest_dml_init_response() -> IngestDmlResponse {
361        IngestDmlResponse {
362            response: Some(ingest_dml_response::Response::Init(
363                IngestDmlInitResponse {},
364            )),
365        }
366    }
367
368    fn ingest_dml_ack_response(dml_batch_id: u64) -> IngestDmlResponse {
369        IngestDmlResponse {
370            response: Some(ingest_dml_response::Response::Ack(IngestDmlAckResponse {
371                dml_batch_id,
372            })),
373        }
374    }
375}