package message import ( "context" "database/sql" "sync" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/broker" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) func TestC01NotReadyDoesNotPush(t *testing.T) { t.Parallel() e := openDeliveryEnv(t, nil) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) e.conns.Set("bob", LiveConn{ConnID: "c-bob"}) // Ready=false ctx := context.Background() if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("nr1", "bob")); err != nil { t.Fatal(err) } e.app.WakePush("bob") if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { t.Fatal(err) } if e.down.FilterType(protocol.TypeMsg) != 0 { t.Fatalf("pushed before handshake: %d", e.down.FilterType(protocol.TypeMsg)) } seq := e.seqOf("alice", "nr1") var pushed sql.NullString _ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed) if pushed.Valid { t.Fatalf("pushed_conn=%s", pushed.String) } live := LiveConn{ConnID: "c-bob", Ready: true} e.conns.Set("bob", live) if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { t.Fatal(err) } if e.down.FilterType(protocol.TypeMsg) != 1 { t.Fatalf("want 1 msg after ready, got %d", e.down.FilterType(protocol.TypeMsg)) } } func TestC01WindowAndOrderWithConcurrentWake(t *testing.T) { t.Parallel() e := openDeliveryEnv(t, func(l *Limits) { l.DeliveryWindow = 2 }) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) ctx := context.Background() e.online("bob", "c-bob") for i := 0; i < 5; i++ { req := baseSend("w"+string(rune('a'+i)), "bob") if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { t.Fatal(err) } } if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { t.Fatal(err) } if n := e.down.FilterType(protocol.TypeMsg); n != 2 { t.Fatalf("window: got %d want 2", n) } var inflight int _ = e.db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries WHERE endpoint_id='bob' AND state='pending' AND pushed_conn IS NOT NULL`).Scan(&inflight) if inflight > 2 { t.Fatalf("inflight=%d", inflight) } } func TestC01WorkerCoalescesWake(t *testing.T) { t.Parallel() e := openDeliveryEnv(t, func(l *Limits) { l.DeliveryWindow = 32 }) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) ctx := context.Background() live := LiveConn{ConnID: "c-bob", Ready: true} e.conns.Set("bob", live) for i := 0; i < 5; i++ { if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("cw"+string(rune('a'+i)), "bob")); err != nil { t.Fatal(err) } } if err := e.app.OnHandshakeComplete(ctx, "bob", live); err != nil { t.Fatal(err) } var wg sync.WaitGroup for i := 0; i < 50; i++ { wg.Add(1) go func() { defer wg.Done() e.app.WakePush("bob") }() } wg.Wait() deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if e.down.FilterType(protocol.TypeMsg) >= 5 { break } time.Sleep(10 * time.Millisecond) } n := e.down.FilterType(protocol.TypeMsg) if n != 5 { t.Fatalf("got %d msg frames want 5", n) } } func TestC01BrokerPublishErrorsClearClaim(t *testing.T) { t.Parallel() e := openDeliveryEnv(t, nil) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) e.online("bob", "c-bob") e.down.FailNext = 1 e.down.FailErr = broker.ErrBackpressure ctx := context.Background() if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("bp1", "bob")); err != nil { t.Fatal(err) } if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { t.Fatal(err) } seq := e.seqOf("alice", "bp1") var pushed sql.NullString _ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed) if pushed.Valid { t.Fatalf("claim left after backpressure: %s", pushed.String) } } type gatedDown struct { blockBob chan struct{} inner *RecordingDownlink } func (g *gatedDown) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { if endpointID == "bob" && g.blockBob != nil { <-g.blockBob } return g.inner.PublishDown(ctx, endpointID, connID, payload, opts) } func TestC01BlockedPushDoesNotFreezeDispatch(t *testing.T) { t.Parallel() e := openDeliveryEnv(t, nil) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) insertEndpoint(t, e.db, "carol", "", 1, 0) ctx := context.Background() block := make(chan struct{}) gate := &gatedDown{blockBob: block, inner: e.down} e.app.down = gate live := LiveConn{ConnID: "c-bob", Ready: true} e.conns.Set("bob", live) if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("blk1", "bob")); err != nil { t.Fatal(err) } if err := e.app.OnHandshakeComplete(ctx, "bob", live); err != nil { t.Fatal(err) } delay := int64(5_000) req := baseSend("duex", "carol") req.DelayMs = &delay if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { t.Fatal(err) } e.setNow(e.nowMs + 5_000) done := make(chan error, 1) go func() { _, err := e.app.DispatchDue(ctx, e.nowMs, 10) done <- err }() select { case err := <-done: if err != nil { t.Fatal(err) } case <-time.After(time.Second): t.Fatal("dispatch blocked by slow push") } st, _ := e.msgState("alice", "duex") if st != StateDispatched && st != StateCompleted { t.Fatalf("state=%s", st) } close(block) } func TestC01DispatchDueBudgetAndSkipError(t *testing.T) { t.Parallel() e := openDeliveryEnv(t, nil) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) ctx := context.Background() delay := int64(10_000) const n = 250 for i := 0; i < n; i++ { req := baseSend("d"+itoa(i), "bob") req.DelayMs = &delay req.RID = itoa(i) if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { t.Fatal(err) } } bad := baseSend("bad1", "bob") bad.DelayMs = &delay if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, bad); err != nil { t.Fatal(err) } _ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error { _, err := tx.Exec(`UPDATE messages SET dest_kind='nope' WHERE sender_id=? AND id=?`, "alice", "bad1") return err }) e.setNow(e.nowMs + 10_000) start := time.Now() total := 0 for time.Since(start) < time.Second { k, err := e.app.DispatchDue(ctx, e.nowMs, 0) if err != nil { t.Fatal(err) } total += k if k == 0 { break } } if total < n { t.Fatalf("dispatched %d want %d in 1s", total, n) } var badState string _ = e.db.Read.QueryRow(`SELECT state FROM messages WHERE sender_id=? AND id=?`, "alice", "bad1").Scan(&badState) if badState != StateScheduled { t.Fatalf("bad message state=%s", badState) } } func itoa(i int) string { if i == 0 { return "0" } var b [16]byte pos := len(b) for i > 0 { pos-- b[pos] = byte('0' + i%10) i /= 10 } return string(b[pos:]) }