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 1c81ce7e..2f78fb84 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 @@ -4,6 +4,13 @@ import org.dromara.aihr.personal.support.PersonalKnowledgeProperties; import org.dromara.common.core.exception.ServiceException; import org.springframework.stereotype.Service; +import javax.naming.Context; +import javax.naming.NamingEnumeration; +import javax.naming.NamingException; +import javax.naming.directory.Attribute; +import javax.naming.directory.Attributes; +import javax.naming.directory.DirContext; +import javax.naming.directory.InitialDirContext; import javax.net.ssl.SNIHostName; import javax.net.ssl.SSLParameters; import javax.net.ssl.SSLSocket; @@ -21,13 +28,12 @@ import java.net.InetSocketAddress; import java.net.Socket; import java.net.URI; import java.net.URISyntaxException; -import java.net.UnknownHostException; import java.nio.charset.StandardCharsets; import java.security.MessageDigest; import java.security.NoSuchAlgorithmException; import java.time.Instant; import java.util.ArrayList; -import java.util.Arrays; +import java.util.Hashtable; import java.util.HashSet; import java.util.HexFormat; import java.util.LinkedHashMap; @@ -35,15 +41,7 @@ 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 { @@ -61,13 +59,6 @@ public class PersonalUrlFetchService { 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( @@ -89,8 +80,7 @@ public class PersonalUrlFetchService { private final Fetcher fetcher; public PersonalUrlFetchService(PersonalKnowledgeProperties properties) { - this(properties, new DeadlineDnsResolver(DNS_EXECUTOR, - host -> Arrays.asList(InetAddress.getAllByName(host))), new RawSocketFetcher()); + this(properties, new DeadlineDnsResolver(new JndiDnsQuery()), new RawSocketFetcher()); } private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) { @@ -136,7 +126,7 @@ public class PersonalUrlFetchService { enforceDeclaredLength(response.headers(), maxBodyBytes); if (isRedirect(response.status())) { if (redirects >= MAX_REDIRECTS) throw new ServiceException(REDIRECT_LIMIT); - String location = firstHeader(response.headers(), "location"); + String location = strictSingletonHeader(response.headers(), "location"); if (location == null || location.isBlank()) throw new ServiceException(RESPONSE_INVALID); URI next; try { @@ -152,7 +142,7 @@ public class PersonalUrlFetchService { if (response.status() < 200 || response.status() >= 300) { throw new ServiceException(RESPONSE_INVALID); } - String contentType = normalizeContentType(firstHeader(response.headers(), "content-type")); + String contentType = normalizeContentType(strictSingletonHeader(response.headers(), "content-type")); if (!ALLOWED_CONTENT_TYPES.contains(contentType)) { throw new ServiceException(CONTENT_TYPE_UNSUPPORTED); } @@ -184,7 +174,7 @@ public class PersonalUrlFetchService { throw new ServiceException(BLOCKED); } try { - URI parsed = new URI(rawUrl.trim()).normalize(); + URI parsed = new URI(rawUrl.trim()); String scheme = parsed.getScheme() == null ? "" : parsed.getScheme().toLowerCase(Locale.ROOT); if (!("http".equals(scheme) || "https".equals(scheme)) || parsed.getRawUserInfo() != null || parsed.getHost() == null || parsed.getHost().isBlank()) { @@ -199,8 +189,12 @@ public class PersonalUrlFetchService { : IDN.toASCII(rawHost, IDN.USE_STD3_ASCII_RULES).toLowerCase(Locale.ROOT); if (host.isBlank() || host.length() > 253) throw new ServiceException(BLOCKED); if ("localhost".equals(host) || host.endsWith(".localhost")) throw new ServiceException(BLOCKED); - URI normalized = new URI(scheme, null, host, port, - parsed.getRawPath().isEmpty() ? "/" : parsed.getRawPath(), parsed.getRawQuery(), null).normalize(); + String authority = host.indexOf(':') >= 0 ? "[" + host + "]" : host; + if (port >= 0) authority += ":" + port; + String rawPath = parsed.getRawPath().isEmpty() ? "/" : parsed.getRawPath(); + StringBuilder rebuilt = new StringBuilder(scheme).append("://").append(authority).append(rawPath); + if (parsed.getRawQuery() != null) rebuilt.append('?').append(parsed.getRawQuery()); + URI normalized = new URI(new URI(rebuilt.toString()).normalize().toASCIIString()); if (normalized.toASCIIString().length() > MAX_URL_LENGTH) throw new ServiceException(BLOCKED); return normalized; } catch (URISyntaxException | IllegalArgumentException ex) { @@ -263,6 +257,12 @@ public class PersonalUrlFetchService { } private static String strictFramingHeader(Map> headers, String name) { + String found = strictSingletonHeader(headers, name); + if (found != null && found.indexOf(',') >= 0) throw new ServiceException(RESPONSE_INVALID); + return found; + } + + private static String strictSingletonHeader(Map> headers, String name) { if (headers == null) return null; String found = null; for (Map.Entry> entry : headers.entrySet()) { @@ -271,22 +271,13 @@ public class PersonalUrlFetchService { throw new ServiceException(RESPONSE_INVALID); } found = entry.getValue().get(0); - if (found == null || found.isBlank() || found.indexOf(',') >= 0) { + if (found == null || found.isBlank()) { 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()) { - if (entry.getKey() != null && entry.getKey().equalsIgnoreCase(name) - && entry.getValue() != null && !entry.getValue().isEmpty()) return entry.getValue().get(0); - } - return null; - } - static boolean isGloballyRoutable(InetAddress address) { if (address.isAnyLocalAddress() || address.isLoopbackAddress() || address.isLinkLocalAddress() || address.isSiteLocalAddress() || address.isMulticastAddress()) return false; @@ -307,7 +298,7 @@ public class PersonalUrlFetchService { if (a == 100 && b >= 64 && b <= 127) return false; if (a == 169 && b == 254) return false; if (a == 172 && b >= 16 && b <= 31) return false; - if (a == 192 && (b == 168 || (b == 0 && (c == 0 || c == 2)))) return false; + if (a == 192 && (b == 168 || b == 0)) return false; if (a == 192 && ((b == 31 && c == 196) || (b == 52 && c == 193) || (b == 88 && c == 99) || (b == 175 && c == 48))) return false; if (a == 198 && (b == 18 || b == 19 || (b == 51 && c == 100))) return false; @@ -339,48 +330,93 @@ public class PersonalUrlFetchService { } @FunctionalInterface - interface HostLookup { - List resolve(String host) throws UnknownHostException; + interface DnsQuery { + List resolve(String host, int timeoutMillis, int retries) throws NamingException; } static final class DeadlineDnsResolver implements Resolver { - private final ExecutorService executor; - private final HostLookup lookup; + private final DnsQuery query; - DeadlineDnsResolver(ExecutorService executor, HostLookup lookup) { - this.executor = executor; - this.lookup = lookup; + DeadlineDnsResolver(DnsQuery query) { + this.query = query; } @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; + long remainingMillis = TimeUnit.NANOSECONDS.toMillis(remaining); + if (remainingMillis <= 0) throw new IOException("resolution deadline exceeded"); + int timeoutMillis = (int) Math.min(5_000L, remainingMillis); + List literals; 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); + literals = query.resolve(host, timeoutMillis, 0); + } catch (NamingException | RuntimeException ex) { throw new IOException("resolution failed"); } + if (System.nanoTime() >= deadlineNanos) throw new IOException("resolution deadline exceeded"); + List addresses = new ArrayList<>(); + if (literals != null) { + for (String literal : literals) addresses.add(numericAddress(literal)); + } + return List.copyOf(addresses); + } + } + + static final class JndiDnsQuery implements DnsQuery { + @Override + public List resolve(String host, int timeoutMillis, int retries) throws NamingException { + Hashtable environment = environment(timeoutMillis); + environment.put("com.sun.jndi.dns.timeout.retries", Integer.toString(Math.max(0, retries))); + DirContext context = new InitialDirContext(environment); + try { + Attributes attributes = context.getAttributes(host, new String[] {"A", "AAAA"}); + List values = new ArrayList<>(); + collect(attributes.get("A"), values); + collect(attributes.get("AAAA"), values); + return values; + } finally { + context.close(); + } } - private void cancelAndPurge(Future future) { - future.cancel(true); - if (executor instanceof ThreadPoolExecutor pool) pool.purge(); + static Hashtable environment(int timeoutMillis) { + Hashtable environment = new Hashtable<>(); + environment.put(Context.INITIAL_CONTEXT_FACTORY, "com.sun.jndi.dns.DnsContextFactory"); + environment.put("com.sun.jndi.dns.timeout.initial", Integer.toString(Math.max(1, timeoutMillis))); + environment.put("com.sun.jndi.dns.timeout.retries", "0"); + return environment; } + + private static void collect(Attribute attribute, List values) throws NamingException { + if (attribute == null) return; + NamingEnumeration all = attribute.getAll(); + while (all.hasMore()) values.add(String.valueOf(all.next()).trim()); + } + } + + private static InetAddress numericAddress(String literal) throws IOException { + if (literal == null || literal.isBlank() || literal.indexOf('%') >= 0) throw new IOException("invalid DNS answer"); + String value = literal.trim(); + if (value.indexOf(':') < 0) { + String[] parts = value.split("\\.", -1); + if (parts.length != 4) throw new IOException("invalid DNS answer"); + byte[] bytes = new byte[4]; + for (int i = 0; i < parts.length; i++) { + if (parts[i].isEmpty() || !parts[i].chars().allMatch(Character::isDigit)) { + throw new IOException("invalid DNS answer"); + } + int octet; + try { octet = Integer.parseInt(parts[i]); } + catch (NumberFormatException ex) { throw new IOException("invalid DNS answer"); } + if (octet > 255) throw new IOException("invalid DNS answer"); + bytes[i] = (byte) octet; + } + return InetAddress.getByAddress(bytes); + } + if (!value.matches("[0-9A-Fa-f:.]+")) throw new IOException("invalid DNS answer"); + InetAddress address = InetAddress.getByName(value); + if (!(address instanceof Inet6Address)) throw new IOException("invalid DNS answer"); + return address; } @FunctionalInterface @@ -420,7 +456,7 @@ public class PersonalUrlFetchService { @FunctionalInterface interface ConnectionFactory { - Connection connect(URI uri, InetAddress address, int port, int connectTimeout, int readTimeout) throws IOException; + Connection connect(URI uri, InetAddress address, int port, long deadlineNanos) throws IOException; } static final class RawSocketFetcher implements Fetcher { @@ -450,9 +486,7 @@ public class PersonalUrlFetchService { 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); - int connectTimeout = timeout(request.deadlineNanos(), 5_000); - int readTimeout = timeout(request.deadlineNanos(), 5_000); - try (Connection connection = connections.connect(uri, address, port, connectTimeout, readTimeout)) { + try (Connection connection = connections.connect(uri, address, port, request.deadlineNanos())) { writeRequest(connection.output(), request); connection.setReadTimeout(timeout(request.deadlineNanos(), 5_000)); TransportResponse response = parseHttpResponse( @@ -505,25 +539,24 @@ public class PersonalUrlFetchService { } } - private static final class JvmConnectionFactory implements ConnectionFactory { + static final class JvmConnectionFactory implements ConnectionFactory { @Override - public Connection connect(URI uri, InetAddress address, int port, - int connectTimeout, int readTimeout) throws IOException { + public Connection connect(URI uri, InetAddress address, int port, long deadlineNanos) 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); + plain.connect(new InetSocketAddress(address, port), RawSocketFetcher.timeout(deadlineNanos, 5_000)); + plain.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000)); 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()) + SSLSocket ssl = (SSLSocket) defaultSslSocketFactory() .createSocket(plain, tlsHost, port, true); SSLParameters parameters = ssl.getSSLParameters(); RawSocketFetcher.configureTlsParameters(parameters, tlsHost); ssl.setSSLParameters(parameters); - ssl.setSoTimeout(readTimeout); + ssl.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000)); ssl.startHandshake(); active = ssl; } @@ -533,6 +566,10 @@ public class PersonalUrlFetchService { throw ex; } } + + static SSLSocketFactory defaultSslSocketFactory() { + return (SSLSocketFactory) SSLSocketFactory.getDefault(); + } } private record SocketConnection(Socket socket) implements Connection { @@ -570,51 +607,17 @@ public class PersonalUrlFetchService { if (maxBodyBytes < 0) throw new ServiceException(RESPONSE_TOO_LARGE); try { BufferedInputStream buffered = input instanceof BufferedInputStream value ? value : new BufferedInputStream(input); - int[] headerBytes = {0}; - String statusLine = readLine(buffered, headerBytes); - if (statusLine == null || !statusLine.startsWith("HTTP/1.")) throw new ServiceException(RESPONSE_INVALID); - String[] statusParts = statusLine.split(" ", 3); - if (statusParts.length < 2) throw new ServiceException(RESPONSE_INVALID); - int status; - try { status = Integer.parseInt(statusParts[1]); } - catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); } - - Map> headers = new LinkedHashMap<>(); + int interimCount = 0; while (true) { - String line = readLine(buffered, headerBytes); - if (line == null) throw new ServiceException(RESPONSE_INVALID); - if (line.isEmpty()) break; - int colon = line.indexOf(':'); - if (colon <= 0) throw new ServiceException(RESPONSE_INVALID); - String name = line.substring(0, colon).trim().toLowerCase(Locale.ROOT); - String value = line.substring(colon + 1).trim(); - if (name.isEmpty()) throw new ServiceException(RESPONSE_INVALID); - headers.computeIfAbsent(name, ignored -> new ArrayList<>()).add(value); - } - 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 (hasNoBody(status)) { - if (transferEncoding != null) throw new ServiceException(RESPONSE_INVALID); - if (status != 304 && contentLength != null) { - throw new ServiceException(RESPONSE_INVALID); + TransportResponse response = parseOneHttpResponse(buffered, maxBodyBytes); + if (response.status() == 101) throw new ServiceException(RESPONSE_INVALID); + if (response.status() == 100 || response.status() == 102 || response.status() == 103) { + if (++interimCount > 3) throw new ServiceException(RESPONSE_INVALID); + continue; } - 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 = parseContentLength(contentLength); - if (length > maxBodyBytes || length > Integer.MAX_VALUE) { - throw new ServiceException(RESPONSE_TOO_LARGE); - } - body = readExactly(buffered, (int) length); - } else { - body = readUntilEof(buffered, maxBodyBytes); + if (response.status() >= 100 && response.status() < 200) throw new ServiceException(RESPONSE_INVALID); + return response; } - return new TransportResponse(status, headers, body); } catch (ServiceException ex) { throw ex; } catch (IOException ex) { @@ -622,6 +625,63 @@ public class PersonalUrlFetchService { } } + private static TransportResponse parseOneHttpResponse(BufferedInputStream buffered, long maxBodyBytes) throws IOException { + int[] headerBytes = {0}; + String statusLine = readLine(buffered, headerBytes); + if (statusLine == null || !(statusLine.startsWith("HTTP/1.0 ") || statusLine.startsWith("HTTP/1.1 "))) { + throw new ServiceException(RESPONSE_INVALID); + } + String[] statusParts = statusLine.split(" ", 3); + if (statusParts.length < 2) throw new ServiceException(RESPONSE_INVALID); + int status; + try { status = Integer.parseInt(statusParts[1]); } + catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); } + + Map> headers = new LinkedHashMap<>(); + while (true) { + String line = readLine(buffered, headerBytes); + if (line == null) throw new ServiceException(RESPONSE_INVALID); + if (line.isEmpty()) break; + int colon = line.indexOf(':'); + if (colon <= 0) throw new ServiceException(RESPONSE_INVALID); + String rawName = line.substring(0, colon); + if (!validHeaderName(rawName)) throw new ServiceException(RESPONSE_INVALID); + String name = rawName.toLowerCase(Locale.ROOT); + String value = line.substring(colon + 1).trim(); + headers.computeIfAbsent(name, ignored -> new ArrayList<>()).add(value); + } + 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 (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 = parseContentLength(contentLength); + if (length > maxBodyBytes || length > Integer.MAX_VALUE) throw new ServiceException(RESPONSE_TOO_LARGE); + body = readExactly(buffered, (int) length); + } else { + body = readUntilEof(buffered, maxBodyBytes); + } + return new TransportResponse(status, headers, body); + } + + private static boolean validHeaderName(String name) { + if (name.isEmpty()) return false; + for (int i = 0; i < name.length(); i++) { + char ch = name.charAt(i); + boolean token = Character.isLetterOrDigit(ch) || "!#$%&'*+-.^_`|~".indexOf(ch) >= 0; + if (!token || ch > 127) return false; + } + return true; + } + private static boolean hasNoBody(int status) { return status >= 100 && status < 200 || status == 204 || status == 304; } @@ -643,7 +703,10 @@ public class PersonalUrlFetchService { String trailer = readLine(input, framingBytes); if (trailer == null) throw new ServiceException(RESPONSE_INVALID); if (trailer.isEmpty()) return body.toByteArray(); - if (trailer.indexOf(':') <= 0) throw new ServiceException(RESPONSE_INVALID); + int colon = trailer.indexOf(':'); + if (colon <= 0 || !validHeaderName(trailer.substring(0, colon))) { + throw new ServiceException(RESPONSE_INVALID); + } } } if ((long) body.size() + size > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE); 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 76b63309..8e896e5f 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 @@ -7,21 +7,23 @@ import org.junit.jupiter.api.Test; import javax.net.ssl.SNIHostName; import javax.net.ssl.SSLParameters; +import javax.net.ssl.SSLSocketFactory; import java.io.ByteArrayInputStream; import java.io.ByteArrayOutputStream; import java.io.IOException; import java.io.InputStream; import java.net.InetAddress; +import java.net.ServerSocket; +import java.net.Socket; import java.net.URI; import java.nio.charset.StandardCharsets; import java.time.Instant; import java.util.ArrayDeque; import java.util.ArrayList; import java.util.Arrays; +import java.util.Hashtable; 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; @@ -40,6 +42,7 @@ class PersonalUrlFetchServiceTest { "file:///etc/passwd", "ftp://example.com/a", "data:text/plain,hello", "http://user:secret@example.com", "http:///missing", "not a url", "http://localhost/admin", "http://service.localhost/admin", + "http://[fe80::1%25en0]/admin", "http://example.com/" + "x".repeat(5000))) { assertCode("PERSONAL_URL_BLOCKED", () -> service.validate(raw), raw); } @@ -59,6 +62,15 @@ class PersonalUrlFetchServiceTest { } } + @Test + void rejectsEntireIanaSpecial192Dot0Dot0Slash24() { + for (int last : List.of(0, 8, 9, 10, 170, 171, 255)) { + String ip = "192.0.0." + last; + var service = fixture((host, deadline) -> List.of(address(ip)), request -> ok("text/plain", "ok")); + assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://example.com"), ip); + } + } + @Test void rejectsEmptyOrMixedDnsAnswersAndAcceptsPublicResolution() { assertCode("PERSONAL_URL_BLOCKED", @@ -204,21 +216,39 @@ class PersonalUrlFetchServiceTest { @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(); - } + AtomicInteger timeoutSeen = new AtomicInteger(); + AtomicInteger retriesSeen = new AtomicInteger(-1); + var resolver = new PersonalUrlFetchService.DeadlineDnsResolver((host, timeoutMillis, retries) -> { + timeoutSeen.set(timeoutMillis); retriesSeen.set(retries); + return List.of("93.184.216.34", "2606:2800:220:1:248:1893:25c8:1946"); + }); + long deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(200); + assertEquals(2, assertDoesNotThrow(() -> resolver.resolve("example.com", deadline)).size()); + assertTrue(timeoutSeen.get() > 0 && timeoutSeen.get() <= 200); + assertEquals(0, retriesSeen.get()); + assertThrows(IOException.class, () -> resolver.resolve("example.com", System.nanoTime() - 1)); + var nonNumeric = new PersonalUrlFetchService.DeadlineDnsResolver( + (host, timeoutMillis, retries) -> List.of("internal.example", "fe80::1%en0")); + assertThrows(IOException.class, () -> nonNumeric.resolve("example.com", + System.nanoTime() + TimeUnit.SECONDS.toNanos(1))); + + Hashtable environment = PersonalUrlFetchService.JndiDnsQuery.environment(123); + assertEquals("123", environment.get("com.sun.jndi.dns.timeout.initial")); + assertEquals("0", environment.get("com.sun.jndi.dns.timeout.retries")); + } + + @Test + void preservesRawPathQueryAndUnicodeEncodingAcrossNormalizationAndRedirects() { + var seen = new ArrayList(); + var responses = new ArrayDeque(); + responses.add(response(302, Map.of("location", List.of("../%E4%B8%AD%2Fnext?sig=a%252Fb%2Fz")), new byte[0])); + responses.add(ok("text/plain", "ok")); + var service = fixture((host, deadline) -> List.of(PUBLIC), request -> { seen.add(request); return responses.remove(); }); + service.fetch("https://example.com/a/%2Fkeep?x=%25&u=%E4%B8%AD#fragment"); + assertEquals("/a/%2Fkeep", seen.get(0).uri().getRawPath()); + assertEquals("x=%25&u=%E4%B8%AD", seen.get(0).uri().getRawQuery()); + assertEquals("/%E4%B8%AD%2Fnext", seen.get(1).uri().getRawPath()); + assertEquals("sig=a%252Fb%2Fz", seen.get(1).uri().getRawQuery()); } @Test @@ -240,7 +270,7 @@ class PersonalUrlFetchServiceTest { @Test void noBodyStatusesDoNotWaitForPayload() { - for (int status : List.of(100, 204, 304)) { + for (int status : List.of(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); @@ -251,6 +281,31 @@ class PersonalUrlFetchServiceTest { } } + @Test + void consumesLimitedInterimResponsesAndRejectsSwitchingProtocols() { + String finalResponse = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nok"; + String interim = "HTTP/1.1 100 Continue\r\n\r\nHTTP/1.1 103 Early Hints\r\nLink: \r\n\r\n" + finalResponse; + assertEquals("ok", new String(PersonalUrlFetchService.parseHttpResponse(stream(interim), 100).body(), StandardCharsets.US_ASCII)); + assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(stream( + "HTTP/1.1 101 Switching Protocols\r\n\r\n"), 100), "101"); + String tooMany = "HTTP/1.1 100 Continue\r\n\r\n".repeat(4) + finalResponse; + assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(stream(tooMany), 100), "interim limit"); + } + + @Test + void rejectsInvalidHeaderNamesObsFoldAndDuplicateSemanticHeaders() { + for (String line : List.of("Bad Header: x", "Content-Type : text/plain", "\tcontinued")) { + assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> PersonalUrlFetchService.parseHttpResponse(stream( + "HTTP/1.1 200 OK\r\n" + line + "\r\n\r\n"), 100), line); + } + assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> fixture((host, deadline) -> List.of(PUBLIC), request -> + response(200, Map.of("content-type", List.of("text/plain", "text/html")), new byte[0])) + .fetch("https://example.com"), "duplicate content type"); + assertCode("PERSONAL_URL_RESPONSE_INVALID", () -> fixture((host, deadline) -> List.of(PUBLIC), request -> + response(302, Map.of("location", List.of("/a", "/b")), new byte[0])) + .fetch("https://example.com"), "duplicate location"); + } + @Test void rejectsOversizedHeaderBlockAndLine() { String longLine = "X-Large: " + "x".repeat(8 * 1024 + 1); @@ -272,9 +327,10 @@ class PersonalUrlFetchServiceTest { 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) -> { + PersonalUrlFetchService.ConnectionFactory factory = (uri, address, port, deadlineNanos) -> { connected.set(address); host.set(uri.getHost()); - connectTimeoutSeen.set(connectTimeout); readTimeoutSeen.set(readTimeout); + int remaining = (int) TimeUnit.NANOSECONDS.toMillis(deadlineNanos - System.nanoTime()); + connectTimeoutSeen.set(remaining); readTimeoutSeen.set(remaining); return new FakeConnection(new ByteArrayInputStream(rawResponse), requestBytes); }; var fetcher = new PersonalUrlFetchService.RawSocketFetcher(factory); @@ -302,13 +358,31 @@ class PersonalUrlFetchServiceTest { SSLParameters parameters = PersonalUrlFetchService.RawSocketFetcher.tlsParameters("origin.example"); assertEquals("HTTPS", parameters.getEndpointIdentificationAlgorithm()); assertEquals("origin.example", ((SNIHostName) parameters.getServerNames().get(0)).getAsciiName()); + assertEquals(SSLSocketFactory.getDefault().getClass(), + PersonalUrlFetchService.JvmConnectionFactory.defaultSslSocketFactory().getClass()); + } + + @Test + void rawTransportWritesBracketedIpv6Host() throws Exception { + 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, deadlineNanos) -> + new FakeConnection(new ByteArrayInputStream(rawResponse), requestBytes); + var fetcher = new PersonalUrlFetchService.RawSocketFetcher(factory); + URI uri = URI.create("http://[2606:2800:220:1:248:1893:25c8:1946]:8080/a"); + fetcher.fetch(new PersonalUrlFetchService.FetchRequest(uri, + List.of(address("2606:2800:220:1:248:1893:25c8:1946")), + System.nanoTime() + TimeUnit.SECONDS.toNanos(1), 100, Map.of())); + assertTrue(requestBytes.toString(StandardCharsets.US_ASCII) + .contains("Host: [2606:2800:220:1:248:1893:25c8:1946]:8080\r\n")); } @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) -> + PersonalUrlFetchService.ConnectionFactory factory = (uri, address, port, deadlineNanos) -> 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), @@ -316,6 +390,39 @@ class PersonalUrlFetchServiceTest { assertThrows(IOException.class, () -> fetcher.fetch(request)); } + @Test + void rawTransportUsesOriginFormOverRealLoopbackSocket() throws Exception { + InetAddress loopback = InetAddress.getLoopbackAddress(); + try (ServerSocket server = new ServerSocket(0, 1, loopback)) { + AtomicReference wire = new AtomicReference<>(); + Thread peer = new Thread(() -> { + try (Socket socket = server.accept()) { + socket.setSoTimeout(2_000); + ByteArrayOutputStream bytes = new ByteArrayOutputStream(); + int value; + while ((value = socket.getInputStream().read()) >= 0) { + bytes.write(value); + byte[] data = bytes.toByteArray(); + int size = data.length; + if (size >= 4 && data[size - 4] == '\r' && data[size - 3] == '\n' + && data[size - 2] == '\r' && data[size - 1] == '\n') break; + } + wire.set(bytes.toString(StandardCharsets.US_ASCII)); + socket.getOutputStream().write("HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nok" + .getBytes(StandardCharsets.US_ASCII)); + } catch (IOException ex) { throw new AssertionError(ex); } + }, "url-fetch-loopback-peer"); + peer.start(); + var fetcher = new PersonalUrlFetchService.RawSocketFetcher(); + URI uri = URI.create("http://public.example:" + server.getLocalPort() + "/raw/%2F?a=%25"); + var response = fetcher.fetch(new PersonalUrlFetchService.FetchRequest(uri, List.of(loopback), + System.nanoTime() + TimeUnit.SECONDS.toNanos(2), 100, Map.of())); + peer.join(2_000); + assertEquals("ok", new String(response.body(), StandardCharsets.US_ASCII)); + assertTrue(wire.get().startsWith("GET /raw/%2F?a=%25 HTTP/1.1\r\nHost: public.example:" + server.getLocalPort())); + } + } + private static PersonalUrlFetchService fixture(PersonalUrlFetchService.Resolver resolver, PersonalUrlFetchService.Fetcher fetcher) { return PersonalUrlFetchService.forTest(properties(), resolver, fetcher);