1use std::future::Future;
16use std::sync::Arc;
17use std::time::Duration;
18
19use api::v1::meta::ddl_task_request::Task;
20use api::v1::meta::procedure_service_client::ProcedureServiceClient;
21use api::v1::meta::{
22 DdlTaskRequest, DdlTaskResponse, GcRegionsRequest, GcRegionsResponse, GcTableRequest,
23 GcTableResponse, MigrateRegionRequest, MigrateRegionResponse, ProcedureDetailRequest,
24 ProcedureDetailResponse, ProcedureId, ProcedureStateResponse, QueryProcedureRequest,
25 ReconcileRequest, ReconcileResponse, RequestHeader, ResponseHeader, Role,
26};
27use common_grpc::channel_manager::ChannelManager;
28use common_meta::rpc::ddl::{
29 CREATE_DATABASE_CREATOR_EXTENSION_KEY, CREATE_DATABASE_CREATOR_METADATA_KEY,
30};
31use common_meta::rpc::procedure::{
32 GcRegionsRequest as MetaGcRegionsRequest, GcResponse as MetaGcResponse,
33 GcTableRequest as MetaGcTableRequest,
34};
35use common_telemetry::tracing_context::TracingContext;
36use common_telemetry::{error, info, warn};
37use snafu::{ResultExt, ensure};
38use tokio::sync::RwLock;
39use tonic::transport::Channel;
40use tonic::{Request, Status};
41
42use crate::client::{Id, LeaderProviderRef, util};
43use crate::error;
44use crate::error::Result;
45
46#[derive(Clone, Debug)]
47pub struct Client {
48 inner: Arc<RwLock<Inner>>,
49}
50
51impl Client {
52 pub fn new(
53 id: Id,
54 role: Role,
55 channel_manager: ChannelManager,
56 max_retry: usize,
57 timeout: Duration,
58 ) -> Self {
59 let inner = Arc::new(RwLock::new(Inner {
60 id,
61 role,
62 channel_manager,
63 leader_provider: None,
64 max_retry,
65 timeout,
66 }));
67
68 Self { inner }
69 }
70
71 pub(crate) async fn start_with(&self, leader_provider: LeaderProviderRef) -> Result<()> {
73 let mut inner = self.inner.write().await;
74 inner.start_with(leader_provider)
75 }
76
77 pub async fn submit_ddl_task(&self, req: DdlTaskRequest) -> Result<DdlTaskResponse> {
78 let inner = self.inner.read().await;
79 inner.submit_ddl_task(req).await
80 }
81
82 pub async fn query_procedure_state(&self, pid: &str) -> Result<ProcedureStateResponse> {
84 let inner = self.inner.read().await;
85 inner.query_procedure_state(pid).await
86 }
87
88 pub async fn migrate_region(
94 &self,
95 region_id: u64,
96 from_peer: u64,
97 to_peer: u64,
98 timeout: Duration,
99 ) -> Result<MigrateRegionResponse> {
100 let inner = self.inner.read().await;
101 inner
102 .migrate_region(region_id, from_peer, to_peer, timeout)
103 .await
104 }
105
106 pub async fn reconcile(&self, request: ReconcileRequest) -> Result<ReconcileResponse> {
108 let inner = self.inner.read().await;
109 inner.reconcile(request).await
110 }
111
112 pub async fn list_procedures(&self) -> Result<ProcedureDetailResponse> {
113 let inner = self.inner.read().await;
114 inner.list_procedures().await
115 }
116
117 pub async fn gc_regions(&self, request: MetaGcRegionsRequest) -> Result<MetaGcResponse> {
118 let inner = self.inner.read().await;
119 inner.gc_regions(request).await
120 }
121
122 pub async fn gc_table(&self, request: MetaGcTableRequest) -> Result<MetaGcResponse> {
123 let inner = self.inner.read().await;
124 inner.gc_table(request).await
125 }
126}
127
128#[derive(Debug)]
129struct Inner {
130 id: Id,
131 role: Role,
132 channel_manager: ChannelManager,
133 leader_provider: Option<LeaderProviderRef>,
134 max_retry: usize,
135 timeout: Duration,
137}
138
139impl Inner {
140 fn start_with(&mut self, leader_provider: LeaderProviderRef) -> Result<()> {
141 ensure!(
142 !self.is_started(),
143 error::IllegalGrpcClientStateSnafu {
144 err_msg: "DDL client already started",
145 }
146 );
147 self.leader_provider = Some(leader_provider);
148 Ok(())
149 }
150
151 fn make_client(&self, addr: impl AsRef<str>) -> Result<ProcedureServiceClient<Channel>> {
152 let channel = self
153 .channel_manager
154 .get(addr)
155 .context(error::CreateChannelSnafu)?;
156
157 Ok(common_grpc::configure_tonic_client!(
158 ProcedureServiceClient::new(channel),
159 self.channel_manager,
160 ))
161 }
162
163 #[inline]
164 fn is_started(&self) -> bool {
165 self.leader_provider.is_some()
166 }
167
168 async fn with_retry<T, F, R, H>(&self, task: &str, body_fn: F, get_header: H) -> Result<T>
169 where
170 R: Future<Output = std::result::Result<T, Status>>,
171 F: Fn(ProcedureServiceClient<Channel>) -> R,
172 H: Fn(&T) -> &Option<ResponseHeader>,
173 {
174 let Some(leader_provider) = self.leader_provider.as_ref() else {
175 return error::IllegalGrpcClientStateSnafu {
176 err_msg: "not started",
177 }
178 .fail();
179 };
180
181 let mut times = 0;
182 let mut last_error = None;
183
184 while times < self.max_retry {
185 if let Some(leader) = &leader_provider.leader() {
186 let client = self.make_client(leader)?;
187 match body_fn(client).await {
188 Ok(res) => {
189 if util::is_not_leader(get_header(&res)) {
190 last_error = Some(format!("{leader} is not a leader"));
191 warn!("Failed to {task} to {leader}, not a leader");
192 let leader = leader_provider.ask_leader().await?;
193 info!("DDL client updated to new leader addr: {leader}");
194 times += 1;
195 continue;
196 }
197 return Ok(res);
198 }
199 Err(status) => {
200 if util::is_unreachable(&status) {
202 last_error = Some(status.to_string());
203 warn!("Failed to {task} to {leader}, source: {status}");
204 let leader = leader_provider.ask_leader().await?;
205 info!("Procedure client updated to new leader addr: {leader}");
206 times += 1;
207 continue;
208 } else {
209 error!("An error occurred in gRPC, status: {status:?}");
210 return Err(error::Error::from(status));
211 }
212 }
213 }
214 } else {
215 leader_provider.ask_leader().await?;
216 }
217 }
218
219 error::RetryTimesExceededSnafu {
220 msg: format!("Failed to {task}, last error: {:?}", last_error),
221 times: self.max_retry,
222 }
223 .fail()
224 }
225
226 async fn migrate_region(
227 &self,
228 region_id: u64,
229 from_peer: u64,
230 to_peer: u64,
231 timeout: Duration,
232 ) -> Result<MigrateRegionResponse> {
233 let mut req = MigrateRegionRequest {
234 region_id,
235 from_peer,
236 to_peer,
237 timeout_secs: timeout.as_secs() as u32,
238 ..Default::default()
239 };
240
241 req.set_header(
242 self.id,
243 self.role,
244 TracingContext::from_current_span().to_w3c(),
245 );
246
247 self.with_retry(
248 "migrate region",
249 move |mut client| {
250 let mut req = Request::new(req.clone());
251 req.set_timeout(self.timeout);
252
253 async move { client.migrate(req).await.map(|res| res.into_inner()) }
254 },
255 |resp: &MigrateRegionResponse| &resp.header,
256 )
257 .await
258 }
259
260 async fn reconcile(&self, request: ReconcileRequest) -> Result<ReconcileResponse> {
261 let mut req = request;
262 req.set_header(
263 self.id,
264 self.role,
265 TracingContext::from_current_span().to_w3c(),
266 );
267
268 self.with_retry(
269 "reconcile",
270 move |mut client| {
271 let mut req = Request::new(req.clone());
272 req.set_timeout(self.timeout);
273
274 async move { client.reconcile(req).await.map(|res| res.into_inner()) }
275 },
276 |resp: &ReconcileResponse| &resp.header,
277 )
278 .await
279 }
280
281 async fn gc_regions(&self, request: MetaGcRegionsRequest) -> Result<MetaGcResponse> {
282 let timeout = request.timeout;
283 let req = GcRegionsRequest {
284 header: Some(RequestHeader {
285 protocol_version: 0,
286 member_id: self.id,
287 role: self.role as i32,
288 tracing_context: TracingContext::from_current_span().to_w3c(),
289 }),
290 region_ids: request.region_ids,
291 full_file_listing: request.full_file_listing,
292 timeout_secs: gc_timeout_secs(timeout),
293 };
294
295 let resp: GcRegionsResponse = self
296 .with_retry(
297 "gc_regions",
298 move |mut client| {
299 let mut req = Request::new(req.clone());
300 if let Some(timeout) = timeout {
301 req.set_timeout(timeout);
302 }
303 async move { client.gc_regions(req).await.map(|res| res.into_inner()) }
304 },
305 |resp: &GcRegionsResponse| &resp.header,
306 )
307 .await?;
308
309 let stats = resp.stats.unwrap_or_default();
310 Ok(MetaGcResponse {
311 processed_regions: stats.processed_regions,
312 need_retry_regions: stats.need_retry_regions,
313 deleted_files: stats.deleted_files,
314 deleted_indexes: stats.deleted_indexes,
315 })
316 }
317
318 async fn gc_table(&self, request: MetaGcTableRequest) -> Result<MetaGcResponse> {
319 let timeout = request.timeout;
320 let req = GcTableRequest {
321 header: Some(RequestHeader {
322 protocol_version: 0,
323 member_id: self.id,
324 role: self.role as i32,
325 tracing_context: TracingContext::from_current_span().to_w3c(),
326 }),
327 catalog_name: request.catalog_name,
328 schema_name: request.schema_name,
329 table_name: request.table_name,
330 full_file_listing: request.full_file_listing,
331 timeout_secs: gc_timeout_secs(timeout),
332 };
333
334 let resp: GcTableResponse = self
335 .with_retry(
336 "gc_table",
337 move |mut client| {
338 let mut req = Request::new(req.clone());
339 if let Some(timeout) = timeout {
340 req.set_timeout(timeout);
341 }
342 async move { client.gc_table(req).await.map(|res| res.into_inner()) }
343 },
344 |resp: &GcTableResponse| &resp.header,
345 )
346 .await?;
347
348 let stats = resp.stats.unwrap_or_default();
349 Ok(MetaGcResponse {
350 processed_regions: stats.processed_regions,
351 need_retry_regions: stats.need_retry_regions,
352 deleted_files: stats.deleted_files,
353 deleted_indexes: stats.deleted_indexes,
354 })
355 }
356
357 async fn query_procedure_state(&self, pid: &str) -> Result<ProcedureStateResponse> {
358 let mut req = QueryProcedureRequest {
359 pid: Some(ProcedureId { key: pid.into() }),
360 ..Default::default()
361 };
362
363 req.set_header(
364 self.id,
365 self.role,
366 TracingContext::from_current_span().to_w3c(),
367 );
368
369 self.with_retry(
370 "query procedure state",
371 move |mut client| {
372 let mut req = Request::new(req.clone());
373 req.set_timeout(self.timeout);
374
375 async move { client.query(req).await.map(|res| res.into_inner()) }
376 },
377 |resp: &ProcedureStateResponse| &resp.header,
378 )
379 .await
380 }
381
382 async fn submit_ddl_task(&self, mut req: DdlTaskRequest) -> Result<DdlTaskResponse> {
383 let creator = create_database_creator_metadata_value(&req);
384 req.set_header(
385 self.id,
386 self.role,
387 TracingContext::from_current_span().to_w3c(),
388 );
389 let timeout = Duration::from_secs(req.timeout_secs.into());
390
391 self.with_retry(
392 "submit ddl task",
393 move |mut client| {
394 let mut req = Request::new(req.clone());
395 if let Some(value) = creator.as_deref() {
396 req.metadata_mut().insert_bin(
397 CREATE_DATABASE_CREATOR_METADATA_KEY,
398 tonic::metadata::MetadataValue::from_bytes(value.as_bytes()),
399 );
400 }
401 req.set_timeout(timeout);
402 async move { client.ddl(req).await.map(|res| res.into_inner()) }
403 },
404 |resp: &DdlTaskResponse| &resp.header,
405 )
406 .await
407 }
408
409 async fn list_procedures(&self) -> Result<ProcedureDetailResponse> {
410 let mut req = ProcedureDetailRequest::default();
411 req.set_header(
412 self.id,
413 self.role,
414 TracingContext::from_current_span().to_w3c(),
415 );
416
417 self.with_retry(
418 "list procedure",
419 move |mut client| {
420 let mut req = Request::new(req.clone());
421 req.set_timeout(self.timeout);
422 async move { client.details(req).await.map(|res| res.into_inner()) }
423 },
424 |resp: &ProcedureDetailResponse| &resp.header,
425 )
426 .await
427 }
428}
429
430fn create_database_creator_metadata_value(req: &DdlTaskRequest) -> Option<String> {
431 if !matches!(req.task, Some(Task::CreateDatabaseTask(_))) {
433 return None;
434 }
435
436 req.query_context
437 .as_ref()?
438 .extensions
439 .get(CREATE_DATABASE_CREATOR_EXTENSION_KEY)
440 .cloned()
441}
442
443fn gc_timeout_secs(timeout: Option<Duration>) -> u32 {
444 timeout
445 .map(|timeout| timeout.as_secs().max(1).try_into().unwrap_or(u32::MAX))
446 .unwrap_or(0)
447}
448
449#[cfg(test)]
450mod tests {
451 use std::time::{Duration, Instant};
452
453 use api::v1::meta::heartbeat_server::{Heartbeat, HeartbeatServer};
454 use api::v1::meta::procedure_service_server::{ProcedureService, ProcedureServiceServer};
455 use api::v1::meta::{
456 AskLeaderRequest, AskLeaderResponse, DdlTaskRequest, DdlTaskResponse, GcRegionsRequest,
457 GcRegionsResponse, GcTableRequest, GcTableResponse, HeartbeatRequest, HeartbeatResponse,
458 MigrateRegionRequest, MigrateRegionResponse, Peer, ProcedureDetailRequest,
459 ProcedureDetailResponse, ProcedureStateResponse, QueryProcedureRequest, ReconcileRequest,
460 ReconcileResponse, ResponseHeader, Role,
461 };
462 use async_trait::async_trait;
463 use common_error::status_code::StatusCode;
464 use common_meta::rpc::ddl::{
465 CREATE_DATABASE_CREATOR_EXTENSION_KEY, CREATE_DATABASE_CREATOR_METADATA_KEY,
466 CommentObjectType, CommentOnTask, CreatorGrantIntent, DdlTask, QueryContext,
467 SubmitDdlTaskRequest,
468 };
469 use common_telemetry::common_error::ext::ErrorExt;
470 use common_telemetry::info;
471 use tokio::net::TcpListener;
472 use tokio::sync::mpsc;
473 use tokio_stream::wrappers::{ReceiverStream, TcpListenerStream};
474 use tonic::codec::CompressionEncoding;
475 use tonic::{Request, Response, Status};
476
477 use super::gc_timeout_secs;
478 use crate::client::MetaClientBuilder;
479
480 #[test]
481 fn test_gc_timeout_secs() {
482 assert_eq!(gc_timeout_secs(None), 0);
483 assert_eq!(gc_timeout_secs(Some(Duration::from_millis(1))), 1);
484 assert_eq!(gc_timeout_secs(Some(Duration::from_millis(999))), 1);
485 assert_eq!(gc_timeout_secs(Some(Duration::from_secs(1))), 1);
486 assert_eq!(gc_timeout_secs(Some(Duration::from_secs(10))), 10);
487 }
488
489 #[derive(Clone)]
490 struct MockHeartbeat {
491 leader_addr: String,
492 }
493
494 #[async_trait]
495 impl Heartbeat for MockHeartbeat {
496 type HeartbeatStream = ReceiverStream<Result<HeartbeatResponse, Status>>;
497
498 async fn heartbeat(
499 &self,
500 _request: Request<tonic::Streaming<HeartbeatRequest>>,
501 ) -> Result<Response<Self::HeartbeatStream>, Status> {
502 Err(Status::unimplemented(
503 "heartbeat stream is not used in this test",
504 ))
505 }
506
507 async fn ask_leader(
508 &self,
509 _request: Request<AskLeaderRequest>,
510 ) -> Result<Response<AskLeaderResponse>, Status> {
511 Ok(Response::new(AskLeaderResponse {
512 header: Some(ResponseHeader {
513 protocol_version: 0,
514 error: None,
515 }),
516 leader: Some(Peer {
517 id: 1,
518 addr: self.leader_addr.clone(),
519 }),
520 }))
521 }
522 }
523
524 #[derive(Clone)]
525 struct MockProcedure {
526 delay: Duration,
527 request_tx: Option<mpsc::UnboundedSender<Request<DdlTaskRequest>>>,
528 }
529
530 #[async_trait]
531 impl ProcedureService for MockProcedure {
532 async fn query(
533 &self,
534 _request: Request<QueryProcedureRequest>,
535 ) -> Result<Response<ProcedureStateResponse>, Status> {
536 Err(Status::unimplemented("query is not used in this test"))
537 }
538
539 async fn ddl(
540 &self,
541 request: Request<DdlTaskRequest>,
542 ) -> Result<Response<DdlTaskResponse>, Status> {
543 if let Some(request_tx) = &self.request_tx {
544 request_tx.send(request).unwrap();
545 }
546 tokio::time::sleep(self.delay).await;
547 Ok(Response::new(DdlTaskResponse {
548 header: Some(ResponseHeader {
549 protocol_version: 0,
550 error: None,
551 }),
552 ..Default::default()
553 }))
554 }
555
556 async fn reconcile(
557 &self,
558 _request: Request<ReconcileRequest>,
559 ) -> Result<Response<ReconcileResponse>, Status> {
560 Err(Status::unimplemented("reconcile is not used in this test"))
561 }
562
563 async fn migrate(
564 &self,
565 _request: Request<MigrateRegionRequest>,
566 ) -> Result<Response<MigrateRegionResponse>, Status> {
567 Err(Status::unimplemented("migrate is not used in this test"))
568 }
569
570 async fn details(
571 &self,
572 _request: Request<ProcedureDetailRequest>,
573 ) -> Result<Response<ProcedureDetailResponse>, Status> {
574 Err(Status::unimplemented("details is not used in this test"))
575 }
576
577 async fn gc_regions(
578 &self,
579 _request: Request<GcRegionsRequest>,
580 ) -> Result<Response<GcRegionsResponse>, Status> {
581 Err(Status::unimplemented("gc_regions is not used in this test"))
582 }
583
584 async fn gc_table(
585 &self,
586 _request: Request<GcTableRequest>,
587 ) -> Result<Response<GcTableResponse>, Status> {
588 Err(Status::unimplemented("gc_table is not used in this test"))
589 }
590 }
591
592 #[tokio::test(flavor = "multi_thread")]
593 async fn test_meta_client_forwards_create_database_creator_metadata() {
594 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
595 let addr_str = listener.local_addr().unwrap().to_string();
596 let (request_tx, mut request_rx) = mpsc::unbounded_channel();
597 let heartbeat = MockHeartbeat {
598 leader_addr: addr_str.clone(),
599 };
600 let procedure = MockProcedure {
601 delay: Duration::ZERO,
602 request_tx: Some(request_tx),
603 };
604 let server = tonic::transport::Server::builder()
605 .add_service(
606 HeartbeatServer::new(heartbeat).accept_compressed(CompressionEncoding::Zstd),
607 )
608 .add_service(
609 ProcedureServiceServer::new(procedure).accept_compressed(CompressionEncoding::Zstd),
610 )
611 .serve_with_incoming(TcpListenerStream::new(listener));
612 let server_handle = tokio::spawn(server);
613
614 let mut client = MetaClientBuilder::new(0, Role::Frontend)
615 .enable_heartbeat()
616 .enable_procedure()
617 .build();
618 client.start(&[addr_str.as_str()]).await.unwrap();
619
620 let creator = CreatorGrantIntent {
621 username: "alice".to_string(),
622 created_at_ns: 42,
623 };
624 client
625 .submit_ddl_task(SubmitDdlTaskRequest::new(
626 QueryContext::default(),
627 DdlTask::new_create_database(
628 "greptime".to_string(),
629 "metrics".to_string(),
630 false,
631 Default::default(),
632 Some(creator.clone()),
633 ),
634 ))
635 .await
636 .unwrap();
637
638 let request = request_rx.recv().await.unwrap();
639 let encoded = serde_json::to_string(&creator).unwrap();
640 assert_eq!(
641 request
642 .metadata()
643 .get_bin(CREATE_DATABASE_CREATOR_METADATA_KEY)
644 .unwrap()
645 .to_bytes()
646 .unwrap()
647 .as_ref(),
648 encoded.as_bytes()
649 );
650 let extensions = &request.into_inner().query_context.unwrap().extensions;
651 assert_eq!(extensions[CREATE_DATABASE_CREATOR_EXTENSION_KEY], encoded);
652
653 server_handle.abort();
654 }
655
656 #[tokio::test(flavor = "multi_thread")]
657 async fn test_meta_client_ddl_request_timeout() {
658 common_telemetry::init_default_ut_logging();
659
660 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
661 let addr = listener.local_addr().unwrap();
662 let addr_str = addr.to_string();
663
664 let heartbeat = MockHeartbeat {
665 leader_addr: addr_str.clone(),
666 };
667 let procedure = MockProcedure {
668 delay: Duration::from_secs(4),
669 request_tx: None,
670 };
671
672 let server = tonic::transport::Server::builder()
673 .add_service(
674 HeartbeatServer::new(heartbeat)
675 .accept_compressed(CompressionEncoding::Gzip)
676 .accept_compressed(CompressionEncoding::Zstd),
677 )
678 .add_service(
679 ProcedureServiceServer::new(procedure)
680 .accept_compressed(CompressionEncoding::Gzip)
681 .accept_compressed(CompressionEncoding::Zstd),
682 )
683 .serve_with_incoming(TcpListenerStream::new(listener));
684 let server_handle = tokio::spawn(server);
685
686 let mut client = MetaClientBuilder::new(0, Role::Frontend)
687 .enable_heartbeat()
688 .enable_procedure()
689 .build();
690 client.start(&[addr_str.as_str()]).await.unwrap();
691
692 let mut request = SubmitDdlTaskRequest::new(
693 QueryContext::default(),
694 DdlTask::new_comment_on(CommentOnTask {
695 catalog_name: "greptime".to_string(),
696 schema_name: "public".to_string(),
697 object_type: CommentObjectType::Table,
698 object_name: "test_table".to_string(),
699 column_name: None,
700 object_id: None,
701 comment: Some("timeout".to_string()),
702 }),
703 );
704 request.timeout = Duration::from_secs(1);
705
706 let now = Instant::now();
707 let err = client.submit_ddl_task(request).await.unwrap_err();
708 let elapsed = now.elapsed();
709 assert!(elapsed < Duration::from_secs(2));
711 info!("err: {err:?}, code: {}", err.status_code());
712 assert_eq!(err.status_code(), StatusCode::Cancelled);
713 let err_msg = err.to_string();
714 assert!(
715 err_msg.contains("Timeout expired"),
716 "unexpected error: {err_msg}"
717 );
718
719 server_handle.abort();
720 }
721}