fix(personal): enforce URL fetch deadlines and framing

This commit is contained in:
2026-07-12 04:36:47 +08:00
parent a909806435
commit 2c9c6d5d4a
2 changed files with 386 additions and 65 deletions
@@ -35,6 +35,15 @@ import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.ArrayBlockingQueue;
import java.util.concurrent.ExecutionException;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Future;
import java.util.concurrent.RejectedExecutionException;
import java.util.concurrent.ThreadPoolExecutor;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException;
import java.util.concurrent.atomic.AtomicInteger;
@Service
public class PersonalUrlFetchService {
@@ -51,6 +60,14 @@ public class PersonalUrlFetchService {
private static final int MAX_HEADER_BYTES = 64 * 1024;
private static final int MAX_LINE_BYTES = 8 * 1024;
private static final long TOTAL_TIMEOUT_NANOS = 15_000_000_000L;
private static final long HARD_MAX_BODY_BYTES = 10L * 1024 * 1024;
private static final AtomicInteger DNS_THREAD_SEQUENCE = new AtomicInteger();
private static final ExecutorService DNS_EXECUTOR = new ThreadPoolExecutor(2, 2, 0L, TimeUnit.MILLISECONDS,
new ArrayBlockingQueue<>(8), runnable -> {
Thread thread = new Thread(runnable, "personal-url-dns-" + DNS_THREAD_SEQUENCE.incrementAndGet());
thread.setDaemon(true);
return thread;
}, new ThreadPoolExecutor.AbortPolicy());
private static final String USER_AGENT = "wygj-personal-url-fetch/1.0";
private static final String ACCEPT = "text/html,text/plain,application/pdf,application/msword,application/vnd.ms-excel,application/vnd.ms-powerpoint,application/vnd.openxmlformats-officedocument.wordprocessingml.document,application/vnd.openxmlformats-officedocument.spreadsheetml.sheet,application/vnd.openxmlformats-officedocument.presentationml.presentation";
private static final Map<String, String> SAFE_HEADERS = Map.of(
@@ -72,7 +89,8 @@ public class PersonalUrlFetchService {
private final Fetcher fetcher;
public PersonalUrlFetchService(PersonalKnowledgeProperties properties) {
this(properties, host -> Arrays.asList(InetAddress.getAllByName(host)), new RawSocketFetcher());
this(properties, new DeadlineDnsResolver(DNS_EXECUTOR,
host -> Arrays.asList(InetAddress.getAllByName(host))), new RawSocketFetcher());
}
private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) {
@@ -88,19 +106,20 @@ public class PersonalUrlFetchService {
/** Validate syntax, DNS answers and address policy. */
public URI validate(String rawUrl) {
return validateAndResolve(rawUrl).uri();
return validateAndResolve(rawUrl, System.nanoTime() + TOTAL_TIMEOUT_NANOS).uri();
}
/** Capture a bounded public web resource without persisting it. */
public FetchResult fetch(String rawUrl) {
long maxBodyBytes = maxBodyBytes();
long deadline = System.nanoTime() + TOTAL_TIMEOUT_NANOS;
ValidatedTarget target = validateAndResolve(rawUrl);
ValidatedTarget target = validateAndResolve(rawUrl, deadline);
Set<URI> visited = new HashSet<>();
visited.add(target.uri());
int redirects = 0;
while (true) {
requireTimeRemaining(deadline);
TransportResponse response;
try {
response = fetcher.fetch(new FetchRequest(target.uri(), target.addresses(), deadline,
@@ -110,6 +129,7 @@ public class PersonalUrlFetchService {
} catch (Exception ex) {
throw new ServiceException(FETCH_FAILED);
}
requireTimeRemaining(deadline);
if (response == null || response.body() == null || response.body().length > maxBodyBytes) {
throw new ServiceException(RESPONSE_TOO_LARGE);
}
@@ -124,7 +144,7 @@ public class PersonalUrlFetchService {
} catch (IllegalArgumentException ex) {
throw new ServiceException(BLOCKED);
}
target = validateAndResolve(next.toString());
target = validateAndResolve(next.toString(), deadline);
if (!visited.add(target.uri())) throw new ServiceException(REDIRECT_LOOP);
redirects++;
continue;
@@ -141,11 +161,13 @@ public class PersonalUrlFetchService {
}
}
private ValidatedTarget validateAndResolve(String rawUrl) {
private ValidatedTarget validateAndResolve(String rawUrl, long deadlineNanos) {
URI uri = normalizeUri(rawUrl);
List<InetAddress> addresses;
try {
addresses = resolver.resolve(canonicalHost(uri));
requireTimeRemaining(deadlineNanos);
addresses = resolver.resolve(canonicalHost(uri), deadlineNanos);
requireTimeRemaining(deadlineNanos);
} catch (Exception ex) {
throw new ServiceException(BLOCKED);
}
@@ -191,12 +213,16 @@ public class PersonalUrlFetchService {
try {
long value = Math.multiplyExact(properties.getMaxUrlBodyMb(), 1024L * 1024L);
if (value <= 0) throw new ArithmeticException();
return value;
return Math.min(value, HARD_MAX_BODY_BYTES);
} catch (ArithmeticException ex) {
throw new ServiceException(RESPONSE_TOO_LARGE);
}
}
private static void requireTimeRemaining(long deadlineNanos) {
if (deadlineNanos - System.nanoTime() <= 0) throw new ServiceException(FETCH_FAILED);
}
private static boolean isRedirect(int status) {
return status == 301 || status == 302 || status == 303 || status == 307 || status == 308;
}
@@ -213,16 +239,45 @@ public class PersonalUrlFetchService {
}
private static void enforceDeclaredLength(Map<String, List<String>> headers, long maxBodyBytes) {
String raw = firstHeader(headers, "content-length");
String transferEncoding = strictFramingHeader(headers, "transfer-encoding");
String raw = strictFramingHeader(headers, "content-length");
if (transferEncoding != null && raw != null) throw new ServiceException(RESPONSE_INVALID);
if (transferEncoding != null && !"chunked".equalsIgnoreCase(transferEncoding)) {
throw new ServiceException(RESPONSE_INVALID);
}
if (raw == null) return;
long length = parseContentLength(raw);
if (length > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE);
}
private static long parseContentLength(String raw) {
String value = raw.trim();
if (value.isEmpty() || !value.chars().allMatch(Character::isDigit)) {
throw new ServiceException(RESPONSE_INVALID);
}
try {
long length = Long.parseLong(raw.trim());
if (length < 0 || length > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE);
return Long.parseLong(value);
} catch (NumberFormatException ex) {
throw new ServiceException(RESPONSE_INVALID);
}
}
private static String strictFramingHeader(Map<String, List<String>> headers, String name) {
if (headers == null) return null;
String found = null;
for (Map.Entry<String, List<String>> entry : headers.entrySet()) {
if (entry.getKey() == null || !entry.getKey().equalsIgnoreCase(name)) continue;
if (found != null || entry.getValue() == null || entry.getValue().size() != 1) {
throw new ServiceException(RESPONSE_INVALID);
}
found = entry.getValue().get(0);
if (found == null || found.isBlank() || found.indexOf(',') >= 0) {
throw new ServiceException(RESPONSE_INVALID);
}
}
return found;
}
private static String firstHeader(Map<String, List<String>> headers, String name) {
if (headers == null) return null;
for (Map.Entry<String, List<String>> entry : headers.entrySet()) {
@@ -280,9 +335,54 @@ public class PersonalUrlFetchService {
@FunctionalInterface
public interface Resolver {
List<InetAddress> resolve(String host, long deadlineNanos) throws IOException;
}
@FunctionalInterface
interface HostLookup {
List<InetAddress> resolve(String host) throws UnknownHostException;
}
static final class DeadlineDnsResolver implements Resolver {
private final ExecutorService executor;
private final HostLookup lookup;
DeadlineDnsResolver(ExecutorService executor, HostLookup lookup) {
this.executor = executor;
this.lookup = lookup;
}
@Override
public List<InetAddress> resolve(String host, long deadlineNanos) throws IOException {
long remaining = deadlineNanos - System.nanoTime();
if (remaining <= 0) throw new IOException("resolution deadline exceeded");
Future<List<InetAddress>> future;
try {
future = executor.submit(() -> lookup.resolve(host));
} catch (RejectedExecutionException ex) {
throw new IOException("resolution unavailable");
}
try {
return future.get(remaining, TimeUnit.NANOSECONDS);
} catch (TimeoutException ex) {
cancelAndPurge(future);
throw new IOException("resolution deadline exceeded");
} catch (InterruptedException ex) {
cancelAndPurge(future);
Thread.currentThread().interrupt();
throw new IOException("resolution interrupted");
} catch (ExecutionException ex) {
cancelAndPurge(future);
throw new IOException("resolution failed");
}
}
private void cancelAndPurge(Future<?> future) {
future.cancel(true);
if (executor instanceof ThreadPoolExecutor pool) pool.purge();
}
}
@FunctionalInterface
public interface Fetcher {
TransportResponse fetch(FetchRequest request) throws IOException;
@@ -311,7 +411,29 @@ public class PersonalUrlFetchService {
private record ValidatedTarget(URI uri, List<InetAddress> addresses) { }
private static final class RawSocketFetcher implements Fetcher {
interface Connection extends AutoCloseable {
InputStream input() throws IOException;
OutputStream output() throws IOException;
void setReadTimeout(int millis) throws IOException;
@Override void close() throws IOException;
}
@FunctionalInterface
interface ConnectionFactory {
Connection connect(URI uri, InetAddress address, int port, int connectTimeout, int readTimeout) throws IOException;
}
static final class RawSocketFetcher implements Fetcher {
private final ConnectionFactory connections;
RawSocketFetcher() {
this(new JvmConnectionFactory());
}
RawSocketFetcher(ConnectionFactory connections) {
this.connections = connections;
}
@Override
public TransportResponse fetch(FetchRequest request) throws IOException {
IOException last = null;
@@ -325,37 +447,32 @@ public class PersonalUrlFetchService {
throw last == null ? new IOException("connection failed") : last;
}
private static TransportResponse fetchAddress(FetchRequest request, InetAddress address) throws IOException {
private TransportResponse fetchAddress(FetchRequest request, InetAddress address) throws IOException {
URI uri = request.uri();
int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80);
Socket plain = new Socket();
try {
plain.connect(new InetSocketAddress(address, port), timeout(request.deadlineNanos(), 5_000));
plain.setSoTimeout(timeout(request.deadlineNanos(), 5_000));
Socket active = plain;
if ("https".equals(uri.getScheme())) {
String tlsHost = canonicalHost(uri);
SSLSocket ssl = (SSLSocket) ((SSLSocketFactory) SSLSocketFactory.getDefault())
.createSocket(plain, tlsHost, port, true);
SSLParameters parameters = ssl.getSSLParameters();
parameters.setEndpointIdentificationAlgorithm("HTTPS");
if (!isIpLiteral(tlsHost)) parameters.setServerNames(List.of(new SNIHostName(tlsHost)));
ssl.setSSLParameters(parameters);
ssl.setSoTimeout(timeout(request.deadlineNanos(), 5_000));
ssl.startHandshake();
active = ssl;
}
writeRequest(active.getOutputStream(), request);
active.setSoTimeout(timeout(request.deadlineNanos(), 5_000));
int connectTimeout = timeout(request.deadlineNanos(), 5_000);
int readTimeout = timeout(request.deadlineNanos(), 5_000);
try (Connection connection = connections.connect(uri, address, port, connectTimeout, readTimeout)) {
writeRequest(connection.output(), request);
connection.setReadTimeout(timeout(request.deadlineNanos(), 5_000));
TransportResponse response = parseHttpResponse(
new DeadlineInputStream(active.getInputStream(), active, request.deadlineNanos()), request.maxBodyBytes());
new DeadlineInputStream(connection.input(), connection, request.deadlineNanos()), request.maxBodyBytes());
if (System.nanoTime() >= request.deadlineNanos()) throw new IOException("deadline exceeded");
return response;
} finally {
try { plain.close(); } catch (IOException ignored) { }
}
}
static SSLParameters tlsParameters(String host) {
SSLParameters parameters = new SSLParameters();
configureTlsParameters(parameters, host);
return parameters;
}
private static void configureTlsParameters(SSLParameters parameters, String host) {
parameters.setEndpointIdentificationAlgorithm("HTTPS");
if (!isIpLiteral(host)) parameters.setServerNames(List.of(new SNIHostName(host)));
}
private static void writeRequest(OutputStream output, FetchRequest request) throws IOException {
URI uri = request.uri();
String target = uri.getRawPath().isEmpty() ? "/" : uri.getRawPath();
@@ -363,7 +480,7 @@ public class PersonalUrlFetchService {
String host = hostHeader(uri);
StringBuilder value = new StringBuilder("GET ").append(target).append(" HTTP/1.1\r\nHost: ")
.append(host).append("\r\n");
request.headers().forEach((name, headerValue) -> value.append(name).append(": ")
SAFE_HEADERS.forEach((name, headerValue) -> value.append(name).append(": ")
.append(headerValue).append("\r\n"));
value.append("Connection: close\r\n\r\n");
output.write(value.toString().getBytes(StandardCharsets.US_ASCII));
@@ -388,26 +505,63 @@ public class PersonalUrlFetchService {
}
}
private static final class JvmConnectionFactory implements ConnectionFactory {
@Override
public Connection connect(URI uri, InetAddress address, int port,
int connectTimeout, int readTimeout) throws IOException {
Socket plain = new Socket();
try {
// The socket connects to the exact address already approved by the resolver policy.
plain.connect(new InetSocketAddress(address, port), connectTimeout);
plain.setSoTimeout(readTimeout);
Socket active = plain;
if ("https".equals(uri.getScheme())) {
String tlsHost = canonicalHost(uri);
// JVM defaults preserve the configured trust store; no permissive trust manager is installed.
SSLSocket ssl = (SSLSocket) ((SSLSocketFactory) SSLSocketFactory.getDefault())
.createSocket(plain, tlsHost, port, true);
SSLParameters parameters = ssl.getSSLParameters();
RawSocketFetcher.configureTlsParameters(parameters, tlsHost);
ssl.setSSLParameters(parameters);
ssl.setSoTimeout(readTimeout);
ssl.startHandshake();
active = ssl;
}
return new SocketConnection(active);
} catch (IOException | RuntimeException ex) {
try { plain.close(); } catch (IOException ignored) { }
throw ex;
}
}
}
private record SocketConnection(Socket socket) implements Connection {
@Override public InputStream input() throws IOException { return socket.getInputStream(); }
@Override public OutputStream output() throws IOException { return socket.getOutputStream(); }
@Override public void setReadTimeout(int millis) throws IOException { socket.setSoTimeout(millis); }
@Override public void close() throws IOException { socket.close(); }
}
private static final class DeadlineInputStream extends InputStream {
private final InputStream delegate;
private final Socket socket;
private final Connection connection;
private final long deadlineNanos;
private DeadlineInputStream(InputStream delegate, Socket socket, long deadlineNanos) {
private DeadlineInputStream(InputStream delegate, Connection connection, long deadlineNanos) {
this.delegate = delegate;
this.socket = socket;
this.connection = connection;
this.deadlineNanos = deadlineNanos;
}
@Override
public int read() throws IOException {
socket.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
connection.setReadTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
return delegate.read();
}
@Override
public int read(byte[] bytes, int offset, int length) throws IOException {
socket.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
connection.setReadTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
return delegate.read(bytes, offset, length);
}
}
@@ -437,18 +591,23 @@ public class PersonalUrlFetchService {
if (name.isEmpty()) throw new ServiceException(RESPONSE_INVALID);
headers.computeIfAbsent(name, ignored -> new ArrayList<>()).add(value);
}
String transferEncoding = firstHeader(headers, "transfer-encoding");
String contentLength = firstHeader(headers, "content-length");
String transferEncoding = strictFramingHeader(headers, "transfer-encoding");
String contentLength = strictFramingHeader(headers, "content-length");
if (transferEncoding != null && contentLength != null) throw new ServiceException(RESPONSE_INVALID);
byte[] body;
if (transferEncoding != null) {
if (hasNoBody(status)) {
if (transferEncoding != null) throw new ServiceException(RESPONSE_INVALID);
if (status != 304 && contentLength != null) {
throw new ServiceException(RESPONSE_INVALID);
}
if (contentLength != null) parseContentLength(contentLength);
body = new byte[0];
} else if (transferEncoding != null) {
if (!"chunked".equalsIgnoreCase(transferEncoding.trim())) throw new ServiceException(RESPONSE_INVALID);
body = readChunked(buffered, maxBodyBytes);
} else if (contentLength != null) {
long length;
try { length = Long.parseLong(contentLength.trim()); }
catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); }
if (length < 0 || length > maxBodyBytes || length > Integer.MAX_VALUE) {
long length = parseContentLength(contentLength);
if (length > maxBodyBytes || length > Integer.MAX_VALUE) {
throw new ServiceException(RESPONSE_TOO_LARGE);
}
body = readExactly(buffered, (int) length);
@@ -463,6 +622,10 @@ public class PersonalUrlFetchService {
}
}
private static boolean hasNoBody(int status) {
return status >= 100 && status < 200 || status == 204 || status == 304;
}
private static byte[] readChunked(BufferedInputStream input, long maxBodyBytes) throws IOException {
ByteArrayOutputStream body = new ByteArrayOutputStream();
int[] framingBytes = {0};
@@ -5,7 +5,12 @@ import org.dromara.common.core.exception.ServiceException;
import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test;
import javax.net.ssl.SNIHostName;
import javax.net.ssl.SSLParameters;
import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.net.InetAddress;
import java.net.URI;
import java.nio.charset.StandardCharsets;
@@ -15,6 +20,11 @@ import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ArrayBlockingQueue;
import java.util.concurrent.ThreadPoolExecutor;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference;
import static org.junit.jupiter.api.Assertions.*;
@@ -25,7 +35,7 @@ class PersonalUrlFetchServiceTest {
@Test
void rejectsUnsafeSchemesSyntaxAndHosts() {
var service = fixture(host -> List.of(PUBLIC), request -> ok("text/plain", "ok"));
var service = fixture((host, deadline) -> List.of(PUBLIC), request -> ok("text/plain", "ok"));
for (String raw : List.of(
"file:///etc/passwd", "ftp://example.com/a", "data:text/plain,hello",
"http://user:secret@example.com", "http:///missing", "not a url",
@@ -44,7 +54,7 @@ class PersonalUrlFetchServiceTest {
"::", "::1", "fe80::1", "fc00::1", "fd00::1", "ff02::1",
"2001:db8::1", "2001:2::1", "2001:100::1", "2002:5db8:d822::1",
"3fff::1", "64:ff9b::c0a8:101")) {
var service = fixture(host -> List.of(address(ip)), request -> ok("text/plain", "ok"));
var service = fixture((host, deadline) -> List.of(address(ip)), request -> ok("text/plain", "ok"));
assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://example.com"), ip);
}
}
@@ -52,17 +62,17 @@ class PersonalUrlFetchServiceTest {
@Test
void rejectsEmptyOrMixedDnsAnswersAndAcceptsPublicResolution() {
assertCode("PERSONAL_URL_BLOCKED",
() -> fixture(host -> List.of(), request -> ok("text/plain", "ok")).validate("https://example.com"), "empty");
() -> fixture((host, deadline) -> List.of(), request -> ok("text/plain", "ok")).validate("https://example.com"), "empty");
assertCode("PERSONAL_URL_BLOCKED",
() -> fixture(host -> List.of(PUBLIC, address("127.0.0.1")), request -> ok("text/plain", "ok"))
() -> fixture((host, deadline) -> List.of(PUBLIC, address("127.0.0.1")), request -> ok("text/plain", "ok"))
.validate("https://example.com"), "mixed");
assertCode("PERSONAL_URL_BLOCKED",
() -> fixture(host -> Arrays.asList(PUBLIC, null), request -> ok("text/plain", "ok"))
() -> fixture((host, deadline) -> Arrays.asList(PUBLIC, null), request -> ok("text/plain", "ok"))
.validate("https://example.com"), "null answer");
assertEquals("https://example.com/a", fixture(host -> List.of(PUBLIC), request -> ok("text/plain", "ok"))
assertEquals("https://example.com/a", fixture((host, deadline) -> List.of(PUBLIC), request -> ok("text/plain", "ok"))
.validate("HTTPS://Example.COM/a").toString());
assertEquals("http://[2606:2800:220:1:248:1893:25c8:1946]/", fixture(
host -> List.of(address("2606:2800:220:1:248:1893:25c8:1946")), request -> ok("text/plain", "ok"))
(host, deadline) -> List.of(address("2606:2800:220:1:248:1893:25c8:1946")), request -> ok("text/plain", "ok"))
.validate("http://[2606:2800:220:1:248:1893:25c8:1946]/").toString());
}
@@ -72,7 +82,7 @@ class PersonalUrlFetchServiceTest {
var responses = new ArrayDeque<PersonalUrlFetchService.TransportResponse>();
responses.add(response(302, Map.of("location", List.of("/final")), new byte[0]));
responses.add(ok("text/plain; charset=utf-8", "done"));
var service = fixture(host -> List.of(PUBLIC), request -> { seen.add(request); return responses.remove(); });
var service = fixture((host, deadline) -> List.of(PUBLIC), request -> { seen.add(request); return responses.remove(); });
var result = service.fetch("https://example.com/start");
@@ -87,11 +97,11 @@ class PersonalUrlFetchServiceTest {
@Test
void blocksUnsafeRedirectAndMixedAddressRedirect() {
var redirect = response(302, Map.of("location", List.of("http://metadata.test/latest")), new byte[0]);
var service = fixture(host -> host.equals("metadata.test")
var service = fixture((host, deadline) -> host.equals("metadata.test")
? List.of(address("169.254.169.254")) : List.of(PUBLIC), request -> redirect);
assertCode("PERSONAL_URL_BLOCKED", () -> service.fetch("https://example.com/start"), "redirect private");
var mixed = fixture(host -> host.equals("mixed.test")
var mixed = fixture((host, deadline) -> host.equals("mixed.test")
? List.of(PUBLIC, address("10.0.0.1")) : List.of(PUBLIC), request ->
response(301, Map.of("location", List.of("https://mixed.test/a")), new byte[0]));
assertCode("PERSONAL_URL_BLOCKED", () -> mixed.fetch("https://example.com"), "redirect mixed");
@@ -99,11 +109,11 @@ class PersonalUrlFetchServiceTest {
@Test
void detectsRedirectLoopAndMoreThanThreeRedirects() {
var loop = fixture(host -> List.of(PUBLIC), request ->
var loop = fixture((host, deadline) -> List.of(PUBLIC), request ->
response(302, Map.of("location", List.of(request.uri().toString())), new byte[0]));
assertCode("PERSONAL_URL_REDIRECT_LOOP", () -> loop.fetch("https://example.com/a"), "loop");
var chain = fixture(host -> List.of(PUBLIC), request -> {
var chain = fixture((host, deadline) -> List.of(PUBLIC), request -> {
int n = Integer.parseInt(request.uri().getPath().substring(1));
return response(302, Map.of("location", List.of("/" + (n + 1))), new byte[0]);
});
@@ -113,7 +123,7 @@ class PersonalUrlFetchServiceTest {
@Test
void sendsOnlyFixedSafeHeaders() {
var requests = new ArrayList<PersonalUrlFetchService.FetchRequest>();
var service = fixture(host -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); });
var service = fixture((host, deadline) -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); });
service.fetch("https://example.com/a");
Map<String, String> headers = requests.get(0).headers();
@@ -128,20 +138,29 @@ class PersonalUrlFetchServiceTest {
@Test
void rejectsForbiddenOrMissingMimeAndOversizedBody() {
assertCode("PERSONAL_URL_CONTENT_TYPE_UNSUPPORTED",
() -> fixture(host -> List.of(PUBLIC), request -> ok("image/png", "x")).fetch("https://example.com"), "mime");
() -> fixture((host, deadline) -> List.of(PUBLIC), request -> ok("image/png", "x")).fetch("https://example.com"), "mime");
assertCode("PERSONAL_URL_CONTENT_TYPE_UNSUPPORTED",
() -> fixture(host -> List.of(PUBLIC), request -> response(200, Map.of(), "x".getBytes())).fetch("https://example.com"), "missing mime");
() -> fixture((host, deadline) -> List.of(PUBLIC), request -> response(200, Map.of(), "x".getBytes())).fetch("https://example.com"), "missing mime");
PersonalKnowledgeProperties properties = properties();
byte[] tooLarge = new byte[10 * 1024 * 1024 + 1];
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () ->
PersonalUrlFetchService.forTest(properties, host -> List.of(PUBLIC),
PersonalUrlFetchService.forTest(properties, (host, deadline) -> List.of(PUBLIC),
request -> ok("text/plain", tooLarge)).fetch("https://example.com"), "body cap");
properties.setMaxUrlBodyMb(100);
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () ->
PersonalUrlFetchService.forTest(properties, (host, deadline) -> List.of(PUBLIC),
request -> ok("text/plain", tooLarge)).fetch("https://example.com"), "hard cap");
properties.setMaxUrlBodyMb(0);
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () ->
PersonalUrlFetchService.forTest(properties, (host, deadline) -> List.of(PUBLIC),
request -> ok("text/plain", "ok")).fetch("https://example.com"), "invalid configured cap");
}
@Test
void returnsDigestAndCaptureMetadata() {
var result = fixture(host -> List.of(PUBLIC), request -> ok("application/pdf", "abc"))
var result = fixture((host, deadline) -> List.of(PUBLIC), request -> ok("application/pdf", "abc"))
.fetch("https://example.com/a.pdf");
assertEquals(200, result.status());
assertEquals("application/pdf", result.contentType());
@@ -176,13 +195,127 @@ class PersonalUrlFetchServiceTest {
@Test
void fetchRequestCarriesValidatedIpsAndOriginalTlsHost() {
var requests = new ArrayList<PersonalUrlFetchService.FetchRequest>();
fixture(host -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); })
fixture((host, deadline) -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); })
.fetch("https://example.com/path");
assertEquals("example.com", requests.get(0).uri().getHost());
assertEquals(List.of(PUBLIC), requests.get(0).addresses());
assertTrue(requests.get(0).deadlineNanos() > System.nanoTime());
}
@Test
void boundedProductionDnsResolverTimesOutAndCancels() {
AtomicInteger interrupted = new AtomicInteger();
ThreadPoolExecutor executor = new ThreadPoolExecutor(1, 1, 0, TimeUnit.MILLISECONDS,
new ArrayBlockingQueue<>(1), runnable -> { Thread thread = new Thread(runnable, "dns-test"); thread.setDaemon(true); return thread; },
new ThreadPoolExecutor.AbortPolicy());
try {
var resolver = new PersonalUrlFetchService.DeadlineDnsResolver(executor, host -> {
try { Thread.sleep(5_000); }
catch (InterruptedException ex) { interrupted.incrementAndGet(); Thread.currentThread().interrupt(); }
return List.of(PUBLIC);
});
assertThrows(IOException.class, () -> resolver.resolve("example.com", System.nanoTime() + 20_000_000L));
assertTrue(interrupted.get() > 0 || executor.getActiveCount() == 0);
} finally {
executor.shutdownNow();
}
}
@Test
void rejectsAmbiguousTransferAndContentLengthFraming() {
for (String headers : List.of(
"Content-Length: 1\r\nContent-Length: 1\r\n",
"Content-Length: 1, 1\r\n",
"Content-Length: +1\r\n",
"Content-Length: -1\r\n",
"Content-Length: 999999999999999999999999\r\n",
"Transfer-Encoding: chunked\r\nContent-Length: 1\r\n",
"Transfer-Encoding: gzip\r\n",
"Transfer-Encoding: chunked, gzip\r\n",
"Transfer-Encoding: chunked\r\nTransfer-Encoding: chunked\r\n")) {
assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(
stream("HTTP/1.1 200 OK\r\n" + headers + "\r\nx"), 100), headers);
}
}
@Test
void noBodyStatusesDoNotWaitForPayload() {
for (int status : List.of(100, 204, 304)) {
var response = PersonalUrlFetchService.parseHttpResponse(stream(
"HTTP/1.1 " + status + " No Body\r\nContent-Type: text/plain\r\n\r\n"), 100);
assertEquals(0, response.body().length);
}
for (int status : List.of(100, 204)) {
assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(stream(
"HTTP/1.1 " + status + " No Body\r\nContent-Length: 0\r\n\r\n"), 100), "forbidden content length");
}
}
@Test
void rejectsOversizedHeaderBlockAndLine() {
String longLine = "X-Large: " + "x".repeat(8 * 1024 + 1);
assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(
stream("HTTP/1.1 200 OK\r\n" + longLine + "\r\n\r\n"), 100), "line");
StringBuilder headers = new StringBuilder("HTTP/1.1 200 OK\r\n");
for (int i = 0; i < 9000; i++) headers.append("X-").append(i).append(": x\r\n");
headers.append("\r\n");
assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(
stream(headers.toString()), 100), "block");
}
@Test
void rawTransportConnectsValidatedIpAndWritesOnlySafeRequestIdentity() throws Exception {
AtomicReference<InetAddress> connected = new AtomicReference<>();
AtomicReference<String> host = new AtomicReference<>();
AtomicInteger connectTimeoutSeen = new AtomicInteger();
AtomicInteger readTimeoutSeen = new AtomicInteger();
ByteArrayOutputStream requestBytes = new ByteArrayOutputStream();
byte[] rawResponse = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nok"
.getBytes(StandardCharsets.US_ASCII);
PersonalUrlFetchService.ConnectionFactory factory = (uri, address, port, connectTimeout, readTimeout) -> {
connected.set(address); host.set(uri.getHost());
connectTimeoutSeen.set(connectTimeout); readTimeoutSeen.set(readTimeout);
return new FakeConnection(new ByteArrayInputStream(rawResponse), requestBytes);
};
var fetcher = new PersonalUrlFetchService.RawSocketFetcher(factory);
var response = fetcher.fetch(new PersonalUrlFetchService.FetchRequest(URI.create("http://origin.example/a?b=1"),
List.of(PUBLIC), System.nanoTime() + TimeUnit.SECONDS.toNanos(1), 100, Map.of(
"User-Agent", "evil-agent", "Accept", "*/*", "Accept-Encoding", "gzip",
"Authorization", "Bearer secret", "Cookie", "sid=secret", "Referer", "https://secret.example")));
assertEquals(PUBLIC, connected.get());
assertEquals("origin.example", host.get());
assertTrue(connectTimeoutSeen.get() > 0 && connectTimeoutSeen.get() <= 5_000);
assertTrue(readTimeoutSeen.get() > 0 && readTimeoutSeen.get() <= 5_000);
assertEquals("ok", new String(response.body(), StandardCharsets.US_ASCII));
String request = requestBytes.toString(StandardCharsets.US_ASCII);
assertTrue(request.startsWith("GET /a?b=1 HTTP/1.1\r\nHost: origin.example\r\n"));
assertTrue(request.contains("User-Agent: wygj-personal-url-fetch/1.0\r\n"));
assertTrue(request.contains("Accept-Encoding: identity\r\n"));
assertFalse(request.toLowerCase().contains("cookie:"));
assertFalse(request.toLowerCase().contains("authorization:"));
assertFalse(request.toLowerCase().contains("referer:"));
}
@Test
void tlsParametersRetainOriginalHostnameVerification() {
SSLParameters parameters = PersonalUrlFetchService.RawSocketFetcher.tlsParameters("origin.example");
assertEquals("HTTPS", parameters.getEndpointIdentificationAlgorithm());
assertEquals("origin.example", ((SNIHostName) parameters.getServerNames().get(0)).getAsciiName());
}
@Test
void rawTransportEnforcesTotalDeadlineWhileReadingSlowBody() {
byte[] rawResponse = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nok"
.getBytes(StandardCharsets.US_ASCII);
PersonalUrlFetchService.ConnectionFactory factory = (uri, address, port, connectTimeout, readTimeout) ->
new FakeConnection(new SlowInputStream(rawResponse, 30), new ByteArrayOutputStream());
var fetcher = new PersonalUrlFetchService.RawSocketFetcher(factory);
var request = new PersonalUrlFetchService.FetchRequest(URI.create("http://origin.example/"), List.of(PUBLIC),
System.nanoTime() + 10_000_000L, 100, Map.of());
assertThrows(IOException.class, () -> fetcher.fetch(request));
}
private static PersonalUrlFetchService fixture(PersonalUrlFetchService.Resolver resolver,
PersonalUrlFetchService.Fetcher fetcher) {
return PersonalUrlFetchService.forTest(properties(), resolver, fetcher);
@@ -218,4 +351,29 @@ class PersonalUrlFetchServiceTest {
private static void assertCode(String code, Runnable action, String context) {
assertEquals(code, assertThrows(ServiceException.class, action::run, context).getMessage(), context);
}
private static final class FakeConnection implements PersonalUrlFetchService.Connection {
private final InputStream input;
private final ByteArrayOutputStream output;
private FakeConnection(InputStream input, ByteArrayOutputStream output) { this.input = input; this.output = output; }
@Override public InputStream input() { return input; }
@Override public ByteArrayOutputStream output() { return output; }
@Override public void setReadTimeout(int millis) { assertTrue(millis > 0 && millis <= 5_000); }
@Override public void close() { }
}
private static final class SlowInputStream extends ByteArrayInputStream {
private final long delayMillis;
private SlowInputStream(byte[] bytes, long delayMillis) { super(bytes); this.delayMillis = delayMillis; }
@Override public synchronized int read(byte[] bytes, int offset, int length) {
try { Thread.sleep(delayMillis); }
catch (InterruptedException ex) { Thread.currentThread().interrupt(); }
return super.read(bytes, offset, length);
}
@Override public synchronized int read() {
try { Thread.sleep(delayMillis); }
catch (InterruptedException ex) { Thread.currentThread().interrupt(); }
return super.read();
}
}
}