diff --git a/cmd/unbounded-net-node/status_details.go b/cmd/unbounded-net-node/status_details.go index 5dc5e9a02..c04694d0c 100644 --- a/cmd/unbounded-net-node/status_details.go +++ b/cmd/unbounded-net-node/status_details.go @@ -79,6 +79,12 @@ func (s *nodeDetailState) enqueue(request *statusv1alpha1.DetailRequest, now tim if _, exists := s.replies[request.RequestID]; !exists { s.replies[request.RequestID] = &nodeDetailReply{request: *request} + time.AfterFunc(time.Until(request.Deadline), func() { + s.mu.Lock() + defer s.mu.Unlock() + + s.expireLocked(time.Now()) + }) } s.mu.Unlock() s.wake() @@ -86,6 +92,23 @@ func (s *nodeDetailState) enqueue(request *statusv1alpha1.DetailRequest, now tim return nil } +func (s *nodeDetailState) receive(ack *statusv1alpha1.NodeStatusAck) { + s.acknowledge(ack) + + if ack.DetailRequest != nil { + if err := s.enqueue(ack.DetailRequest, time.Now()); err != nil { + klog.V(2).Infof("Ignoring invalid detail command: %v", err) + } + } +} + +func (s *nodeDetailState) clear() { + s.mu.Lock() + defer s.mu.Unlock() + + clear(s.replies) +} + func (s *nodeDetailState) acknowledge(ack *statusv1alpha1.NodeStatusAck) { if ack == nil || ack.DetailRequestID == "" || ack.Status != "ok" { return diff --git a/cmd/unbounded-net-node/status_details_test.go b/cmd/unbounded-net-node/status_details_test.go index 3355f220f..538a57d2b 100644 --- a/cmd/unbounded-net-node/status_details_test.go +++ b/cmd/unbounded-net-node/status_details_test.go @@ -115,6 +115,39 @@ func TestDetailStateExpiryAndConcurrentClaims(t *testing.T) { } } +func TestDetailDeadlineExpiresDuringCollectionAndDisconnect(t *testing.T) { + state := (&nodeHealthState{}).detailState() + + req := &statusv1alpha1.DetailRequest{RequestID: "slow", Deadline: time.Now().Add(20 * time.Millisecond)} + if err := state.enqueue(req, time.Now()); err != nil { + t.Fatal(err) + } + + delivery := state.take("node", func() *NodeStatusResponse { + <-time.After(time.Until(req.Deadline) + 10*time.Millisecond) + return &NodeStatusResponse{} + }, time.Now()) + if delivery != nil { + t.Fatal("expired collection was delivered") + } + + req = &statusv1alpha1.DetailRequest{RequestID: "disconnected", Deadline: time.Now().Add(20 * time.Millisecond)} + if err := state.enqueue(req, time.Now()); err != nil { + t.Fatal(err) + } + + if state.take("node", func() *NodeStatusResponse { return &NodeStatusResponse{} }, time.Now()) == nil { + t.Fatal("missing unacknowledged reply") + } + + waitForStatusCondition(t, func() bool { + state.mu.Lock() + defer state.mu.Unlock() + + return len(state.replies) == 0 + }) +} + func TestDetailPayloadErrorsAreCorrelated(t *testing.T) { for _, tc := range []struct { name string diff --git a/cmd/unbounded-net-node/status_details_ws_test.go b/cmd/unbounded-net-node/status_details_ws_test.go new file mode 100644 index 000000000..72778b213 --- /dev/null +++ b/cmd/unbounded-net-node/status_details_ws_test.go @@ -0,0 +1,239 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/coder/websocket" + "google.golang.org/protobuf/proto" + + statusproto "github.com/Azure/unbounded/internal/net/status/proto" + statusv1alpha1 "github.com/Azure/unbounded/internal/net/status/v1alpha1" +) + +func sendTestStatusAck(ctx context.Context, conn *websocket.Conn, ack *statusproto.NodeStatusAck) error { + payload, err := proto.Marshal(ack) + if err != nil { + return err + } + + return conn.Write(ctx, websocket.MessageBinary, payload) +} + +func TestWebSocketDetailsWakeWhilePublicationPending(t *testing.T) { + for _, mode := range []string{"summary", "full"} { + t.Run(mode, func(t *testing.T) { + messages := make(chan *statusproto.NodeStatusMessage, 8) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + return + } + defer func() { _ = conn.Close(websocket.StatusNormalClosure, "done") }() + + for { + _, data, err := conn.Read(r.Context()) + if err != nil { + return + } + + var msg statusproto.NodeStatusMessage + if err := proto.Unmarshal(data, &msg); err != nil { + t.Error(err) + return + } + + select { + case messages <- &msg: + case <-r.Context().Done(): + return + } + + if msg.Type == statusv1alpha1.NodeStatusDetailsType { + if err := sendTestStatusAck(r.Context(), conn, &statusproto.NodeStatusAck{ + Status: "ok", DetailRequestId: "request", Revision: 99, + }); err != nil { + return + } + } else { + if err := sendTestStatusAck(r.Context(), conn, &statusproto.NodeStatusAck{ + Status: statusv1alpha1.DetailRequestStatus, + DetailRequest: &statusproto.DetailRequest{RequestId: "expired", DeadlineUnixNs: time.Now().Add(-time.Second).UnixNano()}, + }); err != nil { + return + } + } + // No publication ACK is sent. This command (and its delayed + // duplicate after the detail ACK) cannot release that pending ACK. + if err := sendTestStatusAck(r.Context(), conn, &statusproto.NodeStatusAck{ + Status: statusv1alpha1.DetailRequestStatus, + DetailRequest: &statusproto.DetailRequest{RequestId: "request", DeadlineUnixNs: time.Now().Add(time.Minute).UnixNano()}, + }); err != nil { + return + } + } + })) + defer server.Close() + + s := summaryRouteFixture() + s.cfg.NodeName = "node-a" + + var collections atomic.Int32 + + s.bpfCollector = func() []BpfEntry { + collections.Add(1) + return []BpfEntry{{CIDR: "10.0.0.0/8"}} + } + h := blockedBootstrapHealthState() + h.setStatusServer(s) + + cfg := &config{ + NodeName: "node-a", StatusDetailMode: mode, StatusWSEnabled: true, + StatusWSURL: "ws" + strings.TrimPrefix(server.URL, "http"), StatusWSAPIServerMode: statusWSAPIServerModeNever, + CriticalDeltaEvery: time.Hour, StatsDeltaEvery: time.Hour, FullSyncEvery: time.Hour, + } + ctx, cancel := context.WithCancel(t.Context()) + startStatusPublishers(ctx, cfg, h) + + defer func() { cancel(); h.stopStatusPublishers() }() + + for i := range 2 { + select { + case msg := <-messages: + if !msg.SupportsDetails { + t.Fatal("missing responder capability") + } + + if i == 1 && (msg.Type != statusv1alpha1.NodeStatusDetailsType || msg.DetailRequestId != "request" || + msg.Status == nil || len(msg.Status.BpfEntries) != 1 || msg.BaseRevision != 0) { + t.Fatalf("invalid immediate detail reply: %v", msg) + } + case <-time.After(3 * time.Second): + t.Fatal("detail command waited for a periodic publication tick") + } + } + + waitForStatusCondition(t, func() bool { + state := h.detailState() + state.mu.Lock() + defer state.mu.Unlock() + + reply := state.replies["request"] + + return reply != nil && reply.done && reply.payload == nil + }) + + select { + case msg := <-messages: + t.Fatalf("duplicate or expired command produced another reply: %v", msg) + case <-time.After(100 * time.Millisecond): + } + + want := int32(1) + if mode == "full" { + want++ + } + + if collections.Load() != want { + t.Fatalf("collected %d full snapshots, want %d", collections.Load(), want) + } + }) + } +} + +func TestWebSocketDetailRetrySurvivesReconnect(t *testing.T) { + deliveries := make(chan []byte, 4) + + var connections atomic.Int32 + + deadline := time.Now().Add(time.Minute) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + return + } + + number := connections.Add(1) + + defer func() { _ = conn.Close(websocket.StatusGoingAway, "reconnect") }() + + for { + _, data, err := conn.Read(r.Context()) + if err != nil { + return + } + + var msg statusproto.NodeStatusMessage + if err := proto.Unmarshal(data, &msg); err != nil { + t.Error(err) + return + } + + if msg.Type == statusv1alpha1.NodeStatusDetailsType { + deliveries <- data + + if number == 1 { + return + } + + if err := sendTestStatusAck(r.Context(), conn, &statusproto.NodeStatusAck{Status: "ok", DetailRequestId: "retry"}); err != nil { + return + } + } else if err := sendTestStatusAck(r.Context(), conn, &statusproto.NodeStatusAck{ + Status: "ok", Revision: 1, SummarySupported: true, + DetailRequest: &statusproto.DetailRequest{RequestId: "retry", DeadlineUnixNs: deadline.UnixNano()}, + }); err != nil { + return + } + } + })) + defer server.Close() + + h := blockedBootstrapHealthState() + s := summaryRouteFixture() + s.cfg.NodeName = "node-a" + + var collections atomic.Int32 + + s.bpfCollector = func() []BpfEntry { collections.Add(1); return nil } + h.setStatusServer(s) + + cfg := &config{ + NodeName: "node-a", StatusDetailMode: "summary", StatusWSEnabled: true, + StatusWSURL: "ws" + strings.TrimPrefix(server.URL, "http"), StatusWSAPIServerMode: statusWSAPIServerModeNever, + CriticalDeltaEvery: time.Hour, StatsDeltaEvery: time.Hour, FullSyncEvery: time.Hour, + } + ctx, cancel := context.WithCancel(t.Context()) + startStatusPublishers(ctx, cfg, h) + + defer func() { cancel(); h.stopStatusPublishers() }() + + var first []byte + + for range 2 { + select { + case data := <-deliveries: + if first != nil && string(first) != string(data) { + t.Fatal("reconnect changed the collected reply") + } + + first = data + case <-time.After(4 * time.Second): + t.Fatal("outstanding detail was not retried after reconnect") + } + } + + if collections.Load() != 1 { + t.Fatal("reconnect recollected details") + } +} diff --git a/cmd/unbounded-net-node/status_server.go b/cmd/unbounded-net-node/status_server.go index f63513647..f2dd35ce1 100644 --- a/cmd/unbounded-net-node/status_server.go +++ b/cmd/unbounded-net-node/status_server.go @@ -101,6 +101,14 @@ func (h *nodeHealthState) stopStatusPublishers() { if wg != nil { wg.Wait() } + + h.mu.RLock() + details := h.details + h.mu.RUnlock() + + if details != nil { + details.clear() + } }) } @@ -1049,6 +1057,7 @@ func runStatusWebSocketPusher( dialHTTPClient *http.Client, hmacMgr *hmacTokenManager, ) { + details := healthState.detailState() // Reuse the push client's TLS trust setup so wss://KUBERNETES_SERVICE_HOST // can validate the cluster CA in fallback/preferred API server modes. // Keep timeout disabled for long-lived websocket connections. @@ -1366,7 +1375,14 @@ func runStatusWebSocketPusher( return } - if acks.accept(data) { + ack, err := decodeNodeStatusAck(data) + if err != nil { + continue + } + + details.receive(ack) + + if acks.acceptAck(ack) { lastAckTimeNs.Store(time.Now().UnixNano()) if cfg.StatusDetailMode == "summary" && !acks.summary.Load() { @@ -1390,6 +1406,7 @@ func runStatusWebSocketPusher( msg := &statusproto.NodeStatusMessage{ Type: statusv1alpha1.NodeStatusSummaryType, NodeName: summary.NodeInfo.Name, BaseRevision: acks.revision.Load(), Summary: nodeSummaryToProto(summary), + SupportsDetails: true, } payload, err := proto.Marshal(msg) @@ -1436,9 +1453,10 @@ func runStatusWebSocketPusher( } msg := &statusproto.NodeStatusMessage{ - Type: "node_status_full", - NodeName: status.NodeInfo.Name, - Status: nodeStatusToProto(status), + Type: "node_status_full", + NodeName: status.NodeInfo.Name, + Status: nodeStatusToProto(status), + SupportsDetails: true, } payload, err := proto.Marshal(msg) @@ -1554,6 +1572,20 @@ func runStatusWebSocketPusher( } fallbackCloseTicker := time.NewTicker(500 * time.Millisecond) + sendDetails := func() error { + delivery := details.take(cfg.NodeName, healthState.getStatusSnapshot, time.Now()) + if delivery == nil { + return nil + } + defer details.finish(delivery.id) + + detailCtx, cancel := context.WithDeadline(ctx, delivery.deadline) + defer cancel() + + return conn.Write(detailCtx, websocket.MessageBinary, delivery.payload) + } + + details.wake() loop: for { @@ -1562,6 +1594,11 @@ func runStatusWebSocketPusher( break loop case <-readCtx.Done(): break loop + case <-details.wsWake: + if err := sendDetails(); err != nil { + klog.V(2).Infof("Status websocket: detail reply write failed: %v", err) + break loop + } case <-criticalTicker.C: if acks.pending.Load() { continue @@ -1599,10 +1636,11 @@ func runStatusWebSocketPusher( } message := &statusproto.NodeStatusMessage{ - Type: "node_status_delta", - NodeName: current.NodeInfo.Name, - BaseRevision: acks.revision.Load(), - Delta: delta, + Type: "node_status_delta", + NodeName: current.NodeInfo.Name, + BaseRevision: acks.revision.Load(), + Delta: delta, + SupportsDetails: true, } payload, err := proto.Marshal(message) @@ -1658,10 +1696,11 @@ func runStatusWebSocketPusher( } wsMsg := &statusproto.NodeStatusMessage{ - Type: "node_status_delta", - NodeName: current.NodeInfo.Name, - BaseRevision: acks.revision.Load(), - Delta: delta, + Type: "node_status_delta", + NodeName: current.NodeInfo.Name, + BaseRevision: acks.revision.Load(), + Delta: delta, + SupportsDetails: true, } payload, err := proto.Marshal(wsMsg) @@ -1730,6 +1769,10 @@ func runStatusWebSocketPusher( directRecoveryTimer.Reset(directRecoveryBackoff) } case <-fallbackCloseTicker.C: + if err := sendDetails(); err != nil { + break loop + } + if acks.pending.Load() && time.Since(lastWriteTime) > 30*time.Second { klog.V(2).Info("Status websocket: status acknowledgment timed out") break loop @@ -1769,6 +1812,8 @@ func runStatusWebSocketPusher( if wsMode != nil { wsMode.Store(statusWSModeNone) } + + details.wake() // Send a graceful WebSocket close frame. Use StatusNormalClosure // for clean shutdown and StatusGoingAway for reconnect scenarios. // conn.Close has its own 5s timeout for the close handshake.