1use std::num::NonZeroU64;
4use std::pin::Pin;
5use std::sync::Arc;
6use std::task::{Context, Poll};
7
8use aion_core::{Event, WorkflowFilter, WorkflowId};
9use aion_proto::{
10 FilteredSubscription, FirehoseSubscription, PerWorkflowSubscription, ProtoWorkflowId,
11 SubscriptionRequest, subscription_request,
12};
13use futures::Stream;
14use futures::future::BoxFuture;
15use futures::stream::BoxStream;
16
17use crate::error::ClientError;
18use crate::transport::{SubscriptionAttempt, WorkflowTransport};
19
20pub type EventStream = Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>;
22
23#[derive(Clone, Debug, PartialEq, Eq)]
25pub enum SubscribeTarget {
26 Workflow {
28 workflow_id: WorkflowId,
30 },
31 Filtered {
33 filter: WorkflowFilter,
35 },
36 Firehose,
38}
39
40impl SubscribeTarget {
41 pub(crate) fn request(&self, namespace: &str) -> SubscriptionRequest {
42 match self {
43 Self::Workflow { workflow_id } => SubscriptionRequest {
44 subscription: Some(subscription_request::Subscription::PerWorkflow(
45 PerWorkflowSubscription {
46 namespace: namespace.to_owned(),
47 workflow_id: Some(ProtoWorkflowId::from(workflow_id.clone())),
48 resume_from_seq: None,
49 },
50 )),
51 },
52 Self::Filtered { filter } => SubscriptionRequest {
53 subscription: Some(subscription_request::Subscription::Filtered(
54 FilteredSubscription {
55 namespace: namespace.to_owned(),
56 workflow_type: filter.workflow_type.clone(),
57 status: filter
58 .status
59 .map(|status| aion_proto::ProtoWorkflowStatus::from(status) as i32),
60 namespace_selector: None,
61 },
62 )),
63 },
64 Self::Firehose => SubscriptionRequest {
65 subscription: Some(subscription_request::Subscription::Firehose(
66 FirehoseSubscription {
67 namespace: namespace.to_owned(),
68 },
69 )),
70 },
71 }
72 }
73}
74
75pub struct ResumingEventStream {
94 transport: Arc<dyn WorkflowTransport>,
95 namespace: String,
96 target: SubscribeTarget,
97 last_seq: Option<u64>,
98 delivered_any: bool,
99 current: Option<BoxStream<'static, Result<Event, ClientError>>>,
100 pending_subscribe: Option<BoxFuture<'static, Result<SubscriptionAttempt, ClientError>>>,
101 terminal_error: Option<ClientError>,
102 finished: bool,
103}
104
105impl ResumingEventStream {
106 #[must_use]
108 pub fn new(
109 transport: Arc<dyn WorkflowTransport>,
110 namespace: impl Into<String>,
111 target: SubscribeTarget,
112 ) -> Self {
113 Self {
114 transport,
115 namespace: namespace.into(),
116 target,
117 last_seq: None,
118 delivered_any: false,
119 current: None,
120 pending_subscribe: None,
121 terminal_error: None,
122 finished: false,
123 }
124 }
125
126 #[must_use]
134 pub fn from_sequence(
135 transport: Arc<dyn WorkflowTransport>,
136 namespace: impl Into<String>,
137 workflow_id: WorkflowId,
138 resume_from: NonZeroU64,
139 ) -> Self {
140 let mut stream = Self::new(
141 transport,
142 namespace,
143 SubscribeTarget::Workflow { workflow_id },
144 );
145 stream.last_seq = Some(resume_from.get() - 1);
149 stream
150 }
151
152 fn is_per_workflow(&self) -> bool {
153 matches!(self.target, SubscribeTarget::Workflow { .. })
154 }
155
156 fn start_subscribe(&mut self) {
157 let transport = Arc::clone(&self.transport);
158 let request = self.target.request(&self.namespace);
159 let resume_from_sequence = if self.is_per_workflow() {
162 self.last_seq.map(|seq| seq.saturating_add(1))
163 } else {
164 None
165 };
166 self.pending_subscribe = Some(Box::pin(async move {
167 transport.subscribe(request, resume_from_sequence).await
168 }));
169 }
170}
171
172impl Stream for ResumingEventStream {
173 type Item = Result<Event, ClientError>;
174
175 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
176 let this = self.get_mut();
177 loop {
178 if this.finished {
179 return Poll::Ready(None);
180 }
181
182 if let Some(error) = this.terminal_error.take() {
183 this.finished = true;
184 return Poll::Ready(Some(Err(error)));
185 }
186
187 if this.current.is_none() && this.pending_subscribe.is_none() {
188 this.start_subscribe();
189 }
190
191 if let Some(pending) = this.pending_subscribe.as_mut() {
192 match pending.as_mut().poll(cx) {
193 Poll::Pending => return Poll::Pending,
194 Poll::Ready(Ok(attempt)) => {
195 this.pending_subscribe = None;
196 this.current = Some(attempt.events);
197 }
198 Poll::Ready(Err(error)) => {
199 this.pending_subscribe = None;
206 if is_retryable(&error) && (this.is_per_workflow() || !this.delivered_any) {
207 continue;
208 }
209 this.finished = true;
210 return Poll::Ready(Some(Err(error)));
211 }
212 }
213 }
214
215 let Some(current) = this.current.as_mut() else {
216 continue;
217 };
218 match current.as_mut().poll_next(cx) {
219 Poll::Pending => return Poll::Pending,
220 Poll::Ready(Some(Ok(event))) => {
221 if this.is_per_workflow() {
222 if this.last_seq.is_some_and(|seq| event.seq() <= seq) {
225 continue;
226 }
227 this.last_seq = Some(event.seq());
228 }
229 this.delivered_any = true;
230 return Poll::Ready(Some(Ok(event)));
231 }
232 Poll::Ready(Some(Err(error))) => {
233 this.current = None;
234 if is_retryable(&error) {
235 if this.is_per_workflow() {
236 continue;
237 }
238 if !this.delivered_any {
239 continue;
242 }
243 }
247 this.terminal_error = Some(error);
248 }
249 Poll::Ready(None) => {
250 this.current = None;
251 this.finished = true;
252 return Poll::Ready(None);
253 }
254 }
255 }
256 }
257}
258
259#[must_use]
261pub fn event_stream(
262 transport: Arc<dyn WorkflowTransport>,
263 namespace: impl Into<String>,
264 target: SubscribeTarget,
265) -> EventStream {
266 Box::pin(ResumingEventStream::new(transport, namespace, target))
267}
268
269#[must_use]
271pub fn event_stream_from(
272 transport: Arc<dyn WorkflowTransport>,
273 namespace: impl Into<String>,
274 workflow_id: WorkflowId,
275 resume_from: NonZeroU64,
276) -> EventStream {
277 Box::pin(ResumingEventStream::from_sequence(
278 transport,
279 namespace,
280 workflow_id,
281 resume_from,
282 ))
283}
284
285fn is_retryable(error: &ClientError) -> bool {
286 matches!(error, ClientError::Unavailable { .. })
287}
288
289#[cfg(test)]
290mod tests {
291 use std::collections::VecDeque;
292 use std::sync::Arc;
293
294 use aion_core::{ContentType, Event, EventEnvelope, Payload, WorkflowId};
295 use aion_proto::{
296 ProtoCancelResponse, ProtoDescribeWorkflowResponse, ProtoListWorkflowsResponse,
297 ProtoQueryResponse, ProtoSignalResponse, ProtoStartWorkflowResponse,
298 };
299 use async_trait::async_trait;
300 use chrono::Utc;
301 use futures::StreamExt;
302 use futures::stream;
303 use tokio::sync::Mutex;
304
305 use super::{ResumingEventStream, SubscribeTarget};
306 use crate::error::ClientError;
307 use crate::transport::{SubscriptionAttempt, WorkflowTransport};
308
309 #[derive(Default)]
310 struct SubscribeStub {
311 attach_failures: Mutex<VecDeque<ClientError>>,
314 attempts: Mutex<VecDeque<SubscriptionAttempt>>,
315 resume_points: Mutex<Vec<Option<u64>>>,
316 }
317
318 #[async_trait]
319 impl WorkflowTransport for SubscribeStub {
320 async fn start_workflow(
321 &self,
322 _: aion_proto::ProtoStartWorkflowRequest,
323 ) -> Result<ProtoStartWorkflowResponse, ClientError> {
324 Err(ClientError::unavailable("stub transport"))
325 }
326
327 async fn signal(
328 &self,
329 _: aion_proto::ProtoSignalRequest,
330 ) -> Result<ProtoSignalResponse, ClientError> {
331 Err(ClientError::unavailable("stub transport"))
332 }
333
334 async fn query(
335 &self,
336 _: aion_proto::ProtoQueryRequest,
337 ) -> Result<ProtoQueryResponse, ClientError> {
338 Err(ClientError::unavailable("stub transport"))
339 }
340
341 async fn cancel(
342 &self,
343 _: aion_proto::ProtoCancelRequest,
344 ) -> Result<ProtoCancelResponse, ClientError> {
345 Err(ClientError::unavailable("stub transport"))
346 }
347
348 async fn reopen(
349 &self,
350 _: aion_proto::ProtoReopenRequest,
351 ) -> Result<aion_proto::ProtoReopenResponse, ClientError> {
352 Err(ClientError::unavailable("stub transport"))
353 }
354
355 async fn pause(
356 &self,
357 _: aion_proto::ProtoPauseRequest,
358 ) -> Result<aion_proto::ProtoPauseResponse, ClientError> {
359 Err(ClientError::unavailable("stub transport"))
360 }
361
362 async fn resume(
363 &self,
364 _: aion_proto::ProtoResumeRequest,
365 ) -> Result<aion_proto::ProtoResumeResponse, ClientError> {
366 Err(ClientError::unavailable("stub transport"))
367 }
368
369 async fn list_workflows(
370 &self,
371 _: aion_proto::ProtoListWorkflowsRequest,
372 ) -> Result<ProtoListWorkflowsResponse, ClientError> {
373 Err(ClientError::unavailable("stub transport"))
374 }
375
376 async fn describe_workflow(
377 &self,
378 _: aion_proto::ProtoDescribeWorkflowRequest,
379 ) -> Result<ProtoDescribeWorkflowResponse, ClientError> {
380 Err(ClientError::unavailable("stub transport"))
381 }
382
383 async fn subscribe(
384 &self,
385 _: aion_proto::SubscriptionRequest,
386 resume_from_sequence: Option<u64>,
387 ) -> Result<SubscriptionAttempt, ClientError> {
388 self.resume_points.lock().await.push(resume_from_sequence);
389 if let Some(failure) = self.attach_failures.lock().await.pop_front() {
390 return Err(failure);
391 }
392 self.attempts
393 .lock()
394 .await
395 .pop_front()
396 .ok_or_else(|| ClientError::server("missing subscribe attempt"))
397 }
398 }
399
400 fn event(seq: u64, workflow_id: &WorkflowId) -> Event {
401 Event::WorkflowStarted {
402 envelope: EventEnvelope {
403 seq,
404 recorded_at: Utc::now(),
405 workflow_id: workflow_id.clone(),
406 },
407 workflow_type: String::from("checkout"),
408 input: Payload::new(ContentType::Json, Vec::new()),
409 run_id: aion_core::RunId::new(uuid::Uuid::from_u128(1)),
410 parent_run_id: None,
411 package_version: aion_core::PackageVersion::new("a".repeat(64)),
412 }
413 }
414
415 #[tokio::test]
416 async fn resumes_after_transient_disconnect_without_gaps_or_duplicates() {
417 let workflow_id = WorkflowId::new_v4();
418 let stub = Arc::new(SubscribeStub::default());
419 stub.attempts
420 .lock()
421 .await
422 .push_back(SubscriptionAttempt::new(
423 stream::iter(vec![
424 Ok(event(1, &workflow_id)),
425 Ok(event(2, &workflow_id)),
426 Err(ClientError::unavailable("transient disconnect")),
427 ])
428 .boxed(),
429 ));
430 stub.attempts
431 .lock()
432 .await
433 .push_back(SubscriptionAttempt::new(
434 stream::iter(vec![
435 Ok(event(2, &workflow_id)),
436 Ok(event(3, &workflow_id)),
437 Ok(event(4, &workflow_id)),
438 ])
439 .boxed(),
440 ));
441 let mut events = ResumingEventStream::new(
442 stub.clone(),
443 "tenant-a",
444 SubscribeTarget::Workflow {
445 workflow_id: workflow_id.clone(),
446 },
447 );
448
449 let mut seqs = Vec::new();
450 while let Some(item) = events.next().await {
451 let event = item
452 .map_err(|e| format!("unexpected stream error: {e}"))
453 .ok();
454 if let Some(event) = event {
455 seqs.push(event.seq());
456 }
457 }
458
459 assert_eq!(seqs, vec![1, 2, 3, 4]);
460 assert_eq!(*stub.resume_points.lock().await, vec![None, Some(3)]);
461 }
462
463 #[tokio::test]
464 async fn terminal_failure_is_yielded_before_end() {
465 let workflow_id = WorkflowId::new_v4();
466 let stub = Arc::new(SubscribeStub::default());
467 stub.attempts
468 .lock()
469 .await
470 .push_back(SubscriptionAttempt::new(
471 stream::iter(vec![Err(ClientError::unauthenticated("bad token"))]).boxed(),
472 ));
473 let mut events =
474 ResumingEventStream::new(stub, "tenant-a", SubscribeTarget::Workflow { workflow_id });
475
476 assert_eq!(
477 events.next().await,
478 Some(Err(ClientError::unauthenticated("bad token")))
479 );
480 assert_eq!(events.next().await, None);
481 }
482
483 #[tokio::test]
484 async fn namespace_denied_is_terminal_and_never_retried() {
485 let workflow_id = WorkflowId::new_v4();
486 let stub = Arc::new(SubscribeStub::default());
487 let denied =
488 ClientError::namespace_denied("namespace tenant-b is not granted to this caller");
489 stub.attempts
490 .lock()
491 .await
492 .push_back(SubscriptionAttempt::new(
493 stream::iter(vec![Err(denied.clone())]).boxed(),
494 ));
495 let mut events = ResumingEventStream::new(
496 stub.clone(),
497 "tenant-b",
498 SubscribeTarget::Workflow { workflow_id },
499 );
500
501 assert_eq!(events.next().await, Some(Err(denied)));
502 assert_eq!(events.next().await, None);
503 assert_eq!(stub.resume_points.lock().await.len(), 1);
504 }
505
506 #[tokio::test]
507 async fn from_sequence_passes_the_cursor_on_the_initial_attach() {
508 let workflow_id = WorkflowId::new_v4();
509 let stub = Arc::new(SubscribeStub::default());
510 stub.attempts
511 .lock()
512 .await
513 .push_back(SubscriptionAttempt::new(
514 stream::iter(vec![Ok(event(1, &workflow_id)), Ok(event(2, &workflow_id))]).boxed(),
515 ));
516 let Some(resume_from) = std::num::NonZeroU64::new(1) else {
517 unreachable!("1 is non-zero");
518 };
519 let mut events = super::ResumingEventStream::from_sequence(
520 stub.clone(),
521 "tenant-a",
522 workflow_id,
523 resume_from,
524 );
525
526 let mut seqs = Vec::new();
527 while let Some(item) = events.next().await {
528 if let Ok(event) = item {
529 seqs.push(event.seq());
530 }
531 }
532
533 assert_eq!(seqs, vec![1, 2]);
534 assert_eq!(
535 *stub.resume_points.lock().await,
536 vec![Some(1)],
537 "the initial attach must carry the explicit cursor"
538 );
539 }
540
541 #[tokio::test]
542 async fn live_only_streams_reconnect_only_before_any_delivery() {
543 let workflow_id = WorkflowId::new_v4();
546 let stub = Arc::new(SubscribeStub::default());
547 stub.attempts
548 .lock()
549 .await
550 .push_back(SubscriptionAttempt::new(
551 stream::iter(vec![Err(ClientError::unavailable("transient disconnect"))]).boxed(),
552 ));
553 stub.attempts
554 .lock()
555 .await
556 .push_back(SubscriptionAttempt::new(
557 stream::iter(vec![Ok(event(1, &workflow_id))]).boxed(),
558 ));
559 let mut events = ResumingEventStream::new(
560 stub.clone(),
561 "tenant-a",
562 SubscribeTarget::Filtered {
563 filter: aion_core::WorkflowFilter::default(),
564 },
565 );
566
567 let mut seqs = Vec::new();
568 while let Some(item) = events.next().await {
569 if let Ok(event) = item {
570 seqs.push(event.seq());
571 }
572 }
573
574 assert_eq!(seqs, vec![1]);
575 assert_eq!(
576 *stub.resume_points.lock().await,
577 vec![None, None],
578 "live-only streams never carry a resume cursor"
579 );
580 }
581
582 #[tokio::test]
583 async fn live_only_disconnect_after_delivery_is_honest_unavailable() {
584 for target in [
588 SubscribeTarget::Filtered {
589 filter: aion_core::WorkflowFilter::default(),
590 },
591 SubscribeTarget::Firehose,
592 ] {
593 let workflow_id = WorkflowId::new_v4();
594 let stub = Arc::new(SubscribeStub::default());
595 stub.attempts
596 .lock()
597 .await
598 .push_back(SubscriptionAttempt::new(
599 stream::iter(vec![
600 Ok(event(1, &workflow_id)),
601 Err(ClientError::unavailable("transient disconnect")),
602 ])
603 .boxed(),
604 ));
605 let mut events = ResumingEventStream::new(stub.clone(), "tenant-a", target);
606
607 let first = events.next().await;
608 assert!(matches!(first, Some(Ok(_))), "got {first:?}");
609 assert_eq!(
610 events.next().await,
611 Some(Err(ClientError::unavailable("transient disconnect")))
612 );
613 assert_eq!(events.next().await, None);
614 assert_eq!(
615 stub.resume_points.lock().await.len(),
616 1,
617 "no reattach may follow a post-delivery live-only disconnect"
618 );
619 }
620 }
621
622 #[tokio::test]
623 async fn live_only_streams_do_not_dedupe_sequence_numbers_across_workflows() {
624 let first_workflow = WorkflowId::new_v4();
627 let second_workflow = WorkflowId::new_v4();
628 let stub = Arc::new(SubscribeStub::default());
629 stub.attempts
630 .lock()
631 .await
632 .push_back(SubscriptionAttempt::new(
633 stream::iter(vec![
634 Ok(event(1, &first_workflow)),
635 Ok(event(1, &second_workflow)),
636 ])
637 .boxed(),
638 ));
639 let mut events = ResumingEventStream::new(stub, "tenant-a", SubscribeTarget::Firehose);
640
641 let mut delivered = Vec::new();
642 while let Some(item) = events.next().await {
643 if let Ok(event) = item {
644 delivered.push(event.envelope().workflow_id.clone());
645 }
646 }
647
648 assert_eq!(delivered, vec![first_workflow, second_workflow]);
649 }
650
651 #[tokio::test]
652 async fn not_found_is_terminal_and_never_retried() {
653 let workflow_id = WorkflowId::new_v4();
657 let stub = Arc::new(SubscribeStub::default());
658 stub.attempts
659 .lock()
660 .await
661 .push_back(SubscriptionAttempt::new(
662 stream::iter(vec![Err(ClientError::not_found("workflow was not found"))]).boxed(),
663 ));
664 let mut events = ResumingEventStream::new(
665 stub.clone(),
666 "tenant-a",
667 SubscribeTarget::Workflow { workflow_id },
668 );
669
670 assert_eq!(
671 events.next().await,
672 Some(Err(ClientError::not_found("workflow was not found")))
673 );
674 assert_eq!(events.next().await, None);
675 assert_eq!(stub.resume_points.lock().await.len(), 1);
676 }
677
678 #[tokio::test]
682 async fn unavailable_attach_failure_is_retried_until_attach_succeeds() -> Result<(), ClientError>
683 {
684 let workflow_id = WorkflowId::new_v4();
685 let stub = Arc::new(SubscribeStub::default());
686 stub.attach_failures
687 .lock()
688 .await
689 .push_back(ClientError::unavailable("connection refused"));
690 stub.attach_failures
691 .lock()
692 .await
693 .push_back(ClientError::unavailable("connection refused"));
694 stub.attempts
695 .lock()
696 .await
697 .push_back(SubscriptionAttempt::new(
698 stream::iter(vec![Ok(event(1, &workflow_id)), Ok(event(2, &workflow_id))]).boxed(),
699 ));
700 let mut events = ResumingEventStream::new(
701 stub.clone(),
702 "tenant-a",
703 SubscribeTarget::Workflow { workflow_id },
704 );
705
706 let mut seqs = Vec::new();
707 while let Some(item) = events.next().await {
708 seqs.push(item?.seq());
711 }
712
713 assert_eq!(seqs, vec![1, 2]);
714 assert_eq!(
715 *stub.resume_points.lock().await,
716 vec![None, None, None],
717 "every retried initial attach is still a live tail (no cursor)"
718 );
719 Ok(())
720 }
721
722 #[tokio::test]
725 async fn unavailable_reconnect_failure_retries_with_the_same_cursor() -> Result<(), ClientError>
726 {
727 let workflow_id = WorkflowId::new_v4();
728 let stub = Arc::new(SubscribeStub::default());
729 stub.attempts
730 .lock()
731 .await
732 .push_back(SubscriptionAttempt::new(
733 stream::iter(vec![
734 Ok(event(1, &workflow_id)),
735 Err(ClientError::unavailable("transient disconnect")),
736 ])
737 .boxed(),
738 ));
739 let mut events = ResumingEventStream::new(
740 stub.clone(),
741 "tenant-a",
742 SubscribeTarget::Workflow {
743 workflow_id: workflow_id.clone(),
744 },
745 );
746 let first = events.next().await;
747 assert!(matches!(first, Some(Ok(_))), "got {first:?}");
748 stub.attach_failures
750 .lock()
751 .await
752 .push_back(ClientError::unavailable("connection refused"));
753 stub.attempts
754 .lock()
755 .await
756 .push_back(SubscriptionAttempt::new(
757 stream::iter(vec![Ok(event(2, &workflow_id))]).boxed(),
758 ));
759
760 let mut seqs = vec![1];
761 while let Some(item) = events.next().await {
762 seqs.push(item?.seq());
765 }
766
767 assert_eq!(seqs, vec![1, 2]);
768 assert_eq!(
769 *stub.resume_points.lock().await,
770 vec![None, Some(2), Some(2)],
771 "the failed reconnect and the successful retry carry the same cursor"
772 );
773 Ok(())
774 }
775
776 #[tokio::test]
779 async fn non_retryable_attach_failure_is_terminal() {
780 let workflow_id = WorkflowId::new_v4();
781 let stub = Arc::new(SubscribeStub::default());
782 stub.attach_failures
783 .lock()
784 .await
785 .push_back(ClientError::unauthenticated("bad token"));
786 let mut events = ResumingEventStream::new(
787 stub.clone(),
788 "tenant-a",
789 SubscribeTarget::Workflow { workflow_id },
790 );
791
792 assert_eq!(
793 events.next().await,
794 Some(Err(ClientError::unauthenticated("bad token")))
795 );
796 assert_eq!(events.next().await, None);
797 assert_eq!(stub.resume_points.lock().await.len(), 1);
798 }
799
800 #[tokio::test]
803 async fn live_only_unavailable_attach_failure_is_retried_before_any_delivery() {
804 let workflow_id = WorkflowId::new_v4();
805 let stub = Arc::new(SubscribeStub::default());
806 stub.attach_failures
807 .lock()
808 .await
809 .push_back(ClientError::unavailable("connection refused"));
810 stub.attempts
811 .lock()
812 .await
813 .push_back(SubscriptionAttempt::new(
814 stream::iter(vec![Ok(event(1, &workflow_id))]).boxed(),
815 ));
816 let mut events =
817 ResumingEventStream::new(stub.clone(), "tenant-a", SubscribeTarget::Firehose);
818
819 let mut seqs = Vec::new();
820 while let Some(item) = events.next().await {
821 if let Ok(event) = item {
822 seqs.push(event.seq());
823 }
824 }
825
826 assert_eq!(seqs, vec![1]);
827 assert_eq!(*stub.resume_points.lock().await, vec![None, None]);
828 }
829}