package live_test import ( "context" "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/coder/websocket" "github.com/coder/websocket/wsjson" "github.com/itworx/pulse/internal/auth" "github.com/itworx/pulse/internal/live" "github.com/itworx/pulse/internal/metriccatalog" "github.com/itworx/pulse/internal/queryplan" ) type sampler struct{} func (sampler) Sample(_ context.Context, request queryplan.Request) ([]live.Sample, error) { value := float64(request.MaxPoints) return []live.Sample{{Series: "test", Timestamp: time.Now().UTC(), Value: &value, Freshness: "fresh"}}, nil } func plannerForTest(t *testing.T) *queryplan.Planner { t.Helper() registry, err := metriccatalog.DefaultRegistry() if err != nil { t.Fatal(err) } planner := queryplan.NewPlanner(registry, queryplan.Limits{}) return &planner } func validQuery() map[string]any { now := time.Now().UTC().Truncate(time.Second) return map[string]any{ "metric": "container.cpu.utilization", "scope": map[string]string{"containerId": "media_server"}, "range": map[string]any{"from": now.Add(-time.Minute).Format(time.RFC3339), "to": now.Format(time.RFC3339), "stepSeconds": 15}, "aggregation": "avg", "maxSeries": 1, "maxPoints": 60, } } func authorizedServer(t *testing.T, handler live.Handler) *httptest.Server { t.Helper() return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { r = r.WithContext(auth.WithPrincipal(r.Context(), auth.Principal{Subject: "viewer", Role: auth.RoleViewer})) handler.ServeHTTP(w, r) })) } func TestHandlerRejectsUnauthenticatedBeforeUpgrade(t *testing.T) { server := httptest.NewServer(live.Handler{}) defer server.Close() _, response, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", nil) if err == nil { t.Fatal("expected unauthenticated upgrade rejection") } if response == nil || response.StatusCode != http.StatusUnauthorized { t.Fatalf("response=%v err=%v", response, err) } } func TestHandlerRejectsCrossOriginUpgrade(t *testing.T) { server := authorizedServer(t, live.Handler{Planner: plannerForTest(t)}) defer server.Close() options := &websocket.DialOptions{HTTPHeader: http.Header{"Origin": []string{"https://evil.example"}}} _, response, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", options) if err == nil { t.Fatal("expected cross-origin rejection") } if response == nil || response.StatusCode != http.StatusForbidden { t.Fatalf("response=%v err=%v", response, err) } } func TestHandlerClosesWhenAuthenticatedRequestIsCancelled(t *testing.T) { cancelled := make(chan context.CancelFunc, 1) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ctx, cancel := context.WithCancel(auth.WithPrincipal(r.Context(), auth.Principal{Subject: "viewer", Role: auth.RoleViewer})) cancelled <- cancel live.Handler{}.ServeHTTP(w, r.WithContext(ctx)) })) defer server.Close() conn, _, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", nil) if err != nil { t.Fatal(err) } defer conn.CloseNow() (<-cancelled)() readCtx, cancelRead := context.WithTimeout(context.Background(), time.Second) defer cancelRead() if _, _, err := conn.Read(readCtx); err == nil { t.Fatal("live connection survived authenticated request cancellation") } } func TestHandlerAuthorizesSubscriptionAndSequencesSamples(t *testing.T) { server := authorizedServer(t, live.Handler{Planner: plannerForTest(t), Sampler: sampler{}}) defer server.Close() conn, _, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", nil) if err != nil { t.Fatal(err) } defer conn.Close(websocket.StatusNormalClosure, "") ctx, cancel := context.WithTimeout(context.Background(), 8*time.Second) defer cancel() if err := wsjson.Write(ctx, conn, map[string]any{"schemaVersion": 1, "type": "subscribe", "subscriptionId": "sub-1", "query": validQuery(), "intervalSeconds": 1}); err != nil { t.Fatal(err) } var status live.StatusMessage if err := wsjson.Read(ctx, conn, &status); err != nil { t.Fatal(err) } if status.State != "subscribed" { t.Fatalf("status=%+v", status) } var first, second live.SamplesMessage if err := wsjson.Read(ctx, conn, &first); err != nil { t.Fatal(err) } if err := wsjson.Read(ctx, conn, &second); err != nil { t.Fatal(err) } if first.Sequence != 1 || second.Sequence != 2 || second.Sequence <= first.Sequence { t.Fatalf("sequences=%d,%d", first.Sequence, second.Sequence) } if len(first.Samples) != 1 || first.Samples[0].Freshness != "fresh" { t.Fatalf("samples=%+v", first.Samples) } } func TestHandlerReportsMissingSamplerInsteadOfEmptySamples(t *testing.T) { server := authorizedServer(t, live.Handler{Planner: plannerForTest(t)}) defer server.Close() conn, _, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", nil) if err != nil { t.Fatal(err) } defer func() { _ = conn.Close(websocket.StatusNormalClosure, "") }() ctx, cancel := context.WithTimeout(context.Background(), 8*time.Second) defer cancel() if err := wsjson.Write(ctx, conn, map[string]any{"schemaVersion": 1, "type": "subscribe", "subscriptionId": "sub-1", "query": validQuery(), "intervalSeconds": 1}); err != nil { t.Fatal(err) } var status live.StatusMessage if err := wsjson.Read(ctx, conn, &status); err != nil { t.Fatal(err) } if status.State != "subscribed" { t.Fatalf("status=%+v", status) } var failure live.ErrorMessage if err := wsjson.Read(ctx, conn, &failure); err != nil { t.Fatal(err) } if failure.Type != "error" || failure.Code != "LIVE_SAMPLE_UNAVAILABLE" || failure.SubscriptionID != "sub-1" { t.Fatalf("failure=%+v", failure) } } func TestHandlerDeniesInvalidSubscriptionAndLimitsMessages(t *testing.T) { server := authorizedServer(t, live.Handler{Planner: plannerForTest(t), MaxMessages: 2, RateWindow: time.Minute}) defer server.Close() conn, _, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", nil) if err != nil { t.Fatal(err) } defer conn.Close(websocket.StatusNormalClosure, "") ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() if err := wsjson.Write(ctx, conn, map[string]any{"schemaVersion": 1, "type": "subscribe", "subscriptionId": "bad", "query": map[string]any{"metric": "unknown"}, "intervalSeconds": 1}); err != nil { t.Fatal(err) } var problem live.ErrorMessage if err := wsjson.Read(ctx, conn, &problem); err != nil { t.Fatal(err) } if problem.Code != "LIVE_SUBSCRIPTION_DENIED" { t.Fatalf("problem=%+v", problem) } if err := wsjson.Write(ctx, conn, map[string]any{"schemaVersion": 1, "type": "ping", "nonce": "one"}); err != nil { t.Fatal(err) } var pong live.PongMessage if err := wsjson.Read(ctx, conn, &pong); err != nil { t.Fatal(err) } if err := wsjson.Write(ctx, conn, map[string]any{"schemaVersion": 1, "type": "ping", "nonce": "two"}); err != nil { t.Fatal(err) } var limited live.ErrorMessage if err := wsjson.Read(ctx, conn, &limited); err != nil { t.Fatal(err) } if limited.Code != "LIVE_RATE_LIMIT" { t.Fatalf("problem=%+v", limited) } } func TestHandlerClosesOversizedMessage(t *testing.T) { server := authorizedServer(t, live.Handler{Planner: plannerForTest(t)}) defer server.Close() conn, _, err := websocket.Dial(context.Background(), "ws"+server.URL[4:]+"/api/v1/live", nil) if err != nil { t.Fatal(err) } defer conn.Close(websocket.StatusNormalClosure, "") ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() if err := wsjson.Write(ctx, conn, map[string]any{"schemaVersion": 1, "type": "ping", "nonce": strings.Repeat("x", 70000)}); err != nil { t.Fatal(err) } var response map[string]any if err := wsjson.Read(ctx, conn, &response); err == nil { t.Fatal("expected oversized message connection to close") } }