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:]) }