package nixmsg import ( "context" "encoding/json" "errors" "sync" "testing" "time" ) func TestK00FirstConnectTimeout(t *testing.T) { fake := NewFakeTransport() fake.AutoHello = false c := New() opts := Options{transport: fake, ConnectTimeout: 150 * time.Millisecond} err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeNotConnected { t.Fatalf("err=%v", err) } fake.AutoHello = true if err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts); err != nil { t.Fatalf("reconnect after timeout: %v", err) } c.Close() } func TestK00AuthErrorCodes(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) fake.SimulateAuthFail(AuthBadCredentials) deadline := time.Now().Add(time.Second) for time.Now().Before(deadline) { if c.LastStopCodeForTest() == CodeBadCredentials { break } time.Sleep(5 * time.Millisecond) } if c.LastStopCodeForTest() != CodeBadCredentials { t.Fatalf("stop=%s", c.LastStopCodeForTest()) } _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeBadCredentials { t.Fatalf("send after auth: %v", err) } } func TestK00TakenOverReason(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) var got string c.OnConnection(func(ev ConnectionEvent) { if ev.State == StateKicked { got = ev.Reason } }) fake.SimulateKick() deadline := time.Now().Add(time.Second) for time.Now().Before(deadline) && got != CodeTakenOver { time.Sleep(5 * time.Millisecond) } if got != CodeTakenOver { t.Fatalf("reason=%q", got) } if c.LastStopCodeForTest() != CodeTakenOver { t.Fatalf("stop=%s", c.LastStopCodeForTest()) } } func TestK00Disconnect8BRetryable(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) fake.SimulateServerDisconnect(0x8B) time.Sleep(30 * time.Millisecond) if c.LastStopCodeForTest() == CodeTakenOver { t.Fatal("0x8B should not kick") } if err := fake.SimulateConnectOK(); err != nil { t.Fatal(err) } } func TestK00QueueFull(t *testing.T) { fake := NewFakeTransport() c := New() opts := Options{transport: fake, ConnectTimeout: 5 * time.Second, SendQueueSize: 1} if err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts); err != nil { t.Fatal(err) } defer c.Close() ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) defer cancel() var wg sync.WaitGroup wg.Add(1) go func() { defer wg.Done() _, _ = c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "1"}, SendOptions{}) }() time.Sleep(20 * time.Millisecond) _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "2"}, SendOptions{}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeQueueFull { t.Fatalf("err=%v", err) } cancel() wg.Wait() } func TestK00RequestReturnsData(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() go func() { for i := 0; i < 40; i++ { for _, fr := range fake.FindUp("self.get") { rid, _ := fr["rid"].(string) fake.ReplyOK(rid, map[string]any{"id": "ep1", "name": "n", "default_delay_ms": 0}) } time.Sleep(5 * time.Millisecond) } }() info, err := c.GetSelf(context.Background()) if err != nil { t.Fatal(err) } if info.ID != "ep1" || info.Name != "n" { t.Fatalf("%+v", info) } } func TestK00SendAtAndDelayConflict(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() at := time.UnixMilli(1) d := time.Second _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{SendAt: &at, Delay: &d}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeBadRequest { t.Fatalf("err=%v", err) } } func TestK00SendAfterStopped(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) fake.SimulateKick() time.Sleep(30 * time.Millisecond) _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeTakenOver { t.Fatalf("err=%v", err) } } func TestK00LogoutReturnsError(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) fake.SimulateServerDisconnect(0x8B) time.Sleep(20 * time.Millisecond) err := c.Logout(context.Background()) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeNotConnected { t.Fatalf("logout err=%v", err) } _, err2 := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) if err2 == nil { t.Fatal("expected send fail after logout") } } func TestK00DurationInt64(t *testing.T) { raw := []byte(`{"id":"m1","send_at_ms":123,"state":"scheduled"}`) var sd SendResult if err := json.Unmarshal(raw, &sd); err != nil { t.Fatal(err) } if sd.ID != "m1" || sd.SendAtMs != 123 || sd.State != "scheduled" { t.Fatalf("%+v", sd) } ms := int64(30) * 24 * 3600 * 1000 if ms != 2592000000 { t.Fatal(ms) } } func TestK00MaxReceiveBytesMin(t *testing.T) { fake := NewFakeTransport() c := New() err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, Options{transport: fake, MaxReceiveBytes: 512}) var ae *APIError if !errors.As(err, &ae) || ae.Code != CodeBadRequest { t.Fatalf("err=%v", err) } } func TestK00URLMapping(t *testing.T) { u, err := normalizeMQTTURL("https://host:7443/", false) if err != nil || u.Scheme != "wss" || u.Path != "/mqtt" { t.Fatalf("%v %v", u, err) } u, err = normalizeMQTTURL("http://host/app", false) if err != nil || u.Scheme != "ws" || u.Path != "/app" { t.Fatalf("%v %v", u, err) } if _, err := normalizeMQTTURL("mqtt://host:1883", false); err == nil { t.Fatal("mqtt without AllowTCP") } if _, err := normalizeMQTTURL("mqtt://host:1883", true); err != nil { t.Fatal(err) } } func TestK00CancelUnsent(t *testing.T) { fake := NewFakeTransport() fake.AutoHello = false c := New() go func() { _ = c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, Options{transport: fake, ConnectTimeout: 2 * time.Second}) }() time.Sleep(40 * time.Millisecond) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) defer cancel() _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) if err == nil { t.Fatal("expected cancel") } if p := c.ResendPayloadForTest(); p != nil { t.Fatalf("still queued %s", p) } c.Close() } func TestK00RateLimitedBackoff(t *testing.T) { DisableJitterForTest(t) fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() var rids []string var id0 string var sendAt any 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) if len(rids) == 0 { id0, _ = last["id"].(string) sendAt = last["send_at_ms"] rids = append(rids, rid) fake.ReplyErr(rid, CodeRateLimited, "slow") continue } if rid == rids[len(rids)-1] { time.Sleep(5 * time.Millisecond) continue } rids = append(rids, rid) if len(rids) < 3 { fake.ReplyErr(rid, CodeRateLimited, "slow") continue } if last["id"] != id0 || last["send_at_ms"] != sendAt { t.Errorf("id/send_at changed") } fake.ReplyOK(rid, map[string]any{"id": id0, "send_at_ms": sendAt, "state": "scheduled"}) return } }() at := time.UnixMilli(1_700_000_000_000) ctx, cancel := context.WithTimeout(context.Background(), 8*time.Second) defer cancel() if _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "hi"}, SendOptions{SendAt: &at}); err != nil { t.Fatal(err) } close(done) if len(rids) != 3 { t.Fatalf("rids=%v", rids) } if rids[0] == rids[1] || rids[1] == rids[2] || rids[0] == rids[2] { t.Fatalf("duplicate rid %v", rids) } } func TestK00ReconnectBackoff(t *testing.T) { b := newReconnectBackoff() if d := b.NextWaitNoJitterForTest(); d != 0 { t.Fatalf("first wait %v", d) } var got []time.Duration for i := 0; i < 6; i++ { b.MarkOffline() got = append(got, b.NextWaitNoJitterForTest()) } want := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second, 16 * time.Second, 30 * time.Second} for i := range want { if got[i] != want[i] { t.Fatalf("i=%d got=%v want=%v", i, got, want) } } b.MarkOnline() b.SetOnlineAtForTest(time.Now()) b.MarkOffline() if d := b.NextWaitNoJitterForTest(); d != 30*time.Second { // flash continues rising: n was 6, +1 = 7 capped 30 if d != 30*time.Second { t.Fatalf("flash %v", d) } } b2 := newReconnectBackoff() _ = b2.NextWaitNoJitterForTest() b2.MarkOnline() b2.SetOnlineAtForTest(time.Now().Add(-61 * time.Second)) b2.MarkOffline() if d := b2.NextWaitNoJitterForTest(); d != time.Second { t.Fatalf("stable reset %v", d) } } func TestK00KeepaliveDefault(t *testing.T) { if DefaultKeepAliveSecondsForTest() != 30 { t.Fatal(DefaultKeepAliveSecondsForTest()) } } func TestK00NoReceiveMaximum(t *testing.T) { fake := NewFakeTransport() c := connectFake(t, fake) defer c.Close() cs := fake.Connects() if len(cs) == 0 || cs[0].ReceiveMaximumSet { t.Fatalf("%+v", cs) } }