Skip to main content

mavlink_core/connection/udp/
async.rs

1//! Async UDP MAVLink connection
2
3use core::{ops::DerefMut, task::Poll};
4use std::io;
5use std::{collections::VecDeque, io::Read, sync::Arc};
6
7use async_trait::async_trait;
8use futures::lock::Mutex;
9use tokio::{
10    io::{AsyncRead, AsyncWrite, ReadBuf},
11    net::UdpSocket,
12};
13
14use crate::MAVLinkMessageRaw;
15use crate::connection::udp::config::{UdpConfig, UdpMode};
16use crate::connection_shared::{
17    ConnectionState, next_send_header, read_message_async, read_raw_message_async,
18    write_message_async, write_raw_message_async,
19};
20use crate::{MavHeader, MavlinkVersion, Message, async_peek_reader::AsyncPeekReader};
21
22use crate::connection::{AsyncConnectable, AsyncMavConnection, get_socket_addr};
23
24#[cfg(feature = "mav2-message-signing")]
25use crate::SigningConfig;
26
27struct UdpRead {
28    socket: Arc<UdpSocket>,
29    buffer: VecDeque<u8>,
30    last_recv_address: Option<std::net::SocketAddr>,
31}
32
33const MTU_SIZE: usize = 1500;
34impl AsyncRead for UdpRead {
35    fn poll_read(
36        mut self: core::pin::Pin<&mut Self>,
37        cx: &mut core::task::Context<'_>,
38        buf: &mut ReadBuf<'_>,
39    ) -> Poll<io::Result<()>> {
40        if self.buffer.is_empty() {
41            let mut read_buffer = [0u8; MTU_SIZE];
42            let mut read_buffer = ReadBuf::new(&mut read_buffer);
43
44            match self.socket.poll_recv_from(cx, &mut read_buffer) {
45                Poll::Ready(Ok(address)) => {
46                    let n_buffer = read_buffer.filled().len();
47
48                    let n = (&read_buffer.filled()[0..n_buffer]).read(buf.initialize_unfilled())?;
49                    buf.advance(n);
50
51                    self.buffer.extend(&read_buffer.filled()[n..n_buffer]);
52                    self.last_recv_address = Some(address);
53                    Poll::Ready(Ok(()))
54                }
55                Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
56                Poll::Pending => Poll::Pending,
57            }
58        } else {
59            let read_result = self.buffer.read(buf.initialize_unfilled());
60            let result = match read_result {
61                Ok(n) => {
62                    buf.advance(n);
63                    Ok(())
64                }
65                Err(err) => Err(err),
66            };
67            Poll::Ready(result)
68        }
69    }
70}
71
72struct UdpWrite {
73    socket: Arc<UdpSocket>,
74    dest: Option<std::net::SocketAddr>,
75    connected: bool,
76    sequence: u8,
77}
78
79impl AsyncWrite for UdpWrite {
80    fn poll_write(
81        self: core::pin::Pin<&mut Self>,
82        cx: &mut core::task::Context<'_>,
83        buf: &[u8],
84    ) -> Poll<io::Result<usize>> {
85        let this = self.get_mut();
86        let addr = this.dest.expect("`dest` is checked before write");
87
88        let result = if this.connected {
89            this.socket.poll_send(cx, buf)
90        } else {
91            this.socket.poll_send_to(cx, buf, addr)
92        };
93        match result {
94            Poll::Ready(Ok(written)) if written == buf.len() => Poll::Ready(Ok(written)),
95            Poll::Ready(Ok(_)) => Poll::Ready(Err(io::Error::new(
96                io::ErrorKind::WriteZero,
97                "failed to send complete UDP datagram",
98            ))),
99            Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
100            Poll::Pending => Poll::Pending,
101        }
102    }
103
104    fn poll_flush(
105        self: core::pin::Pin<&mut Self>,
106        _cx: &mut core::task::Context<'_>,
107    ) -> Poll<io::Result<()>> {
108        Poll::Ready(Ok(()))
109    }
110
111    fn poll_shutdown(
112        self: core::pin::Pin<&mut Self>,
113        _cx: &mut core::task::Context<'_>,
114    ) -> Poll<io::Result<()>> {
115        Poll::Ready(Ok(()))
116    }
117}
118
119pub struct AsyncUdpConnection {
120    reader: Mutex<AsyncPeekReader<UdpRead>>,
121    writer: Mutex<UdpWrite>,
122    state: ConnectionState,
123    server: bool,
124}
125
126impl AsyncUdpConnection {
127    fn new(
128        socket: UdpSocket,
129        server: bool,
130        dest: Option<std::net::SocketAddr>,
131        connected: bool,
132    ) -> io::Result<Self> {
133        let socket = Arc::new(socket);
134        Ok(Self {
135            server,
136            reader: Mutex::new(AsyncPeekReader::new(UdpRead {
137                socket: socket.clone(),
138                buffer: VecDeque::new(),
139                last_recv_address: None,
140            })),
141            writer: Mutex::new(UdpWrite {
142                socket,
143                dest,
144                connected,
145                sequence: 0,
146            }),
147            state: ConnectionState::new(),
148        })
149    }
150
151    async fn update_reply_destination(&self, reader: &mut AsyncPeekReader<UdpRead>) {
152        if self.server {
153            if let addr @ Some(_) = reader.reader_ref().last_recv_address {
154                self.writer.lock().await.dest = addr;
155            }
156        }
157    }
158}
159
160#[async_trait::async_trait]
161impl<M: Message + Sync + Send> AsyncMavConnection<M> for AsyncUdpConnection {
162    async fn recv(&self) -> Result<(MavHeader, M), crate::error::MessageReadError> {
163        let mut reader = self.reader.lock().await;
164        loop {
165            let result = read_message_async::<M, _>(reader.deref_mut(), &self.state).await;
166            self.update_reply_destination(reader.deref_mut()).await;
167            if let ok @ Ok(..) = result {
168                return ok;
169            }
170        }
171    }
172
173    async fn recv_raw(&self) -> Result<MAVLinkMessageRaw, crate::error::MessageReadError> {
174        let mut reader = self.reader.lock().await;
175        loop {
176            let result = read_raw_message_async::<M, _>(reader.deref_mut(), &self.state).await;
177            self.update_reply_destination(reader.deref_mut()).await;
178            if let ok @ Ok(..) = result {
179                return ok;
180            }
181        }
182    }
183
184    async fn try_recv(&self) -> Result<(MavHeader, M), crate::error::MessageReadError> {
185        let mut reader = self.reader.lock().await;
186        let result = read_message_async::<M, _>(reader.deref_mut(), &self.state).await;
187        self.update_reply_destination(reader.deref_mut()).await;
188
189        result
190    }
191
192    async fn send(
193        &self,
194        header: &MavHeader,
195        data: &M,
196    ) -> Result<usize, crate::error::MessageWriteError> {
197        let mut guard = self.writer.lock().await;
198        let writer = &mut *guard;
199
200        let header = next_send_header(&mut writer.sequence, header);
201
202        let len = if writer.dest.is_some() {
203            write_message_async(writer, &self.state, header, data).await?
204        } else {
205            0
206        };
207
208        Ok(len)
209    }
210
211    async fn send_raw(
212        &self,
213        data: &MAVLinkMessageRaw,
214    ) -> Result<usize, crate::error::MessageWriteError> {
215        let mut guard = self.writer.lock().await;
216        let writer = &mut *guard;
217
218        let len = if writer.dest.is_some() {
219            write_raw_message_async(writer, data).await?
220        } else {
221            0
222        };
223
224        Ok(len)
225    }
226
227    fn set_protocol_version(&mut self, version: MavlinkVersion) {
228        self.state.set_protocol_version(version);
229    }
230
231    fn protocol_version(&self) -> MavlinkVersion {
232        self.state.protocol_version()
233    }
234
235    fn set_allow_recv_any_version(&mut self, allow: bool) {
236        self.state.set_allow_recv_any_version(allow);
237    }
238
239    fn allow_recv_any_version(&self) -> bool {
240        self.state.allow_recv_any_version()
241    }
242
243    #[cfg(feature = "mav2-message-signing")]
244    fn setup_signing(&mut self, signing_data: Option<SigningConfig>) {
245        self.state.setup_signing(signing_data);
246    }
247}
248
249#[async_trait]
250impl AsyncConnectable for UdpConfig {
251    async fn connect_async<M>(&self) -> io::Result<Box<dyn AsyncMavConnection<M> + Sync + Send>>
252    where
253        M: Message + Sync + Send,
254    {
255        let (addr, server, dest): (&str, _, _) = match self.mode {
256            UdpMode::Udpin => (&self.address, true, None),
257            _ => ("0.0.0.0:0", false, Some(get_socket_addr(&self.address)?)),
258        };
259        let (socket, connected) = match self.take_socket()? {
260            Some(socket) => {
261                socket.set_nonblocking(true)?;
262                (UdpSocket::from_std(socket)?, !server)
263            }
264            None => (UdpSocket::bind(addr).await?, false),
265        };
266        if matches!(self.mode, UdpMode::UdpBroadcast) {
267            socket.set_broadcast(true)?;
268        }
269        Ok(Box::new(AsyncUdpConnection::new(
270            socket, server, dest, connected,
271        )?))
272    }
273}
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278    use tokio::io::AsyncReadExt;
279
280    #[tokio::test]
281    async fn test_datagram_buffering() {
282        let receiver_socket = Arc::new(UdpSocket::bind("127.0.0.1:5001").await.unwrap());
283        let mut udp_reader = UdpRead {
284            socket: receiver_socket.clone(),
285            buffer: VecDeque::new(),
286            last_recv_address: None,
287        };
288        let sender_socket = UdpSocket::bind("0.0.0.0:0").await.unwrap();
289        sender_socket.connect("127.0.0.1:5001").await.unwrap();
290
291        let datagram: Vec<u8> = (0..50).collect::<Vec<_>>();
292
293        let mut n_sent = sender_socket.send(&datagram).await.unwrap();
294        assert_eq!(n_sent, datagram.len());
295        n_sent = sender_socket.send(&datagram).await.unwrap();
296        assert_eq!(n_sent, datagram.len());
297
298        let mut buf = [0u8; 30];
299
300        let mut n_read = udp_reader.read(&mut buf).await.unwrap();
301        assert_eq!(n_read, 30);
302        assert_eq!(&buf[0..n_read], (0..30).collect::<Vec<_>>().as_slice());
303
304        n_read = udp_reader.read(&mut buf).await.unwrap();
305        assert_eq!(n_read, 20);
306        assert_eq!(&buf[0..n_read], (30..50).collect::<Vec<_>>().as_slice());
307
308        n_read = udp_reader.read(&mut buf).await.unwrap();
309        assert_eq!(n_read, 30);
310        assert_eq!(&buf[0..n_read], (0..30).collect::<Vec<_>>().as_slice());
311
312        n_read = udp_reader.read(&mut buf).await.unwrap();
313        assert_eq!(n_read, 20);
314        assert_eq!(&buf[0..n_read], (30..50).collect::<Vec<_>>().as_slice());
315    }
316}