From c3077b2922d4ade12c77684d78502df156c0ccd1 Mon Sep 17 00:00:00 2001 From: qaq <1937228092@qq.com> Date: Wed, 18 Jun 2025 13:47:11 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E6=B5=81=E9=87=8F=E7=BB=9F?= =?UTF-8?q?=E8=AE=A1=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- go-gost/main.go | 32 +- go-gost/traffic_reporter.go | 192 ++++++++ go-gost/websocket_reporter.go | 450 ++++++++++++++++++ go-gost/x/handler/forward/local/handler.go | 109 +++-- .../x/handler/forward/local/traffic_bridge.go | 48 -- .../com/admin/controller/FlowController.java | 12 +- 6 files changed, 718 insertions(+), 125 deletions(-) create mode 100644 go-gost/traffic_reporter.go create mode 100644 go-gost/websocket_reporter.go delete mode 100644 go-gost/x/handler/forward/local/traffic_bridge.go diff --git a/go-gost/main.go b/go-gost/main.go index 9f523ad..c7ffbf5 100644 --- a/go-gost/main.go +++ b/go-gost/main.go @@ -11,12 +11,10 @@ import ( "runtime" "strings" "sync" - "time" - - "github.com/go-gost/gost/traffic" "github.com/go-gost/core/logger" xlogger "github.com/go-gost/x/logger" + "github.com/go-gost/x/traffic" "github.com/judwhite/go-svc" ) @@ -115,33 +113,21 @@ func main() { os.Exit(1) } - fmt.Println("✅ 配置加载成功 - addr: %s", config.Addr) + fmt.Printf("✅ 配置加载成功 - addr: %s", config.Addr) log := xlogger.NewLogger() logger.SetDefault(log) - // 创建流量管理器 - trafficMgr := traffic.NewTrafficManager() + // 使用内存流量管理器 + trafficMgr := traffic.GetGlobalManager() defer trafficMgr.Close() - // 设置全局流量管理器(供 handler 使用) - traffic.SetGlobalTrafficManager(trafficMgr) + fmt.Println("✅ 使用内存流量管理器") + logger.Default().Info("Using memory traffic manager") - // 设置流量记录器给 handler 使用 - traffic.SetupTrafficRecorder(trafficMgr) - - // 设置实时流量记录器 - traffic.SetupRealtimeTrafficRecorder(trafficMgr) - - fmt.Println("✅ 流量管理器已初始化(使用内存存储)") - logger.Default().Info("Traffic manager initialized (using memory storage)") - - // 启动实时流量统计(每5秒收集一次) - traffic.StartRealtimeTrafficStatistics(5 * time.Second) - - traffic.SetHTTPReportURL(config.Addr, config.Secret) - traffic.StartTrafficReporter(trafficMgr) - wsReporter := traffic.StartWebSocketReporterWithConfig(config.Addr, config.Secret) + SetHTTPReportURL(config.Addr, config.Secret) + StartTrafficReporter(trafficMgr) + wsReporter := StartWebSocketReporterWithConfig(config.Addr, config.Secret) defer wsReporter.Stop() p := &program{} diff --git a/go-gost/traffic_reporter.go b/go-gost/traffic_reporter.go new file mode 100644 index 0000000..7d258fc --- /dev/null +++ b/go-gost/traffic_reporter.go @@ -0,0 +1,192 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "net/http" + "strings" + "time" + + "github.com/go-gost/x/traffic" +) + +// 全局变量存储HTTP地址 +var httpReportURL string + +// TrafficReportItem 流量报告项(压缩格式) +type TrafficReportItem struct { + N string `json:"n"` // 服务名(name缩写) + T string `json:"t"` // 连接类型:conn, cc(type缩写) + U int64 `json:"u"` // 上行流量(up缩写) + D int64 `json:"d"` // 下行流量(down缩写) +} + +// SetHTTPReportURL 设置HTTP报告地址 +func SetHTTPReportURL(addr string, secret string) { + httpReportURL = "http://" + addr + "/flow/upload?secret=" + secret +} + +// StartTrafficReporter 启动流量报告任务 +func StartTrafficReporter(trafficMgr traffic.Manager) { + // 检查是否设置了HTTP地址 + if httpReportURL == "" { + fmt.Println("❌ HTTP报告地址未设置,无法启动流量报告任务") + return + } + + ticker := time.NewTicker(1 * time.Second) + + go func() { + defer ticker.Stop() + + for range ticker.C { + ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second) + + // 先获取流量统计(不清零) + stats, err := trafficMgr.GetAllServicesStats(ctx) + if err != nil { + fmt.Printf("获取流量统计失败: %v\n", err) + cancel() + continue + } + + if len(stats) == 0 { + cancel() + continue + } + + // 构建报告数据为数组格式,分别处理不同类型的流量 + var reportItems []TrafficReportItem + totalServices := 0 + totalTraffic := int64(0) + + // 解析服务和类型,分组处理 + serviceGroups := make(map[string]map[string]map[string]int64) // service -> type -> direction -> bytes + + for serviceKey, serviceStats := range stats { + upload := serviceStats["upload"] + download := serviceStats["download"] + + // 解析服务名和类型 (format: service:type) + var serviceName, serviceType string + parts := strings.Split(serviceKey, ":") + if len(parts) >= 2 { + serviceName = parts[0] + serviceType = parts[1] + } else { + serviceName = serviceKey + serviceType = "unknown" + } + + if serviceGroups[serviceName] == nil { + serviceGroups[serviceName] = make(map[string]map[string]int64) + } + if serviceGroups[serviceName][serviceType] == nil { + serviceGroups[serviceName][serviceType] = make(map[string]int64) + } + + serviceGroups[serviceName][serviceType]["upload"] = upload + serviceGroups[serviceName][serviceType]["download"] = download + + totalTraffic += upload + download + } + + // 为每个服务的每种类型创建报告项 + for serviceName, types := range serviceGroups { + for serviceType, trafficData := range types { + // 过滤掉total类型,只保留conn和cc + if serviceType == "total" { + continue + } + + upload := trafficData["upload"] + download := trafficData["download"] + + // 只有当有流量时才报告 + if upload > 0 || download > 0 { + reportItems = append(reportItems, TrafficReportItem{ + N: serviceName, + T: serviceType, + U: upload, + D: download, + }) + totalServices++ + } + } + } + + // 只有当有流量数据时才发送 + if len(reportItems) > 0 && totalTraffic > 0 { + // 发送到HTTP接口 + success, err := sendTrafficReport(ctx, reportItems) + if err != nil { + fmt.Printf("发送流量报告失败: %v\n", err) + } else if success { + // 只有收到"ok"响应才清零流量 + err = trafficMgr.ClearAllTrafficStats(ctx) + if err != nil { + fmt.Printf("清零流量统计失败: %v\n", err) + } else { + fmt.Printf("✅ 流量报告已发送并清零: %d个记录, 总流量: %d bytes\n", + totalServices, totalTraffic) + } + } else { + fmt.Printf("⚠️ 服务器未确认(非ok响应),保留流量数据: %d个记录, 总流量: %d bytes\n", + totalServices, totalTraffic) + } + } + + cancel() + } + }() + + fmt.Printf("🚀 流量报告任务已启动 (5秒间隔),目标地址: %s\n", httpReportURL) +} + +// sendTrafficReport 发送流量报告到HTTP接口 +func sendTrafficReport(ctx context.Context, reportItems []TrafficReportItem) (bool, error) { + jsonData, err := json.Marshal(reportItems) + if err != nil { + return false, fmt.Errorf("序列化报告数据失败: %v", err) + } + + req, err := http.NewRequestWithContext(ctx, "POST", httpReportURL, bytes.NewBuffer(jsonData)) + if err != nil { + return false, fmt.Errorf("创建HTTP请求失败: %v", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("User-Agent", "GOST-Traffic-Reporter/1.0") + + client := &http.Client{ + Timeout: 5 * time.Second, + } + + resp, err := client.Do(req) + if err != nil { + return false, fmt.Errorf("发送HTTP请求失败: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return false, fmt.Errorf("HTTP响应错误: %d %s", resp.StatusCode, resp.Status) + } + + // 读取响应内容 + var responseBytes bytes.Buffer + _, err = responseBytes.ReadFrom(resp.Body) + if err != nil { + return false, fmt.Errorf("读取响应内容失败: %v", err) + } + + responseText := strings.TrimSpace(responseBytes.String()) + + // 检查响应是否为"ok" + if responseText == "ok" { + return true, nil + } else { + return false, fmt.Errorf("服务器响应: %s (期望: ok)", responseText) + } +} diff --git a/go-gost/websocket_reporter.go b/go-gost/websocket_reporter.go new file mode 100644 index 0000000..eb8e556 --- /dev/null +++ b/go-gost/websocket_reporter.go @@ -0,0 +1,450 @@ +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/handler/forward/local/handler.go b/go-gost/x/handler/forward/local/handler.go index 32f2c3d..c1b7d5d 100644 --- a/go-gost/x/handler/forward/local/handler.go +++ b/go-gost/x/handler/forward/local/handler.go @@ -4,8 +4,6 @@ import ( "bufio" "bytes" "context" - "crypto/rand" - "encoding/hex" "errors" "net" "time" @@ -14,6 +12,7 @@ import ( "github.com/go-gost/core/handler" "github.com/go-gost/core/hop" md "github.com/go-gost/core/metadata" + "github.com/go-gost/core/observer/stats" "github.com/go-gost/core/recorder" ctxvalue "github.com/go-gost/x/ctx" xnet "github.com/go-gost/x/internal/net" @@ -25,6 +24,7 @@ import ( stats_wrapper "github.com/go-gost/x/observer/stats/wrapper" xrecorder "github.com/go-gost/x/recorder" "github.com/go-gost/x/registry" + "github.com/go-gost/x/traffic" ) func init() { @@ -33,18 +33,13 @@ func init() { registry.HandlerRegistry().Register("forward", NewHandler) } -// TrafficRecorder 流量记录器接口 -type TrafficRecorder interface { - RecordTraffic(ctx context.Context, service string, upload, download int64) error -} - type forwardHandler struct { - hop hop.Hop - md metadata - options handler.Options - recorder recorder.RecorderObject - certPool tls_util.CertPool - trafficRecorder TrafficRecorder + hop hop.Hop + md metadata + options handler.Options + recorder recorder.RecorderObject + certPool tls_util.CertPool + trafficManager traffic.Manager } func NewHandler(opts ...handler.Option) handler.Handler { @@ -54,17 +49,11 @@ func NewHandler(opts ...handler.Option) handler.Handler { } return &forwardHandler{ - options: options, + options: options, + trafficManager: traffic.GetGlobalManager(), } } -// generateConnectionID 生成连接唯一标识 -func generateConnectionID() string { - bytes := make([]byte, 8) - rand.Read(bytes) - return hex.EncodeToString(bytes) -} - func (h *forwardHandler) Init(md md.Metadata) (err error) { if err = h.parseMetadata(md); err != nil { return @@ -81,9 +70,6 @@ func (h *forwardHandler) Init(md md.Metadata) (err error) { h.certPool = tls_util.NewMemoryCertPool() } - // 获取全局流量记录器 - h.trafficRecorder = GetGlobalTrafficRecorder() - return } @@ -96,7 +82,6 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand defer conn.Close() start := time.Now() - connID := generateConnectionID() ro := &xrecorder.HandlerRecorderObject{ Service: h.options.Service, @@ -123,7 +108,6 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand "local": conn.LocalAddr().String(), "sid": ro.SID, "client": ro.ClientIP, - "connID": connID, }) network := "tcp" @@ -132,15 +116,54 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand } ro.Network = network - connStats := xstats.Stats{} - ccStats := xstats.Stats{} - conn = stats_wrapper.WrapConn(conn, &connStats) + connStats := xstats.NewStats(false) // false表示不在Get时清零 + ccStats := xstats.NewStats(false) // false表示不在Get时清零 + conn = stats_wrapper.WrapConn(conn, connStats) - // 获取实时流量管理器并注册连接 - rtm := GetGlobalRealtimeTrafficManager() - if rtm != nil { - rtm.RegisterConnection(connID+":conn", h.options.Service+":conn", &connStats) - defer rtm.UnregisterConnection(connID + ":conn") + // 启动定期上报流量的goroutine + trafficCtx, trafficCancel := context.WithCancel(context.Background()) + defer trafficCancel() + + if h.trafficManager != nil { + go func() { + ticker := time.NewTicker(1 * time.Second) + defer ticker.Stop() + + // 记录上次的流量值 + var lastConnOutput, lastConnInput, lastCCOutput, lastCCInput uint64 + + for { + select { + case <-ticker.C: + // 获取当前流量统计(不清零) + connOutput := connStats.Get(stats.KindInputBytes) + connInput := connStats.Get(stats.KindOutputBytes) + ccOutput := ccStats.Get(stats.KindOutputBytes) + ccInput := ccStats.Get(stats.KindInputBytes) + + // 计算增量 + connOutputDelta := connOutput - lastConnOutput + connInputDelta := connInput - lastConnInput + ccOutputDelta := ccOutput - lastCCOutput + ccInputDelta := ccInput - lastCCInput + + if connInputDelta > 0 || connOutputDelta > 0 { + h.trafficManager.RecordTraffic(ctx, ro.Service+":conn", int64(connInputDelta), int64(connOutputDelta)) + } + if ccOutputDelta > 0 || ccInputDelta > 0 { + h.trafficManager.RecordTraffic(ctx, ro.Service+":cc", int64(ccOutputDelta), int64(ccInputDelta)) + } + + // 更新上次的值 + lastConnOutput = connOutput + lastConnInput = connInput + lastCCOutput = ccOutput + lastCCInput = ccInput + case <-trafficCtx.Done(): + return + } + } + }() } defer func() { @@ -148,7 +171,7 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand ro.Err = err.Error() } ro.Duration = time.Since(start) - + // 流量统计已经在定期上报中处理,这里不需要再次记录 }() if !h.checkRateLimit(conn.RemoteAddr()) { @@ -174,13 +197,7 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand cc, err := h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), "tcp", address) ro.Route = buf.String() - cc = stats_wrapper.WrapConn(cc, &ccStats) - - // 为目标连接也注册到实时流量统计 - if rtm != nil && err == nil { - rtm.RegisterConnection(connID+":cc", h.options.Service+":cc", &ccStats) - // 注意:这里不能defer UnregisterConnection,因为dial函数可能被多次调用 - } + cc = stats_wrapper.WrapConn(cc, ccStats) return cc, err } @@ -256,13 +273,7 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand marker.Reset() } - cc = stats_wrapper.WrapConn(cc, &ccStats) - - // 为目标连接注册到实时流量统计 - if rtm != nil { - rtm.RegisterConnection(connID+":cc", h.options.Service+":cc", &ccStats) - defer rtm.UnregisterConnection(connID + ":cc") - } + cc = stats_wrapper.WrapConn(cc, ccStats) defer cc.Close() diff --git a/go-gost/x/handler/forward/local/traffic_bridge.go b/go-gost/x/handler/forward/local/traffic_bridge.go deleted file mode 100644 index f8c0101..0000000 --- a/go-gost/x/handler/forward/local/traffic_bridge.go +++ /dev/null @@ -1,48 +0,0 @@ -package local - -import ( - "sync" - - "github.com/go-gost/core/observer/stats" -) - -// RealtimeTrafficManager 实时流量管理器接口(简化版) -type RealtimeTrafficManager interface { - RegisterConnection(id, service string, stats stats.Stats) - UnregisterConnection(id string) - GetActiveConnectionsCount() int -} - -var ( - globalTrafficRecorder TrafficRecorder - globalRealtimeTrafficManager RealtimeTrafficManager - trafficMutex sync.RWMutex -) - -// SetGlobalTrafficRecorder 设置全局流量记录器 -func SetGlobalTrafficRecorder(recorder TrafficRecorder) { - trafficMutex.Lock() - defer trafficMutex.Unlock() - globalTrafficRecorder = recorder -} - -// GetGlobalTrafficRecorder 获取全局流量记录器 -func GetGlobalTrafficRecorder() TrafficRecorder { - trafficMutex.RLock() - defer trafficMutex.RUnlock() - return globalTrafficRecorder -} - -// SetGlobalRealtimeTrafficManager 设置全局实时流量管理器 -func SetGlobalRealtimeTrafficManager(manager RealtimeTrafficManager) { - trafficMutex.Lock() - defer trafficMutex.Unlock() - globalRealtimeTrafficManager = manager -} - -// GetGlobalRealtimeTrafficManager 获取全局实时流量管理器 -func GetGlobalRealtimeTrafficManager() RealtimeTrafficManager { - trafficMutex.RLock() - defer trafficMutex.RUnlock() - return globalRealtimeTrafficManager -} 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 d2a7667..3c902fd 100644 --- a/springboot-backend/src/main/java/com/admin/controller/FlowController.java +++ b/springboot-backend/src/main/java/com/admin/controller/FlowController.java @@ -74,11 +74,13 @@ public class FlowController extends BaseController { return ERROR_RESPONSE; } - // 2. 过滤有效流量数据 - List validFlowData = filterValidFlowData(flowDataList); - if (validFlowData.isEmpty()) { - return SUCCESS_RESPONSE; - } + List validFlowData = flowDataList; +// // 2. 过滤有效流量数据 +// List validFlowData = filterValidFlowData(flowDataList); +// if (validFlowData.isEmpty()) { +// return SUCCESS_RESPONSE; +// } + // 3. 解析服务名称获取ID信息 String[] serviceIds = parseServiceName(validFlowData.get(0).getN());