mavlink_core/connection/udp/
async.rs1use 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}