1pub(crate) mod providers;
6
7#[cfg(feature = "tokio")]
8pub use providers::default_providers;
9pub use providers::{default_blocking_providers, provider_names};
10
11use crate::error::ProviderError;
12#[cfg(feature = "tokio")]
13use crate::provider::Provider;
14use crate::provider::{BlockingProvider, BoxedBlockingProvider};
15use crate::types::{IpVersion, Protocol};
16use std::net::IpAddr;
17use std::str::FromStr;
18use std::time::Duration;
19
20#[cfg(feature = "tokio")]
21use std::future::Future;
22#[cfg(feature = "tokio")]
23use std::pin::Pin;
24
25pub type ResponseParser = fn(&str) -> Option<IpAddr>;
27
28pub fn parse_plain_text(text: &str) -> Option<IpAddr> {
30 IpAddr::from_str(text.trim()).ok()
31}
32
33pub fn parse_cloudflare_trace(text: &str) -> Option<IpAddr> {
35 for line in text.lines() {
36 if let Some(ip_str) = line.strip_prefix("ip=") {
37 return IpAddr::from_str(ip_str.trim()).ok();
38 }
39 }
40 None
41}
42
43#[derive(Clone)]
45pub struct HttpProvider {
46 name: String,
47 url_v4: Option<String>,
48 url_v6: Option<String>,
49 parser: ResponseParser,
50 #[cfg(feature = "tokio")]
51 client: reqwest::Client,
52}
53
54impl HttpProvider {
55 pub fn new(name: impl Into<String>, url: impl Into<String>) -> Self {
57 #[cfg(feature = "tokio")]
58 let client = reqwest::Client::builder()
59 .user_agent(concat!("ip-discovery/", env!("CARGO_PKG_VERSION")))
60 .build()
61 .unwrap_or_default();
62
63 Self {
64 name: name.into(),
65 url_v4: Some(url.into()),
66 url_v6: None,
67 parser: parse_plain_text,
68 #[cfg(feature = "tokio")]
69 client,
70 }
71 }
72
73 pub fn with_parser(mut self, parser: ResponseParser) -> Self {
75 self.parser = parser;
76 self
77 }
78
79 pub fn with_v6_url(mut self, url: impl Into<String>) -> Self {
81 self.url_v6 = Some(url.into());
82 self
83 }
84
85 fn get_url(&self, version: IpVersion) -> Option<&str> {
87 match version {
88 IpVersion::V6 => self.url_v6.as_deref().or(self.url_v4.as_deref()),
89 _ => self.url_v4.as_deref(),
90 }
91 }
92
93 pub fn fetch_blocking(
95 &self,
96 version: IpVersion,
97 timeout: Duration,
98 ) -> Result<IpAddr, ProviderError> {
99 let url = self
100 .get_url(version)
101 .ok_or_else(|| ProviderError::message(&self.name, "no URL for IP version"))?;
102
103 let client = reqwest::blocking::Client::builder()
104 .timeout(timeout)
105 .user_agent(concat!("ip-discovery/", env!("CARGO_PKG_VERSION")))
106 .build()
107 .unwrap_or_default();
108
109 let response = client
110 .get(url)
111 .send()
112 .map_err(|e| ProviderError::new(&self.name, e))?;
113
114 if !response.status().is_success() {
115 return Err(ProviderError::message(
116 &self.name,
117 format!("HTTP error: {}", response.status()),
118 ));
119 }
120
121 let text = response
122 .text()
123 .map_err(|e| ProviderError::new(&self.name, e))?;
124
125 let ip = (self.parser)(&text)
126 .ok_or_else(|| ProviderError::message(&self.name, "failed to parse response"))?;
127 if version.matches(ip) {
128 Ok(ip)
129 } else {
130 Err(ProviderError::message(
131 &self.name,
132 "provider returned unexpected IP version",
133 ))
134 }
135 }
136
137 #[cfg(feature = "tokio")]
139 async fn fetch(&self, version: IpVersion) -> Result<IpAddr, ProviderError> {
140 let url = self
141 .get_url(version)
142 .ok_or_else(|| ProviderError::message(&self.name, "no URL for IP version"))?;
143
144 let response = self
145 .client
146 .get(url)
147 .send()
148 .await
149 .map_err(|e| ProviderError::new(&self.name, e))?;
150
151 if !response.status().is_success() {
152 return Err(ProviderError::message(
153 &self.name,
154 format!("HTTP error: {}", response.status()),
155 ));
156 }
157
158 let text = response
159 .text()
160 .await
161 .map_err(|e| ProviderError::new(&self.name, e))?;
162
163 let ip = (self.parser)(&text)
164 .ok_or_else(|| ProviderError::message(&self.name, "failed to parse response"))?;
165 if version.matches(ip) {
166 Ok(ip)
167 } else {
168 Err(ProviderError::message(
169 &self.name,
170 "provider returned unexpected IP version",
171 ))
172 }
173 }
174}
175
176impl std::fmt::Debug for HttpProvider {
177 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
178 f.debug_struct("HttpProvider")
179 .field("name", &self.name)
180 .field("url_v4", &self.url_v4)
181 .field("url_v6", &self.url_v6)
182 .finish()
183 }
184}
185
186impl BlockingProvider for HttpProvider {
187 fn name(&self) -> &str {
188 &self.name
189 }
190
191 fn protocol(&self) -> Protocol {
192 Protocol::Http
193 }
194
195 fn supports_v4(&self) -> bool {
196 self.url_v4.is_some()
197 }
198
199 fn supports_v6(&self) -> bool {
200 self.url_v6.is_some()
201 }
202
203 fn get_ip(&self, version: IpVersion, timeout: Duration) -> Result<IpAddr, ProviderError> {
204 self.fetch_blocking(version, timeout)
205 }
206
207 fn clone_box(&self) -> BoxedBlockingProvider {
208 Box::new(self.clone())
209 }
210}
211
212#[cfg(feature = "tokio")]
213impl Provider for HttpProvider {
214 fn name(&self) -> &str {
215 &self.name
216 }
217
218 fn protocol(&self) -> Protocol {
219 Protocol::Http
220 }
221
222 fn supports_v4(&self) -> bool {
223 self.url_v4.is_some()
224 }
225
226 fn supports_v6(&self) -> bool {
227 self.url_v6.is_some()
228 }
229
230 fn get_ip(
231 &self,
232 version: IpVersion,
233 ) -> Pin<Box<dyn Future<Output = Result<IpAddr, ProviderError>> + Send + '_>> {
234 Box::pin(self.fetch(version))
235 }
236}