3838import org .mockito .Captor ;
3939import org .mockito .Mock ;
4040import org .mockito .junit .jupiter .MockitoExtension ;
41- import org .springframework .security .oauth2 .client .AuthorizedClientServiceOAuth2AuthorizedClientManager ;
42- import org .springframework .security .oauth2 .client .OAuth2AuthorizedClientService ;
4341import reactor .core .publisher .Mono ;
4442import reactor .util .context .Context ;
4543
6967import org .springframework .security .core .context .SecurityContextHolder ;
7068import org .springframework .security .core .context .SecurityContextHolderStrategy ;
7169import org .springframework .security .core .context .SecurityContextImpl ;
70+ import org .springframework .security .oauth2 .client .AuthorizedClientServiceOAuth2AuthorizedClientManager ;
7271import org .springframework .security .oauth2 .client .ClientAuthorizationException ;
7372import org .springframework .security .oauth2 .client .JwtBearerOAuth2AuthorizedClientProvider ;
7473import org .springframework .security .oauth2 .client .OAuth2AuthorizationFailureHandler ;
7574import org .springframework .security .oauth2 .client .OAuth2AuthorizedClient ;
7675import org .springframework .security .oauth2 .client .OAuth2AuthorizedClientProvider ;
7776import org .springframework .security .oauth2 .client .OAuth2AuthorizedClientProviderBuilder ;
77+ import org .springframework .security .oauth2 .client .OAuth2AuthorizedClientService ;
7878import org .springframework .security .oauth2 .client .RefreshTokenOAuth2AuthorizedClientProvider ;
7979import org .springframework .security .oauth2 .client .authentication .OAuth2AuthenticationToken ;
8080import org .springframework .security .oauth2 .client .endpoint .JwtBearerGrantRequest ;
@@ -137,7 +137,7 @@ public class ServletOAuth2AuthorizedClientExchangeFilterFunctionTests {
137137 private OAuth2AuthorizedClientRepository authorizedClientRepository ;
138138
139139 @ Mock
140- private OAuth2AuthorizedClientService oAuth2AuthorizedClientService ;
140+ private OAuth2AuthorizedClientService authorizedClientService ;
141141
142142 @ Mock
143143 private ClientRegistrationRepository clientRegistrationRepository ;
@@ -666,11 +666,12 @@ public void filterWhenClientRegistrationIdFromAuthenticationAndCustomPrincipalRe
666666 authentication , servletRequest );
667667 }
668668
669+ // gh-19421
669670 @ Test
670- public void filterWhenServletRequestNullClientRegistrationIdFromAuthenticationAndCustomPrincipalResolverThenAuthorizedClientResolved () {
671+ public void filterWhenServletRequestNullAndClientRegistrationIdFromAuthenticationAndCustomPrincipalResolverThenAuthorizedClientResolved () {
671672 this .function = new ServletOAuth2AuthorizedClientExchangeFilterFunction (
672673 new AuthorizedClientServiceOAuth2AuthorizedClientManager (this .clientRegistrationRepository ,
673- oAuth2AuthorizedClientService ));
674+ this . authorizedClientService ));
674675 this .function .setDefaultOAuth2AuthorizedClient (true );
675676 OAuth2User user = mock (OAuth2User .class );
676677 List <GrantedAuthority > authorities = AuthorityUtils .createAuthorityList ("ROLE_USER" );
@@ -680,8 +681,9 @@ public void filterWhenServletRequestNullClientRegistrationIdFromAuthenticationAn
680681 this .registration .getRegistrationId ());
681682 OAuth2AuthorizedClient authorizedClient = new OAuth2AuthorizedClient (this .registration , "principalName" ,
682683 this .accessToken );
683- given (this .clientRegistrationRepository .findByRegistrationId (any ())).willReturn (this .registration );
684- given (this .oAuth2AuthorizedClientService .loadAuthorizedClient (this .registration .getRegistrationId (),
684+ given (this .clientRegistrationRepository .findByRegistrationId (this .registration .getRegistrationId ()))
685+ .willReturn (this .registration );
686+ given (this .authorizedClientService .loadAuthorizedClient (this .registration .getRegistrationId (),
685687 initialAuthentication .getName ()))
686688 .willReturn (authorizedClient );
687689 final ClientRequest clientRequest = ClientRequest .create (HttpMethod .GET , URI .create ("https://example.com" ))
@@ -697,7 +699,7 @@ public void filterWhenServletRequestNullClientRegistrationIdFromAuthenticationAn
697699 assertThat (request .url ().toASCIIString ()).isEqualTo ("https://example.com" );
698700 assertThat (request .method ()).isEqualTo (HttpMethod .GET );
699701 assertThat (getBody (request )).isEmpty ();
700- verify (this .oAuth2AuthorizedClientService ).loadAuthorizedClient (this .registration .getRegistrationId (),
702+ verify (this .authorizedClientService ).loadAuthorizedClient (this .registration .getRegistrationId (),
701703 authentication .getName ());
702704 }
703705
0 commit comments