feat: 实现登录、会话令牌、握手与顶号
EOF
This commit is contained in:
@@ -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
@@ -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]
|
||||
|
||||
@@ -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:])
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user