1use aion_proto::{
8 ProtoRouteVersionRequest, ProtoUnloadVersionRequest,
9 generated::{self, deploy_service_server::DeployServiceServer},
10};
11use tonic::{Code, Request, Response, Status};
12
13use super::grpc::{caller_from_metadata, status_from_wire_error, status_with_code};
14use crate::api::handlers::deploy::{self, DeployApiError};
15use crate::api::handlers::managed_workers;
16use crate::config::DEPLOY_MAX_ARCHIVE_BYTES_REQUIRED;
17use crate::worker::supervisor::SupervisionError;
18use crate::{CallerIdentity, ServerState};
19
20const TRANSPORT: &str = "grpc";
21
22const LOAD_PACKAGE_FRAMING_ALLOWANCE: usize = 64;
31
32#[derive(Clone)]
34pub struct DeployGrpcService {
35 state: ServerState,
36}
37
38impl DeployGrpcService {
39 #[must_use]
41 pub const fn new(state: ServerState) -> Self {
42 Self { state }
43 }
44
45 async fn caller<T>(&self, request: &Request<T>) -> Result<CallerIdentity, Status> {
46 caller_from_metadata(request.metadata(), &self.state).await
47 }
48}
49
50pub fn deploy_service(
59 state: ServerState,
60) -> Result<DeployServiceServer<DeployGrpcService>, crate::ServerError> {
61 let Some(limit) = state.runtime_config().deploy.max_archive_bytes else {
62 return Err(crate::ServerError::Config {
63 message: DEPLOY_MAX_ARCHIVE_BYTES_REQUIRED.to_owned(),
64 });
65 };
66 let limit = usize::try_from(limit).unwrap_or(usize::MAX);
69 Ok(DeployServiceServer::new(DeployGrpcService::new(state))
70 .max_decoding_message_size(limit.saturating_add(LOAD_PACKAGE_FRAMING_ALLOWANCE)))
71}
72
73#[tonic::async_trait]
74impl generated::deploy_service_server::DeployService for DeployGrpcService {
75 async fn load_package(
76 &self,
77 request: Request<generated::LoadPackageRequest>,
78 ) -> Result<Response<generated::LoadPackageResponse>, Status> {
79 let caller = self.caller(&request).await?;
80 let response = deploy::load_package(
81 &self.state,
82 &caller,
83 TRANSPORT,
84 request.into_inner().archive,
85 )
86 .await
87 .map_err(status_from_deploy_error)?;
88 Ok(Response::new(generated::LoadPackageResponse {
89 workflow_type: response.workflow_type,
90 content_hash: response.content_hash,
91 deployed_entry_module: response.deployed_entry_module,
92 entry_function: response.entry_function,
93 freshly_loaded: response.freshly_loaded,
94 route_changed: response.route_changed,
95 superseded_versions: response.superseded_versions,
96 auto_workers: response
97 .auto_workers
98 .into_iter()
99 .map(|worker| generated::AutoWorker {
100 task_queue: worker.task_queue,
101 decision: worker.decision,
102 deployment: worker.deployment,
103 detail: worker.detail,
104 })
105 .collect(),
106 }))
107 }
108
109 async fn list_versions(
110 &self,
111 request: Request<generated::ListVersionsRequest>,
112 ) -> Result<Response<generated::ListVersionsResponse>, Status> {
113 let caller = self.caller(&request).await?;
114 let response = deploy::list_versions(&self.state, &caller, TRANSPORT)
115 .map_err(status_from_deploy_error)?;
116 Ok(Response::new(generated::ListVersionsResponse {
117 versions: response
118 .versions
119 .into_iter()
120 .map(|version| generated::WorkflowVersion {
121 workflow_type: version.workflow_type,
122 content_hash: version.content_hash,
123 deployed_entry_module: version.deployed_entry_module,
124 entry_function: version.entry_function,
125 manifest_version: version.manifest_version,
126 loaded_at: version.loaded_at,
127 route_active: version.route_active,
128 })
129 .collect(),
130 }))
131 }
132
133 async fn route_version(
134 &self,
135 request: Request<generated::RouteVersionRequest>,
136 ) -> Result<Response<generated::RouteVersionResponse>, Status> {
137 let caller = self.caller(&request).await?;
138 let inner = request.into_inner();
139 deploy::route_version(
140 &self.state,
141 &caller,
142 TRANSPORT,
143 ProtoRouteVersionRequest {
144 workflow_type: inner.workflow_type,
145 content_hash: inner.content_hash,
146 },
147 )
148 .await
149 .map_err(status_from_deploy_error)?;
150 Ok(Response::new(generated::RouteVersionResponse {}))
151 }
152
153 async fn unload_version(
154 &self,
155 request: Request<generated::UnloadVersionRequest>,
156 ) -> Result<Response<generated::UnloadVersionResponse>, Status> {
157 let caller = self.caller(&request).await?;
158 let inner = request.into_inner();
159 deploy::unload_version(
160 &self.state,
161 &caller,
162 TRANSPORT,
163 ProtoUnloadVersionRequest {
164 workflow_type: inner.workflow_type,
165 content_hash: inner.content_hash,
166 },
167 )
168 .await
169 .map_err(status_from_deploy_error)?;
170 Ok(Response::new(generated::UnloadVersionResponse {}))
171 }
172
173 async fn list_managed_workers(
174 &self,
175 request: Request<generated::ListManagedWorkersRequest>,
176 ) -> Result<Response<generated::ListManagedWorkersResponse>, Status> {
177 let caller = self.caller(&request).await?;
178 require_deploy_grant(&caller)?;
179 let report = self
180 .state
181 .worker_supervisor()
182 .report()
183 .await
184 .map_err(|error| status_from_supervision_error(&error))?;
185 Ok(Response::new(managed_workers::managed_worker_report(
186 report,
187 )))
188 }
189
190 async fn start_managed_worker(
191 &self,
192 request: Request<generated::ManagedWorkerRequest>,
193 ) -> Result<Response<generated::ManagedWorkerResponse>, Status> {
194 let caller = self.caller(&request).await?;
195 require_deploy_grant(&caller)?;
196 let name = request.into_inner().name;
197 let status = self
198 .state
199 .worker_supervisor()
200 .start(&name)
201 .await
202 .map_err(|error| status_from_supervision_error(&error))?;
203 Ok(Response::new(generated::ManagedWorkerResponse {
204 worker: Some(managed_workers::managed_worker(status)),
205 }))
206 }
207
208 async fn stop_managed_worker(
209 &self,
210 request: Request<generated::ManagedWorkerRequest>,
211 ) -> Result<Response<generated::ManagedWorkerResponse>, Status> {
212 let caller = self.caller(&request).await?;
213 require_deploy_grant(&caller)?;
214 let name = request.into_inner().name;
215 let status = self
216 .state
217 .worker_supervisor()
218 .stop(&name)
219 .await
220 .map_err(|error| status_from_supervision_error(&error))?;
221 Ok(Response::new(generated::ManagedWorkerResponse {
222 worker: Some(managed_workers::managed_worker(status)),
223 }))
224 }
225
226 async fn restart_managed_worker(
227 &self,
228 request: Request<generated::ManagedWorkerRequest>,
229 ) -> Result<Response<generated::ManagedWorkerResponse>, Status> {
230 let caller = self.caller(&request).await?;
231 require_deploy_grant(&caller)?;
232 let name = request.into_inner().name;
233 let status = self
234 .state
235 .worker_supervisor()
236 .restart(&name)
237 .await
238 .map_err(|error| status_from_supervision_error(&error))?;
239 Ok(Response::new(generated::ManagedWorkerResponse {
240 worker: Some(managed_workers::managed_worker(status)),
241 }))
242 }
243}
244
245fn require_deploy_grant(caller: &CallerIdentity) -> Result<(), Status> {
249 if caller.deploy_granted() {
250 Ok(())
251 } else {
252 Err(status_from_wire_error(
253 aion_proto::WireError::deploy_denied(
254 "managed-worker administration requires the deployment-wide deploy grant",
255 ),
256 ))
257 }
258}
259
260fn status_from_supervision_error(error: &SupervisionError) -> Status {
263 status_from_wire_error(managed_workers::wire_error(error))
264}
265
266fn status_from_deploy_error(error: DeployApiError) -> Status {
270 match error {
271 DeployApiError::Unavailable(wire) => status_with_code(Code::Unavailable, wire),
272 DeployApiError::ArchiveTooLarge(wire) => status_with_code(Code::InvalidArgument, wire),
273 DeployApiError::Wire(wire) => status_from_wire_error(wire),
274 }
275}
276
277#[cfg(test)]
278mod tests {
279 use std::sync::Arc;
280
281 use aion::EngineBuilder;
282 use aion_proto::{ProtoWireError, WireError, WireErrorCode, generated};
283 use aion_store::{EventStore, InMemoryStore};
284 use prost::Message as _;
285 use tonic::{Code, Request, Status};
286
287 use super::DeployGrpcService;
288 use crate::config::{
289 AuthConfig, AuthoringConfig, DeployConfig, ListenConfig, MetricsConfig, NamespaceConfig,
290 NamespaceMode, OpsConsoleAssetSource, OpsConsoleConfig, RuntimeConfig, WebSocketConfig,
291 WorkerConfig,
292 };
293 use crate::test_support::{EngineUnderTest, StateUnderTest};
294 use crate::{
295 NamespaceResolver, ServerState, StaticScheduleNamespaces, StaticWorkflowNamespaces,
296 };
297
298 fn decode_detail(status: &Status) -> Result<WireError, Box<dyn std::error::Error>> {
300 let proto = ProtoWireError::decode(status.details())?;
301 Ok(WireError::try_from(proto)?)
302 }
303
304 fn runtime_config() -> RuntimeConfig {
305 RuntimeConfig {
306 listen: ListenConfig {
307 grpc: std::net::SocketAddr::from(([127, 0, 0, 1], 50051)),
308 http: std::net::SocketAddr::from(([127, 0, 0, 1], 8080)),
309 },
310 tls: None,
311 auth: AuthConfig {
312 enabled: false,
313 jwks_url: None,
314 jwks_refresh_seconds: 300,
315 },
316 ops_console: OpsConsoleConfig {
317 source: OpsConsoleAssetSource::Embedded,
318 },
319 namespace: NamespaceConfig {
320 mode: NamespaceMode::SharedEngine,
321 },
322 worker: WorkerConfig {
323 heartbeat_window: std::time::Duration::from_secs(30),
324 ..WorkerConfig::default()
325 },
326 websocket: WebSocketConfig {
327 outbound_buffer_bound: 32,
328 event_broadcast_capacity: Some(64),
329 cluster_broadcast_capacity: Some(64),
330 },
331 workflow_packages: Vec::new(),
332 deploy: DeployConfig::default(),
333 authoring: AuthoringConfig::default(),
334 dev: crate::config::DevConfig::default(),
335 outbox: crate::config::OutboxConfig::default(),
336 observability: crate::config::ObservabilityConfig::with_flush_policy(64, 0),
337 mcp: crate::config::ResolvedMcpConfig::default(),
338 assistant: crate::config::ResolvedAssistantConfig::default(),
339 scheduler_threads: 1,
340 stop_drain_timeout: Some(std::time::Duration::from_secs(5)),
341 jit_threshold: None,
342 query_timeout: Some(std::time::Duration::from_secs(10)),
343 workloop_sweep_interval: Some(std::time::Duration::from_millis(50)),
344 default_namespace: "default".to_owned(),
345 auto_create: crate::config::AutoCreate::Open,
346 max_in_flight_activities: crate::config::DEFAULT_MAX_IN_FLIGHT_ACTIVITIES,
347 drain_timeout: std::time::Duration::from_secs(30),
348 metrics: MetricsConfig { enabled: true },
349 owned_shards: Vec::new(),
350 cors_allowed_origins: Vec::new(),
351 }
352 }
353
354 async fn deploy_state() -> Result<StateUnderTest, Box<dyn std::error::Error>> {
355 let store: Arc<dyn EventStore> = Arc::new(InMemoryStore::default());
356 let engine = EngineUnderTest::new(Arc::new(
357 EngineBuilder::new()
358 .stop_drain_timeout(std::time::Duration::from_secs(5))
359 .store_arc(store)
360 .in_memory_visibility()
361 .scheduler_threads(1)
362 .build()
363 .await?,
364 ));
365 let resolver = NamespaceResolver::from_parts(
366 NamespaceMode::SharedEngine,
367 Some(engine.handle()),
368 Arc::new(StaticWorkflowNamespaces::default()),
369 Arc::new(StaticScheduleNamespaces::default()),
370 );
371 let mut config = runtime_config();
372 config.deploy = DeployConfig {
373 enabled: true,
374 max_archive_bytes: Some(1024),
375 max_inflated_bytes: Some(2048),
376 };
377 Ok(StateUnderTest::over(
378 engine,
379 ServerState::from_parts(resolver, config),
380 ))
381 }
382
383 #[cfg(not(feature = "auth"))]
386 const AUTH_TOKEN: &str = "deploy-secret";
387
388 #[cfg(not(feature = "auth"))]
391 async fn auth_on_deploy_state() -> Result<StateUnderTest, Box<dyn std::error::Error>> {
392 let store: Arc<dyn EventStore> = Arc::new(InMemoryStore::default());
393 let engine = EngineUnderTest::new(Arc::new(
394 EngineBuilder::new()
395 .stop_drain_timeout(std::time::Duration::from_secs(5))
396 .store_arc(store)
397 .in_memory_visibility()
398 .scheduler_threads(1)
399 .build()
400 .await?,
401 ));
402 let resolver = NamespaceResolver::from_parts(
403 NamespaceMode::SharedEngine,
404 Some(engine.handle()),
405 Arc::new(StaticWorkflowNamespaces::default()),
406 Arc::new(StaticScheduleNamespaces::default()),
407 );
408 let mut config = runtime_config();
409 config.auth = AuthConfig {
410 enabled: true,
411 jwks_url: Some(AUTH_TOKEN.to_owned()),
412 jwks_refresh_seconds: 300,
413 };
414 config.deploy = DeployConfig {
415 enabled: true,
416 max_archive_bytes: Some(1024),
417 max_inflated_bytes: Some(2048),
418 };
419 Ok(StateUnderTest::over(
420 engine,
421 ServerState::from_parts(resolver, config),
422 ))
423 }
424
425 fn granted_request<T>(message: T) -> Result<Request<T>, Box<dyn std::error::Error>> {
426 let mut request = Request::new(message);
427 request
428 .metadata_mut()
429 .insert("x-aion-subject", "ci".parse()?);
430 request
431 .metadata_mut()
432 .insert("x-aion-deploy", "true".parse()?);
433 Ok(request)
434 }
435
436 #[tokio::test]
440 async fn auth_off_operator_is_deploy_granted_without_metadata()
441 -> Result<(), Box<dyn std::error::Error>> {
442 use generated::deploy_service_server::DeployService as _;
443
444 let state = deploy_state().await?;
445 let service = DeployGrpcService::new(state.clone());
446 let mut request = Request::new(generated::ListVersionsRequest {});
447 request
448 .metadata_mut()
449 .insert("x-aion-subject", "ci".parse()?);
450
451 let response = service.list_versions(request).await?;
452 assert!(response.into_inner().versions.is_empty());
453 Ok(())
454 }
455
456 #[cfg(not(feature = "auth"))]
460 #[tokio::test]
461 async fn auth_on_denies_caller_without_deploy_grant() -> Result<(), Box<dyn std::error::Error>>
462 {
463 use generated::deploy_service_server::DeployService as _;
464
465 let state = auth_on_deploy_state().await?;
466 let service = DeployGrpcService::new(state.clone());
467 let mut request = Request::new(generated::ListVersionsRequest {});
468 request
470 .metadata_mut()
471 .insert("authorization", format!("Bearer {AUTH_TOKEN}").parse()?);
472 request
473 .metadata_mut()
474 .insert("x-aion-subject", "ci".parse()?);
475
476 let status = service
477 .list_versions(request)
478 .await
479 .err()
480 .ok_or("expected denial")?;
481 assert_eq!(status.code(), Code::PermissionDenied);
482 let detail = decode_detail(&status)?;
483 assert_eq!(detail.code, WireErrorCode::DeployDenied);
484 assert!(
485 detail.message.contains("x-aion-deploy"),
486 "denial must hint the dev header: {}",
487 detail.message
488 );
489 Ok(())
490 }
491
492 #[tokio::test]
493 async fn granted_metadata_lists_versions() -> Result<(), Box<dyn std::error::Error>> {
494 use generated::deploy_service_server::DeployService as _;
495
496 let state = deploy_state().await?;
497 let service = DeployGrpcService::new(state.clone());
498 let response = service
499 .list_versions(granted_request(generated::ListVersionsRequest {})?)
500 .await?;
501 assert!(response.into_inner().versions.is_empty());
502 Ok(())
503 }
504
505 #[tokio::test]
506 async fn oversized_archive_is_invalid_argument_naming_the_key()
507 -> Result<(), Box<dyn std::error::Error>> {
508 use generated::deploy_service_server::DeployService as _;
509
510 let state = deploy_state().await?;
511 let service = DeployGrpcService::new(state.clone());
512 let status = service
513 .load_package(granted_request(generated::LoadPackageRequest {
514 archive: vec![0_u8; 2048],
515 })?)
516 .await
517 .err()
518 .ok_or("expected oversize refusal")?;
519
520 assert_eq!(status.code(), Code::InvalidArgument);
521 assert!(
522 status.message().contains("deploy.max_archive_bytes"),
523 "refusal must name the config key: {}",
524 status.message()
525 );
526 Ok(())
527 }
528
529 #[tokio::test]
530 async fn route_to_unknown_version_is_not_found() -> Result<(), Box<dyn std::error::Error>> {
531 use generated::deploy_service_server::DeployService as _;
532
533 let state = deploy_state().await?;
534 let service = DeployGrpcService::new(state.clone());
535 let status = service
536 .route_version(granted_request(generated::RouteVersionRequest {
537 workflow_type: "order".to_owned(),
538 content_hash: "a".repeat(64),
539 })?)
540 .await
541 .err()
542 .ok_or("expected unknown-version refusal")?;
543
544 assert_eq!(status.code(), Code::NotFound);
545 let detail = decode_detail(&status)?;
546 assert_eq!(detail.code, WireErrorCode::NotFound);
547 assert_eq!(detail.error_type.as_deref(), Some("UnknownVersion"));
548 Ok(())
549 }
550
551 #[tokio::test]
552 async fn malformed_hash_is_invalid_argument() -> Result<(), Box<dyn std::error::Error>> {
553 use generated::deploy_service_server::DeployService as _;
554
555 let state = deploy_state().await?;
556 let service = DeployGrpcService::new(state.clone());
557 let status = service
558 .unload_version(granted_request(generated::UnloadVersionRequest {
559 workflow_type: "order".to_owned(),
560 content_hash: "not-a-hash".to_owned(),
561 })?)
562 .await
563 .err()
564 .ok_or("expected malformed-hash refusal")?;
565
566 assert_eq!(status.code(), Code::InvalidArgument);
567 assert!(
568 status.message().contains("not-a-hash"),
569 "refusal must name the malformed hash: {}",
570 status.message()
571 );
572 Ok(())
573 }
574
575 #[tokio::test]
578 async fn drain_refuses_mutations_but_serves_listing() -> Result<(), Box<dyn std::error::Error>>
579 {
580 use generated::deploy_service_server::DeployService as _;
581
582 let state = deploy_state().await?;
583 assert!(state.drain_state().begin());
584 let service = DeployGrpcService::new(state.clone());
585
586 let status = service
587 .route_version(granted_request(generated::RouteVersionRequest {
588 workflow_type: "order".to_owned(),
589 content_hash: "a".repeat(64),
590 })?)
591 .await
592 .err()
593 .ok_or("expected drain refusal")?;
594 assert_eq!(status.code(), Code::Unavailable);
595 assert!(
596 status.message().contains("draining"),
597 "drain refusal must be explicit: {}",
598 status.message()
599 );
600
601 let listing = service
602 .list_versions(granted_request(generated::ListVersionsRequest {})?)
603 .await?;
604 assert!(listing.into_inner().versions.is_empty());
605 Ok(())
606 }
607}