From a06df8d752841f63bd225ea909f31b0c44624a91 Mon Sep 17 00:00:00 2001 From: Andrea Marziali Date: Tue, 29 Sep 2026 12:22:55 +0200 Subject: [PATCH 1/3] Fix Scala firstCompletedOf continuation cleanup --- .../scala/ScalaPromiseContinuationHelper.java | 31 ++++ .../scala-promise-2.13/build.gradle | 17 +++ ...DefaultPromiseCallbackInstrumentation.java | 137 ++++++++++++++++++ .../ScalaPromiseUnregisterModule.java | 51 +++++++ .../groovy/ScalaFirstCompletedOfTest.groovy | 74 ++++++++++ .../scala/FirstCompletedOfUtils.scala | 32 ++++ 6 files changed, 342 insertions(+) create mode 100644 dd-java-agent/agent-bootstrap/src/main/java/datadog/trace/bootstrap/instrumentation/scala/ScalaPromiseContinuationHelper.java create mode 100644 dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java create mode 100644 dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java create mode 100644 dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy create mode 100644 dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala diff --git a/dd-java-agent/agent-bootstrap/src/main/java/datadog/trace/bootstrap/instrumentation/scala/ScalaPromiseContinuationHelper.java b/dd-java-agent/agent-bootstrap/src/main/java/datadog/trace/bootstrap/instrumentation/scala/ScalaPromiseContinuationHelper.java new file mode 100644 index 00000000000..2f46c6cec40 --- /dev/null +++ b/dd-java-agent/agent-bootstrap/src/main/java/datadog/trace/bootstrap/instrumentation/scala/ScalaPromiseContinuationHelper.java @@ -0,0 +1,31 @@ +package datadog.trace.bootstrap.instrumentation.scala; + +import static datadog.trace.bootstrap.FieldBackedContextStores.getContextStore; + +import datadog.trace.bootstrap.instrumentation.java.concurrent.State; +import java.util.concurrent.atomic.AtomicReference; + +public final class ScalaPromiseContinuationHelper { + private ScalaPromiseContinuationHelper() {} + + /** + * Atomically replace a callback collection and release a continuation removed with its callback. + * + *

The identity check distinguishes a successful removal from a successful no-op CAS. + */ + public static boolean compareAndSetAndCancel( + AtomicReference reference, + Object expected, + Object replacement, + int contextStoreId, + Object callback) { + boolean updated = reference.compareAndSet(expected, replacement); + if (updated && expected != replacement) { + State state = (State) getContextStore(contextStoreId).get(callback); + if (null != state) { + state.closeContinuation(); + } + } + return updated; + } +} diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle index 6f6515ece54..f08fa186d0b 100644 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle @@ -5,11 +5,19 @@ plugins { muzzle { pass { + name = 'scala-promise-2.13' group = 'org.scala-lang' module = "scala-library" versions = "[2.13,)" assertInverse = true } + pass { + name = 'scala-promise-unregister' + group = 'org.scala-lang' + module = "scala-library" + versions = "[2.13.17,)" + assertInverse = true + } } // Keep Spotless scoped to this project; shared source sets are pulled in from scala-promise-common. @@ -28,6 +36,7 @@ evaluationDependsOn(':dd-java-agent:instrumentation:scala:scala-promise:scala-pr addTestSuiteForDir('latestDepTest', 'test') addTestSuiteExtendingForDir('latestDepForkedTest', 'latestDepTest', 'forkedTest') +addTestSuite('scala21317Test') tasks.named("latestDepTest", Test) { finalizedBy 'latestDepForkedTest' @@ -37,10 +46,12 @@ sourceSets { test.groovy.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.groovy test.groovy.srcDir sourceSets.latestDepForkedTest.groovy latestDepTest.groovy.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.groovy + latestDepTest.groovy.srcDir 'src/scala21317Test/groovy' latestDepForkedTest.groovy.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.groovy test.scala.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.scala latestDepTest.scala.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.scala + latestDepTest.scala.srcDir 'src/scala21317Test/scala' latestDepForkedTest.scala.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.scala } @@ -56,6 +67,10 @@ tasks.named("compileLatestDepForkedTestGroovy", GroovyCompile) { classpath += files(sourceSets.latestDepForkedTest.scala.classesDirectory) } +tasks.named("compileScala21317TestGroovy", GroovyCompile) { + classpath += files(sourceSets.scala21317Test.scala.classesDirectory) +} + tasks.withType(ScalaCompile).configureEach { // Scala compilers here can't run on modern JDKs; pinned to JDK 8. // * https://docs.scala-lang.org/overviews/jdk-compatibility/overview.html @@ -71,4 +86,6 @@ dependencies { latestDepTestImplementation group: 'org.scala-lang', name: 'scala-library', version: '2.+' // scala-lang 3.x requires scala 3.x compiler latestDepTestImplementation project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common') + + scala21317TestImplementation group: 'org.scala-lang', name: 'scala-library', version: '2.13.17' } diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java new file mode 100644 index 00000000000..71d247969e2 --- /dev/null +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java @@ -0,0 +1,137 @@ +package datadog.trace.instrumentation.scala213.concurrent; + +import static datadog.trace.bootstrap.FieldBackedContextStores.getContextStoreId; + +import datadog.trace.agent.tooling.Instrumenter; +import datadog.trace.bootstrap.instrumentation.java.concurrent.State; +import datadog.trace.bootstrap.instrumentation.scala.ScalaPromiseContinuationHelper; +import java.util.concurrent.atomic.AtomicReference; +import net.bytebuddy.asm.AsmVisitorWrapper; +import net.bytebuddy.description.field.FieldDescription; +import net.bytebuddy.description.field.FieldList; +import net.bytebuddy.description.method.MethodList; +import net.bytebuddy.description.type.TypeDescription; +import net.bytebuddy.implementation.Implementation; +import net.bytebuddy.jar.asm.ClassVisitor; +import net.bytebuddy.jar.asm.ClassWriter; +import net.bytebuddy.jar.asm.MethodVisitor; +import net.bytebuddy.jar.asm.Opcodes; +import net.bytebuddy.jar.asm.Type; +import net.bytebuddy.pool.TypePool; + +public final class DefaultPromiseCallbackInstrumentation + implements Instrumenter.ForSingleType, Instrumenter.HasTypeAdvice { + + private static final String TRANSFORMATION = "scala.concurrent.impl.Promise$Transformation"; + + @Override + public String instrumentedType() { + return "scala.concurrent.impl.Promise$DefaultPromise"; + } + + @Override + public void typeAdvice(TypeTransformer transformer) { + transformer.applyAdvice( + new UnregisterCallbackVisitorWrapper( + getContextStoreId(TRANSFORMATION, State.class.getName()))); + } + + static final class UnregisterCallbackVisitorWrapper extends AsmVisitorWrapper.AbstractBase { + private final int contextStoreId; + + UnregisterCallbackVisitorWrapper(int contextStoreId) { + this.contextStoreId = contextStoreId; + } + + @Override + public int mergeWriter(int flags) { + return flags | ClassWriter.COMPUTE_MAXS; + } + + @Override + public ClassVisitor wrap( + TypeDescription instrumentedType, + ClassVisitor classVisitor, + Implementation.Context implementationContext, + TypePool typePool, + FieldList fields, + MethodList methods, + int writerFlags, + int readerFlags) { + return new UnregisterCallbackClassVisitor(classVisitor, contextStoreId); + } + } + + static final class UnregisterCallbackClassVisitor extends ClassVisitor { + private static final String UNREGISTER_DESCRIPTOR = + "(Lscala/concurrent/impl/Promise$Transformation;)V"; + + private final int contextStoreId; + + UnregisterCallbackClassVisitor(ClassVisitor classVisitor, int contextStoreId) { + super(Opcodes.ASM9, classVisitor); + this.contextStoreId = contextStoreId; + } + + @Override + public MethodVisitor visitMethod( + int access, String name, String descriptor, String signature, String[] exceptions) { + MethodVisitor methodVisitor = + super.visitMethod(access, name, descriptor, signature, exceptions); + if ("unregisterCallback".equals(name) && UNREGISTER_DESCRIPTOR.equals(descriptor)) { + return new UnregisterCallbackMethodVisitor(methodVisitor, contextStoreId); + } + return methodVisitor; + } + } + + static final class UnregisterCallbackMethodVisitor extends MethodVisitor { + private static final String DEFAULT_PROMISE = "scala/concurrent/impl/Promise$DefaultPromise"; + private static final String COMPARE_AND_SET_DESCRIPTOR = + "(Ljava/lang/Object;Ljava/lang/Object;)Z"; + private static final String CONTINUATION_HELPER = + Type.getInternalName(ScalaPromiseContinuationHelper.class); + private static final String CANCEL_DESCRIPTOR = + Type.getMethodDescriptor( + Type.BOOLEAN_TYPE, + Type.getType(AtomicReference.class), + Type.getType(Object.class), + Type.getType(Object.class), + Type.INT_TYPE, + Type.getType(Object.class)); + + private final int contextStoreId; + private int rewrittenCallSites; + + UnregisterCallbackMethodVisitor(MethodVisitor methodVisitor, int contextStoreId) { + super(Opcodes.ASM9, methodVisitor); + this.contextStoreId = contextStoreId; + } + + @Override + public void visitMethodInsn( + int opcode, String owner, String name, String descriptor, boolean isInterface) { + if (rewrittenCallSites < 2 + && opcode == Opcodes.INVOKEVIRTUAL + && DEFAULT_PROMISE.equals(owner) + && "compareAndSet".equals(name) + && COMPARE_AND_SET_DESCRIPTOR.equals(descriptor)) { + rewrittenCallSites++; + replaceCompareAndSet(); + } else { + super.visitMethodInsn(opcode, owner, name, descriptor, isInterface); + } + } + + private void replaceCompareAndSet() { + super.visitLdcInsn(contextStoreId); + super.visitVarInsn(Opcodes.ALOAD, 1); + super.visitMethodInsn( + Opcodes.INVOKESTATIC, + CONTINUATION_HELPER, + "compareAndSetAndCancel", + CANCEL_DESCRIPTOR, + false); + } + } +} diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java new file mode 100644 index 00000000000..e69679f314f --- /dev/null +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java @@ -0,0 +1,51 @@ +package datadog.trace.instrumentation.scala213.concurrent; + +import static datadog.trace.agent.tooling.muzzle.Reference.EXPECTS_NON_STATIC; +import static java.util.Collections.singletonList; + +import com.google.auto.service.AutoService; +import datadog.trace.agent.tooling.Instrumenter; +import datadog.trace.agent.tooling.InstrumenterModule; +import datadog.trace.agent.tooling.muzzle.Reference; +import datadog.trace.bootstrap.instrumentation.java.concurrent.State; +import java.util.Collections; +import java.util.List; +import java.util.Map; + +@AutoService(InstrumenterModule.class) +public final class ScalaPromiseUnregisterModule extends InstrumenterModule.ContextTracking { + + public ScalaPromiseUnregisterModule() { + super("scala_concurrent"); + } + + @Override + public String muzzleDirective() { + return "scala-promise-unregister"; + } + + @Override + public Map contextStore() { + return Collections.singletonMap( + "scala.concurrent.impl.Promise$Transformation", State.class.getName()); + } + + @Override + public Reference[] additionalMuzzleReferences() { + return new Reference[] { + new Reference.Builder("scala.concurrent.impl.Promise$DefaultPromise") + .withMethod( + new String[0], + EXPECTS_NON_STATIC, + "unregisterCallback", + "V", + "Lscala/concurrent/impl/Promise$Transformation;") + .build() + }; + } + + @Override + public List typeInstrumentations() { + return singletonList(new DefaultPromiseCallbackInstrumentation()); + } +} diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy new file mode 100644 index 00000000000..29cbecf9c5d --- /dev/null +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy @@ -0,0 +1,74 @@ +import datadog.trace.agent.test.InstrumentationSpecification + +import static datadog.trace.agent.test.utils.TraceUtils.basicSpan +import static datadog.trace.agent.test.utils.TraceUtils.runUnderTrace + +class ScalaFirstCompletedOfTest extends InstrumentationSpecification { + + def "releases callback continuation when firstCompletedOf unregisters it"() { + setup: + def executionContext = new QueuingExecutionContext() + def utils = new FirstCompletedOfUtils(executionContext) + def first = utils.newPromise() + def second = utils.newPromise() + def result + + when: + runUnderTrace("parent") { + result = utils.firstCompleted(first, second) + first.success("first") + } + + then: + executionContext.queuedTaskCount() == 1 + executionContext.runNext() + result.value().get().get() == "first" + + when: + second.success("second") + + then: + executionContext.queuedTaskCount() == 0 + assertTraces(1) { + trace(1) { + basicSpan(it, "parent") + } + } + } + + def "does not release callback continuation when completion wins unregister race"() { + setup: + def executionContext = new QueuingExecutionContext() + def utils = new FirstCompletedOfUtils(executionContext) + def first = utils.newPromise() + def second = utils.newPromise() + + when: + runUnderTrace("parent") { + utils.firstCompleted(first, second) + first.success("first") + second.success("second") + } + + then: + executionContext.queuedTaskCount() == 2 + + when: + executionContext.runNext() + + then: + executionContext.queuedTaskCount() == 1 + TEST_WRITER.size() == 0 + + when: + executionContext.runNext() + + then: + executionContext.queuedTaskCount() == 0 + assertTraces(1) { + trace(1) { + basicSpan(it, "parent") + } + } + } +} diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala new file mode 100644 index 00000000000..2f522065191 --- /dev/null +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala @@ -0,0 +1,32 @@ +import java.util.concurrent.ConcurrentLinkedQueue + +import scala.concurrent.{ExecutionContext, Future, Promise} + +final class QueuingExecutionContext extends ExecutionContext { + private val tasks = new ConcurrentLinkedQueue[Runnable]() + + override def execute(task: Runnable): Unit = tasks.add(task) + + override def reportFailure(cause: Throwable): Unit = throw cause + + def queuedTaskCount: Int = tasks.size() + + def runNext(): Boolean = { + val task = tasks.poll() + if (task eq null) { + false + } else { + task.run() + true + } + } +} + +final class FirstCompletedOfUtils(executionContext: ExecutionContext) { + private implicit val ec: ExecutionContext = executionContext + + def newPromise[T](): Promise[T] = Promise[T]() + + def firstCompleted[T](first: Promise[T], second: Promise[T]): Future[T] = + Future.firstCompletedOf(List(first.future, second.future)) +} From 188f30f37dba06df9513de00e321c05935b4d545 Mon Sep 17 00:00:00 2001 From: Andrea Marziali Date: Tue, 29 Sep 2026 13:43:18 +0200 Subject: [PATCH 2/3] simplify --- ...DefaultPromiseCallbackInstrumentation.java | 42 +++++++++++++-- .../ScalaPromiseUnregisterModule.java | 51 ------------------- 2 files changed, 39 insertions(+), 54 deletions(-) delete mode 100644 dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java index 71d247969e2..5f4c0532bd5 100644 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java @@ -1,10 +1,16 @@ package datadog.trace.instrumentation.scala213.concurrent; +import static datadog.trace.agent.tooling.muzzle.Reference.EXPECTS_NON_STATIC; import static datadog.trace.bootstrap.FieldBackedContextStores.getContextStoreId; +import static java.util.Collections.singletonMap; +import com.google.auto.service.AutoService; import datadog.trace.agent.tooling.Instrumenter; +import datadog.trace.agent.tooling.InstrumenterModule; +import datadog.trace.agent.tooling.muzzle.Reference; import datadog.trace.bootstrap.instrumentation.java.concurrent.State; import datadog.trace.bootstrap.instrumentation.scala.ScalaPromiseContinuationHelper; +import java.util.Map; import java.util.concurrent.atomic.AtomicReference; import net.bytebuddy.asm.AsmVisitorWrapper; import net.bytebuddy.description.field.FieldDescription; @@ -19,11 +25,40 @@ import net.bytebuddy.jar.asm.Type; import net.bytebuddy.pool.TypePool; -public final class DefaultPromiseCallbackInstrumentation +@AutoService(InstrumenterModule.class) +public final class DefaultPromiseCallbackInstrumentation extends InstrumenterModule.ContextTracking implements Instrumenter.ForSingleType, Instrumenter.HasTypeAdvice { private static final String TRANSFORMATION = "scala.concurrent.impl.Promise$Transformation"; + public DefaultPromiseCallbackInstrumentation() { + super("scala_concurrent"); + } + + @Override + public String muzzleDirective() { + return "scala-promise-unregister"; + } + + @Override + public Map contextStore() { + return singletonMap(TRANSFORMATION, State.class.getName()); + } + + @Override + public Reference[] additionalMuzzleReferences() { + return new Reference[] { + new Reference.Builder("scala.concurrent.impl.Promise$DefaultPromise") + .withMethod( + new String[0], + EXPECTS_NON_STATIC, + "unregisterCallback", + "V", + "Lscala/concurrent/impl/Promise$Transformation;") + .build() + }; + } + @Override public String instrumentedType() { return "scala.concurrent.impl.Promise$DefaultPromise"; @@ -36,10 +71,11 @@ public void typeAdvice(TypeTransformer transformer) { getContextStoreId(TRANSFORMATION, State.class.getName()))); } - static final class UnregisterCallbackVisitorWrapper extends AsmVisitorWrapper.AbstractBase { + public static final class UnregisterCallbackVisitorWrapper + extends AsmVisitorWrapper.AbstractBase { private final int contextStoreId; - UnregisterCallbackVisitorWrapper(int contextStoreId) { + public UnregisterCallbackVisitorWrapper(int contextStoreId) { this.contextStoreId = contextStoreId; } diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java deleted file mode 100644 index e69679f314f..00000000000 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/ScalaPromiseUnregisterModule.java +++ /dev/null @@ -1,51 +0,0 @@ -package datadog.trace.instrumentation.scala213.concurrent; - -import static datadog.trace.agent.tooling.muzzle.Reference.EXPECTS_NON_STATIC; -import static java.util.Collections.singletonList; - -import com.google.auto.service.AutoService; -import datadog.trace.agent.tooling.Instrumenter; -import datadog.trace.agent.tooling.InstrumenterModule; -import datadog.trace.agent.tooling.muzzle.Reference; -import datadog.trace.bootstrap.instrumentation.java.concurrent.State; -import java.util.Collections; -import java.util.List; -import java.util.Map; - -@AutoService(InstrumenterModule.class) -public final class ScalaPromiseUnregisterModule extends InstrumenterModule.ContextTracking { - - public ScalaPromiseUnregisterModule() { - super("scala_concurrent"); - } - - @Override - public String muzzleDirective() { - return "scala-promise-unregister"; - } - - @Override - public Map contextStore() { - return Collections.singletonMap( - "scala.concurrent.impl.Promise$Transformation", State.class.getName()); - } - - @Override - public Reference[] additionalMuzzleReferences() { - return new Reference[] { - new Reference.Builder("scala.concurrent.impl.Promise$DefaultPromise") - .withMethod( - new String[0], - EXPECTS_NON_STATIC, - "unregisterCallback", - "V", - "Lscala/concurrent/impl/Promise$Transformation;") - .build() - }; - } - - @Override - public List typeInstrumentations() { - return singletonList(new DefaultPromiseCallbackInstrumentation()); - } -} From 9dd16ff7dd3e0d2aa736fedb041987b3b8cc0670 Mon Sep 17 00:00:00 2001 From: Andrea Marziali Date: Wed, 30 Sep 2026 11:10:45 +0200 Subject: [PATCH 3/3] suggestions --- .../scala-promise-2.13/build.gradle | 2 + ...DefaultPromiseCallbackInstrumentation.java | 119 +++++++- .../groovy/ScalaFirstCompletedOfTest.groovy | 65 ++++ ...ultPromiseCallbackInstrumentationTest.java | 283 ++++++++++++++++++ .../scala/FirstCompletedOfUtils.scala | 3 + 5 files changed, 470 insertions(+), 2 deletions(-) create mode 100644 dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentationTest.java diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle index f08fa186d0b..ddabd8e099d 100644 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/build.gradle @@ -43,6 +43,8 @@ tasks.named("latestDepTest", Test) { } sourceSets { + latestDepTest.java.srcDir 'src/scala21317Test/java' + test.groovy.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.groovy test.groovy.srcDir sourceSets.latestDepForkedTest.groovy latestDepTest.groovy.srcDir project(':dd-java-agent:instrumentation:scala:scala-promise:scala-promise-common').sourceSets.test.groovy diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java index 5f4c0532bd5..eaafdb8aff0 100644 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/main/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentation.java @@ -20,6 +20,8 @@ import net.bytebuddy.implementation.Implementation; import net.bytebuddy.jar.asm.ClassVisitor; import net.bytebuddy.jar.asm.ClassWriter; +import net.bytebuddy.jar.asm.Handle; +import net.bytebuddy.jar.asm.Label; import net.bytebuddy.jar.asm.MethodVisitor; import net.bytebuddy.jar.asm.Opcodes; import net.bytebuddy.jar.asm.Type; @@ -103,6 +105,7 @@ static final class UnregisterCallbackClassVisitor extends ClassVisitor { "(Lscala/concurrent/impl/Promise$Transformation;)V"; private final int contextStoreId; + private boolean foundUnregisterCallback; UnregisterCallbackClassVisitor(ClassVisitor classVisitor, int contextStoreId) { super(Opcodes.ASM9, classVisitor); @@ -115,10 +118,22 @@ public MethodVisitor visitMethod( MethodVisitor methodVisitor = super.visitMethod(access, name, descriptor, signature, exceptions); if ("unregisterCallback".equals(name) && UNREGISTER_DESCRIPTOR.equals(descriptor)) { + if ((access & Opcodes.ACC_STATIC) != 0) { + throw new IllegalStateException("Expected an instance unregisterCallback method"); + } + foundUnregisterCallback = true; return new UnregisterCallbackMethodVisitor(methodVisitor, contextStoreId); } return methodVisitor; } + + @Override + public void visitEnd() { + if (!foundUnregisterCallback) { + throw new IllegalStateException("Missing expected unregisterCallback method"); + } + super.visitEnd(); + } } static final class UnregisterCallbackMethodVisitor extends MethodVisitor { @@ -138,6 +153,7 @@ static final class UnregisterCallbackMethodVisitor extends MethodVisitor { private final int contextStoreId; private int rewrittenCallSites; + private boolean loadedCallback; UnregisterCallbackMethodVisitor(MethodVisitor methodVisitor, int contextStoreId) { super(Opcodes.ASM9, methodVisitor); @@ -147,8 +163,8 @@ static final class UnregisterCallbackMethodVisitor extends MethodVisitor { @Override public void visitMethodInsn( int opcode, String owner, String name, String descriptor, boolean isInterface) { - if (rewrittenCallSites < 2 - && opcode == Opcodes.INVOKEVIRTUAL + loadedCallback = false; + if (opcode == Opcodes.INVOKEVIRTUAL && DEFAULT_PROMISE.equals(owner) && "compareAndSet".equals(name) && COMPARE_AND_SET_DESCRIPTOR.equals(descriptor)) { @@ -159,6 +175,105 @@ public void visitMethodInsn( } } + @Override + public void visitVarInsn(int opcode, int var) { + boolean writesCallback = + (var == 1 && opcode >= Opcodes.ISTORE && opcode <= Opcodes.ASTORE) + || (var == 0 && (opcode == Opcodes.LSTORE || opcode == Opcodes.DSTORE)); + if (writesCallback && !(opcode == Opcodes.ASTORE && var == 1 && loadedCallback)) { + throw new IllegalStateException("unregisterCallback overwrites its callback argument"); + } + loadedCallback = opcode == Opcodes.ALOAD && var == 1; + super.visitVarInsn(opcode, var); + } + + @Override + public void visitLabel(Label label) { + // Only allow adjacent ALOAD 1 / ASTORE 1, with no entry point into the store. + loadedCallback = false; + super.visitLabel(label); + } + + @Override + public void visitInsn(int opcode) { + loadedCallback = false; + super.visitInsn(opcode); + } + + @Override + public void visitIntInsn(int opcode, int operand) { + loadedCallback = false; + super.visitIntInsn(opcode, operand); + } + + @Override + public void visitTypeInsn(int opcode, String type) { + loadedCallback = false; + super.visitTypeInsn(opcode, type); + } + + @Override + public void visitFieldInsn(int opcode, String owner, String name, String descriptor) { + loadedCallback = false; + super.visitFieldInsn(opcode, owner, name, descriptor); + } + + @Override + public void visitInvokeDynamicInsn( + String name, String descriptor, Handle bootstrapMethod, Object... bootstrapArguments) { + loadedCallback = false; + super.visitInvokeDynamicInsn(name, descriptor, bootstrapMethod, bootstrapArguments); + } + + @Override + public void visitJumpInsn(int opcode, Label label) { + loadedCallback = false; + super.visitJumpInsn(opcode, label); + } + + @Override + public void visitLdcInsn(Object value) { + loadedCallback = false; + super.visitLdcInsn(value); + } + + @Override + public void visitIincInsn(int var, int increment) { + if (var == 1) { + throw new IllegalStateException("unregisterCallback overwrites its callback argument"); + } + loadedCallback = false; + super.visitIincInsn(var, increment); + } + + @Override + public void visitTableSwitchInsn(int min, int max, Label defaultLabel, Label... labels) { + loadedCallback = false; + super.visitTableSwitchInsn(min, max, defaultLabel, labels); + } + + @Override + public void visitLookupSwitchInsn(Label defaultLabel, int[] keys, Label[] labels) { + loadedCallback = false; + super.visitLookupSwitchInsn(defaultLabel, keys, labels); + } + + @Override + public void visitMultiANewArrayInsn(String descriptor, int dimensions) { + loadedCallback = false; + super.visitMultiANewArrayInsn(descriptor, dimensions); + } + + @Override + public void visitEnd() { + if (rewrittenCallSites != 2) { + // Reject the combined class transformation instead of installing a partial rewrite. + throw new IllegalStateException( + "Expected 2 unregisterCallback compareAndSet sites, found " + rewrittenCallSites); + } + super.visitEnd(); + } + private void replaceCompareAndSet() { super.visitLdcInsn(contextStoreId); super.visitVarInsn(Opcodes.ALOAD, 1); diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy index 29cbecf9c5d..0775c1f99a9 100644 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/groovy/ScalaFirstCompletedOfTest.groovy @@ -71,4 +71,69 @@ class ScalaFirstCompletedOfTest extends InstrumentationSpecification { } } } + + def "preserves another callback when unregistering the #position of ManyCallbacks"() { + setup: + def executionContext = new QueuingExecutionContext() + def utils = new FirstCompletedOfUtils(executionContext) + def first = utils.newPromise() + def second = utils.newPromise() + def result + def registerSurvivor = { + runUnderTrace("survivor-parent") { + utils.onComplete(second, { + runUnderTrace("survivor-child") {} + } as Runnable) + } + } + + when: + if (survivorFirst) { + registerSurvivor() + } + runUnderTrace("first-completed-parent") { + result = utils.firstCompleted(first, second) + } + if (!survivorFirst) { + registerSurvivor() + } + first.success("first") + + then: + executionContext.queuedTaskCount() == 1 + executionContext.runNext() + result.value().get().get() == "first" + assertTraces(1) { + trace(1) { + basicSpan(it, "first-completed-parent") + } + } + + when: + second.success("second") + + then: + executionContext.queuedTaskCount() == 1 + TEST_WRITER.size() == 1 + + when: + executionContext.runNext() + + then: + executionContext.queuedTaskCount() == 0 + assertTraces(2, SORT_TRACES_BY_NAMES) { + trace(1) { + basicSpan(it, "first-completed-parent") + } + trace(2, true) { + basicSpan(it, "survivor-child", it.span(1)) + basicSpan(it, "survivor-parent") + } + } + + where: + position | survivorFirst + "head" | true + "tail" | false + } } diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentationTest.java b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentationTest.java new file mode 100644 index 00000000000..f1897339825 --- /dev/null +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/java/datadog/trace/instrumentation/scala213/concurrent/DefaultPromiseCallbackInstrumentationTest.java @@ -0,0 +1,283 @@ +package datadog.trace.instrumentation.scala213.concurrent; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import datadog.trace.bootstrap.instrumentation.scala.ScalaPromiseContinuationHelper; +import java.io.IOException; +import java.io.InputStream; +import net.bytebuddy.jar.asm.ClassReader; +import net.bytebuddy.jar.asm.ClassVisitor; +import net.bytebuddy.jar.asm.ClassWriter; +import net.bytebuddy.jar.asm.Label; +import net.bytebuddy.jar.asm.MethodVisitor; +import net.bytebuddy.jar.asm.Opcodes; +import net.bytebuddy.jar.asm.Type; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import org.tabletest.junit.TableTest; + +class DefaultPromiseCallbackInstrumentationTest { + private static final String DEFAULT_PROMISE = "scala/concurrent/impl/Promise$DefaultPromise"; + private static final String CAS_DESCRIPTOR = "(Ljava/lang/Object;Ljava/lang/Object;)Z"; + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void rewritesSupportedScalaBytecode(boolean stripDebug) throws IOException { + byte[] original = scalaBytecode(); + assertEquals(2, countCalls(original, DEFAULT_PROMISE, "compareAndSet")); + + byte[] transformed = transform(original, stripDebug ? ClassReader.SKIP_DEBUG : 0); + + assertEquals(0, countCalls(transformed, DEFAULT_PROMISE, "compareAndSet")); + assertEquals( + 2, + countCalls( + transformed, + Type.getInternalName(ScalaPromiseContinuationHelper.class), + "compareAndSetAndCancel")); + } + + @TableTest({ + "Scenario | Mutation | Cas Count | Message ", + "additional CAS | EXTRA_CAS | 3 | Expected 2 unregisterCallback compareAndSet sites, found 3", + "one CAS removed | MISSING_CAS | 1 | Expected 2 unregisterCallback compareAndSet sites, found 1", + "all CAS removed | NO_CAS | 0 | Expected 2 unregisterCallback compareAndSet sites, found 0", + "method removed | MISSING_METHOD | 0 | Missing expected unregisterCallback method ", + "method made static | STATIC_METHOD | 0 | Expected an instance unregisterCallback method ", + "callback reassigned | REASSIGN_CALLBACK | 2 | unregisterCallback overwrites its callback argument ", + "loaded callback replaced | REPLACE_LOADED | 2 | unregisterCallback overwrites its callback argument ", + "callback slot reused | REUSE_SLOT | 2 | unregisterCallback overwrites its callback argument ", + "long overlaps callback | LONG_OVERLAP | 2 | unregisterCallback overwrites its callback argument ", + "double overlaps callback | DOUBLE_OVERLAP | 2 | unregisterCallback overwrites its callback argument ", + "branch enters store | BRANCH_TO_STORE | 2 | unregisterCallback overwrites its callback argument ", + "handler enters store | HANDLER_AT_STORE | 2 | unregisterCallback overwrites its callback argument " + }) + void rejectsChangedAssumptions(Mutation mutation, int casCount, String message) + throws IOException { + byte[] changed = mutate(scalaBytecode(), mutation); + assertEquals(casCount, countCalls(changed, DEFAULT_PROMISE, "compareAndSet")); + + IllegalStateException failure = + assertThrows(IllegalStateException.class, () -> transform(changed, 0)); + + assertEquals(message, failure.getMessage()); + } + + private static byte[] scalaBytecode() throws IOException { + try (InputStream input = + DefaultPromiseCallbackInstrumentationTest.class + .getClassLoader() + .getResourceAsStream(DEFAULT_PROMISE + ".class")) { + assertNotNull(input); + ClassWriter writer = new ClassWriter(0); + new ClassReader(input).accept(writer, 0); + return writer.toByteArray(); + } + } + + private static byte[] transform(byte[] input, int readerFlags) { + ClassWriter writer = new ClassWriter(ClassWriter.COMPUTE_MAXS); + new ClassReader(input) + .accept( + new DefaultPromiseCallbackInstrumentation.UnregisterCallbackClassVisitor(writer, 17), + readerFlags); + return writer.toByteArray(); + } + + private static int countCalls(byte[] bytes, String expectedOwner, String expectedName) { + int[] calls = {0}; + new ClassReader(bytes) + .accept( + new ClassVisitor(Opcodes.ASM9) { + @Override + public MethodVisitor visitMethod( + int access, + String name, + String descriptor, + String signature, + String[] exceptions) { + if (!"unregisterCallback".equals(name)) { + return null; + } + return new MethodVisitor(Opcodes.ASM9) { + @Override + public void visitMethodInsn( + int opcode, + String owner, + String name, + String descriptor, + boolean isInterface) { + if (expectedOwner.equals(owner) && expectedName.equals(name)) { + calls[0]++; + } + } + }; + } + }, + ClassReader.SKIP_DEBUG | ClassReader.SKIP_FRAMES); + return calls[0]; + } + + private static byte[] mutate(byte[] original, Mutation mutation) { + ClassWriter writer = new ClassWriter(ClassWriter.COMPUTE_FRAMES); + new ClassReader(original) + .accept( + new ClassVisitor(Opcodes.ASM9, writer) { + @Override + public MethodVisitor visitMethod( + int access, + String name, + String descriptor, + String signature, + String[] exceptions) { + if ("unregisterCallback".equals(name)) { + if (mutation == Mutation.MISSING_METHOD) { + return null; + } + if (mutation == Mutation.STATIC_METHOD) { + MethodVisitor method = + super.visitMethod( + Opcodes.ACC_PRIVATE | Opcodes.ACC_STATIC, + name, + descriptor, + signature, + exceptions); + method.visitCode(); + method.visitInsn(Opcodes.RETURN); + method.visitMaxs(0, 1); + method.visitEnd(); + return null; + } + } + MethodVisitor method = + super.visitMethod(access, name, descriptor, signature, exceptions); + if (!"unregisterCallback".equals(name)) { + return method; + } + return new MethodVisitor(Opcodes.ASM9, method) { + private boolean removedCas; + + @Override + public void visitCode() { + super.visitCode(); + mutation.insert(method); + } + + @Override + public void visitMethodInsn( + int opcode, + String owner, + String name, + String descriptor, + boolean isInterface) { + if (DEFAULT_PROMISE.equals(owner) + && "compareAndSet".equals(name) + && (mutation == Mutation.NO_CAS + || (mutation == Mutation.MISSING_CAS && !removedCas))) { + removedCas = true; + super.visitInsn(Opcodes.POP2); + super.visitInsn(Opcodes.POP); + super.visitInsn(Opcodes.ICONST_1); + } else { + super.visitMethodInsn(opcode, owner, name, descriptor, isInterface); + } + } + }; + } + }, + ClassReader.SKIP_FRAMES); + return writer.toByteArray(); + } + + enum Mutation { + EXTRA_CAS, + MISSING_CAS, + NO_CAS, + MISSING_METHOD, + STATIC_METHOD, + REASSIGN_CALLBACK, + REPLACE_LOADED, + REUSE_SLOT, + LONG_OVERLAP, + DOUBLE_OVERLAP, + BRANCH_TO_STORE, + HANDLER_AT_STORE; + + void insert(MethodVisitor method) { + switch (this) { + case EXTRA_CAS: + method.visitVarInsn(Opcodes.ALOAD, 0); + method.visitInsn(Opcodes.ACONST_NULL); + method.visitInsn(Opcodes.ACONST_NULL); + method.visitMethodInsn( + Opcodes.INVOKEVIRTUAL, DEFAULT_PROMISE, "compareAndSet", CAS_DESCRIPTOR, false); + method.visitInsn(Opcodes.POP); + break; + case REASSIGN_CALLBACK: + method.visitInsn(Opcodes.ACONST_NULL); + method.visitVarInsn(Opcodes.ASTORE, 1); + break; + case REPLACE_LOADED: + method.visitVarInsn(Opcodes.ALOAD, 1); + method.visitInsn(Opcodes.POP); + method.visitInsn(Opcodes.ACONST_NULL); + method.visitVarInsn(Opcodes.ASTORE, 1); + break; + case REUSE_SLOT: + method.visitVarInsn(Opcodes.ALOAD, 1); + method.visitVarInsn(Opcodes.ASTORE, 3); + method.visitInsn(Opcodes.ICONST_0); + method.visitVarInsn(Opcodes.ISTORE, 1); + method.visitVarInsn(Opcodes.ALOAD, 3); + method.visitVarInsn(Opcodes.ASTORE, 1); + break; + case LONG_OVERLAP: + case DOUBLE_OVERLAP: + method.visitVarInsn(Opcodes.ALOAD, 0); + method.visitVarInsn(Opcodes.ASTORE, 3); + method.visitVarInsn(Opcodes.ALOAD, 1); + method.visitVarInsn(Opcodes.ASTORE, 4); + method.visitInsn(this == LONG_OVERLAP ? Opcodes.LCONST_0 : Opcodes.DCONST_0); + method.visitVarInsn(this == LONG_OVERLAP ? Opcodes.LSTORE : Opcodes.DSTORE, 0); + method.visitVarInsn(Opcodes.ALOAD, 3); + method.visitVarInsn(Opcodes.ASTORE, 0); + method.visitVarInsn(Opcodes.ALOAD, 4); + method.visitVarInsn(Opcodes.ASTORE, 1); + break; + case BRANCH_TO_STORE: + Label load = new Label(); + Label store = new Label(); + method.visitVarInsn(Opcodes.ALOAD, 0); + method.visitJumpInsn(Opcodes.IFNONNULL, load); + method.visitInsn(Opcodes.ACONST_NULL); + method.visitJumpInsn(Opcodes.GOTO, store); + method.visitLabel(load); + method.visitVarInsn(Opcodes.ALOAD, 1); + method.visitLabel(store); + method.visitVarInsn(Opcodes.ASTORE, 1); + break; + case HANDLER_AT_STORE: + Label start = new Label(); + Label end = new Label(); + Label handler = new Label(); + method.visitTryCatchBlock(start, end, handler, "java/lang/Throwable"); + method.visitVarInsn(Opcodes.ALOAD, 1); + method.visitVarInsn(Opcodes.ASTORE, 3); + method.visitLabel(start); + method.visitVarInsn(Opcodes.ALOAD, 0); + method.visitInsn(Opcodes.POP); + method.visitLabel(end); + method.visitVarInsn(Opcodes.ALOAD, 1); + method.visitLabel(handler); + method.visitVarInsn(Opcodes.ASTORE, 1); + method.visitVarInsn(Opcodes.ALOAD, 3); + method.visitVarInsn(Opcodes.ASTORE, 1); + break; + default: + break; + } + } + } +} diff --git a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala index 2f522065191..c01d77eda4d 100644 --- a/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala +++ b/dd-java-agent/instrumentation/scala/scala-promise/scala-promise-2.13/src/scala21317Test/scala/FirstCompletedOfUtils.scala @@ -29,4 +29,7 @@ final class FirstCompletedOfUtils(executionContext: ExecutionContext) { def firstCompleted[T](first: Promise[T], second: Promise[T]): Future[T] = Future.firstCompletedOf(List(first.future, second.future)) + + def onComplete[T](promise: Promise[T], callback: Runnable): Unit = + promise.future.onComplete(_ => callback.run()) }