Skip to main content

reinhardt_server/server/
rate_limit.rs

1#![allow(deprecated)] // RateLimitConfig is deprecated but still used by the handler and tests in this module.
2
3use reinhardt_http::Handler;
4use reinhardt_http::{Request, Response};
5use std::collections::HashMap;
6use std::net::IpAddr;
7use std::sync::Arc;
8use std::time::{Duration, Instant};
9use tokio::sync::RwLock;
10
11/// Rate limiting strategy
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum RateLimitStrategy {
14	/// Fixed window rate limiting
15	FixedWindow,
16	/// Sliding window rate limiting
17	SlidingWindow,
18}
19
20/// Rate limiter configuration
21#[deprecated(
22	since = "0.2.0",
23	note = "Use `RateLimitSettings` with the `#[settings]` macro instead."
24)]
25#[derive(Debug, Clone)]
26pub struct RateLimitConfig {
27	/// Maximum requests allowed in the window
28	pub max_requests: usize,
29	/// Time window duration
30	pub window_duration: Duration,
31	/// Rate limiting strategy
32	pub strategy: RateLimitStrategy,
33	/// Trusted proxy IP addresses/CIDRs.
34	/// Only requests from these IPs will have their X-Forwarded-For/X-Real-IP headers trusted.
35	pub trusted_proxies: Vec<String>,
36}
37
38impl RateLimitConfig {
39	/// Create a new rate limit configuration
40	///
41	/// # Examples
42	///
43	/// ```
44	/// use std::time::Duration;
45	/// use reinhardt_server::server::{RateLimitConfig, RateLimitStrategy};
46	///
47	/// let config = RateLimitConfig::new(
48	///     100,
49	///     Duration::from_secs(60),
50	///     RateLimitStrategy::FixedWindow,
51	/// );
52	/// ```
53	pub fn new(
54		max_requests: usize,
55		window_duration: Duration,
56		strategy: RateLimitStrategy,
57	) -> Self {
58		Self {
59			max_requests,
60			window_duration,
61			strategy,
62			trusted_proxies: Vec::new(),
63		}
64	}
65
66	/// Create a per-minute rate limit
67	///
68	/// # Examples
69	///
70	/// ```
71	/// use reinhardt_server::server::RateLimitConfig;
72	///
73	/// let config = RateLimitConfig::per_minute(60);
74	/// ```
75	pub fn per_minute(max_requests: usize) -> Self {
76		Self::new(
77			max_requests,
78			Duration::from_secs(60),
79			RateLimitStrategy::FixedWindow,
80		)
81	}
82
83	/// Create a per-hour rate limit
84	///
85	/// # Examples
86	///
87	/// ```
88	/// use reinhardt_server::server::RateLimitConfig;
89	///
90	/// let config = RateLimitConfig::per_hour(1000);
91	/// ```
92	pub fn per_hour(max_requests: usize) -> Self {
93		Self::new(
94			max_requests,
95			Duration::from_secs(3600),
96			RateLimitStrategy::FixedWindow,
97		)
98	}
99
100	/// Set trusted proxy addresses.
101	///
102	/// Only requests originating from these IP addresses will have their
103	/// `X-Forwarded-For` and `X-Real-IP` headers trusted for client IP extraction.
104	///
105	/// # Examples
106	///
107	/// ```
108	/// use reinhardt_server::server::RateLimitConfig;
109	///
110	/// let config = RateLimitConfig::per_minute(60)
111	///     .with_trusted_proxies(vec!["10.0.0.0/8".to_string()]);
112	/// ```
113	pub fn with_trusted_proxies(mut self, proxies: Vec<String>) -> Self {
114		self.trusted_proxies = proxies;
115		self
116	}
117}
118
119/// Rate limit entry for tracking requests (fixed window strategy)
120#[derive(Debug, Clone)]
121struct RateLimitEntry {
122	count: usize,
123	window_start: Instant,
124}
125
126/// Rate limit entry for tracking requests (sliding window strategy)
127///
128/// Stores individual request timestamps to enable true sliding window behavior,
129/// where only requests within the most recent `window_duration` are counted.
130#[derive(Debug, Clone)]
131struct SlidingWindowEntry {
132	timestamps: Vec<Instant>,
133}
134
135/// Middleware that implements rate limiting
136///
137/// Tracks requests by client IP address and enforces rate limits.
138/// Only trusts proxy headers (X-Forwarded-For, X-Real-IP) when the
139/// request comes from a configured trusted proxy address.
140///
141/// # Examples
142///
143/// ```
144/// use std::sync::Arc;
145/// use std::time::Duration;
146/// use reinhardt_server::server::{RateLimitHandler, RateLimitConfig};
147/// use reinhardt_http::Handler;
148/// use reinhardt_http::{Request, Response};
149///
150/// struct MyHandler;
151///
152/// #[async_trait::async_trait]
153/// impl Handler for MyHandler {
154///     async fn handle(&self, _req: Request) -> reinhardt_core::exception::Result<Response> {
155///         Ok(Response::ok())
156///     }
157/// }
158///
159/// let handler = Arc::new(MyHandler);
160/// let config = RateLimitConfig::per_minute(60);
161/// let rate_limit_handler = RateLimitHandler::new(handler, config);
162/// ```
163pub struct RateLimitHandler {
164	inner: Arc<dyn Handler>,
165	config: RateLimitConfig,
166	limits: Arc<RwLock<HashMap<IpAddr, RateLimitEntry>>>,
167	sliding_limits: Arc<RwLock<HashMap<IpAddr, SlidingWindowEntry>>>,
168}
169
170impl RateLimitHandler {
171	/// Create a new rate limit handler
172	///
173	/// # Arguments
174	///
175	/// * `inner` - The inner handler to wrap
176	/// * `config` - Rate limit configuration
177	///
178	/// # Examples
179	///
180	/// ```
181	/// use std::sync::Arc;
182	/// use reinhardt_server::server::{RateLimitHandler, RateLimitConfig};
183	/// use reinhardt_http::Handler;
184	/// use reinhardt_http::{Request, Response};
185	///
186	/// struct MyHandler;
187	///
188	/// #[async_trait::async_trait]
189	/// impl Handler for MyHandler {
190	///     async fn handle(&self, _req: Request) -> reinhardt_core::exception::Result<Response> {
191	///         Ok(Response::ok())
192	///     }
193	/// }
194	///
195	/// let handler = Arc::new(MyHandler);
196	/// let config = RateLimitConfig::per_minute(100);
197	/// let rate_limit_handler = RateLimitHandler::new(handler, config);
198	/// ```
199	pub fn new(inner: Arc<dyn Handler>, config: RateLimitConfig) -> Self {
200		Self {
201			inner,
202			config,
203			limits: Arc::new(RwLock::new(HashMap::new())),
204			sliding_limits: Arc::new(RwLock::new(HashMap::new())),
205		}
206	}
207
208	/// Check if a request is allowed for the given IP
209	///
210	/// Dispatches to the appropriate rate limiting algorithm based on the
211	/// configured strategy.
212	async fn is_allowed(&self, ip: IpAddr) -> bool {
213		match self.config.strategy {
214			RateLimitStrategy::FixedWindow => self.is_allowed_fixed_window(ip).await,
215			RateLimitStrategy::SlidingWindow => self.is_allowed_sliding_window(ip).await,
216		}
217	}
218
219	/// Fixed window rate limiting: resets the counter when the window expires.
220	///
221	/// Also performs periodic eviction of stale entries to prevent
222	/// unbounded memory growth from accumulated per-IP state.
223	async fn is_allowed_fixed_window(&self, ip: IpAddr) -> bool {
224		let now = Instant::now();
225		let mut limits = self.limits.write().await;
226
227		// Periodically evict stale entries (entries whose window has expired)
228		// to prevent unbounded memory growth.
229		if limits.len() > 1024 {
230			limits.retain(|_, entry| {
231				now.duration_since(entry.window_start) < self.config.window_duration * 2
232			});
233		}
234
235		let entry = limits.entry(ip).or_insert(RateLimitEntry {
236			count: 0,
237			window_start: now,
238		});
239
240		// Check if window has expired
241		if now.duration_since(entry.window_start) >= self.config.window_duration {
242			// Reset window
243			entry.count = 0;
244			entry.window_start = now;
245		}
246
247		// Check if under limit
248		if entry.count < self.config.max_requests {
249			entry.count += 1;
250			true
251		} else {
252			false
253		}
254	}
255
256	/// Sliding window rate limiting: counts requests within the most recent
257	/// `window_duration` period, allowing smoother rate distribution.
258	///
259	/// Unlike fixed window, this approach does not have boundary spikes where
260	/// `2 * max_requests` could be served across a window boundary.
261	async fn is_allowed_sliding_window(&self, ip: IpAddr) -> bool {
262		let now = Instant::now();
263		let window = self.config.window_duration;
264		let mut limits = self.sliding_limits.write().await;
265
266		// Periodically evict stale entries to prevent unbounded memory growth.
267		if limits.len() > 1024 {
268			limits.retain(|_, entry| {
269				entry
270					.timestamps
271					.last()
272					.is_some_and(|&last| now.duration_since(last) < window * 2)
273			});
274		}
275
276		let entry = limits.entry(ip).or_insert(SlidingWindowEntry {
277			timestamps: Vec::new(),
278		});
279
280		// Remove timestamps outside the current window
281		entry
282			.timestamps
283			.retain(|&ts| now.duration_since(ts) < window);
284
285		// Check if under limit
286		if entry.timestamps.len() < self.config.max_requests {
287			entry.timestamps.push(now);
288			true
289		} else {
290			false
291		}
292	}
293
294	/// Extract client IP from request
295	///
296	/// Only trusts proxy headers (X-Forwarded-For, X-Real-IP) when the request
297	/// originates from a configured trusted proxy address. Otherwise, uses the
298	/// direct connection IP (remote_addr) or falls back to localhost.
299	fn extract_client_ip(&self, request: &Request) -> IpAddr {
300		let peer_ip = request.remote_addr.map(|addr| addr.ip());
301
302		// Only trust proxy headers if the direct connection is from a trusted proxy
303		let from_trusted_proxy = peer_ip.map(|ip| self.is_trusted_proxy(ip)).unwrap_or(false);
304
305		if from_trusted_proxy {
306			// Check X-Forwarded-For header
307			if let Some(xff) = request.headers.get("X-Forwarded-For")
308				&& let Ok(xff_str) = xff.to_str()
309				&& let Some(first_ip) = xff_str.split(',').next()
310				&& let Ok(ip) = first_ip.trim().parse()
311			{
312				return ip;
313			}
314
315			// Check X-Real-IP header
316			if let Some(xri) = request.headers.get("X-Real-IP")
317				&& let Ok(ip_str) = xri.to_str()
318				&& let Ok(ip) = ip_str.parse()
319			{
320				return ip;
321			}
322		}
323
324		// Use remote_addr (direct connection IP)
325		if let Some(ip) = peer_ip {
326			return ip;
327		}
328
329		// Fallback to localhost
330		"127.0.0.1".parse().unwrap()
331	}
332
333	/// Check if an IP address belongs to a trusted proxy
334	fn is_trusted_proxy(&self, ip: IpAddr) -> bool {
335		self.config.trusted_proxies.iter().any(|proxy| {
336			// Try parsing as CIDR network
337			if let Ok(network) = proxy.parse::<ipnet::IpNet>() {
338				return network.contains(&ip);
339			}
340			// Try parsing as single IP
341			if let Ok(proxy_ip) = proxy.parse::<IpAddr>() {
342				return proxy_ip == ip;
343			}
344			false
345		})
346	}
347}
348
349#[async_trait::async_trait]
350impl Handler for RateLimitHandler {
351	async fn handle(&self, request: Request) -> reinhardt_core::exception::Result<Response> {
352		let client_ip = self.extract_client_ip(&request);
353
354		if self.is_allowed(client_ip).await {
355			self.inner.handle(request).await
356		} else {
357			Ok(Response::new(http::StatusCode::TOO_MANY_REQUESTS).with_body("Rate limit exceeded"))
358		}
359	}
360}
361
362#[cfg(test)]
363mod tests {
364	use super::*;
365	use std::time::Duration;
366
367	/// Polls a condition until it returns true or timeout is reached.
368	async fn poll_until<F, Fut>(
369		timeout: std::time::Duration,
370		interval: std::time::Duration,
371		mut condition: F,
372	) -> Result<(), String>
373	where
374		F: FnMut() -> Fut,
375		Fut: std::future::Future<Output = bool>,
376	{
377		let start = std::time::Instant::now();
378		while start.elapsed() < timeout {
379			if condition().await {
380				return Ok(());
381			}
382			tokio::time::sleep(interval).await;
383		}
384		Err(format!("Timeout after {:?} waiting for condition", timeout))
385	}
386
387	struct TestHandler;
388
389	#[async_trait::async_trait]
390	impl Handler for TestHandler {
391		async fn handle(&self, _request: Request) -> reinhardt_core::exception::Result<Response> {
392			Ok(Response::ok().with_body("Success"))
393		}
394	}
395
396	#[tokio::test]
397	async fn test_rate_limit_config_creation() {
398		// Arrange / Act
399		let config = RateLimitConfig::per_minute(60);
400
401		// Assert
402		assert_eq!(config.max_requests, 60);
403		assert_eq!(config.window_duration, Duration::from_secs(60));
404
405		let config = RateLimitConfig::per_hour(1000);
406		assert_eq!(config.max_requests, 1000);
407		assert_eq!(config.window_duration, Duration::from_secs(3600));
408	}
409
410	#[tokio::test]
411	async fn test_rate_limit_handler_creation() {
412		// Arrange / Act
413		let handler = Arc::new(TestHandler);
414		let config = RateLimitConfig::per_minute(10);
415		let _rate_limit_handler = RateLimitHandler::new(handler, config);
416	}
417
418	#[tokio::test]
419	async fn test_requests_within_limit() {
420		// Arrange
421		let handler = Arc::new(TestHandler);
422		let config = RateLimitConfig::per_minute(5);
423		let rate_limit_handler = RateLimitHandler::new(handler, config);
424
425		// Act / Assert
426		for _ in 0..5 {
427			let request = Request::builder()
428				.method(http::Method::GET)
429				.uri("/")
430				.version(http::Version::HTTP_11)
431				.headers(http::HeaderMap::new())
432				.body(bytes::Bytes::new())
433				.build()
434				.unwrap();
435
436			let response = rate_limit_handler.handle(request).await.unwrap();
437			assert_eq!(response.status, http::StatusCode::OK);
438		}
439	}
440
441	#[tokio::test]
442	async fn test_requests_exceed_limit() {
443		// Arrange
444		let handler = Arc::new(TestHandler);
445		let config = RateLimitConfig::per_minute(3);
446		let rate_limit_handler = RateLimitHandler::new(handler, config);
447
448		// Act - first 3 requests should succeed
449		for _ in 0..3 {
450			let request = Request::builder()
451				.method(http::Method::GET)
452				.uri("/")
453				.version(http::Version::HTTP_11)
454				.headers(http::HeaderMap::new())
455				.body(bytes::Bytes::new())
456				.build()
457				.unwrap();
458
459			let response = rate_limit_handler.handle(request).await.unwrap();
460			assert_eq!(response.status, http::StatusCode::OK);
461		}
462
463		// 4th request should be rate limited
464		let request = Request::builder()
465			.method(http::Method::GET)
466			.uri("/")
467			.version(http::Version::HTTP_11)
468			.headers(http::HeaderMap::new())
469			.body(bytes::Bytes::new())
470			.build()
471			.unwrap();
472
473		let response = rate_limit_handler.handle(request).await.unwrap();
474
475		// Assert
476		assert_eq!(response.status, http::StatusCode::TOO_MANY_REQUESTS);
477	}
478
479	#[tokio::test]
480	async fn test_rate_limit_window_reset() {
481		// Arrange
482		let handler = Arc::new(TestHandler);
483		let config = RateLimitConfig::new(
484			2,
485			Duration::from_millis(100),
486			RateLimitStrategy::FixedWindow,
487		);
488		let rate_limit_handler = RateLimitHandler::new(handler, config);
489
490		// Act - use up the limit
491		for _ in 0..2 {
492			let request = Request::builder()
493				.method(http::Method::GET)
494				.uri("/")
495				.version(http::Version::HTTP_11)
496				.headers(http::HeaderMap::new())
497				.body(bytes::Bytes::new())
498				.build()
499				.unwrap();
500			let response = rate_limit_handler.handle(request).await.unwrap();
501			assert_eq!(response.status, http::StatusCode::OK);
502		}
503
504		// Assert - poll until rate limit window resets (100ms window duration)
505		poll_until(
506			Duration::from_millis(200),
507			Duration::from_millis(10),
508			|| async {
509				let test_request = Request::builder()
510					.method(http::Method::GET)
511					.uri("/")
512					.version(http::Version::HTTP_11)
513					.headers(http::HeaderMap::new())
514					.body(bytes::Bytes::new())
515					.build()
516					.unwrap();
517				let test_response = rate_limit_handler.handle(test_request).await.unwrap();
518				test_response.status == http::StatusCode::OK
519			},
520		)
521		.await
522		.expect("Window should reset within 200ms");
523	}
524
525	// Client IP extraction tests
526
527	#[tokio::test]
528	async fn test_sliding_window_requests_within_limit() {
529		// Arrange
530		let handler = Arc::new(TestHandler);
531		let config = RateLimitConfig::new(
532			3,
533			Duration::from_millis(200),
534			RateLimitStrategy::SlidingWindow,
535		);
536		let rate_limit_handler = RateLimitHandler::new(handler, config);
537
538		// Act / Assert - first 3 requests should succeed
539		for _ in 0..3 {
540			let request = Request::builder()
541				.method(http::Method::GET)
542				.uri("/")
543				.version(http::Version::HTTP_11)
544				.headers(http::HeaderMap::new())
545				.body(bytes::Bytes::new())
546				.build()
547				.unwrap();
548
549			let response = rate_limit_handler.handle(request).await.unwrap();
550			assert_eq!(response.status, http::StatusCode::OK);
551		}
552	}
553
554	#[tokio::test]
555	async fn test_sliding_window_requests_exceed_limit() {
556		// Arrange
557		let handler = Arc::new(TestHandler);
558		let config = RateLimitConfig::new(
559			2,
560			Duration::from_millis(200),
561			RateLimitStrategy::SlidingWindow,
562		);
563		let rate_limit_handler = RateLimitHandler::new(handler, config);
564
565		// Act - first 2 requests should succeed
566		for _ in 0..2 {
567			let request = Request::builder()
568				.method(http::Method::GET)
569				.uri("/")
570				.version(http::Version::HTTP_11)
571				.headers(http::HeaderMap::new())
572				.body(bytes::Bytes::new())
573				.build()
574				.unwrap();
575			let response = rate_limit_handler.handle(request).await.unwrap();
576			assert_eq!(response.status, http::StatusCode::OK);
577		}
578
579		// 3rd request should be rate limited
580		let request = Request::builder()
581			.method(http::Method::GET)
582			.uri("/")
583			.version(http::Version::HTTP_11)
584			.headers(http::HeaderMap::new())
585			.body(bytes::Bytes::new())
586			.build()
587			.unwrap();
588		let response = rate_limit_handler.handle(request).await.unwrap();
589
590		// Assert
591		assert_eq!(response.status, http::StatusCode::TOO_MANY_REQUESTS);
592	}
593
594	#[tokio::test]
595	async fn test_sliding_window_expires_old_requests() {
596		// Arrange
597		let handler = Arc::new(TestHandler);
598		let config = RateLimitConfig::new(
599			2,
600			Duration::from_millis(100),
601			RateLimitStrategy::SlidingWindow,
602		);
603		let rate_limit_handler = RateLimitHandler::new(handler, config);
604
605		// Act - use up the limit
606		for _ in 0..2 {
607			let request = Request::builder()
608				.method(http::Method::GET)
609				.uri("/")
610				.version(http::Version::HTTP_11)
611				.headers(http::HeaderMap::new())
612				.body(bytes::Bytes::new())
613				.build()
614				.unwrap();
615			let response = rate_limit_handler.handle(request).await.unwrap();
616			assert_eq!(response.status, http::StatusCode::OK);
617		}
618
619		// Assert - poll until old timestamps expire (sliding window)
620		poll_until(
621			Duration::from_millis(200),
622			Duration::from_millis(10),
623			|| async {
624				let test_request = Request::builder()
625					.method(http::Method::GET)
626					.uri("/")
627					.version(http::Version::HTTP_11)
628					.headers(http::HeaderMap::new())
629					.body(bytes::Bytes::new())
630					.build()
631					.unwrap();
632				let test_response = rate_limit_handler.handle(test_request).await.unwrap();
633				test_response.status == http::StatusCode::OK
634			},
635		)
636		.await
637		.expect("Sliding window should allow requests after old timestamps expire");
638	}
639
640	#[test]
641	fn test_extract_client_ip_from_trusted_xff() {
642		// Arrange
643		let handler = Arc::new(TestHandler);
644		let config =
645			RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.1".to_string()]);
646		let rate_limit_handler = RateLimitHandler::new(handler, config);
647
648		let mut headers = http::HeaderMap::new();
649		headers.insert(
650			"X-Forwarded-For",
651			"192.168.1.100, 10.0.0.1, 172.16.0.1".parse().unwrap(),
652		);
653
654		let mut request = Request::builder()
655			.method(http::Method::GET)
656			.uri("/")
657			.version(http::Version::HTTP_11)
658			.headers(headers)
659			.body(bytes::Bytes::new())
660			.build()
661			.unwrap();
662		request.remote_addr = Some("10.0.0.1:12345".parse().unwrap());
663
664		// Act
665		let ip = rate_limit_handler.extract_client_ip(&request);
666
667		// Assert
668		assert_eq!(ip, "192.168.1.100".parse::<IpAddr>().unwrap());
669	}
670
671	#[test]
672	fn test_extract_client_ip_ignores_untrusted_xff() {
673		// Arrange
674		let handler = Arc::new(TestHandler);
675		let config =
676			RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.1".to_string()]);
677		let rate_limit_handler = RateLimitHandler::new(handler, config);
678
679		let mut headers = http::HeaderMap::new();
680		headers.insert("X-Forwarded-For", "192.168.1.100".parse().unwrap());
681
682		let mut request = Request::builder()
683			.method(http::Method::GET)
684			.uri("/")
685			.version(http::Version::HTTP_11)
686			.headers(headers)
687			.body(bytes::Bytes::new())
688			.build()
689			.unwrap();
690		// Untrusted source
691		request.remote_addr = Some("203.0.113.42:54321".parse().unwrap());
692
693		// Act
694		let ip = rate_limit_handler.extract_client_ip(&request);
695
696		// Assert - should use remote_addr, not spoofed header
697		assert_eq!(ip, "203.0.113.42".parse::<IpAddr>().unwrap());
698	}
699
700	#[test]
701	fn test_extract_client_ip_from_trusted_x_real_ip() {
702		// Arrange
703		let handler = Arc::new(TestHandler);
704		let config =
705			RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.0/8".to_string()]);
706		let rate_limit_handler = RateLimitHandler::new(handler, config);
707
708		let mut headers = http::HeaderMap::new();
709		headers.insert("X-Real-IP", "203.0.113.42".parse().unwrap());
710
711		let mut request = Request::builder()
712			.method(http::Method::GET)
713			.uri("/")
714			.version(http::Version::HTTP_11)
715			.headers(headers)
716			.body(bytes::Bytes::new())
717			.build()
718			.unwrap();
719		request.remote_addr = Some("10.0.0.5:8080".parse().unwrap());
720
721		// Act
722		let ip = rate_limit_handler.extract_client_ip(&request);
723
724		// Assert
725		assert_eq!(ip, "203.0.113.42".parse::<IpAddr>().unwrap());
726	}
727
728	#[test]
729	fn test_extract_client_ip_fallback_to_localhost() {
730		// Arrange
731		let handler = Arc::new(TestHandler);
732		let config = RateLimitConfig::per_minute(10);
733		let rate_limit_handler = RateLimitHandler::new(handler, config);
734
735		let headers = http::HeaderMap::new();
736		let request = Request::builder()
737			.method(http::Method::GET)
738			.uri("/")
739			.version(http::Version::HTTP_11)
740			.headers(headers)
741			.body(bytes::Bytes::new())
742			.build()
743			.unwrap();
744
745		// Act
746		let ip = rate_limit_handler.extract_client_ip(&request);
747
748		// Assert
749		assert_eq!(ip, "127.0.0.1".parse::<IpAddr>().unwrap());
750	}
751
752	#[test]
753	fn test_extract_client_ip_no_trusted_proxies() {
754		// Arrange - no trusted proxies, proxy headers should be ignored
755		let handler = Arc::new(TestHandler);
756		let config = RateLimitConfig::per_minute(10);
757		let rate_limit_handler = RateLimitHandler::new(handler, config);
758
759		let mut headers = http::HeaderMap::new();
760		headers.insert("X-Forwarded-For", "192.168.1.100".parse().unwrap());
761
762		let mut request = Request::builder()
763			.method(http::Method::GET)
764			.uri("/")
765			.version(http::Version::HTTP_11)
766			.headers(headers)
767			.body(bytes::Bytes::new())
768			.build()
769			.unwrap();
770		request.remote_addr = Some("203.0.113.1:8080".parse().unwrap());
771
772		// Act
773		let ip = rate_limit_handler.extract_client_ip(&request);
774
775		// Assert - uses remote_addr since no proxies are trusted
776		assert_eq!(ip, "203.0.113.1".parse::<IpAddr>().unwrap());
777	}
778
779	#[test]
780	fn test_extract_client_ip_with_invalid_header() {
781		// Arrange
782		let handler = Arc::new(TestHandler);
783		let config =
784			RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.1".to_string()]);
785		let rate_limit_handler = RateLimitHandler::new(handler, config);
786
787		let mut headers = http::HeaderMap::new();
788		headers.insert("X-Forwarded-For", "invalid-ip".parse().unwrap());
789
790		let mut request = Request::builder()
791			.method(http::Method::GET)
792			.uri("/")
793			.version(http::Version::HTTP_11)
794			.headers(headers)
795			.body(bytes::Bytes::new())
796			.build()
797			.unwrap();
798		request.remote_addr = Some("10.0.0.1:8080".parse().unwrap());
799
800		// Act
801		let ip = rate_limit_handler.extract_client_ip(&request);
802
803		// Assert - falls back to remote_addr when header is invalid
804		assert_eq!(ip, "10.0.0.1".parse::<IpAddr>().unwrap());
805	}
806}