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