fix: 合入后修正迁移计数并唤醒回执推送

This commit is contained in:
Nixevol
2026-09-30 16:30:27 +08:00
parent cb9512fe11
commit 8f2ebc7d21
6 changed files with 71 additions and 16 deletions
+8 -5
View File
@@ -384,14 +384,16 @@ func TestEndpointResetPasswordKickFailureFallsBack(t *testing.T) {
t.Fatal(seedErr) t.Fatal(seedErr)
} }
kick := &kickRecorder{} kick := &kickRecorder{}
var logBuf bytes.Buffer var logBuf, auditBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(&logBuf, nil)) logger := slog.New(slog.NewTextHandler(&logBuf, nil))
auditLog := slog.New(slog.NewTextHandler(&auditBuf, nil))
h := admin.New(admin.Deps{ h := admin.New(admin.Deps{
DB: db, DB: db,
Hash: hash, Hash: hash,
Tokens: admin.NewRandomAPITokens(), Tokens: admin.NewRandomAPITokens(),
Locks: admin.NewMemoryLoginLocks(), Locks: admin.NewMemoryLoginLocks(),
Logger: logger, Logger: logger,
AuditLogger: auditLog,
KickEndpoint: kick.Kick, KickEndpoint: kick.Kick,
PasswordResetKick: func(_ context.Context, _ string) (bool, error) { PasswordResetKick: func(_ context.Context, _ string) (bool, error) {
return false, errors.New("fatal kick failed") 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()) t.Fatalf("want KickEndpoint fallback once, got %d", kick.count())
} }
logs := logBuf.String() logs := logBuf.String()
if !strings.Contains(logs, "action=endpoint_reset_login_password") { auditLogs := auditBuf.String()
t.Fatalf("want reset password audit, logs=%s", logs) 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") { if !strings.Contains(auditLogs, "result=ok_kick_failed") {
t.Fatalf("audit want ok_kick_failed, logs=%s", logs) t.Fatalf("audit want ok_kick_failed, logs=%s", auditLogs)
} }
if !strings.Contains(logs, "password reset kick failed") || !strings.Contains(logs, "fatal kick failed") { 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) t.Fatalf("want error log for reset kick, logs=%s", logs)
+2
View File
@@ -226,6 +226,7 @@ LIMIT ?`, endpointID, room)
if rejErr := a.rejectTooLarge(ctx, it.seq, endpointID, it.senderID, nowMs); rejErr != nil { if rejErr := a.rejectTooLarge(ctx, it.seq, endpointID, it.senderID, nowMs); rejErr != nil {
return rejErr return rejErr
} }
a.WakePush(it.senderID)
continue continue
} }
it.payload = payload it.payload = payload
@@ -412,6 +413,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`,
if err != nil { if err != nil {
return err return err
} }
a.WakePush(t.senderID)
} }
return nil return nil
} }
+12 -3
View File
@@ -50,10 +50,13 @@ func (a *App) ExpireOnce(ctx context.Context, nowMs int64) error {
if time.Now().After(deadline) { if time.Now().After(deadline) {
return nil return nil
} }
n, err := a.expireOnceBatch(ctx, nowMs) n, senders, err := a.expireOnceBatch(ctx, nowMs)
if err != nil { if err != nil {
return err return err
} }
for _, s := range senders {
a.WakePush(s)
}
if n == 0 { if n == 0 {
break break
} }
@@ -62,8 +65,9 @@ func (a *App) ExpireOnce(ctx context.Context, nowMs int64) error {
return nil 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 n int
var senders []string
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if err := a.fillZombieExpireTx(tx, nowMs); err != nil { if err := a.fillZombieExpireTx(tx, nowMs); err != nil {
return err return err
@@ -94,6 +98,7 @@ LIMIT ?`, nowMs, expireBatch)
} }
_ = rows.Close() _ = rows.Close()
n = len(list) n = len(list)
seen := map[string]struct{}{}
for _, it := range list { for _, it := range list {
state := DeliveryDropped state := DeliveryDropped
reason := ReasonOffline 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 { if err := a.finishDeliveryTx(tx, it.seq, it.endpointID, it.senderID, it.msgID, state, reason, false, nowMs); err != nil {
return err 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 finalizeStuckDispatchedTx(tx, nowMs, a.lim.RecordRetentionDays)
}) })
return n, err return n, senders, err
} }
func (a *App) fillZombieExpireTx(tx *sql.Tx, nowMs int64) error { func (a *App) fillZombieExpireTx(tx *sql.Tx, nowMs int64) error {
+40
View File
@@ -250,3 +250,43 @@ func itoa(i int) string {
} }
return string(b[pos:]) 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))
}
+1
View File
@@ -273,6 +273,7 @@ INSERT INTO messages(
if err != nil { if err != nil {
return SubmitResult{}, err return SubmitResult{}, err
} }
a.WakePush(senderID)
if result.State == StateDispatched { if result.State == StateDispatched {
a.wakeReceivers(ctx, result.ID, senderID) a.wakeReceivers(ctx, result.ID, senderID)
} }
+4 -4
View File
@@ -19,7 +19,7 @@ func TestOpenEmptyDirCreatesTables(t *testing.T) {
} }
defer func() { _ = db.Close() }() defer func() { _ = db.Close() }()
assertMigrationCount(t, db.Write, 3) assertMigrationCount(t, db.Write, 4)
for _, table := range []string{ for _, table := range []string{
"endpoints", "settings", "talk_grants", "groups", "group_members", "endpoints", "settings", "talk_grants", "groups", "group_members",
"messages", "message_bodies", "deliveries", "receipts", "send_keys", "messages", "message_bodies", "deliveries", "receipts", "send_keys",
@@ -41,7 +41,7 @@ func TestOpenIdempotentNoRemigrate(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
assertMigrationCount(t, db1.Write, 3) assertMigrationCount(t, db1.Write, 4)
_ = db1.Close() _ = db1.Close()
db2, err := Open(dir, "FULL") db2, err := Open(dir, "FULL")
@@ -49,7 +49,7 @@ func TestOpenIdempotentNoRemigrate(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
defer func() { _ = db2.Close() }() 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 { if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
t.Fatalf("idempotent reopen should not backup: %v", entries) t.Fatalf("idempotent reopen should not backup: %v", entries)
@@ -83,7 +83,7 @@ CREATE TABLE schema_migrations (
} }
defer func() { _ = db.Close() }() defer func() { _ = db.Close() }()
assertMigrationCount(t, db.Write, 3) assertMigrationCount(t, db.Write, 4)
assertTableExists(t, db.Write, "endpoints") assertTableExists(t, db.Write, "endpoints")
assertTableExists(t, db.Write, "api_tokens") assertTableExists(t, db.Write, "api_tokens")