95 lines
2.6 KiB
Go
95 lines
2.6 KiB
Go
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
|
|
}
|