fix: 补提交校验、停用检查、入群过滤与按完成时刻清理
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user