fix: 合入后修正迁移计数并唤醒回执推送
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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))
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user