package broker import ( "bytes" "context" "encoding/base64" "io" "log/slog" "net" "strings" "testing" "time" "github.com/mochi-mqtt/server/v2/packets" ) func TestFailedAuthDoesNotLeakConnTable(t *testing.T) { secret := "s3cret-token-xyz" var logBuf bytes.Buffer log := slog.New(slog.NewTextHandler(&logBuf, &slog.HandlerOptions{Level: slog.LevelDebug})) b, err := New(Options{Authenticator: RejectAuthenticator{}, Logger: log}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() const n = 200 for i := 0; i < n; i++ { dialFailedCONNECT(t, b, func(w net.Conn) { writeConnect(t, w, "ep-rej", 30, 0) }) } for i := 0; i < n; i++ { dialFailedCONNECT(t, b, func(w net.Conn) { writeConnectMismatch(t, w) }) } b2, err := New(Options{Authenticator: &errAuthenticator{err: context.DeadlineExceeded}}) if err != nil { t.Fatal(err) } defer func() { _ = b2.Close() }() for i := 0; i < n; i++ { dialFailedCONNECT(t, b2, func(w net.Conn) { writeConnect(t, w, "ep-err", 30, 0) }) } if got := len(b.byClient); got != 0 { t.Fatalf("reject/mismatch leaked %d", got) } if got := len(b2.byClient); got != 0 { t.Fatalf("internal error leaked %d", got) } // B-07:拒绝路径的 mochi 日志不能带密码 b3, err := New(Options{Authenticator: AllowAuthenticator{}, Logger: log}) if err != nil { t.Fatal(err) } defer func() { _ = b3.Close() }() r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b3.AttachTCP(r) }() writeConnectWithPassword(t, w, "ep-log", secret) readExactPacket(t, w, packets.Connack, 3*time.Second) writeConnectWithPassword(t, w, "ep-log", secret) // 同一连接第二个 CONNECT _ = w.Close() select { case <-done: case <-time.After(2 * time.Second): } out := logBuf.String() if strings.Contains(out, secret) { t.Fatalf("log contains password: %s", out) } if strings.Contains(out, base64.StdEncoding.EncodeToString([]byte(secret))) { t.Fatalf("log contains password base64: %s", out) } } func TestSweepUnestablishedClosedConn(t *testing.T) { b, err := New(Options{Authenticator: AllowAuthenticator{}}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b.AttachTCP(r) }() writeConnect(t, w, "ep-sweep", 30, 0) _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } b.sweepUnestablished(0) if got := len(b.byClient); got != 0 { t.Fatalf("after sweep byClient=%d", got) } } func TestLookupByConnIDIndependentOfFailedConns(t *testing.T) { b, err := New(Options{Authenticator: AllowAuthenticator{}}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b.AttachTCP(r) }() connectAndSubscribe(t, w, "ep-ok", 0) waitSession(t, b, "ep-ok") info, ok := b.ConnInfoOf("ep-ok") if !ok { t.Fatal("missing session") } st := b.lookupConn("ep-ok", info.ConnID) if st == nil { t.Fatal("lookup by conn id") } _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } } func dialFailedCONNECT(t *testing.T, b *Broker, write func(net.Conn)) { t.Helper() r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b.AttachTCP(r) }() write(w) _ = w.Close() select { case <-done: case <-time.After(2 * time.Second): t.Fatal("attach did not return") } } func writeConnectMismatch(t *testing.T, w net.Conn) { t.Helper() pk := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Connect}, ProtocolVersion: 5, Connect: packets.ConnectParams{ ProtocolName: []byte("MQTT"), Clean: true, ClientIdentifier: "id-a", Keepalive: 30, UsernameFlag: true, Username: []byte("id-b"), PasswordFlag: true, Password: []byte("nope"), }, } var buf bytes.Buffer if err := pk.ConnectEncode(&buf); err != nil { t.Fatal(err) } if _, err := w.Write(buf.Bytes()); err != nil { t.Fatal(err) } } func writeConnectWithPassword(t *testing.T, w net.Conn, endpoint, password string) { t.Helper() pk := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Connect}, ProtocolVersion: 5, Connect: packets.ConnectParams{ ProtocolName: []byte("MQTT"), Clean: true, ClientIdentifier: endpoint, Keepalive: 30, UsernameFlag: true, Username: []byte(endpoint), PasswordFlag: true, Password: []byte(password), }, } var buf bytes.Buffer if err := pk.ConnectEncode(&buf); err != nil { t.Fatal(err) } if _, err := w.Write(buf.Bytes()); err != nil { t.Fatal(err) } } func TestMQTT311UnauthorizedPublishOmitsPayloadInLogs(t *testing.T) { var logBuf bytes.Buffer log := slog.New(slog.NewTextHandler(&logBuf, &slog.HandlerOptions{Level: slog.LevelDebug})) b, err := New(Options{Authenticator: AllowAuthenticator{}, Logger: log}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b.AttachTCP(r) }() pk := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Connect}, ProtocolVersion: 4, Connect: packets.ConnectParams{ ProtocolName: []byte("MQTT"), Clean: true, ClientIdentifier: "ep311", Keepalive: 30, UsernameFlag: true, Username: []byte("ep311"), PasswordFlag: true, Password: []byte("test"), }, } var buf bytes.Buffer if err := pk.ConnectEncode(&buf); err != nil { t.Fatal(err) } if _, err := w.Write(buf.Bytes()); err != nil { t.Fatal(err) } _ = w.SetReadDeadline(time.Now().Add(3 * time.Second)) raw := make([]byte, 256) if _, err := io.ReadAtLeast(w, raw, 2); err != nil { t.Fatal(err) } body := []byte(`{"talk_password":"super-secret-body"}`) pub := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1}, ProtocolVersion: 4, TopicName: "nix/c/other/up", PacketID: 7, Payload: body, } buf.Reset() if err := pub.PublishEncode(&buf); err != nil { t.Fatal(err) } _, _ = w.Write(buf.Bytes()) time.Sleep(50 * time.Millisecond) _ = w.Close() select { case <-done: case <-time.After(2 * time.Second): } out := logBuf.String() if strings.Contains(out, "super-secret-body") { t.Fatalf("log contains publish payload: %s", out) } }