diff --git a/go-gost/x/socket/websocket_reporter.go b/go-gost/x/socket/websocket_reporter.go index c416d91..8e2d9aa 100644 --- a/go-gost/x/socket/websocket_reporter.go +++ b/go-gost/x/socket/websocket_reporter.go @@ -1,10 +1,16 @@ package socket import ( + "bytes" + "compress/gzip" "context" "encoding/json" "fmt" + "net" "net/url" + "os/exec" + "runtime" + "strconv" "strings" "time" @@ -57,6 +63,23 @@ type CommandResponse struct { RequestId string `json:"requestId,omitempty"` } +// PingRequest ping请求结构体 +type PingRequest struct { + IP string `json:"ip"` + Count int `json:"count"` + RequestId string `json:"requestId,omitempty"` +} + +// PingResponse ping响应结构体 +type PingResponse struct { + IP string `json:"ip"` + Success bool `json:"success"` + AverageTime float64 `json:"averageTime"` // 平均延迟(ms) + PacketLoss float64 `json:"packetLoss"` // 丢包率(%) + ErrorMessage string `json:"errorMessage,omitempty"` + RequestId string `json:"requestId,omitempty"` +} + type WebSocketReporter struct { url string conn *websocket.Conn @@ -137,6 +160,9 @@ func (w *WebSocketReporter) connect() error { w.conn = conn w.connected = true + // 设置最大消息大小为 16MB (默认是 1024 * 1024) + w.conn.SetReadLimit(100 * 1024 * 1024) + // 设置关闭处理器来检测连接状态 w.conn.SetCloseHandler(func(code int, text string) error { w.connected = false @@ -257,16 +283,61 @@ func (w *WebSocketReporter) receiveMessages() { func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byte) { switch messageType { case websocket.TextMessage: - - // 解析命令消息 - var cmdMsg CommandMessage - if err := json.Unmarshal(message, &cmdMsg); err != nil { - fmt.Printf("❌ 解析命令消息失败: %v\n", err) - w.sendErrorResponse("ParseError", fmt.Sprintf("解析命令失败: %v", err)) - return + // 先尝试解析是否是压缩消息 + var compressedMsg struct { + Type string `json:"type"` + Compressed bool `json:"compressed"` + Data json.RawMessage `json:"data"` + RequestId string `json:"requestId,omitempty"` } - if cmdMsg.Type != "call" { - w.routeCommand(cmdMsg) + + if err := json.Unmarshal(message, &compressedMsg); err == nil && compressedMsg.Compressed { + // 处理压缩消息 + fmt.Printf("📥 收到压缩消息,正在解压...\n") + + // 解压数据 + gzipReader, err := gzip.NewReader(bytes.NewReader(compressedMsg.Data)) + if err != nil { + fmt.Printf("❌ 创建解压读取器失败: %v\n", err) + w.sendErrorResponse("DecompressError", fmt.Sprintf("解压失败: %v", err)) + return + } + defer gzipReader.Close() + + var decompressedData bytes.Buffer + if _, err := decompressedData.ReadFrom(gzipReader); err != nil { + fmt.Printf("❌ 解压数据失败: %v\n", err) + w.sendErrorResponse("DecompressError", fmt.Sprintf("解压失败: %v", err)) + return + } + + // 使用解压后的数据继续处理 + message = decompressedData.Bytes() + + // 构建解压后的命令消息 + var cmdMsg CommandMessage + cmdMsg.Type = compressedMsg.Type + cmdMsg.RequestId = compressedMsg.RequestId + if err := json.Unmarshal(message, &cmdMsg.Data); err != nil { + fmt.Printf("❌ 解析解压后的命令数据失败: %v\n", err) + w.sendErrorResponse("ParseError", fmt.Sprintf("解析命令失败: %v", err)) + return + } + + if cmdMsg.Type != "call" { + w.routeCommand(cmdMsg) + } + } else { + // 处理普通消息 + var cmdMsg CommandMessage + if err := json.Unmarshal(message, &cmdMsg); err != nil { + fmt.Printf("❌ 解析命令消息失败: %v\n", err) + w.sendErrorResponse("ParseError", fmt.Sprintf("解析命令失败: %v", err)) + return + } + if cmdMsg.Type != "call" { + w.routeCommand(cmdMsg) + } } default: @@ -321,6 +392,14 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) { case "DeleteLimiters": err = w.handleDeleteLimiter(cmd.Data) response.Type = "DeleteLimitersResponse" + + // Ping 诊断命令 + case "Ping": + var pingResult PingResponse + pingResult, err = w.handlePing(cmd.Data) + response.Type = "PingResponse" + response.Data = pingResult + default: err = fmt.Errorf("未知命令类型: %s", cmd.Type) response.Type = "UnknownCommandResponse" @@ -621,29 +700,78 @@ func (w *WebSocketReporter) sendConfigReport() { return } - // 构建配置报告消息 - configMsg := struct { - Type string `json:"type"` - Data interface{} `json:"data"` - }{ - Type: "config_report", - Data: json.RawMessage(configData), - } + // 检查数据大小,如果超过1MB则压缩 + if len(configData) > 1024*1024 { + fmt.Printf("📦 配置数据较大 (%.2f MB),进行压缩处理\n", float64(len(configData))/(1024*1024)) - // 转换为JSON - jsonData, err := json.Marshal(configMsg) - if err != nil { - fmt.Printf("❌ 序列化配置报告失败: %v\n", err) - return - } + // 压缩配置数据 + var compressedBuf bytes.Buffer + gzipWriter := gzip.NewWriter(&compressedBuf) + if _, err := gzipWriter.Write(configData); err != nil { + fmt.Printf("❌ 压缩配置数据失败: %v\n", err) + return + } + gzipWriter.Close() - // 设置写入超时 - w.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + // 构建压缩后的配置报告消息 + configMsg := struct { + Type string `json:"type"` + Compressed bool `json:"compressed"` + Data []byte `json:"data"` + }{ + Type: "config_report", + Compressed: true, + Data: compressedBuf.Bytes(), + } - if err := w.conn.WriteMessage(websocket.TextMessage, jsonData); err != nil { - fmt.Printf("❌ 发送配置报告失败: %v\n", err) - w.connected = false - return + // 转换为JSON + jsonData, err := json.Marshal(configMsg) + if err != nil { + fmt.Printf("❌ 序列化压缩配置报告失败: %v\n", err) + return + } + + fmt.Printf("✅ 压缩后大小: %.2f MB -> %.2f MB (压缩率: %.1f%%)\n", + float64(len(configData))/(1024*1024), + float64(len(jsonData))/(1024*1024), + 100.0-float64(len(jsonData))/float64(len(configData))*100) + + // 设置写入超时 + w.conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + + if err := w.conn.WriteMessage(websocket.TextMessage, jsonData); err != nil { + fmt.Printf("❌ 发送压缩配置报告失败: %v\n", err) + w.connected = false + return + } + } else { + // 数据较小,直接发送 + // 构建配置报告消息 + configMsg := struct { + Type string `json:"type"` + Compressed bool `json:"compressed"` + Data interface{} `json:"data"` + }{ + Type: "config_report", + Compressed: false, + Data: json.RawMessage(configData), + } + + // 转换为JSON + jsonData, err := json.Marshal(configMsg) + if err != nil { + fmt.Printf("❌ 序列化配置报告失败: %v\n", err) + return + } + + // 设置写入超时 + w.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + + if err := w.conn.WriteMessage(websocket.TextMessage, jsonData); err != nil { + fmt.Printf("❌ 发送配置报告失败: %v\n", err) + w.connected = false + return + } } } @@ -661,7 +789,18 @@ func (w *WebSocketReporter) sendResponse(response CommandResponse) { return } - w.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + // 检查消息大小,如果超过10MB则记录警告 + if len(jsonData) > 10*1024*1024 { + fmt.Printf("⚠️ 响应消息过大 (%.2f MB),可能会被拒绝\n", float64(len(jsonData))/(1024*1024)) + } + + // 设置较长的写入超时,以应对大消息 + timeout := 5 * time.Second + if len(jsonData) > 1024*1024 { + timeout = 30 * time.Second + } + + w.conn.SetWriteDeadline(time.Now().Add(timeout)) if err := w.conn.WriteMessage(websocket.TextMessage, jsonData); err != nil { fmt.Printf("❌ 发送响应失败: %v\n", err) w.connected = false @@ -750,3 +889,200 @@ func StartWebSocketReporterWithConfig(Addr string, Secret string) *WebSocketRepo reporter.Start() return reporter } + +// handlePing 处理ping诊断命令 +func (w *WebSocketReporter) handlePing(data interface{}) (PingResponse, error) { + jsonData, err := json.Marshal(data) + if err != nil { + return PingResponse{}, fmt.Errorf("序列化ping数据失败: %v", err) + } + + var req PingRequest + if err := json.Unmarshal(jsonData, &req); err != nil { + return PingResponse{}, fmt.Errorf("解析ping请求失败: %v", err) + } + + // 验证IP地址格式 + if net.ParseIP(req.IP) == nil && !isValidHostname(req.IP) { + return PingResponse{ + IP: req.IP, + Success: false, + ErrorMessage: "无效的IP地址或主机名", + RequestId: req.RequestId, + }, nil + } + + // 设置默认ping次数 + if req.Count <= 0 { + req.Count = 4 + } + + // 执行ping操作 + avgTime, packetLoss, err := pingHost(req.IP, req.Count) + + response := PingResponse{ + IP: req.IP, + RequestId: req.RequestId, + } + + if err != nil { + response.Success = false + response.ErrorMessage = err.Error() + } else { + response.Success = true + response.AverageTime = avgTime + response.PacketLoss = packetLoss + } + + return response, nil +} + +// pingHost 执行ping操作,返回平均延迟和丢包率 +func pingHost(ip string, count int) (float64, float64, error) { + var cmd *exec.Cmd + + // 根据操作系统选择不同的ping命令 + switch runtime.GOOS { + case "windows": + cmd = exec.Command("ping", "-n", strconv.Itoa(count), ip) + case "darwin", "linux": + cmd = exec.Command("ping", "-c", strconv.Itoa(count), ip) + default: + return 0, 0, fmt.Errorf("不支持的操作系统: %s", runtime.GOOS) + } + + output, err := cmd.Output() + if err != nil { + return 0, 0, fmt.Errorf("ping命令执行失败: %v", err) + } + + // 解析ping输出 + return parsePingOutput(string(output), runtime.GOOS) +} + +// parsePingOutput 解析ping命令输出,提取平均延迟和丢包率 +func parsePingOutput(output, osType string) (float64, float64, error) { + lines := strings.Split(output, "\n") + + switch osType { + case "windows": + return parsePingOutputWindows(lines) + case "darwin", "linux": + return parsePingOutputUnix(lines) + default: + return 0, 0, fmt.Errorf("不支持的操作系统类型") + } +} + +// parsePingOutputWindows 解析Windows系统的ping输出 +func parsePingOutputWindows(lines []string) (float64, float64, error) { + var avgTime float64 + var packetLoss float64 + + for _, line := range lines { + line = strings.TrimSpace(line) + + // 查找平均延迟 (例如: "最短 = 1ms,最长 = 2ms,平均 = 1ms") + if strings.Contains(line, "平均") && strings.Contains(line, "ms") { + parts := strings.Split(line, "平均 = ") + if len(parts) > 1 { + avgPart := strings.Split(parts[1], "ms")[0] + if avg, err := strconv.ParseFloat(avgPart, 64); err == nil { + avgTime = avg + } + } + } + + // 查找丢包率 (例如: "丢失 = 0 (0% 丢失)") + if strings.Contains(line, "丢失") && strings.Contains(line, "%") { + if strings.Contains(line, "(0%") { + packetLoss = 0 + } else { + // 提取百分比 + start := strings.Index(line, "(") + end := strings.Index(line, "%") + if start != -1 && end != -1 && start < end { + lossStr := line[start+1 : end] + if loss, err := strconv.ParseFloat(lossStr, 64); err == nil { + packetLoss = loss + } + } + } + } + } + + return avgTime, packetLoss, nil +} + +// parsePingOutputUnix 解析Unix系统(Linux/macOS)的ping输出 +func parsePingOutputUnix(lines []string) (float64, float64, error) { + var avgTime float64 + var packetLoss float64 + + for _, line := range lines { + line = strings.TrimSpace(line) + + // 查找统计行 (例如: "4 packets transmitted, 4 received, 0% packet loss") + if strings.Contains(line, "packet loss") { + parts := strings.Split(line, "%") + if len(parts) > 0 { + // 查找百分比前的数字 + lossStr := strings.Fields(parts[0]) + if len(lossStr) > 0 { + if loss, err := strconv.ParseFloat(lossStr[len(lossStr)-1], 64); err == nil { + packetLoss = loss + } + } + } + } + + // 查找往返时间统计 (例如: "round-trip min/avg/max/stddev = 0.123/0.456/0.789/0.012 ms") + if strings.Contains(line, "round-trip") && strings.Contains(line, "=") { + parts := strings.Split(line, "=") + if len(parts) > 1 { + times := strings.TrimSpace(parts[1]) + times = strings.Split(times, " ")[0] // 去掉末尾的"ms" + timeValues := strings.Split(times, "/") + if len(timeValues) >= 2 { + if avg, err := strconv.ParseFloat(timeValues[1], 64); err == nil { + avgTime = avg + } + } + } + } + + // macOS的格式可能不同,查找avg (例如: "min/avg/max/stddev = 0.123/0.456/0.789/0.012 ms") + if strings.Contains(line, "min/avg/max") && strings.Contains(line, "=") { + parts := strings.Split(line, "=") + if len(parts) > 1 { + times := strings.TrimSpace(parts[1]) + times = strings.Split(times, " ")[0] // 去掉末尾的"ms" + timeValues := strings.Split(times, "/") + if len(timeValues) >= 2 { + if avg, err := strconv.ParseFloat(timeValues[1], 64); err == nil { + avgTime = avg + } + } + } + } + } + + return avgTime, packetLoss, nil +} + +// isValidHostname 验证主机名格式 +func isValidHostname(hostname string) bool { + if len(hostname) == 0 || len(hostname) > 253 { + return false + } + + // 简单的主机名验证 + for _, r := range hostname { + if !((r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || + (r >= '0' && r <= '9') || r == '-' || r == '.') { + return false + } + } + + return true +} diff --git a/gost.sql b/gost.sql index e098dc4..c4f4cab 100644 --- a/gost.sql +++ b/gost.sql @@ -3,7 +3,7 @@ -- https://www.phpmyadmin.net/ -- -- 主机: localhost --- 生成日期: 2025-06-25 10:51:18 +-- 生成日期: 2025-06-26 14:11:52 -- 服务器版本: 5.7.40-log -- PHP 版本: 7.4.33 @@ -86,6 +86,7 @@ CREATE TABLE `speed_limit` ( CREATE TABLE `tunnel` ( `id` int(10) NOT NULL, `name` varchar(100) NOT NULL, + `traffic_ratio` decimal(10,1) NOT NULL DEFAULT '1.0', `in_node_id` int(10) NOT NULL, `in_ip` varchar(100) NOT NULL, `in_port_sta` int(10) NOT NULL, @@ -131,7 +132,7 @@ CREATE TABLE `user` ( -- INSERT INTO `user` (`id`, `user`, `pwd`, `role_id`, `exp_time`, `flow`, `in_flow`, `out_flow`, `flow_reset_time`, `num`, `created_time`, `updated_time`, `status`) VALUES -(1, 'admin_user', '3c85cdebade1c51cf64ca9f3c09d182d', 0, 1780480500000, 99999, 0, 0, 1, 99999, 1748914865000, 1750228795866, 1); +(1, 'admin_user', '3c85cdebade1c51cf64ca9f3c09d182d', 0, 1780480500000, 99999, 0, 0, 1, 99999, 1748914865000, 1750910756282, 1); -- -------------------------------------------------------- @@ -201,37 +202,37 @@ ALTER TABLE `user_tunnel` -- 使用表AUTO_INCREMENT `forward` -- ALTER TABLE `forward` - MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=110; + MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=128; -- -- 使用表AUTO_INCREMENT `node` -- ALTER TABLE `node` - MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=19; + MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=22; -- -- 使用表AUTO_INCREMENT `speed_limit` -- ALTER TABLE `speed_limit` - MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=66; + MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=67; -- -- 使用表AUTO_INCREMENT `tunnel` -- ALTER TABLE `tunnel` - MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=33; + MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=45; -- -- 使用表AUTO_INCREMENT `user` -- ALTER TABLE `user` - MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=29; + MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=31; -- -- 使用表AUTO_INCREMENT `user_tunnel` -- ALTER TABLE `user_tunnel` - MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=47; + MODIFY `id` int(10) NOT NULL AUTO_INCREMENT, AUTO_INCREMENT=48; COMMIT; /*!40101 SET CHARACTER_SET_CLIENT=@OLD_CHARACTER_SET_CLIENT */; diff --git a/panel_install.sh b/panel_install.sh index 6b11053..5444b58 100755 --- a/panel_install.sh +++ b/panel_install.sh @@ -498,6 +498,29 @@ SET @sql = ( PREPARE stmt FROM @sql; EXECUTE stmt; DEALLOCATE PREPARE stmt; + +-- traffic_ratio (流量倍率) +SET @sql = ( + SELECT IF( + NOT EXISTS ( + SELECT 1 + FROM information_schema.COLUMNS + WHERE table_schema = DATABASE() + AND table_name = 'tunnel' + AND column_name = 'traffic_ratio' + ), + 'ALTER TABLE \`tunnel\` ADD COLUMN \`traffic_ratio\` DECIMAL(5,1) DEFAULT 1.0 COMMENT "流量倍率";', + 'SELECT "Column \`traffic_ratio\` already exists in \`tunnel\`";' + ) +); +PREPARE stmt FROM @sql; +EXECUTE stmt; +DEALLOCATE PREPARE stmt; + +-- 为现有数据设置默认流量倍率 +UPDATE \`tunnel\` +SET \`traffic_ratio\` = 1.0 +WHERE \`traffic_ratio\` IS NULL; EOF # 检查数据库容器 diff --git a/springboot-backend/src/main/java/com/admin/common/dto/GostDto.java b/springboot-backend/src/main/java/com/admin/common/dto/GostDto.java index 5249583..a57f2b7 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/GostDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/GostDto.java @@ -7,4 +7,6 @@ public class GostDto { private Integer code; private String msg; + + private Object data; // 添加数据字段,用于存储响应的详细数据 } diff --git a/springboot-backend/src/main/java/com/admin/common/dto/TunnelDto.java b/springboot-backend/src/main/java/com/admin/common/dto/TunnelDto.java index 4109417..c8c18e3 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/TunnelDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/TunnelDto.java @@ -5,6 +5,9 @@ import javax.validation.constraints.NotBlank; import javax.validation.constraints.NotNull; import javax.validation.constraints.Min; import javax.validation.constraints.Max; +import javax.validation.constraints.DecimalMin; +import javax.validation.constraints.DecimalMax; +import java.math.BigDecimal; @Data public class TunnelDto { @@ -44,6 +47,11 @@ public class TunnelDto { @NotNull(message = "流量计算类型不能为空") private Integer flow; + // 流量倍率,默认为1.0 + @DecimalMin(value = "0.0", message = "流量倍率不能小于0.0") + @DecimalMax(value = "100.0", message = "流量倍率不能大于100.0") + private BigDecimal trafficRatio = new BigDecimal("1.0"); + // 协议类型(隧道转发时使用:tls、tcp、mtls),默认为tls private String protocol; diff --git a/springboot-backend/src/main/java/com/admin/common/dto/TunnelUpdateDto.java b/springboot-backend/src/main/java/com/admin/common/dto/TunnelUpdateDto.java index a78866c..62e7bd6 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/TunnelUpdateDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/TunnelUpdateDto.java @@ -5,6 +5,9 @@ import javax.validation.constraints.NotBlank; import javax.validation.constraints.NotNull; import javax.validation.constraints.Min; import javax.validation.constraints.Max; +import javax.validation.constraints.DecimalMin; +import javax.validation.constraints.DecimalMax; +import java.math.BigDecimal; @Data public class TunnelUpdateDto { @@ -18,6 +21,11 @@ public class TunnelUpdateDto { @NotNull(message = "流量计算类型不能为空") private Integer flow; + // 流量倍率 + @DecimalMin(value = "0.0", message = "流量倍率不能小于0.0") + @DecimalMax(value = "100.0", message = "流量倍率不能大于100.0") + private BigDecimal trafficRatio; + @NotNull(message = "入口端口开始不能为空") @Min(value = 1, message = "入口端口开始必须大于0") @Max(value = 65535, message = "入口端口开始不能超过65535") diff --git a/springboot-backend/src/main/java/com/admin/common/utils/WebSocketServer.java b/springboot-backend/src/main/java/com/admin/common/utils/WebSocketServer.java index 88bb323..703d153 100644 --- a/springboot-backend/src/main/java/com/admin/common/utils/WebSocketServer.java +++ b/springboot-backend/src/main/java/com/admin/common/utils/WebSocketServer.java @@ -70,12 +70,27 @@ public class WebSocketServer extends TextWebSocketHandler { JSONObject responseJson = JSONObject.parseObject(message.getPayload()); String requestId = responseJson.getString("requestId"); String responseMessage = responseJson.getString("message"); + String responseType = responseJson.getString("type"); + JSONObject responseData = responseJson.getJSONObject("data"); if (requestId != null) { CompletableFuture future = pendingRequests.remove(requestId); if (future != null) { GostDto result = new GostDto(); - result.setMsg(responseMessage != null ? responseMessage : "无响应消息"); + + // 根据响应类型处理不同的数据 + if ("PingResponse".equals(responseType) && responseData != null) { + // 特殊处理ping响应,将完整的响应数据返回 + result.setMsg(responseMessage != null ? responseMessage : "OK"); + result.setData(responseData); // 保存ping详细结果 + } else { + // 其他类型的响应 + result.setMsg(responseMessage != null ? responseMessage : "无响应消息"); + if (responseData != null) { + result.setData(responseData); + } + } + future.complete(result); } } diff --git a/springboot-backend/src/main/java/com/admin/config/WebSocketConfig.java b/springboot-backend/src/main/java/com/admin/config/WebSocketConfig.java index 64335dc..5dee9d2 100644 --- a/springboot-backend/src/main/java/com/admin/config/WebSocketConfig.java +++ b/springboot-backend/src/main/java/com/admin/config/WebSocketConfig.java @@ -7,7 +7,12 @@ import org.springframework.web.socket.WebSocketHandler; import org.springframework.web.socket.config.annotation.EnableWebSocket; import org.springframework.web.socket.config.annotation.WebSocketConfigurer; import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry; +import org.springframework.boot.web.servlet.ServletContextInitializer; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import javax.servlet.ServletContext; +import javax.servlet.ServletException; import javax.annotation.Resource; @@ -18,6 +23,16 @@ public class WebSocketConfig implements WebSocketConfigurer { @Resource private WebSocketInterceptor webSocketInterceptor; + @Bean + public ServletContextInitializer websocketBufferConfig() { + return new ServletContextInitializer() { + @Override + public void onStartup(ServletContext servletContext) throws ServletException { + servletContext.setInitParameter("org.apache.tomcat.websocket.textBufferSize", String.valueOf(100 * 1024 * 1024)); + servletContext.setInitParameter("org.apache.tomcat.websocket.binaryBufferSize", String.valueOf(100 * 1024 * 1024)); + } + }; + } @Override public void registerWebSocketHandlers(WebSocketHandlerRegistry webSocketHandlerRegistry) { diff --git a/springboot-backend/src/main/java/com/admin/controller/FlowController.java b/springboot-backend/src/main/java/com/admin/controller/FlowController.java index 10a5a68..facde58 100644 --- a/springboot-backend/src/main/java/com/admin/controller/FlowController.java +++ b/springboot-backend/src/main/java/com/admin/controller/FlowController.java @@ -12,6 +12,7 @@ import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RestController; +import java.math.BigDecimal; import java.util.List; import java.util.Objects; import java.util.concurrent.ConcurrentHashMap; @@ -75,24 +76,15 @@ public class FlowController extends BaseController { return SUCCESS_RESPONSE; } - List validFlowData = flowDataList; -// // 2. 过滤有效流量数据 -// List validFlowData = filterValidFlowData(flowDataList); -// if (validFlowData.isEmpty()) { -// return SUCCESS_RESPONSE; -// } - - // 3. 解析服务名称获取ID信息 - String[] serviceIds = parseServiceName(validFlowData.get(0).getN()); + // 2. 解析服务名称获取ID信息 + String[] serviceIds = parseServiceName(flowDataList.get(0).getN()); String forwardId = serviceIds[0]; String userId = serviceIds[1]; String userTunnelId = serviceIds[2]; - // 4. 计算总流量 - FlowStatistics flowStats = calculateTotalFlow(validFlowData); - // 5. 一次性查询相关实体,避免后续重复查询 + // 3. 一次性查询相关实体,避免后续重复查询 Forward forward = forwardService.getById(forwardId); User user = userService.getById(userId); UserTunnel userTunnel = null; @@ -100,6 +92,20 @@ public class FlowController extends BaseController { userTunnel = userTunnelService.getById(userTunnelId); } + // 4. 处理流量倍率 + List validFlowData = flowDataList; + if (forward != null) { + validFlowData = filterFlowData(flowDataList, forward.getTunnelId()); + } + + + + + // 5. 计算总流量 + FlowStatistics flowStats = calculateTotalFlow(validFlowData); + + + // 6. 获取流量计费类型 int flowType = getFlowType(forward); @@ -138,6 +144,27 @@ public class FlowController extends BaseController { return SUCCESS_RESPONSE; } + + + private List filterFlowData(List flowDataList, Integer tunnel_id) { + Tunnel tunnel = tunnelService.getById(tunnel_id); + if (tunnel != null && tunnel.getTrafficRatio() != null){ + BigDecimal trafficRatio = tunnel.getTrafficRatio(); + for (FlowDto flowDto : flowDataList) { + // 将Long转为BigDecimal进行计算,然后转回Long + BigDecimal originalD = BigDecimal.valueOf(flowDto.getD()); + BigDecimal originalU = BigDecimal.valueOf(flowDto.getU()); + + BigDecimal newD = originalD.multiply(trafficRatio); + BigDecimal newU = originalU.multiply(trafficRatio); + + flowDto.setD(newD.longValue()); + flowDto.setU(newU.longValue()); + } + } + return flowDataList; + } + /** * 验证节点是否有效 */ diff --git a/springboot-backend/src/main/java/com/admin/controller/TunnelController.java b/springboot-backend/src/main/java/com/admin/controller/TunnelController.java index 0fb90dd..203269b 100644 --- a/springboot-backend/src/main/java/com/admin/controller/TunnelController.java +++ b/springboot-backend/src/main/java/com/admin/controller/TunnelController.java @@ -138,4 +138,17 @@ public class TunnelController extends BaseController { return tunnelService.userTunnel(); } + /** + * 隧道诊断功能 + * @param params 包含tunnelId的参数 + * @return 诊断结果 + */ + @LogAnnotation + @RequireRole + @PostMapping("/diagnose") + public R diagnoseTunnel(@RequestBody Map params) { + Long tunnelId = Long.valueOf(params.get("tunnelId").toString()); + return tunnelService.diagnoseTunnel(tunnelId); + } + } diff --git a/springboot-backend/src/main/java/com/admin/entity/Tunnel.java b/springboot-backend/src/main/java/com/admin/entity/Tunnel.java index 76695ae..c24b9c9 100644 --- a/springboot-backend/src/main/java/com/admin/entity/Tunnel.java +++ b/springboot-backend/src/main/java/com/admin/entity/Tunnel.java @@ -1,6 +1,7 @@ package com.admin.entity; import java.io.Serializable; +import java.math.BigDecimal; import lombok.Data; import lombok.EqualsAndHashCode; @@ -78,6 +79,12 @@ public class Tunnel extends BaseEntity { */ private String protocol; + /** + * 流量倍率 + */ + private BigDecimal trafficRatio; + + private String tcpListenAddr; private String udpListenAddr; diff --git a/springboot-backend/src/main/java/com/admin/service/TunnelService.java b/springboot-backend/src/main/java/com/admin/service/TunnelService.java index 327a2b0..a62706a 100644 --- a/springboot-backend/src/main/java/com/admin/service/TunnelService.java +++ b/springboot-backend/src/main/java/com/admin/service/TunnelService.java @@ -45,5 +45,16 @@ public interface TunnelService extends IService { */ R deleteTunnel(Long id); + /** + * 获取用户可用的隧道列表 + * @return 结果 + */ R userTunnel(); + + /** + * 隧道诊断功能 + * @param tunnelId 隧道ID + * @return 诊断结果 + */ + R diagnoseTunnel(Long tunnelId); } diff --git a/springboot-backend/src/main/java/com/admin/service/impl/TunnelServiceImpl.java b/springboot-backend/src/main/java/com/admin/service/impl/TunnelServiceImpl.java index d5b7980..37ab7f5 100644 --- a/springboot-backend/src/main/java/com/admin/service/impl/TunnelServiceImpl.java +++ b/springboot-backend/src/main/java/com/admin/service/impl/TunnelServiceImpl.java @@ -1,12 +1,15 @@ package com.admin.service.impl; import cn.hutool.core.util.StrUtil; +import com.admin.common.dto.GostDto; import com.admin.common.dto.TunnelDto; import com.admin.common.dto.TunnelListDto; import com.admin.common.dto.TunnelUpdateDto; import com.admin.common.lang.R; +import com.admin.common.utils.GostUtil; import com.admin.common.utils.JwtUtil; +import com.admin.common.utils.WebSocketServer; import com.admin.entity.Forward; import com.admin.entity.Node; import com.admin.entity.Tunnel; @@ -18,6 +21,7 @@ import com.admin.service.ForwardService; import com.admin.service.NodeService; import com.admin.service.TunnelService; import com.admin.service.UserTunnelService; +import com.alibaba.fastjson.JSONObject; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl; import lombok.Data; @@ -25,7 +29,11 @@ import org.springframework.beans.BeanUtils; import org.springframework.stereotype.Service; import javax.annotation.Resource; +import java.math.BigDecimal; +import java.util.ArrayList; +import java.util.HashMap; import java.util.List; +import java.util.Map; import java.util.stream.Collectors; /** @@ -193,6 +201,11 @@ public class TunnelServiceImpl extends ServiceImpl impleme existingTunnel.setInPortSta(tunnelUpdateDto.getInPortSta()); existingTunnel.setInPortEnd(tunnelUpdateDto.getInPortEnd()); + // 更新流量倍率 + if (tunnelUpdateDto.getTrafficRatio() != null) { + existingTunnel.setTrafficRatio(tunnelUpdateDto.getTrafficRatio()); + } + // 更新TCP和UDP监听地址 if (StrUtil.isNotBlank(tunnelUpdateDto.getTcpListenAddr())) { existingTunnel.setTcpListenAddr(tunnelUpdateDto.getTcpListenAddr()); @@ -413,6 +426,13 @@ public class TunnelServiceImpl extends ServiceImpl impleme // 设置流量计算类型 tunnel.setFlow(tunnelDto.getFlow()); + // 设置流量倍率,如果为空则设置默认值1.0 + if (tunnelDto.getTrafficRatio() != null) { + tunnel.setTrafficRatio(tunnelDto.getTrafficRatio()); + } else { + tunnel.setTrafficRatio(new BigDecimal("1.0")); + } + // 设置协议类型(仅隧道转发需要) if (tunnelDto.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { // 隧道转发时,设置协议类型,默认为tls @@ -527,8 +547,6 @@ public class TunnelServiceImpl extends ServiceImpl impleme tunnel.setUpdatedTime(currentTime); } - - /** * 检查隧道是否存在 * @@ -676,6 +694,140 @@ public class TunnelServiceImpl extends ServiceImpl impleme return dto; } + /** + * 隧道诊断功能 + * + * @param tunnelId 隧道ID + * @return 诊断结果响应 + */ + @Override + public R diagnoseTunnel(Long tunnelId) { + // 1. 验证隧道是否存在 + Tunnel tunnel = this.getById(tunnelId); + if (tunnel == null) { + return R.err(ERROR_TUNNEL_NOT_FOUND); + } + + // 2. 获取入口和出口节点信息 + Node inNode = nodeService.getById(tunnel.getInNodeId()); + if (inNode == null) { + return R.err(ERROR_IN_NODE_NOT_FOUND); + } + + Node outNode = null; + if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { + outNode = nodeService.getById(tunnel.getOutNodeId()); + if (outNode == null) { + return R.err(ERROR_OUT_NODE_NOT_FOUND); + } + } + + List results = new ArrayList<>(); + + // 3. 根据隧道类型执行不同的诊断策略 + if (tunnel.getType() == TUNNEL_TYPE_PORT_FORWARD) { + // 端口转发:只给入口节点发送诊断指令,ping谷歌DNS + DiagnosisResult inResult = performPingDiagnosis(inNode, "8.8.8.8", "入口->外网"); + results.add(inResult); + } else { + // 隧道转发:入口ping出口,出口ping谷歌DNS + DiagnosisResult inToOutResult = performPingDiagnosis(inNode, outNode.getServerIp(), "入口->出口"); + results.add(inToOutResult); + + DiagnosisResult outToExternalResult = performPingDiagnosis(outNode, "8.8.8.8", "出口->外网"); + results.add(outToExternalResult); + } + + // 4. 构建诊断报告 + Map diagnosisReport = new HashMap<>(); + diagnosisReport.put("tunnelId", tunnelId); + diagnosisReport.put("tunnelName", tunnel.getName()); + diagnosisReport.put("tunnelType", tunnel.getType() == TUNNEL_TYPE_PORT_FORWARD ? "端口转发" : "隧道转发"); + diagnosisReport.put("results", results); + diagnosisReport.put("timestamp", System.currentTimeMillis()); + + return R.ok(diagnosisReport); + } + + /** + * 执行ping诊断 + * + * @param node 执行ping的节点 + * @param targetIp 目标IP地址 + * @param description 诊断描述 + * @return 诊断结果 + */ + private DiagnosisResult performPingDiagnosis(Node node, String targetIp, String description) { + try { + // 构建ping请求数据 + JSONObject pingData = new JSONObject(); + pingData.put("ip", targetIp); + pingData.put("count", 4); + + // 发送ping命令到节点 + GostDto gostResult = WebSocketServer.send_msg(node.getId(), pingData, "Ping"); + + DiagnosisResult result = new DiagnosisResult(); + result.setNodeId(node.getId()); + result.setNodeName(node.getName()); + result.setTargetIp(targetIp); + result.setDescription(description); + result.setTimestamp(System.currentTimeMillis()); + + if (gostResult != null && "OK".equals(gostResult.getMsg())) { + // 尝试解析ping响应数据 + try { + if (gostResult.getData() != null) { + JSONObject pingResponse = (JSONObject) gostResult.getData(); + boolean success = pingResponse.getBooleanValue("success"); + + result.setSuccess(success); + if (success) { + result.setMessage("ping成功"); + result.setAverageTime(pingResponse.getDoubleValue("averageTime")); + result.setPacketLoss(pingResponse.getDoubleValue("packetLoss")); + } else { + result.setMessage(pingResponse.getString("errorMessage")); + result.setAverageTime(-1.0); + result.setPacketLoss(100.0); + } + } else { + // 没有详细数据,使用默认值 + result.setSuccess(true); + result.setMessage("ping成功"); + result.setAverageTime(0.0); + result.setPacketLoss(0.0); + } + } catch (Exception e) { + // 解析响应数据失败,但ping命令本身成功了 + result.setSuccess(true); + result.setMessage("ping成功,但无法解析详细数据"); + result.setAverageTime(0.0); + result.setPacketLoss(0.0); + } + } else { + result.setSuccess(false); + result.setMessage(gostResult != null ? gostResult.getMsg() : "节点无响应"); + result.setAverageTime(-1.0); + result.setPacketLoss(100.0); + } + + return result; + } catch (Exception e) { + DiagnosisResult result = new DiagnosisResult(); + result.setNodeId(node.getId()); + result.setNodeName(node.getName()); + result.setTargetIp(targetIp); + result.setDescription(description); + result.setSuccess(false); + result.setMessage("诊断执行异常: " + e.getMessage()); + result.setTimestamp(System.currentTimeMillis()); + result.setAverageTime(-1.0); + result.setPacketLoss(100.0); + return result; + } + } + // ========== 内部数据类 ========== /** @@ -710,4 +862,20 @@ public class TunnelServiceImpl extends ServiceImpl impleme return new NodeValidationResult(true, errorMessage, null); } } + + /** + * 诊断结果数据类 + */ + @Data + public static class DiagnosisResult { + private Long nodeId; + private String nodeName; + private String targetIp; + private String description; + private boolean success; + private String message; + private double averageTime; + private double packetLoss; + private long timestamp; + } } diff --git a/springboot-backend/src/main/resources/application.yml b/springboot-backend/src/main/resources/application.yml index 3176ded..019f812 100644 --- a/springboot-backend/src/main/resources/application.yml +++ b/springboot-backend/src/main/resources/application.yml @@ -26,6 +26,7 @@ server: uri-encoding: UTF-8 max-thread: 800 max-connections: 2000 + max-swallow-size: 100MB shutdown: graceful mybatis-plus: diff --git a/vue-frontend/src/api/index.js b/vue-frontend/src/api/index.js index bf82b38..934ebe0 100644 --- a/vue-frontend/src/api/index.js +++ b/vue-frontend/src/api/index.js @@ -22,6 +22,7 @@ export const getTunnelList = () => Network.post("/tunnel/list") export const getTunnelById = (id) => Network.post("/tunnel/get", { id }) export const updateTunnel = (data) => Network.post("/tunnel/update", data) export const deleteTunnel = (id) => Network.post("/tunnel/delete", { id }) +export const diagnoseTunnel = (tunnelId) => Network.post("/tunnel/diagnose", { tunnelId }) // 用户隧道权限管理操作 - 全部使用POST请求 export const assignUserTunnel = (data) => Network.post("/tunnel/user/assign", data) diff --git a/vue-frontend/src/views/Forward.vue b/vue-frontend/src/views/Forward.vue index 9483e8b..b88ab97 100644 --- a/vue-frontend/src/views/Forward.vue +++ b/vue-frontend/src/views/Forward.vue @@ -1,16 +1,18 @@