1use std::sync::Arc;
4use std::time::Duration;
5
6use aion_core::{Payload, WorkflowId};
7use tokio::sync::oneshot;
8use tokio::time;
9
10use crate::engine_seam::{
11 EngineHandle, EngineSeamError, WorkflowMailboxMessage, WorkflowProcessHandle, WorkflowResidency,
12};
13
14pub type QueryResult = Result<Payload, QueryError>;
16
17pub type QueryServiceResult = Result<Payload, QueryError>;
19
20#[derive(thiserror::Error, Debug, Clone, PartialEq, Eq)]
22pub enum QueryError {
23 #[error("unknown query {0}")]
25 UnknownQuery(String),
26
27 #[error("query reply timed out")]
29 Timeout,
30
31 #[error("workflow {0} is not running")]
33 NotRunning(WorkflowId),
34
35 #[error("workflow {0} is unknown")]
37 Unknown(WorkflowId),
38
39 #[error("query reply channel closed before a handler response was sent")]
41 ReplyDropped,
42
43 #[error("query handler failed: {message}")]
45 HandlerFailed {
46 message: String,
48 },
49
50 #[error("query arguments are invalid: {reason}")]
57 InvalidArguments {
58 reason: String,
60 },
61
62 #[error("query engine seam failed: {0}")]
64 Engine(#[from] EngineSeamError),
65}
66
67#[derive(Debug)]
76pub struct QueryService<H: ?Sized> {
77 engine: Arc<H>,
78 query_timeout: Duration,
79}
80
81impl<H> QueryService<H>
82where
83 H: EngineHandle + ?Sized,
84{
85 #[must_use]
87 pub fn new(engine: Arc<H>, query_timeout: Duration) -> Self {
88 Self {
89 engine,
90 query_timeout,
91 }
92 }
93
94 pub async fn query(
103 &self,
104 workflow_id: &WorkflowId,
105 name: impl Into<String>,
106 args: Payload,
107 ) -> QueryServiceResult {
108 let process = match self.engine.resolve_workflow(workflow_id)? {
109 WorkflowResidency::Resident(process) => process,
110 WorkflowResidency::NonResident | WorkflowResidency::Terminal => {
111 return Err(QueryError::NotRunning(workflow_id.clone()));
112 }
113 WorkflowResidency::Unknown => return Err(QueryError::Unknown(workflow_id.clone())),
114 };
115 self.query_process(process, name, args).await
116 }
117
118 pub async fn query_process(
133 &self,
134 process: WorkflowProcessHandle,
135 name: impl Into<String>,
136 args: Payload,
137 ) -> QueryServiceResult {
138 let (reply_to, reply_from) = oneshot::channel();
139 self.engine.deliver_workflow_message(
140 process,
141 WorkflowMailboxMessage::Query {
142 name: name.into(),
143 payload: args,
144 reply_to,
145 },
146 )?;
147
148 match time::timeout(self.query_timeout, reply_from).await {
149 Ok(Ok(reply)) => reply,
150 Ok(Err(_)) => Err(QueryError::ReplyDropped),
151 Err(_) => Err(QueryError::Timeout),
152 }
153 }
154
155 #[must_use]
157 pub const fn query_timeout(&self) -> Duration {
158 self.query_timeout
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use std::collections::HashMap;
165 use std::sync::{Arc, Mutex, MutexGuard};
166 use std::time::Duration;
167
168 use aion_core::{ContentType, Event, Payload, TimerId, WorkflowId};
169 use aion_store::{InMemoryStore, ReadableEventStore};
170
171 use super::{QueryError, QueryService};
172 use crate::Pid;
173 use crate::engine_seam::{
174 ChildWorkflowSpawnRequest, ChildWorkflowSpawnResult, EngineHandle, EngineSeamError,
175 TimerWheelEntry, WorkflowMailboxMessage, WorkflowProcessHandle, WorkflowResidency,
176 };
177
178 const QUERY_TIMEOUT: Duration = Duration::from_millis(10);
179
180 #[derive(Clone)]
181 enum QueryBehavior {
182 Reply(Payload),
183 Fail(String),
184 HoldSender,
185 }
186
187 #[derive(Default)]
188 struct FakeQueryWorkflow {
189 handlers: HashMap<String, QueryBehavior>,
190 query_count: usize,
191 last_payload: Option<Payload>,
192 }
193
194 #[derive(Default)]
195 struct FakeQueryEngineState {
196 residency: HashMap<WorkflowId, WorkflowResidency>,
197 workflows: HashMap<WorkflowProcessHandle, FakeQueryWorkflow>,
198 held_replies: Vec<crate::engine_seam::QueryReplySender>,
199 }
200
201 #[derive(Default)]
202 struct FakeQueryEngine {
203 state: Mutex<FakeQueryEngineState>,
204 }
205
206 impl FakeQueryEngine {
207 fn set_resident_workflow(
208 &self,
209 workflow_id: WorkflowId,
210 process: WorkflowProcessHandle,
211 workflow: FakeQueryWorkflow,
212 ) -> Result<(), EngineSeamError> {
213 let mut state = self.state()?;
214 state
215 .residency
216 .insert(workflow_id, WorkflowResidency::Resident(process));
217 state.workflows.insert(process, workflow);
218 Ok(())
219 }
220
221 fn set_residency(
222 &self,
223 workflow_id: WorkflowId,
224 residency: WorkflowResidency,
225 ) -> Result<(), EngineSeamError> {
226 self.state()?.residency.insert(workflow_id, residency);
227 Ok(())
228 }
229
230 fn query_count(&self, process: WorkflowProcessHandle) -> Result<usize, EngineSeamError> {
231 Ok(self
232 .state()?
233 .workflows
234 .get(&process)
235 .map_or(0, |workflow| workflow.query_count))
236 }
237
238 fn last_payload(
239 &self,
240 process: WorkflowProcessHandle,
241 ) -> Result<Option<Payload>, EngineSeamError> {
242 Ok(self
243 .state()?
244 .workflows
245 .get(&process)
246 .and_then(|workflow| workflow.last_payload.clone()))
247 }
248
249 fn state(&self) -> Result<MutexGuard<'_, FakeQueryEngineState>, EngineSeamError> {
250 self.state.lock().map_err(|_| EngineSeamError::Delivery {
251 reason: "fake query engine state lock was poisoned".to_owned(),
252 })
253 }
254 }
255
256 impl EngineHandle for FakeQueryEngine {
257 fn resolve_workflow(
258 &self,
259 workflow_id: &WorkflowId,
260 ) -> Result<WorkflowResidency, EngineSeamError> {
261 Ok(self
262 .state()?
263 .residency
264 .get(workflow_id)
265 .copied()
266 .unwrap_or(WorkflowResidency::Unknown))
267 }
268
269 fn deliver_workflow_message(
270 &self,
271 process: WorkflowProcessHandle,
272 message: WorkflowMailboxMessage,
273 ) -> Result<(), EngineSeamError> {
274 match message {
275 WorkflowMailboxMessage::Query {
276 name,
277 payload,
278 reply_to,
279 } => {
280 let mut state = self.state()?;
281 let behavior = {
282 let workflow = state.workflows.get_mut(&process).ok_or_else(|| {
283 EngineSeamError::Delivery {
284 reason: "query target process was not registered".to_owned(),
285 }
286 })?;
287 workflow.last_payload = Some(payload);
288 workflow.query_count += 1;
289 workflow.handlers.get(&name).cloned()
290 };
291
292 match behavior {
293 Some(QueryBehavior::Reply(payload)) => {
294 if reply_to.send(Ok(payload)).is_err() {
295 return Err(EngineSeamError::Delivery {
296 reason: "query caller dropped reply receiver".to_owned(),
297 });
298 }
299 }
300 Some(QueryBehavior::Fail(message)) => {
301 if reply_to
302 .send(Err(QueryError::HandlerFailed { message }))
303 .is_err()
304 {
305 return Err(EngineSeamError::Delivery {
306 reason: "query caller dropped reply receiver".to_owned(),
307 });
308 }
309 }
310 None => {
311 if reply_to.send(Err(QueryError::UnknownQuery(name))).is_err() {
312 return Err(EngineSeamError::Delivery {
313 reason: "query caller dropped reply receiver".to_owned(),
314 });
315 }
316 }
317 Some(QueryBehavior::HoldSender) => state.held_replies.push(reply_to),
318 }
319 Ok(())
320 }
321 _ => Err(EngineSeamError::Delivery {
322 reason: "fake query engine only accepts query messages".to_owned(),
323 }),
324 }
325 }
326
327 fn spawn_child_workflow(
328 &self,
329 request: ChildWorkflowSpawnRequest,
330 ) -> Result<ChildWorkflowSpawnResult, EngineSeamError> {
331 Err(EngineSeamError::ChildSpawn {
332 reason: format!(
333 "fake query engine does not spawn child workflow {}",
334 request.workflow_type
335 ),
336 })
337 }
338
339 fn terminate_linked_child_workflow(
340 &self,
341 parent_workflow_id: &WorkflowId,
342 child_process: WorkflowProcessHandle,
343 correlation: u64,
344 ) -> Result<(), EngineSeamError> {
345 Err(EngineSeamError::ChildTermination {
346 reason: format!(
347 "fake query engine does not terminate child workflow process {} for parent {parent_workflow_id} with correlation {correlation}",
348 child_process.pid()
349 ),
350 })
351 }
352
353 fn terminate_linked_activity(
354 &self,
355 parent_workflow_id: &WorkflowId,
356 activity_process: Pid,
357 correlation: u64,
358 ) -> Result<(), EngineSeamError> {
359 Err(EngineSeamError::ChildTermination {
360 reason: format!(
361 "fake query engine does not terminate activity process {activity_process} for parent {parent_workflow_id} with correlation {correlation}"
362 ),
363 })
364 }
365
366 fn arm_timer(&self, entry: TimerWheelEntry) -> Result<(), EngineSeamError> {
367 Err(EngineSeamError::TimerWheel {
368 reason: format!("fake query engine does not arm timer {}", entry.timer_id),
369 })
370 }
371
372 fn disarm_timer(
373 &self,
374 process: WorkflowProcessHandle,
375 timer_id: &TimerId,
376 ) -> Result<(), EngineSeamError> {
377 Err(EngineSeamError::TimerWheel {
378 reason: format!(
379 "fake query engine does not disarm timer {timer_id} for process {}",
380 process.pid()
381 ),
382 })
383 }
384
385 fn record_workflow_event(
386 &self,
387 workflow_id: &WorkflowId,
388 event: Event,
389 ) -> Result<crate::engine_seam::RecordOutcome, EngineSeamError> {
390 Err(EngineSeamError::Recorder {
391 reason: format!(
392 "queries must not record event {} for workflow {workflow_id}",
393 event.seq()
394 ),
395 })
396 }
397
398 fn record_redelivered_timer_fire(
399 &self,
400 workflow_id: &WorkflowId,
401 timer_id: &aion_core::TimerId,
402 ) -> Result<crate::engine_seam::RedeliveredFire, EngineSeamError> {
403 Err(EngineSeamError::Recorder {
404 reason: format!(
405 "queries must not answer timer redelivery of `{timer_id}` for {workflow_id}"
406 ),
407 })
408 }
409 }
410
411 fn payload(label: &str) -> Payload {
412 Payload::new(
413 ContentType::Json,
414 format!("{{\"label\":\"{label}\"}}").into_bytes(),
415 )
416 }
417
418 fn known_workflow(reply: Payload) -> FakeQueryWorkflow {
419 let mut handlers = HashMap::new();
420 handlers.insert("state".to_owned(), QueryBehavior::Reply(reply));
421 FakeQueryWorkflow {
422 handlers,
423 query_count: 0,
424 last_payload: None,
425 }
426 }
427
428 #[tokio::test]
429 async fn query_returns_registered_handler_reply() -> Result<(), Box<dyn std::error::Error>> {
430 let engine = Arc::new(FakeQueryEngine::default());
431 let workflow_id = WorkflowId::new_v4();
432 let process = WorkflowProcessHandle::new(7);
433 let reply = payload("answer");
434 engine.set_resident_workflow(
435 workflow_id.clone(),
436 process,
437 known_workflow(reply.clone()),
438 )?;
439 let service = QueryService::new(Arc::clone(&engine), QUERY_TIMEOUT);
440
441 let returned = service
442 .query(&workflow_id, "state", payload("args"))
443 .await?;
444
445 assert_eq!(returned, reply);
446 assert_eq!(engine.query_count(process)?, 1);
447 assert_eq!(engine.last_payload(process)?, Some(payload("args")));
448 Ok(())
449 }
450
451 #[tokio::test]
452 async fn query_does_not_record_events() -> Result<(), Box<dyn std::error::Error>> {
453 let store = InMemoryStore::default();
454 let engine = Arc::new(FakeQueryEngine::default());
455 let workflow_id = WorkflowId::new_v4();
456 let process = WorkflowProcessHandle::new(8);
457 engine.set_resident_workflow(
458 workflow_id.clone(),
459 process,
460 known_workflow(payload("visible-state")),
461 )?;
462 let service = QueryService::new(engine, QUERY_TIMEOUT);
463
464 let reply = service
465 .query(&workflow_id, "state", payload("args"))
466 .await?;
467 assert_eq!(reply, payload("visible-state"));
468
469 let history = store.read_history(&workflow_id).await?;
470 assert!(history.is_empty());
471 Ok(())
472 }
473
474 #[tokio::test]
475 async fn unknown_query_returns_typed_error_and_workflow_remains_live()
476 -> Result<(), Box<dyn std::error::Error>> {
477 let engine = Arc::new(FakeQueryEngine::default());
478 let workflow_id = WorkflowId::new_v4();
479 let process = WorkflowProcessHandle::new(9);
480 engine.set_resident_workflow(
481 workflow_id.clone(),
482 process,
483 known_workflow(payload("known")),
484 )?;
485 let service = QueryService::new(Arc::clone(&engine), QUERY_TIMEOUT);
486
487 let result = service
488 .query(&workflow_id, "missing", payload("args"))
489 .await;
490
491 assert_eq!(result, Err(QueryError::UnknownQuery("missing".to_owned())));
492 assert_eq!(
493 engine.resolve_workflow(&workflow_id)?,
494 WorkflowResidency::Resident(process)
495 );
496 assert_eq!(engine.query_count(process)?, 1);
497 Ok(())
498 }
499
500 #[tokio::test]
501 async fn non_replying_workflow_times_out() -> Result<(), Box<dyn std::error::Error>> {
502 let engine = Arc::new(FakeQueryEngine::default());
503 let workflow_id = WorkflowId::new_v4();
504 let process = WorkflowProcessHandle::new(10);
505 let mut handlers = HashMap::new();
506 handlers.insert("slow".to_owned(), QueryBehavior::HoldSender);
507 engine.set_resident_workflow(
508 workflow_id.clone(),
509 process,
510 FakeQueryWorkflow {
511 handlers,
512 query_count: 0,
513 last_payload: None,
514 },
515 )?;
516 let service = QueryService::new(engine, QUERY_TIMEOUT);
517
518 let result = service.query(&workflow_id, "slow", payload("args")).await;
519
520 assert_eq!(result, Err(QueryError::Timeout));
521 Ok(())
522 }
523
524 #[tokio::test]
525 async fn terminal_and_non_resident_workflows_are_not_running()
526 -> Result<(), Box<dyn std::error::Error>> {
527 let engine = Arc::new(FakeQueryEngine::default());
528 let terminal_id = WorkflowId::new_v4();
529 let non_resident_id = WorkflowId::new_v4();
530 engine.set_residency(terminal_id.clone(), WorkflowResidency::Terminal)?;
531 engine.set_residency(non_resident_id.clone(), WorkflowResidency::NonResident)?;
532 let service = QueryService::new(engine, QUERY_TIMEOUT);
533
534 let terminal_result = service.query(&terminal_id, "state", payload("args")).await;
535 let non_resident_result = service
536 .query(&non_resident_id, "state", payload("args"))
537 .await;
538
539 assert_eq!(terminal_result, Err(QueryError::NotRunning(terminal_id)));
540 assert_eq!(
541 non_resident_result,
542 Err(QueryError::NotRunning(non_resident_id))
543 );
544 Ok(())
545 }
546
547 #[tokio::test]
548 async fn query_process_dispatches_to_the_resolved_process_without_resolving()
549 -> Result<(), Box<dyn std::error::Error>> {
550 let engine = Arc::new(FakeQueryEngine::default());
551 let workflow_id = WorkflowId::new_v4();
554 let process = WorkflowProcessHandle::new(11);
555 let reply = payload("run-exact");
556 engine.set_resident_workflow(
557 workflow_id.clone(),
558 process,
559 known_workflow(reply.clone()),
560 )?;
561 engine.set_residency(workflow_id, WorkflowResidency::Unknown)?;
562 let service = QueryService::new(Arc::clone(&engine), QUERY_TIMEOUT);
563
564 let returned = service
565 .query_process(process, "state", payload("args"))
566 .await?;
567
568 assert_eq!(returned, reply);
569 assert_eq!(engine.query_count(process)?, 1);
570 Ok(())
571 }
572
573 #[tokio::test]
574 async fn handler_failure_propagates_as_typed_handler_failed()
575 -> Result<(), Box<dyn std::error::Error>> {
576 let engine = Arc::new(FakeQueryEngine::default());
577 let workflow_id = WorkflowId::new_v4();
578 let process = WorkflowProcessHandle::new(12);
579 let mut handlers = HashMap::new();
580 handlers.insert(
581 "state".to_owned(),
582 QueryBehavior::Fail("handler raised".to_owned()),
583 );
584 engine.set_resident_workflow(
585 workflow_id.clone(),
586 process,
587 FakeQueryWorkflow {
588 handlers,
589 query_count: 0,
590 last_payload: None,
591 },
592 )?;
593 let service = QueryService::new(Arc::clone(&engine), QUERY_TIMEOUT);
594
595 let resolved = service.query(&workflow_id, "state", payload("args")).await;
596 let run_exact = service
597 .query_process(process, "state", payload("args"))
598 .await;
599
600 let expected = Err(QueryError::HandlerFailed {
601 message: "handler raised".to_owned(),
602 });
603 assert_eq!(resolved, expected);
604 assert_eq!(run_exact, expected);
605 Ok(())
606 }
607
608 #[tokio::test]
609 async fn unknown_workflow_returns_typed_unknown_error() -> Result<(), Box<dyn std::error::Error>>
610 {
611 let engine = Arc::new(FakeQueryEngine::default());
612 let workflow_id = WorkflowId::new_v4();
613 let service = QueryService::new(engine, QUERY_TIMEOUT);
614
615 let result = service.query(&workflow_id, "state", payload("args")).await;
616
617 assert_eq!(result, Err(QueryError::Unknown(workflow_id)));
618 Ok(())
619 }
620}