Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -26,9 +26,11 @@
import jakarta.servlet.http.HttpServletResponse;
import org.jspecify.annotations.Nullable;

import org.springframework.cloud.gateway.server.mvc.common.MvcUtils;
import org.springframework.http.HttpHeaders;
import org.springframework.http.HttpMethod;
import org.springframework.http.HttpStatusCode;
import org.springframework.http.client.ClientHttpResponse;
import org.springframework.util.LinkedMultiValueMap;
import org.springframework.util.MultiValueMap;
import org.springframework.web.context.request.ServletWebRequest;
Expand Down Expand Up @@ -91,6 +93,11 @@ public MultiValueMap<String, Cookie> cookies() {
HttpMethod httpMethod = HttpMethod.valueOf(request.getMethod());
if (SAFE_METHODS.contains(httpMethod)
&& servletWebRequest.checkNotModified(headers().getETag(), lastModified)) {
// Not-modified short-circuit skips writeToInternal, which is where
// RestClientProxyExchange / ClientHttpRequestFactoryProxyExchange
// close the upstream ClientHttpResponse and return the connection
// to the pool (GH-4259).
closeClientResponse(request);
return null;
}
else {
Expand All @@ -102,6 +109,19 @@ public MultiValueMap<String, Cookie> cookies() {
}
}

/**
* Closes a proxied {@link ClientHttpResponse} stored on the request when the response
* body will not be written (for example HTTP 304 Not Modified).
* @param request the current servlet request
*/
private static void closeClientResponse(HttpServletRequest request) {
Object clientResponse = request.getAttribute(MvcUtils.CLIENT_RESPONSE_ATTR);
if (clientResponse instanceof ClientHttpResponse clientHttpResponse) {
clientHttpResponse.close();
request.removeAttribute(MvcUtils.CLIENT_RESPONSE_ATTR);
}
}

private void writeStatusAndHeaders(HttpServletResponse response) {
response.setStatus(this.statusCode.value());
writeHeaders(response);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,44 @@ void exchangeWhenStreamingResponseCopyFailsThenDoesNotCloseClientResponse() {
assertThat(responseBody.closed).isTrue();
}

@Test
void exchangeWhenNotModifiedThenClosesClientResponse() throws Exception {
RestClient restClient = mock(RestClient.class);
RestClient.RequestBodyUriSpec requestSpec = mock(RestClient.RequestBodyUriSpec.class);
CloseAwareInputStream responseBody = new CloseAwareInputStream();
TestClientHttpResponse clientResponse = new TestClientHttpResponse(responseBody);
clientResponse.getHeaders().setContentType(MediaType.TEXT_PLAIN);
clientResponse.getHeaders().setETag("\"abc\"");

when(restClient.method(HttpMethod.GET)).thenReturn(requestSpec);
when(requestSpec.uri(any(URI.class))).thenReturn(requestSpec);
when(requestSpec.headers(any())).thenReturn(requestSpec);
when(requestSpec.exchange(any(), eq(false))).thenAnswer((invocation) -> {
RestClient.RequestHeadersSpec.ExchangeFunction<ServerResponse> exchangeFunction = invocation.getArgument(0);
return exchangeFunction.exchange(mock(HttpRequest.class), clientResponse);
});

RestClientProxyExchange proxyExchange = new RestClientProxyExchange(restClient, new GatewayMvcProperties());
MockHttpServletRequest servletRequest = MockMvcRequestBuilders.get("http://localhost/resource")
.header(HttpHeaders.IF_NONE_MATCH, "\"abc\"")
.buildRequest(null);
ServerRequest serverRequest = ServerRequest.create(servletRequest, Collections.emptyList());
ProxyExchange.Request request = proxyExchange.request(serverRequest)
.uri(URI.create("http://localhost:8781/resource"))
.build();

ServerResponse serverResponse = proxyExchange.exchange(request);
// Proxy responses carry the upstream ETag; apply it so checkNotModified can
// short-circuit.
serverResponse.headers().setETag("\"abc\"");

MockHttpServletResponse servletResponse = new MockHttpServletResponse();
serverResponse.writeTo(servletRequest, servletResponse, Collections::emptyList);

assertThat(servletResponse.getStatus()).isEqualTo(HttpStatus.NOT_MODIFIED.value());
assertThat(clientResponse.closed).isTrue();
}

private static final class ClientDisconnectedResponse extends MockHttpServletResponse {

private final ServletOutputStream outputStream = new ServletOutputStream() {
Expand Down Expand Up @@ -155,6 +193,8 @@ private static final class TestClientHttpResponse

private TestClientHttpResponse(CloseAwareInputStream body) {
this.body = body;
// Default content type used by existing streaming tests; other tests may
// override.
this.headers.setContentType(MediaType.TEXT_EVENT_STREAM);
}

Expand Down