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.Locale;
import java.util.Map; import java.util.Map;
import java.util.Set; 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 @Service
public class PersonalUrlFetchService { public class PersonalUrlFetchService {
@@ -51,6 +60,14 @@ public class PersonalUrlFetchService {
private static final int MAX_HEADER_BYTES = 64 * 1024; private static final int MAX_HEADER_BYTES = 64 * 1024;
private static final int MAX_LINE_BYTES = 8 * 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 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 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 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( private static final Map<String, String> SAFE_HEADERS = Map.of(
@@ -72,7 +89,8 @@ public class PersonalUrlFetchService {
private final Fetcher fetcher; private final Fetcher fetcher;
public PersonalUrlFetchService(PersonalKnowledgeProperties properties) { 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) { private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) {
@@ -88,19 +106,20 @@ public class PersonalUrlFetchService {
/** Validate syntax, DNS answers and address policy. */ /** Validate syntax, DNS answers and address policy. */
public URI validate(String rawUrl) { 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. */ /** Capture a bounded public web resource without persisting it. */
public FetchResult fetch(String rawUrl) { public FetchResult fetch(String rawUrl) {
long maxBodyBytes = maxBodyBytes(); long maxBodyBytes = maxBodyBytes();
long deadline = System.nanoTime() + TOTAL_TIMEOUT_NANOS; long deadline = System.nanoTime() + TOTAL_TIMEOUT_NANOS;
ValidatedTarget target = validateAndResolve(rawUrl); ValidatedTarget target = validateAndResolve(rawUrl, deadline);
Set<URI> visited = new HashSet<>(); Set<URI> visited = new HashSet<>();
visited.add(target.uri()); visited.add(target.uri());
int redirects = 0; int redirects = 0;
while (true) { while (true) {
requireTimeRemaining(deadline);
TransportResponse response; TransportResponse response;
try { try {
response = fetcher.fetch(new FetchRequest(target.uri(), target.addresses(), deadline, response = fetcher.fetch(new FetchRequest(target.uri(), target.addresses(), deadline,
@@ -110,6 +129,7 @@ public class PersonalUrlFetchService {
} catch (Exception ex) { } catch (Exception ex) {
throw new ServiceException(FETCH_FAILED); throw new ServiceException(FETCH_FAILED);
} }
requireTimeRemaining(deadline);
if (response == null || response.body() == null || response.body().length > maxBodyBytes) { if (response == null || response.body() == null || response.body().length > maxBodyBytes) {
throw new ServiceException(RESPONSE_TOO_LARGE); throw new ServiceException(RESPONSE_TOO_LARGE);
} }
@@ -124,7 +144,7 @@ public class PersonalUrlFetchService {
} catch (IllegalArgumentException ex) { } catch (IllegalArgumentException ex) {
throw new ServiceException(BLOCKED); throw new ServiceException(BLOCKED);
} }
target = validateAndResolve(next.toString()); target = validateAndResolve(next.toString(), deadline);
if (!visited.add(target.uri())) throw new ServiceException(REDIRECT_LOOP); if (!visited.add(target.uri())) throw new ServiceException(REDIRECT_LOOP);
redirects++; redirects++;
continue; continue;
@@ -141,11 +161,13 @@ public class PersonalUrlFetchService {
} }
} }
private ValidatedTarget validateAndResolve(String rawUrl) { private ValidatedTarget validateAndResolve(String rawUrl, long deadlineNanos) {
URI uri = normalizeUri(rawUrl); URI uri = normalizeUri(rawUrl);
List<InetAddress> addresses; List<InetAddress> addresses;
try { try {
addresses = resolver.resolve(canonicalHost(uri)); requireTimeRemaining(deadlineNanos);
addresses = resolver.resolve(canonicalHost(uri), deadlineNanos);
requireTimeRemaining(deadlineNanos);
} catch (Exception ex) { } catch (Exception ex) {
throw new ServiceException(BLOCKED); throw new ServiceException(BLOCKED);
} }
@@ -191,12 +213,16 @@ public class PersonalUrlFetchService {
try { try {
long value = Math.multiplyExact(properties.getMaxUrlBodyMb(), 1024L * 1024L); long value = Math.multiplyExact(properties.getMaxUrlBodyMb(), 1024L * 1024L);
if (value <= 0) throw new ArithmeticException(); if (value <= 0) throw new ArithmeticException();
return value; return Math.min(value, HARD_MAX_BODY_BYTES);
} catch (ArithmeticException ex) { } catch (ArithmeticException ex) {
throw new ServiceException(RESPONSE_TOO_LARGE); 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) { private static boolean isRedirect(int status) {
return status == 301 || status == 302 || status == 303 || status == 307 || status == 308; 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) { 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; 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 { try {
long length = Long.parseLong(raw.trim()); return Long.parseLong(value);
if (length < 0 || length > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE);
} catch (NumberFormatException ex) { } catch (NumberFormatException ex) {
throw new ServiceException(RESPONSE_INVALID); 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) { private static String firstHeader(Map<String, List<String>> headers, String name) {
if (headers == null) return null; if (headers == null) return null;
for (Map.Entry<String, List<String>> entry : headers.entrySet()) { for (Map.Entry<String, List<String>> entry : headers.entrySet()) {
@@ -280,9 +335,54 @@ public class PersonalUrlFetchService {
@FunctionalInterface @FunctionalInterface
public interface Resolver { public interface Resolver {
List<InetAddress> resolve(String host, long deadlineNanos) throws IOException;
}
@FunctionalInterface
interface HostLookup {
List<InetAddress> resolve(String host) throws UnknownHostException; 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 @FunctionalInterface
public interface Fetcher { public interface Fetcher {
TransportResponse fetch(FetchRequest request) throws IOException; TransportResponse fetch(FetchRequest request) throws IOException;
@@ -311,7 +411,29 @@ public class PersonalUrlFetchService {
private record ValidatedTarget(URI uri, List<InetAddress> addresses) { } 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 @Override
public TransportResponse fetch(FetchRequest request) throws IOException { public TransportResponse fetch(FetchRequest request) throws IOException {
IOException last = null; IOException last = null;
@@ -325,37 +447,32 @@ public class PersonalUrlFetchService {
throw last == null ? new IOException("connection failed") : last; 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(); URI uri = request.uri();
int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80); int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80);
Socket plain = new Socket(); int connectTimeout = timeout(request.deadlineNanos(), 5_000);
try { int readTimeout = timeout(request.deadlineNanos(), 5_000);
plain.connect(new InetSocketAddress(address, port), timeout(request.deadlineNanos(), 5_000)); try (Connection connection = connections.connect(uri, address, port, connectTimeout, readTimeout)) {
plain.setSoTimeout(timeout(request.deadlineNanos(), 5_000)); writeRequest(connection.output(), request);
Socket active = plain; connection.setReadTimeout(timeout(request.deadlineNanos(), 5_000));
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));
TransportResponse response = parseHttpResponse( 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"); if (System.nanoTime() >= request.deadlineNanos()) throw new IOException("deadline exceeded");
return response; 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 { private static void writeRequest(OutputStream output, FetchRequest request) throws IOException {
URI uri = request.uri(); URI uri = request.uri();
String target = uri.getRawPath().isEmpty() ? "/" : uri.getRawPath(); String target = uri.getRawPath().isEmpty() ? "/" : uri.getRawPath();
@@ -363,7 +480,7 @@ public class PersonalUrlFetchService {
String host = hostHeader(uri); String host = hostHeader(uri);
StringBuilder value = new StringBuilder("GET ").append(target).append(" HTTP/1.1\r\nHost: ") StringBuilder value = new StringBuilder("GET ").append(target).append(" HTTP/1.1\r\nHost: ")
.append(host).append("\r\n"); .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")); .append(headerValue).append("\r\n"));
value.append("Connection: close\r\n\r\n"); value.append("Connection: close\r\n\r\n");
output.write(value.toString().getBytes(StandardCharsets.US_ASCII)); 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 static final class DeadlineInputStream extends InputStream {
private final InputStream delegate; private final InputStream delegate;
private final Socket socket; private final Connection connection;
private final long deadlineNanos; private final long deadlineNanos;
private DeadlineInputStream(InputStream delegate, Socket socket, long deadlineNanos) { private DeadlineInputStream(InputStream delegate, Connection connection, long deadlineNanos) {
this.delegate = delegate; this.delegate = delegate;
this.socket = socket; this.connection = connection;
this.deadlineNanos = deadlineNanos; this.deadlineNanos = deadlineNanos;
} }
@Override @Override
public int read() throws IOException { public int read() throws IOException {
socket.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000)); connection.setReadTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
return delegate.read(); return delegate.read();
} }
@Override @Override
public int read(byte[] bytes, int offset, int length) throws IOException { 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); return delegate.read(bytes, offset, length);
} }
} }
@@ -437,18 +591,23 @@ public class PersonalUrlFetchService {
if (name.isEmpty()) throw new ServiceException(RESPONSE_INVALID); if (name.isEmpty()) throw new ServiceException(RESPONSE_INVALID);
headers.computeIfAbsent(name, ignored -> new ArrayList<>()).add(value); headers.computeIfAbsent(name, ignored -> new ArrayList<>()).add(value);
} }
String transferEncoding = firstHeader(headers, "transfer-encoding"); String transferEncoding = strictFramingHeader(headers, "transfer-encoding");
String contentLength = firstHeader(headers, "content-length"); String contentLength = strictFramingHeader(headers, "content-length");
if (transferEncoding != null && contentLength != null) throw new ServiceException(RESPONSE_INVALID); if (transferEncoding != null && contentLength != null) throw new ServiceException(RESPONSE_INVALID);
byte[] body; 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); if (!"chunked".equalsIgnoreCase(transferEncoding.trim())) throw new ServiceException(RESPONSE_INVALID);
body = readChunked(buffered, maxBodyBytes); body = readChunked(buffered, maxBodyBytes);
} else if (contentLength != null) { } else if (contentLength != null) {
long length; long length = parseContentLength(contentLength);
try { length = Long.parseLong(contentLength.trim()); } if (length > maxBodyBytes || length > Integer.MAX_VALUE) {
catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); }
if (length < 0 || length > maxBodyBytes || length > Integer.MAX_VALUE) {
throw new ServiceException(RESPONSE_TOO_LARGE); throw new ServiceException(RESPONSE_TOO_LARGE);
} }
body = readExactly(buffered, (int) length); 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 { private static byte[] readChunked(BufferedInputStream input, long maxBodyBytes) throws IOException {
ByteArrayOutputStream body = new ByteArrayOutputStream(); ByteArrayOutputStream body = new ByteArrayOutputStream();
int[] framingBytes = {0}; 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.Tag;
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Test;
import javax.net.ssl.SNIHostName;
import javax.net.ssl.SSLParameters;
import java.io.ByteArrayInputStream; import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.net.InetAddress; import java.net.InetAddress;
import java.net.URI; import java.net.URI;
import java.nio.charset.StandardCharsets; import java.nio.charset.StandardCharsets;
@@ -15,6 +20,11 @@ import java.util.ArrayList;
import java.util.Arrays; import java.util.Arrays;
import java.util.List; import java.util.List;
import java.util.Map; 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.*; import static org.junit.jupiter.api.Assertions.*;
@@ -25,7 +35,7 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void rejectsUnsafeSchemesSyntaxAndHosts() { 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( for (String raw : List.of(
"file:///etc/passwd", "ftp://example.com/a", "data:text/plain,hello", "file:///etc/passwd", "ftp://example.com/a", "data:text/plain,hello",
"http://user:secret@example.com", "http:///missing", "not a url", "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", "::", "::1", "fe80::1", "fc00::1", "fd00::1", "ff02::1",
"2001:db8::1", "2001:2::1", "2001:100::1", "2002:5db8:d822::1", "2001:db8::1", "2001:2::1", "2001:100::1", "2002:5db8:d822::1",
"3fff::1", "64:ff9b::c0a8:101")) { "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); assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://example.com"), ip);
} }
} }
@@ -52,17 +62,17 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void rejectsEmptyOrMixedDnsAnswersAndAcceptsPublicResolution() { void rejectsEmptyOrMixedDnsAnswersAndAcceptsPublicResolution() {
assertCode("PERSONAL_URL_BLOCKED", 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", 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"); .validate("https://example.com"), "mixed");
assertCode("PERSONAL_URL_BLOCKED", 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"); .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()); .validate("HTTPS://Example.COM/a").toString());
assertEquals("http://[2606:2800:220:1:248:1893:25c8:1946]/", fixture( 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()); .validate("http://[2606:2800:220:1:248:1893:25c8:1946]/").toString());
} }
@@ -72,7 +82,7 @@ class PersonalUrlFetchServiceTest {
var responses = new ArrayDeque<PersonalUrlFetchService.TransportResponse>(); var responses = new ArrayDeque<PersonalUrlFetchService.TransportResponse>();
responses.add(response(302, Map.of("location", List.of("/final")), new byte[0])); responses.add(response(302, Map.of("location", List.of("/final")), new byte[0]));
responses.add(ok("text/plain; charset=utf-8", "done")); 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"); var result = service.fetch("https://example.com/start");
@@ -87,11 +97,11 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void blocksUnsafeRedirectAndMixedAddressRedirect() { void blocksUnsafeRedirectAndMixedAddressRedirect() {
var redirect = response(302, Map.of("location", List.of("http://metadata.test/latest")), new byte[0]); 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); ? List.of(address("169.254.169.254")) : List.of(PUBLIC), request -> redirect);
assertCode("PERSONAL_URL_BLOCKED", () -> service.fetch("https://example.com/start"), "redirect private"); 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 -> ? 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])); 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"); assertCode("PERSONAL_URL_BLOCKED", () -> mixed.fetch("https://example.com"), "redirect mixed");
@@ -99,11 +109,11 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void detectsRedirectLoopAndMoreThanThreeRedirects() { 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])); 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"); 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)); int n = Integer.parseInt(request.uri().getPath().substring(1));
return response(302, Map.of("location", List.of("/" + (n + 1))), new byte[0]); return response(302, Map.of("location", List.of("/" + (n + 1))), new byte[0]);
}); });
@@ -113,7 +123,7 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void sendsOnlyFixedSafeHeaders() { void sendsOnlyFixedSafeHeaders() {
var requests = new ArrayList<PersonalUrlFetchService.FetchRequest>(); 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"); service.fetch("https://example.com/a");
Map<String, String> headers = requests.get(0).headers(); Map<String, String> headers = requests.get(0).headers();
@@ -128,20 +138,29 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void rejectsForbiddenOrMissingMimeAndOversizedBody() { void rejectsForbiddenOrMissingMimeAndOversizedBody() {
assertCode("PERSONAL_URL_CONTENT_TYPE_UNSUPPORTED", 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", 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(); PersonalKnowledgeProperties properties = properties();
byte[] tooLarge = new byte[10 * 1024 * 1024 + 1]; byte[] tooLarge = new byte[10 * 1024 * 1024 + 1];
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () -> 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"); 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 @Test
void returnsDigestAndCaptureMetadata() { 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"); .fetch("https://example.com/a.pdf");
assertEquals(200, result.status()); assertEquals(200, result.status());
assertEquals("application/pdf", result.contentType()); assertEquals("application/pdf", result.contentType());
@@ -176,13 +195,127 @@ class PersonalUrlFetchServiceTest {
@Test @Test
void fetchRequestCarriesValidatedIpsAndOriginalTlsHost() { void fetchRequestCarriesValidatedIpsAndOriginalTlsHost() {
var requests = new ArrayList<PersonalUrlFetchService.FetchRequest>(); 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"); .fetch("https://example.com/path");
assertEquals("example.com", requests.get(0).uri().getHost()); assertEquals("example.com", requests.get(0).uri().getHost());
assertEquals(List.of(PUBLIC), requests.get(0).addresses()); assertEquals(List.of(PUBLIC), requests.get(0).addresses());
assertTrue(requests.get(0).deadlineNanos() > System.nanoTime()); 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, private static PersonalUrlFetchService fixture(PersonalUrlFetchService.Resolver resolver,
PersonalUrlFetchService.Fetcher fetcher) { PersonalUrlFetchService.Fetcher fetcher) {
return PersonalUrlFetchService.forTest(properties(), resolver, fetcher); return PersonalUrlFetchService.forTest(properties(), resolver, fetcher);
@@ -218,4 +351,29 @@ class PersonalUrlFetchServiceTest {
private static void assertCode(String code, Runnable action, String context) { private static void assertCode(String code, Runnable action, String context) {
assertEquals(code, assertThrows(ServiceException.class, action::run, context).getMessage(), 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();
}
}
} }