fix(personal): close URL deadline and literal gaps

This commit is contained in:
2026-07-12 04:58:57 +08:00
parent 86e87be6f5
commit eff8ac2787
2 changed files with 207 additions and 11 deletions
@@ -41,7 +41,15 @@ import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.ArrayBlockingQueue;
import java.util.concurrent.ExecutionException;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Future;
import java.util.concurrent.RejectedExecutionException;
import java.util.concurrent.ThreadPoolExecutor;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException;
import java.util.concurrent.atomic.AtomicInteger;
@Service
public class PersonalUrlFetchService {
@@ -59,6 +67,8 @@ 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 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(
@@ -80,7 +90,7 @@ public class PersonalUrlFetchService {
private final Fetcher fetcher;
public PersonalUrlFetchService(PersonalKnowledgeProperties properties) {
this(properties, new DeadlineDnsResolver(new JndiDnsQuery()), new RawSocketFetcher());
this(properties, new DeadlineDnsResolver(DNS_EXECUTOR, new JndiDnsQuery()), new RawSocketFetcher());
}
private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) {
@@ -156,7 +166,9 @@ public class PersonalUrlFetchService {
List<InetAddress> addresses;
try {
requireTimeRemaining(deadlineNanos);
addresses = resolver.resolve(canonicalHost(uri), deadlineNanos);
String host = canonicalHost(uri);
InetAddress literal = literalHostAddress(host);
addresses = literal == null ? resolver.resolve(host, deadlineNanos) : List.of(literal);
requireTimeRemaining(deadlineNanos);
} catch (Exception ex) {
throw new ServiceException(BLOCKED);
@@ -217,6 +229,21 @@ public class PersonalUrlFetchService {
if (deadlineNanos - System.nanoTime() <= 0) throw new ServiceException(FETCH_FAILED);
}
private static ExecutorService boundedExecutor(String prefix, int threads, int queueCapacity) {
AtomicInteger sequence = new AtomicInteger();
return new ThreadPoolExecutor(threads, threads, 0L, TimeUnit.MILLISECONDS,
new ArrayBlockingQueue<>(queueCapacity), runnable -> {
Thread thread = new Thread(runnable, prefix + "-" + sequence.incrementAndGet());
thread.setDaemon(true);
return thread;
}, new ThreadPoolExecutor.AbortPolicy());
}
private static void cancelAndPurge(ExecutorService executor, Future<?> future) {
future.cancel(true);
if (executor instanceof ThreadPoolExecutor pool) pool.purge();
}
private static boolean isRedirect(int status) {
return status == 301 || status == 302 || status == 303 || status == 307 || status == 308;
}
@@ -298,7 +325,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)) return false;
if (a == 192 && (b == 168 || (b == 0 && (c == 0 || c == 2)))) 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;
@@ -335,22 +362,44 @@ public class PersonalUrlFetchService {
}
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();
long remainingMillis = TimeUnit.NANOSECONDS.toMillis(remaining);
if (remainingMillis <= 0) throw new IOException("resolution deadline exceeded");
int timeoutMillis = (int) Math.min(5_000L, remainingMillis);
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 = query.resolve(host, timeoutMillis, 0);
} catch (NamingException | RuntimeException ex) {
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");
@@ -405,6 +454,7 @@ public class PersonalUrlFetchService {
if (parts[i].isEmpty() || !parts[i].chars().allMatch(Character::isDigit)) {
throw new IOException("invalid DNS answer");
}
if (parts[i].length() > 1 && parts[i].charAt(0) == '0') throw new IOException("invalid DNS answer");
int octet;
try { octet = Integer.parseInt(parts[i]); }
catch (NumberFormatException ex) { throw new IOException("invalid DNS answer"); }
@@ -419,6 +469,15 @@ public class PersonalUrlFetchService {
return address;
}
private static InetAddress literalHostAddress(String host) throws IOException {
String lower = host.toLowerCase(Locale.ROOT);
if (host.indexOf(':') >= 0 || host.chars().allMatch(ch -> Character.isDigit(ch) || ch == '.')) {
return numericAddress(host);
}
if (lower.startsWith("0x") || lower.contains(".0x")) throw new IOException("invalid numeric host");
return null;
}
@FunctionalInterface
public interface Fetcher {
TransportResponse fetch(FetchRequest request) throws IOException;
@@ -461,13 +520,19 @@ public class PersonalUrlFetchService {
static final class RawSocketFetcher implements Fetcher {
private final ConnectionFactory connections;
private final ExecutorService writes;
RawSocketFetcher() {
this(new JvmConnectionFactory());
this(new JvmConnectionFactory(), WRITE_EXECUTOR);
}
RawSocketFetcher(ConnectionFactory connections) {
this(connections, WRITE_EXECUTOR);
}
RawSocketFetcher(ConnectionFactory connections, ExecutorService writes) {
this.connections = connections;
this.writes = writes;
}
@Override
@@ -487,7 +552,7 @@ public class PersonalUrlFetchService {
URI uri = request.uri();
int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80);
try (Connection connection = connections.connect(uri, address, port, request.deadlineNanos())) {
writeRequest(connection.output(), request);
writeRequestWithDeadline(connection, request);
connection.setReadTimeout(timeout(request.deadlineNanos(), 5_000));
TransportResponse response = parseHttpResponse(
new DeadlineInputStream(connection.input(), connection, request.deadlineNanos()), request.maxBodyBytes());
@@ -496,6 +561,41 @@ public class PersonalUrlFetchService {
}
}
private void writeRequestWithDeadline(Connection connection, FetchRequest request) throws IOException {
long remaining = request.deadlineNanos() - System.nanoTime();
if (remaining <= 0) throw new IOException("deadline exceeded");
Future<?> future;
try {
future = writes.submit(() -> {
writeRequest(connection.output(), request);
return null;
});
} catch (RejectedExecutionException ex) {
closeQuietly(connection);
throw new IOException("request writer unavailable");
}
try {
future.get(remaining, TimeUnit.NANOSECONDS);
} catch (TimeoutException ex) {
closeQuietly(connection);
cancelAndPurge(writes, future);
throw new IOException("request write deadline exceeded");
} catch (InterruptedException ex) {
closeQuietly(connection);
cancelAndPurge(writes, future);
Thread.currentThread().interrupt();
throw new IOException("request write interrupted");
} catch (ExecutionException ex) {
closeQuietly(connection);
cancelAndPurge(writes, future);
throw new IOException("request write failed");
}
}
private static void closeQuietly(Connection connection) {
try { connection.close(); } catch (IOException ignored) { }
}
static SSLParameters tlsParameters(String host) {
SSLParameters parameters = new SSLParameters();
configureTlsParameters(parameters, host);
@@ -12,6 +12,7 @@ import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.InetAddress;
import java.net.ServerSocket;
import java.net.Socket;
@@ -24,7 +25,12 @@ 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.CountDownLatch;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.ThreadPoolExecutor;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference;
@@ -215,7 +221,7 @@ class PersonalUrlFetchServiceTest {
}
@Test
void boundedProductionDnsResolverTimesOutAndCancels() {
void nativeDnsQueryReceivesRemainingTimeoutAndNumericAnswersOnly() {
AtomicInteger timeoutSeen = new AtomicInteger();
AtomicInteger retriesSeen = new AtomicInteger(-1);
var resolver = new PersonalUrlFetchService.DeadlineDnsResolver((host, timeoutMillis, retries) -> {
@@ -237,6 +243,52 @@ class PersonalUrlFetchServiceTest {
assertEquals("0", environment.get("com.sun.jndi.dns.timeout.retries"));
}
@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));
}
}
@Test
void literalHostsBypassDnsAndStillApplyAddressPolicy() {
AtomicInteger dnsCalls = new AtomicInteger();
var service = fixture((host, deadline) -> { dnsCalls.incrementAndGet(); return List.of(PUBLIC); },
request -> ok("text/plain", "ok"));
assertEquals("http://8.8.8.8/", service.validate("http://8.8.8.8").toString());
assertEquals("http://[2606:4700:4700::1111]/",
service.validate("http://[2606:4700:4700::1111]").toString());
assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://127.0.0.1"), "private literal");
assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://2130706433"), "integer literal");
assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://0177.0.0.1"), "octal literal");
assertEquals(0, dnsCalls.get());
}
@Test
void public192Dot0Dot1AddressIsNotCaughtBySpecialSlash24Rule() {
var service = fixture((host, deadline) -> List.of(address("192.0.1.1")), request -> ok("text/plain", "ok"));
assertEquals("http://example.com/", service.validate("http://example.com").toString());
}
@Test
void preservesRawPathQueryAndUnicodeEncodingAcrossNormalizationAndRedirects() {
var seen = new ArrayList<PersonalUrlFetchService.FetchRequest>();
@@ -390,6 +442,26 @@ class PersonalUrlFetchServiceTest {
assertThrows(IOException.class, () -> fetcher.fetch(request));
}
@Test
void rawTransportClosesConnectionWhenRequestWriteMissesDeadline() throws Exception {
ExecutorService executor = boundedExecutor("write-wall-test");
BlockingConnection connection = new BlockingConnection();
try {
var fetcher = new PersonalUrlFetchService.RawSocketFetcher(
(uri, address, port, deadlineNanos) -> connection, executor);
var request = new PersonalUrlFetchService.FetchRequest(URI.create("http://origin.example/"), List.of(PUBLIC),
System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(40), 100, Map.of());
long started = System.nanoTime();
assertThrows(IOException.class, () -> fetcher.fetch(request));
assertTrue(TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - started) < 500);
assertTrue(connection.closed.get());
} finally {
connection.close();
executor.shutdownNow();
assertTrue(executor.awaitTermination(1, TimeUnit.SECONDS));
}
}
@Test
void rawTransportUsesOriginFormOverRealLoopbackSocket() throws Exception {
InetAddress loopback = InetAddress.getLoopbackAddress();
@@ -459,6 +531,12 @@ class PersonalUrlFetchServiceTest {
assertEquals(code, assertThrows(ServiceException.class, action::run, context).getMessage(), context);
}
private static ExecutorService boundedExecutor(String name) {
return new ThreadPoolExecutor(1, 1, 0, TimeUnit.MILLISECONDS, new ArrayBlockingQueue<>(1), runnable -> {
Thread thread = new Thread(runnable, name); thread.setDaemon(true); return thread;
}, new ThreadPoolExecutor.AbortPolicy());
}
private static final class FakeConnection implements PersonalUrlFetchService.Connection {
private final InputStream input;
private final ByteArrayOutputStream output;
@@ -483,4 +561,22 @@ class PersonalUrlFetchServiceTest {
return super.read();
}
}
private static final class BlockingConnection implements PersonalUrlFetchService.Connection {
private final CountDownLatch release = new CountDownLatch(1);
private final AtomicBoolean closed = new AtomicBoolean();
private final OutputStream output = new OutputStream() {
@Override public void write(int value) {
boolean done = false;
while (!done) {
try { release.await(); done = true; }
catch (InterruptedException ignored) { }
}
}
};
@Override public InputStream input() { return new ByteArrayInputStream(new byte[0]); }
@Override public OutputStream output() { return output; }
@Override public void setReadTimeout(int millis) { }
@Override public void close() { closed.set(true); release.countDown(); }
}
}