fix(personal): bound TLS handshake duration
This commit is contained in:
+44
-1
@@ -69,6 +69,7 @@ public class PersonalUrlFetchService {
|
||||
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";
|
||||
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(
|
||||
@@ -657,7 +658,7 @@ public class PersonalUrlFetchService {
|
||||
RawSocketFetcher.configureTlsParameters(parameters, tlsHost);
|
||||
ssl.setSSLParameters(parameters);
|
||||
ssl.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
|
||||
ssl.startHandshake();
|
||||
runTlsHandshake(ssl, deadlineNanos);
|
||||
active = ssl;
|
||||
}
|
||||
return new SocketConnection(active);
|
||||
@@ -670,6 +671,48 @@ public class PersonalUrlFetchService {
|
||||
static SSLSocketFactory defaultSslSocketFactory() {
|
||||
return (SSLSocketFactory) SSLSocketFactory.getDefault();
|
||||
}
|
||||
|
||||
static void runTlsHandshake(SSLSocket socket, long deadlineNanos) throws IOException {
|
||||
runTlsHandshake(socket, deadlineNanos, HANDSHAKE_EXECUTOR);
|
||||
}
|
||||
|
||||
static void runTlsHandshake(SSLSocket socket, long deadlineNanos, ExecutorService executor) throws IOException {
|
||||
long remaining = deadlineNanos - System.nanoTime();
|
||||
if (remaining <= 0) {
|
||||
closeTlsSocket(socket);
|
||||
throw new IOException("TLS handshake deadline exceeded");
|
||||
}
|
||||
Future<?> future;
|
||||
try {
|
||||
future = executor.submit(() -> {
|
||||
socket.startHandshake();
|
||||
return null;
|
||||
});
|
||||
} catch (RejectedExecutionException ex) {
|
||||
closeTlsSocket(socket);
|
||||
throw new IOException("TLS handshake unavailable");
|
||||
}
|
||||
try {
|
||||
future.get(remaining, TimeUnit.NANOSECONDS);
|
||||
} catch (TimeoutException ex) {
|
||||
closeTlsSocket(socket);
|
||||
cancelAndPurge(executor, future);
|
||||
throw new IOException("TLS handshake deadline exceeded");
|
||||
} catch (InterruptedException ex) {
|
||||
closeTlsSocket(socket);
|
||||
cancelAndPurge(executor, future);
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IOException("TLS handshake interrupted");
|
||||
} catch (ExecutionException ex) {
|
||||
closeTlsSocket(socket);
|
||||
cancelAndPurge(executor, future);
|
||||
throw new IOException("TLS handshake failed");
|
||||
}
|
||||
}
|
||||
|
||||
private static void closeTlsSocket(SSLSocket socket) {
|
||||
try { socket.close(); } catch (IOException ignored) { }
|
||||
}
|
||||
}
|
||||
|
||||
private record SocketConnection(Socket socket) implements Connection {
|
||||
|
||||
+51
@@ -7,6 +7,7 @@ import org.junit.jupiter.api.Test;
|
||||
|
||||
import javax.net.ssl.SNIHostName;
|
||||
import javax.net.ssl.SSLParameters;
|
||||
import javax.net.ssl.SSLSocket;
|
||||
import javax.net.ssl.SSLSocketFactory;
|
||||
import java.io.ByteArrayInputStream;
|
||||
import java.io.ByteArrayOutputStream;
|
||||
@@ -35,6 +36,7 @@ import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.Mockito.*;
|
||||
|
||||
@Tag("dev")
|
||||
class PersonalUrlFetchServiceTest {
|
||||
@@ -414,6 +416,55 @@ class PersonalUrlFetchServiceTest {
|
||||
PersonalUrlFetchService.JvmConnectionFactory.defaultSslSocketFactory().getClass());
|
||||
}
|
||||
|
||||
@Test
|
||||
void tlsHandshakeDeadlineClosesSocketAndReleasesIgnoringTask() throws Exception {
|
||||
ExecutorService executor = boundedExecutor("tls-wall-test");
|
||||
CountDownLatch release = new CountDownLatch(1);
|
||||
AtomicBoolean closed = new AtomicBoolean();
|
||||
SSLSocket socket = mock(SSLSocket.class);
|
||||
doAnswer(invocation -> {
|
||||
boolean done = false;
|
||||
while (!done) {
|
||||
try { release.await(); done = true; }
|
||||
catch (InterruptedException ignored) { }
|
||||
}
|
||||
return null;
|
||||
}).when(socket).startHandshake();
|
||||
doAnswer(invocation -> { closed.set(true); release.countDown(); return null; }).when(socket).close();
|
||||
try {
|
||||
long started = System.nanoTime();
|
||||
assertThrows(IOException.class, () -> PersonalUrlFetchService.JvmConnectionFactory.runTlsHandshake(
|
||||
socket, started + TimeUnit.MILLISECONDS.toNanos(40), executor));
|
||||
assertTrue(TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - started) < 500);
|
||||
assertTrue(closed.get());
|
||||
} finally {
|
||||
release.countDown();
|
||||
executor.shutdownNow();
|
||||
assertTrue(executor.awaitTermination(1, TimeUnit.SECONDS));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void successfulTlsHandshakeKeepsSocketOpenAndRejectedQueueFailsClosed() throws Exception {
|
||||
ExecutorService successExecutor = boundedExecutor("tls-success-test");
|
||||
SSLSocket success = mock(SSLSocket.class);
|
||||
try {
|
||||
PersonalUrlFetchService.JvmConnectionFactory.runTlsHandshake(success,
|
||||
System.nanoTime() + TimeUnit.SECONDS.toNanos(1), successExecutor);
|
||||
verify(success).startHandshake();
|
||||
verify(success, never()).close();
|
||||
} finally {
|
||||
successExecutor.shutdownNow();
|
||||
}
|
||||
|
||||
ExecutorService rejected = boundedExecutor("tls-rejected-test");
|
||||
rejected.shutdownNow();
|
||||
SSLSocket socket = mock(SSLSocket.class);
|
||||
assertThrows(IOException.class, () -> PersonalUrlFetchService.JvmConnectionFactory.runTlsHandshake(
|
||||
socket, System.nanoTime() + TimeUnit.SECONDS.toNanos(1), rejected));
|
||||
verify(socket).close();
|
||||
}
|
||||
|
||||
@Test
|
||||
void rawTransportWritesBracketedIpv6Host() throws Exception {
|
||||
ByteArrayOutputStream requestBytes = new ByteArrayOutputStream();
|
||||
|
||||
Reference in New Issue
Block a user