Fix to allow null servlet request/response in ServletOAuth2AuthorizedClientExchangeFilterFunction
Issue gh-17819 Closes gh-19421 Signed-off-by: Peter Phillips <5099053+petergphillips@users.noreply.github.com>
This commit is contained in:
committed by
Joe Grandja
parent
3cf0867955
commit
d90714f07d
+2
-8
@@ -511,9 +511,6 @@ public final class ServletOAuth2AuthorizedClientExchangeFilterFunction implement
|
||||
}
|
||||
HttpServletRequest servletRequest = getRequest(attrs);
|
||||
HttpServletResponse servletResponse = getResponse(attrs);
|
||||
if (servletRequest == null || servletResponse == null) {
|
||||
return Mono.empty();
|
||||
}
|
||||
OAuth2AuthorizeRequest.Builder builder = OAuth2AuthorizeRequest.withClientRegistrationId(clientRegistrationId)
|
||||
.principal(authentication);
|
||||
builder.attributes((attributes) -> addToAttributes(attributes, servletRequest, servletResponse));
|
||||
@@ -538,9 +535,6 @@ public final class ServletOAuth2AuthorizedClientExchangeFilterFunction implement
|
||||
}
|
||||
HttpServletRequest servletRequest = getRequest(attrs);
|
||||
HttpServletResponse servletResponse = getResponse(attrs);
|
||||
if (servletRequest == null || servletResponse == null) {
|
||||
return Mono.just(authorizedClient);
|
||||
}
|
||||
OAuth2AuthorizeRequest.Builder builder = OAuth2AuthorizeRequest.withAuthorizedClient(authorizedClient)
|
||||
.principal(authentication);
|
||||
builder.attributes((attributes) -> addToAttributes(attributes, servletRequest, servletResponse));
|
||||
@@ -552,8 +546,8 @@ public final class ServletOAuth2AuthorizedClientExchangeFilterFunction implement
|
||||
.subscribeOn(Schedulers.boundedElastic());
|
||||
}
|
||||
|
||||
private void addToAttributes(Map<String, Object> attributes, HttpServletRequest servletRequest,
|
||||
HttpServletResponse servletResponse) {
|
||||
private void addToAttributes(Map<String, Object> attributes, @Nullable HttpServletRequest servletRequest,
|
||||
@Nullable HttpServletResponse servletResponse) {
|
||||
if (servletRequest != null) {
|
||||
attributes.put(HTTP_SERVLET_REQUEST_ATTR_NAME, servletRequest);
|
||||
}
|
||||
|
||||
+63
@@ -40,8 +40,10 @@ import org.springframework.mock.web.MockHttpServletResponse;
|
||||
import org.springframework.security.authentication.TestingAuthenticationToken;
|
||||
import org.springframework.security.core.Authentication;
|
||||
import org.springframework.security.core.context.SecurityContextHolder;
|
||||
import org.springframework.security.oauth2.client.AuthorizedClientServiceOAuth2AuthorizedClientManager;
|
||||
import org.springframework.security.oauth2.client.InMemoryOAuth2AuthorizedClientService;
|
||||
import org.springframework.security.oauth2.client.OAuth2AuthorizedClient;
|
||||
import org.springframework.security.oauth2.client.OAuth2AuthorizedClientService;
|
||||
import org.springframework.security.oauth2.client.registration.ClientRegistration;
|
||||
import org.springframework.security.oauth2.client.registration.ClientRegistrationRepository;
|
||||
import org.springframework.security.oauth2.client.registration.TestClientRegistrations;
|
||||
@@ -166,6 +168,67 @@ public class ServletOAuth2AuthorizedClientExchangeFilterFunctionITests {
|
||||
assertThat(authorizedClientCaptor.getValue().getClientRegistration()).isSameAs(clientRegistration);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void requestWhenNoServletRequestThenAuthorizeAndSendRequest() {
|
||||
RequestContextHolder.resetRequestAttributes();
|
||||
InMemoryOAuth2AuthorizedClientService delegate = new InMemoryOAuth2AuthorizedClientService(
|
||||
this.clientRegistrationRepository);
|
||||
OAuth2AuthorizedClientService clientService = spy(new OAuth2AuthorizedClientService() {
|
||||
@Override
|
||||
public <T extends OAuth2AuthorizedClient> T loadAuthorizedClient(String clientRegistrationId,
|
||||
String principal) {
|
||||
return delegate.loadAuthorizedClient(clientRegistrationId, principal);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void saveAuthorizedClient(OAuth2AuthorizedClient authorizedClient, Authentication principal) {
|
||||
delegate.saveAuthorizedClient(authorizedClient, principal);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void removeAuthorizedClient(String clientRegistrationId, String principal) {
|
||||
delegate.removeAuthorizedClient(clientRegistrationId, principal);
|
||||
}
|
||||
});
|
||||
this.authorizedClientFilter = new ServletOAuth2AuthorizedClientExchangeFilterFunction(
|
||||
new AuthorizedClientServiceOAuth2AuthorizedClientManager(this.clientRegistrationRepository,
|
||||
clientService));
|
||||
this.webClient = WebClient.builder().apply(this.authorizedClientFilter.oauth2Configuration()).build();
|
||||
|
||||
// @formatter:off
|
||||
String accessTokenResponse = "{\n"
|
||||
+ " \"access_token\": \"access-token-1234\",\n"
|
||||
+ " \"token_type\": \"bearer\",\n"
|
||||
+ " \"expires_in\": \"3600\",\n"
|
||||
+ " \"scope\": \"read write\"\n"
|
||||
+ "}\n";
|
||||
String clientResponse = "{\n"
|
||||
+ " \"attribute1\": \"value1\",\n"
|
||||
+ " \"attribute2\": \"value2\"\n"
|
||||
+ "}\n";
|
||||
// @formatter:on
|
||||
this.server.enqueue(jsonResponse(accessTokenResponse));
|
||||
this.server.enqueue(jsonResponse(clientResponse));
|
||||
ClientRegistration clientRegistration = TestClientRegistrations.clientCredentials()
|
||||
.tokenUri(this.serverUrl)
|
||||
.build();
|
||||
given(this.clientRegistrationRepository.findByRegistrationId(eq(clientRegistration.getRegistrationId())))
|
||||
.willReturn(clientRegistration);
|
||||
|
||||
this.webClient.get()
|
||||
.uri(this.serverUrl)
|
||||
.attributes(ServletOAuth2AuthorizedClientExchangeFilterFunction
|
||||
.clientRegistrationId(clientRegistration.getRegistrationId()))
|
||||
.retrieve()
|
||||
.bodyToMono(String.class)
|
||||
.block();
|
||||
assertThat(this.server.getRequestCount()).isEqualTo(2);
|
||||
ArgumentCaptor<OAuth2AuthorizedClient> authorizedClientCaptor = ArgumentCaptor
|
||||
.forClass(OAuth2AuthorizedClient.class);
|
||||
verify(clientService).saveAuthorizedClient(authorizedClientCaptor.capture(), eq(this.authentication));
|
||||
assertThat(authorizedClientCaptor.getValue().getClientRegistration()).isSameAs(clientRegistration);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void requestWhenAuthorizedButExpiredThenRefreshAndSendRequest() {
|
||||
// @formatter:off
|
||||
|
||||
+40
@@ -38,6 +38,8 @@ import org.mockito.ArgumentCaptor;
|
||||
import org.mockito.Captor;
|
||||
import org.mockito.Mock;
|
||||
import org.mockito.junit.jupiter.MockitoExtension;
|
||||
import org.springframework.security.oauth2.client.AuthorizedClientServiceOAuth2AuthorizedClientManager;
|
||||
import org.springframework.security.oauth2.client.OAuth2AuthorizedClientService;
|
||||
import reactor.core.publisher.Mono;
|
||||
import reactor.util.context.Context;
|
||||
|
||||
@@ -134,6 +136,9 @@ public class ServletOAuth2AuthorizedClientExchangeFilterFunctionTests {
|
||||
@Mock
|
||||
private OAuth2AuthorizedClientRepository authorizedClientRepository;
|
||||
|
||||
@Mock
|
||||
private OAuth2AuthorizedClientService oAuth2AuthorizedClientService;
|
||||
|
||||
@Mock
|
||||
private ClientRegistrationRepository clientRegistrationRepository;
|
||||
|
||||
@@ -661,6 +666,41 @@ public class ServletOAuth2AuthorizedClientExchangeFilterFunctionTests {
|
||||
authentication, servletRequest);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void filterWhenServletRequestNullClientRegistrationIdFromAuthenticationAndCustomPrincipalResolverThenAuthorizedClientResolved() {
|
||||
this.function = new ServletOAuth2AuthorizedClientExchangeFilterFunction(
|
||||
new AuthorizedClientServiceOAuth2AuthorizedClientManager(this.clientRegistrationRepository,
|
||||
oAuth2AuthorizedClientService));
|
||||
this.function.setDefaultOAuth2AuthorizedClient(true);
|
||||
OAuth2User user = mock(OAuth2User.class);
|
||||
List<GrantedAuthority> authorities = AuthorityUtils.createAuthorityList("ROLE_USER");
|
||||
OAuth2AuthenticationToken initialAuthentication = new OAuth2AuthenticationToken(user, authorities,
|
||||
"initial-registration-id");
|
||||
OAuth2AuthenticationToken authentication = new OAuth2AuthenticationToken(user, authorities,
|
||||
this.registration.getRegistrationId());
|
||||
OAuth2AuthorizedClient authorizedClient = new OAuth2AuthorizedClient(this.registration, "principalName",
|
||||
this.accessToken);
|
||||
given(this.clientRegistrationRepository.findByRegistrationId(any())).willReturn(this.registration);
|
||||
given(this.oAuth2AuthorizedClientService.loadAuthorizedClient(this.registration.getRegistrationId(),
|
||||
initialAuthentication.getName()))
|
||||
.willReturn(authorizedClient);
|
||||
final ClientRequest clientRequest = ClientRequest.create(HttpMethod.GET, URI.create("https://example.com"))
|
||||
.build();
|
||||
this.function.setPrincipalResolver((request) -> authentication);
|
||||
this.function.filter(clientRequest, this.exchange)
|
||||
.contextWrite(context(null, null, initialAuthentication))
|
||||
.block();
|
||||
List<ClientRequest> requests = this.exchange.getRequests();
|
||||
assertThat(requests).hasSize(1);
|
||||
ClientRequest request = requests.get(0);
|
||||
assertThat(request.headers().getFirst(HttpHeaders.AUTHORIZATION)).isEqualTo("Bearer token-0");
|
||||
assertThat(request.url().toASCIIString()).isEqualTo("https://example.com");
|
||||
assertThat(request.method()).isEqualTo(HttpMethod.GET);
|
||||
assertThat(getBody(request)).isEmpty();
|
||||
verify(this.oAuth2AuthorizedClientService).loadAuthorizedClient(this.registration.getRegistrationId(),
|
||||
authentication.getName());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void filterWhenUnauthorizedThenInvokeFailureHandler() {
|
||||
assertHttpStatusInvokesFailureHandler(HttpStatus.UNAUTHORIZED, OAuth2ErrorCodes.INVALID_TOKEN);
|
||||
|
||||
Reference in New Issue
Block a user