package broker import ( "io" "net" "net/http" "net/http/httptest" "strings" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/metrics" dto "github.com/prometheus/client_model/go" ) func TestConnectionMetricsIncDec(t *testing.T) { reg := metrics.New() b, err := New(Options{Authenticator: AllowAuthenticator{}, Metrics: reg}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() if got := gaugeValue(t, reg, "nixmsg_connections", "tcp"); got != 0 { t.Fatalf("before connect tcp=%v", got) } r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b.AttachTCP(r) }() connectAndSubscribe(t, w, "ep-metrics", 0) deadline := time.Now().Add(3 * time.Second) for { if _, ok := b.ConnInfoOf("ep-metrics"); ok { break } if time.Now().After(deadline) { t.Fatal("session not established") } time.Sleep(10 * time.Millisecond) } if got := gaugeValue(t, reg, "nixmsg_connections", "tcp"); got != 1 { t.Fatalf("after connect tcp=%v want 1", got) } _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): t.Fatal("attach did not finish") } deadline = time.Now().Add(3 * time.Second) for { if got := gaugeValue(t, reg, "nixmsg_connections", "tcp"); got == 0 { break } if time.Now().After(deadline) { t.Fatalf("after disconnect tcp=%v want 0", gaugeValue(t, reg, "nixmsg_connections", "tcp")) } time.Sleep(10 * time.Millisecond) } } func TestWSConnectionMetrics(t *testing.T) { reg := metrics.New() b, err := New(Options{Authenticator: AllowAuthenticator{}, Metrics: reg}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() mux := http.NewServeMux() mux.Handle("/mqtt", b.WSHandler(nil)) srv := httptest.NewServer(mux) defer srv.Close() // 仅验证 handler 暴露指标文本仍含初始标签;建连用 TCP 测即可。 req := httptest.NewRequest(http.MethodGet, "/metrics", nil) rec := httptest.NewRecorder() reg.Handler().ServeHTTP(rec, req) body, _ := io.ReadAll(rec.Body) if !strings.Contains(string(body), `nixmsg_connections{transport="ws"} 0`) { t.Fatalf("missing ws series: %s", body) } } func gaugeValue(t *testing.T, reg *metrics.Registry, name, transport string) float64 { t.Helper() mfs, err := reg.Gatherer().Gather() if err != nil { t.Fatal(err) } for _, mf := range mfs { if mf.GetName() != name { continue } for _, m := range mf.GetMetric() { if matchLabel(m, "transport", transport) { return m.GetGauge().GetValue() } } } t.Fatalf("metric %s transport=%s not found", name, transport) return 0 } func matchLabel(m *dto.Metric, key, val string) bool { for _, lp := range m.GetLabel() { if lp.GetName() == key && lp.GetValue() == val { return true } } return false }