use anyhow::{anyhow, Result};
use bytes::Bytes;
use futures::sink::SinkExt;
use futures::stream::StreamExt;
use std::net::SocketAddr;
use tokio::net::TcpStream;
use tokio_util::codec::Framed;
use tracing::{debug, error, info};
use theater_server::{FragmentingCodec, ManagementCommand, ManagementResponse};
#[derive(Debug)]
pub struct TheaterConnection {
pub address: SocketAddr,
connection: Option<Framed<TcpStream, FragmentingCodec>>,
}
impl TheaterConnection {
pub fn new(address: SocketAddr) -> Self {
Self {
address,
connection: None,
}
}
pub async fn connect(&mut self) -> Result<()> {
if self.connection.is_some() {
return Ok(());
}
info!("Connecting to Theater server at {}", self.address);
let socket = TcpStream::connect(self.address).await?;
let codec = FragmentingCodec::new();
let framed = Framed::new(socket, codec);
self.connection = Some(framed);
info!("Connected to Theater server");
Ok(())
}
pub async fn send(&mut self, command: ManagementCommand) -> Result<()> {
if self.connection.is_none() {
self.connect().await?;
}
debug!("Sending command: {:?}", command);
let command_bytes = serde_json::to_vec(&command)?;
let connection = self
.connection
.as_mut()
.ok_or_else(|| anyhow!("Connection lost"))?;
connection.send(Bytes::from(command_bytes)).await?;
debug!("Command sent");
Ok(())
}
pub async fn receive(&mut self) -> Result<ManagementResponse> {
if self.connection.is_none() {
return Err(anyhow!("Not connected"));
}
let connection = self
.connection
.as_mut()
.ok_or_else(|| anyhow!("Connection lost"))?;
match connection.next().await {
Some(Ok(bytes)) => {
let response: ManagementResponse = serde_json::from_slice(&bytes)?;
debug!("Received response: {:?}", response);
Ok(response)
}
Some(Err(e)) => {
error!("Error receiving response: {}", e);
self.connection = None;
Err(anyhow!("Connection error: {}", e))
}
None => {
debug!("Connection closed by server");
self.connection = None;
Err(anyhow!("Connection closed by server"))
}
}
}
pub async fn send_and_receive(
&mut self,
command: ManagementCommand,
) -> Result<ManagementResponse> {
self.send(command).await?;
self.receive().await
}
pub fn is_connected(&self) -> bool {
self.connection.is_some()
}
pub async fn close(&mut self) -> Result<()> {
if let Some(mut connection) = self.connection.take() {
connection.close().await?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
use tokio::net::TcpListener;
use tokio::sync::oneshot;
async fn run_mock_server(addr: SocketAddr, shutdown_rx: oneshot::Receiver<()>) -> Result<()> {
let listener = TcpListener::bind(addr).await?;
let shutdown_future = shutdown_rx;
tokio::select! {
_ = async {
while let Ok((socket, _)) = listener.accept().await {
let mut framed = Framed::new(socket, FragmentingCodec::new());
while let Some(Ok(bytes)) = framed.next().await {
framed.send(bytes.into()).await?;
}
}
Ok::<(), anyhow::Error>(())
} => {},
_ = shutdown_future => {},
}
Ok(())
}
#[tokio::test]
async fn test_connection() -> Result<()> {
let listener = TcpListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
drop(listener);
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server_handle = tokio::spawn(run_mock_server(addr, shutdown_rx));
tokio::time::sleep(Duration::from_millis(100)).await;
let mut client = TheaterConnection::new(addr);
client.connect().await?;
assert!(client.is_connected());
client.close().await?;
let _ = shutdown_tx.send(());
let _ = server_handle.await;
Ok(())
}
}