use crate::error::{ZitiError, ZitiResult};
use crate::transport::tls::TlsConfig;
use futures_util::{SinkExt, StreamExt};
use tokio::net::TcpStream;
use tokio_rustls::{TlsConnector, client::TlsStream};
use tokio_rustls::rustls::pki_types::ServerName;
use tokio_tungstenite::{
tungstenite::Message,
WebSocketStream,
};
use url::Url;
#[derive(Debug)]
pub struct WebSocketTransport {
stream: WebSocketStream<TlsStream<TcpStream>>,
url: Url,
closed: bool,
}
impl WebSocketTransport {
pub async fn connect(url: Url, tls_config: TlsConfig) -> ZitiResult<Self> {
let host = url
.host_str()
.ok_or_else(|| ZitiError::ConfigError("Invalid URL: no host specified".to_string()))?;
let port = url.port().unwrap_or(443);
let tls_connector = TlsConnector::from(tls_config.client_config());
let server_name = ServerName::try_from(host)
.map_err(|e| {
ZitiError::ConfigError(format!("Invalid server name '{}': {}", host, e))
})?
.to_owned();
let tcp_stream = TcpStream::connect((host, port))
.await
.map_err(|e| ZitiError::ConnectionFailed(format!("TCP connection failed: {}", e)))?;
let tls_stream = tls_connector
.connect(server_name, tcp_stream)
.await
.map_err(|e| ZitiError::ConnectionFailed(format!("TLS handshake failed: {}", e)))?;
let (ws_stream, _response) = tokio_tungstenite::client_async(url.as_str(), tls_stream)
.await
.map_err(|e| {
ZitiError::ConnectionFailed(format!("WebSocket handshake failed: {}", e))
})?;
Ok(Self {
stream: ws_stream,
url,
closed: false,
})
}
pub async fn send(&mut self, message: Message) -> ZitiResult<()> {
self.stream.send(message).await.map_err(|e| {
ZitiError::ConnectionFailed(format!("Failed to send WebSocket message: {}", e))
})
}
pub async fn receive(&mut self) -> ZitiResult<Option<Message>> {
match self.stream.next().await {
Some(Ok(message)) => Ok(Some(message)),
Some(Err(e)) => Err(ZitiError::ConnectionFailed(format!(
"Failed to receive WebSocket message: {}",
e
))),
None => Ok(None), }
}
pub async fn close(&mut self) -> ZitiResult<()> {
self.stream.close(None).await.map_err(|e| {
ZitiError::ConnectionFailed(format!("Failed to close WebSocket connection: {}", e))
})?;
self.closed = true;
Ok(())
}
pub fn url(&self) -> &Url {
&self.url
}
pub fn stream_mut(&mut self) -> &mut WebSocketStream<TlsStream<TcpStream>> {
&mut self.stream
}
pub fn is_closed(&self) -> bool {
self.closed
}
}
pub async fn connect_websocket(
url: Url,
tls_config: TlsConfig,
) -> ZitiResult<WebSocketStream<TlsStream<TcpStream>>> {
let transport = WebSocketTransport::connect(url, tls_config).await?;
Ok(transport.stream)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transport::tls::TlsConfig;
#[tokio::test]
async fn test_invalid_url() {
let url = Url::parse("wss:///path").unwrap(); let tls_config = TlsConfig::new();
let result = WebSocketTransport::connect(url, tls_config).await;
assert!(result.is_err());
match result {
Err(ZitiError::ConfigError(msg)) => {
assert!(msg.contains("Invalid URL: no host specified"));
}
Err(ZitiError::ConnectionFailed(msg)) => {
assert!(msg.contains("TCP connection failed") || msg.contains("connection failed"));
}
Err(_) => {
}
Ok(_) => panic!("Expected error but got success"),
}
}
#[test]
fn test_url_parsing() {
let url = Url::parse("wss://example.com:8443/ws").unwrap();
assert_eq!(url.host_str().unwrap(), "example.com");
assert_eq!(url.port().unwrap(), 8443);
}
}