1use std::{
7 fs::File,
8 io::{Cursor, Error, ErrorKind, Read, Write},
9 net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket},
10 time::Duration,
11};
12
13macro_rules! ensure {
14 ($cond:expr, $msg:literal $(,)?) => {
15 if !$cond {
16 return Err(Error::new(ErrorKind::InvalidInput, $msg.to_string()));
17 }
18 };
19}
20
21const CLASS_IN: u16 = 0x0001;
22const TYPE_A: u16 = 0x0001;
23const TYPE_AAAA: u16 = 0x001C;
24
25static QNAME: &[&str] = &["whatismyip", "on", "quad9", "net"];
27const IPV4_ADDRESS: Ipv4Addr = Ipv4Addr::new(9, 9, 9, 9);
28const IPV6_ADDRESS: Ipv6Addr = Ipv6Addr::new(0x2620, 0xfe, 0, 0, 0, 0, 0, 0xfe);
29
30pub enum Preference {
31 Ipv4,
32 Ipv6,
33}
34
35pub fn get_both() -> (Option<Ipv4Addr>, Option<Ipv6Addr>) {
36 let v4 = Request::start(IPV4_ADDRESS.into());
37 let v6 = Request::start(IPV6_ADDRESS.into());
38 (
39 v4.and_then(Request::read_response).map(Ipv4Addr::from).ok(),
40 v6.and_then(Request::read_response).map(Ipv6Addr::from).ok(),
41 )
42}
43
44pub fn get_any(preference: Preference) -> Option<IpAddr> {
45 let (v4, v6) = get_both();
46 let (v4, v6) = (v4.map(IpAddr::from), v6.map(IpAddr::from));
47 match preference {
48 Preference::Ipv4 => v4.or(v6),
49 Preference::Ipv6 => v6.or(v4),
50 }
51}
52
53struct Request {
54 socket: UdpSocket,
55 id: [u8; 2],
56 buf: [u8; 1500],
57 record_type: u16,
58}
59
60impl Request {
61 fn start(resolver_ip: IpAddr) -> Result<Self, Error> {
62 let (addr, record_type) = if resolver_ip.is_ipv4() {
63 (Ipv4Addr::UNSPECIFIED.into(), TYPE_A)
64 } else {
65 (Ipv6Addr::UNSPECIFIED.into(), TYPE_AAAA)
66 };
67 let socket = UdpSocket::bind(SocketAddr::new(addr, 0))?;
68 socket.set_read_timeout(Some(Duration::from_millis(500)))?;
69 let endpoint = SocketAddr::new(resolver_ip, 53);
70
71 let id = get_id()?;
72 let mut buf = [0u8; 1500];
73 let mut cursor = Cursor::new(&mut buf[..]);
74 cursor.write_all(&id)?;
75 cursor.write_all(&0x0100u16.to_be_bytes())?; cursor.write_all(&0x0001u16.to_be_bytes())?; cursor.write_all(&0x0000u16.to_be_bytes())?; cursor.write_all(&0x0000u16.to_be_bytes())?; cursor.write_all(&0x0000u16.to_be_bytes())?; for atom in QNAME {
81 cursor.write_all(&[atom.len() as u8])?;
83 cursor.write_all(atom.as_bytes())?;
84 }
85 cursor.write_all(&[0x00])?;
87 cursor.write_all(&record_type.to_be_bytes())?;
88 cursor.write_all(&CLASS_IN.to_be_bytes())?;
89
90 let len = cursor.position() as usize;
91 socket.connect(endpoint)?;
92 socket.send(&buf[..len])?;
93
94 Ok(Self {
95 socket,
96 id,
97 buf,
98 record_type,
99 })
100 }
101
102 fn read_response<const N: usize>(mut self) -> Result<[u8; N], Error> {
103 let len = self.socket.recv(&mut self.buf)?;
104 ensure!(self.buf[..2] == self.id, "question/answer IDs don't match");
105 let response = &self.buf[..len];
106 let mut buf = Cursor::new(response);
107 let _id = buf.read_u16()?;
108
109 let flags = buf.read_u16()?;
110 ensure!(flags & 0x8000 != 0, "not a response");
111 ensure!(flags & 0x000f == 0, "non-zero DNS error code");
112
113 let qd = buf.read_u16()?;
114 ensure!(qd <= 1, "unexpected number of questions");
115 ensure!(buf.read_u16()? == 1, "unexpected number of answers");
116 ensure!(buf.read_u16()? == 0, "unexpected NS value");
117 ensure!(buf.read_u16()? == 0, "unexpected AR value"); if qd != 0 {
121 loop {
122 let len = buf.read_u8()?;
123 if len == 0 {
124 break;
125 }
126 buf.set_position(buf.position() + len as u64);
127 }
128 buf.set_position(buf.position() + 4);
130 }
131
132 let qname_len = buf.read_u16()?;
133 if qname_len & 0xc000 != 0xc000 {
135 buf.set_position(buf.position() + qname_len as u64);
136 }
137 ensure!(
138 buf.read_u16()? == self.record_type,
139 "answer is not expected type"
140 );
141 ensure!(buf.read_u16()? == CLASS_IN, "answer is not IN class");
142 buf.set_position(buf.position() + 4); let mut output = [0u8; N];
145 let data_len = buf.read_u16()? as usize;
146 let start = buf.position() as usize;
147 ensure!(data_len == N, "unexpected record data length");
148 output.copy_from_slice(&response[start..(start + data_len)]);
149 Ok(output)
150 }
151}
152
153fn get_id() -> Result<[u8; 2], Error> {
155 let mut id = [0u8; 2];
156 File::open("/dev/urandom")?.read_exact(&mut id)?;
157 Ok(id)
158}
159
160trait ReadExt {
161 fn read_u16(&mut self) -> Result<u16, std::io::Error>;
162 fn read_u8(&mut self) -> Result<u8, std::io::Error>;
163}
164
165impl ReadExt for Cursor<&[u8]> {
166 fn read_u16(&mut self) -> Result<u16, std::io::Error> {
167 let mut u16_buf = [0; 2];
168 self.read_exact(&mut u16_buf)?;
169 Ok(u16::from_be_bytes(u16_buf))
170 }
171
172 fn read_u8(&mut self) -> Result<u8, std::io::Error> {
173 let mut u8_buf = [0];
174 self.read_exact(&mut u8_buf)?;
175 Ok(u8_buf[0])
176 }
177}
178
179#[cfg(test)]
180mod tests {
181 use std::time::Instant;
182
183 use crate::*;
184
185 #[test]
186 #[ignore]
187 fn it_works() -> Result<(), Error> {
188 let now = Instant::now();
189 let (v4, v6) = get_both();
190 println!("Done in {}ms", now.elapsed().as_millis());
191 println!("v4: {v4:?}, v6: {v6:?}");
192 assert!(v4.is_some() || v6.is_some());
193 Ok(())
194 }
195}