risingwave_batch/rpc/service/
task_service.rs1use 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 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 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 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}