diff --git a/spring-core/src/main/java/org/springframework/core/task/SimpleAsyncTaskExecutor.java b/spring-core/src/main/java/org/springframework/core/task/SimpleAsyncTaskExecutor.java index 3f1c53f7ab7..5ae2201f4b0 100644 --- a/spring-core/src/main/java/org/springframework/core/task/SimpleAsyncTaskExecutor.java +++ b/spring-core/src/main/java/org/springframework/core/task/SimpleAsyncTaskExecutor.java @@ -211,10 +211,12 @@ public class SimpleAsyncTaskExecutor extends CustomizableThreadCreator /** * Specify whether to reject tasks when the concurrency limit has been reached, - * throwing {@link TaskRejectedException} on any further submission attempts. + * throwing {@link TaskRejectedException} (which extends the common + * {@link java.util.concurrent.RejectedExecutionException}) *

The default is {@code false}, blocking the caller until the submission can * be accepted. Switch this to {@code true} for immediate rejection instead. * @since 6.2.6 + * @see #setConcurrencyLimit */ public void setRejectTasksWhenLimitReached(boolean rejectTasksWhenLimitReached) { this.rejectTasksWhenLimitReached = rejectTasksWhenLimitReached; diff --git a/spring-core/src/main/java/org/springframework/core/task/SyncTaskExecutor.java b/spring-core/src/main/java/org/springframework/core/task/SyncTaskExecutor.java index 5b6e633da09..56e90dd2fa2 100644 --- a/spring-core/src/main/java/org/springframework/core/task/SyncTaskExecutor.java +++ b/spring-core/src/main/java/org/springframework/core/task/SyncTaskExecutor.java @@ -42,6 +42,24 @@ import org.springframework.util.ConcurrencyThrottleSupport; @SuppressWarnings("serial") public class SyncTaskExecutor extends ConcurrencyThrottleSupport implements TaskExecutor, Serializable { + private boolean rejectTasksWhenLimitReached = false; + + + /** + * Specify whether to reject tasks when the concurrency limit has been reached, + * throwing {@link TaskRejectedException} (which extends the common + * {@link java.util.concurrent.RejectedExecutionException}) + * on any further execution attempts. + *

The default is {@code false}, blocking the caller until the submission can + * be accepted. Switch this to {@code true} for immediate rejection instead. + * @since 7.0.3 + * @see #setConcurrencyLimit + */ + public void setRejectTasksWhenLimitReached(boolean rejectTasksWhenLimitReached) { + this.rejectTasksWhenLimitReached = rejectTasksWhenLimitReached; + } + + /** * Execute the given {@code task} synchronously, through direct * invocation of its {@link Runnable#run() run()} method. @@ -88,4 +106,12 @@ public class SyncTaskExecutor extends ConcurrencyThrottleSupport implements Task } } + @Override + protected void onLimitReached() { + if (this.rejectTasksWhenLimitReached) { + throw new TaskRejectedException("Concurrency limit reached: " + getConcurrencyLimit()); + } + super.onLimitReached(); + } + } diff --git a/spring-core/src/test/java/org/springframework/core/task/SyncTaskExecutorTests.java b/spring-core/src/test/java/org/springframework/core/task/SyncTaskExecutorTests.java index 4fe8385ab53..b94da01be96 100644 --- a/spring-core/src/test/java/org/springframework/core/task/SyncTaskExecutorTests.java +++ b/spring-core/src/test/java/org/springframework/core/task/SyncTaskExecutorTests.java @@ -25,6 +25,7 @@ import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.Test; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; import static org.assertj.core.api.Assertions.assertThatIOException; import static org.assertj.core.api.Assertions.assertThatNoException; @@ -36,23 +37,23 @@ class SyncTaskExecutorTests { @Test void plainExecution() { - SyncTaskExecutor taskExecutor = new SyncTaskExecutor(); + SyncTaskExecutor executor = new SyncTaskExecutor(); ConcurrentClass target = new ConcurrentClass(); - assertThatNoException().isThrownBy(() -> taskExecutor.execute(target::concurrentOperation)); - assertThat(taskExecutor.execute(target::concurrentOperationWithResult)).isEqualTo("result"); - assertThatIOException().isThrownBy(() -> taskExecutor.execute(target::concurrentOperationWithException)); + assertThatNoException().isThrownBy(() -> executor.execute(target::concurrentOperation)); + assertThat(executor.execute(target::concurrentOperationWithResult)).isEqualTo("result"); + assertThatIOException().isThrownBy(() -> executor.execute(target::concurrentOperationWithException)); } @Test void withConcurrencyLimit() { - SyncTaskExecutor taskExecutor = new SyncTaskExecutor(); - taskExecutor.setConcurrencyLimit(2); + SyncTaskExecutor executor = new SyncTaskExecutor(); + executor.setConcurrencyLimit(2); ConcurrentClass target = new ConcurrentClass(); List> futures = new ArrayList<>(10); for (int i = 0; i < 10; i++) { - futures.add(CompletableFuture.runAsync(() -> taskExecutor.execute(target::concurrentOperation))); + futures.add(CompletableFuture.runAsync(() -> executor.execute(target::concurrentOperation))); } CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])).join(); assertThat(target.current).hasValue(0); @@ -61,14 +62,14 @@ class SyncTaskExecutorTests { @Test void withConcurrencyLimitAndResult() { - SyncTaskExecutor taskExecutor = new SyncTaskExecutor(); - taskExecutor.setConcurrencyLimit(2); + SyncTaskExecutor executor = new SyncTaskExecutor(); + executor.setConcurrencyLimit(2); ConcurrentClass target = new ConcurrentClass(); List> futures = new ArrayList<>(10); for (int i = 0; i < 10; i++) { futures.add(CompletableFuture.runAsync(() -> - assertThat(taskExecutor.execute(target::concurrentOperationWithResult)).isEqualTo("result"))); + assertThat(executor.execute(target::concurrentOperationWithResult)).isEqualTo("result"))); } CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])).join(); assertThat(target.current).hasValue(0); @@ -77,20 +78,41 @@ class SyncTaskExecutorTests { @Test void withConcurrencyLimitAndException() { - SyncTaskExecutor taskExecutor = new SyncTaskExecutor(); - taskExecutor.setConcurrencyLimit(2); + SyncTaskExecutor executor = new SyncTaskExecutor(); + executor.setConcurrencyLimit(2); ConcurrentClass target = new ConcurrentClass(); List> futures = new ArrayList<>(10); for (int i = 0; i < 10; i++) { futures.add(CompletableFuture.runAsync(() -> - assertThatIOException().isThrownBy(() -> taskExecutor.execute(target::concurrentOperationWithException)))); + assertThatIOException().isThrownBy(() -> executor.execute(target::concurrentOperationWithException)))); } CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])).join(); assertThat(target.current).hasValue(0); assertThat(target.counter).hasValue(10); } + @Test + void taskRejectedWhenConcurrencyLimitReached() throws Exception { + SyncTaskExecutor executor = new SyncTaskExecutor(); + executor.setConcurrencyLimit(2); + executor.setRejectTasksWhenLimitReached(true); + + ConcurrentClass target = new ConcurrentClass(); + List> futures = new ArrayList<>(10); + for (int i = 0; i < 2; i++) { + futures.add(CompletableFuture.runAsync(() -> executor.execute(target::concurrentOperation))); + } + Thread.sleep(10); + for (int i = 2; i < 10; i++) { + futures.add(CompletableFuture.runAsync(() -> + assertThatExceptionOfType(TaskRejectedException.class).isThrownBy(() -> executor.execute(target::concurrentOperation)))); + } + CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])).join(); + assertThat(target.current).hasValue(0); + assertThat(target.counter).hasValue(2); + } + static class ConcurrentClass { @@ -103,7 +125,7 @@ class SyncTaskExecutorTests { throw new IllegalStateException(); } try { - Thread.sleep(10); + Thread.sleep(100); } catch (InterruptedException ex) { throw new IllegalStateException(ex);