1use crate::runtime::with_ambient_tokio_runtime;
2use crate::sockets::{
3 ErrorCode, SocketAddrCheck, SocketAddrUse, SocketAddressFamily, WasiSocketsCtx,
4 get_receive_buffer_size, get_send_buffer_size, get_unicast_hop_limit, is_valid_address_family,
5 is_valid_remote_address, set_receive_buffer_size, set_send_buffer_size, set_unicast_hop_limit,
6 unspecified_addr,
7};
8use rustix::fd::AsFd;
9use rustix::io::Errno;
10use std::net::SocketAddr;
11use std::sync::Arc;
12use tracing::debug;
13
14pub(crate) const MAX_DATAGRAM_SIZE: usize = u16::MAX as usize;
18
19pub struct UdpSocket {
24 socket: Arc<tokio::net::UdpSocket>,
25 family: SocketAddressFamily,
26
27 permissions: SocketAddrCheck,
29
30 is_bound: bool,
33
34 remote_addr: Option<SocketAddr>,
37}
38
39impl UdpSocket {
40 pub(crate) async fn new(
42 cx: &WasiSocketsCtx,
43 family: SocketAddressFamily,
44 ) -> Result<Self, ErrorCode> {
45 cx.allowed_network_uses.check_allowed_udp()?;
46
47 let socket = with_ambient_tokio_runtime(|| socket(family))?;
48
49 socket.writable().await?;
59
60 Ok(Self {
61 socket: Arc::new(socket),
62 is_bound: false,
63 remote_addr: None,
64 permissions: cx.socket_addr_check.clone(),
65 family,
66 })
67 }
68
69 pub(crate) fn is_bound(&mut self) -> bool {
70 if !self.is_bound {
74 self.is_bound = self
75 .socket
76 .local_addr()
77 .is_ok_and(|addr| addr != unspecified_addr(self.family));
78 }
79 self.is_bound
80 }
81
82 pub(crate) fn local_address(&mut self) -> Result<SocketAddr, ErrorCode> {
83 if !self.is_bound() {
84 return Err(ErrorCode::InvalidState);
85 }
86 self.socket.local_addr().map_err(|e| e.into())
87 }
88
89 pub(crate) fn remote_address(&mut self) -> Result<SocketAddr, ErrorCode> {
90 self.remote_addr.ok_or(ErrorCode::InvalidState)
91 }
92
93 pub(crate) fn is_connected(&mut self) -> bool {
94 self.remote_addr.is_some()
95 }
96
97 pub(crate) async fn bind(&mut self, addr: SocketAddr) -> Result<(), ErrorCode> {
98 if self.is_bound() {
99 return Err(ErrorCode::InvalidState);
100 }
101 if !is_valid_address_family(addr.ip(), self.family) {
102 return Err(ErrorCode::InvalidArgument);
103 }
104
105 self.permissions.check(addr, SocketAddrUse::UdpBind).await?;
106
107 bind(&self.socket, addr)?;
108 Ok(())
109 }
110
111 pub(crate) async fn connect(&mut self, addr: SocketAddr) -> Result<(), ErrorCode> {
112 if !is_valid_address_family(addr.ip(), self.family) || !is_valid_remote_address(addr) {
113 return Err(ErrorCode::InvalidArgument);
114 }
115
116 {
118 if !self.is_bound() {
119 let implicit_bind_addr = unspecified_addr(self.family);
122 self.permissions
123 .check(implicit_bind_addr, SocketAddrUse::UdpBind)
124 .await?;
125 }
126
127 if self
132 .permissions
133 .check(addr, SocketAddrUse::UdpSend)
134 .await
135 .is_err()
136 {
137 self.permissions
138 .check(addr, SocketAddrUse::UdpReceive)
139 .await?;
140 }
141 }
142
143 let result = connect(&self.socket, addr);
144 self.update_remote_address();
145 result.map_err(|e| e.into())
146 }
147
148 pub(crate) fn disconnect(&mut self) -> Result<(), ErrorCode> {
149 if !self.is_connected() {
150 return Err(ErrorCode::InvalidState);
151 }
152
153 self.is_bound = true;
160
161 let result = disconnect(&self.socket);
162 self.update_remote_address();
163 result.map_err(|e| e.into())
164 }
165
166 fn update_remote_address(&mut self) {
170 self.remote_addr = if let Ok(addr) = self.socket.peer_addr()
171 && addr != unspecified_addr(self.family)
172 {
173 Some(addr)
174 } else {
175 None
176 }
177 }
178
179 pub(crate) fn send(
180 &mut self,
181 data: Vec<u8>,
182 addr: Option<SocketAddr>,
183 ) -> impl Future<Output = Result<(), ErrorCode>> + Send + use<> {
184 let family = self.family;
185 let socket = self.socket.clone();
186 let permissions = self.permissions.clone();
187 let connected_addr = self.remote_address().ok();
188 let is_bound = self.is_bound();
189
190 async move {
191 if data.len() > MAX_DATAGRAM_SIZE {
192 return Err(ErrorCode::DatagramTooLarge);
193 }
194
195 let effective_addr = if let Some(addr) = addr {
196 if !is_valid_remote_address(addr) || !is_valid_address_family(addr.ip(), family) {
197 return Err(ErrorCode::InvalidArgument);
198 }
199
200 if connected_addr.is_some() && connected_addr != Some(addr) {
203 return Err(ErrorCode::InvalidArgument);
204 }
205
206 addr
207 } else if let Some(connected_addr) = connected_addr {
208 connected_addr
209 } else {
210 return Err(ErrorCode::InvalidArgument);
211 };
212
213 {
215 if !is_bound {
216 let implicit_bind_addr = unspecified_addr(family);
219 permissions
220 .check(implicit_bind_addr, SocketAddrUse::UdpBind)
221 .await?;
222 }
223
224 permissions
225 .check(effective_addr, SocketAddrUse::UdpSend)
226 .await?;
227 }
228
229 if connected_addr == Some(effective_addr) {
230 socket.send(&data).await?;
231 } else {
232 socket.send_to(&data, effective_addr).await?;
233 }
234
235 Ok(())
236 }
237 }
238
239 pub(crate) fn recv(
240 &mut self,
241 ) -> impl Future<Output = Result<(Vec<u8>, SocketAddr), ErrorCode>> + Send + use<> {
242 let socket = self.socket.clone();
243 let permissions = self.permissions.clone();
244 let is_bound = self.is_bound();
245
246 async move {
247 if !is_bound {
248 return Err(ErrorCode::InvalidState);
249 }
250
251 loop {
252 let mut data = vec![0; MAX_DATAGRAM_SIZE];
253 let (len, addr) = socket.recv_from(&mut data).await?;
254 data.truncate(len);
255
256 match permissions.check(addr, SocketAddrUse::UdpReceive).await {
257 Ok(()) => return Ok((data, addr)),
258 Err(_) => {
259 continue;
261 }
262 }
263 }
264 }
265 }
266
267 pub(crate) fn address_family(&self) -> SocketAddressFamily {
268 self.family
269 }
270
271 pub(crate) fn unicast_hop_limit(&self) -> Result<u8, ErrorCode> {
272 let n = get_unicast_hop_limit(&self.socket, self.family)?;
273 Ok(n)
274 }
275
276 pub(crate) fn set_unicast_hop_limit(&self, value: u8) -> Result<(), ErrorCode> {
277 set_unicast_hop_limit(&self.socket, self.family, value)?;
278 Ok(())
279 }
280
281 pub(crate) fn receive_buffer_size(&self) -> Result<u64, ErrorCode> {
282 let n = get_receive_buffer_size(&self.socket)?;
283 Ok(n)
284 }
285
286 pub(crate) fn set_receive_buffer_size(&self, value: u64) -> Result<(), ErrorCode> {
287 set_receive_buffer_size(&self.socket, value)?;
288 Ok(())
289 }
290
291 pub(crate) fn send_buffer_size(&self) -> Result<u64, ErrorCode> {
292 let n = get_send_buffer_size(&self.socket)?;
293 Ok(n)
294 }
295
296 pub(crate) fn set_send_buffer_size(&self, value: u64) -> Result<(), ErrorCode> {
297 set_send_buffer_size(&self.socket, value)?;
298 Ok(())
299 }
300}
301
302fn socket(family: SocketAddressFamily) -> std::io::Result<tokio::net::UdpSocket> {
304 #[cfg(windows)]
306 static INIT: std::sync::Once = std::sync::Once::new();
307 #[cfg(windows)]
308 INIT.call_once(|| {
309 let _ = std::net::TcpStream::connect(std::net::SocketAddrV4::new(
310 std::net::Ipv4Addr::UNSPECIFIED,
311 0,
312 ));
313 });
314
315 #[cfg(not(any(windows, target_vendor = "apple")))]
316 let flags = rustix::net::SocketFlags::CLOEXEC | rustix::net::SocketFlags::NONBLOCK;
317 #[cfg(any(windows, target_vendor = "apple"))]
318 let flags = rustix::net::SocketFlags::empty();
319
320 let socket = rustix::net::socket_with(
321 match family {
322 SocketAddressFamily::Ipv4 => rustix::net::AddressFamily::INET,
323 SocketAddressFamily::Ipv6 => rustix::net::AddressFamily::INET6,
324 },
325 rustix::net::SocketType::DGRAM,
326 flags,
327 None,
328 )?;
329 #[cfg(target_vendor = "apple")]
330 rustix::io::ioctl_fioclex(&socket)?;
331 #[cfg(any(windows, target_vendor = "apple"))]
332 rustix::io::ioctl_fionbio(&socket, true)?;
333
334 if family == SocketAddressFamily::Ipv6 {
338 rustix::net::sockopt::set_ipv6_v6only(&socket, true)?;
339 }
340
341 Ok(tokio::net::UdpSocket::try_from(std::net::UdpSocket::from(
342 socket,
343 ))?)
344}
345
346fn bind(sockfd: impl AsFd, addr: SocketAddr) -> Result<(), Errno> {
347 rustix::net::bind(sockfd, &addr).map_err(|err| match err {
348 #[cfg(windows)]
351 Errno::NOBUFS => Errno::ADDRINUSE,
352 Errno::AFNOSUPPORT => Errno::INVAL,
359 _ => err,
360 })
361}
362
363fn connect(sockfd: impl AsFd, addr: SocketAddr) -> Result<(), Errno> {
364 match rustix::net::connect(sockfd.as_fd(), &addr) {
365 #[cfg(target_os = "linux")]
375 Err(Errno::INVAL) => {
376 _ = disconnect(sockfd.as_fd());
377 return rustix::net::connect(sockfd.as_fd(), &addr);
378 }
379 Err(Errno::AFNOSUPPORT) => Err(Errno::INVAL),
384 Err(Errno::INPROGRESS) => {
387 debug!("UDP connect returned EINPROGRESS, which should never happen");
388 Ok(())
389 }
390 r => r,
391 }
392}
393
394fn disconnect(sockfd: impl AsFd) -> Result<(), Errno> {
395 match rustix::net::connect_unspec(sockfd) {
396 #[cfg(target_os = "macos")]
408 Err(Errno::INVAL | Errno::AFNOSUPPORT) => Ok(()),
409 r => r,
410 }
411}