package main import ( "context" "database/sql" "io" "log/slog" "path/filepath" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/app/message" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/config" "git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/store" ) func TestHandleUplinkRateLimitStatusAndAckExempt(t *testing.T) { t.Parallel() db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) nowMs := int64(1_700_000_000_000) err = 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('alice','alice','stub$login',NULL,0,0,1,?)`, nowMs) return e }) if err != nil { t.Fatal(err) } lim := message.LimitsFromFullConfig(config.Default()) lim.RequestsPerSecond = 50 lim.RequestBurst = 100 app := message.New(db, lim, auth.NewStubHashPool(), message.WithNow(func() time.Time { return time.UnixMilli(nowMs) }), ) down := &message.RecordingDownlink{} conns := message.NewMemoryConns() conns.Set("alice", message.LiveConn{ConnID: "c1"}) u := &appUplink{msg: app, conns: conns, down: down, log: slog.New(slog.NewTextHandler(io.Discard, nil))} conn := port.ConnInfo{EndpointID: "alice", ConnID: "c1"} ctx := context.Background() ackPayload, err := protocol.Marshal(&protocol.Ack{ V: protocol.Version, Type: protocol.TypeAck, RID: "a", From: "alice", ID: "missing", }) if err != nil { t.Fatal(err) } for i := 0; i < 150; i++ { if e := u.HandleUplink(ctx, conn, ackPayload); e != nil { t.Fatal(e) } } if n := countRespCode(down, protocol.CodeRateLimited); n != 0 { t.Fatalf("ack should not count, rate_limited=%d", n) } statusPayload, err := protocol.Marshal(&protocol.Status{ V: protocol.Version, Type: protocol.TypeStatus, RID: "s", ID: "no-such", }) if err != nil { t.Fatal(err) } for i := 0; i < 150; i++ { if e := u.HandleUplink(ctx, conn, statusPayload); e != nil { t.Fatal(e) } } limited := countRespCode(down, protocol.CodeRateLimited) if limited != 50 { t.Fatalf("status rate_limited=%d want 50 (burst 100 of 150)", limited) } } func countRespCode(down *message.RecordingDownlink, code string) int { n := 0 for _, p := range down.Snapshots() { var resp protocol.Resp if err := protocol.Unmarshal(p.Payload, &resp); err != nil { continue } if !resp.OK && resp.Error != nil && resp.Error.Code == code { n++ } } return n }