package nixmsg import ( "context" "encoding/json" "errors" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" ) func connectFake(t *testing.T, fake *FakeTransport) *Client { t.Helper() c := New() opts := Options{transport: fake, ConnectTimeout: 5 * time.Second} ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if err := c.Connect(ctx, "ws://example.test/mqtt", "ep1", Credential{Password: "secret"}, opts); err != nil { t.Fatalf("connect: %v", err) } return c } func TestCleanStartEveryConnect(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() if err := fake.SimulateReconnect(); err != nil { t.Fatal(err) } if err := fake.SimulateReconnect(); err != nil { t.Fatal(err) } cs := fake.Connects() if len(cs) < 3 { t.Fatalf("connects=%d want >=3", len(cs)) } for i, c0 := range cs { if !c0.CleanStart { t.Fatalf("connect %d CleanStart=false", i) } if c0.SessionExpiry != 0 { t.Fatalf("connect %d SessionExpiry=%d", i, c0.SessionExpiry) } } } func TestSessionTokenCallback(t *testing.T) { fake := NewFakeTransport() fake.HelloToken = "nst_abc" c := New() var got string c.OnSession(func(tok string) { got = tok }) opts := Options{transport: fake, ConnectTimeout: 5 * time.Second} ctx := context.Background() if err := c.Connect(ctx, "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts); err != nil { t.Fatal(err) } defer c.Close() if got != "nst_abc" { t.Fatalf("token=%q", got) } } func TestDedupReack(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() var calls int var mu sync.Mutex c.OnMessage(func(msg Message) error { mu.Lock() calls++ mu.Unlock() return nil }) // 自动回复 ack go func() { for i := 0; i < 50; i++ { for _, fr := range fake.FindUp("ack") { rid, _ := fr["rid"].(string) if rid != "" { fake.ReplyOK(rid, map[string]any{"result": "accepted"}) } } time.Sleep(5 * time.Millisecond) } }() msg, _ := marshalJSON(map[string]any{ "v": 1, "type": "msg", "id": "m1", "from": "a", "to": map[string]any{"kind": "endpoint", "id": "ep1"}, "body": map[string]any{"enc": "utf8", "data": "hi"}, "send_at_ms": 1, }) fake.InjectDown(msg) fake.InjectDown(msg) // 已交给未确认:忽略 time.Sleep(50 * time.Millisecond) // 等 ack 发出并标记已确认后再推一次 deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if len(fake.FindUp("ack")) >= 1 { break } time.Sleep(10 * time.Millisecond) } time.Sleep(30 * time.Millisecond) fake.InjectDown(msg) // 已确认:再 ack,不回调 time.Sleep(80 * time.Millisecond) mu.Lock() n := calls mu.Unlock() if n != 1 { t.Fatalf("callbacks=%d want 1", n) } acks := fake.FindUp("ack") if len(acks) < 2 { t.Fatalf("acks=%d want >=2 (re-ack)", len(acks)) } } func TestBodyTooLargeLocal(t *testing.T) { fake := NewFakeTransport() fake.MaxBodyBytes = 16 c := connectFake(t, fake) defer c.Close() body := Body{Enc: "utf8", Data: strings.Repeat("x", 64)} _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, body, SendOptions{}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeBodyTooLarge { t.Fatalf("err=%v", err) } } func TestFrameTooLargeLocal(t *testing.T) { fake := NewFakeTransport() fake.MaxBodyBytes = 1 << 20 fake.MaxFrameBytes = 200 c := connectFake(t, fake) defer c.Close() body := Body{Enc: "utf8", Data: strings.Repeat("y", 180)} _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, body, SendOptions{}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeFrameTooLarge { t.Fatalf("err=%v", err) } } func TestResendKeepsIDAndSendAt(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() at := time.UnixMilli(1_700_000_000_000) var firstID string var firstSendAt any var once sync.Once done := make(chan struct{}) go func() { for { select { case <-done: return default: } sends := fake.FindUp("send") if len(sends) == 0 { time.Sleep(5 * time.Millisecond) continue } last := sends[len(sends)-1] rid, _ := last["rid"].(string) once.Do(func() { firstID, _ = last["id"].(string) firstSendAt = last["send_at_ms"] fake.ReplyErr(rid, CodeRateLimited, "slow") }) if len(sends) >= 2 { second := sends[1] if second["id"] != firstID { t.Errorf("id changed %v -> %v", firstID, second["id"]) } if second["send_at_ms"] != firstSendAt { t.Errorf("send_at_ms changed %v -> %v", firstSendAt, second["send_at_ms"]) } fake.ReplyOK(second["rid"].(string), map[string]any{ "id": firstID, "send_at_ms": firstSendAt, "state": "scheduled", }) return } time.Sleep(5 * time.Millisecond) } }() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() res, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "hi"}, SendOptions{SendAt: &at}) close(done) if err != nil { t.Fatal(err) } if res.ID == "" || res.ID != firstID { t.Fatalf("result id=%q first=%q", res.ID, firstID) } } func TestRegisterHTTP(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/api/client/register" { t.Fatalf("path %s", r.URL.Path) } _ = json.NewEncoder(w).Encode(map[string]any{ "ok": true, "data": map[string]any{"id": "e_1", "login_password": "genpass"}, }) })) defer srv.Close() // 从 ws 地址推出 ws := "ws" + strings.TrimPrefix(srv.URL, "http") + "/mqtt" res, err := Register(context.Background(), ws, "code", RegisterOptions{Name: "n"}) if err != nil { t.Fatal(err) } if res.ID != "e_1" || res.LoginPassword != "genpass" { t.Fatalf("%+v", res) } } func TestBuildCleanConnectFlags(t *testing.T) { clean, exp := buildCleanConnectFlags() if !clean || exp != 0 { t.Fatalf("clean=%v exp=%d", clean, exp) } } func TestBackoffNominal(t *testing.T) { if d := computeBackoffDelay(time.Second, 1); d != time.Second { t.Fatal(d) } if d := computeBackoffDelay(time.Second, 6); d != 30*time.Second { t.Fatal(d) } }