1use std::collections::HashMap;
16use std::collections::hash_map::Entry;
17use std::fmt::{Debug, Display, Formatter};
18use std::sync::atomic::{AtomicU32, Ordering};
19use std::sync::{Arc, RwLock};
20use std::time::{Duration, Instant, UNIX_EPOCH};
21
22use api::v1::frontend::{KillProcessRequest, ListProcessRequest, ProcessInfo};
23use common_base::cancellation::CancellationHandle;
24use common_event_recorder::EventRecorderRef;
25use common_frontend::selector::{FrontendSelector, MetaClientSelector};
26use common_frontend::slow_query_event::SlowQueryEvent;
27use common_telemetry::logging::SlowQueriesRecordType;
28use common_telemetry::{debug, info, slow, warn};
29use common_time::util::current_time_millis;
30use meta_client::MetaClientRef;
31use promql_parser::parser::EvalStmt;
32use rand::random;
33use snafu::{OptionExt, ResultExt, ensure};
34use sql::statements::statement::Statement;
35
36use crate::error;
37use crate::metrics::{PROCESS_KILL_COUNT, PROCESS_LIST_COUNT};
38
39pub type ProcessId = u32;
40pub type ProcessManagerRef = Arc<ProcessManager>;
41
42pub struct ProcessManager {
44 server_addr: String,
46 next_id: AtomicU32,
48 catalogs: RwLock<HashMap<String, HashMap<ProcessId, CancellableProcess>>>,
50 frontend_selector: Option<MetaClientSelector>,
52}
53
54#[derive(Debug, Clone)]
57pub enum QueryStatement {
58 Sql(Statement),
59 Promql(EvalStmt, Option<String>),
61 Plan(String),
63}
64
65impl Display for QueryStatement {
66 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
67 match self {
68 QueryStatement::Sql(stmt) => write!(f, "{}", stmt),
69 QueryStatement::Promql(eval_stmt, alias) => {
70 if let Some(alias) = alias {
71 write!(f, "{} AS {}", eval_stmt, alias)
72 } else {
73 write!(f, "{}", eval_stmt)
74 }
75 }
76 QueryStatement::Plan(query) => write!(f, "{}", query),
77 }
78 }
79}
80
81impl ProcessManager {
82 pub fn new(server_addr: String, meta_client: Option<MetaClientRef>) -> Self {
84 let frontend_selector = meta_client.map(MetaClientSelector::new);
85 Self {
86 server_addr,
87 next_id: Default::default(),
88 catalogs: Default::default(),
89 frontend_selector,
90 }
91 }
92}
93
94impl ProcessManager {
95 #[must_use]
97 pub fn register_query(
98 self: &Arc<Self>,
99 catalog: String,
100 schemas: Vec<String>,
101 query: String,
102 client: String,
103 query_id: Option<ProcessId>,
104 _slow_query_timer: Option<SlowQueryTimer>,
105 ) -> Ticket {
106 let id = query_id.unwrap_or_else(|| self.next_id.fetch_add(1, Ordering::Relaxed));
107 let process = ProcessInfo {
108 id,
109 catalog: catalog.clone(),
110 schemas,
111 query,
112 start_timestamp: current_time_millis(),
113 client,
114 frontend: self.server_addr.clone(),
115 };
116 let cancellation_handle = Arc::new(CancellationHandle::default());
117 let cancellable_process = CancellableProcess::new(cancellation_handle.clone(), process);
118
119 self.catalogs
120 .write()
121 .unwrap()
122 .entry(catalog.clone())
123 .or_default()
124 .insert(id, cancellable_process);
125
126 Ticket {
127 catalog,
128 manager: self.clone(),
129 id,
130 cancellation_handle,
131 _slow_query_timer,
132 }
133 }
134
135 pub fn next_id(&self) -> u32 {
137 self.next_id.fetch_add(1, Ordering::Relaxed)
138 }
139
140 pub fn deregister_query(&self, catalog: String, id: ProcessId) {
142 if let Entry::Occupied(mut o) = self.catalogs.write().unwrap().entry(catalog) {
143 let process = o.get_mut().remove(&id);
144 debug!("Deregister process: {:?}", process);
145 if o.get().is_empty() {
146 o.remove();
147 }
148 }
149 }
150
151 pub fn local_processes(&self, catalog: Option<&str>) -> error::Result<Vec<ProcessInfo>> {
153 let catalogs = self.catalogs.read().unwrap();
154 let result = if let Some(catalog) = catalog {
155 if let Some(catalogs) = catalogs.get(catalog) {
156 catalogs.values().map(|p| p.process.clone()).collect()
157 } else {
158 vec![]
159 }
160 } else {
161 catalogs
162 .values()
163 .flat_map(|v| v.values().map(|p| p.process.clone()))
164 .collect()
165 };
166 Ok(result)
167 }
168
169 pub async fn list_all_processes(
170 &self,
171 catalog: Option<&str>,
172 ) -> error::Result<Vec<ProcessInfo>> {
173 let mut processes = vec![];
174 if let Some(remote_frontend_selector) = self.frontend_selector.as_ref() {
175 let frontends = remote_frontend_selector
176 .select(|peer| peer.addr != self.server_addr)
177 .await
178 .context(error::InvokeFrontendSnafu)?;
179 for mut f in frontends {
180 let result = f
181 .list_process(ListProcessRequest {
182 catalog: catalog.unwrap_or_default().to_string(),
183 })
184 .await
185 .context(error::InvokeFrontendSnafu);
186 match result {
187 Ok(resp) => {
188 processes.extend(resp.processes);
189 }
190 Err(e) => {
191 warn!(e; "Skipping failing node: {:?}", f)
192 }
193 }
194 }
195 }
196 processes.extend(self.local_processes(catalog)?);
197 Ok(processes)
198 }
199
200 pub async fn kill_process(
202 &self,
203 server_addr: String,
204 catalog: String,
205 id: ProcessId,
206 ) -> error::Result<bool> {
207 if server_addr == self.server_addr {
208 self.kill_local_process(catalog, id).await
209 } else {
210 let mut nodes = self
211 .frontend_selector
212 .as_ref()
213 .context(error::MetaClientMissingSnafu)?
214 .select(|peer| peer.addr == server_addr)
215 .await
216 .context(error::InvokeFrontendSnafu)?;
217 ensure!(
218 !nodes.is_empty(),
219 error::FrontendNotFoundSnafu { addr: server_addr }
220 );
221
222 let request = KillProcessRequest {
223 server_addr,
224 catalog,
225 process_id: id,
226 };
227 nodes[0]
228 .kill_process(request)
229 .await
230 .context(error::InvokeFrontendSnafu)?;
231 Ok(true)
232 }
233 }
234
235 pub async fn kill_local_process(&self, catalog: String, id: ProcessId) -> error::Result<bool> {
237 if let Some(catalogs) = self.catalogs.write().unwrap().get_mut(&catalog) {
238 if let Some(process) = catalogs.remove(&id) {
239 process.handle.cancel();
240 info!(
241 "Killed process, catalog: {}, id: {:?}",
242 process.process.catalog, process.process.id
243 );
244 PROCESS_KILL_COUNT.with_label_values(&[&catalog]).inc();
245 Ok(true)
246 } else {
247 debug!("Failed to kill process, id not found: {}", id);
248 Ok(false)
249 }
250 } else {
251 debug!("Failed to kill process, catalog not found: {}", catalog);
252 Ok(false)
253 }
254 }
255}
256
257pub struct Ticket {
258 pub(crate) catalog: String,
259 pub(crate) manager: ProcessManagerRef,
260 pub(crate) id: ProcessId,
261 pub cancellation_handle: Arc<CancellationHandle>,
262
263 _slow_query_timer: Option<SlowQueryTimer>,
265}
266
267impl Drop for Ticket {
268 fn drop(&mut self) {
269 self.manager
270 .deregister_query(std::mem::take(&mut self.catalog), self.id);
271 }
272}
273
274struct CancellableProcess {
275 handle: Arc<CancellationHandle>,
276 process: ProcessInfo,
277}
278
279impl Drop for CancellableProcess {
280 fn drop(&mut self) {
281 PROCESS_LIST_COUNT
282 .with_label_values(&[&self.process.catalog])
283 .dec();
284 }
285}
286
287impl CancellableProcess {
288 fn new(handle: Arc<CancellationHandle>, process: ProcessInfo) -> Self {
289 PROCESS_LIST_COUNT
290 .with_label_values(&[&process.catalog])
291 .inc();
292 Self { handle, process }
293 }
294}
295
296impl Debug for CancellableProcess {
297 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
298 f.debug_struct("CancellableProcess")
299 .field("cancelled", &self.handle.is_cancelled())
300 .field("process", &self.process)
301 .finish()
302 }
303}
304
305pub struct SlowQueryTimer {
308 start: Instant,
309 stmt: QueryStatement,
310 schema_name: String,
311 threshold: Duration,
312 sample_ratio: f64,
313 record_type: SlowQueriesRecordType,
314 recorder: EventRecorderRef,
315 record_state: Arc<RwLock<SlowQueryRecordState>>,
316}
317
318#[derive(Default)]
319struct SlowQueryRecordState {
320 force_record: bool,
321 payload: serde_json::Value,
322}
323
324#[derive(Clone)]
326pub struct SlowQueryRecorder {
327 record_state: Arc<RwLock<SlowQueryRecordState>>,
328}
329
330impl SlowQueryRecorder {
331 pub fn force_record_with_payload(&self, payload: serde_json::Value) {
333 let mut state = self.record_state.write().unwrap();
334 state.payload = payload;
335 state.force_record = true;
336 }
337}
338
339impl SlowQueryTimer {
340 pub fn new(
341 stmt: QueryStatement,
342 schema_name: String,
343 threshold: Duration,
344 sample_ratio: f64,
345 record_type: SlowQueriesRecordType,
346 recorder: EventRecorderRef,
347 ) -> Self {
348 Self {
349 start: Instant::now(),
350 stmt,
351 schema_name,
352 threshold,
353 sample_ratio,
354 record_type,
355 recorder,
356 record_state: Arc::default(),
357 }
358 }
359
360 pub fn recorder(&self) -> SlowQueryRecorder {
362 SlowQueryRecorder {
363 record_state: self.record_state.clone(),
364 }
365 }
366}
367
368impl SlowQueryTimer {
369 fn send_slow_query_event(&self, elapsed: Duration, payload: serde_json::Value) {
370 let mut slow_query_event = SlowQueryEvent {
371 cost: elapsed.as_millis() as u64,
372 threshold: self.threshold.as_millis() as u64,
373 query: "".to_string(),
374 schema_name: self.schema_name.clone(),
375
376 is_promql: false,
378 promql_range: None,
379 promql_step: None,
380 promql_start: None,
381 promql_end: None,
382 payload,
383 };
384
385 match &self.stmt {
386 QueryStatement::Promql(stmt, _alias) => {
387 slow_query_event.is_promql = true;
388 slow_query_event.query = self.stmt.to_string();
389 slow_query_event.promql_step = Some(stmt.interval.as_millis() as u64);
390
391 let start = stmt
392 .start
393 .duration_since(UNIX_EPOCH)
394 .unwrap_or_default()
395 .as_millis() as i64;
396
397 let end = stmt
398 .end
399 .duration_since(UNIX_EPOCH)
400 .unwrap_or_default()
401 .as_millis() as i64;
402
403 slow_query_event.promql_range = Some((end - start) as u64);
404 slow_query_event.promql_start = Some(start);
405 slow_query_event.promql_end = Some(end);
406 }
407 QueryStatement::Sql(stmt) => {
408 slow_query_event.query = stmt.to_string();
409 }
410 QueryStatement::Plan(query) => {
411 slow_query_event.query = query.clone();
412 }
413 }
414
415 match self.record_type {
416 SlowQueriesRecordType::SystemTable => {
418 self.recorder.record(Box::new(slow_query_event));
419 }
420 SlowQueriesRecordType::Log => {
422 slow!(
423 cost = slow_query_event.cost,
424 threshold = slow_query_event.threshold,
425 query = slow_query_event.query,
426 schema_name = slow_query_event.schema_name,
427 is_promql = slow_query_event.is_promql,
428 promql_range = slow_query_event.promql_range,
429 promql_step = slow_query_event.promql_step,
430 promql_start = slow_query_event.promql_start,
431 promql_end = slow_query_event.promql_end,
432 payload = slow_query_event.payload.to_string(),
433 );
434 }
435 }
436 }
437}
438
439impl Drop for SlowQueryTimer {
440 fn drop(&mut self) {
441 let elapsed = self.start.elapsed();
442 let (force_record, payload) = {
443 let state = self.record_state.read().unwrap();
444 (state.force_record, state.payload.clone())
445 };
446
447 if force_record
448 || (elapsed > self.threshold
449 && (self.sample_ratio >= 1.0 || random::<f64>() <= self.sample_ratio))
450 {
451 self.send_slow_query_event(elapsed, payload);
452 }
453 }
454}
455
456#[cfg(test)]
457mod tests {
458 use std::sync::{Arc, Mutex};
459 use std::time::Duration;
460
461 use common_event_recorder::{Event, EventRecorder, EventTypeFilter, EventTypeFilterRef};
462 use common_frontend::slow_query_event::SlowQueryEvent;
463 use common_telemetry::logging::SlowQueriesRecordType;
464 use serde_json::Value;
465
466 use crate::process_manager::{ProcessManager, QueryStatement, SlowQueryTimer};
467
468 #[derive(Debug, Default)]
469 struct RecordingEventRecorder {
470 events: Mutex<Vec<(String, Value)>>,
471 }
472
473 impl EventRecorder for RecordingEventRecorder {
474 fn record(&self, event: Box<dyn Event>) {
475 let event = event
476 .as_any()
477 .downcast_ref::<SlowQueryEvent>()
478 .expect("expected a slow query event");
479 self.events
480 .lock()
481 .unwrap()
482 .push((event.query.clone(), event.payload.clone()));
483 }
484
485 fn event_type_filter(&self) -> EventTypeFilterRef {
486 Arc::new(EventTypeFilter::All)
487 }
488
489 fn close(&self) {}
490 }
491
492 #[test]
493 fn test_forced_slow_query_bypasses_threshold_and_sampling() {
494 let event_recorder = Arc::new(RecordingEventRecorder::default());
495 let timer = SlowQueryTimer::new(
496 QueryStatement::Plan("EXPLAIN ANALYZE VERBOSE SELECT 1".to_string()),
497 "public".to_string(),
498 Duration::from_secs(3600),
499 0.0,
500 SlowQueriesRecordType::SystemTable,
501 event_recorder.clone(),
502 );
503 let payload = serde_json::json!({
504 "timed_out": true,
505 "metrics": [{"stage": 0}],
506 });
507 timer.recorder().force_record_with_payload(payload.clone());
508
509 drop(timer);
510
511 let events = event_recorder.events.lock().unwrap();
512 assert_eq!(events.len(), 1);
513 assert_eq!(events[0].0, "EXPLAIN ANALYZE VERBOSE SELECT 1");
514 assert_eq!(events[0].1, payload);
515 }
516
517 #[test]
518 fn test_unforced_fast_query_is_not_recorded() {
519 let event_recorder = Arc::new(RecordingEventRecorder::default());
520 let timer = SlowQueryTimer::new(
521 QueryStatement::Plan("SELECT 1".to_string()),
522 "public".to_string(),
523 Duration::from_secs(3600),
524 0.0,
525 SlowQueriesRecordType::SystemTable,
526 event_recorder.clone(),
527 );
528
529 drop(timer);
530
531 assert!(event_recorder.events.lock().unwrap().is_empty());
532 }
533
534 #[tokio::test]
535 async fn test_register_query() {
536 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
537 let ticket = process_manager.clone().register_query(
538 "public".to_string(),
539 vec!["test".to_string()],
540 "SELECT * FROM table".to_string(),
541 "".to_string(),
542 None,
543 None,
544 );
545
546 let running_processes = process_manager.local_processes(None).unwrap();
547 assert_eq!(running_processes.len(), 1);
548 assert_eq!(&running_processes[0].frontend, "127.0.0.1:8000");
549 assert_eq!(running_processes[0].id, ticket.id);
550 assert_eq!(&running_processes[0].query, "SELECT * FROM table");
551
552 drop(ticket);
553 assert_eq!(process_manager.local_processes(None).unwrap().len(), 0);
554 }
555
556 #[tokio::test]
557 async fn test_register_query_with_custom_id() {
558 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
559 let custom_id = 12345;
560
561 let ticket = process_manager.clone().register_query(
562 "public".to_string(),
563 vec!["test".to_string()],
564 "SELECT * FROM table".to_string(),
565 "client1".to_string(),
566 Some(custom_id),
567 None,
568 );
569
570 assert_eq!(ticket.id, custom_id);
571
572 let running_processes = process_manager.local_processes(None).unwrap();
573 assert_eq!(running_processes.len(), 1);
574 assert_eq!(running_processes[0].id, custom_id);
575 assert_eq!(&running_processes[0].client, "client1");
576 }
577
578 #[tokio::test]
579 async fn test_multiple_queries_same_catalog() {
580 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
581
582 let ticket1 = process_manager.clone().register_query(
583 "public".to_string(),
584 vec!["schema1".to_string()],
585 "SELECT * FROM table1".to_string(),
586 "client1".to_string(),
587 None,
588 None,
589 );
590
591 let ticket2 = process_manager.clone().register_query(
592 "public".to_string(),
593 vec!["schema2".to_string()],
594 "SELECT * FROM table2".to_string(),
595 "client2".to_string(),
596 None,
597 None,
598 );
599
600 let running_processes = process_manager.local_processes(Some("public")).unwrap();
601 assert_eq!(running_processes.len(), 2);
602
603 let ids: Vec<u32> = running_processes.iter().map(|p| p.id).collect();
605 assert!(ids.contains(&ticket1.id));
606 assert!(ids.contains(&ticket2.id));
607 }
608
609 #[tokio::test]
610 async fn test_multiple_catalogs() {
611 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
612
613 let _ticket1 = process_manager.clone().register_query(
614 "catalog1".to_string(),
615 vec!["schema1".to_string()],
616 "SELECT * FROM table1".to_string(),
617 "client1".to_string(),
618 None,
619 None,
620 );
621
622 let _ticket2 = process_manager.clone().register_query(
623 "catalog2".to_string(),
624 vec!["schema2".to_string()],
625 "SELECT * FROM table2".to_string(),
626 "client2".to_string(),
627 None,
628 None,
629 );
630
631 let catalog1_processes = process_manager.local_processes(Some("catalog1")).unwrap();
633 assert_eq!(catalog1_processes.len(), 1);
634 assert_eq!(&catalog1_processes[0].catalog, "catalog1");
635
636 let catalog2_processes = process_manager.local_processes(Some("catalog2")).unwrap();
637 assert_eq!(catalog2_processes.len(), 1);
638 assert_eq!(&catalog2_processes[0].catalog, "catalog2");
639
640 let all_processes = process_manager.local_processes(None).unwrap();
642 assert_eq!(all_processes.len(), 2);
643 }
644
645 #[tokio::test]
646 async fn test_deregister_query() {
647 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
648
649 let ticket = process_manager.clone().register_query(
650 "public".to_string(),
651 vec!["test".to_string()],
652 "SELECT * FROM table".to_string(),
653 "client1".to_string(),
654 None,
655 None,
656 );
657 assert_eq!(process_manager.local_processes(None).unwrap().len(), 1);
658 process_manager.deregister_query("public".to_string(), ticket.id);
659 assert_eq!(process_manager.local_processes(None).unwrap().len(), 0);
660 }
661
662 #[tokio::test]
663 async fn test_cancellation_handle() {
664 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
665
666 let ticket = process_manager.clone().register_query(
667 "public".to_string(),
668 vec!["test".to_string()],
669 "SELECT * FROM table".to_string(),
670 "client1".to_string(),
671 None,
672 None,
673 );
674
675 assert!(!ticket.cancellation_handle.is_cancelled());
676 ticket.cancellation_handle.cancel();
677 assert!(ticket.cancellation_handle.is_cancelled());
678 }
679
680 #[tokio::test]
681 async fn test_kill_local_process() {
682 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
683
684 let ticket = process_manager.clone().register_query(
685 "public".to_string(),
686 vec!["test".to_string()],
687 "SELECT * FROM table".to_string(),
688 "client1".to_string(),
689 None,
690 None,
691 );
692 assert!(!ticket.cancellation_handle.is_cancelled());
693 let killed = process_manager
694 .kill_process(
695 "127.0.0.1:8000".to_string(),
696 "public".to_string(),
697 ticket.id,
698 )
699 .await
700 .unwrap();
701
702 assert!(killed);
703 assert_eq!(process_manager.local_processes(None).unwrap().len(), 0);
704 }
705
706 #[tokio::test]
707 async fn test_kill_nonexistent_process() {
708 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
709 let killed = process_manager
710 .kill_process("127.0.0.1:8000".to_string(), "public".to_string(), 999)
711 .await
712 .unwrap();
713 assert!(!killed);
714 }
715
716 #[tokio::test]
717 async fn test_kill_process_nonexistent_catalog() {
718 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
719 let killed = process_manager
720 .kill_process("127.0.0.1:8000".to_string(), "nonexistent".to_string(), 1)
721 .await
722 .unwrap();
723 assert!(!killed);
724 }
725
726 #[tokio::test]
727 async fn test_process_info_fields() {
728 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
729
730 let _ticket = process_manager.clone().register_query(
731 "test_catalog".to_string(),
732 vec!["schema1".to_string(), "schema2".to_string()],
733 "SELECT COUNT(*) FROM users WHERE age > 18".to_string(),
734 "test_client".to_string(),
735 Some(42),
736 None,
737 );
738
739 let processes = process_manager.local_processes(None).unwrap();
740 assert_eq!(processes.len(), 1);
741
742 let process = &processes[0];
743 assert_eq!(process.id, 42);
744 assert_eq!(&process.catalog, "test_catalog");
745 assert_eq!(process.schemas, vec!["schema1", "schema2"]);
746 assert_eq!(&process.query, "SELECT COUNT(*) FROM users WHERE age > 18");
747 assert_eq!(&process.client, "test_client");
748 assert_eq!(&process.frontend, "127.0.0.1:8000");
749 assert!(process.start_timestamp > 0);
750 }
751
752 #[tokio::test]
753 async fn test_ticket_drop_deregisters_process() {
754 let process_manager = Arc::new(ProcessManager::new("127.0.0.1:8000".to_string(), None));
755
756 {
757 let _ticket = process_manager.clone().register_query(
758 "public".to_string(),
759 vec!["test".to_string()],
760 "SELECT * FROM table".to_string(),
761 "client1".to_string(),
762 None,
763 None,
764 );
765
766 assert_eq!(process_manager.local_processes(None).unwrap().len(), 1);
768 } assert_eq!(process_manager.local_processes(None).unwrap().len(), 0);
772 }
773}