fix: 仅向已握手连接推送并拆分调度循环

This commit is contained in:
Nixevol
2026-09-30 16:22:43 +08:00
parent ad4f13193c
commit 83521f4b87
14 changed files with 1192 additions and 290 deletions
+31 -8
View File
@@ -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
}
+17 -9
View File
@@ -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)
+21 -2
View File
@@ -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{
+2 -2
View File
@@ -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))
+28 -14
View File
@@ -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 {
+118
View File
@@ -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
}
+311 -171
View File
@@ -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)
}
+206 -44
View File
@@ -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
+252
View File
@@ -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:])
}
+17 -5
View File
@@ -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 {
+3
View File
@@ -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
}
+133
View File
@@ -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:
}
}