diff --git a/src/main/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerability.java b/src/main/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerability.java index d23f59e2e..47578e796 100644 --- a/src/main/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerability.java +++ b/src/main/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerability.java @@ -116,8 +116,8 @@ public ResponseEntity> level2( String actualToken = cookieToken; try { if (actualToken != null && loggedInUser != null) { - idorLoginService.decodeToken(actualToken); - User profile = fetchUserById(loggedInUser); + User decodedUser = idorLoginService.decodeToken(actualToken); + User profile = fetchUserById(decodedUser.getUserId()); if (profile == null) { return response(USER_NOT_FOUND, false); } diff --git a/src/test/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerabilityTest.java b/src/test/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerabilityTest.java index 12b3b1cbd..107fb7a7b 100644 --- a/src/test/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerabilityTest.java +++ b/src/test/java/org/sasanlabs/service/vulnerability/idor/IDORVulnerabilityTest.java @@ -5,6 +5,7 @@ import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.argThat; import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -55,24 +56,27 @@ void level1_ShouldAllowAccessToAnyId() { } @Test - void level2_ShouldAllowCookieTampering() { + void level2_ShouldUseTokenOwnerInsteadOfTamperedCookie() { String validToken = "valid-token-level2"; User decoded = new User(); decoded.setUserId(1); decoded.setRole("USER"); - User bob = new User(2, "Bob", 60000, "USER"); + User alice = new User(1, "Alice", 50000, "USER"); when(idorLoginService.decodeToken(validToken)).thenReturn(decoded); when(jdbcTemplate.query( - anyString(), - any(Object[].class), + eq(SQL_PROFILE_BY_ID), + argThat( + (Object[] args) -> + args.length == 1 && Integer.valueOf(1).equals(args[0])), any(org.springframework.jdbc.core.RowMapper.class))) - .thenReturn(Arrays.asList(bob)); + .thenReturn(Arrays.asList(alice)); ResponseEntity> response = - idor.level2(validToken, 1); + idor.level2(validToken, 2); assertTrue(response.getBody().getIsValid()); + assertEquals(1, ((User) response.getBody().getContent()).getUserId()); } @Test