From f7e82cdc8ba5ae04f014803e56df24048bb6e3e1 Mon Sep 17 00:00:00 2001 From: qaq <1937228092@qq.com> Date: Wed, 25 Jun 2025 10:49:06 +0800 Subject: [PATCH] =?UTF-8?q?gost=E9=80=9A=E8=AE=AF=E6=94=B9=E4=B8=BAws?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 35 +- docker-compose.yml | 2 +- go-gost/go.mod | 4 +- go-gost/main.go | 8 +- go-gost/websocket_reporter.go | 450 ----------- go-gost/x/go.mod | 4 + go-gost/x/go.sum | 1 + go-gost/x/socket/chain.go | 111 +++ go-gost/x/socket/config.go | 83 ++ go-gost/x/socket/limiter.go | 102 +++ go-gost/x/socket/service.go | 519 ++++++++++++ go-gost/x/socket/websocket_reporter.go | 752 ++++++++++++++++++ go-gost/x/traffic/memory_manager.go | 27 + go-gost/x/traffic/traffic.go | 1 + go-gost/{ => x/traffic}/traffic_reporter.go | 14 +- install.sh | 313 +++++--- panel_install.sh | 490 ++++++++++-- .../java/com/admin/common/dto/ForwardDto.java | 9 + .../admin/common/dto/ForwardUpdateDto.java | 9 + .../java/com/admin/common/dto/NodeDto.java | 9 +- .../com/admin/common/dto/NodeUpdateDto.java | 5 +- .../java/com/admin/common/dto/TunnelDto.java | 6 + .../com/admin/common/dto/TunnelListDto.java | 42 +- .../com/admin/common/dto/TunnelUpdateDto.java | 6 + .../common/task/CheckGostConfigAsync.java | 65 +- .../admin/common/task/DelayQueueManager.java | 23 +- .../admin/common/task/SaveConfigAsync.java | 22 - .../java/com/admin/common/utils/GostUtil.java | 269 ++----- .../com/admin/common/utils/HttpUtils.java | 262 +----- .../java/com/admin/common/utils/JwtUtil.java | 2 +- .../admin/common/utils/WebSocketServer.java | 184 ++++- .../admin/config/WebSocketInterceptor.java | 4 - .../com/admin/controller/FlowController.java | 19 +- .../src/main/java/com/admin/entity/Node.java | 3 +- .../main/java/com/admin/entity/Tunnel.java | 4 + .../src/main/java/com/admin/entity/User.java | 2 - .../service/impl/ForwardServiceImpl.java | 449 ++++++++--- .../admin/service/impl/NodeServiceImpl.java | 62 +- .../service/impl/SpeedLimitServiceImpl.java | 22 +- .../admin/service/impl/TunnelServiceImpl.java | 26 +- .../admin/service/impl/UserServiceImpl.java | 16 +- .../service/impl/UserTunnelServiceImpl.java | 26 +- vue-frontend/src/views/Forward.vue | 77 +- vue-frontend/src/views/Home.vue | 10 + vue-frontend/src/views/Limit.vue | 2 +- vue-frontend/src/views/Tunnel.vue | 76 +- vue-frontend/src/views/node.vue | 414 ++++++++-- 47 files changed, 3528 insertions(+), 1513 deletions(-) delete mode 100644 go-gost/websocket_reporter.go create mode 100644 go-gost/x/socket/chain.go create mode 100644 go-gost/x/socket/config.go create mode 100644 go-gost/x/socket/limiter.go create mode 100644 go-gost/x/socket/service.go create mode 100644 go-gost/x/socket/websocket_reporter.go rename go-gost/{ => x/traffic}/traffic_reporter.go (93%) mode change 100644 => 100755 install.sh delete mode 100644 springboot-backend/src/main/java/com/admin/common/task/SaveConfigAsync.java diff --git a/README.md b/README.md index 17a8faf..7b1dbbc 100644 --- a/README.md +++ b/README.md @@ -73,43 +73,10 @@ ```bash -ipv6需要面板端支持,同时开启docker的ipv6服务和composer中的ipv6 - - -如果以前安装过需要重新安装 -推荐先删除本地上次下载的文件 - panel_install.sh - gost.sql - docker-compose.yml -在执行下面的安装命令 - -github -curl -L https://github.com/bqlpfy/forward-panel/raw/refs/heads/main/panel_install.sh -o panel_install.sh && chmod +x panel_install.sh && ./panel_install.sh - -gitee -curl -L https://gitee.com/bqlpfy/forward-panel/raw/master/panel_install.sh -o panel_install.sh && chmod +x panel_install.sh && ./panel_install.sh - -节点端安装时可以手动将install.sh换成下方的连接 -https://gitee.com/bqlpfy/forward-panel/raw/master/install.sh +curl -L https://ghproxy.com/https://github.com/bqlpfy/forward-panel/raw/refs/heads/main/panel_install.sh -o panel_install.sh && chmod +x panel_install.sh && ./panel_install.sh ``` -### 卸载 -```bash -面板端 - cd到compose所在位置执行下面的命令 - 改操作会删除所有数据 包括数据库文件 - docker compose down --rmi all --volumes --remove-orphans - 或 - docker-compose down --rmi all --volumes --remove-orphans -节点端 - systemctl stop gost - systemctl disable gost - rm -f /etc/systemd/system/gost.service - rm -rf /etc/gost - systemctl daemon-reload - -``` diff --git a/docker-compose.yml b/docker-compose.yml index f484cb4..9f5f538 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -37,7 +37,7 @@ services: JWT_SECRET: ${JWT_SECRET} LOG_DIR: /app/logs SERVER_ADDR: ${SERVER_HOST} - JAVA_OPTS: "-Xms256m -Xmx512m -Dfile.encoding=UTF-8 -Duser.timezone=Asia/Shanghai" + JAVA_OPTS: "-Xms128m -Xmx512m -Dfile.encoding=UTF-8 -Duser.timezone=Asia/Shanghai" ports: - "${BACKEND_PORT}:6365" volumes: diff --git a/go-gost/go.mod b/go-gost/go.mod index 382dcbe..9a6da8f 100644 --- a/go-gost/go.mod +++ b/go-gost/go.mod @@ -7,9 +7,7 @@ toolchain go1.23.4 require ( github.com/go-gost/core v0.3.1 github.com/go-gost/x v0.5.3 - github.com/gorilla/websocket v1.5.3 github.com/judwhite/go-svc v1.2.1 - github.com/shirou/gopsutil/v3 v3.24.5 ) require ( @@ -49,6 +47,7 @@ require ( github.com/google/gopacket v1.1.19 // indirect github.com/google/pprof v0.0.0-20241210010833-40e02aabc2ad // indirect github.com/google/uuid v1.6.0 // indirect + github.com/gorilla/websocket v1.5.3 // indirect github.com/gravitational/trace v1.1.16-0.20220114165159-14a9a7dd6aaf // indirect github.com/hashicorp/hcl v1.0.0 // indirect github.com/jonboulle/clockwork v0.2.2 // indirect @@ -87,6 +86,7 @@ require ( github.com/sagikazarmark/slog-shim v0.1.0 // indirect github.com/shadowsocks/go-shadowsocks2 v0.1.5 // indirect github.com/shadowsocks/shadowsocks-go v0.0.0-20200409064450-3e585ff90601 // indirect + github.com/shirou/gopsutil/v3 v3.24.5 // indirect github.com/shoenig/go-m1cpu v0.1.6 // indirect github.com/sirupsen/logrus v1.8.1 // indirect github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8 // indirect diff --git a/go-gost/main.go b/go-gost/main.go index b84749f..687d5b2 100644 --- a/go-gost/main.go +++ b/go-gost/main.go @@ -14,6 +14,7 @@ import ( "github.com/go-gost/core/logger" xlogger "github.com/go-gost/x/logger" + "github.com/go-gost/x/socket" "github.com/go-gost/x/traffic" "github.com/judwhite/go-svc" ) @@ -124,9 +125,10 @@ func main() { fmt.Println("✅ 使用内存流量管理器") logger.Default().Info("Using memory traffic manager") - SetHTTPReportURL(config.Addr, config.Secret) - StartTrafficReporter(trafficMgr) - wsReporter := StartWebSocketReporterWithConfig(config.Addr, config.Secret) + traffic.SetHTTPReportURL(config.Addr, config.Secret) + traffic.StartTrafficReporter(trafficMgr) + + wsReporter := socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret) defer wsReporter.Stop() p := &program{} diff --git a/go-gost/websocket_reporter.go b/go-gost/websocket_reporter.go deleted file mode 100644 index eb8e556..0000000 --- a/go-gost/websocket_reporter.go +++ /dev/null @@ -1,450 +0,0 @@ -package main - -import ( - "context" - "encoding/json" - "fmt" - "io" - "net" - "net/http" - "net/url" - "strings" - "time" - - "github.com/gorilla/websocket" - "github.com/shirou/gopsutil/v3/cpu" - "github.com/shirou/gopsutil/v3/host" - "github.com/shirou/gopsutil/v3/mem" - psnet "github.com/shirou/gopsutil/v3/net" -) - -// SystemInfo 系统信息结构体 -type SystemInfo struct { - Uptime uint64 `json:"uptime"` // 开机时间 (秒) - BytesReceived uint64 `json:"bytes_received"` // 接收字节数 - BytesTransmitted uint64 `json:"bytes_transmitted"` // 发送字节数 - CPUUsage float64 `json:"cpu_usage"` // CPU使用率(百分比) - MemoryUsage float64 `json:"memory_usage"` // 内存使用率(百分比) -} - -// NetworkStats 网络统计信息 -type NetworkStats struct { - BytesReceived uint64 `json:"bytes_received"` // 接收字节数 - BytesTransmitted uint64 `json:"bytes_transmitted"` // 发送字节数 -} - -// CPUInfo CPU信息 -type CPUInfo struct { - Usage float64 `json:"usage"` // CPU使用率(百分比) -} - -// MemoryInfo 内存信息 -type MemoryInfo struct { - Usage float64 `json:"usage"` // 内存使用率(百分比) -} - -type WebSocketReporter struct { - url string - conn *websocket.Conn - reconnectTime time.Duration - pingInterval time.Duration - ctx context.Context - cancel context.CancelFunc - connected bool -} - -// NewWebSocketReporter 创建一个新的WebSocket报告器 -func NewWebSocketReporter(serverURL string) *WebSocketReporter { - ctx, cancel := context.WithCancel(context.Background()) - return &WebSocketReporter{ - url: serverURL, - reconnectTime: 5 * time.Second, // 重连间隔 - pingInterval: 2 * time.Second, // 发送间隔改为2秒 - ctx: ctx, - cancel: cancel, - connected: false, - } -} - -// Start 启动WebSocket报告器 -func (w *WebSocketReporter) Start() { - go w.run() -} - -// Stop 停止WebSocket报告器 -func (w *WebSocketReporter) Stop() { - w.cancel() - if w.conn != nil { - w.conn.Close() - } - -} - -// run 主运行循环 -func (w *WebSocketReporter) run() { - for { - select { - case <-w.ctx.Done(): - return - default: - if err := w.connect(); err != nil { - fmt.Printf("❌ WebSocket连接失败: %v,%v后重试\n", err, w.reconnectTime) - select { - case <-time.After(w.reconnectTime): - continue - case <-w.ctx.Done(): - return - } - } - - // 连接成功,开始发送消息 - w.handleConnection() - } - } -} - -// connect 建立WebSocket连接 -func (w *WebSocketReporter) connect() error { - u, err := url.Parse(w.url) - if err != nil { - return fmt.Errorf("解析URL失败: %v", err) - } - - dialer := websocket.DefaultDialer - dialer.HandshakeTimeout = 10 * time.Second - - conn, _, err := dialer.Dial(u.String(), nil) - if err != nil { - return fmt.Errorf("连接WebSocket失败: %v", err) - } - - w.conn = conn - w.connected = true - - // 设置关闭处理器来检测连接状态 - w.conn.SetCloseHandler(func(code int, text string) error { - w.connected = false - return nil - }) - - return nil -} - -// handleConnection 处理WebSocket连接 -func (w *WebSocketReporter) handleConnection() { - defer func() { - if w.conn != nil { - w.conn.Close() - w.conn = nil - } - w.connected = false - }() - - // 主发送循环 - ticker := time.NewTicker(w.pingInterval) - defer ticker.Stop() - - for { - select { - case <-w.ctx.Done(): - return - case <-ticker.C: - // 检查连接状态 - if !w.connected { - return - } - - // 获取系统信息并发送 - sysInfo := w.collectSystemInfo() - if err := w.sendSystemInfo(sysInfo); err != nil { - fmt.Printf("❌ 发送系统信息失败: %v,准备重连\n", err) - return - } - } - } -} - -// collectSystemInfo 收集系统信息 -func (w *WebSocketReporter) collectSystemInfo() SystemInfo { - networkStats := getNetworkStats() - cpuInfo := getCPUInfo() - memoryInfo := getMemoryInfo() - - return SystemInfo{ - Uptime: getUptime(), - BytesReceived: networkStats.BytesReceived, - BytesTransmitted: networkStats.BytesTransmitted, - CPUUsage: cpuInfo.Usage, - MemoryUsage: memoryInfo.Usage, - } -} - -// sendSystemInfo 发送系统信息 -func (w *WebSocketReporter) sendSystemInfo(sysInfo SystemInfo) error { - if w.conn == nil || !w.connected { - return fmt.Errorf("连接未建立") - } - - // 转换为JSON - jsonData, err := json.Marshal(sysInfo) - if err != nil { - return fmt.Errorf("序列化系统信息失败: %v", err) - } - - // 设置写入超时 - w.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) - - if err := w.conn.WriteMessage(websocket.TextMessage, jsonData); err != nil { - w.connected = false // 标记连接已断开 - return fmt.Errorf("写入消息失败: %v", err) - } - - return nil -} - -// getHostIP 获取主机IP地址(IPv4优先策略:公网v4 → 公网v6 → 内网v4 → 内网v6) -func getHostIP() string { - // 先尝试获取IPv4地址 - ipv4 := getLocalIPv4() - if ipv4 != "unknown" && !isPrivateIP(ipv4) { - return ipv4 - } - - // 尝试获取IPv6地址 - ipv6 := getLocalIPv6() - if ipv6 != "unknown" && !isPrivateIP(ipv6) { - return ipv6 - } - - // IPv4公网IP查询服务 - ipv4Services := []string{ - "https://ipv4.icanhazip.com", - "https://api.ipify.org", - "https://checkip.amazonaws.com", - "https://myip.biturl.top", - } - - // IPv6公网IP查询服务 - ipv6Services := []string{ - "https://ipv6.icanhazip.com", - "https://v6.ident.me", - "https://ipv6.myip.biturl.top", - } - - client := &http.Client{ - Timeout: 10 * time.Second, - } - - // 优先尝试获取IPv4公网地址 - for _, service := range ipv4Services { - if ip := getIPFromService(client, service); ip != "" && net.ParseIP(ip).To4() != nil { - return strings.TrimSpace(ip) - } - } - - // IPv4公网IP获取失败,尝试获取IPv6公网地址 - if ipv6 != "unknown" { - for _, service := range ipv6Services { - if ip := getIPFromService(client, service); ip != "" && net.ParseIP(ip).To4() == nil { - return strings.TrimSpace(ip) - } - } - } - - // 如果所有公网IP服务都失败,按优先级返回本地IP:IPv4 → IPv6 - if ipv4 != "unknown" { - return ipv4 - } - if ipv6 != "unknown" { - return ipv6 - } - - return "unknown" -} - -// getIPFromService 从指定服务获取IP地址 -func getIPFromService(client *http.Client, serviceURL string) string { - resp, err := client.Get(serviceURL) - if err != nil { - return "" - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return "" - } - - body, err := io.ReadAll(resp.Body) - if err != nil { - return "" - } - - ip := strings.TrimSpace(string(body)) - // 简单验证是否为有效的IP地址 - if net.ParseIP(ip) != nil { - return ip - } - - return "" -} - -// getLocalIP 获取本地接口IP地址(作为备用方案,保持向后兼容) -func getLocalIP() string { - return getLocalIPv4() -} - -// getLocalIPv4 获取本地IPv4接口地址 -func getLocalIPv4() string { - conn, err := net.Dial("udp", "8.8.8.8:80") - if err != nil { - return "unknown" - } - defer conn.Close() - - localAddr := conn.LocalAddr().(*net.UDPAddr) - return localAddr.IP.String() -} - -// getLocalIPv6 获取本地IPv6接口地址 -func getLocalIPv6() string { - // 使用Google的IPv6 DNS服务器 - conn, err := net.Dial("udp", "[2001:4860:4860::8888]:80") - if err != nil { - return "unknown" - } - defer conn.Close() - - localAddr := conn.LocalAddr().(*net.UDPAddr) - return localAddr.IP.String() -} - -// isPrivateIP 判断IP地址是否为内网地址 -func isPrivateIP(ipStr string) bool { - ip := net.ParseIP(ipStr) - if ip == nil { - return false - } - - // 检查IPv4私有地址范围 - if ip.To4() != nil { - // 10.0.0.0/8 - if ip[12] == 10 { - return true - } - // 172.16.0.0/12 - if ip[12] == 172 && ip[13] >= 16 && ip[13] <= 31 { - return true - } - // 192.168.0.0/16 - if ip[12] == 192 && ip[13] == 168 { - return true - } - // 127.0.0.0/8 (回环地址) - if ip[12] == 127 { - return true - } - } - - // 检查IPv6私有地址 - if ip.To4() == nil { - // ::1 (回环地址) - if ip.IsLoopback() { - return true - } - // fc00::/7 (唯一本地地址 Unique Local Addresses) - if len(ip) >= 2 && (ip[0]&0xfe) == 0xfc { - return true - } - // fe80::/10 (链路本地地址 Link Local) - if len(ip) >= 2 && ip[0] == 0xfe && (ip[1]&0xc0) == 0x80 { - return true - } - // ::ffff:0:0/96 (IPv4映射地址) - if len(ip) >= 12 && ip[10] == 0xff && ip[11] == 0xff { - // 检查映射的IPv4地址是否为私有地址 - ipv4 := net.IPv4(ip[12], ip[13], ip[14], ip[15]) - return isPrivateIP(ipv4.String()) - } - // ::/128 (未指定地址) - if ip.IsUnspecified() { - return true - } - } - - return false -} - -// getUptime 获取系统开机时间(秒) -func getUptime() uint64 { - uptime, err := host.Uptime() - if err != nil { - return 0 - } - return uptime -} - -// getNetworkStats 获取网络统计信息 -func getNetworkStats() NetworkStats { - var stats NetworkStats - - ioCounters, err := psnet.IOCounters(true) - if err != nil { - fmt.Printf("获取网络统计失败: %v\n", err) - return stats - } - - // 汇总所有非回环接口的流量 - for _, io := range ioCounters { - // 跳过回环接口 - if io.Name == "lo" || strings.HasPrefix(io.Name, "lo") { - continue - } - - stats.BytesReceived += io.BytesRecv - stats.BytesTransmitted += io.BytesSent - } - - return stats -} - -// getCPUInfo 获取CPU信息 -func getCPUInfo() CPUInfo { - var cpuInfo CPUInfo - - // 获取CPU使用率 - percentages, err := cpu.Percent(time.Second, false) - if err == nil && len(percentages) > 0 { - cpuInfo.Usage = percentages[0] - } - - return cpuInfo -} - -// getMemoryInfo 获取内存信息 -func getMemoryInfo() MemoryInfo { - var memInfo MemoryInfo - - vmStat, err := mem.VirtualMemory() - if err != nil { - return memInfo - } - - memInfo.Usage = vmStat.UsedPercent - - return memInfo -} - -// StartWebSocketReporterWithConfig 使用配置启动WebSocket报告器 -func StartWebSocketReporterWithConfig(Addr string, Secret string) *WebSocketReporter { - // 获取本机IP地址 - localIP := getHostIP() - - // 构建包含本机IP的WebSocket URL - var fullURL = "ws://" + Addr + "/system-info?type=1&secret=" + Secret + "&client_ip=" + localIP - - fmt.Printf("🔗 WebSocket连接URL: %s\n", fullURL) - - reporter := NewWebSocketReporter(fullURL) - reporter.Start() - return reporter -} diff --git a/go-gost/x/go.mod b/go-gost/x/go.mod index fd85c0c..1d0e18b 100644 --- a/go-gost/x/go.mod +++ b/go-gost/x/go.mod @@ -4,6 +4,10 @@ go 1.22.0 toolchain go1.23.4 +require ( + github.com/shirou/gopsutil/v3 v3.24.5 +) + require ( github.com/alecthomas/units v0.0.0-20211218093645-b94a6e3cc137 github.com/asaskevich/govalidator v0.0.0-20210307081110-f21760c49a8d diff --git a/go-gost/x/go.sum b/go-gost/x/go.sum index 5aee788..1c3081c 100644 --- a/go-gost/x/go.sum +++ b/go-gost/x/go.sum @@ -204,6 +204,7 @@ github.com/shadowsocks/go-shadowsocks2 v0.1.5 h1:PDSQv9y2S85Fl7VBeOMF9StzeXZyK1H github.com/shadowsocks/go-shadowsocks2 v0.1.5/go.mod h1:AGGpIoek4HRno4xzyFiAtLHkOpcoznZEkAccaI/rplM= github.com/shadowsocks/shadowsocks-go v0.0.0-20200409064450-3e585ff90601 h1:XU9hik0exChEmY92ALW4l9WnDodxLVS9yOSNh2SizaQ= github.com/shadowsocks/shadowsocks-go v0.0.0-20200409064450-3e585ff90601/go.mod h1:mttDPaeLm87u74HMrP+n2tugXvIKWcwff/cqSX0lehY= +github.com/shirou/gopsutil/v3 v3.24.5/go.mod h1:bsoOS1aStSs9ErQ1WWfxllSeS1K5D+U30r2NfcubMVk= github.com/sirupsen/logrus v1.7.0/go.mod h1:yWOB1SBYBC5VeMP7gHvWumXLIWorT60ONWic61uBYv0= github.com/sirupsen/logrus v1.8.1 h1:dJKuHgqk1NNQlqoA6BTlM1Wf9DOH3NBjQyu0h9+AZZE= github.com/sirupsen/logrus v1.8.1/go.mod h1:yWOB1SBYBC5VeMP7gHvWumXLIWorT60ONWic61uBYv0= diff --git a/go-gost/x/socket/chain.go b/go-gost/x/socket/chain.go new file mode 100644 index 0000000..5d6149e --- /dev/null +++ b/go-gost/x/socket/chain.go @@ -0,0 +1,111 @@ +package socket + +import ( + "errors" + "strings" + + "github.com/go-gost/core/logger" + "github.com/go-gost/x/config" + parser "github.com/go-gost/x/config/parsing/chain" + "github.com/go-gost/x/registry" +) + +func createChain(req createChainRequest) error { + + name := strings.TrimSpace(req.Data.Name) + if name == "" { + return errors.New("chain name is required") + } + req.Data.Name = name + + if registry.ChainRegistry().IsRegistered(name) { + return errors.New("chain " + name + " already exists") + } + + v, err := parser.ParseChain(&req.Data, logger.Default()) + if err != nil { + return errors.New("create chain " + name + " failed: " + err.Error()) + } + + if err := registry.ChainRegistry().Register(name, v); err != nil { + return errors.New("chain " + name + " already exists") + } + + config.OnUpdate(func(c *config.Config) error { + c.Chains = append(c.Chains, &req.Data) + return nil + }) + + return nil +} + +func updateChain(req updateChainRequest) error { + + name := strings.TrimSpace(req.Chain) + + if !registry.ChainRegistry().IsRegistered(name) { + return errors.New("chain " + name + " not found") + } + + req.Data.Name = name + + v, err := parser.ParseChain(&req.Data, logger.Default()) + if err != nil { + return errors.New("create chain " + name + " failed: " + err.Error()) + } + + registry.ChainRegistry().Unregister(name) + + if err := registry.ChainRegistry().Register(name, v); err != nil { + return errors.New("chain " + name + " already exists") + } + + config.OnUpdate(func(c *config.Config) error { + for i := range c.Chains { + if c.Chains[i].Name == name { + c.Chains[i] = &req.Data + break + } + } + return nil + }) + + return nil +} + +func deleteChain(req deleteChainRequest) error { + + name := strings.TrimSpace(req.Chain) + + if !registry.ChainRegistry().IsRegistered(name) { + return errors.New("chain " + name + " not found") + } + registry.ChainRegistry().Unregister(name) + + config.OnUpdate(func(c *config.Config) error { + chains := c.Chains + c.Chains = nil + for _, s := range chains { + if s.Name == name { + continue + } + c.Chains = append(c.Chains, s) + } + return nil + }) + + return nil +} + +type createChainRequest struct { + Data config.ChainConfig `json:"data"` +} + +type updateChainRequest struct { + Chain string `json:"chain"` + Data config.ChainConfig `json:"data"` +} + +type deleteChainRequest struct { + Chain string `json:"chain"` +} diff --git a/go-gost/x/socket/config.go b/go-gost/x/socket/config.go new file mode 100644 index 0000000..e3d37db --- /dev/null +++ b/go-gost/x/socket/config.go @@ -0,0 +1,83 @@ +package socket + +import ( + "bytes" + "os" + + "github.com/go-gost/core/observer/stats" + "github.com/go-gost/x/config" + "github.com/go-gost/x/registry" + "github.com/go-gost/x/service" +) + +func saveConfig() { + + file := "gost.json" + + f, err := os.Create(file) + if err != nil { + return + } + defer f.Close() + + if err := config.Global().Write(f, "json"); err != nil { + + return + } + + return +} + +type serviceStatus interface { + Status() *service.Status +} + +type getConfigResponse struct { + Config *config.Config +} + +func getConfig() ([]byte, error) { + + config.OnUpdate(func(c *config.Config) error { + for _, svc := range c.Services { + if svc == nil { + continue + } + s := registry.ServiceRegistry().Get(svc.Name) + ss, ok := s.(serviceStatus) + if ok && ss != nil { + status := ss.Status() + svc.Status = &config.ServiceStatus{ + CreateTime: status.CreateTime().Unix(), + State: string(status.State()), + } + if st := status.Stats(); st != nil { + svc.Status.Stats = &config.ServiceStats{ + TotalConns: st.Get(stats.KindTotalConns), + CurrentConns: st.Get(stats.KindCurrentConns), + TotalErrs: st.Get(stats.KindTotalErrs), + InputBytes: st.Get(stats.KindInputBytes), + OutputBytes: st.Get(stats.KindOutputBytes), + } + } + for _, ev := range status.Events() { + if !ev.Time.IsZero() { + svc.Status.Events = append(svc.Status.Events, config.ServiceEvent{ + Time: ev.Time.Unix(), + Msg: ev.Message, + }) + } + } + } + } + return nil + }) + + var resp getConfigResponse + resp.Config = config.Global() + + buf := &bytes.Buffer{} + + resp.Config.Write(buf, "json") + return buf.Bytes(), nil +} diff --git a/go-gost/x/socket/limiter.go b/go-gost/x/socket/limiter.go new file mode 100644 index 0000000..a41db66 --- /dev/null +++ b/go-gost/x/socket/limiter.go @@ -0,0 +1,102 @@ +package socket + +import ( + "errors" + "github.com/go-gost/x/config" + parser "github.com/go-gost/x/config/parsing/limiter" + "github.com/go-gost/x/registry" + "strings" +) + +func createLimiter(req createLimiterRequest) error { + name := strings.TrimSpace(req.Data.Name) + if name == "" { + return errors.New("limiter name is required") + } + req.Data.Name = name + + if registry.TrafficLimiterRegistry().IsRegistered(name) { + return errors.New("limiter " + name + " already exists") + } + + v := parser.ParseTrafficLimiter(&req.Data) + + if err := registry.TrafficLimiterRegistry().Register(name, v); err != nil { + return errors.New("limiter " + name + " already exists") + } + + config.OnUpdate(func(c *config.Config) error { + c.Limiters = append(c.Limiters, &req.Data) + return nil + }) + + return nil +} + +func updateLimiter(req updateLimiterRequest) error { + + name := strings.TrimSpace(req.Limiter) + + if !registry.TrafficLimiterRegistry().IsRegistered(name) { + return errors.New("limiter " + name + " not found") + } + + req.Data.Name = name + + v := parser.ParseTrafficLimiter(&req.Data) + + registry.TrafficLimiterRegistry().Unregister(name) + + if err := registry.TrafficLimiterRegistry().Register(name, v); err != nil { + return errors.New("limiter " + name + " already exists") + } + + config.OnUpdate(func(c *config.Config) error { + for i := range c.Limiters { + if c.Limiters[i].Name == name { + c.Limiters[i] = &req.Data + break + } + } + return nil + }) + + return nil +} + +func deleteLimiter(req deleteLimiterRequest) error { + + name := strings.TrimSpace(req.Limiter) + + if !registry.TrafficLimiterRegistry().IsRegistered(name) { + return errors.New("limiter " + name + " not found") + } + registry.TrafficLimiterRegistry().Unregister(name) + + config.OnUpdate(func(c *config.Config) error { + limiteres := c.Limiters + c.Limiters = nil + for _, s := range limiteres { + if s.Name == name { + continue + } + c.Limiters = append(c.Limiters, s) + } + return nil + }) + + return nil +} + +type createLimiterRequest struct { + Data config.LimiterConfig `json:"data"` +} + +type updateLimiterRequest struct { + Limiter string `json:"limiter"` + Data config.LimiterConfig `json:"data"` +} + +type deleteLimiterRequest struct { + Limiter string `json:"limiter"` +} diff --git a/go-gost/x/socket/service.go b/go-gost/x/socket/service.go new file mode 100644 index 0000000..d644b06 --- /dev/null +++ b/go-gost/x/socket/service.go @@ -0,0 +1,519 @@ +package socket + +import ( + "errors" + "fmt" + "strings" + "time" + + "github.com/go-gost/core/service" + "github.com/go-gost/x/config" + parser "github.com/go-gost/x/config/parsing/service" + "github.com/go-gost/x/registry" +) + +func createServices(req createServicesRequest) error { + + if len(req.Data) == 0 { + return errors.New("services list cannot be empty") + } + + // 第一阶段:验证所有服务配置 + var parsedServices []struct { + config config.ServiceConfig + service service.Service + } + + for _, serviceConfig := range req.Data { + name := strings.TrimSpace(serviceConfig.Name) + if name == "" { + return errors.New("service name is required") + } + serviceConfig.Name = name + + if registry.ServiceRegistry().IsRegistered(name) { + return errors.New("service " + name + " already exists") + } + + svc, err := parser.ParseService(&serviceConfig) + if err != nil { + return errors.New("create service " + name + " failed: " + err.Error()) + } + + parsedServices = append(parsedServices, struct { + config config.ServiceConfig + service service.Service + }{serviceConfig, svc}) + } + + // 第二阶段:注册所有服务 + var registeredServices []string + for _, ps := range parsedServices { + if err := registry.ServiceRegistry().Register(ps.config.Name, ps.service); err != nil { + // 如果注册失败,回滚已注册的服务 + for _, regName := range registeredServices { + if svc := registry.ServiceRegistry().Get(regName); svc != nil { + registry.ServiceRegistry().Unregister(regName) + svc.Close() + } + } + return errors.New("service " + ps.config.Name + " already exists") + } + registeredServices = append(registeredServices, ps.config.Name) + } + + // 第三阶段:启动所有服务 + for _, ps := range parsedServices { + if svc := registry.ServiceRegistry().Get(ps.config.Name); svc != nil { + go svc.Serve() + } + } + + // 第四阶段:更新配置 + config.OnUpdate(func(c *config.Config) error { + for _, ps := range parsedServices { + c.Services = append(c.Services, &ps.config) + } + return nil + }) + + return nil +} + +func updateServices(req updateServicesRequest) error { + + if len(req.Data) == 0 { + return errors.New("services list cannot be empty") + } + + // 第一阶段:验证所有服务存在 + for _, serviceConfig := range req.Data { + name := strings.TrimSpace(serviceConfig.Name) + if name == "" { + return errors.New("service name is required") + } + serviceConfig.Name = name + + old := registry.ServiceRegistry().Get(name) + if old == nil { + return errors.New("service " + name + " not found") + } + } + + // 第二阶段:按照原来的updateService逻辑,逐个更新服务 + for _, serviceConfig := range req.Data { + name := strings.TrimSpace(serviceConfig.Name) + serviceConfig.Name = name + + // 1. 获取旧服务 + old := registry.ServiceRegistry().Get(name) + + // 2. 关闭旧服务 + old.Close() + + // 3. 从注册表移除旧服务 + registry.ServiceRegistry().Unregister(name) + + // 4. 解析新服务配置 + svc, err := parser.ParseService(&serviceConfig) + if err != nil { + return errors.New("create service " + name + " failed: " + err.Error()) + } + + // 5. 注册新服务 + if err := registry.ServiceRegistry().Register(name, svc); err != nil { + svc.Close() + return errors.New("service " + name + " already exists") + } + + // 6. 启动新服务 + go svc.Serve() + } + + // 第三阶段:更新配置 + config.OnUpdate(func(c *config.Config) error { + for _, serviceConfig := range req.Data { + for i := range c.Services { + if c.Services[i].Name == serviceConfig.Name { + c.Services[i] = &serviceConfig + break + } + } + } + return nil + }) + + return nil +} + +func deleteServices(req deleteServicesRequest) error { + + if len(req.Services) == 0 { + return errors.New("services list cannot be empty") + } + + // 第一阶段:验证所有服务是否存在 + var servicesToDelete []struct { + name string + service service.Service + } + + for _, serviceName := range req.Services { + name := strings.TrimSpace(serviceName) + if name == "" { + return errors.New("service name is required") + } + + svc := registry.ServiceRegistry().Get(name) + if svc == nil { + return errors.New("service " + name + " not found") + } + + servicesToDelete = append(servicesToDelete, struct { + name string + service service.Service + }{name, svc}) + } + + // 第二阶段:删除所有服务 + for _, std := range servicesToDelete { + registry.ServiceRegistry().Unregister(std.name) + std.service.Close() + } + + // 第三阶段:更新配置 + config.OnUpdate(func(c *config.Config) error { + services := c.Services + c.Services = nil + for _, s := range services { + shouldDelete := false + for _, std := range servicesToDelete { + if s.Name == std.name { + shouldDelete = true + break + } + } + if !shouldDelete { + c.Services = append(c.Services, s) + } + } + return nil + }) + + return nil +} + +func pauseServices(req pauseServicesRequest) error { + + if len(req.Services) == 0 { + return errors.New("services list cannot be empty") + } + + // 第一阶段:验证所有服务是否存在,并筛选需要暂停的服务 + var servicesToPause []struct { + name string + service service.Service + } + var skippedServices []string + + cfg := config.Global() + for _, serviceName := range req.Services { + name := strings.TrimSpace(serviceName) + if name == "" { + return errors.New("service name is required") + } + + svc := registry.ServiceRegistry().Get(name) + if svc == nil { + return errors.New(fmt.Sprintf("service %s not found", name)) + } + + // 检查服务是否已经暂停 + var serviceConfig *config.ServiceConfig + for _, s := range cfg.Services { + if s.Name == name { + serviceConfig = s + break + } + } + + // 如果服务已经暂停,跳过 + if serviceConfig != nil && serviceConfig.Metadata != nil { + if pausedVal, exists := serviceConfig.Metadata["paused"]; exists && pausedVal == true { + skippedServices = append(skippedServices, name) + continue + } + } + + servicesToPause = append(servicesToPause, struct { + name string + service service.Service + }{name, svc}) + } + + // 第二阶段:事务性暂停所有服务 + var pausedServices []struct { + name string + service service.Service + serviceConfig *config.ServiceConfig + } + + // 获取服务配置 + serviceConfigs := make(map[string]*config.ServiceConfig) + for _, s := range cfg.Services { + serviceConfigs[s.Name] = s + } + + // 逐个暂停服务,如果失败则回滚 + for _, stp := range servicesToPause { + serviceConfig := serviceConfigs[stp.name] + if serviceConfig == nil { + // 找不到配置,回滚已暂停的服务 + rollbackPausedServices(pausedServices) + return errors.New(fmt.Sprintf("service %s configuration not found", stp.name)) + } + + // 暂停服务 + stp.service.Close() + + // 记录已暂停的服务 + pausedServices = append(pausedServices, struct { + name string + service service.Service + serviceConfig *config.ServiceConfig + }{stp.name, stp.service, serviceConfig}) + } + + // 第三阶段:更新配置,标记暂停状态 + err := config.OnUpdate(func(c *config.Config) error { + for _, stp := range servicesToPause { + for i := range c.Services { + if c.Services[i].Name == stp.name { + if c.Services[i].Metadata == nil { + c.Services[i].Metadata = make(map[string]any) + } + c.Services[i].Metadata["paused"] = true + break + } + } + } + return nil + }) + + if err != nil { + // 配置更新失败,需要回滚所有暂停的服务 + rollbackPausedServices(pausedServices) + return errors.New(fmt.Sprintf("Failed to update config, rolling back paused services: %v", err)) + } + + return nil +} + +func resumeServices(req resumeServicesRequest) error { + if len(req.Services) == 0 { + return errors.New("services list cannot be empty") + } + + // 第一阶段:验证所有服务是否存在,并筛选需要恢复的服务 + var servicesToResume []struct { + name string + service service.Service + serviceConfig *config.ServiceConfig + } + var skippedServices []string + + cfg := config.Global() + for _, serviceName := range req.Services { + name := strings.TrimSpace(serviceName) + if name == "" { + return errors.New("service name is required") + } + + // 检查服务是否存在 + svc := registry.ServiceRegistry().Get(name) + if svc == nil { + return errors.New(fmt.Sprintf("service %s not found", name)) + } + + // 查找配置中的服务 + var serviceConfig *config.ServiceConfig + for _, s := range cfg.Services { + if s.Name == name { + serviceConfig = s + break + } + } + + if serviceConfig == nil { + return errors.New(fmt.Sprintf("service %s configuration not found", name)) + } + + // 检查是否处于暂停状态 + paused := false + if serviceConfig.Metadata != nil { + if pausedVal, exists := serviceConfig.Metadata["paused"]; exists && pausedVal == true { + paused = true + } + } + + // 如果服务没有暂停(即正在运行),跳过 + if !paused { + skippedServices = append(skippedServices, name) + continue + } + + servicesToResume = append(servicesToResume, struct { + name string + service service.Service + serviceConfig *config.ServiceConfig + }{name, svc, serviceConfig}) + } + + // 第二阶段:事务性恢复所有服务 + var resumedServices []struct { + name string + service service.Service + serviceConfig *config.ServiceConfig + } + + // 逐个恢复服务,如果失败则回滚 + for _, str := range servicesToResume { + // 先关闭现有服务 + str.service.Close() + registry.ServiceRegistry().Unregister(str.name) + + // 等待端口释放 + time.Sleep(100 * time.Millisecond) + + // 重新解析并启动服务 + svc, err := parser.ParseService(str.serviceConfig) + if err != nil { + // 恢复失败,回滚已恢复的服务 + rollbackResumedServices(resumedServices) + return errors.New(fmt.Sprintf("resume service %s failed: %s", str.name, err.Error())) + } + + if err := registry.ServiceRegistry().Register(str.name, svc); err != nil { + svc.Close() + // 恢复失败,回滚已恢复的服务 + rollbackResumedServices(resumedServices) + return errors.New(fmt.Sprintf("service %s already exists", str.name)) + } + + go svc.Serve() + + // 记录已成功恢复的服务 + resumedServices = append(resumedServices, str) + } + + // 第三阶段:更新配置,移除暂停状态 + err := config.OnUpdate(func(c *config.Config) error { + for _, str := range servicesToResume { + for i := range c.Services { + if c.Services[i].Name == str.name { + if c.Services[i].Metadata != nil { + delete(c.Services[i].Metadata, "paused") + // 如果 metadata 为空,设置为 nil + if len(c.Services[i].Metadata) == 0 { + c.Services[i].Metadata = nil + } + } + break + } + } + } + return nil + }) + + if err != nil { + // 配置更新失败,回滚所有已恢复的服务 + rollbackResumedServices(resumedServices) + return errors.New(fmt.Sprintf("Failed to update config, rolling back resumed services: %v", err)) + } + + return nil +} + +func rollbackPausedServices(pausedServices []struct { + name string + service service.Service + serviceConfig *config.ServiceConfig +}) { + for _, pss := range pausedServices { + // 重新解析并启动服务 + svc, err := parser.ParseService(pss.serviceConfig) + if err != nil { + continue // 回滚失败,记录日志但继续处理其他服务 + } + + if err := registry.ServiceRegistry().Register(pss.name, svc); err != nil { + svc.Close() + continue // 回滚失败,记录日志但继续处理其他服务 + } + + go svc.Serve() + + // 移除暂停状态标记 + config.OnUpdate(func(c *config.Config) error { + for i := range c.Services { + if c.Services[i].Name == pss.name { + if c.Services[i].Metadata != nil { + delete(c.Services[i].Metadata, "paused") + if len(c.Services[i].Metadata) == 0 { + c.Services[i].Metadata = nil + } + } + break + } + } + return nil + }) + } +} + +func rollbackResumedServices(resumedServices []struct { + name string + service service.Service + serviceConfig *config.ServiceConfig +}) { + for _, rss := range resumedServices { + // 关闭已恢复的服务 + if svc := registry.ServiceRegistry().Get(rss.name); svc != nil { + svc.Close() + } + + // 重新标记为暂停状态 + config.OnUpdate(func(c *config.Config) error { + for i := range c.Services { + if c.Services[i].Name == rss.name { + if c.Services[i].Metadata == nil { + c.Services[i].Metadata = make(map[string]any) + } + c.Services[i].Metadata["paused"] = true + break + } + } + return nil + }) + } +} + +type resumeServicesRequest struct { + Services []string `json:"services"` +} + +type pauseServicesRequest struct { + Services []string `json:"services"` +} + +type deleteServicesRequest struct { + Services []string `json:"services"` +} + +type updateServicesRequest struct { + Data []config.ServiceConfig `json:"data"` +} + +type createServicesRequest struct { + Data []config.ServiceConfig `json:"data"` +} diff --git a/go-gost/x/socket/websocket_reporter.go b/go-gost/x/socket/websocket_reporter.go new file mode 100644 index 0000000..c416d91 --- /dev/null +++ b/go-gost/x/socket/websocket_reporter.go @@ -0,0 +1,752 @@ +package socket + +import ( + "context" + "encoding/json" + "fmt" + "net/url" + "strings" + "time" + + "github.com/go-gost/x/config" + "github.com/gorilla/websocket" + "github.com/shirou/gopsutil/v3/cpu" + "github.com/shirou/gopsutil/v3/host" + "github.com/shirou/gopsutil/v3/mem" + psnet "github.com/shirou/gopsutil/v3/net" +) + +// SystemInfo 系统信息结构体 +type SystemInfo struct { + Uptime uint64 `json:"uptime"` // 开机时间 (秒) + BytesReceived uint64 `json:"bytes_received"` // 接收字节数 + BytesTransmitted uint64 `json:"bytes_transmitted"` // 发送字节数 + CPUUsage float64 `json:"cpu_usage"` // CPU使用率(百分比) + MemoryUsage float64 `json:"memory_usage"` // 内存使用率(百分比) +} + +// NetworkStats 网络统计信息 +type NetworkStats struct { + BytesReceived uint64 `json:"bytes_received"` // 接收字节数 + BytesTransmitted uint64 `json:"bytes_transmitted"` // 发送字节数 +} + +// CPUInfo CPU信息 +type CPUInfo struct { + Usage float64 `json:"usage"` // CPU使用率(百分比) +} + +// MemoryInfo 内存信息 +type MemoryInfo struct { + Usage float64 `json:"usage"` // 内存使用率(百分比) +} + +// CommandMessage 命令消息结构体 +type CommandMessage struct { + Type string `json:"type"` + Data interface{} `json:"data"` + RequestId string `json:"requestId,omitempty"` +} + +// CommandResponse 命令响应结构体 +type CommandResponse struct { + Type string `json:"type"` + Success bool `json:"success"` + Message string `json:"message"` + Data interface{} `json:"data,omitempty"` + RequestId string `json:"requestId,omitempty"` +} + +type WebSocketReporter struct { + url string + conn *websocket.Conn + reconnectTime time.Duration + pingInterval time.Duration + configInterval time.Duration + ctx context.Context + cancel context.CancelFunc + connected bool +} + +// NewWebSocketReporter 创建一个新的WebSocket报告器 +func NewWebSocketReporter(serverURL string) *WebSocketReporter { + ctx, cancel := context.WithCancel(context.Background()) + return &WebSocketReporter{ + url: serverURL, + reconnectTime: 5 * time.Second, // 重连间隔 + pingInterval: 2 * time.Second, // 发送间隔改为2秒 + configInterval: 10 * time.Minute, // 配置上报间隔 + ctx: ctx, + cancel: cancel, + connected: false, + } +} + +// Start 启动WebSocket报告器 +func (w *WebSocketReporter) Start() { + go w.run() +} + +// Stop 停止WebSocket报告器 +func (w *WebSocketReporter) Stop() { + w.cancel() + if w.conn != nil { + w.conn.Close() + } + +} + +// run 主运行循环 +func (w *WebSocketReporter) run() { + for { + select { + case <-w.ctx.Done(): + return + default: + if err := w.connect(); err != nil { + fmt.Printf("❌ WebSocket连接失败: %v,%v后重试\n", err, w.reconnectTime) + select { + case <-time.After(w.reconnectTime): + continue + case <-w.ctx.Done(): + return + } + } + + // 连接成功,开始发送消息 + w.handleConnection() + } + } +} + +// connect 建立WebSocket连接 +func (w *WebSocketReporter) connect() error { + u, err := url.Parse(w.url) + if err != nil { + return fmt.Errorf("解析URL失败: %v", err) + } + + dialer := websocket.DefaultDialer + dialer.HandshakeTimeout = 10 * time.Second + + conn, _, err := dialer.Dial(u.String(), nil) + if err != nil { + return fmt.Errorf("连接WebSocket失败: %v", err) + } + + w.conn = conn + w.connected = true + + // 设置关闭处理器来检测连接状态 + w.conn.SetCloseHandler(func(code int, text string) error { + w.connected = false + return nil + }) + + return nil +} + +// handleConnection 处理WebSocket连接 +func (w *WebSocketReporter) handleConnection() { + defer func() { + if w.conn != nil { + w.conn.Close() + w.conn = nil + } + w.connected = false + }() + + // 启动消息接收goroutine + go w.receiveMessages() + + // 启动配置上报goroutine + go w.reportConfig() + + // 主发送循环 + ticker := time.NewTicker(w.pingInterval) + defer ticker.Stop() + + for { + select { + case <-w.ctx.Done(): + return + case <-ticker.C: + // 检查连接状态 + if !w.connected { + return + } + + // 获取系统信息并发送 + sysInfo := w.collectSystemInfo() + if err := w.sendSystemInfo(sysInfo); err != nil { + fmt.Printf("❌ 发送系统信息失败: %v,准备重连\n", err) + return + } + } + } +} + +// collectSystemInfo 收集系统信息 +func (w *WebSocketReporter) collectSystemInfo() SystemInfo { + networkStats := getNetworkStats() + cpuInfo := getCPUInfo() + memoryInfo := getMemoryInfo() + + return SystemInfo{ + Uptime: getUptime(), + BytesReceived: networkStats.BytesReceived, + BytesTransmitted: networkStats.BytesTransmitted, + CPUUsage: cpuInfo.Usage, + MemoryUsage: memoryInfo.Usage, + } +} + +// sendSystemInfo 发送系统信息 +func (w *WebSocketReporter) sendSystemInfo(sysInfo SystemInfo) error { + if w.conn == nil || !w.connected { + return fmt.Errorf("连接未建立") + } + + // 转换为JSON + jsonData, err := json.Marshal(sysInfo) + if err != nil { + return fmt.Errorf("序列化系统信息失败: %v", err) + } + + // 设置写入超时 + w.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + + if err := w.conn.WriteMessage(websocket.TextMessage, jsonData); err != nil { + w.connected = false // 标记连接已断开 + return fmt.Errorf("写入消息失败: %v", err) + } + + return nil +} + +// receiveMessages 接收服务端发送的消息 +func (w *WebSocketReporter) receiveMessages() { + for { + select { + case <-w.ctx.Done(): + return + default: + if w.conn == nil || !w.connected { + return + } + + // 设置读取超时 + w.conn.SetReadDeadline(time.Now().Add(30 * time.Second)) + + messageType, message, err := w.conn.ReadMessage() + if err != nil { + if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) { + fmt.Printf("❌ WebSocket读取消息错误: %v\n", err) + } + w.connected = false + return + } + + // 处理接收到的消息 + w.handleReceivedMessage(messageType, message) + } + } +} + +// handleReceivedMessage 处理接收到的消息 +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 + } + if cmdMsg.Type != "call" { + w.routeCommand(cmdMsg) + } + + default: + fmt.Printf("📨 收到未知类型消息: %d\n", messageType) + } +} + +// routeCommand 路由命令到对应的处理函数 +func (w *WebSocketReporter) routeCommand(cmd CommandMessage) { + var err error + var response CommandResponse + + // 传递 requestId + response.RequestId = cmd.RequestId + + switch cmd.Type { + // Service 相关命令 + case "AddService": + err = w.handleAddService(cmd.Data) + response.Type = "AddServiceResponse" + case "UpdateService": + err = w.handleUpdateService(cmd.Data) + response.Type = "UpdateServiceResponse" + case "DeleteService": + err = w.handleDeleteService(cmd.Data) + response.Type = "DeleteServiceResponse" + case "PauseService": + err = w.handlePauseService(cmd.Data) + response.Type = "PauseServiceResponse" + case "ResumeService": + err = w.handleResumeService(cmd.Data) + response.Type = "ResumeServiceResponse" + + // Chain 相关命令 + case "AddChains": + err = w.handleAddChain(cmd.Data) + response.Type = "AddChainsResponse" + case "UpdateChains": + err = w.handleUpdateChain(cmd.Data) + response.Type = "UpdateChainsResponse" + case "DeleteChains": + err = w.handleDeleteChain(cmd.Data) + response.Type = "DeleteChainsResponse" + + // Limiter 相关命令 + case "AddLimiters": + err = w.handleAddLimiter(cmd.Data) + response.Type = "AddLimitersResponse" + case "UpdateLimiters": + err = w.handleUpdateLimiter(cmd.Data) + response.Type = "UpdateLimitersResponse" + case "DeleteLimiters": + err = w.handleDeleteLimiter(cmd.Data) + response.Type = "DeleteLimitersResponse" + default: + err = fmt.Errorf("未知命令类型: %s", cmd.Type) + response.Type = "UnknownCommandResponse" + } + + // 发送响应 + if err != nil { + saveConfig() + response.Success = false + response.Message = err.Error() + } else { + saveConfig() + response.Success = true + response.Message = "OK" + } + + w.sendResponse(response) +} + +// Service 命令处理函数 +func (w *WebSocketReporter) handleAddService(data interface{}) error { + // 将 interface{} 转换为 JSON 再解析为具体类型 + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var services []config.ServiceConfig + if err := json.Unmarshal(jsonData, &services); err != nil { + return fmt.Errorf("解析服务配置失败: %v", err) + } + + req := createServicesRequest{Data: services} + return createServices(req) +} + +func (w *WebSocketReporter) handleUpdateService(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var services []config.ServiceConfig + if err := json.Unmarshal(jsonData, &services); err != nil { + return fmt.Errorf("解析服务配置失败: %v", err) + } + + req := updateServicesRequest{Data: services} + return updateServices(req) +} + +func (w *WebSocketReporter) handleDeleteService(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var req deleteServicesRequest + if err := json.Unmarshal(jsonData, &req); err != nil { + return fmt.Errorf("解析删除请求失败: %v", err) + } + + return deleteServices(req) +} + +func (w *WebSocketReporter) handlePauseService(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var req pauseServicesRequest + if err := json.Unmarshal(jsonData, &req); err != nil { + return fmt.Errorf("解析暂停请求失败: %v", err) + } + + return pauseServices(req) +} + +func (w *WebSocketReporter) handleResumeService(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var req resumeServicesRequest + if err := json.Unmarshal(jsonData, &req); err != nil { + return fmt.Errorf("解析恢复请求失败: %v", err) + } + + return resumeServices(req) +} + +// Chain 命令处理函数 +func (w *WebSocketReporter) handleAddChain(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var chainConfig config.ChainConfig + if err := json.Unmarshal(jsonData, &chainConfig); err != nil { + return fmt.Errorf("解析链配置失败: %v", err) + } + + req := createChainRequest{Data: chainConfig} + return createChain(req) +} + +func (w *WebSocketReporter) handleUpdateChain(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + // 对于更新操作,Java端发送的格式可能是: {"chain": "name", "data": {...}} + var updateReq struct { + Chain string `json:"chain"` + Data config.ChainConfig `json:"data"` + } + + // 尝试解析为更新请求格式 + if err := json.Unmarshal(jsonData, &updateReq); err != nil { + // 如果失败,可能是直接的ChainConfig,从name字段获取chain名称 + var chainConfig config.ChainConfig + if err := json.Unmarshal(jsonData, &chainConfig); err != nil { + return fmt.Errorf("解析链配置失败: %v", err) + } + updateReq.Chain = chainConfig.Name + updateReq.Data = chainConfig + } + + req := updateChainRequest{ + Chain: updateReq.Chain, + Data: updateReq.Data, + } + return updateChain(req) +} + +func (w *WebSocketReporter) handleDeleteChain(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + // 删除操作可能是: {"chain": "name"} 或者直接是链名称字符串 + var deleteReq deleteChainRequest + + // 尝试解析为删除请求格式 + if err := json.Unmarshal(jsonData, &deleteReq); err != nil { + // 如果失败,可能是字符串格式的名称 + var chainName string + if err := json.Unmarshal(jsonData, &chainName); err != nil { + return fmt.Errorf("解析链删除请求失败: %v", err) + } + deleteReq.Chain = chainName + } + + return deleteChain(deleteReq) +} + +// Limiter 命令处理函数 +func (w *WebSocketReporter) handleAddLimiter(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + var limiterConfig config.LimiterConfig + if err := json.Unmarshal(jsonData, &limiterConfig); err != nil { + return fmt.Errorf("解析限流器配置失败: %v", err) + } + + req := createLimiterRequest{Data: limiterConfig} + return createLimiter(req) +} + +func (w *WebSocketReporter) handleUpdateLimiter(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + // 对于更新操作,Java端发送的格式可能是: {"limiter": "name", "data": {...}} + var updateReq struct { + Limiter string `json:"limiter"` + Data config.LimiterConfig `json:"data"` + } + + // 尝试解析为更新请求格式 + if err := json.Unmarshal(jsonData, &updateReq); err != nil { + // 如果失败,可能是直接的LimiterConfig,从name字段获取limiter名称 + var limiterConfig config.LimiterConfig + if err := json.Unmarshal(jsonData, &limiterConfig); err != nil { + return fmt.Errorf("解析限流器配置失败: %v", err) + } + updateReq.Limiter = limiterConfig.Name + updateReq.Data = limiterConfig + } + + req := updateLimiterRequest{ + Limiter: updateReq.Limiter, + Data: updateReq.Data, + } + return updateLimiter(req) +} + +func (w *WebSocketReporter) handleDeleteLimiter(data interface{}) error { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化数据失败: %v", err) + } + + // 删除操作可能是: {"limiter": "name"} 或者直接是限流器名称字符串 + var deleteReq deleteLimiterRequest + + // 尝试解析为删除请求格式 + if err := json.Unmarshal(jsonData, &deleteReq); err != nil { + // 如果失败,可能是字符串格式的名称 + var limiterName string + if err := json.Unmarshal(jsonData, &limiterName); err != nil { + return fmt.Errorf("解析限流器删除请求失败: %v", err) + } + deleteReq.Limiter = limiterName + } + + return deleteLimiter(deleteReq) +} + +// handleCall 处理服务端的call回调消息 +func (w *WebSocketReporter) handleCall(data interface{}) error { + // 解析call数据 + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化call数据失败: %v", err) + } + + // 可以根据call的具体内容进行不同的处理 + var callData map[string]interface{} + if err := json.Unmarshal(jsonData, &callData); err != nil { + return fmt.Errorf("解析call数据失败: %v", err) + } + + fmt.Printf("🔔 收到服务端call回调: %v\n", callData) + + // 根据call的类型执行不同的操作 + if callType, exists := callData["type"]; exists { + switch callType { + case "ping": + fmt.Printf("📡 收到ping,发送pong回应\n") + // 可以在这里发送pong响应 + case "info_request": + fmt.Printf("📊 服务端请求额外信息\n") + // 可以在这里发送额外的系统信息 + case "command": + fmt.Printf("⚡ 服务端发送执行命令\n") + // 可以在这里执行特定命令 + default: + fmt.Printf("❓ 未知的call类型: %v\n", callType) + } + } + + // 简单返回成功,表示call已被处理 + return nil +} + +// reportConfig 定时上报配置信息 +func (w *WebSocketReporter) reportConfig() { + // 立即发送一次配置 + w.sendConfigReport() + + // 启动定时器 + ticker := time.NewTicker(w.configInterval) + defer ticker.Stop() + + for { + select { + case <-w.ctx.Done(): + return + case <-ticker.C: + if w.connected { + w.sendConfigReport() + } + } + } +} + +// sendConfigReport 发送配置报告 +func (w *WebSocketReporter) sendConfigReport() { + if w.conn == nil || !w.connected { + return + } + + // 获取配置数据 + configData, err := getConfig() + if err != nil { + fmt.Printf("❌ 获取配置失败: %v\n", err) + return + } + + // 构建配置报告消息 + configMsg := struct { + Type string `json:"type"` + Data interface{} `json:"data"` + }{ + Type: "config_report", + 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 + } + +} + +// sendResponse 发送响应消息到服务端 +func (w *WebSocketReporter) sendResponse(response CommandResponse) { + if w.conn == nil || !w.connected { + fmt.Printf("❌ 无法发送响应:连接未建立\n") + return + } + + jsonData, err := json.Marshal(response) + 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 + } +} + +// sendErrorResponse 发送错误响应 +func (w *WebSocketReporter) sendErrorResponse(responseType, message string) { + response := CommandResponse{ + Type: responseType, + Success: false, + Message: message, + } + w.sendResponse(response) +} + +// getUptime 获取系统开机时间(秒) +func getUptime() uint64 { + uptime, err := host.Uptime() + if err != nil { + return 0 + } + return uptime +} + +// getNetworkStats 获取网络统计信息 +func getNetworkStats() NetworkStats { + var stats NetworkStats + + ioCounters, err := psnet.IOCounters(true) + if err != nil { + fmt.Printf("获取网络统计失败: %v\n", err) + return stats + } + + // 汇总所有非回环接口的流量 + for _, io := range ioCounters { + // 跳过回环接口 + if io.Name == "lo" || strings.HasPrefix(io.Name, "lo") { + continue + } + + stats.BytesReceived += io.BytesRecv + stats.BytesTransmitted += io.BytesSent + } + + return stats +} + +// getCPUInfo 获取CPU信息 +func getCPUInfo() CPUInfo { + var cpuInfo CPUInfo + + // 获取CPU使用率 + percentages, err := cpu.Percent(time.Second, false) + if err == nil && len(percentages) > 0 { + cpuInfo.Usage = percentages[0] + } + + return cpuInfo +} + +// getMemoryInfo 获取内存信息 +func getMemoryInfo() MemoryInfo { + var memInfo MemoryInfo + + vmStat, err := mem.VirtualMemory() + if err != nil { + return memInfo + } + + memInfo.Usage = vmStat.UsedPercent + + return memInfo +} + +// StartWebSocketReporterWithConfig 使用配置启动WebSocket报告器 +func StartWebSocketReporterWithConfig(Addr string, Secret string) *WebSocketReporter { + + // 构建包含本机IP的WebSocket URL + var fullURL = "ws://" + Addr + "/system-info?type=1&secret=" + Secret + + fmt.Printf("🔗 WebSocket连接URL: %s\n", fullURL) + + reporter := NewWebSocketReporter(fullURL) + reporter.Start() + return reporter +} diff --git a/go-gost/x/traffic/memory_manager.go b/go-gost/x/traffic/memory_manager.go index 0137091..b5a1584 100644 --- a/go-gost/x/traffic/memory_manager.go +++ b/go-gost/x/traffic/memory_manager.go @@ -80,6 +80,33 @@ func (m *MemoryManager) ClearAllTrafficStats(ctx context.Context) error { return nil } +// SubtractTrafficStats 从指定服务的流量统计中减去给定的值 +func (m *MemoryManager) SubtractTrafficStats(ctx context.Context, stats map[string]map[string]int64) error { + m.mu.RLock() + defer m.mu.RUnlock() + + for service, serviceStats := range stats { + if managerStats, exists := m.stats[service]; exists { + if upload, ok := serviceStats["upload"]; ok && upload > 0 { + managerStats.upload.Add(-upload) + // 确保不会变成负数 + if managerStats.upload.Load() < 0 { + managerStats.upload.Store(0) + } + } + if download, ok := serviceStats["download"]; ok && download > 0 { + managerStats.download.Add(-download) + // 确保不会变成负数 + if managerStats.download.Load() < 0 { + managerStats.download.Store(0) + } + } + } + } + + return nil +} + // Close 关闭管理器(内存管理器无需特殊清理) func (m *MemoryManager) Close() error { return nil diff --git a/go-gost/x/traffic/traffic.go b/go-gost/x/traffic/traffic.go index 2265571..49a470e 100644 --- a/go-gost/x/traffic/traffic.go +++ b/go-gost/x/traffic/traffic.go @@ -10,6 +10,7 @@ type Manager interface { RecordTraffic(ctx context.Context, service string, upload, download int64) error GetAllServicesStats(ctx context.Context) (map[string]map[string]int64, error) ClearAllTrafficStats(ctx context.Context) error + SubtractTrafficStats(ctx context.Context, stats map[string]map[string]int64) error Close() error TestConnection(ctx context.Context) error } diff --git a/go-gost/traffic_reporter.go b/go-gost/x/traffic/traffic_reporter.go similarity index 93% rename from go-gost/traffic_reporter.go rename to go-gost/x/traffic/traffic_reporter.go index 259e0c5..37b6fa7 100644 --- a/go-gost/traffic_reporter.go +++ b/go-gost/x/traffic/traffic_reporter.go @@ -1,4 +1,4 @@ -package main +package traffic import ( "bytes" @@ -8,8 +8,6 @@ import ( "net/http" "strings" "time" - - "github.com/go-gost/x/traffic" ) // 全局变量存储HTTP地址 @@ -29,7 +27,7 @@ func SetHTTPReportURL(addr string, secret string) { } // StartTrafficReporter 启动流量报告任务 -func StartTrafficReporter(trafficMgr traffic.Manager) { +func StartTrafficReporter(trafficMgr Manager) { // 检查是否设置了HTTP地址 if httpReportURL == "" { fmt.Println("❌ HTTP报告地址未设置,无法启动流量报告任务") @@ -124,12 +122,12 @@ func StartTrafficReporter(trafficMgr traffic.Manager) { if err != nil { fmt.Printf("发送流量报告失败: %v\n", err) } else if success { - // 只有收到"ok"响应才清零流量 - err = trafficMgr.ClearAllTrafficStats(ctx) + // 只有收到"ok"响应才减去已上报的流量 + err = trafficMgr.SubtractTrafficStats(ctx, stats) if err != nil { - fmt.Printf("清零流量统计失败: %v\n", err) + fmt.Printf("减去已上报流量失败: %v\n", err) } else { - fmt.Printf("✅ 流量报告已发送并清零: %d个记录, 总流量: %d bytes\n", + fmt.Printf("✅ 流量报告已发送并减去已上报流量: %d个记录, 总流量: %d bytes\n", totalServices, totalTraffic) } } else { diff --git a/install.sh b/install.sh old mode 100644 new mode 100755 index 836464e..c5633f9 --- a/install.sh +++ b/install.sh @@ -1,113 +1,112 @@ #!/bin/bash -ARCH=$(uname -m) -if [[ "$ARCH" != "x86_64" ]]; then - echo "❌ 不支持的架构: $ARCH,仅支持 x86_64。" - exit 1 -fi - # 下载地址 -DOWNLOAD_URL="https://github.com/bqlpfy/forward-panel/releases/download/gost/gost" +DOWNLOAD_URL="https://ghproxy.com/https://github.com/bqlpfy/forward-panel/releases/download/gost/gost" +INSTALL_DIR="/etc/gost" -# 解析参数 -while getopts "a:p:s:" opt; do +# 显示菜单 +show_menu() { + echo "===============================================" + echo " 管理脚本" + echo "===============================================" + echo "请选择操作:" + echo "1. 安装" + echo "2. 更新" + echo "3. 卸载" + echo "4. 退出" + echo "===============================================" +} + +# 获取用户输入的配置参数 +get_config_params() { + if [[ -z "$SERVER_ADDR" || -z "$SECRET" ]]; then + echo "请输入配置参数:" + + if [[ -z "$SERVER_ADDR" ]]; then + read -p "服务器地址: " SERVER_ADDR + fi + + if [[ -z "$SECRET" ]]; then + read -p "密钥: " SECRET + fi + + if [[ -z "$SERVER_ADDR" || -z "$SECRET" ]]; then + echo "❌ 参数不完整,操作取消。" + exit 1 + fi + fi +} + +# 解析命令行参数 +while getopts "a:s:" opt; do case $opt in a) SERVER_ADDR="$OPTARG" ;; - p) PORT="$OPTARG" ;; s) SECRET="$OPTARG" ;; *) echo "❌ 无效参数"; exit 1 ;; esac done -if [[ -z "$SERVER_ADDR" || -z "$PORT" || -z "$SECRET" ]]; then - echo "用法: $0 -a 服务器地址 -p 端口 -s 密钥" - exit 1 -fi +# 安装功能 +install_gost() { + echo "🚀 开始安装 GOST..." + get_config_params + + mkdir -p "$INSTALL_DIR" -INSTALL_DIR="/etc/gost" -mkdir -p "$INSTALL_DIR" + # 停止并禁用已有服务 + if systemctl list-units --full -all | grep -Fq "gost.service"; then + echo "🔍 检测到已存在的gost服务" + systemctl stop gost 2>/dev/null && echo "🛑 停止服务" + systemctl disable gost 2>/dev/null && echo "🚫 禁用自启" + fi -# 停止并禁用已有服务 -if systemctl list-units --full -all | grep -Fq "gost.service"; then - echo "🔍 检测到已存在的gost服务" - systemctl stop gost 2>/dev/null && echo "🛑 停止服务" - systemctl disable gost 2>/dev/null && echo "🚫 禁用自启" -fi + # 删除旧文件 + [[ -f "$INSTALL_DIR/gost" ]] && echo "🧹 删除旧文件 gost" && rm -f "$INSTALL_DIR/gost" -# 删除旧文件 -[[ -f "$INSTALL_DIR/gost" ]] && echo "🧹 删除旧文件 gost" && rm -f "$INSTALL_DIR/gost" + # 下载 gost + echo "⬇️ 下载 gost 中..." + curl -L "$DOWNLOAD_URL" -o "$INSTALL_DIR/gost" + if [[ ! -f "$INSTALL_DIR/gost" || ! -s "$INSTALL_DIR/gost" ]]; then + echo "❌ 下载失败,请检查网络或下载链接。" + exit 1 + fi + chmod +x "$INSTALL_DIR/gost" + echo "✅ 下载完成" -# 下载 gost -echo "⬇️ 下载 gost 中..." -curl -L "$DOWNLOAD_URL" -o "$INSTALL_DIR/gost" -if [[ ! -f "$INSTALL_DIR/gost" || ! -s "$INSTALL_DIR/gost" ]]; then - echo "❌ 下载失败,请检查网络或下载链接。" - exit 1 -fi -chmod +x "$INSTALL_DIR/gost" -echo "✅ 下载完成" + # 打印版本 + echo "🔎 gost 版本:$($INSTALL_DIR/gost -V)" -# 打印版本 -echo "🔎 gost 版本:$($INSTALL_DIR/gost -V)" - -# 写入 config.json -CONFIG_FILE="$INSTALL_DIR/config.json" -if [[ -f "$CONFIG_FILE" ]]; then - echo "📝 更新配置: config.json" - sed -i.bak "s|\"addr\": \".*\"|\"addr\": \"$SERVER_ADDR\"|g" "$CONFIG_FILE" - sed -i.bak "s|\"secret\": \".*\"|\"secret\": \"$SECRET\"|g" "$CONFIG_FILE" - rm -f "$CONFIG_FILE.bak" -else - echo "📄 创建新配置: config.json" - cat > "$CONFIG_FILE" < "$CONFIG_FILE" < "$GOST_CONFIG" < "$GOST_CONFIG" < "$SERVICE_FILE" < "$SERVICE_FILE" </dev/null && echo "✨ 安装脚本已自动清理" || echo "⚠️ 安装脚本清理失败,请手动删除" +# 更新功能 +update_gost() { + echo "🔄 开始更新 GOST..." + + if [[ ! -d "$INSTALL_DIR" ]]; then + echo "❌ GOST 未安装,请先选择安装。" + return 1 + fi + + # 先下载新版本 + echo "⬇️ 下载最新版本..." + curl -L "$DOWNLOAD_URL" -o "$INSTALL_DIR/gost.new" + if [[ ! -f "$INSTALL_DIR/gost.new" || ! -s "$INSTALL_DIR/gost.new" ]]; then + echo "❌ 下载失败。" + return 1 + fi + + # 停止服务 + if systemctl list-units --full -all | grep -Fq "gost.service"; then + echo "🛑 停止 gost 服务..." + systemctl stop gost + fi + + # 替换文件 + mv "$INSTALL_DIR/gost.new" "$INSTALL_DIR/gost" + chmod +x "$INSTALL_DIR/gost" + + # 打印版本 + echo "🔎 新版本:$($INSTALL_DIR/gost -V)" + + # 重启服务 + echo "🔄 重启服务..." + systemctl start gost + + echo "✅ 更新完成,服务已重新启动。" +} + +# 卸载功能 +uninstall_gost() { + echo "🗑️ 开始卸载 GOST..." + + read -p "确认卸载 GOST 吗?此操作将删除所有相关文件 (y/N): " confirm + if [[ "$confirm" != "y" && "$confirm" != "Y" ]]; then + echo "❌ 取消卸载" + return 0 + fi + + # 停止并禁用服务 + if systemctl list-units --full -all | grep -Fq "gost.service"; then + echo "🛑 停止并禁用服务..." + systemctl stop gost 2>/dev/null + systemctl disable gost 2>/dev/null + fi + + # 删除服务文件 + if [[ -f "/etc/systemd/system/gost.service" ]]; then + rm -f "/etc/systemd/system/gost.service" + echo "🧹 删除服务文件" + fi + + # 删除安装目录 + if [[ -d "$INSTALL_DIR" ]]; then + rm -rf "$INSTALL_DIR" + echo "🧹 删除安装目录: $INSTALL_DIR" + fi + + # 重载 systemd + systemctl daemon-reload + + echo "✅ 卸载完成" +} + +# 主逻辑 +main() { + # 如果提供了命令行参数,直接执行安装 + if [[ -n "$SERVER_ADDR" && -n "$SECRET" ]]; then + install_gost + exit 0 + fi + + # 显示交互式菜单 + while true; do + show_menu + read -p "请输入选项 (1-4): " choice + + case $choice in + 1) + install_gost + break + ;; + 2) + update_gost + break + ;; + 3) + uninstall_gost + break + ;; + 4) + echo "👋 退出脚本" + exit 0 + ;; + *) + echo "❌ 无效选项,请输入 1-4" + echo "" + ;; + esac + done +} + +# 执行主函数 +main \ No newline at end of file diff --git a/panel_install.sh b/panel_install.sh index aa5e498..d85710a 100755 --- a/panel_install.sh +++ b/panel_install.sh @@ -5,62 +5,84 @@ set -e export LANG=en_US.UTF-8 export LC_ALL=C +# 全局下载地址配置 +DOCKER_COMPOSE_URL="https://github.com/bqlpfy/forward-panel/raw/refs/heads/main/docker-compose.yml" +GOST_SQL_URL="https://github.com/bqlpfy/forward-panel/raw/refs/heads/main/gost.sql" + # 检查 docker-compose 或 docker compose 命令 -if command -v docker-compose &> /dev/null; then - DOCKER_CMD="docker-compose" -elif command -v docker &> /dev/null; then - if docker compose version &> /dev/null; then - DOCKER_CMD="docker compose" +check_docker() { + if command -v docker-compose &> /dev/null; then + DOCKER_CMD="docker-compose" + elif command -v docker &> /dev/null; then + if docker compose version &> /dev/null; then + DOCKER_CMD="docker compose" + else + echo "错误:检测到 docker,但不支持 'docker compose' 命令。请安装 docker-compose 或更新 docker 版本。" + exit 1 + fi else - echo "错误:检测到 docker,但不支持 'docker compose' 命令。请安装 docker-compose 或更新 docker 版本。" + echo "错误:未检测到 docker 或 docker-compose 命令。请先安装 Docker。" exit 1 fi -else - echo "错误:未检测到 docker 或 docker-compose 命令。请先安装 Docker。" - exit 1 -fi + echo "检测到 Docker 命令:$DOCKER_CMD" +} -echo "检测到 Docker 命令:$DOCKER_CMD" - -echo "🔽 下载必要文件..." -curl -L -o docker-compose.yml https://github.com/bqlpfy/forward-panel/raw/refs/heads/main/docker-compose.yml -curl -L -o gost.sql https://github.com/bqlpfy/forward-panel/raw/refs/heads/main/gost.sql -echo "✅ 下载完成" +# 显示菜单 +show_menu() { + echo "===============================================" + echo " 面板管理脚本" + echo "===============================================" + echo "请选择操作:" + echo "1. 安装面板" + echo "2. 更新面板" + echo "3. 卸载面板" + echo "4. 退出" + echo "===============================================" +} generate_random() { LC_ALL=C tr -dc 'A-Za-z0-9' .env < .env </dev/null && echo "✨ 下载文件已清理" || echo "⚠️ 下载文件清理失败" -#rm -f "$0" 2>/dev/null && echo "✨ 安装脚本已自动清理" || echo "⚠️ 安装脚本清理失败,请手动删除" +# 更新功能 +update_panel() { + echo "🔄 开始更新面板..." + check_docker + + echo "🔽 下载最新配置文件..." + curl -L -o docker-compose.yml "$DOCKER_COMPOSE_URL" + echo "✅ 下载完成" + + echo "🛑 停止当前服务..." + $DOCKER_CMD down + + echo "⬇️ 拉取最新镜像..." + $DOCKER_CMD pull + + echo "🚀 启动更新后的服务..." + $DOCKER_CMD up -d + + # 等待服务启动 + echo "⏳ 等待服务启动..." + + # 检查后端容器健康状态 + echo "🔍 检查后端服务状态..." + for i in {1..90}; do + if docker ps --format "{{.Names}}" | grep -q "^springboot-backend$"; then + BACKEND_HEALTH=$(docker inspect -f '{{.State.Health.Status}}' springboot-backend 2>/dev/null || echo "unknown") + if [[ "$BACKEND_HEALTH" == "healthy" ]]; then + echo "✅ 后端服务健康检查通过" + break + elif [[ "$BACKEND_HEALTH" == "starting" ]]; then + # 继续等待 + : + elif [[ "$BACKEND_HEALTH" == "unhealthy" ]]; then + echo "⚠️ 后端健康状态:$BACKEND_HEALTH" + fi + else + echo "⚠️ 后端容器未找到或未运行" + BACKEND_HEALTH="not_running" + fi + if [ $i -eq 90 ]; then + echo "❌ 后端服务启动超时(90秒)" + echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' springboot-backend 2>/dev/null || echo '容器不存在')" + echo "🛑 更新终止" + return 1 + fi + # 每15秒显示一次进度 + if [ $((i % 15)) -eq 1 ]; then + echo "⏳ 等待后端服务启动... ($i/90) 状态:${BACKEND_HEALTH:-unknown}" + fi + sleep 1 + done + + # 检查数据库容器健康状态 + echo "🔍 检查数据库服务状态..." + for i in {1..60}; do + if docker ps --format "{{.Names}}" | grep -q "^gost-mysql$"; then + DB_HEALTH=$(docker inspect -f '{{.State.Health.Status}}' gost-mysql 2>/dev/null || echo "unknown") + if [[ "$DB_HEALTH" == "healthy" ]]; then + echo "✅ 数据库服务健康检查通过" + break + elif [[ "$DB_HEALTH" == "starting" ]]; then + # 继续等待 + : + elif [[ "$DB_HEALTH" == "unhealthy" ]]; then + echo "⚠️ 数据库健康状态:$DB_HEALTH" + fi + else + echo "⚠️ 数据库容器未找到或未运行" + DB_HEALTH="not_running" + fi + if [ $i -eq 60 ]; then + echo "❌ 数据库服务启动超时(60秒)" + echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' gost-mysql 2>/dev/null || echo '容器不存在')" + echo "🛑 更新终止" + return 1 + fi + # 每10秒显示一次进度 + if [ $((i % 10)) -eq 1 ]; then + echo "⏳ 等待数据库服务启动... ($i/60) 状态:${DB_HEALTH:-unknown}" + fi + sleep 1 + done + + # 从容器环境变量获取数据库信息 + echo "🔍 获取数据库配置信息..." + + # 等待一下让服务完全就绪 + echo "⏳ 等待服务完全就绪..." + sleep 5 + + # 先检查后端容器是否在运行 + if ! docker ps --format "{{.Names}}" | grep -q "^springboot-backend$"; then + echo "❌ 后端容器未运行,无法获取数据库配置" + echo "🔍 当前运行的容器:" + docker ps --format "table {{.Names}}\t{{.Status}}" + echo "🛑 更新终止" + return 1 + fi + + DB_INFO=$(docker exec springboot-backend env | grep "^DB_" 2>/dev/null || echo "") + + if [[ -n "$DB_INFO" ]]; then + DB_NAME=$(echo "$DB_INFO" | grep "^DB_NAME=" | cut -d'=' -f2) + DB_PASSWORD=$(echo "$DB_INFO" | grep "^DB_PASSWORD=" | cut -d'=' -f2) + DB_USER=$(echo "$DB_INFO" | grep "^DB_USER=" | cut -d'=' -f2) + DB_HOST=$(echo "$DB_INFO" | grep "^DB_HOST=" | cut -d'=' -f2) + + echo "📋 数据库配置:" + echo " 数据库名: $DB_NAME" + echo " 用户名: $DB_USER" + echo " 主机: $DB_HOST" + else + echo "❌ 无法获取数据库配置信息" + echo "🔍 尝试诊断问题:" + echo " 容器状态: $(docker inspect -f '{{.State.Status}}' springboot-backend 2>/dev/null || echo '容器不存在')" + echo " 健康状态: $(docker inspect -f '{{.State.Health.Status}}' springboot-backend 2>/dev/null || echo '无健康检查')" + + # 尝试从 .env 文件读取配置 + if [[ -f ".env" ]]; then + echo "🔄 尝试从 .env 文件读取配置..." + DB_NAME=$(grep "^DB_NAME=" .env | cut -d'=' -f2 2>/dev/null) + DB_PASSWORD=$(grep "^DB_PASSWORD=" .env | cut -d'=' -f2 2>/dev/null) + DB_USER=$(grep "^DB_USER=" .env | cut -d'=' -f2 2>/dev/null) + + if [[ -n "$DB_NAME" && -n "$DB_PASSWORD" && -n "$DB_USER" ]]; then + echo "✅ 从 .env 文件成功读取数据库配置" + echo "📋 数据库配置:" + echo " 数据库名: $DB_NAME" + echo " 用户名: $DB_USER" + else + echo "❌ .env 文件中的数据库配置不完整" + echo "🛑 更新终止" + return 1 + fi + else + echo "❌ 未找到 .env 文件" + echo "🛑 更新终止" + return 1 + fi + fi + + # 检查必要的数据库配置 + if [[ -z "$DB_PASSWORD" || -z "$DB_USER" || -z "$DB_NAME" ]]; then + echo "❌ 数据库配置不完整(缺少必要参数)" + echo "🛑 更新终止" + return 1 + fi + + # 执行数据库字段变更 + echo "🔄 执行数据库结构更新..." + + # 创建临时迁移文件(现在有了数据库信息) + cat > temp_migration.sql </dev/null; then + echo "✅ 数据库结构更新完成" + else + echo "⚠️ 使用用户密码失败,尝试root密码..." + if docker exec -i gost-mysql mysql -u root -p"$DB_PASSWORD" < temp_migration.sql 2>/dev/null; then + echo "✅ 数据库结构更新完成" + else + echo "❌ 数据库结构更新失败,请手动执行 temp_migration.sql" + echo "📁 迁移文件已保存为 temp_migration.sql" + echo "🔍 数据库容器状态: $(docker inspect -f '{{.State.Status}}' gost-mysql 2>/dev/null || echo '容器不存在')" + echo "🛑 更新终止" + return 1 + fi + fi + + # 清理临时文件 + rm -f temp_migration.sql + + echo "✅ 更新完成" +} + +# 卸载功能 +uninstall_panel() { + echo "🗑️ 开始卸载面板..." + check_docker + + if [[ ! -f "docker-compose.yml" ]]; then + echo "⚠️ 未找到 docker-compose.yml 文件,正在下载以完成卸载..." + curl -L -o docker-compose.yml "$DOCKER_COMPOSE_URL" + echo "✅ docker-compose.yml 下载完成" + fi + + read -p "确认卸载面板吗?此操作将停止并删除所有容器和数据 (y/N): " confirm + if [[ "$confirm" != "y" && "$confirm" != "Y" ]]; then + echo "❌ 取消卸载" + return 0 + fi + + echo "🛑 停止并删除容器、镜像、卷..." + $DOCKER_CMD down --rmi all --volumes --remove-orphans + echo "🧹 删除配置文件..." + rm -f docker-compose.yml gost.sql .env + echo "✅ 卸载完成" +} + +# 主逻辑 +main() { + # 显示交互式菜单 + while true; do + show_menu + read -p "请输入选项 (1-4): " choice + + case $choice in + 1) + install_panel + break + ;; + 2) + update_panel + break + ;; + 3) + uninstall_panel + break + ;; + 4) + echo "👋 退出脚本" + exit 0 + ;; + *) + echo "❌ 无效选项,请输入 1-4" + echo "" + ;; + esac + done +} + +# 执行主函数 +main \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/dto/ForwardDto.java b/springboot-backend/src/main/java/com/admin/common/dto/ForwardDto.java index 74b8f7f..b39cc79 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/ForwardDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/ForwardDto.java @@ -3,6 +3,8 @@ package com.admin.common.dto; import lombok.Data; import javax.validation.constraints.NotBlank; import javax.validation.constraints.NotNull; +import javax.validation.constraints.Min; +import javax.validation.constraints.Max; @Data public class ForwardDto { @@ -15,4 +17,11 @@ public class ForwardDto { @NotBlank(message = "远程地址不能为空") private String remoteAddr; + + /** + * 入口端口(可选,为空时自动分配) + */ + @Min(value = 1, message = "端口号不能小于1") + @Max(value = 65535, message = "端口号不能大于65535") + private Integer inPort; } \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/dto/ForwardUpdateDto.java b/springboot-backend/src/main/java/com/admin/common/dto/ForwardUpdateDto.java index 7e37e27..c75d4c7 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/ForwardUpdateDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/ForwardUpdateDto.java @@ -3,6 +3,8 @@ package com.admin.common.dto; import lombok.Data; import javax.validation.constraints.NotBlank; import javax.validation.constraints.NotNull; +import javax.validation.constraints.Min; +import javax.validation.constraints.Max; @Data public class ForwardUpdateDto { @@ -21,4 +23,11 @@ public class ForwardUpdateDto { @NotBlank(message = "远程地址不能为空") private String remoteAddr; + + /** + * 入口端口(可选,为空时自动分配) + */ + @Min(value = 1, message = "端口号不能小于1") + @Max(value = 65535, message = "端口号不能大于65535") + private Integer inPort; } \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/dto/NodeDto.java b/springboot-backend/src/main/java/com/admin/common/dto/NodeDto.java index ea5cc61..b55a57f 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/NodeDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/NodeDto.java @@ -13,8 +13,9 @@ public class NodeDto { @NotBlank(message = "节点名称不能为空") private String name; - @NotNull(message = "控制端口不能为空") - @Min(value = 1, message = "端口号必须在1-65535之间") - @Max(value = 65535, message = "端口号必须在1-65535之间") - private Integer port; + @NotBlank(message = "入口IP不能为空") + private String ip; + + @NotBlank(message = "服务器ip不能为空") + private String serverIp; } \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/dto/NodeUpdateDto.java b/springboot-backend/src/main/java/com/admin/common/dto/NodeUpdateDto.java index fe6e776..59b8983 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/NodeUpdateDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/NodeUpdateDto.java @@ -14,6 +14,9 @@ public class NodeUpdateDto { @NotBlank(message = "节点名称不能为空") private String name; - @NotBlank(message = "节点IP不能为空") + @NotBlank(message = "入口IP不能为空") private String ip; + + @NotBlank(message = "服务器ip不能为空") + private String serverIp; } \ No newline at end of file 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 9f5ada7..4109417 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 @@ -46,4 +46,10 @@ public class TunnelDto { // 协议类型(隧道转发时使用:tls、tcp、mtls),默认为tls private String protocol; + + // TCP监听地址,默认为0.0.0.0 + private String tcpListenAddr = "0.0.0.0"; + + // UDP监听地址,默认为0.0.0.0 + private String udpListenAddr = "0.0.0.0"; } \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/dto/TunnelListDto.java b/springboot-backend/src/main/java/com/admin/common/dto/TunnelListDto.java index e1d8e7f..af10338 100644 --- a/springboot-backend/src/main/java/com/admin/common/dto/TunnelListDto.java +++ b/springboot-backend/src/main/java/com/admin/common/dto/TunnelListDto.java @@ -8,6 +8,44 @@ public class TunnelListDto { private Integer id; private String name; - - + + /** + * 入口IP + */ + private String ip; + + /** + * 入口端口范围开始 + */ + private Integer inPortSta; + + /** + * 入口端口范围结束 + */ + private Integer inPortEnd; + + /** + * 出口IP + */ + private String outIp; + + /** + * 出口端口范围开始 + */ + private Integer outIpSta; + + /** + * 出口端口范围结束 + */ + private Integer outIpEnd; + + /** + * 隧道类型(1-端口转发,2-隧道转发) + */ + private Integer type; + + /** + * 协议类型 + */ + 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 49fed26..a78866c 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 @@ -37,4 +37,10 @@ public class TunnelUpdateDto { @Min(value = 1, message = "出口端口结束必须大于等于0") @Max(value = 65535, message = "出口端口结束不能超过65535") private Integer outIpEnd; + + // TCP监听地址 + private String tcpListenAddr; + + // UDP监听地址 + private String udpListenAddr; } \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/task/CheckGostConfigAsync.java b/springboot-backend/src/main/java/com/admin/common/task/CheckGostConfigAsync.java index f0b3204..fce9bb7 100644 --- a/springboot-backend/src/main/java/com/admin/common/task/CheckGostConfigAsync.java +++ b/springboot-backend/src/main/java/com/admin/common/task/CheckGostConfigAsync.java @@ -12,16 +12,17 @@ import com.admin.service.SpeedLimitService; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; import lombok.extern.slf4j.Slf4j; import org.springframework.context.annotation.Configuration; +import org.springframework.scheduling.annotation.Async; import org.springframework.scheduling.annotation.EnableScheduling; import org.springframework.scheduling.annotation.Scheduled; +import org.springframework.stereotype.Service; import javax.annotation.Resource; import java.util.List; import java.util.Objects; @Slf4j -@Configuration -@EnableScheduling +@Service public class CheckGostConfigAsync { @Resource @@ -33,39 +34,17 @@ public class CheckGostConfigAsync { @Resource private SpeedLimitService speedLimitService; - /** - * 启动后10秒执行一次,然后每10分钟执行一次 - * 清理孤立的Gost配置项 - */ - @Scheduled(initialDelay = 10000, fixedRate = 600000) - public void cleanOrphanedGostConfigs() { - log.info("开始清理孤立的Gost配置项"); - - List activeNodes = nodeService.list(new QueryWrapper().eq("status", 1)); - log.info("找到 {} 个活跃节点", activeNodes.size()); - - for (Node node : activeNodes) { - cleanNodeConfigs(node); - } - - log.info("Gost配置清理任务完成"); - } /** - * 清理单个节点的配置 + * 清理孤立的Gost配置项 */ - private void cleanNodeConfigs(Node node) { - String nodeAddress = node.getIp() + ":" + node.getPort(); - - try { - GostConfigDto gostConfig = GostUtil.GetConfig(nodeAddress, node.getSecret()); - + @Async + public void cleanNodeConfigs(String node_id, GostConfigDto gostConfig) { + Node node = nodeService.getById(node_id); + if (node != null) { cleanOrphanedServices(gostConfig, node); cleanOrphanedChains(gostConfig, node); cleanOrphanedLimiters(gostConfig, node); - - } catch (Exception e) { - log.error("清理节点 {} 配置时发生错误", nodeAddress, e); } } @@ -76,8 +55,7 @@ public class CheckGostConfigAsync { if (gostConfig.getServices() == null) { return; } - - String nodeAddress = node.getIp() + ":" + node.getPort(); + for (ConfigItem service : gostConfig.getServices()) { safeExecute(() -> { @@ -95,8 +73,15 @@ public class CheckGostConfigAsync { if (Objects.equals(type, "tcp")) { // 只处理TCP,避免重复处理 Forward forward = forwardService.getById(forwardId); if (forward == null) { - log.warn("删除孤立的服务: {} (节点: {})", service.getName(), nodeAddress); - GostUtil.DeleteService(nodeAddress, forwardId+"_"+userId+"_"+userTunnelId, node.getSecret()); + log.warn("删除孤立的服务: {} (节点: {})", service.getName(), node.getId()); + GostUtil.DeleteService(node.getId(), forwardId+"_"+userId+"_"+userTunnelId); + } + } + if (Objects.equals(type, "tls")) { + Forward forward = forwardService.getById(forwardId); + if (forward == null) { + log.warn("删除孤立的服务: {} (节点: {})", service.getName(), node.getId()); + GostUtil.DeleteRemoteService(node.getId(), forwardId+"_"+userId+"_"+userTunnelId); } } } @@ -112,8 +97,7 @@ public class CheckGostConfigAsync { return; } - String nodeAddress = node.getIp() + ":" + node.getPort(); - + for (ConfigItem chain : gostConfig.getChains()) { safeExecute(() -> { String[] serviceIds = parseServiceName(chain.getName()); @@ -126,8 +110,8 @@ public class CheckGostConfigAsync { if (Objects.equals(type, "chains")) { Forward forward = forwardService.getById(forwardId); if (forward == null) { - log.warn("删除孤立的链: {} (节点: {})", chain.getName(), nodeAddress); - GostUtil.DeleteChains(nodeAddress, forwardId+"_"+userId+"_"+userTunnelId, node.getSecret()); + log.warn("删除孤立的链: {} (节点: {})", chain.getName(), node.getId()); + GostUtil.DeleteChains(node.getId(), forwardId+"_"+userId+"_"+userTunnelId); } } } @@ -143,14 +127,13 @@ public class CheckGostConfigAsync { return; } - String nodeAddress = node.getIp() + ":" + node.getPort(); - + for (ConfigItem limiter : gostConfig.getLimiters()) { safeExecute(() -> { SpeedLimit speedLimit = speedLimitService.getById(limiter.getName()); if (speedLimit == null) { - log.warn("删除孤立的限流器: {} (节点: {})", limiter.getName(), nodeAddress); - GostUtil.DeleteLimiters(nodeAddress, Long.parseLong(limiter.getName()), node.getSecret()); + log.warn("删除孤立的限流器: {} (节点: {})", limiter.getName(), node.getId()); + GostUtil.DeleteLimiters(node.getId(), Long.parseLong(limiter.getName())); } }, "清理限流器 " + limiter.getName()); } diff --git a/springboot-backend/src/main/java/com/admin/common/task/DelayQueueManager.java b/springboot-backend/src/main/java/com/admin/common/task/DelayQueueManager.java index baea91e..2c357cb 100644 --- a/springboot-backend/src/main/java/com/admin/common/task/DelayQueueManager.java +++ b/springboot-backend/src/main/java/com/admin/common/task/DelayQueueManager.java @@ -235,17 +235,15 @@ public class DelayQueueManager implements CommandLineRunner { } String serviceName = buildServiceName(forward.getId(), userId, userTunnel.getId()); - String nodeAddress = buildNodeAddress(inNode); // 暂停主服务 - GostDto result = GostUtil.PauseService(nodeAddress, serviceName, inNode.getSecret()); + GostDto result = GostUtil.PauseService(inNode.getId(), serviceName); // 隧道转发需要同时暂停远端服务 if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD Node outNode = nodeMapper.selectById(tunnel.getOutNodeId()); if (outNode != null) { - String outNodeAddress = buildNodeAddress(outNode); - GostDto remoteResult = GostUtil.PauseRemoteService(outNodeAddress, serviceName, outNode.getSecret()); + GostDto remoteResult = GostUtil.PauseRemoteService(outNode.getId(), serviceName); if (!"OK".equals(remoteResult.getMsg())) { log.warn("暂停远端服务失败,转发ID:{},用户ID:{},服务名:{},结果:{}", forward.getId(), userId, serviceName, remoteResult.getMsg()); @@ -297,16 +295,7 @@ public class DelayQueueManager implements CommandLineRunner { return forwardId + "_" + userId + "_" + userTunnelId; } - /** - * 构建节点地址 - * - * @param node 节点对象 - * @return 节点地址字符串 - */ - private String buildNodeAddress(Node node) { - return node.getIp() + ":" + node.getPort(); - } - + /** * 初始化用户账号到期延时任务 * 查询所有非管理员的正常用户,为有到期时间且未过期的用户创建延时任务 @@ -455,17 +444,15 @@ public class DelayQueueManager implements CommandLineRunner { } String serviceName = buildServiceName(forward.getId(), Long.valueOf(userTunnel.getUserId()), userTunnel.getId()); - String nodeAddress = buildNodeAddress(inNode); // 暂停服务 - GostDto result = GostUtil.PauseService(nodeAddress, serviceName, inNode.getSecret()); + GostDto result = GostUtil.PauseService(inNode.getId(), serviceName); // 隧道转发需要同时暂停远端服务 if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD Node outNode = nodeMapper.selectById(tunnel.getOutNodeId()); if (outNode != null) { - String outNodeAddress = buildNodeAddress(outNode); - GostDto remoteResult = GostUtil.PauseRemoteService(outNodeAddress, serviceName, outNode.getSecret()); + GostDto remoteResult = GostUtil.PauseRemoteService(outNode.getId(), serviceName); if (!"OK".equals(remoteResult.getMsg())) { log.warn("暂停远端服务失败,转发ID:{},用户ID:{},隧道ID:{},服务名:{},结果:{}", forward.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), serviceName, remoteResult.getMsg()); diff --git a/springboot-backend/src/main/java/com/admin/common/task/SaveConfigAsync.java b/springboot-backend/src/main/java/com/admin/common/task/SaveConfigAsync.java deleted file mode 100644 index 4a1c4f7..0000000 --- a/springboot-backend/src/main/java/com/admin/common/task/SaveConfigAsync.java +++ /dev/null @@ -1,22 +0,0 @@ -package com.admin.common.task; - - -import com.admin.common.dto.GostDto; -import com.admin.common.utils.GostUtil; -import org.springframework.scheduling.annotation.Async; -import org.springframework.stereotype.Service; - -@Service -public class SaveConfigAsync { - - - @Async - public void run(String addr, String secret){ - try { - GostUtil.SaveConfig(addr, secret); - }catch (Exception e){ - e.printStackTrace(); - } - } - -} diff --git a/springboot-backend/src/main/java/com/admin/common/utils/GostUtil.java b/springboot-backend/src/main/java/com/admin/common/utils/GostUtil.java index 06be1c0..8618c5c 100644 --- a/springboot-backend/src/main/java/com/admin/common/utils/GostUtil.java +++ b/springboot-backend/src/main/java/com/admin/common/utils/GostUtil.java @@ -2,212 +2,99 @@ package com.admin.common.utils; import com.admin.common.dto.GostConfigDto; import com.admin.common.dto.GostDto; +import com.admin.entity.Tunnel; import com.alibaba.fastjson.JSONArray; import com.alibaba.fastjson.JSONObject; import org.aspectj.apache.bcel.generic.RET; +import java.util.Objects; + public class GostUtil { - private static final String API_BASE_URL = "/api/config/"; - private static final String LIMITERS_ENDPOINT = "limiters"; - private static final String SERVICES_ENDPOINT = "services"; - private static final String CHAINS_ENDPOINT = "chains"; - - public static GostDto SaveConfig(String addr, String secret) { - JSONObject data = new JSONObject(); - data.put("format", "json"); - - if (!addr.contains("[") && addr.indexOf(':') != addr.lastIndexOf(':')) { - // 这是IPv6地址,找到最后一个冒号(端口分隔符) - int lastColonIndex = addr.lastIndexOf(':'); - String ipPart = addr.substring(0, lastColonIndex); - String portPart = addr.substring(lastColonIndex); - addr = "[" + ipPart + "]" + portPart; - } - - String url = "https://" + addr + "/api/config?format=json"; - return HttpUtils.post(url, data, secret); - } - - public static GostConfigDto GetConfig(String addr, String secret) { - JSONObject data = new JSONObject(); - data.put("format", "json"); - - if (!addr.contains("[") && addr.indexOf(':') != addr.lastIndexOf(':')) { - // 这是IPv6地址,找到最后一个冒号(端口分隔符) - int lastColonIndex = addr.lastIndexOf(':'); - String ipPart = addr.substring(0, lastColonIndex); - String portPart = addr.substring(lastColonIndex); - addr = "[" + ipPart + "]" + portPart; - } - - String url = "https://" + addr + "/api/config?format=json"; - return HttpUtils.get(url, secret); - } - - - /** - * 添加限流器配置 - * - * @param addr 服务器地址 - * @param name 限流器名称 - * @param speed 限速值(MB) - * @param secret 认证密钥 - * @return 请求结果 - */ - public static GostDto AddLimiters(String addr, Long name, String speed, String secret) { + public static GostDto AddLimiters(Long node_id, Long name, String speed) { JSONObject data = createLimiterData(name, speed); - String url = buildUrl(addr, LIMITERS_ENDPOINT); - return HttpUtils.post(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "AddLimiters"); } - /** - * 更新限流器配置 - * - * @param addr 服务器地址 - * @param name 限流器名称 - * @param speed 限速值(MB) - * @param secret 认证密钥 - * @return 请求结果 - */ - public static GostDto UpdateLimiters(String addr, Long name, String speed, String secret) { + public static GostDto UpdateLimiters(Long node_id, Long name, String speed) { JSONObject data = createLimiterData(name, speed); - String url = buildUrl(addr, LIMITERS_ENDPOINT + "/" + name); - return HttpUtils.put(url, data, secret); + JSONObject req = new JSONObject(); + req.put("limiter", name + ""); + req.put("data", data); + return WebSocketServer.send_msg(node_id, req, "UpdateLimiters"); } - /** - * 删除限流器配置 - * - * @param addr 服务器地址 - * @param name 限流器名称 - * @param secret 认证密钥 - * @return 请求结果 - */ - public static GostDto DeleteLimiters(String addr, Long name, String secret) { - String url = buildUrl(addr, LIMITERS_ENDPOINT + "/" + name); - return HttpUtils.delete(url, secret); + public static GostDto DeleteLimiters(Long node_id, Long name) { + JSONObject req = new JSONObject(); + req.put("limiter", name + ""); + return WebSocketServer.send_msg(node_id, req, "DeleteLimiters"); } - /** - * 创建限流器数据 - */ - private static JSONObject createLimiterData(Long name, String speed) { - JSONObject data = new JSONObject(); - data.put("name", name.toString()); - JSONArray limits = new JSONArray(); - limits.add("$ " + speed + "MB " + speed + "MB"); - data.put("limits", limits); - return data; - } - - - /** - * 添加服务配置(支持端口转发和隧道转发) - * - * @param addr 服务器地址 - * @param name 服务名称 - * @param in_port 监听端口 - * @param limiter 限流器ID - * @param remoteAddr 远程地址(端口转发时使用) - * @param secret 认证密钥 - * @param fow_type 转发类型:1=端口转发,2=隧道转发 - * @return 请求结果 - */ - public static GostDto AddService(String addr, String name, Integer in_port, Integer limiter, String remoteAddr, String secret, Integer fow_type) { + public static GostDto AddService(Long node_id, String name, Integer in_port, Integer limiter, String remoteAddr, Integer fow_type, Tunnel tunnel) { JSONArray services = new JSONArray(); String[] protocols = {"tcp", "udp"}; for (String protocol : protocols) { - JSONObject service = createServiceConfig(name, in_port, limiter, remoteAddr, protocol, fow_type); + JSONObject service = createServiceConfig(name, in_port, limiter, remoteAddr, protocol, fow_type, tunnel); services.add(service); } - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch"); - return HttpUtils.post(url, services, secret); + return WebSocketServer.send_msg(node_id, services, "AddService"); } - /** - * 更新服务配置(批量更新TCP和UDP服务) - * - * @param addr 服务器地址 - * @param name 服务名称 - * @param in_port 监听端口 - * @param limiter 限流器ID - * @param remoteAddr 远程地址(端口转发时使用) - * @param secret 认证密钥 - * @param fow_type 转发类型:1=端口转发,2=隧道转发 - * @return 请求结果 - */ - public static GostDto UpdateService(String addr, String name, Integer in_port, Integer limiter, String remoteAddr, String secret, Integer fow_type) { + public static GostDto UpdateService(Long node_id, String name, Integer in_port, Integer limiter, String remoteAddr, Integer fow_type, Tunnel tunnel) { JSONArray services = new JSONArray(); String[] protocols = {"tcp", "udp"}; for (String protocol : protocols) { - JSONObject service = createServiceConfig(name, in_port, limiter, remoteAddr, protocol, fow_type); + JSONObject service = createServiceConfig(name, in_port, limiter, remoteAddr, protocol, fow_type, tunnel); services.add(service); } - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch"); - return HttpUtils.put(url, services, secret); + return WebSocketServer.send_msg(node_id, services, "UpdateService"); } - /** - * 删除服务配置(批量删除TCP和UDP服务) - * - * @param addr 服务器地址 - * @param name 服务名称 - * @param secret 认证密钥 - * @return 请求结果 - */ - public static GostDto DeleteService(String addr, String name, String secret) { + public static GostDto DeleteService(Long node_id, String name) { JSONObject data = new JSONObject(); JSONArray services = new JSONArray(); services.add(name + "_tcp"); services.add(name + "_udp"); data.put("services", services); - - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch"); - return HttpUtils.delete(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "DeleteService"); } - - public static GostDto PauseService(String addr, String name, String secret) { + public static GostDto PauseService(Long node_id, String name) { JSONObject data = new JSONObject(); JSONArray services = new JSONArray(); services.add(name + "_tcp"); services.add(name + "_udp"); data.put("services", services); - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/pause"); - return HttpUtils.post(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "PauseService"); } - public static GostDto ResumeService(String addr, String name, String secret) { + public static GostDto ResumeService(Long node_id, String name) { JSONObject data = new JSONObject(); JSONArray services = new JSONArray(); services.add(name + "_tcp"); services.add(name + "_udp"); data.put("services", services); - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/resume"); - return HttpUtils.post(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "ResumeService"); } - public static GostDto PauseRemoteService(String addr, String name, String secret) { + public static GostDto PauseRemoteService(Long node_id, String name) { JSONObject data = new JSONObject(); JSONArray services = new JSONArray(); services.add(name + "_tls"); data.put("services", services); - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/pause"); - return HttpUtils.post(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "PauseRemoteService"); } - public static GostDto ResumeRemoteService(String addr, String name, String secret) { + public static GostDto ResumeRemoteService(Long node_id, String name) { JSONObject data = new JSONObject(); JSONArray services = new JSONArray(); services.add(name + "_tls"); data.put("services", services); - String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/resume"); - return HttpUtils.post(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "ResumeRemoteService"); } - public static GostDto AddChains(String addr, String name, String remoteAddr, String secret, String protocol) { + public static GostDto AddChains(Long node_id, String name, String remoteAddr, String protocol) { JSONObject dialer = new JSONObject(); dialer.put("type", protocol); @@ -234,11 +121,10 @@ public class GostUtil { data.put("name", name + "_chains"); data.put("hops", hops); - String url = buildUrl(addr, CHAINS_ENDPOINT); - return HttpUtils.post(url, data, secret); + return WebSocketServer.send_msg(node_id, data, "AddChains"); } - public static GostDto UpdateChains(String addr, String name, String remoteAddr, String secret, String protocol) { + public static GostDto UpdateChains(Long node_id, String name, String remoteAddr, String protocol) { JSONObject dialer = new JSONObject(); dialer.put("type", protocol); @@ -264,18 +150,19 @@ public class GostUtil { JSONObject data = new JSONObject(); data.put("name", name + "_chains"); data.put("hops", hops); - - String url = buildUrl(addr, CHAINS_ENDPOINT + "/" + name + "_chains"); - return HttpUtils.put(url, data, secret); + JSONObject req = new JSONObject(); + req.put("chain", name + "_chains"); + req.put("data", data); + return WebSocketServer.send_msg(node_id, req, "UpdateChains"); } - - public static GostDto DeleteChains(String addr, String name, String secret) { - String url = buildUrl(addr, CHAINS_ENDPOINT + "/" + name + "_chains"); - return HttpUtils.delete(url, secret); + public static GostDto DeleteChains(Long node_id, String name) { + JSONObject data = new JSONObject(); + data.put("chain", name + "_chains"); + return WebSocketServer.send_msg(node_id, data, "DeleteChains"); } - public static GostDto AddRemoteService(String addr, String name, Integer out_port, String remoteAddr, String secret, String protocol) { + public static GostDto AddRemoteService(Long node_id, String name, Integer out_port, String remoteAddr, String protocol) { JSONObject data = new JSONObject(); data.put("name", name + "_tls"); data.put("addr", ":" + out_port); @@ -293,12 +180,12 @@ public class GostUtil { nodes.add(node); forwarder.put("nodes", nodes); data.put("forwarder", forwarder); - String url = buildUrl(addr, SERVICES_ENDPOINT); - return HttpUtils.post(url, data, secret); + JSONArray services = new JSONArray(); + services.add(data); + return WebSocketServer.send_msg(node_id, services, "AddService"); } - - public static GostDto UpdateRemoteService(String addr, String name, Integer out_port, String remoteAddr, String secret) { + public static GostDto UpdateRemoteService(Long node_id, String name, Integer out_port, String remoteAddr) { JSONObject data = new JSONObject(); data.put("name", name + "_tls"); data.put("addr", ":" + out_port); @@ -316,24 +203,36 @@ public class GostUtil { nodes.add(node); forwarder.put("nodes", nodes); data.put("forwarder", forwarder); - String url = buildUrl(addr, SERVICES_ENDPOINT + "/" + name + "_tls"); - return HttpUtils.put(url, data, secret); + JSONArray services = new JSONArray(); + services.add(data); + return WebSocketServer.send_msg(node_id, services, "UpdateService"); } - - public static GostDto DeleteRemoteService(String addr, String name, String secret) { - String url = buildUrl(addr, SERVICES_ENDPOINT + "/" + name + "_tls"); - return HttpUtils.delete(url, secret); + public static GostDto DeleteRemoteService(Long node_id, String name) { + JSONArray data = new JSONArray(); + data.add(name + "_tls"); + JSONObject req = new JSONObject(); + req.put("services", data); + return WebSocketServer.send_msg(node_id, req, "DeleteService"); } + private static JSONObject createLimiterData(Long name, String speed) { + JSONObject data = new JSONObject(); + data.put("name", name.toString()); + JSONArray limits = new JSONArray(); + limits.add("$ " + speed + "MB " + speed + "MB"); + data.put("limits", limits); + return data; + } - /** - * 创建单个服务配置 - */ - private static JSONObject createServiceConfig(String name, Integer in_port, Integer limiter, String remoteAddr, String protocol, Integer fow_type) { + private static JSONObject createServiceConfig(String name, Integer in_port, Integer limiter, String remoteAddr, String protocol, Integer fow_type, Tunnel tunnel) { JSONObject service = new JSONObject(); service.put("name", name + "_" + protocol); - service.put("addr", ":" + in_port); + if (Objects.equals(protocol, "tcp")){ + service.put("addr", tunnel.getTcpListenAddr() + ":" + in_port); + }else { + service.put("addr", tunnel.getUdpListenAddr() + ":" + in_port); + } // 添加限流器配置 if (limiter != null) { @@ -357,9 +256,6 @@ public class GostUtil { return service; } - /** - * 创建处理器配置 - */ private static JSONObject createHandler(String protocol, String name, Integer fow_type) { JSONObject handler = new JSONObject(); handler.put("type", protocol); @@ -372,18 +268,12 @@ public class GostUtil { return handler; } - /** - * 创建监听器配置 - */ private static JSONObject createListener(String protocol) { JSONObject listener = new JSONObject(); listener.put("type", protocol); return listener; } - /** - * 创建转发器配置 - */ private static JSONObject createForwarder(String protocol, String remoteAddr) { JSONObject forwarder = new JSONObject(); JSONArray nodes = new JSONArray(); @@ -395,33 +285,12 @@ public class GostUtil { return forwarder; } - /** - * 判断是否为端口转发 - */ private static boolean isPortForwarding(Integer fow_type) { return fow_type != null && fow_type == 1; } - /** - * 判断是否为隧道转发 - */ private static boolean isTunnelForwarding(Integer fow_type) { return fow_type != null && fow_type != 1; } - - /** - * 构建API URL - */ - private static String buildUrl(String addr, String endpoint) { - // 如果是IPv6地址(包含多个冒号且不包含方括号),需要用方括号包裹IP部分 - if (!addr.contains("[") && addr.indexOf(':') != addr.lastIndexOf(':')) { - // 这是IPv6地址,找到最后一个冒号(端口分隔符) - int lastColonIndex = addr.lastIndexOf(':'); - String ipPart = addr.substring(0, lastColonIndex); - String portPart = addr.substring(lastColonIndex); - addr = "[" + ipPart + "]" + portPart; - } - return "https://" + addr + API_BASE_URL + endpoint; - } } diff --git a/springboot-backend/src/main/java/com/admin/common/utils/HttpUtils.java b/springboot-backend/src/main/java/com/admin/common/utils/HttpUtils.java index 77f65e9..b0036e3 100644 --- a/springboot-backend/src/main/java/com/admin/common/utils/HttpUtils.java +++ b/springboot-backend/src/main/java/com/admin/common/utils/HttpUtils.java @@ -2,24 +2,9 @@ package com.admin.common.utils; import com.admin.common.dto.GostConfigDto; import com.admin.common.dto.GostDto; -import com.admin.common.task.SaveConfigAsync; import com.admin.config.RestTemplateConfig; import com.alibaba.fastjson.JSONObject; import lombok.SneakyThrows; -import org.apache.http.HttpResponse; -import org.apache.http.NameValuePair; -import org.apache.http.client.config.RequestConfig; -import org.apache.http.client.entity.UrlEncodedFormEntity; -import org.apache.http.client.methods.CloseableHttpResponse; -import org.apache.http.client.methods.HttpGet; -import org.apache.http.client.methods.HttpPost; -import org.apache.http.client.utils.URIBuilder; -import org.apache.http.entity.ContentType; -import org.apache.http.entity.StringEntity; -import org.apache.http.impl.client.CloseableHttpClient; -import org.apache.http.impl.client.HttpClients; -import org.apache.http.message.BasicNameValuePair; -import org.apache.http.util.EntityUtils; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.context.ApplicationContext; @@ -29,7 +14,6 @@ import org.springframework.http.client.ClientHttpResponse; import org.springframework.stereotype.Component; import org.springframework.web.client.ResponseErrorHandler; import org.springframework.web.client.RestTemplate; -import org.springframework.util.StreamUtils; import java.io.IOException; import java.net.URI; @@ -41,251 +25,7 @@ import java.util.*; * 支持GET和POST请求,支持表单和JSON格式的请求体 */ @Component -public class HttpUtils implements ApplicationContextAware { - - private static final Logger logger = LoggerFactory.getLogger(HttpUtils.class); - - // 10秒超时配置 - private static final int TIMEOUT_SECONDS = 10; - private static final int TIMEOUT_MILLISECONDS = TIMEOUT_SECONDS * 1000; - - private static ApplicationContext applicationContext; - - @Override - public void setApplicationContext(ApplicationContext context) { - HttpUtils.applicationContext = context; - } - - /** - * 获取SaveConfigAsync Bean - */ - private static SaveConfigAsync getSaveConfigAsync() { - try { - return applicationContext.getBean(SaveConfigAsync.class); - } catch (Exception e) { - logger.warn("无法获取SaveConfigAsync Bean: {}", e.getMessage()); - return null; - } - } - - /** - * 从URL中提取IP和端口 - */ - private static String extractIpAndPortFromUrl(String url) { - try { - URI uri = URI.create(url); - String host = uri.getHost(); - int port = uri.getPort(); - if (port == -1) { - port = uri.getScheme().equals("https") ? 443 : 80; - } - return host + ":" + port; - } catch (Exception e) { - logger.warn("无法从URL提取IP和端口: {}", url); - return ""; - } - } - - /** - * 异步保存配置 - */ - private static void asyncSaveConfig(String url, String secret) { - try { - SaveConfigAsync saveConfigAsync = getSaveConfigAsync(); - if (saveConfigAsync != null) { - String ipAndPort = extractIpAndPortFromUrl(url); - saveConfigAsync.run(ipAndPort, secret); - } - } catch (Exception e) { - logger.warn("异步保存配置失败: {}", e.getMessage()); - } - } - - /** - * 自定义错误处理器,不抛出异常,允许获取所有状态码的响应 - */ - private static class NoOpResponseErrorHandler implements ResponseErrorHandler { - @Override - public boolean hasError(ClientHttpResponse response) throws IOException { - // 返回 false,让 RestTemplate 不认为任何状态码是错误 - // 这样就可以正常获取 4xx 和 5xx 的响应体 - return false; - } - - @Override - public void handleError(ClientHttpResponse response) throws IOException { - - } - } - - /** - * 创建带超时配置的RestTemplate - */ - @SneakyThrows - private static RestTemplate createRestTemplateWithTimeout() { - - // 创建RestTemplate - RestTemplate restTemplate = new RestTemplate(RestTemplateConfig.generateHttpRequestFactory()); - restTemplate.setErrorHandler(new NoOpResponseErrorHandler()); - - return restTemplate; - } +public class HttpUtils{ - @SneakyThrows - public static GostConfigDto get(String url, String secret) { - HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.APPLICATION_JSON); - headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON)); - String auth = secret + ":" + secret; - String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8)); - headers.set("Authorization", "Basic " + encodedAuth); - RestTemplate restTemplate = createRestTemplateWithTimeout(); - HttpEntity entity = new HttpEntity<>("", headers); - try { - ResponseEntity response = restTemplate.exchange( - url, - HttpMethod.GET, - entity, - GostConfigDto.class - ); - return response.getBody(); - } catch (Exception e) { - e.printStackTrace(); - GostConfigDto gostDto = new GostConfigDto(); - return gostDto; - } - } - - @SneakyThrows - public static GostDto post(String url, Object requestBody, String secret) { - HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.APPLICATION_JSON); - headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON)); - String auth = secret + ":" + secret; - String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8)); - headers.set("Authorization", "Basic " + encodedAuth); - HttpEntity entity = new HttpEntity<>(requestBody, headers); - RestTemplate restTemplate = createRestTemplateWithTimeout(); - try { - ResponseEntity response = restTemplate.postForEntity(url, entity, GostDto.class); - GostDto body = response.getBody(); - if (body.getMsg() != null && body.getMsg().contains("exists")) { - body.setMsg("OK"); - } - - if (!url.contains("/api/config?format=json")) { - asyncSaveConfig(url, secret); - } - - return body; - } catch (Exception e) { - e.printStackTrace(); - GostDto gostDto = new GostDto(); - gostDto.setCode(500); - gostDto.setMsg("请求失败"); - return gostDto; - } - } - - @SneakyThrows - public static GostDto put(String url, Object requestBody, String secret) { - HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.APPLICATION_JSON); - headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON)); - String auth = secret + ":" + secret; - String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8)); - headers.set("Authorization", "Basic " + encodedAuth); - HttpEntity entity = new HttpEntity<>(requestBody, headers); - RestTemplate restTemplate = createRestTemplateWithTimeout(); - try { - ResponseEntity response = restTemplate.exchange( - url, - HttpMethod.PUT, - entity, - GostDto.class - ); - GostDto body = response.getBody(); - asyncSaveConfig(url, secret); - return body; - } catch (Exception e) { - e.printStackTrace(); - GostDto gostDto = new GostDto(); - gostDto.setCode(500); - gostDto.setMsg("请求失败"); - return gostDto; - } - } - - @SneakyThrows - public static GostDto delete(String url, String secret) { - HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.APPLICATION_JSON); - headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON)); - - // Basic Auth - String auth = secret + ":" + secret; - String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8)); - headers.set("Authorization", "Basic " + encodedAuth); - - HttpEntity entity = new HttpEntity<>(headers); - RestTemplate restTemplate = createRestTemplateWithTimeout(); - - try { - ResponseEntity response = restTemplate.exchange( - url, - HttpMethod.DELETE, - entity, - GostDto.class - ); - GostDto body = response.getBody(); - if (body != null && body.getMsg() != null && body.getMsg().contains("not found")) { - body.setMsg("OK"); - } - asyncSaveConfig(url, secret); - return body; - } catch (Exception e) { - e.printStackTrace(); - GostDto gostDto = new GostDto(); - gostDto.setCode(500); - gostDto.setMsg("请求失败"); - return gostDto; - } - } - - @SneakyThrows - public static GostDto delete(String url, JSONObject data, String secret) { - HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.APPLICATION_JSON); - headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON)); - - // Basic Auth - String auth = secret + ":" + secret; - String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8)); - headers.set("Authorization", "Basic " + encodedAuth); - - HttpEntity entity = new HttpEntity<>(data, headers); - RestTemplate restTemplate = createRestTemplateWithTimeout(); - - try { - ResponseEntity response = restTemplate.exchange( - url, - HttpMethod.DELETE, - entity, - GostDto.class - ); - GostDto body = response.getBody(); - if (body != null && body.getMsg() != null && body.getMsg().contains("not found")) { - body.setMsg("OK"); - } - asyncSaveConfig(url, secret); - return body; - } catch (Exception e) { - e.printStackTrace(); - GostDto gostDto = new GostDto(); - gostDto.setCode(500); - gostDto.setMsg("请求失败"); - return gostDto; - } - } } \ No newline at end of file diff --git a/springboot-backend/src/main/java/com/admin/common/utils/JwtUtil.java b/springboot-backend/src/main/java/com/admin/common/utils/JwtUtil.java index 1ceedb1..bfdfa5f 100644 --- a/springboot-backend/src/main/java/com/admin/common/utils/JwtUtil.java +++ b/springboot-backend/src/main/java/com/admin/common/utils/JwtUtil.java @@ -62,7 +62,7 @@ public class JwtUtil { payload.put("iat", now.getTime() / 1000); // 发布时间 payload.put("exp", expireDate.getTime() / 1000); // 过期时间 payload.put("user", user.getUser()); - payload.put("name", user.getName()); + payload.put("name", user.getUser()); payload.put("role_id", user.getRoleId()); String payloadJson = JSON.toJSONString(payload); 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 c16160e..88bb323 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 @@ -1,6 +1,9 @@ package com.admin.common.utils; +import com.admin.common.dto.GostConfigDto; +import com.admin.common.dto.GostDto; +import com.admin.common.task.CheckGostConfigAsync; import com.admin.entity.Node; import com.admin.service.NodeService; import com.alibaba.fastjson.JSONObject; @@ -14,8 +17,11 @@ import org.springframework.web.socket.handler.TextWebSocketHandler; import javax.annotation.Resource; import java.util.Objects; +import java.util.concurrent.CompletableFuture; import java.util.concurrent.CopyOnWriteArraySet; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.TimeUnit; +import java.util.UUID; @Slf4j @@ -24,25 +30,62 @@ public class WebSocketServer extends TextWebSocketHandler { @Resource NodeService nodeService; - // 存储所有活跃的 WebSocket 连接 + @Resource + CheckGostConfigAsync checkGostConfigAsync; + + // 存储所有活跃的 WebSocket 连接( private static final CopyOnWriteArraySet activeSessions = new CopyOnWriteArraySet<>(); + // 存储节点ID和对应的WebSocket session映射 + private static final ConcurrentHashMap nodeSessions = new ConcurrentHashMap<>(); + // 为每个session提供锁对象,防止并发发送消息 private static final ConcurrentHashMap sessionLocks = new ConcurrentHashMap<>(); + + // 存储等待响应的请求,key为requestId,value为CompletableFuture + private static final ConcurrentHashMap> pendingRequests = new ConcurrentHashMap<>(); //接受客户端消息 @Override public void handleTextMessage(WebSocketSession session, TextMessage message) { try { if (StringUtils.isNoneBlank(message.getPayload())) { - //log.info("收到消息: {}", message.getPayload()); String id = session.getAttributes().get("id").toString(); String type = session.getAttributes().get("type").toString(); - // 先发送确认消息 - sendToUser(session, "ok"); - + if (message.getPayload().contains("memory_usage")){ + // 先发送确认消息 + sendToUser(session, "{\"type\":\"call\"}"); + } else if (message.getPayload().contains("config_report")) { + log.info("收到消息: {}", message.getPayload()); + JSONObject jsonObject = JSONObject.parseObject(message.getPayload()); + String string = jsonObject.getString("data"); + GostConfigDto gostConfigDto = JSONObject.parseObject(string, GostConfigDto.class); + checkGostConfigAsync.cleanNodeConfigs(id, gostConfigDto); + } else if (message.getPayload().contains("requestId")) { + log.info("收到消息: {}", message.getPayload()); + // 处理命令响应消息 + try { + JSONObject responseJson = JSONObject.parseObject(message.getPayload()); + String requestId = responseJson.getString("requestId"); + String responseMessage = responseJson.getString("message"); + + if (requestId != null) { + CompletableFuture future = pendingRequests.remove(requestId); + if (future != null) { + GostDto result = new GostDto(); + result.setMsg(responseMessage != null ? responseMessage : "无响应消息"); + future.complete(result); + } + } + } catch (Exception e) { + log.error("处理响应消息失败: {}", e.getMessage(), e); + } + } else { + log.info("收到消息: {}", message.getPayload()); + } + // 如果是节点类型,转发消息给其他会话 if (Objects.equals(type, "1")) { JSONObject jsonObject = new JSONObject(); @@ -70,21 +113,30 @@ public class WebSocketServer extends TextWebSocketHandler { try { String id = session.getAttributes().get("id").toString(); String type = session.getAttributes().get("type").toString(); + if (!Objects.equals(type, "1")) { + // 网页管理员连接 activeSessions.add(session); - }else { - Node byId = nodeService.getById(id); + } else { + // 客户端节点连接 + Long nodeId = Long.valueOf(id); + nodeSessions.put(nodeId, session); + + // 更新节点状态为在线 + Node byId = nodeService.getById(nodeId); if (byId != null) { byId.setStatus(1); nodeService.updateById(byId); + + // 广播节点上线状态给所有管理员 JSONObject res = new JSONObject(); res.put("id", id); res.put("type", "status"); res.put("data", 1); broadcastMessage(res.toJSONString()); } + } - log.info("WebSocket 连接建立成功 - id: {}, type: {}, 当前连接数: {}", id, type, activeSessions.size()); } catch (Exception e) { log.error("建立连接时发生异常: {}", e.getMessage(), e); @@ -100,25 +152,35 @@ public class WebSocketServer extends TextWebSocketHandler { String sessionId = session.getId(); if (!Objects.equals(type, "1")) { + // 连接关闭 activeSessions.remove(session); - }else { - Node byId = nodeService.getById(id); + } else { + // 客户端节点连接关闭 + Long nodeId = Long.valueOf(id); + nodeSessions.remove(nodeId); + + // 更新节点状态为离线 + Node byId = nodeService.getById(nodeId); if (byId != null) { byId.setStatus(0); nodeService.updateById(byId); + JSONObject res = new JSONObject(); res.put("id", id); res.put("type", "status"); res.put("data", 0); broadcastMessage(res.toJSONString()); } + } // 清理session锁对象 sessionLocks.remove(sessionId); - - log.info("WebSocket 连接关闭 - id: {}, sessionId: {}, 关闭状态: {}, 当前连接数: {}", - id, sessionId, status, activeSessions.size()); + + // 清理该节点的待处理请求 + if (Objects.equals(type, "1")) { + clearPendingRequestsForNode(Long.valueOf(id)); + } } catch (Exception e) { log.error("关闭连接时发生异常: {}", e.getMessage(), e); @@ -139,15 +201,34 @@ public class WebSocketServer extends TextWebSocketHandler { } } catch (Exception e) { log.error("发送WebSocket消息失败 [sessionId={}]: {}", sessionId, e.getMessage()); - activeSessions.remove(socketSession); - sessionLocks.remove(sessionId); + cleanupSession(socketSession); } } } else { - activeSessions.remove(socketSession); - if (socketSession != null) { - sessionLocks.remove(socketSession.getId()); - } + cleanupSession(socketSession); + } + } + + /** + * 清理失效的session,自动识别是节点session还是管理员session + */ + private static void cleanupSession(WebSocketSession session) { + if (session == null) return; + + String sessionId = session.getId(); + + // 清理session锁 + sessionLocks.remove(sessionId); + + boolean removedFromAdmin = activeSessions.remove(session); + + if (!removedFromAdmin) { + nodeSessions.entrySet().removeIf(entry -> { + if (entry.getValue() == session) { + return true; + } + return false; + }); } } @@ -157,4 +238,69 @@ public class WebSocketServer extends TextWebSocketHandler { sendToUser(session, message); } } + + /** + * 清理指定节点的待处理请求 + */ + private static void clearPendingRequestsForNode(Long nodeId) { + // 完成所有待处理的请求,设置为连接断开错误 + pendingRequests.entrySet().removeIf(entry -> { + CompletableFuture future = entry.getValue(); + if (!future.isDone()) { + GostDto errorResult = new GostDto(); + errorResult.setMsg("节点连接已断开"); + future.complete(errorResult); + } + return true; // 移除所有请求 + }); + } + + + public static GostDto send_msg(Long node_id, Object msg, String type) { + WebSocketSession nodeSession = nodeSessions.get(node_id); + + if (nodeSession == null) { + GostDto result = new GostDto(); + result.setMsg("节点不在线"); + return result; + } + + if (!nodeSession.isOpen()) { + nodeSessions.remove(node_id); + sessionLocks.remove(nodeSession.getId()); + GostDto result = new GostDto(); + result.setMsg("节点连接已断开"); + return result; + } + + // 生成唯一的请求ID + String requestId = UUID.randomUUID().toString(); + + // 创建CompletableFuture用于等待响应 + CompletableFuture future = new CompletableFuture<>(); + pendingRequests.put(requestId, future); + + try { + JSONObject data = new JSONObject(); + data.put("type", type); + data.put("data", msg); + data.put("requestId", requestId); + sendToUser(nodeSession, data.toJSONString()); + GostDto result = future.get(10, TimeUnit.SECONDS); + return result; + + } catch (Exception e) { + pendingRequests.remove(requestId); + GostDto result = new GostDto(); + if (e instanceof java.util.concurrent.TimeoutException) { + result.setMsg("等待响应超时"); + } else { + result.setMsg("发送消息失败: " + e.getMessage()); + } + log.error("发送消息到节点{}失败: {}", node_id, e.getMessage(), e); + return result; + } + } + + } diff --git a/springboot-backend/src/main/java/com/admin/config/WebSocketInterceptor.java b/springboot-backend/src/main/java/com/admin/config/WebSocketInterceptor.java index 859b426..bb79a0a 100644 --- a/springboot-backend/src/main/java/com/admin/config/WebSocketInterceptor.java +++ b/springboot-backend/src/main/java/com/admin/config/WebSocketInterceptor.java @@ -35,14 +35,10 @@ public class WebSocketInterceptor extends HttpSessionHandshakeInterceptor { String secret = serverHttpRequest.getServletRequest().getParameter("secret"); String type = serverHttpRequest.getServletRequest().getParameter("type"); if (Objects.equals(type, "1")) { - String client_ip = serverHttpRequest.getServletRequest().getParameter("client_ip"); Node node = nodeService.getOne(new QueryWrapper().eq("secret", secret)); if (node == null) return false; attributes.put("id", node.getId()); node.setStatus(1); - if (node.getIp() == null){ - node.setIp(client_ip); - } nodeService.updateById(node); }else { boolean b = JwtUtil.validateToken(secret); 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 2d82947..10a5a68 100644 --- a/springboot-backend/src/main/java/com/admin/controller/FlowController.java +++ b/springboot-backend/src/main/java/com/admin/controller/FlowController.java @@ -44,7 +44,6 @@ public class FlowController extends BaseController { // 常量定义 private static final String SUCCESS_RESPONSE = "ok"; - private static final String ERROR_RESPONSE = "err1"; private static final String DEFAULT_USER_TUNNEL_ID = "0"; private static final int FLOW_TYPE_UPLOAD_ONLY = 1; private static final int FLOW_TYPE_BIDIRECTIONAL = 2; @@ -73,7 +72,7 @@ public class FlowController extends BaseController { public String uploadFlowData(@RequestBody List flowDataList, String secret) { // 1. 验证节点权限 if (!isValidNode(secret)) { - return ERROR_RESPONSE; + return SUCCESS_RESPONSE; } List validFlowData = flowDataList; @@ -266,13 +265,13 @@ public class FlowController extends BaseController { Node node = nodeService.getNodeById(tunnel.getInNodeId()); if (node != null) { String serviceName = buildServiceName(forwardId, userId, userTunnelId); - GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret()); + GostUtil.PauseService(node.getId(), serviceName); // 隧道转发需要同时暂停远端服务 if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); if (outNode != null) { - GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); + GostUtil.PauseRemoteService(outNode.getId(), serviceName); } } } @@ -296,13 +295,13 @@ public class FlowController extends BaseController { Node node = nodeService.getNodeById(tunnel.getInNodeId()); if (node != null) { String serviceName = buildServiceName(forwardId, userId, userTunnelId); - GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret()); + GostUtil.PauseService(node.getId(), serviceName); // 隧道转发需要同时暂停远端服务 if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); if (outNode != null) { - GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); + GostUtil.PauseRemoteService(outNode.getId(), serviceName); } } } @@ -392,13 +391,13 @@ public class FlowController extends BaseController { Node node = nodeService.getNodeById(tunnel.getInNodeId()); if (node != null) { String serviceName = buildServiceName(String.valueOf(forward.getId()), userId, userTunnelId); - GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret()); + GostUtil.PauseService(node.getId(), serviceName); // 隧道转发需要同时暂停远端服务 if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); if (outNode != null) { - GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); + GostUtil.PauseRemoteService(outNode.getId(), serviceName); } } } @@ -423,13 +422,13 @@ public class FlowController extends BaseController { // 查找该转发对应的正确userTunnelId String actualUserTunnelId = findActualUserTunnelId(userId, forward.getTunnelId().toString()); String serviceName = buildServiceName(String.valueOf(forward.getId()), userId, actualUserTunnelId); - GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret()); + GostUtil.PauseService(node.getId(), serviceName); // 隧道转发需要同时暂停远端服务 if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); if (outNode != null) { - GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); + GostUtil.PauseRemoteService(outNode.getId(), serviceName); } } } diff --git a/springboot-backend/src/main/java/com/admin/entity/Node.java b/springboot-backend/src/main/java/com/admin/entity/Node.java index 5551507..702d5dd 100644 --- a/springboot-backend/src/main/java/com/admin/entity/Node.java +++ b/springboot-backend/src/main/java/com/admin/entity/Node.java @@ -24,7 +24,6 @@ public class Node extends BaseEntity { private String ip; - private Integer port; - + private String serverIp; } 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 3f8a03d..76695ae 100644 --- a/springboot-backend/src/main/java/com/admin/entity/Tunnel.java +++ b/springboot-backend/src/main/java/com/admin/entity/Tunnel.java @@ -77,4 +77,8 @@ public class Tunnel extends BaseEntity { * 协议类型(隧道转发时使用:tls、tcp、mtls) */ private String protocol; + + private String tcpListenAddr; + + private String udpListenAddr; } diff --git a/springboot-backend/src/main/java/com/admin/entity/User.java b/springboot-backend/src/main/java/com/admin/entity/User.java index 70213d4..a858205 100644 --- a/springboot-backend/src/main/java/com/admin/entity/User.java +++ b/springboot-backend/src/main/java/com/admin/entity/User.java @@ -18,8 +18,6 @@ public class User extends BaseEntity { private static final long serialVersionUID = 1L; - private String name; - private String user; private String pwd; diff --git a/springboot-backend/src/main/java/com/admin/service/impl/ForwardServiceImpl.java b/springboot-backend/src/main/java/com/admin/service/impl/ForwardServiceImpl.java index 5bd51a4..1fa5cf2 100644 --- a/springboot-backend/src/main/java/com/admin/service/impl/ForwardServiceImpl.java +++ b/springboot-backend/src/main/java/com/admin/service/impl/ForwardServiceImpl.java @@ -1,12 +1,10 @@ package com.admin.service.impl; -import cn.hutool.core.util.StrUtil; import com.admin.common.dto.ForwardDto; import com.admin.common.dto.ForwardUpdateDto; import com.admin.common.dto.ForwardWithTunnelDto; import com.admin.common.dto.GostDto; import com.admin.common.lang.R; -import com.admin.common.task.SaveConfigAsync; import com.admin.common.utils.GostUtil; import com.admin.common.utils.JwtUtil; import com.admin.entity.*; @@ -15,13 +13,12 @@ import com.admin.service.*; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl; import lombok.Data; +import lombok.extern.slf4j.Slf4j; import org.springframework.beans.BeanUtils; -import org.springframework.beans.factory.annotation.Autowired; import org.springframework.context.annotation.Lazy; import org.springframework.stereotype.Service; import javax.annotation.Resource; -import javax.swing.*; import java.util.List; import java.util.Objects; import java.util.Set; @@ -35,6 +32,7 @@ import java.util.stream.Collectors; * @author QAQ * @since 2025-06-03 */ +@Slf4j @Service public class ForwardServiceImpl extends ServiceImpl implements ForwardService { @@ -86,7 +84,7 @@ public class ForwardServiceImpl extends ServiceImpl impl } // 4. 分配端口 - PortAllocation portAllocation = allocatePorts(tunnel); + PortAllocation portAllocation = allocatePorts(tunnel, forwardDto.getInPort()); if (portAllocation.isHasError()) { return R.err(portAllocation.getErrorMessage()); } @@ -97,15 +95,22 @@ public class ForwardServiceImpl extends ServiceImpl impl return R.err("端口转发创建失败"); } - // 6. 调用Gost服务创建转发 - R gostResult = createGostServices(forward, tunnel, permissionResult.getLimiter()); + // 6. 获取所需的节点信息 + NodeInfo nodeInfo = getRequiredNodes(tunnel); + if (nodeInfo.isHasError()) { + this.removeById(forward.getId()); + return R.err(nodeInfo.getErrorMessage()); + } + + // 7. 调用Gost服务创建转发 + R gostResult = createGostServices(forward, tunnel, permissionResult.getLimiter(), + nodeInfo, permissionResult.getUserTunnel()); if (gostResult.getCode() != 0) { this.removeById(forward.getId()); return gostResult; } - return R.ok(); } @@ -155,14 +160,31 @@ public class ForwardServiceImpl extends ServiceImpl impl // 5. 更新Forward对象 Forward updatedForward = updateForwardEntity(forwardUpdateDto, existForward, tunnel); - // 6. 调用Gost服务更新转发 - R gostResult = updateGostServices(updatedForward, tunnel, - permissionResult != null ? permissionResult.getLimiter() : null); + // 6. 获取所需的节点信息 + NodeInfo nodeInfo = getRequiredNodes(tunnel); + if (nodeInfo.isHasError()) { + return R.err(nodeInfo.getErrorMessage()); + } + + // 7. 调用Gost服务更新转发 + R gostResult; + if (isTunnelChanged(existForward, forwardUpdateDto)) { + // 隧道变化时:先删除原配置,再创建新配置 + gostResult = updateGostServicesWithTunnelChange(existForward, updatedForward, tunnel, + permissionResult != null ? permissionResult.getLimiter() : null, + nodeInfo, permissionResult != null ? permissionResult.getUserTunnel() : null); + } else { + // 隧道未变化时:直接更新配置 + gostResult = updateGostServices(updatedForward, tunnel, + permissionResult != null ? permissionResult.getLimiter() : null, + nodeInfo, permissionResult != null ? permissionResult.getUserTunnel() : null); + } + if (gostResult.getCode() != 0) { return gostResult; } - - // 7. 保存更新 + updatedForward.setStatus(1); + // 8. 保存更新 boolean result = this.updateById(updatedForward); return result ? R.ok("端口转发更新成功") : R.err("端口转发更新失败"); } @@ -185,19 +207,27 @@ public class ForwardServiceImpl extends ServiceImpl impl } // 4. 权限检查(仅普通用户需要) + UserTunnel userTunnel = null; if (currentUser.getRoleId() != ADMIN_ROLE_ID) { - if (!hasUserTunnelPermission(currentUser.getUserId(), tunnel.getId().intValue())) { + userTunnel = getUserTunnel(currentUser.getUserId(), tunnel.getId().intValue()); + if (userTunnel == null) { return R.err("你没有该隧道权限"); } } - // 5. 调用Gost服务删除转发 - R gostResult = deleteGostServices(forward, tunnel); + // 5. 获取所需的节点信息 + NodeInfo nodeInfo = getRequiredNodes(tunnel); + if (nodeInfo.isHasError()) { + return R.err(nodeInfo.getErrorMessage()); + } + + // 6. 调用Gost服务删除转发 + R gostResult = deleteGostServices(forward, tunnel, nodeInfo, userTunnel); if (gostResult.getCode() != 0) { return gostResult; } - // 6. 删除转发记录 + // 7. 删除转发记录 boolean result = this.removeById(id); if (result) { // 归还用户转发条数(普通用户才需要归还) @@ -260,6 +290,7 @@ public class ForwardServiceImpl extends ServiceImpl impl } // 4. 恢复服务时需要额外检查 + UserTunnel userTunnel = null; if (targetStatus == FORWARD_STATUS_ACTIVE) { if (tunnel.getStatus() != TUNNEL_STATUS_ACTIVE) { return R.err("隧道已禁用,无法恢复服务"); @@ -271,49 +302,50 @@ public class ForwardServiceImpl extends ServiceImpl impl if (flowCheckResult.getCode() != 0) { return flowCheckResult; } + + userTunnel = getUserTunnel(currentUser.getUserId(), tunnel.getId().intValue()); + if (userTunnel == null) { + return R.err("你没有该隧道权限"); + } } } // 5. 权限检查(仅普通用户需要) - if (currentUser.getRoleId() != ADMIN_ROLE_ID) { - if (!hasUserTunnelPermission(currentUser.getUserId(), tunnel.getId().intValue())) { + if (currentUser.getRoleId() != ADMIN_ROLE_ID && userTunnel == null) { + userTunnel = getUserTunnel(currentUser.getUserId(), tunnel.getId().intValue()); + if (userTunnel == null) { return R.err("你没有该隧道权限"); } } - // 6. 调用Gost服务 - Node node = nodeService.getNodeById(tunnel.getInNodeId()); - if (node == null) { - return R.err("节点不存在"); + // 6. 获取所需的节点信息 + NodeInfo nodeInfo = getRequiredNodes(tunnel); + if (nodeInfo.isHasError()) { + return R.err(nodeInfo.getErrorMessage()); } - String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId()); + // 7. 调用Gost服务 + String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId(), userTunnel); GostDto gostResult; if ("PauseService".equals(gostMethod)) { - gostResult = GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret()); + gostResult = GostUtil.PauseService(nodeInfo.getInNode().getId(), serviceName); // 隧道转发需要同时暂停远端服务 - if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { - Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); - if (outNode != null) { - GostDto remoteResult = GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); - if (!isGostOperationSuccess(remoteResult)) { - return R.err(operation + "远端服务失败:" + remoteResult.getMsg()); - } + if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD && nodeInfo.getOutNode() != null) { + GostDto remoteResult = GostUtil.PauseRemoteService(nodeInfo.getOutNode().getId(), serviceName); + if (!isGostOperationSuccess(remoteResult)) { + return R.err(operation + "远端服务失败:" + remoteResult.getMsg()); } } } else { - gostResult = GostUtil.ResumeService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret()); + gostResult = GostUtil.ResumeService(nodeInfo.getInNode().getId(), serviceName); // 隧道转发需要同时恢复远端服务 - if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { - Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); - if (outNode != null) { - GostDto remoteResult = GostUtil.ResumeRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); - if (!isGostOperationSuccess(remoteResult)) { - return R.err(operation + "远端服务失败:" + remoteResult.getMsg()); - } + if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD && nodeInfo.getOutNode() != null) { + GostDto remoteResult = GostUtil.ResumeRemoteService(nodeInfo.getOutNode().getId(), serviceName); + if (!isGostOperationSuccess(remoteResult)) { + return R.err(operation + "远端服务失败:" + remoteResult.getMsg()); } } } @@ -322,7 +354,7 @@ public class ForwardServiceImpl extends ServiceImpl impl return R.err(operation + "服务失败:" + gostResult.getMsg()); } - // 7. 更新转发状态 + // 8. 更新转发状态 forward.setStatus(targetStatus); forward.setUpdatedTime(System.currentTimeMillis()); boolean result = this.updateById(forward); @@ -365,12 +397,32 @@ public class ForwardServiceImpl extends ServiceImpl impl return forward; } + /** + * 获取所需的节点信息 + */ + private NodeInfo getRequiredNodes(Tunnel tunnel) { + Node inNode = nodeService.getNodeById(tunnel.getInNodeId()); + if (inNode == null) { + return NodeInfo.error("入口节点不存在"); + } + + Node outNode = null; + if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { + outNode = nodeService.getNodeById(tunnel.getOutNodeId()); + if (outNode == null) { + return NodeInfo.error("出口节点不存在"); + } + } + + return NodeInfo.success(inNode, outNode); + } + /** * 检查用户权限和限制 */ private UserPermissionResult checkUserPermissions(UserInfo currentUser, Tunnel tunnel, Long excludeForwardId) { if (currentUser.getRoleId() == ADMIN_ROLE_ID) { - return UserPermissionResult.success(null); + return UserPermissionResult.success(null, null); } // 获取用户信息 @@ -404,7 +456,7 @@ public class ForwardServiceImpl extends ServiceImpl impl return UserPermissionResult.error(quotaCheckResult.getMsg()); } - return UserPermissionResult.success(userTunnel.getSpeedId()); + return UserPermissionResult.success(userTunnel.getSpeedId(), userTunnel); } /** @@ -473,15 +525,33 @@ public class ForwardServiceImpl extends ServiceImpl impl /** * 分配端口 */ - private PortAllocation allocatePorts(Tunnel tunnel) { - Integer inPort = allocateInPort(tunnel); - if (inPort == null) { - return PortAllocation.error("隧道入口端口已满,无法分配新端口"); + private PortAllocation allocatePorts(Tunnel tunnel, Integer specifiedInPort) { + return allocatePorts(tunnel, specifiedInPort, null); + } + + /** + * 分配端口 + */ + private PortAllocation allocatePorts(Tunnel tunnel, Integer specifiedInPort, Long excludeForwardId) { + Integer inPort; + + if (specifiedInPort != null) { + // 用户指定了入口端口,需要检查是否可用 + if (!isInPortAvailable(tunnel, specifiedInPort, excludeForwardId)) { + return PortAllocation.error("指定的入口端口 " + specifiedInPort + " 已被占用或不在允许范围内"); + } + inPort = specifiedInPort; + } else { + // 用户未指定端口时自动分配 + inPort = allocateInPort(tunnel, excludeForwardId); + if (inPort == null) { + return PortAllocation.error("隧道入口端口已满,无法分配新端口"); + } } Integer outPort = null; if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { - outPort = allocateOutPort(tunnel); + outPort = allocateOutPort(tunnel, excludeForwardId); if (outPort == null) { return PortAllocation.error("隧道出口端口已满,无法分配新端口"); } @@ -513,13 +583,27 @@ public class ForwardServiceImpl extends ServiceImpl impl Forward forward = new Forward(); BeanUtils.copyProperties(forwardUpdateDto, forward); - // 如果隧道ID发生变化,需要重新分配端口 - if (!existForward.getTunnelId().equals(forwardUpdateDto.getTunnelId())) { - PortAllocation portAllocation = allocatePorts(tunnel); + // 处理端口分配逻辑 + boolean tunnelChanged = !existForward.getTunnelId().equals(forwardUpdateDto.getTunnelId()); + boolean inPortChanged = forwardUpdateDto.getInPort() != null && + !Objects.equals(forwardUpdateDto.getInPort(), existForward.getInPort()); + + if (tunnelChanged || inPortChanged) { + // 隧道变化或入口端口变化时需要重新分配 + Integer specifiedInPort = forwardUpdateDto.getInPort(); + // 如果没有指定新端口但隧道未变化,保持原端口 + if (specifiedInPort == null && !tunnelChanged) { + specifiedInPort = existForward.getInPort(); + } + + PortAllocation portAllocation = allocatePorts(tunnel, specifiedInPort, forwardUpdateDto.getId()); + if (portAllocation.isHasError()) { + throw new RuntimeException(portAllocation.getErrorMessage()); + } forward.setInPort(portAllocation.getInPort()); forward.setOutPort(portAllocation.getOutPort()); } else { - // 隧道未变化,保持原端口 + // 隧道和端口都未变化,保持原端口 forward.setInPort(existForward.getInPort()); forward.setOutPort(existForward.getOutPort()); } @@ -531,27 +615,33 @@ public class ForwardServiceImpl extends ServiceImpl impl /** * 创建Gost服务 */ - private R createGostServices(Forward forward, Tunnel tunnel, Integer limiter) { - String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId()); - Node inNode = nodeService.getNodeById(tunnel.getInNodeId()); + private R createGostServices(Forward forward, Tunnel tunnel, Integer limiter, + NodeInfo nodeInfo, UserTunnel userTunnel) { + String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId(), userTunnel); // 隧道转发需要创建链和远程服务 if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { - R chainResult = createChainService(inNode, serviceName, tunnel.getOutIp(), forward.getOutPort(), tunnel.getProtocol()); + R chainResult = createChainService(nodeInfo.getInNode(), serviceName, tunnel.getOutIp(), forward.getOutPort(), tunnel.getProtocol()); if (chainResult.getCode() != 0) { + GostUtil.DeleteChains(nodeInfo.getInNode().getId(), serviceName); return chainResult; } - R remoteResult = createRemoteService(tunnel.getOutNodeId().intValue(), serviceName, forward, tunnel.getProtocol()); + R remoteResult = createRemoteService(nodeInfo.getOutNode(), serviceName, forward, tunnel.getProtocol()); if (remoteResult.getCode() != 0) { + GostUtil.DeleteChains(nodeInfo.getInNode().getId(), serviceName); + GostUtil.DeleteRemoteService(nodeInfo.getOutNode().getId(), serviceName); return remoteResult; } - } // 创建主服务 - R serviceResult = createMainService(inNode, serviceName, forward, limiter, tunnel.getType()); + R serviceResult = createMainService(nodeInfo.getInNode(), serviceName, forward, limiter, tunnel.getType(), tunnel); if (serviceResult.getCode() != 0) { + GostUtil.DeleteChains(nodeInfo.getInNode().getId(), serviceName); + if (nodeInfo.getOutNode() != null) { + GostUtil.DeleteRemoteService(nodeInfo.getOutNode().getId(), serviceName); + } return serviceResult; } return R.ok(); @@ -560,19 +650,19 @@ public class ForwardServiceImpl extends ServiceImpl impl /** * 更新Gost服务 */ - private R updateGostServices(Forward forward, Tunnel tunnel, Integer limiter) { - String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId()); - Node inNode = nodeService.getNodeById(tunnel.getInNodeId()); + private R updateGostServices(Forward forward, Tunnel tunnel, Integer limiter, + NodeInfo nodeInfo, UserTunnel userTunnel) { + String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId(), userTunnel); // 隧道转发需要更新链和远程服务 if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { - R chainResult = updateChainService(inNode, serviceName, tunnel.getOutIp(), forward.getOutPort(), tunnel.getProtocol()); + R chainResult = updateChainService(nodeInfo.getInNode(), serviceName, tunnel.getOutIp(), forward.getOutPort(), tunnel.getProtocol()); if (chainResult.getCode() != 0) { updateForwardStatusToError(forward); return chainResult; } - R remoteResult = updateRemoteService(tunnel.getOutNodeId().intValue(), serviceName, forward, tunnel.getProtocol()); + R remoteResult = updateRemoteService(nodeInfo.getOutNode(), serviceName, forward, tunnel.getProtocol()); if (remoteResult.getCode() != 0) { updateForwardStatusToError(forward); return remoteResult; @@ -580,7 +670,7 @@ public class ForwardServiceImpl extends ServiceImpl impl } // 更新主服务 - R serviceResult = updateMainService(inNode, serviceName, forward, limiter, tunnel.getType()); + R serviceResult = updateMainService(nodeInfo.getInNode(), serviceName, forward, limiter, tunnel.getType(), tunnel); if (serviceResult.getCode() != 0) { updateForwardStatusToError(forward); return serviceResult; @@ -589,30 +679,64 @@ public class ForwardServiceImpl extends ServiceImpl impl return R.ok(); } + /** + * 隧道变化时更新Gost服务:先删除原配置,再创建新配置 + */ + private R updateGostServicesWithTunnelChange(Forward existForward, Forward updatedForward, Tunnel newTunnel, + Integer limiter, NodeInfo nodeInfo, UserTunnel userTunnel) { + // 1. 获取原隧道信息 + Tunnel oldTunnel = tunnelService.getById(existForward.getTunnelId()); + if (oldTunnel == null) { + return R.err("原隧道不存在,无法删除旧配置"); + } + + // 2. 获取原隧道的节点信息 + NodeInfo oldNodeInfo = getRequiredNodes(oldTunnel); + if (oldNodeInfo.isHasError()) { + log.warn("获取原隧道{}的节点信息失败: {}", oldTunnel.getId(), oldNodeInfo.getErrorMessage()); + } else { + // 3. 删除原有的Gost服务配置 + R deleteResult = deleteGostServices(existForward, oldTunnel, oldNodeInfo, userTunnel); + if (deleteResult.getCode() != 0) { + // 删除失败时记录日志,但不影响后续创建(可能原配置已不存在) + log.warn("删除原隧道{}的Gost配置失败: {}", oldTunnel.getId(), deleteResult.getMsg()); + } + } + + // 4. 创建新的Gost服务配置 + R createResult = createGostServices(updatedForward, newTunnel, limiter, nodeInfo, userTunnel); + if (createResult.getCode() != 0) { + updateForwardStatusToError(updatedForward); + return R.err("创建新隧道配置失败: " + createResult.getMsg()); + } + + return R.ok(); + } + /** * 删除Gost服务 */ - private R deleteGostServices(Forward forward, Tunnel tunnel) { - String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId()); - Node inNode = nodeService.getNodeById(tunnel.getInNodeId()); + private R deleteGostServices(Forward forward, Tunnel tunnel, NodeInfo nodeInfo, UserTunnel userTunnel) { + String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId(), userTunnel); // 删除主服务 - GostDto serviceResult = GostUtil.DeleteService(inNode.getIp() + ":" + inNode.getPort(), serviceName, inNode.getSecret()); + GostDto serviceResult = GostUtil.DeleteService(nodeInfo.getInNode().getId(), serviceName); if (!isGostOperationSuccess(serviceResult)) { return R.err(serviceResult.getMsg()); } // 隧道转发需要删除链和远程服务 if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { - GostDto chainResult = GostUtil.DeleteChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, inNode.getSecret()); + GostDto chainResult = GostUtil.DeleteChains(nodeInfo.getInNode().getId(), serviceName); if (!isGostOperationSuccess(chainResult)) { return R.err(chainResult.getMsg()); } - Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); - GostDto remoteResult = GostUtil.DeleteRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret()); - if (!isGostOperationSuccess(remoteResult)) { - return R.err(remoteResult.getMsg()); + if (nodeInfo.getOutNode() != null) { + GostDto remoteResult = GostUtil.DeleteRemoteService(nodeInfo.getOutNode().getId(), serviceName); + if (!isGostOperationSuccess(remoteResult)) { + return R.err(remoteResult.getMsg()); + } } } @@ -624,28 +748,29 @@ public class ForwardServiceImpl extends ServiceImpl impl */ private R createChainService(Node inNode, String serviceName, String outIp, Integer outPort, String protocol) { String remoteAddr = outIp + ":" + outPort; - GostDto result = GostUtil.AddChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, remoteAddr, inNode.getSecret(), protocol); + if (outIp.contains(":")) { + remoteAddr = "[" + outIp + "]:" + outPort; + } + GostDto result = GostUtil.AddChains(inNode.getId(), serviceName, remoteAddr, protocol); return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg()); } /** * 创建远程服务 */ - private R createRemoteService(Integer outNodeId, String serviceName, Forward forward, String protocol) { - Node outNode = nodeService.getNodeById(outNodeId.longValue()); - GostDto result = GostUtil.AddRemoteService(outNode.getIp() + ":" + outNode.getPort(), + private R createRemoteService(Node outNode, String serviceName, Forward forward, String protocol) { + GostDto result = GostUtil.AddRemoteService(outNode.getId(), serviceName, forward.getOutPort(), - forward.getRemoteAddr(), outNode.getSecret(), protocol); + forward.getRemoteAddr(), protocol); return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg()); } /** * 创建主服务 */ - private R createMainService(Node inNode, String serviceName, Forward forward, Integer limiter, Integer tunnelType) { - GostDto result = GostUtil.AddService(inNode.getIp() + ":" + inNode.getPort(), serviceName, - forward.getInPort(), limiter, forward.getRemoteAddr(), - inNode.getSecret(), tunnelType); + private R createMainService(Node inNode, String serviceName, Forward forward, Integer limiter, Integer tunnelType, Tunnel tunnel) { + GostDto result = GostUtil.AddService(inNode.getId(), serviceName, + forward.getInPort(), limiter, forward.getRemoteAddr(), tunnelType, tunnel); return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg()); } @@ -653,12 +778,14 @@ public class ForwardServiceImpl extends ServiceImpl impl * 更新链服务 */ private R updateChainService(Node inNode, String serviceName, String outIp, Integer outPort, String protocol) { - // 创建新链 String remoteAddr = outIp + ":" + outPort; - GostDto createResult = GostUtil.UpdateChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, remoteAddr, inNode.getSecret(), protocol); + if (outIp.contains(":")) { + remoteAddr = "[" + outIp + "]:" + outPort; + } + GostDto createResult = GostUtil.UpdateChains(inNode.getId(), serviceName, remoteAddr, protocol); if (createResult.getMsg().contains(GOST_NOT_FOUND_MSG)) { - createResult = GostUtil.AddChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, remoteAddr, inNode.getSecret(), protocol); + createResult = GostUtil.AddChains(inNode.getId(), serviceName, remoteAddr, protocol); } return isGostOperationSuccess(createResult) ? R.ok() : R.err(createResult.getMsg()); } @@ -666,16 +793,15 @@ public class ForwardServiceImpl extends ServiceImpl impl /** * 更新远程服务 */ - private R updateRemoteService(Integer outNodeId, String serviceName, Forward forward, String protocol) { - Node outNode = nodeService.getNodeById(outNodeId.longValue()); + private R updateRemoteService(Node outNode, String serviceName, Forward forward, String protocol) { // 创建新远程服务 - GostDto createResult = GostUtil.UpdateRemoteService(outNode.getIp() + ":" + outNode.getPort(), + GostDto createResult = GostUtil.UpdateRemoteService(outNode.getId(), serviceName, forward.getOutPort(), - forward.getRemoteAddr(), outNode.getSecret()); + forward.getRemoteAddr()); if (createResult.getMsg().contains(GOST_NOT_FOUND_MSG)) { - createResult = GostUtil.AddRemoteService(outNode.getIp() + ":" + outNode.getPort(), + createResult = GostUtil.AddRemoteService(outNode.getId(), serviceName, forward.getOutPort(), - forward.getRemoteAddr(), outNode.getSecret(),protocol); + forward.getRemoteAddr(),protocol); } return isGostOperationSuccess(createResult) ? R.ok() : R.err(createResult.getMsg()); } @@ -683,15 +809,14 @@ public class ForwardServiceImpl extends ServiceImpl impl /** * 更新主服务 */ - private R updateMainService(Node inNode, String serviceName, Forward forward, Integer limiter, Integer tunnelType) { - GostDto result = GostUtil.UpdateService(inNode.getIp() + ":" + inNode.getPort(), serviceName, - forward.getInPort(), limiter, forward.getRemoteAddr(), - inNode.getSecret(), tunnelType); + private R updateMainService(Node inNode, String serviceName, Forward forward, Integer limiter, Integer tunnelType, Tunnel tunnel) { + GostDto result = GostUtil.UpdateService(inNode.getId(), serviceName, + forward.getInPort(), limiter, forward.getRemoteAddr(), tunnelType, tunnel); if (result.getMsg().contains(GOST_NOT_FOUND_MSG)) { - result = GostUtil.AddService(inNode.getIp() + ":" + inNode.getPort(), serviceName, + result = GostUtil.AddService(inNode.getId(), serviceName, forward.getInPort(), limiter, forward.getRemoteAddr(), - inNode.getSecret(), tunnelType); + tunnelType, tunnel); } return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg()); @@ -705,13 +830,6 @@ public class ForwardServiceImpl extends ServiceImpl impl this.updateById(forward); } - /** - * 检查是否有用户隧道权限 - */ - private boolean hasUserTunnelPermission(Integer userId, Integer tunnelId) { - return getUserTunnel(userId, tunnelId) != null; - } - /** * 获取用户隧道关系 */ @@ -750,17 +868,67 @@ public class ForwardServiceImpl extends ServiceImpl impl } /** - * 为隧道分配一个可用的入口端口 + * 检查指定的入口端口是否可用 */ - private Integer allocateInPort(Tunnel tunnel) { + private boolean isInPortAvailable(Tunnel tunnel, Integer port) { + return isInPortAvailable(tunnel, port, null); + } + + /** + * 检查指定的入口端口是否可用(可排除指定的转发ID) + */ + private boolean isInPortAvailable(Tunnel tunnel, Integer port, Long excludeForwardId) { + // 检查端口是否在隧道允许的范围内 + if (port < tunnel.getInPortSta() || port > tunnel.getInPortEnd()) { + return false; + } + // 获取所有使用相同入口节点的隧道 List tunnelsWithSameInNode = tunnelService.list(new QueryWrapper().eq("in_node_id", tunnel.getInNodeId())); Set tunnelIds = tunnelsWithSameInNode.stream() .map(Tunnel::getId) .collect(Collectors.toSet()); - // 获取这些隧道的所有转发已使用的入口端口 - List usedForwards = this.list(new QueryWrapper().in("tunnel_id", tunnelIds)); + // 获取这些隧道的所有转发已使用的入口端口(排除指定的转发ID) + QueryWrapper queryWrapper = new QueryWrapper().in("tunnel_id", tunnelIds); + if (excludeForwardId != null) { + queryWrapper.ne("id", excludeForwardId); + } + + List usedForwards = this.list(queryWrapper); + Set usedInPorts = usedForwards.stream() + .map(Forward::getInPort) + .filter(portNum -> portNum != null) + .collect(Collectors.toSet()); + + // 检查端口是否已被占用 + return !usedInPorts.contains(port); + } + + /** + * 为隧道分配一个可用的入口端口 + */ + private Integer allocateInPort(Tunnel tunnel) { + return allocateInPort(tunnel, null); + } + + /** + * 为隧道分配一个可用的入口端口(可排除指定的转发ID) + */ + private Integer allocateInPort(Tunnel tunnel, Long excludeForwardId) { + // 获取所有使用相同入口节点的隧道 + List tunnelsWithSameInNode = tunnelService.list(new QueryWrapper().eq("in_node_id", tunnel.getInNodeId())); + Set tunnelIds = tunnelsWithSameInNode.stream() + .map(Tunnel::getId) + .collect(Collectors.toSet()); + + // 获取这些隧道的所有转发已使用的入口端口(排除指定的转发ID) + QueryWrapper queryWrapper = new QueryWrapper().in("tunnel_id", tunnelIds); + if (excludeForwardId != null) { + queryWrapper.ne("id", excludeForwardId); + } + + List usedForwards = this.list(queryWrapper); Set usedInPorts = usedForwards.stream() .map(Forward::getInPort) .filter(port -> port != null) @@ -779,14 +947,26 @@ public class ForwardServiceImpl extends ServiceImpl impl * 为隧道分配一个可用的出口端口 */ private Integer allocateOutPort(Tunnel tunnel) { + return allocateOutPort(tunnel, null); + } + + /** + * 为隧道分配一个可用的出口端口(可排除指定的转发ID) + */ + private Integer allocateOutPort(Tunnel tunnel, Long excludeForwardId) { // 获取所有使用相同出口节点的隧道 List tunnelsWithSameOutNode = tunnelService.list(new QueryWrapper().eq("out_node_id", tunnel.getOutNodeId())); Set tunnelIds = tunnelsWithSameOutNode.stream() .map(Tunnel::getId) .collect(Collectors.toSet()); - // 获取这些隧道的所有转发已使用的出口端口 - List usedForwards = this.list(new QueryWrapper().in("tunnel_id", tunnelIds)); + // 获取这些隧道的所有转发已使用的出口端口(排除指定的转发ID) + QueryWrapper queryWrapper = new QueryWrapper().in("tunnel_id", tunnelIds); + if (excludeForwardId != null) { + queryWrapper.ne("id", excludeForwardId); + } + + List usedForwards = this.list(queryWrapper); Set usedOutPorts = usedForwards.stream() .map(Forward::getOutPort) .filter(port -> port != null) @@ -802,16 +982,10 @@ public class ForwardServiceImpl extends ServiceImpl impl } /** - * 构建服务名称,确保管理员和用户操作的一致性 + * 构建服务名称,优化后减少重复查询 */ - private String buildServiceName(Long forwardId, Integer userId, Integer tunnelId) { - // 根据userId和tunnelId查询UserTunnel获取正确的user_tunnel_id - UserTunnel userTunnel = userTunnelService.getOne(new QueryWrapper() - .eq("user_id", userId) - .eq("tunnel_id", tunnelId)); - + private String buildServiceName(Long forwardId, Integer userId, Integer tunnelId, UserTunnel userTunnel) { int userTunnelId = (userTunnel != null) ? userTunnel.getId() : 0; - return forwardId + "_" + userId + "_" + userTunnelId; } @@ -825,7 +999,6 @@ public class ForwardServiceImpl extends ServiceImpl impl private final Integer userId; private final Integer roleId; private final String userName; - } /** @@ -836,19 +1009,21 @@ public class ForwardServiceImpl extends ServiceImpl impl private final boolean hasError; private final String errorMessage; private final Integer limiter; + private final UserTunnel userTunnel; - private UserPermissionResult(boolean hasError, String errorMessage, Integer limiter) { + private UserPermissionResult(boolean hasError, String errorMessage, Integer limiter, UserTunnel userTunnel) { this.hasError = hasError; this.errorMessage = errorMessage; this.limiter = limiter; + this.userTunnel = userTunnel; } - public static UserPermissionResult success(Integer limiter) { - return new UserPermissionResult(false, null, limiter); + public static UserPermissionResult success(Integer limiter, UserTunnel userTunnel) { + return new UserPermissionResult(false, null, limiter, userTunnel); } public static UserPermissionResult error(String errorMessage) { - return new UserPermissionResult(true, errorMessage, null); + return new UserPermissionResult(true, errorMessage, null, null); } } @@ -877,4 +1052,30 @@ public class ForwardServiceImpl extends ServiceImpl impl return new PortAllocation(true, errorMessage, null, null); } } + + /** + * 节点信息封装类 + */ + @Data + private static class NodeInfo { + private final boolean hasError; + private final String errorMessage; + private final Node inNode; + private final Node outNode; + + private NodeInfo(boolean hasError, String errorMessage, Node inNode, Node outNode) { + this.hasError = hasError; + this.errorMessage = errorMessage; + this.inNode = inNode; + this.outNode = outNode; + } + + public static NodeInfo success(Node inNode, Node outNode) { + return new NodeInfo(false, null, inNode, outNode); + } + + public static NodeInfo error(String errorMessage) { + return new NodeInfo(true, errorMessage, null, null); + } + } } diff --git a/springboot-backend/src/main/java/com/admin/service/impl/NodeServiceImpl.java b/springboot-backend/src/main/java/com/admin/service/impl/NodeServiceImpl.java index ebd0ac9..f366099 100644 --- a/springboot-backend/src/main/java/com/admin/service/impl/NodeServiceImpl.java +++ b/springboot-backend/src/main/java/com/admin/service/impl/NodeServiceImpl.java @@ -187,6 +187,7 @@ public class NodeServiceImpl extends ServiceImpl implements No node.setId(nodeUpdateDto.getId()); node.setName(nodeUpdateDto.getName()); node.setIp(nodeUpdateDto.getIp()); + node.setServerIp(nodeUpdateDto.getServerIp()); node.setUpdatedTime(System.currentTimeMillis()); return node; } @@ -297,15 +298,68 @@ public class NodeServiceImpl extends ServiceImpl implements No StringBuilder command = new StringBuilder(); // 第一部分:下载安装脚本 - command.append("curl -L https://raw.githubusercontent.com/bqlpfy/forward-panel/refs/heads/main/install.sh") + command.append("curl -L https://ghproxy.com/https://raw.githubusercontent.com/bqlpfy/forward-panel/refs/heads/main/install.sh") .append(" -o ./install.sh && chmod +x ./install.sh && "); + // 处理服务器地址,如果是IPv6需要添加方括号 + String processedServerAddr = processServerAddress(serverAddr); + // 第二部分:执行安装脚本(去掉-u参数) command.append("./install.sh") - .append(" -a ").append(serverAddr) // 服务器地址 - .append(" -p ").append(node.getPort()) // 节点端口 - .append(" -s ").append(node.getSecret()); // 节点密钥 + .append(" -a ").append(processedServerAddr) // 服务器地址 + .append(" -s ").append(node.getSecret()); // 节点密钥 return command.toString(); } + + /** + * 处理服务器地址,确保IPv6地址被方括号包裹 + * + * @param serverAddr 原始服务器地址,格式可能为 host:port + * @return 处理后的服务器地址 + */ + private String processServerAddress(String serverAddr) { + if (StrUtil.isBlank(serverAddr)) { + return serverAddr; + } + + // 如果已经被方括号包裹,直接返回 + if (serverAddr.startsWith("[")) { + return serverAddr; + } + + // 查找最后一个冒号,分离主机和端口 + int lastColonIndex = serverAddr.lastIndexOf(':'); + if (lastColonIndex == -1) { + // 没有端口号,直接检查是否需要包裹 + return isIPv6Address(serverAddr) ? "[" + serverAddr + "]" : serverAddr; + } + + String host = serverAddr.substring(0, lastColonIndex); + String port = serverAddr.substring(lastColonIndex); + + // 检查主机部分是否为IPv6地址 + if (isIPv6Address(host)) { + return "[" + host + "]" + port; + } + + return serverAddr; + } + + /** + * 判断是否为IPv6地址 + * + * @param address 地址字符串(不包含端口号) + * @return 是否为IPv6地址 + */ + private boolean isIPv6Address(String address) { + // IPv6地址包含多个冒号,至少2个 + if (!address.contains(":")) { + return false; + } + + // 计算冒号数量,IPv6地址至少有2个冒号 + long colonCount = address.chars().filter(ch -> ch == ':').count(); + return colonCount >= 2; + } } diff --git a/springboot-backend/src/main/java/com/admin/service/impl/SpeedLimitServiceImpl.java b/springboot-backend/src/main/java/com/admin/service/impl/SpeedLimitServiceImpl.java index b489f7b..927a781 100644 --- a/springboot-backend/src/main/java/com/admin/service/impl/SpeedLimitServiceImpl.java +++ b/springboot-backend/src/main/java/com/admin/service/impl/SpeedLimitServiceImpl.java @@ -289,10 +289,9 @@ public class SpeedLimitServiceImpl extends ServiceImpl impleme existingTunnel.setInPortSta(tunnelUpdateDto.getInPortSta()); existingTunnel.setInPortEnd(tunnelUpdateDto.getInPortEnd()); + // 更新TCP和UDP监听地址 + if (StrUtil.isNotBlank(tunnelUpdateDto.getTcpListenAddr())) { + existingTunnel.setTcpListenAddr(tunnelUpdateDto.getTcpListenAddr()); + } + if (StrUtil.isNotBlank(tunnelUpdateDto.getUdpListenAddr())) { + existingTunnel.setUdpListenAddr(tunnelUpdateDto.getUdpListenAddr()); + } + if (existingTunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { existingTunnel.setOutIpSta(tunnelUpdateDto.getOutIpSta()); existingTunnel.setOutIpEnd(tunnelUpdateDto.getOutIpEnd()); @@ -400,7 +408,7 @@ public class TunnelServiceImpl extends ServiceImpl impleme // 设置入口节点信息 tunnel.setInNodeId(tunnelDto.getInNodeId()); - tunnel.setInIp(inNode.getIp()); + tunnel.setInIp(inNode.getServerIp()); // 设置流量计算类型 tunnel.setFlow(tunnelDto.getFlow()); @@ -415,6 +423,12 @@ public class TunnelServiceImpl extends ServiceImpl impleme tunnel.setProtocol(null); } + // 设置TCP和UDP监听地址 + tunnel.setTcpListenAddr(StrUtil.isNotBlank(tunnelDto.getTcpListenAddr()) ? + tunnelDto.getTcpListenAddr() : "0.0.0.0"); + tunnel.setUdpListenAddr(StrUtil.isNotBlank(tunnelDto.getUdpListenAddr()) ? + tunnelDto.getUdpListenAddr() : "0.0.0.0"); + return tunnel; } @@ -496,7 +510,7 @@ public class TunnelServiceImpl extends ServiceImpl impleme // 设置出口参数 tunnel.setOutNodeId(tunnelDto.getOutNodeId()); - tunnel.setOutIp(outNode.getIp()); + tunnel.setOutIp(outNode.getServerIp()); return R.ok(); } @@ -651,6 +665,14 @@ public class TunnelServiceImpl extends ServiceImpl impleme TunnelListDto dto = new TunnelListDto(); dto.setId(tunnel.getId().intValue()); dto.setName(tunnel.getName()); + dto.setIp(tunnel.getInIp()); + dto.setInPortSta(tunnel.getInPortSta()); + dto.setInPortEnd(tunnel.getInPortEnd()); + dto.setOutIp(tunnel.getOutIp()); + dto.setOutIpSta(tunnel.getOutIpSta()); + dto.setOutIpEnd(tunnel.getOutIpEnd()); + dto.setType(tunnel.getType()); + dto.setProtocol(tunnel.getProtocol()); return dto; } diff --git a/springboot-backend/src/main/java/com/admin/service/impl/UserServiceImpl.java b/springboot-backend/src/main/java/com/admin/service/impl/UserServiceImpl.java index eda5e3b..283bfd6 100644 --- a/springboot-backend/src/main/java/com/admin/service/impl/UserServiceImpl.java +++ b/springboot-backend/src/main/java/com/admin/service/impl/UserServiceImpl.java @@ -573,7 +573,7 @@ public class UserServiceImpl extends ServiceImpl implements Us String serviceName = buildServiceName(forward.getId(), userId, userTunnel.getId()); // 删除主服务 - GostUtil.DeleteService(buildNodeAddress(inNode), serviceName, inNode.getSecret()); + GostUtil.DeleteService(inNode.getId(), serviceName); // 如果是隧道转发,还需要删除链和远程服务 if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) { @@ -591,8 +591,8 @@ public class UserServiceImpl extends ServiceImpl implements Us private void deleteGostTunnelForwardServices(Tunnel tunnel, String serviceName, Node inNode) { Node outNode = nodeService.getNodeById(tunnel.getOutNodeId()); if (outNode != null) { - GostUtil.DeleteChains(buildNodeAddress(inNode), serviceName, inNode.getSecret()); - GostUtil.DeleteRemoteService(buildNodeAddress(outNode), serviceName, outNode.getSecret()); + GostUtil.DeleteChains(inNode.getId(), serviceName); + GostUtil.DeleteRemoteService(outNode.getId(), serviceName); } } @@ -621,15 +621,6 @@ public class UserServiceImpl extends ServiceImpl implements Us return forwardId + "_" + userId + "_" + userTunnelId; } - /** - * 构建节点地址 - * - * @param node 节点对象 - * @return 节点地址字符串 - */ - private String buildNodeAddress(Node node) { - return node.getIp() + ":" + node.getPort(); - } /** * 删除用户隧道权限 @@ -700,7 +691,6 @@ public class UserServiceImpl extends ServiceImpl implements Us private UserPackageDto.UserInfoDto buildUserInfoDto(User user) { UserPackageDto.UserInfoDto userInfo = new UserPackageDto.UserInfoDto(); userInfo.setId(user.getId()); - userInfo.setName(user.getName()); userInfo.setUser(user.getUser()); userInfo.setStatus(user.getStatus()); userInfo.setFlow(user.getFlow()); diff --git a/springboot-backend/src/main/java/com/admin/service/impl/UserTunnelServiceImpl.java b/springboot-backend/src/main/java/com/admin/service/impl/UserTunnelServiceImpl.java index add1dbd..74b4c6a 100644 --- a/springboot-backend/src/main/java/com/admin/service/impl/UserTunnelServiceImpl.java +++ b/springboot-backend/src/main/java/com/admin/service/impl/UserTunnelServiceImpl.java @@ -368,9 +368,8 @@ public class UserTunnelServiceImpl extends ServiceImpl + + + + +
+ 允许范围: {{ selectedTunnel.inPortSta }}-{{ selectedTunnel.inPortEnd }},留空将自动分配可用端口 +
+
+ 请先选择隧道以查看端口范围,留空将自动分配可用端口 +
+
+ { + if (value !== null && value !== undefined && value !== '') { + // 检查端口号范围 + if (value < 1 || value > 65535) { + callback(new Error('端口号必须在1-65535之间')); + return; + } + + // 检查是否在隧道允许范围内 + if (this.selectedTunnel) { + if (value < this.selectedTunnel.inPortSta || value > this.selectedTunnel.inPortEnd) { + callback(new Error(`端口号必须在${this.selectedTunnel.inPortSta}-${this.selectedTunnel.inPortEnd}范围内`)); + return; + } + } + } + callback(); + }, + trigger: 'blur' + } + ], remoteAddr: [ { required: true, message: '请输入远程地址', trigger: 'blur' }, { @@ -431,7 +474,8 @@ export default { // 隧道列表加载完成后,设置表单数据并弹出对话框 this.forwardForm = { ...row, - userId: row.userId // 确保userId被正确设置 + userId: row.userId, // 确保userId被正确设置 + inPort: row.inPort || null // 设置入口端口 }; this.handleTunnelChange(row.tunnelId); this.dialogVisible = true; @@ -511,8 +555,18 @@ export default { // 隧道选择变化处理 handleTunnelChange(tunnelId) { this.selectedTunnel = this.tunnelList.find(tunnel => - tunnel.id === tunnelId || tunnel.tunnelId === tunnelId + tunnel.id === tunnelId ) || null; + + // 清空端口输入,避免与新隧道的端口范围冲突 + this.forwardForm.inPort = null; + + // 触发端口字段重新验证 + this.$nextTick(() => { + if (this.$refs.forwardForm) { + this.$refs.forwardForm.clearValidate('inPort'); + } + }); }, // 提交表单 @@ -530,6 +584,7 @@ export default { userId: this.forwardForm.userId, name: this.forwardForm.name, tunnelId: this.forwardForm.tunnelId, + inPort: this.forwardForm.inPort || null, remoteAddr: this.forwardForm.remoteAddr }; res = await updateForward(updateData); @@ -538,6 +593,7 @@ export default { const createData = { name: this.forwardForm.name, tunnelId: this.forwardForm.tunnelId, + inPort: this.forwardForm.inPort || null, remoteAddr: this.forwardForm.remoteAddr }; res = await createForward(createData); @@ -567,6 +623,7 @@ export default { userId: null, name: '', tunnelId: null, + inPort: null, remoteAddr: '' }; this.selectedTunnel = null; @@ -581,19 +638,11 @@ export default { getTunnelDisplayName(tunnel) { if (!tunnel) return '未知隧道'; - // 处理用户隧道权限列表的数据结构 - if (tunnel.tunnelId) { - const tunnelInfo = this.tunnelList.find(t => t.id === tunnel.tunnelId); - if (tunnelInfo && tunnelInfo.ip && tunnelInfo.port) { - return `${tunnelInfo.name || tunnel.tunnelId} (${tunnelInfo.ip}:${tunnelInfo.port})`; - } - return `隧道ID: ${tunnel.tunnelId}`; - } - - // 处理直接隧道数据结构 + // 处理隧道数据结构 if (tunnel.name) { - if (tunnel.ip && tunnel.port) { - return `${tunnel.name} (${tunnel.ip}:${tunnel.port})`; + // 显示隧道名称和IP信息 + if (tunnel.ip) { + return `${tunnel.name} (${tunnel.ip})`; } return tunnel.name; } diff --git a/vue-frontend/src/views/Home.vue b/vue-frontend/src/views/Home.vue index 9a02400..484cf8a 100644 --- a/vue-frontend/src/views/Home.vue +++ b/vue-frontend/src/views/Home.vue @@ -112,6 +112,9 @@ :width="isMobile ? '90%' : '400px'" :before-close="handlePasswordDialogClose"> + + + @@ -147,11 +150,17 @@ export default { passwordDialogVisible: false, passwordLoading: false, passwordForm: { + newUsername: '', currentPassword: '', newPassword: '', confirmPassword: '' }, passwordRules: { + newUsername: [ + { required: true, message: '请输入新用户名', trigger: 'blur' }, + { min: 3, message: '用户名长度至少3位', trigger: 'blur' }, + { max: 20, message: '用户名长度不能超过20位', trigger: 'blur' } + ], currentPassword: [ { required: true, message: '请输入当前密码', trigger: 'blur' }, { min: 1, message: '密码不能为空', trigger: 'blur' } @@ -224,6 +233,7 @@ export default { // 重置修改密码表单 resetPasswordForm() { this.passwordForm = { + newUsername: '', currentPassword: '', newPassword: '', confirmPassword: '' diff --git a/vue-frontend/src/views/Limit.vue b/vue-frontend/src/views/Limit.vue index 67de730..84d05ba 100644 --- a/vue-frontend/src/views/Limit.vue +++ b/vue-frontend/src/views/Limit.vue @@ -138,7 +138,7 @@ -
+
diff --git a/vue-frontend/src/views/Tunnel.vue b/vue-frontend/src/views/Tunnel.vue index f0688f9..01513d4 100644 --- a/vue-frontend/src/views/Tunnel.vue +++ b/vue-frontend/src/views/Tunnel.vue @@ -256,6 +256,30 @@ > + + + + + + + + + + + +
+ 部分专线需要指定才能转发udp +
+
+