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
- Full test of client (Improve existing one)
- Write correct readme.md
- Make client dsp408 thread safe
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}")]
OperationFailed(String),
#[error("internal error: {0}")]
InternalError(String),
#[error("mutex poisoned: {0}")]
MutexPoisoned(String),
#[error("I/O error: {0}")]
Io(#[from] std::io::Error),
@@ -134,3 +140,9 @@ pub enum DSPError {
#[error("frame error: {0}")]
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 {
connection: types::DeviceConnection,
socket: Arc<Mutex<Option<TcpStream>>>,
connection: types::DeviceConnection, // read-only so thread protected
socket: Arc<Mutex<Option<TcpStream>>>, // thread protected
connected: bool,
state: Option<types::DSPState>,
}
@@ -93,7 +93,10 @@ impl DSP408 {
}
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)?;
@@ -124,58 +127,7 @@ impl DSP408 {
self.state.as_mut().ok_or(DSPError::NotConnected)
}
// ---------- Commands ----------
// 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;
fn initialize_state(&mut self) -> Result<(), DSPError> {
let name = self.get_device_name()?;
let flags = self.get_device_flags()?;
let current_preset = self.get_current_preset()?;
@@ -199,29 +151,114 @@ impl DSP408 {
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(());
}
}
self.disconnect();
self.disconnect()?;
Err(DSPError::HandshakeFailed)
}
// Command `0x11`
pub fn disconnect(&mut self) {
pub fn disconnect(&mut self) -> Result<(), DSPError> {
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() {
let _ = socket.write_all(&commands::build_disconnect());
let _ = socket.shutdown(std::net::Shutdown::Both);
socket
.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;
Ok(())
}
// Command `0x12`
@@ -1428,7 +1465,9 @@ impl DSP408 {
impl Drop for DSP408 {
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()?);
println!("Disconnecting...");
dsp.disconnect();
dsp.disconnect()?;
println!("Disconnected");
Ok(())