552 lines
14 KiB
Go
552 lines
14 KiB
Go
package message
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"encoding/json"
|
|
"fmt"
|
|
"time"
|
|
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
|
)
|
|
|
|
// DispatchDue 分发已到点的 scheduled 消息(按 send_at、seq)。
|
|
func (a *App) DispatchDue(ctx context.Context, nowMs int64, limit int) (int, error) {
|
|
if limit <= 0 {
|
|
limit = 64
|
|
}
|
|
type due struct {
|
|
seq int64
|
|
senderID string
|
|
destKind string
|
|
destID string
|
|
sendAt int64
|
|
keep int
|
|
ttl int64
|
|
receipt int
|
|
}
|
|
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 []due
|
|
for rows.Next() {
|
|
var d due
|
|
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()
|
|
|
|
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, `
|
|
SELECT DISTINCT endpoint_id FROM deliveries WHERE seq = ? AND state = 'pending'`, d.seq)
|
|
if qErr == nil {
|
|
for rows2.Next() {
|
|
var ep string
|
|
if rows2.Scan(&ep) == nil {
|
|
wake[ep] = struct{}{}
|
|
}
|
|
}
|
|
_ = rows2.Close()
|
|
}
|
|
}
|
|
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 {
|
|
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
|
|
}
|
|
|
|
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
|
|
}
|
|
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
|
|
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 {
|
|
_ = rows.Close()
|
|
return err
|
|
}
|
|
_ = body
|
|
it.body = bodyBlob
|
|
items = append(items, it)
|
|
}
|
|
_ = rows.Close()
|
|
if err := rows.Err(); err != nil {
|
|
return err
|
|
}
|
|
|
|
for _, it := range items {
|
|
if it.body == nil {
|
|
// 正文已删则跳过(异常)
|
|
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, err := protocol.Marshal(msg)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
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
|
|
}
|
|
continue
|
|
}
|
|
|
|
claimed := false
|
|
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
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()
|
|
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
|
|
}
|
|
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)
|
|
}
|
|
a.scheduleRepush(endpointID, time.Second)
|
|
}
|
|
}
|
|
return a.pushReceipts(ctx, endpointID, connID, nowMs)
|
|
}
|
|
|
|
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 {
|
|
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
|
|
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 {
|
|
_ = 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
|
|
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)
|
|
if err != nil {
|
|
if err == sql.ErrNoRows {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
if !pushedConn.Valid || pushedConn.String != string(connID) {
|
|
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.releaseLarge(t.seq, endpointID)
|
|
}
|
|
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
|
|
}
|
|
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
|
|
}
|
|
_ = a.down.PublishDown(ctx, j.endpointID, j.connID, payload, port.PublishOpts{QoS: 1})
|
|
}
|
|
}
|
|
|
|
// 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"`
|
|
}
|
|
if err := json.Unmarshal(payload, &head); err != nil || head.Type != protocol.TypeMsg {
|
|
return nil
|
|
}
|
|
nowMs := a.now().UnixMilli()
|
|
err := a.db.Queue.Do(ctx, 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 {
|
|
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))
|
|
a.releaseLarge(seq, endpointID)
|
|
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
|
|
}
|
|
// 简化:未单独记 inflight 回执,按未确认回执取窗口条数
|
|
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)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() { _ = 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
|
|
}
|
|
frame := protocol.Receipt{
|
|
V: protocol.Version, Type: protocol.TypeReceipt,
|
|
ReceiptID: fmt.Sprintf("%d", rid), ID: msgID, EndpointID: epID,
|
|
State: state, Reason: reason, AtMs: created,
|
|
}
|
|
payload, err := protocol.Marshal(frame)
|
|
if err != nil {
|
|
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:
|
|
}
|
|
}
|
|
}
|
|
|
|
func largeKey(seq int64, endpointID string) string {
|
|
return fmt.Sprintf("%d:%s", seq, endpointID)
|
|
}
|
|
|
|
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 唤醒推送;若有登记的连接则异步 PushPending。
|
|
func (a *App) WakePush(endpointID string) {
|
|
live, ok := a.lookupConn(endpointID)
|
|
if !ok || a.down == nil {
|
|
return
|
|
}
|
|
go func() {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
|
defer cancel()
|
|
_ = a.PushPending(ctx, endpointID, live.ConnID)
|
|
a.flushRevokes(ctx)
|
|
}()
|
|
}
|