1use std::future::Future;
19use std::pin::Pin;
20use std::sync::Arc;
21
22use a2a_protocol_types::error::{A2aError, A2aResult};
23use a2a_protocol_types::events::StreamResponse;
24use tokio::sync::{broadcast, mpsc};
25
26use super::{EventQueueReader, EventQueueWriter};
27
28struct CountingWriter(usize);
34
35impl std::io::Write for CountingWriter {
36 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
37 self.0 += buf.len();
38 Ok(buf.len())
39 }
40
41 fn flush(&mut self) -> std::io::Result<()> {
42 Ok(())
43 }
44}
45
46#[derive(Debug, Clone)]
58pub struct InMemoryQueueWriter {
59 tx: broadcast::Sender<A2aResult<StreamResponse>>,
60 persistence_tx: Option<mpsc::Sender<A2aResult<StreamResponse>>>,
64 max_event_size: usize,
66 #[allow(dead_code)]
68 write_timeout: std::time::Duration,
69}
70
71impl InMemoryQueueWriter {
72 pub(super) const fn new(
74 tx: broadcast::Sender<A2aResult<StreamResponse>>,
75 max_event_size: usize,
76 write_timeout: std::time::Duration,
77 ) -> Self {
78 Self {
79 tx,
80 persistence_tx: None,
81 max_event_size,
82 write_timeout,
83 }
84 }
85
86 pub(super) const fn new_with_persistence(
88 tx: broadcast::Sender<A2aResult<StreamResponse>>,
89 persistence_tx: mpsc::Sender<A2aResult<StreamResponse>>,
90 max_event_size: usize,
91 write_timeout: std::time::Duration,
92 ) -> Self {
93 Self {
94 tx,
95 persistence_tx: Some(persistence_tx),
96 max_event_size,
97 write_timeout,
98 }
99 }
100
101 #[must_use]
106 pub fn subscribe(&self) -> InMemoryQueueReader {
107 InMemoryQueueReader::new(self.tx.subscribe())
108 }
109
110 pub(crate) fn raw_subscribe(&self) -> broadcast::Receiver<A2aResult<StreamResponse>> {
115 self.tx.subscribe()
116 }
117}
118
119#[allow(clippy::manual_async_fn)]
120impl EventQueueWriter for InMemoryQueueWriter {
121 fn write<'a>(
122 &'a self,
123 event: StreamResponse,
124 ) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
125 Box::pin(async move {
126 let serialized_size = {
131 let mut counter = CountingWriter(0);
132 serde_json::to_writer(&mut counter, &event)
133 .map_err(|e| A2aError::internal(format!("event serialization failed: {e}")))?;
134 counter.0
135 };
136 if serialized_size > self.max_event_size {
137 return Err(A2aError::internal(format!(
138 "event size {serialized_size} bytes exceeds maximum {} bytes",
139 self.max_event_size
140 )));
141 }
142 if let Some(ref persistence_tx) = self.persistence_tx {
145 if let Err(_e) = persistence_tx.send(Ok(event.clone())).await {
146 trace_warn!("persistence channel closed, event not persisted");
147 }
148 }
149 match self.tx.send(Ok(event)) {
158 Ok(_) => Ok(()),
159 Err(_) if self.persistence_tx.is_some() => {
160 trace_warn!("no live event subscribers; event persisted only");
161 Ok(())
162 }
163 Err(_) => Err(A2aError::internal("event queue: no active receivers")),
164 }
165 })
166 }
167
168 fn close<'a>(&'a self) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
169 Box::pin(async move {
170 Ok(())
173 })
174 }
175}
176
177pub struct InMemoryQueueReader {
188 rx: broadcast::Receiver<A2aResult<StreamResponse>>,
189 pending_first: Option<A2aResult<StreamResponse>>,
190 reattach: Option<ReattachFn>,
192 saw_terminal: bool,
196}
197
198#[allow(clippy::missing_fields_in_debug)]
202impl std::fmt::Debug for InMemoryQueueReader {
203 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
204 f.debug_struct("InMemoryQueueReader")
205 .field("pending_first", &self.pending_first.is_some())
206 .field("reattach", &self.reattach.is_some())
207 .field("saw_terminal", &self.saw_terminal)
208 .finish()
209 }
210}
211
212#[allow(clippy::large_enum_variant)]
216pub enum Reattached {
217 Channel(broadcast::Receiver<A2aResult<StreamResponse>>),
219 Final(StreamResponse),
224 End,
226}
227
228pub type ReattachFn =
231 Arc<dyn Fn() -> Pin<Box<dyn Future<Output = Reattached> + Send>> + Send + Sync>;
232
233const fn carries_terminal_state(event: &StreamResponse) -> bool {
235 match event {
236 StreamResponse::Task(t) => t.status.state.is_terminal(),
237 StreamResponse::StatusUpdate(u) => u.status.state.is_terminal(),
238 _ => false,
239 }
240}
241
242impl InMemoryQueueReader {
243 pub(crate) fn with_reattach(mut self, reattach: ReattachFn) -> Self {
258 self.reattach = Some(reattach);
259 self
260 }
261
262 pub(crate) const fn new(rx: broadcast::Receiver<A2aResult<StreamResponse>>) -> Self {
264 Self {
265 rx,
266 pending_first: None,
267 reattach: None,
268 saw_terminal: false,
269 }
270 }
271
272 pub fn set_first_event(&mut self, event: StreamResponse) {
274 self.pending_first = Some(Ok(event));
275 }
276
277 pub(crate) const fn with_first_event(
279 rx: broadcast::Receiver<A2aResult<StreamResponse>>,
280 first: StreamResponse,
281 ) -> Self {
282 Self {
283 rx,
284 pending_first: Some(Ok(first)),
285 reattach: None,
286 saw_terminal: false,
287 }
288 }
289
290 pub(crate) fn snapshot_then_end(first: StreamResponse) -> Self {
297 let (tx, rx) = broadcast::channel(1);
300 drop(tx);
301 Self {
302 rx,
303 pending_first: Some(Ok(first)),
304 reattach: None,
305 saw_terminal: false,
306 }
307 }
308}
309
310fn lag_error(dropped: u64) -> a2a_protocol_types::error::A2aError {
326 a2a_protocol_types::error::A2aError::stream_lagged(dropped)
327}
328
329#[allow(clippy::redundant_pub_crate)] pub(crate) fn is_lag_error(err: &a2a_protocol_types::error::A2aError) -> bool {
333 err.is_stream_lagged()
334}
335
336impl EventQueueReader for InMemoryQueueReader {
337 fn read(
338 &mut self,
339 ) -> Pin<Box<dyn Future<Output = Option<A2aResult<StreamResponse>>> + Send + '_>> {
340 Box::pin(async move {
341 if let Some(first) = self.pending_first.take() {
344 if let Ok(ref ev) = first {
345 self.saw_terminal |= carries_terminal_state(ev);
346 }
347 return Some(first);
348 }
349 loop {
350 match self.rx.recv().await {
351 Ok(event) => {
352 if let Ok(ref ev) = event {
353 self.saw_terminal |= carries_terminal_state(ev);
354 }
355 return Some(event);
356 }
357 Err(broadcast::error::RecvError::Lagged(n)) => {
358 trace_warn!(
359 dropped_events = n,
360 "event queue reader lagged, {n} events dropped"
361 );
362 return Some(Err(lag_error(n)));
363 }
364 Err(broadcast::error::RecvError::Closed) => {
365 if self.saw_terminal {
370 return None;
371 }
372 let reattach = self.reattach.as_ref()?;
373 match reattach().await {
374 Reattached::Channel(rx) => self.rx = rx,
375 Reattached::Final(event) => {
376 self.saw_terminal = true;
377 return Some(Ok(event));
378 }
379 Reattached::End => return None,
380 }
381 }
382 }
383 }
384 })
385 }
386}
387
388#[cfg(test)]
389mod tests {
390 use super::*;
391 use crate::streaming::event_queue::{
392 new_in_memory_queue, new_in_memory_queue_with_options, DEFAULT_MAX_EVENT_SIZE,
393 DEFAULT_WRITE_TIMEOUT,
394 };
395 use a2a_protocol_types::events::{StreamResponse, TaskStatusUpdateEvent};
396 use a2a_protocol_types::task::{ContextId, TaskId, TaskState, TaskStatus};
397
398 fn make_status_event(task_id: &str, state: TaskState) -> StreamResponse {
400 StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
401 task_id: TaskId::new(task_id),
402 context_id: ContextId::new("ctx-test"),
403 status: TaskStatus {
404 state,
405 message: None,
406 timestamp: None,
407 },
408 metadata: None,
409 })
410 }
411
412 fn counting_reattach(flag: &Arc<std::sync::atomic::AtomicBool>) -> ReattachFn {
424 let flag = Arc::clone(flag);
425 Arc::new(move || {
426 let flag = Arc::clone(&flag);
427 Box::pin(async move {
428 flag.store(true, std::sync::atomic::Ordering::SeqCst);
429 Reattached::End
430 })
431 })
432 }
433
434 #[tokio::test]
439 async fn terminal_status_update_ends_the_stream_without_reattaching() {
440 use std::sync::atomic::{AtomicBool, Ordering};
441
442 let (writer, reader) = new_in_memory_queue();
443 let called = Arc::new(AtomicBool::new(false));
444 let mut reader = reader.with_reattach(counting_reattach(&called));
445
446 writer
447 .write(make_status_event("t-term", TaskState::Completed))
448 .await
449 .unwrap();
450 drop(writer);
451
452 assert!(reader.read().await.is_some(), "the terminal frame arrives");
453 assert!(
454 reader.read().await.is_none(),
455 "a stream that delivered a terminal state ends at close"
456 );
457 assert!(
458 !called.load(Ordering::SeqCst),
459 "the reattach hook must not be consulted once a terminal state has been seen"
460 );
461 }
462
463 #[tokio::test]
467 async fn terminal_task_snapshot_ends_the_stream_without_reattaching() {
468 use a2a_protocol_types::task::Task;
469 use std::sync::atomic::{AtomicBool, Ordering};
470
471 let (writer, reader) = new_in_memory_queue();
472 let called = Arc::new(AtomicBool::new(false));
473 let mut reader = reader.with_reattach(counting_reattach(&called));
474 reader.set_first_event(StreamResponse::Task(Task {
475 id: TaskId::new("t-snap"),
476 context_id: ContextId::new("ctx-test"),
477 status: TaskStatus {
478 state: TaskState::Completed,
479 message: None,
480 timestamp: None,
481 },
482 history: None,
483 artifacts: None,
484 metadata: None,
485 }));
486 drop(writer);
487
488 assert!(reader.read().await.is_some(), "the snapshot arrives first");
489 assert!(
490 reader.read().await.is_none(),
491 "a terminal Task snapshot ends the stream at close"
492 );
493 assert!(
494 !called.load(Ordering::SeqCst),
495 "the reattach hook must not be consulted after a terminal Task snapshot"
496 );
497 }
498
499 #[tokio::test]
504 async fn non_terminal_event_still_consults_the_reattach_hook() {
505 use std::sync::atomic::{AtomicBool, Ordering};
506
507 let (writer, reader) = new_in_memory_queue();
508 let called = Arc::new(AtomicBool::new(false));
509 let mut reader = reader.with_reattach(counting_reattach(&called));
510
511 writer
512 .write(make_status_event("t-working", TaskState::Working))
513 .await
514 .unwrap();
515 drop(writer);
516
517 assert!(reader.read().await.is_some(), "the working frame arrives");
518 assert!(reader.read().await.is_none(), "the hook here returns End");
519 assert!(
520 called.load(Ordering::SeqCst),
521 "without a terminal state the reader must ask the hook whether the task is done"
522 );
523 }
524
525 #[tokio::test]
531 async fn event_of_exactly_max_size_is_accepted() {
532 let event = make_status_event("t-exact", TaskState::Working);
533 let exact = serde_json::to_vec(&event).expect("serializes").len();
534
535 let (writer, _reader) = new_in_memory_queue_with_options(16, exact, DEFAULT_WRITE_TIMEOUT);
536 assert!(
537 writer.write(event).await.is_ok(),
538 "an event of exactly max_event_size ({exact} bytes) is within the \
539 limit and must be accepted"
540 );
541
542 let event = make_status_event("t-exact", TaskState::Working);
545 let (writer, _reader) =
546 new_in_memory_queue_with_options(16, exact - 1, DEFAULT_WRITE_TIMEOUT);
547 assert!(
548 writer.write(event).await.is_err(),
549 "one byte over the limit must still be rejected"
550 );
551 }
552
553 #[tokio::test]
558 async fn reader_debug_reports_the_decision_carrying_fields() {
559 let (_writer, mut reader) = new_in_memory_queue();
560 reader.set_first_event(make_status_event("t-dbg", TaskState::Working));
561
562 let rendered = format!("{reader:?}");
563 assert!(
564 rendered.contains("InMemoryQueueReader"),
565 "the type name must appear: {rendered}"
566 );
567 assert!(
568 rendered.contains("pending_first: true"),
569 "a pending snapshot must be visible: {rendered}"
570 );
571 assert!(
572 rendered.contains("saw_terminal: false"),
573 "terminal tracking must be visible: {rendered}"
574 );
575 }
576
577 #[tokio::test]
585 async fn write_with_no_subscribers_succeeds_when_persistence_attached() {
586 let (writer, reader, mut persistence_rx) =
587 crate::streaming::event_queue::new_in_memory_queue_with_persistence(
588 8,
589 1024 * 1024,
590 std::time::Duration::from_secs(1),
591 );
592 drop(reader); writer
595 .write(make_status_event("t1", TaskState::Working))
596 .await
597 .expect("write must succeed with persistence attached");
598
599 let persisted = persistence_rx
600 .recv()
601 .await
602 .expect("persistence channel should have the event")
603 .expect("event should be Ok");
604 match persisted {
605 StreamResponse::StatusUpdate(evt) => {
606 assert_eq!(evt.status.state, TaskState::Working);
607 }
608 other => panic!("expected StatusUpdate, got: {other:?}"),
609 }
610 }
611
612 #[tokio::test]
616 async fn write_with_no_subscribers_fails_without_persistence() {
617 let (writer, reader) = new_in_memory_queue();
618 drop(reader);
619
620 let result = writer
621 .write(make_status_event("t1", TaskState::Working))
622 .await;
623 assert!(
624 result.is_err(),
625 "sync-mode write with no receivers must fail"
626 );
627 }
628
629 #[tokio::test]
630 async fn write_then_read_single_event() {
631 let (writer, mut reader) = new_in_memory_queue();
632 let event = make_status_event("t1", TaskState::Working);
633
634 writer.write(event).await.expect("write should succeed");
635 drop(writer);
636
637 let received = reader.read().await;
638 assert!(received.is_some(), "reader should return the written event");
639 let result = received.unwrap();
640 let event = result.expect("event should be Ok");
641 match &event {
642 StreamResponse::StatusUpdate(evt) => {
643 assert_eq!(
644 evt.status.state,
645 TaskState::Working,
646 "should be Working event"
647 );
648 }
649 other => panic!("expected StatusUpdate, got: {other:?}"),
650 }
651
652 let eof = reader.read().await;
654 assert!(
655 eof.is_none(),
656 "reader should return None after writer is dropped"
657 );
658 }
659
660 #[tokio::test]
661 async fn write_multiple_events_read_in_order() {
662 let (writer, mut reader) = new_in_memory_queue();
663
664 let e1 = make_status_event("t1", TaskState::Working);
665 let e2 = make_status_event("t1", TaskState::Completed);
666
667 writer.write(e1).await.expect("first write should succeed");
668 writer.write(e2).await.expect("second write should succeed");
669 drop(writer);
670
671 let r1 = reader.read().await.expect("should read first event");
673 let sr1 = r1.expect("first event should be Ok");
674 match &sr1 {
675 StreamResponse::StatusUpdate(evt) => {
676 assert_eq!(
677 evt.status.state,
678 TaskState::Working,
679 "first event should be Working"
680 );
681 }
682 other => panic!("expected StatusUpdate, got: {other:?}"),
683 }
684
685 let r2 = reader.read().await.expect("should read second event");
687 let sr2 = r2.expect("second event should be Ok");
688 match &sr2 {
689 StreamResponse::StatusUpdate(evt) => {
690 assert_eq!(
691 evt.status.state,
692 TaskState::Completed,
693 "second event should be Completed"
694 );
695 }
696 other => panic!("expected StatusUpdate, got: {other:?}"),
697 }
698
699 assert!(
701 reader.read().await.is_none(),
702 "should be EOF after all events"
703 );
704 }
705
706 #[tokio::test]
709 async fn read_returns_none_on_empty_closed_queue() {
710 let (writer, mut reader) = new_in_memory_queue();
711 drop(writer); let result = reader.read().await;
714 assert!(
715 result.is_none(),
716 "reading from an empty closed queue should return None"
717 );
718 }
719
720 #[tokio::test]
721 async fn write_after_all_readers_dropped_returns_error() {
722 let (writer, reader) = new_in_memory_queue();
723 drop(reader);
724
725 let result = writer
726 .write(make_status_event("t1", TaskState::Working))
727 .await;
728 assert!(
729 result.is_err(),
730 "writing with no active receivers should return an error"
731 );
732 }
733
734 #[tokio::test]
735 async fn close_is_no_op_and_succeeds() {
736 let (writer, _reader) = new_in_memory_queue();
737 let result = writer.close().await;
738 assert!(result.is_ok(), "close() should succeed");
739 }
740
741 #[tokio::test]
744 async fn subscribe_creates_independent_reader() {
745 let (writer, mut reader1) = new_in_memory_queue();
746 let mut reader2 = writer.subscribe();
747
748 let event = make_status_event("t1", TaskState::Working);
749 writer.write(event).await.expect("write should succeed");
750 drop(writer);
751
752 let r1 = reader1.read().await;
754 assert!(r1.is_some(), "reader1 should receive the event");
755
756 let r2 = reader2.read().await;
757 assert!(r2.is_some(), "reader2 should receive the event");
758
759 assert!(reader1.read().await.is_none(), "reader1 should see EOF");
761 assert!(reader2.read().await.is_none(), "reader2 should see EOF");
762 }
763
764 #[tokio::test]
765 async fn subscriber_only_sees_events_after_subscribe() {
766 let (writer, mut reader1) = new_in_memory_queue();
767
768 writer
770 .write(make_status_event("t1", TaskState::Submitted))
771 .await
772 .expect("write should succeed");
773
774 let mut reader2 = writer.subscribe();
776
777 writer
779 .write(make_status_event("t1", TaskState::Working))
780 .await
781 .expect("write should succeed");
782 drop(writer);
783
784 let r1a = reader1
786 .read()
787 .await
788 .expect("reader1 should see first event");
789 let evt1a = r1a.expect("first event should be Ok");
790 assert!(
791 matches!(&evt1a, StreamResponse::StatusUpdate(e) if e.status.state == TaskState::Submitted),
792 "reader1 first event should be Submitted"
793 );
794 let r1b = reader1
795 .read()
796 .await
797 .expect("reader1 should see second event");
798 let evt_1b = r1b.expect("second event should be Ok");
799 assert!(
800 matches!(&evt_1b, StreamResponse::StatusUpdate(e) if e.status.state == TaskState::Working),
801 "reader1 second event should be Working"
802 );
803 assert!(reader1.read().await.is_none());
804
805 let r2a = reader2
807 .read()
808 .await
809 .expect("reader2 should see second event");
810 let evt2a = r2a.expect("event should be Ok");
811 assert!(
812 matches!(&evt2a, StreamResponse::StatusUpdate(e) if e.status.state == TaskState::Working),
813 "reader2 should see Working event"
814 );
815 assert!(
816 reader2.read().await.is_none(),
817 "reader2 should see EOF after the one event it received"
818 );
819 }
820
821 #[tokio::test]
824 async fn oversized_event_is_rejected() {
825 let (writer, _reader) = new_in_memory_queue_with_options(
827 16,
828 10, DEFAULT_WRITE_TIMEOUT,
830 );
831
832 let event = make_status_event("t1", TaskState::Working);
833 let result = writer.write(event).await;
834 assert!(
835 result.is_err(),
836 "event exceeding max_event_size should be rejected"
837 );
838 let err = result.unwrap_err();
839 let msg = format!("{err}");
840 assert!(
841 msg.contains("exceeds maximum"),
842 "error message should mention size limit, got: {msg}"
843 );
844 }
845
846 #[test]
848 fn counting_writer_flush_is_noop() {
849 use std::io::Write;
850 let mut cw = super::CountingWriter(0);
851 cw.write_all(b"hello").unwrap();
852 assert_eq!(cw.0, 5);
853 cw.flush().unwrap();
855 assert_eq!(cw.0, 5, "flush should not change the count");
856 }
857
858 #[tokio::test]
859 async fn event_within_size_limit_is_accepted() {
860 let (writer, mut reader) =
862 new_in_memory_queue_with_options(16, DEFAULT_MAX_EVENT_SIZE, DEFAULT_WRITE_TIMEOUT);
863
864 let event = make_status_event("t1", TaskState::Working);
865 writer
866 .write(event)
867 .await
868 .expect("event within size limit should be accepted");
869 drop(writer);
870
871 let r = reader.read().await;
872 assert!(r.is_some(), "reader should receive the event");
873 }
874}