fix(personal): bound DNS fallback and smoke redirects

This commit is contained in:
2026-07-12 17:00:20 +08:00
parent fed32939e6
commit c8a0de82b7
5 changed files with 333 additions and 254 deletions
@@ -5,13 +5,6 @@ import org.dromara.common.core.exception.ServiceException;
import org.springframework.beans.factory.annotation.Autowired;
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;
@@ -22,20 +15,23 @@ import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.IDN;
import java.net.DatagramPacket;
import java.net.DatagramSocket;
import java.net.Inet4Address;
import java.net.Inet6Address;
import java.net.InetAddress;
import java.net.InetSocketAddress;
import java.net.Socket;
import java.net.SocketTimeoutException;
import java.net.URI;
import java.net.URISyntaxException;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
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;
@@ -51,7 +47,9 @@ import java.util.concurrent.RejectedExecutionException;
import java.util.concurrent.ThreadPoolExecutor;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException;
import java.util.concurrent.ThreadLocalRandom;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.function.IntSupplier;
@Service
public class PersonalUrlFetchService {
@@ -69,7 +67,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 ExecutorService DNS_EXECUTOR = boundedExecutor("personal-url-dns", 2, 8);
private static final ExecutorService WRITE_EXECUTOR = boundedExecutor("personal-url-write", 2, 8);
private static final ExecutorService HANDSHAKE_EXECUTOR = boundedExecutor("personal-url-tls", 2, 8);
private static final String USER_AGENT = "wygj-personal-url-fetch/1.0";
@@ -94,9 +91,8 @@ public class PersonalUrlFetchService {
@Autowired
public PersonalUrlFetchService(PersonalKnowledgeProperties properties) {
this(properties, new FallbackResolver(
new DeadlineDnsResolver(DNS_EXECUTOR, new JndiDnsQuery()),
new DeadlineSystemResolver(DNS_EXECUTOR, InetAddress::getAllByName)), new RawSocketFetcher());
this(properties, new UdpDnsResolver(configuredDnsServers(), PersonalUrlFetchService::exchangeDns,
() -> ThreadLocalRandom.current().nextInt(0x10000)), new RawSocketFetcher());
}
private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) {
@@ -362,154 +358,189 @@ public class PersonalUrlFetchService {
List<InetAddress> resolve(String host, long deadlineNanos) throws IOException;
}
/** Falls back only when the primary resolver is unavailable; empty or unsafe answers remain fail-closed. */
static final class FallbackResolver implements Resolver {
private final Resolver primary;
private final Resolver fallback;
FallbackResolver(Resolver primary, Resolver fallback) {
this.primary = primary;
this.fallback = fallback;
}
@Override
public List<InetAddress> resolve(String host, long deadlineNanos) throws IOException {
try {
return primary.resolve(host, deadlineNanos);
} catch (IOException primaryFailure) {
if (deadlineNanos - System.nanoTime() <= 0) throw primaryFailure;
return fallback.resolve(host, deadlineNanos);
}
}
}
@FunctionalInterface
interface SystemAddressQuery {
InetAddress[] resolve(String host) throws IOException;
interface DnsExchange {
byte[] exchange(InetSocketAddress server, byte[] request, int timeoutMillis) throws IOException;
}
/** Bounds the JVM/system resolver with the same end-to-end deadline used by the fetch. */
static final class DeadlineSystemResolver implements Resolver {
private final ExecutorService executor;
private final SystemAddressQuery query;
/** Direct bounded UDP resolver. Closing the socket terminates every timed-out query without worker threads. */
static final class UdpDnsResolver implements Resolver {
private static final int TYPE_A = 1;
private static final int TYPE_AAAA = 28;
private final List<InetSocketAddress> servers;
private final DnsExchange exchange;
private final IntSupplier transactionIds;
DeadlineSystemResolver(ExecutorService executor, SystemAddressQuery query) {
this.executor = executor;
this.query = query;
UdpDnsResolver(List<InetSocketAddress> servers, DnsExchange exchange, IntSupplier transactionIds) {
this.servers = List.copyOf(servers);
this.exchange = exchange;
this.transactionIds = transactionIds;
}
@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<InetAddress[]> future;
try {
future = executor.submit(() -> query.resolve(host));
} catch (RejectedExecutionException ex) {
throw new IOException("resolution unavailable");
}
try {
InetAddress[] addresses = future.get(remaining, TimeUnit.NANOSECONDS);
return addresses == null ? List.of() : List.copyOf(Arrays.asList(addresses));
} catch (TimeoutException ex) {
cancelAndPurge(executor, future);
throw new IOException("resolution deadline exceeded");
} catch (InterruptedException ex) {
cancelAndPurge(executor, future);
Thread.currentThread().interrupt();
throw new IOException("resolution interrupted");
} catch (ExecutionException ex) {
cancelAndPurge(executor, future);
throw new IOException("resolution failed");
if (servers.isEmpty()) throw new IOException("DNS resolver unavailable");
IOException last = null;
for (int serverIndex = 0; serverIndex < servers.size(); serverIndex++) {
List<InetAddress> addresses = new ArrayList<>();
boolean received = false;
for (int typeIndex = 0; typeIndex < 2; typeIndex++) {
int type = typeIndex == 0 ? TYPE_A : TYPE_AAAA;
int operationsLeft = (servers.size() - serverIndex) * 2 - typeIndex;
int timeout = dnsTimeout(deadlineNanos, operationsLeft);
int transactionId = transactionIds.getAsInt() & 0xffff;
byte[] request = dnsQuery(host, type, transactionId);
try {
byte[] response = exchange.exchange(servers.get(serverIndex), request, timeout);
addresses.addAll(dnsAnswers(response, transactionId, type));
received = true;
} catch (SocketTimeoutException ex) {
last = ex;
} catch (IOException ex) {
last = ex;
}
}
if (received) return addresses.stream().distinct().toList();
}
throw last == null ? new IOException("DNS resolution failed") : last;
}
}
@FunctionalInterface
interface DnsQuery {
List<String> resolve(String host, int timeoutMillis, int retries) throws NamingException;
private static List<InetSocketAddress> configuredDnsServers() {
String configured = System.getProperty("aihr.personal.dns-servers");
if (configured == null || configured.isBlank()) configured = System.getenv("AIHR_PERSONAL_DNS_SERVERS");
List<String> literals = new ArrayList<>();
if (configured != null && !configured.isBlank()) {
for (String value : configured.split("[,\\s]+")) if (!value.isBlank()) literals.add(value.trim());
} else {
try {
for (String line : Files.readAllLines(Path.of("/etc/resolv.conf"), StandardCharsets.US_ASCII)) {
String value = line.replaceFirst("#.*$", "").trim();
if (!value.startsWith("nameserver")) continue;
String[] parts = value.split("\\s+");
if (parts.length == 2) literals.add(parts[1]);
}
} catch (IOException ignored) {
return List.of();
}
}
List<InetSocketAddress> servers = new ArrayList<>();
for (String literal : literals) {
if (servers.size() >= 4) break;
try {
servers.add(new InetSocketAddress(numericAddress(literal), 53));
} catch (IOException ignored) {
// Invalid configured resolver entries are not resolved as hostnames.
}
}
return List.copyOf(servers);
}
static final class DeadlineDnsResolver implements Resolver {
private final ExecutorService executor;
private final DnsQuery query;
DeadlineDnsResolver(DnsQuery query) {
this(DNS_EXECUTOR, query);
}
DeadlineDnsResolver(ExecutorService executor, DnsQuery query) {
this.executor = executor;
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<String>> future;
try {
future = executor.submit(() -> {
long taskRemainingMillis = TimeUnit.NANOSECONDS.toMillis(deadlineNanos - System.nanoTime());
if (taskRemainingMillis <= 0) throw new NamingException("resolution deadline exceeded");
return query.resolve(host, (int) Math.min(5_000L, taskRemainingMillis), 0);
});
} catch (RejectedExecutionException ex) {
throw new IOException("resolution unavailable");
}
List<String> literals;
try {
literals = future.get(remaining, TimeUnit.NANOSECONDS);
} catch (TimeoutException ex) {
cancelAndPurge(executor, future);
throw new IOException("resolution deadline exceeded");
} catch (InterruptedException ex) {
cancelAndPurge(executor, future);
Thread.currentThread().interrupt();
throw new IOException("resolution interrupted");
} catch (ExecutionException ex) {
cancelAndPurge(executor, future);
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);
private static byte[] exchangeDns(InetSocketAddress server, byte[] request, int timeoutMillis) throws IOException {
try (DatagramSocket socket = new DatagramSocket()) {
socket.connect(server);
socket.setSoTimeout(timeoutMillis);
socket.send(new DatagramPacket(request, request.length));
byte[] buffer = new byte[4096];
DatagramPacket response = new DatagramPacket(buffer, buffer.length);
socket.receive(response);
validateDnsSource(server, response);
return java.util.Arrays.copyOf(response.getData(), response.getLength());
}
}
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();
static void validateDnsSource(InetSocketAddress server, DatagramPacket response) throws IOException {
if (!server.getAddress().equals(response.getAddress()) || server.getPort() != response.getPort()) {
throw new IOException("DNS response source mismatch");
}
}
private static int dnsTimeout(long deadlineNanos, int operationsLeft) throws IOException {
long remaining = deadlineNanos - System.nanoTime();
if (remaining <= 0) throw new IOException("DNS resolution deadline exceeded");
long millis = Math.max(1, TimeUnit.NANOSECONDS.toMillis(remaining) / Math.max(1, operationsLeft));
return (int) Math.min(2_000, millis);
}
private static byte[] dnsQuery(String host, int type, int transactionId) throws IOException {
ByteArrayOutputStream output = new ByteArrayOutputStream();
output.write((transactionId >>> 8) & 0xff);
output.write(transactionId & 0xff);
output.write(new byte[]{1, 0, 0, 1, 0, 0, 0, 0, 0, 0});
for (String label : host.split("\\.")) {
byte[] bytes = label.getBytes(StandardCharsets.US_ASCII);
if (bytes.length == 0 || bytes.length > 63) throw new IOException("invalid DNS name");
output.write(bytes.length);
output.write(bytes);
}
output.write(0);
output.write((type >>> 8) & 0xff);
output.write(type & 0xff);
output.write(new byte[]{0, 1});
return output.toByteArray();
}
private static List<InetAddress> dnsAnswers(byte[] response, int transactionId, int expectedType)
throws IOException {
if (response == null || response.length < 12 || response.length > 4096
|| unsigned16(response, 0) != transactionId) throw new IOException("invalid DNS response");
int flags = unsigned16(response, 2);
if ((flags & 0x8000) == 0 || (flags & 0x0200) != 0 || (flags & 0x000f) != 0
|| unsigned16(response, 4) != 1) throw new IOException("invalid DNS response");
int answerCount = unsigned16(response, 6);
int totalRecords = answerCount + unsigned16(response, 8) + unsigned16(response, 10);
if (answerCount > 64 || totalRecords > 128) throw new IOException("invalid DNS response");
int position = skipDnsName(response, 12);
requireDnsBytes(response, position, 4);
int questionType = unsigned16(response, position);
int questionClass = unsigned16(response, position + 2);
if (questionType != expectedType || questionClass != 1) throw new IOException("invalid DNS response");
position += 4;
List<InetAddress> addresses = new ArrayList<>();
for (int index = 0; index < answerCount; index++) {
position = skipDnsName(response, position);
requireDnsBytes(response, position, 10);
int type = unsigned16(response, position);
int recordClass = unsigned16(response, position + 2);
int length = unsigned16(response, position + 8);
position += 10;
requireDnsBytes(response, position, length);
if (recordClass == 1 && type == expectedType
&& ((type == UdpDnsResolver.TYPE_A && length == 4)
|| (type == UdpDnsResolver.TYPE_AAAA && length == 16))) {
addresses.add(InetAddress.getByAddress(java.util.Arrays.copyOfRange(response, position, position + length)));
}
position += length;
}
return List.copyOf(addresses);
}
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 int skipDnsName(byte[] message, int position) throws IOException {
for (int labels = 0; labels < 128; labels++) {
requireDnsBytes(message, position, 1);
int length = message[position] & 0xff;
if (length == 0) return position + 1;
if ((length & 0xc0) == 0xc0) {
requireDnsBytes(message, position, 2);
int pointer = ((length & 0x3f) << 8) | (message[position + 1] & 0xff);
if (pointer >= message.length) throw new IOException("invalid DNS compression pointer");
return position + 2;
}
if ((length & 0xc0) != 0 || length > 63) throw new IOException("invalid DNS label");
position++;
requireDnsBytes(message, position, length);
position += length;
}
throw new IOException("DNS name too deep");
}
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 int unsigned16(byte[] value, int offset) throws IOException {
requireDnsBytes(value, offset, 2);
return ((value[offset] & 0xff) << 8) | (value[offset + 1] & 0xff);
}
private static void requireDnsBytes(byte[] value, int offset, int length) throws IOException {
if (offset < 0 || length < 0 || offset > value.length - length) throw new IOException("truncated DNS response");
}
private static InetAddress numericAddress(String literal) throws IOException {
@@ -14,9 +14,12 @@ import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.DatagramPacket;
import java.net.InetAddress;
import java.net.InetSocketAddress;
import java.net.ServerSocket;
import java.net.Socket;
import java.net.SocketTimeoutException;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.time.Instant;
@@ -223,112 +226,86 @@ class PersonalUrlFetchServiceTest {
}
@Test
void nativeDnsQueryReceivesRemainingTimeoutAndNumericAnswersOnly() {
void udpDnsMovesPastTwoSilentResolversWithoutWorkerPoolExhaustion() throws Exception {
List<InetSocketAddress> servers = List.of(
new InetSocketAddress("127.0.0.1", 5301),
new InetSocketAddress("127.0.0.1", 5302),
new InetSocketAddress("127.0.0.1", 5303));
AtomicInteger exchanges = new AtomicInteger();
var resolver = new PersonalUrlFetchService.UdpDnsResolver(servers, (server, request, timeoutMillis) -> {
exchanges.incrementAndGet();
if (server.getPort() != 5303) throw new SocketTimeoutException("silent resolver");
return dnsResponse(request, request[request.length - 3] == 1 ? PUBLIC : null);
}, () -> 0x1234);
List<InetAddress> result = resolver.resolve("example.com", System.nanoTime() + TimeUnit.SECONDS.toNanos(1));
assertEquals(List.of(PUBLIC), result);
assertEquals(6, exchanges.get());
}
@Test
void udpDnsFallbackAddressesStillUsePublicPolicyAndHonorDeadline() {
InetSocketAddress server = new InetSocketAddress("127.0.0.1", 5301);
var privateResolver = new PersonalUrlFetchService.UdpDnsResolver(List.of(server),
(ignored, request, timeoutMillis) -> dnsResponse(request, address("127.0.0.1")), () -> 7);
assertCode("PERSONAL_URL_BLOCKED", () -> fixture(privateResolver, request -> ok("text/plain", "ok"))
.validate("https://example.com/"), "private UDP answer");
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",
var silent = new PersonalUrlFetchService.UdpDnsResolver(List.of(server),
(ignored, request, timeoutMillis) -> {
timeoutSeen.set(timeoutMillis);
throw new SocketTimeoutException("silent resolver");
}, () -> 8);
long started = System.nanoTime();
assertThrows(IOException.class, () -> silent.resolve("example.com",
started + TimeUnit.MILLISECONDS.toNanos(40)));
assertTrue(timeoutSeen.get() > 0 && timeoutSeen.get() <= 40);
assertTrue(TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - started) < 500);
}
@Test
void udpDnsParsesAAndAaaaAndAcceptsValidEmptyAnswer() throws Exception {
InetAddress ipv6 = address("2606:4700:4700::1111");
InetSocketAddress server = new InetSocketAddress("127.0.0.1", 5301);
var resolver = new PersonalUrlFetchService.UdpDnsResolver(List.of(server),
(ignored, request, timeoutMillis) -> dnsResponse(request,
request[request.length - 3] == 1 ? PUBLIC : ipv6), () -> 0x2211);
assertEquals(List.of(PUBLIC, ipv6), resolver.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"));
var empty = new PersonalUrlFetchService.UdpDnsResolver(List.of(server),
(ignored, request, timeoutMillis) -> dnsResponse(request, null), () -> 0x2212);
assertEquals(List.of(), empty.resolve("example.com",
System.nanoTime() + TimeUnit.SECONDS.toNanos(1)));
}
@Test
void fallsBackToBoundedSystemDnsOnlyWhenJndiResolutionFails() throws Exception {
AtomicInteger fallbackCalls = new AtomicInteger();
var resolver = new PersonalUrlFetchService.FallbackResolver(
(host, deadline) -> { throw new IOException("JNDI unavailable"); },
(host, deadline) -> { fallbackCalls.incrementAndGet(); return List.of(PUBLIC); });
var service = fixture(resolver, request -> ok("text/plain", "ok"));
assertEquals("https://example.com/", service.validate("https://example.com/").toString());
assertEquals(1, fallbackCalls.get());
fallbackCalls.set(0);
var emptyPrimary = new PersonalUrlFetchService.FallbackResolver(
(host, deadline) -> List.of(),
(host, deadline) -> { fallbackCalls.incrementAndGet(); return List.of(PUBLIC); });
assertCode("PERSONAL_URL_BLOCKED", () -> fixture(emptyPrimary, request -> ok("text/plain", "ok"))
.validate("https://example.com/"), "empty primary result must fail closed");
assertEquals(0, fallbackCalls.get());
void udpDnsRejectsTransactionMismatchTruncationAndInvalidCompressionPointer() {
InetSocketAddress server = new InetSocketAddress("127.0.0.1", 5301);
assertMalformedDns(server, response -> response[1] ^= 1, "transaction mismatch");
assertMalformedDns(server, response -> response[2] |= 0x02, "truncated response flag");
assertMalformedDns(server, response -> {
int answerOffset = dnsQuestionEnd(response);
response[answerOffset] = (byte) 0xff;
response[answerOffset + 1] = (byte) 0xff;
}, "compression pointer out of bounds");
var emptyPacket = new PersonalUrlFetchService.UdpDnsResolver(List.of(server),
(ignored, request, timeoutMillis) -> new byte[0], () -> 0x3311);
assertThrows(IOException.class, () -> emptyPacket.resolve("example.com",
System.nanoTime() + TimeUnit.SECONDS.toNanos(1)), "empty packet");
}
@Test
void validatesEveryFallbackAddressAndFailsClosedWhenFallbackFails() {
var privateFallback = new PersonalUrlFetchService.FallbackResolver(
(host, deadline) -> { throw new IOException("JNDI unavailable"); },
(host, deadline) -> List.of(PUBLIC, address("127.0.0.1")));
assertCode("PERSONAL_URL_BLOCKED", () -> fixture(privateFallback, request -> ok("text/plain", "ok"))
.validate("https://example.com/"), "mixed fallback addresses");
var failedFallback = new PersonalUrlFetchService.FallbackResolver(
(host, deadline) -> { throw new IOException("JNDI unavailable"); },
(host, deadline) -> { throw new IOException("system DNS unavailable"); });
assertCode("PERSONAL_URL_BLOCKED", () -> fixture(failedFallback, request -> ok("text/plain", "ok"))
.validate("https://example.com/"), "fallback failure");
}
@Test
void boundsSystemDnsFallbackByTheSharedDeadline() throws Exception {
ExecutorService executor = boundedExecutor("system-dns-wall-test");
CountDownLatch entered = new CountDownLatch(1);
CountDownLatch release = new CountDownLatch(1);
try {
var resolver = new PersonalUrlFetchService.DeadlineSystemResolver(executor, host -> {
entered.countDown();
boolean done = false;
while (!done) {
try { release.await(); done = true; }
catch (InterruptedException ignored) { }
}
return new InetAddress[]{PUBLIC};
});
long deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(30);
assertThrows(IOException.class, () -> resolver.resolve("example.com", deadline));
assertTrue(entered.await(1, TimeUnit.SECONDS));
} finally {
release.countDown();
executor.shutdownNow();
assertTrue(executor.awaitTermination(1, TimeUnit.SECONDS));
}
}
@Test
void outerDnsDeadlineReturnsWhenQueryIgnoresInterrupt() throws Exception {
ExecutorService executor = boundedExecutor("dns-wall-test");
CountDownLatch entered = new CountDownLatch(1);
CountDownLatch release = new CountDownLatch(1);
try {
var resolver = new PersonalUrlFetchService.DeadlineDnsResolver(executor, (host, timeoutMillis, retries) -> {
entered.countDown();
boolean done = false;
while (!done) {
try { release.await(); done = true; }
catch (InterruptedException ignored) { }
}
return List.of("93.184.216.34");
});
long started = System.nanoTime();
assertThrows(IOException.class, () -> resolver.resolve("example.com", started + TimeUnit.MILLISECONDS.toNanos(40)));
assertTrue(entered.await(200, TimeUnit.MILLISECONDS));
assertTrue(TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - started) < 500);
} finally {
release.countDown();
executor.shutdownNow();
assertTrue(executor.awaitTermination(1, TimeUnit.SECONDS));
}
void udpDnsRejectsUnexpectedResponseSource() throws Exception {
InetSocketAddress expected = new InetSocketAddress(address("127.0.0.1"), 5301);
DatagramPacket wrongAddress = new DatagramPacket(new byte[1], 1,
address("127.0.0.2"), 5301);
DatagramPacket wrongPort = new DatagramPacket(new byte[1], 1,
address("127.0.0.1"), 5302);
assertThrows(IOException.class, () -> PersonalUrlFetchService.validateDnsSource(expected, wrongAddress));
assertThrows(IOException.class, () -> PersonalUrlFetchService.validateDnsSource(expected, wrongPort));
}
@Test
@@ -634,6 +611,40 @@ class PersonalUrlFetchServiceTest {
catch (Exception ex) { throw new AssertionError(ex); }
}
private static byte[] dnsResponse(byte[] request, InetAddress answer) throws IOException {
ByteArrayOutputStream output = new ByteArrayOutputStream();
output.write(request, 0, 2);
output.write(new byte[]{(byte) 0x81, (byte) 0x80, 0, 1, 0, (byte) (answer == null ? 0 : 1), 0, 0, 0, 0});
output.write(request, 12, request.length - 12);
if (answer != null) {
byte[] address = answer.getAddress();
output.write(new byte[]{(byte) 0xc0, 0x0c});
output.write(request, request.length - 4, 2);
output.write(new byte[]{0, 1, 0, 0, 0, 30, 0, (byte) address.length});
output.write(address);
}
return output.toByteArray();
}
private static void assertMalformedDns(InetSocketAddress server,
java.util.function.Consumer<byte[]> mutation,
String context) {
var resolver = new PersonalUrlFetchService.UdpDnsResolver(List.of(server),
(ignored, request, timeoutMillis) -> {
byte[] response = dnsResponse(request, PUBLIC);
mutation.accept(response);
return response;
}, () -> 0x3311);
assertThrows(IOException.class, () -> resolver.resolve("example.com",
System.nanoTime() + TimeUnit.SECONDS.toNanos(1)), context);
}
private static int dnsQuestionEnd(byte[] response) {
int position = 12;
while ((response[position] & 0xff) != 0) position += 1 + (response[position] & 0xff);
return position + 5;
}
private static ByteArrayInputStream stream(String value) {
return new ByteArrayInputStream(value.getBytes(StandardCharsets.US_ASCII));
}