1use 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
26pub struct TaskApiService<H> {
33 handler: Arc<H>,
34 metrics: ApiMetricsHandle,
35}
36
37impl<H> TaskApiService<H>
38where
39 H: ApiHandler,
40{
41 pub fn new(handler: Arc<H>) -> Self {
43 Self::new_with_metrics(handler, noop_api_metrics())
44 }
45
46 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
71pub 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
93pub 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#[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
135fn 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
143pub 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
154pub 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 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}