Skip to main content

solti_api/
grpc.rs

1//! # gRPC Transport
2//!
3//! Tonic service for protobuf package `solti.task.v1`.
4//! Every RPC delegates to [`ApiHandler`].
5//!
6//! [`GrpcApi`] installs message limits, optional metrics, and bearer authentication.
7//! [`wire`] exposes the generated client, server, and message types.
8//!
9//! ## RPCs
10//!
11//! | RPC              | Shape            | Operation   |
12//! |------------------|------------------|-------------|
13//! | `CreateTask`     | Unary            | Create      |
14//! | `ApplyTask`      | Unary            | Apply       |
15//! | `GetTask`        | Unary            | Get         |
16//! | `ListTasks`      | Unary            | List        |
17//! | `WatchTasks`     | Server streaming | Watch       |
18//! | `ListTaskRuns`   | Unary            | Run history |
19//! | `DeleteTask`     | Unary            | Delete      |
20//! | `StreamTaskLogs` | Server streaming | Live output |
21//!
22//! Domain failures become `tonic::Status`.
23//! Stream failures terminate the corresponding stream.
24
25use std::pin::Pin;
26use std::sync::Arc;
27
28use tokio_stream::{Stream, StreamExt};
29use tonic::service::Interceptor;
30use tonic::service::interceptor::InterceptedService;
31use tonic::{Request, Response, Status};
32use tracing::debug;
33
34use solti_model::{LabelSelector, TaskFilter, TaskQuery, Token};
35
36use crate::GRPC_API_SERVICE;
37use crate::auth::bearer_value;
38use crate::handler::ApiHandler;
39use crate::metrics::{
40    ApiMetricsHandle, InFlightGuard, RequestMetrics, Transport, noop_api_metrics,
41};
42use crate::proto_api::{
43    self, task_service_server::TaskService, task_service_server::TaskServiceServer,
44};
45use crate::validate::{parse_list_limit, parse_task_id, validate_slot};
46
47mod convert;
48use convert::{
49    output_event_to_proto, proto_to_domain_phase, task_watch_event_to_proto, tasks_page_to_proto,
50    write_preconditions_from_proto,
51};
52
53/// Generated protobuf types for the current Task API.
54///
55/// This module includes messages, enums, [`crate::grpc::wire::TaskServiceClient`], and [`crate::grpc::wire::TaskServiceServer`].
56///
57/// ## Example
58///
59/// ```rust,no_run
60/// use solti_api::grpc::wire::TaskServiceClient;
61///
62/// async fn connect() -> Result<(), solti_api::tonic::transport::Error> {
63///     let client = TaskServiceClient::connect("http://127.0.0.1:50052").await?;
64///     let _ = client;
65///     Ok(())
66/// }
67/// ```
68pub mod wire {
69    pub use crate::proto_api::task_service_client::TaskServiceClient;
70    pub use crate::proto_api::task_service_server::TaskServiceServer;
71    pub use crate::proto_api::*;
72}
73
74type ServerStream<T> = Pin<Box<dyn Stream<Item = Result<T, Status>> + Send + 'static>>;
75
76struct GrpcMetricsStream<T> {
77    inner: ServerStream<T>,
78    request: RequestMetrics,
79}
80
81impl<T> Stream for GrpcMetricsStream<T> {
82    type Item = Result<T, Status>;
83
84    fn poll_next(
85        mut self: Pin<&mut Self>,
86        context: &mut std::task::Context<'_>,
87    ) -> std::task::Poll<Option<Self::Item>> {
88        match self.inner.as_mut().poll_next(context) {
89            std::task::Poll::Ready(Some(Err(status))) => {
90                self.request.complete(status.code() as u16);
91                std::task::Poll::Ready(Some(Err(status)))
92            }
93            std::task::Poll::Ready(None) => {
94                self.request.complete(tonic::Code::Ok as u16);
95                std::task::Poll::Ready(None)
96            }
97            poll => poll,
98        }
99    }
100}
101
102/// Generated `TaskService` implementation over an [`ApiHandler`].
103///
104/// This is the lower-level service implementation.
105/// It converts protobuf values and records optional metrics.
106/// Use [`GrpcApi`] to also install message limits and authentication.
107///
108/// ## See Also
109///
110/// - [`GrpcApi`] builds the configured public service.
111/// - [`ApiError`](crate::ApiError) defines gRPC error mapping.
112pub struct TaskApiService<H> {
113    handler: Arc<H>,
114    metrics: ApiMetricsHandle,
115}
116
117impl<H> TaskApiService<H>
118where
119    H: ApiHandler,
120{
121    /// Creates a service with the no-op metrics backend.
122    pub fn new(handler: Arc<H>) -> Self {
123        Self::new_with_metrics(handler, noop_api_metrics())
124    }
125
126    /// Creates a service with an explicit metrics backend.
127    pub fn new_with_metrics(handler: Arc<H>, metrics: ApiMetricsHandle) -> Self {
128        Self { handler, metrics }
129    }
130
131    async fn instrument<F, T>(&self, method: &'static str, fut: F) -> Result<Response<T>, Status>
132    where
133        F: Future<Output = Result<Response<T>, Status>>,
134    {
135        let path = format!("/{GRPC_API_SERVICE}/{method}");
136        let mut request = RequestMetrics::enter(&self.metrics, Transport::Grpc, method, path);
137        let result = fut.await;
138        let status = match &result {
139            Ok(_) => 0u16,
140            Err(s) => s.code() as u16,
141        };
142        request.complete(status);
143        result
144    }
145
146    async fn instrument_stream<F, T>(
147        &self,
148        method: &'static str,
149        fut: F,
150    ) -> Result<Response<ServerStream<T>>, Status>
151    where
152        F: Future<Output = Result<ServerStream<T>, Status>>,
153        T: 'static,
154    {
155        let path = format!("/{GRPC_API_SERVICE}/{method}");
156        let mut request = RequestMetrics::enter(&self.metrics, Transport::Grpc, method, path);
157        match fut.await {
158            Ok(stream) => {
159                let stream: ServerStream<T> = Box::pin(GrpcMetricsStream {
160                    inner: stream,
161                    request,
162                });
163                Ok(Response::new(stream))
164            }
165            Err(status) => {
166                request.complete(status.code() as u16);
167                Err(status)
168            }
169        }
170    }
171}
172
173/// Complete gRPC service returned by [`GrpcApi::server`].
174///
175/// The [`BearerAuth`] interceptor is always present.
176/// It passes calls through when no token is configured.
177pub type GrpcServer<H> = InterceptedService<TaskServiceServer<TaskApiService<H>>, BearerAuth>;
178
179/// Builder for the tonic task API.
180///
181/// Authentication and metrics are optional.
182/// [`server`](Self::server) installs the public message-size limit.
183///
184/// ## Example
185///
186/// ```rust,no_run
187/// # use std::sync::Arc;
188/// # use solti_api::{ApiHandler, GrpcApi};
189/// # async fn example<H: ApiHandler>(adapter: Arc<H>) -> Result<(), Box<dyn std::error::Error>> {
190/// let svc = GrpcApi::new(adapter).server();
191/// solti_api::tonic::transport::Server::builder()
192///     .add_service(svc)
193///     .serve("0.0.0.0:50052".parse()?)
194///     .await?;
195/// # Ok(()) }
196/// ```
197///
198/// ## See Also
199///
200/// - [`ApiHandler`] defines the backend operations.
201/// - [`ApiError`](crate::ApiError) defines the gRPC status mapping.
202pub struct GrpcApi<H> {
203    handler: Arc<H>,
204    metrics: ApiMetricsHandle,
205    auth: Option<Token>,
206}
207
208impl<H> GrpcApi<H>
209where
210    H: ApiHandler,
211{
212    /// Creates a gRPC API for one handler.
213    pub fn new(handler: Arc<H>) -> Self {
214        Self {
215            handler,
216            metrics: noop_api_metrics(),
217            auth: None,
218        }
219    }
220
221    /// Requires a bearer token on every call.
222    ///
223    /// The expected metadata is `authorization: Bearer <token>`.
224    /// Missing or invalid credentials return `Unauthenticated`.
225    /// Rejected calls do not reach the handler.
226    /// Authentication is disabled when this method is not called.
227    pub fn with_auth(mut self, token: Token) -> Self {
228        self.auth = Some(token);
229        self
230    }
231
232    /// Attaches a metrics backend.
233    ///
234    /// The default backend ignores every update.
235    pub fn with_metrics(mut self, metrics: ApiMetricsHandle) -> Self {
236        self.metrics = metrics;
237        self
238    }
239
240    /// Builds the configured gRPC service.
241    ///
242    /// Encoded and decoded messages are limited to
243    /// [`MAX_REQUEST_BYTES`](crate::MAX_REQUEST_BYTES).
244    /// The returned service always contains [`BearerAuth`].
245    pub fn server(self) -> GrpcServer<H> {
246        let inner = TaskServiceServer::new(TaskApiService::new_with_metrics(
247            self.handler,
248            Arc::clone(&self.metrics),
249        ))
250        .max_decoding_message_size(crate::MAX_REQUEST_BYTES)
251        .max_encoding_message_size(crate::MAX_REQUEST_BYTES);
252        InterceptedService::new(
253            inner,
254            BearerAuth {
255                expected: self.auth,
256                metrics: self.metrics,
257            },
258        )
259    }
260}
261
262/// Bearer interceptor used by [`GrpcServer`].
263///
264/// It verifies `authorization: Bearer <token>` metadata.
265/// Token comparison uses [`Token::verify`].
266/// Without a configured token, every call passes through.
267/// Configure it through [`GrpcApi::with_auth`].
268#[derive(Clone)]
269pub struct BearerAuth {
270    expected: Option<Token>,
271    metrics: ApiMetricsHandle,
272}
273
274impl Interceptor for BearerAuth {
275    fn call(&mut self, request: Request<()>) -> Result<Request<()>, Status> {
276        let Some(expected) = &self.expected else {
277            return Ok(request);
278        };
279        let ok = request
280            .metadata()
281            .get("authorization")
282            .and_then(|v| v.to_str().ok())
283            .and_then(bearer_value)
284            .map(|presented| expected.verify(presented))
285            .unwrap_or(false);
286
287        if ok {
288            Ok(request)
289        } else {
290            record_auth_failure(&self.metrics, &request);
291            Err(Status::unauthenticated("missing or invalid bearer token"))
292        }
293    }
294}
295
296fn record_auth_failure(metrics: &ApiMetricsHandle, request: &Request<()>) {
297    let method = request
298        .extensions()
299        .get::<tonic::GrpcMethod<'static>>()
300        .map(tonic::GrpcMethod::method)
301        .unwrap_or("<unknown>");
302    let path = request
303        .extensions()
304        .get::<tonic::GrpcMethod<'static>>()
305        .map(|grpc| format!("/{}/{}", grpc.service(), grpc.method()))
306        .unwrap_or_else(|| "<unknown>".to_owned());
307    let _in_flight = InFlightGuard::enter(metrics, Transport::Grpc);
308    metrics.record_request(
309        Transport::Grpc,
310        method,
311        &path,
312        tonic::Code::Unauthenticated as u16,
313        0,
314    );
315}
316
317fn task_filter_from_wire(
318    slot: Option<String>,
319    phases: Vec<i32>,
320    label_selector: String,
321) -> Result<TaskFilter, crate::ApiError> {
322    let mut filter = TaskFilter::new();
323
324    if let Some(slot) = slot {
325        filter = filter.with_slot(validate_slot(slot)?);
326    }
327
328    for phase_raw in phases {
329        filter = filter.with_phase(proto_to_domain_phase(phase_raw)?);
330    }
331
332    if !label_selector.is_empty() {
333        let selector = label_selector.parse::<LabelSelector>().map_err(|error| {
334            crate::ApiError::InvalidRequest(format!("invalid labelSelector: {error}"))
335        })?;
336        filter = filter
337            .with_label_selector(selector)
338            .map_err(|error| crate::ApiError::InvalidRequest(error.to_string()))?;
339    }
340
341    Ok(filter)
342}
343
344#[tonic::async_trait]
345impl<H> TaskService for TaskApiService<H>
346where
347    H: ApiHandler,
348{
349    async fn create_task(
350        &self,
351        request: Request<proto_api::CreateTaskRequest>,
352    ) -> Result<Response<proto_api::CreateTaskResponse>, Status> {
353        self.instrument("CreateTask", async move {
354            let req = request.into_inner();
355
356            let manifest = req
357                .manifest
358                .ok_or_else(|| Status::invalid_argument("missing manifest"))?;
359            let manifest = convert::task_manifest_from_proto(manifest).map_err(Status::from)?;
360            debug!(name = %manifest.name(), "grpc: creating task");
361            let task = self
362                .handler
363                .create_task(manifest)
364                .await
365                .map_err(Status::from)?;
366            let task = proto_api::Task::try_from(task).map_err(Status::from)?;
367
368            Ok(Response::new(proto_api::CreateTaskResponse {
369                task: Some(task),
370            }))
371        })
372        .await
373    }
374
375    async fn apply_task(
376        &self,
377        request: Request<proto_api::ApplyTaskRequest>,
378    ) -> Result<Response<proto_api::ApplyTaskResponse>, Status> {
379        self.instrument("ApplyTask", async move {
380            let req = request.into_inner();
381
382            let manifest = req
383                .manifest
384                .ok_or_else(|| Status::invalid_argument("missing manifest"))?;
385            let manifest = convert::task_manifest_from_proto(manifest).map_err(Status::from)?;
386            let preconditions =
387                write_preconditions_from_proto(req.preconditions).map_err(Status::from)?;
388            debug!(name = %manifest.name(), "grpc: applying task");
389            let task = self
390                .handler
391                .apply_task(manifest, preconditions)
392                .await
393                .map_err(Status::from)?;
394            let task = proto_api::Task::try_from(task).map_err(Status::from)?;
395
396            Ok(Response::new(proto_api::ApplyTaskResponse {
397                task: Some(task),
398            }))
399        })
400        .await
401    }
402
403    async fn get_task(
404        &self,
405        request: Request<proto_api::GetTaskRequest>,
406    ) -> Result<Response<proto_api::GetTaskResponse>, Status> {
407        self.instrument("GetTask", async move {
408            let req = request.into_inner();
409
410            let task_id = parse_task_id("task name", req.name).map_err(Status::from)?;
411            debug!(%task_id, "grpc: getting task status");
412
413            let task = self
414                .handler
415                .get_task(&task_id)
416                .await
417                .map_err(Status::from)?
418                .ok_or_else(|| Status::from(crate::ApiError::TaskNotFound(task_id.to_string())))?;
419
420            let task = proto_api::Task::try_from(task).map_err(Status::from)?;
421
422            Ok(Response::new(proto_api::GetTaskResponse {
423                task: Some(task),
424            }))
425        })
426        .await
427    }
428
429    async fn list_tasks(
430        &self,
431        request: Request<proto_api::ListTasksRequest>,
432    ) -> Result<Response<proto_api::ListTasksResponse>, Status> {
433        self.instrument("ListTasks", async move {
434            let req = request.into_inner();
435
436            let filter = task_filter_from_wire(req.slot, req.phases, req.label_selector)
437                .map_err(Status::from)?;
438            let mut query = TaskQuery::from_filter(filter);
439
440            query = query.with_limit(parse_list_limit(req.limit).map_err(Status::from)?);
441            if !req.r#continue.is_empty() {
442                query = query.with_continuation(
443                    crate::continuation::decode(&req.r#continue).map_err(Status::from)?,
444                );
445            }
446
447            let page_filter = query.filter().clone();
448            let page_limit = query.limit();
449            let page = self
450                .handler
451                .query_tasks(query)
452                .await
453                .map_err(Status::from)?;
454            crate::continuation::validate_page(&page, &page_filter, page_limit)
455                .map_err(Status::from)?;
456
457            debug!(
458                count = page.items.len(),
459                remaining = page.remaining_item_count,
460                "grpc: tasks listed"
461            );
462
463            let response = tasks_page_to_proto(page).map_err(Status::from)?;
464            Ok(Response::new(response))
465        })
466        .await
467    }
468
469    /// Server-streaming Task collection watch.
470    type WatchTasksStream = ServerStream<proto_api::WatchTasksResponse>;
471
472    async fn watch_tasks(
473        &self,
474        request: Request<proto_api::WatchTasksRequest>,
475    ) -> Result<Response<Self::WatchTasksStream>, Status> {
476        self.instrument_stream("WatchTasks", async move {
477            let req = request.into_inner();
478            let filter = task_filter_from_wire(req.slot, req.phases, req.label_selector)
479                .map_err(Status::from)?;
480            if req
481                .resource_version
482                .as_deref()
483                .is_some_and(|value| value.trim().is_empty())
484            {
485                return Err(Status::invalid_argument(
486                    "resource_version must not be empty",
487                ));
488            }
489
490            let domain_stream = self
491                .handler
492                .watch_tasks(filter, req.resource_version)
493                .await
494                .map_err(Status::from)?;
495            let proto_stream = domain_stream.map(|event| match event {
496                Ok(event) => task_watch_event_to_proto(event).map_err(Status::from),
497                Err(error) => Err(Status::from(error)),
498            });
499            let stream: Self::WatchTasksStream = Box::pin(proto_stream);
500            Ok(stream)
501        })
502        .await
503    }
504
505    async fn list_task_runs(
506        &self,
507        request: Request<proto_api::ListTaskRunsRequest>,
508    ) -> Result<Response<proto_api::ListTaskRunsResponse>, Status> {
509        self.instrument("ListTaskRuns", async move {
510            let req = request.into_inner();
511
512            let task_id = parse_task_id("task name", req.name).map_err(Status::from)?;
513            debug!(%task_id, "grpc: listing task runs");
514
515            let runs = self
516                .handler
517                .list_task_runs(&task_id)
518                .await
519                .map_err(Status::from)?;
520
521            let runs = runs
522                .into_iter()
523                .map(proto_api::TaskRunInfo::try_from)
524                .collect::<Result<_, _>>()
525                .map_err(Status::from)?;
526
527            Ok(Response::new(proto_api::ListTaskRunsResponse { runs }))
528        })
529        .await
530    }
531
532    async fn delete_task(
533        &self,
534        request: Request<proto_api::DeleteTaskRequest>,
535    ) -> Result<Response<proto_api::DeleteTaskResponse>, Status> {
536        self.instrument("DeleteTask", async move {
537            let req = request.into_inner();
538
539            let task_id = parse_task_id("task name", req.name).map_err(Status::from)?;
540            let preconditions =
541                write_preconditions_from_proto(req.preconditions).map_err(Status::from)?;
542            debug!(%task_id, "grpc: deleting task");
543
544            self.handler
545                .delete_task(&task_id, preconditions)
546                .await
547                .map_err(Status::from)?;
548
549            debug!(%task_id, "grpc: task deleted");
550            Ok(Response::new(proto_api::DeleteTaskResponse {}))
551        })
552        .await
553    }
554
555    /// Server-streaming RPC.
556    type StreamTaskLogsStream = ServerStream<proto_api::StreamTaskLogsResponse>;
557
558    async fn stream_task_logs(
559        &self,
560        request: Request<proto_api::StreamTaskLogsRequest>,
561    ) -> Result<Response<Self::StreamTaskLogsStream>, Status> {
562        self.instrument_stream("StreamTaskLogs", async move {
563            let req = request.into_inner();
564            let task_id = parse_task_id("task name", req.name).map_err(Status::from)?;
565            debug!(%task_id, "grpc: subscribing to task log stream");
566
567            let domain_stream = self
568                .handler
569                .stream_task_logs(&task_id)
570                .await
571                .map_err(Status::from)?;
572
573            let proto_stream =
574                domain_stream.map(|event| output_event_to_proto(event).map_err(Status::from));
575            let stream: Self::StreamTaskLogsStream = Box::pin(proto_stream);
576            Ok(stream)
577        })
578        .await
579    }
580}
581
582#[cfg(test)]
583mod tests {
584    use super::*;
585
586    use std::time::{Duration, UNIX_EPOCH};
587
588    use async_trait::async_trait;
589    use bytes::Bytes;
590    use solti_model::{
591        ExtensionWorkload, OutputChunk, OutputEvent, StreamKind as ModelStreamKind, Task,
592        TaskContinuation, TaskFilter, TaskId, TaskManifest, TaskPage, TaskPhase, TaskQuery,
593        TaskRun, TaskSpec, TaskWatchEvent, TaskWorkload, WORKLOAD_API_VERSION, WorkloadTypeMeta,
594        WritePreconditions,
595    };
596
597    use crate::error::ApiError;
598    use crate::handler::{ApiHandler, OutputEventStream, TaskWatchEventStream};
599
600    #[derive(Default)]
601    struct StreamMock {
602        last_preconditions: std::sync::Mutex<Option<WritePreconditions>>,
603        last_query: std::sync::Mutex<Option<TaskQuery>>,
604        last_watch_filter: std::sync::Mutex<Option<TaskFilter>>,
605        last_watch_resource_version: std::sync::Mutex<Option<Option<String>>>,
606        watch_expired: bool,
607        watch_stream_expired: bool,
608        log_stream_pending: bool,
609    }
610
611    #[async_trait]
612    impl ApiHandler for StreamMock {
613        async fn create_task(&self, _manifest: TaskManifest) -> Result<Task, ApiError> {
614            unreachable!()
615        }
616        async fn apply_task(
617            &self,
618            manifest: TaskManifest,
619            preconditions: WritePreconditions,
620        ) -> Result<Task, ApiError> {
621            *self.last_preconditions.lock().unwrap() = Some(preconditions);
622            Task::from_manifest(manifest).map_err(|error| ApiError::Internal(error.to_string()))
623        }
624        async fn get_task(&self, _id: &TaskId) -> Result<Option<Task>, ApiError> {
625            Ok(None)
626        }
627        async fn query_tasks(&self, query: TaskQuery) -> Result<TaskPage<Task>, ApiError> {
628            *self.last_query.lock().unwrap() = Some(query);
629            Ok(TaskPage {
630                items: vec![],
631                resource_version: "test:1".into(),
632                continuation: None,
633                remaining_item_count: 0,
634            })
635        }
636        async fn watch_tasks(
637            &self,
638            filter: TaskFilter,
639            resource_version: Option<String>,
640        ) -> Result<TaskWatchEventStream, ApiError> {
641            *self.last_watch_filter.lock().unwrap() = Some(filter);
642            *self.last_watch_resource_version.lock().unwrap() = Some(resource_version);
643            if self.watch_expired {
644                return Err(ApiError::ResourceVersionExpired(
645                    "requested resourceVersion is no longer retained".into(),
646                ));
647            }
648            let mut events = vec![Ok(TaskWatchEvent::Added(watch_task()))];
649            if self.watch_stream_expired {
650                events.push(Err(ApiError::ResourceVersionExpired(
651                    "watch position is no longer retained".into(),
652                )));
653            }
654            Ok(Box::pin(tokio_stream::iter(events)))
655        }
656        async fn list_task_runs(&self, id: &TaskId) -> Result<Vec<TaskRun>, ApiError> {
657            let workload = if id.as_str() == "embedded-run" {
658                WorkloadTypeMeta::new(WORKLOAD_API_VERSION, "Embedded").unwrap()
659            } else {
660                WorkloadTypeMeta::new("workloads.example.io/v1", "DatabaseBackup").unwrap()
661            };
662            Ok(vec![TaskRun::starting(2, 1, workload).unwrap()])
663        }
664        async fn delete_task(
665            &self,
666            _id: &TaskId,
667            preconditions: WritePreconditions,
668        ) -> Result<(), ApiError> {
669            *self.last_preconditions.lock().unwrap() = Some(preconditions);
670            Ok(())
671        }
672        async fn stream_task_logs(&self, id: &TaskId) -> Result<OutputEventStream, ApiError> {
673            if id.as_str() == "missing" {
674                return Err(ApiError::TaskNotFound(id.to_string()));
675            }
676            if self.log_stream_pending {
677                return Ok(Box::pin(tokio_stream::pending()));
678            }
679            let events = vec![
680                OutputEvent::RunStarted {
681                    generation: 2,
682                    attempt: 1,
683                    started_at: UNIX_EPOCH + Duration::from_millis(1000),
684                },
685                OutputEvent::Chunk(OutputChunk {
686                    generation: 2,
687                    attempt: 1,
688                    stream: ModelStreamKind::Stdout,
689                    seq: 0,
690                    ts: UNIX_EPOCH + Duration::from_millis(1100),
691                    line: Bytes::from_static(b"hello-grpc"),
692                }),
693                OutputEvent::RunFinished {
694                    generation: 2,
695                    attempt: 1,
696                    exit_code: Some(0),
697                    finished_at: UNIX_EPOCH + Duration::from_millis(1500),
698                },
699            ];
700            Ok(Box::pin(tokio_stream::iter(events)))
701        }
702    }
703
704    fn service() -> TaskApiService<StreamMock> {
705        TaskApiService::new(Arc::new(StreamMock::default()))
706    }
707
708    fn watch_task() -> Task {
709        let workload = TaskWorkload::Extension(
710            ExtensionWorkload::new(
711                "workloads.example.io/v1",
712                "ExampleJob",
713                serde_json::json!({"value": 1}),
714            )
715            .unwrap(),
716        );
717        let spec = TaskSpec::builder("primary", workload, 5_000_u64)
718            .build()
719            .unwrap();
720        let mut task = Task::new("watch-task", spec).unwrap();
721        task.set_resource_version("test:2").unwrap();
722        task
723    }
724
725    #[tokio::test]
726    async fn get_task_maps_missing_resource_to_not_found_status() {
727        let status = service()
728            .get_task(Request::new(proto_api::GetTaskRequest {
729                name: "missing".into(),
730            }))
731            .await
732            .unwrap_err();
733
734        assert_eq!(status.code(), tonic::Code::NotFound);
735    }
736
737    #[tokio::test]
738    async fn delete_task_forwards_write_preconditions() {
739        let handler = Arc::new(StreamMock::default());
740        let service = TaskApiService::new(Arc::clone(&handler));
741
742        service
743            .delete_task(Request::new(proto_api::DeleteTaskRequest {
744                name: "task-1".into(),
745                preconditions: Some(proto_api::WritePreconditions {
746                    uid: Some("uid-1".into()),
747                    resource_version: Some("17".into()),
748                }),
749            }))
750            .await
751            .unwrap();
752
753        let preconditions = handler
754            .last_preconditions
755            .lock()
756            .unwrap()
757            .clone()
758            .expect("handler received preconditions");
759        assert_eq!(preconditions.uid().unwrap().as_str(), "uid-1");
760        assert_eq!(preconditions.resource_version(), Some("17"));
761    }
762
763    #[tokio::test]
764    async fn delete_task_rejects_empty_write_precondition() {
765        let handler = Arc::new(StreamMock::default());
766        let service = TaskApiService::new(Arc::clone(&handler));
767
768        let status = service
769            .delete_task(Request::new(proto_api::DeleteTaskRequest {
770                name: "task-1".into(),
771                preconditions: Some(proto_api::WritePreconditions {
772                    uid: None,
773                    resource_version: Some(String::new()),
774                }),
775            }))
776            .await
777            .unwrap_err();
778
779        assert_eq!(status.code(), tonic::Code::InvalidArgument);
780        assert!(handler.last_preconditions.lock().unwrap().is_none());
781    }
782
783    #[tokio::test]
784    async fn list_tasks_forwards_filters_and_continuation() {
785        let handler = Arc::new(StreamMock::default());
786        let service = TaskApiService::new(Arc::clone(&handler));
787        let phases = vec![
788            proto_api::TaskPhase::Pending as i32,
789            proto_api::TaskPhase::Running as i32,
790            proto_api::TaskPhase::Pending as i32,
791        ];
792        let label_selector = "environment=production,tier in (frontend,backend)";
793        let filter = task_filter_from_wire(
794            Some("primary".into()),
795            phases.clone(),
796            label_selector.into(),
797        )
798        .unwrap();
799        let continuation =
800            TaskContinuation::new("test:7", filter.clone(), TaskId::new("task-20").unwrap())
801                .unwrap();
802
803        service
804            .list_tasks(Request::new(proto_api::ListTasksRequest {
805                slot: Some("primary".into()),
806                phases,
807                limit: 25,
808                label_selector: label_selector.into(),
809                r#continue: crate::continuation::encode(continuation.clone()).unwrap(),
810            }))
811            .await
812            .unwrap();
813
814        let query = handler
815            .last_query
816            .lock()
817            .unwrap()
818            .take()
819            .expect("handler received query");
820        assert_eq!(query.slot().unwrap().as_str(), "primary");
821        assert_eq!(query.phases(), &[TaskPhase::Pending, TaskPhase::Running]);
822        assert_eq!(query.limit(), 25);
823        assert_eq!(query.continuation(), Some(&continuation));
824        assert_eq!(query.filter(), &filter);
825        assert!(query.matches_labels(&{
826            let mut labels = solti_model::Labels::new();
827            labels
828                .insert("environment", "production")
829                .insert("tier", "backend");
830            labels
831        }));
832    }
833
834    #[tokio::test]
835    async fn list_tasks_rejects_invalid_phase_or_label_selector_before_handler() {
836        let handler = Arc::new(StreamMock::default());
837        let service = TaskApiService::new(Arc::clone(&handler));
838
839        let phase = service
840            .list_tasks(Request::new(proto_api::ListTasksRequest {
841                phases: vec![proto_api::TaskPhase::Unspecified as i32],
842                ..Default::default()
843            }))
844            .await
845            .unwrap_err();
846        assert_eq!(phase.code(), tonic::Code::InvalidArgument);
847
848        let selector = service
849            .list_tasks(Request::new(proto_api::ListTasksRequest {
850                label_selector: "tier in (".into(),
851                ..Default::default()
852            }))
853            .await
854            .unwrap_err();
855        assert_eq!(selector.code(), tonic::Code::InvalidArgument);
856
857        let continuation = service
858            .list_tasks(Request::new(proto_api::ListTasksRequest {
859                r#continue: "not-a-token".into(),
860                ..Default::default()
861            }))
862            .await
863            .unwrap_err();
864        assert_eq!(continuation.code(), tonic::Code::InvalidArgument);
865        assert!(handler.last_query.lock().unwrap().is_none());
866    }
867
868    #[tokio::test]
869    async fn watch_tasks_forwards_filters_and_resource_version() {
870        let handler = Arc::new(StreamMock::default());
871        let service = TaskApiService::new(Arc::clone(&handler));
872
873        let mut stream = service
874            .watch_tasks(Request::new(proto_api::WatchTasksRequest {
875                slot: Some("primary".into()),
876                phases: vec![
877                    proto_api::TaskPhase::Pending as i32,
878                    proto_api::TaskPhase::Running as i32,
879                ],
880                label_selector: "environment=production".into(),
881                resource_version: Some("test:1".into()),
882            }))
883            .await
884            .unwrap()
885            .into_inner();
886
887        let event = stream.next().await.unwrap().unwrap();
888        assert_eq!(event.r#type, proto_api::TaskWatchEventType::Added as i32);
889        assert_eq!(event.object.unwrap().metadata.unwrap().name, "watch-task");
890        assert!(stream.next().await.is_none());
891
892        let filter = handler
893            .last_watch_filter
894            .lock()
895            .unwrap()
896            .take()
897            .expect("handler received watch filter");
898        assert_eq!(filter.slot().unwrap().as_str(), "primary");
899        assert_eq!(filter.phases(), &[TaskPhase::Pending, TaskPhase::Running]);
900        let mut labels = solti_model::Labels::new();
901        labels.insert("environment", "production");
902        assert!(filter.matches_labels(&labels));
903        assert_eq!(
904            handler.last_watch_resource_version.lock().unwrap().take(),
905            Some(Some("test:1".into()))
906        );
907    }
908
909    #[tokio::test]
910    async fn watch_tasks_maps_initial_expiration_to_out_of_range() {
911        use std::sync::atomic::Ordering;
912
913        let (probe, service) = probed_service_with(StreamMock {
914            watch_expired: true,
915            ..StreamMock::default()
916        });
917
918        let status = service
919            .watch_tasks(Request::new(proto_api::WatchTasksRequest {
920                resource_version: Some("old:1".into()),
921                ..Default::default()
922            }))
923            .await
924            .err()
925            .expect("expired watch must fail");
926
927        assert_eq!(status.code(), tonic::Code::OutOfRange);
928        assert_eq!(probe.completed.load(Ordering::SeqCst), 1);
929        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 0);
930        assert_eq!(
931            probe.last_status.load(Ordering::SeqCst),
932            tonic::Code::OutOfRange as u16
933        );
934    }
935
936    #[tokio::test]
937    async fn watch_tasks_maps_stream_expiration_to_out_of_range() {
938        use std::sync::atomic::Ordering;
939
940        let (probe, service) = probed_service_with(StreamMock {
941            watch_stream_expired: true,
942            ..StreamMock::default()
943        });
944        let mut stream = service
945            .watch_tasks(Request::new(proto_api::WatchTasksRequest::default()))
946            .await
947            .unwrap()
948            .into_inner();
949
950        assert_eq!(probe.completed.load(Ordering::SeqCst), 0);
951        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 1);
952        assert!(stream.next().await.unwrap().is_ok());
953        assert_eq!(probe.completed.load(Ordering::SeqCst), 0);
954        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 1);
955        let status = stream.next().await.unwrap().unwrap_err();
956        assert_eq!(status.code(), tonic::Code::OutOfRange);
957        assert_eq!(probe.completed.load(Ordering::SeqCst), 1);
958        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 0);
959        assert_eq!(
960            probe.last_status.load(Ordering::SeqCst),
961            tonic::Code::OutOfRange as u16
962        );
963        assert!(stream.next().await.is_none());
964        assert_eq!(probe.completed.load(Ordering::SeqCst), 1);
965    }
966
967    #[tokio::test]
968    async fn watch_tasks_rejects_invalid_input_before_handler() {
969        for request in [
970            proto_api::WatchTasksRequest {
971                resource_version: Some(String::new()),
972                ..Default::default()
973            },
974            proto_api::WatchTasksRequest {
975                phases: vec![proto_api::TaskPhase::Unspecified as i32],
976                ..Default::default()
977            },
978            proto_api::WatchTasksRequest {
979                label_selector: "tier in (".into(),
980                ..Default::default()
981            },
982        ] {
983            let handler = Arc::new(StreamMock::default());
984            let service = TaskApiService::new(Arc::clone(&handler));
985            let status = service
986                .watch_tasks(Request::new(request))
987                .await
988                .err()
989                .expect("invalid watch must fail");
990            assert_eq!(status.code(), tonic::Code::InvalidArgument);
991            assert!(handler.last_watch_filter.lock().unwrap().is_none());
992        }
993    }
994
995    #[tokio::test]
996    async fn list_task_runs_exposes_historical_workload_gvk() {
997        let response = service()
998            .list_task_runs(Request::new(proto_api::ListTaskRunsRequest {
999                name: "extension-run".into(),
1000            }))
1001            .await
1002            .unwrap()
1003            .into_inner();
1004
1005        assert_eq!(response.runs.len(), 1);
1006        assert_eq!(
1007            response.runs[0].workload_api_version,
1008            "workloads.example.io/v1"
1009        );
1010        assert_eq!(response.runs[0].workload_kind, "DatabaseBackup");
1011    }
1012
1013    #[tokio::test]
1014    async fn list_task_runs_guards_embedded_history_from_custom_handler() {
1015        let status = service()
1016            .list_task_runs(Request::new(proto_api::ListTaskRunsRequest {
1017                name: "embedded-run".into(),
1018            }))
1019            .await
1020            .unwrap_err();
1021
1022        assert_eq!(status.code(), tonic::Code::Internal);
1023    }
1024
1025    #[tokio::test]
1026    async fn stream_task_logs_returns_three_proto_events_in_order() {
1027        let svc = service();
1028        let req = Request::new(proto_api::StreamTaskLogsRequest {
1029            name: "task-1".into(),
1030        });
1031
1032        let response = svc.stream_task_logs(req).await.expect("stream Ok");
1033        let mut stream = response.into_inner();
1034
1035        match stream.next().await.unwrap().unwrap().kind.unwrap() {
1036            proto_api::stream_task_logs_response::Kind::RunStarted(r) => {
1037                assert_eq!(r.generation, 2);
1038                assert_eq!(r.attempt, 1);
1039                assert_eq!(r.started_at, 1000);
1040            }
1041            other => panic!("expected RunStarted, got {other:?}"),
1042        }
1043
1044        match stream.next().await.unwrap().unwrap().kind.unwrap() {
1045            proto_api::stream_task_logs_response::Kind::Chunk(c) => {
1046                assert_eq!(c.generation, 2);
1047                assert_eq!(c.attempt, 1);
1048                assert_eq!(c.stream, proto_api::OutputStreamKind::Stdout as i32);
1049                assert_eq!(c.seq, 0);
1050                assert_eq!(&c.line[..], b"hello-grpc");
1051            }
1052            other => panic!("expected Chunk, got {other:?}"),
1053        }
1054
1055        match stream.next().await.unwrap().unwrap().kind.unwrap() {
1056            proto_api::stream_task_logs_response::Kind::RunFinished(r) => {
1057                assert_eq!(r.generation, 2);
1058                assert_eq!(r.attempt, 1);
1059                assert_eq!(r.exit_code, Some(0));
1060                assert_eq!(r.finished_at, 1500);
1061            }
1062            other => panic!("expected RunFinished, got {other:?}"),
1063        }
1064        assert!(stream.next().await.is_none(), "stream must terminate");
1065    }
1066
1067    #[tokio::test]
1068    async fn stream_task_logs_rejects_every_invalid_model_name() {
1069        let svc = service();
1070        for invalid in ["  ", "a/b", "a b", ".", "bad$name"] {
1071            let req = Request::new(proto_api::StreamTaskLogsRequest {
1072                name: invalid.into(),
1073            });
1074            let status = match svc.stream_task_logs(req).await {
1075                Err(s) => s,
1076                Ok(_) => panic!("expected error status for {invalid:?}"),
1077            };
1078            assert_eq!(status.code(), tonic::Code::InvalidArgument);
1079        }
1080    }
1081
1082    #[tokio::test]
1083    async fn stream_task_logs_maps_task_not_found_to_not_found_status() {
1084        let svc = service();
1085        let req = Request::new(proto_api::StreamTaskLogsRequest {
1086            name: "missing".into(),
1087        });
1088        let status = match svc.stream_task_logs(req).await {
1089            Err(s) => s,
1090            Ok(_) => panic!("expected error status"),
1091        };
1092        assert_eq!(status.code(), tonic::Code::NotFound);
1093    }
1094
1095    fn auth_interceptor(secret: &str) -> BearerAuth {
1096        BearerAuth {
1097            expected: Some(Token::new(secret).unwrap()),
1098            metrics: noop_api_metrics(),
1099        }
1100    }
1101
1102    fn request_with_authorization(value: &str) -> Request<()> {
1103        let mut req = Request::new(());
1104        req.metadata_mut()
1105            .insert("authorization", value.parse().expect("ascii metadata"));
1106        req
1107    }
1108
1109    #[test]
1110    fn bearer_auth_rejects_invalid_credentials() {
1111        let requests = [
1112            Request::new(()),
1113            request_with_authorization("Bearer not-the-secret"),
1114            request_with_authorization("sekret"),
1115            request_with_authorization("Basic sekret"),
1116        ];
1117
1118        for request in requests {
1119            let status = auth_interceptor("sekret").call(request).unwrap_err();
1120            assert_eq!(status.code(), tonic::Code::Unauthenticated);
1121        }
1122    }
1123
1124    #[test]
1125    fn bearer_auth_accepts_valid_token_scheme_case_insensitively() {
1126        for header in ["Bearer sekret", "bearer sekret", "BEARER sekret"] {
1127            let mut auth = auth_interceptor("sekret");
1128            assert!(
1129                auth.call(request_with_authorization(header)).is_ok(),
1130                "header {header:?} must pass"
1131            );
1132        }
1133    }
1134
1135    #[test]
1136    fn bearer_auth_passes_through_when_no_token_configured() {
1137        let mut auth = BearerAuth {
1138            expected: None,
1139            metrics: noop_api_metrics(),
1140        };
1141        assert!(auth.call(Request::new(())).is_ok());
1142        assert!(
1143            auth.call(request_with_authorization("Bearer anything"))
1144                .is_ok()
1145        );
1146    }
1147
1148    #[derive(Debug, Default)]
1149    struct GaugeProbe {
1150        in_flight: std::sync::atomic::AtomicI64,
1151        completed: std::sync::atomic::AtomicUsize,
1152        last_status: std::sync::atomic::AtomicU16,
1153    }
1154
1155    impl crate::metrics::ApiMetricsBackend for GaugeProbe {
1156        fn record_request(
1157            &self,
1158            _transport: crate::metrics::Transport,
1159            _method: &str,
1160            _path: &str,
1161            status: u16,
1162            _duration_ms: u64,
1163        ) {
1164            self.last_status
1165                .store(status, std::sync::atomic::Ordering::SeqCst);
1166            self.completed
1167                .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1168        }
1169
1170        fn record_in_flight_delta(&self, _transport: crate::metrics::Transport, delta: i64) {
1171            self.in_flight
1172                .fetch_add(delta, std::sync::atomic::Ordering::SeqCst);
1173        }
1174    }
1175
1176    fn probed_service() -> (Arc<GaugeProbe>, TaskApiService<StreamMock>) {
1177        probed_service_with(StreamMock::default())
1178    }
1179
1180    fn probed_service_with(handler: StreamMock) -> (Arc<GaugeProbe>, TaskApiService<StreamMock>) {
1181        let probe = Arc::new(GaugeProbe::default());
1182        let handle: ApiMetricsHandle = probe.clone();
1183        (
1184            probe,
1185            TaskApiService::new_with_metrics(Arc::new(handler), handle),
1186        )
1187    }
1188
1189    #[test]
1190    fn rejected_auth_is_recorded_and_balances_gauge() {
1191        use std::sync::atomic::Ordering;
1192
1193        let probe = Arc::new(GaugeProbe::default());
1194        let metrics: ApiMetricsHandle = probe.clone();
1195        let mut auth = BearerAuth {
1196            expected: Some(Token::new("secret").unwrap()),
1197            metrics,
1198        };
1199        let mut request = Request::new(());
1200        request
1201            .extensions_mut()
1202            .insert(tonic::GrpcMethod::new(GRPC_API_SERVICE, "GetTask"));
1203
1204        let status = auth.call(request).unwrap_err();
1205
1206        assert_eq!(status.code(), tonic::Code::Unauthenticated);
1207        assert_eq!(probe.completed.load(Ordering::SeqCst), 1);
1208        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 0);
1209    }
1210
1211    #[tokio::test]
1212    async fn instrument_records_completed_request_and_balances_gauge() {
1213        use std::sync::atomic::Ordering;
1214
1215        let (probe, svc) = probed_service();
1216        let result = svc
1217            .instrument("Probe", async { Ok(Response::new(())) })
1218            .await;
1219
1220        assert!(result.is_ok());
1221        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 0);
1222        assert_eq!(probe.completed.load(Ordering::SeqCst), 1);
1223    }
1224
1225    #[tokio::test]
1226    async fn stream_subscription_is_instrumented() {
1227        use std::sync::atomic::Ordering;
1228
1229        let (probe, service) = probed_service();
1230        let mut stream = service
1231            .stream_task_logs(Request::new(proto_api::StreamTaskLogsRequest {
1232                name: "task-a".into(),
1233            }))
1234            .await
1235            .unwrap()
1236            .into_inner();
1237
1238        assert_eq!(probe.completed.load(Ordering::SeqCst), 0);
1239        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 1);
1240
1241        while let Some(event) = stream.next().await {
1242            event.unwrap();
1243        }
1244
1245        assert_eq!(probe.completed.load(Ordering::SeqCst), 1);
1246        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 0);
1247        assert_eq!(
1248            probe.last_status.load(Ordering::SeqCst),
1249            tonic::Code::Ok as u16
1250        );
1251    }
1252
1253    #[tokio::test]
1254    async fn dropping_server_stream_releases_gauge_without_completion() {
1255        use std::sync::atomic::Ordering;
1256
1257        let (probe, service) = probed_service_with(StreamMock {
1258            log_stream_pending: true,
1259            ..StreamMock::default()
1260        });
1261        let response = service
1262            .stream_task_logs(Request::new(proto_api::StreamTaskLogsRequest {
1263                name: "task-a".into(),
1264            }))
1265            .await
1266            .unwrap();
1267
1268        assert_eq!(probe.completed.load(Ordering::SeqCst), 0);
1269        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 1);
1270
1271        drop(response);
1272
1273        assert_eq!(probe.completed.load(Ordering::SeqCst), 0);
1274        assert_eq!(probe.in_flight.load(Ordering::SeqCst), 0);
1275    }
1276
1277    #[test]
1278    fn in_flight_gauge_recovers_when_rpc_future_is_dropped() {
1279        use std::future::Future;
1280        use std::sync::atomic::Ordering;
1281        use std::task::{Context, Poll, Waker};
1282
1283        let (probe, svc) = probed_service();
1284
1285        let mut fut = Box::pin(svc.instrument(
1286            "Probe",
1287            std::future::pending::<Result<Response<()>, Status>>(),
1288        ));
1289
1290        let mut cx = Context::from_waker(Waker::noop());
1291        assert!(matches!(fut.as_mut().poll(&mut cx), Poll::Pending));
1292        assert_eq!(
1293            probe.in_flight.load(Ordering::SeqCst),
1294            1,
1295            "gauge must be armed after the first poll"
1296        );
1297
1298        drop(fut);
1299        assert_eq!(
1300            probe.in_flight.load(Ordering::SeqCst),
1301            0,
1302            "dropping the future must release the in-flight slot"
1303        );
1304        assert_eq!(
1305            probe.completed.load(Ordering::SeqCst),
1306            0,
1307            "a cancelled RPC must not be recorded as completed"
1308        );
1309    }
1310}