diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManager.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManager.java index 2ea74c435b..984a03b246 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManager.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManager.java @@ -32,9 +32,11 @@ import java.util.ArrayList; import java.util.Collections; import java.util.Date; +import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.Set; import java.util.concurrent.ArrayBlockingQueue; import java.util.concurrent.FutureTask; import java.util.concurrent.RejectedExecutionException; @@ -65,7 +67,9 @@ public class LocalTaskManager { private final ReentrantLock lifecycleLock = new ReentrantLock(); - private boolean preparingForExit; + private boolean applicationExiting; + + private final Set exitingOwners = new HashSet<>(); public LocalTaskManager(TaskStorage taskStorage, TaskExecutorRegistry taskExecutorRegistry, ArtifactService artifactService, ConnectionContextConverter connectionContextConverter, @@ -102,7 +106,7 @@ Task submit(Task task, TaskEvent createdEvent, S spec, Cont ConnectInfo connectInfo) { lifecycleLock.lock(); try { - if (preparingForExit) { + if (applicationExiting || exitingOwners.contains(ownerOf(task))) { throw new RejectedExecutionException("The application is preparing to exit"); } Task persistedTask = taskStorage.create(task, createdEvent); @@ -193,10 +197,10 @@ void prepareForUserExit(Long userId, Long organizationId) { new TaskOwner(userId, organizationId)); } - void abortUserExit() { + void abortUserExit(Long userId, Long organizationId) { lifecycleLock.lock(); try { - preparingForExit = false; + exitingOwners.remove(new TaskOwner(userId, organizationId)); } finally { lifecycleLock.unlock(); } @@ -213,10 +217,14 @@ void shutdown() { private void terminateActiveTasks(String errorCode, String eventCode, String message, TaskOwner owner) { lifecycleLock.lock(); try { - if (preparingForExit && owner != null) { + if (applicationExiting && owner != null) { return; } - preparingForExit = true; + if (owner == null) { + applicationExiting = true; + } else { + exitingOwners.add(owner); + } List activeTasks = taskStorage.listNonTerminalTasks(); List tasksToAwait = new ArrayList<>(); List tasksToCleanup = new ArrayList<>(); @@ -345,6 +353,10 @@ private boolean belongsTo(Task task, Long userId, Long organizationId) { && Objects.equals(task.getOrganizationId(), organizationId); } + private TaskOwner ownerOf(Task task) { + return new TaskOwner(task.getUserId(), task.getOrganizationId()); + } + private TaskSubmissionContext extensionContext(Task task, TaskSpec spec, ConnectInfo connectInfo) { TaskType taskType = TaskType.valueOf(spec.getTaskType()); TaskTargetSnapshot target = spec.getTarget(); diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskServiceImpl.java index 735127640a..3f7b008356 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/task/TaskServiceImpl.java @@ -143,7 +143,8 @@ public void prepareForUserExit() { @Override public void abortUserExit() { - localTaskManager.abortUserExit(); + TaskOwner owner = currentOwner(); + localTaskManager.abortUserExit(owner.userId(), owner.organizationId()); } @Override diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManagerTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManagerTest.java index c6ea4b5779..1e287cf25a 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManagerTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/task/LocalTaskManagerTest.java @@ -313,7 +313,7 @@ void abortedUserExitAllowsNewTasksAgain() throws Exception { taskManager = manager(storage, (spec, context) -> {}); taskManager.prepareForUserExit(null, null); - taskManager.abortUserExit(); + taskManager.abortUserExit(null, null); Task submitted = taskManager.submit(newTask(), event(TaskEventCode.TASK_CREATED.name()), spec(), null, null); @@ -322,6 +322,59 @@ void abortedUserExitAllowsNewTasksAgain() throws Exception { assertEquals(TaskStatus.SUCCESS.name(), storage.get(submitted.getId()).orElseThrow().getStatus()); } + @Test + void userExitPreparationOnlyRejectsTasksOfTheExitingOwner() throws Exception { + TestTaskStorage storage = new TestTaskStorage(); + taskManager = manager(storage, (spec, context) -> {}); + + taskManager.prepareForUserExit(1L, 10L); + + assertThrows(RejectedExecutionException.class, + () -> taskManager.submit(newTask(1L, 10L), event(TaskEventCode.TASK_CREATED.name()), + spec(), null, null)); + Task otherOwnerTask = taskManager.submit(newTask(2L, 10L), + event(TaskEventCode.TASK_CREATED.name()), spec(), null, null); + + assertNotNull(otherOwnerTask); + assertTrue(storage.awaitTerminal()); + assertEquals(TaskStatus.SUCCESS.name(), + storage.get(otherOwnerTask.getId()).orElseThrow().getStatus()); + assertTrue(storage.listNonTerminalTasks().isEmpty()); + } + + @Test + void abortUserExitUnblocksOnlyTheRequestingOwner() throws Exception { + TestTaskStorage storage = new TestTaskStorage(); + taskManager = manager(storage, (spec, context) -> {}); + taskManager.prepareForUserExit(1L, 10L); + taskManager.prepareForUserExit(2L, 10L); + + taskManager.abortUserExit(1L, 10L); + + Task firstOwnerTask = taskManager.submit(newTask(1L, 10L), + event(TaskEventCode.TASK_CREATED.name()), spec(), null, null); + assertNotNull(firstOwnerTask); + assertThrows(RejectedExecutionException.class, + () -> taskManager.submit(newTask(2L, 10L), event(TaskEventCode.TASK_CREATED.name()), + spec(), null, null)); + } + + @Test + void applicationExitRejectsTasksOfAllOwners() { + TestTaskStorage storage = new TestTaskStorage(); + taskManager = manager(storage, (spec, context) -> {}); + + taskManager.shutdown(); + + assertThrows(RejectedExecutionException.class, + () -> taskManager.submit(newTask(1L, 10L), event(TaskEventCode.TASK_CREATED.name()), + spec(), null, null)); + assertThrows(RejectedExecutionException.class, + () -> taskManager.submit(newTask(2L, 10L), event(TaskEventCode.TASK_CREATED.name()), + spec(), null, null)); + assertTrue(storage.listNonTerminalTasks().isEmpty()); + } + @Test void executionContextBindingFailureFailsTaskAndCleansRunningRegistry() { TestTaskStorage storage = new TestTaskStorage(); @@ -558,10 +611,16 @@ private TaskExtensionManager emptyExtensionManager() { } private Task newTask() { + return newTask(null, null); + } + + private Task newTask(Long userId, Long organizationId) { return Task.builder() .type(TaskType.QUERY_RESULT_EXPORT.name()) .name("Export result") .target(TaskTargetSnapshot.builder().dataSourceId(1L).build()) + .userId(userId) + .organizationId(organizationId) .build(); }