Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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.
*
* <p>The identity check distinguishes a successful removal from a successful no-op CAS.
*/
public static boolean compareAndSetAndCancel(
AtomicReference<Object> 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;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -28,19 +36,24 @@ 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'
}

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
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
}

Expand All @@ -56,6 +69,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
Expand All @@ -71,4 +88,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'
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,288 @@
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;
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.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;
import net.bytebuddy.pool.TypePool;

@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<String, String> 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";
}

@Override
public void typeAdvice(TypeTransformer transformer) {
transformer.applyAdvice(
new UnregisterCallbackVisitorWrapper(
getContextStoreId(TRANSFORMATION, State.class.getName())));
}

public static final class UnregisterCallbackVisitorWrapper
extends AsmVisitorWrapper.AbstractBase {
private final int contextStoreId;

public 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<FieldDescription.InDefinedShape> 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;
private boolean foundUnregisterCallback;

UnregisterCallbackClassVisitor(ClassVisitor classVisitor, int contextStoreId) {
super(Opcodes.ASM9, classVisitor);
Comment thread
amarziali marked this conversation as resolved.
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)) {
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 {
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";
Comment thread
amarziali marked this conversation as resolved.
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;
private boolean loadedCallback;

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) {
loadedCallback = false;
if (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);
}
}

@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);
Comment thread
amarziali marked this conversation as resolved.
super.visitMethodInsn(
Opcodes.INVOKESTATIC,
CONTINUATION_HELPER,
"compareAndSetAndCancel",
CANCEL_DESCRIPTOR,
false);
}
}
}
Loading
Loading