1use aion_proto::{
4 ProtoCancelRequest, ProtoCancelResponse, ProtoCountWorkflowsRequest,
5 ProtoCountWorkflowsResponse, ProtoCreateScheduleRequest, ProtoCreateScheduleResponse,
6 ProtoDeleteScheduleResponse, ProtoDescribeScheduleResponse, ProtoDescribeWorkflowRequest,
7 ProtoDescribeWorkflowResponse, ProtoListSchedulesRequest, ProtoListSchedulesResponse,
8 ProtoListWorkflowsRequest, ProtoListWorkflowsResponse, ProtoPauseScheduleResponse,
9 ProtoQueryRequest, ProtoQueryResponse, ProtoResumeScheduleResponse, ProtoScheduleIdRequest,
10 ProtoSignalRequest, ProtoSignalResponse, ProtoStartWorkflowRequest, ProtoStartWorkflowResponse,
11 ProtoUpdateScheduleRequest, ProtoUpdateScheduleResponse, ProtoWireError, WireError,
12 generated::{self, workflow_service_server::WorkflowServiceServer},
13};
14use prost::Message;
15use tonic::{Code, Request, Response, Status};
16
17use crate::{CallerIdentity, ServerState, api::handlers, api::schedule_handlers};
18
19#[derive(Clone)]
21pub struct WorkflowGrpcService {
22 state: ServerState,
23}
24
25impl WorkflowGrpcService {
26 #[must_use]
28 pub const fn new(state: ServerState) -> Self {
29 Self { state }
30 }
31
32 async fn caller<T>(&self, request: &Request<T>) -> Result<CallerIdentity, Status> {
33 caller_from_metadata(request.metadata(), &self.state).await
34 }
35
36 #[cfg(feature = "haematite-backend")]
53 async fn resolve_route(
54 &self,
55 workflow_id: Option<aion_proto::ProtoWorkflowId>,
56 metadata: &tonic::metadata::MetadataMap,
57 request: crate::routing::ForwardRequest,
58 ) -> RouteResolution {
59 use crate::routing::{RouteDecision, route_mutation};
60 let Some(cluster_store) = self.state.cluster_store() else {
61 return RouteResolution::Local;
62 };
63 let Some(proto) = workflow_id else {
64 return RouteResolution::Local;
65 };
66 let Ok(workflow_id) = aion_core::WorkflowId::try_from(proto) else {
67 return RouteResolution::Local;
68 };
69 let directory = self
70 .state
71 .shard_directory()
72 .map(|directory| directory.as_ref() as &dyn crate::routing::ShardDirectory);
73 match route_mutation(Some(cluster_store.as_ref()), directory, &workflow_id) {
74 RouteDecision::Local => RouteResolution::Local,
75 RouteDecision::NotOwner { shard } => RouteResolution::Reject(not_owner_status(shard)),
76 RouteDecision::Forward { owner, shard } => {
77 self.forward_or_reject(owner, shard, metadata, request)
78 .await
79 }
80 }
81 }
82
83 #[cfg(feature = "haematite-backend")]
88 async fn forward_or_reject(
89 &self,
90 owner: crate::routing::NodeRef,
91 shard: usize,
92 metadata: &tonic::metadata::MetadataMap,
93 request: crate::routing::ForwardRequest,
94 ) -> RouteResolution {
95 use crate::routing::{MAX_FORWARD_HOPS, current_hops};
96 if current_hops(metadata) >= MAX_FORWARD_HOPS {
100 return RouteResolution::Reject(not_owner_status(shard));
101 }
102 let Some(target) = owner.grpc_addr else {
104 return RouteResolution::Reject(not_owner_status(shard));
105 };
106 let Some(forwarder) = self.state.request_forwarder() else {
107 return RouteResolution::Reject(not_owner_status(shard));
108 };
109 match forwarder.forward(target, metadata.clone(), request).await {
110 Ok(reply) => RouteResolution::Reply(reply),
111 Err(_status) => RouteResolution::Reject(not_owner_status(shard)),
115 }
116 }
117
118 #[cfg(feature = "haematite-backend")]
122 fn start_placement(&self) -> Option<aion_core::WorkflowId> {
123 use crate::routing::{RemintOutcome, route_start};
124 match route_start(self.state.cluster_store().map(AsRef::as_ref)) {
125 RemintOutcome::UseId(workflow_id) => Some(workflow_id),
126 RemintOutcome::EngineMint => None,
127 }
128 }
129
130 #[cfg(feature = "haematite-backend")]
139 async fn resolve_start(
140 &self,
141 request: &generated::StartWorkflowRequest,
142 metadata: &tonic::metadata::MetadataMap,
143 ) -> StartResolution {
144 use crate::routing::{SteerDecision, route_start_steered};
145 let routing_key = request.routing_key.as_deref().filter(|key| !key.is_empty());
146 let Some(routing_key) = routing_key else {
147 return StartResolution::Local(self.start_placement());
149 };
150 let Some(cluster_store) = self.state.cluster_store() else {
151 return StartResolution::Local(None);
154 };
155 let directory = self
156 .state
157 .shard_directory()
158 .map(|directory| directory.as_ref() as &dyn crate::routing::ShardDirectory);
159 match route_start_steered(cluster_store.as_ref(), directory, routing_key) {
160 SteerDecision::Local(workflow_id) => StartResolution::Local(Some(workflow_id)),
161 SteerDecision::NotOwner { shard } => StartResolution::Reject(not_owner_status(shard)),
162 SteerDecision::Forward { owner, shard } => {
163 self.forward_or_reject_start(owner, shard, metadata, request.clone())
164 .await
165 }
166 }
167 }
168
169 #[cfg(feature = "haematite-backend")]
174 async fn forward_or_reject_start(
175 &self,
176 owner: crate::routing::NodeRef,
177 shard: usize,
178 metadata: &tonic::metadata::MetadataMap,
179 request: generated::StartWorkflowRequest,
180 ) -> StartResolution {
181 use crate::routing::{ForwardReply, ForwardRequest, MAX_FORWARD_HOPS, current_hops};
182 if current_hops(metadata) >= MAX_FORWARD_HOPS {
183 return StartResolution::Reject(not_owner_status(shard));
184 }
185 let Some(target) = owner.grpc_addr else {
186 return StartResolution::Reject(not_owner_status(shard));
187 };
188 let Some(forwarder) = self.state.request_forwarder() else {
189 return StartResolution::Reject(not_owner_status(shard));
190 };
191 match forwarder
192 .forward(target, metadata.clone(), ForwardRequest::Start(request))
193 .await
194 {
195 Ok(ForwardReply::Start(reply)) => StartResolution::Reply(reply),
196 Ok(_) => {
197 StartResolution::Reject(Status::internal("forwarder returned a mismatched reply"))
198 }
199 Err(_status) => StartResolution::Reject(not_owner_status(shard)),
201 }
202 }
203}
204
205#[cfg(feature = "haematite-backend")]
207enum StartResolution {
208 Local(Option<aion_core::WorkflowId>),
210 Reply(generated::StartWorkflowResponse),
212 Reject(Status),
214}
215
216#[cfg(feature = "haematite-backend")]
218enum RouteResolution {
219 Local,
221 Reply(crate::routing::ForwardReply),
223 Reject(Status),
225}
226
227#[cfg(feature = "haematite-backend")]
229fn not_owner_status(shard: usize) -> Status {
230 let wire = WireError::not_owner(format!(
231 "workflow shard {shard} is owned by another cluster node"
232 ))
233 .with_error_type("NotOwner");
234 status_from_wire_error(wire)
235}
236
237#[must_use]
239pub fn workflow_service(state: ServerState) -> WorkflowServiceServer<WorkflowGrpcService> {
240 WorkflowServiceServer::new(WorkflowGrpcService::new(state))
241}
242
243#[tonic::async_trait]
244impl generated::workflow_service_server::WorkflowService for WorkflowGrpcService {
245 async fn start_workflow(
246 &self,
247 request: Request<generated::StartWorkflowRequest>,
248 ) -> Result<Response<generated::StartWorkflowResponse>, Status> {
249 if self.state.drain_state().is_draining() {
250 return Err(Status::unavailable(
251 "server is draining and not accepting new workflow starts",
252 ));
253 }
254 let caller = self.caller(&request).await?;
255 #[cfg(feature = "haematite-backend")]
262 {
263 let (metadata, _ext, inner) = request.into_parts();
264 let placement = match self.resolve_start(&inner, &metadata).await {
265 StartResolution::Reject(status) => return Err(status),
266 StartResolution::Reply(reply) => return Ok(Response::new(reply)),
267 StartResolution::Local(placement) => placement,
268 };
269 let response = handlers::start_with_placement(
270 self.state.namespace_guard(),
271 &caller,
272 decode_start_request(inner),
273 placement,
274 )
275 .await
276 .map_err(status_from_wire_error)?;
277 return Ok(Response::new(encode_start_response(response)));
278 }
279 #[cfg(not(feature = "haematite-backend"))]
280 {
281 let placement: Option<aion_core::WorkflowId> = None;
282 let response = handlers::start_with_placement(
283 self.state.namespace_guard(),
284 &caller,
285 decode_start_request(request.into_inner()),
286 placement,
287 )
288 .await
289 .map_err(status_from_wire_error)?;
290 Ok(Response::new(encode_start_response(response)))
291 }
292 }
293
294 async fn signal(
295 &self,
296 request: Request<generated::SignalRequest>,
297 ) -> Result<Response<generated::SignalResponse>, Status> {
298 let caller = self.caller(&request).await?;
299 #[cfg(feature = "haematite-backend")]
300 {
301 use crate::routing::{ForwardReply, ForwardRequest};
302 let (metadata, _ext, inner) = request.into_parts();
303 let workflow_id = inner.workflow_id.clone().map(decode_workflow_id);
304 match self
305 .resolve_route(
306 workflow_id,
307 &metadata,
308 ForwardRequest::Signal(inner.clone()),
309 )
310 .await
311 {
312 RouteResolution::Reject(status) => return Err(status),
313 RouteResolution::Reply(ForwardReply::Signal(reply)) => {
314 return Ok(Response::new(reply));
315 }
316 RouteResolution::Reply(_) => {
317 return Err(Status::internal("forwarder returned a mismatched reply"));
318 }
319 RouteResolution::Local => {
320 let response = handlers::signal(
321 self.state.namespace_guard(),
322 &caller,
323 decode_signal_request(inner),
324 )
325 .await
326 .map_err(status_from_wire_error)?;
327 return Ok(Response::new(encode_signal_response(response)));
328 }
329 }
330 }
331 #[cfg(not(feature = "haematite-backend"))]
332 {
333 let response = handlers::signal(
334 self.state.namespace_guard(),
335 &caller,
336 decode_signal_request(request.into_inner()),
337 )
338 .await
339 .map_err(status_from_wire_error)?;
340 Ok(Response::new(encode_signal_response(response)))
341 }
342 }
343
344 async fn query(
345 &self,
346 request: Request<generated::QueryRequest>,
347 ) -> Result<Response<generated::QueryResponse>, Status> {
348 let caller = self.caller(&request).await?;
349 #[cfg(feature = "haematite-backend")]
350 {
351 use crate::routing::{ForwardReply, ForwardRequest};
352 let (metadata, _ext, inner) = request.into_parts();
353 let workflow_id = inner.workflow_id.clone().map(decode_workflow_id);
354 match self
355 .resolve_route(workflow_id, &metadata, ForwardRequest::Query(inner.clone()))
356 .await
357 {
358 RouteResolution::Reject(status) => return Err(status),
359 RouteResolution::Reply(ForwardReply::Query(reply)) => {
360 return Ok(Response::new(reply));
361 }
362 RouteResolution::Reply(_) => {
363 return Err(Status::internal("forwarder returned a mismatched reply"));
364 }
365 RouteResolution::Local => {
366 let response = handlers::query(
367 self.state.namespace_guard(),
368 &caller,
369 decode_query_request(inner),
370 )
371 .await
372 .map_err(status_from_wire_error)?;
373 return Ok(Response::new(encode_query_response(response)));
374 }
375 }
376 }
377 #[cfg(not(feature = "haematite-backend"))]
378 {
379 let response = handlers::query(
380 self.state.namespace_guard(),
381 &caller,
382 decode_query_request(request.into_inner()),
383 )
384 .await
385 .map_err(status_from_wire_error)?;
386 Ok(Response::new(encode_query_response(response)))
387 }
388 }
389
390 async fn cancel(
391 &self,
392 request: Request<generated::CancelRequest>,
393 ) -> Result<Response<generated::CancelResponse>, Status> {
394 let caller = self.caller(&request).await?;
395 #[cfg(feature = "haematite-backend")]
396 {
397 use crate::routing::{ForwardReply, ForwardRequest};
398 let (metadata, _ext, inner) = request.into_parts();
399 let workflow_id = inner.workflow_id.clone().map(decode_workflow_id);
400 match self
401 .resolve_route(
402 workflow_id,
403 &metadata,
404 ForwardRequest::Cancel(inner.clone()),
405 )
406 .await
407 {
408 RouteResolution::Reject(status) => return Err(status),
409 RouteResolution::Reply(ForwardReply::Cancel(reply)) => {
410 return Ok(Response::new(reply));
411 }
412 RouteResolution::Reply(_) => {
413 return Err(Status::internal("forwarder returned a mismatched reply"));
414 }
415 RouteResolution::Local => {
416 let response = handlers::cancel(
417 self.state.namespace_guard(),
418 &caller,
419 decode_cancel_request(inner),
420 )
421 .await
422 .map_err(status_from_wire_error)?;
423 return Ok(Response::new(encode_cancel_response(response)));
424 }
425 }
426 }
427 #[cfg(not(feature = "haematite-backend"))]
428 {
429 let response = handlers::cancel(
430 self.state.namespace_guard(),
431 &caller,
432 decode_cancel_request(request.into_inner()),
433 )
434 .await
435 .map_err(status_from_wire_error)?;
436 Ok(Response::new(encode_cancel_response(response)))
437 }
438 }
439
440 async fn list_workflows(
441 &self,
442 request: Request<generated::ListWorkflowsRequest>,
443 ) -> Result<Response<generated::ListWorkflowsResponse>, Status> {
444 let caller = self.caller(&request).await?;
445 let response = handlers::list(
446 self.state.namespace_guard(),
447 &caller,
448 decode_list_request(request.into_inner()),
449 )
450 .await
451 .map_err(status_from_wire_error)?;
452 Ok(Response::new(encode_list_response(response)))
453 }
454
455 async fn count_workflows(
456 &self,
457 request: Request<generated::CountWorkflowsRequest>,
458 ) -> Result<Response<generated::CountWorkflowsResponse>, Status> {
459 let caller = self.caller(&request).await?;
460 let response = handlers::count(
461 self.state.namespace_guard(),
462 &caller,
463 decode_count_request(request.into_inner()),
464 )
465 .await
466 .map_err(status_from_wire_error)?;
467 Ok(Response::new(encode_count_response(response)))
468 }
469
470 async fn describe_workflow(
471 &self,
472 request: Request<generated::DescribeWorkflowRequest>,
473 ) -> Result<Response<generated::DescribeWorkflowResponse>, Status> {
474 let caller = self.caller(&request).await?;
475 let response = handlers::describe(
476 self.state.namespace_guard(),
477 &caller,
478 decode_describe_request(request.into_inner()),
479 )
480 .await
481 .map_err(status_from_wire_error)?;
482 Ok(Response::new(encode_describe_response(response)))
483 }
484
485 async fn create_schedule(
486 &self,
487 request: Request<generated::CreateScheduleRequest>,
488 ) -> Result<Response<generated::CreateScheduleResponse>, Status> {
489 let caller = self.caller(&request).await?;
490 let response = schedule_handlers::create_schedule(
491 self.state.namespace_guard(),
492 &caller,
493 decode_create_schedule_request(request.into_inner()),
494 )
495 .await
496 .map_err(status_from_wire_error)?;
497 Ok(Response::new(encode_create_schedule_response(response)))
498 }
499
500 async fn update_schedule(
501 &self,
502 request: Request<generated::UpdateScheduleRequest>,
503 ) -> Result<Response<generated::UpdateScheduleResponse>, Status> {
504 let caller = self.caller(&request).await?;
505 let response = schedule_handlers::update_schedule(
506 self.state.namespace_guard(),
507 &caller,
508 decode_update_schedule_request(request.into_inner()),
509 )
510 .await
511 .map_err(status_from_wire_error)?;
512 Ok(Response::new(encode_update_schedule_response(response)))
513 }
514
515 async fn pause_schedule(
516 &self,
517 request: Request<generated::ScheduleIdRequest>,
518 ) -> Result<Response<generated::PauseScheduleResponse>, Status> {
519 let caller = self.caller(&request).await?;
520 let response = schedule_handlers::pause_schedule(
521 self.state.namespace_guard(),
522 &caller,
523 decode_schedule_id_request(request.into_inner()),
524 )
525 .await
526 .map_err(status_from_wire_error)?;
527 Ok(Response::new(encode_pause_schedule_response(response)))
528 }
529
530 async fn resume_schedule(
531 &self,
532 request: Request<generated::ScheduleIdRequest>,
533 ) -> Result<Response<generated::ResumeScheduleResponse>, Status> {
534 let caller = self.caller(&request).await?;
535 let response = schedule_handlers::resume_schedule(
536 self.state.namespace_guard(),
537 &caller,
538 decode_schedule_id_request(request.into_inner()),
539 )
540 .await
541 .map_err(status_from_wire_error)?;
542 Ok(Response::new(encode_resume_schedule_response(response)))
543 }
544
545 async fn delete_schedule(
546 &self,
547 request: Request<generated::ScheduleIdRequest>,
548 ) -> Result<Response<generated::DeleteScheduleResponse>, Status> {
549 let caller = self.caller(&request).await?;
550 let response = schedule_handlers::delete_schedule(
551 self.state.namespace_guard(),
552 &caller,
553 decode_schedule_id_request(request.into_inner()),
554 )
555 .await
556 .map_err(status_from_wire_error)?;
557 Ok(Response::new(encode_delete_schedule_response(response)))
558 }
559
560 async fn list_schedules(
561 &self,
562 request: Request<generated::ListSchedulesRequest>,
563 ) -> Result<Response<generated::ListSchedulesResponse>, Status> {
564 let caller = self.caller(&request).await?;
565 let response = schedule_handlers::list_schedules(
566 self.state.namespace_guard(),
567 &caller,
568 decode_list_schedules_request(request.into_inner()),
569 )
570 .await
571 .map_err(status_from_wire_error)?;
572 Ok(Response::new(encode_list_schedules_response(response)))
573 }
574
575 async fn describe_schedule(
576 &self,
577 request: Request<generated::ScheduleIdRequest>,
578 ) -> Result<Response<generated::DescribeScheduleResponse>, Status> {
579 let caller = self.caller(&request).await?;
580 let response = schedule_handlers::describe_schedule(
581 self.state.namespace_guard(),
582 &caller,
583 decode_schedule_id_request(request.into_inner()),
584 )
585 .await
586 .map_err(status_from_wire_error)?;
587 Ok(Response::new(encode_describe_schedule_response(response)))
588 }
589}
590
591pub(crate) async fn caller_from_metadata(
592 metadata: &tonic::metadata::MetadataMap,
593 state: &ServerState,
594) -> Result<CallerIdentity, Status> {
595 if !state.runtime_config().auth.enabled {
596 return Ok(development_caller_from_metadata(metadata));
597 }
598 #[cfg(feature = "auth")]
599 {
600 let bearer = metadata
601 .get("authorization")
602 .and_then(|value| value.to_str().ok())
603 .and_then(parse_bearer)
604 .ok_or_else(|| Status::unauthenticated("missing bearer token"))?;
605 let Some(cache) = state.jwks_cache() else {
606 return Err(Status::unauthenticated("invalid bearer token"));
607 };
608 return cache
609 .validate(&bearer)
610 .await
611 .map(|claims| claims.caller_identity())
612 .map_err(|_error| Status::unauthenticated("invalid bearer token"));
613 }
614 #[cfg(not(feature = "auth"))]
615 {
616 tokio::task::yield_now().await;
618 Ok(development_token_caller_from_metadata(
619 metadata,
620 &state.runtime_config().auth,
621 ))
622 }
623}
624
625fn development_caller_from_metadata(metadata: &tonic::metadata::MetadataMap) -> CallerIdentity {
626 let subject = metadata
627 .get("x-aion-subject")
628 .and_then(|value| value.to_str().ok())
629 .filter(|value| !value.is_empty())
630 .unwrap_or("anonymous");
631 let namespaces = metadata
632 .get("x-aion-namespaces")
633 .and_then(|value| value.to_str().ok())
634 .map(parse_namespaces)
635 .unwrap_or_default();
636 CallerIdentity::new(subject, namespaces).with_deploy(deploy_metadata_granted(metadata))
637}
638
639fn deploy_metadata_granted(metadata: &tonic::metadata::MetadataMap) -> bool {
642 metadata
643 .get("x-aion-deploy")
644 .and_then(|value| value.to_str().ok())
645 .is_some_and(|value| value.trim().eq_ignore_ascii_case("true"))
646}
647
648#[cfg(not(feature = "auth"))]
652fn development_token_caller_from_metadata(
653 metadata: &tonic::metadata::MetadataMap,
654 auth: &crate::config::AuthConfig,
655) -> CallerIdentity {
656 let subject = metadata
657 .get("x-aion-subject")
658 .and_then(|value| value.to_str().ok())
659 .filter(|value| !value.is_empty());
660 let namespaces = metadata
661 .get("x-aion-namespaces")
662 .and_then(|value| value.to_str().ok())
663 .map(parse_namespaces)
664 .unwrap_or_default();
665
666 let bearer_token = auth.jwks_url.as_deref().unwrap_or_default();
667 let expected = format!("Bearer {bearer_token}");
668 let Some(authorization) = metadata.get("authorization") else {
669 return CallerIdentity::denied(subject.unwrap_or("anonymous"), "missing bearer token");
670 };
671 let authorization = authorization.to_str().ok();
672 if authorization != Some(expected.as_str()) {
673 return CallerIdentity::denied(subject.unwrap_or("anonymous"), "invalid bearer token");
674 }
675
676 let Some(subject) = subject else {
677 return CallerIdentity::denied("anonymous", "missing required metadata: x-aion-subject");
678 };
679
680 CallerIdentity::new(subject, namespaces).with_deploy(deploy_metadata_granted(metadata))
681}
682
683#[cfg(feature = "auth")]
684fn parse_bearer(value: &str) -> Option<String> {
685 let token = value.strip_prefix("Bearer ")?.trim();
686 if token.is_empty() {
687 return None;
688 }
689 Some(token.to_owned())
690}
691
692fn parse_namespaces(value: &str) -> Vec<String> {
693 value
694 .split(',')
695 .map(str::trim)
696 .filter(|namespace| !namespace.is_empty())
697 .map(str::to_owned)
698 .collect()
699}
700
701pub(crate) fn status_from_wire_error(error: WireError) -> Status {
702 status_with_code(grpc_code(error.code), error)
703}
704
705pub(crate) fn status_with_code(code: Code, error: WireError) -> Status {
708 let message = error.message.clone();
709 let mut details = Vec::new();
710 let proto_error = ProtoWireError::from(error);
711 if proto_error.encode(&mut details).is_ok() {
712 Status::with_details(code, message, details.into())
713 } else {
714 Status::new(code, message)
715 }
716}
717
718fn grpc_code(code: aion_proto::WireErrorCode) -> Code {
719 match code {
720 aion_proto::WireErrorCode::NotFound => Code::NotFound,
721 aion_proto::WireErrorCode::NamespaceDenied | aion_proto::WireErrorCode::DeployDenied => {
722 Code::PermissionDenied
723 }
724 aion_proto::WireErrorCode::SequenceConflict | aion_proto::WireErrorCode::NotOwner => {
728 Code::Aborted
729 }
730 aion_proto::WireErrorCode::UnknownQuery | aion_proto::WireErrorCode::InvalidInput => {
731 Code::InvalidArgument
732 }
733 aion_proto::WireErrorCode::QueryTimeout => Code::DeadlineExceeded,
734 aion_proto::WireErrorCode::NotRunning | aion_proto::WireErrorCode::VersionPinned => {
735 Code::FailedPrecondition
736 }
737 aion_proto::WireErrorCode::Lagged => Code::ResourceExhausted,
738 aion_proto::WireErrorCode::Backend | aion_proto::WireErrorCode::QueryFailed => {
742 Code::Internal
743 }
744 }
745}
746
747fn decode_workflow_id(value: generated::WorkflowId) -> aion_proto::ProtoWorkflowId {
748 aion_proto::ProtoWorkflowId { uuid: value.uuid }
749}
750
751fn encode_workflow_id(value: aion_proto::ProtoWorkflowId) -> generated::WorkflowId {
752 generated::WorkflowId { uuid: value.uuid }
753}
754
755fn decode_run_id(value: generated::RunId) -> aion_proto::ProtoRunId {
756 aion_proto::ProtoRunId { uuid: value.uuid }
757}
758
759fn encode_run_id(value: aion_proto::ProtoRunId) -> generated::RunId {
760 generated::RunId { uuid: value.uuid }
761}
762
763fn decode_schedule_id(value: generated::ScheduleId) -> aion_proto::ProtoScheduleId {
764 aion_proto::ProtoScheduleId { uuid: value.uuid }
765}
766
767fn encode_schedule_id(value: aion_proto::ProtoScheduleId) -> generated::ScheduleId {
768 generated::ScheduleId { uuid: value.uuid }
769}
770
771fn decode_payload(value: generated::Payload) -> aion_proto::ProtoPayload {
772 aion_proto::ProtoPayload {
773 content_type: value.content_type,
774 bytes: value.bytes,
775 }
776}
777
778fn encode_payload(value: aion_proto::ProtoPayload) -> generated::Payload {
779 generated::Payload {
780 content_type: value.content_type,
781 bytes: value.bytes,
782 }
783}
784
785fn decode_envelope(value: generated::WireEnvelope) -> aion_proto::WireEnvelope {
786 aion_proto::WireEnvelope {
787 namespace: value.namespace,
788 request_id: value.request_id,
789 payload: value.payload.map(decode_payload),
790 }
791}
792
793fn encode_envelope(value: aion_proto::WireEnvelope) -> generated::WireEnvelope {
794 generated::WireEnvelope {
795 namespace: value.namespace,
796 request_id: value.request_id,
797 payload: value.payload.map(encode_payload),
798 }
799}
800
801fn decode_start_request(value: generated::StartWorkflowRequest) -> ProtoStartWorkflowRequest {
802 ProtoStartWorkflowRequest {
803 namespace: value.namespace,
804 workflow_type: value.workflow_type,
805 input: value.input.map(decode_payload),
806 routing_key: value.routing_key,
807 }
808}
809
810fn encode_start_response(value: ProtoStartWorkflowResponse) -> generated::StartWorkflowResponse {
811 generated::StartWorkflowResponse {
812 workflow_id: value.workflow_id.map(encode_workflow_id),
813 run_id: value.run_id.map(encode_run_id),
814 }
815}
816
817fn decode_signal_request(value: generated::SignalRequest) -> ProtoSignalRequest {
818 ProtoSignalRequest {
819 namespace: value.namespace,
820 workflow_id: value.workflow_id.map(decode_workflow_id),
821 run_id: value.run_id.map(decode_run_id),
822 signal_name: value.signal_name,
823 payload: value.payload.map(decode_payload),
824 }
825}
826
827fn encode_signal_response(_: ProtoSignalResponse) -> generated::SignalResponse {
828 generated::SignalResponse {}
829}
830
831fn decode_query_request(value: generated::QueryRequest) -> ProtoQueryRequest {
832 ProtoQueryRequest {
833 namespace: value.namespace,
834 workflow_id: value.workflow_id.map(decode_workflow_id),
835 run_id: value.run_id.map(decode_run_id),
836 query_name: value.query_name,
837 }
838}
839
840fn encode_query_response(value: ProtoQueryResponse) -> generated::QueryResponse {
841 generated::QueryResponse {
842 outcome: value.outcome.map(encode_query_outcome),
843 }
844}
845
846fn encode_query_outcome(
847 value: aion_proto::proto_query_response::Outcome,
848) -> generated::query_response::Outcome {
849 match value {
850 aion_proto::proto_query_response::Outcome::Result(payload) => {
851 generated::query_response::Outcome::Result(encode_payload(payload))
852 }
853 aion_proto::proto_query_response::Outcome::Error(error) => {
854 generated::query_response::Outcome::Error(encode_wire_error(error))
855 }
856 }
857}
858
859fn encode_wire_error(value: ProtoWireError) -> generated::WireError {
860 generated::WireError {
861 code: value.code,
862 message: value.message,
863 error_type: value.error_type,
864 }
865}
866
867fn decode_cancel_request(value: generated::CancelRequest) -> ProtoCancelRequest {
868 ProtoCancelRequest {
869 namespace: value.namespace,
870 workflow_id: value.workflow_id.map(decode_workflow_id),
871 run_id: value.run_id.map(decode_run_id),
872 reason: value.reason,
873 }
874}
875
876fn encode_cancel_response(_: ProtoCancelResponse) -> generated::CancelResponse {
877 generated::CancelResponse {}
878}
879
880fn decode_list_request(value: generated::ListWorkflowsRequest) -> ProtoListWorkflowsRequest {
881 ProtoListWorkflowsRequest {
882 namespace: value.namespace,
883 filter: value.filter.map(decode_envelope),
884 }
885}
886
887fn encode_list_response(value: ProtoListWorkflowsResponse) -> generated::ListWorkflowsResponse {
888 generated::ListWorkflowsResponse {
889 summaries: value.summaries.into_iter().map(encode_envelope).collect(),
890 }
891}
892
893fn decode_count_request(value: generated::CountWorkflowsRequest) -> ProtoCountWorkflowsRequest {
894 ProtoCountWorkflowsRequest {
895 namespace: value.namespace,
896 filter: value.filter.map(decode_envelope),
897 }
898}
899
900fn encode_count_response(value: ProtoCountWorkflowsResponse) -> generated::CountWorkflowsResponse {
901 generated::CountWorkflowsResponse { count: value.count }
902}
903
904fn decode_describe_request(
905 value: generated::DescribeWorkflowRequest,
906) -> ProtoDescribeWorkflowRequest {
907 ProtoDescribeWorkflowRequest {
908 namespace: value.namespace,
909 workflow_id: value.workflow_id.map(decode_workflow_id),
910 run_id: value.run_id.map(decode_run_id),
911 include_history: value.include_history,
912 }
913}
914
915fn encode_describe_response(
916 value: ProtoDescribeWorkflowResponse,
917) -> generated::DescribeWorkflowResponse {
918 generated::DescribeWorkflowResponse {
919 summary: value.summary.map(encode_envelope),
920 history: value.history.into_iter().map(encode_envelope).collect(),
921 }
922}
923
924fn decode_create_schedule_request(
925 value: generated::CreateScheduleRequest,
926) -> ProtoCreateScheduleRequest {
927 ProtoCreateScheduleRequest {
928 namespace: value.namespace,
929 config: value.config.map(decode_envelope),
930 }
931}
932
933fn encode_create_schedule_response(
934 value: ProtoCreateScheduleResponse,
935) -> generated::CreateScheduleResponse {
936 generated::CreateScheduleResponse {
937 schedule_id: value.schedule_id.map(encode_schedule_id),
938 state: value.state.map(encode_envelope),
939 }
940}
941
942fn decode_update_schedule_request(
943 value: generated::UpdateScheduleRequest,
944) -> ProtoUpdateScheduleRequest {
945 ProtoUpdateScheduleRequest {
946 namespace: value.namespace,
947 schedule_id: value.schedule_id.map(decode_schedule_id),
948 config: value.config.map(decode_envelope),
949 }
950}
951
952fn encode_update_schedule_response(
953 value: ProtoUpdateScheduleResponse,
954) -> generated::UpdateScheduleResponse {
955 generated::UpdateScheduleResponse {
956 state: value.state.map(encode_envelope),
957 }
958}
959
960fn decode_schedule_id_request(value: generated::ScheduleIdRequest) -> ProtoScheduleIdRequest {
961 ProtoScheduleIdRequest {
962 namespace: value.namespace,
963 schedule_id: value.schedule_id.map(decode_schedule_id),
964 }
965}
966
967fn encode_pause_schedule_response(
968 value: ProtoPauseScheduleResponse,
969) -> generated::PauseScheduleResponse {
970 generated::PauseScheduleResponse {
971 state: value.state.map(encode_envelope),
972 }
973}
974
975fn encode_resume_schedule_response(
976 value: ProtoResumeScheduleResponse,
977) -> generated::ResumeScheduleResponse {
978 generated::ResumeScheduleResponse {
979 state: value.state.map(encode_envelope),
980 }
981}
982
983fn encode_delete_schedule_response(
984 _: ProtoDeleteScheduleResponse,
985) -> generated::DeleteScheduleResponse {
986 generated::DeleteScheduleResponse {}
987}
988
989fn decode_list_schedules_request(
990 value: generated::ListSchedulesRequest,
991) -> ProtoListSchedulesRequest {
992 ProtoListSchedulesRequest {
993 namespace: value.namespace,
994 }
995}
996
997fn encode_list_schedules_response(
998 value: ProtoListSchedulesResponse,
999) -> generated::ListSchedulesResponse {
1000 generated::ListSchedulesResponse {
1001 schedules: value.schedules.into_iter().map(encode_envelope).collect(),
1002 }
1003}
1004
1005fn encode_describe_schedule_response(
1006 value: ProtoDescribeScheduleResponse,
1007) -> generated::DescribeScheduleResponse {
1008 generated::DescribeScheduleResponse {
1009 state: value.state.map(encode_envelope),
1010 }
1011}
1012
1013#[cfg(test)]
1014mod tests {
1015 use std::{net::SocketAddr, sync::Arc};
1016
1017 use aion::EngineBuilder;
1018 use aion_core::{Event, EventEnvelope, Payload, WorkflowId, WorkflowStatus};
1019 use aion_proto::{
1020 WireErrorCode,
1021 convert::{decode_core_value, encode_core_value},
1022 generated::workflow_service_server::WorkflowService,
1023 };
1024 use aion_store::{
1025 EventStore, InMemoryStore, WriteToken,
1026 visibility::{VisibilityRecord, VisibilityStore},
1027 };
1028 use chrono::Utc;
1029 use serde_json::json;
1030 use tonic::Request;
1031
1032 use super::*;
1033 use crate::{
1034 NamespaceResolver,
1035 config::{
1036 AuthConfig, AuthoringConfig, DashboardAssetSource, DashboardConfig, DeployConfig,
1037 ListenConfig, MetricsConfig, NamespaceConfig, NamespaceMode, RuntimeConfig,
1038 WebSocketConfig, WorkerConfig,
1039 },
1040 };
1041
1042 const NAMESPACE: &str = "tenant-a";
1043 const TOKEN: &str = "test-token";
1044
1045 async fn server_state(
1050 resolver: NamespaceResolver,
1051 runtime: RuntimeConfig,
1052 ) -> Result<ServerState, Box<dyn std::error::Error>> {
1053 #[cfg(feature = "auth")]
1054 {
1055 let url = crate::auth::test_support::serve_jwks()?;
1056 let refresh = std::time::Duration::from_secs(runtime.auth.jwks_refresh_seconds);
1057 let cache = crate::auth::JwksCache::new(url, refresh).await?;
1058 Ok(ServerState::from_parts_with_jwks(resolver, runtime, cache))
1059 }
1060 #[cfg(not(feature = "auth"))]
1061 {
1062 tokio::task::yield_now().await;
1064 Ok(ServerState::from_parts(resolver, runtime))
1065 }
1066 }
1067
1068 #[test]
1072 fn not_owner_wire_code_maps_to_retryable_aborted() {
1073 assert_eq!(grpc_code(WireErrorCode::NotOwner), Code::Aborted);
1074 }
1075
1076 #[tokio::test]
1077 async fn in_process_tonic_start_and_list_use_shared_handlers()
1078 -> Result<(), Box<dyn std::error::Error>> {
1079 let backing = Arc::new(InMemoryStore::default());
1080 let store: Arc<dyn EventStore> = backing.clone();
1081 let visibility_store: Arc<dyn VisibilityStore> = backing;
1082 let engine = Arc::new(
1083 EngineBuilder::new()
1084 .store_arc(Arc::clone(&store))
1085 .visibility_store_arc(Arc::clone(&visibility_store))
1086 .scheduler_threads(1)
1087 .build()
1088 .await?,
1089 );
1090 store
1091 .append(
1092 WriteToken::recorder(),
1093 &workflow_id(),
1094 &[started_event()?],
1095 0,
1096 )
1097 .await?;
1098 visibility_store
1099 .record_visibility(VisibilityRecord {
1100 workflow_id: workflow_id(),
1101 run_id: aion_core::RunId::new(uuid::Uuid::from_u128(2)),
1102 workflow_type: String::from("fixture"),
1103 status: WorkflowStatus::Running,
1104 start_time: Utc::now(),
1105 close_time: None,
1106 search_attributes: std::collections::HashMap::from([(
1107 crate::namespace::NAMESPACE_ATTRIBUTE.to_owned(),
1108 aion_core::SearchAttributeValue::String(NAMESPACE.to_owned()),
1109 )]),
1110 })
1111 .await?;
1112 let resolver = NamespaceResolver::from_config(
1113 crate::config::NamespaceConfig {
1114 mode: NamespaceMode::SharedEngine,
1115 },
1116 engine,
1117 );
1118 let state = server_state(resolver.clone(), runtime_config()).await?;
1119 let service = WorkflowGrpcService::new(state);
1120
1121 let mut start = Request::new(generated::StartWorkflowRequest {
1122 namespace: NAMESPACE.to_owned(),
1123 workflow_type: "missing-workflow".to_owned(),
1124 input: Some(encode_payload(proto_payload()?)),
1125 routing_key: None,
1126 });
1127 apply_metadata(start.metadata_mut())?;
1128 let start_error = service.start_workflow(start).await;
1129 let status = start_error
1130 .err()
1131 .ok_or_else(|| WireError::backend("expected error"))?;
1132 assert_eq!(status.code(), Code::NotFound);
1133 let detail = ProtoWireError::decode(status.details())?;
1134 assert_eq!(detail.error_type.as_deref(), Some("WorkflowTypeNotFound"));
1135 assert!(detail.message.contains("missing-workflow"));
1136
1137 let list_filter = encode_core_value(
1138 NAMESPACE,
1139 None,
1140 &aion_store::visibility::ListWorkflowsFilter {
1141 workflow_type: Some(String::from("fixture")),
1142 status: Some(WorkflowStatus::Running),
1143 ..aion_store::visibility::ListWorkflowsFilter::default()
1144 },
1145 )?;
1146 let mut list = Request::new(generated::ListWorkflowsRequest {
1147 namespace: NAMESPACE.to_owned(),
1148 filter: Some(encode_envelope(list_filter)),
1149 });
1150 apply_metadata(list.metadata_mut())?;
1151 let response = service.list_workflows(list).await?.into_inner();
1152
1153 assert_eq!(response.summaries.len(), 1);
1154 let summary = response
1155 .summaries
1156 .into_iter()
1157 .next()
1158 .map(decode_envelope)
1159 .map(|envelope| decode_core_value::<aion_store::visibility::WorkflowSummary>(&envelope))
1160 .transpose()?
1161 .ok_or_else(|| WireError::backend("summary missing"))?;
1162 assert_eq!(summary.workflow_id, workflow_id());
1163 assert_eq!(
1169 resolver
1170 .verify_workflow_ownership(NAMESPACE, &workflow_id())
1171 .await
1172 .err()
1173 .map(|error| error.to_wire_error().code),
1174 Some(WireErrorCode::NotFound)
1175 );
1176 Ok(())
1177 }
1178
1179 fn apply_metadata(
1180 metadata: &mut tonic::metadata::MetadataMap,
1181 ) -> Result<(), Box<dyn std::error::Error>> {
1182 #[cfg(feature = "auth")]
1186 let bearer = crate::auth::test_support::mint_token("alice", NAMESPACE)?;
1187 #[cfg(not(feature = "auth"))]
1188 let bearer = TOKEN.to_owned();
1189 metadata.insert("authorization", format!("Bearer {bearer}").parse()?);
1190 metadata.insert("x-aion-subject", "alice".parse()?);
1191 metadata.insert("x-aion-namespaces", NAMESPACE.parse()?);
1192 Ok(())
1193 }
1194
1195 fn runtime_config() -> RuntimeConfig {
1199 RuntimeConfig {
1200 listen: ListenConfig {
1201 grpc: SocketAddr::from(([127, 0, 0, 1], 50051)),
1202 http: SocketAddr::from(([127, 0, 0, 1], 8080)),
1203 },
1204 tls: None,
1205 auth: AuthConfig {
1206 enabled: true,
1207 jwks_url: Some(TOKEN.to_owned()),
1208 jwks_refresh_seconds: 300,
1209 },
1210 dashboard: DashboardConfig {
1211 source: DashboardAssetSource::Embedded,
1212 },
1213 namespace: NamespaceConfig {
1214 mode: NamespaceMode::SharedEngine,
1215 },
1216 worker: WorkerConfig {
1217 heartbeat_window: std::time::Duration::from_millis(30_000),
1218 },
1219 websocket: WebSocketConfig {
1220 outbound_buffer_bound: 32,
1221 event_broadcast_capacity: Some(64),
1222 },
1223 workflow_packages: Vec::new(),
1224 deploy: DeployConfig::default(),
1225 authoring: AuthoringConfig::default(),
1226 dev: crate::config::DevConfig::default(),
1227 outbox: crate::config::OutboxConfig::default(),
1228 scheduler_threads: 1,
1229 query_timeout: Some(std::time::Duration::from_millis(10_000)),
1230 default_namespace: "default".to_owned(),
1231 drain_timeout: std::time::Duration::from_secs(30),
1232 metrics: MetricsConfig { enabled: true },
1233 owned_shards: Vec::new(),
1234 cors_allowed_origins: Vec::new(),
1235 }
1236 }
1237
1238 fn started_event() -> Result<Event, aion_core::PayloadError> {
1239 Ok(Event::WorkflowStarted {
1240 envelope: EventEnvelope {
1241 seq: 1,
1242 recorded_at: Utc::now(),
1243 workflow_id: workflow_id(),
1244 },
1245 workflow_type: "fixture".to_owned(),
1246 input: payload()?,
1247 run_id: aion_core::RunId::new(uuid::Uuid::from_u128(1)),
1248 parent_run_id: None,
1249 package_version: aion_core::PackageVersion::new("a".repeat(64)),
1250 })
1251 }
1252
1253 fn proto_payload() -> Result<aion_proto::ProtoPayload, aion_core::PayloadError> {
1254 Ok(payload()?.into())
1255 }
1256
1257 fn payload() -> Result<Payload, aion_core::PayloadError> {
1258 Payload::from_json(&json!({ "fixture": "input" }))
1259 }
1260
1261 fn workflow_id() -> WorkflowId {
1262 WorkflowId::new(uuid::Uuid::from_u128(1))
1263 }
1264}