diff --git a/src/main/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandler.java b/src/main/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandler.java index f4dac9d..c586acf 100644 --- a/src/main/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandler.java +++ b/src/main/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandler.java @@ -117,7 +117,7 @@ public class RecoverAnnotationRecoveryHandler implements MethodInvocationReco private boolean compareParameters(Object[] args, int argCount, Class[] parameterTypes) { if (argCount == (args.length + 1)) { int startingIndex = 0; - if (parameterTypes.length > 0 && parameterTypes[0] == Throwable.class) { + if (parameterTypes.length > 0 && Throwable.class.isAssignableFrom(parameterTypes[0])) { startingIndex = 1; } for (int i = startingIndex; i < parameterTypes.length; i++) { diff --git a/src/test/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandlerTests.java b/src/test/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandlerTests.java index d00da5c..3c764e3 100644 --- a/src/test/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandlerTests.java +++ b/src/test/java/org/springframework/retry/annotation/RecoverAnnotationRecoveryHandlerTests.java @@ -146,6 +146,19 @@ public class RecoverAnnotationRecoveryHandlerTests { handler.recover(new Object[] { "Randell" }, new RuntimeException("Planned"))); } + + @Test + public void multipleQualifyingRecoverMethodsExtendsThrowable(){ + Method foo = ReflectionUtils.findMethod(MultipleQualifyingRecoversExtendsThrowable.class, + "foo", String.class); + RecoverAnnotationRecoveryHandler handler = new RecoverAnnotationRecoveryHandler( + new MultipleQualifyingRecoversExtendsThrowable(), foo); + assertEquals(2, + handler.recover(new Object[] { "Kevin" }, new IllegalArgumentException("Planned"))); + assertEquals(3, + handler.recover(new Object[] { "Kevin" }, new UnsupportedOperationException("Planned"))); + + } private static class InAccessibleRecover { @@ -310,5 +323,29 @@ public class RecoverAnnotationRecoveryHandlerTests { } } + + protected static class MultipleQualifyingRecoversExtendsThrowable { + + @Retryable + public int foo(String name) { + return 0; + } + + @Recover + public int fooRecover(IllegalArgumentException e, String name) { + return 1; + } + + @Recover + public int barRecover(IllegalArgumentException e, String name) { + return 2; + } + + @Recover + public int bazRecover(UnsupportedOperationException e, String name) { + return 3; + } + + } }