From 6e6cf19e009fad41d69fbbc2c337a02c825c4a07 Mon Sep 17 00:00:00 2001 From: XiaoSeS <87064762+XiaoSeS@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:12:30 +0800 Subject: [PATCH] feat(auth): let providers override OAuth userinfo loading Some providers do not return a flat, standard userinfo payload, so DefaultOAuth2UserService cannot read them. Add ProviderOAuth2UserService so a provider can claim its own registration id and supply the loading step, while everything after it stays shared. The override runs inside the RemoteIdentityIoExecutor boundary added in R1-A, so a provider's HTTP call does not hold the surrounding transaction open. Registrations without an override keep using the default user service unchanged. Part of R1-A2 (public Provider adapters) per openspec/changes/enterprise-identity-platform/rollout-plan.md. Signed-off-by: XiaoSeS <87064762+XiaoSeS@users.noreply.github.com> --- .../auth/oauth/OAuthLoginFlowService.java | 27 ++- .../auth/oauth/ProviderOAuth2UserService.java | 14 ++ .../auth/oauth/OAuthLoginFlowServiceTest.java | 154 +++++++++++++++++- 3 files changed, 192 insertions(+), 3 deletions(-) create mode 100644 server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/ProviderOAuth2UserService.java diff --git a/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowService.java b/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowService.java index 9a2c2824..3b4d265e 100644 --- a/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowService.java +++ b/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowService.java @@ -36,6 +36,7 @@ import org.springframework.stereotype.Service; public class OAuthLoginFlowService { private final Map extractors; + private final Map userServiceOverrides; private final AccessPolicy accessPolicy; private final IdentityBindingService identityBindingService; private final LegacyPlatformIdentityCore identityCore; @@ -44,12 +45,14 @@ public class OAuthLoginFlowService { @Autowired public OAuthLoginFlowService(List extractorList, + List userServiceList, AccessPolicy accessPolicy, IdentityBindingService identityBindingService, LegacyPlatformIdentityCore identityCore, RemoteIdentityIoExecutor remoteIdentityIo) { this( extractorList, + userServiceList, accessPolicy, identityBindingService, identityCore, @@ -59,6 +62,7 @@ public class OAuthLoginFlowService { } OAuthLoginFlowService(List extractorList, + List userServiceList, AccessPolicy accessPolicy, IdentityBindingService identityBindingService, LegacyPlatformIdentityCore identityCore, @@ -66,6 +70,8 @@ public class OAuthLoginFlowService { RemoteIdentityIoExecutor remoteIdentityIo) { this.extractors = extractorList.stream() .collect(Collectors.toMap(OAuthClaimsExtractor::getProvider, Function.identity())); + this.userServiceOverrides = userServiceList.stream() + .collect(Collectors.toMap(ProviderOAuth2UserService::getProvider, Function.identity())); this.accessPolicy = accessPolicy; this.identityBindingService = identityBindingService; this.identityCore = identityCore; @@ -73,12 +79,30 @@ public class OAuthLoginFlowService { this.remoteIdentityIo = remoteIdentityIo; } + OAuthLoginFlowService(List extractorList, + AccessPolicy accessPolicy, + IdentityBindingService identityBindingService, + LegacyPlatformIdentityCore identityCore, + OAuth2UserService delegate, + RemoteIdentityIoExecutor remoteIdentityIo) { + this( + extractorList, + List.of(), + accessPolicy, + identityBindingService, + identityCore, + delegate, + remoteIdentityIo + ); + } + OAuthLoginFlowService(List extractorList, AccessPolicy accessPolicy, IdentityBindingService identityBindingService, LegacyPlatformIdentityCore identityCore) { this( extractorList, + List.of(), accessPolicy, identityBindingService, identityCore, @@ -95,8 +119,9 @@ public class OAuthLoginFlowService { public AuthenticatedLoginContext loadLoginContext(OAuth2UserRequest request) { LoadedProviderIdentity loadedIdentity = remoteIdentityIo.execute(() -> { - OAuth2User upstreamUser = delegate.loadUser(request); String registrationId = request.getClientRegistration().getRegistrationId(); + ProviderOAuth2UserService override = userServiceOverrides.get(registrationId); + OAuth2User upstreamUser = (override != null ? override : delegate).loadUser(request); OAuthClaimsExtractor extractor = extractors.get(registrationId); if (extractor == null) { throw new OAuth2AuthenticationException( diff --git a/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/ProviderOAuth2UserService.java b/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/ProviderOAuth2UserService.java new file mode 100644 index 00000000..bf587e81 --- /dev/null +++ b/server/skillhub-auth/src/main/java/com/iflytek/skillhub/auth/oauth/ProviderOAuth2UserService.java @@ -0,0 +1,14 @@ +package com.iflytek.skillhub.auth.oauth; + +import org.springframework.security.oauth2.client.userinfo.OAuth2UserRequest; +import org.springframework.security.oauth2.client.userinfo.OAuth2UserService; +import org.springframework.security.oauth2.core.user.OAuth2User; + +/** + * Strategy interface for provider-specific OAuth user loading. Implementations override the + * default user info loading for providers whose endpoints deviate from the standard + * flat-attribute response format. + */ +public interface ProviderOAuth2UserService extends OAuth2UserService { + String getProvider(); +} diff --git a/server/skillhub-auth/src/test/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowServiceTest.java b/server/skillhub-auth/src/test/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowServiceTest.java index 29a280f1..e9046ef7 100644 --- a/server/skillhub-auth/src/test/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowServiceTest.java +++ b/server/skillhub-auth/src/test/java/com/iflytek/skillhub/auth/oauth/OAuthLoginFlowServiceTest.java @@ -102,6 +102,148 @@ class OAuthLoginFlowServiceTest { verify(delegate).loadUser(request); } + @Test + void loadLoginContext_prefersProviderUserServiceOverrideInsideRemoteIoBoundary() { + OAuthClaims claims = claims("feishu", "ou_1"); + OAuthClaimsExtractor extractor = new OAuthClaimsExtractor() { + @Override + public String getProvider() { + return "feishu"; + } + + @Override + public OAuthClaims extract(OAuth2UserRequest request, OAuth2User user) { + return claims; + } + }; + OAuth2User overrideUser = new DefaultOAuth2User( + List.of(new SimpleGrantedAuthority("OAUTH_USER")), + Map.of("open_id", "ou_1"), + "open_id" + ); + AtomicInteger boundaryCalls = new AtomicInteger(); + AtomicInteger overrideCallsInsideBoundary = new AtomicInteger(); + RemoteIdentityIoExecutor remoteIdentityIo = new RemoteIdentityIoExecutor() { + @Override + public T execute(java.util.function.Supplier operation) { + boundaryCalls.incrementAndGet(); + return operation.get(); + } + }; + ProviderOAuth2UserService override = new ProviderOAuth2UserService() { + @Override + public String getProvider() { + return "feishu"; + } + + @Override + public OAuth2User loadUser(OAuth2UserRequest request) { + // Records the boundary state at call time: a provider override must run inside the + // remote-IO boundary, otherwise its HTTP call would hold the surrounding transaction. + if (boundaryCalls.get() == 1) { + overrideCallsInsideBoundary.incrementAndGet(); + } + return overrideUser; + } + }; + AccessPolicy accessPolicy = mock(AccessPolicy.class); + IdentityBindingService identityBindingService = mock(IdentityBindingService.class); + LegacyPlatformIdentityCore identityCore = mock(LegacyPlatformIdentityCore.class); + OAuth2UserService delegate = mock(); + PlatformPrincipal principal = new PlatformPrincipal( + "usr_2", "zhangsan", null, null, "feishu", Set.of("USER") + ); + OAuthLoginFlowService service = new OAuthLoginFlowService( + List.of(extractor), + List.of(override), + accessPolicy, + identityBindingService, + identityCore, + delegate, + remoteIdentityIo + ); + OAuth2UserRequest request = oauthUserRequest("feishu"); + when(accessPolicy.evaluate(claims)).thenReturn(AccessDecision.ALLOW); + when(identityCore.evaluate(claims)).thenReturn(LegacyPlatformIdentityDecision.legacy()); + when(identityBindingService.bindOrCreate(claims, UserStatus.ACTIVE)).thenReturn(principal); + + OAuthLoginFlowService.AuthenticatedLoginContext result = service.loadLoginContext(request); + + assertThat(result.upstreamUser()).isSameAs(overrideUser); + assertThat(result.principal()).isSameAs(principal); + assertThat(boundaryCalls).hasValue(1); + assertThat(overrideCallsInsideBoundary).hasValue(1); + // The default user service must not be consulted when an override claims the registration. + verify(delegate, never()).loadUser(request); + } + + @Test + void loadLoginContext_fallsBackToDefaultUserServiceForUnclaimedProviders() { + OAuthClaims claims = claims(); + OAuthClaimsExtractor extractor = new OAuthClaimsExtractor() { + @Override + public String getProvider() { + return "github"; + } + + @Override + public OAuthClaims extract(OAuth2UserRequest request, OAuth2User user) { + return claims; + } + }; + ProviderOAuth2UserService unrelatedOverride = new ProviderOAuth2UserService() { + @Override + public String getProvider() { + return "feishu"; + } + + @Override + public OAuth2User loadUser(OAuth2UserRequest request) { + throw new AssertionError("Feishu override must not handle a GitHub login"); + } + }; + AccessPolicy accessPolicy = mock(AccessPolicy.class); + IdentityBindingService identityBindingService = mock(IdentityBindingService.class); + LegacyPlatformIdentityCore identityCore = mock(LegacyPlatformIdentityCore.class); + OAuth2UserService delegate = mock(); + OAuth2User upstreamUser = new DefaultOAuth2User( + List.of(new SimpleGrantedAuthority("OAUTH_USER")), + Map.of("id", "gh_1"), + "id" + ); + PlatformPrincipal principal = new PlatformPrincipal( + "usr_1", "alice", "alice@example.com", null, "github", Set.of("USER") + ); + OAuthLoginFlowService service = new OAuthLoginFlowService( + List.of(extractor), + List.of(unrelatedOverride), + accessPolicy, + identityBindingService, + identityCore, + delegate, + directRemoteIo() + ); + OAuth2UserRequest request = oauthUserRequest(); + when(delegate.loadUser(request)).thenReturn(upstreamUser); + when(accessPolicy.evaluate(claims)).thenReturn(AccessDecision.ALLOW); + when(identityCore.evaluate(claims)).thenReturn(LegacyPlatformIdentityDecision.legacy()); + when(identityBindingService.bindOrCreate(claims, UserStatus.ACTIVE)).thenReturn(principal); + + OAuthLoginFlowService.AuthenticatedLoginContext result = service.loadLoginContext(request); + + assertThat(result.upstreamUser()).isSameAs(upstreamUser); + verify(delegate).loadUser(request); + } + + private static RemoteIdentityIoExecutor directRemoteIo() { + return new RemoteIdentityIoExecutor() { + @Override + public T execute(java.util.function.Supplier operation) { + return operation.get(); + } + }; + } + @ParameterizedTest @EnumSource(IdentityCoreMode.class) void authenticate_preservesPrincipalAcrossLegacyShadowAndActiveModes(IdentityCoreMode mode) { @@ -320,12 +462,20 @@ class OAuthLoginFlowServiceTest { ); } + private static OAuthClaims claims(String provider, String subject) { + return new OAuthClaims(provider, subject, null, false, subject, Map.of()); + } + private static OAuth2UserRequest oauthUserRequest() { - ClientRegistration registration = ClientRegistration.withRegistrationId("github") + return oauthUserRequest("github"); + } + + private static OAuth2UserRequest oauthUserRequest(String registrationId) { + ClientRegistration registration = ClientRegistration.withRegistrationId(registrationId) .clientId("client") .clientSecret("secret") .authorizationGrantType(AuthorizationGrantType.AUTHORIZATION_CODE) - .redirectUri("https://skillhub.example/login/oauth2/code/github") + .redirectUri("https://skillhub.example/login/oauth2/code/" + registrationId) .authorizationUri("https://github.example/oauth/authorize") .tokenUri("https://github.example/oauth/token") .userInfoUri("https://github.example/user")