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) } else { a.observeDispatchToPush(it.sendAt, nowMs) } } return a.pushReceipts(ctx, 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 { 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) }() }