diff --git a/src/main/java/org/example/ConnectionHandler.java b/src/main/java/org/example/ConnectionHandler.java index 1fcc29fb..8b72c8fb 100644 --- a/src/main/java/org/example/ConnectionHandler.java +++ b/src/main/java/org/example/ConnectionHandler.java @@ -35,6 +35,9 @@ public ConnectionHandler(Socket client, String webRoot) { private List buildFilters() { List list = new ArrayList<>(); + + list.add(new org.example.filter.RateLimitingFilter()); + AppConfig config = ConfigLoader.get(); AppConfig.IpFilterConfig ipFilterConfig = config.ipFilter(); if (Boolean.TRUE.equals(ipFilterConfig.enabled())) { @@ -73,7 +76,8 @@ public void runConnectionHandler() throws IOException { int statusCode = response.getStatusCode(); if (statusCode == HttpResponseBuilder.SC_FORBIDDEN || - statusCode == HttpResponseBuilder.SC_BAD_REQUEST) { + statusCode == HttpResponseBuilder.SC_BAD_REQUEST || + statusCode == HttpResponseBuilder.SC_TOO_MANY_REQUESTS) { byte[] responseBytes = response.build(); client.getOutputStream().write(responseBytes); client.getOutputStream().flush(); diff --git a/src/main/java/org/example/filter/RateLimitingFilter.java b/src/main/java/org/example/filter/RateLimitingFilter.java new file mode 100644 index 00000000..e5a9667b --- /dev/null +++ b/src/main/java/org/example/filter/RateLimitingFilter.java @@ -0,0 +1,162 @@ +package org.example.filter; + +import io.github.bucket4j.Bandwidth; +import io.github.bucket4j.Bucket; +import org.example.http.HttpResponseBuilder; +import org.example.httpparser.HttpRequest; +import java.time.Duration; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.logging.Logger; +import java.util.concurrent.atomic.AtomicBoolean; + +/** + * Rate Limiting Filter responsible for limiting the number of requests per client IP. + * Implements the Token Bucket algorithm using the Bucket4j library. + * How it works: + * A "bucket" hold a fixed number of tokens (capacity) + * Each incoming request attempts to consume exactly one token + * If a token is available, the request is processed and the token is removed + * If the bucket is empty, the request is rejected with an HTTP 429 (Too Many Requests) status + * Tokens are replenished at a fixed rate over time (Refill Rate), up to the maximum capaci + * This allows for occasional bursts of traffic while maintaining a steady long-term rate limit + * The capacity of the bucket is 10, and it refills one token per 10 seconds + */ + +public class RateLimitingFilter implements Filter { + private static final Logger logger = Logger.getLogger(RateLimitingFilter.class.getName()); + private static final Map buckets = new ConcurrentHashMap<>(); + private static final long CAPACITY = 10; + private static final long REFILL_TOKENS = 1; + private final Duration refillPeriod = Duration.ofSeconds(10); + private static final int MAX_BUCKETS_THRESHOLD = 1000; + private final AtomicBoolean cleanupStarted = new AtomicBoolean(false); + private volatile Thread cleanupThread; + + @Override + public void init() { + logger.info("RateLimitingFilter initialized with capacity: " + CAPACITY); + if (cleanupStarted.compareAndSet(false, true)) { + cleanupThread = startCleanupThread(); + } + } + + /** + * Intercepts the request and checks if the client has enough tokens. + */ + @Override + public void doFilter(HttpRequest request, HttpResponseBuilder response, FilterChain chain) { + + String clientIp = resolveClientIp(request, response); + + if (clientIp == null) return; + + BucketWrapper wrapper = buckets.computeIfAbsent(clientIp, k -> new BucketWrapper(createNewBucket())); + + wrapper.updateAccess(); + + if (wrapper.bucket.tryConsume(1)) { + chain.doFilter(request, response); + } else { + logger.warning("Limit exceeded per IP: " + clientIp); + response.setStatusCode(HttpResponseBuilder.SC_TOO_MANY_REQUESTS); + response.setBody("

429 Too Many Requests

Limit of requests exceeded.

\n"); + } + } + + @Override + public void destroy() { + Thread t = cleanupThread; + if (t != null) { + t.interrupt(); + cleanupThread = null; + } + cleanupStarted.set(false); + buckets.clear(); + } + + /** + * Configures a new Bucket with the specified bandwidth. + */ + private Bucket createNewBucket() { + return Bucket.builder() + .addLimit(Bandwidth.builder() + .capacity(CAPACITY) + .refillGreedy(REFILL_TOKENS, refillPeriod) + .build()) + .build(); + } + + /** + * Track the last access time of every bucket + */ + private static class BucketWrapper { + private final Bucket bucket; + private volatile long lastAccessTime; + + BucketWrapper(Bucket bucket) { + this.bucket = bucket; + this.lastAccessTime = System.currentTimeMillis(); + } + + void updateAccess() { + this.lastAccessTime = System.currentTimeMillis(); + } + } + + public static String resolveClientIp(HttpRequest request, HttpResponseBuilder response) { + + Object clientIpAttr = request.getAttribute("clientIp"); + + if (!(clientIpAttr instanceof String clientIp) || (clientIp.isBlank())) { + response.setStatusCode(HttpResponseBuilder.SC_BAD_REQUEST); + response.setBody("

400 Bad Request

Missing client IP.

\n"); + return null; + } + + String xForwardedFor = request.getHeaders().get("X-Forwarded-For"); + + if( xForwardedFor != null && !xForwardedFor.isBlank() ) { + clientIp = xForwardedFor.split(",")[0].trim(); + } + + return clientIp; + } + + public Thread startCleanupThread() { + return Thread.ofVirtual().name("rate-limit-cleanup").start(() -> { + while (!Thread.currentThread().isInterrupted()) { + try { + //it checks every 10 minutes + Thread.sleep(Duration.ofMinutes(10).toMillis()); + + cleanupIdleBuckets(); + + } catch (InterruptedException _) { + Thread.currentThread().interrupt(); + break; + } + } + }); + } + + public void cleanupIdleBuckets() { + //it will only clean when the size of the buckets is more than 1000 + if (buckets.size() > MAX_BUCKETS_THRESHOLD) { + long idleThreshold = System.currentTimeMillis() - Duration.ofMinutes(30).toMillis(); + buckets.entrySet().removeIf(entry -> entry.getValue().lastAccessTime < idleThreshold); + } + } + + public int getBucketsCount() { + return buckets.size(); + } + + public void ageBucketsForTesting(long millisToSubtract) { + for (BucketWrapper wrapper : buckets.values()) { + long oldTime = wrapper.lastAccessTime; + wrapper.lastAccessTime = oldTime - millisToSubtract; + } + } + +} diff --git a/src/main/java/org/example/http/HttpResponseBuilder.java b/src/main/java/org/example/http/HttpResponseBuilder.java index bd4026af..8480b22c 100644 --- a/src/main/java/org/example/http/HttpResponseBuilder.java +++ b/src/main/java/org/example/http/HttpResponseBuilder.java @@ -26,6 +26,7 @@ public class HttpResponseBuilder { public static final int SC_UNAUTHORIZED = 401; public static final int SC_FORBIDDEN = 403; public static final int SC_NOT_FOUND = 404; + public static final int SC_TOO_MANY_REQUESTS = 429; // SERVER ERROR public static final int SC_INTERNAL_SERVER_ERROR = 500; @@ -57,6 +58,7 @@ public class HttpResponseBuilder { Map.entry(SC_UNAUTHORIZED, "Unauthorized"), Map.entry(SC_FORBIDDEN, "Forbidden"), Map.entry(SC_NOT_FOUND, "Not Found"), + Map.entry(SC_TOO_MANY_REQUESTS, "Too Many Requests"), Map.entry(SC_INTERNAL_SERVER_ERROR, "Internal Server Error"), Map.entry(SC_BAD_GATEWAY, "Bad Gateway"), Map.entry(SC_SERVICE_UNAVAILABLE, "Service Unavailable"), diff --git a/src/test/java/org/example/filter/RateLimitingFilterIpTest.java b/src/test/java/org/example/filter/RateLimitingFilterIpTest.java new file mode 100644 index 00000000..f27cd8ec --- /dev/null +++ b/src/test/java/org/example/filter/RateLimitingFilterIpTest.java @@ -0,0 +1,48 @@ +package org.example.filter; + +import org.example.http.HttpResponseBuilder; +import org.example.httpparser.HttpRequest; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import java.util.HashMap; +import java.util.Map; + +import static org.example.filter.RateLimitingFilter.resolveClientIp; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +class RateLimitingFilterIpTest { + + @Mock HttpRequest request; + @Mock HttpResponseBuilder response; + + @Test + void shouldUseXForwarded_WhenPresent(){ + + Map headers = new HashMap<>(); + headers.put("X-Forwarded-For", "203.0.113.195"); + + when(request.getHeaders()).thenReturn(headers); + when(request.getAttribute("clientIp")).thenReturn("127.0.0.1"); + + String finalIp = resolveClientIp(request, response); + + assertEquals("203.0.113.195", finalIp); + } + + @Test + void shouldFallbackToAttribute_WhenXForwardedForIsNotPresent(){ + + Map headers = new HashMap<>(); + when(request.getHeaders()).thenReturn(headers); + when(request.getAttribute("clientIp")).thenReturn("10.0.0.5"); + + String finalIp = resolveClientIp(request, response); + + assertEquals("10.0.0.5", finalIp); + } +} diff --git a/src/test/java/org/example/filter/RateLimitingFilterTest.java b/src/test/java/org/example/filter/RateLimitingFilterTest.java new file mode 100644 index 00000000..d115d9c4 --- /dev/null +++ b/src/test/java/org/example/filter/RateLimitingFilterTest.java @@ -0,0 +1,138 @@ +package org.example.filter; + +import org.example.http.HttpResponseBuilder; +import org.example.httpparser.HttpRequest; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import java.util.HashMap; + +import static org.example.http.HttpResponseBuilder.SC_OK; +import static org.example.http.HttpResponseBuilder.SC_TOO_MANY_REQUESTS; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.*; + +@ExtendWith(MockitoExtension.class) +class RateLimitingFilterTest { + + @Mock + FilterChain filterChain; + + private RateLimitingFilter filter; + private HttpRequest request; + private HttpResponseBuilder response; + + @BeforeEach + void setUp(){ + filter = new RateLimitingFilter(); + request = new HttpRequest("GET", "/", "HTTP/1.1", new HashMap<>(), ""); + request.setAttribute("clientIp", "127.0.0.1"); + response = new HttpResponseBuilder(); + filter.destroy(); + } + + @Test + void shouldAllowRequest_WhenTokensAreAvailable(){ + + filter.doFilter(request, response, filterChain); + + verify(filterChain, times(1)).doFilter(request, response); + assertEquals(SC_OK, response.getStatusCode()); + } + + @Test + void shouldNotAllowRequest_WhenTokensAreNotAvailable(){ + + //capacity of the bucket is 10 + for(int i = 0; i < 11; i++ ) + filter.doFilter(request, response, filterChain); + + assertEquals(SC_TOO_MANY_REQUESTS, response.getStatusCode()); + verify(filterChain, times(10)).doFilter(any(), any()); + } + + @Test + void shouldHaveSeparateBucketsPerIp(){ + + //Request 1 + for(int i = 0; i < 11; i++) + filter.doFilter(request, response, filterChain); + + //Request 2 with a different Ip + HttpRequest request2 = new HttpRequest("GET", "/", "HTTP/1.1", new HashMap<>(), ""); + request2.setAttribute("clientIp", "127.2.2.2"); + HttpResponseBuilder response2 = new HttpResponseBuilder(); + + filter.doFilter(request2, response2, filterChain); + + //First request should be 429 because it exceeded the capacity of the bucket (10) + assertEquals(SC_TOO_MANY_REQUESTS, response.getStatusCode()); + //Second request should be 200 + assertEquals(SC_OK, response2.getStatusCode()); + } + + @Test + void shouldDeleteOldBuckets_WhenSizeIsMoreThanThreshold(){ + + filter.init(); + + for(int i = 0; i < 1001; i++ ){ + String fakeIp = "192.168.1." + i; + request.setAttribute("clientIp", fakeIp); + filter.doFilter(request, response, filterChain); + } + + assertEquals(1001, filter.getBucketsCount()); + + filter.ageBucketsForTesting(3600000); + filter.cleanupIdleBuckets(); + + assertEquals(0, filter.getBucketsCount()); + } + + @Test + void shouldNotDeleteOldBuckets_WhenSizeIsLessThanThreshold(){ + filter.init(); + + for(int i = 0; i < 1000; i++ ){ + String fakeIp = "192.168.1." + i; + request.setAttribute("clientIp", fakeIp); + filter.doFilter(request, response, filterChain); + } + + assertEquals(1000, filter.getBucketsCount()); + + filter.ageBucketsForTesting(3600000); + filter.cleanupIdleBuckets(); + + assertEquals(1000, filter.getBucketsCount()); + + } + + @Test + void shouldDeleteOnlyExpiredBuckets_WhenAreOld(){ + filter.init(); + + for(int i = 0; i < 1001; i++ ){ + String fakeIp = "192.168.1." + i; + request.setAttribute("clientIp", fakeIp); + filter.doFilter(request, response, filterChain); + } + + assertEquals(1001, filter.getBucketsCount()); + + filter.ageBucketsForTesting(3600000); + + for(int i = 0; i < 500; i++ ){ + String fakeIp = "192.168.1." + i; + request.setAttribute("clientIp", fakeIp); + filter.doFilter(request, response, filterChain); + } + + filter.cleanupIdleBuckets(); + + assertEquals(500, filter.getBucketsCount()); + } +}