diff --git a/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/AccountServiceImpl.java b/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/AccountServiceImpl.java index aa6a11284b9..3b744697645 100644 --- a/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/AccountServiceImpl.java +++ b/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/AccountServiceImpl.java @@ -25,10 +25,12 @@ import com.usthe.sureness.util.Md5Util; import com.usthe.sureness.util.SurenessContextHolder; import io.jsonwebtoken.Claims; + import java.util.HashMap; import java.util.List; import java.util.Map; import javax.naming.AuthenticationException; + import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; import org.apache.hertzbeat.common.util.JsonUtil; @@ -46,14 +48,26 @@ @Order(value = Ordered.HIGHEST_PRECEDENCE) @Slf4j public class AccountServiceImpl implements AccountService { + + private static final String REFRESH_CLAIM = "refresh"; + /** * Token validity time in seconds */ private static final long PERIOD_TIME = 3600L; + /** * account data provider */ - private final SurenessAccountProvider accountProvider = new DocumentAccountProvider(); + private final SurenessAccountProvider accountProvider; + + public AccountServiceImpl() { + this(new DocumentAccountProvider()); + } + + public AccountServiceImpl(SurenessAccountProvider accountProvider) { + this.accountProvider = accountProvider; + } @Override public Map authGetToken(LoginDto loginDto) throws AuthenticationException { @@ -75,10 +89,8 @@ public Map authGetToken(LoginDto loginDto) throws Authentication // Get the roles the user has - rbac List roles = account.getOwnRoles(); // Issue TOKEN - String issueToken = JsonWebTokenUtil.issueJwt(loginDto.getIdentifier(), PERIOD_TIME, roles); - Map customClaimMap = new HashMap<>(1); - customClaimMap.put("refresh", true); - String issueRefresh = JsonWebTokenUtil.issueJwt(loginDto.getIdentifier(), PERIOD_TIME << 5, customClaimMap); + String issueToken = issueAccessToken(loginDto.getIdentifier(), roles, PERIOD_TIME); + String issueRefresh = issueRefreshToken(loginDto.getIdentifier(), PERIOD_TIME << 5); Map resp = new HashMap<>(2); resp.put("token", issueToken); resp.put("refreshToken", issueRefresh); @@ -90,18 +102,21 @@ public Map authGetToken(LoginDto loginDto) throws Authentication @Override public RefreshTokenResponse refreshToken(String refreshToken) throws Exception { Claims claims = JsonWebTokenUtil.parseJwt(refreshToken); - String userId = String.valueOf(claims.getSubject()); - boolean isRefresh = claims.get("refresh", Boolean.class); - if (StringUtils.isBlank(userId) || !isRefresh) { + String userId = claims.getSubject(); + Boolean isRefresh = claims.get(REFRESH_CLAIM, Boolean.class); + if (StringUtils.isBlank(userId) || !Boolean.TRUE.equals(isRefresh)) { throw new AuthenticationException("Illegal Refresh Token"); } SurenessAccount account = accountProvider.loadAccount(userId); if (account == null) { throw new AuthenticationException("Not Exists This Token Mapping Account"); } + if (account.isDisabledAccount() || account.isExcessiveAttempts()) { + throw new AuthenticationException("Expired or Illegal Account"); + } List roles = account.getOwnRoles(); - String issueToken = issueToken(userId, roles, PERIOD_TIME); - String issueRefresh = issueToken(userId, roles, PERIOD_TIME << 5); + String issueToken = issueAccessToken(userId, roles, PERIOD_TIME); + String issueRefresh = issueRefreshToken(userId, PERIOD_TIME << 5); return new RefreshTokenResponse(issueToken, issueRefresh); } @@ -113,13 +128,24 @@ public String generateToken() throws AuthenticationException { if (account == null) { throw new AuthenticationException("Not Exists This Token Mapping Account"); } + if (account.isDisabledAccount() || account.isExcessiveAttempts()) { + throw new AuthenticationException("Expired or Illegal Account"); + } List roles = account.getOwnRoles(); - return issueToken(userId, roles, null); + return issueApiToken(userId, roles); } - private String issueToken(String userId, List roles, Long expirationMillis) { + private String issueAccessToken(String userId, List roles, Long expirationMillis) { + return JsonWebTokenUtil.issueJwt(userId, expirationMillis, roles, new HashMap<>(0)); + } + + private String issueRefreshToken(String userId, Long expirationMillis) { Map customClaimMap = new HashMap<>(1); - customClaimMap.put("refresh", true); - return JsonWebTokenUtil.issueJwt(userId, expirationMillis, roles, customClaimMap); + customClaimMap.put(REFRESH_CLAIM, true); + return JsonWebTokenUtil.issueJwt(userId, expirationMillis, customClaimMap); + } + + private String issueApiToken(String userId, List roles) { + return issueAccessToken(userId, roles, null); } } diff --git a/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/AccountServiceTest.java b/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/AccountServiceTest.java index c4804a116a6..3848af58c4b 100644 --- a/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/AccountServiceTest.java +++ b/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/AccountServiceTest.java @@ -19,14 +19,18 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.mockito.Mockito.mockStatic; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; import com.usthe.sureness.provider.DefaultAccount; import com.usthe.sureness.provider.SurenessAccount; import com.usthe.sureness.provider.SurenessAccountProvider; -import com.usthe.sureness.provider.ducument.DocumentAccountProvider; +import com.usthe.sureness.subject.SubjectSum; import com.usthe.sureness.util.JsonWebTokenUtil; import com.usthe.sureness.util.Md5Util; +import com.usthe.sureness.util.SurenessContextHolder; +import io.jsonwebtoken.Claims; import io.jsonwebtoken.MalformedJwtException; import java.util.Collections; import java.util.List; @@ -65,8 +69,8 @@ class AccountServiceTest { @BeforeEach void setUp() { - accountProvider = mock(DocumentAccountProvider.class); - accountService = new AccountServiceImpl(); + accountProvider = mock(SurenessAccountProvider.class); + accountService = new AccountServiceImpl(accountProvider); JsonWebTokenUtil.setDefaultSecretKey(jwt); } @@ -136,6 +140,89 @@ void testRefreshTokenWithValidToken() throws Exception { assertNotNull(response); assertNotNull(response.getToken()); assertNotNull(response.getRefreshToken()); + Claims accessClaims = JsonWebTokenUtil.parseJwt(response.getToken()); + Claims refreshClaims = JsonWebTokenUtil.parseJwt(response.getRefreshToken()); + assertNull(accessClaims.get("refresh", Boolean.class)); + assertEquals(Boolean.TRUE, refreshClaims.get("refresh", Boolean.class)); + } + + @Test + void testRefreshTokenRejectsAccessToken() { + String userId = "admin"; + String accessToken = JsonWebTokenUtil.issueJwt(userId, 3600L, roles); + + Assertions.assertThrows( + AuthenticationException.class, + () -> accountService.refreshToken(accessToken) + ); + } + + @Test + void testRefreshTokenRejectsDisabledAccount() { + String userId = "admin"; + String refreshToken = JsonWebTokenUtil.issueJwt(userId, 3600L, Collections.singletonMap("refresh", true)); + SurenessAccount account = DefaultAccount.builder("app1") + .setPassword(Md5Util.md5(password + salt)) + .setSalt(salt) + .setOwnRoles(roles) + .setDisabledAccount(Boolean.TRUE) + .setExcessiveAttempts(Boolean.FALSE) + .build(); + when(accountProvider.loadAccount(userId)).thenReturn(account); + + Assertions.assertThrows( + AuthenticationException.class, + () -> accountService.refreshToken(refreshToken) + ); + } + + @Test + void testGenerateTokenCannotRefresh() throws Exception { + SurenessAccount account = DefaultAccount.builder("app1") + .setPassword(Md5Util.md5(password + salt)) + .setSalt(salt) + .setOwnRoles(roles) + .setDisabledAccount(Boolean.FALSE) + .setExcessiveAttempts(Boolean.FALSE) + .build(); + when(accountProvider.loadAccount(identifier)).thenReturn(account); + SubjectSum subjectSum = mock(SubjectSum.class); + when(subjectSum.getPrincipal()).thenReturn(identifier); + + try (var mockedStatic = mockStatic(SurenessContextHolder.class)) { + mockedStatic.when(SurenessContextHolder::getBindSubject).thenReturn(subjectSum); + + String token = accountService.generateToken(); + + assertNull(JsonWebTokenUtil.parseJwt(token).get("refresh", Boolean.class)); + Assertions.assertThrows( + AuthenticationException.class, + () -> accountService.refreshToken(token) + ); + } + } + + @Test + void testGenerateTokenRejectsDisabledAccount() { + SurenessAccount account = DefaultAccount.builder("app1") + .setPassword(Md5Util.md5(password + salt)) + .setSalt(salt) + .setOwnRoles(roles) + .setDisabledAccount(Boolean.TRUE) + .setExcessiveAttempts(Boolean.FALSE) + .build(); + when(accountProvider.loadAccount(identifier)).thenReturn(account); + SubjectSum subjectSum = mock(SubjectSum.class); + when(subjectSum.getPrincipal()).thenReturn(identifier); + + try (var mockedStatic = mockStatic(SurenessContextHolder.class)) { + mockedStatic.when(SurenessContextHolder::getBindSubject).thenReturn(subjectSum); + + Assertions.assertThrows( + AuthenticationException.class, + () -> accountService.generateToken() + ); + } } @Test