use std::time::Duration;
use async_trait::async_trait;
use futures_util::{SinkExt, StreamExt};
use tokio::net::TcpStream;
use tokio_tungstenite::{
connect_async,
tungstenite::{client::IntoClientRequest, http::header::AUTHORIZATION, protocol::Message},
MaybeTlsStream, WebSocketStream,
};
use url::Url;
use uuid::Uuid;
use crate::{
transport::{Transport, TransportError},
EventData,
};
pub struct WebSocketTransport {
url: Url,
auth_token: Option<String>,
ws_stream: Option<WebSocketStream<MaybeTlsStream<TcpStream>>>,
}
impl WebSocketTransport {
pub fn new(
base_url: Url,
org_id: Uuid,
app_id: Uuid,
auth_token: Option<String>,
) -> Result<Self, TransportError> {
let mut url = base_url.clone();
match url.scheme() {
"http" => url.set_scheme("ws").map_err(|_| {
TransportError::Configuration("Failed to set WebSocket scheme".to_string())
})?,
"https" => url.set_scheme("wss").map_err(|_| {
TransportError::Configuration("Failed to set WebSocket scheme".to_string())
})?,
"ws" | "wss" => {}
_ => {
return Err(TransportError::Configuration(
"Invalid URL scheme for WebSocket".to_string(),
));
}
}
url.set_path(&format!("/api/orgs/{}/apps/{}/ws", org_id, app_id));
Ok(Self {
url,
auth_token,
ws_stream: None,
})
}
}
#[async_trait]
impl Transport for WebSocketTransport {
async fn connect(&mut self) -> Result<(), TransportError> {
let mut request = self
.url
.as_str()
.into_client_request()
.map_err(|e| TransportError::Connection(format!("invalid WS request: {}", e)))?;
if let Some(token) = &self.auth_token {
let value = format!("Bearer {token}").parse().map_err(|_| {
TransportError::Configuration("invalid token for Authorization header".to_string())
})?;
request.headers_mut().insert(AUTHORIZATION, value);
}
let (ws_stream, _) = connect_async(request).await.map_err(|e| {
TransportError::Connection(format!("WebSocket connection failed: {}", e))
})?;
self.ws_stream = Some(ws_stream);
Ok(())
}
async fn send(&mut self, event: EventData) -> Result<(), TransportError> {
let ws_stream = self
.ws_stream
.as_mut()
.ok_or_else(|| TransportError::Send("WebSocket not connected".to_string()))?;
let json = serde_json::to_string(&event)
.map_err(|e| TransportError::Send(format!("Failed to serialize event: {}", e)))?;
ws_stream.send(Message::Text(json)).await.map_err(|e| {
TransportError::Send(format!("Failed to send WebSocket message: {}", e))
})?;
Ok(())
}
async fn close(&mut self) -> Result<(), TransportError> {
if let Some(mut ws_stream) = self.ws_stream.take() {
ws_stream
.close(None)
.await
.map_err(|e| TransportError::Send(format!("Failed to close WebSocket: {}", e)))?;
let drain = async {
while let Some(msg) = ws_stream.next().await {
match msg {
Ok(Message::Close(_)) | Err(_) => break,
Ok(_) => {}
}
}
};
let _ = tokio::time::timeout(Duration::from_secs(2), drain).await;
}
Ok(())
}
}