fix:审计npe

This commit is contained in:
waner 2026-04-25 11:44:47 +08:00
parent 57ebf819a3
commit 779052142b
3 changed files with 243 additions and 9 deletions

View File

@ -1,7 +1,12 @@
package com.cisd.tms.modules.log.aspect;
import com.cisd.tms.common.api.ApiResponse;
import com.cisd.tms.modules.auth.dto.LoginResponse;
import com.cisd.tms.modules.auth.dto.PasswordLoginRequest;
import com.cisd.tms.modules.auth.dto.UkeyLoginRequest;
import com.cisd.tms.modules.log.annotation.AuditedOperation;
import com.cisd.tms.modules.log.dto.OperationAuditCommand;
import com.cisd.tms.modules.log.enums.ActionType;
import com.cisd.tms.modules.log.enums.AuthLevel;
import com.cisd.tms.modules.log.enums.OperationResult;
import com.cisd.tms.modules.log.enums.OperatorRoleCode;
@ -16,6 +21,8 @@ import org.springframework.stereotype.Component;
import org.springframework.web.context.request.RequestContextHolder;
import org.springframework.web.context.request.ServletRequestAttributes;
import java.util.stream.Collectors;
@Aspect
@Component
public class OperationAuditAspect {
@ -33,13 +40,18 @@ public class OperationAuditAspect {
command.setActionType(auditedOperation.action());
command.setSummary(auditedOperation.summary());
//todo 后续测
// 获取上下文信息
// 获取已认证请求上下文登录接口尚未有会话会在下面从登录请求体补充
getContextInfo(command);
if (ActionType.LOGIN == auditedOperation.action()) {
getLoginRequestContext(command, joinPoint.getArgs());
}
Object result = null;
try {
result = joinPoint.proceed();
if (ActionType.LOGIN == auditedOperation.action()) {
getLoginResponseContext(command, result);
}
command.setOperationResult(OperationResult.SUCCESS);
return result;
@ -69,12 +81,98 @@ public class OperationAuditAspect {
Object roleObj = request.getAttribute(InternalApiAuthInterceptor.ATTR_ROLE_CODE);
command.setOperatorRoleCode(OperatorRoleCode.valueOf(roleObj.toString()));
setOperatorRoleCode(command, roleObj);
Object levelObj = request.getAttribute(InternalApiAuthInterceptor.ATTR_AUTH_LEVEL);
command.setOperatorAuthLevel(AuthLevel.valueOf(levelObj.toString()));
setOperatorAuthLevel(command, levelObj);
command.setRemoteIp(request.getRemoteAddr());
}
}
}
private void getLoginRequestContext(OperationAuditCommand command, Object[] args) {
if (args == null) {
return;
}
for (Object arg : args) {
if (arg instanceof PasswordLoginRequest request) {
setOperatorRoleCode(command, request.getRoleCode());
appendLoginSummary(command, request.getRoleCode(), passwordLoginAccounts(request));
return;
}
if (arg instanceof UkeyLoginRequest request) {
setOperatorRoleCode(command, request.getRoleCode());
appendLoginSummary(command, request.getRoleCode(), ukeyLoginUids(request));
return;
}
}
}
private void getLoginResponseContext(OperationAuditCommand command, Object result) {
if (result instanceof ApiResponse<?> apiResponse && apiResponse.getData() instanceof LoginResponse loginResponse) {
setOperatorRoleCode(command, loginResponse.getRoleCode());
setOperatorAuthLevel(command, loginResponse.getAuthLevel());
}
}
private void setOperatorRoleCode(OperationAuditCommand command, Object roleObj) {
if (roleObj == null || roleObj.toString().isBlank()) {
return;
}
try {
command.setOperatorRoleCode(OperatorRoleCode.valueOf(roleObj.toString()));
} catch (IllegalArgumentException ignored) {
// 保留空值避免审计上下文解析异常反向阻断登录或业务请求
}
}
private void setOperatorAuthLevel(OperationAuditCommand command, Object levelObj) {
if (levelObj == null || levelObj.toString().isBlank()) {
return;
}
try {
command.setOperatorAuthLevel(AuthLevel.valueOf(levelObj.toString()));
} catch (IllegalArgumentException ignored) {
// 保留空值避免审计上下文解析异常反向阻断登录或业务请求
}
}
private String passwordLoginAccounts(PasswordLoginRequest request) {
if (request.getAccounts() == null || request.getAccounts().isEmpty()) {
return "";
}
return request.getAccounts().stream()
.map(account -> account.getUsername() == null ? "" : account.getUsername().trim())
.filter(value -> !value.isBlank())
.collect(Collectors.joining(","));
}
private String ukeyLoginUids(UkeyLoginRequest request) {
if (request.getLoginFactors() == null || request.getLoginFactors().isEmpty()) {
return "";
}
return request.getLoginFactors().stream()
.map(factor -> factor.getUid() == null ? "" : "uid=" + factor.getUid())
.filter(value -> !value.isBlank())
.collect(Collectors.joining(","));
}
private void appendLoginSummary(OperationAuditCommand command, String roleCode, String principalText) {
StringBuilder builder = new StringBuilder(command.getSummary() == null ? "" : command.getSummary());
if ((roleCode == null || roleCode.isBlank()) && (principalText == null || principalText.isBlank())) {
return;
}
builder.append(" [");
if (roleCode != null && !roleCode.isBlank()) {
builder.append("role=").append(roleCode);
}
if (principalText != null && !principalText.isBlank()) {
if (roleCode != null && !roleCode.isBlank()) {
builder.append(", ");
}
builder.append("principal=").append(principalText);
}
builder.append("]");
command.setSummary(builder.toString());
}
}

View File

@ -89,11 +89,22 @@ public class OperationAuditService {
// 从上下文中获取刚才解析出的用户信息
Object roleObj = request.getAttribute(ATTR_ROLE_CODE);
// todo 是否考虑未知情况
command.setOperatorRoleCode(OperatorRoleCode.valueOf(roleObj.toString()));
if (roleObj != null && !roleObj.toString().isBlank()) {
try {
command.setOperatorRoleCode(OperatorRoleCode.valueOf(roleObj.toString()));
} catch (IllegalArgumentException ignored) {
// 保留空值避免审计上下文解析失败影响原始越权响应
}
}
Object levelObj = request.getAttribute(ATTR_AUTH_LEVEL);
command.setOperatorAuthLevel(AuthLevel.valueOf(levelObj.toString()));
if (levelObj != null && !levelObj.toString().isBlank()) {
try {
command.setOperatorAuthLevel(AuthLevel.valueOf(levelObj.toString()));
} catch (IllegalArgumentException ignored) {
// 保留空值避免审计上下文解析失败影响原始越权响应
}
}
command.setRemoteIp(request.getRemoteAddr());
command.setErrorMessage(errorMsg);
@ -238,4 +249,4 @@ public class OperationAuditService {
private String trim(String s){
return s == null ? "" : s.trim();
}
}
}

View File

@ -0,0 +1,125 @@
package com.cisd.tms.modules.log.aspect;
import com.cisd.tms.common.api.ApiResponse;
import com.cisd.tms.modules.auth.controller.AuthController;
import com.cisd.tms.modules.auth.dto.LoginResponse;
import com.cisd.tms.modules.auth.dto.PasswordLoginAccountRequest;
import com.cisd.tms.modules.auth.dto.PasswordLoginRequest;
import com.cisd.tms.modules.auth.dto.UkeyLoginProof;
import com.cisd.tms.modules.auth.dto.UkeyLoginRequest;
import com.cisd.tms.modules.log.annotation.AuditedOperation;
import com.cisd.tms.modules.log.dto.OperationAuditCommand;
import com.cisd.tms.modules.log.enums.AuthLevel;
import com.cisd.tms.modules.log.enums.OperationResult;
import com.cisd.tms.modules.log.enums.OperatorRoleCode;
import com.cisd.tms.modules.log.service.OperationAuditService;
import org.aspectj.lang.ProceedingJoinPoint;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.mockito.Mockito;
import org.springframework.mock.web.MockHttpServletRequest;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.web.context.request.RequestContextHolder;
import org.springframework.web.context.request.ServletRequestAttributes;
import java.util.List;
import static org.mockito.ArgumentMatchers.eq;
class OperationAuditAspectLoginTest {
@AfterEach
void tearDown() {
RequestContextHolder.resetRequestAttributes();
}
@Test
void shouldRecordLoginRoleAndAuthLevelFromRequestAndResponseWithoutSessionContext() throws Throwable {
OperationAuditService auditService = Mockito.mock(OperationAuditService.class);
OperationAuditAspect aspect = newAspect(auditService);
PasswordLoginRequest request = new PasswordLoginRequest();
request.setRoleCode("AUDIT_ADMIN");
PasswordLoginAccountRequest account = new PasswordLoginAccountRequest();
account.setUsername("audit-admin-01");
request.setAccounts(List.of(account));
LoginResponse response = new LoginResponse();
response.setRoleCode("AUDIT_ADMIN");
response.setAuthLevel("LIMITED");
ProceedingJoinPoint joinPoint = joinPoint(request, ApiResponse.success(response));
bindRequest("10.0.0.8");
Object result = aspect.doAround(joinPoint, auditedOperation("passwordLogin", PasswordLoginRequest.class));
Assertions.assertTrue(result instanceof ApiResponse<?>);
ArgumentCaptor<OperationAuditCommand> captor = ArgumentCaptor.forClass(OperationAuditCommand.class);
Mockito.verify(auditService).record(captor.capture(), eq(true));
OperationAuditCommand command = captor.getValue();
Assertions.assertEquals(OperatorRoleCode.AUDIT_ADMIN, command.getOperatorRoleCode());
Assertions.assertEquals(AuthLevel.LIMITED, command.getOperatorAuthLevel());
Assertions.assertEquals(OperationResult.SUCCESS, command.getOperationResult());
Assertions.assertEquals("口令登录 [role=AUDIT_ADMIN, principal=audit-admin-01]", command.getSummary());
Assertions.assertEquals("10.0.0.8", command.getRemoteIp());
}
@Test
void shouldRecordFailedLoginRoleFromRequestWithoutSessionContext() throws Throwable {
OperationAuditService auditService = Mockito.mock(OperationAuditService.class);
OperationAuditAspect aspect = newAspect(auditService);
UkeyLoginRequest request = new UkeyLoginRequest();
request.setRoleCode("KEY_ADMIN");
UkeyLoginProof proof = new UkeyLoginProof();
proof.setUid(1);
request.setLoginFactors(List.of(proof));
RuntimeException failure = new RuntimeException("ukey verification failed");
ProceedingJoinPoint joinPoint = joinPointThrowing(request, failure);
bindRequest("10.0.0.9");
RuntimeException actual = Assertions.assertThrows(
RuntimeException.class,
() -> aspect.doAround(joinPoint, auditedOperation("ukeyLogin", UkeyLoginRequest.class))
);
Assertions.assertSame(failure, actual);
ArgumentCaptor<OperationAuditCommand> captor = ArgumentCaptor.forClass(OperationAuditCommand.class);
Mockito.verify(auditService).record(captor.capture(), eq(true));
OperationAuditCommand command = captor.getValue();
Assertions.assertEquals(OperatorRoleCode.KEY_ADMIN, command.getOperatorRoleCode());
Assertions.assertNull(command.getOperatorAuthLevel());
Assertions.assertEquals(OperationResult.FAILED, command.getOperationResult());
Assertions.assertEquals("UKey 登录 [role=KEY_ADMIN, principal=uid=1]", command.getSummary());
Assertions.assertEquals("ukey verification failed", command.getErrorMessage());
Assertions.assertEquals("10.0.0.9", command.getRemoteIp());
}
private OperationAuditAspect newAspect(OperationAuditService auditService) {
OperationAuditAspect aspect = new OperationAuditAspect();
ReflectionTestUtils.setField(aspect, "auditService", auditService);
return aspect;
}
private ProceedingJoinPoint joinPoint(Object request, Object result) throws Throwable {
ProceedingJoinPoint joinPoint = Mockito.mock(ProceedingJoinPoint.class);
Mockito.when(joinPoint.getArgs()).thenReturn(new Object[] { request });
Mockito.when(joinPoint.proceed()).thenReturn(result);
return joinPoint;
}
private ProceedingJoinPoint joinPointThrowing(Object request, RuntimeException failure) throws Throwable {
ProceedingJoinPoint joinPoint = Mockito.mock(ProceedingJoinPoint.class);
Mockito.when(joinPoint.getArgs()).thenReturn(new Object[] { request });
Mockito.when(joinPoint.proceed()).thenThrow(failure);
return joinPoint;
}
private AuditedOperation auditedOperation(String methodName, Class<?> requestType) throws NoSuchMethodException {
return AuthController.class.getMethod(methodName, requestType).getAnnotation(AuditedOperation.class);
}
private void bindRequest(String remoteAddr) {
MockHttpServletRequest request = new MockHttpServletRequest();
request.setRemoteAddr(remoteAddr);
RequestContextHolder.setRequestAttributes(new ServletRequestAttributes(request));
}
}