Files
NixMsg/internal/app/message/push.go
T

711 lines
18 KiB
Go

package message
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) {
budgeted := limit <= 0
batch := limit
if batch <= 0 {
batch = 64
}
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
WHERE state = 'scheduled' AND send_at <= ?
ORDER BY send_at ASC, seq ASC
LIMIT ?`, nowMs, limit)
if err != nil {
return 0, err
}
var list []dueMsg
for rows.Next() {
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
}
list = append(list, d)
}
if err := rows.Err(); err != nil {
_ = rows.Close()
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 {
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++
wake[d.senderID] = struct{}{}
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 {
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 投递与回执。
func (a *App) PushPending(ctx context.Context, endpointID string, connID port.ConnID) error {
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
}
window := a.lim.DeliveryWindow
if window <= 0 {
window = defaultDeliveryWindow
}
var inflight int
if err := a.db.Read.QueryRowContext(ctx, `
SELECT COUNT(*) FROM deliveries
WHERE endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`, endpointID, string(connID)).Scan(&inflight); err != nil {
return err
}
room := window - inflight
if room <= 0 {
return a.pushReceipts(ctx, endpointID, connID, nowMs)
}
rows, err := a.db.Read.QueryContext(ctx, `
SELECT d.seq, d.send_at, d.keep, d.expire_at, m.id, m.sender_id, m.dest_kind, m.dest_id, m.meta, m.content_type, m.body_enc, b.body
FROM deliveries d
JOIN messages m ON m.seq = d.seq
LEFT JOIN message_bodies b ON b.seq = d.seq
WHERE d.endpoint_id = ? AND d.state = 'pending' AND d.pushed_conn IS NULL
ORDER BY d.send_at ASC, d.seq ASC
LIMIT ?`, endpointID, room)
if err != nil {
return err
}
var items []pushItem
for rows.Next() {
var it pushItem
if scanErr := 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); scanErr != nil {
_ = rows.Close()
return scanErr
}
items = append(items, it)
}
_ = rows.Close()
if rowsErr := rows.Err(); rowsErr != nil {
return rowsErr
}
var toClaim []pushItem
skipped := false
for _, it := range items {
if it.body == nil {
skipped = true
continue
}
msg := protocol.Msg{
V: protocol.Version,
Type: protocol.TypeMsg,
ID: it.msgID,
From: it.senderID,
To: protocol.Target{Kind: it.destKind, ID: it.destID},
Body: encodeStoredBody(it.bodyEnc, it.contentType, it.body),
Meta: decodeMetaJSON(it.meta),
SendAtMs: it.sendAt,
}
payload, marshErr := protocol.Marshal(msg)
if marshErr != nil {
return marshErr
}
limit := effectivePayloadLimit(live.MaxPacketSize, live.MaxReceiveBytes)
if limit > 0 && len(payload) > limit {
if rejErr := a.rejectTooLarge(ctx, it.seq, endpointID, it.senderID, nowMs); rejErr != nil {
return rejErr
}
a.WakePush(it.senderID)
skipped = true
continue
}
it.payload = payload
toClaim = append(toClaim, it)
}
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)
}
}
// 超限等跳过未占满窗口时再唤醒本端,让 worker 下一轮取后续 pending(不在本轮递归扫表)。
if skipped && len(claimed) < room {
a.WakePush(endpointID)
}
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`,
string(connID), nowMs, nowMs, it.seq, endpointID)
if e != nil {
return e
}
aff, _ := res.RowsAffected()
if aff > 0 {
claimed = append(claimed, it)
}
}
return nil
})
if err != nil {
if ctx.Err() != nil {
short, cancel := shortWriteCtx()
a.clearIfClaimed(short, items, endpointID, connID)
cancel()
a.scheduleRepush(endpointID, time.Second)
}
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)
}
}
}
func (a *App) observeDispatchToPush(sendAtMs, pushedAtMs int64) {
if a.met == nil || pushedAtMs < sendAtMs {
return
}
a.met.DispatchToPushSeconds.Observe(float64(pushedAtMs-sendAtMs) / 1000.0)
}
func (a *App) rejectTooLarge(ctx context.Context, seq int64, endpointID, senderID string, nowMs int64) error {
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
res, err := tx.Exec(`
UPDATE deliveries SET state = ?, reason = ?, pushed_conn = NULL, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL`,
DeliveryRejected, ReasonTooLarge, nowMs, seq, endpointID)
if err != nil {
return err
}
aff, _ := res.RowsAffected()
if aff == 0 {
return nil
}
if _, err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, ReasonTooLarge, nowMs); err != nil {
return err
}
return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
})
}
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 = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`,
nowMs, seq, endpointID, string(connID))
return err
})
}
func (a *App) processAckTimeouts(ctx context.Context, endpointID string, connID port.ConnID, nowMs int64) error {
timeoutMs := a.lim.AckTimeoutSeconds * 1000
if timeoutMs <= 0 {
timeoutMs = 300 * 1000
}
rows, err := a.db.Read.QueryContext(ctx, `
SELECT d.seq, d.keep, d.expire_at, d.pushed_at, m.sender_id, m.id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.endpoint_id = ? AND d.state = 'pending' AND d.pushed_conn = ?
AND d.pushed_at IS NOT NULL AND d.pushed_at <= ?`, endpointID, string(connID), nowMs-timeoutMs)
if err != nil {
return err
}
type to struct {
seq int64
keep int
expireAt sql.NullInt64
pushedAt int64
senderID string
msgID string
}
var list []to
for rows.Next() {
var t to
if err := rows.Scan(&t.seq, &t.keep, &t.expireAt, &t.pushedAt, &t.senderID, &t.msgID); err != nil {
_ = rows.Close()
return err
}
list = append(list, t)
}
_ = rows.Close()
for _, t := range list {
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
var keep int
var expireAt sql.NullInt64
var pushedConn sql.NullString
var pushedAt sql.NullInt64
err := tx.QueryRow(`
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 errors.Is(err, sql.ErrNoRows) {
return nil
}
return err
}
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 = ?`,
nowMs, t.seq, endpointID, string(connID))
return err
})
if err != nil {
return err
}
a.WakePush(t.senderID)
}
return nil
}
func (a *App) finishDeliveryTx(tx *sql.Tx, seq int64, endpointID, senderID, msgID, state, reason string, sendRevoked bool, nowMs int64) error {
res, err := tx.Exec(`
UPDATE deliveries SET state = ?, reason = ?, pushed_conn = NULL, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
state, reason, nowMs, seq, endpointID)
if err != nil {
return err
}
aff, _ := res.RowsAffected()
if aff == 0 {
return nil
}
if state != DeliveryRecalled {
if _, err := insertReceiptTx(tx, senderID, seq, endpointID, state, reason, nowMs); err != nil {
return err
}
}
if err := TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays); err != nil {
return err
}
if sendRevoked {
a.mu.Lock()
a.pendingRevoke = append(a.pendingRevoke, revokeJob{
endpointID: endpointID,
msgID: msgID,
from: senderID,
reason: reasonForRevoked(state, reason),
})
a.mu.Unlock()
}
return nil
}
func reasonForRevoked(state, reason string) string {
switch state {
case DeliveryRecalled:
return ReasonRecalled
case DeliveryExpired:
return "expired"
case DeliveryDropped:
return "dropped"
default:
return reason
}
}
type revokeJob struct {
endpointID string
connID port.ConnID
msgID string
from string
reason string
}
func (a *App) flushRevokes(ctx context.Context) {
a.mu.Lock()
jobs := a.pendingRevoke
a.pendingRevoke = nil
a.mu.Unlock()
if a.down == nil {
return
}
var retry []revokeJob
for _, j := range jobs {
frame := protocol.Revoked{
V: protocol.Version, Type: protocol.TypeRevoked,
ID: j.msgID, From: j.from, Reason: j.reason,
}
payload, err := protocol.Marshal(frame)
if err != nil {
continue
}
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"`
ReceiptID string `json:"receipt_id"`
}
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()
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 errors.Is(err, sql.ErrNoRows) {
return nil
}
return err
}
_, err = tx.Exec(`
UPDATE deliveries SET pushed_conn = NULL, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn = ?`,
nowMs, seq, endpointID, string(connID))
return err
})
if err != nil {
return err
}
a.scheduleRepush(endpointID, time.Second)
return nil
}
func (a *App) pushReceipts(ctx context.Context, endpointID string, connID port.ConnID, nowMs int64) error {
window := a.lim.ReceiptWindow
if window <= 0 {
window = defaultReceiptWindow
}
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+len(stale))
if err != nil {
return err
}
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
}
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", 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
}
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 (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) {
a.mu.Lock()
defer a.mu.Unlock()
if a.repushTimers == nil {
a.repushTimers = make(map[string]*time.Timer)
}
if t, ok := a.repushTimers[endpointID]; ok {
t.Stop()
}
a.repushTimers[endpointID] = time.AfterFunc(d, func() {
a.WakePush(endpointID)
})
}
// WakePush 唤醒该端已握手连接的推送 worker(合并唤醒)。
func (a *App) WakePush(endpointID string) {
_, _, ok := a.canPush(endpointID, "")
if !ok {
return
}
a.signalWorker(endpointID)
}