1use crate::config::{Config, Strategy};
29use crate::error::Error;
30use crate::provider::BoxedBlockingProvider;
31use crate::types::{IpVersion, Protocol, ProviderResult};
32use std::collections::HashMap;
33use std::net::IpAddr;
34use std::sync::mpsc::{Receiver, RecvTimeoutError, Sender};
35use std::time::{Duration, Instant};
36
37type WorkerMessage = (
38 usize,
39 String,
40 Protocol,
41 Result<IpAddr, crate::ProviderError>,
42 Duration,
43);
44
45fn spawn_provider_worker(
46 id: usize,
47 provider: BoxedBlockingProvider,
48 version: IpVersion,
49 timeout: Duration,
50 tx: Sender<WorkerMessage>,
51) -> Result<(), crate::ProviderError> {
52 let name = provider.name().to_string();
53 let spawn_error_name = name.clone();
54 let protocol = provider.protocol();
55 let thread_name = format!("ip-discovery-{id}");
56 std::thread::Builder::new()
57 .name(thread_name)
58 .spawn(move || {
59 let start = Instant::now();
60 let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
61 provider.get_ip(version, timeout)
62 }))
63 .unwrap_or_else(|_| Err(crate::ProviderError::message(&name, "provider panicked")));
64 let latency = start.elapsed();
65 let _ = tx.send((id, name, protocol, result, latency));
66 })
67 .map(|_| ())
68 .map_err(|error| crate::ProviderError::new(spawn_error_name, error))
69}
70
71fn timeout_error(name: impl Into<String>) -> crate::ProviderError {
72 crate::ProviderError::message(name, "timeout")
73}
74
75fn validate_ip(
76 name: impl Into<String>,
77 version: IpVersion,
78 ip: IpAddr,
79) -> Result<IpAddr, crate::ProviderError> {
80 if version.matches(ip) {
81 Ok(ip)
82 } else {
83 Err(crate::ProviderError::message(
84 name,
85 "provider returned unexpected IP version",
86 ))
87 }
88}
89
90fn receive_with_timeout(
91 rx: &Receiver<WorkerMessage>,
92 timeout: Duration,
93) -> Result<WorkerMessage, RecvTimeoutError> {
94 rx.recv_timeout(timeout)
95}
96
97pub struct Resolver {
102 config: Config,
103}
104
105impl Resolver {
106 pub fn new(config: Config) -> Self {
108 Self { config }
109 }
110
111 #[inline]
113 fn matching_providers(&self) -> impl Iterator<Item = &BoxedBlockingProvider> {
114 self.config
115 .blocking_providers
116 .iter()
117 .filter(|p| p.supports_version(self.config.version))
118 }
119
120 pub fn resolve(&self) -> Result<ProviderResult, Error> {
128 if self.matching_providers().next().is_none() {
129 return Err(Error::NoProvidersForVersion);
130 }
131
132 match self.config.strategy {
133 Strategy::First => self.resolve_first(),
134 Strategy::Race => self.resolve_race(),
135 Strategy::Consensus { min_agree } => self.resolve_consensus(min_agree),
136 }
137 }
138
139 fn resolve_first(&self) -> Result<ProviderResult, Error> {
141 let mut errors = Vec::new();
142
143 for (id, provider) in self.matching_providers().enumerate() {
144 let provider_name = provider.name().to_string();
145 if self.config.timeout.is_zero() {
146 errors.push(timeout_error(provider_name));
147 continue;
148 }
149
150 let (tx, rx) = std::sync::mpsc::channel();
151 if let Err(error) = spawn_provider_worker(
152 id,
153 provider.clone(),
154 self.config.version,
155 self.config.timeout,
156 tx,
157 ) {
158 errors.push(error);
159 continue;
160 }
161
162 match receive_with_timeout(&rx, self.config.timeout) {
163 Ok((_, name, protocol, Ok(ip), latency)) => {
164 match validate_ip(&name, self.config.version, ip) {
165 Ok(ip) => {
166 return Ok(ProviderResult {
167 ip,
168 provider: name,
169 protocol,
170 latency,
171 });
172 }
173 Err(error) => errors.push(error),
174 }
175 }
176 Ok((_, _, _, Err(error), _)) => errors.push(error),
177 Err(RecvTimeoutError::Timeout) => errors.push(timeout_error(provider_name)),
178 Err(RecvTimeoutError::Disconnected) => errors.push(crate::ProviderError::message(
179 provider_name,
180 "provider worker disconnected",
181 )),
182 }
183 }
184
185 Err(Error::AllProvidersFailed(errors))
186 }
187
188 fn resolve_race(&self) -> Result<ProviderResult, Error> {
194 let providers: Vec<BoxedBlockingProvider> = self.matching_providers().cloned().collect();
195 if providers.is_empty() {
196 return Err(Error::NoProvidersForVersion);
197 }
198
199 if self.config.timeout.is_zero() {
200 return Err(Error::AllProvidersFailed(
201 providers
202 .iter()
203 .map(|provider| timeout_error(provider.name()))
204 .collect(),
205 ));
206 }
207
208 let (tx, rx) = std::sync::mpsc::channel();
209 let mut pending = HashMap::new();
210 let mut errors = Vec::new();
211
212 for (id, provider) in providers.into_iter().enumerate() {
213 pending.insert(id, provider.name().to_string());
214 if let Err(error) = spawn_provider_worker(
215 id,
216 provider,
217 self.config.version,
218 self.config.timeout,
219 tx.clone(),
220 ) {
221 pending.remove(&id);
222 errors.push(error);
223 }
224 }
225 drop(tx);
226
227 let deadline = Instant::now().checked_add(self.config.timeout);
228
229 while !pending.is_empty() {
230 let remaining = deadline
231 .and_then(|deadline| deadline.checked_duration_since(Instant::now()))
232 .unwrap_or(Duration::ZERO);
233 if remaining.is_zero() {
234 break;
235 }
236 match receive_with_timeout(&rx, remaining) {
237 Ok((id, name, protocol, res, latency)) => {
238 pending.remove(&id);
239 match res {
240 Ok(ip) => match validate_ip(&name, self.config.version, ip) {
241 Ok(ip) => {
242 return Ok(ProviderResult {
243 ip,
244 provider: name,
245 protocol,
246 latency,
247 });
248 }
249 Err(error) => errors.push(error),
250 },
251 Err(e) => errors.push(e),
252 }
253 }
254 Err(RecvTimeoutError::Timeout | RecvTimeoutError::Disconnected) => break,
255 }
256 }
257
258 errors.extend(pending.into_values().map(timeout_error));
259
260 Err(Error::AllProvidersFailed(errors))
261 }
262
263 fn resolve_consensus(&self, min_agree: usize) -> Result<ProviderResult, Error> {
265 let providers: Vec<BoxedBlockingProvider> = self.matching_providers().cloned().collect();
266 if providers.is_empty() {
267 return Err(Error::NoProvidersForVersion);
268 }
269
270 if self.config.timeout.is_zero() {
271 return Err(Error::ConsensusNotReached {
272 required: min_agree,
273 got: 0,
274 errors: providers
275 .iter()
276 .map(|provider| timeout_error(provider.name()))
277 .collect(),
278 });
279 }
280
281 let (tx, rx) = std::sync::mpsc::channel();
282 let mut pending = HashMap::new();
283 let mut errors = Vec::new();
284
285 for (id, provider) in providers.into_iter().enumerate() {
286 pending.insert(id, provider.name().to_string());
287 if let Err(error) = spawn_provider_worker(
288 id,
289 provider,
290 self.config.version,
291 self.config.timeout,
292 tx.clone(),
293 ) {
294 pending.remove(&id);
295 errors.push(error);
296 }
297 }
298 drop(tx);
299
300 let mut ip_results: HashMap<IpAddr, Vec<ProviderResult>> = HashMap::new();
301 let deadline = Instant::now().checked_add(self.config.timeout);
302
303 while !pending.is_empty() {
304 let remaining = deadline
305 .and_then(|deadline| deadline.checked_duration_since(Instant::now()))
306 .unwrap_or(Duration::ZERO);
307 if remaining.is_zero() {
308 break;
309 }
310 match receive_with_timeout(&rx, remaining) {
311 Ok((id, name, protocol, res, latency)) => {
312 pending.remove(&id);
313 match res {
314 Ok(ip) => match validate_ip(&name, self.config.version, ip) {
315 Ok(ip) => {
316 let pr = ProviderResult {
317 ip,
318 provider: name,
319 protocol,
320 latency,
321 };
322 ip_results.entry(ip).or_default().push(pr);
323 }
324 Err(error) => errors.push(error),
325 },
326 Err(e) => errors.push(e),
327 }
328 }
329 Err(RecvTimeoutError::Timeout | RecvTimeoutError::Disconnected) => break,
330 }
331 }
332
333 errors.extend(pending.into_values().map(timeout_error));
334
335 let mut best: Option<(IpAddr, usize)> = None;
336 for (ip, providers) in &ip_results {
337 if providers.len() >= min_agree {
338 match &best {
339 None => best = Some((*ip, providers.len())),
340 Some((_, current_len)) if providers.len() > *current_len => {
341 best = Some((*ip, providers.len()))
342 }
343 _ => {}
344 }
345 }
346 }
347
348 match best {
349 Some((ip, _)) => {
350 if let Some(providers) = ip_results.remove(&ip) {
351 if let Some(fastest) = providers.into_iter().min_by_key(|p| p.latency) {
352 return Ok(fastest);
353 }
354 }
355 Err(Error::ConsensusNotReached {
356 required: min_agree,
357 got: 0,
358 errors,
359 })
360 }
361 None => {
362 let max_agreement = ip_results.values().map(|v| v.len()).max().unwrap_or(0);
363 Err(Error::ConsensusNotReached {
364 required: min_agree,
365 got: max_agreement,
366 errors,
367 })
368 }
369 }
370 }
371}
372
373pub fn get_ip() -> Result<ProviderResult, Error> {
379 let config = Config::default();
380 get_ip_with(config)
381}
382
383pub fn get_ipv4() -> Result<ProviderResult, Error> {
389 let config = Config::builder().version(IpVersion::V4).build();
390 get_ip_with(config)
391}
392
393pub fn get_ipv6() -> Result<ProviderResult, Error> {
399 let config = Config::builder().version(IpVersion::V6).build();
400 get_ip_with(config)
401}
402
403pub fn get_ip_with(config: Config) -> Result<ProviderResult, Error> {
409 let resolver = Resolver::new(config);
410 resolver.resolve()
411}