Skip to main content

aion_server/api/
grpc.rs

1//! tonic workflow service adapter.
2
3use 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/// Cloneable tonic implementation for workflow management.
20#[derive(Clone)]
21pub struct WorkflowGrpcService {
22    state: ServerState,
23}
24
25impl WorkflowGrpcService {
26    /// Build a tonic workflow service from shared server state.
27    #[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
37/// Construct the generated tonic server wrapper.
38#[must_use]
39pub fn workflow_service(state: ServerState) -> WorkflowServiceServer<WorkflowGrpcService> {
40    WorkflowServiceServer::new(WorkflowGrpcService::new(state))
41}
42
43#[tonic::async_trait]
44impl generated::workflow_service_server::WorkflowService for WorkflowGrpcService {
45    async fn start_workflow(
46        &self,
47        request: Request<generated::StartWorkflowRequest>,
48    ) -> Result<Response<generated::StartWorkflowResponse>, Status> {
49        if self.state.drain_state().is_draining() {
50            return Err(Status::unavailable(
51                "server is draining and not accepting new workflow starts",
52            ));
53        }
54        let caller = self.caller(&request).await?;
55        let response = handlers::start(
56            self.state.namespace_guard(),
57            &caller,
58            decode_start_request(request.into_inner()),
59        )
60        .await
61        .map_err(status_from_wire_error)?;
62        Ok(Response::new(encode_start_response(response)))
63    }
64
65    async fn signal(
66        &self,
67        request: Request<generated::SignalRequest>,
68    ) -> Result<Response<generated::SignalResponse>, Status> {
69        let caller = self.caller(&request).await?;
70        let response = handlers::signal(
71            self.state.namespace_guard(),
72            &caller,
73            decode_signal_request(request.into_inner()),
74        )
75        .await
76        .map_err(status_from_wire_error)?;
77        Ok(Response::new(encode_signal_response(response)))
78    }
79
80    async fn query(
81        &self,
82        request: Request<generated::QueryRequest>,
83    ) -> Result<Response<generated::QueryResponse>, Status> {
84        let caller = self.caller(&request).await?;
85        let response = handlers::query(
86            self.state.namespace_guard(),
87            &caller,
88            decode_query_request(request.into_inner()),
89        )
90        .await
91        .map_err(status_from_wire_error)?;
92        Ok(Response::new(encode_query_response(response)))
93    }
94
95    async fn cancel(
96        &self,
97        request: Request<generated::CancelRequest>,
98    ) -> Result<Response<generated::CancelResponse>, Status> {
99        let caller = self.caller(&request).await?;
100        let response = handlers::cancel(
101            self.state.namespace_guard(),
102            &caller,
103            decode_cancel_request(request.into_inner()),
104        )
105        .await
106        .map_err(status_from_wire_error)?;
107        Ok(Response::new(encode_cancel_response(response)))
108    }
109
110    async fn list_workflows(
111        &self,
112        request: Request<generated::ListWorkflowsRequest>,
113    ) -> Result<Response<generated::ListWorkflowsResponse>, Status> {
114        let caller = self.caller(&request).await?;
115        let response = handlers::list(
116            self.state.namespace_guard(),
117            &caller,
118            decode_list_request(request.into_inner()),
119        )
120        .await
121        .map_err(status_from_wire_error)?;
122        Ok(Response::new(encode_list_response(response)))
123    }
124
125    async fn count_workflows(
126        &self,
127        request: Request<generated::CountWorkflowsRequest>,
128    ) -> Result<Response<generated::CountWorkflowsResponse>, Status> {
129        let caller = self.caller(&request).await?;
130        let response = handlers::count(
131            self.state.namespace_guard(),
132            &caller,
133            decode_count_request(request.into_inner()),
134        )
135        .await
136        .map_err(status_from_wire_error)?;
137        Ok(Response::new(encode_count_response(response)))
138    }
139
140    async fn describe_workflow(
141        &self,
142        request: Request<generated::DescribeWorkflowRequest>,
143    ) -> Result<Response<generated::DescribeWorkflowResponse>, Status> {
144        let caller = self.caller(&request).await?;
145        let response = handlers::describe(
146            self.state.namespace_guard(),
147            &caller,
148            decode_describe_request(request.into_inner()),
149        )
150        .await
151        .map_err(status_from_wire_error)?;
152        Ok(Response::new(encode_describe_response(response)))
153    }
154
155    async fn create_schedule(
156        &self,
157        request: Request<generated::CreateScheduleRequest>,
158    ) -> Result<Response<generated::CreateScheduleResponse>, Status> {
159        let caller = self.caller(&request).await?;
160        let response = schedule_handlers::create_schedule(
161            self.state.namespace_guard(),
162            &caller,
163            decode_create_schedule_request(request.into_inner()),
164        )
165        .await
166        .map_err(status_from_wire_error)?;
167        Ok(Response::new(encode_create_schedule_response(response)))
168    }
169
170    async fn update_schedule(
171        &self,
172        request: Request<generated::UpdateScheduleRequest>,
173    ) -> Result<Response<generated::UpdateScheduleResponse>, Status> {
174        let caller = self.caller(&request).await?;
175        let response = schedule_handlers::update_schedule(
176            self.state.namespace_guard(),
177            &caller,
178            decode_update_schedule_request(request.into_inner()),
179        )
180        .await
181        .map_err(status_from_wire_error)?;
182        Ok(Response::new(encode_update_schedule_response(response)))
183    }
184
185    async fn pause_schedule(
186        &self,
187        request: Request<generated::ScheduleIdRequest>,
188    ) -> Result<Response<generated::PauseScheduleResponse>, Status> {
189        let caller = self.caller(&request).await?;
190        let response = schedule_handlers::pause_schedule(
191            self.state.namespace_guard(),
192            &caller,
193            decode_schedule_id_request(request.into_inner()),
194        )
195        .await
196        .map_err(status_from_wire_error)?;
197        Ok(Response::new(encode_pause_schedule_response(response)))
198    }
199
200    async fn resume_schedule(
201        &self,
202        request: Request<generated::ScheduleIdRequest>,
203    ) -> Result<Response<generated::ResumeScheduleResponse>, Status> {
204        let caller = self.caller(&request).await?;
205        let response = schedule_handlers::resume_schedule(
206            self.state.namespace_guard(),
207            &caller,
208            decode_schedule_id_request(request.into_inner()),
209        )
210        .await
211        .map_err(status_from_wire_error)?;
212        Ok(Response::new(encode_resume_schedule_response(response)))
213    }
214
215    async fn delete_schedule(
216        &self,
217        request: Request<generated::ScheduleIdRequest>,
218    ) -> Result<Response<generated::DeleteScheduleResponse>, Status> {
219        let caller = self.caller(&request).await?;
220        let response = schedule_handlers::delete_schedule(
221            self.state.namespace_guard(),
222            &caller,
223            decode_schedule_id_request(request.into_inner()),
224        )
225        .await
226        .map_err(status_from_wire_error)?;
227        Ok(Response::new(encode_delete_schedule_response(response)))
228    }
229
230    async fn list_schedules(
231        &self,
232        request: Request<generated::ListSchedulesRequest>,
233    ) -> Result<Response<generated::ListSchedulesResponse>, Status> {
234        let caller = self.caller(&request).await?;
235        let response = schedule_handlers::list_schedules(
236            self.state.namespace_guard(),
237            &caller,
238            decode_list_schedules_request(request.into_inner()),
239        )
240        .await
241        .map_err(status_from_wire_error)?;
242        Ok(Response::new(encode_list_schedules_response(response)))
243    }
244
245    async fn describe_schedule(
246        &self,
247        request: Request<generated::ScheduleIdRequest>,
248    ) -> Result<Response<generated::DescribeScheduleResponse>, Status> {
249        let caller = self.caller(&request).await?;
250        let response = schedule_handlers::describe_schedule(
251            self.state.namespace_guard(),
252            &caller,
253            decode_schedule_id_request(request.into_inner()),
254        )
255        .await
256        .map_err(status_from_wire_error)?;
257        Ok(Response::new(encode_describe_schedule_response(response)))
258    }
259}
260
261pub(crate) async fn caller_from_metadata(
262    metadata: &tonic::metadata::MetadataMap,
263    state: &ServerState,
264) -> Result<CallerIdentity, Status> {
265    if !state.runtime_config().auth.enabled {
266        return Ok(development_caller_from_metadata(metadata));
267    }
268    #[cfg(feature = "auth")]
269    {
270        let bearer = metadata
271            .get("authorization")
272            .and_then(|value| value.to_str().ok())
273            .and_then(parse_bearer)
274            .ok_or_else(|| Status::unauthenticated("missing bearer token"))?;
275        let Some(cache) = state.jwks_cache() else {
276            return Err(Status::unauthenticated("invalid bearer token"));
277        };
278        return cache
279            .validate(&bearer)
280            .await
281            .map(|claims| claims.caller_identity())
282            .map_err(|_error| Status::unauthenticated("invalid bearer token"));
283    }
284    #[cfg(not(feature = "auth"))]
285    {
286        // Yield to preserve the async signature required by the auth-feature branch.
287        tokio::task::yield_now().await;
288        Ok(development_token_caller_from_metadata(
289            metadata,
290            &state.runtime_config().auth,
291        ))
292    }
293}
294
295fn development_caller_from_metadata(metadata: &tonic::metadata::MetadataMap) -> CallerIdentity {
296    let subject = metadata
297        .get("x-aion-subject")
298        .and_then(|value| value.to_str().ok())
299        .filter(|value| !value.is_empty())
300        .unwrap_or("anonymous");
301    let namespaces = metadata
302        .get("x-aion-namespaces")
303        .and_then(|value| value.to_str().ok())
304        .map(parse_namespaces)
305        .unwrap_or_default();
306    CallerIdentity::new(subject, namespaces).with_deploy(deploy_metadata_granted(metadata))
307}
308
309/// Deployment-wide deploy grant from the development `x-aion-deploy`
310/// metadata entry, the dev-mode analog of the JWT `deploy` claim.
311fn deploy_metadata_granted(metadata: &tonic::metadata::MetadataMap) -> bool {
312    metadata
313        .get("x-aion-deploy")
314        .and_then(|value| value.to_str().ok())
315        .is_some_and(|value| value.trim().eq_ignore_ascii_case("true"))
316}
317
318/// Development-mode token authentication for gRPC metadata, mirroring the HTTP
319/// development token auth.  Used when `auth.enabled` is `true` but the `auth`
320/// crate feature is not compiled.
321#[cfg(not(feature = "auth"))]
322fn development_token_caller_from_metadata(
323    metadata: &tonic::metadata::MetadataMap,
324    auth: &crate::config::AuthConfig,
325) -> CallerIdentity {
326    let subject = metadata
327        .get("x-aion-subject")
328        .and_then(|value| value.to_str().ok())
329        .filter(|value| !value.is_empty());
330    let namespaces = metadata
331        .get("x-aion-namespaces")
332        .and_then(|value| value.to_str().ok())
333        .map(parse_namespaces)
334        .unwrap_or_default();
335
336    let bearer_token = auth.jwks_url.as_deref().unwrap_or_default();
337    let expected = format!("Bearer {bearer_token}");
338    let Some(authorization) = metadata.get("authorization") else {
339        return CallerIdentity::denied(subject.unwrap_or("anonymous"), "missing bearer token");
340    };
341    let authorization = authorization.to_str().ok();
342    if authorization != Some(expected.as_str()) {
343        return CallerIdentity::denied(subject.unwrap_or("anonymous"), "invalid bearer token");
344    }
345
346    let Some(subject) = subject else {
347        return CallerIdentity::denied("anonymous", "missing required metadata: x-aion-subject");
348    };
349
350    CallerIdentity::new(subject, namespaces).with_deploy(deploy_metadata_granted(metadata))
351}
352
353#[cfg(feature = "auth")]
354fn parse_bearer(value: &str) -> Option<String> {
355    let token = value.strip_prefix("Bearer ")?.trim();
356    if token.is_empty() {
357        return None;
358    }
359    Some(token.to_owned())
360}
361
362fn parse_namespaces(value: &str) -> Vec<String> {
363    value
364        .split(',')
365        .map(str::trim)
366        .filter(|namespace| !namespace.is_empty())
367        .map(str::to_owned)
368        .collect()
369}
370
371pub(crate) fn status_from_wire_error(error: WireError) -> Status {
372    status_with_code(grpc_code(error.code), error)
373}
374
375/// Build a tonic status with an explicit code, carrying the typed
376/// `ProtoWireError` detail payload when it encodes.
377pub(crate) fn status_with_code(code: Code, error: WireError) -> Status {
378    let message = error.message.clone();
379    let mut details = Vec::new();
380    let proto_error = ProtoWireError::from(error);
381    if proto_error.encode(&mut details).is_ok() {
382        Status::with_details(code, message, details.into())
383    } else {
384        Status::new(code, message)
385    }
386}
387
388fn grpc_code(code: aion_proto::WireErrorCode) -> Code {
389    match code {
390        aion_proto::WireErrorCode::NotFound => Code::NotFound,
391        aion_proto::WireErrorCode::NamespaceDenied | aion_proto::WireErrorCode::DeployDenied => {
392            Code::PermissionDenied
393        }
394        aion_proto::WireErrorCode::SequenceConflict => Code::Aborted,
395        aion_proto::WireErrorCode::UnknownQuery | aion_proto::WireErrorCode::InvalidInput => {
396            Code::InvalidArgument
397        }
398        aion_proto::WireErrorCode::QueryTimeout => Code::DeadlineExceeded,
399        aion_proto::WireErrorCode::NotRunning | aion_proto::WireErrorCode::VersionPinned => {
400            Code::FailedPrecondition
401        }
402        aion_proto::WireErrorCode::Lagged => Code::ResourceExhausted,
403        // query_failed normally rides QueryResponse.error inside an OK
404        // response; a transport-level carrier still attaches the typed
405        // ProtoWireError detail, so detail-aware clients keep QueryFailed.
406        aion_proto::WireErrorCode::Backend | aion_proto::WireErrorCode::QueryFailed => {
407            Code::Internal
408        }
409    }
410}
411
412fn decode_workflow_id(value: generated::WorkflowId) -> aion_proto::ProtoWorkflowId {
413    aion_proto::ProtoWorkflowId { uuid: value.uuid }
414}
415
416fn encode_workflow_id(value: aion_proto::ProtoWorkflowId) -> generated::WorkflowId {
417    generated::WorkflowId { uuid: value.uuid }
418}
419
420fn decode_run_id(value: generated::RunId) -> aion_proto::ProtoRunId {
421    aion_proto::ProtoRunId { uuid: value.uuid }
422}
423
424fn encode_run_id(value: aion_proto::ProtoRunId) -> generated::RunId {
425    generated::RunId { uuid: value.uuid }
426}
427
428fn decode_schedule_id(value: generated::ScheduleId) -> aion_proto::ProtoScheduleId {
429    aion_proto::ProtoScheduleId { uuid: value.uuid }
430}
431
432fn encode_schedule_id(value: aion_proto::ProtoScheduleId) -> generated::ScheduleId {
433    generated::ScheduleId { uuid: value.uuid }
434}
435
436fn decode_payload(value: generated::Payload) -> aion_proto::ProtoPayload {
437    aion_proto::ProtoPayload {
438        content_type: value.content_type,
439        bytes: value.bytes,
440    }
441}
442
443fn encode_payload(value: aion_proto::ProtoPayload) -> generated::Payload {
444    generated::Payload {
445        content_type: value.content_type,
446        bytes: value.bytes,
447    }
448}
449
450fn decode_envelope(value: generated::WireEnvelope) -> aion_proto::WireEnvelope {
451    aion_proto::WireEnvelope {
452        namespace: value.namespace,
453        request_id: value.request_id,
454        payload: value.payload.map(decode_payload),
455    }
456}
457
458fn encode_envelope(value: aion_proto::WireEnvelope) -> generated::WireEnvelope {
459    generated::WireEnvelope {
460        namespace: value.namespace,
461        request_id: value.request_id,
462        payload: value.payload.map(encode_payload),
463    }
464}
465
466fn decode_start_request(value: generated::StartWorkflowRequest) -> ProtoStartWorkflowRequest {
467    ProtoStartWorkflowRequest {
468        namespace: value.namespace,
469        workflow_type: value.workflow_type,
470        input: value.input.map(decode_payload),
471    }
472}
473
474fn encode_start_response(value: ProtoStartWorkflowResponse) -> generated::StartWorkflowResponse {
475    generated::StartWorkflowResponse {
476        workflow_id: value.workflow_id.map(encode_workflow_id),
477        run_id: value.run_id.map(encode_run_id),
478    }
479}
480
481fn decode_signal_request(value: generated::SignalRequest) -> ProtoSignalRequest {
482    ProtoSignalRequest {
483        namespace: value.namespace,
484        workflow_id: value.workflow_id.map(decode_workflow_id),
485        run_id: value.run_id.map(decode_run_id),
486        signal_name: value.signal_name,
487        payload: value.payload.map(decode_payload),
488    }
489}
490
491fn encode_signal_response(_: ProtoSignalResponse) -> generated::SignalResponse {
492    generated::SignalResponse {}
493}
494
495fn decode_query_request(value: generated::QueryRequest) -> ProtoQueryRequest {
496    ProtoQueryRequest {
497        namespace: value.namespace,
498        workflow_id: value.workflow_id.map(decode_workflow_id),
499        run_id: value.run_id.map(decode_run_id),
500        query_name: value.query_name,
501    }
502}
503
504fn encode_query_response(value: ProtoQueryResponse) -> generated::QueryResponse {
505    generated::QueryResponse {
506        outcome: value.outcome.map(encode_query_outcome),
507    }
508}
509
510fn encode_query_outcome(
511    value: aion_proto::proto_query_response::Outcome,
512) -> generated::query_response::Outcome {
513    match value {
514        aion_proto::proto_query_response::Outcome::Result(payload) => {
515            generated::query_response::Outcome::Result(encode_payload(payload))
516        }
517        aion_proto::proto_query_response::Outcome::Error(error) => {
518            generated::query_response::Outcome::Error(encode_wire_error(error))
519        }
520    }
521}
522
523fn encode_wire_error(value: ProtoWireError) -> generated::WireError {
524    generated::WireError {
525        code: value.code,
526        message: value.message,
527        error_type: value.error_type,
528    }
529}
530
531fn decode_cancel_request(value: generated::CancelRequest) -> ProtoCancelRequest {
532    ProtoCancelRequest {
533        namespace: value.namespace,
534        workflow_id: value.workflow_id.map(decode_workflow_id),
535        run_id: value.run_id.map(decode_run_id),
536        reason: value.reason,
537    }
538}
539
540fn encode_cancel_response(_: ProtoCancelResponse) -> generated::CancelResponse {
541    generated::CancelResponse {}
542}
543
544fn decode_list_request(value: generated::ListWorkflowsRequest) -> ProtoListWorkflowsRequest {
545    ProtoListWorkflowsRequest {
546        namespace: value.namespace,
547        filter: value.filter.map(decode_envelope),
548    }
549}
550
551fn encode_list_response(value: ProtoListWorkflowsResponse) -> generated::ListWorkflowsResponse {
552    generated::ListWorkflowsResponse {
553        summaries: value.summaries.into_iter().map(encode_envelope).collect(),
554    }
555}
556
557fn decode_count_request(value: generated::CountWorkflowsRequest) -> ProtoCountWorkflowsRequest {
558    ProtoCountWorkflowsRequest {
559        namespace: value.namespace,
560        filter: value.filter.map(decode_envelope),
561    }
562}
563
564fn encode_count_response(value: ProtoCountWorkflowsResponse) -> generated::CountWorkflowsResponse {
565    generated::CountWorkflowsResponse { count: value.count }
566}
567
568fn decode_describe_request(
569    value: generated::DescribeWorkflowRequest,
570) -> ProtoDescribeWorkflowRequest {
571    ProtoDescribeWorkflowRequest {
572        namespace: value.namespace,
573        workflow_id: value.workflow_id.map(decode_workflow_id),
574        run_id: value.run_id.map(decode_run_id),
575        include_history: value.include_history,
576    }
577}
578
579fn encode_describe_response(
580    value: ProtoDescribeWorkflowResponse,
581) -> generated::DescribeWorkflowResponse {
582    generated::DescribeWorkflowResponse {
583        summary: value.summary.map(encode_envelope),
584        history: value.history.into_iter().map(encode_envelope).collect(),
585    }
586}
587
588fn decode_create_schedule_request(
589    value: generated::CreateScheduleRequest,
590) -> ProtoCreateScheduleRequest {
591    ProtoCreateScheduleRequest {
592        namespace: value.namespace,
593        config: value.config.map(decode_envelope),
594    }
595}
596
597fn encode_create_schedule_response(
598    value: ProtoCreateScheduleResponse,
599) -> generated::CreateScheduleResponse {
600    generated::CreateScheduleResponse {
601        schedule_id: value.schedule_id.map(encode_schedule_id),
602        state: value.state.map(encode_envelope),
603    }
604}
605
606fn decode_update_schedule_request(
607    value: generated::UpdateScheduleRequest,
608) -> ProtoUpdateScheduleRequest {
609    ProtoUpdateScheduleRequest {
610        namespace: value.namespace,
611        schedule_id: value.schedule_id.map(decode_schedule_id),
612        config: value.config.map(decode_envelope),
613    }
614}
615
616fn encode_update_schedule_response(
617    value: ProtoUpdateScheduleResponse,
618) -> generated::UpdateScheduleResponse {
619    generated::UpdateScheduleResponse {
620        state: value.state.map(encode_envelope),
621    }
622}
623
624fn decode_schedule_id_request(value: generated::ScheduleIdRequest) -> ProtoScheduleIdRequest {
625    ProtoScheduleIdRequest {
626        namespace: value.namespace,
627        schedule_id: value.schedule_id.map(decode_schedule_id),
628    }
629}
630
631fn encode_pause_schedule_response(
632    value: ProtoPauseScheduleResponse,
633) -> generated::PauseScheduleResponse {
634    generated::PauseScheduleResponse {
635        state: value.state.map(encode_envelope),
636    }
637}
638
639fn encode_resume_schedule_response(
640    value: ProtoResumeScheduleResponse,
641) -> generated::ResumeScheduleResponse {
642    generated::ResumeScheduleResponse {
643        state: value.state.map(encode_envelope),
644    }
645}
646
647fn encode_delete_schedule_response(
648    _: ProtoDeleteScheduleResponse,
649) -> generated::DeleteScheduleResponse {
650    generated::DeleteScheduleResponse {}
651}
652
653fn decode_list_schedules_request(
654    value: generated::ListSchedulesRequest,
655) -> ProtoListSchedulesRequest {
656    ProtoListSchedulesRequest {
657        namespace: value.namespace,
658    }
659}
660
661fn encode_list_schedules_response(
662    value: ProtoListSchedulesResponse,
663) -> generated::ListSchedulesResponse {
664    generated::ListSchedulesResponse {
665        schedules: value.schedules.into_iter().map(encode_envelope).collect(),
666    }
667}
668
669fn encode_describe_schedule_response(
670    value: ProtoDescribeScheduleResponse,
671) -> generated::DescribeScheduleResponse {
672    generated::DescribeScheduleResponse {
673        state: value.state.map(encode_envelope),
674    }
675}
676
677#[cfg(test)]
678mod tests {
679    use std::{net::SocketAddr, sync::Arc};
680
681    use aion::EngineBuilder;
682    use aion_core::{Event, EventEnvelope, Payload, WorkflowId, WorkflowStatus};
683    use aion_proto::{
684        WireErrorCode,
685        convert::{decode_core_value, encode_core_value},
686        generated::workflow_service_server::WorkflowService,
687    };
688    use aion_store::{
689        EventStore, InMemoryStore, WriteToken,
690        visibility::{VisibilityRecord, VisibilityStore},
691    };
692    use chrono::Utc;
693    use serde_json::json;
694    use tonic::Request;
695
696    use super::*;
697    use crate::{
698        NamespaceResolver,
699        config::{
700            AuthConfig, DashboardAssetSource, DashboardConfig, DeployConfig, ListenConfig,
701            MetricsConfig, NamespaceConfig, NamespaceMode, RuntimeConfig, WebSocketConfig,
702            WorkerConfig,
703        },
704    };
705
706    const NAMESPACE: &str = "tenant-a";
707    const TOKEN: &str = "test-token";
708
709    /// Server state whose bearer validation matches the compiled auth path:
710    /// under `feature = "auth"` a real [`crate::auth::JwksCache`] is fetched
711    /// from a live fixture JWKS endpoint; otherwise the development token path
712    /// needs no cache.
713    async fn server_state(
714        resolver: NamespaceResolver,
715        runtime: RuntimeConfig,
716    ) -> Result<ServerState, Box<dyn std::error::Error>> {
717        #[cfg(feature = "auth")]
718        {
719            let url = crate::auth::test_support::serve_jwks()?;
720            let refresh = std::time::Duration::from_secs(runtime.auth.jwks_refresh_seconds);
721            let cache = crate::auth::JwksCache::new(url, refresh).await?;
722            Ok(ServerState::from_parts_with_jwks(resolver, runtime, cache))
723        }
724        #[cfg(not(feature = "auth"))]
725        {
726            // Yield to preserve the async signature required by the auth-feature branch.
727            tokio::task::yield_now().await;
728            Ok(ServerState::from_parts(resolver, runtime))
729        }
730    }
731
732    #[tokio::test]
733    async fn in_process_tonic_start_and_list_use_shared_handlers()
734    -> Result<(), Box<dyn std::error::Error>> {
735        let backing = Arc::new(InMemoryStore::default());
736        let store: Arc<dyn EventStore> = backing.clone();
737        let visibility_store: Arc<dyn VisibilityStore> = backing;
738        let engine = Arc::new(
739            EngineBuilder::new()
740                .store_arc(Arc::clone(&store))
741                .visibility_store_arc(Arc::clone(&visibility_store))
742                .scheduler_threads(1)
743                .build()
744                .await?,
745        );
746        store
747            .append(
748                WriteToken::recorder(),
749                &workflow_id(),
750                &[started_event()?],
751                0,
752            )
753            .await?;
754        visibility_store
755            .record_visibility(VisibilityRecord {
756                workflow_id: workflow_id(),
757                run_id: aion_core::RunId::new(uuid::Uuid::from_u128(2)),
758                workflow_type: String::from("fixture"),
759                status: WorkflowStatus::Running,
760                start_time: Utc::now(),
761                close_time: None,
762                search_attributes: std::collections::HashMap::from([(
763                    crate::namespace::NAMESPACE_ATTRIBUTE.to_owned(),
764                    aion_core::SearchAttributeValue::String(NAMESPACE.to_owned()),
765                )]),
766            })
767            .await?;
768        let resolver = NamespaceResolver::from_config(
769            crate::config::NamespaceConfig {
770                mode: NamespaceMode::SharedEngine,
771            },
772            engine,
773        );
774        let state = server_state(resolver.clone(), runtime_config()).await?;
775        let service = WorkflowGrpcService::new(state);
776
777        let mut start = Request::new(generated::StartWorkflowRequest {
778            namespace: NAMESPACE.to_owned(),
779            workflow_type: "missing-workflow".to_owned(),
780            input: Some(encode_payload(proto_payload()?)),
781        });
782        apply_metadata(start.metadata_mut())?;
783        let start_error = service.start_workflow(start).await;
784        let status = start_error
785            .err()
786            .ok_or_else(|| WireError::backend("expected error"))?;
787        assert_eq!(status.code(), Code::NotFound);
788        let detail = ProtoWireError::decode(status.details())?;
789        assert_eq!(detail.error_type.as_deref(), Some("WorkflowTypeNotFound"));
790        assert!(detail.message.contains("missing-workflow"));
791
792        let list_filter = encode_core_value(
793            NAMESPACE,
794            None,
795            &aion_store::visibility::ListWorkflowsFilter {
796                workflow_type: Some(String::from("fixture")),
797                status: Some(WorkflowStatus::Running),
798                ..aion_store::visibility::ListWorkflowsFilter::default()
799            },
800        )?;
801        let mut list = Request::new(generated::ListWorkflowsRequest {
802            namespace: NAMESPACE.to_owned(),
803            filter: Some(encode_envelope(list_filter)),
804        });
805        apply_metadata(list.metadata_mut())?;
806        let response = service.list_workflows(list).await?.into_inner();
807
808        assert_eq!(response.summaries.len(), 1);
809        let summary = response
810            .summaries
811            .into_iter()
812            .next()
813            .map(decode_envelope)
814            .map(|envelope| decode_core_value::<aion_store::visibility::WorkflowSummary>(&envelope))
815            .transpose()?
816            .ok_or_else(|| WireError::backend("summary missing"))?;
817        assert_eq!(summary.workflow_id, workflow_id());
818        // The seeded history records no namespace attribute, so durable
819        // ownership verification must reject targeted access with NotFound:
820        // a missing ownership attribute is indistinguishable from a
821        // nonexistent workflow (anti-existence-leak), and NamespaceDenied is
822        // reserved for callers without a grant for the requested namespace.
823        assert_eq!(
824            resolver
825                .verify_workflow_ownership(NAMESPACE, &workflow_id())
826                .await
827                .err()
828                .map(|error| error.to_wire_error().code),
829            Some(WireErrorCode::NotFound)
830        );
831        Ok(())
832    }
833
834    fn apply_metadata(
835        metadata: &mut tonic::metadata::MetadataMap,
836    ) -> Result<(), Box<dyn std::error::Error>> {
837        // Bearer credential accepted by the compiled authentication path: a
838        // JWT minted against the fixture JWKS under `feature = "auth"`, the
839        // development shared-secret token otherwise.
840        #[cfg(feature = "auth")]
841        let bearer = crate::auth::test_support::mint_token("alice", NAMESPACE)?;
842        #[cfg(not(feature = "auth"))]
843        let bearer = TOKEN.to_owned();
844        metadata.insert("authorization", format!("Bearer {bearer}").parse()?);
845        metadata.insert("x-aion-subject", "alice".parse()?);
846        metadata.insert("x-aion-namespaces", NAMESPACE.parse()?);
847        Ok(())
848    }
849
850    /// Test runtime settings with authentication enabled; under
851    /// `feature = "auth"` validation runs against the [`server_state`]-injected
852    /// JWKS cache, so the configured dev-secret `jwks_url` is never fetched.
853    fn runtime_config() -> RuntimeConfig {
854        RuntimeConfig {
855            listen: ListenConfig {
856                grpc: SocketAddr::from(([127, 0, 0, 1], 50051)),
857                http: SocketAddr::from(([127, 0, 0, 1], 8080)),
858            },
859            tls: None,
860            auth: AuthConfig {
861                enabled: true,
862                jwks_url: Some(TOKEN.to_owned()),
863                jwks_refresh_seconds: 300,
864            },
865            dashboard: DashboardConfig {
866                source: DashboardAssetSource::Embedded,
867            },
868            namespace: NamespaceConfig {
869                mode: NamespaceMode::SharedEngine,
870            },
871            worker: WorkerConfig {
872                heartbeat_window: std::time::Duration::from_millis(30_000),
873            },
874            websocket: WebSocketConfig {
875                outbound_buffer_bound: 32,
876                event_broadcast_capacity: Some(64),
877            },
878            workflow_packages: Vec::new(),
879            deploy: DeployConfig::default(),
880            scheduler_threads: 1,
881            query_timeout: Some(std::time::Duration::from_millis(10_000)),
882            default_namespace: "default".to_owned(),
883            drain_timeout: std::time::Duration::from_secs(30),
884            metrics: MetricsConfig { enabled: true },
885        }
886    }
887
888    fn started_event() -> Result<Event, aion_core::PayloadError> {
889        Ok(Event::WorkflowStarted {
890            envelope: EventEnvelope {
891                seq: 1,
892                recorded_at: Utc::now(),
893                workflow_id: workflow_id(),
894            },
895            workflow_type: "fixture".to_owned(),
896            input: payload()?,
897            run_id: aion_core::RunId::new(uuid::Uuid::from_u128(1)),
898            parent_run_id: None,
899            package_version: aion_core::PackageVersion::new("a".repeat(64)),
900        })
901    }
902
903    fn proto_payload() -> Result<aion_proto::ProtoPayload, aion_core::PayloadError> {
904        Ok(payload()?.into())
905    }
906
907    fn payload() -> Result<Payload, aion_core::PayloadError> {
908        Payload::from_json(&json!({ "fixture": "input" }))
909    }
910
911    fn workflow_id() -> WorkflowId {
912        WorkflowId::new(uuid::Uuid::from_u128(1))
913    }
914}