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