Skip to main content

ip_discovery/stun/
mod.rs

1//! STUN protocol implementation for public IP detection
2//!
3//! Implements a minimal RFC 5389 STUN client for detecting public IP addresses.
4//!
5//! # Limitations
6//!
7//! - **No retransmission**: RFC 5389 §7.2.1 recommends retransmitting requests
8//!   with exponential backoff (RTO ≥ 500 ms). This implementation sends a single
9//!   binding request; packet loss is handled by the resolver's fallback strategy.
10//!
11//! # Security
12//!
13//! Transaction IDs are generated with [`getrandom`] (OS-level CSPRNG),
14//! per RFC 8489 requirements.
15
16mod message;
17pub(crate) mod providers;
18
19#[cfg(feature = "tokio")]
20pub use providers::default_providers;
21pub use providers::{default_blocking_providers, provider_names};
22
23use crate::error::ProviderError;
24#[cfg(feature = "tokio")]
25use crate::provider::Provider;
26use crate::provider::{BlockingProvider, BoxedBlockingProvider};
27use crate::types::{IpVersion, Protocol};
28use message::{StunMessage, StunMethod};
29use std::net::{IpAddr, SocketAddr};
30use std::time::Duration;
31
32#[cfg(feature = "tokio")]
33use std::future::Future;
34#[cfg(feature = "tokio")]
35use std::pin::Pin;
36
37/// STUN provider for IP detection
38#[derive(Debug, Clone)]
39pub struct StunProvider {
40    name: String,
41    server: String,
42    port: u16,
43}
44
45impl StunProvider {
46    /// Create a new STUN provider
47    pub fn new(name: impl Into<String>, server: impl Into<String>, port: u16) -> Self {
48        Self {
49            name: name.into(),
50            server: server.into(),
51            port,
52        }
53    }
54
55    fn select_addr(addrs: &[SocketAddr], version: IpVersion) -> Option<SocketAddr> {
56        match version {
57            IpVersion::V4 => addrs.iter().copied().find(SocketAddr::is_ipv4),
58            IpVersion::V6 => addrs.iter().copied().find(SocketAddr::is_ipv6),
59            IpVersion::Any => addrs
60                .iter()
61                .copied()
62                .find(SocketAddr::is_ipv4)
63                .or_else(|| addrs.first().copied()),
64        }
65    }
66
67    fn process_response(
68        request: &StunMessage,
69        buf: &[u8],
70        name: &str,
71    ) -> Result<IpAddr, ProviderError> {
72        let response = StunMessage::decode(buf).map_err(|e| ProviderError::message(name, e))?;
73
74        if response.transaction_id() != request.transaction_id() {
75            return Err(ProviderError::message(name, "transaction ID mismatch"));
76        }
77
78        if response.method() != StunMethod::Response {
79            return Err(ProviderError::message(
80                name,
81                "not a binding success response",
82            ));
83        }
84
85        response
86            .get_mapped_address()
87            .ok_or_else(|| ProviderError::message(name, "no mapped address in response"))
88    }
89
90    /// Perform STUN binding request synchronously using standard UDP sockets with a timeout
91    pub fn binding_request_blocking(
92        &self,
93        version: IpVersion,
94        timeout: Duration,
95    ) -> Result<IpAddr, ProviderError> {
96        use std::net::ToSocketAddrs;
97
98        if timeout.is_zero() {
99            return Err(ProviderError::message(&self.name, "timeout"));
100        }
101
102        let server_addr = format!("{}:{}", self.server, self.port);
103        let addrs: Vec<SocketAddr> = server_addr
104            .to_socket_addrs()
105            .map_err(|e| ProviderError::new(&self.name, e))?
106            .collect();
107
108        let addr = Self::select_addr(&addrs, version).ok_or_else(|| {
109            ProviderError::message(&self.name, "no suitable address for IP version")
110        })?;
111
112        let local_addr = if addr.is_ipv4() {
113            SocketAddr::from(([0, 0, 0, 0], 0))
114        } else {
115            SocketAddr::from(([0u16; 8], 0))
116        };
117
118        let socket =
119            std::net::UdpSocket::bind(local_addr).map_err(|e| ProviderError::new(&self.name, e))?;
120
121        socket
122            .set_read_timeout(Some(timeout))
123            .map_err(|e| ProviderError::new(&self.name, e))?;
124        socket
125            .set_write_timeout(Some(timeout))
126            .map_err(|e| ProviderError::new(&self.name, e))?;
127
128        socket
129            .connect(addr)
130            .map_err(|e| ProviderError::new(&self.name, e))?;
131
132        let request =
133            StunMessage::new(StunMethod::Request).map_err(|e| ProviderError::new(&self.name, e))?;
134        let request_bytes = request.encode();
135
136        socket
137            .send(&request_bytes)
138            .map_err(|e| ProviderError::new(&self.name, e))?;
139
140        let mut buf = [0u8; 576]; // Minimum MTU
141        let len = socket
142            .recv(&mut buf)
143            .map_err(|e| ProviderError::new(&self.name, e))?;
144
145        Self::process_response(&request, &buf[..len], &self.name)
146    }
147
148    /// Perform STUN binding request asynchronously
149    #[cfg(feature = "tokio")]
150    async fn binding_request(&self, version: IpVersion) -> Result<IpAddr, ProviderError> {
151        let server_addr = format!("{}:{}", self.server, self.port);
152        let addrs: Vec<SocketAddr> = tokio::net::lookup_host(&server_addr)
153            .await
154            .map_err(|e| ProviderError::new(&self.name, e))?
155            .collect();
156
157        let addr = Self::select_addr(&addrs, version).ok_or_else(|| {
158            ProviderError::message(&self.name, "no suitable address for IP version")
159        })?;
160
161        let local_addr = if addr.is_ipv4() {
162            SocketAddr::from(([0, 0, 0, 0], 0))
163        } else {
164            SocketAddr::from(([0u16; 8], 0))
165        };
166
167        let socket = tokio::net::UdpSocket::bind(local_addr)
168            .await
169            .map_err(|e| ProviderError::new(&self.name, e))?;
170
171        socket
172            .connect(addr)
173            .await
174            .map_err(|e| ProviderError::new(&self.name, e))?;
175
176        let request =
177            StunMessage::new(StunMethod::Request).map_err(|e| ProviderError::new(&self.name, e))?;
178        let request_bytes = request.encode();
179
180        socket
181            .send(&request_bytes)
182            .await
183            .map_err(|e| ProviderError::new(&self.name, e))?;
184
185        let mut buf = [0u8; 576]; // Minimum MTU
186        let len = socket
187            .recv(&mut buf)
188            .await
189            .map_err(|e| ProviderError::new(&self.name, e))?;
190
191        Self::process_response(&request, &buf[..len], &self.name)
192    }
193}
194
195impl BlockingProvider for StunProvider {
196    fn name(&self) -> &str {
197        &self.name
198    }
199
200    fn protocol(&self) -> Protocol {
201        Protocol::Stun
202    }
203
204    fn supports_v4(&self) -> bool {
205        true
206    }
207
208    fn supports_v6(&self) -> bool {
209        true
210    }
211
212    fn get_ip(&self, version: IpVersion, timeout: Duration) -> Result<IpAddr, ProviderError> {
213        self.binding_request_blocking(version, timeout)
214    }
215
216    fn clone_box(&self) -> BoxedBlockingProvider {
217        Box::new(self.clone())
218    }
219}
220
221#[cfg(feature = "tokio")]
222impl Provider for StunProvider {
223    fn name(&self) -> &str {
224        &self.name
225    }
226
227    fn protocol(&self) -> Protocol {
228        Protocol::Stun
229    }
230
231    fn supports_v4(&self) -> bool {
232        true
233    }
234
235    fn supports_v6(&self) -> bool {
236        true
237    }
238
239    fn get_ip(
240        &self,
241        version: IpVersion,
242    ) -> Pin<Box<dyn Future<Output = Result<IpAddr, ProviderError>> + Send + '_>> {
243        Box::pin(self.binding_request(version))
244    }
245}
246
247#[cfg(test)]
248mod blocking_tests {
249    use super::{StunMessage, StunMethod, StunProvider};
250    use crate::types::IpVersion;
251    use std::net::{IpAddr, Ipv4Addr, SocketAddr, UdpSocket};
252    use std::time::{Duration, Instant};
253
254    #[test]
255    fn any_prefers_ipv4_address() {
256        let v6: SocketAddr = "[::1]:3478".parse().unwrap();
257        let v4: SocketAddr = "127.0.0.1:3478".parse().unwrap();
258
259        assert_eq!(
260            StunProvider::select_addr(&[v6, v4], IpVersion::Any),
261            Some(v4)
262        );
263    }
264
265    #[test]
266    fn zero_timeout_fails_before_waiting_for_stun_response() {
267        let server = UdpSocket::bind("127.0.0.1:0").unwrap();
268        let port = server.local_addr().unwrap().port();
269        server
270            .set_read_timeout(Some(Duration::from_millis(200)))
271            .unwrap();
272        std::thread::spawn(move || {
273            let mut request = [0u8; 512];
274            if let Ok((_, peer)) = server.recv_from(&mut request) {
275                std::thread::sleep(Duration::from_millis(100));
276                let _ = server.send_to(&[0u8; 20], peer);
277            }
278        });
279
280        let provider = StunProvider::new("local-stun", "127.0.0.1", port);
281        let started = Instant::now();
282        let result = provider.binding_request_blocking(IpVersion::V4, Duration::ZERO);
283
284        assert!(result.is_err());
285        assert!(started.elapsed() < Duration::from_millis(50));
286    }
287
288    #[test]
289    fn version_filter_selects_requested_family() {
290        let v4 = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 3478);
291        let v6: SocketAddr = "[::1]:3478".parse().unwrap();
292
293        assert_eq!(
294            StunProvider::select_addr(&[v6, v4], IpVersion::V4),
295            Some(v4)
296        );
297        assert_eq!(
298            StunProvider::select_addr(&[v4, v6], IpVersion::V6),
299            Some(v6)
300        );
301    }
302
303    #[test]
304    fn rejects_non_success_message_with_mapped_address() {
305        let request = StunMessage::new(StunMethod::Request).unwrap();
306        let mut response = request.encode();
307        response[2..4].copy_from_slice(&12u16.to_be_bytes());
308        response.extend_from_slice(&0x0001u16.to_be_bytes());
309        response.extend_from_slice(&8u16.to_be_bytes());
310        response.extend_from_slice(&[0, 1, 0, 0, 203, 0, 113, 1]);
311
312        let error = StunProvider::process_response(&request, &response, "test").unwrap_err();
313        assert!(error.to_string().contains("not a binding success response"));
314
315        response[0..2].copy_from_slice(&0x0111u16.to_be_bytes());
316        let error = StunProvider::process_response(&request, &response, "test").unwrap_err();
317        assert!(error.to_string().contains("not a binding success response"));
318    }
319}