Skip to content

Commit 24c1a34

Browse files
authored
41-add-cors-filter-to-enable-cross-origin-requests (#120)
* Add CORS filter to enable cross-origin requests * Make HttpResponse mutable to support CORS and response-modifying filters HttpResponse was immutable, which prevented filters from adding headers or modifying status/body. Since our filter chain relies on mutating the shared response object, immutability blocked correct CORS implementation. Added setter methods and removed unmodifiable headers wrapper. HttpResponseWriter remains unaffected. * Add unit tests for CORS filter functionality * Add tests for CORS filter handling of origin headers and preflight requests * Make HttpResponse mutable and update CORS filter to set status and headers correctly * Enhance CORS filter to manage 'Vary' header and fix 'Access-Control-Allow-Headers' setting * Add default constructor to HttpResponse for easier instantiation * Refactor CORS filter tests to use mock requests and responses
1 parent 1e443d8 commit 24c1a34

3 files changed

Lines changed: 201 additions & 1 deletion

File tree

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,74 @@
1+
package org.juv25d.filter;
2+
3+
import org.juv25d.http.HttpRequest;
4+
import org.juv25d.http.HttpResponse;
5+
6+
import java.io.IOException;
7+
import java.util.Map;
8+
import java.util.Set;
9+
10+
public class CorsFilter implements Filter {
11+
12+
// Whitelist, allow known origins:
13+
private static final Set<String> ALLOWED_ORIGINS = Set.of(
14+
"http://localhost:3000"
15+
);
16+
17+
// Supported methods
18+
private static final String ALLOWED_METHODS = "GET, POST, PUT, PATCH, DELETE, OPTIONS";
19+
20+
@Override
21+
public void doFilter(HttpRequest req, HttpResponse res, FilterChain chain) throws IOException {
22+
String origin = header(req.headers(), "Origin");
23+
24+
// No Origin header || no browser cross-origin req, No CORS headers needed
25+
if (origin == null || origin.isBlank()) {
26+
chain.doFilter(req, res);
27+
return;
28+
}
29+
30+
// Origin exists but are not allowed, return no CORS headers
31+
if (!ALLOWED_ORIGINS.contains(origin)) {
32+
chain.doFilter(req, res);
33+
return;
34+
}
35+
36+
// Allowed origin => AC-AO on all res, even GET
37+
res.setHeader("Access-Control-Allow-Origin", origin);
38+
String vary = res.getHeader("Vary");
39+
if (vary == null || vary.isBlank()) {
40+
res.setHeader("Vary", "Origin");
41+
} else if (!vary.toLowerCase().contains("origin")) {
42+
res.setHeader("Vary", vary + ", Origin");
43+
}
44+
45+
// Preflight, OPTIONS
46+
if ("OPTIONS".equalsIgnoreCase(req.method())) {
47+
res.setHeader("Access-Control-Allow-Methods", ALLOWED_METHODS);
48+
49+
// If browser requests specific headers, mirror
50+
String requestedHeaders = header(req.headers(), "Access-Control-Request-Headers");
51+
if (requestedHeaders != null && !requestedHeaders.isBlank()) {
52+
res.setHeader("Access-Control-Allow-Headers", requestedHeaders);
53+
} else {
54+
res.setHeader("Access-Control-Allow-Headers", "Content-Type");
55+
}
56+
res.setHeader("Access-Control-Max-Age", "3600");
57+
res.setStatusText("No Content");
58+
res.setStatusCode(204);
59+
res.setBody(new byte[0]);
60+
return;
61+
}
62+
// Regular request (GET/POST)
63+
chain.doFilter(req, res);
64+
}
65+
66+
private String header(Map<String, String> headers, String key) {
67+
for (var entry : headers.entrySet()) {
68+
if (entry.getKey() != null && entry.getKey().equalsIgnoreCase(key)) {
69+
return entry.getValue();
70+
}
71+
}
72+
return null;
73+
}
74+
}

src/main/java/org/juv25d/http/HttpResponse.java

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@ public class HttpResponse {
1111

1212
private int statusCode;
1313
private String statusText;
14-
private final Map<String, String> headers;
14+
private Map<String, String> headers;
1515
private byte[] body;
1616

1717
public HttpResponse() {
@@ -49,6 +49,18 @@ public Map<String, String> headers() {
4949
return headers;
5050
}
5151

52+
public String getHeader(String name) {
53+
if (name == null) {
54+
return null;
55+
}
56+
for (var entry : headers.entrySet()) {
57+
if (entry.getKey() != null && entry.getKey().equalsIgnoreCase(name)) {
58+
return entry.getValue();
59+
}
60+
}
61+
return null;
62+
}
63+
5264
public void setHeader(String name, String value) {
5365
headers.put(name, value);
5466
}
Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,114 @@
1+
package org.juv25d.filter;
2+
3+
import org.juv25d.http.HttpRequest;
4+
import org.juv25d.http.HttpResponse;
5+
import org.junit.jupiter.api.Test;
6+
7+
import java.util.HashMap;
8+
import java.util.Map;
9+
10+
import static org.junit.jupiter.api.Assertions.*;
11+
import static org.mockito.Mockito.*;
12+
13+
public class CorsFilterTest {
14+
15+
private final CorsFilter filter = new CorsFilter();
16+
17+
@Test
18+
void shouldAllowConfiguredOrigin_onGet() throws Exception {
19+
HttpRequest req = request("GET", "/api/test", Map.of("Origin", "http://localhost:3000"));
20+
HttpResponse res = response();
21+
FilterChain chain = mock(FilterChain.class);
22+
23+
filter.doFilter(req, res, chain);
24+
25+
verify(chain).doFilter(req, res);
26+
assertEquals("http://localhost:3000", res.getHeader("Access-Control-Allow-Origin"));
27+
assertEquals("Origin", res.getHeader("Vary"));
28+
}
29+
30+
@Test
31+
void shouldNotAddCorsHeaders_whenNoOriginHeader() throws Exception {
32+
HttpRequest req = request("GET", "/api/test", Map.of());
33+
HttpResponse res = response();
34+
FilterChain chain = mock(FilterChain.class);
35+
36+
filter.doFilter(req, res, chain);
37+
38+
verify(chain).doFilter(req, res);
39+
assertNull(res.getHeader("Access-Control-Allow-Origin"));
40+
}
41+
42+
@Test
43+
void shouldHandlePreflightOptionsRequest() throws Exception {
44+
HttpRequest req = request(
45+
"OPTIONS",
46+
"/api/test",
47+
Map.of(
48+
"Origin", "http://localhost:3000",
49+
"Access-Control-Request-Method", "GET",
50+
"Access-Control-Request-Headers", "Content-Type"
51+
)
52+
);
53+
HttpResponse res = response();
54+
FilterChain chain = mock(FilterChain.class);
55+
56+
filter.doFilter(req, res, chain);
57+
58+
verify(chain, never()).doFilter(any(), any());
59+
assertEquals(204, res.statusCode());
60+
assertEquals("No Content", res.statusText());
61+
assertEquals("http://localhost:3000", res.getHeader("Access-Control-Allow-Origin"));
62+
assertTrue(res.getHeader("Access-Control-Allow-Methods").contains("GET"));
63+
assertEquals("Content-Type", res.getHeader("Access-Control-Allow-Headers"));
64+
assertArrayEquals(new byte[0], res.body());
65+
}
66+
67+
@Test
68+
void shouldNotAllowUnknownOrigin() throws Exception {
69+
HttpRequest req = request("GET", "/api/test", Map.of("Origin", "http://evil.com"));
70+
HttpResponse res = response();
71+
FilterChain chain = mock(FilterChain.class);
72+
73+
filter.doFilter(req, res, chain);
74+
75+
verify(chain).doFilter(req, res);
76+
assertNull(res.getHeader("Access-Control-Allow-Origin"));
77+
}
78+
79+
@Test
80+
void shouldFallbackToDefaultAllowHeaders_onPreflightWithoutRequestHeaders() throws Exception {
81+
HttpRequest req = request(
82+
"OPTIONS",
83+
"/api/test",
84+
Map.of(
85+
"Origin", "http://localhost:3000",
86+
"Access-Control-Request-Method", "GET"
87+
)
88+
);
89+
HttpResponse res = response();
90+
FilterChain chain = mock(FilterChain.class);
91+
92+
filter.doFilter(req, res, chain);
93+
94+
verify(chain, never()).doFilter(any(), any());
95+
assertEquals(204, res.statusCode());
96+
assertEquals("Content-Type", res.getHeader("Access-Control-Allow-Headers"));
97+
}
98+
99+
private static HttpRequest request(String method, String path, Map<String, String> headers) {
100+
return new HttpRequest(
101+
method,
102+
path,
103+
"",
104+
"HTTP/1.1",
105+
new HashMap<>(headers),
106+
new byte[0],
107+
"127.0.0.1"
108+
);
109+
}
110+
111+
private static HttpResponse response() {
112+
return new HttpResponse(200, "OK", new HashMap<>(), new byte[0]);
113+
}
114+
}

0 commit comments

Comments
 (0)