From 34e7c2827f0321a18768786a6904af4b3f450211 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 07:25:35 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=9E=E7=8E=B0=E7=99=BB=E5=BD=95?= =?UTF-8?q?=E3=80=81=E4=BC=9A=E8=AF=9D=E4=BB=A4=E7=89=8C=E3=80=81=E6=8F=A1?= =?UTF-8?q?=E6=89=8B=E4=B8=8E=E9=A1=B6=E5=8F=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit EOF --- docs/DEVIATIONS.md | 37 ++ internal/broker/authn.go | 296 +++++++++++++++ internal/broker/broker.go | 130 ++++++- internal/broker/f02_test.go | 734 ++++++++++++++++++++++++++++++++++++ internal/broker/hooks.go | 44 ++- internal/broker/session.go | 363 ++++++++++++++++++ 6 files changed, 1590 insertions(+), 14 deletions(-) create mode 100644 internal/broker/authn.go create mode 100644 internal/broker/f02_test.go create mode 100644 internal/broker/session.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 3a99bc7..47a963c 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -276,6 +276,43 @@ - 备选方案:N2 暴露回调给 M 注册。 - 影响:接线后 M 需订阅或包装该钩子;当前接口可后续加 `OnPublishDropped` 回调字段。 +### N3 2026-09-30 + +1. **仍未接线 `cmd/nixmsg`** + - 原条款:serve 最终应挂上真实 Authenticator / Session。 + - 实际做法:交付 `broker.Login`、`broker.Session` 与 F02 测试;不改 `cmd/nixmsg`/`wire.go`。 + - 原因:与总控/其他线并行改 wire 冲突;N1/N2 已约定合并时接线。 + - 备选方案:本分支改 wire(与隔离指令冲突)。 + - 影响:进程默认仍 RejectAuthenticator,需接线注入 `Login`+`Session`。 + +2. **`session_hash` 存十六进制文本** + - 原条款:库中存 SHA-256;列为 TEXT,未规定编码。 + - 实际做法:存 32 字节哈希的小写 hex(与后台 API 令牌存法一致)。 + - 原因:TEXT 列无法直接存原始字节;hex 便于排查。 + - 备选方案:BLOB 列或 base64。 + - 影响:其他线读写 `session_hash` 需按 hex 编解码。 + +3. **上下线通知走 `PresenceSink` + port 回调** + - 原条款:写 `online_since`/`offline_since` 并通知;通过现有 port 接口供身份线订阅。 + - 实际做法:N3 自己写时间戳;可选注入 `PresenceSink`(对齐 `presence.Service.SetOnline/SetOffline`);并继续调用 `OnHandshakeComplete`/`OnDisconnect`。旧连接断开用连接代号判断,只有当时仍是 current 才标离线。 + - 原因:I3 尚未合入,不能依赖具体 presence 实现;双通道便于接线。 + - 备选方案:只靠 port、由 I 线写库(与「N3 写 online_since」字面不符)。 + - 影响:接线时避免 I 线重复写同一时间戳即可。 + +4. **InlineClient 的 `OnPublish` 必须放行** + - 原条款:客户端上行 `OnPublish` 返回 `CodeSuccessIgnore`。 + - 实际做法:`cl.Net.Inline` 时原样返回,不 Ignore,否则 `PublishDown` 无法送达订阅者。 + - 原因:mochi `Publish` 经 InlineClient `InjectPacket` 再进 `OnPublish`。 + - 备选方案:不用 InlineClient,改直接 `publishToClient`(偏离文档装配)。 + - 影响:N2 既有 PublishDown 测试此前未读回包,此缺陷在 N3 才暴露并修复。 + +5. **管理员踢线类入口挂在 `Session`** + - 原条款:停用/删除/重置密码先 fatal 再断开;踢下线只断开。 + - 实际做法:`Session.Disable`/`Deleted`/`ResetPassword`/`Kick` 可调用;管理 HTTP 未接。 + - 原因:A2 管理接口尚未接线。 + - 备选方案:放到 `internal/admin`(超出 N 目录)。 + - 影响:A/I 接线时调用这些方法即可。 + ## 消息 M ### M1 2026-09-30 diff --git a/internal/broker/authn.go b/internal/broker/authn.go new file mode 100644 index 0000000..29de669 --- /dev/null +++ b/internal/broker/authn.go @@ -0,0 +1,296 @@ +package broker + +import ( + "context" + "database/sql" + "encoding/hex" + "errors" + "strings" + "sync" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +// ErrEndpointNotFound 编号不存在(业务拒绝,非内部故障)。 +var ErrEndpointNotFound = errors.New("broker: endpoint not found") + +// ErrEndpointDisabled 端已停用(业务拒绝)。 +var ErrEndpointDisabled = errors.New("broker: endpoint disabled") + +// Login 实现 Authenticator:会话令牌或登录密码(含锁定)。 +type Login struct { + DB *store.DB + Pool auth.HashPool + Tokens auth.SessionTokens + Locks auth.LoginLocks + IdleDays int + Now func() time.Time + + usedMu sync.Mutex + // 内存中的 session_used_at(毫秒)与上次落库时间。 + usedAt map[string]int64 + lastFlush map[string]int64 +} + +// LoginOptions 装配 Login。 +type LoginOptions struct { + DB *store.DB + Pool auth.HashPool + Tokens auth.SessionTokens + Locks auth.LoginLocks + IdleDays int + Now func() time.Time +} + +// NewLogin 创建登录校验器。 +func NewLogin(opts LoginOptions) *Login { + now := opts.Now + if now == nil { + now = time.Now + } + tokens := opts.Tokens + if tokens == nil { + tokens = auth.NewSessionTokens() + } + locks := opts.Locks + if locks == nil { + locks = auth.NewLoginLocks() + } + return &Login{ + DB: opts.DB, + Pool: opts.Pool, + Tokens: tokens, + Locks: locks, + IdleDays: opts.IdleDays, + Now: now, + usedAt: make(map[string]int64), + lastFlush: make(map[string]int64), + } +} + +// Authenticate 按 DEVELOPMENT 第 5 节校验;内部故障返回 error。 +func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) { + if l == nil || l.DB == nil { + return AuthResult{}, errors.New("broker: login not configured") + } + if endpointID == "" { + return AuthResult{OK: false}, nil + } + + row, err := l.loadEndpoint(ctx, endpointID) + if err != nil { + if errors.Is(err, ErrEndpointNotFound) || errors.Is(err, ErrEndpointDisabled) { + return AuthResult{OK: false}, nil + } + return AuthResult{}, err + } + + cred := string(password) + if l.Tokens.LooksLikeSessionToken(cred) { + ok, authErr := l.authSession(ctx, endpointID, cred, row) + if authErr != nil { + return AuthResult{}, authErr + } + return AuthResult{OK: ok}, nil + } + + ok, token, authErr := l.authPassword(ctx, endpointID, cred, remoteIP, row) + if authErr != nil { + return AuthResult{}, authErr + } + return AuthResult{OK: ok, SessionToken: token}, nil +} + +type endpointAuthRow struct { + loginHash string + sessionHash []byte // 原始 32 字节;无令牌时 nil + sessionUsedAt int64 // 毫秒;无则 0 +} + +func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) { + var ( + loginHash string + enabled int + sessHex sql.NullString + usedAt sql.NullInt64 + ) + err := l.DB.Read.QueryRowContext(ctx, ` +SELECT login_hash, enabled, session_hash, session_used_at +FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return endpointAuthRow{}, ErrEndpointNotFound + } + return endpointAuthRow{}, err + } + if enabled == 0 { + return endpointAuthRow{}, ErrEndpointDisabled + } + row := endpointAuthRow{loginHash: loginHash} + if usedAt.Valid { + row.sessionUsedAt = usedAt.Int64 + } + if sessHex.Valid && sessHex.String != "" { + raw, decErr := hex.DecodeString(sessHex.String) + if decErr != nil || len(raw) != 32 { + // 损坏的哈希视为无有效会话(令牌校验失败),不是内部故障 + row.sessionHash = nil + } else { + row.sessionHash = raw + } + } + return row, nil +} + +func (l *Login) authSession(ctx context.Context, endpointID, token string, row endpointAuthRow) (bool, error) { + if len(row.sessionHash) == 0 { + return false, nil + } + got := l.Tokens.HashToken(token) + if !auth.EqualHash(got, row.sessionHash) { + return false, nil + } + now := l.Now() + nowMs := now.UnixMilli() + usedAt := row.sessionUsedAt + l.usedMu.Lock() + if mem, ok := l.usedAt[endpointID]; ok && mem > usedAt { + usedAt = mem + } + l.usedMu.Unlock() + if l.IdleDays > 0 { + idle := time.Duration(l.IdleDays) * 24 * time.Hour + if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle { + return false, nil + } + } + if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil { + return false, err + } + return true, nil +} + +func (l *Login) touchSessionUsed(ctx context.Context, endpointID string, nowMs int64) error { + const flushEvery = int64(time.Hour / time.Millisecond) + l.usedMu.Lock() + l.usedAt[endpointID] = nowMs + last := l.lastFlush[endpointID] + needFlush := last == 0 || nowMs-last >= flushEvery + if needFlush { + l.lastFlush[endpointID] = nowMs + } + l.usedMu.Unlock() + if !needFlush { + return nil + } + return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { + _, err := tx.Exec(`UPDATE endpoints SET session_used_at = ? WHERE id = ? AND session_hash IS NOT NULL AND session_hash != ''`, + nowMs, endpointID) + return err + }) +} + +func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP string, row endpointAuthRow) (ok bool, token string, err error) { + ipKey := auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP} + epKey := auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID} + if locked, _ := l.Locks.Check(ipKey); locked { + return false, "", nil + } + if locked, _ := l.Locks.Check(epKey); locked { + return false, "", nil + } + if l.Pool == nil { + return false, "", errors.New("broker: password pool not configured") + } + match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash) + if verErr != nil { + return false, "", verErr + } + if !match { + l.Locks.Fail(ipKey) + l.Locks.Fail(epKey) + return false, "", nil + } + + tok, hash, issErr := l.Tokens.Issue(ctx) + if issErr != nil { + return false, "", issErr + } + nowMs := l.Now().UnixMilli() + hashHex := hex.EncodeToString(hash) + writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { + _, e := tx.Exec(` +UPDATE endpoints +SET session_hash = ?, session_issued_at = ?, session_used_at = ? +WHERE id = ?`, hashHex, nowMs, nowMs, endpointID) + return e + }) + if writeErr != nil { + return false, "", writeErr + } + l.usedMu.Lock() + l.usedAt[endpointID] = nowMs + l.lastFlush[endpointID] = nowMs + l.usedMu.Unlock() + return true, tok, nil +} + +// ClearSession 清空会话令牌(logout / 停用 / 删除 / 重置密码)。 +func (l *Login) ClearSession(ctx context.Context, endpointID string) error { + if l == nil || l.DB == nil { + return errors.New("broker: login not configured") + } + err := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { + _, e := tx.Exec(` +UPDATE endpoints +SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL +WHERE id = ?`, endpointID) + return e + }) + if err != nil { + return err + } + l.usedMu.Lock() + delete(l.usedAt, endpointID) + delete(l.lastFlush, endpointID) + l.usedMu.Unlock() + return nil +} + +// SetOnlineSince 握手完成时写入 online_since。 +func (l *Login) SetOnlineSince(ctx context.Context, endpointID string, atMs int64) error { + return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { + _, err := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, atMs, endpointID) + return err + }) +} + +// SetOfflineSince 当前连接断开时写入 offline_since。 +func (l *Login) SetOfflineSince(ctx context.Context, endpointID string, atMs int64) error { + return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { + _, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, atMs, endpointID) + return err + }) +} + +// SessionHashOf 返回当前库中的会话哈希(测试用);无则 nil。 +func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, error) { + var sessHex sql.NullString + err := l.DB.Read.QueryRowContext(ctx, `SELECT session_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&sessHex) + if err != nil { + return nil, err + } + if !sessHex.Valid || sessHex.String == "" { + return nil, nil + } + return hex.DecodeString(sessHex.String) +} + +// LooksLikeSessionToken 暴露给测试。 +func (l *Login) LooksLikeSessionToken(s string) bool { + return strings.HasPrefix(s, "nst_") +} + +var _ Authenticator = (*Login)(nil) diff --git a/internal/broker/broker.go b/internal/broker/broker.go index 41e103b..bc1fd8f 100644 --- a/internal/broker/broker.go +++ b/internal/broker/broker.go @@ -9,6 +9,7 @@ import ( "net" "sync" "sync/atomic" + "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" mqtt "github.com/mochi-mqtt/server/v2" @@ -85,18 +86,22 @@ type Broker struct { } type connState struct { - connID port.ConnID - endpointID string - transport port.Transport - remoteIP string - client *mqtt.Client - maxPacketSize uint32 - maxRecvBytes int - authOK bool - authErr error - sessionToken string - largeHeld int - mu sync.Mutex + connID port.ConnID + endpointID string + transport port.Transport + remoteIP string + client *mqtt.Client + maxPacketSize uint32 + maxRecvBytes int + authOK bool + authErr error + sessionToken string + handshook bool + subscribedDown bool + largeHeld int + mu sync.Mutex + + handshakeTimer *time.Timer } // New 创建并 Serve mochi(无监听器)。 @@ -265,7 +270,12 @@ func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.Co case port.DisconnectKicked, port.DisconnectFatal: code = packets.ErrAdministrativeAction } - return b.server.DisconnectClient(st.client, code) + err := b.server.DisconnectClient(st.client, code) + // mochi 对错误类原因码会把 Code 当作 error 返回,表示已按该原因断开,不算失败。 + if _, ok := err.(packets.Code); ok { + return nil + } + return err } func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState { @@ -362,6 +372,100 @@ func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) { }, true } +// IsHandshook 当前连接是否已完成握手。 +func (b *Broker) IsHandshook(endpointID string) bool { + b.connsMu.RLock() + st := b.current[endpointID] + b.connsMu.RUnlock() + if st == nil { + return false + } + st.mu.Lock() + defer st.mu.Unlock() + return st.handshook +} + +// CurrentConnID 返回端的当前连接代号。 +func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) { + b.connsMu.RLock() + st := b.current[endpointID] + b.connsMu.RUnlock() + if st == nil { + return "", false + } + return st.connID, true +} + +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 + } + } + return nil +} + +func (b *Broker) hasDownSub(st *connState) bool { + if st == nil { + return false + } + st.mu.Lock() + defer st.mu.Unlock() + if st.subscribedDown { + return true + } + // 回退:直接看 mochi 订阅表 + if st.client != nil && st.client.State.Subscriptions != nil { + _, ok := st.client.State.Subscriptions.Get(downTopic(st.endpointID)) + return ok + } + return false +} + +func (b *Broker) startHandshakeDeadline(endpointID string, connID port.ConnID, d time.Duration) { + st := b.connStateOf(endpointID, connID) + if st == nil { + return + } + st.mu.Lock() + if st.handshook { + st.mu.Unlock() + return + } + if st.handshakeTimer != nil { + st.handshakeTimer.Stop() + } + st.handshakeTimer = time.AfterFunc(d, func() { + cur := b.connStateOf(endpointID, connID) + if cur == nil { + return + } + cur.mu.Lock() + done := cur.handshook + cur.mu.Unlock() + if done { + return + } + _ = b.Disconnect(context.Background(), endpointID, connID, port.DisconnectIdle) + }) + st.mu.Unlock() +} + +func (b *Broker) cancelHandshakeDeadline(endpointID string, connID port.ConnID) { + st := b.connStateOf(endpointID, connID) + if st == nil { + return + } + st.mu.Lock() + if st.handshakeTimer != nil { + st.handshakeTimer.Stop() + st.handshakeTimer = nil + } + st.mu.Unlock() +} + func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) { b.queuesMu.Lock() q, ok := b.queues[endpointID] diff --git a/internal/broker/f02_test.go b/internal/broker/f02_test.go new file mode 100644 index 0000000..cd2425b --- /dev/null +++ b/internal/broker/f02_test.go @@ -0,0 +1,734 @@ +package broker_test + +import ( + "bytes" + "context" + "database/sql" + "encoding/json" + "io" + "net" + "sync" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/broker" + "git.asio.asia/nixevol/NixMsg/internal/protocol" + "git.asio.asia/nixevol/NixMsg/internal/store" + "github.com/mochi-mqtt/server/v2/packets" +) + +type presenceRec struct { + mu sync.Mutex + online []string + offline []string +} + +func (p *presenceRec) SetOnline(_ context.Context, endpointID string, _ port.ConnID, _ int64) error { + p.mu.Lock() + defer p.mu.Unlock() + p.online = append(p.online, endpointID) + return nil +} + +func (p *presenceRec) SetOffline(_ context.Context, endpointID string, _ port.ConnID, _ int64) error { + p.mu.Lock() + defer p.mu.Unlock() + p.offline = append(p.offline, endpointID) + return nil +} + +type uplinkRec struct { + port.StubUplinkHandler + mu sync.Mutex + handshakes int + disconnects []port.DisconnectReason +} + +func (u *uplinkRec) OnHandshakeComplete(context.Context, port.HandshakeInfo) error { + u.mu.Lock() + defer u.mu.Unlock() + u.handshakes++ + return nil +} + +func (u *uplinkRec) OnDisconnect(_ context.Context, _ port.ConnInfo, reason port.DisconnectReason) { + u.mu.Lock() + defer u.mu.Unlock() + u.disconnects = append(u.disconnects, reason) +} + +type testEnv struct { + t *testing.T + db *store.DB + login *broker.Login + sess *broker.Session + b *broker.Broker + presence *presenceRec + uplink *uplinkRec + pool auth.HashPool + dir string +} + +func openEnv(t *testing.T, idleDays int) *testEnv { + t.Helper() + dir := t.TempDir() + db, err := store.Open(dir, "FULL") + if err != nil { + t.Fatal(err) + } + pool := auth.NewStubHashPool() + locks := auth.NewLoginLocks() + login := broker.NewLogin(broker.LoginOptions{ + DB: db, + Pool: pool, + Tokens: auth.NewSessionTokens(), + Locks: locks, + IdleDays: idleDays, + }) + pres := &presenceRec{} + up := &uplinkRec{} + sess := broker.NewSession(broker.SessionOptions{ + Login: login, + Inner: up, + Presence: pres, + Limits: broker.HelloLimits{ServerVersion: "0.1.0-test"}, + }) + b, err := broker.New(broker.Options{Authenticator: login, Uplink: sess}) + if err != nil { + t.Fatal(err) + } + sess.Attach(b) + t.Cleanup(func() { + _ = b.Close() + _ = db.Close() + }) + return &testEnv{t: t, db: db, login: login, sess: sess, b: b, presence: pres, uplink: up, pool: pool, dir: dir} +} + +func (e *testEnv) insertEndpoint(id, password string) { + e.t.Helper() + phc, err := e.pool.Hash(context.Background(), auth.PasswordLogin, password) + if err != nil { + e.t.Fatal(err) + } + now := time.Now().UnixMilli() + err = e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, execErr := tx.Exec(` +INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) +VALUES (?, '', ?, NULL, 0, 0, 1, ?)`, id, phc, now) + return execErr + }) + if err != nil { + e.t.Fatal(err) + } +} + +type pipeClient struct { + t *testing.T + conn net.Conn + done chan struct{} + packet uint16 +} + +func (e *testEnv) dial() *pipeClient { + e.t.Helper() + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = e.b.AttachTCP(r) + }() + return &pipeClient{t: e.t, conn: w, done: done, packet: 1} +} + +func (c *pipeClient) close() { + _ = c.conn.Close() + select { + case <-c.done: + case <-time.After(3 * time.Second): + } +} + +func (c *pipeClient) connect(endpoint, password string, maxPacket uint32) (connack byte, ok bool) { + c.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), + }, + Properties: packets.Properties{MaximumPacketSize: maxPacket}, + } + var buf bytes.Buffer + if err := pk.ConnectEncode(&buf); err != nil { + c.t.Fatal(err) + } + if _, err := c.conn.Write(buf.Bytes()); err != nil { + c.t.Fatal(err) + } + _ = c.conn.SetReadDeadline(time.Now().Add(3 * time.Second)) + raw := make([]byte, 256) + n, err := io.ReadAtLeast(c.conn, raw, 2) + if err != nil { + return 0, false + } + if raw[0]>>4 != packets.Connack { + c.t.Fatalf("want connack got %x", raw[:n]) + } + // MQTT5 CONNACK: type, remaining len, flags, reason + reason := byte(0) + if n >= 4 { + reason = raw[3] + } + return reason, reason == 0 +} + +func (c *pipeClient) expectNoConnack() { + c.t.Helper() + _ = c.conn.SetReadDeadline(time.Now().Add(400 * time.Millisecond)) + buf := make([]byte, 64) + n, err := c.conn.Read(buf) + if err == nil && n > 0 && buf[0]>>4 == packets.Connack { + c.t.Fatalf("unexpected connack %x", buf[:n]) + } +} + +func (c *pipeClient) subscribe(endpoint string) { + c.t.Helper() + c.packet++ + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1}, + ProtocolVersion: 5, + PacketID: c.packet, + Filters: packets.Subscriptions{ + {Filter: "nix/c/" + endpoint + "/down", Qos: 1}, + }, + } + var buf bytes.Buffer + if err := pk.SubscribeEncode(&buf); err != nil { + c.t.Fatal(err) + } + if _, err := c.conn.Write(buf.Bytes()); err != nil { + c.t.Fatal(err) + } + _ = c.conn.SetReadDeadline(time.Now().Add(3 * time.Second)) + raw := make([]byte, 256) + n, err := io.ReadAtLeast(c.conn, raw, 2) + if err != nil { + c.t.Fatal(err) + } + if raw[0]>>4 != packets.Suback { + c.t.Fatalf("want suback got %x", raw[:n]) + } +} + +func (c *pipeClient) publishUp(endpoint string, payload []byte) { + c.t.Helper() + c.packet++ + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1}, + ProtocolVersion: 5, + TopicName: "nix/c/" + endpoint + "/up", + PacketID: c.packet, + Payload: payload, + } + var buf bytes.Buffer + if err := pk.PublishEncode(&buf); err != nil { + c.t.Fatal(err) + } + if _, err := c.conn.Write(buf.Bytes()); err != nil { + c.t.Fatal(err) + } +} + +func (c *pipeClient) readDownJSON(timeout time.Duration) map[string]any { + c.t.Helper() + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + _ = c.conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond)) + hdr := make([]byte, 1) + if _, err := io.ReadFull(c.conn, hdr); err != nil { + continue + } + rem, err := readRemainingLength(c.conn) + if err != nil { + continue + } + body := make([]byte, rem) + if _, err := io.ReadFull(c.conn, body); err != nil { + continue + } + typ := hdr[0] >> 4 + switch typ { + case packets.Publish: + pk := new(packets.Packet) + pk.ProtocolVersion = 5 + pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem} + fhQos := (hdr[0] >> 1) & 0x3 + pk.FixedHeader.Qos = fhQos + if err := pk.PublishDecode(body); err != nil { + c.t.Fatalf("publish decode: %v", err) + } + if fhQos > 0 { + ack := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Puback}, + ProtocolVersion: 5, + PacketID: pk.PacketID, + } + var ab bytes.Buffer + _ = ack.PubackEncode(&ab) + _, _ = c.conn.Write(ab.Bytes()) + } + var m map[string]any + if err := json.Unmarshal(pk.Payload, &m); err != nil { + c.t.Fatalf("json: %v payload=%s", err, pk.Payload) + } + return m + case packets.Puback, packets.Pingresp, packets.Disconnect: + continue + default: + continue + } + } + c.t.Fatal("timeout waiting down json") + return nil +} + +func readRemainingLength(r io.Reader) (int, error) { + var mul uint32 = 1 + var value uint32 + for i := 0; i < 4; i++ { + var b [1]byte + if _, err := io.ReadFull(r, b[:]); err != nil { + return 0, err + } + value += uint32(b[0]&127) * mul + if b[0]&128 == 0 { + return int(value), nil + } + mul *= 128 + } + return 0, io.ErrUnexpectedEOF +} + +func helloPayload(rid string) []byte { + b, _ := protocol.Marshal(protocol.Hello{ + V: protocol.Version, Type: protocol.TypeHello, RID: rid, + }) + return b +} + +func waitHandshook(t *testing.T, b *broker.Broker, endpoint string) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if b.IsHandshook(endpoint) { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatal("not handshook") +} + +func TestF02PasswordLoginReturnsTokenAndHandshake(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep1", "password1") + c := e.dial() + defer c.close() + reason, ok := c.connect("ep1", "password1", 0) + if !ok { + t.Fatalf("connack reason=%d", reason) + } + c.subscribe("ep1") + c.publishUp("ep1", helloPayload("1")) + m := c.readDownJSON(3 * time.Second) + if m["type"] != "resp" || m["ok"] != true { + t.Fatalf("hello resp=%v", m) + } + data, _ := m["data"].(map[string]any) + tok, _ := data["session_token"].(string) + if tok == "" || tok[:4] != "nst_" { + t.Fatalf("session_token=%v", data["session_token"]) + } + waitHandshook(t, e.b, "ep1") + e.presence.mu.Lock() + nOnline := len(e.presence.online) + e.presence.mu.Unlock() + if nOnline < 1 { + t.Fatal("expected presence online") + } +} + +func TestF02TakenOverByPasswordLogin(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep2", "password1") + + a := e.dial() + defer a.close() + if _, ok := a.connect("ep2", "password1", 0); !ok { + t.Fatal("A connect") + } + a.subscribe("ep2") + a.publishUp("ep2", helloPayload("1")) + _ = a.readDownJSON(3 * time.Second) + waitHandshook(t, e.b, "ep2") + infoA, _ := e.b.ConnInfoOf("ep2") + + // 后台排空 A,避免顶号写 DISCONNECT 时 pipe 阻塞 + go func() { + buf := make([]byte, 512) + for { + _ = a.conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, err := a.conn.Read(buf) + if err != nil { + return + } + } + }() + + b := e.dial() + defer b.close() + if _, ok := b.connect("ep2", "password1", 0); !ok { + t.Fatal("B connect") + } + b.subscribe("ep2") + b.publishUp("ep2", helloPayload("2")) + m := b.readDownJSON(3 * time.Second) + data, _ := m["data"].(map[string]any) + tokB, _ := data["session_token"].(string) + if tokB == "" { + t.Fatal("B should get new token") + } + waitHandshook(t, e.b, "ep2") + infoB, ok := e.b.ConnInfoOf("ep2") + if !ok || infoB.ConnID == infoA.ConnID { + t.Fatalf("current should be B, got %+v old=%s", infoB, infoA.ConnID) + } +} + +func TestF02OldTokenRejectedAfterPasswordLogin(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep3", "password1") + + a := e.dial() + if _, ok := a.connect("ep3", "password1", 0); !ok { + t.Fatal("A") + } + a.subscribe("ep3") + a.publishUp("ep3", helloPayload("1")) + m := a.readDownJSON(3 * time.Second) + data, _ := m["data"].(map[string]any) + oldTok, _ := data["session_token"].(string) + a.close() + + // 另一处密码登录换令牌 + b := e.dial() + if _, ok := b.connect("ep3", "password1", 0); !ok { + t.Fatal("B") + } + b.subscribe("ep3") + b.publishUp("ep3", helloPayload("2")) + _ = b.readDownJSON(3 * time.Second) + b.close() + + c := e.dial() + defer c.close() + reason, ok := c.connect("ep3", oldTok, 0) + if ok { + t.Fatal("old token should fail") + } + if reason != 0x86 { + t.Fatalf("want 0x86 got %#x", reason) + } +} + +func TestF02TokenReconnectDifferentIPKeepsToken(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep4", "password1") + + a := e.dial() + if _, ok := a.connect("ep4", "password1", 0); !ok { + t.Fatal("A") + } + a.subscribe("ep4") + a.publishUp("ep4", helloPayload("1")) + m := a.readDownJSON(3 * time.Second) + data, _ := m["data"].(map[string]any) + tok, _ := data["session_token"].(string) + hash1, err := e.login.SessionHashOf(context.Background(), "ep4") + if err != nil || hash1 == nil { + t.Fatalf("hash1=%v err=%v", hash1, err) + } + a.close() + + b := e.dial() + defer b.close() + if _, ok := b.connect("ep4", tok, 0); !ok { + t.Fatal("token reconnect") + } + b.subscribe("ep4") + b.publishUp("ep4", helloPayload("2")) + m2 := b.readDownJSON(3 * time.Second) + data2, _ := m2["data"].(map[string]any) + if _, has := data2["session_token"]; has { + t.Fatalf("token reconnect must not return session_token: %v", data2) + } + hash2, _ := e.login.SessionHashOf(context.Background(), "ep4") + if !auth.EqualHash(hash1, hash2) { + t.Fatal("session hash changed on token reconnect") + } +} + +func TestF02IPLockDoesNotAffectOtherIP(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep5", "password1") + login := e.login + for i := 0; i < 10; i++ { + res, err := login.Authenticate(context.Background(), "ep5", []byte("wrong-pass"), "1.1.1.1") + if err != nil { + t.Fatal(err) + } + if res.OK { + t.Fatal("should fail") + } + } + res, err := login.Authenticate(context.Background(), "ep5", []byte("password1"), "1.1.1.1") + if err != nil || res.OK { + t.Fatalf("locked same IP ok=%v err=%v", res.OK, err) + } + res, err = login.Authenticate(context.Background(), "ep5", []byte("password1"), "2.2.2.2") + if err != nil || !res.OK { + t.Fatalf("other IP ok=%v err=%v", res.OK, err) + } +} + +func TestF02EndpointLockAllowsTokenReconnect(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep6", "password1") + login := e.login + + // 先拿到令牌 + res, err := login.Authenticate(context.Background(), "ep6", []byte("password1"), "10.0.0.1") + if err != nil || !res.OK || res.SessionToken == "" { + t.Fatalf("login=%+v err=%v", res, err) + } + tok := res.SessionToken + + // 多 IP 累计 50 次失败 + for i := 0; i < 50; i++ { + ip := "203.0.113." + itoa(i%250+1) + r, e2 := login.Authenticate(context.Background(), "ep6", []byte("bad"), ip) + if e2 != nil { + t.Fatal(e2) + } + if r.OK { + t.Fatal("unexpected ok") + } + } + // 密码登录暂停 + r, err := login.Authenticate(context.Background(), "ep6", []byte("password1"), "198.51.100.1") + if err != nil || r.OK { + t.Fatalf("password should be locked ok=%v err=%v", r.OK, err) + } + // 令牌仍可 + r, err = login.Authenticate(context.Background(), "ep6", []byte(tok), "198.51.100.9") + if err != nil || !r.OK { + t.Fatalf("token should work ok=%v err=%v", r.OK, err) + } + if r.SessionToken != "" { + t.Fatal("token auth must not issue new token") + } +} + +func TestF02DBErrorClosesWithout086(t *testing.T) { + dir := t.TempDir() + db, err := store.Open(dir, "FULL") + if err != nil { + t.Fatal(err) + } + pool := auth.NewStubHashPool() + login := broker.NewLogin(broker.LoginOptions{ + DB: db, Pool: pool, Tokens: auth.NewSessionTokens(), Locks: auth.NewLoginLocks(), IdleDays: 30, + }) + phc, _ := pool.Hash(context.Background(), auth.PasswordLogin, "password1") + _ = db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(`INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) +VALUES ('ep7', '', ?, NULL, 0, 0, 1, ?)`, phc, time.Now().UnixMilli()) + return e + }) + _ = db.Read.Close() + + b, err := broker.New(broker.Options{Authenticator: login}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + + r, w := net.Pipe() + errCh := make(chan error, 1) + go func() { errCh <- b.AttachTCP(r) }() + c := &pipeClient{t: t, conn: w, done: make(chan struct{}), packet: 1} + c.expectNoConnack() + _ = w.Close() + select { + case <-errCh: + case <-time.After(2 * time.Second): + } + _ = db.Close() +} + +func TestF02NotReadyBeforeHello(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep8", "password1") + c := e.dial() + defer c.close() + if _, ok := c.connect("ep8", "password1", 0); !ok { + t.Fatal("connect") + } + c.subscribe("ep8") + payload, _ := protocol.Marshal(map[string]any{ + "v": 1, "type": "self.get", "rid": "9", + }) + c.publishUp("ep8", payload) + m := c.readDownJSON(3 * time.Second) + if m["ok"] != false { + t.Fatalf("want not_ready resp got %v", m) + } + errObj, _ := m["error"].(map[string]any) + if errObj["code"] != protocol.CodeNotReady { + t.Fatalf("code=%v", errObj) + } +} + +func TestF02LogoutClearsToken(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep9", "password1") + c := e.dial() + defer c.close() + if _, ok := c.connect("ep9", "password1", 0); !ok { + t.Fatal("connect") + } + c.subscribe("ep9") + c.publishUp("ep9", helloPayload("1")) + m := c.readDownJSON(3 * time.Second) + data, _ := m["data"].(map[string]any) + tok, _ := data["session_token"].(string) + waitHandshook(t, e.b, "ep9") + + logout, _ := protocol.Marshal(protocol.SelfLogout{V: protocol.Version, Type: protocol.TypeSelfLogout, RID: "24"}) + c.publishUp("ep9", logout) + m2 := c.readDownJSON(3 * time.Second) + if m2["ok"] != true { + t.Fatalf("logout resp=%v", m2) + } + + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + h, _ := e.login.SessionHashOf(context.Background(), "ep9") + if h == nil { + break + } + time.Sleep(20 * time.Millisecond) + } + h, _ := e.login.SessionHashOf(context.Background(), "ep9") + if h != nil { + t.Fatal("session should be cleared") + } + + c2 := e.dial() + defer c2.close() + if _, ok := c2.connect("ep9", tok, 0); ok { + t.Fatal("token after logout should fail") + } +} + +func TestF02AdminResetPasswordFatal(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep10", "password1") + c := e.dial() + defer c.close() + if _, ok := c.connect("ep10", "password1", 0); !ok { + t.Fatal("connect") + } + c.subscribe("ep10") + c.publishUp("ep10", helloPayload("1")) + m := c.readDownJSON(3 * time.Second) + data, _ := m["data"].(map[string]any) + tok, _ := data["session_token"].(string) + waitHandshook(t, e.b, "ep10") + + if err := e.sess.ResetPassword(context.Background(), "ep10"); err != nil { + t.Fatal(err) + } + fatal := c.readDownJSON(3 * time.Second) + if fatal["type"] != "fatal" || fatal["reason"] != "password_reset" { + t.Fatalf("fatal=%v", fatal) + } + + c2 := e.dial() + defer c2.close() + if _, ok := c2.connect("ep10", tok, 0); ok { + t.Fatal("token after reset should fail") + } +} + +func TestF02KickKeepsToken(t *testing.T) { + e := openEnv(t, 30) + e.insertEndpoint("ep11", "password1") + c := e.dial() + defer c.close() + if _, ok := c.connect("ep11", "password1", 0); !ok { + t.Fatal("connect") + } + c.subscribe("ep11") + c.publishUp("ep11", helloPayload("1")) + m := c.readDownJSON(3 * time.Second) + data, _ := m["data"].(map[string]any) + tok, _ := data["session_token"].(string) + waitHandshook(t, e.b, "ep11") + + go func() { + buf := make([]byte, 512) + for { + _ = c.conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, err := c.conn.Read(buf) + if err != nil { + return + } + } + }() + + if err := e.sess.Kick(context.Background(), "ep11"); err != nil { + t.Fatal(err) + } + time.Sleep(100 * time.Millisecond) + + c2 := e.dial() + defer c2.close() + if _, ok := c2.connect("ep11", tok, 0); !ok { + t.Fatal("token should still work after kick") + } +} + +func itoa(n int) string { + if n == 0 { + return "0" + } + var b [16]byte + i := len(b) + for n > 0 { + i-- + b[i] = byte('0' + n%10) + n /= 10 + } + return string(b[i:]) +} diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go index 1becf62..58b4b57 100644 --- a/internal/broker/hooks.go +++ b/internal/broker/hooks.go @@ -26,13 +26,15 @@ func (h *nixHook) Provides(b byte) bool { mqtt.OnSessionEstablished, mqtt.OnDisconnect, mqtt.OnQosComplete, + mqtt.OnSubscribed, }, []byte{b}) } func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { endpointID := string(pk.Connect.Username) + clientID := pk.Connect.ClientIdentifier if endpointID == "" { - endpointID = pk.Connect.ClientIdentifier + endpointID = clientID } remoteIP := remoteIPOf(cl) @@ -45,6 +47,13 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { maxPacketSize: pk.Properties.MaximumPacketSize, } + // ClientID、Username 都必须等于端编号 + if clientID == "" || endpointID == "" || clientID != endpointID { + st.authOK = false + h.rememberPending(cl, st) + return nil + } + // 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive ka := pk.Connect.Keepalive if ka < keepaliveMin || ka > keepaliveMax { @@ -103,6 +112,10 @@ func (h *nixHook) OnACLCheck(cl *mqtt.Client, topic string, write bool) bool { } func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, error) { + // InlineClient 的 PublishDown 走 InjectPacket → OnPublish;必须放行才能分发给订阅者。 + if cl != nil && cl.Net.Inline { + return pk, nil + } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() @@ -122,6 +135,28 @@ func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, return pk, packets.CodeSuccessIgnore } +func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes []byte) { + h.b.connsMu.RLock() + st := h.b.byClient[cl] + h.b.connsMu.RUnlock() + if st == nil { + return + } + down := downTopic(st.endpointID) + for i, sub := range pk.Filters { + if sub.Filter != down { + continue + } + if i < len(reasonCodes) && reasonCodes[i] >= 0x80 { + continue + } + st.mu.Lock() + st.subscribedDown = true + st.mu.Unlock() + return + } +} + func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) } @@ -151,14 +186,17 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { h.b.connsMu.Lock() st := h.b.byClient[cl] delete(h.b.byClient, cl) + isCurrent := false if st != nil && h.b.current[st.endpointID] == st { delete(h.b.current, st.endpointID) + isCurrent = true } h.b.connsMu.Unlock() if st == nil { return } h.b.releaseAllLarge(st) + h.b.cancelHandshakeDeadline(st.endpointID, st.connID) reason := port.DisconnectNormal if err != nil { @@ -179,6 +217,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { SessionToken: st.sessionToken, MaxPacketSize: st.maxPacketSize, } + if sess, ok := h.b.uplink.(*Session); ok { + sess.HandleDisconnect(context.Background(), info, reason, isCurrent) + return + } h.b.uplink.OnDisconnect(context.Background(), info, reason) } diff --git a/internal/broker/session.go b/internal/broker/session.go new file mode 100644 index 0000000..d714751 --- /dev/null +++ b/internal/broker/session.go @@ -0,0 +1,363 @@ +package broker + +import ( + "context" + "encoding/json" + "log/slog" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +const handshakeTimeout = 30 * time.Second + +// PresenceSink 供身份线订阅上下线(与 presence.Service 的 SetOnline/SetOffline 对齐)。 +type PresenceSink interface { + SetOnline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error + SetOffline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error +} + +// HelloLimits 握手响应里的服务器限制。 +type HelloLimits struct { + MaxBodyBytes int + MaxMetaBytes int + MaxFrameBytes int + MaxTTLSeconds int64 + MaxScheduleSeconds int64 + AckTimeoutSeconds int64 + ServerVersion string +} + +// Session 处理握手、logout、上下线落库,并转发其余上行给 Inner。 +type Session struct { + b *Broker + login *Login + inner port.UplinkHandler + presence PresenceSink + limits HelloLimits + log *slog.Logger + now func() time.Time +} + +// SessionOptions 装配 Session。 +type SessionOptions struct { + Login *Login + Inner port.UplinkHandler + Presence PresenceSink + Limits HelloLimits + Logger *slog.Logger + Now func() time.Time +} + +// NewSession 创建会话层;调用 Attach 绑定 Broker 后再接连接。 +func NewSession(opts SessionOptions) *Session { + inner := opts.Inner + if inner == nil { + inner = port.StubUplinkHandler{} + } + log := opts.Logger + if log == nil { + log = slog.Default() + } + now := opts.Now + if now == nil { + now = time.Now + } + lim := opts.Limits + if lim.ServerVersion == "" { + lim.ServerVersion = "0.1.0" + } + if lim.MaxBodyBytes == 0 { + lim.MaxBodyBytes = protocol.DefaultMaxBodyBytes + } + if lim.MaxMetaBytes == 0 { + lim.MaxMetaBytes = protocol.DefaultMaxMetaBytes + } + if lim.MaxFrameBytes == 0 { + lim.MaxFrameBytes = protocol.DefaultMaxFrameBytes + } + if lim.MaxTTLSeconds == 0 { + lim.MaxTTLSeconds = 2592000 + } + if lim.MaxScheduleSeconds == 0 { + lim.MaxScheduleSeconds = 31536000 + } + if lim.AckTimeoutSeconds == 0 { + lim.AckTimeoutSeconds = 300 + } + return &Session{ + login: opts.Login, + inner: inner, + presence: opts.Presence, + limits: lim, + log: log, + now: now, + } +} + +// Attach 绑定 Broker(PublishDown / Disconnect / 连接表)。 +func (s *Session) Attach(b *Broker) { + s.b = b +} + +func (s *Session) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error { + if s.b != nil { + s.b.startHandshakeDeadline(conn.EndpointID, conn.ConnID, handshakeTimeout) + } + return s.inner.OnSessionEstablished(ctx, conn) +} + +func (s *Session) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error { + return s.inner.OnHandshakeComplete(ctx, hs) +} + +func (s *Session) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) { + // 正常路径由 hooks 调 HandleDisconnect(带 isCurrent)。 + // 此方法满足 UplinkHandler;直接调用时按非当前处理,避免误标离线。 + s.HandleDisconnect(ctx, conn, reason, false) +} + +// HandleDisconnect 由 hooks 在确知 isCurrent 后调用(含落库与 presence)。 +func (s *Session) HandleDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason, isCurrent bool) { + if s.b != nil { + s.b.cancelHandshakeDeadline(conn.EndpointID, conn.ConnID) + } + if isCurrent && s.login != nil { + atMs := s.now().UnixMilli() + if err := s.login.SetOfflineSince(ctx, conn.EndpointID, atMs); err != nil { + s.log.Error("set offline_since", "endpoint", conn.EndpointID, "err", err) + } + if s.presence != nil { + if err := s.presence.SetOffline(ctx, conn.EndpointID, conn.ConnID, atMs); err != nil { + s.log.Error("presence offline", "endpoint", conn.EndpointID, "err", err) + } + } + } + s.inner.OnDisconnect(ctx, conn, reason) +} + +func (s *Session) HandleUplink(ctx context.Context, conn port.ConnInfo, payload []byte) error { + if s.b == nil { + return nil + } + st := s.b.connStateOf(conn.EndpointID, conn.ConnID) + if st == nil { + return nil + } + + frame, err := protocol.Decode(payload) + if err != nil { + s.replyErr(ctx, conn, peekRID(payload), protocol.CodeBadRequest, err.Error()) + return nil + } + + st.mu.Lock() + ready := st.handshook + st.mu.Unlock() + + switch f := frame.(type) { + case *protocol.Hello: + return s.handleHello(ctx, conn, st, f) + case *protocol.SelfLogout: + if !ready { + s.replyErr(ctx, conn, f.RID, protocol.CodeNotReady, "handshake required") + return nil + } + return s.handleLogout(ctx, conn, f) + default: + if !ready { + rid := peekRID(payload) + s.replyErr(ctx, conn, rid, protocol.CodeNotReady, "handshake required") + return nil + } + return s.inner.HandleUplink(ctx, conn, payload) + } +} + +func (s *Session) handleHello(ctx context.Context, conn port.ConnInfo, st *connState, hello *protocol.Hello) error { + if err := hello.Validate(); err != nil { + code := protocol.CodeBadRequest + if pe, ok := err.(*protocol.Error); ok { + code = pe.Code + } + s.replyErr(ctx, conn, hello.RID, code, err.Error()) + return nil + } + st.mu.Lock() + if st.handshook { + st.mu.Unlock() + s.replyErr(ctx, conn, hello.RID, protocol.CodeBadRequest, "already handshook") + return nil + } + st.mu.Unlock() + + if !s.b.hasDownSub(st) { + go func() { + _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectIdle) + }() + return nil + } + + maxRecv := 0 + if hello.MaxReceiveBytes != nil { + maxRecv = *hello.MaxReceiveBytes + } + s.b.SetMaxReceiveBytes(conn.EndpointID, conn.ConnID, maxRecv) + + data := protocol.HelloData{ + ServerTimeMs: s.now().UnixMilli(), + ServerVersion: s.limits.ServerVersion, + MaxBodyBytes: s.limits.MaxBodyBytes, + MaxMetaBytes: s.limits.MaxMetaBytes, + MaxFrameBytes: s.limits.MaxFrameBytes, + MaxTTLSeconds: s.limits.MaxTTLSeconds, + MaxScheduleSeconds: s.limits.MaxScheduleSeconds, + AckTimeoutSeconds: s.limits.AckTimeoutSeconds, + } + if conn.SessionToken != "" { + data.SessionToken = conn.SessionToken + } + raw, err := protocol.Marshal(data) + if err != nil { + return err + } + resp := protocol.Resp{ + V: protocol.Version, + Type: protocol.TypeResp, + RID: hello.RID, + OK: true, + Data: raw, + } + if err := s.publishJSON(ctx, conn, resp, 1); err != nil { + return err + } + + atMs := s.now().UnixMilli() + if s.login != nil { + if err := s.login.SetOnlineSince(ctx, conn.EndpointID, atMs); err != nil { + s.log.Error("set online_since", "endpoint", conn.EndpointID, "err", err) + } + } + if s.presence != nil { + if err := s.presence.SetOnline(ctx, conn.EndpointID, conn.ConnID, atMs); err != nil { + s.log.Error("presence online", "endpoint", conn.EndpointID, "err", err) + } + } + + st.mu.Lock() + st.handshook = true + st.mu.Unlock() + s.b.cancelHandshakeDeadline(conn.EndpointID, conn.ConnID) + + hs := port.HandshakeInfo{ + ConnInfo: conn, + MaxReceiveBytes: maxRecv, + Client: hello.Client, + } + return s.inner.OnHandshakeComplete(ctx, hs) +} + +func (s *Session) handleLogout(ctx context.Context, conn port.ConnInfo, req *protocol.SelfLogout) error { + if err := req.Validate(); err != nil { + code := protocol.CodeBadRequest + if pe, ok := err.(*protocol.Error); ok { + code = pe.Code + } + s.replyErr(ctx, conn, req.RID, code, err.Error()) + return nil + } + if s.login != nil { + if err := s.login.ClearSession(ctx, conn.EndpointID); err != nil { + s.replyErr(ctx, conn, req.RID, protocol.CodeBusy, "clear session failed") + return nil + } + } + resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: req.RID, OK: true} + if err := s.publishJSON(ctx, conn, resp, 1); err != nil { + s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", err) + } + go func() { + // 稍等让 QoS1 resp 写入连接,再断开 + time.Sleep(50 * time.Millisecond) + _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal) + }() + return nil +} + +// Kick 只断开当前连接,令牌不变。 +func (s *Session) Kick(ctx context.Context, endpointID string) error { + if s.b == nil { + return ErrNoConnection + } + return s.b.Disconnect(ctx, endpointID, "", port.DisconnectKicked) +} + +// Disable 清空令牌,发 fatal(disabled) 后断开。 +func (s *Session) Disable(ctx context.Context, endpointID string) error { + return s.fatalKick(ctx, endpointID, "disabled") +} + +// Deleted 清空令牌,发 fatal(deleted) 后断开。 +func (s *Session) Deleted(ctx context.Context, endpointID string) error { + return s.fatalKick(ctx, endpointID, "deleted") +} + +// ResetPassword 清空令牌,发 fatal(password_reset) 后断开。 +func (s *Session) ResetPassword(ctx context.Context, endpointID string) error { + return s.fatalKick(ctx, endpointID, "password_reset") +} + +func (s *Session) fatalKick(ctx context.Context, endpointID, reason string) error { + if s.login != nil { + if err := s.login.ClearSession(ctx, endpointID); err != nil { + return err + } + } + if s.b == nil { + return nil + } + info, ok := s.b.ConnInfoOf(endpointID) + if !ok { + return nil + } + fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason} + _ = s.publishJSON(ctx, info, fatal, 1) + go func() { + time.Sleep(20 * time.Millisecond) + _ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal) + }() + return nil +} + +func (s *Session) replyErr(ctx context.Context, conn port.ConnInfo, rid, code, message string) { + if rid == "" { + rid = "0" + } + resp := protocol.Resp{ + V: protocol.Version, + Type: protocol.TypeResp, + RID: rid, + OK: false, + Error: &protocol.ErrorBody{Code: code, Message: message}, + } + _ = s.publishJSON(ctx, conn, resp, 1) +} + +func (s *Session) publishJSON(ctx context.Context, conn port.ConnInfo, v any, qos byte) error { + b, err := protocol.Marshal(v) + if err != nil { + return err + } + return s.b.PublishDown(ctx, conn.EndpointID, conn.ConnID, b, port.PublishOpts{QoS: qos}) +} + +func peekRID(payload []byte) string { + var peek struct { + RID string `json:"rid"` + } + _ = json.Unmarshal(payload, &peek) + return peek.RID +} + +var _ port.UplinkHandler = (*Session)(nil)