Skip to main content

autd3_rs/client/
mod.rs

1mod completion;
2mod config;
3mod pool;
4mod rt;
5
6#[cfg(test)]
7mod tests;
8
9pub use completion::ResponseFuture;
10pub use config::{ClientConfig, MAX_DEVICES};
11
12use std::sync::atomic::{AtomicBool, Ordering};
13use std::sync::{Arc, PoisonError};
14use std::thread::JoinHandle;
15
16use std::sync::mpsc;
17
18use autd3_rs_core::rt::oneshot;
19
20use autd3_cpu_wire::payload::ReadTelemetryPayload;
21use zerocopy::FromBytes;
22
23use crate::commands::Pattern;
24use crate::commands::operation::{Clear, Distribution, Synchronize};
25use crate::datagram::{Datagram, DatagramBuilder, Frame, Mirror, MirrorHandle};
26use crate::error::{Error, LinkCause, PayloadError};
27use crate::firmware_version::{FirmwareVersion, Version};
28use crate::fpga_state::FpgaState;
29use crate::geometry::Geometry;
30use crate::link::{DcClock, IntoLink, Link, LinkStats};
31use crate::mirror::FirmwareState;
32use crate::protocol::{Cmd, DeviceErrorCode};
33use crate::telemetry::Telemetry;
34use crate::value::{DcSysTime, Emission};
35
36use completion::{CompletionPool, Reply};
37use pool::SlotPool;
38use rt::CmdMessage;
39
40pub struct Client {
41    cmd_tx: mpsc::Sender<CmdMessage>,
42    geometry: Arc<Geometry>,
43    num_devices: usize,
44    pool: Arc<SlotPool>,
45    completions: Arc<CompletionPool>,
46    join: std::sync::Mutex<Option<JoinHandle<()>>>,
47    done: std::sync::Mutex<Option<oneshot::Receiver<Option<LinkCause>>>>,
48    closed: Arc<AtomicBool>,
49    stopping: AtomicBool,
50    mirror: MirrorHandle,
51    dc_clock: Option<DcClock>,
52    stats: LinkStats,
53}
54
55impl Client {
56    pub fn open<'g, T: IntoLink + 'g>(
57        geometry: &'g Geometry,
58        link: T,
59        config: ClientConfig,
60    ) -> impl Future<Output = Result<Self, Error>> + Send + 'g {
61        Box::pin(async move {
62            Self::open_impl(geometry, link, config)
63                .await
64                .map(|(client, _checker)| client)
65        })
66    }
67
68    pub fn open_with_checker<'g, T: IntoLink + 'g>(
69        geometry: &'g Geometry,
70        link: T,
71        config: ClientConfig,
72    ) -> impl Future<Output = Result<(Self, <T::Link as Link>::Checker), Error>> + Send + 'g {
73        Box::pin(Self::open_impl(geometry, link, config))
74    }
75
76    async fn open_impl<T: IntoLink>(
77        geometry: &Geometry,
78        link: T,
79        config: ClientConfig,
80    ) -> Result<(Self, <T::Link as Link>::Checker), Error> {
81        let config = config.validate()?;
82        let mut link = link.into_link(geometry)?;
83        let num_devices = link.num_devices();
84        if num_devices == 0 || num_devices > MAX_DEVICES {
85            close_unopened(&mut link);
86            return Err(PayloadError::DeviceCountOutOfRange {
87                got: num_devices,
88                max: MAX_DEVICES,
89            }
90            .into());
91        }
92        if geometry.num_devices() != num_devices {
93            close_unopened(&mut link);
94            return Err(PayloadError::GeometryDeviceMismatch {
95                geometry: geometry.num_devices(),
96                link: num_devices,
97            }
98            .into());
99        }
100
101        let checker = link.state_checker();
102        let dc_clock = link.dc_clock();
103        let stats = link.stats();
104        let pool = SlotPool::new(num_devices, config.max_inflight.get());
105        let completions = CompletionPool::new(config.max_inflight.get());
106
107        let (cmd_tx, cmd_rx) = mpsc::channel::<CmdMessage>();
108        let (hs_done_tx, hs_done_rx) = oneshot::channel::<Result<(), LinkCause>>();
109        let (done_tx, done_rx) = oneshot::channel::<Option<LinkCause>>();
110        let closed = Arc::new(AtomicBool::new(false));
111        let closed_for_rt = Arc::clone(&closed);
112
113        let join = std::thread::Builder::new()
114            .name("autd3-rs-rt".to_owned())
115            .spawn(move || {
116                rt::run_rt_thread(link, cmd_rx, config, hs_done_tx, done_tx, closed_for_rt);
117            })
118            .map_err(|e| Error::Link(LinkCause::new(e)))?;
119
120        match hs_done_rx.await {
121            Ok(Ok(())) => {
122                tracing::debug!("RT thread handshake complete");
123                let client = Self {
124                    cmd_tx,
125                    geometry: Arc::new(geometry.clone()),
126                    num_devices,
127                    pool,
128                    completions,
129                    join: std::sync::Mutex::new(Some(join)),
130                    done: std::sync::Mutex::new(Some(done_rx)),
131                    closed,
132                    stopping: AtomicBool::new(false),
133                    mirror: MirrorHandle {
134                        state: Arc::new(std::sync::Mutex::new(Mirror::Desynced)),
135                        enabled: config.validate_state,
136                    },
137                    dc_clock,
138                    stats,
139                };
140                if let Err(e) = client
141                    .check_firmware_version(config.require_supported_firmware)
142                    .await
143                {
144                    let _ = client.close_impl(false).await;
145                    return Err(e);
146                }
147                if let Err(e) = client.clear().await {
148                    let _ = client.close_impl(false).await;
149                    return Err(e);
150                }
151                if let Err(e) = client.synchronize().await {
152                    let _ = client.close_impl(false).await;
153                    return Err(e);
154                }
155                tracing::info!(num_devices, "client opened");
156                Ok((client, checker))
157            }
158            Ok(Err(cause)) => {
159                let _ = wait_rt(done_rx, join).await;
160                Err(Error::Link(cause))
161            }
162            Err(_) => {
163                let _ = wait_rt(done_rx, join).await;
164                Err(Error::RtClosed)
165            }
166        }
167    }
168
169    #[must_use]
170    pub fn num_devices(&self) -> usize {
171        self.num_devices
172    }
173
174    #[must_use]
175    pub fn geometry(&self) -> &Geometry {
176        &self.geometry
177    }
178
179    #[must_use]
180    pub fn link_stats(&self) -> LinkStats {
181        self.stats.clone()
182    }
183
184    #[must_use]
185    pub fn dc_offset_ns(&self) -> i64 {
186        self.dc_clock
187            .as_ref()
188            .and_then(DcClock::offset_ns)
189            .unwrap_or(0)
190    }
191
192    pub fn bus_time_now(&self) -> Result<DcSysTime, Error> {
193        Ok(DcSysTime::now()?.with_dc_offset(self.dc_offset_ns()))
194    }
195
196    #[must_use]
197    pub fn datagram_builder<'a>(&self) -> DatagramBuilder<'a> {
198        DatagramBuilder::with_mirror(
199            Arc::clone(&self.geometry),
200            self.mirror.clone(),
201            self.dc_clock.clone().into(),
202        )
203    }
204
205    fn mark_desynced(&self) {
206        self.mirror.desync();
207    }
208
209    fn mirror_for_response(&self) -> Option<MirrorHandle> {
210        self.mirror.enabled.then(|| self.mirror.clone())
211    }
212
213    async fn clear(&self) -> Result<(), Error> {
214        let datagrams = self.datagram_builder().push(Clear).build()?;
215        for frame in &datagrams {
216            self.send_checked(frame).await?;
217        }
218        self.mirror.set(Mirror::Synced(vec![
219            FirmwareState::boot_default();
220            self.num_devices
221        ]));
222        Ok(())
223    }
224
225    async fn send_datagrams(&self, datagrams: &[Datagram]) -> Result<ResponseFuture, Error> {
226        if datagrams.len() != self.num_devices {
227            self.mark_desynced();
228            return Err(PayloadError::DatagramCountMismatch {
229                expected: self.num_devices,
230                got: datagrams.len(),
231            }
232            .into());
233        }
234        tracing::trace!(cmd = ?datagrams[0].cmd, "sending per-device frame");
235        let mut slot = self.pool.acquire().await;
236        slot.reset(Distribution::PerDevice);
237        for (device, datagram) in datagrams.iter().enumerate() {
238            slot.payload_mut(device).copy_from_slice(&datagram.payload);
239            slot.set_cmd(device, datagram.cmd);
240        }
241        self.dispatch(slot, Reply::Ack)
242    }
243
244    async fn send_broadcast(&self, datagram: &Datagram) -> Result<ResponseFuture, Error> {
245        tracing::trace!(cmd = ?datagram.cmd, "sending broadcast frame");
246        let mut slot = self.pool.acquire().await;
247        slot.reset(Distribution::Broadcast);
248        slot.payload_mut(0).copy_from_slice(&datagram.payload);
249        slot.set_cmd(0, datagram.cmd);
250        self.dispatch(slot, Reply::Ack)
251    }
252
253    async fn send_broadcast_exclusive(&self, datagram: &Datagram) -> Result<ResponseFuture, Error> {
254        tracing::trace!(cmd = ?datagram.cmd, "sending exclusive broadcast frame");
255        let mut slot = self.pool.acquire().await;
256        slot.reset(Distribution::Broadcast);
257        slot.payload_mut(0).copy_from_slice(&datagram.payload);
258        slot.set_cmd(0, datagram.cmd);
259        self.dispatch(slot, Reply::Value)
260    }
261
262    pub async fn send(&self, frame: Frame<'_>) -> Result<ResponseFuture, Error> {
263        match frame.distribution() {
264            Distribution::Broadcast => self.send_broadcast(&frame.datagrams()[0]).await,
265            Distribution::PerDevice => self.send_datagrams(frame.datagrams()).await,
266        }
267    }
268
269    pub async fn send_checked(&self, frame: Frame<'_>) -> Result<(), Error> {
270        self.send(frame).await?.await?.check()
271    }
272
273    fn dispatch(&self, slot: pool::Slot, reply: Reply) -> Result<ResponseFuture, Error> {
274        let (response_tx, response_rx) =
275            self.completions.channel(self.mirror_for_response(), reply);
276        if self
277            .cmd_tx
278            .send(CmdMessage {
279                frame: slot,
280                response_tx,
281                exclusive: reply.exclusive(),
282            })
283            .is_err()
284        {
285            tracing::warn!("RT thread is closed; frame dropped");
286            self.mark_desynced();
287            return Err(Error::RtClosed);
288        }
289        Ok(response_rx)
290    }
291
292    async fn synchronize(&self) -> Result<(), Error> {
293        let datagrams = self.datagram_builder().push(Synchronize).build()?;
294        for frame in &datagrams {
295            self.send_checked(frame).await?;
296        }
297        Ok(())
298    }
299
300    pub async fn stop(&self) -> Result<(), Error> {
301        tracing::debug!("sending stop");
302        let buf: Vec<Vec<Emission>> = self
303            .geometry
304            .iter()
305            .map(|d| vec![Emission::NULL; d.num_transducers()])
306            .collect();
307        let datagrams = self.datagram_builder().push(Pattern::new(&buf)).build()?;
308        for frame in &datagrams {
309            self.send_checked(frame).await?;
310        }
311        Ok(())
312    }
313
314    async fn read_broadcast(&self, cmd: Cmd) -> Result<Vec<u8>, Error> {
315        self.read_broadcast_with(&Datagram::no_payload(cmd)).await
316    }
317
318    async fn read_broadcast_with(&self, datagram: &Datagram) -> Result<Vec<u8>, Error> {
319        Ok(self
320            .send_broadcast_exclusive(datagram)
321            .await?
322            .await?
323            .data()
324            .to_vec())
325    }
326
327    pub async fn read_firmware_version(&self) -> Result<Vec<FirmwareVersion>, Error> {
328        const UNKNOWN_CMD: u8 = DeviceErrorCode::UnknownCmd as u8;
329
330        let cpu_major = self.read_broadcast(Cmd::ReadCpuFwVersionMajor).await?;
331        let cpu_minor = self.read_broadcast(Cmd::ReadCpuFwVersionMinor).await?;
332        let cpu_patch = self.read_broadcast(Cmd::ReadCpuFwVersionPatch).await?;
333
334        let err_before = self.read_broadcast(Cmd::ReadErrorDetail).await?;
335        let fpga_major = self.read_broadcast(Cmd::ReadFpgaFwVersionMajor).await?;
336        let fpga_minor = self.read_broadcast(Cmd::ReadFpgaFwVersionMinor).await?;
337        let fpga_patch = self.read_broadcast(Cmd::ReadFpgaFwVersionPatch).await?;
338        let err_after_version = self.read_broadcast(Cmd::ReadErrorDetail).await?;
339        let fpga_functions = self.read_broadcast(Cmd::ReadFpgaFunctions).await?;
340        let err_after_functions = self.read_broadcast(Cmd::ReadErrorDetail).await?;
341
342        let versions: Vec<_> = (0..cpu_major.len())
343            .map(|i| {
344                let fpga = if err_after_version[i] == UNKNOWN_CMD {
345                    warn_unknown(i, "FPGA firmware version", err_before[i] == UNKNOWN_CMD);
346                    Version::UNKNOWN
347                } else {
348                    Version {
349                        major: fpga_major[i],
350                        minor: fpga_minor[i],
351                        patch: fpga_patch[i],
352                    }
353                };
354                let function_bits = if err_after_functions[i] == UNKNOWN_CMD {
355                    if err_after_version[i] != UNKNOWN_CMD {
356                        warn_unknown(i, "FPGA function bits", false);
357                    }
358                    0
359                } else {
360                    fpga_functions[i]
361                };
362                FirmwareVersion {
363                    cpu: Version {
364                        major: cpu_major[i],
365                        minor: cpu_minor[i],
366                        patch: cpu_patch[i],
367                    },
368                    fpga,
369                    function_bits,
370                }
371            })
372            .collect();
373
374        let (major, minor) = FirmwareVersion::SUPPORTED_SERIES;
375        versions
376            .iter()
377            .enumerate()
378            .filter(|(_, v)| !v.is_supported())
379            .for_each(|(device, version)| {
380                tracing::warn!(
381                    device,
382                    "firmware {version} is outside the series supported by this SDK ({major}.{minor}.x); correct operation is not guaranteed"
383                );
384            });
385
386        Ok(versions)
387    }
388
389    async fn check_firmware_version(&self, require_supported: bool) -> Result<(), Error> {
390        let versions = match self.read_firmware_version().await {
391            Ok(versions) => versions,
392            Err(e) if require_supported => return Err(e),
393            Err(e) => {
394                tracing::warn!(
395                    error = %e,
396                    "could not read the firmware version, so the series check was skipped"
397                );
398                return Ok(());
399            }
400        };
401        if !require_supported {
402            return Ok(());
403        }
404        versions
405            .into_iter()
406            .enumerate()
407            .find(|(_, version)| !version.is_supported())
408            .map_or(Ok(()), |(device, version)| {
409                Err(Error::UnsupportedFirmware { device, version })
410            })
411    }
412
413    pub async fn read_error_detail(&self) -> Result<Vec<u8>, Error> {
414        self.read_broadcast(Cmd::ReadErrorDetail).await
415    }
416
417    pub async fn read_fpga_state(&self) -> Result<Vec<FpgaState>, Error> {
418        Ok(self
419            .read_broadcast(Cmd::ReadFpgaState)
420            .await?
421            .into_iter()
422            .map(FpgaState)
423            .collect())
424    }
425
426    pub async fn read_telemetry(&self, counter: Telemetry) -> Result<Vec<u8>, Error> {
427        const INVALID_PAYLOAD: u8 = DeviceErrorCode::InvalidPayload as u8;
428
429        let mut datagram = Datagram::no_payload(Cmd::ReadTelemetry);
430        let (p, _) = ReadTelemetryPayload::mut_from_prefix(&mut datagram.payload).unwrap();
431        p.counter_id = counter.as_u8();
432
433        let err_before = self.read_broadcast(Cmd::ReadErrorDetail).await?;
434        let counters = self.read_broadcast_with(&datagram).await?;
435        let err_after = self.read_broadcast(Cmd::ReadErrorDetail).await?;
436
437        if let Some(device) = (0..counters.len()).find(|&i| err_after[i] == INVALID_PAYLOAD) {
438            if err_before[device] == INVALID_PAYLOAD {
439                warn_unknown(device, "telemetry counter", true);
440            }
441            return Err(Error::UnsupportedTelemetry { device, counter });
442        }
443        Ok(counters)
444    }
445
446    pub async fn close(&self) -> Result<(), Error> {
447        self.close_impl(true).await
448    }
449
450    async fn close_impl(&self, stop: bool) -> Result<(), Error> {
451        tracing::debug!("closing client");
452        let stopped = if stop && !self.stopping.swap(true, Ordering::AcqRel) {
453            self.stop().await
454        } else {
455            Ok(())
456        };
457        self.closed.store(true, Ordering::Release);
458        let done = self
459            .done
460            .lock()
461            .unwrap_or_else(PoisonError::into_inner)
462            .take();
463        let joined = if let Some(done) = done {
464            rt_outcome(done.await)
465        } else {
466            Ok(())
467        };
468        stopped.and(joined)
469    }
470}
471
472impl Drop for Client {
473    fn drop(&mut self) {
474        self.closed.store(true, Ordering::Release);
475        let join = self
476            .join
477            .lock()
478            .unwrap_or_else(PoisonError::into_inner)
479            .take();
480        if let Some(join) = join {
481            let _ = join.join();
482        }
483    }
484}
485
486fn close_unopened<L: Link>(link: &mut L) {
487    if let Err(e) = link.close() {
488        tracing::warn!(error = %e, "failed to close the link that never opened");
489    }
490}
491
492fn warn_unknown(device: usize, what: &str, pre_latched: bool) {
493    if pre_latched {
494        tracing::warn!(
495            device,
496            "{what} is unknown: an error was already latched before the query, so it cannot be attributed to it"
497        );
498    } else {
499        tracing::warn!(
500            device,
501            "{what} is unknown; device firmware may be out of date"
502        );
503    }
504}
505
506fn rt_outcome(done: Result<Option<LinkCause>, oneshot::Canceled>) -> Result<(), Error> {
507    match done {
508        Ok(None) => Ok(()),
509        Ok(Some(cause)) => Err(Error::Link(cause)),
510        Err(oneshot::Canceled) => Err(Error::RtPanicked),
511    }
512}
513
514async fn wait_rt(
515    done: oneshot::Receiver<Option<LinkCause>>,
516    join: JoinHandle<()>,
517) -> Result<(), Error> {
518    let outcome = rt_outcome(done.await);
519    let _ = join.join();
520    outcome
521}