1
use crate::{ Result, Error };
2
use std::io::IoSlice;
3
use std::os::unix::io::RawFd;
4
use std::net::IpAddr;
5
use std::os::unix::prelude::AsRawFd;
6
use nix::sys::socket::{self, SockaddrLike, SockaddrStorage};
7
use tokio::io::unix::AsyncFd;
8

            
9
use super::SocketAddr;
10

            
11
fn sockaddrlike_to_storage(addr: &dyn SockaddrLike) -> SockaddrStorage
12
{
13
    unsafe { SockaddrStorage::from_raw(addr.as_ptr(), Some(addr.len())) }.unwrap()
14
}
15

            
16
trait OptUnwrap<T: Sized> {
17
    fn unwrap_opt(self) -> Result<T>;
18
}
19

            
20
impl <T: Sized> OptUnwrap<T> for Option<T>
21
{
22
    fn unwrap_opt(self) -> Result<T> {
23
        self.ok_or(Error::Internal("Option is None"))
24
    }
25
}
26

            
27
#[derive(Default, Debug)]
28
struct RecvInfoOpt {
29
    size:	usize,
30
    if_idx:	Option<libc::c_int>,
31
    local:	Option<IpAddr>,
32
    remote:	Option<SocketAddr>,
33
    spec_dest:	Option<IpAddr>,
34
}
35

            
36
impl RecvInfoOpt {
37
    fn convert_v4(ip_raw: nix::libc::in_addr) -> Option<std::net::IpAddr>
38
    {
39
	Some(std::net::Ipv4Addr::from(ip_raw.s_addr.to_be()).into())
40
    }
41

            
42
    fn convert_v6(ip_raw: nix::libc::in6_addr) -> Option<std::net::IpAddr>
43
    {
44
	Some(std::net::Ipv6Addr::from(ip_raw.s6_addr).into())
45
    }
46

            
47
    pub fn set_local_v4(&mut self, ip_raw: nix::libc::in_addr) {
48
	self.local = Self::convert_v4(ip_raw);
49
    }
50

            
51
    pub fn set_local_v6(&mut self, ip_raw: nix::libc::in6_addr) {
52
	self.local = Self::convert_v6(ip_raw);
53
    }
54

            
55
    pub fn set_spec_dest_v4(&mut self, ip_raw: nix::libc::in_addr) {
56
	self.spec_dest = Self::convert_v4(ip_raw);
57
    }
58
}
59

            
60
#[derive(Clone, Copy)]
61
pub struct RecvInfo {
62
    pub size:	usize,
63
    if_idx:	libc::c_int,
64
    pub local:	IpAddr,
65
    pub remote:	SocketAddr,
66
}
67

            
68

            
69
impl RecvInfo {
70
    pub fn local(&self) -> IpAddr {
71
	// TODO: handle spec_dest?
72
	self.local
73
    }
74
}
75

            
76
impl TryFrom<RecvInfoOpt> for RecvInfo {
77
    type Error = Error;
78

            
79
    fn try_from(v: RecvInfoOpt) -> std::result::Result<Self, Self::Error> {
80
        Ok(Self {
81
	    size:	v.size,
82
	    if_idx:	v.if_idx.unwrap_opt()?,
83
	    // TODO: prefer spec_dest when set?
84
	    local:	v.local.unwrap_opt()?,
85
	    remote:	v.remote.unwrap_opt()?,
86
	})
87
    }
88
}
89

            
90
pub struct UdpSocket {
91
    fd:		RawFd,
92
    af:		socket::AddressFamily,
93
    // must be an `Option` so that we can control the destruction order of
94
    // 'fd' itself and 'async_fd' in drop()
95
    async_fd:	Option<AsyncFd<RawFd>>,
96
}
97

            
98
impl UdpSocket {
99
    fn get_fd(&self) -> &AsyncFd<RawFd>
100
    {
101
	self.async_fd.as_ref().unwrap()
102
    }
103

            
104
    pub async fn sendto(&self, buf: &[u8], addr: SocketAddr) -> Result<()>
105
    {
106
	use socket::MsgFlags as M;
107
	use nix::Error as E;
108

            
109
	let addr = addr.as_nix();
110

            
111
	loop {
112
	    let mut async_guard = self.get_fd().writable().await?;
113

            
114
	    match self.sendto_sync(buf, &*addr, M::MSG_NOSIGNAL | M::MSG_DONTWAIT) {
115
		Ok(_)			=> break Ok(()),
116
		Err(E::EAGAIN)		=> async_guard.clear_ready(),
117
		Err(e)			=> break Err(e.into())
118
	    };
119
	}
120
    }
121

            
122
    fn sendto_sync(&self, buf: &[u8], addr: &dyn SockaddrLike,
123
		   flags: socket::MsgFlags) -> nix::Result<()>
124
    {
125
	use nix::Error as E;
126

            
127
	match socket::sendto(self.fd, buf, addr, flags) {
128
	    Ok(sz) if sz == buf.len()	=> Ok(()),
129
	    Ok(sz)			=> {
130
		error!("sent only {} bytes out of {} ones", sz, buf.len());
131
		Err(E::ENOPKG)
132
	    },
133
	    Err(e)			=> Err(e)
134
	}
135
    }
136

            
137
    pub async fn sendmsg(&self, iov: &[IoSlice<'_>], addr: SocketAddr) -> Result<()>
138
    {
139
	use socket::MsgFlags as M;
140
	use nix::Error as E;
141

            
142
	let addr = addr.as_nix();
143

            
144
	loop {
145
	    let mut async_guard = self.get_fd().writable().await?;
146

            
147
	    match self.sendmsg_sync(iov, &*addr, M::MSG_NOSIGNAL | M::MSG_DONTWAIT) {
148
		Ok(_)			=> break Ok(()),
149
		Err(E::EAGAIN)		=> async_guard.clear_ready(),
150
		Err(e)			=> break Err(e.into())
151
	    }
152
	}
153
    }
154

            
155
    fn sendmsg_sync(&self, iov: &[IoSlice<'_>], addr: &dyn SockaddrLike,
156
		    flags: socket::MsgFlags) -> nix::Result<()>
157
    {
158
	use nix::Error as E;
159

            
160
	let total_sz = iov.iter().map(|v| v.len()).sum();
161

            
162
	// TODO: this is too expensive but nix api makes it difficulty/impossible
163
	// to use the `dyn SockaddrLike` object directly
164
	let addr = sockaddrlike_to_storage(addr);
165

            
166
	match socket::sendmsg(self.fd, iov, &[], flags, Some(&addr)) {
167
	    Ok(sz) if sz == total_sz	=> Ok(()),
168
	    Ok(sz)			=> {
169
		error!("sent only {} bytes out of {} ones", sz, total_sz);
170
		Err(E::ENOPKG)
171
	    },
172
	    Err(e)			=> Err(e),
173
	}
174
    }
175

            
176
    pub async fn recvfrom(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr)>
177
    {
178
	use nix::Error as E;
179

            
180
	loop {
181
	    let mut async_guard = self.get_fd().readable().await?;
182

            
183
	    match socket::recvfrom::<SockaddrStorage>(self.fd, buf) {
184
		Ok((sz, Some(addr)))	=> break Ok((sz, addr.try_into()?)),
185
		Ok((_, None))		=> break Err(Error::Internal("no address from recvfrom")),
186
		Err(E::EAGAIN)		=> async_guard.clear_ready(),
187
		Err(e)			=> break Err(e.into())
188
	    }
189
	}
190
    }
191

            
192
    pub async fn recvmsg(&self, buf: &mut [u8]) -> Result<RecvInfo>
193
    {
194
	use socket::MsgFlags as M;
195
	use nix::Error as E;
196

            
197
	loop {
198
	    let mut async_guard = self.get_fd().readable().await?;
199

            
200
	    match self.recvmsg_sync(buf, M::MSG_DONTWAIT) {
201
		Ok(info)			=> break Ok(info),
202
		Err(Error::Nix(E::EAGAIN))	=> async_guard.clear_ready(),
203
		Err(e)				=> break Err(e),
204
	    }
205
	}
206
    }
207

            
208
    fn recvmsg_sync(&self, buf: &mut [u8], flags: socket::MsgFlags) -> Result<RecvInfo>
209
    {
210

            
211
	let mut iov = [std::io::IoSliceMut::new(buf)];
212
	let mut cmsg = nix::cmsg_space!(libc::in6_pktinfo,
213
					libc::in_pktinfo);
214

            
215
	let recv = socket::recvmsg::<SockaddrStorage>(self.fd, &mut iov, Some(&mut cmsg), flags)?;
216

            
217
	let mut res = RecvInfoOpt {
218
	    size:	recv.bytes,
219
	    ..Default::default()
220
	};
221

            
222
	for msg in recv.cmsgs() {
223
	    use socket::ControlMessageOwned as C;
224

            
225
	    match msg {
226
		C::Ipv4PacketInfo(i)	=> {
227
		    res.set_local_v4(i.ipi_addr);
228
		    res.set_spec_dest_v4(i.ipi_spec_dst);
229
		    res.if_idx = Some(i.ipi_ifindex);
230
		},
231

            
232
		C::Ipv6PacketInfo(i)	=> {
233
		    res.set_local_v6(i.ipi6_addr);
234
		    res.spec_dest = None;
235
		    res.if_idx = Some(i.ipi6_ifindex as libc::c_int);
236
		},
237

            
238
		m			=> {
239
		    debug!("unhandled msg {:?}", m);
240
		},
241
	    }
242
	}
243

            
244
	match recv.address {
245
	    Some(addr)	=> res.remote = Some(SocketAddr::try_from(addr)?),
246
	    None	=> {
247
		warn!("missing remote address");
248
		return Err(Error::Internal("missing remote address"));
249
	    },
250
	};
251

            
252
	res.try_into()
253
    }
254

            
255
    pub fn bind(addr: SocketAddr) -> Result<Self> {
256
	let fd = unsafe { addr.socket() }?;
257

            
258
	let af = addr.get_af();
259
	let addr = addr.as_nix();
260

            
261
	match socket::bind(fd, &*addr) {
262
	    Ok(_)	=> Ok(Self {
263
		fd:		fd,
264
		af:		af,
265
		async_fd:	Some(AsyncFd::new(fd)?),
266
	    }),
267

            
268
	    Err(e)	=> {
269
		unsafe { libc::close(fd) };
270
		Err(std::io::Error::from(e).into())
271
	    }
272
	}
273
    }
274

            
275
    pub fn from_raw(fd: RawFd) -> Result<Self> {
276
	let addr = SocketAddr::from_raw_fd(fd)?;
277

            
278
	Ok(Self {
279
	    fd:		fd,
280
	    af:		addr.get_af(),
281
	    async_fd:	Some(AsyncFd::new(fd)?),
282
	})
283
    }
284

            
285
    pub fn local_addr(&self) -> Result<SocketAddr> {
286
	let addr: SockaddrStorage = socket::getsockname(self.fd.as_raw_fd())?;
287

            
288
	addr.try_into()
289
    }
290

            
291
    pub fn set_request_pktinfo(&mut self) -> Result<()> {
292
	use socket::AddressFamily as AF;
293
	use nix::sys::socket::sockopt as O;
294

            
295
	match self.af {
296
	    AF::Inet	=> socket::setsockopt(self.fd, O::Ipv4PacketInfo, &true),
297
	    AF::Inet6	=> socket::setsockopt(self.fd, O::Ipv6RecvPacketInfo, &true),
298
	    _		=> return Err(Error::Internal("unexpected af")),
299
	}?;
300

            
301
	Ok(())
302
    }
303

            
304
    pub fn set_nonblocking(&self) -> Result<()> {
305
	let rc = unsafe { libc::fcntl(self.fd, libc::F_GETFL) };
306

            
307
	if rc < 0 {
308
	    return Err(std::io::Error::last_os_error().into());
309
	}
310

            
311
	let flags = rc as u32;
312

            
313
	if flags & (libc::O_NONBLOCK as u32) != 0 {
314
	    return Ok(());
315
	}
316

            
317
	let rc = unsafe { libc::fcntl(self.fd, libc::F_SETFL, flags | (libc::O_NONBLOCK as u32)) };
318

            
319
	if rc < 0 {
320
	    return Err(std::io::Error::last_os_error().into());
321
	}
322

            
323
	Ok(())
324
    }
325
}
326

            
327
impl Drop for UdpSocket {
328
    fn drop(&mut self) {
329
	self.async_fd = None;
330
        unsafe { libc::close(self.fd) };
331
    }
332
}