Skip to main content

datafusion_distributed/protocol/grpc/observability/
service.rs

1use super::{
2    GetTaskProgressResponse, ObservabilityService, TaskProgress, TaskStatus, WorkerMetrics,
3    generated::observability::{GetTaskProgressRequest, PingRequest, PingResponse},
4};
5use crate::common::serialize_uuid;
6use crate::grpc::{GetClusterWorkersRequest, GetClusterWorkersResponse};
7use crate::protocol::grpc::generated::worker as worker_pb;
8use crate::worker::{SingleWriteMultiRead, TaskData};
9use crate::{TaskKey, WorkerResolver};
10use datafusion::error::DataFusionError;
11use datafusion::physical_plan::ExecutionPlan;
12use moka::future::Cache;
13use std::sync::Arc;
14#[cfg(feature = "system-metrics")]
15use std::time::Duration;
16#[cfg(feature = "system-metrics")]
17use sysinfo::{Pid, ProcessRefreshKind};
18#[cfg(feature = "system-metrics")]
19use tokio::sync::watch;
20use tonic::{Request, Response, Status};
21
22type ResultTaskData = Result<TaskData, Arc<DataFusionError>>;
23
24pub struct ObservabilityServiceImpl {
25    task_data_entries: Arc<Cache<TaskKey, Arc<SingleWriteMultiRead<ResultTaskData>>>>,
26    worker_resolver: Arc<dyn WorkerResolver + Send + Sync>,
27    #[cfg(feature = "system-metrics")]
28    system: watch::Receiver<WorkerMetrics>,
29}
30
31impl ObservabilityServiceImpl {
32    pub fn new(
33        task_data_entries: Arc<Cache<TaskKey, Arc<SingleWriteMultiRead<ResultTaskData>>>>,
34        worker_resolver: Arc<dyn WorkerResolver + Send + Sync>,
35    ) -> Self {
36        #[cfg(feature = "system-metrics")]
37        let (tx, rx) = tokio::sync::watch::channel(WorkerMetrics::default());
38
39        #[cfg(feature = "system-metrics")]
40        {
41            let pid = Pid::from_u32(std::process::id());
42            let mut sys = sysinfo::System::new_all();
43
44            // Spawn background task to periodically collect and send system metrics.
45            #[allow(clippy::disallowed_methods)]
46            tokio::task::spawn(async move {
47                loop {
48                    sys.refresh_process_specifics(
49                        pid,
50                        ProcessRefreshKind::new().with_cpu().with_memory(),
51                    );
52
53                    if let Some(process) = sys.process(pid) {
54                        let num_cpus = std::thread::available_parallelism()
55                            .map(|n| n.get() as f64)
56                            .unwrap_or(1.0);
57                        let metrics = WorkerMetrics {
58                            rss_bytes: process.memory(),
59                            cpu_usage_percent: process.cpu_usage() as f64 / num_cpus,
60                        };
61                        if tx.send(metrics).is_err() {
62                            break;
63                        }
64                    } else if tx.send(WorkerMetrics::default()).is_err() {
65                        break;
66                    };
67
68                    tokio::time::sleep(Duration::from_millis(100)).await;
69                }
70            });
71        }
72        Self {
73            task_data_entries,
74            worker_resolver,
75            #[cfg(feature = "system-metrics")]
76            system: rx,
77        }
78    }
79}
80
81#[tonic::async_trait]
82impl ObservabilityService for ObservabilityServiceImpl {
83    async fn ping(&self, _request: Request<PingRequest>) -> Result<Response<PingResponse>, Status> {
84        Ok(Response::new(PingResponse { value: 1 }))
85    }
86
87    async fn get_task_progress(
88        &self,
89        _request: Request<GetTaskProgressRequest>,
90    ) -> Result<Response<GetTaskProgressResponse>, Status> {
91        let mut tasks = Vec::new();
92
93        for entry in self.task_data_entries.iter() {
94            let (internal_key, task_data_cell) = entry;
95
96            // Only include initialized tasks
97            if let Some(Ok(task_data)) = task_data_cell.read_now() {
98                let output_rows = output_rows_from_plan(&task_data.base_plan);
99
100                tasks.push(TaskProgress {
101                    task_key: Some(task_key_to_proto(&internal_key)),
102                    status: TaskStatus::Running as i32,
103                    output_rows,
104                });
105            }
106        }
107
108        let worker_metrics = Some(self.collect_worker_metrics());
109
110        Ok(Response::new(GetTaskProgressResponse {
111            tasks,
112            worker_metrics,
113        }))
114    }
115
116    async fn get_cluster_workers(
117        &self,
118        _request: Request<GetClusterWorkersRequest>,
119    ) -> Result<Response<GetClusterWorkersResponse>, Status> {
120        let urls = self
121            .worker_resolver
122            .get_urls()
123            .map_err(|e| Status::internal(format!("Failed to resolve workers: {e}")))?;
124
125        let worker_urls = urls.into_iter().map(|url| url.to_string()).collect();
126
127        Ok(Response::new(GetClusterWorkersResponse { worker_urls }))
128    }
129}
130
131impl ObservabilityServiceImpl {
132    fn collect_worker_metrics(&self) -> WorkerMetrics {
133        #[cfg(not(feature = "system-metrics"))]
134        {
135            WorkerMetrics::default()
136        }
137
138        #[cfg(feature = "system-metrics")]
139        *self.system.borrow()
140    }
141}
142
143/// Extracts output rows from the root plan node's metrics.
144fn output_rows_from_plan(plan: &Arc<dyn ExecutionPlan>) -> u64 {
145    plan.metrics().and_then(|m| m.output_rows()).unwrap_or(0) as u64
146}
147
148fn task_key_to_proto(task_key: &TaskKey) -> worker_pb::TaskKey {
149    worker_pb::TaskKey {
150        query_id: serialize_uuid(&task_key.query_id),
151        stage_id: task_key.stage_id as u64,
152        task_number: task_key.task_number as u64,
153    }
154}