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.Locale;
import java.util.Map; import java.util.Map;
import java.util.Set; 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.TimeUnit;
import java.util.concurrent.TimeoutException;
import java.util.concurrent.atomic.AtomicInteger;
@Service @Service
public class PersonalUrlFetchService { public class PersonalUrlFetchService {
@@ -59,6 +67,8 @@ public class PersonalUrlFetchService {
private static final int MAX_LINE_BYTES = 8 * 1024; private static final int MAX_LINE_BYTES = 8 * 1024;
private static final long TOTAL_TIMEOUT_NANOS = 15_000_000_000L; private static final long TOTAL_TIMEOUT_NANOS = 15_000_000_000L;
private static final long HARD_MAX_BODY_BYTES = 10L * 1024 * 1024; 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 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 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( private static final Map<String, String> SAFE_HEADERS = Map.of(
@@ -80,7 +90,7 @@ public class PersonalUrlFetchService {
private final Fetcher fetcher; private final Fetcher fetcher;
public PersonalUrlFetchService(PersonalKnowledgeProperties properties) { 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) { private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) {
@@ -156,7 +166,9 @@ public class PersonalUrlFetchService {
List<InetAddress> addresses; List<InetAddress> addresses;
try { try {
requireTimeRemaining(deadlineNanos); 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); requireTimeRemaining(deadlineNanos);
} catch (Exception ex) { } catch (Exception ex) {
throw new ServiceException(BLOCKED); throw new ServiceException(BLOCKED);
@@ -217,6 +229,21 @@ public class PersonalUrlFetchService {
if (deadlineNanos - System.nanoTime() <= 0) throw new ServiceException(FETCH_FAILED); 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) { private static boolean isRedirect(int status) {
return status == 301 || status == 302 || status == 303 || status == 307 || status == 308; 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 == 100 && b >= 64 && b <= 127) return false;
if (a == 169 && b == 254) return false; if (a == 169 && b == 254) return false;
if (a == 172 && b >= 16 && b <= 31) 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) if (a == 192 && ((b == 31 && c == 196) || (b == 52 && c == 193)
|| (b == 88 && c == 99) || (b == 175 && c == 48))) return false; || (b == 88 && c == 99) || (b == 175 && c == 48))) return false;
if (a == 198 && (b == 18 || b == 19 || (b == 51 && c == 100))) 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 { static final class DeadlineDnsResolver implements Resolver {
private final ExecutorService executor;
private final DnsQuery query; private final DnsQuery query;
DeadlineDnsResolver(DnsQuery query) { DeadlineDnsResolver(DnsQuery query) {
this(DNS_EXECUTOR, query);
}
DeadlineDnsResolver(ExecutorService executor, DnsQuery query) {
this.executor = executor;
this.query = query; this.query = query;
} }
@Override @Override
public List<InetAddress> resolve(String host, long deadlineNanos) throws IOException { public List<InetAddress> resolve(String host, long deadlineNanos) throws IOException {
long remaining = deadlineNanos - System.nanoTime(); long remaining = deadlineNanos - System.nanoTime();
long remainingMillis = TimeUnit.NANOSECONDS.toMillis(remaining); if (remaining <= 0) throw new IOException("resolution deadline exceeded");
if (remainingMillis <= 0) throw new IOException("resolution deadline exceeded"); Future<List<String>> future;
int timeoutMillis = (int) Math.min(5_000L, remainingMillis); 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; List<String> literals;
try { try {
literals = query.resolve(host, timeoutMillis, 0); literals = future.get(remaining, TimeUnit.NANOSECONDS);
} catch (NamingException | RuntimeException ex) { } 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"); throw new IOException("resolution failed");
} }
if (System.nanoTime() >= deadlineNanos) throw new IOException("resolution deadline exceeded"); 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)) { if (parts[i].isEmpty() || !parts[i].chars().allMatch(Character::isDigit)) {
throw new IOException("invalid DNS answer"); throw new IOException("invalid DNS answer");
} }
if (parts[i].length() > 1 && parts[i].charAt(0) == '0') throw new IOException("invalid DNS answer");
int octet; int octet;
try { octet = Integer.parseInt(parts[i]); } try { octet = Integer.parseInt(parts[i]); }
catch (NumberFormatException ex) { throw new IOException("invalid DNS answer"); } catch (NumberFormatException ex) { throw new IOException("invalid DNS answer"); }
@@ -419,6 +469,15 @@ public class PersonalUrlFetchService {
return address; 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 @FunctionalInterface
public interface Fetcher { public interface Fetcher {
TransportResponse fetch(FetchRequest request) throws IOException; TransportResponse fetch(FetchRequest request) throws IOException;
@@ -461,13 +520,19 @@ public class PersonalUrlFetchService {
static final class RawSocketFetcher implements Fetcher { static final class RawSocketFetcher implements Fetcher {
private final ConnectionFactory connections; private final ConnectionFactory connections;
private final ExecutorService writes;
RawSocketFetcher() { RawSocketFetcher() {
this(new JvmConnectionFactory()); this(new JvmConnectionFactory(), WRITE_EXECUTOR);
} }
RawSocketFetcher(ConnectionFactory connections) { RawSocketFetcher(ConnectionFactory connections) {
this(connections, WRITE_EXECUTOR);
}
RawSocketFetcher(ConnectionFactory connections, ExecutorService writes) {
this.connections = connections; this.connections = connections;
this.writes = writes;
} }
@Override @Override
@@ -487,7 +552,7 @@ public class PersonalUrlFetchService {
URI uri = request.uri(); URI uri = request.uri();
int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80); int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80);
try (Connection connection = connections.connect(uri, address, port, request.deadlineNanos())) { try (Connection connection = connections.connect(uri, address, port, request.deadlineNanos())) {
writeRequest(connection.output(), request); writeRequestWithDeadline(connection, request);
connection.setReadTimeout(timeout(request.deadlineNanos(), 5_000)); connection.setReadTimeout(timeout(request.deadlineNanos(), 5_000));
TransportResponse response = parseHttpResponse( TransportResponse response = parseHttpResponse(
new DeadlineInputStream(connection.input(), connection, request.deadlineNanos()), request.maxBodyBytes()); 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) { static SSLParameters tlsParameters(String host) {
SSLParameters parameters = new SSLParameters(); SSLParameters parameters = new SSLParameters();
configureTlsParameters(parameters, host); configureTlsParameters(parameters, host);
@@ -12,6 +12,7 @@ import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream; import java.io.ByteArrayOutputStream;
import java.io.IOException; import java.io.IOException;
import java.io.InputStream; import java.io.InputStream;
import java.io.OutputStream;
import java.net.InetAddress; import java.net.InetAddress;
import java.net.ServerSocket; import java.net.ServerSocket;
import java.net.Socket; import java.net.Socket;
@@ -24,7 +25,12 @@ import java.util.Arrays;
import java.util.Hashtable; import java.util.Hashtable;
import java.util.List; import java.util.List;
import java.util.Map; 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.TimeUnit;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference; import java.util.concurrent.atomic.AtomicReference;
@@ -215,7 +221,7 @@ class PersonalUrlFetchServiceTest {
} }
@Test @Test
void boundedProductionDnsResolverTimesOutAndCancels() { void nativeDnsQueryReceivesRemainingTimeoutAndNumericAnswersOnly() {
AtomicInteger timeoutSeen = new AtomicInteger(); AtomicInteger timeoutSeen = new AtomicInteger();
AtomicInteger retriesSeen = new AtomicInteger(-1); AtomicInteger retriesSeen = new AtomicInteger(-1);
var resolver = new PersonalUrlFetchService.DeadlineDnsResolver((host, timeoutMillis, retries) -> { var resolver = new PersonalUrlFetchService.DeadlineDnsResolver((host, timeoutMillis, retries) -> {
@@ -237,6 +243,52 @@ class PersonalUrlFetchServiceTest {
assertEquals("0", environment.get("com.sun.jndi.dns.timeout.retries")); 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 @Test
void preservesRawPathQueryAndUnicodeEncodingAcrossNormalizationAndRedirects() { void preservesRawPathQueryAndUnicodeEncodingAcrossNormalizationAndRedirects() {
var seen = new ArrayList<PersonalUrlFetchService.FetchRequest>(); var seen = new ArrayList<PersonalUrlFetchService.FetchRequest>();
@@ -390,6 +442,26 @@ class PersonalUrlFetchServiceTest {
assertThrows(IOException.class, () -> fetcher.fetch(request)); 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 @Test
void rawTransportUsesOriginFormOverRealLoopbackSocket() throws Exception { void rawTransportUsesOriginFormOverRealLoopbackSocket() throws Exception {
InetAddress loopback = InetAddress.getLoopbackAddress(); InetAddress loopback = InetAddress.getLoopbackAddress();
@@ -459,6 +531,12 @@ class PersonalUrlFetchServiceTest {
assertEquals(code, assertThrows(ServiceException.class, action::run, context).getMessage(), context); 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 static final class FakeConnection implements PersonalUrlFetchService.Connection {
private final InputStream input; private final InputStream input;
private final ByteArrayOutputStream output; private final ByteArrayOutputStream output;
@@ -483,4 +561,22 @@ class PersonalUrlFetchServiceTest {
return super.read(); 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(); }
}
} }