mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-04 17:16:37 +08:00
gost通讯改为ws
This commit is contained in:
+2
-2
@@ -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
|
||||
|
||||
+5
-3
@@ -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{}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
Reference in New Issue
Block a user