use std::time::Duration;
use syncable_ag_ui_core::{AgentState, Event, JsonValue};
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::response::IntoResponse;
use futures::{SinkExt, StreamExt};
use tokio::sync::mpsc;
use tokio::time::interval;
use crate::error::ServerError;
pub const DEFAULT_PING_INTERVAL: Duration = Duration::from_secs(30);
#[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, "WebSocket channel closed")
}
}
impl<T: std::fmt::Debug> std::error::Error for SendError<T> {}
#[derive(Debug, Clone)]
pub struct WsConfig {
pub ping_interval: Duration,
pub enable_ping: bool,
}
impl Default for WsConfig {
fn default() -> Self {
Self {
ping_interval: DEFAULT_PING_INTERVAL,
enable_ping: true,
}
}
}
impl WsConfig {
pub fn new() -> Self {
Self::default()
}
pub fn ping_interval(mut self, interval: Duration) -> Self {
self.ping_interval = interval;
self
}
pub fn disable_ping(mut self) -> Self {
self.enable_ping = false;
self
}
}
#[derive(Debug, Clone)]
pub struct WsSender<StateT: AgentState = JsonValue> {
sender: mpsc::Sender<Event<StateT>>,
}
impl<StateT: AgentState> WsSender<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 WsHandler<StateT: AgentState = JsonValue> {
receiver: mpsc::Receiver<Event<StateT>>,
config: WsConfig,
}
impl<StateT: AgentState> WsHandler<StateT> {
pub fn into_response(self, upgrade: WebSocketUpgrade) -> impl IntoResponse {
upgrade.on_upgrade(move |socket| self.handle_socket(socket))
}
async fn handle_socket(self, socket: WebSocket) {
let (mut ws_sender, mut ws_receiver) = socket.split();
let mut event_receiver = self.receiver;
let mut ping_interval = if self.config.enable_ping {
Some(interval(self.config.ping_interval))
} else {
None
};
loop {
tokio::select! {
event = event_receiver.recv() => {
match event {
Some(event) => {
let json = match serde_json::to_string(&event) {
Ok(json) => json,
Err(e) => {
eprintln!("WebSocket serialization error: {}", e);
continue;
}
};
if ws_sender.send(Message::Text(json.into())).await.is_err() {
break;
}
}
None => {
let _ = ws_sender.send(Message::Close(None)).await;
break;
}
}
}
_ = async {
if let Some(ref mut interval) = ping_interval {
interval.tick().await;
} else {
std::future::pending::<()>().await;
}
} => {
if ws_sender.send(Message::Ping(vec![].into())).await.is_err() {
break;
}
}
msg = ws_receiver.next() => {
match msg {
Some(Ok(Message::Pong(_))) => {
}
Some(Ok(Message::Close(_))) | None => {
break;
}
Some(Ok(_)) => {
}
Some(Err(_)) => {
break;
}
}
}
}
}
}
}
pub fn channel<StateT: AgentState>(buffer: usize) -> (WsSender<StateT>, WsHandler<StateT>) {
channel_with_config(buffer, WsConfig::default())
}
pub fn channel_with_config<StateT: AgentState>(
buffer: usize,
config: WsConfig,
) -> (WsSender<StateT>, WsHandler<StateT>) {
let (tx, rx) = mpsc::channel(buffer);
(
WsSender { sender: tx },
WsHandler {
receiver: rx,
config,
},
)
}
pub fn format_ws_message<StateT: AgentState>(event: &Event<StateT>) -> Result<String, ServerError> {
serde_json::to_string(event).map_err(|e| ServerError::Serialization(e.to_string()))
}
#[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_channel_with_config() {
let config = WsConfig::new()
.ping_interval(Duration::from_secs(10))
.disable_ping();
let (sender, handler) = channel_with_config::<JsonValue>(10, config);
assert!(!sender.is_closed());
assert!(!handler.config.enable_ping);
assert_eq!(handler.config.ping_interval, Duration::from_secs(10));
}
#[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_ws_message() {
let event: Event = Event::RunError(RunErrorEvent::new("test error"));
let message = format_ws_message(&event).unwrap();
assert!(message.contains("\"type\":\"RUN_ERROR\""));
assert!(message.contains("\"message\":\"test error\""));
}
#[test]
fn test_format_ws_message_complex() {
let event: Event =
Event::TextMessageStart(TextMessageStartEvent::new(MessageId::random()));
let message = format_ws_message(&event).unwrap();
assert!(message.contains("\"type\":\"TEXT_MESSAGE_START\""));
assert!(message.contains("\"messageId\":"));
assert!(message.contains("\"role\":\"assistant\""));
}
#[test]
fn test_ws_config_default() {
let config = WsConfig::default();
assert!(config.enable_ping);
assert_eq!(config.ping_interval, DEFAULT_PING_INTERVAL);
}
#[test]
fn test_ws_config_builder() {
let config = WsConfig::new()
.ping_interval(Duration::from_secs(60))
.disable_ping();
assert!(!config.enable_ping);
assert_eq!(config.ping_interval, Duration::from_secs(60));
}
#[test]
fn test_send_error_display() {
let error: SendError<i32> = SendError(42);
assert_eq!(format!("{}", error), "WebSocket channel closed");
}
}