fix: 认证失败不泄漏连接表并脱敏 mochi 整包日志
This commit is contained in:
@@ -1232,3 +1232,21 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
|
|||||||
- 原因:mochi 传给 `OnQosComplete` 的是 PUBACK,没有载荷,旧实现从未归还。
|
- 原因:mochi 传给 `OnQosComplete` 的是 PUBACK,没有载荷,旧实现从未归还。
|
||||||
- 备选方案:发布后立即归还(会把卡死点挪到消息包那份名额)。
|
- 备选方案:发布后立即归还(会把卡死点挪到消息包那份名额)。
|
||||||
- 影响:只在 broker 保留一份全局 64 名额;确认超时仍由消息线踢线/清标记触发断线归还。
|
- 影响:只在 broker 保留一份全局 64 名额;确认超时仍由消息线踢线/清标记触发断线归还。
|
||||||
|
|
||||||
|
### 复审修复 B-05
|
||||||
|
|
||||||
|
- 日期:2026-09-30
|
||||||
|
- 原条款:DEVELOPMENT 第 5 节连接表;Gitea #12。
|
||||||
|
- 实际做法:`OnConnect` 只在认证通过时写入 `byClient`/`byConnID`;拒绝与内部错误不登记。`connState` 增加 `established` 与 `createdAt`,每分钟清扫未建立且已关闭超过 1 分钟的条目。按连接代号查找改为 O(1)。
|
||||||
|
- 原因:mochi 在认证失败路径不调用 `OnDisconnect`,旧实现会永久泄漏。
|
||||||
|
- 备选方案:失败路径也登记再在 Authenticate 返回 false 时删除(仍覆盖不了 CONNACK 失败)。
|
||||||
|
- 影响:失败连接不再占用查找路径;行为对客户端不变(仍回 0x86 或不回 CONNACK)。
|
||||||
|
|
||||||
|
### 复审修复 B-07
|
||||||
|
|
||||||
|
- 日期:2026-09-30
|
||||||
|
- 原条款:PRD §8 日志无正文、无密码、无令牌。Gitea #14。
|
||||||
|
- 实际做法:`broker.New` 给 mochi 包一层 slog.Handler,把 `packets.Packet` / `*packets.Packet` 换成类型、QoS、包号、主题、正文长度。
|
||||||
|
- 原因:默认 info 下第二个 CONNECT、3.1.1 发到错误主题等会把整包写入 JSON 日志。
|
||||||
|
- 备选方案:改 mochi 日志调用点(需 fork)。
|
||||||
|
- 影响:排障时看不到载荷与密码,只见摘要。
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -92,6 +92,8 @@ type Broker struct {
|
|||||||
connsMu sync.RWMutex
|
connsMu sync.RWMutex
|
||||||
current map[string]*connState
|
current map[string]*connState
|
||||||
byClient map[*mqtt.Client]*connState
|
byClient map[*mqtt.Client]*connState
|
||||||
|
byConnID map[port.ConnID]*connState
|
||||||
|
closedCh chan struct{}
|
||||||
|
|
||||||
queuesMu sync.Mutex
|
queuesMu sync.Mutex
|
||||||
queues map[string]*uplinkQueue
|
queues map[string]*uplinkQueue
|
||||||
@@ -116,6 +118,8 @@ type connState struct {
|
|||||||
largePIDs map[uint16]struct{}
|
largePIDs map[uint16]struct{}
|
||||||
largePending int
|
largePending int
|
||||||
metricsCounted bool
|
metricsCounted bool
|
||||||
|
established bool
|
||||||
|
createdAt time.Time
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
|
|
||||||
handshakeTimer *time.Timer
|
handshakeTimer *time.Timer
|
||||||
@@ -135,6 +139,7 @@ func New(opts Options) (*Broker, error) {
|
|||||||
if log == nil {
|
if log == nil {
|
||||||
log = slog.Default()
|
log = slog.Default()
|
||||||
}
|
}
|
||||||
|
log = slog.New(newRedactHandler(log.Handler()))
|
||||||
|
|
||||||
caps := mqtt.NewDefaultServerCapabilities()
|
caps := mqtt.NewDefaultServerCapabilities()
|
||||||
caps.MaximumClients = maxClients
|
caps.MaximumClients = maxClients
|
||||||
@@ -165,6 +170,8 @@ func New(opts Options) (*Broker, error) {
|
|||||||
metrics: opts.Metrics,
|
metrics: opts.Metrics,
|
||||||
current: make(map[string]*connState),
|
current: make(map[string]*connState),
|
||||||
byClient: make(map[*mqtt.Client]*connState),
|
byClient: make(map[*mqtt.Client]*connState),
|
||||||
|
byConnID: make(map[port.ConnID]*connState),
|
||||||
|
closedCh: make(chan struct{}),
|
||||||
queues: make(map[string]*uplinkQueue),
|
queues: make(map[string]*uplinkQueue),
|
||||||
largeSem: make(chan struct{}, largeFrameSlots),
|
largeSem: make(chan struct{}, largeFrameSlots),
|
||||||
}
|
}
|
||||||
@@ -175,6 +182,7 @@ func New(opts Options) (*Broker, error) {
|
|||||||
if err := srv.Serve(); err != nil {
|
if err := srv.Serve(); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
go b.sweepLoop()
|
||||||
return b, nil
|
return b, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -186,6 +194,11 @@ func (b *Broker) Close() error {
|
|||||||
if b.closed.Swap(true) {
|
if b.closed.Swap(true) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
select {
|
||||||
|
case <-b.closedCh:
|
||||||
|
default:
|
||||||
|
close(b.closedCh)
|
||||||
|
}
|
||||||
b.queuesMu.Lock()
|
b.queuesMu.Lock()
|
||||||
for _, q := range b.queues {
|
for _, q := range b.queues {
|
||||||
q.close()
|
q.close()
|
||||||
@@ -345,11 +358,10 @@ func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
|
|||||||
b.connsMu.RLock()
|
b.connsMu.RLock()
|
||||||
defer b.connsMu.RUnlock()
|
defer b.connsMu.RUnlock()
|
||||||
if connID != "" {
|
if connID != "" {
|
||||||
for _, st := range b.byClient {
|
st := b.byConnID[connID]
|
||||||
if st.endpointID == endpointID && st.connID == connID {
|
if st != nil && st.endpointID == endpointID {
|
||||||
return st
|
return st
|
||||||
}
|
}
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return b.current[endpointID]
|
return b.current[endpointID]
|
||||||
@@ -462,12 +474,46 @@ func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) {
|
|||||||
func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState {
|
func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState {
|
||||||
b.connsMu.RLock()
|
b.connsMu.RLock()
|
||||||
defer b.connsMu.RUnlock()
|
defer b.connsMu.RUnlock()
|
||||||
for _, st := range b.byClient {
|
st := b.byConnID[connID]
|
||||||
if st.endpointID == endpointID && st.connID == connID {
|
if st == nil || st.endpointID != endpointID {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return st
|
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 {
|
func (b *Broker) hasDownSub(st *connState) bool {
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package broker
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
mqtt "github.com/mochi-mqtt/server/v2"
|
mqtt "github.com/mochi-mqtt/server/v2"
|
||||||
@@ -47,12 +48,11 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
|||||||
remoteIP: remoteIP,
|
remoteIP: remoteIP,
|
||||||
client: cl,
|
client: cl,
|
||||||
maxPacketSize: pk.Properties.MaximumPacketSize,
|
maxPacketSize: pk.Properties.MaximumPacketSize,
|
||||||
|
createdAt: time.Now(),
|
||||||
}
|
}
|
||||||
|
|
||||||
// ClientID、Username 都必须等于端编号
|
// ClientID、Username 都必须等于端编号
|
||||||
if clientID == "" || endpointID == "" || clientID != endpointID {
|
if clientID == "" || endpointID == "" || clientID != endpointID {
|
||||||
st.authOK = false
|
|
||||||
h.rememberPending(cl, st)
|
|
||||||
return nil
|
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)
|
res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
st.authErr = err
|
return err // mochi 不回 CONNACK,直接断开;不登记连接表
|
||||||
h.rememberPending(cl, st)
|
|
||||||
return err // mochi 不回 CONNACK,直接断开
|
|
||||||
}
|
}
|
||||||
st.authOK = res.OK
|
if !res.OK {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
st.authOK = true
|
||||||
st.sessionToken = res.SessionToken
|
st.sessionToken = res.SessionToken
|
||||||
h.rememberPending(cl, st)
|
h.rememberPending(cl, st)
|
||||||
return nil
|
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) {
|
func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) {
|
||||||
h.b.connsMu.Lock()
|
h.b.connsMu.Lock()
|
||||||
h.b.byClient[cl] = st
|
h.b.byClient[cl] = st
|
||||||
|
h.b.byConnID[st.connID] = st
|
||||||
h.b.connsMu.Unlock()
|
h.b.connsMu.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,6 +190,7 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
|
|||||||
st := h.b.byClient[cl]
|
st := h.b.byClient[cl]
|
||||||
if st != nil {
|
if st != nil {
|
||||||
h.b.current[st.endpointID] = st
|
h.b.current[st.endpointID] = st
|
||||||
|
st.established = true
|
||||||
}
|
}
|
||||||
h.b.connsMu.Unlock()
|
h.b.connsMu.Unlock()
|
||||||
if st == nil {
|
if st == nil {
|
||||||
@@ -209,6 +212,9 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
|||||||
h.b.connsMu.Lock()
|
h.b.connsMu.Lock()
|
||||||
st := h.b.byClient[cl]
|
st := h.b.byClient[cl]
|
||||||
delete(h.b.byClient, cl)
|
delete(h.b.byClient, cl)
|
||||||
|
if st != nil {
|
||||||
|
delete(h.b.byConnID, st.connID)
|
||||||
|
}
|
||||||
isCurrent := false
|
isCurrent := false
|
||||||
if st != nil && h.b.current[st.endpointID] == st {
|
if st != nil && h.b.current[st.endpointID] == st {
|
||||||
delete(h.b.current, st.endpointID)
|
delete(h.b.current, st.endpointID)
|
||||||
|
|||||||
@@ -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),
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user