From 8f2ebc7d2123df5c4195eeeb843384c0b1f9c053 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 16:30:27 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E5=90=88=E5=85=A5=E5=90=8E=E4=BF=AE?= =?UTF-8?q?=E6=AD=A3=E8=BF=81=E7=A7=BB=E8=AE=A1=E6=95=B0=E5=B9=B6=E5=94=A4?= =?UTF-8?q?=E9=86=92=E5=9B=9E=E6=89=A7=E6=8E=A8=E9=80=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/admin/endpoints_test.go | 13 ++++---- internal/app/message/push.go | 10 ++++--- internal/app/message/recover.go | 15 ++++++++-- internal/app/message/review_c01_test.go | 40 +++++++++++++++++++++++++ internal/app/message/submit.go | 1 + internal/store/db_test.go | 8 ++--- 6 files changed, 71 insertions(+), 16 deletions(-) diff --git a/internal/admin/endpoints_test.go b/internal/admin/endpoints_test.go index 6b53aed..d204ed9 100644 --- a/internal/admin/endpoints_test.go +++ b/internal/admin/endpoints_test.go @@ -384,14 +384,16 @@ func TestEndpointResetPasswordKickFailureFallsBack(t *testing.T) { t.Fatal(seedErr) } kick := &kickRecorder{} - var logBuf bytes.Buffer + var logBuf, auditBuf bytes.Buffer logger := slog.New(slog.NewTextHandler(&logBuf, nil)) + auditLog := slog.New(slog.NewTextHandler(&auditBuf, nil)) h := admin.New(admin.Deps{ DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(), Locks: admin.NewMemoryLoginLocks(), Logger: logger, + AuditLogger: auditLog, KickEndpoint: kick.Kick, PasswordResetKick: func(_ context.Context, _ string) (bool, error) { return false, errors.New("fatal kick failed") @@ -443,11 +445,12 @@ func TestEndpointResetPasswordKickFailureFallsBack(t *testing.T) { t.Fatalf("want KickEndpoint fallback once, got %d", kick.count()) } logs := logBuf.String() - if !strings.Contains(logs, "action=endpoint_reset_login_password") { - t.Fatalf("want reset password audit, logs=%s", logs) + auditLogs := auditBuf.String() + if !strings.Contains(auditLogs, "action=endpoint_reset_login_password") { + t.Fatalf("want reset password audit, logs=%s", auditLogs) } - if !strings.Contains(logs, "result=ok_kick_failed") { - t.Fatalf("audit want ok_kick_failed, logs=%s", logs) + if !strings.Contains(auditLogs, "result=ok_kick_failed") { + t.Fatalf("audit want ok_kick_failed, logs=%s", auditLogs) } if !strings.Contains(logs, "password reset kick failed") || !strings.Contains(logs, "fatal kick failed") { t.Fatalf("want error log for reset kick, logs=%s", logs) diff --git a/internal/app/message/push.go b/internal/app/message/push.go index 3e48505..1a8881c 100644 --- a/internal/app/message/push.go +++ b/internal/app/message/push.go @@ -226,6 +226,7 @@ LIMIT ?`, endpointID, room) if rejErr := a.rejectTooLarge(ctx, it.seq, endpointID, it.senderID, nowMs); rejErr != nil { return rejErr } + a.WakePush(it.senderID) continue } it.payload = payload @@ -412,6 +413,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, if err != nil { return err } + a.WakePush(t.senderID) } return nil } @@ -580,10 +582,10 @@ LIMIT ?`, endpointID, window+len(stale)) return err } type rcpt struct { - rid int64 - msgID, epID string - state, reason string - created int64 + rid int64 + msgID, epID string + state, reason string + created int64 } var list []rcpt for rows.Next() { diff --git a/internal/app/message/recover.go b/internal/app/message/recover.go index e65b3db..d169554 100644 --- a/internal/app/message/recover.go +++ b/internal/app/message/recover.go @@ -50,10 +50,13 @@ func (a *App) ExpireOnce(ctx context.Context, nowMs int64) error { if time.Now().After(deadline) { return nil } - n, err := a.expireOnceBatch(ctx, nowMs) + n, senders, err := a.expireOnceBatch(ctx, nowMs) if err != nil { return err } + for _, s := range senders { + a.WakePush(s) + } if n == 0 { break } @@ -62,8 +65,9 @@ func (a *App) ExpireOnce(ctx context.Context, nowMs int64) error { return nil } -func (a *App) expireOnceBatch(ctx context.Context, nowMs int64) (int, error) { +func (a *App) expireOnceBatch(ctx context.Context, nowMs int64) (int, []string, error) { var n int + var senders []string err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { if err := a.fillZombieExpireTx(tx, nowMs); err != nil { return err @@ -94,6 +98,7 @@ LIMIT ?`, nowMs, expireBatch) } _ = rows.Close() n = len(list) + seen := map[string]struct{}{} for _, it := range list { state := DeliveryDropped reason := ReasonOffline @@ -104,10 +109,14 @@ LIMIT ?`, nowMs, expireBatch) if err := a.finishDeliveryTx(tx, it.seq, it.endpointID, it.senderID, it.msgID, state, reason, false, nowMs); err != nil { return err } + if _, ok := seen[it.senderID]; !ok { + seen[it.senderID] = struct{}{} + senders = append(senders, it.senderID) + } } return finalizeStuckDispatchedTx(tx, nowMs, a.lim.RecordRetentionDays) }) - return n, err + return n, senders, err } func (a *App) fillZombieExpireTx(tx *sql.Tx, nowMs int64) error { diff --git a/internal/app/message/review_c01_test.go b/internal/app/message/review_c01_test.go index cddb7f1..bdb1097 100644 --- a/internal/app/message/review_c01_test.go +++ b/internal/app/message/review_c01_test.go @@ -250,3 +250,43 @@ func itoa(i int) string { } return string(b[pos:]) } + +func TestC01RejectTooLargeWakesSenderReceipt(t *testing.T) { + t.Parallel() + e := openDeliveryEnv(t, nil) + insertEndpoint(t, e.db, "alice", "", 1, 0) + insertEndpoint(t, e.db, "bob", "", 1, 0) + ctx := context.Background() + aliceConn := LiveConn{ConnID: "c-alice", Ready: true} + bobConn := LiveConn{ConnID: "c-bob", Ready: true, MaxReceiveBytes: 50} + e.conns.Set("alice", aliceConn) + e.conns.Set("bob", bobConn) + if err := e.app.OnHandshakeComplete(ctx, "alice", aliceConn); err != nil { + t.Fatal(err) + } + if err := e.app.OnHandshakeComplete(ctx, "bob", bobConn); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + e.app.stopPushWorker("alice", "") + e.app.stopPushWorker("bob", "") + }) + req := baseSend("big-wake", "bob") + b := make([]byte, 80) + for i := range b { + b[i] = 'a' + } + req.Body.Data = string(b) + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if e.down.FilterType(protocol.TypeReceipt) > 0 { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("sender got no receipt after too_large reject, receipts=%d msgs=%d", + e.down.FilterType(protocol.TypeReceipt), e.down.FilterType(protocol.TypeMsg)) +} diff --git a/internal/app/message/submit.go b/internal/app/message/submit.go index c318617..2ab5b42 100644 --- a/internal/app/message/submit.go +++ b/internal/app/message/submit.go @@ -273,6 +273,7 @@ INSERT INTO messages( if err != nil { return SubmitResult{}, err } + a.WakePush(senderID) if result.State == StateDispatched { a.wakeReceivers(ctx, result.ID, senderID) } diff --git a/internal/store/db_test.go b/internal/store/db_test.go index 8864f20..b968550 100644 --- a/internal/store/db_test.go +++ b/internal/store/db_test.go @@ -19,7 +19,7 @@ func TestOpenEmptyDirCreatesTables(t *testing.T) { } defer func() { _ = db.Close() }() - assertMigrationCount(t, db.Write, 3) + assertMigrationCount(t, db.Write, 4) for _, table := range []string{ "endpoints", "settings", "talk_grants", "groups", "group_members", "messages", "message_bodies", "deliveries", "receipts", "send_keys", @@ -41,7 +41,7 @@ func TestOpenIdempotentNoRemigrate(t *testing.T) { if err != nil { t.Fatal(err) } - assertMigrationCount(t, db1.Write, 3) + assertMigrationCount(t, db1.Write, 4) _ = db1.Close() db2, err := Open(dir, "FULL") @@ -49,7 +49,7 @@ func TestOpenIdempotentNoRemigrate(t *testing.T) { t.Fatal(err) } defer func() { _ = db2.Close() }() - assertMigrationCount(t, db2.Write, 3) + assertMigrationCount(t, db2.Write, 4) if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 { t.Fatalf("idempotent reopen should not backup: %v", entries) @@ -83,7 +83,7 @@ CREATE TABLE schema_migrations ( } defer func() { _ = db.Close() }() - assertMigrationCount(t, db.Write, 3) + assertMigrationCount(t, db.Write, 4) assertTableExists(t, db.Write, "endpoints") assertTableExists(t, db.Write, "api_tokens")