1use std::collections::HashSet;
4use std::error::Error as _;
5use std::time::{Duration, Instant};
6
7use bytes::Bytes;
8use codoseo_core::crawl::{AddressPolicy, CrawlLimits, USER_AGENT};
9use reqwest::header::{CONTENT_LENGTH, CONTENT_TYPE, HeaderMap, LOCATION};
10use reqwest::redirect::Policy;
11use url::Url;
12
13use crate::guard::{GuardError, GuardedResolver, SystemLookup, check_url};
14
15#[derive(Debug, Clone)]
16pub struct FetcherConfig {
17 pub user_agent: String,
18 pub address_policy: AddressPolicy,
19 pub request_timeout: Duration,
20 pub connect_timeout: Duration,
21 pub max_redirects: u8,
22 pub max_body_bytes: usize,
23}
24
25impl FetcherConfig {
26 pub fn new(address_policy: AddressPolicy) -> Self {
27 let limits = CrawlLimits::default();
28 FetcherConfig {
29 user_agent: USER_AGENT.to_owned(),
30 address_policy,
31 request_timeout: limits.request_timeout,
32 connect_timeout: Duration::from_secs(10),
33 max_redirects: limits.max_redirects,
34 max_body_bytes: limits.max_page_bytes,
35 }
36 }
37}
38
39#[derive(Debug, Clone, PartialEq, Eq)]
41pub struct Hop {
42 pub status: u16,
43 pub url: Url,
44}
45
46#[derive(Debug, Clone)]
47pub struct FetchResult {
48 pub final_url: Url,
49 pub status: u16,
50 pub chain: Vec<Hop>,
51 pub headers: HeaderMap,
52 pub content_type: Option<String>,
53 pub x_robots_tag: Option<String>,
54 pub response_ms: u32,
56 pub size_bytes: u64,
57 pub body: Option<Bytes>,
59 pub truncated: bool,
60}
61
62#[derive(Debug, thiserror::Error)]
63pub enum FetchError {
64 #[error("blocked: {0}")]
65 Blocked(String),
66 #[error("request timed out")]
67 Timeout,
68 #[error("connection failed: {0}")]
69 Connect(String),
70 #[error("more than the allowed number of redirects")]
71 TooManyRedirects { chain: Vec<Hop> },
72 #[error("redirect loop")]
73 RedirectLoop { chain: Vec<Hop> },
74 #[error("invalid redirect target: {location}")]
75 InvalidRedirect { location: String },
76 #[error("request failed: {0}")]
77 Http(String),
78 #[error("could not build HTTP client: {0}")]
79 Client(String),
80}
81
82impl From<GuardError> for FetchError {
83 fn from(e: GuardError) -> Self {
84 FetchError::Blocked(e.to_string())
85 }
86}
87
88pub struct Fetcher {
89 client: reqwest::Client,
90 cfg: FetcherConfig,
91}
92
93enum BodyMode {
94 HtmlOnly,
95 Any(usize),
96}
97
98impl Fetcher {
99 pub fn new(cfg: FetcherConfig) -> Result<Self, FetchError> {
100 let mut builder = reqwest::Client::builder()
101 .redirect(Policy::none())
102 .user_agent(cfg.user_agent.clone())
103 .timeout(cfg.request_timeout)
104 .connect_timeout(cfg.connect_timeout)
105 .no_proxy();
106 if cfg.address_policy == AddressPolicy::Public {
107 builder = builder.dns_resolver(GuardedResolver::new(SystemLookup));
108 }
109 let client = builder
110 .build()
111 .map_err(|e| FetchError::Client(e.to_string()))?;
112 Ok(Fetcher { client, cfg })
113 }
114
115 pub async fn fetch(&self, url: &Url) -> Result<FetchResult, FetchError> {
117 self.run(url, BodyMode::HtmlOnly).await
118 }
119
120 pub async fn fetch_raw(&self, url: &Url, max_bytes: usize) -> Result<FetchResult, FetchError> {
122 self.run(url, BodyMode::Any(max_bytes)).await
123 }
124
125 async fn run(&self, start: &Url, mode: BodyMode) -> Result<FetchResult, FetchError> {
126 let mut current = start.clone();
127 let mut chain: Vec<Hop> = Vec::new();
128 let mut seen: HashSet<String> = HashSet::new();
129 loop {
130 check_url(¤t, self.cfg.address_policy)?;
131 seen.insert(current.as_str().to_owned());
132
133 let sent = Instant::now();
134 let resp = self
135 .client
136 .get(current.clone())
137 .send()
138 .await
139 .map_err(map_err)?;
140 let response_ms = sent.elapsed().as_millis().min(u32::MAX as u128) as u32;
141 let status = resp.status().as_u16();
142
143 if resp.status().is_redirection()
144 && let Some(location) = resp.headers().get(LOCATION)
145 {
146 let location = String::from_utf8_lossy(location.as_bytes()).into_owned();
147 let next = current
148 .join(&location)
149 .ok()
150 .filter(|u| matches!(u.scheme(), "http" | "https"))
151 .ok_or(FetchError::InvalidRedirect { location })?;
152 chain.push(Hop {
153 status,
154 url: current,
155 });
156 if seen.contains(next.as_str()) {
157 return Err(FetchError::RedirectLoop { chain });
158 }
159 if chain.len() > usize::from(self.cfg.max_redirects) {
160 return Err(FetchError::TooManyRedirects { chain });
161 }
162 current = next;
163 continue;
164 }
165
166 return self.finish(current, chain, resp, response_ms, mode).await;
167 }
168 }
169
170 async fn finish(
171 &self,
172 final_url: Url,
173 chain: Vec<Hop>,
174 mut resp: reqwest::Response,
175 response_ms: u32,
176 mode: BodyMode,
177 ) -> Result<FetchResult, FetchError> {
178 let headers = resp.headers().clone();
179 let content_type = header_str(&headers, CONTENT_TYPE.as_str());
180 let x_robots_tag = joined(&headers, "x-robots-tag");
181 let declared_len = headers
182 .get(CONTENT_LENGTH)
183 .and_then(|v| v.to_str().ok())
184 .and_then(|v| v.parse::<u64>().ok());
185
186 let cap = match mode {
187 BodyMode::Any(max) => Some(max),
188 BodyMode::HtmlOnly if is_html(content_type.as_deref()) => Some(self.cfg.max_body_bytes),
189 BodyMode::HtmlOnly => None,
190 };
191
192 let (body, truncated, read) = match cap {
193 None => (None, false, 0),
194 Some(cap) => {
195 let initial = declared_len.map_or(64 * 1024, |len| len as usize).min(cap);
198 let mut buf: Vec<u8> = Vec::with_capacity(initial);
199 let mut truncated = false;
200 while let Some(chunk) = resp.chunk().await.map_err(map_err)? {
201 let take = chunk.len().min(cap - buf.len());
202 let needed = buf.len() + take;
203 if needed > buf.capacity() {
204 let target = (buf.capacity() * 2).max(needed).min(cap);
205 buf.reserve_exact(target - buf.len());
206 }
207 buf.extend_from_slice(&chunk[..take]);
208 if take < chunk.len() {
209 truncated = true;
210 break;
211 }
212 }
213 let read = buf.len() as u64;
214 (Some(Bytes::from(buf)), truncated, read)
215 }
216 };
217
218 Ok(FetchResult {
219 final_url,
220 status: resp.status().as_u16(),
221 chain,
222 headers,
223 content_type,
224 x_robots_tag,
225 response_ms,
226 size_bytes: declared_len.unwrap_or(read),
227 body,
228 truncated,
229 })
230 }
231}
232
233fn is_html(content_type: Option<&str>) -> bool {
234 match content_type {
235 None => true,
236 Some(ct) => {
237 let mime = ct
238 .split(';')
239 .next()
240 .unwrap_or("")
241 .trim()
242 .to_ascii_lowercase();
243 mime == "text/html" || mime == "application/xhtml+xml"
244 }
245 }
246}
247
248fn header_str(headers: &HeaderMap, name: &str) -> Option<String> {
249 headers
250 .get(name)
251 .and_then(|v| v.to_str().ok())
252 .map(str::to_owned)
253}
254
255fn joined(headers: &HeaderMap, name: &str) -> Option<String> {
256 let values: Vec<&str> = headers
257 .get_all(name)
258 .iter()
259 .filter_map(|v| v.to_str().ok())
260 .collect();
261 (!values.is_empty()).then(|| values.join(", "))
262}
263
264fn map_err(e: reqwest::Error) -> FetchError {
265 let mut source = e.source();
266 while let Some(s) = source {
267 if let Some(guard) = s.downcast_ref::<GuardError>() {
268 return FetchError::Blocked(guard.to_string());
269 }
270 source = s.source();
271 }
272 if e.is_timeout() {
273 FetchError::Timeout
274 } else if e.is_connect() {
275 FetchError::Connect(error_chain(&e))
276 } else {
277 FetchError::Http(error_chain(&e))
278 }
279}
280
281fn error_chain(e: &reqwest::Error) -> String {
282 let mut msg = e.to_string();
283 let mut source = e.source();
284 while let Some(s) = source {
285 msg.push_str(": ");
286 msg.push_str(&s.to_string());
287 source = s.source();
288 }
289 msg
290}