1use 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
53pub 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
102pub struct TaskApiService<H> {
113 handler: Arc<H>,
114 metrics: ApiMetricsHandle,
115}
116
117impl<H> TaskApiService<H>
118where
119 H: ApiHandler,
120{
121 pub fn new(handler: Arc<H>) -> Self {
123 Self::new_with_metrics(handler, noop_api_metrics())
124 }
125
126 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
173pub type GrpcServer<H> = InterceptedService<TaskServiceServer<TaskApiService<H>>, BearerAuth>;
178
179pub 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 pub fn new(handler: Arc<H>) -> Self {
214 Self {
215 handler,
216 metrics: noop_api_metrics(),
217 auth: None,
218 }
219 }
220
221 pub fn with_auth(mut self, token: Token) -> Self {
228 self.auth = Some(token);
229 self
230 }
231
232 pub fn with_metrics(mut self, metrics: ApiMetricsHandle) -> Self {
236 self.metrics = metrics;
237 self
238 }
239
240 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#[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 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 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}