Compare commits

...
Author SHA1 Message Date
Nixevol 34e7c2827f feat: 实现登录、会话令牌、握手与顶号
EOF
2026-09-30 07:25:35 +08:00
6 changed files with 1590 additions and 14 deletions
+37
View File
@@ -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
+296
View File
@@ -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)
+117 -13
View File
@@ -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]
+734
View File
@@ -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:])
}
+43 -1
View File
@@ -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)
}
+363
View File
@@ -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)