Skip to content
Merged
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
Expand Up @@ -4,12 +4,12 @@
import lombok.extern.slf4j.Slf4j;

import java.util.concurrent.CountDownLatch;
import java.util.concurrent.Executor;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.Future;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference;
import java.util.concurrent.locks.ReentrantLock;

@Slf4j
Expand All @@ -26,9 +26,13 @@ final class RunningTask {

private final Long taskId;

private final Executor cancellationExecutor;

private final CancellationToken cancellationToken = new CancellationToken();

private final AtomicReference<TaskCancelable> cancelable = new AtomicReference<>();
private final Object cancellationLock = new Object();

private TaskCancelable cancelable;

private final ReentrantLock completionLock = new ReentrantLock();

Expand All @@ -39,7 +43,12 @@ final class RunningTask {
private volatile boolean closed;

RunningTask(Long taskId) {
this(taskId, CANCELLATION_EXECUTOR);
}

RunningTask(Long taskId, Executor cancellationExecutor) {
this.taskId = taskId;
this.cancellationExecutor = cancellationExecutor;
}

Long taskId() {
Expand All @@ -59,38 +68,53 @@ void setFuture(Future<?> future) {
}

boolean requestCancellation(boolean mayInterruptIfRunning) {
if (closed) {
return false;
}
if (!cancellationToken.cancel()) {
return false;
Future<?> currentFuture;
TaskCancelable currentCancelable;
synchronized (cancellationLock) {
if (closed) {
return false;
}
if (!cancellationToken.cancel()) {
return false;
}
currentFuture = future;
currentCancelable = cancelable;
}
Future<?> currentFuture = future;
if (currentFuture != null) {
currentFuture.cancel(mayInterruptIfRunning);
}
cancelRegisteredResourceAsync(cancelable.get());
cancelRegisteredResourceAsync(currentCancelable);
return true;
}

void registerCancelable(TaskCancelable resource) {
cancelable.set(resource);
if (resource != null && cancellationToken.isCancelled()) {
boolean cancelImmediately;
synchronized (cancellationLock) {
cancelable = resource;
cancelImmediately = resource != null && cancellationToken.isCancelled();
}
if (cancelImmediately) {
cancelRegisteredResourceAsync(resource);
}
}

void clearCancelable(TaskCancelable resource) {
cancelable.compareAndSet(resource, null);
synchronized (cancellationLock) {
if (cancelable == resource) {
cancelable = null;
}
}
}

boolean isClosed() {
return closed;
}

void close() {
closed = true;
cancelable.set(null);
synchronized (cancellationLock) {
closed = true;
cancelable = null;
}
}

void markFinished() {
Expand All @@ -105,7 +129,7 @@ private void cancelRegisteredResourceAsync(TaskCancelable resource) {
if (resource == null) {
return;
}
CANCELLATION_EXECUTOR.execute(() -> {
cancellationExecutor.execute(() -> {
try {
resource.cancel();
} catch (Exception e) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -139,13 +139,6 @@ public void onStatementCreated(Statement statement) {
TaskCancelable cancelable = statement::cancel;
activeStatement.set(new StatementRegistration(statement, cancelable));
runningTask.registerCancelable(cancelable);
if (runningTask.cancellationToken().isCancelled()) {
try {
statement.cancel();
} catch (Exception ignored) {
// The runner will still observe the cancellation token.
}
}
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,20 @@

import org.junit.jupiter.api.Test;

import java.lang.reflect.Proxy;
import java.sql.Statement;
import java.time.Duration;
import java.util.List;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.CopyOnWriteArrayList;
import java.util.concurrent.FutureTask;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference;

import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNotSame;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.junit.jupiter.api.Assertions.assertTimeoutPreemptively;

Expand Down Expand Up @@ -35,4 +43,78 @@ void blockingJdbcCancellationDoesNotBlockTheCancellationRequest() throws Excepti
releaseCancel.countDown();
}
}

@Test
void preCancelledTaskCancelsNewStatementOnceWithoutBlockingRegistration() throws Exception {
RunningTask runningTask = new RunningTask(42L);
assertTrue(runningTask.requestCancellation(true));
TaskExecutionContextImpl context = new TaskExecutionContextImpl(42L, runningTask, null, null);
AtomicInteger cancellationCount = new AtomicInteger();
AtomicReference<Thread> registrationThread = new AtomicReference<>();
AtomicReference<Thread> cancellationThread = new AtomicReference<>();
CountDownLatch cancelStarted = new CountDownLatch(1);
CountDownLatch releaseCancel = new CountDownLatch(1);
Statement statement = (Statement) Proxy.newProxyInstance(Statement.class.getClassLoader(),
new Class<?>[] {Statement.class}, (proxy, method, args) -> {
if ("cancel".equals(method.getName())) {
cancellationCount.incrementAndGet();
cancellationThread.set(Thread.currentThread());
cancelStarted.countDown();
releaseCancel.await();
}
return null;
});

try {
assertTimeoutPreemptively(Duration.ofSeconds(1), () -> {
registrationThread.set(Thread.currentThread());
context.onStatementCreated(statement);
});
assertTrue(cancelStarted.await(1, TimeUnit.SECONDS));
assertEquals(1, cancellationCount.get());
assertNotSame(registrationThread.get(), cancellationThread.get());
} finally {
releaseCancel.countDown();
}
}

@Test
void concurrentCancellationAndRegistrationScheduleResourceOnce() throws Exception {
List<Runnable> scheduledCancellations = new CopyOnWriteArrayList<>();
RunningTask runningTask = new RunningTask(42L, scheduledCancellations::add);
CountDownLatch futureCancellationStarted = new CountDownLatch(1);
CountDownLatch releaseFutureCancellation = new CountDownLatch(1);
FutureTask<Void> future = new FutureTask<>(() -> null) {
@Override
public boolean cancel(boolean mayInterruptIfRunning) {
futureCancellationStarted.countDown();
try {
if (!releaseFutureCancellation.await(1, TimeUnit.SECONDS)) {
throw new AssertionError("Timed out waiting to resume future cancellation");
}
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
throw new AssertionError(e);
}
return super.cancel(mayInterruptIfRunning);
}
};
runningTask.setFuture(future);
AtomicInteger cancellationCount = new AtomicInteger();
FutureTask<Boolean> cancellationRequest = new FutureTask<>(() -> runningTask.requestCancellation(true));
Thread cancellationThread = new Thread(cancellationRequest, "running-task-cancellation-test");
cancellationThread.start();

try {
assertTrue(futureCancellationStarted.await(1, TimeUnit.SECONDS));
runningTask.registerCancelable(cancellationCount::incrementAndGet);
} finally {
releaseFutureCancellation.countDown();
}

assertTrue(cancellationRequest.get(1, TimeUnit.SECONDS));
assertEquals(1, scheduledCancellations.size());
scheduledCancellations.get(0).run();
assertEquals(1, cancellationCount.get());
}
}
Loading