1use std::collections::HashMap;
13use std::pin::Pin;
14use std::sync::Arc;
15use std::task::{Context, Poll};
16
17use async_trait::async_trait;
18use futures_util::Stream;
19use tokio::sync::{broadcast, mpsc};
20use tokio::task::JoinHandle;
21use tracing::warn;
22
23use super::stream::{BoxRuntimeEventStream, RuntimeEventEnvelope, RuntimeRoom, RuntimeStreamError};
24
25#[cfg(feature = "redis")]
26#[path = "stream_adapter/redis.rs"]
27pub mod redis;
28
29#[cfg(feature = "nats")]
30#[path = "stream_adapter/nats_jetstream.rs"]
31pub mod nats_jetstream;
32
33const ROOM_CHANNEL_CAPACITY: usize = 256;
35
36#[async_trait]
42pub trait RuntimeStreamAdapter: Send + Sync {
43 async fn publish(
45 &self,
46 room: RuntimeRoom,
47 event: RuntimeEventEnvelope,
48 ) -> Result<(), RuntimeStreamError>;
49
50 async fn subscribe(
52 &self,
53 room: RuntimeRoom,
54 ) -> Result<BoxRuntimeEventStream, RuntimeStreamError>;
55}
56
57#[derive(Debug, Default)]
63pub struct MemoryRuntimeStreamAdapter {
64 rooms: tokio::sync::Mutex<HashMap<RuntimeRoom, broadcast::Sender<RuntimeEventEnvelope>>>,
65}
66
67impl MemoryRuntimeStreamAdapter {
68 #[must_use]
70 pub fn new() -> Self {
71 Self::default()
72 }
73
74 async fn sender_for(&self, room: &RuntimeRoom) -> broadcast::Sender<RuntimeEventEnvelope> {
75 let mut rooms = self.rooms.lock().await;
76 rooms
77 .entry(room.clone())
78 .or_insert_with(|| broadcast::channel(ROOM_CHANNEL_CAPACITY).0)
79 .clone()
80 }
81}
82
83#[async_trait]
84impl RuntimeStreamAdapter for MemoryRuntimeStreamAdapter {
85 async fn publish(
86 &self,
87 room: RuntimeRoom,
88 event: RuntimeEventEnvelope,
89 ) -> Result<(), RuntimeStreamError> {
90 let sender = self.sender_for(&room).await;
91 if let Err(broadcast::error::SendError(_envelope)) = sender.send(event) {
94 tracing::trace!(
95 room = %room,
96 "runtime stream publish had no live subscribers"
97 );
98 }
99 Ok(())
100 }
101
102 async fn subscribe(
103 &self,
104 room: RuntimeRoom,
105 ) -> Result<BoxRuntimeEventStream, RuntimeStreamError> {
106 let sender = self.sender_for(&room).await;
107 let mut broadcast_rx = sender.subscribe();
108
109 let (mpsc_tx, mpsc_rx) = mpsc::channel::<Result<RuntimeEventEnvelope, RuntimeStreamError>>(
110 ROOM_CHANNEL_CAPACITY,
111 );
112 let handle = tokio::spawn(async move {
113 loop {
114 match broadcast_rx.recv().await {
115 Ok(envelope) => {
116 if mpsc_tx.send(Ok(envelope)).await.is_err() {
117 break;
118 }
119 }
120 Err(broadcast::error::RecvError::Closed) => break,
121 Err(broadcast::error::RecvError::Lagged(skipped)) => {
122 warn!(
123 skipped,
124 "runtime stream subscriber lagged behind live fanout"
125 );
126 if mpsc_tx
127 .send(Err(RuntimeStreamError::Lagged { skipped }))
128 .await
129 .is_err()
130 {
131 break;
132 }
133 }
134 }
135 }
136 });
137
138 Ok(Box::pin(BroadcastEnvelopeStream {
139 rx: mpsc_rx,
140 handle,
141 }))
142 }
143}
144
145pub struct BroadcastEnvelopeStream {
151 rx: mpsc::Receiver<Result<RuntimeEventEnvelope, RuntimeStreamError>>,
152 handle: JoinHandle<()>,
153}
154
155impl Stream for BroadcastEnvelopeStream {
156 type Item = Result<RuntimeEventEnvelope, RuntimeStreamError>;
157
158 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
159 self.rx.poll_recv(cx)
160 }
161}
162
163impl Drop for BroadcastEnvelopeStream {
164 fn drop(&mut self) {
165 self.handle.abort();
166 }
167}
168
169#[derive(Debug, Default, Clone, Copy)]
172pub struct FailingRuntimeStreamAdapter;
173
174impl FailingRuntimeStreamAdapter {
175 #[must_use]
177 pub fn new() -> Self {
178 Self
179 }
180}
181
182#[async_trait]
183impl RuntimeStreamAdapter for FailingRuntimeStreamAdapter {
184 async fn publish(
185 &self,
186 _room: RuntimeRoom,
187 _event: RuntimeEventEnvelope,
188 ) -> Result<(), RuntimeStreamError> {
189 Err(RuntimeStreamError::Publish {
190 message: "failing runtime stream adapter always rejects publish".to_owned(),
191 })
192 }
193
194 async fn subscribe(
195 &self,
196 _room: RuntimeRoom,
197 ) -> Result<BoxRuntimeEventStream, RuntimeStreamError> {
198 Err(RuntimeStreamError::Subscribe {
199 message: "failing runtime stream adapter never subscribes".to_owned(),
200 })
201 }
202}
203
204pub type DynRuntimeStreamAdapter = Arc<dyn RuntimeStreamAdapter>;
206
207#[cfg(test)]
208mod tests {
209 #![allow(clippy::unwrap_used, clippy::expect_used)]
210
211 use std::time::Duration;
212
213 use chrono::Utc;
214 use futures_util::StreamExt;
215 use uuid::Uuid;
216
217 use super::*;
218 use crate::event::{AgentEvent, RunCompleted, RunStarted};
219 use crate::run::RunId;
220 use crate::stream::RuntimeEventId;
221 use behest_provider::{ModelName, ProviderId};
222
223 fn envelope(run: RunId, seq: u64, session_id: Option<Uuid>) -> RuntimeEventEnvelope {
224 let event = if seq == 1 {
225 AgentEvent::RunStarted(RunStarted {
226 run_id: run,
227 session_id: session_id.unwrap_or_default(),
228 provider: ProviderId::new("acme"),
229 model: ModelName::new("gpt-test"),
230 timestamp: Utc::now(),
231 })
232 } else {
233 AgentEvent::RunCompleted(RunCompleted {
234 run_id: run,
235 finish_reason: behest_provider::FinishReason::Stop,
236 iterations: usize::try_from(seq).unwrap_or(usize::MAX),
237 timestamp: Utc::now(),
238 })
239 };
240 RuntimeEventEnvelope {
241 event_id: RuntimeEventId::new(),
242 seq,
243 run_id: run,
244 session_id,
245 event,
246 emitted_at: Utc::now(),
247 }
248 }
249
250 #[tokio::test]
251 async fn publish_reaches_subscriber() {
252 let adapter = MemoryRuntimeStreamAdapter::new();
253 let run = RunId::new();
254 let room = RuntimeRoom::Run(run);
255
256 let mut stream = adapter.subscribe(room.clone()).await.unwrap();
257 adapter.publish(room, envelope(run, 1, None)).await.unwrap();
258
259 let received = tokio::time::timeout(Duration::from_secs(1), stream.next())
260 .await
261 .expect("timed out waiting for live event")
262 .expect("stream ended")
263 .expect("lagged");
264 assert_eq!(received.seq, 1);
265 }
266
267 #[tokio::test]
268 async fn different_rooms_do_not_cross_talk() {
269 let adapter = MemoryRuntimeStreamAdapter::new();
270 let run_a = RunId::new();
271 let run_b = RunId::new();
272
273 let mut stream_a = adapter.subscribe(RuntimeRoom::Run(run_a)).await.unwrap();
274
275 adapter
276 .publish(RuntimeRoom::Run(run_b), envelope(run_b, 1, None))
277 .await
278 .unwrap();
279 adapter
280 .publish(RuntimeRoom::Run(run_a), envelope(run_a, 1, None))
281 .await
282 .unwrap();
283
284 let received = tokio::time::timeout(Duration::from_secs(1), stream_a.next())
285 .await
286 .expect("timed out waiting for live event")
287 .expect("stream ended")
288 .expect("lagged");
289 assert_eq!(received.run_id, run_a);
290 }
291
292 #[tokio::test]
293 async fn publish_without_subscribers_is_not_an_error() {
294 let adapter = MemoryRuntimeStreamAdapter::new();
295 let run = RunId::new();
296 adapter
297 .publish(RuntimeRoom::Run(run), envelope(run, 1, None))
298 .await
299 .expect("publish with no subscribers must not error");
300 }
301
302 #[tokio::test]
303 async fn failing_adapter_publish_returns_error_without_panic() {
304 let adapter = FailingRuntimeStreamAdapter::new();
305 let run = RunId::new();
306 let err = adapter
307 .publish(RuntimeRoom::Run(run), envelope(run, 1, None))
308 .await
309 .unwrap_err();
310 assert!(matches!(err, RuntimeStreamError::Publish { .. }));
311 }
312}