diff --git a/server/skillhub-app/src/main/java/com/iflytek/skillhub/filter/AuthContextFilter.java b/server/skillhub-app/src/main/java/com/iflytek/skillhub/filter/AuthContextFilter.java index 60d7ac6d..aeced23f 100644 --- a/server/skillhub-app/src/main/java/com/iflytek/skillhub/filter/AuthContextFilter.java +++ b/server/skillhub-app/src/main/java/com/iflytek/skillhub/filter/AuthContextFilter.java @@ -4,15 +4,18 @@ import com.iflytek.skillhub.auth.rbac.PlatformPrincipal; import com.iflytek.skillhub.domain.namespace.NamespaceMember; import com.iflytek.skillhub.domain.namespace.NamespaceMemberRepository; import com.iflytek.skillhub.domain.namespace.NamespaceRole; +import com.iflytek.skillhub.domain.user.UserAccountRepository; import jakarta.servlet.FilterChain; import jakarta.servlet.ServletException; import jakarta.servlet.http.HttpServletRequest; import jakarta.servlet.http.HttpServletResponse; +import jakarta.servlet.http.HttpSession; import java.io.IOException; import java.util.Map; import java.util.stream.Collectors; import org.springframework.security.core.Authentication; import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.security.web.context.HttpSessionSecurityContextRepository; import org.springframework.stereotype.Component; import org.springframework.web.filter.OncePerRequestFilter; @@ -20,9 +23,12 @@ import org.springframework.web.filter.OncePerRequestFilter; public class AuthContextFilter extends OncePerRequestFilter { private final NamespaceMemberRepository namespaceMemberRepository; + private final UserAccountRepository userAccountRepository; - public AuthContextFilter(NamespaceMemberRepository namespaceMemberRepository) { + public AuthContextFilter(NamespaceMemberRepository namespaceMemberRepository, + UserAccountRepository userAccountRepository) { this.namespaceMemberRepository = namespaceMemberRepository; + this.userAccountRepository = userAccountRepository; } @Override @@ -32,6 +38,11 @@ public class AuthContextFilter extends OncePerRequestFilter { FilterChain filterChain) throws ServletException, IOException { PlatformPrincipal principal = resolvePrincipal(request); if (principal != null) { + if (isInactiveUser(principal.userId())) { + clearAuthentication(request); + response.sendError(HttpServletResponse.SC_UNAUTHORIZED); + return; + } request.setAttribute("userId", principal.userId()); Map userNsRoles = namespaceMemberRepository.findByUserId(principal.userId()).stream() .collect(Collectors.toMap( @@ -44,6 +55,23 @@ public class AuthContextFilter extends OncePerRequestFilter { filterChain.doFilter(request, response); } + private boolean isInactiveUser(String userId) { + return userAccountRepository.findById(userId) + .map(user -> !user.isActive()) + .orElse(true); + } + + private void clearAuthentication(HttpServletRequest request) { + SecurityContextHolder.clearContext(); + HttpSession session = request.getSession(false); + if (session == null) { + return; + } + session.removeAttribute("platformPrincipal"); + session.removeAttribute(HttpSessionSecurityContextRepository.SPRING_SECURITY_CONTEXT_KEY); + session.invalidate(); + } + private PlatformPrincipal resolvePrincipal(HttpServletRequest request) { Authentication authentication = SecurityContextHolder.getContext().getAuthentication(); if (authentication != null) { diff --git a/server/skillhub-app/src/test/java/com/iflytek/skillhub/filter/AuthContextFilterTest.java b/server/skillhub-app/src/test/java/com/iflytek/skillhub/filter/AuthContextFilterTest.java new file mode 100644 index 00000000..e6050ca4 --- /dev/null +++ b/server/skillhub-app/src/test/java/com/iflytek/skillhub/filter/AuthContextFilterTest.java @@ -0,0 +1,92 @@ +package com.iflytek.skillhub.filter; + +import com.iflytek.skillhub.auth.rbac.PlatformPrincipal; +import com.iflytek.skillhub.domain.namespace.NamespaceMember; +import com.iflytek.skillhub.domain.namespace.NamespaceMemberRepository; +import com.iflytek.skillhub.domain.namespace.NamespaceRole; +import com.iflytek.skillhub.domain.user.UserAccount; +import com.iflytek.skillhub.domain.user.UserAccountRepository; +import com.iflytek.skillhub.domain.user.UserStatus; +import jakarta.servlet.FilterChain; +import jakarta.servlet.http.HttpSession; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.springframework.mock.web.MockHttpServletRequest; +import org.springframework.mock.web.MockHttpServletResponse; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.context.SecurityContextHolder; + +import java.util.List; +import java.util.Set; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +class AuthContextFilterTest { + + private final NamespaceMemberRepository namespaceMemberRepository = mock(NamespaceMemberRepository.class); + private final UserAccountRepository userAccountRepository = mock(UserAccountRepository.class); + private final AuthContextFilter filter = new AuthContextFilter(namespaceMemberRepository, userAccountRepository); + + @AfterEach + void clearSecurityContext() { + SecurityContextHolder.clearContext(); + } + + @Test + void disabledSessionUser_shouldInvalidateSessionAndBlockRequest() throws Exception { + PlatformPrincipal principal = new PlatformPrincipal("user-1", "Alice", "alice@example.com", null, "local", Set.of("USER")); + UserAccount user = new UserAccount("Alice", "alice@example.com"); + user.setStatus(UserStatus.DISABLED); + + MockHttpServletRequest request = new MockHttpServletRequest(); + HttpSession session = request.getSession(true); + session.setAttribute("platformPrincipal", principal); + SecurityContextHolder.getContext().setAuthentication( + new UsernamePasswordAuthenticationToken(principal, null, List.of()) + ); + + MockHttpServletResponse response = new MockHttpServletResponse(); + FilterChain filterChain = mock(FilterChain.class); + + when(userAccountRepository.findById("user-1")).thenReturn(java.util.Optional.of(user)); + + filter.doFilter(request, response, filterChain); + + assertEquals(401, response.getStatus()); + assertTrue(!request.isRequestedSessionIdValid() || request.getSession(false) == null); + assertNull(SecurityContextHolder.getContext().getAuthentication()); + verify(filterChain, never()).doFilter(request, response); + } + + @Test + void activeSessionUser_shouldPopulateRequestContextAndContinue() throws Exception { + PlatformPrincipal principal = new PlatformPrincipal("user-2", "Bob", "bob@example.com", null, "local", Set.of("USER")); + UserAccount user = new UserAccount("Bob", "bob@example.com"); + user.setStatus(UserStatus.ACTIVE); + NamespaceMember member = new NamespaceMember(9L, "user-2", NamespaceRole.ADMIN); + + MockHttpServletRequest request = new MockHttpServletRequest(); + request.getSession(true).setAttribute("platformPrincipal", principal); + SecurityContextHolder.getContext().setAuthentication( + new UsernamePasswordAuthenticationToken(principal, null, List.of()) + ); + + MockHttpServletResponse response = new MockHttpServletResponse(); + FilterChain filterChain = mock(FilterChain.class); + + when(userAccountRepository.findById("user-2")).thenReturn(java.util.Optional.of(user)); + when(namespaceMemberRepository.findByUserId("user-2")).thenReturn(List.of(member)); + + filter.doFilter(request, response, filterChain); + + assertEquals("user-2", request.getAttribute("userId")); + assertEquals(NamespaceRole.ADMIN, ((java.util.Map) request.getAttribute("userNsRoles")).get(9L)); + verify(filterChain).doFilter(request, response); + } +}