package rpc import ( "context" "errors" "fmt" "log" "net" "sync" "time" "github.com/jinzhu/copier" geoipx "github.com/nezhahq/nezha/pkg/geoip" "github.com/nezhahq/nezha/pkg/grpcx" "github.com/nezhahq/nezha/pkg/tsdb" "github.com/nezhahq/nezha/model" pb "github.com/nezhahq/nezha/proto" "github.com/nezhahq/nezha/service/singleton" ) var _ pb.NezhaServiceServer = (*NezhaHandler)(nil) var NezhaHandlerSingleton *NezhaHandler type NezhaHandler struct { Auth *authHandler ioStreams map[string]*ioStreamContext ioStreamMutex *sync.RWMutex } func NewNezhaHandler() *NezhaHandler { return &NezhaHandler{ Auth: &authHandler{}, ioStreamMutex: new(sync.RWMutex), ioStreams: make(map[string]*ioStreamContext), } } func (s *NezhaHandler) RequestTask(stream pb.NezhaService_RequestTaskServer) error { var clientID uint64 var err error if clientID, err = s.Auth.CheckRequestTask(stream.Context()); err != nil { return err } server, _ := singleton.ServerShared.Get(clientID) server.SetTaskStream(stream) defer server.ClearTaskStreamIfCurrent(stream) // If a transfer is mid-flight for this server, the agent has just brought // up a fresh bidi stream — this is the moment to (re)deliver the // ApplyConfig task carrying the new owner's AgentSecret. Pushes from // dashboard mutation time are best-effort; this hook is the reliable // re-delivery point that closes the offline-during-transfer gap. if singleton.ServerTransferShared != nil { singleton.ServerTransferShared.OnAgentReconnect(clientID) } var result *pb.TaskResult for { result, err = stream.Recv() if err != nil { log.Printf("NEZHA>> RequestTask error: %v, clientID: %d\n", err, clientID) return err } switch result.GetType() { case model.TaskTypeCommand: // 处理上报的计划任务 cr, _ := singleton.CronShared.Get(result.GetId()) // 任务结果 ID 来自 agent,必须确认该 cron 本应派发给当前 reporter。 if singleton.CanReportCronResult(cr, server) { // 保存当前服务器状态信息 var curServer model.Server copier.Copy(&curServer, server) if cr.PushSuccessful && result.GetSuccessful() { singleton.NotificationShared.SendNotification(cr.NotificationGroupID, fmt.Sprintf("[%s] %s, %s\n%s", singleton.Localizer.T("Scheduled Task Executed Successfully"), cr.Name, server.Name, result.GetData()), "", &curServer) } if !result.GetSuccessful() { singleton.NotificationShared.SendNotification(cr.NotificationGroupID, fmt.Sprintf("[%s] %s, %s\n%s", singleton.Localizer.T("Scheduled Task Executed Failed"), cr.Name, server.Name, result.GetData()), "", &curServer) } singleton.DB.Model(cr).Updates(model.Cron{ LastExecutedAt: time.Now().Add(time.Second * -1 * time.Duration(result.GetDelay())), LastResult: result.GetSuccessful(), }) } case model.TaskTypeReportConfig: if len(server.ConfigCache) < 1 { if !result.GetSuccessful() { server.ConfigCache <- errors.New(result.Data) continue } server.ConfigCache <- result.Data } case model.TaskTypeServerTransferApply: // Authorization: TaskResult.Id is attacker-controlled. Without // the pending.ID == result.Id check below, agent A could cancel // server B's in-flight transfer by spoofing B's transfer ID — // same class of bug as commit 02129f1 in the cron path. // Successful=true here is best-effort only; the authoritative // verification is the agent's reconnect under the new secret. if singleton.ServerTransferShared == nil { continue } pending, ok := singleton.ServerTransferShared.LookupPending(clientID) if !ok || pending.ID != result.GetId() { log.Printf("NEZHA>> ServerTransferApply result ignored: clientID=%d reported transferID=%d but no matching pending transfer", clientID, result.GetId()) continue } if result.GetSuccessful() { continue } if _, err := singleton.ServerTransferShared.MarkFailed(result.GetId(), result.GetData()); err != nil { log.Printf("NEZHA>> ServerTransfer MarkFailed(%d) failed: %v", result.GetId(), err) } default: if model.IsServiceSentinelNeeded(result.GetType()) { singleton.ServiceSentinelShared.Dispatch(singleton.ReportData{ Data: result, Reporter: clientID, }) } } } } func (s *NezhaHandler) ReportSystemState(stream pb.NezhaService_ReportSystemStateServer) error { clientID, err := s.Auth.Check(stream.Context()) if err != nil { return err } var state *pb.State for { state, err = stream.Recv() if err != nil { log.Printf("NEZHA>> ReportSystemState error: %v, clientID: %d\n", err, clientID) return err } innerState := model.PB2State(state) server, ok := singleton.ServerShared.Get(clientID) if !ok || server == nil { return errors.New("server not found") } server.LastActive = time.Now() server.State = &innerState if singleton.TSDBEnabled() { maxTemp := 0.0 for _, t := range innerState.Temperatures { if t.Temperature > maxTemp { maxTemp = t.Temperature } } maxGPU := 0.0 for _, g := range innerState.GPU { if g > maxGPU { maxGPU = g } } if err := singleton.TSDBShared.WriteServerMetrics(&tsdb.ServerMetrics{ ServerID: clientID, Timestamp: time.Now(), CPU: innerState.CPU, MemUsed: innerState.MemUsed, SwapUsed: innerState.SwapUsed, DiskUsed: innerState.DiskUsed, NetInSpeed: innerState.NetInSpeed, NetOutSpeed: innerState.NetOutSpeed, NetInTransfer: innerState.NetInTransfer, NetOutTransfer: innerState.NetOutTransfer, Load1: innerState.Load1, Load5: innerState.Load5, Load15: innerState.Load15, TCPConnCount: innerState.TcpConnCount, UDPConnCount: innerState.UdpConnCount, ProcessCount: innerState.ProcessCount, Temperature: maxTemp, Uptime: innerState.Uptime, GPU: maxGPU, }); err != nil { log.Printf("NEZHA>> Failed to write server metrics to TSDB: %v", err) } } // 应对 dashboard / agent 重启的情况,如果从未记录过,先打点,等到小时时间点时入库 if server.PrevTransferInSnapshot == 0 || server.PrevTransferOutSnapshot == 0 { server.PrevTransferInSnapshot = state.NetInTransfer server.PrevTransferOutSnapshot = state.NetOutTransfer } if err = stream.Send(&pb.Receipt{Proced: true}); err != nil { return err } } } func (s *NezhaHandler) onReportSystemInfo(c context.Context, r *pb.Host) error { var clientID uint64 var err error if clientID, err = s.Auth.Check(c); err != nil { return err } host := model.PB2Host(r) server, ok := singleton.ServerShared.Get(clientID) if !ok || server == nil { return errors.New("server not found") } /** * 这里的 singleton 中的数据都是关机前的旧数据 * 当 agent 重启时,bootTime 变大,agent 端会先上报 host 信息,然后上报 state 信息 * 这时可以借助上报顺序的空档,立即记录停机前的数据并重置 Prev* 数据,并由接下来的 state 方法重新赋值 */ if !server.LastActive.IsZero() && host.BootTime > server.Host.BootTime { singleton.RecordTransferHourlyUsage(server) server.PrevTransferInSnapshot = 0 server.PrevTransferOutSnapshot = 0 } server.Host = &host return nil } func (s *NezhaHandler) ReportSystemInfo(c context.Context, r *pb.Host) (*pb.Receipt, error) { if err := s.onReportSystemInfo(c, r); err != nil { return nil, err } return &pb.Receipt{Proced: true}, nil } func (s *NezhaHandler) ReportSystemInfo2(c context.Context, r *pb.Host) (*pb.Uint64Receipt, error) { if err := s.onReportSystemInfo(c, r); err != nil { return nil, err } return &pb.Uint64Receipt{Data: singleton.DashboardBootTime}, nil } func (s *NezhaHandler) IOStream(stream pb.NezhaService_IOStreamServer) error { clientID, err := s.Auth.Check(stream.Context()) if err != nil { return err } id, err := stream.Recv() if err != nil { return err } // ff05ff05 是 Nezha 的魔数,用于标识流 ID。校验由 isValidIOStreamMagic 完成, // 历史 inline 检查曾因 && 短路放过几乎全部非魔数 payload (byte0==0xff 即通过)。 if id == nil || !isValidIOStreamMagic(id.Data) { return fmt.Errorf("invalid stream id") } streamId := string(id.Data[4:]) // agent 侧归属校验:只有 createTerminal / createFM / ServeNAT 选定的目标 server // 才能接管该 stream。漏掉这一步等同于把 terminal / fm / NAT 会话向所有合法 agent // 开放(任何获得 streamId 的 agent 都能抢答),构成 session-hijack RCE 中介。 // 这是 commit 6661d6a(user 侧归属校验)的对偶补丁。先校验后启 keepalive, // 避免未授权 agent 触发悬空 goroutine 持续向其发心跳。 if !s.IsStreamAuthorizedForAgent(streamId, clientID) { return fmt.Errorf("stream not authorized for agent") } go func() { for { if err := stream.Send(&pb.IOStreamData{Data: []byte{}}); err != nil { log.Printf("NEZHA>> IOStream keepAlive error: %v\n", err) return } time.Sleep(time.Second * 30) } }() if _, err := s.GetStream(streamId); err != nil { return err } iw := grpcx.NewIOStreamWrapper(stream) if err := s.AgentConnected(streamId, iw); err != nil { return err } iw.Wait() return nil } func (s *NezhaHandler) ReportGeoIP(c context.Context, r *pb.GeoIP) (*pb.GeoIP, error) { var clientID uint64 var err error if clientID, err = s.Auth.Check(c); err != nil { return nil, err } geoip := model.PB2GeoIP(r) use6 := r.GetUse6() if geoip.IP.IPv4Addr == "" && geoip.IP.IPv6Addr == "" { ip, _ := c.Value(model.CtxKeyRealIP{}).(string) if ip == "" { ip, _ = c.Value(model.CtxKeyConnectingIP{}).(string) } geoip.IP.IPv4Addr = ip } joinedIP := geoip.IP.Join() server, ok := singleton.ServerShared.Get(clientID) if !ok || server == nil { return nil, fmt.Errorf("server not found") } // 检查并更新DDNS if server.EnableDDNS && joinedIP != "" && (server.GeoIP == nil || server.GeoIP.IP != geoip.IP) { ipv4 := geoip.IP.IPv4Addr ipv6 := geoip.IP.IPv6Addr if err := singleton.ServerShared.UpdateDDNS(server, &model.IP{IPv4Addr: ipv4, IPv6Addr: ipv6}); err != nil { log.Printf("NEZHA>> Failed to update DDNS for server %d: %v", err, server.ID) } } // 发送IP变动通知 if server.GeoIP != nil && singleton.Conf.EnableIPChangeNotification && ((singleton.Conf.Cover == model.ConfigCoverAll && !singleton.Conf.IgnoredIPNotificationServerIDs[clientID]) || (singleton.Conf.Cover == model.ConfigCoverIgnoreAll && singleton.Conf.IgnoredIPNotificationServerIDs[clientID])) && server.GeoIP.IP.Join() != "" && joinedIP != "" && server.GeoIP.IP != geoip.IP { singleton.NotificationShared.SendNotification(singleton.Conf.IPChangeNotificationGroupID, fmt.Sprintf( "[%s] %s, %s => %s", singleton.Localizer.T("IP Changed"), server.Name, singleton.IPDesensitize(server.GeoIP.IP.Join()), singleton.IPDesensitize(joinedIP), ), "") } // 根据内置数据库查询 IP 地理位置 var ip string if geoip.IP.IPv6Addr != "" && (use6 || geoip.IP.IPv4Addr == "") { ip = geoip.IP.IPv6Addr } else { ip = geoip.IP.IPv4Addr } netIP := net.ParseIP(ip) location, err := geoipx.Lookup(netIP) if err != nil { log.Printf("NEZHA>> geoip.Lookup: %v", err) } geoip.CountryCode = location // 将地区码写入到 Host server.GeoIP = &geoip return &pb.GeoIP{Ip: nil, CountryCode: location, DashboardBootTime: singleton.DashboardBootTime}, nil }