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

This commit is contained in:
Nixevol
2026-09-30 15:08:03 +08:00
parent b5b63ed070
commit 42160720bd
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)
}
}
+54 -8
View File
@@ -92,6 +92,8 @@ type Broker struct {
connsMu sync.RWMutex
current map[string]*connState
byClient map[*mqtt.Client]*connState
byConnID map[port.ConnID]*connState
closedCh chan struct{}
queuesMu sync.Mutex
queues map[string]*uplinkQueue
@@ -116,6 +118,8 @@ type connState struct {
largePIDs map[uint16]struct{}
largePending int
metricsCounted bool
established bool
createdAt time.Time
mu sync.Mutex
handshakeTimer *time.Timer
@@ -135,6 +139,7 @@ func New(opts Options) (*Broker, error) {
if log == nil {
log = slog.Default()
}
log = slog.New(newRedactHandler(log.Handler()))
caps := mqtt.NewDefaultServerCapabilities()
caps.MaximumClients = maxClients
@@ -165,6 +170,8 @@ func New(opts Options) (*Broker, error) {
metrics: opts.Metrics,
current: make(map[string]*connState),
byClient: make(map[*mqtt.Client]*connState),
byConnID: make(map[port.ConnID]*connState),
closedCh: make(chan struct{}),
queues: make(map[string]*uplinkQueue),
largeSem: make(chan struct{}, largeFrameSlots),
}
@@ -175,6 +182,7 @@ func New(opts Options) (*Broker, error) {
if err := srv.Serve(); err != nil {
return nil, err
}
go b.sweepLoop()
return b, nil
}
@@ -186,6 +194,11 @@ func (b *Broker) Close() error {
if b.closed.Swap(true) {
return nil
}
select {
case <-b.closedCh:
default:
close(b.closedCh)
}
b.queuesMu.Lock()
for _, q := range b.queues {
q.close()
@@ -345,10 +358,9 @@ func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
if connID != "" {
for _, st := range b.byClient {
if st.endpointID == endpointID && st.connID == connID {
return st
}
st := b.byConnID[connID]
if st != nil && st.endpointID == endpointID {
return st
}
return nil
}
@@ -462,12 +474,46 @@ func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) {
func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
for _, st := range b.byClient {
if st.endpointID == endpointID && st.connID == connID {
return st
st := b.byConnID[connID]
if st == nil || st.endpointID != endpointID {
return nil
}
return st
}
func (b *Broker) sweepLoop() {
tick := time.NewTicker(time.Minute)
defer tick.Stop()
for {
select {
case <-tick.C:
b.sweepUnestablished(time.Minute)
case <-b.closedCh:
return
}
}
}
func (b *Broker) sweepUnestablished(minAge time.Duration) {
now := time.Now()
b.connsMu.Lock()
defer b.connsMu.Unlock()
for cl, st := range b.byClient {
if st.established {
continue
}
if cl != nil && !cl.Closed() {
continue
}
if minAge > 0 && now.Sub(st.createdAt) < minAge {
continue
}
delete(b.byClient, cl)
delete(b.byConnID, st.connID)
if b.current[st.endpointID] == st {
delete(b.current, st.endpointID)
}
}
return nil
}
func (b *Broker) hasDownSub(st *connState) bool {
+12 -6
View File
@@ -3,6 +3,7 @@ package broker
import (
"bytes"
"context"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
mqtt "github.com/mochi-mqtt/server/v2"
@@ -47,12 +48,11 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
remoteIP: remoteIP,
client: cl,
maxPacketSize: pk.Properties.MaximumPacketSize,
createdAt: time.Now(),
}
// ClientID、Username 都必须等于端编号
if clientID == "" || endpointID == "" || clientID != endpointID {
st.authOK = false
h.rememberPending(cl, st)
return nil
}
@@ -81,11 +81,12 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP)
if err != nil {
st.authErr = err
h.rememberPending(cl, st)
return err // mochi 不回 CONNACK,直接断开
return err // mochi 不回 CONNACK,直接断开;不登记连接表
}
st.authOK = res.OK
if !res.OK {
return nil
}
st.authOK = true
st.sessionToken = res.SessionToken
h.rememberPending(cl, st)
return nil
@@ -94,6 +95,7 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) {
h.b.connsMu.Lock()
h.b.byClient[cl] = st
h.b.byConnID[st.connID] = st
h.b.connsMu.Unlock()
}
@@ -188,6 +190,7 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
st := h.b.byClient[cl]
if st != nil {
h.b.current[st.endpointID] = st
st.established = true
}
h.b.connsMu.Unlock()
if st == nil {
@@ -209,6 +212,9 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
h.b.connsMu.Lock()
st := h.b.byClient[cl]
delete(h.b.byClient, cl)
if st != nil {
delete(h.b.byConnID, st.connID)
}
isCurrent := false
if st != nil && h.b.current[st.endpointID] == st {
delete(h.b.current, st.endpointID)
+88
View File
@@ -0,0 +1,88 @@
package broker
import (
"context"
"log/slog"
"github.com/mochi-mqtt/server/v2/packets"
)
type redactHandler struct {
inner slog.Handler
}
func newRedactHandler(inner slog.Handler) slog.Handler {
if inner == nil {
inner = slog.Default().Handler()
}
return &redactHandler{inner: inner}
}
func (h *redactHandler) Enabled(ctx context.Context, level slog.Level) bool {
return h.inner.Enabled(ctx, level)
}
func (h *redactHandler) Handle(ctx context.Context, r slog.Record) error {
rec := slog.NewRecord(r.Time, r.Level, r.Message, r.PC)
r.Attrs(func(a slog.Attr) bool {
rec.AddAttrs(redactSlogAttr(a))
return true
})
return h.inner.Handle(ctx, rec)
}
func (h *redactHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
out := make([]slog.Attr, len(attrs))
for i, a := range attrs {
out[i] = redactSlogAttr(a)
}
return &redactHandler{inner: h.inner.WithAttrs(out)}
}
func (h *redactHandler) WithGroup(name string) slog.Handler {
return &redactHandler{inner: h.inner.WithGroup(name)}
}
func redactSlogAttr(a slog.Attr) slog.Attr {
a.Value = a.Value.Resolve()
switch v := a.Value.Any().(type) {
case packets.Packet:
return slog.Any(a.Key, summarizePacket(v))
case *packets.Packet:
if v == nil {
return a
}
return slog.Any(a.Key, summarizePacket(*v))
}
if a.Value.Kind() == slog.KindGroup {
group := a.Value.Group()
out := make([]slog.Attr, len(group))
for i, g := range group {
out[i] = redactSlogAttr(g)
}
return slog.Attr{Key: a.Key, Value: slog.GroupValue(out...)}
}
return a
}
type mqttPacketLog struct {
Type string `json:"type"`
QoS byte `json:"qos"`
PacketID uint16 `json:"packet_id"`
Topic string `json:"topic,omitempty"`
PayloadLen int `json:"payload_len"`
}
func summarizePacket(pk packets.Packet) mqttPacketLog {
name := packets.PacketNames[pk.FixedHeader.Type]
if name == "" {
name = "unknown"
}
return mqttPacketLog{
Type: name,
QoS: pk.FixedHeader.Qos,
PacketID: pk.PacketID,
Topic: pk.TopicName,
PayloadLen: len(pk.Payload),
}
}