Skip to main content

mavlink_core/connection/udp/
sync.rs

1//! UDP MAVLink connection
2
3use crate::Connectable;
4use crate::MAVLinkMessageRaw;
5use crate::connection::get_socket_addr;
6use crate::connection::{Connection, MavConnection};
7use crate::connection_shared::{
8    ConnectionState, next_send_header, read_message, read_raw_message, write_message,
9    write_raw_message,
10};
11use crate::peek_reader::PeekReader;
12use crate::{MavHeader, MavlinkVersion, Message};
13use core::ops::DerefMut;
14use std::collections::VecDeque;
15use std::io::{self, Read, Write};
16use std::net::{SocketAddr, UdpSocket};
17use std::sync::Mutex;
18
19#[cfg(feature = "mav2-message-signing")]
20use crate::SigningConfig;
21
22use super::config::{UdpConfig, UdpMode};
23
24struct UdpRead {
25    socket: UdpSocket,
26    buffer: VecDeque<u8>,
27    last_recv_address: Option<SocketAddr>,
28}
29
30const MTU_SIZE: usize = 1500;
31impl Read for UdpRead {
32    fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
33        if !self.buffer.is_empty() {
34            self.buffer.read(buf)
35        } else {
36            let mut read_buffer = [0u8; MTU_SIZE];
37            let (n_buffer, address) = self.socket.recv_from(&mut read_buffer)?;
38            let n = (&read_buffer[0..n_buffer]).read(buf)?;
39            self.buffer.extend(&read_buffer[n..n_buffer]);
40
41            self.last_recv_address = Some(address);
42            Ok(n)
43        }
44    }
45}
46
47struct UdpWrite {
48    socket: UdpSocket,
49    dest: Option<SocketAddr>,
50    connected: bool,
51    sequence: u8,
52}
53
54impl Write for UdpWrite {
55    fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
56        let addr = self.dest.expect("`dest` is checked before write");
57        if self.connected {
58            self.socket.send(buf)
59        } else {
60            self.socket.send_to(buf, addr)
61        }
62    }
63
64    fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
65        if self.write(buf)? != buf.len() {
66            return Err(io::Error::new(
67                io::ErrorKind::WriteZero,
68                "failed to send complete UDP datagram",
69            ));
70        }
71
72        Ok(())
73    }
74
75    fn flush(&mut self) -> io::Result<()> {
76        Ok(())
77    }
78}
79
80pub struct UdpConnection {
81    reader: Mutex<PeekReader<UdpRead>>,
82    writer: Mutex<UdpWrite>,
83    state: ConnectionState,
84    server: bool,
85}
86
87impl UdpConnection {
88    fn new(
89        socket: UdpSocket,
90        server: bool,
91        dest: Option<SocketAddr>,
92        connected: bool,
93    ) -> io::Result<Self> {
94        Ok(Self {
95            server,
96            reader: Mutex::new(PeekReader::new(UdpRead {
97                socket: socket.try_clone()?,
98                buffer: VecDeque::new(),
99                last_recv_address: None,
100            })),
101            writer: Mutex::new(UdpWrite {
102                socket,
103                dest,
104                connected,
105                sequence: 0,
106            }),
107            state: ConnectionState::new(),
108        })
109    }
110
111    fn update_reply_destination(&self, reader: &PeekReader<UdpRead>) {
112        if self.server {
113            if let addr @ Some(_) = reader.reader_ref().last_recv_address {
114                self.writer.lock().unwrap().dest = addr;
115            }
116        }
117    }
118}
119
120impl<M: Message> MavConnection<M> for UdpConnection {
121    fn recv(&self) -> Result<(MavHeader, M), crate::error::MessageReadError> {
122        let mut reader = self.reader.lock().unwrap();
123
124        let result = read_message::<M, _>(reader.deref_mut(), &self.state);
125        self.update_reply_destination(&reader);
126        result
127    }
128
129    fn recv_raw(&self) -> Result<MAVLinkMessageRaw, crate::error::MessageReadError> {
130        let mut reader = self.reader.lock().unwrap();
131
132        let result = read_raw_message::<M, _>(reader.deref_mut(), &self.state);
133        self.update_reply_destination(&reader);
134        result
135    }
136
137    fn try_recv(&self) -> Result<(MavHeader, M), crate::error::MessageReadError> {
138        let mut reader = self.reader.lock().unwrap();
139        reader.reader_mut().socket.set_nonblocking(true)?;
140
141        let result = read_message::<M, _>(reader.deref_mut(), &self.state);
142        self.update_reply_destination(&reader);
143
144        reader.reader_mut().socket.set_nonblocking(false)?;
145
146        result
147    }
148
149    fn send(&self, header: &MavHeader, data: &M) -> Result<usize, crate::error::MessageWriteError> {
150        let mut guard = self.writer.lock().unwrap();
151        let writer = &mut *guard;
152
153        let header = next_send_header(&mut writer.sequence, header);
154
155        let len = if writer.dest.is_some() {
156            write_message(writer, &self.state, header, data)?
157        } else {
158            0
159        };
160
161        Ok(len)
162    }
163
164    fn send_raw(&self, data: &MAVLinkMessageRaw) -> Result<usize, crate::error::MessageWriteError> {
165        let mut guard = self.writer.lock().unwrap();
166        let writer = &mut *guard;
167
168        let len = if writer.dest.is_some() {
169            write_raw_message(writer, data)?
170        } else {
171            0
172        };
173
174        Ok(len)
175    }
176
177    fn set_protocol_version(&mut self, version: MavlinkVersion) {
178        self.state.set_protocol_version(version);
179    }
180
181    fn protocol_version(&self) -> MavlinkVersion {
182        self.state.protocol_version()
183    }
184
185    fn set_allow_recv_any_version(&mut self, allow: bool) {
186        self.state.set_allow_recv_any_version(allow);
187    }
188
189    fn allow_recv_any_version(&self) -> bool {
190        self.state.allow_recv_any_version()
191    }
192
193    #[cfg(feature = "mav2-message-signing")]
194    fn setup_signing(&mut self, signing_data: Option<SigningConfig>) {
195        self.state.setup_signing(signing_data);
196    }
197}
198
199impl Connectable for UdpConfig {
200    fn connect<M: Message>(&self) -> io::Result<Connection<M>> {
201        let (addr, server, dest): (&str, _, _) = match self.mode {
202            UdpMode::Udpin => (&self.address, true, None),
203            _ => ("0.0.0.0:0", false, Some(get_socket_addr(&self.address)?)),
204        };
205        let (socket, connected) = match self.take_socket()? {
206            Some(socket) => (socket, !server),
207            None => (UdpSocket::bind(addr)?, false),
208        };
209        if let Some(timeout) = self.read_timeout {
210            socket.set_read_timeout(Some(timeout))?;
211        }
212        if matches!(self.mode, UdpMode::UdpBroadcast) {
213            socket.set_broadcast(true)?;
214        }
215        Ok(UdpConnection::new(socket, server, dest, connected)?.into())
216    }
217}
218
219#[cfg(test)]
220mod tests {
221    use super::*;
222
223    #[test]
224    fn test_datagram_buffering() {
225        let receiver_socket = UdpSocket::bind("127.0.0.1:5000").unwrap();
226        let mut udp_reader = UdpRead {
227            socket: receiver_socket.try_clone().unwrap(),
228            buffer: VecDeque::new(),
229            last_recv_address: None,
230        };
231        let sender_socket = UdpSocket::bind("0.0.0.0:0").unwrap();
232        sender_socket.connect("127.0.0.1:5000").unwrap();
233
234        let datagram: Vec<u8> = (0..50).collect::<Vec<_>>();
235
236        let mut n_sent = sender_socket.send(&datagram).unwrap();
237        assert_eq!(n_sent, datagram.len());
238        n_sent = sender_socket.send(&datagram).unwrap();
239        assert_eq!(n_sent, datagram.len());
240
241        let mut buf = [0u8; 30];
242
243        let mut n_read = udp_reader.read(&mut buf).unwrap();
244        assert_eq!(n_read, 30);
245        assert_eq!(&buf[0..n_read], (0..30).collect::<Vec<_>>().as_slice());
246
247        n_read = udp_reader.read(&mut buf).unwrap();
248        assert_eq!(n_read, 20);
249        assert_eq!(&buf[0..n_read], (30..50).collect::<Vec<_>>().as_slice());
250
251        n_read = udp_reader.read(&mut buf).unwrap();
252        assert_eq!(n_read, 30);
253        assert_eq!(&buf[0..n_read], (0..30).collect::<Vec<_>>().as_slice());
254
255        n_read = udp_reader.read(&mut buf).unwrap();
256        assert_eq!(n_read, 20);
257        assert_eq!(&buf[0..n_read], (30..50).collect::<Vec<_>>().as_slice());
258    }
259}