301 lines
7.4 KiB
Go
301 lines
7.4 KiB
Go
package nixmsg
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestK01InflightResendAfterDisconnect(t *testing.T) {
|
|
fake := NewFakeTransport()
|
|
c := connectFake(t, fake)
|
|
defer c.Close()
|
|
|
|
var firstID string
|
|
var firstSendAt any
|
|
var firstRID string
|
|
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 firstRID == "" {
|
|
firstRID = rid
|
|
firstID, _ = last["id"].(string)
|
|
firstSendAt = last["send_at_ms"]
|
|
if err := fake.SimulateReconnect(); err != nil {
|
|
t.Error(err)
|
|
}
|
|
continue
|
|
}
|
|
if rid != firstRID {
|
|
if last["id"] != firstID || last["send_at_ms"] != firstSendAt {
|
|
t.Errorf("changed id/send_at")
|
|
}
|
|
fake.ReplyOK(rid, map[string]any{"id": firstID, "send_at_ms": firstSendAt, "state": "accepted"})
|
|
return
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
}()
|
|
at := time.UnixMilli(1_700_000_000_111)
|
|
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)
|
|
|
|
stop150 := make(chan struct{})
|
|
go replyAllSends(fake, stop150)
|
|
for i := 0; i < 150; i++ {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
if _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}); err != nil {
|
|
cancel()
|
|
close(stop150)
|
|
t.Fatalf("i=%d %v", i, err)
|
|
}
|
|
cancel()
|
|
}
|
|
close(stop150)
|
|
}
|
|
|
|
func replyAllSends(fake *FakeTransport, stop <-chan struct{}) {
|
|
seen := map[string]struct{}{}
|
|
for {
|
|
select {
|
|
case <-stop:
|
|
return
|
|
default:
|
|
}
|
|
for _, fr := range fake.FindUp("send") {
|
|
rid, _ := fr["rid"].(string)
|
|
if rid == "" {
|
|
continue
|
|
}
|
|
if _, ok := seen[rid]; ok {
|
|
continue
|
|
}
|
|
seen[rid] = struct{}{}
|
|
id, _ := fr["id"].(string)
|
|
fake.ReplyOK(rid, map[string]any{"id": id, "state": "accepted"})
|
|
}
|
|
time.Sleep(3 * time.Millisecond)
|
|
}
|
|
}
|
|
|
|
func TestK01CallbackNoDeadlock(t *testing.T) {
|
|
fake := NewFakeTransport()
|
|
c := connectFake(t, fake)
|
|
defer c.Close()
|
|
go func() {
|
|
for i := 0; i < 80; i++ {
|
|
for _, typ := range []string{"ack", "self.login_password"} {
|
|
for _, fr := range fake.FindUp(typ) {
|
|
rid, _ := fr["rid"].(string)
|
|
if typ == "ack" {
|
|
fake.ReplyOK(rid, map[string]any{"result": "accepted"})
|
|
} else {
|
|
fake.ReplyOK(rid, map[string]any{})
|
|
}
|
|
}
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
}()
|
|
c.opts.ManualAck = true
|
|
started := make(chan struct{})
|
|
done := make(chan error, 1)
|
|
c.OnMessage(func(msg Message) error {
|
|
close(started)
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
if err := c.ChangeLoginPassword(ctx, "old", "newpass12"); err != nil {
|
|
done <- err
|
|
return nil
|
|
}
|
|
done <- c.Ack(msg)
|
|
return nil
|
|
})
|
|
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)
|
|
select {
|
|
case <-started:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("callback not entered")
|
|
}
|
|
select {
|
|
case err := <-done:
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("deadlock")
|
|
}
|
|
}
|
|
|
|
func TestK01PresenceFloodAck(t *testing.T) {
|
|
fake := NewFakeTransport()
|
|
c := connectFake(t, fake)
|
|
defer c.Close()
|
|
var acked atomic.Bool
|
|
go func() {
|
|
for i := 0; i < 100; i++ {
|
|
for _, fr := range fake.FindUp("ack") {
|
|
rid, _ := fr["rid"].(string)
|
|
fake.ReplyOK(rid, map[string]any{"result": "accepted"})
|
|
acked.Store(true)
|
|
}
|
|
time.Sleep(2 * 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)
|
|
for i := 0; i < 1000; i++ {
|
|
p, _ := marshalJSON(map[string]any{
|
|
"v": 1, "type": "presence", "id": "e", "online": true, "at_ms": i,
|
|
})
|
|
fake.InjectDown(p)
|
|
}
|
|
deadline := time.Now().Add(200 * time.Millisecond)
|
|
for time.Now().Before(deadline) {
|
|
if acked.Load() {
|
|
return
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
if !acked.Load() {
|
|
t.Fatal("ack not finished in 200ms")
|
|
}
|
|
}
|
|
|
|
func TestK01WatchRestored(t *testing.T) {
|
|
fake := NewFakeTransport()
|
|
c := connectFake(t, fake)
|
|
defer c.Close()
|
|
go func() {
|
|
for i := 0; i < 80; i++ {
|
|
for _, fr := range fake.FindUp("presence.watch") {
|
|
rid, _ := fr["rid"].(string)
|
|
fake.ReplyOK(rid, map[string]any{})
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
}()
|
|
if err := c.WatchPresence(context.Background(), []string{"a", "b"}, false); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
n1 := len(fake.FindUp("presence.watch"))
|
|
if err := fake.SimulateReconnect(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
deadline := time.Now().Add(2 * time.Second)
|
|
for time.Now().Before(deadline) {
|
|
if len(fake.FindUp("presence.watch")) > n1 {
|
|
return
|
|
}
|
|
time.Sleep(10 * time.Millisecond)
|
|
}
|
|
t.Fatalf("watch not restored, had %d", n1)
|
|
}
|
|
|
|
func TestK01FatalOnce(t *testing.T) {
|
|
fake := NewFakeTransport()
|
|
c := connectFake(t, fake)
|
|
var n atomic.Int32
|
|
c.OnConnection(func(ev ConnectionEvent) {
|
|
if ev.State == StateAuthFailed && ev.Reason == "disabled" {
|
|
n.Add(1)
|
|
}
|
|
})
|
|
fatal, _ := marshalJSON(map[string]any{"v": 1, "type": "fatal", "reason": "disabled"})
|
|
fake.InjectDown(fatal)
|
|
fake.InjectDown(fatal)
|
|
time.Sleep(50 * time.Millisecond)
|
|
if n.Load() != 1 {
|
|
t.Fatalf("reason reports=%d", n.Load())
|
|
}
|
|
}
|
|
|
|
func TestK01SendResultJSON(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)
|
|
}
|
|
}
|
|
|
|
func TestK01DedupLRUKeepsReinserted(t *testing.T) {
|
|
c := New()
|
|
c.opts.DedupCapacity = 10000
|
|
c.store = newLRU(10000)
|
|
key := "m\x00a\x00id1"
|
|
c.DedupPutForTest(key, dedupDelivered)
|
|
c.DedupDeleteForTest(key)
|
|
c.DedupPutForTest(key, dedupAcked)
|
|
for i := 0; i < 9999; i++ {
|
|
c.DedupPutForTest(fmt.Sprintf("n:%d", i), dedupAcked)
|
|
}
|
|
if !c.DedupHasForTest(key) {
|
|
t.Fatal("key evicted too early")
|
|
}
|
|
}
|
|
|
|
func TestK01FailAuthFast(t *testing.T) {
|
|
fake := NewFakeTransport()
|
|
c := connectFake(t, fake)
|
|
start := time.Now()
|
|
c.OnMessage(func(msg Message) error {
|
|
fake.SimulateAuthFail(AuthBadCredentials)
|
|
return nil
|
|
})
|
|
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)
|
|
deadline := time.Now().Add(200 * time.Millisecond)
|
|
for time.Now().Before(deadline) {
|
|
if c.LastStopCodeForTest() == CodeBadCredentials {
|
|
if time.Since(start) > 100*time.Millisecond {
|
|
t.Fatalf("too slow %v", time.Since(start))
|
|
}
|
|
return
|
|
}
|
|
time.Sleep(2 * time.Millisecond)
|
|
}
|
|
t.Fatal("auth fail not observed")
|
|
}
|
|
|
|
func TestK01HelloDelay15s(t *testing.T) {
|
|
if testing.Short() {
|
|
t.Skip()
|
|
}
|
|
t.Skip("optional 15s handshake; covered by ConnectTimeout=PacketTimeout")
|
|
}
|