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++ 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 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 } items = append(items, it) } _ = rows.Close() if err := rows.Err(); err != nil { return err } var toClaim []pushItem 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 } 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) } } 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 } } 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) }