Skip to main content

autd3_link_remote/link/
async.rs

1use std::net::SocketAddr;
2
3use tokio::{
4    io::{AsyncReadExt, AsyncWriteExt},
5    net::TcpStream,
6};
7
8use crate::{
9    MSG_CLOSE, MSG_CONFIG_GEOMETRY, MSG_ERROR, MSG_OK, MSG_READ_DATA, MSG_SEND_DATA,
10    MSG_UPDATE_GEOMETRY,
11    link::blocking::{REMOTE_HANDSHAKE_LEN, handshake_payload},
12};
13
14use autd3_core::{
15    geometry::Geometry,
16    link::{AsyncLink, LinkError, RxMessage, TxBufferPoolSync, TxMessage},
17};
18
19struct RemoteInner {
20    stream: TcpStream,
21    last_geometry_version: usize,
22    tx_buffer_pool: TxBufferPoolSync,
23    buffer: Vec<u8>,
24}
25
26impl RemoteInner {
27    async fn open(addr: &SocketAddr, geometry: &Geometry) -> Result<RemoteInner, LinkError> {
28        let mut stream = TcpStream::connect(addr).await?;
29
30        Self::perform_handshake(&mut stream).await?;
31        Self::send_geometry(&mut stream, MSG_CONFIG_GEOMETRY, geometry).await?;
32        Self::wait_response(&mut stream).await?;
33
34        let mut tx_buffer_pool = TxBufferPoolSync::default();
35        tx_buffer_pool.init(geometry);
36
37        Ok(Self {
38            stream,
39            last_geometry_version: geometry.version(),
40            tx_buffer_pool,
41            buffer: Vec::new(),
42        })
43    }
44
45    async fn send_geometry(
46        stream: &mut TcpStream,
47        msg_type: u8,
48        geometry: &autd3_core::geometry::Geometry,
49    ) -> Result<(), LinkError> {
50        let num_devices = geometry.len() as u32;
51
52        let mut buffer = Vec::with_capacity(
53            size_of::<u8>()
54                + size_of::<u32>()
55                + (size_of::<f32>() * 3 + size_of::<f32>() * 4) * geometry.len(),
56        );
57        buffer.push(msg_type);
58        buffer.extend_from_slice(&num_devices.to_le_bytes());
59        geometry.iter().for_each(|dev| {
60            let pos = dev[0].position();
61            buffer.extend_from_slice(&pos.x.to_le_bytes());
62            buffer.extend_from_slice(&pos.y.to_le_bytes());
63            buffer.extend_from_slice(&pos.z.to_le_bytes());
64
65            let rot = dev.rotation();
66            buffer.extend_from_slice(&rot.w.to_le_bytes());
67            buffer.extend_from_slice(&rot.i.to_le_bytes());
68            buffer.extend_from_slice(&rot.j.to_le_bytes());
69            buffer.extend_from_slice(&rot.k.to_le_bytes());
70        });
71
72        stream.write_all(&buffer).await?;
73
74        Ok(())
75    }
76
77    async fn perform_handshake(stream: &mut TcpStream) -> Result<(), LinkError> {
78        const PAYLOAD: [u8; REMOTE_HANDSHAKE_LEN] = handshake_payload();
79        stream.write_all(&PAYLOAD).await?;
80        Self::wait_response(stream).await
81    }
82
83    async fn wait_response(stream: &mut TcpStream) -> Result<(), LinkError> {
84        let mut status = [0u8; size_of::<u8>()];
85        stream.read_exact(&mut status).await?;
86
87        match status[0] {
88            MSG_OK => Ok(()),
89            MSG_ERROR => {
90                let mut error_len_buf = [0u8; size_of::<u32>()];
91                stream.read_exact(&mut error_len_buf).await?;
92                let error_len = u32::from_le_bytes(error_len_buf) as usize;
93
94                let mut error_msg = vec![0u8; error_len];
95                stream.read_exact(&mut error_msg).await?;
96
97                let error_str = String::from_utf8_lossy(&error_msg);
98                Err(LinkError::new(format!("Server error: {}", error_str)))
99            }
100            msg => Err(LinkError::new(format!("Unknown response status: {}", msg))),
101        }
102    }
103
104    async fn close(&mut self) -> Result<(), LinkError> {
105        self.stream.write_all(&[MSG_CLOSE]).await?;
106        Self::wait_response(&mut self.stream).await?;
107        Ok(())
108    }
109
110    async fn update(&mut self, geometry: &autd3_core::geometry::Geometry) -> Result<(), LinkError> {
111        if self.last_geometry_version == geometry.version() {
112            return Ok(());
113        }
114        self.last_geometry_version = geometry.version();
115        Self::send_geometry(&mut self.stream, MSG_UPDATE_GEOMETRY, geometry).await?;
116        Self::wait_response(&mut self.stream).await?;
117        Ok(())
118    }
119
120    fn alloc_tx_buffer(&mut self) -> Vec<TxMessage> {
121        self.tx_buffer_pool.borrow()
122    }
123
124    async fn send(&mut self, tx: Vec<TxMessage>) -> Result<(), LinkError> {
125        let buffer_size = size_of::<u8>() + size_of::<TxMessage>() * tx.len();
126        if self.buffer.len() < buffer_size {
127            self.buffer.resize(buffer_size, 0);
128        }
129
130        self.buffer[0] = MSG_SEND_DATA;
131        unsafe {
132            std::ptr::copy_nonoverlapping(
133                tx.as_ptr() as *const u8,
134                self.buffer.as_mut_ptr().add(1),
135                size_of::<TxMessage>() * tx.len(),
136            );
137        }
138        self.tx_buffer_pool.return_buffer(tx);
139
140        self.stream.write_all(&self.buffer).await?;
141        Self::wait_response(&mut self.stream).await?;
142
143        Ok(())
144    }
145
146    async fn receive(&mut self, rx: &mut [RxMessage]) -> Result<(), LinkError> {
147        self.stream.write_all(&[MSG_READ_DATA]).await?;
148        Self::wait_response(&mut self.stream).await?;
149        for bytes in rx.iter_mut().map(|msg| unsafe {
150            std::slice::from_raw_parts_mut(msg as *mut RxMessage as *mut u8, size_of::<RxMessage>())
151        }) {
152            self.stream.read_exact(bytes).await?;
153        }
154        Ok(())
155    }
156}
157
158/// A [`AsyncLink`] for a remote server or [`AUTD3 Simulator`].
159///
160/// [`AUTD3 Simulator`]: https://github.com/shinolab/autd3-server
161pub struct AsyncRemote {
162    addr: SocketAddr,
163    inner: Option<RemoteInner>,
164}
165
166impl AsyncRemote {
167    /// Creates a new [`AsyncRemote`].
168    #[must_use]
169    pub const fn new(addr: SocketAddr) -> AsyncRemote {
170        AsyncRemote { addr, inner: None }
171    }
172}
173
174impl AsyncLink for AsyncRemote {
175    async fn open(&mut self, geometry: &autd3_core::geometry::Geometry) -> Result<(), LinkError> {
176        self.inner = Some(RemoteInner::open(&self.addr, geometry).await?);
177        Ok(())
178    }
179
180    async fn close(&mut self) -> Result<(), LinkError> {
181        if let Some(mut inner) = self.inner.take() {
182            inner.close().await?;
183        }
184        Ok(())
185    }
186
187    async fn update(&mut self, geometry: &autd3_core::geometry::Geometry) -> Result<(), LinkError> {
188        if let Some(inner) = self.inner.as_mut() {
189            inner.update(geometry).await
190        } else {
191            Err(LinkError::closed())
192        }
193    }
194
195    async fn alloc_tx_buffer(&mut self) -> Result<Vec<TxMessage>, LinkError> {
196        if let Some(inner) = self.inner.as_mut() {
197            Ok(inner.alloc_tx_buffer())
198        } else {
199            Err(LinkError::closed())
200        }
201    }
202
203    async fn send(&mut self, tx: Vec<TxMessage>) -> Result<(), LinkError> {
204        if let Some(inner) = self.inner.as_mut() {
205            inner.send(tx).await
206        } else {
207            Err(LinkError::closed())
208        }
209    }
210
211    async fn receive(&mut self, rx: &mut [RxMessage]) -> Result<(), LinkError> {
212        if let Some(inner) = self.inner.as_mut() {
213            inner.receive(rx).await
214        } else {
215            Err(LinkError::closed())
216        }
217    }
218
219    fn is_open(&self) -> bool {
220        self.inner.is_some()
221    }
222
223    fn ensure_is_open(&self) -> Result<(), LinkError> {
224        if self.is_open() {
225            Ok(())
226        } else {
227            Err(LinkError::closed())
228        }
229    }
230}