llm_agent_runtime/
streaming.rs1use crate::error::AgentRuntimeError;
30use serde::{Deserialize, Serialize};
31use tokio::sync::broadcast;
32
33#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
48#[serde(tag = "type", rename_all = "lowercase")]
49pub enum AgentEvent {
50 Thought(
52 #[serde(rename = "content")]
54 String,
55 ),
56 Action {
58 tool: String,
60 input: String,
62 },
63 Observation(
65 #[serde(rename = "content")]
67 String,
68 ),
69 Result(
71 #[serde(rename = "content")]
73 String,
74 ),
75 Error(
77 #[serde(rename = "content")]
79 String,
80 ),
81}
82
83impl AgentEvent {
84 pub fn to_sse(&self) -> Result<String, AgentRuntimeError> {
89 let json =
90 serde_json::to_string(self).map_err(|e| AgentRuntimeError::AgentLoop(e.to_string()))?;
91 Ok(format!("data: {json}\n\n"))
92 }
93
94 pub fn event_type(&self) -> &'static str {
96 match self {
97 AgentEvent::Thought(_) => "thought",
98 AgentEvent::Action { .. } => "action",
99 AgentEvent::Observation(_) => "observation",
100 AgentEvent::Result(_) => "result",
101 AgentEvent::Error(_) => "error",
102 }
103 }
104
105 pub fn is_terminal(&self) -> bool {
107 matches!(self, AgentEvent::Result(_) | AgentEvent::Error(_))
108 }
109}
110
111#[derive(Debug, Clone)]
124pub struct StreamBroadcaster {
125 sender: broadcast::Sender<String>,
126}
127
128impl StreamBroadcaster {
129 pub fn new(capacity: usize) -> Result<Self, AgentRuntimeError> {
135 if capacity == 0 {
136 return Err(AgentRuntimeError::AgentLoop(
137 "broadcast channel capacity must be greater than zero".into(),
138 ));
139 }
140 let (sender, _) = broadcast::channel(capacity);
141 Ok(Self { sender })
142 }
143
144 pub fn send(&self, sse_line: String) -> Result<usize, AgentRuntimeError> {
150 match self.sender.send(sse_line) {
151 Ok(n) => Ok(n),
152 Err(_) => Ok(0),
155 }
156 }
157
158 pub fn send_event(&self, event: &AgentEvent) -> Result<usize, AgentRuntimeError> {
160 let sse = event.to_sse()?;
161 self.send(sse)
162 }
163
164 pub fn subscribe(&self) -> broadcast::Receiver<String> {
168 self.sender.subscribe()
169 }
170
171 pub fn receiver_count(&self) -> usize {
173 self.sender.receiver_count()
174 }
175}
176
177#[derive(Debug, Clone)]
200pub struct AgentEventStream {
201 broadcaster: StreamBroadcaster,
202}
203
204impl AgentEventStream {
205 pub fn new(capacity: usize) -> Result<Self, AgentRuntimeError> {
210 Ok(Self {
211 broadcaster: StreamBroadcaster::new(capacity)?,
212 })
213 }
214
215 pub fn broadcaster(&self) -> &StreamBroadcaster {
218 &self.broadcaster
219 }
220
221 pub fn emit_thought(&self, thought: impl Into<String>) -> Result<usize, AgentRuntimeError> {
223 self.broadcaster
224 .send_event(&AgentEvent::Thought(thought.into()))
225 }
226
227 pub fn emit_action(
229 &self,
230 tool: impl Into<String>,
231 input: impl Into<String>,
232 ) -> Result<usize, AgentRuntimeError> {
233 self.broadcaster.send_event(&AgentEvent::Action {
234 tool: tool.into(),
235 input: input.into(),
236 })
237 }
238
239 pub fn emit_observation(
241 &self,
242 observation: impl Into<String>,
243 ) -> Result<usize, AgentRuntimeError> {
244 self.broadcaster
245 .send_event(&AgentEvent::Observation(observation.into()))
246 }
247
248 pub fn emit_result(&self, result: impl Into<String>) -> Result<usize, AgentRuntimeError> {
250 self.broadcaster
251 .send_event(&AgentEvent::Result(result.into()))
252 }
253
254 pub fn emit_error(&self, error: impl Into<String>) -> Result<usize, AgentRuntimeError> {
256 self.broadcaster
257 .send_event(&AgentEvent::Error(error.into()))
258 }
259
260 pub fn receiver_count(&self) -> usize {
262 self.broadcaster.receiver_count()
263 }
264}
265
266#[cfg(test)]
269mod tests {
270 use super::*;
271
272 #[test]
275 fn test_thought_event_serializes_with_type_tag() {
276 let e = AgentEvent::Thought("think".into());
277 let json = serde_json::to_string(&e).unwrap();
278 assert!(json.contains(r#""type":"thought""#));
279 assert!(json.contains(r#""content":"think""#));
280 }
281
282 #[test]
283 fn test_action_event_serializes_tool_and_input() {
284 let e = AgentEvent::Action {
285 tool: "search".into(),
286 input: r#"{"q":"rust"}"#.into(),
287 };
288 let json = serde_json::to_string(&e).unwrap();
289 assert!(json.contains(r#""type":"action""#));
290 assert!(json.contains(r#""tool":"search""#));
291 }
292
293 #[test]
294 fn test_observation_event_serializes() {
295 let e = AgentEvent::Observation("some observation".into());
296 let json = serde_json::to_string(&e).unwrap();
297 assert!(json.contains(r#""type":"observation""#));
298 }
299
300 #[test]
301 fn test_result_event_serializes() {
302 let e = AgentEvent::Result("42".into());
303 let json = serde_json::to_string(&e).unwrap();
304 assert!(json.contains(r#""type":"result""#));
305 assert!(json.contains(r#""content":"42""#));
306 }
307
308 #[test]
309 fn test_error_event_serializes() {
310 let e = AgentEvent::Error("oops".into());
311 let json = serde_json::to_string(&e).unwrap();
312 assert!(json.contains(r#""type":"error""#));
313 }
314
315 #[test]
316 fn test_event_deserializes_roundtrip() {
317 let orig = AgentEvent::Action {
318 tool: "double".into(),
319 input: "21".into(),
320 };
321 let json = serde_json::to_string(&orig).unwrap();
322 let back: AgentEvent = serde_json::from_str(&json).unwrap();
323 assert_eq!(back, orig);
324 }
325
326 #[test]
329 fn test_to_sse_format() {
330 let e = AgentEvent::Thought("hi".into());
331 let sse = e.to_sse().unwrap();
332 assert!(sse.starts_with("data: "));
333 assert!(sse.ends_with("\n\n"));
334 }
335
336 #[test]
339 fn test_event_type_labels() {
340 assert_eq!(AgentEvent::Thought("x".into()).event_type(), "thought");
341 assert_eq!(
342 AgentEvent::Action { tool: "t".into(), input: "i".into() }.event_type(),
343 "action"
344 );
345 assert_eq!(AgentEvent::Observation("o".into()).event_type(), "observation");
346 assert_eq!(AgentEvent::Result("r".into()).event_type(), "result");
347 assert_eq!(AgentEvent::Error("e".into()).event_type(), "error");
348 }
349
350 #[test]
353 fn test_result_is_terminal() {
354 assert!(AgentEvent::Result("done".into()).is_terminal());
355 }
356
357 #[test]
358 fn test_error_is_terminal() {
359 assert!(AgentEvent::Error("bad".into()).is_terminal());
360 }
361
362 #[test]
363 fn test_thought_is_not_terminal() {
364 assert!(!AgentEvent::Thought("thinking".into()).is_terminal());
365 }
366
367 #[test]
368 fn test_action_is_not_terminal() {
369 assert!(!AgentEvent::Action { tool: "t".into(), input: "i".into() }.is_terminal());
370 }
371
372 #[test]
373 fn test_observation_is_not_terminal() {
374 assert!(!AgentEvent::Observation("obs".into()).is_terminal());
375 }
376
377 #[test]
380 fn test_broadcaster_zero_capacity_returns_err() {
381 assert!(StreamBroadcaster::new(0).is_err());
382 }
383
384 #[test]
385 fn test_broadcaster_new_succeeds_with_nonzero_capacity() {
386 assert!(StreamBroadcaster::new(16).is_ok());
387 }
388
389 #[test]
390 fn test_broadcaster_receiver_count_after_subscribe() {
391 let b = StreamBroadcaster::new(8).unwrap();
392 assert_eq!(b.receiver_count(), 0);
393 let _rx = b.subscribe();
394 assert_eq!(b.receiver_count(), 1);
395 }
396
397 #[test]
398 fn test_broadcaster_send_with_no_receivers_returns_ok() {
399 let b = StreamBroadcaster::new(8).unwrap();
400 let result = b.send("data: test\n\n".to_string());
401 assert!(result.is_ok());
402 }
403
404 #[tokio::test]
405 async fn test_broadcaster_subscriber_receives_message() {
406 let b = StreamBroadcaster::new(8).unwrap();
407 let mut rx = b.subscribe();
408 b.send("data: hello\n\n".to_string()).unwrap();
409 let msg = rx.recv().await.unwrap();
410 assert_eq!(msg, "data: hello\n\n");
411 }
412
413 #[tokio::test]
414 async fn test_broadcaster_send_event_delivers_sse() {
415 let b = StreamBroadcaster::new(8).unwrap();
416 let mut rx = b.subscribe();
417 let event = AgentEvent::Thought("test thought".into());
418 b.send_event(&event).unwrap();
419 let msg = rx.recv().await.unwrap();
420 assert!(msg.starts_with("data: "));
421 assert!(msg.contains("thought"));
422 }
423
424 #[tokio::test]
425 async fn test_broadcaster_multiple_subscribers_all_receive() {
426 let b = StreamBroadcaster::new(8).unwrap();
427 let mut rx1 = b.subscribe();
428 let mut rx2 = b.subscribe();
429 b.send("data: ping\n\n".to_string()).unwrap();
430 assert_eq!(rx1.recv().await.unwrap(), "data: ping\n\n");
431 assert_eq!(rx2.recv().await.unwrap(), "data: ping\n\n");
432 }
433
434 #[test]
437 fn test_stream_zero_capacity_fails() {
438 assert!(AgentEventStream::new(0).is_err());
439 }
440
441 #[test]
442 fn test_stream_nonzero_capacity_succeeds() {
443 assert!(AgentEventStream::new(32).is_ok());
444 }
445
446 #[tokio::test]
447 async fn test_stream_emit_thought_received_by_subscriber() {
448 let stream = AgentEventStream::new(16).unwrap();
449 let mut rx = stream.broadcaster().subscribe();
450 stream.emit_thought("some thought").unwrap();
451 let msg = rx.recv().await.unwrap();
452 assert!(msg.contains("thought"));
453 assert!(msg.contains("some thought"));
454 }
455
456 #[tokio::test]
457 async fn test_stream_emit_action_received() {
458 let stream = AgentEventStream::new(16).unwrap();
459 let mut rx = stream.broadcaster().subscribe();
460 stream.emit_action("mytool", r#"{"x":1}"#).unwrap();
461 let msg = rx.recv().await.unwrap();
462 assert!(msg.contains("action"));
463 assert!(msg.contains("mytool"));
464 }
465
466 #[tokio::test]
467 async fn test_stream_emit_observation_received() {
468 let stream = AgentEventStream::new(16).unwrap();
469 let mut rx = stream.broadcaster().subscribe();
470 stream.emit_observation("found result").unwrap();
471 let msg = rx.recv().await.unwrap();
472 assert!(msg.contains("observation"));
473 }
474
475 #[tokio::test]
476 async fn test_stream_emit_result_received() {
477 let stream = AgentEventStream::new(16).unwrap();
478 let mut rx = stream.broadcaster().subscribe();
479 stream.emit_result("final answer").unwrap();
480 let msg = rx.recv().await.unwrap();
481 assert!(msg.contains("result"));
482 assert!(msg.contains("final answer"));
483 }
484
485 #[tokio::test]
486 async fn test_stream_emit_error_received() {
487 let stream = AgentEventStream::new(16).unwrap();
488 let mut rx = stream.broadcaster().subscribe();
489 stream.emit_error("something went wrong").unwrap();
490 let msg = rx.recv().await.unwrap();
491 assert!(msg.contains("error"));
492 assert!(msg.contains("something went wrong"));
493 }
494
495 #[test]
496 fn test_stream_receiver_count() {
497 let stream = AgentEventStream::new(8).unwrap();
498 assert_eq!(stream.receiver_count(), 0);
499 let _rx = stream.broadcaster().subscribe();
500 assert_eq!(stream.receiver_count(), 1);
501 }
502
503 #[test]
504 fn test_stream_clone_shares_broadcaster() {
505 let stream = AgentEventStream::new(8).unwrap();
506 let clone = stream.clone();
507 let _rx = stream.broadcaster().subscribe();
508 assert_eq!(clone.receiver_count(), 1);
509 }
510
511 #[test]
512 fn test_event_is_send_sync() {
513 fn assert_send_sync<T: Send + Sync>() {}
514 assert_send_sync::<AgentEvent>();
515 assert_send_sync::<StreamBroadcaster>();
516 assert_send_sync::<AgentEventStream>();
517 }
518}