diff --git a/dd-trace-core/build.gradle b/dd-trace-core/build.gradle index 46695a38059..f1b42962a50 100644 --- a/dd-trace-core/build.gradle +++ b/dd-trace-core/build.gradle @@ -120,6 +120,8 @@ dependencies { testImplementation group: 'commons-codec', name: 'commons-codec', version: '1.3' testImplementation group: 'com.amazonaws', name: 'aws-lambda-java-events', version:'3.11.0' testImplementation group: 'com.google.protobuf', name: 'protobuf-java', version: '3.14.0' + testImplementation libs.jnr.unixsocket + testImplementation libs.okhttp3.mockwebserver testImplementation libs.testcontainers testImplementation project(':utils:test-junit-utils') testImplementation project(':utils:test-junit-converter-utils') diff --git a/dd-trace-core/gradle.lockfile b/dd-trace-core/gradle.lockfile index 9540290ce41..6c8603df7d4 100644 --- a/dd-trace-core/gradle.lockfile +++ b/dd-trace-core/gradle.lockfile @@ -21,14 +21,14 @@ com.github.docker-java:docker-java-api:3.4.2=jmhRuntimeClasspath,testCompileClas com.github.docker-java:docker-java-transport-zerodep:3.4.2=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.github.docker-java:docker-java-transport:3.4.2=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.github.javaparser:javaparser-core:3.25.6=codenarc -com.github.jnr:jffi:1.3.15=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-a64asm:1.0.0=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-constants:0.10.4=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-enxio:0.32.20=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-ffi:2.2.19=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-posix:3.1.22=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-unixsocket:0.38.25=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath -com.github.jnr:jnr-x86asm:1.0.2=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jffi:1.3.15=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-a64asm:1.0.0=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-constants:0.10.4=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-enxio:0.32.20=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-ffi:2.2.19=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-posix:3.1.22=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-unixsocket:0.38.25=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.github.jnr:jnr-x86asm:1.0.2=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.github.spotbugs:spotbugs-annotations:4.10.4=compileClasspath,jmhCompileClasspath,spotbugs com.github.spotbugs:spotbugs:4.10.4=spotbugs com.github.stephenc.jcip:jcip-annotations:1.0-1=spotbugs @@ -51,6 +51,7 @@ com.google.protobuf:protobuf-java:3.14.0=jmhRuntimeClasspath,testCompileClasspat com.google.re2j:re2j:1.8=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.squareup.moshi:moshi:1.11.0=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.squareup.okhttp3:logging-interceptor:3.12.12=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath +com.squareup.okhttp3:mockwebserver:3.12.12=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.squareup.okhttp3:okhttp:3.12.12=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.squareup.okio:okio:1.17.5=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath com.thoughtworks.qdox:qdox:1.12.1=codenarc @@ -127,13 +128,13 @@ org.openjdk.jmh:jmh-generator-reflection:1.37=jmh,jmhCompileClasspath,jmhRuntime org.openjdk.jol:jol-core:0.17=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.opentest4j:opentest4j:1.3.0=jmhRuntimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.ow2.asm:asm-analysis:9.10.1=spotbugs -org.ow2.asm:asm-analysis:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath +org.ow2.asm:asm-analysis:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.ow2.asm:asm-commons:9.10.1=jacocoAnt,spotbugs -org.ow2.asm:asm-commons:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath +org.ow2.asm:asm-commons:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.ow2.asm:asm-tree:9.10.1=jacocoAnt,spotbugs -org.ow2.asm:asm-tree:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath +org.ow2.asm:asm-tree:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.ow2.asm:asm-util:9.10.1=spotbugs -org.ow2.asm:asm-util:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath,traceAgentTestRuntimeClasspath +org.ow2.asm:asm-util:9.7.1=jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.ow2.asm:asm:9.0=jmh,jmhCompileClasspath org.ow2.asm:asm:9.10.1=jacocoAnt,jmhRuntimeClasspath,spotbugs,testCompileClasspath,testRuntimeClasspath,traceAgentTestCompileClasspath,traceAgentTestRuntimeClasspath org.ow2.asm:asm:9.7.1=runtimeClasspath diff --git a/dd-trace-core/src/test/java/datadog/common/socket/UnixDomainServerSocketFactory.java b/dd-trace-core/src/test/java/datadog/common/socket/UnixDomainServerSocketFactory.java new file mode 100644 index 00000000000..6cad92ba460 --- /dev/null +++ b/dd-trace-core/src/test/java/datadog/common/socket/UnixDomainServerSocketFactory.java @@ -0,0 +1,109 @@ +package datadog.common.socket; + +import java.io.File; +import java.io.IOException; +import java.net.InetAddress; +import java.net.InetSocketAddress; +import java.net.ServerSocket; +import java.net.Socket; +import java.net.SocketAddress; +import java.net.SocketException; +import java.nio.channels.ClosedChannelException; +import javax.net.ServerSocketFactory; +import jnr.unixsocket.UnixServerSocketChannel; +import jnr.unixsocket.UnixSocketAddress; +import jnr.unixsocket.UnixSocketChannel; + +/** + * Adapts a JNR Unix-domain server channel to APIs such as MockWebServer that require a {@link + * ServerSocket}. Adapted from OkHttp's Unix-domain + * socket sample. + */ +public final class UnixDomainServerSocketFactory extends ServerSocketFactory { + private final File path; + + public UnixDomainServerSocketFactory(File path) { + this.path = path; + } + + @Override + public ServerSocket createServerSocket() throws IOException { + return new UnixDomainServerSocket(); + } + + @Override + public ServerSocket createServerSocket(int port) throws IOException { + return createServerSocket(); + } + + @Override + public ServerSocket createServerSocket(int port, int backlog) throws IOException { + return createServerSocket(); + } + + @Override + public ServerSocket createServerSocket(int port, int backlog, InetAddress inetAddress) + throws IOException { + return createServerSocket(); + } + + private final class UnixDomainServerSocket extends ServerSocket { + private UnixServerSocketChannel serverSocketChannel; + private InetSocketAddress endpoint; + + private UnixDomainServerSocket() throws IOException {} + + @Override + public void bind(SocketAddress endpoint, int backlog) throws IOException { + this.endpoint = (InetSocketAddress) endpoint; + UnixServerSocketChannel channel = UnixServerSocketChannel.open(); + boolean bound = false; + try { + channel.configureBlocking(true); + channel.socket().bind(new UnixSocketAddress(path)); + serverSocketChannel = channel; + bound = true; + } finally { + if (!bound) { + channel.close(); + } + } + } + + @Override + public void setReuseAddress(boolean on) { + // MockWebServer configures this TCP option before binding. It has no UDS equivalent. + } + + @Override + public int getLocalPort() { + return 1; // MockWebServer requires a port even though a UDS has none. + } + + @Override + public SocketAddress getLocalSocketAddress() { + return endpoint; + } + + @Override + public Socket accept() throws IOException { + try { + UnixSocketChannel channel = serverSocketChannel.accept(); + return new TunnelingUnixSocket(path, channel, endpoint); + } catch (ClosedChannelException e) { + SocketException socketException = new SocketException("Socket is closed"); + socketException.initCause(e); + throw socketException; + } + } + + @Override + public void close() throws IOException { + super.close(); + if (serverSocketChannel != null) { + serverSocketChannel.close(); + } + } + } +} diff --git a/dd-trace-core/src/test/java/datadog/trace/common/writer/DDAgentWriterCombinedTest.java b/dd-trace-core/src/test/java/datadog/trace/common/writer/DDAgentWriterCombinedTest.java index 3f1bc9d1324..8f3a8c40add 100644 --- a/dd-trace-core/src/test/java/datadog/trace/common/writer/DDAgentWriterCombinedTest.java +++ b/dd-trace-core/src/test/java/datadog/trace/common/writer/DDAgentWriterCombinedTest.java @@ -2,9 +2,17 @@ import static datadog.trace.api.ProtocolVersion.V0_5; import static datadog.trace.api.config.GeneralConfig.EXPERIMENTAL_PROPAGATE_PROCESS_TAGS_ENABLED; +import static datadog.trace.api.config.GeneralConfig.JDK_SOCKET_ENABLED; import static datadog.trace.common.writer.ddagent.Prioritization.ENSURE_TRACE; +import static okhttp3.mockwebserver.SocketPolicy.DISCONNECT_AT_END; +import static okhttp3.mockwebserver.SocketPolicy.NO_RESPONSE; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.condition.JRE.JAVA_16; +import static org.junit.jupiter.api.condition.OS.LINUX; +import static org.junit.jupiter.api.condition.OS.MAC; import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyInt; import static org.mockito.ArgumentMatchers.anyLong; @@ -18,6 +26,7 @@ import static org.mockito.Mockito.verifyNoMoreInteractions; import static org.mockito.Mockito.when; +import datadog.common.socket.UnixDomainServerSocketFactory; import datadog.communication.ddagent.DDAgentFeaturesDiscovery; import datadog.communication.http.OkHttpUtils; import datadog.communication.serialization.FlushingBuffer; @@ -39,6 +48,8 @@ import datadog.trace.test.junit.utils.config.WithConfig; import datadog.trace.test.util.Flaky; import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; import java.util.Collections; import java.util.List; import java.util.concurrent.CountDownLatch; @@ -46,11 +57,17 @@ import java.util.concurrent.Semaphore; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; import okhttp3.HttpUrl; +import okhttp3.mockwebserver.MockResponse; +import okhttp3.mockwebserver.MockWebServer; +import okhttp3.mockwebserver.RecordedRequest; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Timeout; +import org.junit.jupiter.api.condition.EnabledForJreRange; +import org.junit.jupiter.api.condition.EnabledOnOs; import org.mockito.Mockito; import org.tabletest.junit.TableTest; @@ -339,12 +356,11 @@ void monitorHappyPath(String agentVersion) { List minimalTrace = createMinimalTrace(); // DQH -- need to set-up a dummy agent for the final send callback to work - JavaTestHttpServer agent = + try (JavaTestHttpServer agent = JavaTestHttpServer.httpServer( server -> server.handlers( - h -> h.put(agentVersion, api -> api.getResponse().status(200).send()))); - try { + h -> h.put(agentVersion, api -> api.getResponse().status(200).send())))) { HttpUrl agentUrl = HttpUrl.get(agent.getAddress()); okhttp3.OkHttpClient client = OkHttpUtils.buildHttpClient(agentUrl, 1000); DDAgentFeaturesDiscovery discovery = @@ -383,8 +399,6 @@ void monitorHappyPath(String agentVersion) { writer.close(); verify(healthMetrics, times(1)).onShutdown(true); - } finally { - agent.close(); } } @@ -399,25 +413,11 @@ void monitorAgentReturnsError(String agentVersion) { List minimalTrace = createMinimalTrace(); // DQH -- need to set-up a dummy agent for the final send callback to work - final boolean[] first = {true}; - JavaTestHttpServer agent = + try (JavaTestHttpServer agent = JavaTestHttpServer.httpServer( server -> server.handlers( - h -> - h.put( - agentVersion, - api -> { - // DQH - DDApi sniffs for end point existence, so respond with 200 the - // first time - if (first[0]) { - api.getResponse().status(200).send(); - first[0] = false; - } else { - api.getResponse().status(500).send(); - } - }))); - try { + h -> h.put(agentVersion, api -> api.getResponse().status(500).send())))) { HttpUrl agentUrl = HttpUrl.get(agent.getAddress()); okhttp3.OkHttpClient client = OkHttpUtils.buildHttpClient(agentUrl, 1000); DDAgentFeaturesDiscovery discovery = @@ -456,8 +456,86 @@ void monitorAgentReturnsError(String agentVersion) { writer.close(); verify(healthMetrics, times(1)).onShutdown(true); + } + } + + @Test + @WithConfig(key = JDK_SOCKET_ENABLED, value = "true") + @EnabledForJreRange(min = JAVA_16) + @EnabledOnOs({LINUX, MAC}) + void unixSocketTimeoutKeepsWorkerAliveAndReconnects() throws Exception { + assertTrue(Config.get().isJdkSocketEnabled()); + + Path socketPath = Files.createTempFile("dd-trace-agent-", ".sock"); + Files.delete(socketPath); + + HealthMetrics healthMetrics = mock(HealthMetrics.class); + AtomicReference failedSendThread = new AtomicReference<>(); + AtomicReference successfulSendThread = new AtomicReference<>(); + CountDownLatch failedSend = new CountDownLatch(1); + CountDownLatch successfulSend = new CountDownLatch(1); + doAnswer( + invocation -> { + failedSendThread.set(Thread.currentThread()); + failedSend.countDown(); + return null; + }) + .when(healthMetrics) + .onFailedSend(anyInt(), anyInt(), any()); + doAnswer( + invocation -> { + successfulSendThread.set(Thread.currentThread()); + successfulSend.countDown(); + return null; + }) + .when(healthMetrics) + .onSend(anyInt(), anyInt(), any()); + + DDAgentFeaturesDiscovery discovery = mock(DDAgentFeaturesDiscovery.class); + when(discovery.getTraceEndpoint()).thenReturn("v0.4/traces"); + + try (MockWebServer server = new MockWebServer(); + DDAgentWriter writer = + DDAgentWriter.builder() + .featureDiscovery(discovery) + .unixDomainSocket(socketPath.toString()) + .timeoutMillis(500) + .monitoring(monitoring) + .healthMetrics(healthMetrics) + .flushIntervalMilliseconds(10) + .flushTimeout(5, TimeUnit.SECONDS) + .build()) { + server.setServerSocketFactory(new UnixDomainServerSocketFactory(socketPath.toFile())); + // Read the first request fully, then withhold the response to trigger a header-read timeout. + server.enqueue(new MockResponse().setSocketPolicy(NO_RESPONSE)); + server.enqueue(new MockResponse().setResponseCode(200).setSocketPolicy(DISCONNECT_AT_END)); + server.start(); + + writer.start(); + + writer.write(createMinimalTrace()); + assertTrue( + failedSend.await(5, TimeUnit.SECONDS), "The periodic flush did not report failure"); + + RecordedRequest failedRequest = server.takeRequest(5, TimeUnit.SECONDS); + assertNotNull(failedRequest); + assertEquals(0, failedRequest.getSequenceNumber()); + assertEquals(1, server.getRequestCount()); + verify(healthMetrics, times(1)).onFailedSend(anyInt(), anyInt(), any()); + + writer.write(createMinimalTrace()); + assertTrue( + successfulSend.await(5, TimeUnit.SECONDS), "The worker did not send the next payload"); + + RecordedRequest successfulRequest = server.takeRequest(5, TimeUnit.SECONDS); + assertNotNull(successfulRequest); + // Sequence numbers are per connection; zero again proves that this used a fresh socket. + assertEquals(0, successfulRequest.getSequenceNumber()); + assertEquals(2, server.getRequestCount()); + verify(healthMetrics, times(1)).onSend(anyInt(), anyInt(), any()); + assertSame(failedSendThread.get(), successfulSendThread.get()); } finally { - agent.close(); + Files.deleteIfExists(socketPath); } } diff --git a/utils/socket-utils/build.gradle.kts b/utils/socket-utils/build.gradle.kts index 7e5e21200e5..75775b6d30e 100644 --- a/utils/socket-utils/build.gradle.kts +++ b/utils/socket-utils/build.gradle.kts @@ -4,6 +4,7 @@ plugins { `java-library` idea id("dd-trace-java.module.internal-library") + id("dd-trace-java.jmh-conventions") } extensions.getByName("tracerJava").withGroovyBuilder { @@ -11,14 +12,23 @@ extensions.getByName("tracerJava").withGroovyBuilder { } dependencies { + add("main_java17CompileOnly", project(":components:annotations")) implementation(project(":components:environment")) implementation(project(":utils:logging-utils")) implementation(libs.slf4j) implementation(libs.jnr.unixsocket) testImplementation(files(sourceSets["main_java17"].output)) + jmhImplementation(files(sourceSets["main_java17"].output)) } -listOf("compileMain_java17Java", "compileTestJava").forEach { +jmh { + jmhVersion = libs.versions.jmh.get() + includeTests = false + resultFormat = "JSON" + failOnError = true +} + +listOf("compileMain_java17Java", "compileTestJava", "compileJmhJava").forEach { tasks.named(it) { // The Java 17 implementation can lift this offset, but compileTestJava must first be split if // the remaining socket tests still need to run on Java 8. diff --git a/utils/socket-utils/gradle.lockfile b/utils/socket-utils/gradle.lockfile index 5c10a22161c..f7cf6691d0f 100644 --- a/utils/socket-utils/gradle.lockfile +++ b/utils/socket-utils/gradle.lockfile @@ -4,35 +4,37 @@ # To regenerate this file, run: ./gradlew :utils:socket-utils:dependencies --write-locks ch.qos.logback:logback-classic:1.2.13=testCompileClasspath,testRuntimeClasspath ch.qos.logback:logback-core:1.2.13=testCompileClasspath,testRuntimeClasspath -com.datadoghq:dd-instrument-java:0.0.5=runtimeClasspath,testRuntimeClasspath -com.datadoghq:dd-javac-plugin-client:0.2.2=runtimeClasspath,testRuntimeClasspath +com.datadoghq:dd-instrument-java:0.0.5=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath +com.datadoghq:dd-javac-plugin-client:0.2.2=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath com.github.javaparser:javaparser-core:3.25.6=codenarc -com.github.jnr:jffi:1.3.15=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-a64asm:1.0.0=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-constants:0.10.4=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-enxio:0.32.20=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-ffi:2.2.19=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-posix:3.1.22=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-unixsocket:0.38.25=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.jnr:jnr-x86asm:1.0.2=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath -com.github.spotbugs:spotbugs-annotations:4.10.4=compileClasspath,spotbugs +com.github.jnr:jffi:1.3.15=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-a64asm:1.0.0=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-constants:0.10.4=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-enxio:0.32.20=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-ffi:2.2.19=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-posix:3.1.22=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-unixsocket:0.38.25=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.jnr:jnr-x86asm:1.0.2=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +com.github.spotbugs:spotbugs-annotations:4.10.4=compileClasspath,jmhCompileClasspath,spotbugs com.github.spotbugs:spotbugs:4.10.4=spotbugs com.github.stephenc.jcip:jcip-annotations:1.0-1=spotbugs -com.google.code.findbugs:jsr305:3.0.2=compileClasspath,spotbugs,testCompileClasspath,testRuntimeClasspath +com.google.code.findbugs:jsr305:3.0.2=compileClasspath,jmhCompileClasspath,spotbugs,testCompileClasspath,testRuntimeClasspath com.google.code.gson:gson:2.14.0=spotbugs com.google.errorprone:error_prone_annotations:2.48.0=spotbugs com.thoughtworks.qdox:qdox:1.12.1=codenarc commons-io:commons-io:2.21.0=spotbugs -de.thetaphi:forbiddenapis:3.11=compileClasspath +de.thetaphi:forbiddenapis:3.11=compileClasspath,jmhCompileClasspath io.leangen.geantyref:geantyref:1.3.16=testRuntimeClasspath jaxen:jaxen:2.0.6=spotbugs net.bytebuddy:byte-buddy-agent:1.12.8=testRuntimeClasspath net.bytebuddy:byte-buddy:1.12.8=testRuntimeClasspath +net.sf.jopt-simple:jopt-simple:5.0.4=jmh,jmhCompileClasspath,jmhRuntimeClasspath net.sf.saxon:Saxon-HE:12.10=spotbugs org.apache.ant:ant-antlr:1.10.14=codenarc org.apache.ant:ant-junit:1.10.14=codenarc org.apache.bcel:bcel:6.12.0=spotbugs org.apache.commons:commons-lang3:3.20.0=spotbugs +org.apache.commons:commons-math3:3.6.1=jmh,jmhCompileClasspath,jmhRuntimeClasspath org.apache.commons:commons-text:1.15.0=spotbugs org.apache.logging.log4j:log4j-api:2.26.1=spotbugs org.apache.logging.log4j:log4j-core:2.26.1=spotbugs @@ -64,29 +66,34 @@ org.junit.platform:junit-platform-launcher:1.14.1=testRuntimeClasspath org.junit:junit-bom:5.14.1=testCompileClasspath,testRuntimeClasspath org.mockito:mockito-core:4.4.0=testRuntimeClasspath org.objenesis:objenesis:3.3=testCompileClasspath,testRuntimeClasspath +org.openjdk.jmh:jmh-core:1.37=jmh,jmhCompileClasspath,jmhRuntimeClasspath +org.openjdk.jmh:jmh-generator-asm:1.37=jmh,jmhCompileClasspath,jmhRuntimeClasspath +org.openjdk.jmh:jmh-generator-bytecode:1.37=jmh,jmhCompileClasspath,jmhRuntimeClasspath +org.openjdk.jmh:jmh-generator-reflection:1.37=jmh,jmhCompileClasspath,jmhRuntimeClasspath org.opentest4j:opentest4j:1.3.0=testCompileClasspath,testRuntimeClasspath org.ow2.asm:asm-analysis:9.10.1=spotbugs -org.ow2.asm:asm-analysis:9.7.1=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +org.ow2.asm:asm-analysis:9.7.1=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath org.ow2.asm:asm-commons:9.10.1=jacocoAnt,spotbugs -org.ow2.asm:asm-commons:9.7.1=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +org.ow2.asm:asm-commons:9.7.1=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath org.ow2.asm:asm-tree:9.10.1=jacocoAnt,spotbugs -org.ow2.asm:asm-tree:9.7.1=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +org.ow2.asm:asm-tree:9.7.1=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath org.ow2.asm:asm-util:9.10.1=spotbugs -org.ow2.asm:asm-util:9.7.1=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +org.ow2.asm:asm-util:9.7.1=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +org.ow2.asm:asm:9.0=jmh org.ow2.asm:asm:9.10.1=jacocoAnt,spotbugs -org.ow2.asm:asm:9.7.1=compileClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath +org.ow2.asm:asm:9.7.1=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath,testCompileClasspath,testRuntimeClasspath org.slf4j:jcl-over-slf4j:1.7.30=testCompileClasspath,testRuntimeClasspath org.slf4j:jul-to-slf4j:1.7.30=testCompileClasspath,testRuntimeClasspath org.slf4j:log4j-over-slf4j:1.7.30=testCompileClasspath,testRuntimeClasspath -org.slf4j:slf4j-api:1.7.30=compileClasspath,runtimeClasspath +org.slf4j:slf4j-api:1.7.30=compileClasspath,jmhCompileClasspath,jmhRuntimeClasspath,runtimeClasspath org.slf4j:slf4j-api:1.7.32=testCompileClasspath,testRuntimeClasspath org.slf4j:slf4j-api:2.0.17=spotbugsSlf4j org.slf4j:slf4j-api:2.0.18=spotbugs org.slf4j:slf4j-simple:2.0.17=spotbugsSlf4j -org.snakeyaml:snakeyaml-engine:2.9=runtimeClasspath,testRuntimeClasspath +org.snakeyaml:snakeyaml-engine:2.9=jmhRuntimeClasspath,runtimeClasspath,testRuntimeClasspath org.spockframework:spock-bom:2.4-groovy-3.0=testCompileClasspath,testRuntimeClasspath org.spockframework:spock-core:2.4-groovy-3.0=testCompileClasspath,testRuntimeClasspath org.tabletest:tabletest-junit:1.2.2=testCompileClasspath,testRuntimeClasspath org.tabletest:tabletest-parser:1.2.1=testCompileClasspath,testRuntimeClasspath org.xmlresolver:xmlresolver:5.3.3=spotbugs -empty=annotationProcessor,main_java17AnnotationProcessor,main_java17CompileClasspath,main_java17RuntimeClasspath,spotbugsPlugins,testAnnotationProcessor +empty=annotationProcessor,jmhAnnotationProcessor,main_java17AnnotationProcessor,main_java17CompileClasspath,main_java17RuntimeClasspath,spotbugsPlugins,testAnnotationProcessor diff --git a/utils/socket-utils/src/jmh/java/datadog/common/socket/TunnelingJdkSocketBenchmark.java b/utils/socket-utils/src/jmh/java/datadog/common/socket/TunnelingJdkSocketBenchmark.java new file mode 100644 index 00000000000..c977ed3db8d --- /dev/null +++ b/utils/socket-utils/src/jmh/java/datadog/common/socket/TunnelingJdkSocketBenchmark.java @@ -0,0 +1,184 @@ +package datadog.common.socket; + +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.net.InetSocketAddress; +import java.net.Socket; +import java.net.StandardProtocolFamily; +import java.net.UnixDomainSocketAddress; +import java.nio.ByteBuffer; +import java.nio.channels.ServerSocketChannel; +import java.nio.channels.SocketChannel; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.Arrays; +import java.util.concurrent.TimeUnit; +import org.openjdk.jmh.annotations.Benchmark; +import org.openjdk.jmh.annotations.BenchmarkMode; +import org.openjdk.jmh.annotations.Fork; +import org.openjdk.jmh.annotations.Measurement; +import org.openjdk.jmh.annotations.Mode; +import org.openjdk.jmh.annotations.OutputTimeUnit; +import org.openjdk.jmh.annotations.Param; +import org.openjdk.jmh.annotations.Scope; +import org.openjdk.jmh.annotations.Setup; +import org.openjdk.jmh.annotations.State; +import org.openjdk.jmh.annotations.TearDown; +import org.openjdk.jmh.annotations.Warmup; +import org.openjdk.jmh.infra.Blackhole; + +/** + * Real UDS connection setup and bulk I/O on reused connections. Requires JDK 17+ on Linux or macOS. + * + *

Run {@code :utils:socket-utils:jmh -PtestJvm=17 -Pjmh.profilers=gc}. Compare revisions on the + * same JVM and host. Round-trip time includes the peer thread and OS scheduling; GC results include + * peer allocations. This does not measure HTTP encoding or timeout/backpressure recovery. + */ +@BenchmarkMode(Mode.AverageTime) +@OutputTimeUnit(TimeUnit.MICROSECONDS) +@Warmup(iterations = 3, time = 1) +@Measurement(iterations = 5, time = 1) +@Fork(2) +public class TunnelingJdkSocketBenchmark { + private static final InetSocketAddress ENDPOINT = new InetSocketAddress("localhost", 0); + + @Benchmark + public void connectAndClose(Listener listener, Blackhole blackhole) throws IOException { + try (TunnelingJdkSocket socket = new TunnelingJdkSocket(listener.path)) { + socket.connect(ENDPOINT); + try (SocketChannel peer = listener.server.accept()) { + blackhole.consume(socket.getInputStream()); + blackhole.consume(socket.getOutputStream()); + } + } + } + + @Benchmark + public byte roundTrip(Connection connection) throws IOException { + connection.output.write(connection.request); + int received = 0; + while (received < connection.response.length) { + int count = + connection.input.read( + connection.response, received, connection.response.length - received); + if (count <= 0) { + throw new IOException("Peer closed or timed out", connection.peerFailure); + } + received += count; + } + return connection.response[received - 1]; + } + + @State(Scope.Thread) + public static class Listener { + private Path path; + private ServerSocketChannel server; + + @Setup + public void setup() throws IOException { + // Gradle's JMH temp directory can exceed the Unix-domain socket path limit. + path = Files.createTempFile(Paths.get("/tmp"), "uds-jmh-", ".sock"); + Files.delete(path); + server = ServerSocketChannel.open(StandardProtocolFamily.UNIX); + try { + server.bind(UnixDomainSocketAddress.of(path)); + } catch (IOException | RuntimeException e) { + close(); + throw e; + } + } + + @TearDown + public void close() throws IOException { + try { + server.close(); + } finally { + Files.deleteIfExists(path); + } + } + } + + @State(Scope.Thread) + public static class Connection { + @Param({"256", "8192", "65536"}) + public int bytes; + + private Socket socket; + private SocketChannel peer; + private InputStream input; + private OutputStream output; + private byte[] request; + private byte[] response; + private Thread responder; + private volatile IOException peerFailure; + + @Setup + public void setup(Listener listener) throws IOException { + request = new byte[bytes]; + Arrays.fill(request, (byte) 1); + response = new byte[bytes]; + socket = new TunnelingJdkSocket(listener.path); + try { + socket.connect(ENDPOINT); + peer = listener.server.accept(); + socket.setSoTimeout(5000); + input = socket.getInputStream(); + output = socket.getOutputStream(); + } catch (IOException | RuntimeException e) { + socket.close(); + if (peer != null) { + peer.close(); + } + throw e; + } + responder = new Thread(this::echo, "uds-jmh-peer"); + responder.setDaemon(true); + responder.start(); + } + + private void echo() { + ByteBuffer buffer = ByteBuffer.allocate(bytes); + try { + while (true) { + buffer.clear(); + while (buffer.hasRemaining()) { + if (peer.read(buffer) == -1) { + return; + } + } + buffer.flip(); + while (buffer.hasRemaining()) { + peer.write(buffer); + } + } + } catch (IOException e) { + if (peer.isOpen()) { + peerFailure = e; + try { + socket.close(); + } catch (IOException closeFailure) { + e.addSuppressed(closeFailure); + } + } + } + } + + @TearDown + public void close() throws IOException, InterruptedException { + try { + peer.close(); + } finally { + socket.close(); + } + responder.join(TimeUnit.SECONDS.toMillis(5)); + if (responder.isAlive()) { + throw new IllegalStateException("Peer did not terminate"); + } + if (peerFailure != null) { + throw peerFailure; + } + } + } +} diff --git a/utils/socket-utils/src/main/java/datadog/common/socket/UnixDomainSocketFactory.java b/utils/socket-utils/src/main/java/datadog/common/socket/UnixDomainSocketFactory.java index 1df84896c56..09d03d1325d 100644 --- a/utils/socket-utils/src/main/java/datadog/common/socket/UnixDomainSocketFactory.java +++ b/utils/socket-utils/src/main/java/datadog/common/socket/UnixDomainSocketFactory.java @@ -45,7 +45,7 @@ public Socket createSocket() throws IOException { if (this.useJdkUdsSocket) { try { return new TunnelingJdkSocket(this.path.toPath()); - } catch (Throwable ignore) { + } catch (IOException | UnsupportedOperationException ignore) { // fall back to jnr-unixsocket library } } diff --git a/utils/socket-utils/src/main/java17/datadog/common/socket/TunnelingJdkSocket.java b/utils/socket-utils/src/main/java17/datadog/common/socket/TunnelingJdkSocket.java index fb25228e0f6..a0487a4e944 100644 --- a/utils/socket-utils/src/main/java17/datadog/common/socket/TunnelingJdkSocket.java +++ b/utils/socket-utils/src/main/java17/datadog/common/socket/TunnelingJdkSocket.java @@ -1,5 +1,6 @@ package datadog.common.socket; +import datadog.trace.api.internal.VisibleForTesting; import java.io.IOException; import java.io.InputStream; import java.io.OutputStream; @@ -8,8 +9,12 @@ import java.net.Socket; import java.net.SocketAddress; import java.net.SocketException; +import java.net.StandardProtocolFamily; +import java.net.StandardSocketOptions; import java.net.UnixDomainSocketAddress; import java.nio.ByteBuffer; +import java.nio.channels.CancelledKeyException; +import java.nio.channels.ClosedSelectorException; import java.nio.channels.SelectionKey; import java.nio.channels.Selector; import java.nio.channels.SocketChannel; @@ -25,34 +30,30 @@ * 16. */ final class TunnelingJdkSocket extends Socket { - private final SocketAddress unixSocketAddress; - private InetSocketAddress inetSocketAddress; + private final UnixDomainSocketAddress unixSocketAddress; + private final SocketChannel unixSocketChannel; - private SocketChannel unixSocketChannel; - private Selector selector; + private volatile InetSocketAddress inetSocketAddress; + @VisibleForTesting Selector selector; - private int timeout; - private boolean shutIn; - private boolean shutOut; - private boolean closed; + private volatile int timeout; + private volatile boolean shutIn; + private volatile boolean shutOut; + private volatile boolean closed; static final int DEFAULT_BUFFER_SIZE = 8192; // Indicate that the buffer size is not set by initializing to -1 private int sendBufferSize = -1; private int receiveBufferSize = -1; - TunnelingJdkSocket(final Path path) { + TunnelingJdkSocket(final Path path) throws IOException, UnsupportedOperationException { this.unixSocketAddress = UnixDomainSocketAddress.of(path); - } - - TunnelingJdkSocket(final Path path, final InetSocketAddress address) { - this(path); - inetSocketAddress = address; + this.unixSocketChannel = SocketChannel.open(StandardProtocolFamily.UNIX); } @Override public boolean isConnected() { - return null != unixSocketChannel; + return inetSocketAddress != null; } @Override @@ -71,7 +72,7 @@ public boolean isClosed() { } @Override - public synchronized void setSoTimeout(int timeout) throws SocketException { + public void setSoTimeout(int timeout) throws SocketException { if (isClosed()) { throw new SocketException("Socket is closed"); } @@ -82,7 +83,7 @@ public synchronized void setSoTimeout(int timeout) throws SocketException { } @Override - public synchronized int getSoTimeout() throws SocketException { + public int getSoTimeout() throws SocketException { if (isClosed()) { throw new SocketException("Socket is closed"); } @@ -91,17 +92,7 @@ public synchronized int getSoTimeout() throws SocketException { @Override public void connect(final SocketAddress endpoint) throws IOException { - if (endpoint == null) { - throw new IllegalArgumentException("Endpoint cannot be null"); - } - if (isClosed()) { - throw new SocketException("Socket is closed"); - } - if (isConnected()) { - throw new SocketException("Socket is already connected"); - } - inetSocketAddress = (InetSocketAddress) endpoint; - unixSocketChannel = SocketChannel.open(unixSocketAddress); + connect(endpoint, 0); } // `timeout` is intentionally ignored here, like in the jnr-unixsocket implementation. @@ -121,8 +112,14 @@ public void connect(final SocketAddress endpoint, final int timeout) throws IOEx if (isConnected()) { throw new SocketException("Socket is already connected"); } - inetSocketAddress = (InetSocketAddress) endpoint; - unixSocketChannel = SocketChannel.open(unixSocketAddress); + InetSocketAddress inetSocketAddress = (InetSocketAddress) endpoint; + try { + unixSocketChannel.connect(unixSocketAddress); + this.inetSocketAddress = inetSocketAddress; + } catch (IOException e) { + close(); + throw e; + } } @Override @@ -140,7 +137,7 @@ public void setSendBufferSize(int size) throws SocketException { } sendBufferSize = size; try { - unixSocketChannel.setOption(java.net.StandardSocketOptions.SO_SNDBUF, size); + unixSocketChannel.setOption(StandardSocketOptions.SO_SNDBUF, size); } catch (IOException e) { SocketException se = new SocketException("Failed to set send buffer size socket option"); se.initCause(e); @@ -169,7 +166,7 @@ public void setReceiveBufferSize(int size) throws SocketException { } receiveBufferSize = size; try { - unixSocketChannel.setOption(java.net.StandardSocketOptions.SO_RCVBUF, size); + unixSocketChannel.setOption(StandardSocketOptions.SO_RCVBUF, size); } catch (IOException e) { SocketException se = new SocketException("Failed to set receive buffer size socket option"); se.initCause(e); @@ -200,69 +197,98 @@ public int getStreamBufferSize() throws SocketException { @Override public InputStream getInputStream() throws IOException { - if (isClosed()) { - throw new SocketException("Socket is closed"); - } - if (!isConnected()) { - throw new SocketException("Socket is not connected"); - } - if (isInputShutdown()) { - throw new SocketException("Socket input is shutdown"); - } - - if (selector == null) { - selector = Selector.open(); + // Configuring the channel can wait for a blocked write. Let close() interrupt that write. + if (!isClosed() && isConnected() && !isInputShutdown() && unixSocketChannel.isBlocking()) { unixSocketChannel.configureBlocking(false); - unixSocketChannel.register(selector, SelectionKey.OP_READ); } - - return new InputStream() { - private final ByteBuffer buffer = ByteBuffer.allocate(getStreamBufferSize()); - - @Override - public int read() throws IOException { - byte[] nextByte = new byte[1]; - return (read(nextByte, 0, 1) == -1) ? -1 : (nextByte[0] & 0xFF); + // Serialize validation and selector publication with close() so close cannot miss a selector + // that is still being initialized. + synchronized (this) { + if (isClosed()) { + throw new SocketException("Socket is closed"); + } + if (!isConnected()) { + throw new SocketException("Socket is not connected"); + } + if (isInputShutdown()) { + throw new SocketException("Socket input is shutdown"); } - @Override - public int read(byte[] b, int off, int len) throws IOException { - if (isInputShutdown()) { - return -1; + if (selector == null) { + Selector newSelector = Selector.open(); + try { + unixSocketChannel.register(newSelector, SelectionKey.OP_READ); + selector = newSelector; + } catch (IOException | RuntimeException e) { + try { + newSelector.close(); + } catch (IOException closeException) { + e.addSuppressed(closeException); + } + throw e; } - buffer.clear(); + } + + final Selector readSelector = selector; + return new InputStream() { + private final ByteBuffer buffer = ByteBuffer.allocate(getStreamBufferSize()); - int readyChannels = selector.select(timeout); - if (readyChannels == 0) { - return 0; + @Override + public int read() throws IOException { + byte[] nextByte = new byte[1]; + return (read(nextByte, 0, 1) == -1) ? -1 : (nextByte[0] & 0xFF); } - Set selectedKeys = selector.selectedKeys(); - synchronized (selectedKeys) { - Iterator keyIterator = selectedKeys.iterator(); - while (keyIterator.hasNext()) { - SelectionKey key = keyIterator.next(); - keyIterator.remove(); - if (key.isReadable()) { - int r = unixSocketChannel.read(buffer); - if (r == -1) { - return -1; + @Override + public int read(byte[] b, int off, int len) throws IOException { + if (isInputShutdown()) { + return -1; + } + buffer.clear(); + + try { + int readyChannels = readSelector.select(timeout); + if (readyChannels == 0) { + if (isClosed() || !readSelector.isOpen()) { + throw new SocketException("Socket is closed"); } - buffer.flip(); - len = Math.min(r, len); - buffer.get(b, off, len); - return len; + return 0; } + + Set selectedKeys = readSelector.selectedKeys(); + // Multiple input streams share this selector, so serialize iteration and removal from + // its non-thread-safe selected-key set. + synchronized (selectedKeys) { + Iterator keyIterator = selectedKeys.iterator(); + while (keyIterator.hasNext()) { + SelectionKey key = keyIterator.next(); + keyIterator.remove(); + if (key.isReadable()) { + int r = unixSocketChannel.read(buffer); + if (r == -1) { + return -1; + } + buffer.flip(); + len = Math.min(r, len); + buffer.get(b, off, len); + return len; + } + } + } + return 0; + } catch (ClosedSelectorException | CancelledKeyException e) { + SocketException socketException = new SocketException("Socket is closed"); + socketException.initCause(e); + throw socketException; } } - return 0; - } - @Override - public void close() throws IOException { - TunnelingJdkSocket.this.close(); - } - }; + @Override + public void close() throws IOException { + TunnelingJdkSocket.this.close(); + } + }; + } } @Override @@ -304,32 +330,38 @@ public void close() throws IOException { @Override public void shutdownInput() throws IOException { - if (isClosed()) { - throw new SocketException("Socket is closed"); - } - if (!isConnected()) { - throw new SocketException("Socket is not connected"); - } - if (isInputShutdown()) { - throw new SocketException("Socket input is already shutdown"); + // Keep validation, channel shutdown, and state publication atomic with close(). + synchronized (this) { + if (isClosed()) { + throw new SocketException("Socket is closed"); + } + if (!isConnected()) { + throw new SocketException("Socket is not connected"); + } + if (isInputShutdown()) { + throw new SocketException("Socket input is already shutdown"); + } + unixSocketChannel.shutdownInput(); + shutIn = true; } - unixSocketChannel.shutdownInput(); - shutIn = true; } @Override public void shutdownOutput() throws IOException { - if (isClosed()) { - throw new SocketException("Socket is closed"); - } - if (!isConnected()) { - throw new SocketException("Socket is not connected"); - } - if (isOutputShutdown()) { - throw new SocketException("Socket output is already shutdown"); + // Keep validation, channel shutdown, and state publication atomic with close(). + synchronized (this) { + if (isClosed()) { + throw new SocketException("Socket is closed"); + } + if (!isConnected()) { + throw new SocketException("Socket is not connected"); + } + if (isOutputShutdown()) { + throw new SocketException("Socket output is already shutdown"); + } + unixSocketChannel.shutdownOutput(); + shutOut = true; } - unixSocketChannel.shutdownOutput(); - shutOut = true; } @Override @@ -341,36 +373,29 @@ public InetAddress getInetAddress() { } @Override - public void close() throws IOException { - if (isClosed()) { - return; - } - // Ignore possible exceptions so that we continue closing the socket - try { - if (!isInputShutdown()) { - shutdownInput(); + public void close() { + Selector currentSelector; + // Publish the terminal state and snapshot the selector atomically with selector creation and + // half-close operations. The resources are closed after releasing this monitor. + synchronized (this) { + if (isClosed()) { + return; } - } catch (IOException e) { - } - try { - if (!isOutputShutdown()) { - shutdownOutput(); - } - } catch (IOException e) { + shutIn = true; + shutOut = true; + closed = true; + currentSelector = selector; } + // Ignore possible exceptions so that we continue closing the socket try { - if (selector != null) { - selector.close(); - selector = null; + if (currentSelector != null) { + currentSelector.close(); } - } catch (IOException e) { + } catch (IOException ignored) { } try { - if (unixSocketChannel != null) { - unixSocketChannel.close(); - } - } catch (IOException e) { + unixSocketChannel.close(); + } catch (IOException ignored) { } - closed = true; } } diff --git a/utils/socket-utils/src/test/java/datadog/common/socket/TunnelingJdkSocketTest.java b/utils/socket-utils/src/test/java/datadog/common/socket/TunnelingJdkSocketTest.java index 1c7c7ec7a19..ef6741e35bb 100644 --- a/utils/socket-utils/src/test/java/datadog/common/socket/TunnelingJdkSocketTest.java +++ b/utils/socket-utils/src/test/java/datadog/common/socket/TunnelingJdkSocketTest.java @@ -2,6 +2,8 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTimeoutPreemptively; import static org.junit.jupiter.api.Assertions.assertTrue; @@ -11,32 +13,65 @@ import java.io.IOException; import java.io.InputStream; import java.io.OutputStream; -import java.lang.management.ManagementFactory; import java.net.InetSocketAddress; import java.net.SocketException; import java.net.StandardProtocolFamily; import java.net.UnixDomainSocketAddress; +import java.nio.channels.CancelledKeyException; +import java.nio.channels.ClosedSelectorException; +import java.nio.channels.SelectableChannel; +import java.nio.channels.SelectionKey; +import java.nio.channels.Selector; import java.nio.channels.ServerSocketChannel; import java.nio.channels.SocketChannel; +import java.nio.channels.spi.SelectorProvider; import java.nio.file.Files; import java.nio.file.Path; import java.time.Duration; -import java.util.concurrent.atomic.AtomicBoolean; +import java.util.Collections; +import java.util.HashSet; +import java.util.Set; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Assumptions; +import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledForJreRange; +@EnabledForJreRange(min = JAVA_16) public class TunnelingJdkSocketTest { - private static final AtomicBoolean isServerRunning = new AtomicBoolean(false); + private TestUnixSocketServer server; + + @BeforeAll + static void assumeUnixDomainSocketsAreSupported() { + Assumptions.assumeTrue(udsSupported()); + } + + @AfterEach + void closeServer() throws Exception { + if (server != null) { + server.close(); + server = null; + } + } @Test - @EnabledForJreRange(min = JAVA_16) public void testSocketConnectAndClose() throws Exception { Path socketPath = getSocketPath(); UnixDomainSocketAddress socketAddress = UnixDomainSocketAddress.of(socketPath); startServer(socketAddress); TunnelingJdkSocket clientSocket = new TunnelingJdkSocket(socketPath); + assertNotNull(clientSocket.getChannel()); + assertTrue(clientSocket.getChannel().isOpen()); assertFalse(clientSocket.isConnected()); assertFalse(clientSocket.isClosed()); @@ -55,6 +90,7 @@ public void testSocketConnectAndClose() throws Exception { assertTrue(clientSocket.isConnected()); assertTrue(clientSocket.isClosed()); + assertFalse(clientSocket.getChannel().isOpen()); assertTrue(clientSocket.isInputShutdown()); assertTrue(clientSocket.isOutputShutdown()); assertEquals(-1, inputStream.read()); @@ -62,12 +98,26 @@ public void testSocketConnectAndClose() throws Exception { assertThrows(SocketException.class, clientSocket::getInputStream); assertThrows(SocketException.class, clientSocket::getOutputStream); clientSocket.close(); + } + + @Test + public void testConnectFailureClosesSocket() throws Exception { + Path missingSocketPath = getSocketPath(); + TunnelingJdkSocket clientSocket = new TunnelingJdkSocket(missingSocketPath); + + assertThrows( + IOException.class, () -> clientSocket.connect(new InetSocketAddress("localhost", 0))); - isServerRunning.set(false); + assertFalse(clientSocket.isConnected()); + assertTrue(clientSocket.isClosed()); + assertTrue(clientSocket.isInputShutdown()); + assertTrue(clientSocket.isOutputShutdown()); + assertFalse(clientSocket.getChannel().isOpen()); + assertThrows( + SocketException.class, () -> clientSocket.connect(new InetSocketAddress("localhost", 0))); } @Test - @EnabledForJreRange(min = JAVA_16) public void testInputStreamClose() throws Exception { TunnelingJdkSocket clientSocket = createClient(); InputStream inputStream = clientSocket.getInputStream(); @@ -86,12 +136,9 @@ public void testInputStreamClose() throws Exception { assertThrows(IOException.class, () -> outputStream.write(1)); assertThrows(SocketException.class, clientSocket::getInputStream); assertThrows(SocketException.class, clientSocket::getOutputStream); - - isServerRunning.set(false); } @Test - @EnabledForJreRange(min = JAVA_16) public void testOutputStreamClose() throws Exception { TunnelingJdkSocket clientSocket = createClient(); InputStream inputStream = clientSocket.getInputStream(); @@ -110,12 +157,9 @@ public void testOutputStreamClose() throws Exception { assertThrows(IOException.class, () -> outputStream.write(1)); assertThrows(SocketException.class, clientSocket::getInputStream); assertThrows(SocketException.class, clientSocket::getOutputStream); - - isServerRunning.set(false); } @Test - @EnabledForJreRange(min = JAVA_16) public void testTimeout() throws Exception { TunnelingJdkSocket clientSocket = createClient(); InputStream inputStream = clientSocket.getInputStream(); @@ -155,12 +199,9 @@ public void testTimeout() throws Exception { clientSocket.close(); assertThrows(SocketException.class, () -> clientSocket.setSoTimeout(testTimeout)); assertThrows(SocketException.class, clientSocket::getSoTimeout); - - isServerRunning.set(false); } @Test - @EnabledForJreRange(min = JAVA_16) public void testBufferSizes() throws Exception { TunnelingJdkSocket clientSocket = createClient(); @@ -191,33 +232,175 @@ public void testBufferSizes() throws Exception { assertThrows(SocketException.class, clientSocket::getSendBufferSize); assertThrows(SocketException.class, clientSocket::getReceiveBufferSize); assertThrows(SocketException.class, clientSocket::getStreamBufferSize); - - isServerRunning.set(false); } @Test - @EnabledForJreRange(min = JAVA_16) public void testFileDescriptorLeak() throws Exception { long initialCount = getFileDescriptorCount(); TunnelingJdkSocket clientSocket = createClient(); for (int i = 0; i < 100; i++) { + @SuppressWarnings("unused") InputStream inputStream = clientSocket.getInputStream(); long currentCount = getFileDescriptorCount(); assertTrue(currentCount <= initialCount + 7); } clientSocket.close(); - isServerRunning.set(false); + closeServer(); long finalCount = getFileDescriptorCount(); assertTrue(finalCount <= initialCount + 3); } + @Test + public void testClosedSelectorIsReportedAsSocketException() throws Exception { + try (TunnelingJdkSocket clientSocket = createClient()) { + InputStream inputStream = clientSocket.getInputStream(); + clientSocket.selector.close(); + + SocketException exception = assertThrows(SocketException.class, inputStream::read); + + assertInstanceOf(ClosedSelectorException.class, exception.getCause()); + } + } + + @Test + public void testSelectorClosedBetweenSelectAndSelectedKeysIsReportedAsSocketException() + throws Exception { + try (TunnelingJdkSocket clientSocket = createClient()) { + clientSocket.selector = new ClosedAfterSelectSelector(); + assertTrue( + clientSocket.getChannel().isBlocking(), + "The channel should be blocking before getInputStream()"); + InputStream inputStream = clientSocket.getInputStream(); + assertFalse( + clientSocket.getChannel().isBlocking(), + "getInputStream() should switch the channel to nonblocking mode"); + + SocketException exception = assertThrows(SocketException.class, inputStream::read); + + assertInstanceOf(ClosedSelectorException.class, exception.getCause()); + } + } + + @Test + public void testCancelledKeyIsReportedAsSocketException() throws Exception { + try (TunnelingJdkSocket clientSocket = createClient()) { + clientSocket.selector = new CancelledKeyAfterSelectionSelector(); + InputStream inputStream = clientSocket.getInputStream(); + + SocketException exception = assertThrows(SocketException.class, inputStream::read); + + assertInstanceOf(CancelledKeyException.class, exception.getCause()); + } + } + + @Test + public void testAsynchronousCloseInterruptsBlockedReadWithIOException() throws Exception { + TunnelingJdkSocket clientSocket = createClient(); + Thread reader = null; + try { + BlockingCloseSelector blockingCloseSelector = new BlockingCloseSelector(); + clientSocket.selector = blockingCloseSelector; + InputStream inputStream = clientSocket.getInputStream(); + AtomicReference readFailure = new AtomicReference<>(); + + reader = + new Thread( + () -> { + try { + inputStream.read(); + } catch (Throwable t) { + readFailure.set(t); + } + }, + "tunneling-jdk-socket-reader"); + reader.setDaemon(true); + reader.start(); + + assertTrue( + blockingCloseSelector.awaitSelectStarted(5, TimeUnit.SECONDS), + "The reader did not block in Selector.select"); + clientSocket.close(); + reader.join(TimeUnit.SECONDS.toMillis(5)); + + assertFalse(reader.isAlive(), "The blocked read did not terminate after close"); + Throwable failure = readFailure.get(); + assertNotNull(failure, "The blocked read should fail when the socket is closed"); + assertInstanceOf( + IOException.class, + failure, + () -> "Expected an IOException, but got " + failure.getClass().getName()); + } finally { + clientSocket.close(); + if (reader != null && reader.isAlive()) { + reader.interrupt(); + reader.join(TimeUnit.SECONDS.toMillis(5)); + } + } + } + + @Test + public void testCloseInterruptsWriteWhileInputStreamIsBeingCreated() throws Exception { + TunnelingJdkSocket clientSocket = createClient(); + ExecutorService executor = Executors.newFixedThreadPool(3); + try { + clientSocket.setSendBufferSize(4096); + OutputStream output = clientSocket.getOutputStream(); + CountDownLatch writing = new CountDownLatch(1); + Future writer = + executor.submit( + () -> { + writing.countDown(); + // The server accepts but does not read, so this exceeds the socket's send buffer. + output.write(new byte[1024 * 1024]); + return null; + }); + assertTrue(writing.await(5, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> writer.get(100, TimeUnit.MILLISECONDS)); + + AtomicReference inputThread = new AtomicReference<>(); + Future input = + executor.submit( + () -> { + inputThread.set(Thread.currentThread()); + return clientSocket.getInputStream(); + }); + // Wait until input setup is waiting for the channel's blocked writer before closing. + assertTimeoutPreemptively( + Duration.ofSeconds(5), + () -> { + while (inputThread.get() == null + || inputThread.get().getState() != Thread.State.WAITING) { + assertFalse(input.isDone(), "Input setup should wait for the blocked writer"); + Thread.sleep(1); + } + }); + + executor.submit(clientSocket::close).get(5, TimeUnit.SECONDS); + + assertTrue(clientSocket.isClosed()); + assertFalse(clientSocket.getChannel().isOpen()); + assertInstanceOf( + IOException.class, + assertThrows(ExecutionException.class, () -> writer.get(5, TimeUnit.SECONDS)).getCause()); + assertInstanceOf( + IOException.class, + assertThrows(ExecutionException.class, () -> input.get(5, TimeUnit.SECONDS)).getCause()); + } finally { + // Also release the writer if the regression prevents socket.close() from taking its monitor. + clientSocket.getChannel().close(); + executor.shutdownNow(); + assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + clientSocket.close(); + } + } + private long getFileDescriptorCount() { try { - Process process = Runtime.getRuntime().exec("lsof -p " + getPid()); + Process process = Runtime.getRuntime().exec("lsof -p " + ProcessHandle.current().pid()); int count = 0; try (java.io.BufferedReader reader = new java.io.BufferedReader(new java.io.InputStreamReader(process.getInputStream()))) { @@ -231,41 +414,8 @@ private long getFileDescriptorCount() { } } - private String getPid() { - return ManagementFactory.getRuntimeMXBean().getName().split("@")[0]; - } - - private static void startServer(UnixDomainSocketAddress socketAddress) { - Thread serverThread = - new Thread( - () -> { - try (ServerSocketChannel serverChannel = - ServerSocketChannel.open(StandardProtocolFamily.UNIX)) { - serverChannel.bind(socketAddress); - isServerRunning.set(true); - - synchronized (isServerRunning) { - isServerRunning.notifyAll(); - } - - while (isServerRunning.get()) { - SocketChannel clientChannel = serverChannel.accept(); - } - } catch (IOException e) { - throw new RuntimeException(e); - } - }); - serverThread.start(); - - synchronized (isServerRunning) { - while (!isServerRunning.get()) { - try { - isServerRunning.wait(); - } catch (InterruptedException e) { - throw new RuntimeException(e); - } - } - } + private void startServer(UnixDomainSocketAddress socketAddress) throws IOException { + server = new TestUnixSocketServer(socketAddress); } private Path getSocketPath() throws IOException { @@ -283,4 +433,286 @@ private TunnelingJdkSocket createClient() throws IOException { clientSocket.connect(new InetSocketAddress("localhost", 0)); return clientSocket; } + + private static final class TestUnixSocketServer implements AutoCloseable { + private final Path socketPath; + private final ServerSocketChannel serverChannel; + private final AtomicReference acceptedChannel = new AtomicReference<>(); + private final AtomicReference serverFailure = new AtomicReference<>(); + private final Thread serverThread; + + private TestUnixSocketServer(UnixDomainSocketAddress socketAddress) throws IOException { + socketPath = socketAddress.getPath(); + serverChannel = ServerSocketChannel.open(StandardProtocolFamily.UNIX); + boolean bound = false; + try { + serverChannel.bind(socketAddress); + bound = true; + } finally { + if (!bound) { + serverChannel.close(); + } + } + + serverThread = + new Thread( + () -> { + try { + acceptedChannel.set(serverChannel.accept()); + } catch (IOException e) { + if (serverChannel.isOpen()) { + serverFailure.set(e); + } + } + }, + "tunneling-jdk-socket-test-server"); + serverThread.setDaemon(true); + serverThread.start(); + } + + @Override + public void close() throws Exception { + serverChannel.close(); + serverThread.join(TimeUnit.SECONDS.toMillis(5)); + if (serverThread.isAlive()) { + serverThread.interrupt(); + throw new AssertionError("The Unix-domain test server did not terminate"); + } + + SocketChannel clientChannel = acceptedChannel.get(); + if (clientChannel != null) { + clientChannel.close(); + } + Files.deleteIfExists(socketPath); + + Throwable failure = serverFailure.get(); + if (failure != null) { + throw new AssertionError("The Unix-domain test server failed", failure); + } + } + } + + private abstract static class SelectorAdapter extends Selector { + private volatile boolean open = true; + + @Override + public final boolean isOpen() { + return open; + } + + @Override + public final SelectorProvider provider() { + return SelectorProvider.provider(); + } + + @Override + public final int selectNow() { + return doSelect(); + } + + @Override + public final int select(long timeout) { + return doSelect(); + } + + @Override + public final int select() { + return doSelect(); + } + + @Override + public final Selector wakeup() { + return this; + } + + @Override + public final void close() { + if (open) { + open = false; + onClose(); + } + } + + abstract int doSelect(); + + void onClose() {} + } + + /** + * Models another thread closing a selector while a socket read is blocked: + * + *

    + *
  1. Signals when {@code select()} is entered. + *
  2. Blocks until another thread closes the selector. + *
  3. Throws {@link ClosedSelectorException} after closure. + *
  4. Lets the test verify that the reader receives an {@link IOException} and terminates. + *
+ */ + private static final class BlockingCloseSelector extends SelectorAdapter { + private final CountDownLatch selectStarted = new CountDownLatch(1); + private final CountDownLatch closed = new CountDownLatch(1); + + @Override + public Set keys() { + return Collections.emptySet(); + } + + @Override + public Set selectedKeys() { + return Collections.emptySet(); + } + + @Override + int doSelect() { + selectStarted.countDown(); + try { + if (!closed.await(5, TimeUnit.SECONDS)) { + throw new AssertionError("Selector was not closed"); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new AssertionError("Interrupted while waiting for the selector to close", e); + } + throw new ClosedSelectorException(); + } + + @Override + void onClose() { + closed.countDown(); + } + + boolean awaitSelectStarted(long timeout, TimeUnit unit) throws InterruptedException { + return selectStarted.await(timeout, unit); + } + } + + /** + * Models a selector closing between selection and selected-key processing: + * + *
    + *
  1. Closes itself during {@code select()} and reports one ready channel. + *
  2. Throws {@link ClosedSelectorException} when the selected keys are requested. + *
  3. Lets the test verify that the exception is reported as a {@link SocketException}. + *
+ */ + private static final class ClosedAfterSelectSelector extends SelectorAdapter { + @Override + public Set keys() { + return Collections.emptySet(); + } + + @Override + public Set selectedKeys() { + if (!isOpen()) { + throw new ClosedSelectorException(); + } + return Collections.emptySet(); + } + + @Override + int doSelect() { + close(); + return 1; + } + } + + /** + * Models a key being canceled after selection, before selected-key processing: + * + *
    + *
  1. Reports one ready channel from {@code select()}. + *
  2. Cancels the selected key before returning it. + *
  3. Throws {@link CancelledKeyException} when the read checks whether the key is readable. + *
  4. Lets the test verify that the exception is reported as a {@link SocketException}. + *
+ */ + private static final class CancelledKeyAfterSelectionSelector extends SelectorAdapter { + private final CancelledSelectionKey key = new CancelledSelectionKey(this); + private final Set selectedKeys = new HashSet<>(Collections.singleton(key)); + + @Override + public Set keys() { + return selectedKeys; + } + + @Override + public Set selectedKeys() { + key.cancel(); + return selectedKeys; + } + + @Override + int doSelect() { + return 1; + } + + private static final class CancelledSelectionKey extends SelectionKey { + private final Selector selector; + private boolean valid = true; + + private CancelledSelectionKey(Selector selector) { + this.selector = selector; + } + + @Override + public SelectableChannel channel() { + return null; + } + + @Override + public Selector selector() { + return selector; + } + + @Override + public boolean isValid() { + return valid; + } + + @Override + public void cancel() { + valid = false; + } + + @Override + public int interestOps() { + return OP_READ; + } + + @Override + public SelectionKey interestOps(int ops) { + return this; + } + + @Override + public int readyOps() { + if (!valid) { + throw new CancelledKeyException(); + } + return OP_READ; + } + } + } + + private static boolean udsSupported() { + Path socketPath = null; + try { + socketPath = Files.createTempFile("testSocketSupport", null); + Files.delete(socketPath); + try (ServerSocketChannel serverChannel = + ServerSocketChannel.open(StandardProtocolFamily.UNIX)) { + serverChannel.bind(UnixDomainSocketAddress.of(socketPath)); + } + return true; + } catch (IOException | UnsupportedOperationException e) { + return false; + } finally { + if (socketPath != null) { + try { + Files.deleteIfExists(socketPath); + } catch (IOException ignored) { + } + } + } + } }