1mod 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#[derive(Debug, Clone)]
39pub struct StunProvider {
40 name: String,
41 server: String,
42 port: u16,
43}
44
45impl StunProvider {
46 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 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]; 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 #[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]; 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}