diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index 26400ad..9cd7725 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -299,7 +299,7 @@ func runServe(ctx context.Context, cfg config.Config) error { loopCtx, loopCancel := context.WithCancel(ctx) defer loopCancel() - go messageLoops(loopCtx, msgApp, memConns, db, hashPool, metricsReg) + go messageLoops(loopCtx, msgApp, db, hashPool, metricsReg) <-ctx.Done() loopCancel() @@ -316,30 +316,25 @@ func runServe(ctx context.Context, cfg config.Config) error { return nil } -func messageLoops(ctx context.Context, msgApp *message.App, conns *message.MemoryConns, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) { - t := time.NewTicker(time.Second) +func messageLoops(ctx context.Context, msgApp *message.App, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) { + msgApp.StartLoops(ctx) + t := time.NewTicker(15 * time.Second) defer t.Stop() + sample := func() { + opCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + if err := metrics.SampleStoreGauges(opCtx, met, db.Read); err != nil { + slog.Debug("sample store gauges", "err", err) + } + metrics.SampleQueues(met, db.Queue.Len(), hashPool.QueueLen()) + } + sample() for { select { case <-ctx.Done(): return case <-t.C: - nowMs := time.Now().UnixMilli() - if _, err := msgApp.DispatchDue(ctx, nowMs, 100); err != nil { - slog.Error("dispatch due", "err", err) - } - for ep, live := range conns.Snapshot() { - if err := msgApp.PushPending(ctx, ep, live.ConnID); err != nil { - slog.Debug("push pending", "endpoint", ep, "err", err) - } - } - if err := msgApp.CleanupOnce(ctx, nowMs); err != nil { - slog.Error("cleanup once", "err", err) - } - if err := metrics.SampleStoreGauges(ctx, met, db.Read); err != nil { - slog.Debug("sample store gauges", "err", err) - } - metrics.SampleQueues(met, db.Queue.Len(), hashPool.QueueLen()) + sample() } } } diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index a5ca39d..ec30578 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -511,12 +511,12 @@ - 备选方案:直接依赖 `internal/broker.Broker`。 - 影响:接线方需在握手/断线时调用 `OnHandshakeComplete`/`OnDisconnect`,登记连接,并把 `OnPublishDropped` 转到 `App`。 -2. **大帧并发名额在 message 包再管一份** +2. **大帧并发名额只由 broker 按 PacketID 归还** - 原条款:大于 64KiB 全局同时不超过 64(DEVELOPMENT 7.5);N2 broker 已有信号量。 - - 实际做法:`App` 内另有容量 64 的 `largeSem`,发布前申请,确认/超时/清标记时释放。 - - 原因:假 `Downlink` 不经 broker 时仍要满足上限。 - - 备选方案:只依赖 broker,测试也走真实 PublishDown。 - - 影响:接线真实 broker 后可能双重限流(更严,不破坏语义)。 + - 实际做法(C-01):删除 message 包 `largeSem`/`largeHeld`,发布失败(含 `ErrBackpressure`/`ErrNotSubscribed`/`ErrNoConnection`)清标记并 1 秒后重推;假下行如需限流在其实现里模拟。 + - 原因:消息包名额在断线/撤回/作废时泄漏;broker 已按 PacketID 在 PUBACK/断线归还。 + - 备选方案:message 内改用 (connID, seq) 计数并对账。 + - 影响:单元测试的 `RecordingDownlink` 不再限制大帧并发。 3. **确认超时按库内 `pushed_at` 判定,不另开每连接计时器 goroutine** - 原条款:推送循环在内存里计时。 @@ -525,19 +525,19 @@ - 备选方案:每连接 `time.AfterFunc`。 - 影响:需周期性调用 `PushPending`(或 `WakePush`)才会触发超时。 -4. **后台调度/清理循环未在 App 内自启** +4. **调度/到期/清理循环由 App.StartLoops 启动** - 原条款:调度按 `send_at` 唤醒;清理约每秒;推送每连接一循环。 - - 实际做法:导出 `DispatchDue`、`PushPending`、`CleanupOnce`、`RecoverOnStart`、`WakePush`;由接线方起 goroutine。`WakePush` 在有连接时异步 `PushPending`。 - - 原因:未改 `cmd/nixmsg`;避免无 context 的后台泄漏。 - - 备选方案:`App.Start(ctx)` 内启三循环。 - - 影响:未接线则定时消息不会自动到点,需外部调用 `DispatchDue`。 + - 实际做法(C-01):`StartLoops` 起分发(最早 `send_at` 定时 + `NotifyDispatch`)、到期每秒、清理每小时;握手启动每连接 worker(容量 1 唤醒通道)。`cmd/nixmsg` 的 `messageLoops` 只调 `StartLoops` 并 15 秒采指标。停机等待留给 L-03。 + - 原因:原先单协程串行、握手前也推送、WakePush 每次新协程。 + - 备选方案:继续由 serve 逐秒扫全部连接。 + - 影响:未握手连接不再 claim;测试需 `LiveConn.Ready` 或 `OnHandshakeComplete`。 -5. **回执推送窗口未单独记 inflight** - - 原条款:回执窗口默认 64,确认一笔再推下一笔。 - - 实际做法:按 `acked=0` 取最多 `ReceiptWindow` 条尽力发布;不因未 `receipt_ack` 停推后续。 - - 原因:简化;回执可重复、SDK 按 `receipt_id` 去重。 - - 备选方案:内存记已推未确认回执数。 - - 影响:发送方慢确认时可能多推几条回执(协议允许重复)。 +5. **回执按连接记在途,产品仍允许断线后重复** + - 原条款:回执窗口默认 64,确认一笔再推下一笔;可能重复,SDK 按 `receipt_id` 去重。 + - 实际做法(C-02):每连接 `rcptInflight`;发布前标记、失败撤销;超过确认超时才允许重发;先收集查询结果再 `PublishDown`。重复 ack 不唤醒;`ReceiptAck` 成功后移出在途并唤醒。 + - 原因:原先每次推送全量重发最早 64 条,并发必重复,第 65 条饿死。 + - 备选方案:把在途写入 receipts 表。 + - 影响:断线重连后未确认回执仍各重发一次(协议允许);裸设备须 `receipt_ack` 或 `receipt:false`。 6. **`Status` 返回自建 map,非独立协议类型** - 原条款:6.4 状态响应字段。 @@ -582,6 +582,29 @@ - 备选方案:等 C-03 迁移后改用 `completed_at`;发送方停用改用 `endpoint_disabled`(与目标停用混用)。 - 影响:发送方停用错误码为 `unauthorized`;保留期口径在 C-03 合入前对无投递的 scheduled 作废行用 `send_at` 近似。 +### 复审修复 C-01 + +1. **只对已握手连接推送,每连接一个 worker** + - 原条款:DEVELOPMENT 6.1 / 7.5 握手完成才推送;每连接一循环。 + - 实际做法:`LiveConn.Ready` / `handshook`;`PushPending`/`WakePush` 未就绪直接返回。`OnHandshakeComplete` 清本连接旧 `pushed_conn` 后启动 worker。一轮写操作 claim 全部条目再按 `(send_at, seq)` 发布。`RecoverOnStart` 只做 SQL 修正。不改 `PublishDown` 签名,不改 uplink 生命周期,serve 停机段留给 L-03。 + - 原因:握手前推送会被 mochi 静默丢弃却占窗口。 + - 备选方案:只靠 broker `ErrNotSubscribed` 兜底。 + - 影响:分发在线判定仍含握手中(7.4)。 + +### 复审修复 C-02 + +1. **回执在途表** + - 见上文 M2/M3/M4 第 5 条。不改「可能重复」的产品约定。 + +### 复审修复 C-03 + +1. **到期与小时清理拆分,迁移避开 U-02 的 0003** + - 原条款:DEVELOPMENT 7.5 每秒到期;7.6 每小时分批删并 `wal_checkpoint`。 + - 实际做法:`ExpireOnce` 每写最多 500 条,补僵尸不保留投递的宽限;`PurgeOnce` 三类删除独立写操作,结束后 `Queue.Checkpoint`。迁移 `0004_cleanup_indexes.sql`(U-02 已占用 0003,原计划 0003/0005 合并改号为 0004):`messages.completed_at`、回执/防重/完成时刻索引。收尾写入 `completed_at`;清理按 `COALESCE(completed_at, send_at)`。保留 0 天删全部 completed。断线 UPDATE 带 `endpoint_id`。读池 `MaxOpenConns=64`、`MaxIdleConns=16`、DSN `query_only`。 + - 原因:每秒全表扫描加级联大删除会卡住写队列;按 `created_at` 会误删长定时刚完成的记录(C-07)。 + - 备选方案:继续用 `MAX(deliveries.updated_at)` 近似完成时刻。 + - 影响:旧 completed 行由迁移回填;无投递的作废行仍可能用 `send_at` 兜底。 + ## 身份 I ### I1 2026-09-30 diff --git a/internal/app/message/ack.go b/internal/app/message/ack.go index 3085147..0be436b 100644 --- a/internal/app/message/ack.go +++ b/internal/app/message/ack.go @@ -19,6 +19,8 @@ func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (Ac var seq int64 var ackLatencySec float64 var observeAck bool + var changed bool + var wroteReceipt bool err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { err := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id = ? AND id = ?`, req.From, req.ID).Scan(&seq) if err == sql.ErrNoRows { @@ -40,14 +42,17 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, } aff, _ := res.RowsAffected() if aff > 0 { + changed = true out.Result = DeliveryAccepted if pushedAt.Valid && pushedAt.Int64 > 0 && nowMs >= pushedAt.Int64 { ackLatencySec = float64(nowMs-pushedAt.Int64) / 1000.0 observeAck = true } - if e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs); e != nil { + wrote, e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs) + if e != nil { return e } + wroteReceipt = wrote return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays) } var state string @@ -68,11 +73,12 @@ SELECT state FROM deliveries WHERE seq = ? AND endpoint_id = ?`, seq, endpointID if observeAck && a.met != nil { a.met.AckSeconds.Observe(ackLatencySec) } - if out.Result == DeliveryAccepted { - a.releaseLarge(seq, endpointID) + if changed { + a.WakePush(endpointID) + } + if wroteReceipt { + a.WakePush(req.From) } - a.WakePush(endpointID) - a.WakePush(req.From) return out, nil } @@ -294,8 +300,25 @@ func (a *App) ReceiptAck(ctx context.Context, endpointID string, req *protocol.R if err != nil { return errCode(protocol.CodeBadRequest, "invalid receipt_id") } - return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - _, err := tx.Exec(`UPDATE receipts SET acked = 1 WHERE receipt_id = ? AND sender_id = ?`, rid, endpointID) - return err + var acked bool + err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + res, err := tx.Exec(`UPDATE receipts SET acked = 1 WHERE receipt_id = ? AND sender_id = ? AND acked = 0`, rid, endpointID) + if err != nil { + return err + } + aff, _ := res.RowsAffected() + acked = aff > 0 + return nil }) + if err != nil { + return err + } + if acked { + live, ok := a.lookupConn(endpointID) + if ok { + a.unmarkReceipt(string(live.ConnID), rid) + } + a.WakePush(endpointID) + } + return nil } diff --git a/internal/app/message/app.go b/internal/app/message/app.go index 6e55ea9..095f0dc 100644 --- a/internal/app/message/app.go +++ b/internal/app/message/app.go @@ -90,10 +90,16 @@ type App struct { met *metrics.Registry mu sync.Mutex - largeSem chan struct{} - largeHeld map[string]bool pendingRevoke []revokeJob repushTimers map[string]*time.Timer + workers map[string]*pushWorker + handshook map[string]port.ConnID + rcptInflight map[string]map[int64]int64 // connID → receiptID → 推送时刻 ms + + dispatchCh chan struct{} + loopWG sync.WaitGroup + lastPurge time.Time + purgeMu sync.Mutex } // Option 配置 App。 @@ -142,13 +148,15 @@ func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App { lim.GraceSeconds = 60 } a := &App{ - db: db, - lim: lim, - hash: hash, - nowFn: time.Now, - rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst), - largeSem: make(chan struct{}, maxLargeInflight), - largeHeld: make(map[string]bool), + db: db, + lim: lim, + hash: hash, + nowFn: time.Now, + rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst), + workers: make(map[string]*pushWorker), + handshook: make(map[string]port.ConnID), + rcptInflight: make(map[string]map[int64]int64), + dispatchCh: make(chan struct{}, 1), } for _, opt := range opts { opt(a) diff --git a/internal/app/message/conn.go b/internal/app/message/conn.go index a828040..6cfaa63 100644 --- a/internal/app/message/conn.go +++ b/internal/app/message/conn.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "sync" + "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" ) @@ -19,6 +20,8 @@ type LiveConn struct { ConnID port.ConnID MaxReceiveBytes int MaxPacketSize uint32 + // Ready 为 true 表示已完成 hello,可以推送。分发在线判定不看此字段。 + Ready bool } // ConnRegistry 查询端是否有连接(由 N 线或测试假实现注入)。 @@ -82,8 +85,11 @@ func (c *MemoryConns) Snapshot() map[string]LiveConn { type RecordingDownlink struct { mu sync.Mutex Published []DownPublish - FailNext int // 接下来 N 次 PublishDown 返回错误 - MaxSize int // >0 时超限返回错误 + FailNext int // 接下来 N 次 PublishDown 返回错误 + FailErr error // FailNext 时返回的错误;空则用 errPublishFailed + MaxSize int // >0 时超限返回错误 + Delay time.Duration // 每次发布前休眠 + Block <-chan struct{} } // DownPublish 是一次下行记录。 @@ -96,6 +102,16 @@ type DownPublish struct { // PublishDown 实现 port.Downlink。 func (d *RecordingDownlink) PublishDown(_ context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { + d.mu.Lock() + block := d.Block + delay := d.Delay + d.mu.Unlock() + if block != nil { + <-block + } + if delay > 0 { + time.Sleep(delay) + } d.mu.Lock() defer d.mu.Unlock() if d.MaxSize > 0 && len(payload) > d.MaxSize { @@ -103,6 +119,9 @@ func (d *RecordingDownlink) PublishDown(_ context.Context, endpointID string, co } if d.FailNext > 0 { d.FailNext-- + if d.FailErr != nil { + return d.FailErr + } return errPublishFailed } d.Published = append(d.Published, DownPublish{ diff --git a/internal/app/message/delivery_test.go b/internal/app/message/delivery_test.go index 537dffb..e3aa128 100644 --- a/internal/app/message/delivery_test.go +++ b/internal/app/message/delivery_test.go @@ -57,7 +57,7 @@ func (e *deliveryEnv) setNow(ms int64) { } func (e *deliveryEnv) online(id string, connID port.ConnID) { - e.conns.Set(id, LiveConn{ConnID: connID, MaxReceiveBytes: 0, MaxPacketSize: 0}) + e.conns.Set(id, LiveConn{ConnID: connID, MaxReceiveBytes: 0, MaxPacketSize: 0, Ready: true}) } func (e *deliveryEnv) deliveryState(seq int64, endpointID string) (state, reason string) { @@ -586,7 +586,7 @@ func TestDeliveryStateMachine(t *testing.T) { e := openDeliveryEnv(t, nil) insertEndpoint(t, e.db, "alice", "", 1, 0) insertEndpoint(t, e.db, "bob", "", 1, 0) - e.conns.Set("bob", LiveConn{ConnID: "c-bob", MaxReceiveBytes: 50}) + e.conns.Set("bob", LiveConn{ConnID: "c-bob", MaxReceiveBytes: 50, Ready: true}) ctx := context.Background() req := baseSend("big1", "bob") req.Body.Data = string(make([]byte, 200)) diff --git a/internal/app/message/dispatch.go b/internal/app/message/dispatch.go index e68f8ee..4ad1775 100644 --- a/internal/app/message/dispatch.go +++ b/internal/app/message/dispatch.go @@ -3,6 +3,7 @@ package message import ( "database/sql" "encoding/base64" + "time" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) @@ -36,10 +37,16 @@ const ( const ( packetOverheadBudget = 128 - largeFrameBytes = 64 * 1024 - maxLargeInflight = 64 defaultDeliveryWindow = 32 defaultReceiptWindow = 64 + dispatchConcurrency = 8 + dispatchBudget = 250 * time.Millisecond + expireBatch = 500 + expireBudget = 200 * time.Millisecond + purgeRowBatch = 2000 + purgeMsgBatch = 80 + clearPushedTimeout = 5 * time.Second + pushOpTimeout = 30 * time.Second ) // dispatchFullTx 按 DEVELOPMENT 7.4 完整分发一条已到点的 scheduled 消息。 @@ -157,8 +164,13 @@ WHERE endpoint_id = ? AND state = 'pending'`, r.id).Scan(&n); err != nil { case !online && keepBool: expireAt = sql.NullInt64{Int64: nowMs + ttlSeconds*1000, Valid: true} default: - var offlineSince sql.NullInt64 - _ = tx.QueryRow(`SELECT offline_since FROM endpoints WHERE id = ?`, r.id).Scan(&offlineSince) + var onlineSince, offlineSince sql.NullInt64 + _ = tx.QueryRow(`SELECT online_since, offline_since FROM endpoints WHERE id = ?`, r.id).Scan(&onlineSince, &offlineSince) + dbShowsOnline := onlineSince.Valid && (!offlineSince.Valid || onlineSince.Int64 > offlineSince.Int64) + if dbShowsOnline { + expireAt = sql.NullInt64{Int64: nowMs + graceMs, Valid: true} + break + } if !offlineSince.Valid { dState = DeliveryDropped reason = ReasonOffline @@ -182,7 +194,7 @@ VALUES(?,?,?,?,?,?,?,NULL,NULL,0,?)`, if dState == DeliveryPending { pendingAny = true } else if wantReceipt { - if err := insertReceiptTx(tx, senderID, seq, r.id, dState, reason, nowMs); err != nil { + if _, err := insertReceiptTx(tx, senderID, seq, r.id, dState, reason, nowMs); err != nil { return "", true, err } } @@ -223,7 +235,7 @@ func FinalizeMessageTx(tx *sql.Tx, seq int64, wantReceipt bool, senderID, endpoi return err } if msgReason != "" && wantReceipt && receipt != 0 { - if err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, msgReason, nowMs); err != nil { + if _, err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, msgReason, nowMs); err != nil { return err } } @@ -298,35 +310,37 @@ WHERE seq = ? AND endpoint_id = ? AND state = ?`, if err := tx.QueryRow(`SELECT sender_id FROM messages WHERE seq = ?`, seq).Scan(&senderID); err != nil { return false, err } - if err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, reason, nowMs); err != nil { + if _, err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, reason, nowMs); err != nil { return false, err } } return pushedAt.Valid, nil } -func insertReceiptTx(tx *sql.Tx, senderID string, seq int64, endpointID, state, reason string, nowMs int64) error { +func insertReceiptTx(tx *sql.Tx, senderID string, seq int64, endpointID, state, reason string, nowMs int64) (bool, error) { var msgID string var want int if err := tx.QueryRow(`SELECT id, receipt FROM messages WHERE seq = ?`, seq).Scan(&msgID, &want); err != nil { - return err + return false, err } if want == 0 { - return nil + return false, nil } - // 发送方仍存在 var one int err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, senderID).Scan(&one) if err == sql.ErrNoRows { - return nil + return false, nil } if err != nil { - return err + return false, err } _, err = tx.Exec(` INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked) VALUES(?,?,?,?,?,?,0)`, senderID, msgID, endpointID, state, reason, nowMs) - return err + if err != nil { + return false, err + } + return true, nil } func encodeStoredBody(enc, contentType string, raw []byte) protocol.Body { diff --git a/internal/app/message/loops.go b/internal/app/message/loops.go new file mode 100644 index 0000000..90305f4 --- /dev/null +++ b/internal/app/message/loops.go @@ -0,0 +1,118 @@ +package message + +import ( + "context" + "log/slog" + "time" +) + +// StartLoops 启动到点分发、到期处理与小时清理循环。推送 worker 在握手时启动。 +// 循环随 ctx 取消退出;停机等待由 L-03 调用 WaitLoops。 +func (a *App) StartLoops(ctx context.Context) { + a.loopWG.Add(3) + go func() { + defer a.loopWG.Done() + a.dispatchLoop(ctx) + }() + go func() { + defer a.loopWG.Done() + a.expireLoop(ctx) + }() + go func() { + defer a.loopWG.Done() + a.purgeLoop(ctx) + }() +} + +// WaitLoops 等待 StartLoops 启动的 goroutine 退出。 +func (a *App) WaitLoops() { + a.loopWG.Wait() +} + +func (a *App) dispatchLoop(ctx context.Context) { + timer := time.NewTimer(time.Millisecond) + defer timer.Stop() + for { + select { + case <-ctx.Done(): + return + case <-a.dispatchCh: + case <-timer.C: + } + opCtx, cancel := context.WithTimeout(ctx, time.Second) + nowMs := a.now().UnixMilli() + if _, err := a.DispatchDue(opCtx, nowMs, 0); err != nil && ctx.Err() == nil { + slog.Error("dispatch due", "err", err) + } + cancel() + delay := time.Second + if next, ok := a.earliestScheduled(ctx); ok { + d := time.Until(time.UnixMilli(next)) + switch { + case d < 0: + d = 0 + case d > time.Minute: + d = time.Minute + } + delay = d + } + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(delay) + } +} + +func (a *App) expireLoop(ctx context.Context) { + t := time.NewTicker(time.Second) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + opCtx, cancel := context.WithTimeout(ctx, time.Second) + if err := a.ExpireOnce(opCtx, a.now().UnixMilli()); err != nil && ctx.Err() == nil { + slog.Error("expire once", "err", err) + } + cancel() + } + } +} + +func (a *App) purgeLoop(ctx context.Context) { + t := time.NewTicker(time.Hour) + defer t.Stop() + run := func() { + opCtx, cancel := context.WithTimeout(ctx, 30*time.Second) + if err := a.PurgeOnce(opCtx, a.now().UnixMilli()); err != nil && ctx.Err() == nil { + slog.Error("purge once", "err", err) + } + cancel() + } + run() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + run() + } + } +} + +func (a *App) earliestScheduled(ctx context.Context) (int64, bool) { + if a.db == nil || a.db.Read == nil { + return 0, false + } + var sendAt int64 + err := a.db.Read.QueryRowContext(ctx, ` +SELECT MIN(send_at) FROM messages WHERE state = 'scheduled'`).Scan(&sendAt) + if err != nil || sendAt == 0 { + return 0, false + } + return sendAt, true +} diff --git a/internal/app/message/push.go b/internal/app/message/push.go index ec4a2cb..3e48505 100644 --- a/internal/app/message/push.go +++ b/internal/app/message/push.go @@ -4,28 +4,75 @@ import ( "context" "database/sql" "encoding/json" + "errors" "fmt" + "log/slog" + "sync" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) +type dueMsg struct { + seq int64 + senderID string + destKind string + destID string + sendAt int64 + keep int + ttl int64 + receipt int +} + +type pushItem struct { + seq int64 + sendAt int64 + keep int + expireAt sql.NullInt64 + msgID string + senderID string + destKind string + destID string + meta string + contentType string + bodyEnc string + body []byte + payload []byte +} + // DispatchDue 分发已到点的 scheduled 消息(按 send_at、seq)。 +// limit<=0 时循环到取空或用完时间预算,单条失败只记日志并跳过。 func (a *App) DispatchDue(ctx context.Context, nowMs int64, limit int) (int, error) { - if limit <= 0 { - limit = 64 + budgeted := limit <= 0 + batch := limit + if batch <= 0 { + batch = 64 } - type due struct { - seq int64 - senderID string - destKind string - destID string - sendAt int64 - keep int - ttl int64 - receipt int + deadline := time.Now().Add(time.Hour) + if budgeted { + deadline = time.Now().Add(dispatchBudget) } + total := 0 + for { + if err := ctx.Err(); err != nil { + return total, err + } + if budgeted && time.Now().After(deadline) { + return total, nil + } + n, err := a.dispatchDueBatch(ctx, nowMs, batch) + total += n + if err != nil { + return total, err + } + if n == 0 || !budgeted { + return total, nil + } + } +} + +func (a *App) dispatchDueBatch(ctx context.Context, nowMs int64, limit int) (int, error) { rows, err := a.db.Read.QueryContext(ctx, ` SELECT seq, sender_id, dest_kind, dest_id, send_at, keep, ttl_seconds, receipt FROM messages @@ -35,9 +82,9 @@ LIMIT ?`, nowMs, limit) if err != nil { return 0, err } - var list []due + var list []dueMsg for rows.Next() { - var d due + var d dueMsg if err := rows.Scan(&d.seq, &d.senderID, &d.destKind, &d.destID, &d.sendAt, &d.keep, &d.ttl, &d.receipt); err != nil { _ = rows.Close() return 0, err @@ -49,55 +96,68 @@ LIMIT ?`, nowMs, limit) return 0, err } _ = rows.Close() + if len(list) == 0 { + return 0, nil + } + sem := make(chan struct{}, dispatchConcurrency) + var wg sync.WaitGroup + var mu sync.Mutex n := 0 wake := map[string]struct{}{} for _, d := range list { - var claimed bool - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - _, c, err := a.dispatchFullTx(tx, d.seq, d.senderID, d.destKind, d.destID, d.sendAt, d.keep, d.ttl, d.receipt != 0, nowMs) - claimed = c - return err - }) - if err != nil { - return n, err - } - if !claimed { - continue - } - n++ - rows2, qErr := a.db.Read.QueryContext(ctx, ` + d := d + wg.Add(1) + sem <- struct{}{} + go func() { + defer wg.Done() + defer func() { <-sem }() + var claimed bool + err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + _, c, err := a.dispatchFullTx(tx, d.seq, d.senderID, d.destKind, d.destID, d.sendAt, d.keep, d.ttl, d.receipt != 0, nowMs) + claimed = c + return err + }) + if err != nil { + slog.Error("dispatch due item", "seq", d.seq, "err", err) + return + } + if !claimed { + return + } + mu.Lock() + n++ + mu.Unlock() + rows2, qErr := a.db.Read.QueryContext(ctx, ` SELECT DISTINCT endpoint_id FROM deliveries WHERE seq = ? AND state = 'pending'`, d.seq) - if qErr == nil { + if qErr != nil { + return + } for rows2.Next() { var ep string if rows2.Scan(&ep) == nil { + mu.Lock() wake[ep] = struct{}{} + mu.Unlock() } } _ = rows2.Close() - } + }() } + wg.Wait() for ep := range wake { a.WakePush(ep) } return n, nil } -// PushPending 向指定连接推送 pending 投递与回执。 +// PushPending 向已握手连接推送 pending 投递与回执。 func (a *App) PushPending(ctx context.Context, endpointID string, connID port.ConnID) error { - nowMs := a.now().UnixMilli() - live, ok := a.lookupConn(endpointID) - if !ok || (connID != "" && live.ConnID != connID) { - // 仍处理该代号上的确认超时与清标记场景:用传入 connID - if connID == "" { - return nil - } - live = LiveConn{ConnID: connID} - } else if connID == "" { - connID = live.ConnID + live, connID, ok := a.canPush(endpointID, connID) + if !ok { + return nil } - + nowMs := a.now().UnixMilli() if err := a.processAckTimeouts(ctx, endpointID, connID, nowMs); err != nil { return err } @@ -128,31 +188,13 @@ LIMIT ?`, endpointID, room) if err != nil { return err } - type item struct { - seq int64 - sendAt int64 - keep int - expireAt sql.NullInt64 - msgID string - senderID string - destKind string - destID string - meta string - contentType string - bodyEnc string - body []byte - } - var items []item + var items []pushItem for rows.Next() { - var it item - var body sql.NullString - var bodyBlob []byte - if err := rows.Scan(&it.seq, &it.sendAt, &it.keep, &it.expireAt, &it.msgID, &it.senderID, &it.destKind, &it.destID, &it.meta, &it.contentType, &it.bodyEnc, &bodyBlob); err != nil { + var it pushItem + if err := rows.Scan(&it.seq, &it.sendAt, &it.keep, &it.expireAt, &it.msgID, &it.senderID, &it.destKind, &it.destID, &it.meta, &it.contentType, &it.bodyEnc, &it.body); err != nil { _ = rows.Close() return err } - _ = body - it.body = bodyBlob items = append(items, it) } _ = rows.Close() @@ -160,9 +202,9 @@ LIMIT ?`, endpointID, room) return err } + var toClaim []pushItem for _, it := range items { if it.body == nil { - // 正文已删则跳过(异常) continue } msg := protocol.Msg{ @@ -186,9 +228,40 @@ LIMIT ?`, endpointID, room) } continue } + it.payload = payload + toClaim = append(toClaim, it) + } - claimed := false - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + claimed, err := a.claimPushBatch(ctx, endpointID, connID, nowMs, toClaim) + if err != nil { + return err + } + + for _, it := range claimed { + if a.down == nil { + continue + } + pubErr := a.down.PublishDown(ctx, endpointID, connID, it.payload, port.PublishOpts{QoS: 1}) + if pubErr != nil { + short, cancel := shortWriteCtx() + _ = a.clearPushed(short, it.seq, endpointID, connID, a.now().UnixMilli()) + cancel() + a.scheduleRepush(endpointID, time.Second) + } else { + a.observeDispatchToPush(it.sendAt, nowMs) + } + } + return a.pushReceipts(ctx, endpointID, connID, nowMs) +} + +func (a *App) claimPushBatch(ctx context.Context, endpointID string, connID port.ConnID, nowMs int64, items []pushItem) ([]pushItem, error) { + if len(items) == 0 { + return nil, nil + } + var claimed []pushItem + err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + claimed = claimed[:0] + for _, it := range items { res, e := tx.Exec(` UPDATE deliveries SET pushed_conn = ?, pushed_at = ?, attempts = attempts + 1, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL`, @@ -197,43 +270,35 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL` return e } aff, _ := res.RowsAffected() - claimed = aff > 0 - return nil - }) - if err != nil { - return err - } - if !claimed { - continue - } - - large := len(payload) > largeFrameBytes - if large { - if !a.acquireLarge(ctx) { - _ = a.clearPushed(ctx, it.seq, endpointID, connID, nowMs) - continue + if aff > 0 { + claimed = append(claimed, it) } - a.trackLarge(it.seq, endpointID, true) } - - if a.down == nil { - if large { - a.releaseLarge(it.seq, endpointID) - } - continue - } - pubErr := a.down.PublishDown(ctx, endpointID, connID, payload, port.PublishOpts{QoS: 1}) - if pubErr != nil { - _ = a.clearPushed(ctx, it.seq, endpointID, connID, nowMs) - if large { - a.releaseLarge(it.seq, endpointID) - } + return nil + }) + if err != nil { + if ctx.Err() != nil { + short, cancel := shortWriteCtx() + a.clearIfClaimed(short, items, endpointID, connID) + cancel() a.scheduleRepush(endpointID, time.Second) - } else { - a.observeDispatchToPush(it.sendAt, nowMs) + } + return nil, err + } + return claimed, nil +} + +func (a *App) clearIfClaimed(ctx context.Context, items []pushItem, endpointID string, connID port.ConnID) { + nowMs := a.now().UnixMilli() + for _, it := range items { + var pushed sql.NullString + _ = a.db.Read.QueryRowContext(ctx, ` +SELECT pushed_conn FROM deliveries WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, + it.seq, endpointID).Scan(&pushed) + if pushed.Valid && pushed.String == string(connID) { + _ = a.clearPushed(ctx, it.seq, endpointID, connID, nowMs) } } - return a.pushReceipts(ctx, endpointID, connID, nowMs) } func (a *App) observeDispatchToPush(sendAtMs, pushedAtMs int64) { @@ -256,7 +321,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL` if aff == 0 { return nil } - if err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, ReasonTooLarge, nowMs); err != nil { + if _, err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, ReasonTooLarge, nowMs); err != nil { return err } return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays) @@ -264,6 +329,11 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL` } func (a *App) clearPushed(ctx context.Context, seq int64, endpointID string, connID port.ConnID, nowMs int64) error { + if err := ctx.Err(); err != nil { + var cancel context.CancelFunc + ctx, cancel = shortWriteCtx() + defer cancel() + } return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { _, err := tx.Exec(` UPDATE deliveries SET pushed_conn = NULL, updated_at = ? @@ -291,14 +361,14 @@ WHERE d.endpoint_id = ? AND d.state = 'pending' AND d.pushed_conn = ? seq int64 keep int expireAt sql.NullInt64 + pushedAt int64 senderID string msgID string } var list []to for rows.Next() { var t to - var pushedAt int64 - if err := rows.Scan(&t.seq, &t.keep, &t.expireAt, &pushedAt, &t.senderID, &t.msgID); err != nil { + if err := rows.Scan(&t.seq, &t.keep, &t.expireAt, &t.pushedAt, &t.senderID, &t.msgID); err != nil { _ = rows.Close() return err } @@ -311,11 +381,12 @@ WHERE d.endpoint_id = ? AND d.state = 'pending' AND d.pushed_conn = ? var keep int var expireAt sql.NullInt64 var pushedConn sql.NullString + var pushedAt sql.NullInt64 err := tx.QueryRow(` -SELECT keep, expire_at, pushed_conn FROM deliveries -WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, t.seq, endpointID).Scan(&keep, &expireAt, &pushedConn) +SELECT keep, expire_at, pushed_conn, pushed_at FROM deliveries +WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, t.seq, endpointID).Scan(&keep, &expireAt, &pushedConn, &pushedAt) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil } return err @@ -323,13 +394,15 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, t.seq, endpointID).Sca if !pushedConn.Valid || pushedConn.String != string(connID) { return nil } + if !pushedAt.Valid || pushedAt.Int64 > nowMs-timeoutMs { + return nil + } if keep == 0 { return a.finishDeliveryTx(tx, t.seq, endpointID, t.senderID, t.msgID, DeliveryDropped, ReasonNotAcked, true, nowMs) } if expireAt.Valid && expireAt.Int64 <= nowMs { return a.finishDeliveryTx(tx, t.seq, endpointID, t.senderID, t.msgID, DeliveryExpired, ReasonTTL, true, nowMs) } - // 重推:清标记 _, err = tx.Exec(` UPDATE deliveries SET pushed_conn = NULL, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, @@ -339,7 +412,6 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, if err != nil { return err } - a.releaseLarge(t.seq, endpointID) } return nil } @@ -357,7 +429,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, return nil } if state != DeliveryRecalled { - if err := insertReceiptTx(tx, senderID, seq, endpointID, state, reason, nowMs); err != nil { + if _, err := insertReceiptTx(tx, senderID, seq, endpointID, state, reason, nowMs); err != nil { return err } } @@ -406,6 +478,7 @@ func (a *App) flushRevokes(ctx context.Context) { if a.down == nil { return } + var retry []revokeJob for _, j := range jobs { frame := protocol.Revoked{ V: protocol.Version, Type: protocol.TypeRevoked, @@ -415,26 +488,48 @@ func (a *App) flushRevokes(ctx context.Context) { if err != nil { continue } - _ = a.down.PublishDown(ctx, j.endpointID, j.connID, payload, port.PublishOpts{QoS: 1}) + if err := a.down.PublishDown(ctx, j.endpointID, j.connID, payload, port.PublishOpts{QoS: 1}); err != nil { + retry = append(retry, j) + a.scheduleRepush(j.endpointID, time.Second) + } + } + if len(retry) > 0 { + a.mu.Lock() + a.pendingRevoke = append(retry, a.pendingRevoke...) + a.mu.Unlock() } } // OnPublishDropped 清推送标记并 1 秒后重推。 func (a *App) OnPublishDropped(ctx context.Context, endpointID string, connID port.ConnID, payload []byte) error { var head struct { - Type string `json:"type"` - ID string `json:"id"` - From string `json:"from"` + Type string `json:"type"` + ID string `json:"id"` + From string `json:"from"` + ReceiptID string `json:"receipt_id"` } - if err := json.Unmarshal(payload, &head); err != nil || head.Type != protocol.TypeMsg { + if err := json.Unmarshal(payload, &head); err != nil { + return nil + } + switch head.Type { + case protocol.TypeReceipt: + if rid, err := parseReceiptID(head.ReceiptID); err == nil { + a.unmarkReceipt(string(connID), rid) + } + a.scheduleRepush(endpointID, time.Second) + return nil + case protocol.TypeMsg: + default: return nil } nowMs := a.now().UnixMilli() - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + short, cancel := shortWriteCtx() + defer cancel() + err := a.db.Queue.Do(short, func(tx *sql.Tx) error { var seq int64 err := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id = ? AND id = ?`, head.From, head.ID).Scan(&seq) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil } return err @@ -443,7 +538,6 @@ func (a *App) OnPublishDropped(ctx context.Context, endpointID string, connID po UPDATE deliveries SET pushed_conn = NULL, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, nowMs, seq, endpointID, string(connID)) - a.releaseLarge(seq, endpointID) return err }) if err != nil { @@ -458,77 +552,128 @@ func (a *App) pushReceipts(ctx context.Context, endpointID string, connID port.C if window <= 0 { window = defaultReceiptWindow } - // 简化:未单独记 inflight 回执,按未确认回执取窗口条数 + retryAfter := a.receiptRetryMs() + a.mu.Lock() + held := a.rcptInflight[string(connID)] + inflightN := 0 + stale := map[int64]struct{}{} + for rid, at := range held { + if nowMs-at >= retryAfter { + stale[rid] = struct{}{} + continue + } + inflightN++ + } + a.mu.Unlock() + room := window - inflightN + if room <= 0 { + return nil + } + rows, err := a.db.Read.QueryContext(ctx, ` SELECT receipt_id, msg_id, endpoint_id, state, reason, created_at FROM receipts WHERE sender_id = ? AND acked = 0 ORDER BY receipt_id ASC -LIMIT ?`, endpointID, window) +LIMIT ?`, endpointID, window+len(stale)) if err != nil { return err } - defer func() { _ = rows.Close() }() + type rcpt struct { + rid int64 + msgID, epID string + state, reason string + created int64 + } + var list []rcpt + for rows.Next() { + var r rcpt + if err := rows.Scan(&r.rid, &r.msgID, &r.epID, &r.state, &r.reason, &r.created); err != nil { + _ = rows.Close() + return err + } + list = append(list, r) + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return err + } + _ = rows.Close() + if a.down == nil { return nil } - for rows.Next() { - var rid int64 - var msgID, epID, state, reason string - var created int64 - if err := rows.Scan(&rid, &msgID, &epID, &state, &reason, &created); err != nil { - return err + sent := 0 + for _, r := range list { + if sent >= room { + break } + a.mu.Lock() + m := a.rcptInflight[string(connID)] + at, in := m[r.rid] + fresh := in && nowMs-at < retryAfter + a.mu.Unlock() + if fresh { + continue + } + a.markReceipt(string(connID), r.rid, nowMs) frame := protocol.Receipt{ V: protocol.Version, Type: protocol.TypeReceipt, - ReceiptID: fmt.Sprintf("%d", rid), ID: msgID, EndpointID: epID, - State: state, Reason: reason, AtMs: created, + ReceiptID: fmt.Sprintf("%d", r.rid), ID: r.msgID, EndpointID: r.epID, + State: r.state, Reason: r.reason, AtMs: r.created, } payload, err := protocol.Marshal(frame) if err != nil { + a.unmarkReceipt(string(connID), r.rid) return err } - _ = a.down.PublishDown(ctx, endpointID, connID, payload, port.PublishOpts{QoS: 1}) - } - _ = nowMs - return rows.Err() -} - -func (a *App) acquireLarge(ctx context.Context) bool { - select { - case a.largeSem <- struct{}{}: - return true - case <-ctx.Done(): - return false - default: - return false - } -} - -func (a *App) trackLarge(seq int64, endpointID string, hold bool) { - a.mu.Lock() - defer a.mu.Unlock() - key := largeKey(seq, endpointID) - if hold { - a.largeHeld[key] = true - } -} - -func (a *App) releaseLarge(seq int64, endpointID string) { - a.mu.Lock() - defer a.mu.Unlock() - key := largeKey(seq, endpointID) - if a.largeHeld[key] { - delete(a.largeHeld, key) - select { - case <-a.largeSem: - default: + if err := a.down.PublishDown(ctx, endpointID, connID, payload, port.PublishOpts{QoS: 1}); err != nil { + a.unmarkReceipt(string(connID), r.rid) + a.scheduleRepush(endpointID, time.Second) + continue } + sent++ + } + return nil +} + +func (a *App) receiptRetryMs() int64 { + sec := a.lim.AckTimeoutSeconds + if sec <= 0 { + sec = 60 + } + return sec * 1000 +} + +func (a *App) markReceipt(connID string, rid, nowMs int64) { + a.mu.Lock() + defer a.mu.Unlock() + m := a.rcptInflight[connID] + if m == nil { + m = make(map[int64]int64) + a.rcptInflight[connID] = m + } + m[rid] = nowMs +} + +func (a *App) unmarkReceipt(connID string, rid int64) { + a.mu.Lock() + defer a.mu.Unlock() + if m := a.rcptInflight[connID]; m != nil { + delete(m, rid) } } -func largeKey(seq int64, endpointID string) string { - return fmt.Sprintf("%d:%s", seq, endpointID) +func (a *App) clearReceiptInflight(connID string) { + a.mu.Lock() + defer a.mu.Unlock() + delete(a.rcptInflight, connID) +} + +func parseReceiptID(s string) (int64, error) { + var n int64 + _, err := fmt.Sscan(s, &n) + return n, err } func (a *App) scheduleRepush(endpointID string, d time.Duration) { @@ -545,16 +690,11 @@ func (a *App) scheduleRepush(endpointID string, d time.Duration) { }) } -// WakePush 唤醒推送;若有登记的连接则异步 PushPending。 +// WakePush 唤醒该端已握手连接的推送 worker(合并唤醒)。 func (a *App) WakePush(endpointID string) { - live, ok := a.lookupConn(endpointID) - if !ok || a.down == nil { + _, _, ok := a.canPush(endpointID, "") + if !ok { return } - go func() { - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - _ = a.PushPending(ctx, endpointID, live.ConnID) - a.flushRevokes(ctx) - }() + a.signalWorker(endpointID) } diff --git a/internal/app/message/recover.go b/internal/app/message/recover.go index 2754928..e01e6ac 100644 --- a/internal/app/message/recover.go +++ b/internal/app/message/recover.go @@ -3,15 +3,16 @@ package message import ( "context" "database/sql" + "time" ) -// RecoverOnStart 启动恢复(DEVELOPMENT 7.8)。 +// RecoverOnStart 启动恢复(DEVELOPMENT 7.8):只做 SQL 修正,分发交给调度循环。 func (a *App) RecoverOnStart(ctx context.Context) error { nowMs := a.now().UnixMilli() graceMs := a.lim.GraceSeconds * 1000 minExpire := nowMs + graceMs - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { if _, err := tx.Exec(` UPDATE deliveries SET pushed_conn = NULL, @@ -23,30 +24,57 @@ UPDATE deliveries SET WHERE state = 'pending'`, minExpire, minExpire, nowMs); err != nil { return err } - // 停机前在线:online_since 晚于 offline_since,或 offline_since 空而 online_since 非空 _, err := tx.Exec(` UPDATE endpoints SET offline_since = ? WHERE online_since IS NOT NULL AND (offline_since IS NULL OR online_since > offline_since)`, nowMs) return err }) - if err != nil { - return err - } - // 停机期间到点的 scheduled 立即分发 - _, err = a.DispatchDue(ctx, nowMs, 1000) - return err } -// CleanupOnce 处理未推送且到期的 pending,并做记录/回执/防重清理。 +// CleanupOnce 兼容旧调用:先到期处理再做一次保留清理。 func (a *App) CleanupOnce(ctx context.Context, nowMs int64) error { + if err := a.ExpireOnce(ctx, nowMs); err != nil { + return err + } + return a.PurgeOnce(ctx, nowMs) +} + +// ExpireOnce 处理未推送且到期的 pending,并给僵尸不保留投递补宽限。每写操作最多 expireBatch 条。 +func (a *App) ExpireOnce(ctx context.Context, nowMs int64) error { + deadline := time.Now().Add(expireBudget) + for { + if err := ctx.Err(); err != nil { + return err + } + if time.Now().After(deadline) { + return nil + } + n, err := a.expireOnceBatch(ctx, nowMs) + if err != nil { + return err + } + if n == 0 { + break + } + } + a.flushRevokes(ctx) + return nil +} + +func (a *App) expireOnceBatch(ctx context.Context, nowMs int64) (int, error) { + var n int err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + if err := a.fillZombieExpireTx(tx, nowMs); err != nil { + return err + } rows, err := tx.Query(` SELECT d.seq, d.endpoint_id, d.keep, m.sender_id, m.id FROM deliveries d JOIN messages m ON m.seq = d.seq WHERE d.state = 'pending' AND d.pushed_conn IS NULL - AND d.expire_at IS NOT NULL AND d.expire_at <= ?`, nowMs) + AND d.expire_at IS NOT NULL AND d.expire_at <= ? +LIMIT ?`, nowMs, expireBatch) if err != nil { return err } @@ -65,7 +93,7 @@ WHERE d.state = 'pending' AND d.pushed_conn IS NULL list = append(list, it) } _ = rows.Close() - + n = len(list) for _, it := range list { state := DeliveryDropped reason := ReasonOffline @@ -77,63 +105,197 @@ WHERE d.state = 'pending' AND d.pushed_conn IS NULL return err } } + return finalizeStuckDispatchedTx(tx, nowMs, a.lim.RecordRetentionDays) + }) + return n, err +} - if err := finalizeStuckDispatchedTx(tx, nowMs, a.lim.RecordRetentionDays); err != nil { +func (a *App) fillZombieExpireTx(tx *sql.Tx, nowMs int64) error { + graceMs := a.lim.GraceSeconds * 1000 + if graceMs < 0 { + graceMs = 0 + } + deadline := nowMs + graceMs + rows, err := tx.Query(` +SELECT DISTINCT endpoint_id FROM deliveries +WHERE state = 'pending' AND keep = 0 AND pushed_conn IS NULL AND expire_at IS NULL +LIMIT 500`) + if err != nil { + return err + } + seen := map[string]struct{}{} + var eps []string + for rows.Next() { + var ep string + if err := rows.Scan(&ep); err != nil { + _ = rows.Close() return err } + if _, ok := seen[ep]; ok { + continue + } + seen[ep] = struct{}{} + if a.isReadyEndpoint(ep) { + continue + } + eps = append(eps, ep) + } + _ = rows.Close() + for _, ep := range eps { + if _, err := tx.Exec(` +UPDATE deliveries SET expire_at = ?, updated_at = ? +WHERE endpoint_id = ? AND state = 'pending' AND keep = 0 AND pushed_conn IS NULL AND expire_at IS NULL`, + deadline, nowMs, ep); err != nil { + return err + } + } + return nil +} - if a.lim.RecordRetentionDays > 0 { - cutoff := nowMs - int64(a.lim.RecordRetentionDays)*24*3600*1000 - if _, err := tx.Exec(` -DELETE FROM messages WHERE seq IN ( - 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 { +// PurgeOnce 分批删除过期记录、回执和防重行,然后 wal_checkpoint + optimize。 +func (a *App) PurgeOnce(ctx context.Context, nowMs int64) error { + a.purgeMu.Lock() + a.lastPurge = a.now() + a.purgeMu.Unlock() + + if err := a.purgeCompletedMessages(ctx, nowMs); err != nil { + return err + } + if err := a.purgeReceipts(ctx, nowMs); err != nil { + return err + } + if err := a.purgeSendKeys(ctx, nowMs); err != nil { + return err + } + if a.db != nil && a.db.Queue != nil { + if err := a.db.Queue.Checkpoint(ctx); err != nil { + return err + } + _ = a.db.Queue.Optimize(ctx) + } + return nil +} + +func (a *App) purgeCompletedMessages(ctx context.Context, nowMs int64) error { + days := a.lim.RecordRetentionDays + if days < 0 { + return nil + } + for { + if err := ctx.Err(); err != nil { + return err + } + var seqs []int64 + err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + q := ` +SELECT m.seq FROM messages m +WHERE m.state = 'completed'` + args := []any{} + if days > 0 { + cutoff := nowMs - int64(days)*24*3600*1000 + q += ` AND COALESCE( + (SELECT MAX(d.updated_at) FROM deliveries d WHERE d.seq = m.seq), + m.send_at +) < ?` + args = append(args, cutoff) + } + q += ` LIMIT ?` + args = append(args, purgeMsgBatch) + rows, err := tx.Query(q, args...) + if err != nil { return err } + for rows.Next() { + var seq int64 + if err := rows.Scan(&seq); err != nil { + _ = rows.Close() + return err + } + seqs = append(seqs, seq) + } + _ = rows.Close() + if len(seqs) == 0 { + return nil + } + for _, seq := range seqs { + if _, err := tx.Exec(`DELETE FROM deliveries WHERE seq = ?`, seq); err != nil { + return err + } + if _, err := tx.Exec(`DELETE FROM messages WHERE seq = ?`, seq); err != nil { + return err + } + } + return nil + }) + if err != nil { + return err } + if len(seqs) == 0 { + return nil + } + } +} - if a.lim.ReceiptRetentionDays > 0 { - cutoff := nowMs - int64(a.lim.ReceiptRetentionDays)*24*3600*1000 - if _, err := tx.Exec(` +func (a *App) purgeReceipts(ctx context.Context, nowMs int64) error { + if a.lim.ReceiptRetentionDays <= 0 { + return nil + } + cutoff := nowMs - int64(a.lim.ReceiptRetentionDays)*24*3600*1000 + for { + var n int64 + err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + res, err := tx.Exec(` DELETE FROM receipts WHERE receipt_id IN ( - SELECT receipt_id FROM receipts WHERE created_at < ? LIMIT 5000 -)`, cutoff); err != nil { + SELECT receipt_id FROM receipts WHERE created_at < ? LIMIT ? +)`, cutoff, purgeRowBatch) + if err != nil { return err } + n, _ = res.RowsAffected() + return nil + }) + if err != nil { + return err } + if n == 0 { + return nil + } + } +} - if a.lim.IdempotencyHours > 0 { - cutoff := nowMs - int64(a.lim.IdempotencyHours)*3600*1000 - if _, err := tx.Exec(` +func (a *App) purgeSendKeys(ctx context.Context, nowMs int64) error { + if a.lim.IdempotencyHours <= 0 { + return nil + } + cutoff := nowMs - int64(a.lim.IdempotencyHours)*3600*1000 + for { + var n int64 + err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + res, err := tx.Exec(` DELETE FROM send_keys WHERE rowid IN ( SELECT sk.rowid FROM send_keys sk WHERE sk.created_at < ? AND NOT EXISTS ( SELECT 1 FROM messages m WHERE m.sender_id = sk.sender_id AND m.id = sk.msg_id ) - LIMIT 5000 -)`, cutoff); err != nil { + LIMIT ? +)`, cutoff, purgeRowBatch) + if err != nil { return err } + n, _ = res.RowsAffected() + return nil + }) + if err != nil { + return err + } + if n == 0 { + return nil } - return nil - }) - if err != nil { - return err } - a.flushRevokes(ctx) - return nil } -// finalizeStuckDispatchedTx 收尾「dispatched 且已无 pending 投递」的消息(C-04 兜底,修复已卡住的数据)。 +// finalizeStuckDispatchedTx 收尾「dispatched 且已无 pending 投递」的消息(C-04 兜底)。 func finalizeStuckDispatchedTx(tx *sql.Tx, nowMs int64, recordDays int) error { rows, err := tx.Query(` SELECT seq FROM messages diff --git a/internal/app/message/review_c01_test.go b/internal/app/message/review_c01_test.go new file mode 100644 index 0000000..cddb7f1 --- /dev/null +++ b/internal/app/message/review_c01_test.go @@ -0,0 +1,252 @@ +package message + +import ( + "context" + "database/sql" + "sync" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "git.asio.asia/nixevol/NixMsg/internal/broker" + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +func TestC01NotReadyDoesNotPush(t *testing.T) { + t.Parallel() + e := openDeliveryEnv(t, nil) + insertEndpoint(t, e.db, "alice", "", 1, 0) + insertEndpoint(t, e.db, "bob", "", 1, 0) + e.conns.Set("bob", LiveConn{ConnID: "c-bob"}) // Ready=false + ctx := context.Background() + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("nr1", "bob")); err != nil { + t.Fatal(err) + } + e.app.WakePush("bob") + if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { + t.Fatal(err) + } + if e.down.FilterType(protocol.TypeMsg) != 0 { + t.Fatalf("pushed before handshake: %d", e.down.FilterType(protocol.TypeMsg)) + } + seq := e.seqOf("alice", "nr1") + var pushed sql.NullString + _ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed) + if pushed.Valid { + t.Fatalf("pushed_conn=%s", pushed.String) + } + live := LiveConn{ConnID: "c-bob", Ready: true} + e.conns.Set("bob", live) + if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { + t.Fatal(err) + } + if e.down.FilterType(protocol.TypeMsg) != 1 { + t.Fatalf("want 1 msg after ready, got %d", e.down.FilterType(protocol.TypeMsg)) + } +} + +func TestC01WindowAndOrderWithConcurrentWake(t *testing.T) { + t.Parallel() + e := openDeliveryEnv(t, func(l *Limits) { l.DeliveryWindow = 2 }) + insertEndpoint(t, e.db, "alice", "", 1, 0) + insertEndpoint(t, e.db, "bob", "", 1, 0) + ctx := context.Background() + e.online("bob", "c-bob") + for i := 0; i < 5; i++ { + req := baseSend("w"+string(rune('a'+i)), "bob") + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { + t.Fatal(err) + } + } + if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { + t.Fatal(err) + } + if n := e.down.FilterType(protocol.TypeMsg); n != 2 { + t.Fatalf("window: got %d want 2", n) + } + var inflight int + _ = e.db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries WHERE endpoint_id='bob' AND state='pending' AND pushed_conn IS NOT NULL`).Scan(&inflight) + if inflight > 2 { + t.Fatalf("inflight=%d", inflight) + } +} + +func TestC01WorkerCoalescesWake(t *testing.T) { + t.Parallel() + e := openDeliveryEnv(t, func(l *Limits) { l.DeliveryWindow = 32 }) + insertEndpoint(t, e.db, "alice", "", 1, 0) + insertEndpoint(t, e.db, "bob", "", 1, 0) + ctx := context.Background() + live := LiveConn{ConnID: "c-bob", Ready: true} + e.conns.Set("bob", live) + for i := 0; i < 5; i++ { + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("cw"+string(rune('a'+i)), "bob")); err != nil { + t.Fatal(err) + } + } + if err := e.app.OnHandshakeComplete(ctx, "bob", live); err != nil { + t.Fatal(err) + } + var wg sync.WaitGroup + for i := 0; i < 50; i++ { + wg.Add(1) + go func() { + defer wg.Done() + e.app.WakePush("bob") + }() + } + wg.Wait() + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if e.down.FilterType(protocol.TypeMsg) >= 5 { + break + } + time.Sleep(10 * time.Millisecond) + } + n := e.down.FilterType(protocol.TypeMsg) + if n != 5 { + t.Fatalf("got %d msg frames want 5", n) + } +} + +func TestC01BrokerPublishErrorsClearClaim(t *testing.T) { + t.Parallel() + e := openDeliveryEnv(t, nil) + insertEndpoint(t, e.db, "alice", "", 1, 0) + insertEndpoint(t, e.db, "bob", "", 1, 0) + e.online("bob", "c-bob") + e.down.FailNext = 1 + e.down.FailErr = broker.ErrBackpressure + ctx := context.Background() + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("bp1", "bob")); err != nil { + t.Fatal(err) + } + if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil { + t.Fatal(err) + } + seq := e.seqOf("alice", "bp1") + var pushed sql.NullString + _ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed) + if pushed.Valid { + t.Fatalf("claim left after backpressure: %s", pushed.String) + } +} + +type gatedDown struct { + blockBob chan struct{} + inner *RecordingDownlink +} + +func (g *gatedDown) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { + if endpointID == "bob" && g.blockBob != nil { + <-g.blockBob + } + return g.inner.PublishDown(ctx, endpointID, connID, payload, opts) +} + +func TestC01BlockedPushDoesNotFreezeDispatch(t *testing.T) { + t.Parallel() + e := openDeliveryEnv(t, nil) + insertEndpoint(t, e.db, "alice", "", 1, 0) + insertEndpoint(t, e.db, "bob", "", 1, 0) + insertEndpoint(t, e.db, "carol", "", 1, 0) + ctx := context.Background() + block := make(chan struct{}) + gate := &gatedDown{blockBob: block, inner: e.down} + e.app.down = gate + live := LiveConn{ConnID: "c-bob", Ready: true} + e.conns.Set("bob", live) + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("blk1", "bob")); err != nil { + t.Fatal(err) + } + if err := e.app.OnHandshakeComplete(ctx, "bob", live); err != nil { + t.Fatal(err) + } + delay := int64(5_000) + req := baseSend("duex", "carol") + req.DelayMs = &delay + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { + t.Fatal(err) + } + e.setNow(e.nowMs + 5_000) + done := make(chan error, 1) + go func() { + _, err := e.app.DispatchDue(ctx, e.nowMs, 10) + done <- err + }() + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("dispatch blocked by slow push") + } + st, _ := e.msgState("alice", "duex") + if st != StateDispatched && st != StateCompleted { + t.Fatalf("state=%s", st) + } + close(block) +} + +func TestC01DispatchDueBudgetAndSkipError(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() + delay := int64(10_000) + const n = 250 + for i := 0; i < n; i++ { + req := baseSend("d"+itoa(i), "bob") + req.DelayMs = &delay + req.RID = itoa(i) + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { + t.Fatal(err) + } + } + bad := baseSend("bad1", "bob") + bad.DelayMs = &delay + if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, bad); err != nil { + t.Fatal(err) + } + _ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error { + _, err := tx.Exec(`UPDATE messages SET dest_kind='nope' WHERE sender_id=? AND id=?`, "alice", "bad1") + return err + }) + e.setNow(e.nowMs + 10_000) + start := time.Now() + total := 0 + for time.Since(start) < time.Second { + k, err := e.app.DispatchDue(ctx, e.nowMs, 0) + if err != nil { + t.Fatal(err) + } + total += k + if k == 0 { + break + } + } + if total < n { + t.Fatalf("dispatched %d want %d in 1s", total, n) + } + var badState string + _ = e.db.Read.QueryRow(`SELECT state FROM messages WHERE sender_id=? AND id=?`, "alice", "bad1").Scan(&badState) + if badState != StateScheduled { + t.Fatalf("bad message state=%s", badState) + } +} + +func itoa(i int) string { + if i == 0 { + return "0" + } + var b [16]byte + pos := len(b) + for i > 0 { + pos-- + b[pos] = byte('0' + i%10) + i /= 10 + } + return string(b[pos:]) +} diff --git a/internal/app/message/session.go b/internal/app/message/session.go index 2b259f9..7b12c82 100644 --- a/internal/app/message/session.go +++ b/internal/app/message/session.go @@ -7,7 +7,7 @@ import ( "git.asio.asia/nixevol/NixMsg/internal/app/port" ) -// OnHandshakeComplete 握手完成:写 online_since、清空不 keep 的 expire_at,并推送。 +// OnHandshakeComplete 握手完成:写 online_since、清空不 keep 的 expire_at,启动推送 worker。 // 调用方须先把连接登记进 ConnRegistry(MemoryConns.Set)。 func (a *App) OnHandshakeComplete(ctx context.Context, endpointID string, conn LiveConn) error { nowMs := a.now().UnixMilli() @@ -15,15 +15,23 @@ func (a *App) OnHandshakeComplete(ctx context.Context, endpointID string, conn L if _, err := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, nowMs, endpointID); err != nil { return err } - _, err := tx.Exec(` + if _, err := tx.Exec(` UPDATE deliveries SET expire_at = NULL, updated_at = ? -WHERE endpoint_id = ? AND state = 'pending' AND keep = 0`, nowMs, endpointID) +WHERE endpoint_id = ? AND state = 'pending' AND keep = 0`, nowMs, endpointID); err != nil { + return err + } + _, err := tx.Exec(` +UPDATE deliveries SET pushed_conn = NULL, updated_at = ? +WHERE endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, nowMs, endpointID, string(conn.ConnID)) return err }) if err != nil { return err } - return a.PushPending(ctx, endpointID, conn.ConnID) + a.markHandshook(endpointID, conn) + a.startPushWorker(endpointID, conn.ConnID) + a.WakePush(endpointID) + return nil } // OnDisconnect 连接断开:当前连接则延长宽限;按代号清 pushed_conn。 @@ -33,6 +41,10 @@ func (a *App) OnDisconnect(ctx context.Context, endpointID string, connID port.C graceMs := a.lim.GraceSeconds * 1000 deadline := nowMs + graceMs + a.stopPushWorker(endpointID, connID) + a.clearHandshook(endpointID, connID) + a.clearReceiptInflight(string(connID)) + err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { if isCurrent { if _, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, nowMs, endpointID); err != nil { @@ -62,7 +74,7 @@ WHERE endpoint_id = ? AND state = 'pending' AND keep = 1 AND pushed_conn = ?`, } _, err := tx.Exec(` UPDATE deliveries SET pushed_conn = NULL, updated_at = ? -WHERE state = 'pending' AND pushed_conn = ?`, nowMs, string(connID)) +WHERE endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, nowMs, endpointID, string(connID)) return err }) if err != nil { diff --git a/internal/app/message/submit.go b/internal/app/message/submit.go index f69fa83..0d63252 100644 --- a/internal/app/message/submit.go +++ b/internal/app/message/submit.go @@ -276,6 +276,9 @@ INSERT INTO messages( if result.State == StateDispatched { a.wakeReceivers(ctx, result.ID, senderID) } + if result.State == StateScheduled { + a.NotifyDispatch() + } return result, nil } diff --git a/internal/app/message/worker.go b/internal/app/message/worker.go new file mode 100644 index 0000000..68b21c8 --- /dev/null +++ b/internal/app/message/worker.go @@ -0,0 +1,133 @@ +package message + +import ( + "context" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" +) + +type pushWorker struct { + endpointID string + connID port.ConnID + wake chan struct{} + stop chan struct{} + done chan struct{} +} + +func (a *App) startPushWorker(endpointID string, connID port.ConnID) { + a.stopPushWorker(endpointID, "") + w := &pushWorker{ + endpointID: endpointID, + connID: connID, + wake: make(chan struct{}, 1), + stop: make(chan struct{}), + done: make(chan struct{}), + } + a.mu.Lock() + a.workers[endpointID] = w + a.mu.Unlock() + go a.runPushWorker(w) +} + +func (a *App) stopPushWorker(endpointID string, connID port.ConnID) { + a.mu.Lock() + w, ok := a.workers[endpointID] + if !ok || (connID != "" && w.connID != connID) { + a.mu.Unlock() + return + } + delete(a.workers, endpointID) + a.mu.Unlock() + close(w.stop) + <-w.done +} + +func (a *App) runPushWorker(w *pushWorker) { + defer close(w.done) + for { + select { + case <-w.stop: + return + case <-w.wake: + ctx, cancel := context.WithTimeout(context.Background(), pushOpTimeout) + _ = a.PushPending(ctx, w.endpointID, w.connID) + a.flushRevokes(ctx) + cancel() + } + } +} + +func (a *App) signalWorker(endpointID string) { + a.mu.Lock() + w := a.workers[endpointID] + a.mu.Unlock() + if w == nil { + return + } + select { + case w.wake <- struct{}{}: + default: + } +} + +func (a *App) markHandshook(endpointID string, conn LiveConn) { + conn.Ready = true + a.mu.Lock() + a.handshook[endpointID] = conn.ConnID + a.mu.Unlock() + if setter, ok := a.conns.(interface { + Set(string, LiveConn) + }); ok { + setter.Set(endpointID, conn) + } +} + +func (a *App) clearHandshook(endpointID string, connID port.ConnID) { + a.mu.Lock() + defer a.mu.Unlock() + if cur, ok := a.handshook[endpointID]; ok && (connID == "" || cur == connID) { + delete(a.handshook, endpointID) + } +} + +func (a *App) canPush(endpointID string, connID port.ConnID) (LiveConn, port.ConnID, bool) { + live, ok := a.lookupConn(endpointID) + if !ok { + return LiveConn{}, "", false + } + if connID != "" && live.ConnID != connID { + return LiveConn{}, "", false + } + if connID == "" { + connID = live.ConnID + } + if live.Ready { + return live, connID, true + } + a.mu.Lock() + hs, marked := a.handshook[endpointID] + a.mu.Unlock() + if marked && hs == live.ConnID { + return live, connID, true + } + return LiveConn{}, "", false +} + +func (a *App) isReadyEndpoint(endpointID string) bool { + _, _, ok := a.canPush(endpointID, "") + return ok +} + +func shortWriteCtx() (context.Context, context.CancelFunc) { + return context.WithTimeout(context.Background(), clearPushedTimeout) +} + +func (a *App) NotifyDispatch() { + if a.dispatchCh == nil { + return + } + select { + case a.dispatchCh <- struct{}{}: + default: + } +}