1use super::resolver::SsrfSafeResolver;
2use super::response::FeedHttpResponse;
3use super::validation::validate_url;
4use crate::error::{FeedError, Result};
5use reqwest::blocking::{Client, Response};
6use reqwest::header::{
7 ACCEPT, ACCEPT_ENCODING, HeaderMap, HeaderName, HeaderValue, IF_MODIFIED_SINCE, IF_NONE_MATCH,
8 USER_AGENT,
9};
10use std::collections::HashMap;
11use std::sync::Arc;
12use std::time::Duration;
13
14const MAX_REDIRECTS: usize = 10;
19
20const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
24
25pub struct FeedHttpClient {
27 client: Client,
28 user_agent: String,
29 timeout: Duration,
30}
31
32impl FeedHttpClient {
33 pub fn new() -> Result<Self> {
51 let client = Client::builder()
52 .timeout(DEFAULT_TIMEOUT)
53 .gzip(true)
54 .deflate(true)
55 .brotli(true)
56 .redirect(Self::redirect_policy())
57 .dns_resolver(Arc::new(SsrfSafeResolver))
58 .no_proxy()
59 .build()
60 .map_err(|e| FeedError::Http {
61 message: format!("Failed to create HTTP client: {e}"),
62 })?;
63
64 Ok(Self {
65 client,
66 user_agent: format!(
67 "feedparser-rs/{} (+https://github.com/bug-ops/feedparser-rs)",
68 env!("CARGO_PKG_VERSION")
69 ),
70 timeout: DEFAULT_TIMEOUT,
71 })
72 }
73
74 fn should_follow_redirect(next_url: &str, previous_hops: usize) -> Result<()> {
86 if previous_hops > MAX_REDIRECTS {
87 return Err(FeedError::Http {
88 message: format!("Too many redirects (max {MAX_REDIRECTS})"),
89 });
90 }
91
92 validate_url(next_url)?;
93 Ok(())
94 }
95
96 fn describe_request_error(error: &reqwest::Error) -> String {
105 use std::fmt::Write as _;
106
107 let mut message = error.to_string();
108 let mut source = std::error::Error::source(error);
109 while let Some(err) = source {
110 let _ = write!(message, ": {err}");
111 source = err.source();
112 }
113 message
114 }
115
116 fn redirect_policy() -> reqwest::redirect::Policy {
127 reqwest::redirect::Policy::custom(|attempt| {
128 match Self::should_follow_redirect(attempt.url().as_str(), attempt.previous().len()) {
129 Ok(()) => attempt.follow(),
130 Err(e) => attempt.error(e),
131 }
132 })
133 }
134
135 #[must_use]
141 pub fn with_user_agent(mut self, agent: String) -> Self {
142 const MAX_USER_AGENT_LEN: usize = 512;
144 self.user_agent = if agent.len() > MAX_USER_AGENT_LEN {
145 agent.chars().take(MAX_USER_AGENT_LEN).collect()
146 } else {
147 agent
148 };
149 self
150 }
151
152 #[must_use]
163 pub const fn with_timeout(mut self, timeout: Duration) -> Self {
164 const MAX_TIMEOUT_SECS: u64 = 3600;
165 self.timeout = if timeout.as_secs() > MAX_TIMEOUT_SECS {
166 Duration::from_secs(MAX_TIMEOUT_SECS)
167 } else {
168 timeout
169 };
170 self
171 }
172
173 #[inline]
177 fn insert_header(
178 headers: &mut HeaderMap,
179 name: HeaderName,
180 value: &str,
181 field_name: &str,
182 ) -> Result<()> {
183 headers.insert(
184 name,
185 HeaderValue::from_str(value).map_err(|e| FeedError::Http {
186 message: format!("Invalid {field_name}: {e}"),
187 })?,
188 );
189 Ok(())
190 }
191
192 pub fn get(
207 &self,
208 url: &str,
209 etag: Option<&str>,
210 modified: Option<&str>,
211 extra_headers: Option<&HeaderMap>,
212 ) -> Result<FeedHttpResponse> {
213 let validated_url = validate_url(url)?;
215 let url_str = validated_url.as_str();
216
217 let mut headers = HeaderMap::new();
218
219 Self::insert_header(&mut headers, USER_AGENT, &self.user_agent, "User-Agent")?;
221
222 headers.insert(
223 ACCEPT,
224 HeaderValue::from_static(
225 "application/rss+xml, application/atom+xml, application/xml, text/xml, */*",
226 ),
227 );
228
229 headers.insert(
230 ACCEPT_ENCODING,
231 HeaderValue::from_static("gzip, deflate, br"),
232 );
233
234 if let Some(etag_val) = etag {
236 const MAX_ETAG_LEN: usize = 1024;
238 let sanitized_etag = if etag_val.len() > MAX_ETAG_LEN {
239 &etag_val[..MAX_ETAG_LEN]
240 } else {
241 etag_val
242 };
243 Self::insert_header(&mut headers, IF_NONE_MATCH, sanitized_etag, "ETag")?;
244 }
245
246 if let Some(modified_val) = modified {
247 const MAX_MODIFIED_LEN: usize = 64;
249 let sanitized_modified = if modified_val.len() > MAX_MODIFIED_LEN {
250 &modified_val[..MAX_MODIFIED_LEN]
251 } else {
252 modified_val
253 };
254 Self::insert_header(
255 &mut headers,
256 IF_MODIFIED_SINCE,
257 sanitized_modified,
258 "Last-Modified",
259 )?;
260 }
261
262 if let Some(extra) = extra_headers {
264 headers.extend(extra.clone());
265 }
266
267 let request = self.build_request(url_str, headers)?;
268
269 let response = self.client.execute(request).map_err(|e| FeedError::Http {
270 message: format!("HTTP request failed: {}", Self::describe_request_error(&e)),
271 })?;
272
273 Self::build_response(response, url_str)
274 }
275
276 fn build_request(
287 &self,
288 url_str: &str,
289 headers: HeaderMap,
290 ) -> Result<reqwest::blocking::Request> {
291 self.client
292 .get(url_str)
293 .headers(headers)
294 .timeout(self.timeout)
295 .build()
296 .map_err(|e| FeedError::Http {
297 message: format!(
298 "Failed to build request: {}",
299 Self::describe_request_error(&e)
300 ),
301 })
302 }
303
304 fn build_response(response: Response, _original_url: &str) -> Result<FeedHttpResponse> {
306 let status = response.status().as_u16();
307 let url = response.url().to_string();
308
309 let mut headers_map = HashMap::with_capacity(response.headers().len());
311 for (name, value) in response.headers() {
312 if let Ok(val_str) = value.to_str() {
313 headers_map.insert(name.to_string(), val_str.to_string());
314 }
315 }
316
317 let etag = headers_map.get("etag").cloned();
319 let last_modified = headers_map.get("last-modified").cloned();
320 let content_type = headers_map.get("content-type").cloned();
321
322 let encoding = content_type
324 .as_ref()
325 .and_then(|ct| FeedHttpResponse::extract_charset_from_content_type(ct));
326
327 let body = if status == 304 {
329 Vec::new()
331 } else {
332 response
333 .bytes()
334 .map_err(|e| FeedError::Http {
335 message: format!("Failed to read response body: {e}"),
336 })?
337 .to_vec()
338 };
339
340 Ok(FeedHttpResponse {
341 status,
342 url,
343 headers: headers_map,
344 body,
345 etag,
346 last_modified,
347 content_type,
348 encoding,
349 })
350 }
351}
352
353#[cfg(test)]
354mod tests {
355 use super::*;
356
357 #[test]
358 fn test_client_creation() {
359 let client = FeedHttpClient::new();
360 assert!(client.is_ok());
361 }
362
363 #[test]
366 fn test_redirect_rejects_metadata_endpoint() {
367 let result =
368 FeedHttpClient::should_follow_redirect("http://169.254.169.254/latest/meta-data/", 1);
369 assert!(result.is_err());
370 }
371
372 #[test]
373 fn test_redirect_rejects_private_ip() {
374 let result = FeedHttpClient::should_follow_redirect("http://10.0.0.5/admin", 1);
375 assert!(result.is_err());
376 }
377
378 #[test]
379 fn test_redirect_allows_public_ip_within_hop_limit() {
380 let result = FeedHttpClient::should_follow_redirect("http://8.8.8.8/", 1);
381 assert!(result.is_ok());
382 }
383
384 #[test]
385 fn test_redirect_rejects_over_hop_limit_even_for_safe_url() {
386 let result = FeedHttpClient::should_follow_redirect("http://8.8.8.8/", MAX_REDIRECTS + 1);
387 assert!(result.is_err());
388 }
389
390 #[test]
391 fn test_redirect_allows_at_hop_limit() {
392 let result = FeedHttpClient::should_follow_redirect("http://8.8.8.8/", MAX_REDIRECTS);
393 assert!(result.is_ok());
394 }
395
396 #[test]
397 #[allow(clippy::significant_drop_tightening)]
398 fn test_redirect_to_metadata_endpoint_rejected_end_to_end() {
399 let mut server = mockito::Server::new();
400 let mock = server
401 .mock("GET", "/redirect")
402 .with_status(302)
403 .with_header("location", "http://169.254.169.254/latest/meta-data/")
404 .create();
405
406 let client = Client::builder()
410 .redirect(FeedHttpClient::redirect_policy())
411 .build()
412 .unwrap();
413
414 let url = format!("{}/redirect", server.url());
415 let result = client.get(&url).send();
416
417 let err = result.expect_err("redirect to a metadata IP must be rejected");
418 let description = FeedHttpClient::describe_request_error(&err);
424 assert!(
425 description.contains("Link-local address not allowed"),
426 "expected the SSRF rejection reason in the error chain, got: {description}"
427 );
428 mock.assert();
429 }
430
431 #[test]
432 #[allow(clippy::significant_drop_tightening)]
433 fn test_redirect_follows_legitimate_multi_hop_chain() {
434 let mut server = mockito::Server::new();
435 let addr = server.socket_address();
436
437 let mock1 = server
438 .mock("GET", "/hop1")
439 .with_status(302)
440 .with_header("location", "http://public.test/hop2")
441 .create();
442 let mock2 = server
443 .mock("GET", "/hop2")
444 .with_status(302)
445 .with_header("location", "http://public.test/hop3")
446 .create();
447 let mock3 = server
448 .mock("GET", "/hop3")
449 .with_status(200)
450 .with_body("ok")
451 .create();
452
453 let client = Client::builder()
459 .redirect(FeedHttpClient::redirect_policy())
460 .dns_resolver(Arc::new(SsrfSafeResolver))
461 .resolve("public.test", addr)
462 .build()
463 .unwrap();
464
465 let url = format!("http://public.test:{}/hop1", addr.port());
466 let response = client.get(&url).send().unwrap();
467
468 assert_eq!(response.status().as_u16(), 200);
469 mock1.assert();
470 mock2.assert();
471 mock3.assert();
472 }
473
474 #[test]
475 #[allow(clippy::significant_drop_tightening)]
476 fn test_redirect_chain_rejects_metadata_after_legitimate_hops() {
477 let mut server = mockito::Server::new();
478 let addr = server.socket_address();
479
480 let mock1 = server
481 .mock("GET", "/hop1")
482 .with_status(302)
483 .with_header("location", "http://public.test/hop2")
484 .create();
485 let mock2 = server
486 .mock("GET", "/hop2")
487 .with_status(302)
488 .with_header("location", "http://169.254.169.254/latest/meta-data/")
489 .create();
490
491 let client = Client::builder()
492 .redirect(FeedHttpClient::redirect_policy())
493 .dns_resolver(Arc::new(SsrfSafeResolver))
494 .resolve("public.test", addr)
495 .build()
496 .unwrap();
497
498 let url = format!("http://public.test:{}/hop1", addr.port());
502 let err = client.get(&url).send().expect_err(
503 "a chain ending at a metadata IP must be rejected, even after legitimate hops",
504 );
505
506 let description = FeedHttpClient::describe_request_error(&err);
507 assert!(
508 description.contains("Link-local address not allowed"),
509 "expected the SSRF rejection reason to survive a multi-hop chain, got: {description}"
510 );
511 mock1.assert();
512 mock2.assert();
513 }
514
515 #[test]
516 fn test_dns_resolver_wired_into_client_rejects_loopback() {
517 let client = Client::builder()
518 .dns_resolver(Arc::new(SsrfSafeResolver))
519 .build()
520 .unwrap();
521
522 let result = client.get("http://localhost/").send();
527 assert!(result.is_err());
528 }
529
530 #[test]
531 fn test_custom_user_agent() {
532 let client = FeedHttpClient::new()
533 .unwrap()
534 .with_user_agent("CustomBot/1.0".to_string());
535 assert_eq!(client.user_agent, "CustomBot/1.0");
536 }
537
538 #[test]
539 fn test_custom_timeout() {
540 let timeout = Duration::from_secs(60);
541 let client = FeedHttpClient::new().unwrap().with_timeout(timeout);
542 assert_eq!(client.timeout, timeout);
543 }
544
545 #[test]
546 fn test_with_timeout_clamps_absurd_duration() {
547 let client = FeedHttpClient::new().unwrap().with_timeout(Duration::MAX);
550 assert_eq!(client.timeout, Duration::from_secs(3600));
551 }
552
553 #[test]
554 fn test_build_request_applies_configured_timeout() {
555 let timeout = Duration::from_secs(7);
559 let client = FeedHttpClient::new().unwrap().with_timeout(timeout);
560 let request = client
561 .build_request("http://example.test/feed.xml", HeaderMap::new())
562 .unwrap();
563 assert_eq!(request.timeout(), Some(&timeout));
564 }
565
566 #[test]
569 #[allow(clippy::significant_drop_tightening)]
570 fn test_with_timeout_enforced_on_slow_response() {
571 let mut server = mockito::Server::new();
572 let addr = server.socket_address();
573
574 let mock = server
575 .mock("GET", "/slow")
576 .with_chunked_body(|w| {
577 std::thread::sleep(Duration::from_millis(500));
578 w.write_all(b"too slow")
579 })
580 .create();
581
582 let raw_client = Client::builder()
587 .dns_resolver(Arc::new(SsrfSafeResolver))
588 .resolve("public.test", addr)
589 .build()
590 .unwrap();
591
592 let client = FeedHttpClient {
593 client: raw_client,
594 user_agent: "test-agent".to_string(),
595 timeout: Duration::from_millis(100),
596 };
597
598 let url = format!("http://public.test:{}/slow", addr.port());
599 let started = std::time::Instant::now();
600 client
601 .get(&url, None, None, None)
602 .expect_err("a response slower than the configured 100ms timeout must fail");
603 let elapsed = started.elapsed();
604
605 assert!(
610 elapsed < Duration::from_millis(400),
611 "request did not fail until {elapsed:?}; timeout was not enforced"
612 );
613 mock.assert();
614 }
615
616 #[test]
617 #[allow(clippy::significant_drop_tightening)]
618 fn test_with_timeout_allows_fast_response() {
619 let mut server = mockito::Server::new();
620 let addr = server.socket_address();
621
622 let mock = server
623 .mock("GET", "/fast")
624 .with_status(200)
625 .with_body("ok")
626 .create();
627
628 let raw_client = Client::builder()
629 .dns_resolver(Arc::new(SsrfSafeResolver))
630 .resolve("public.test", addr)
631 .build()
632 .unwrap();
633
634 let client = FeedHttpClient {
635 client: raw_client,
636 user_agent: "test-agent".to_string(),
637 timeout: Duration::from_secs(5),
638 };
639
640 let url = format!("http://public.test:{}/fast", addr.port());
641 let response = client
642 .get(&url, None, None, None)
643 .expect("a response faster than the configured timeout must succeed");
644
645 assert_eq!(response.status, 200);
646 mock.assert();
647 }
648
649 #[test]
651 fn test_reject_localhost_url() {
652 let client = FeedHttpClient::new().unwrap();
653 let result = client.get("http://localhost/feed.xml", None, None, None);
654 assert!(result.is_err());
655 let err_msg = result.err().unwrap().to_string();
656 assert!(err_msg.contains("Localhost domain not allowed"));
657 }
658
659 #[test]
660 fn test_reject_private_ip() {
661 let client = FeedHttpClient::new().unwrap();
662 let result = client.get("http://192.168.1.1/feed.xml", None, None, None);
663 assert!(result.is_err());
664 let err_msg = result.err().unwrap().to_string();
665 assert!(err_msg.contains("Private IP address not allowed"));
666 }
667
668 #[test]
669 fn test_reject_metadata_endpoint() {
670 let client = FeedHttpClient::new().unwrap();
671 let result = client.get("http://169.254.169.254/latest/meta-data/", None, None, None);
672 assert!(result.is_err());
673 let err_msg = result.err().unwrap().to_string();
674 assert!(err_msg.contains("metadata") || err_msg.contains("Link-local"));
676 }
677
678 #[test]
679 fn test_reject_file_scheme() {
680 let client = FeedHttpClient::new().unwrap();
681 let result = client.get("file:///etc/passwd", None, None, None);
682 assert!(result.is_err());
683 let err_msg = result.err().unwrap().to_string();
684 assert!(err_msg.contains("Unsupported URL scheme"));
685 }
686
687 #[test]
688 fn test_reject_internal_domain() {
689 let client = FeedHttpClient::new().unwrap();
690 let result = client.get("http://server.local/feed.xml", None, None, None);
691 assert!(result.is_err());
692 let err_msg = result.err().unwrap().to_string();
693 assert!(err_msg.contains("Internal domain TLD not allowed"));
694 }
695
696 #[test]
697 fn test_insert_header_valid() {
698 let mut headers = HeaderMap::new();
699 let result =
700 FeedHttpClient::insert_header(&mut headers, USER_AGENT, "TestBot/1.0", "User-Agent");
701 assert!(result.is_ok());
702 assert_eq!(headers.get(USER_AGENT).unwrap(), "TestBot/1.0");
703 }
704
705 #[test]
706 fn test_insert_header_invalid_value() {
707 let mut headers = HeaderMap::new();
708 let result = FeedHttpClient::insert_header(
710 &mut headers,
711 USER_AGENT,
712 "Invalid\nHeader",
713 "User-Agent",
714 );
715 assert!(result.is_err());
716 match result {
717 Err(FeedError::Http { message }) => {
718 assert!(message.contains("Invalid User-Agent"));
719 }
720 _ => panic!("Expected Http error"),
721 }
722 }
723
724 #[test]
725 fn test_insert_header_multiple_headers() {
726 let mut headers = HeaderMap::new();
727
728 FeedHttpClient::insert_header(&mut headers, USER_AGENT, "TestBot/1.0", "User-Agent")
729 .unwrap();
730
731 FeedHttpClient::insert_header(&mut headers, ACCEPT, "application/xml", "Accept").unwrap();
732
733 assert_eq!(headers.len(), 2);
734 assert_eq!(headers.get(USER_AGENT).unwrap(), "TestBot/1.0");
735 assert_eq!(headers.get(ACCEPT).unwrap(), "application/xml");
736 }
737}