1use std::num::NonZeroUsize;
4
5use aion_core::{Event, WorkflowId};
6use aion_proto::SubscriptionRequest;
7use aion_proto::{WireError, WireErrorCode, encode_streamed_event};
8use axum::extract::ws::{CloseFrame, Message, WebSocket, close_code};
9use futures::{SinkExt, StreamExt};
10use tokio::sync::{mpsc, oneshot};
11
12use crate::config::EVENT_BROADCAST_CAPACITY_REQUIRED;
13use crate::error::ServerError;
14use crate::namespace::CallerIdentity;
15use crate::state::ServerState;
16use crate::stream::namespace_filter::{GateVerdict, NamespaceEventGate};
17use crate::stream::selector::SubscriptionSelector;
18use crate::stream::subscribe::{EventSubscription, subscribe_events};
19
20pub type EncodedFrame = String;
22
23pub const SEQUENCE_CONTIGUITY_VIOLATION: &str = "SequenceContiguityViolation";
28
29pub async fn handle_subscription_socket(
46 mut socket: WebSocket,
47 state: &ServerState,
48 caller: &CallerIdentity,
49 request: &SubscriptionRequest,
50) -> Result<(), ServerError> {
51 let subscription = match subscribe_events(state.namespace_guard(), caller, request).await {
52 Ok(subscription) => subscription,
53 Err(error) => {
54 send_wire_error(&mut socket, &error.to_wire_error()).await?;
55 return Err(error);
56 }
57 };
58 let Some(gate_capacity) = state
64 .runtime_config()
65 .websocket
66 .event_broadcast_capacity
67 .and_then(NonZeroUsize::new)
68 else {
69 let error = ServerError::Config {
70 message: EVENT_BROADCAST_CAPACITY_REQUIRED.to_owned(),
71 };
72 send_wire_error(&mut socket, &error.to_wire_error()).await?;
73 return Err(error);
74 };
75 let mut gate = NamespaceEventGate::new(
79 state.namespace_guard().resolver().clone(),
80 subscription.namespace.clone(),
81 gate_capacity,
82 );
83 if let Some(target) = &subscription.workflow_target {
84 gate.allow(target.clone());
85 }
86 let outbound_buffer_bound = state.runtime_config().websocket.outbound_buffer_bound;
87 forward_subscription(socket, subscription, gate, outbound_buffer_bound).await
88}
89
90pub async fn forward_subscription(
101 socket: WebSocket,
102 subscription: EventSubscription,
103 gate: NamespaceEventGate,
104 outbound_buffer_bound: usize,
105) -> Result<(), ServerError> {
106 let EncodedEventStream {
107 mut frames,
108 lagged,
109 reader_done,
110 } = spawn_encoded_event_stream(subscription, gate, outbound_buffer_bound)?;
111 let (mut socket_tx, mut socket_rx) = socket.split();
112 let result = drive_socket(&mut socket_tx, &mut socket_rx, &mut frames, lagged).await;
113 reader_done.abort();
116 result
117}
118
119async fn drive_socket<Tx, Rx>(
129 socket_tx: &mut Tx,
130 socket_rx: &mut Rx,
131 frames: &mut mpsc::Receiver<EncodedFrame>,
132 lagged: oneshot::Receiver<WireError>,
133) -> Result<(), ServerError>
134where
135 Tx: futures::Sink<Message> + Unpin,
136 <Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
137 Rx: futures::Stream<Item = Result<Message, axum::Error>> + Unpin,
138{
139 let mut lagged = lagged;
140 let mut lag_closed = false;
141 loop {
142 tokio::select! {
143 client_message = socket_rx.next() => {
144 match client_message {
145 Some(Ok(Message::Close(_))) | None => return Ok(()),
146 Some(Ok(message)) => drop(message),
147 Some(Err(error)) => {
148 drop(error);
149 return Ok(());
150 }
151 }
152 }
153 lag = &mut lagged, if !lag_closed => {
154 match lag {
155 Ok(error) => {
156 return drain_then_terminal(socket_tx, frames, error).await;
161 }
162 Err(_closed) => {
163 lag_closed = true;
164 }
165 }
166 }
167 frame = frames.recv() => {
168 let Some(frame) = frame else {
169 if !lag_closed {
174 if let Ok(error) = lagged.try_recv() {
175 return deliver_terminal(socket_tx, error).await;
176 }
177 }
178 return send_normal_close(socket_tx).await;
185 };
186 if socket_tx.send(Message::Text(frame.into())).await.is_err() {
187 return Ok(());
188 }
189 }
190 }
191 }
192}
193
194async fn drain_then_terminal<Tx>(
201 socket_tx: &mut Tx,
202 frames: &mut mpsc::Receiver<EncodedFrame>,
203 error: WireError,
204) -> Result<(), ServerError>
205where
206 Tx: futures::Sink<Message> + Unpin,
207 <Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
208{
209 while let Ok(frame) = frames.try_recv() {
210 if socket_tx.send(Message::Text(frame.into())).await.is_err() {
211 return Ok(());
213 }
214 }
215 deliver_terminal(socket_tx, error).await
216}
217
218const SUBSCRIPTION_COMPLETE_REASON: &str = "subscription complete";
220
221async fn send_normal_close<Tx>(socket_tx: &mut Tx) -> Result<(), ServerError>
226where
227 Tx: futures::Sink<Message> + Unpin,
228 <Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
229{
230 let close = CloseFrame {
231 code: close_code::NORMAL,
232 reason: SUBSCRIPTION_COMPLETE_REASON.into(),
233 };
234 let close_result = socket_tx.send(Message::Close(Some(close))).await;
235 drop(close_result);
236 Ok(())
237}
238
239async fn deliver_terminal<Tx>(socket_tx: &mut Tx, error: WireError) -> Result<(), ServerError>
241where
242 Tx: futures::Sink<Message> + Unpin,
243 <Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
244{
245 send_wire_error(socket_tx, &error).await?;
246 Err(ServerError::Wire { wire: error })
247}
248
249pub(crate) async fn send_wire_error<S>(
255 socket_tx: &mut S,
256 error: &WireError,
257) -> Result<(), ServerError>
258where
259 S: futures::Sink<Message> + Unpin,
260 <S as futures::Sink<Message>>::Error: std::fmt::Debug,
261{
262 let frame = serde_json::json!({ "error": error });
263 let payload = serde_json::to_string(&frame).map_err(|source| ServerError::Wire {
264 wire: WireError::backend(format!("failed to serialize stream error: {source}")),
265 })?;
266 if socket_tx.send(Message::Text(payload.into())).await.is_err() {
267 return Ok(());
268 }
269 let close = CloseFrame {
270 code: close_code::ERROR,
271 reason: error.code.as_str().into(),
272 };
273 let close_result = socket_tx.send(Message::Close(Some(close))).await;
274 drop(close_result);
275 Ok(())
276}
277
278pub struct EncodedEventStream {
280 pub frames: mpsc::Receiver<EncodedFrame>,
282 pub lagged: oneshot::Receiver<WireError>,
286 pub reader_done: tokio::task::JoinHandle<()>,
288}
289
290pub fn spawn_encoded_event_stream(
307 subscription: EventSubscription,
308 gate: NamespaceEventGate,
309 outbound_buffer_bound: usize,
310) -> Result<EncodedEventStream, ServerError> {
311 if outbound_buffer_bound == 0 {
312 return Err(ServerError::Config {
313 message: "websocket.outbound_buffer_bound must be greater than zero".to_owned(),
314 });
315 }
316
317 let EventSubscription {
318 namespace,
319 workflow_target,
320 replay,
321 events,
322 selector,
323 filter: _,
324 } = subscription;
325 let (frames_tx, frames) = mpsc::channel(outbound_buffer_bound);
326 let (lag_tx, lagged) = oneshot::channel();
327 let reader = SubscriptionReader {
328 namespace,
329 workflow_target,
330 gate,
331 selector,
332 contiguity: ContiguityGuard::new(),
333 error_tx: Some(lag_tx),
334 frames_tx,
335 };
336 let reader_done = tokio::spawn(reader.run(replay, events));
337
338 Ok(EncodedEventStream {
339 frames,
340 lagged,
341 reader_done,
342 })
343}
344
345enum QueueMode {
347 Awaiting,
349 Bounded,
351}
352
353enum FrameOutcome {
355 Delivered,
357 Filtered,
359 Stop,
361}
362
363enum ReaderStep {
365 Continue,
367 Stop,
370}
371
372struct SubscriptionReader {
374 namespace: String,
375 workflow_target: Option<WorkflowId>,
376 gate: NamespaceEventGate,
377 selector: SubscriptionSelector,
378 contiguity: ContiguityGuard,
379 error_tx: Option<oneshot::Sender<WireError>>,
380 frames_tx: mpsc::Sender<EncodedFrame>,
381}
382
383impl SubscriptionReader {
384 async fn run(
385 mut self,
386 replay: Vec<Event>,
387 mut events: futures::stream::BoxStream<'static, Result<Event, aion::EventStreamLagged>>,
388 ) {
389 for event in replay {
391 if matches!(
392 self.process(&event, QueueMode::Awaiting).await,
393 ReaderStep::Stop
394 ) {
395 return;
396 }
397 }
398
399 while let Some(item) = events.next().await {
401 let Ok(event) = item else {
404 self.send_terminal(ServerError::lagged_stream().to_wire_error());
405 return;
406 };
407 if matches!(
408 self.process(&event, QueueMode::Bounded).await,
409 ReaderStep::Stop
410 ) {
411 return;
412 }
413 }
414 }
415
416 async fn process(&mut self, event: &Event, mode: QueueMode) -> ReaderStep {
417 let is_target = self
418 .workflow_target
419 .as_ref()
420 .is_some_and(|target| event.workflow_id() == target);
421 if is_target {
424 if let Err(error) = self.contiguity.check(event) {
425 self.send_terminal(error);
426 return ReaderStep::Stop;
427 }
428 }
429 match self.queue(event, mode).await {
430 Ok(FrameOutcome::Delivered) => {
431 if is_target {
432 self.contiguity.record_delivered(event);
433 if is_terminal_workflow_event(event) {
434 return ReaderStep::Stop;
435 }
436 }
437 ReaderStep::Continue
438 }
439 Ok(FrameOutcome::Filtered) => ReaderStep::Continue,
440 Ok(FrameOutcome::Stop) | Err(()) => ReaderStep::Stop,
441 }
442 }
443
444 async fn queue(&mut self, event: &Event, mode: QueueMode) -> Result<FrameOutcome, ()> {
447 let workflow_type = match self.gate.admit(event).await {
448 Ok(GateVerdict::Permitted { workflow_type }) => workflow_type,
449 Ok(GateVerdict::Filtered) => return Ok(FrameOutcome::Filtered),
452 Err(error) => {
453 self.send_terminal(error.to_wire_error());
454 return Err(());
455 }
456 };
457 if !self.selector.matches(event, workflow_type.as_deref()) {
460 return Ok(FrameOutcome::Filtered);
461 }
462 let frame = match encode_frame(&self.namespace, event) {
463 Ok(frame) => frame,
464 Err(error) => {
465 self.send_terminal(error);
466 return Err(());
467 }
468 };
469 match mode {
470 QueueMode::Awaiting => {
471 if self.frames_tx.send(frame).await.is_err() {
472 return Ok(FrameOutcome::Stop);
473 }
474 Ok(FrameOutcome::Delivered)
475 }
476 QueueMode::Bounded => match self.frames_tx.try_send(frame) {
477 Ok(()) => Ok(FrameOutcome::Delivered),
478 Err(mpsc::error::TrySendError::Full(frame)) => {
479 drop(frame);
480 self.send_terminal(ServerError::lagged_stream().to_wire_error());
481 Err(())
482 }
483 Err(mpsc::error::TrySendError::Closed(frame)) => {
484 drop(frame);
485 Ok(FrameOutcome::Stop)
486 }
487 },
488 }
489 }
490
491 fn send_terminal(&mut self, error: WireError) {
492 if let Some(sender) = self.error_tx.take() {
493 let send_result = sender.send(error);
494 drop(send_result);
495 }
496 }
497}
498
499struct ContiguityGuard {
509 last_delivered: Option<u64>,
510}
511
512impl ContiguityGuard {
513 const fn new() -> Self {
514 Self {
515 last_delivered: None,
516 }
517 }
518
519 fn check(&self, event: &Event) -> Result<(), WireError> {
522 let Some(last) = self.last_delivered else {
523 return Ok(());
526 };
527 let expected = last.saturating_add(1);
528 let observed = event.seq();
529 if observed == expected {
530 return Ok(());
531 }
532 Err(WireError::new_with_type(
533 WireErrorCode::Lagged,
534 SEQUENCE_CONTIGUITY_VIOLATION,
535 format!(
536 "per-workflow stream contiguity violated: expected seq {expected}, observed seq \
537 {observed}; reconnect with resume_from_seq = {expected} to resume gap-free from \
538 recorded history"
539 ),
540 ))
541 }
542
543 fn record_delivered(&mut self, event: &Event) {
544 self.last_delivered = Some(event.seq());
545 }
546}
547
548fn encode_frame(namespace: &str, event: &Event) -> Result<EncodedFrame, WireError> {
549 let frame = encode_streamed_event(namespace.to_owned(), None, event)?;
550 serde_json::to_string(&frame).map_err(|source| {
551 WireError::backend(format!(
552 "failed to serialize streamed event frame: {source}"
553 ))
554 })
555}
556
557fn is_terminal_workflow_event(event: &Event) -> bool {
558 matches!(
559 event,
560 Event::WorkflowCompleted { .. }
561 | Event::WorkflowFailed { .. }
562 | Event::WorkflowCancelled { .. }
563 | Event::WorkflowTimedOut { .. }
564 | Event::WorkflowContinuedAsNew { .. }
565 )
566}
567
568#[cfg(test)]
569mod tests {
570 use std::num::NonZeroUsize;
571 use std::time::Duration;
572
573 use aion::EventFilter;
574 use aion_core::{Event, EventEnvelope, Payload, WorkflowId, WorkflowStatus};
575 use aion_proto::{WireError, WireErrorCode};
576 use axum::extract::ws::Message;
577 use futures::{StreamExt, stream, stream::BoxStream};
578 use serde_json::json;
579
580 use super::{SEQUENCE_CONTIGUITY_VIOLATION, drive_socket, spawn_encoded_event_stream};
581 use crate::config::NamespaceMode;
582 use crate::error::ServerError;
583 use crate::namespace::{NamespaceResolver, StaticScheduleNamespaces, StaticWorkflowNamespaces};
584 use crate::stream::namespace_filter::NamespaceEventGate;
585 use crate::stream::selector::SubscriptionSelector;
586 use crate::stream::subscribe::EventSubscription;
587
588 fn capacity(value: usize) -> Result<NonZeroUsize, Box<dyn std::error::Error>> {
589 NonZeroUsize::new(value).ok_or_else(|| "capacity must be non-zero".into())
590 }
591
592 fn envelope(seq: u64, workflow_id: &WorkflowId) -> EventEnvelope {
593 EventEnvelope {
594 seq,
595 recorded_at: chrono::Utc::now(),
596 workflow_id: workflow_id.clone(),
597 }
598 }
599
600 fn payload(label: &str) -> Result<Payload, aion_core::PayloadError> {
601 Payload::from_json(&json!({ "label": label }))
602 }
603
604 fn started_with_type(
605 seq: u64,
606 workflow_id: &WorkflowId,
607 workflow_type: &str,
608 ) -> Result<aion_core::Event, aion_core::PayloadError> {
609 Ok(aion_core::Event::WorkflowStarted {
610 envelope: envelope(seq, workflow_id),
611 workflow_type: workflow_type.to_owned(),
612 input: payload("input")?,
613 run_id: aion_core::RunId::new(uuid::Uuid::from_u128(1)),
614 parent_run_id: None,
615 parent_workflow_id: None,
616 package_version: aion_core::PackageVersion::new("a".repeat(64)),
617 })
618 }
619
620 fn started(
621 seq: u64,
622 workflow_id: &WorkflowId,
623 ) -> Result<aion_core::Event, aion_core::PayloadError> {
624 started_with_type(seq, workflow_id, "checkout")
625 }
626
627 fn signal(
628 seq: u64,
629 workflow_id: &WorkflowId,
630 ) -> Result<aion_core::Event, aion_core::PayloadError> {
631 Ok(aion_core::Event::SignalReceived {
632 envelope: envelope(seq, workflow_id),
633 name: format!("signal-{seq}"),
634 payload: payload("signal")?,
635 })
636 }
637
638 fn completed(
639 seq: u64,
640 workflow_id: &WorkflowId,
641 ) -> Result<aion_core::Event, aion_core::PayloadError> {
642 Ok(aion_core::Event::WorkflowCompleted {
643 envelope: envelope(seq, workflow_id),
644 result: payload("result")?,
645 })
646 }
647
648 fn tenant_a_gate(
649 ownership: StaticWorkflowNamespaces,
650 ) -> Result<NamespaceEventGate, Box<dyn std::error::Error>> {
651 let resolver = NamespaceResolver::authorization_only(
652 NamespaceMode::SharedEngine,
653 ownership,
654 StaticScheduleNamespaces::default(),
655 );
656 Ok(NamespaceEventGate::new(
657 resolver,
658 "tenant-a".to_owned(),
659 capacity(16)?,
660 ))
661 }
662
663 fn subscription(
664 workflow_target: Option<WorkflowId>,
665 replay: Vec<Event>,
666 events: BoxStream<'static, Result<Event, aion::EventStreamLagged>>,
667 ) -> EventSubscription {
668 selected_subscription(
669 workflow_target,
670 replay,
671 events,
672 SubscriptionSelector::unrestricted(),
673 )
674 }
675
676 fn selected_subscription(
677 workflow_target: Option<WorkflowId>,
678 replay: Vec<Event>,
679 events: BoxStream<'static, Result<Event, aion::EventStreamLagged>>,
680 selector: SubscriptionSelector,
681 ) -> EventSubscription {
682 EventSubscription {
683 namespace: "tenant-a".to_owned(),
684 filter: EventFilter::default(),
685 selector,
686 workflow_target,
687 replay,
688 events,
689 }
690 }
691
692 fn owned_gate(
693 workflow_ids: &[&WorkflowId],
694 ) -> Result<NamespaceEventGate, Box<dyn std::error::Error>> {
695 let ownership = StaticWorkflowNamespaces::default();
696 for workflow_id in workflow_ids {
697 ownership.record((*workflow_id).clone(), "tenant-a")?;
698 }
699 tenant_a_gate(ownership)
700 }
701
702 async fn next_frame(
703 receiver: &mut tokio::sync::mpsc::Receiver<String>,
704 ) -> Result<Option<String>, tokio::time::error::Elapsed> {
705 tokio::time::timeout(Duration::from_secs(1), receiver.recv()).await
706 }
707
708 #[tokio::test]
709 async fn per_workflow_stream_ends_after_terminal_event()
710 -> Result<(), Box<dyn std::error::Error>> {
711 let workflow_id = WorkflowId::new_v4();
712 let events = stream::iter([
713 Ok(started(1, &workflow_id)?),
714 Ok(completed(2, &workflow_id)?),
715 Ok(started(3, &workflow_id)?),
716 ])
717 .boxed();
718 let mut stream = spawn_encoded_event_stream(
719 subscription(Some(workflow_id.clone()), Vec::new(), events),
720 owned_gate(&[&workflow_id])?,
721 4,
722 )?;
723
724 let first = next_frame(&mut stream.frames).await?;
725 let second = next_frame(&mut stream.frames).await?;
726 let third = next_frame(&mut stream.frames).await?;
727
728 assert!(first.is_some());
729 assert!(second.is_some());
730 assert!(third.is_none());
731 Ok(())
732 }
733
734 #[tokio::test]
735 async fn dropping_receiver_cleans_up_subscription_reader()
736 -> Result<(), Box<dyn std::error::Error>> {
737 let workflow_id = WorkflowId::new_v4();
738 let events =
739 stream::iter([Ok(started(1, &workflow_id)?), Ok(signal(2, &workflow_id)?)]).boxed();
740 let stream = spawn_encoded_event_stream(
741 subscription(None, Vec::new(), events),
742 owned_gate(&[&workflow_id])?,
743 1,
744 )?;
745 drop(stream.frames);
746
747 tokio::time::timeout(Duration::from_secs(1), stream.reader_done).await??;
748 Ok(())
749 }
750
751 #[tokio::test]
752 async fn slow_consumer_lags_without_blocking_fast_consumer()
753 -> Result<(), Box<dyn std::error::Error>> {
754 let workflow_id = WorkflowId::new_v4();
755 let events: Vec<Result<aion_core::Event, aion::EventStreamLagged>> = vec![
756 Ok(started(1, &workflow_id)?),
757 Ok(signal(2, &workflow_id)?),
758 Ok(completed(3, &workflow_id)?),
759 ];
760 let slow = spawn_encoded_event_stream(
761 subscription(None, Vec::new(), stream::iter(events.clone()).boxed()),
762 owned_gate(&[&workflow_id])?,
763 1,
764 )?;
765 let mut fast = spawn_encoded_event_stream(
766 subscription(None, Vec::new(), stream::iter(events).boxed()),
767 owned_gate(&[&workflow_id])?,
768 4,
769 )?;
770
771 let lag = tokio::time::timeout(Duration::from_secs(1), slow.lagged).await??;
772 assert_eq!(lag.code, WireErrorCode::Lagged);
773
774 let mut received = 0_usize;
775 while let Some(frame) = next_frame(&mut fast.frames).await? {
776 drop(frame);
777 received += 1;
778 }
779 assert_eq!(received, 3);
780 Ok(())
781 }
782
783 #[tokio::test]
788 async fn firehose_never_delivers_foreign_namespace_events()
789 -> Result<(), Box<dyn std::error::Error>> {
790 let own = WorkflowId::new(uuid::Uuid::from_u128(1));
791 let foreign = WorkflowId::new(uuid::Uuid::from_u128(2));
792 let unknown = WorkflowId::new(uuid::Uuid::from_u128(3));
793 let ownership = StaticWorkflowNamespaces::default();
794 ownership.record(own.clone(), "tenant-a")?;
795 ownership.record(foreign.clone(), "tenant-b")?;
796 let events = stream::iter([
799 Ok(started(1, &foreign)?),
800 Ok(started(1, &own)?),
801 Ok(started(1, &unknown)?),
802 Ok(signal(2, &foreign)?),
803 Ok(signal(2, &own)?),
804 ])
805 .boxed();
806 let mut stream = spawn_encoded_event_stream(
807 subscription(None, Vec::new(), events),
808 tenant_a_gate(ownership)?,
809 8,
810 )?;
811
812 let mut delivered = Vec::new();
813 while let Some(frame) = next_frame(&mut stream.frames).await? {
814 let streamed: aion_proto::StreamedEvent = serde_json::from_str(&frame)?;
815 assert_eq!(streamed.namespace, "tenant-a");
816 delivered.push(streamed.decode_event()?.workflow_id().clone());
817 }
818 assert_eq!(
819 delivered,
820 vec![own.clone(), own],
821 "only tenant-a workflow events may be delivered"
822 );
823 Ok(())
824 }
825
826 #[tokio::test]
830 async fn type_selector_delivers_only_matching_workflows_events()
831 -> Result<(), Box<dyn std::error::Error>> {
832 let checkout = WorkflowId::new(uuid::Uuid::from_u128(1));
833 let fulfillment = WorkflowId::new(uuid::Uuid::from_u128(2));
834 let untyped = WorkflowId::new(uuid::Uuid::from_u128(3));
835 let ownership = StaticWorkflowNamespaces::default();
836 ownership.record_with_type(checkout.clone(), "tenant-a", "checkout")?;
837 ownership.record_with_type(fulfillment.clone(), "tenant-a", "fulfillment")?;
838 ownership.record(untyped.clone(), "tenant-a")?;
839 let events = stream::iter([
840 Ok(signal(5, &checkout)?),
842 Ok(signal(5, &fulfillment)?),
843 Ok(signal(5, &untyped)?),
844 Ok(started_with_type(6, &checkout, "checkout")?),
845 Ok(started_with_type(6, &fulfillment, "fulfillment")?),
846 ])
847 .boxed();
848 let mut stream = spawn_encoded_event_stream(
849 selected_subscription(
850 None,
851 Vec::new(),
852 events,
853 SubscriptionSelector {
854 workflow_type: Some("checkout".to_owned()),
855 status: None,
856 },
857 ),
858 tenant_a_gate(ownership)?,
859 8,
860 )?;
861
862 let mut delivered = Vec::new();
863 while let Some(frame) = next_frame(&mut stream.frames).await? {
864 let streamed: aion_proto::StreamedEvent = serde_json::from_str(&frame)?;
865 delivered.push(streamed.decode_event()?.workflow_id().clone());
866 }
867 assert_eq!(
868 delivered,
869 vec![checkout.clone(), checkout],
870 "only events of workflows with the selected type may be delivered"
871 );
872 Ok(())
873 }
874
875 #[tokio::test]
879 async fn status_selector_delivers_per_event_kind_rule() -> Result<(), Box<dyn std::error::Error>>
880 {
881 let workflow_id = WorkflowId::new(uuid::Uuid::from_u128(1));
882 let make_events = || -> Result<_, aion_core::PayloadError> {
883 Ok(stream::iter([
884 Ok(started(1, &workflow_id)?),
885 Ok(signal(2, &workflow_id)?),
886 Ok(completed(3, &workflow_id)?),
887 ])
888 .boxed())
889 };
890
891 let mut completed_only = spawn_encoded_event_stream(
892 selected_subscription(
893 None,
894 Vec::new(),
895 make_events()?,
896 SubscriptionSelector {
897 workflow_type: None,
898 status: Some(WorkflowStatus::Completed),
899 },
900 ),
901 owned_gate(&[&workflow_id])?,
902 8,
903 )?;
904 let mut delivered = Vec::new();
905 while let Some(frame) = next_frame(&mut completed_only.frames).await? {
906 let streamed: aion_proto::StreamedEvent = serde_json::from_str(&frame)?;
907 delivered.push(streamed.decode_event()?.seq());
908 }
909 assert_eq!(
910 delivered,
911 vec![3],
912 "status=Completed delivers only the WorkflowCompleted event"
913 );
914
915 let mut running_only = spawn_encoded_event_stream(
916 selected_subscription(
917 None,
918 Vec::new(),
919 make_events()?,
920 SubscriptionSelector {
921 workflow_type: None,
922 status: Some(WorkflowStatus::Running),
923 },
924 ),
925 owned_gate(&[&workflow_id])?,
926 8,
927 )?;
928 let mut delivered = Vec::new();
929 while let Some(frame) = next_frame(&mut running_only.frames).await? {
930 let streamed: aion_proto::StreamedEvent = serde_json::from_str(&frame)?;
931 delivered.push(streamed.decode_event()?.seq());
932 }
933 assert_eq!(
934 delivered,
935 vec![1, 2],
936 "status=Running delivers exactly the non-terminal events"
937 );
938 Ok(())
939 }
940
941 #[tokio::test]
943 async fn combined_selectors_and_together() -> Result<(), Box<dyn std::error::Error>> {
944 let checkout = WorkflowId::new(uuid::Uuid::from_u128(1));
945 let fulfillment = WorkflowId::new(uuid::Uuid::from_u128(2));
946 let ownership = StaticWorkflowNamespaces::default();
947 ownership.record_with_type(checkout.clone(), "tenant-a", "checkout")?;
948 ownership.record_with_type(fulfillment.clone(), "tenant-a", "fulfillment")?;
949 let events = stream::iter([
950 Ok(signal(1, &checkout)?),
951 Ok(completed(2, &fulfillment)?),
952 Ok(completed(2, &checkout)?),
953 ])
954 .boxed();
955 let mut stream = spawn_encoded_event_stream(
956 selected_subscription(
957 None,
958 Vec::new(),
959 events,
960 SubscriptionSelector {
961 workflow_type: Some("checkout".to_owned()),
962 status: Some(WorkflowStatus::Completed),
963 },
964 ),
965 tenant_a_gate(ownership)?,
966 8,
967 )?;
968
969 let mut delivered = Vec::new();
970 while let Some(frame) = next_frame(&mut stream.frames).await? {
971 let streamed: aion_proto::StreamedEvent = serde_json::from_str(&frame)?;
972 let event = streamed.decode_event()?;
973 delivered.push((event.workflow_id().clone(), event.seq()));
974 }
975 assert_eq!(
976 delivered,
977 vec![(checkout, 2)],
978 "only the selected type's terminal event may pass both selectors"
979 );
980 Ok(())
981 }
982
983 #[tokio::test]
986 async fn replay_longer_than_outbound_buffer_is_delivered_without_lag()
987 -> Result<(), Box<dyn std::error::Error>> {
988 let workflow_id = WorkflowId::new_v4();
989 let mut replay: Vec<Event> = vec![started(1, &workflow_id)?];
990 for seq in 2..=6 {
991 replay.push(signal(seq, &workflow_id)?);
992 }
993 let mut stream = spawn_encoded_event_stream(
994 subscription(Some(workflow_id.clone()), replay, stream::empty().boxed()),
995 owned_gate(&[&workflow_id])?,
996 2,
997 )?;
998
999 let mut received = 0_usize;
1000 while let Some(frame) = next_frame(&mut stream.frames).await? {
1001 drop(frame);
1002 received += 1;
1003 }
1004 assert_eq!(received, 6, "all replay frames must arrive despite bound 2");
1005 let lag = tokio::time::timeout(Duration::from_secs(1), stream.lagged).await?;
1006 assert!(lag.is_err(), "replay must not produce a lag error");
1007 Ok(())
1008 }
1009
1010 #[tokio::test]
1013 async fn gapped_per_workflow_stream_is_terminal_error_never_silent_delivery()
1014 -> Result<(), Box<dyn std::error::Error>> {
1015 let workflow_id = WorkflowId::new_v4();
1016 let events = stream::iter([
1018 Ok(started(1, &workflow_id)?),
1019 Ok(signal(2, &workflow_id)?),
1020 Ok(signal(4, &workflow_id)?),
1021 ])
1022 .boxed();
1023 let mut stream = spawn_encoded_event_stream(
1024 subscription(Some(workflow_id.clone()), Vec::new(), events),
1025 owned_gate(&[&workflow_id])?,
1026 8,
1027 )?;
1028
1029 let mut delivered = Vec::new();
1030 while let Some(frame) = next_frame(&mut stream.frames).await? {
1031 let streamed: aion_proto::StreamedEvent = serde_json::from_str(&frame)?;
1032 delivered.push(streamed.decode_event()?.seq());
1033 }
1034 assert_eq!(delivered, vec![1, 2], "the gapped event must never deliver");
1035
1036 let error = tokio::time::timeout(Duration::from_secs(1), stream.lagged).await??;
1037 assert_eq!(error.code, WireErrorCode::Lagged);
1038 assert_eq!(
1039 error.error_type.as_deref(),
1040 Some(SEQUENCE_CONTIGUITY_VIOLATION)
1041 );
1042 Ok(())
1043 }
1044
1045 #[tokio::test]
1048 async fn contiguity_tripwire_spans_replay_live_boundary_and_duplicates()
1049 -> Result<(), Box<dyn std::error::Error>> {
1050 let workflow_id = WorkflowId::new_v4();
1051
1052 let gapped = spawn_encoded_event_stream(
1054 subscription(
1055 Some(workflow_id.clone()),
1056 vec![started(1, &workflow_id)?, signal(2, &workflow_id)?],
1057 stream::iter([Ok(signal(4, &workflow_id)?)]).boxed(),
1058 ),
1059 owned_gate(&[&workflow_id])?,
1060 8,
1061 )?;
1062 let error = tokio::time::timeout(Duration::from_secs(1), gapped.lagged).await??;
1063 assert_eq!(
1064 error.error_type.as_deref(),
1065 Some(SEQUENCE_CONTIGUITY_VIOLATION)
1066 );
1067
1068 let duplicated = spawn_encoded_event_stream(
1070 subscription(
1071 Some(workflow_id.clone()),
1072 vec![started(1, &workflow_id)?, signal(2, &workflow_id)?],
1073 stream::iter([Ok(signal(2, &workflow_id)?)]).boxed(),
1074 ),
1075 owned_gate(&[&workflow_id])?,
1076 8,
1077 )?;
1078 let error = tokio::time::timeout(Duration::from_secs(1), duplicated.lagged).await??;
1079 assert_eq!(
1080 error.error_type.as_deref(),
1081 Some(SEQUENCE_CONTIGUITY_VIOLATION)
1082 );
1083 Ok(())
1084 }
1085
1086 async fn run_drive_socket(
1088 frames: tokio::sync::mpsc::Receiver<String>,
1089 lagged: tokio::sync::oneshot::Receiver<WireError>,
1090 ) -> Result<(Vec<Message>, Result<(), ServerError>), Box<dyn std::error::Error>> {
1091 let mut frames = frames;
1092 let (mut sink, collected) = futures::channel::mpsc::unbounded();
1093 let mut socket_rx = stream::pending::<Result<Message, axum::Error>>();
1094 let outcome = tokio::time::timeout(
1095 Duration::from_secs(1),
1096 drive_socket(&mut sink, &mut socket_rx, &mut frames, lagged),
1097 )
1098 .await?;
1099 drop(sink);
1100 let messages: Vec<Message> = collected.collect().await;
1101 Ok((messages, outcome))
1102 }
1103
1104 fn assert_frames_then_error_then_close(
1105 messages: &[Message],
1106 expected_frames: usize,
1107 expected_code: &str,
1108 ) -> Result<(), Box<dyn std::error::Error>> {
1109 assert_eq!(
1110 messages.len(),
1111 expected_frames + 2,
1112 "expected {expected_frames} event frames + error frame + close, got {messages:?}"
1113 );
1114 for message in &messages[..expected_frames] {
1115 let Message::Text(text) = message else {
1116 return Err(format!("expected an event text frame, got {message:?}").into());
1117 };
1118 let value: serde_json::Value = serde_json::from_str(text.as_str())?;
1119 assert!(
1120 value.get("error").is_none(),
1121 "event frames must precede the error frame"
1122 );
1123 }
1124 let Message::Text(text) = &messages[expected_frames] else {
1125 return Err("expected the terminal error text frame".into());
1126 };
1127 let value: serde_json::Value = serde_json::from_str(text.as_str())?;
1128 assert_eq!(value["error"]["code"], json!(expected_code));
1129 let Message::Close(Some(close)) = &messages[expected_frames + 1] else {
1130 return Err("expected a close frame after the error frame".into());
1131 };
1132 assert_eq!(close.reason.as_str(), expected_code);
1133 Ok(())
1134 }
1135
1136 #[tokio::test]
1143 async fn terminal_error_and_buffered_frames_are_never_lost_regardless_of_select_order()
1144 -> Result<(), Box<dyn std::error::Error>> {
1145 for _ in 0..64 {
1146 let (frames_tx, frames_rx) = tokio::sync::mpsc::channel::<String>(8);
1147 let (lag_tx, lag_rx) = tokio::sync::oneshot::channel::<WireError>();
1148 for seq in 1..=3 {
1152 frames_tx
1153 .send(json!({ "seq": seq }).to_string())
1154 .await
1155 .map_err(|_| "frame channel must accept the fixture frames")?;
1156 }
1157 lag_tx
1158 .send(WireError::lagged("subscriber lagged behind"))
1159 .map_err(|_| "oneshot must accept the terminal error")?;
1160 drop(frames_tx);
1161
1162 let (messages, outcome) = run_drive_socket(frames_rx, lag_rx).await?;
1163 assert_frames_then_error_then_close(&messages, 3, "lagged")?;
1164 let error = outcome.err().ok_or("terminal stream must surface Err")?;
1165 assert_eq!(error.to_wire_error().code, WireErrorCode::Lagged);
1166 }
1167 Ok(())
1168 }
1169
1170 #[tokio::test]
1175 async fn reader_lag_after_events_delivers_all_frames_then_error()
1176 -> Result<(), Box<dyn std::error::Error>> {
1177 let workflow_id = WorkflowId::new_v4();
1178 for _ in 0..32 {
1179 let events = stream::iter([
1180 Ok(started(1, &workflow_id)?),
1181 Ok(signal(2, &workflow_id)?),
1182 Ok(signal(3, &workflow_id)?),
1183 Err(aion::EventStreamLagged { skipped: 7 }),
1184 ])
1185 .boxed();
1186 let encoded = spawn_encoded_event_stream(
1187 subscription(Some(workflow_id.clone()), Vec::new(), events),
1188 owned_gate(&[&workflow_id])?,
1189 8,
1190 )?;
1191 let (messages, outcome) = run_drive_socket(encoded.frames, encoded.lagged).await?;
1192 assert_frames_then_error_then_close(&messages, 3, "lagged")?;
1193 assert!(outcome.is_err(), "lagged stream must surface Err");
1194 encoded.reader_done.abort();
1195 }
1196 Ok(())
1197 }
1198
1199 #[tokio::test]
1203 async fn clean_stream_end_delivers_frames_then_close_1000_without_error_frame()
1204 -> Result<(), Box<dyn std::error::Error>> {
1205 let (frames_tx, frames_rx) = tokio::sync::mpsc::channel::<String>(8);
1206 let (lag_tx, lag_rx) = tokio::sync::oneshot::channel::<WireError>();
1207 for seq in 1..=2 {
1208 frames_tx
1209 .send(json!({ "seq": seq }).to_string())
1210 .await
1211 .map_err(|_| "frame channel must accept the fixture frames")?;
1212 }
1213 drop(frames_tx);
1214 drop(lag_tx);
1215
1216 let (messages, outcome) = run_drive_socket(frames_rx, lag_rx).await?;
1217 assert!(outcome.is_ok(), "clean end must not surface an error");
1218 assert_eq!(
1219 messages.len(),
1220 3,
1221 "exactly the event frames plus the close-1000 handshake frame"
1222 );
1223 for message in &messages[..2] {
1224 let Message::Text(text) = message else {
1225 return Err(format!("expected a text frame, got {message:?}").into());
1226 };
1227 let value: serde_json::Value = serde_json::from_str(text.as_str())?;
1228 assert!(value.get("error").is_none());
1229 }
1230 let Message::Close(Some(close)) = &messages[2] else {
1231 return Err(format!(
1232 "graceful end must finish with a close frame, got {:?}",
1233 messages[2]
1234 )
1235 .into());
1236 };
1237 assert_eq!(close.code, axum::extract::ws::close_code::NORMAL);
1238 assert_eq!(close.reason.as_str(), super::SUBSCRIPTION_COMPLETE_REASON);
1239 Ok(())
1240 }
1241
1242 #[tokio::test]
1243 async fn wire_error_frame_is_wrapped_and_followed_by_close()
1244 -> Result<(), Box<dyn std::error::Error>> {
1245 let (mut sink, collected) = futures::channel::mpsc::unbounded();
1246 let error = crate::error::ServerError::lagged_stream().to_wire_error();
1247
1248 super::send_wire_error(&mut sink, &error).await?;
1249 drop(sink);
1250
1251 let messages: Vec<axum::extract::ws::Message> = collected.collect().await;
1252 assert_eq!(
1253 messages.len(),
1254 2,
1255 "expected exactly one error frame + close"
1256 );
1257
1258 let axum::extract::ws::Message::Text(text) = &messages[0] else {
1259 return Err("expected a text error frame".into());
1260 };
1261 let frame: serde_json::Value = serde_json::from_str(text.as_str())?;
1262 assert_eq!(frame["error"]["code"], json!("lagged"));
1263 assert!(
1264 frame["error"]["message"].is_string(),
1265 "error frame must carry the informational message"
1266 );
1267
1268 let axum::extract::ws::Message::Close(Some(close)) = &messages[1] else {
1269 return Err("expected a close frame after the error frame".into());
1270 };
1271 assert_eq!(close.reason.as_str(), "lagged");
1272 Ok(())
1273 }
1274}