From aa24e974af9989681926740b8b50df13664e462b Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:00:06 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=B8=8A=E8=A1=8C=E5=88=86=E5=8F=91?= =?UTF-8?q?=E5=89=8D=E7=BB=9F=E4=B8=80=E9=99=90=E9=80=9F=EF=BC=8Csend=20?= =?UTF-8?q?=E4=B8=8D=E5=86=8D=E5=8D=95=E7=8B=AC=E6=89=A3=E6=A1=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/nixmsg/uplink.go | 13 ++++ cmd/nixmsg/uplink_rate_test.go | 94 +++++++++++++++++++++++++++++ docs/DEVIATIONS.md | 15 ++++- internal/app/message/rate.go | 8 +++ internal/app/message/rate_test.go | 22 +++++++ internal/app/message/submit.go | 3 - internal/app/message/submit_test.go | 13 ++-- 7 files changed, 155 insertions(+), 13 deletions(-) create mode 100644 cmd/nixmsg/uplink_rate_test.go create mode 100644 internal/app/message/rate_test.go diff --git a/cmd/nixmsg/uplink.go b/cmd/nixmsg/uplink.go index 6f4ab57..04bd426 100644 --- a/cmd/nixmsg/uplink.go +++ b/cmd/nixmsg/uplink.go @@ -107,6 +107,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 { @@ -117,6 +121,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: diff --git a/cmd/nixmsg/uplink_rate_test.go b/cmd/nixmsg/uplink_rate_test.go new file mode 100644 index 0000000..a6bdfa3 --- /dev/null +++ b/cmd/nixmsg/uplink_rate_test.go @@ -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 +} diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index b6eeac3..193c68f 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -428,10 +428,10 @@ 2. **请求频率突发容量写死为 100** - 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。 - - 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 不计入桶(与 6.10 一致)。 + - 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。`message.App.AllowRequest` 导出同一令牌桶;`HandleUplink` 在分发前对 ack/receipt_ack 以外的帧调用。`Submit` 不再单独扣桶,避免 send 计两次。 - 原因:配置无独立 burst 字段。 - - 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。 - - 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。 + - 备选方案:配置增加 `request_burst`。 + - 影响:改 `requests_per_second` 不改突发;非 send 请求也受同一桶限制。 3. **未接线 `cmd/nixmsg`** - 原条款:可替换 T0.4 假实现。 @@ -500,6 +500,15 @@ - 备选方案:在 group/identity 各自补写回执与收尾(继续分叉)。 - 影响:退群/解散/停用后发送方可收到 rejected 回执,配额释放,正文删除。 +### 复审修复 C-05 + +1. **每端请求限速覆盖非 send 帧** + - 原条款:PRD F05 / DEVELOPMENT 6.10:除 ack、receipt_ack 外共用一个桶,默认每秒 50、突发 100。 + - 实际做法:message 导出 `AllowRequest`;`cmd/nixmsg/uplink.go` 的 `HandleUplink` 解码后、分发前检查;超限回 `rate_limited`。去掉 `Submit` 内扣桶。不改 uplink 生命周期与 `publishResp`。 + - 原因:原先只有 send 限速,unlock/status/目录/群等可打满哈希池与读库。 + - 备选方案:把桶挪到 broker 层(B-09 范围)。 + - 影响:开放注册后的非 send 请求也计入配额;直接调 `Submit` 的单测不再覆盖限速。 + ## 身份 I ### I1 2026-09-30 diff --git a/internal/app/message/rate.go b/internal/app/message/rate.go index 17555af..0f40975 100644 --- a/internal/app/message/rate.go +++ b/internal/app/message/rate.go @@ -57,3 +57,11 @@ func (r *rateLimiter) allow(endpointID string, now time.Time) bool { b.tokens-- return true } + +// AllowRequest 消耗该端 1 个请求令牌;允许则 true。rps<=0 时不限速。 +func (a *App) AllowRequest(endpointID string) bool { + if a == nil { + return true + } + return a.rates.allow(endpointID, a.now()) +} diff --git a/internal/app/message/rate_test.go b/internal/app/message/rate_test.go new file mode 100644 index 0000000..e4751a8 --- /dev/null +++ b/internal/app/message/rate_test.go @@ -0,0 +1,22 @@ +package message + +import ( + "testing" +) + +func TestAllowRequestBurstAndAckExemptBucket(t *testing.T) { + t.Parallel() + lim := defaultTestLimits() + lim.RequestsPerSecond = 50 + lim.RequestBurst = 100 + app, _ := openTestApp(t, lim) + allowed := 0 + for i := 0; i < 150; i++ { + if app.AllowRequest("alice") { + allowed++ + } + } + if allowed != 100 { + t.Fatalf("allowed=%d want 100 (burst)", allowed) + } +} diff --git a/internal/app/message/submit.go b/internal/app/message/submit.go index bcee77b..acf8a96 100644 --- a/internal/app/message/submit.go +++ b/internal/app/message/submit.go @@ -21,9 +21,6 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid sender") } now := a.now() - if !a.rates.allow(senderID, now) { - return SubmitResult{}, errCode(protocol.CodeRateLimited, "request rate exceeded") - } if err := req.Validate(a.protocolLimits()); err != nil { return SubmitResult{}, err diff --git a/internal/app/message/submit_test.go b/internal/app/message/submit_test.go index 07ef6e1..208c346 100644 --- a/internal/app/message/submit_test.go +++ b/internal/app/message/submit_test.go @@ -371,16 +371,15 @@ SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice") app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) + if !app.AllowRequest("alice") || !app.AllowRequest("alice") { + t.Fatal("burst should allow first two") + } + if app.AllowRequest("alice") { + t.Fatal("third request should be rate limited") + } ctx := context.Background() if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil { t.Fatal(err) } - if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r2", "bob")); err != nil { - t.Fatal(err) - } - _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r3", "bob")) - if protoCode(err) != protocol.CodeRateLimited { - t.Fatalf("want rate_limited got %v", err) - } }) }