From 2c9c6d5d4a8f76f502ec3368d25e2eebc71cab4f Mon Sep 17 00:00:00 2001 From: let5sne Date: Sun, 12 Jul 2026 04:36:47 +0800 Subject: [PATCH] fix(personal): enforce URL fetch deadlines and framing --- .../service/PersonalUrlFetchService.java | 257 ++++++++++++++---- .../personal/PersonalUrlFetchServiceTest.java | 194 +++++++++++-- 2 files changed, 386 insertions(+), 65 deletions(-) diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java index 834f398e..1c81ce7e 100644 --- a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java +++ b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java @@ -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 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 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 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> 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> headers, String name) { + if (headers == null) return null; + String found = null; + for (Map.Entry> 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> headers, String name) { if (headers == null) return null; for (Map.Entry> entry : headers.entrySet()) { @@ -280,9 +335,54 @@ public class PersonalUrlFetchService { @FunctionalInterface public interface Resolver { + List resolve(String host, long deadlineNanos) throws IOException; + } + + @FunctionalInterface + interface HostLookup { List 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 resolve(String host, long deadlineNanos) throws IOException { + long remaining = deadlineNanos - System.nanoTime(); + if (remaining <= 0) throw new IOException("resolution deadline exceeded"); + Future> 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 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}; diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java index ff5f7a4a..76b63309 100644 --- a/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java +++ b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java @@ -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(); 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(); - 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 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(); - 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 connected = new AtomicReference<>(); + AtomicReference 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(); + } + } }