Skip to main content

innernet_publicip/
lib.rs

1//! Get your public IP address(es) as fast as possible, with no dependencies.
2//!
3//! It uses Quad9's DNS server along with a feature specific to their resolver
4//! which returns your external IP address.
5
6use 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
25// Reference: https://www.quad9.net/service/service-addresses-and-features
26static 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())?; // Request type (query, in this case)
76        cursor.write_all(&0x0001u16.to_be_bytes())?; // Number of queries
77        cursor.write_all(&0x0000u16.to_be_bytes())?; // Number of responses
78        cursor.write_all(&0x0000u16.to_be_bytes())?; // Number of name server records
79        cursor.write_all(&0x0000u16.to_be_bytes())?; // Number of additional records
80        for atom in QNAME {
81            // Write the length of this atom followed by the string itself
82            cursor.write_all(&[atom.len() as u8])?;
83            cursor.write_all(atom.as_bytes())?;
84        }
85        // Finish the qname with a terminating byte (0-length atom).
86        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"); // "Additional Records"
118
119        // Skip past the query section, don't care.
120        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            // Skip type and class information as well.
129            buf.set_position(buf.position() + 4);
130        }
131
132        let qname_len = buf.read_u16()?;
133        // Ignore if it's a pointer, ignore if it's a normal QNAME...
134        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); // Ignore TTL
143
144        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
153/// DNS wants a random-ish ID to be generated per request.
154fn 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}