use std::convert::Infallible;
use std::pin::Pin;
use std::task::{Context, Poll};
use syncable_ag_ui_core::{AgentState, Event, JsonValue};
use axum::response::sse::{Event as AxumSseEvent, KeepAlive, Sse};
use axum::response::IntoResponse;
use futures::Stream;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use crate::error::ServerError;
#[derive(Debug, Clone)]
pub struct SendError<T>(pub T);
impl<T> std::fmt::Display for SendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "channel closed")
}
}
impl<T: std::fmt::Debug> std::error::Error for SendError<T> {}
#[derive(Debug, Clone)]
pub struct SseSender<StateT: AgentState = JsonValue> {
sender: mpsc::Sender<Event<StateT>>,
}
impl<StateT: AgentState> SseSender<StateT> {
pub async fn send(&self, event: Event<StateT>) -> Result<(), SendError<Event<StateT>>> {
self.sender.send(event).await.map_err(|e| SendError(e.0))
}
pub async fn send_many(
&self,
events: impl IntoIterator<Item = Event<StateT>>,
) -> Result<(), SendError<Event<StateT>>> {
for event in events {
self.send(event).await?;
}
Ok(())
}
pub fn try_send(&self, event: Event<StateT>) -> Result<(), SendError<Event<StateT>>> {
self.sender.try_send(event).map_err(|e| SendError(e.into_inner()))
}
pub fn is_closed(&self) -> bool {
self.sender.is_closed()
}
}
pub struct SseHandler<StateT: AgentState = JsonValue> {
receiver: mpsc::Receiver<Event<StateT>>,
}
impl<StateT: AgentState> SseHandler<StateT> {
pub fn into_response(self) -> impl IntoResponse {
let stream = SseEventStream {
inner: ReceiverStream::new(self.receiver),
};
Sse::new(stream).keep_alive(KeepAlive::default())
}
}
struct SseEventStream<StateT: AgentState> {
inner: ReceiverStream<Event<StateT>>,
}
impl<StateT: AgentState> Stream for SseEventStream<StateT> {
type Item = Result<AxumSseEvent, Infallible>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match Pin::new(&mut self.inner).poll_next(cx) {
Poll::Ready(Some(event)) => {
let json = match serde_json::to_string(&event) {
Ok(json) => json,
Err(e) => {
eprintln!("SSE serialization error: {}", e);
format!(r#"{{"type":"RUN_ERROR","message":"Serialization error: {}"}}"#, e)
}
};
let sse_event = AxumSseEvent::default()
.event(event.event_type().as_str())
.data(json);
Poll::Ready(Some(Ok(sse_event)))
}
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
pub fn channel<StateT: AgentState>(buffer: usize) -> (SseSender<StateT>, SseHandler<StateT>) {
let (tx, rx) = mpsc::channel(buffer);
(SseSender { sender: tx }, SseHandler { receiver: rx })
}
pub fn format_sse_event<StateT: AgentState>(event: &Event<StateT>) -> Result<String, ServerError> {
let json = serde_json::to_string(event)
.map_err(|e| ServerError::Serialization(e.to_string()))?;
Ok(format!("data: {}\n\n", json))
}
#[cfg(test)]
mod tests {
use super::*;
use syncable_ag_ui_core::{
MessageId, RunErrorEvent, TextMessageContentEvent, TextMessageStartEvent,
};
#[tokio::test]
async fn test_channel_creation() {
let (sender, _handler) = channel::<JsonValue>(10);
assert!(!sender.is_closed());
}
#[tokio::test]
async fn test_send_event() {
let (sender, mut handler) = channel::<JsonValue>(10);
let event: Event = Event::TextMessageStart(TextMessageStartEvent::new(MessageId::random()));
sender.send(event.clone()).await.unwrap();
let received = handler.receiver.recv().await.unwrap();
assert_eq!(received.event_type(), event.event_type());
}
#[tokio::test]
async fn test_send_many_events() {
let (sender, mut handler) = channel::<JsonValue>(10);
let events: Vec<Event> = vec![
Event::TextMessageStart(TextMessageStartEvent::new(MessageId::random())),
Event::TextMessageContent(TextMessageContentEvent::new_unchecked(
MessageId::random(),
"Hello",
)),
Event::RunError(RunErrorEvent::new("test error")),
];
sender.send_many(events.clone()).await.unwrap();
for expected in &events {
let received = handler.receiver.recv().await.unwrap();
assert_eq!(received.event_type(), expected.event_type());
}
}
#[tokio::test]
async fn test_channel_close_detection() {
let (sender, handler) = channel::<JsonValue>(10);
drop(handler);
assert!(sender.is_closed());
let event: Event = Event::RunError(RunErrorEvent::new("test"));
let result = sender.send(event).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_try_send() {
let (sender, _handler) = channel::<JsonValue>(2);
let event: Event = Event::RunError(RunErrorEvent::new("test"));
assert!(sender.try_send(event.clone()).is_ok());
assert!(sender.try_send(event.clone()).is_ok());
assert!(sender.try_send(event).is_err());
}
#[test]
fn test_format_sse_event() {
let event: Event = Event::RunError(RunErrorEvent::new("test error"));
let formatted = format_sse_event(&event).unwrap();
assert!(formatted.starts_with("data: "));
assert!(formatted.ends_with("\n\n"));
assert!(formatted.contains("\"type\":\"RUN_ERROR\""));
assert!(formatted.contains("\"message\":\"test error\""));
}
#[test]
fn test_format_sse_event_with_complex_event() {
let event: Event = Event::TextMessageStart(TextMessageStartEvent::new(MessageId::random()));
let formatted = format_sse_event(&event).unwrap();
assert!(formatted.contains("\"type\":\"TEXT_MESSAGE_START\""));
assert!(formatted.contains("\"messageId\":"));
assert!(formatted.contains("\"role\":\"assistant\""));
}
}