diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/RunningTask.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/RunningTask.java index 219ce48b4b..46b56387ab 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/RunningTask.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/RunningTask.java @@ -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 @@ -26,9 +26,13 @@ final class RunningTask { private final Long taskId; + private final Executor cancellationExecutor; + private final CancellationToken cancellationToken = new CancellationToken(); - private final AtomicReference cancelable = new AtomicReference<>(); + private final Object cancellationLock = new Object(); + + private TaskCancelable cancelable; private final ReentrantLock completionLock = new ReentrantLock(); @@ -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() { @@ -59,29 +68,42 @@ 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() { @@ -89,8 +111,10 @@ boolean isClosed() { } void close() { - closed = true; - cancelable.set(null); + synchronized (cancellationLock) { + closed = true; + cancelable = null; + } } void markFinished() { @@ -105,7 +129,7 @@ private void cancelRegisteredResourceAsync(TaskCancelable resource) { if (resource == null) { return; } - CANCELLATION_EXECUTOR.execute(() -> { + cancellationExecutor.execute(() -> { try { resource.cancel(); } catch (Exception e) { diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskExecutionContextImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskExecutionContextImpl.java index c37118a01a..56c4efee48 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskExecutionContextImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskExecutionContextImpl.java @@ -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 diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/RunningTaskTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/RunningTaskTest.java index 2a3802f41e..ae5b5a2422 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/RunningTaskTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/RunningTaskTest.java @@ -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; @@ -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 registrationThread = new AtomicReference<>(); + AtomicReference 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 scheduledCancellations = new CopyOnWriteArrayList<>(); + RunningTask runningTask = new RunningTask(42L, scheduledCancellations::add); + CountDownLatch futureCancellationStarted = new CountDownLatch(1); + CountDownLatch releaseFutureCancellation = new CountDownLatch(1); + FutureTask 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 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()); + } }