Skip to main content

autd3_rs/client/
mod.rs

1mod config;
2mod pool;
3mod response_future;
4mod rt;
5
6#[cfg(test)]
7mod tests;
8
9pub use config::{ClientConfig, MAX_DEVICES};
10pub use response_future::ResponseFuture;
11
12use std::sync::atomic::{AtomicBool, Ordering};
13use std::sync::{Arc, PoisonError};
14use std::thread::JoinHandle;
15
16use tokio::sync::{mpsc, oneshot};
17
18use crate::command::Pattern;
19use crate::datagram::{Datagram, DatagramBuilder, Frame, Mirror, MirrorHandle};
20use crate::error::{Error, PayloadError};
21use crate::firmware_version::FirmwareVersion;
22use crate::fpga_state::FpgaState;
23use crate::geometry::Autd3;
24use crate::geometry::Geometry;
25use crate::link::{IntoLink, Link};
26use crate::mirror::FirmwareState;
27use crate::operation::{Clear, Distribution, Synchronize};
28use crate::protocol::Cmd;
29use crate::value::Emission;
30
31use pool::SlotPool;
32use rt::CmdMessage;
33
34pub struct Client {
35    cmd_tx: mpsc::Sender<CmdMessage>,
36    num_devices: usize,
37    pool: Arc<SlotPool>,
38    join: std::sync::Mutex<Option<JoinHandle<()>>>,
39    closed: Arc<AtomicBool>,
40    mirror: Arc<std::sync::Mutex<Mirror>>,
41    validate_state: bool,
42}
43
44impl Client {
45    pub fn open<'g, T: IntoLink + 'g>(
46        geometry: &'g Geometry,
47        link: T,
48        config: ClientConfig,
49    ) -> impl Future<Output = Result<Self, Error>> + Send + 'g {
50        Box::pin(async move {
51            Self::open_impl(geometry, link, config)
52                .await
53                .map(|(client, _checker)| client)
54        })
55    }
56
57    pub fn open_with_checker<'g, T: IntoLink + 'g>(
58        geometry: &'g Geometry,
59        link: T,
60        config: ClientConfig,
61    ) -> impl Future<Output = Result<(Self, <T::Link as Link>::Checker), Error>> + Send + 'g {
62        Box::pin(Self::open_impl(geometry, link, config))
63    }
64
65    async fn open_impl<T: IntoLink>(
66        geometry: &Geometry,
67        link: T,
68        config: ClientConfig,
69    ) -> Result<(Self, <T::Link as Link>::Checker), Error> {
70        let config = config.validate()?;
71        let link = link.into_link(geometry).await?;
72        let num_devices = link.num_devices();
73        if num_devices == 0 || num_devices > MAX_DEVICES {
74            return Err(Error::InvalidPayload(PayloadError::DeviceCountOutOfRange {
75                got: num_devices,
76                max: MAX_DEVICES,
77            }));
78        }
79        if geometry.num_devices() != num_devices {
80            return Err(Error::InvalidPayload(
81                PayloadError::GeometryDeviceMismatch {
82                    geometry: geometry.num_devices(),
83                    link: num_devices,
84                },
85            ));
86        }
87
88        let checker = link.state_checker();
89        let pool = SlotPool::new(num_devices, config.max_inflight.get());
90
91        let (cmd_tx, cmd_rx) = mpsc::channel::<CmdMessage>(1);
92        let (hs_done_tx, hs_done_rx) = oneshot::channel::<Result<(), String>>();
93        let closed = Arc::new(AtomicBool::new(false));
94        let closed_for_rt = Arc::clone(&closed);
95        let pool_for_rt = Arc::clone(&pool);
96
97        let join = std::thread::Builder::new()
98            .name("autd3-rs-rt".to_owned())
99            .spawn(move || {
100                rt::run_rt_thread(link, cmd_rx, pool_for_rt, config, hs_done_tx, closed_for_rt);
101            })
102            .map_err(|e| Error::Link(format!("failed to spawn RT thread: {e}")))?;
103
104        match hs_done_rx.await {
105            Ok(Ok(())) => {
106                let client = Self {
107                    cmd_tx,
108                    num_devices,
109                    pool,
110                    join: std::sync::Mutex::new(Some(join)),
111                    closed,
112                    mirror: Arc::new(std::sync::Mutex::new(Mirror::Desynced)),
113                    validate_state: config.validate_state,
114                };
115                if let Err(e) = client.clear().await {
116                    let _ = client.close().await;
117                    return Err(e);
118                }
119                if let Err(e) = client.synchronize().await {
120                    let _ = client.close().await;
121                    return Err(e);
122                }
123                Ok((client, checker))
124            }
125            Ok(Err(msg)) => {
126                let _ = wait_thread(join).await;
127                Err(Error::Link(msg))
128            }
129            Err(_) => {
130                let _ = wait_thread(join).await;
131                Err(Error::RtClosed)
132            }
133        }
134    }
135
136    #[must_use]
137    pub fn num_devices(&self) -> usize {
138        self.num_devices
139    }
140
141    #[must_use]
142    pub fn datagram_builder<'a>(&self) -> DatagramBuilder<'a> {
143        DatagramBuilder::with_mirror(
144            self.num_devices,
145            MirrorHandle {
146                state: Arc::clone(&self.mirror),
147                enabled: self.validate_state,
148            },
149        )
150    }
151
152    fn mark_desynced(&self) {
153        *self.mirror.lock().unwrap_or_else(PoisonError::into_inner) = Mirror::Desynced;
154    }
155
156    async fn clear(&self) -> Result<(), Error> {
157        let datagrams = self.datagram_builder().push(Clear).build()?;
158        for frame in &datagrams {
159            self.send_checked(frame).await?;
160        }
161        *self.mirror.lock().unwrap_or_else(PoisonError::into_inner) =
162            Mirror::Synced(vec![FirmwareState::boot_default(); self.num_devices]);
163        Ok(())
164    }
165
166    async fn send_datagrams(&self, datagrams: &[Datagram]) -> Result<ResponseFuture, Error> {
167        if datagrams.len() != self.num_devices {
168            return Err(Error::InvalidPayload(PayloadError::DatagramCountMismatch {
169                expected: self.num_devices,
170                got: datagrams.len(),
171            }));
172        }
173        let mut slot = self.pool.acquire().await;
174        slot.reset(Distribution::PerDevice);
175        for (device, datagram) in datagrams.iter().enumerate() {
176            slot.payload_mut(device).copy_from_slice(&datagram.payload);
177            slot.set_cmd(device, datagram.cmd);
178        }
179        self.dispatch(slot, false).await
180    }
181
182    async fn send_broadcast(&self, datagram: &Datagram) -> Result<ResponseFuture, Error> {
183        let mut slot = self.pool.acquire().await;
184        slot.reset(Distribution::Broadcast);
185        slot.payload_mut(0).copy_from_slice(&datagram.payload);
186        slot.set_cmd(0, datagram.cmd);
187        self.dispatch(slot, false).await
188    }
189
190    async fn send_broadcast_exclusive(&self, datagram: &Datagram) -> Result<ResponseFuture, Error> {
191        let mut slot = self.pool.acquire().await;
192        slot.reset(Distribution::Broadcast);
193        slot.payload_mut(0).copy_from_slice(&datagram.payload);
194        slot.set_cmd(0, datagram.cmd);
195        self.dispatch(slot, true).await
196    }
197
198    pub async fn send(&self, frame: Frame<'_>) -> Result<ResponseFuture, Error> {
199        match frame.distribution() {
200            Distribution::Broadcast => self.send_broadcast(&frame.datagrams()[0]).await,
201            Distribution::PerDevice => self.send_datagrams(frame.datagrams()).await,
202        }
203    }
204
205    pub async fn send_checked(&self, frame: Frame<'_>) -> Result<(), Error> {
206        let result = match self.send(frame).await {
207            Ok(future) => match future.await {
208                Ok(response) => response.check(),
209                Err(e) => Err(e),
210            },
211            Err(e) => Err(e),
212        };
213        if result.is_err() {
214            self.mark_desynced();
215        }
216        result
217    }
218
219    async fn dispatch(&self, slot: pool::Slot, exclusive: bool) -> Result<ResponseFuture, Error> {
220        let (response_tx, response_rx) = oneshot::channel();
221        if let Err(e) = self
222            .cmd_tx
223            .send(CmdMessage {
224                frame: slot,
225                response_tx,
226                exclusive,
227            })
228            .await
229        {
230            self.pool.release(e.0.frame);
231            return Err(Error::RtClosed);
232        }
233        Ok(ResponseFuture { rx: response_rx })
234    }
235
236    async fn synchronize(&self) -> Result<(), Error> {
237        let datagrams = self.datagram_builder().push(Synchronize).build()?;
238        for frame in &datagrams {
239            self.send_checked(frame).await?;
240        }
241        Ok(())
242    }
243
244    pub async fn stop(&self) -> Result<(), Error> {
245        let patterns = vec![vec![Emission::default(); Autd3::NUM_TRANSDUCERS]; self.num_devices];
246        let datagrams = self
247            .datagram_builder()
248            .push(Pattern::new(&patterns))
249            .build()?;
250        for frame in &datagrams {
251            self.send_checked(frame).await?;
252        }
253        Ok(())
254    }
255
256    pub async fn read_firmware_version(&self) -> Result<Vec<FirmwareVersion>, Error> {
257        let major = self
258            .send_broadcast_exclusive(&Datagram::no_payload(Cmd::ReadCpuFwVersionMajor))
259            .await?
260            .await?
261            .data;
262        let minor = self
263            .send_broadcast_exclusive(&Datagram::no_payload(Cmd::ReadCpuFwVersionMinor))
264            .await?
265            .await?
266            .data;
267        let patch = self
268            .send_broadcast_exclusive(&Datagram::no_payload(Cmd::ReadCpuFwVersionPatch))
269            .await?
270            .await?
271            .data;
272        Ok(major
273            .into_iter()
274            .zip(minor)
275            .zip(patch)
276            .map(|((major, minor), patch)| FirmwareVersion {
277                major,
278                minor,
279                patch,
280            })
281            .collect())
282    }
283
284    pub async fn read_error_detail(&self) -> Result<Vec<u8>, Error> {
285        Ok(self
286            .send_broadcast_exclusive(&Datagram::no_payload(Cmd::ReadErrorDetail))
287            .await?
288            .await?
289            .data)
290    }
291
292    pub async fn read_fpga_state(&self) -> Result<Vec<FpgaState>, Error> {
293        Ok(self
294            .send_broadcast_exclusive(&Datagram::no_payload(Cmd::ReadFpgaState))
295            .await?
296            .await?
297            .data
298            .into_iter()
299            .map(FpgaState)
300            .collect())
301    }
302
303    pub async fn close(&self) -> Result<(), Error> {
304        self.closed.store(true, Ordering::Release);
305        let join = self
306            .join
307            .lock()
308            .unwrap_or_else(PoisonError::into_inner)
309            .take();
310        if let Some(join) = join {
311            wait_thread(join).await
312        } else {
313            Ok(())
314        }
315    }
316}
317
318impl Drop for Client {
319    fn drop(&mut self) {
320        self.closed.store(true, Ordering::Release);
321        let join = self
322            .join
323            .lock()
324            .unwrap_or_else(PoisonError::into_inner)
325            .take();
326        if let Some(join) = join {
327            let _ = join.join();
328        }
329    }
330}
331
332async fn wait_thread(join: JoinHandle<()>) -> Result<(), Error> {
333    tokio::task::spawn_blocking(move || join.join())
334        .await
335        .map_err(|e| Error::Link(format!("RT thread join failed: {e}")))?
336        .map_err(|_| Error::Link("RT thread panicked".to_owned()))
337}