fix: 补提交校验、停用检查、入群过滤与按完成时刻清理
This commit is contained in:
@@ -162,9 +162,11 @@ func (a *App) now() time.Time {
|
||||
|
||||
func (a *App) protocolLimits() protocol.Limits {
|
||||
return protocol.Limits{
|
||||
MaxBodyBytes: a.lim.MaxBodyBytes,
|
||||
MaxMetaBytes: a.lim.MaxMetaBytes,
|
||||
MaxFrameBytes: a.lim.MaxFrameBytes,
|
||||
MaxBodyBytes: a.lim.MaxBodyBytes,
|
||||
MaxMetaBytes: a.lim.MaxMetaBytes,
|
||||
MaxFrameBytes: a.lim.MaxFrameBytes,
|
||||
MaxTTLSeconds: a.lim.MaxTTLSeconds,
|
||||
MaxScheduleSeconds: a.lim.MaxScheduleSeconds,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -86,7 +86,15 @@ WHERE d.state = 'pending' AND d.pushed_conn IS NULL
|
||||
cutoff := nowMs - int64(a.lim.RecordRetentionDays)*24*3600*1000
|
||||
if _, err := tx.Exec(`
|
||||
DELETE FROM messages WHERE seq IN (
|
||||
SELECT seq FROM messages WHERE state = 'completed' AND created_at < ? LIMIT 5000
|
||||
SELECT seq FROM (
|
||||
SELECT m.seq FROM messages m
|
||||
WHERE m.state = 'completed'
|
||||
AND COALESCE(
|
||||
(SELECT MAX(d.updated_at) FROM deliveries d WHERE d.seq = m.seq),
|
||||
m.send_at
|
||||
) < ?
|
||||
LIMIT 5000
|
||||
)
|
||||
)`, cutoff); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -46,6 +46,9 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
keep := protocol.EffectiveOfflineKeep(req)
|
||||
ttl := protocol.EffectiveOfflineTTL(req)
|
||||
receipt := protocol.EffectiveReceipt(req)
|
||||
if keep && ttl <= 0 {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds must be > 0")
|
||||
}
|
||||
if keep && a.lim.MaxTTLSeconds > 0 && ttl > a.lim.MaxTTLSeconds {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds exceeds max_ttl_seconds")
|
||||
}
|
||||
@@ -66,6 +69,9 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
}
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
if sender.Enabled == 0 {
|
||||
return SubmitResult{}, errCode(protocol.CodeUnauthorized, "sender disabled")
|
||||
}
|
||||
|
||||
sendAt, err := a.computeSendAt(req, sender.DefaultDelayMs, nowMs)
|
||||
if err != nil {
|
||||
@@ -155,6 +161,17 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
return e
|
||||
}
|
||||
|
||||
snd, se := loadEndpointTx(tx, senderID)
|
||||
if se != nil {
|
||||
if errors.Is(se, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeInvalidTarget, "sender not found")
|
||||
}
|
||||
return se
|
||||
}
|
||||
if snd.Enabled == 0 {
|
||||
return errCode(protocol.CodeUnauthorized, "sender disabled")
|
||||
}
|
||||
|
||||
// 写事务内再确认目标与授权(防并发停用/退群)。
|
||||
switch req.To.Kind {
|
||||
case protocol.TargetEndpoint:
|
||||
@@ -197,10 +214,6 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
}
|
||||
}
|
||||
// 发送方设了对话密码且发给别人的单聊:给对方写回复授权。
|
||||
snd, se := loadEndpointTx(tx, senderID)
|
||||
if se != nil {
|
||||
return se
|
||||
}
|
||||
if senderID != req.To.ID && snd.TalkHash != nil && *snd.TalkHash != "" {
|
||||
if ge := upsertGrantTx(tx, req.To.ID, senderID, snd.TalkVersion, GrantKindReply, nowMs); ge != nil {
|
||||
return ge
|
||||
@@ -295,6 +308,12 @@ func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (in
|
||||
if *req.DelayMs < 0 {
|
||||
return 0, errCode(protocol.CodeBadRequest, "delay_ms negative")
|
||||
}
|
||||
if a.lim.MaxScheduleSeconds > 0 {
|
||||
maxDelay := a.lim.MaxScheduleSeconds * 1000
|
||||
if *req.DelayMs > maxDelay {
|
||||
return 0, errCode(protocol.CodeBadRequest, "send time exceeds max_schedule_seconds")
|
||||
}
|
||||
}
|
||||
sendAt = nowMs + *req.DelayMs
|
||||
default:
|
||||
if defaultDelayMs < 0 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -99,4 +99,47 @@ VALUES(?,?,?,0,'rejected','left_group',?)`, seq, "bob", nowMs, nowMs)
|
||||
if bodies != 0 {
|
||||
t.Fatalf("body still present: %d", bodies)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupKeepsRecentlyCompletedOldCreated(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
lim.RecordRetentionDays = 7
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
nowMs := int64(1_700_000_000_000)
|
||||
created := nowMs - int64(30)*24*3600*1000
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, e := tx.Exec(`
|
||||
INSERT INTO messages(id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
|
||||
send_at, keep, ttl_seconds, receipt, state, reason, created_at)
|
||||
VALUES('old-created','alice','endpoint','bob','{}','text/plain; charset=utf-8','utf8',
|
||||
?,0,0,1,'completed','',?)`, nowMs, created)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
seq, e := res.LastInsertId()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
_, e = tx.Exec(`
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at)
|
||||
VALUES(?,?,?,0,'accepted','',?)`, seq, "bob", nowMs, nowMs)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.CleanupOnce(ctx, nowMs); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE id='old-created'`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("recently completed message should remain, n=%d", n)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user