fix: 上行分发前统一限速,send 不再单独扣桶

This commit is contained in:
Nixevol
2026-09-30 15:00:06 +08:00
parent ac90495137
commit af8278d2a9
7 changed files with 155 additions and 13 deletions
+13
View File
@@ -67,6 +67,10 @@ func (u *appUplink) HandleUplink(ctx context.Context, conn port.ConnInfo, payloa
u.replyErr(ctx, conn, peekRID(payload), protocol.CodeBadRequest, err.Error())
return nil
}
if !uplinkRateExempt(frame) && u.msg != nil && !u.msg.AllowRequest(conn.EndpointID) {
u.replyErr(ctx, conn, peekRID(payload), protocol.CodeRateLimited, "request rate exceeded")
return nil
}
rid, data, callErr := u.dispatch(ctx, conn, frame)
if callErr != nil {
@@ -77,6 +81,15 @@ func (u *appUplink) HandleUplink(ctx context.Context, conn port.ConnInfo, payloa
return nil
}
func uplinkRateExempt(frame any) bool {
switch frame.(type) {
case *protocol.Ack, *protocol.ReceiptAck:
return true
default:
return false
}
}
func (u *appUplink) dispatch(ctx context.Context, conn port.ConnInfo, frame any) (rid string, data any, err error) {
switch f := frame.(type) {
case *protocol.Send:
+94
View File
@@ -0,0 +1,94 @@
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
}