diff --git a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/RetryFilterFunctions.java b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/RetryFilterFunctions.java index 764bc664..91380c90 100644 --- a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/RetryFilterFunctions.java +++ b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/RetryFilterFunctions.java @@ -26,6 +26,8 @@ import java.util.Set; import java.util.concurrent.TimeoutException; import java.util.function.Consumer; +import org.springframework.core.NestedRuntimeException; +import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; import org.springframework.http.HttpStatusCode; import org.springframework.retry.RetryContext; @@ -35,8 +37,8 @@ import org.springframework.retry.policy.NeverRetryPolicy; import org.springframework.retry.policy.SimpleRetryPolicy; import org.springframework.retry.support.RetryTemplate; import org.springframework.retry.support.RetryTemplateBuilder; -import org.springframework.web.client.HttpServerErrorException; import org.springframework.web.servlet.function.HandlerFilterFunction; +import org.springframework.web.servlet.function.ServerRequest; import org.springframework.web.servlet.function.ServerResponse; public abstract class RetryFilterFunctions { @@ -56,14 +58,16 @@ public abstract class RetryFilterFunctions { Map, Boolean> retryableExceptions = new HashMap<>(); config.getExceptions().forEach(exception -> retryableExceptions.put(exception, true)); SimpleRetryPolicy simpleRetryPolicy = new SimpleRetryPolicy(config.getRetries(), retryableExceptions); - compositeRetryPolicy.setPolicies( - Arrays.asList(simpleRetryPolicy, new HttpStatusRetryPolicy(config)).toArray(new RetryPolicy[0])); + compositeRetryPolicy + .setPolicies(Arrays.asList(simpleRetryPolicy, new HttpRetryPolicy(config)).toArray(new RetryPolicy[0])); RetryTemplate retryTemplate = retryTemplateBuilder.customPolicy(compositeRetryPolicy).build(); return (request, next) -> retryTemplate.execute(context -> { ServerResponse serverResponse = next.handle(request); - if (isRetryableStatusCode(serverResponse.statusCode(), config)) { - throw new HttpServerErrorException(serverResponse.statusCode()); + if (isRetryableStatusCode(serverResponse.statusCode(), config) + && isRetryableMethod(request.method(), config)) { + // use this to transfer information to HttpStatusRetryPolicy + throw new RetryException(request, serverResponse); } return serverResponse; }); @@ -73,19 +77,24 @@ public abstract class RetryFilterFunctions { return config.getSeries().stream().anyMatch(series -> HttpStatus.Series.resolve(httpStatus.value()) == series); } - public static class HttpStatusRetryPolicy extends NeverRetryPolicy { + private static boolean isRetryableMethod(HttpMethod method, RetryConfig config) { + return config.methods.contains(method); + } + + public static class HttpRetryPolicy extends NeverRetryPolicy { private final RetryConfig config; - public HttpStatusRetryPolicy(RetryConfig config) { + public HttpRetryPolicy(RetryConfig config) { this.config = config; } @Override public boolean canRetry(RetryContext context) { // TODO: custom exception - if (context.getLastThrowable() instanceof HttpServerErrorException e) { - return isRetryableStatusCode(e.getStatusCode(), config); + if (context.getLastThrowable() instanceof RetryException e) { + return isRetryableStatusCode(e.getResponse().statusCode(), config) + && isRetryableMethod(e.getRequest().method(), config); } return super.canRetry(context); } @@ -99,9 +108,12 @@ public abstract class RetryFilterFunctions { private Set series = new HashSet<>(List.of(HttpStatus.Series.SERVER_ERROR)); private Set> exceptions = new HashSet<>( - List.of(IOException.class, TimeoutException.class, HttpServerErrorException.class)); + List.of(IOException.class, TimeoutException.class, RetryException.class)); + + private Set methods = new HashSet<>(List.of(HttpMethod.GET)); // TODO: individual statuses + // TODO: backoff // TODO: support more Spring Retry policies public int getRetries() { @@ -141,6 +153,42 @@ public abstract class RetryFilterFunctions { return this; } + public Set getMethods() { + return methods; + } + + public RetryConfig setMethods(Set methods) { + this.methods = methods; + return this; + } + + public RetryConfig addMethods(HttpMethod... methods) { + this.methods.addAll(Arrays.asList(methods)); + return this; + } + + } + + private static class RetryException extends NestedRuntimeException { + + private final ServerRequest request; + + private final ServerResponse response; + + RetryException(ServerRequest request, ServerResponse response) { + super(null); + this.request = request; + this.response = response; + } + + public ServerRequest getRequest() { + return request; + } + + public ServerResponse getResponse() { + return response; + } + } } diff --git a/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java b/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java index d5da291e..64f87c58 100644 --- a/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java +++ b/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java @@ -802,6 +802,7 @@ public class ServerMvcIntegrationTests { .route(path("/retry"), http()) .before(new LocalServerPortUriResolver()) .filter(retry(3)) + //.filter(retry(config -> config.setRetries(3).setSeries(Set.of(HttpStatus.Series.SERVER_ERROR)).setMethods(Set.of(HttpMethod.GET, HttpMethod.POST)))) .filter(prefixPath("/do")) .build(); // @formatter:on