fix: prevent race condition on socket + remove panicking code

This commit is contained in:
2026-07-21 23:38:35 +02:00
parent a9a49bef89
commit f6b6beda5e
4 changed files with 89 additions and 37 deletions
+1
View File
@@ -13,6 +13,7 @@
- Figure out Meter level encoding - Figure out Meter level encoding
- Full test of client (Improve existing one) - Full test of client (Improve existing one)
- Write correct readme.md - Write correct readme.md
- Make client dsp408 thread safe
Rust library for remote LAN control of Thomann DSP408 digital signal processors. Rust library for remote LAN control of Thomann DSP408 digital signal processors.
+12
View File
@@ -122,6 +122,12 @@ pub enum DSPError {
#[error("operation failed: {0}")] #[error("operation failed: {0}")]
OperationFailed(String), OperationFailed(String),
#[error("internal error: {0}")]
InternalError(String),
#[error("mutex poisoned: {0}")]
MutexPoisoned(String),
#[error("I/O error: {0}")] #[error("I/O error: {0}")]
Io(#[from] std::io::Error), Io(#[from] std::io::Error),
@@ -134,3 +140,9 @@ pub enum DSPError {
#[error("frame error: {0}")] #[error("frame error: {0}")]
Frame(#[from] FrameError), Frame(#[from] FrameError),
} }
impl<T> From<std::sync::PoisonError<T>> for DSPError {
fn from(err: std::sync::PoisonError<T>) -> Self {
DSPError::MutexPoisoned(err.to_string())
}
}
+101 -62
View File
@@ -32,8 +32,8 @@ fn instances() -> &'static Mutex<HashMap<(IpAddr, u8), ()>> {
} }
pub struct DSP408 { pub struct DSP408 {
connection: types::DeviceConnection, connection: types::DeviceConnection, // read-only so thread protected
socket: Arc<Mutex<Option<TcpStream>>>, socket: Arc<Mutex<Option<TcpStream>>>, // thread protected
connected: bool, connected: bool,
state: Option<types::DSPState>, state: Option<types::DSPState>,
} }
@@ -93,7 +93,10 @@ impl DSP408 {
} }
pub fn send_recv(&self, command: &[u8]) -> Result<types::DeviceResponseType, DSPError> { pub fn send_recv(&self, command: &[u8]) -> Result<types::DeviceResponseType, DSPError> {
let mut socket_guard = self.socket.lock().unwrap(); let mut socket_guard = self
.socket
.lock()
.map_err(|e| DSPError::MutexPoisoned(e.to_string()))?;
let socket = socket_guard.as_mut().ok_or(DSPError::NotConnected)?; let socket = socket_guard.as_mut().ok_or(DSPError::NotConnected)?;
@@ -124,58 +127,7 @@ impl DSP408 {
self.state.as_mut().ok_or(DSPError::NotConnected) self.state.as_mut().ok_or(DSPError::NotConnected)
} }
// ---------- Commands ---------- fn initialize_state(&mut self) -> Result<(), DSPError> {
// Command `0x10`
pub fn connect(&mut self) -> Result<(), DSPError> {
let key = (self.connection.host, self.connection.device_id);
{
let map = instances().lock().unwrap();
if map.contains_key(&key) {
return Err(DSPError::AlreadyConnected);
}
}
let addr = SocketAddr::new(self.connection.host, self.connection.port);
let stream = TcpStream::connect_timeout(&addr, SOCKET_TIMEOUT)?;
stream.set_read_timeout(Some(SOCKET_TIMEOUT))?;
{
let mut socket = self.socket.lock().unwrap();
*socket = Some(stream);
}
let handshake = commands::build_handshake();
{
let mut socket = self.socket.lock().unwrap();
let socket = socket.as_mut().ok_or(DSPError::NotConnected)?;
Self::send_raw(socket, &handshake)?;
}
let start = std::time::Instant::now();
while start.elapsed() < SOCKET_TIMEOUT {
let raw = {
let mut socket = self.socket.lock().unwrap();
let socket = socket.as_mut().ok_or(DSPError::NotConnected)?;
Self::recv_raw(socket)?
};
let frame = extract_frame(&raw)?;
let response = map_response(&frame)?;
if matches!(response, types::DeviceResponseType::HandshakeAck(_)) {
instances().lock().unwrap().insert(key, ());
self.connected = true;
let name = self.get_device_name()?; let name = self.get_device_name()?;
let flags = self.get_device_flags()?; let flags = self.get_device_flags()?;
let current_preset = self.get_current_preset()?; let current_preset = self.get_current_preset()?;
@@ -199,29 +151,114 @@ impl DSP408 {
current_config: state, current_config: state,
}); });
Ok(())
}
// ---------- Commands ----------
// Command `0x10`
pub fn connect(&mut self) -> Result<(), DSPError> {
let key = (self.connection.host, self.connection.device_id);
{
let map = instances()
.lock()
.map_err(|e| DSPError::InternalError(e.to_string()))?;
if map.contains_key(&key) {
return Err(DSPError::AlreadyConnected);
}
}
let addr = SocketAddr::new(self.connection.host, self.connection.port);
let stream = TcpStream::connect_timeout(&addr, SOCKET_TIMEOUT)?;
stream.set_read_timeout(Some(SOCKET_TIMEOUT))?;
{
let mut socket = self
.socket
.lock()
.map_err(|e| DSPError::InternalError(e.to_string()))?;
*socket = Some(stream);
}
let handshake = commands::build_handshake();
{
let mut socket = self
.socket
.lock()
.map_err(|e| DSPError::InternalError(e.to_string()))?;
let socket = socket.as_mut().ok_or(DSPError::NotConnected)?;
Self::send_raw(socket, &handshake)?;
}
let start = std::time::Instant::now();
while start.elapsed() < SOCKET_TIMEOUT {
let raw = {
let mut socket = self
.socket
.lock()
.map_err(|e| DSPError::InternalError(e.to_string()))?;
let socket = socket.as_mut().ok_or(DSPError::NotConnected)?;
Self::recv_raw(socket)?
};
let frame = extract_frame(&raw)?;
let response = map_response(&frame)?;
if matches!(response, types::DeviceResponseType::HandshakeAck(_)) {
let result = self.initialize_state();
if result.is_err() {
self.disconnect()?;
return result;
}
instances()
.lock()
.map_err(|e| DSPError::InternalError(e.to_string()))?
.insert(key, ());
self.connected = true;
return Ok(()); return Ok(());
} }
} }
self.disconnect(); self.disconnect()?;
Err(DSPError::HandshakeFailed) Err(DSPError::HandshakeFailed)
} }
// Command `0x11` // Command `0x11`
pub fn disconnect(&mut self) { pub fn disconnect(&mut self) -> Result<(), DSPError> {
let key = (self.connection.host, self.connection.device_id); let key = (self.connection.host, self.connection.device_id);
let mut socket_guard = self.socket.lock().unwrap(); let mut socket_guard = self.socket.lock()?;
if let Some(mut socket) = socket_guard.take() { if let Some(mut socket) = socket_guard.take() {
let _ = socket.write_all(&commands::build_disconnect()); socket
let _ = socket.shutdown(std::net::Shutdown::Both); .write_all(&commands::build_disconnect())
.map_err(DSPError::Io)?;
socket
.shutdown(std::net::Shutdown::Both)
.map_err(DSPError::Io)?;
} }
instances().lock().unwrap().remove(&key); instances().lock()?.remove(&key);
self.connected = false; self.connected = false;
Ok(())
} }
// Command `0x12` // Command `0x12`
@@ -1428,7 +1465,9 @@ impl DSP408 {
impl Drop for DSP408 { impl Drop for DSP408 {
fn drop(&mut self) { fn drop(&mut self) {
self.disconnect(); if self.connected {
let _ = self.disconnect();
}
} }
} }
+1 -1
View File
@@ -27,7 +27,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
dbg!(dsp.state()?); dbg!(dsp.state()?);
println!("Disconnecting..."); println!("Disconnecting...");
dsp.disconnect(); dsp.disconnect()?;
println!("Disconnected"); println!("Disconnected");
Ok(()) Ok(())