Skip to main content

solti_api/
grpc.rs

1//! # gRPC transport.
2//!
3//! [`TaskApiService`] implements the generated `TaskService` trait from `proto/solti/task/v1/api.proto`, delegating to an [`ApiHandler`](crate::ApiHandler).
4
5use std::pin::Pin;
6use std::sync::Arc;
7use std::time::Instant;
8
9use tokio_stream::StreamExt;
10use tonic::service::Interceptor;
11use tonic::service::interceptor::InterceptedService;
12use tonic::{Request, Response, Status};
13use tracing::debug;
14
15use solti_model::{TaskQuery, Token};
16
17use crate::convert::{output_event_to_proto, proto_to_domain_status, tasks_page_to_proto};
18use crate::error::ApiError;
19use crate::handler::ApiHandler;
20use crate::metrics::{ApiMetricsHandle, Transport, noop_api_metrics};
21use crate::proto_api::{
22    self, task_service_server::TaskService, task_service_server::TaskServiceServer,
23};
24use crate::validate::{clamp_list_limit, non_empty_id};
25
26/// gRPC service wrapping an [`ApiHandler`](crate::ApiHandler).
27///
28/// ## Also
29///
30/// - `TaskServiceServer` generated tonic server wrapper.
31/// - [`ApiError`](crate::ApiError) mapped to `tonic::Status`.
32pub struct TaskApiService<H> {
33    handler: Arc<H>,
34    metrics: ApiMetricsHandle,
35}
36
37impl<H> TaskApiService<H>
38where
39    H: ApiHandler,
40{
41    /// Create a new gRPC service with the given handler and no-op metrics.
42    pub fn new(handler: Arc<H>) -> Self {
43        Self::new_with_metrics(handler, noop_api_metrics())
44    }
45
46    /// Create a new gRPC service with an explicit metrics backend.
47    pub fn new_with_metrics(handler: Arc<H>, metrics: ApiMetricsHandle) -> Self {
48        Self { handler, metrics }
49    }
50
51    async fn instrument<F, T>(&self, method: &'static str, fut: F) -> Result<Response<T>, Status>
52    where
53        F: Future<Output = Result<Response<T>, Status>>,
54    {
55        self.metrics.record_in_flight_delta(Transport::Grpc, 1);
56        let start = Instant::now();
57        let result = fut.await;
58        let duration_ms = start.elapsed().as_millis() as u64;
59        let status = match &result {
60            Ok(_) => 0u16,
61            Err(s) => s.code() as u16,
62        };
63        let path = format!("/solti.task.v1.TaskService/{}", method);
64        self.metrics
65            .record_request(Transport::Grpc, method, &path, status, duration_ms);
66        self.metrics.record_in_flight_delta(Transport::Grpc, -1);
67        result
68    }
69}
70
71/// Build a configured `TaskServiceServer` with no-op metrics.
72///
73/// ## Example
74///
75/// ```rust,no_run
76/// # use std::sync::Arc;
77/// # use solti_api::{build_grpc_server, SupervisorApiAdapter};
78/// # async fn example(adapter: Arc<SupervisorApiAdapter>) -> Result<(), Box<dyn std::error::Error>> {
79/// let svc = build_grpc_server(adapter);
80/// tonic::transport::Server::builder()
81///     .add_service(svc)
82///     .serve("0.0.0.0:50052".parse()?)
83///     .await?;
84/// # Ok(()) }
85/// ```
86pub fn build_grpc_server<H>(handler: Arc<H>) -> TaskServiceServer<TaskApiService<H>>
87where
88    H: ApiHandler,
89{
90    build_grpc_server_with_metrics(handler, noop_api_metrics())
91}
92
93/// Build a configured `TaskServiceServer` with an explicit metrics backend.
94pub fn build_grpc_server_with_metrics<H>(
95    handler: Arc<H>,
96    metrics: ApiMetricsHandle,
97) -> TaskServiceServer<TaskApiService<H>>
98where
99    H: ApiHandler,
100{
101    TaskServiceServer::new(TaskApiService::new_with_metrics(handler, metrics))
102        .max_decoding_message_size(crate::MAX_REQUEST_BYTES)
103        .max_encoding_message_size(crate::MAX_REQUEST_BYTES)
104}
105
106/// gRPC interceptor enforcing a bearer token on every call.
107///
108/// Verifies `authorization: Bearer <token>` metadata in constant time and rejects with `Unauthenticated` otherwise.
109///
110/// This is the same shared secret the agent presents to the control plane in discovery, one config value enables both directions.
111/// Orthogonal to TLS. Install via [`build_grpc_server_with_auth`] / [`build_grpc_server_with_metrics_auth`].
112#[derive(Clone)]
113pub struct BearerAuth {
114    expected: Token,
115}
116
117impl Interceptor for BearerAuth {
118    fn call(&mut self, request: Request<()>) -> Result<Request<()>, Status> {
119        let ok = request
120            .metadata()
121            .get("authorization")
122            .and_then(|v| v.to_str().ok())
123            .and_then(bearer_value)
124            .map(|presented| self.expected.verify(presented))
125            .unwrap_or(false);
126
127        if ok {
128            Ok(request)
129        } else {
130            Err(Status::unauthenticated("missing or invalid bearer token"))
131        }
132    }
133}
134
135/// Extract the credential from an `authorization` metadata value, accepting the scheme case-insensitively.
136///
137/// The credential is returned verbatim after the first space; it is never trimmed, so it is matched byte-for-byte by [`Token::verify`].
138fn bearer_value(header: &str) -> Option<&str> {
139    let (scheme, token) = header.split_once(' ')?;
140    scheme.eq_ignore_ascii_case("bearer").then_some(token)
141}
142
143/// Like [`build_grpc_server`] but enforcing a bearer token on every call.
144pub fn build_grpc_server_with_auth<H>(
145    handler: Arc<H>,
146    token: Token,
147) -> InterceptedService<TaskServiceServer<TaskApiService<H>>, BearerAuth>
148where
149    H: ApiHandler,
150{
151    build_grpc_server_with_metrics_auth(handler, noop_api_metrics(), token)
152}
153
154/// Like [`build_grpc_server_with_metrics`] but enforcing a bearer token.
155///
156/// Wraps the configured server (message-size limits preserved) in an [`InterceptedService`] that gates every call on the token.
157pub fn build_grpc_server_with_metrics_auth<H>(
158    handler: Arc<H>,
159    metrics: ApiMetricsHandle,
160    token: Token,
161) -> InterceptedService<TaskServiceServer<TaskApiService<H>>, BearerAuth>
162where
163    H: ApiHandler,
164{
165    InterceptedService::new(
166        build_grpc_server_with_metrics(handler, metrics),
167        BearerAuth { expected: token },
168    )
169}
170
171#[tonic::async_trait]
172impl<H> TaskService for TaskApiService<H>
173where
174    H: ApiHandler,
175{
176    async fn submit_task(
177        &self,
178        request: Request<proto_api::SubmitTaskRequest>,
179    ) -> Result<Response<proto_api::SubmitTaskResponse>, Status> {
180        self.instrument("SubmitTask", async move {
181            let req = request.into_inner();
182
183            let spec = req
184                .spec
185                .ok_or_else(|| Status::invalid_argument("missing spec"))?;
186
187            let spec =
188                crate::convert::convert_create_spec(spec).map_err(|e: ApiError| Status::from(e))?;
189
190            debug!(slot = %spec.slot(), kind = ?spec.kind(), "grpc: submitting task");
191            let task_id = self.handler.submit_task(spec).await.map_err(Status::from)?;
192
193            Ok(Response::new(proto_api::SubmitTaskResponse {
194                task_id: task_id.to_string(),
195            }))
196        })
197        .await
198    }
199
200    async fn apply_task(
201        &self,
202        request: Request<proto_api::ApplyTaskRequest>,
203    ) -> Result<Response<proto_api::ApplyTaskResponse>, Status> {
204        self.instrument("ApplyTask", async move {
205            let req = request.into_inner();
206
207            let spec = req
208                .spec
209                .ok_or_else(|| Status::invalid_argument("missing spec"))?;
210
211            let spec =
212                crate::convert::convert_create_spec(spec).map_err(|e: ApiError| Status::from(e))?;
213
214            debug!(slot = %spec.slot(), kind = ?spec.kind(), "grpc: applying task");
215            let task_id = self.handler.apply_task(spec).await.map_err(Status::from)?;
216
217            Ok(Response::new(proto_api::ApplyTaskResponse {
218                task_id: task_id.to_string(),
219            }))
220        })
221        .await
222    }
223
224    async fn get_task_status(
225        &self,
226        request: Request<proto_api::GetTaskStatusRequest>,
227    ) -> Result<Response<proto_api::GetTaskStatusResponse>, Status> {
228        self.instrument("GetTaskStatus", async move {
229            let req = request.into_inner();
230
231            non_empty_id("task_id", &req.task_id).map_err(Status::from)?;
232
233            let task_id = solti_model::TaskId::from(req.task_id);
234            debug!(%task_id, "grpc: getting task status");
235
236            let info = self
237                .handler
238                .get_task_status(&task_id)
239                .await
240                .map_err(Status::from)?;
241
242            let task = info
243                .map(proto_api::TaskData::try_from)
244                .transpose()
245                .map_err(Status::from)?;
246
247            Ok(Response::new(proto_api::GetTaskStatusResponse { task }))
248        })
249        .await
250    }
251
252    async fn list_tasks(
253        &self,
254        request: Request<proto_api::ListTasksRequest>,
255    ) -> Result<Response<proto_api::ListTasksResponse>, Status> {
256        self.instrument("ListTasks", async move {
257            let req = request.into_inner();
258
259            let mut query = TaskQuery::new();
260
261            if let Some(slot) = req.slot {
262                non_empty_id("slot", &slot).map_err(Status::from)?;
263                query = query.with_slot(slot);
264            }
265
266            if let Some(status_raw) = req.status {
267                let status = proto_to_domain_status(status_raw).map_err(Status::from)?;
268                query = query.with_status(status);
269            }
270
271            query = query.with_limit(clamp_list_limit(req.limit));
272            if req.offset > 0 {
273                query = query.with_offset(req.offset as usize);
274            }
275
276            let page = self
277                .handler
278                .query_tasks(query)
279                .await
280                .map_err(Status::from)?;
281
282            debug!(
283                count = page.items.len(),
284                total = page.total,
285                "grpc: tasks listed"
286            );
287
288            let response = tasks_page_to_proto(page).map_err(Status::from)?;
289            Ok(Response::new(response))
290        })
291        .await
292    }
293
294    async fn list_task_runs(
295        &self,
296        request: Request<proto_api::ListTaskRunsRequest>,
297    ) -> Result<Response<proto_api::ListTaskRunsResponse>, Status> {
298        self.instrument("ListTaskRuns", async move {
299            let req = request.into_inner();
300
301            non_empty_id("task_id", &req.task_id).map_err(Status::from)?;
302
303            let task_id = solti_model::TaskId::from(req.task_id);
304            debug!(%task_id, "grpc: listing task runs");
305
306            let runs = self
307                .handler
308                .list_task_runs(&task_id)
309                .await
310                .map_err(Status::from)?;
311
312            let runs = runs.into_iter().map(proto_api::TaskRunInfo::from).collect();
313
314            Ok(Response::new(proto_api::ListTaskRunsResponse { runs }))
315        })
316        .await
317    }
318
319    async fn delete_task(
320        &self,
321        request: Request<proto_api::DeleteTaskRequest>,
322    ) -> Result<Response<proto_api::DeleteTaskResponse>, Status> {
323        self.instrument("DeleteTask", async move {
324            let req = request.into_inner();
325
326            non_empty_id("task_id", &req.task_id).map_err(Status::from)?;
327
328            let task_id = solti_model::TaskId::from(req.task_id);
329            debug!(%task_id, "grpc: deleting task");
330
331            self.handler
332                .delete_task(&task_id)
333                .await
334                .map_err(Status::from)?;
335
336            debug!(%task_id, "grpc: task deleted");
337            Ok(Response::new(proto_api::DeleteTaskResponse {}))
338        })
339        .await
340    }
341
342    /// Server-streaming RPC.
343    type StreamTaskLogsStream = Pin<
344        Box<
345            dyn tokio_stream::Stream<Item = Result<proto_api::StreamTaskLogsResponse, Status>>
346                + Send
347                + 'static,
348        >,
349    >;
350
351    async fn stream_task_logs(
352        &self,
353        request: Request<proto_api::StreamTaskLogsRequest>,
354    ) -> Result<Response<Self::StreamTaskLogsStream>, Status> {
355        let req = request.into_inner();
356        non_empty_id("task_id", &req.task_id).map_err(Status::from)?;
357
358        let task_id = solti_model::TaskId::from(req.task_id);
359        debug!(%task_id, "grpc: subscribing to task log stream");
360
361        let domain_stream = self
362            .handler
363            .stream_task_logs(&task_id)
364            .await
365            .map_err(Status::from)?;
366
367        let proto_stream = domain_stream.map(|ev| Ok(output_event_to_proto(ev)));
368        Ok(Response::new(Box::pin(proto_stream)))
369    }
370}
371
372#[cfg(test)]
373mod tests {
374    use super::*;
375
376    use std::time::{Duration, UNIX_EPOCH};
377
378    use async_trait::async_trait;
379    use bytes::Bytes;
380    use solti_model::{
381        OutputChunk, OutputEvent, StreamKind as ModelStreamKind, Task, TaskId, TaskPage, TaskQuery,
382        TaskRun, TaskSpec,
383    };
384
385    use crate::error::ApiError;
386    use crate::handler::{ApiHandler, OutputEventStream};
387
388    struct StreamMock;
389
390    #[async_trait]
391    impl ApiHandler for StreamMock {
392        async fn submit_task(&self, _spec: TaskSpec) -> Result<TaskId, ApiError> {
393            unreachable!()
394        }
395        async fn get_task_status(&self, _id: &TaskId) -> Result<Option<Task>, ApiError> {
396            unreachable!()
397        }
398        async fn query_tasks(&self, _q: TaskQuery) -> Result<TaskPage<Task>, ApiError> {
399            unreachable!()
400        }
401        async fn list_task_runs(&self, _id: &TaskId) -> Result<Vec<TaskRun>, ApiError> {
402            unreachable!()
403        }
404        async fn delete_task(&self, _id: &TaskId) -> Result<(), ApiError> {
405            unreachable!()
406        }
407        async fn stream_task_logs(&self, id: &TaskId) -> Result<OutputEventStream, ApiError> {
408            if id.as_str() == "missing" {
409                return Err(ApiError::TaskNotFound(id.to_string()));
410            }
411            let events = vec![
412                OutputEvent::RunStarted {
413                    attempt: 1,
414                    started_at: UNIX_EPOCH + Duration::from_millis(1000),
415                },
416                OutputEvent::Chunk(OutputChunk {
417                    attempt: 1,
418                    stream: ModelStreamKind::Stdout,
419                    seq: 0,
420                    ts: UNIX_EPOCH + Duration::from_millis(1100),
421                    line: Bytes::from_static(b"hello-grpc"),
422                }),
423                OutputEvent::RunFinished {
424                    attempt: 1,
425                    exit_code: Some(0),
426                    finished_at: UNIX_EPOCH + Duration::from_millis(1500),
427                },
428            ];
429            Ok(Box::pin(tokio_stream::iter(events)))
430        }
431    }
432
433    fn service() -> TaskApiService<StreamMock> {
434        TaskApiService::new(Arc::new(StreamMock))
435    }
436
437    #[tokio::test]
438    async fn stream_task_logs_returns_three_proto_events_in_order() {
439        let svc = service();
440        let req = Request::new(proto_api::StreamTaskLogsRequest {
441            task_id: "tsk_1".into(),
442        });
443
444        let response = svc.stream_task_logs(req).await.expect("stream Ok");
445        let mut stream = response.into_inner();
446
447        match stream.next().await.unwrap().unwrap().kind.unwrap() {
448            proto_api::stream_task_logs_response::Kind::RunStarted(r) => {
449                assert_eq!(r.attempt, 1);
450                assert_eq!(r.started_at, 1000);
451            }
452            other => panic!("expected RunStarted, got {other:?}"),
453        }
454
455        match stream.next().await.unwrap().unwrap().kind.unwrap() {
456            proto_api::stream_task_logs_response::Kind::Chunk(c) => {
457                assert_eq!(c.attempt, 1);
458                assert_eq!(c.stream, proto_api::OutputStreamKind::Stdout as i32);
459                assert_eq!(c.seq, 0);
460                assert_eq!(&c.line[..], b"hello-grpc");
461            }
462            other => panic!("expected Chunk, got {other:?}"),
463        }
464
465        match stream.next().await.unwrap().unwrap().kind.unwrap() {
466            proto_api::stream_task_logs_response::Kind::RunFinished(r) => {
467                assert_eq!(r.attempt, 1);
468                assert_eq!(r.exit_code, Some(0));
469                assert_eq!(r.finished_at, 1500);
470            }
471            other => panic!("expected RunFinished, got {other:?}"),
472        }
473        assert!(stream.next().await.is_none(), "stream must terminate");
474    }
475
476    #[tokio::test]
477    async fn stream_task_logs_rejects_empty_task_id() {
478        let svc = service();
479        let req = Request::new(proto_api::StreamTaskLogsRequest {
480            task_id: "  ".into(),
481        });
482        let status = match svc.stream_task_logs(req).await {
483            Err(s) => s,
484            Ok(_) => panic!("expected error status"),
485        };
486        assert_eq!(status.code(), tonic::Code::InvalidArgument);
487    }
488
489    #[tokio::test]
490    async fn stream_task_logs_maps_task_not_found_to_not_found_status() {
491        let svc = service();
492        let req = Request::new(proto_api::StreamTaskLogsRequest {
493            task_id: "missing".into(),
494        });
495        let status = match svc.stream_task_logs(req).await {
496            Err(s) => s,
497            Ok(_) => panic!("expected error status"),
498        };
499        assert_eq!(status.code(), tonic::Code::NotFound);
500    }
501
502    #[test]
503    fn bearer_value_accepts_scheme_case_insensitively() {
504        assert_eq!(bearer_value("Bearer tok"), Some("tok"));
505        assert_eq!(bearer_value("bearer tok"), Some("tok"));
506        assert_eq!(bearer_value("BEARER tok"), Some("tok"));
507        assert_eq!(bearer_value("BeArEr tok"), Some("tok"));
508        assert_eq!(bearer_value("Bearer a b"), Some("a b"));
509        assert_eq!(bearer_value("Basic tok"), None);
510        assert_eq!(bearer_value("tok"), None);
511        assert_eq!(bearer_value(""), None);
512    }
513}