feat: 实现登录、会话令牌、握手与顶号

EOF
This commit is contained in:
Nixevol
2026-09-30 07:25:35 +08:00
parent d357082f3d
commit 34e7c2827f
6 changed files with 1590 additions and 14 deletions
+37
View File
@@ -276,6 +276,43 @@
- 备选方案:N2 暴露回调给 M 注册。 - 备选方案:N2 暴露回调给 M 注册。
- 影响:接线后 M 需订阅或包装该钩子;当前接口可后续加 `OnPublishDropped` 回调字段。 - 影响:接线后 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 ## 消息 M
### M1 2026-09-30 ### 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" "net"
"sync" "sync"
"sync/atomic" "sync/atomic"
"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"
@@ -85,18 +86,22 @@ type Broker struct {
} }
type connState struct { type connState struct {
connID port.ConnID connID port.ConnID
endpointID string endpointID string
transport port.Transport transport port.Transport
remoteIP string remoteIP string
client *mqtt.Client client *mqtt.Client
maxPacketSize uint32 maxPacketSize uint32
maxRecvBytes int maxRecvBytes int
authOK bool authOK bool
authErr error authErr error
sessionToken string sessionToken string
largeHeld int handshook bool
mu sync.Mutex subscribedDown bool
largeHeld int
mu sync.Mutex
handshakeTimer *time.Timer
} }
// New 创建并 Serve mochi(无监听器)。 // New 创建并 Serve mochi(无监听器)。
@@ -265,7 +270,12 @@ func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.Co
case port.DisconnectKicked, port.DisconnectFatal: case port.DisconnectKicked, port.DisconnectFatal:
code = packets.ErrAdministrativeAction 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 { func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
@@ -362,6 +372,100 @@ func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) {
}, true }, 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) { func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) {
b.queuesMu.Lock() b.queuesMu.Lock()
q, ok := b.queues[endpointID] 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.OnSessionEstablished,
mqtt.OnDisconnect, mqtt.OnDisconnect,
mqtt.OnQosComplete, mqtt.OnQosComplete,
mqtt.OnSubscribed,
}, []byte{b}) }, []byte{b})
} }
func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
endpointID := string(pk.Connect.Username) endpointID := string(pk.Connect.Username)
clientID := pk.Connect.ClientIdentifier
if endpointID == "" { if endpointID == "" {
endpointID = pk.Connect.ClientIdentifier endpointID = clientID
} }
remoteIP := remoteIPOf(cl) remoteIP := remoteIPOf(cl)
@@ -45,6 +47,13 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
maxPacketSize: pk.Properties.MaximumPacketSize, maxPacketSize: pk.Properties.MaximumPacketSize,
} }
// ClientID、Username 都必须等于端编号
if clientID == "" || endpointID == "" || clientID != endpointID {
st.authOK = false
h.rememberPending(cl, st)
return nil
}
// 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive // 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive
ka := pk.Connect.Keepalive ka := pk.Connect.Keepalive
if ka < keepaliveMin || ka > keepaliveMax { 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) { 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() h.b.connsMu.RLock()
st := h.b.byClient[cl] st := h.b.byClient[cl]
h.b.connsMu.RUnlock() h.b.connsMu.RUnlock()
@@ -122,6 +135,28 @@ func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet,
return pk, packets.CodeSuccessIgnore 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) { 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)) 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() h.b.connsMu.Lock()
st := h.b.byClient[cl] st := h.b.byClient[cl]
delete(h.b.byClient, cl) delete(h.b.byClient, cl)
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)
isCurrent = true
} }
h.b.connsMu.Unlock() h.b.connsMu.Unlock()
if st == nil { if st == nil {
return return
} }
h.b.releaseAllLarge(st) h.b.releaseAllLarge(st)
h.b.cancelHandshakeDeadline(st.endpointID, st.connID)
reason := port.DisconnectNormal reason := port.DisconnectNormal
if err != nil { if err != nil {
@@ -179,6 +217,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
SessionToken: st.sessionToken, SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize, 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) 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)