package daemon import ( "context" "io" "log/slog" "net" "net/http" "path/filepath" "testing" "time" "git.cloudinit.dev/coreci/orca/internal/store" ) func TestStartPprof_Disabled(t *testing.T) { srv, err := StartPprof("", slog.Default()) if err != nil { t.Fatalf("StartPprof(\"\", _) returned err: %v", err) } if srv != nil { t.Fatalf("StartPprof(\"\", _) returned non-nil server: %v", srv) } } func TestStartPprof_Enabled(t *testing.T) { log := slog.New(slog.NewTextHandler(io.Discard, nil)) ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } addr := ln.Addr().String() _ = ln.Close() srv, err := StartPprof(addr, log) if err != nil { t.Fatalf("StartPprof returned err: %v", err) } if srv == nil { t.Fatal("StartPprof returned nil server for non-empty addr") } t.Cleanup(func() { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() _ = srv.Shutdown(ctx) }) deadline := time.Now().Add(2 * time.Second) var base string for time.Now().Before(deadline) { conn, derr := net.DialTimeout("tcp", addr, 50*time.Millisecond) if derr == nil { _ = conn.Close() base = "http://" + addr break } time.Sleep(20 * time.Millisecond) } if base == "" { t.Fatal("pprof server did not start listening") } client := &http.Client{Timeout: 500 * time.Millisecond} for _, path := range []string{"/debug/pprof/", "/debug/pprof/cmdline", "/debug/pprof/heap"} { resp, gerr := client.Get(base + path) if gerr != nil { t.Errorf("GET %s: %v", path, gerr) continue } _, _ = io.Copy(io.Discard, resp.Body) _ = resp.Body.Close() if resp.StatusCode != 200 { t.Errorf("GET %s: expected 200, got %d", path, resp.StatusCode) } } } func TestStartPprof_Shutdown(t *testing.T) { log := slog.New(slog.NewTextHandler(io.Discard, nil)) ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } addr := ln.Addr().String() _ = ln.Close() srv, err := StartPprof(addr, log) if err != nil { t.Fatalf("StartPprof returned err: %v", err) } if srv == nil { t.Fatal("StartPprof returned nil server") } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { conn, derr := net.DialTimeout("tcp", addr, 50*time.Millisecond) if derr == nil { _ = conn.Close() break } time.Sleep(20 * time.Millisecond) } ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() if err := srv.Shutdown(ctx); err != nil { t.Fatalf("Shutdown: %v", err) } client := &http.Client{Timeout: 300 * time.Millisecond} _, gerr := client.Get("http://" + addr + "/debug/pprof/") if gerr == nil { t.Error("expected GET to fail after Shutdown, but it succeeded") } } func TestStartPprof_MuxIsolated(t *testing.T) { log := slog.New(slog.NewTextHandler(io.Discard, nil)) ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } addr := ln.Addr().String() _ = ln.Close() srv, err := StartPprof(addr, log) if err != nil { t.Fatalf("StartPprof returned err: %v", err) } if srv == nil { t.Fatal("StartPprof returned nil server") } t.Cleanup(func() { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() _ = srv.Shutdown(ctx) }) deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { conn, derr := net.DialTimeout("tcp", addr, 50*time.Millisecond) if derr == nil { _ = conn.Close() break } time.Sleep(20 * time.Millisecond) } client := &http.Client{Timeout: 500 * time.Millisecond} resp, err := client.Get("http://" + addr + "/healthz") if err != nil { t.Fatalf("GET /healthz: %v", err) } _, _ = io.Copy(io.Discard, resp.Body) _ = resp.Body.Close() if resp.StatusCode != 404 { t.Errorf("expected /healthz to 404 on pprof-only mux, got %d", resp.StatusCode) } } func TestServer_WithPprof(t *testing.T) { db, err := store.Open(filepath.Join(t.TempDir(), "pprof.db")) if err != nil { t.Fatalf("open db: %v", err) } defer db.Close() log := slog.New(slog.NewTextHandler(io.Discard, nil)) ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen main: %v", err) } mainAddr := ln.Addr().String() pln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen pprof: %v", err) } pprofAddr := pln.Addr().String() _ = pln.Close() s := NewServer(Options{ DB: db, Log: log, Addr: mainAddr, PprofAddr: pprofAddr, }) s.MarkReady() if s.pprofServer == nil { t.Fatal("expected pprofServer to be non-nil after NewServer with PprofAddr") } errCh := make(chan error, 2) go func() { err := s.httpServer.Serve(ln) if err != nil && err != http.ErrServerClosed { errCh <- err } }() deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { conn, derr := net.DialTimeout("tcp", pprofAddr, 50*time.Millisecond) if derr == nil { _ = conn.Close() break } time.Sleep(20 * time.Millisecond) } client := &http.Client{Timeout: 500 * time.Millisecond} resp, err := client.Get("http://" + mainAddr + "/healthz") if err != nil { t.Fatalf("GET main /healthz: %v", err) } if resp.StatusCode != 200 { t.Errorf("main /healthz: expected 200, got %d", resp.StatusCode) } _, _ = io.Copy(io.Discard, resp.Body) _ = resp.Body.Close() presp, err := client.Get("http://" + pprofAddr + "/debug/pprof/") if err != nil { t.Fatalf("GET pprof /debug/pprof/: %v", err) } if presp.StatusCode != 200 { t.Errorf("pprof /debug/pprof/: expected 200, got %d", presp.StatusCode) } _, _ = io.Copy(io.Discard, presp.Body) _ = presp.Body.Close() presp, err = client.Get("http://" + pprofAddr + "/healthz") if err != nil { t.Fatalf("GET pprof /healthz: %v", err) } _, _ = io.Copy(io.Discard, presp.Body) _ = presp.Body.Close() if presp.StatusCode != 404 { t.Errorf("expected /healthz 404 on pprof mux, got %d", presp.StatusCode) } ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() if err := s.Shutdown(ctx); err != nil { t.Errorf("Shutdown: %v", err) } client = &http.Client{Timeout: 300 * time.Millisecond} _, gerr := client.Get("http://" + pprofAddr + "/debug/pprof/") if gerr == nil { t.Error("expected pprof GET to fail after Shutdown") } _, merr := client.Get("http://" + mainAddr + "/healthz") if merr == nil { t.Error("expected main GET to fail after Shutdown") } }