diff --git a/model/notification.go b/model/notification.go index 59ebb109..635eae93 100644 --- a/model/notification.go +++ b/model/notification.go @@ -179,37 +179,43 @@ func (ns *NotificationServerBundle) replaceParamsInString(str string, message st } if ns.Server != nil { + runtime := ns.Server.RuntimeSnapshot() + if runtime.State == nil || runtime.Host == nil { + return str + } + state := runtime.State + host := runtime.Host replacements = append(replacements, "#SERVER.NAME#", mod(ns.Server.Name), "#SERVER.ID#", mod(fmt.Sprintf("%d", ns.Server.ID)), // Converted metrics - "#SERVER.CPU#", mod(ns.formatUsage(false, ns.Server.State.CPU)), - "#SERVER.MEM#", mod(ns.formatUsage(true, float64(ns.Server.State.MemUsed)/float64(ns.Server.Host.MemTotal))), - "#SERVER.SWAP#", mod(ns.formatUsage(true, float64(ns.Server.State.SwapUsed)/float64(ns.Server.Host.SwapTotal))), - "#SERVER.DISK#", mod(ns.formatUsage(true, float64(ns.Server.State.DiskUsed)/float64(ns.Server.Host.DiskTotal))), - "#SERVER.SPEEDIN#", mod(fmt.Sprintf("%s/s", ns.formatSize(ns.Server.State.NetInSpeed))), - "#SERVER.SPEEDOUT#", mod(fmt.Sprintf("%s/s", ns.formatSize(ns.Server.State.NetOutSpeed))), - "#SERVER.TRANSFERIN#", mod(ns.formatSize(ns.Server.State.NetInTransfer)), - "#SERVER.TRANSFEROUT#", mod(ns.formatSize(ns.Server.State.NetOutTransfer)), + "#SERVER.CPU#", mod(ns.formatUsage(false, state.CPU)), + "#SERVER.MEM#", mod(ns.formatUsage(true, float64(state.MemUsed)/float64(host.MemTotal))), + "#SERVER.SWAP#", mod(ns.formatUsage(true, float64(state.SwapUsed)/float64(host.SwapTotal))), + "#SERVER.DISK#", mod(ns.formatUsage(true, float64(state.DiskUsed)/float64(host.DiskTotal))), + "#SERVER.SPEEDIN#", mod(fmt.Sprintf("%s/s", ns.formatSize(state.NetInSpeed))), + "#SERVER.SPEEDOUT#", mod(fmt.Sprintf("%s/s", ns.formatSize(state.NetOutSpeed))), + "#SERVER.TRANSFERIN#", mod(ns.formatSize(state.NetInTransfer)), + "#SERVER.TRANSFEROUT#", mod(ns.formatSize(state.NetOutTransfer)), // Raw metrics - "#SERVER.CPUUSED#", mod(fmt.Sprintf("%f", ns.Server.State.CPU)), - "#SERVER.MEMUSED#", mod(fmt.Sprintf("%d", ns.Server.State.MemUsed)), - "#SERVER.SWAPUSED#", mod(fmt.Sprintf("%d", ns.Server.State.SwapUsed)), - "#SERVER.DISKUSED#", mod(fmt.Sprintf("%d", ns.Server.State.DiskUsed)), - "#SERVER.MEMTOTAL#", mod(fmt.Sprintf("%d", ns.Server.Host.MemTotal)), - "#SERVER.SWAPTOTAL#", mod(fmt.Sprintf("%d", ns.Server.Host.SwapTotal)), - "#SERVER.DISKTOTAL#", mod(fmt.Sprintf("%d", ns.Server.Host.DiskTotal)), - "#SERVER.NETINSPEED#", mod(fmt.Sprintf("%d", ns.Server.State.NetInSpeed)), - "#SERVER.NETOUTSPEED#", mod(fmt.Sprintf("%d", ns.Server.State.NetOutSpeed)), - "#SERVER.NETINTRANSFER#", mod(fmt.Sprintf("%d", ns.Server.State.NetInTransfer)), - "#SERVER.NETOUTTRANSFER#", mod(fmt.Sprintf("%d", ns.Server.State.NetOutTransfer)), - "#SERVER.LOAD1#", mod(fmt.Sprintf("%f", ns.Server.State.Load1)), - "#SERVER.LOAD5#", mod(fmt.Sprintf("%f", ns.Server.State.Load5)), - "#SERVER.LOAD15#", mod(fmt.Sprintf("%f", ns.Server.State.Load15)), - "#SERVER.TCPCONNCOUNT#", mod(fmt.Sprintf("%d", ns.Server.State.TcpConnCount)), - "#SERVER.UDPCONNCOUNT#", mod(fmt.Sprintf("%d", ns.Server.State.UdpConnCount)), + "#SERVER.CPUUSED#", mod(fmt.Sprintf("%f", state.CPU)), + "#SERVER.MEMUSED#", mod(fmt.Sprintf("%d", state.MemUsed)), + "#SERVER.SWAPUSED#", mod(fmt.Sprintf("%d", state.SwapUsed)), + "#SERVER.DISKUSED#", mod(fmt.Sprintf("%d", state.DiskUsed)), + "#SERVER.MEMTOTAL#", mod(fmt.Sprintf("%d", host.MemTotal)), + "#SERVER.SWAPTOTAL#", mod(fmt.Sprintf("%d", host.SwapTotal)), + "#SERVER.DISKTOTAL#", mod(fmt.Sprintf("%d", host.DiskTotal)), + "#SERVER.NETINSPEED#", mod(fmt.Sprintf("%d", state.NetInSpeed)), + "#SERVER.NETOUTSPEED#", mod(fmt.Sprintf("%d", state.NetOutSpeed)), + "#SERVER.NETINTRANSFER#", mod(fmt.Sprintf("%d", state.NetInTransfer)), + "#SERVER.NETOUTTRANSFER#", mod(fmt.Sprintf("%d", state.NetOutTransfer)), + "#SERVER.LOAD1#", mod(fmt.Sprintf("%f", state.Load1)), + "#SERVER.LOAD5#", mod(fmt.Sprintf("%f", state.Load5)), + "#SERVER.LOAD15#", mod(fmt.Sprintf("%f", state.Load15)), + "#SERVER.TCPCONNCOUNT#", mod(fmt.Sprintf("%d", state.TcpConnCount)), + "#SERVER.UDPCONNCOUNT#", mod(fmt.Sprintf("%d", state.UdpConnCount)), ) var ipv4, ipv6, validIP string diff --git a/model/rule.go b/model/rule.go index 81592196..e069eef5 100644 --- a/model/rule.go +++ b/model/rule.go @@ -62,73 +62,87 @@ func (u *Rule) Snapshot(cycleTransferStats *CycleTransferStats, server *Server, } var src float64 + runtime := server.RuntimeSnapshot() + if runtime.State == nil { + return false + } + state := runtime.State switch u.Type { case "cpu": - src = float64(server.State.CPU) + src = float64(state.CPU) case "gpu_max": - src = slices.Max(server.State.GPU) + src = slices.Max(state.GPU) case "memory": - src = percentage(server.State.MemUsed, server.Host.MemTotal) + if runtime.Host == nil { + return false + } + src = percentage(state.MemUsed, runtime.Host.MemTotal) case "swap": - src = percentage(server.State.SwapUsed, server.Host.SwapTotal) + if runtime.Host == nil { + return false + } + src = percentage(state.SwapUsed, runtime.Host.SwapTotal) case "disk": - src = percentage(server.State.DiskUsed, server.Host.DiskTotal) + if runtime.Host == nil { + return false + } + src = percentage(state.DiskUsed, runtime.Host.DiskTotal) case "net_in_speed": - src = float64(server.State.NetInSpeed) + src = float64(state.NetInSpeed) case "net_out_speed": - src = float64(server.State.NetOutSpeed) + src = float64(state.NetOutSpeed) case "net_all_speed": - src = float64(server.State.NetOutSpeed + server.State.NetOutSpeed) + src = float64(state.NetOutSpeed + state.NetOutSpeed) case "transfer_in": - src = float64(server.State.NetInTransfer) + src = float64(state.NetInTransfer) case "transfer_out": - src = float64(server.State.NetOutTransfer) + src = float64(state.NetOutTransfer) case "transfer_all": - src = float64(server.State.NetOutTransfer + server.State.NetInTransfer) + src = float64(state.NetOutTransfer + state.NetInTransfer) case "offline": - if server.LastActive.IsZero() { + if runtime.LastActive.IsZero() { src = 0 } else { - src = float64(server.LastActive.Unix()) + src = float64(runtime.LastActive.Unix()) } case "transfer_in_cycle": - src = float64(utils.SubUintChecked(server.State.NetInTransfer, server.PrevTransferInSnapshot)) + src = float64(utils.SubUintChecked(state.NetInTransfer, runtime.PrevTransferInSnapshot)) if u.CycleInterval != 0 { var res NResult db.Model(&Transfer{}).Select("SUM(`in`) AS n").Where("datetime(`created_at`) >= datetime(?) AND server_id = ?", u.GetTransferDurationStart().UTC(), server.ID).Scan(&res) src += float64(res.N) } case "transfer_out_cycle": - src = float64(utils.SubUintChecked(server.State.NetOutTransfer, server.PrevTransferOutSnapshot)) + src = float64(utils.SubUintChecked(state.NetOutTransfer, runtime.PrevTransferOutSnapshot)) if u.CycleInterval != 0 { var res NResult db.Model(&Transfer{}).Select("SUM(`out`) AS n").Where("datetime(`created_at`) >= datetime(?) AND server_id = ?", u.GetTransferDurationStart().UTC(), server.ID).Scan(&res) src += float64(res.N) } case "transfer_all_cycle": - src = float64(utils.SubUintChecked(server.State.NetOutTransfer, server.PrevTransferOutSnapshot) + utils.SubUintChecked(server.State.NetInTransfer, server.PrevTransferInSnapshot)) + src = float64(utils.SubUintChecked(state.NetOutTransfer, runtime.PrevTransferOutSnapshot) + utils.SubUintChecked(state.NetInTransfer, runtime.PrevTransferInSnapshot)) if u.CycleInterval != 0 { var res NResult db.Model(&Transfer{}).Select("SUM(`in`+`out`) AS n").Where("datetime(`created_at`) >= datetime(?) AND server_id = ?", u.GetTransferDurationStart().UTC(), server.ID).Scan(&res) src += float64(res.N) } case "load1": - src = server.State.Load1 + src = state.Load1 case "load5": - src = server.State.Load5 + src = state.Load5 case "load15": - src = server.State.Load15 + src = state.Load15 case "tcp_conn_count": - src = float64(server.State.TcpConnCount) + src = float64(state.TcpConnCount) case "udp_conn_count": - src = float64(server.State.UdpConnCount) + src = float64(state.UdpConnCount) case "process_count": - src = float64(server.State.ProcessCount) + src = float64(state.ProcessCount) case "temperature_max": var temp []float64 - if server.State.Temperatures != nil { - for _, tempStat := range server.State.Temperatures { + if state.Temperatures != nil { + for _, tempStat := range state.Temperatures { if tempStat.Temperature != 0 { temp = append(temp, tempStat.Temperature) } diff --git a/model/server.go b/model/server.go index a892f226..fd566199 100644 --- a/model/server.go +++ b/model/server.go @@ -15,6 +15,8 @@ import ( pb "github.com/nezhahq/nezha/proto" ) +var runtimeHolderInitMu sync.Mutex + type Server struct { Common @@ -48,6 +50,7 @@ type Server struct { // two independent mutexes, defeating the "one SendMsg goroutine per stream" // invariant grpc-go requires. taskStream atomic.Pointer[taskStreamHolder] + runtime atomic.Pointer[serverRuntimeHolder] ConfigCache chan any `gorm:"-" json:"-"` PrevTransferInSnapshot uint64 `gorm:"-" json:"-"` // 上次数据点时的入站使用量 @@ -69,6 +72,160 @@ type taskStreamHolder struct { sendMu sync.Mutex } +type serverRuntimeHolder struct { + mu sync.Mutex + canonical *Server + stream pb.NezhaService_ReportSystemStateServer + generation uint64 + state *HostState + host *Host + lastActive time.Time + prevIn uint64 + prevOut uint64 +} + +type StateStreamLease struct { + holder *serverRuntimeHolder + generation uint64 +} + +func (lease StateStreamLease) Generation() uint64 { + return lease.generation +} + +type RuntimeHandle struct { + holder *serverRuntimeHolder +} + +type HostReportResult struct { + ServerID uint64 + UUID string + Applied bool + Initial bool + Equal bool + Stale bool + Restart bool + Transfer Transfer +} + +func (s *Server) RuntimeHandle() RuntimeHandle { + runtimeHolderInitMu.Lock() + holder := s.runtime.Load() + if holder == nil { + holder = &serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), host: cloneHost(s.Host), lastActive: s.LastActive, prevIn: s.PrevTransferInSnapshot, prevOut: s.PrevTransferOutSnapshot} + s.runtime.Store(holder) + } + runtimeHolderInitMu.Unlock() + return RuntimeHandle{holder: holder} +} + +func (handle RuntimeHandle) ApplyHostReport(host *Host, createdAt time.Time, persist func(Transfer) error) (HostReportResult, error) { + if handle.holder == nil || host == nil { + return HostReportResult{}, errors.New("invalid runtime handle") + } + holder := handle.holder + holder.mu.Lock() + defer holder.mu.Unlock() + canonical := holder.canonical + if canonical == nil { + return HostReportResult{}, errors.New("runtime handle has no canonical server") + } + result := HostReportResult{ServerID: canonical.ID, UUID: canonical.UUID} + if holder.host == nil { + holder.host = cloneHost(host) + canonical.Host = cloneHost(host) + result.Applied = true + result.Initial = true + return result, nil + } + if host.BootTime < holder.host.BootTime { + result.Stale = true + return result, nil + } + if host.BootTime == holder.host.BootTime { + holder.host = cloneHost(host) + canonical.Host = cloneHost(host) + result.Applied = true + result.Equal = true + return result, nil + } + result.Restart = true + if holder.state != nil { + result.Transfer = Transfer{Common: Common{CreatedAt: createdAt}, ServerID: canonical.ID, In: holder.state.NetInTransfer - min(holder.state.NetInTransfer, holder.prevIn), Out: holder.state.NetOutTransfer - min(holder.state.NetOutTransfer, holder.prevOut)} + } + if persist != nil { + if err := persist(result.Transfer); err != nil { + return HostReportResult{}, err + } + } + holder.host = cloneHost(host) + holder.state = &HostState{} + holder.lastActive = time.Time{} + holder.prevIn, holder.prevOut = 0, 0 + canonical.Host = cloneHost(host) + canonical.State = &HostState{} + canonical.LastActive = time.Time{} + canonical.PrevTransferInSnapshot = 0 + canonical.PrevTransferOutSnapshot = 0 + result.Applied = true + return result, nil +} + +func (lease StateStreamLease) UpdateState(state *HostState, lastActive time.Time) bool { + return lease.UpdateStateWithSideEffect(state, lastActive, nil) +} + +func (lease StateStreamLease) UpdateStateWithSideEffect(state *HostState, lastActive time.Time, sideEffect func() error) bool { + return lease.updateState(nil, state, lastActive, sideEffect) +} + +func (lease StateStreamLease) updateState(receiver *Server, state *HostState, lastActive time.Time, sideEffect func() error) bool { + if lease.holder == nil { + return false + } + lease.holder.mu.Lock() + defer lease.holder.mu.Unlock() + if lease.holder.generation != lease.generation || lease.holder.stream == nil || lease.holder.canonical == nil || (receiver != nil && lease.holder.canonical != receiver) { + return false + } + canonical := lease.holder.canonical + canonical.State = cloneHostState(state) + canonical.LastActive = lastActive + lease.holder.state = cloneHostState(state) + lease.holder.lastActive = lastActive + if lease.holder.prevIn == 0 || lease.holder.prevOut == 0 { + lease.holder.prevIn = state.NetInTransfer + lease.holder.prevOut = state.NetOutTransfer + } + canonical.PrevTransferInSnapshot = lease.holder.prevIn + canonical.PrevTransferOutSnapshot = lease.holder.prevOut + if sideEffect != nil { + if err := sideEffect(); err != nil { + return false + } + } + return true +} + +func (lease StateStreamLease) Clear() bool { + return lease.clear(nil) +} + +func (lease StateStreamLease) clear(receiver *Server) bool { + if lease.holder == nil { + return false + } + lease.holder.mu.Lock() + defer lease.holder.mu.Unlock() + if lease.holder.generation != lease.generation || lease.holder.stream == nil || lease.holder.canonical == nil || (receiver != nil && lease.holder.canonical != receiver) { + return false + } + lease.holder.stream = nil + lease.holder.lastActive = time.Time{} + lease.holder.canonical.LastActive = time.Time{} + return true +} + // SetTaskStream publishes the agent's RequestTask stream so other goroutines // can deliver tasks to the agent. Pass nil to detach (e.g. on disconnect). func (s *Server) SetTaskStream(stream pb.NezhaService_RequestTaskServer) { @@ -136,6 +293,172 @@ func (s *Server) SendTask(task *pb.Task) error { return h.s.Send(task) } +// AttachStateStream returns the ownership generation used to serialize state +// writes with reconnect and disconnect cleanup. +func (s *Server) AttachStateStream(stream pb.NezhaService_ReportSystemStateServer) StateStreamLease { + if stream == nil { + return StateStreamLease{} + } + runtimeHolderInitMu.Lock() + defer runtimeHolderInitMu.Unlock() + holder := s.runtime.Load() + if holder == nil { + candidate := &serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), host: cloneHost(s.Host), lastActive: s.LastActive, prevIn: s.PrevTransferInSnapshot, prevOut: s.PrevTransferOutSnapshot} + if s.runtime.CompareAndSwap(nil, candidate) { + holder = candidate + } else { + holder = s.runtime.Load() + } + } + holder.mu.Lock() + defer holder.mu.Unlock() + holder.generation++ + holder.stream = stream + return StateStreamLease{holder: holder, generation: holder.generation} +} + +func (s *Server) UpdateStateIfCurrent(lease StateStreamLease, state *HostState, lastActive time.Time) bool { + return s.UpdateStateIfCurrentWithSideEffect(lease, state, lastActive, nil) +} + +func (s *Server) UpdateStateIfCurrentWithSideEffect(lease StateStreamLease, state *HostState, lastActive time.Time, sideEffect func() error) bool { + return lease.updateState(s, state, lastActive, sideEffect) +} + +func (s *Server) ClearStateStreamIfCurrent(lease StateStreamLease) bool { + return lease.clear(s) +} + +// RuntimeSnapshot is a deep copy of the mutable runtime state. +type RuntimeSnapshot struct { + State *HostState + Host *Host + LastActive time.Time + PrevTransferInSnapshot uint64 + PrevTransferOutSnapshot uint64 +} + +func (s *Server) RuntimeSnapshot() RuntimeSnapshot { + runtimeHolderInitMu.Lock() + holder := s.runtime.Load() + if holder == nil { + candidate := &serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), lastActive: s.LastActive, prevIn: s.PrevTransferInSnapshot, prevOut: s.PrevTransferOutSnapshot} + if s.runtime.CompareAndSwap(nil, candidate) { + holder = candidate + } else { + holder = s.runtime.Load() + } + } + runtimeHolderInitMu.Unlock() + holder.mu.Lock() + defer holder.mu.Unlock() + if holder.canonical == s { + if holder.state == nil { + holder.state = cloneHostState(s.State) + } + if holder.host == nil { + holder.host = cloneHost(s.Host) + } + } + return RuntimeSnapshot{State: cloneHostState(holder.state), Host: cloneHost(holder.host), LastActive: holder.lastActive, PrevTransferInSnapshot: holder.prevIn, PrevTransferOutSnapshot: holder.prevOut} +} + +func (s *Server) SetTransferSnapshots(inbound, outbound uint64) bool { + runtimeHolderInitMu.Lock() + holder := s.runtime.Load() + if holder == nil { + holder = &serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), lastActive: s.LastActive} + s.runtime.Store(holder) + } + runtimeHolderInitMu.Unlock() + holder.mu.Lock() + if holder.canonical != s { + holder.mu.Unlock() + return false + } + holder.prevIn = inbound + holder.prevOut = outbound + if holder.canonical != nil { + holder.canonical.PrevTransferInSnapshot = inbound + holder.canonical.PrevTransferOutSnapshot = outbound + } + holder.mu.Unlock() + return true +} + +func (s *Server) TransferSnapshotDelta() (inbound, outbound, snapshotIn, snapshotOut uint64) { + snapshot := s.RuntimeSnapshot() + if snapshot.State == nil { + return 0, 0, snapshot.PrevTransferInSnapshot, snapshot.PrevTransferOutSnapshot + } + return snapshot.State.NetInTransfer, snapshot.State.NetOutTransfer, snapshot.PrevTransferInSnapshot, snapshot.PrevTransferOutSnapshot +} + +func (s *Server) TransferDeltaAndAdvance() (inbound, outbound uint64, deltaIn, deltaOut uint64) { + runtimeHolderInitMu.Lock() + holder := s.runtime.Load() + if holder == nil { + holder = &serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), host: cloneHost(s.Host), lastActive: s.LastActive, prevIn: s.PrevTransferInSnapshot, prevOut: s.PrevTransferOutSnapshot} + s.runtime.Store(holder) + } + runtimeHolderInitMu.Unlock() + holder.mu.Lock() + defer holder.mu.Unlock() + if holder.canonical != s || holder.state == nil { + return 0, 0, 0, 0 + } + inbound, outbound = holder.state.NetInTransfer, holder.state.NetOutTransfer + deltaIn = inbound - min(inbound, holder.prevIn) + deltaOut = outbound - min(outbound, holder.prevOut) + holder.prevIn, holder.prevOut = inbound, outbound + if holder.canonical != nil { + holder.canonical.PrevTransferInSnapshot = inbound + holder.canonical.PrevTransferOutSnapshot = outbound + } + return +} + +func cloneHostState(state *HostState) *HostState { + if state == nil { + return nil + } + clone := *state + clone.GPU = slices.Clone(state.GPU) + clone.Temperatures = slices.Clone(state.Temperatures) + return &clone +} + +func cloneHost(host *Host) *Host { + if host == nil { + return nil + } + clone := *host + clone.CPU = slices.Clone(host.CPU) + clone.GPU = slices.Clone(host.GPU) + return &clone +} + +func (s *Server) SetHost(host *Host) bool { + runtimeHolderInitMu.Lock() + holder := s.runtime.Load() + if holder == nil { + holder = &serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), lastActive: s.LastActive} + s.runtime.Store(holder) + } + runtimeHolderInitMu.Unlock() + holder.mu.Lock() + if holder.canonical != s { + holder.mu.Unlock() + return false + } + holder.host = cloneHost(host) + if holder.canonical != nil { + holder.canonical.Host = cloneHost(host) + } + holder.mu.Unlock() + return true +} + // ErrTaskStreamOffline is returned by SendTask when the agent has no // published RequestTask stream. Defined here (rather than in service/rpc) // so model-layer callers can branch on it without an import cycle. @@ -146,22 +469,35 @@ func InitServer(s *Server) { s.State = &HostState{} s.GeoIP = &GeoIP{} s.ConfigCache = make(chan any, 1) + s.runtime.Store(&serverRuntimeHolder{canonical: s, state: cloneHostState(s.State), host: cloneHost(s.Host)}) } func (s *Server) CopyFromRunningServer(old *Server) { - s.Host = old.Host - s.State = old.State + runtimeHolderInitMu.Lock() + defer runtimeHolderInitMu.Unlock() s.GeoIP = old.GeoIP - s.LastActive = old.LastActive // Adopt the holder pointer verbatim so the new *Server shares the send // mutex AND the stream identity with the old *Server; constructing a fresh // holder via SetTaskStream(GetTaskStream()) would give the new object its // own mutex, letting two *Server pointers race SendMsg on the same stream // during the edit/transfer rotation window. s.adoptTaskStreamHolder(old.taskStream.Load()) + holder := old.runtime.Load() + if holder == nil { + holder = &serverRuntimeHolder{canonical: old, state: cloneHostState(old.State), host: cloneHost(old.Host), lastActive: old.LastActive, prevIn: old.PrevTransferInSnapshot, prevOut: old.PrevTransferOutSnapshot} + old.runtime.CompareAndSwap(nil, holder) + holder = old.runtime.Load() + } + holder.mu.Lock() + holder.canonical = s + s.runtime.Store(holder) + s.State = cloneHostState(holder.state) + s.Host = cloneHost(holder.host) + s.LastActive = holder.lastActive + s.PrevTransferInSnapshot = holder.prevIn + s.PrevTransferOutSnapshot = holder.prevOut + holder.mu.Unlock() s.ConfigCache = old.ConfigCache - s.PrevTransferInSnapshot = old.PrevTransferInSnapshot - s.PrevTransferOutSnapshot = old.PrevTransferOutSnapshot } func (s *Server) AfterFind(tx *gorm.DB) error { @@ -245,6 +581,8 @@ type serverWithOwner struct { // global-secret pseudo-owner and is best surfaced as such by the caller's // translation table on the frontend. func (s *Server) MarshalJSON() ([]byte, error) { + runtime := s.RuntimeSnapshot() + copy := s.RuntimeCopy(runtime) owner := &ServerOwnerInfo{ID: s.GetUserID()} if ServerOwnerLookup != nil { if info, ok := ServerOwnerLookup(owner.ID); ok { @@ -252,11 +590,40 @@ func (s *Server) MarshalJSON() ([]byte, error) { } } return json.Marshal(serverWithOwner{ - serverJSON: (*serverJSON)(s), + serverJSON: (*serverJSON)(copy), Owner: owner, }) } +func (s *Server) RuntimeCopy(runtime RuntimeSnapshot) *Server { + return &Server{ + Common: Common{ + ID: s.ID, + CreatedAt: s.CreatedAt, + UpdatedAt: s.UpdatedAt, + UserID: s.GetUserID(), + }, + Name: s.Name, + UUID: s.UUID, + Note: s.Note, + PublicNote: s.PublicNote, + DisplayIndex: s.DisplayIndex, + HideForGuest: s.HideForGuest, + EnableDDNS: s.EnableDDNS, + DDNSProfilesRaw: s.DDNSProfilesRaw, + OverrideDDNSDomainsRaw: s.OverrideDDNSDomainsRaw, + DDNSProfiles: slices.Clone(s.DDNSProfiles), + OverrideDDNSDomains: s.OverrideDDNSDomains, + Host: runtime.Host, + State: runtime.State, + GeoIP: s.GeoIP, + LastActive: runtime.LastActive, + ConfigCache: s.ConfigCache, + PrevTransferInSnapshot: runtime.PrevTransferInSnapshot, + PrevTransferOutSnapshot: runtime.PrevTransferOutSnapshot, + } +} + func (s *Server) HasPermission(ctx *gin.Context) bool { if !s.Common.HasPermission(ctx) { return false diff --git a/model/server_runtime_ownership_test.go b/model/server_runtime_ownership_test.go new file mode 100644 index 00000000..970ac7df --- /dev/null +++ b/model/server_runtime_ownership_test.go @@ -0,0 +1,369 @@ +package model + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc/metadata" + + pb "github.com/nezhahq/nezha/proto" +) + +type runtimeOwnershipStream struct{} + +func (runtimeOwnershipStream) Send(*pb.Receipt) error { return nil } +func (runtimeOwnershipStream) Recv() (*pb.State, error) { return nil, context.Canceled } +func (runtimeOwnershipStream) SetHeader(metadata.MD) error { return nil } +func (runtimeOwnershipStream) SendHeader(metadata.MD) error { return nil } +func (runtimeOwnershipStream) SetTrailer(metadata.MD) {} +func (runtimeOwnershipStream) Context() context.Context { return context.Background() } +func (runtimeOwnershipStream) SendMsg(any) error { return nil } +func (runtimeOwnershipStream) RecvMsg(any) error { return nil } + +func TestServerRuntimeOwnership_replacementAdoptsHolderBeforeFirstAttach(t *testing.T) { + old := &Server{State: &HostState{Uptime: 1}, Host: &Host{BootTime: 10}} + newServer := &Server{} + var lease StateStreamLease + started := make(chan struct{}) + var waitGroup sync.WaitGroup + waitGroup.Add(1) + go func() { + defer waitGroup.Done() + close(started) + lease = old.AttachStateStream(runtimeOwnershipStream{}) + }() + <-started + newServer.CopyFromRunningServer(old) + waitGroup.Wait() + + require.True(t, newServer.UpdateStateIfCurrent(lease, &HostState{Uptime: 2}, time.Unix(2, 0))) + snapshot := newServer.RuntimeSnapshot() + require.Equal(t, uint64(2), snapshot.State.Uptime) + require.Equal(t, time.Unix(2, 0), snapshot.LastActive) + require.False(t, old.ClearStateStreamIfCurrent(lease)) +} + +func TestServerRuntimeOwnership_oldLeaseMutatesCanonicalAfterReplacement(t *testing.T) { + old := &Server{} + InitServer(old) + lease := old.AttachStateStream(runtimeOwnershipStream{}) + newServer := &Server{} + newServer.CopyFromRunningServer(old) + + require.True(t, newServer.UpdateStateIfCurrent(lease, &HostState{Uptime: 7}, time.Unix(7, 0))) + snapshot := newServer.RuntimeSnapshot() + require.Equal(t, uint64(7), snapshot.State.Uptime) + require.Equal(t, time.Unix(7, 0), snapshot.LastActive) + require.False(t, old.ClearStateStreamIfCurrent(lease)) + require.True(t, newServer.ClearStateStreamIfCurrent(lease)) + require.True(t, newServer.RuntimeSnapshot().LastActive.IsZero()) +} + +func TestServerRuntimeOwnership_leaseMutatesCanonicalWithoutReceiver(t *testing.T) { + // Given + old := &Server{} + InitServer(old) + lease := old.AttachStateStream(runtimeOwnershipStream{}) + canonical := &Server{} + canonical.CopyFromRunningServer(old) + + // When + accepted := lease.UpdateState(&HostState{Uptime: 19}, time.Unix(19, 0)) + + // Then + require.True(t, accepted) + require.Equal(t, uint64(19), canonical.RuntimeSnapshot().State.Uptime) + require.Equal(t, time.Unix(19, 0), canonical.RuntimeSnapshot().LastActive) +} + +func TestServerRuntimeOwnership_oldReceiverMutatorsCannotChangeCanonical(t *testing.T) { + // Given + old := &Server{} + InitServer(old) + lease := old.AttachStateStream(runtimeOwnershipStream{}) + canonical := &Server{} + canonical.CopyFromRunningServer(old) + + // When + hostChanged := old.SetHost(&Host{Version: "stale"}) + snapshotChanged := old.SetTransferSnapshots(91, 92) + inbound, outbound, deltaIn, deltaOut := old.TransferDeltaAndAdvance() + + // Then + require.False(t, hostChanged) + require.False(t, snapshotChanged) + require.Equal(t, uint64(0), inbound) + require.Equal(t, uint64(0), outbound) + require.Equal(t, uint64(0), deltaIn) + require.Equal(t, uint64(0), deltaOut) + require.Empty(t, canonical.RuntimeSnapshot().Host.Version) + require.Equal(t, uint64(0), canonical.RuntimeSnapshot().PrevTransferInSnapshot) + require.Equal(t, uint64(0), canonical.RuntimeSnapshot().PrevTransferOutSnapshot) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 10, NetOutTransfer: 20}, time.Unix(20, 0))) +} + +func TestServerRuntimeOwnership_copyFallbackPreservesHost(t *testing.T) { + // Given + old := &Server{Host: &Host{Version: "fallback"}, State: &HostState{Uptime: 4}, LastActive: time.Unix(4, 0), PrevTransferInSnapshot: 5, PrevTransferOutSnapshot: 6} + canonical := &Server{} + + // When + canonical.CopyFromRunningServer(old) + + // Then + snapshot := canonical.RuntimeSnapshot() + require.Equal(t, "fallback", snapshot.Host.Version) + require.Equal(t, uint64(4), snapshot.State.Uptime) + require.Equal(t, time.Unix(4, 0), snapshot.LastActive) + require.Equal(t, uint64(5), snapshot.PrevTransferInSnapshot) + require.Equal(t, uint64(6), snapshot.PrevTransferOutSnapshot) +} + +func TestServerRuntimeSnapshot_isSafeDuringStateUpdates(t *testing.T) { + server := &Server{} + InitServer(server) + lease := server.AttachStateStream(runtimeOwnershipStream{}) + var waitGroup sync.WaitGroup + waitGroup.Add(2) + go func() { + defer waitGroup.Done() + for index := uint64(1); index <= 500; index++ { + server.UpdateStateIfCurrent(lease, &HostState{Uptime: index, GPU: []float64{float64(index)}}, time.Unix(int64(index), 0)) + } + }() + go func() { + defer waitGroup.Done() + for index := 0; index < 500; index++ { + snapshot := server.RuntimeSnapshot() + require.NotNil(t, snapshot.State) + if snapshot.State.Uptime > 0 { + require.Len(t, snapshot.State.GPU, 1) + } + } + }() + waitGroup.Wait() +} + +func TestServerRuntimeOwnership_restartHostReportUsesCurrentCanonicalOnce(t *testing.T) { + // Given + old := &Server{} + InitServer(old) + lease := old.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 140, NetOutTransfer: 90}, time.Unix(10, 0))) + require.True(t, old.SetTransferSnapshots(100, 70)) + middle := &Server{} + middle.CopyFromRunningServer(old) + current := &Server{} + current.CopyFromRunningServer(middle) + + // When + result, err := old.RuntimeHandle().ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), nil) + + // Then + require.NoError(t, err) + require.True(t, result.Applied) + require.True(t, result.Restart) + require.Equal(t, current.ID, result.ServerID) + require.Equal(t, uint64(40), result.Transfer.In) + require.Equal(t, uint64(20), result.Transfer.Out) + require.Equal(t, uint64(0), current.RuntimeSnapshot().PrevTransferInSnapshot) + require.Equal(t, uint64(0), current.RuntimeSnapshot().PrevTransferOutSnapshot) + require.Equal(t, uint64(20), current.RuntimeSnapshot().Host.BootTime) + + secondResult, secondErr := old.RuntimeHandle().ApplyHostReport(&Host{BootTime: 20}, time.Unix(21, 0), nil) + require.NoError(t, secondErr) + require.True(t, secondResult.Applied) + require.True(t, secondResult.Equal) + require.Zero(t, secondResult.Transfer) +} + +func TestServerRuntimeOwnership_hostReportPersistenceFailurePreservesRuntime(t *testing.T) { + // Given + old := &Server{} + InitServer(old) + lease := old.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 140, NetOutTransfer: 90}, time.Unix(10, 0))) + require.True(t, old.SetTransferSnapshots(100, 70)) + current := &Server{} + current.CopyFromRunningServer(old) + handle := old.RuntimeHandle() + before := current.RuntimeSnapshot() + + // When + _, err := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), func(Transfer) error { + return context.Canceled + }) + + // Then + require.ErrorIs(t, err, context.Canceled) + after := current.RuntimeSnapshot() + require.Equal(t, before.Host, after.Host) + require.Equal(t, before.State, after.State) + require.Equal(t, before.LastActive, after.LastActive) + require.Equal(t, before.PrevTransferInSnapshot, after.PrevTransferInSnapshot) + require.Equal(t, before.PrevTransferOutSnapshot, after.PrevTransferOutSnapshot) +} + +func TestServerRuntimeOwnership_hostReportRetryPersistsExactlyOnce(t *testing.T) { + // Given + server := &Server{Common: Common{ID: 41}, UUID: "server-41"} + InitServer(server) + lease := server.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 20, NetOutTransfer: 30}, time.Unix(10, 0))) + require.True(t, server.SetTransferSnapshots(5, 10)) + callbackCalls := 0 + callback := func(transfer Transfer) error { + callbackCalls++ + if callbackCalls == 1 { + return context.Canceled + } + return nil + } + handle := server.RuntimeHandle() + + // When + first, firstErr := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), callback) + second, secondErr := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), callback) + third, thirdErr := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(21, 0), callback) + + // Then + require.ErrorIs(t, firstErr, context.Canceled) + require.NoError(t, secondErr) + require.NoError(t, thirdErr) + require.Equal(t, 2, callbackCalls) + require.Equal(t, uint64(15), second.Transfer.In) + require.Equal(t, uint64(20), second.Transfer.Out) + require.True(t, third.Equal) + require.Zero(t, third.Transfer) + _ = first +} + +func TestServerRuntimeOwnership_hostReportClassifiesLowerAndEqualWithoutRestart(t *testing.T) { + // Given + server := &Server{Common: Common{ID: 42}, UUID: "server-42"} + InitServer(server) + require.True(t, server.SetHost(&Host{BootTime: 20, Version: "old"})) + lease := server.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{Uptime: 7, NetInTransfer: 30}, time.Unix(7, 0))) + require.True(t, server.SetTransferSnapshots(12, 0)) + persistCalls := 0 + persist := func(Transfer) error { persistCalls++; return nil } + handle := server.RuntimeHandle() + + // When + lower, lowerErr := handle.ApplyHostReport(&Host{BootTime: 19, Version: "stale"}, time.Unix(8, 0), persist) + equal, equalErr := handle.ApplyHostReport(&Host{BootTime: 20, Version: "new"}, time.Unix(9, 0), persist) + + // Then + require.NoError(t, lowerErr) + require.True(t, lower.Stale) + require.NoError(t, equalErr) + require.True(t, equal.Equal) + require.Zero(t, persistCalls) + snapshot := server.RuntimeSnapshot() + require.Equal(t, "new", snapshot.Host.Version) + require.Equal(t, uint64(7), snapshot.State.Uptime) + require.Equal(t, time.Unix(7, 0), snapshot.LastActive) + require.Equal(t, uint64(12), snapshot.PrevTransferInSnapshot) +} + +func TestServerRuntimeOwnership_hostReportReturnsLatestCanonicalIdentity(t *testing.T) { + // Given + old := &Server{Common: Common{ID: 11}, UUID: "old"} + InitServer(old) + middle := &Server{Common: Common{ID: 22}, UUID: "middle"} + middle.CopyFromRunningServer(old) + current := &Server{Common: Common{ID: 33}, UUID: "current"} + current.CopyFromRunningServer(middle) + + // When + result, err := old.RuntimeHandle().ApplyHostReport(&Host{BootTime: 1}, time.Unix(1, 0), nil) + + // Then + require.NoError(t, err) + require.True(t, result.Applied) + require.Equal(t, current.ID, result.ServerID) + require.Equal(t, current.UUID, result.UUID) +} + +func TestServerRuntimeOwnership_transferAndRestartDoNotDuplicateWhenTransferRunsFirst(t *testing.T) { + // Given + server := &Server{Common: Common{ID: 51}, UUID: "server-51"} + InitServer(server) + lease := server.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 100, NetOutTransfer: 200}, time.Unix(10, 0))) + require.True(t, server.SetTransferSnapshots(40, 80)) + handle := server.RuntimeHandle() + holder := handle.holder + holder.mu.Lock() + hourlyDone := make(chan struct{}) + go func() { + server.TransferDeltaAndAdvance() + close(hourlyDone) + }() + holder.mu.Unlock() + <-hourlyDone + + // When + records := 0 + result, err := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), func(Transfer) error { records++; return nil }) + + // Then + require.NoError(t, err) + require.Equal(t, 1, records) + require.Equal(t, uint64(0), result.Transfer.In) + require.Equal(t, uint64(0), result.Transfer.Out) + require.Equal(t, uint64(51), result.ServerID) + require.Equal(t, uint64(0), server.RuntimeSnapshot().PrevTransferInSnapshot) +} + +func TestServerRuntimeOwnership_restartAndTransferDoNotDuplicateWhenRestartRunsFirst(t *testing.T) { + // Given + server := &Server{Common: Common{ID: 52}, UUID: "server-52"} + InitServer(server) + lease := server.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 100, NetOutTransfer: 200}, time.Unix(10, 0))) + require.True(t, server.SetTransferSnapshots(40, 80)) + handle := server.RuntimeHandle() + result, err := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), func(Transfer) error { return nil }) + require.NoError(t, err) + + // When + inbound, outbound, deltaIn, deltaOut := server.TransferDeltaAndAdvance() + + // Then + require.Equal(t, uint64(60), result.Transfer.In) + require.Equal(t, uint64(120), result.Transfer.Out) + require.Equal(t, uint64(0), inbound) + require.Equal(t, uint64(0), outbound) + require.Equal(t, uint64(0), deltaIn) + require.Equal(t, uint64(0), deltaOut) +} + +func TestServerRuntimeOwnership_failedRestartAllowsHourlyRecordThenRetry(t *testing.T) { + // Given + server := &Server{Common: Common{ID: 53}, UUID: "server-53"} + InitServer(server) + lease := server.AttachStateStream(runtimeOwnershipStream{}) + require.True(t, lease.UpdateState(&HostState{NetInTransfer: 90, NetOutTransfer: 110}, time.Unix(10, 0))) + require.True(t, server.SetTransferSnapshots(30, 50)) + handle := server.RuntimeHandle() + _, err := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), func(Transfer) error { return context.Canceled }) + require.ErrorIs(t, err, context.Canceled) + + // When + _, _, hourlyIn, hourlyOut := server.TransferDeltaAndAdvance() + records := 0 + result, retryErr := handle.ApplyHostReport(&Host{BootTime: 20}, time.Unix(20, 0), func(Transfer) error { records++; return nil }) + + // Then + require.NoError(t, retryErr) + require.Equal(t, uint64(60), hourlyIn) + require.Equal(t, uint64(60), hourlyOut) + require.Equal(t, 1, records) + require.Equal(t, uint64(0), result.Transfer.In) + require.Equal(t, uint64(0), result.Transfer.Out) +} diff --git a/pkg/agentcompatcontract/header.go b/pkg/agentcompatcontract/header.go new file mode 100644 index 00000000..26a4e846 --- /dev/null +++ b/pkg/agentcompatcontract/header.go @@ -0,0 +1,3 @@ +package agentcompatcontract + +const IOStreamCapabilityHeader = "X-Nezha-AgentCompat-IOStream-Capability" diff --git a/pkg/agentcompatcontract/header_test.go b/pkg/agentcompatcontract/header_test.go new file mode 100644 index 00000000..21db83e7 --- /dev/null +++ b/pkg/agentcompatcontract/header_test.go @@ -0,0 +1,9 @@ +package agentcompatcontract + +import "testing" + +func TestIOStreamCapabilityHeaderUsesFrozenName(t *testing.T) { + if IOStreamCapabilityHeader != "X-Nezha-AgentCompat-IOStream-Capability" { + t.Fatalf("unexpected capability header name") + } +} diff --git a/service/singleton/server.go b/service/singleton/server.go index dcc18a26..ea758ea8 100644 --- a/service/singleton/server.go +++ b/service/singleton/server.go @@ -30,10 +30,10 @@ func NewServerClass() *ServerClass { var servers []model.Server DB.Find(&servers) - for _, s := range servers { - innerS := s - model.InitServer(&innerS) - sc.list[innerS.ID] = &innerS + for i := range servers { + innerS := &servers[i] + model.InitServer(innerS) + sc.list[innerS.ID] = innerS sc.uuidToID[innerS.UUID] = innerS.ID } sc.sortList() diff --git a/service/singleton/server_transfer.go b/service/singleton/server_transfer.go index 901f802a..fe87aef8 100644 --- a/service/singleton/server_transfer.go +++ b/service/singleton/server_transfer.go @@ -79,9 +79,9 @@ const defaultRevertDeliveryRecoveryWindow = defaultServerTransferTimeout // subscribers. All mutating operations go through methods so DB and in-memory // state stay in sync. type ServerTransferClass struct { - mu sync.RWMutex - pending map[uint64]*model.ServerTransfer - revertDeliveries map[uint64]*model.ServerTransfer + mu sync.RWMutex + pending map[uint64]*model.ServerTransfer + revertDeliveries map[uint64]*model.ServerTransfer // revertRecovery holds RevertHandshakeSecrets the dashboard has pushed // but the agent has not yet acknowledged, in the window between Cancel/ // Fail/Timeout and either the agent's reconnect (which MarkRevertDelivered @@ -179,10 +179,14 @@ var ErrAgentTooOldForTransfer = fmt.Errorf("agent build older than %s does not s // never reported) so callers can defer the decision; PushIfOnline re-checks // at push time. func agentSupportsTransfer(s *model.Server) bool { - if s == nil || s.Host == nil { + if s == nil { return true } - v := strings.TrimSpace(s.Host.Version) + runtime := s.RuntimeSnapshot() + if runtime.Host == nil { + return true + } + v := strings.TrimSpace(runtime.Host.Version) if v == "" { return true } @@ -1103,7 +1107,7 @@ func (c *ServerTransferClass) revertTransition(transferID uint64, newStatus mode // back to a possibly-deleted FromUserID (regression pinned by // TestOnUserDeleteCancelsPendingTransfersAwayFromDeletedUser). var transitionedByThisCall bool - err := DB.Transaction(func(tx *gorm.DB) error { + err := DB.Transaction(func(tx *gorm.DB) error { if err := tx.First(&t, transferID).Error; err != nil { return err }