fix: 认证失败不泄漏连接表并脱敏 mochi 整包日志

This commit is contained in:
Nixevol
2026-09-30 16:21:05 +08:00
parent deb2398e27
commit 0b9ce0359a
5 changed files with 442 additions and 14 deletions
+270
View File
@@ -0,0 +1,270 @@
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)
}
}