fix: 补提交校验、停用检查、入群过滤与按完成时刻清理

This commit is contained in:
Nixevol
2026-09-30 16:21:05 +08:00
parent 2461bec3b6
commit 659373e142
8 changed files with 247 additions and 11 deletions
+136
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"errors"
"math"
"path/filepath"
"testing"
"time"
@@ -363,6 +364,141 @@ SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice")
}
})
t.Run("ttl_zero_rejected", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ttl := int64(0)
req := baseSend("ttl0", "bob")
req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req)
if protoCode(err) != protocol.CodeBadRequest {
t.Fatalf("ttl=0 want bad_request got %v", err)
}
})
t.Run("ttl_negative_rejected", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ttl := int64(-1)
req := baseSend("ttlneg", "bob")
req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req)
if protoCode(err) != protocol.CodeBadRequest {
t.Fatalf("ttl=-1 want bad_request got %v", err)
}
})
t.Run("delay_maxint64_rejected", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
delay := int64(math.MaxInt64)
req := baseSend("delaymax", "bob")
req.DelayMs = &delay
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req)
if protoCode(err) != protocol.CodeBadRequest {
t.Fatalf("delay=MaxInt64 want bad_request got %v", err)
}
})
t.Run("sender_disabled_unauthorized", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 0, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, baseSend("from-off", "bob"))
if protoCode(err) != protocol.CodeUnauthorized {
t.Fatalf("want unauthorized got %v", err)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("deliveries=%d", n)
}
})
t.Run("group_late_joiner_skipped", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
insertEndpoint(t, db, "dave", "", 1, 0)
ctx := context.Background()
nowMs := int64(1_700_000_000_000)
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
"g-late", "g", "alice", nowMs); e != nil {
return e
}
for _, m := range []string{"alice", "bob"} {
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
"g-late", m, nowMs); e != nil {
return e
}
}
return nil
})
if err != nil {
t.Fatal(err)
}
delay := int64(10_000)
req := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "late-1",
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g-late"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
DelayMs: &delay,
Offline: keepTrue(),
}
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if res.State != StateScheduled {
t.Fatalf("state=%s", res.State)
}
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
"g-late", "dave", res.SendAtMs+500)
return e
})
if err != nil {
t.Fatal(err)
}
if _, err := app.DispatchDue(ctx, res.SendAtMs, 10); err != nil {
t.Fatal(err)
}
var daveN int
if err := db.Read.QueryRow(`
SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq
WHERE m.id='late-1' AND d.endpoint_id='dave'`).Scan(&daveN); err != nil {
t.Fatal(err)
}
if daveN != 0 {
t.Fatalf("late joiner deliveries=%d", daveN)
}
var bobN int
if err := db.Read.QueryRow(`
SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq
WHERE m.id='late-1' AND d.endpoint_id='bob'`).Scan(&bobN); err != nil {
t.Fatal(err)
}
if bobN != 1 {
t.Fatalf("bob deliveries=%d", bobN)
}
})
t.Run("rate_limited", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()