fix(personal): harden URL resolution and normalization

This commit is contained in:
2026-07-12 04:51:05 +08:00
parent 2c9c6d5d4a
commit 86e87be6f5
2 changed files with 307 additions and 137 deletions
@@ -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<String, String> 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<String, List<String>> 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<String, List<String>> headers, String name) {
if (headers == null) return null;
String found = null;
for (Map.Entry<String, List<String>> 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<String, List<String>> headers, String name) {
if (headers == null) return null;
for (Map.Entry<String, List<String>> 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<InetAddress> resolve(String host) throws UnknownHostException;
interface DnsQuery {
List<String> 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<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;
long remainingMillis = TimeUnit.NANOSECONDS.toMillis(remaining);
if (remainingMillis <= 0) throw new IOException("resolution deadline exceeded");
int timeoutMillis = (int) Math.min(5_000L, remainingMillis);
List<String> 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<InetAddress> 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<String> resolve(String host, int timeoutMillis, int retries) throws NamingException {
Hashtable<String, String> 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<String> 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<String, String> environment(int timeoutMillis) {
Hashtable<String, String> 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<String> 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<String, List<String>> 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<String, List<String>> 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);
@@ -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<String, String> 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<PersonalUrlFetchService.FetchRequest>();
var responses = new ArrayDeque<PersonalUrlFetchService.TransportResponse>();
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: </a>\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<String> 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);