Commit c37303e7 by tntxia Committed by gdj

安全整改的问题修改

parent 7583b48c
...@@ -85,6 +85,22 @@ ...@@ -85,6 +85,22 @@
</excludes> </excludes>
</configuration> </configuration>
</plugin> </plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-compiler-plugin</artifactId>
<version>3.10.1</version> <!-- JDK 11 推荐用 3.8+ 版本 -->
<configuration>
<!-- 统一为 JDK 11,消除版本不匹配警告 -->
<release>11</release>
<encoding>UTF-8</encoding>
<!-- 可选:关闭严格泛型检查(避免额外报错) -->
<compilerArgs>
<arg>-Xlint:-unchecked</arg>
</compilerArgs>
</configuration>
</plugin>
</plugins> </plugins>
</build> </build>
</project> </project>
...@@ -56,9 +56,71 @@ public class AiInfoServiceImpl extends ServiceImpl<IAiInfoMapper, AiInfoEntity> ...@@ -56,9 +56,71 @@ public class AiInfoServiceImpl extends ServiceImpl<IAiInfoMapper, AiInfoEntity>
String[] arr = param.getOrderBy().split(" "); String[] arr = param.getOrderBy().split(" ");
String column = arr[0]; String column = arr[0];
String desc = arr.length > 1 ? arr[1] : "desc"; String direction = arr.length > 1 ? arr[1] : "desc";
wrapper.last(Objects.nonNull(param.getOrderBy()), " order by " + column + " " + desc); // 避免SQL注入,不直接使用SQL语句,所以判断相应的列
switch (column) {
case "id":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getId);
} else {
wrapper.orderByDesc(AiInfoEntity::getId);
}
break;
case "createTime":
case "create_time":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getCreateTime);
} else {
wrapper.orderByDesc(AiInfoEntity::getCreateTime);
}
break;
case "warnTime":
case "warn_time":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getWarnTime);
} else {
wrapper.orderByDesc(AiInfoEntity::getWarnTime);
}
break;
case "deviceSn":
case "device_sn":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getDeviceSn);
} else {
wrapper.orderByDesc(AiInfoEntity::getDeviceSn);
}
break;
case "warnType":
case "warn_type":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getWarnType);
} else {
wrapper.orderByDesc(AiInfoEntity::getWarnType);
}
break;
case "warnEvent":
case "warn_event":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getWarnEvent);
} else {
wrapper.orderByDesc(AiInfoEntity::getWarnEvent);
}
break;
case "algorithmType":
case "algorithm_type":
if ("asc".equals(direction)) {
wrapper.orderByAsc(AiInfoEntity::getAlgorithmType);
} else {
wrapper.orderByDesc(AiInfoEntity::getAlgorithmType);
}
break;
default:
wrapper.orderByDesc(AiInfoEntity::getId);
break;
}
} else {
wrapper.orderByDesc(AiInfoEntity::getId);
} }
Page<AiInfoEntity> pagination = this.page(new Page<>(page, pageSize), wrapper); Page<AiInfoEntity> pagination = this.page(new Page<>(page, pageSize), wrapper);
......
...@@ -32,19 +32,20 @@ public class AuthPrincipalHandler extends DefaultHandshakeHandler { ...@@ -32,19 +32,20 @@ public class AuthPrincipalHandler extends DefaultHandshakeHandler {
HttpServletRequest servletRequest = ((ServletServerHttpRequest) request).getServletRequest(); HttpServletRequest servletRequest = ((ServletServerHttpRequest) request).getServletRequest();
String token = servletRequest.getParameter(AuthInterceptor.PARAM_TOKEN); String token = servletRequest.getParameter(AuthInterceptor.PARAM_TOKEN);
// 默认让WebSocket的认证都通过
if (!StringUtils.hasText(token)) { if (!StringUtils.hasText(token)) {
return false; return true;
} }
log.debug("token:" + token); log.debug("token:" + token);
Optional<CustomClaim> customClaim = JwtUtil.parseToken(token); Optional<CustomClaim> customClaim = JwtUtil.parseToken(token);
if (customClaim.isEmpty()) { if (customClaim.isEmpty()) {
return false; return true;
} }
servletRequest.setAttribute(AuthInterceptor.TOKEN_CLAIM, customClaim.get()); servletRequest.setAttribute(AuthInterceptor.TOKEN_CLAIM, customClaim.get());
return true; return true;
} }
return false; return true;
} }
...@@ -63,6 +64,10 @@ public class AuthPrincipalHandler extends DefaultHandshakeHandler { ...@@ -63,6 +64,10 @@ public class AuthPrincipalHandler extends DefaultHandshakeHandler {
CustomClaim claim = (CustomClaim) ((ServletServerHttpRequest) request).getServletRequest() CustomClaim claim = (CustomClaim) ((ServletServerHttpRequest) request).getServletRequest()
.getAttribute(AuthInterceptor.TOKEN_CLAIM); .getAttribute(AuthInterceptor.TOKEN_CLAIM);
if (claim == null) {
return () -> null;
}
return () -> claim.getWorkspaceId() + "/" + claim.getUserType() + "/" + claim.getId(); return () -> claim.getWorkspaceId() + "/" + claim.getUserType() + "/" + claim.getId();
} }
return () -> null; return () -> null;
......
package com.dji.sample.component.websocket.config; package com.dji.sample.component.websocket.config;
import com.dji.sample.common.model.CustomClaim;
import com.dji.sample.common.util.JwtUtil;
import com.dji.sample.component.websocket.service.IWebSocketManageService; import com.dji.sample.component.websocket.service.IWebSocketManageService;
import com.dji.sdk.websocket.WebSocketDefaultHandler; import com.dji.sdk.websocket.WebSocketDefaultHandler;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.springframework.util.StringUtils; import org.springframework.util.StringUtils;
import org.springframework.web.socket.CloseStatus; import org.springframework.web.socket.CloseStatus;
...@@ -10,6 +14,9 @@ import org.springframework.web.socket.WebSocketMessage; ...@@ -10,6 +14,9 @@ import org.springframework.web.socket.WebSocketMessage;
import org.springframework.web.socket.WebSocketSession; import org.springframework.web.socket.WebSocketSession;
import java.security.Principal; import java.security.Principal;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.ConcurrentHashMap;
/** /**
* *
...@@ -22,6 +29,10 @@ public class MyWebSocketHandler extends WebSocketDefaultHandler { ...@@ -22,6 +29,10 @@ public class MyWebSocketHandler extends WebSocketDefaultHandler {
private IWebSocketManageService webSocketManageService; private IWebSocketManageService webSocketManageService;
private ObjectMapper objectMapper = new ObjectMapper();
private Map<String, Boolean> authenticatedSessions = new ConcurrentHashMap<>();
private Map<String, CustomClaim> sessionClaims = new ConcurrentHashMap<>();
MyWebSocketHandler(WebSocketHandler delegate, IWebSocketManageService webSocketManageService) { MyWebSocketHandler(WebSocketHandler delegate, IWebSocketManageService webSocketManageService) {
super(delegate); super(delegate);
this.webSocketManageService = webSocketManageService; this.webSocketManageService = webSocketManageService;
...@@ -30,29 +41,104 @@ public class MyWebSocketHandler extends WebSocketDefaultHandler { ...@@ -30,29 +41,104 @@ public class MyWebSocketHandler extends WebSocketDefaultHandler {
@Override @Override
public void afterConnectionEstablished(WebSocketSession session) throws Exception { public void afterConnectionEstablished(WebSocketSession session) throws Exception {
Principal principal = session.getPrincipal(); Principal principal = session.getPrincipal();
if (StringUtils.hasText(principal.getName())) { String principalName = principal.getName();
webSocketManageService.put(principal.getName(), new MyConcurrentWebSocketSession(session));
log.debug("{} is connected. ID: {}. WebSocketSession[current count: {}]", if (StringUtils.hasText(principalName) && !principalName.startsWith("temp-")) {
principal.getName(), session.getId(), webSocketManageService.getConnectedCount()); webSocketManageService.put(principalName, new MyConcurrentWebSocketSession(session));
return; authenticatedSessions.put(session.getId(), true);
log.debug("{} is connected (pre-authenticated). ID: {}. WebSocketSession[current count: {}]",
principalName, session.getId(), webSocketManageService.getConnectedCount());
} else {
log.debug("Unauthenticated connection established. ID: {}, temp principal: {}",
session.getId(), principalName);
} }
session.close();
} }
@Override @Override
public void afterConnectionClosed(WebSocketSession session, CloseStatus closeStatus) throws Exception { public void afterConnectionClosed(WebSocketSession session, CloseStatus closeStatus) throws Exception {
Principal principal = session.getPrincipal(); String sessionId = session.getId();
if (StringUtils.hasText(principal.getName())) { Boolean isAuthenticated = authenticatedSessions.get(sessionId);
webSocketManageService.remove(principal.getName(), session.getId());
log.debug("{} is disconnected. ID: {}. WebSocketSession[current count: {}]", if (Boolean.TRUE.equals(isAuthenticated)) {
principal.getName(), session.getId(), webSocketManageService.getConnectedCount()); CustomClaim claim = sessionClaims.get(sessionId);
if (claim != null) {
String key = claim.getWorkspaceId() + "/" + claim.getUserType() + "/" + claim.getId();
webSocketManageService.remove(key, sessionId);
log.debug("{} is disconnected. ID: {}. WebSocketSession[current count: {}]",
key, sessionId, webSocketManageService.getConnectedCount());
}
} }
authenticatedSessions.remove(sessionId);
sessionClaims.remove(sessionId);
} }
@Override @Override
public void handleMessage(WebSocketSession session, WebSocketMessage<?> message) throws Exception { public void handleMessage(WebSocketSession session, WebSocketMessage<?> message) throws Exception {
log.debug("received message: {}", message.getPayload()); String sessionId = session.getId();
Boolean isAuthenticated = authenticatedSessions.get(sessionId);
String payload = message.getPayload().toString();
log.debug("Received message from session {}: {}", sessionId, payload);
try {
JsonNode jsonNode = objectMapper.readTree(payload);
String type = jsonNode.has("type") ? jsonNode.get("type").asText() : null;
if ("auth".equals(type) && !Boolean.TRUE.equals(isAuthenticated)) {
String token = jsonNode.has("token") ? jsonNode.get("token").asText() : null;
if (StringUtils.hasText(token)) {
Optional<CustomClaim> customClaimOpt = JwtUtil.parseToken(token);
if (customClaimOpt.isPresent()) {
CustomClaim claim = customClaimOpt.get();
String key = claim.getWorkspaceId() + "/" + claim.getUserType() + "/" + claim.getId();
authenticatedSessions.put(sessionId, true);
sessionClaims.put(sessionId, claim);
webSocketManageService.put(key, new MyConcurrentWebSocketSession(session));
log.debug("Session {} authenticated successfully. User: {}", sessionId, key);
String authResponse = objectMapper.writeValueAsString(Map.of(
"type", "auth_success",
"status", "ok"
));
session.sendMessage(new org.springframework.web.socket.TextMessage(authResponse));
return;
}
}
log.warn("Authentication failed for session {}", sessionId);
String authFail = objectMapper.writeValueAsString(Map.of(
"type", "auth_fail",
"status", "error",
"message", "Invalid token"
));
session.sendMessage(new org.springframework.web.socket.TextMessage(authFail));
session.close(CloseStatus.NOT_ACCEPTABLE);
return;
}
if (!Boolean.TRUE.equals(isAuthenticated)) {
log.warn("Unauthenticated session {} tried to send message. Closing.", sessionId);
String authRequired = objectMapper.writeValueAsString(Map.of(
"type", "auth_required",
"message", "Please authenticate first"
));
session.sendMessage(new org.springframework.web.socket.TextMessage(authRequired));
// session.close(CloseStatus.NOT_ACCEPTABLE);
return;
}
super.handleMessage(session, message);
} catch (Exception e) {
log.error("Error handling message from session {}", sessionId, e);
if (!Boolean.TRUE.equals(isAuthenticated)) {
session.close(CloseStatus.BAD_DATA);
}
}
} }
} }
\ No newline at end of file
...@@ -152,7 +152,7 @@ public class UserController { ...@@ -152,7 +152,7 @@ public class UserController {
} }
/** /**
* Admin resets a user's password. * User reset password.
* The new password must comply with all password rules. * The new password must comply with all password rules.
* *
* @param request HTTP request * @param request HTTP request
...@@ -162,6 +162,7 @@ public class UserController { ...@@ -162,6 +162,7 @@ public class UserController {
@PostMapping("/resetPassword") @PostMapping("/resetPassword")
public HttpResultResponse<Object> resetPassword(HttpServletRequest request, public HttpResultResponse<Object> resetPassword(HttpServletRequest request,
@RequestBody ChangePasswordParam param) { @RequestBody ChangePasswordParam param) {
CustomClaim customClaim = (CustomClaim) request.getAttribute(TOKEN_CLAIM); CustomClaim customClaim = (CustomClaim) request.getAttribute(TOKEN_CLAIM);
String userId = customClaim.getId(); String userId = customClaim.getId();
return userService.resetPassword(userId, param); return userService.resetPassword(userId, param);
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment