feat: 实现消息分发推送确认撤回与启动恢复
This commit is contained in:
@@ -0,0 +1,288 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strconv"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
// Ack 处理确认(DEVELOPMENT 7.6)。
|
||||
func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (AckResult, error) {
|
||||
if req == nil || req.ID == "" || req.From == "" {
|
||||
return AckResult{}, errCode(protocol.CodeBadRequest, "invalid ack")
|
||||
}
|
||||
nowMs := a.now().UnixMilli()
|
||||
var out AckResult
|
||||
var seq int64
|
||||
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 {
|
||||
return errCode(protocol.CodeNotFound, "message not found")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
res, err := tx.Exec(`
|
||||
UPDATE deliveries SET state = ?, reason = '', pushed_conn = NULL, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
DeliveryAccepted, nowMs, seq, endpointID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
aff, _ := res.RowsAffected()
|
||||
if aff > 0 {
|
||||
out.Result = DeliveryAccepted
|
||||
if e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs); e != nil {
|
||||
return e
|
||||
}
|
||||
return tryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
}
|
||||
var state string
|
||||
err = tx.QueryRow(`
|
||||
SELECT state FROM deliveries WHERE seq = ? AND endpoint_id = ?`, seq, endpointID).Scan(&state)
|
||||
if err == sql.ErrNoRows {
|
||||
return errCode(protocol.CodeNotFound, "delivery not found")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
out.Result = state
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return out, err
|
||||
}
|
||||
if out.Result == DeliveryAccepted {
|
||||
a.releaseLarge(seq, endpointID)
|
||||
}
|
||||
a.WakePush(endpointID)
|
||||
a.WakePush(req.From)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Recall 处理撤回。
|
||||
func (a *App) Recall(ctx context.Context, senderID string, req *protocol.Recall) (protocol.RecallData, error) {
|
||||
if req == nil || req.ID == "" {
|
||||
return protocol.RecallData{}, errCode(protocol.CodeBadRequest, "invalid recall")
|
||||
}
|
||||
nowMs := a.now().UnixMilli()
|
||||
var data protocol.RecallData
|
||||
var revokes []revokeJob
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var seq int64
|
||||
var state string
|
||||
err := tx.QueryRow(`SELECT seq, state FROM messages WHERE sender_id = ? AND id = ?`, senderID, req.ID).Scan(&seq, &state)
|
||||
if err == sql.ErrNoRows {
|
||||
var one int
|
||||
e2 := tx.QueryRow(`SELECT 1 FROM send_keys WHERE sender_id = ? AND msg_id = ?`, senderID, req.ID).Scan(&one)
|
||||
if e2 == nil {
|
||||
data.Result = "failed"
|
||||
return nil
|
||||
}
|
||||
return errCode(protocol.CodeNotFound, "message not found")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if state == StateScheduled {
|
||||
if _, e := tx.Exec(`UPDATE messages SET state = ?, reason = ? WHERE seq = ? AND state = 'scheduled'`,
|
||||
StateCompleted, ReasonRecalled, seq); e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq); e != nil {
|
||||
return e
|
||||
}
|
||||
if a.lim.RecordRetentionDays == 0 {
|
||||
if _, e := tx.Exec(`DELETE FROM messages WHERE seq = ?`, seq); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
data.Result = "recalled"
|
||||
return nil
|
||||
}
|
||||
rows, err := tx.Query(`
|
||||
SELECT endpoint_id, pushed_conn FROM deliveries WHERE seq = ? AND state = 'pending'`, seq)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type pend struct {
|
||||
ep string
|
||||
pushed sql.NullString
|
||||
}
|
||||
var pending []pend
|
||||
for rows.Next() {
|
||||
var p pend
|
||||
if err := rows.Scan(&p.ep, &p.pushed); err != nil {
|
||||
_ = rows.Close()
|
||||
return err
|
||||
}
|
||||
pending = append(pending, p)
|
||||
}
|
||||
_ = rows.Close()
|
||||
|
||||
recalled := 0
|
||||
for _, p := range pending {
|
||||
res, err := tx.Exec(`
|
||||
UPDATE deliveries SET state = ?, reason = ?, pushed_conn = NULL, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
DeliveryRecalled, ReasonRecalled, nowMs, seq, p.ep)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
aff, _ := res.RowsAffected()
|
||||
if aff == 0 {
|
||||
continue
|
||||
}
|
||||
recalled++
|
||||
if p.pushed.Valid && p.pushed.String != "" {
|
||||
revokes = append(revokes, revokeJob{
|
||||
endpointID: p.ep,
|
||||
connID: port.ConnID(p.pushed.String),
|
||||
msgID: req.ID,
|
||||
from: senderID,
|
||||
reason: ReasonRecalled,
|
||||
})
|
||||
}
|
||||
}
|
||||
var accepted, other int
|
||||
_ = tx.QueryRow(`SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state = 'accepted'`, seq).Scan(&accepted)
|
||||
_ = tx.QueryRow(`
|
||||
SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state IN ('expired','dropped','rejected')`, seq).Scan(&other)
|
||||
data.Recalled = recalled
|
||||
data.Accepted = accepted
|
||||
data.Other = other
|
||||
switch {
|
||||
case recalled > 0 && accepted == 0:
|
||||
data.Result = "recalled"
|
||||
case recalled > 0 && accepted > 0:
|
||||
data.Result = "partial"
|
||||
default:
|
||||
data.Result = "failed"
|
||||
}
|
||||
return tryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
})
|
||||
if err != nil {
|
||||
return data, err
|
||||
}
|
||||
a.mu.Lock()
|
||||
a.pendingRevoke = append(a.pendingRevoke, revokes...)
|
||||
a.mu.Unlock()
|
||||
a.flushRevokes(ctx)
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// Status 查询自己发出的消息状态。
|
||||
func (a *App) Status(ctx context.Context, senderID string, req *protocol.Status) (any, error) {
|
||||
if req == nil || req.ID == "" {
|
||||
return nil, errCode(protocol.CodeBadRequest, "invalid status")
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > 200 {
|
||||
limit = 200
|
||||
}
|
||||
var seq int64
|
||||
var state, reason string
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT seq, state, reason FROM messages WHERE sender_id = ? AND id = ?`, senderID, req.ID).Scan(&seq, &state, &reason)
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, errCode(protocol.CodeNotFound, "message not found")
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type counts struct {
|
||||
Pending int `json:"pending"`
|
||||
Accepted int `json:"accepted"`
|
||||
Recalled int `json:"recalled"`
|
||||
Expired int `json:"expired"`
|
||||
Dropped int `json:"dropped"`
|
||||
Rejected int `json:"rejected"`
|
||||
}
|
||||
var c counts
|
||||
rows, err := a.db.Read.QueryContext(ctx, `
|
||||
SELECT state, COUNT(*) FROM deliveries WHERE seq = ? GROUP BY state`, seq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for rows.Next() {
|
||||
var st string
|
||||
var n int
|
||||
if scanErr := rows.Scan(&st, &n); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return nil, scanErr
|
||||
}
|
||||
switch st {
|
||||
case DeliveryPending:
|
||||
c.Pending = n
|
||||
case DeliveryAccepted:
|
||||
c.Accepted = n
|
||||
case DeliveryRecalled:
|
||||
c.Recalled = n
|
||||
case DeliveryExpired:
|
||||
c.Expired = n
|
||||
case DeliveryDropped:
|
||||
c.Dropped = n
|
||||
case DeliveryRejected:
|
||||
c.Rejected = n
|
||||
}
|
||||
}
|
||||
_ = rows.Close()
|
||||
|
||||
q := `SELECT endpoint_id, state, reason FROM deliveries WHERE seq = ?`
|
||||
args := []any{seq}
|
||||
if req.Cursor != "" {
|
||||
q += ` AND endpoint_id > ?`
|
||||
args = append(args, req.Cursor)
|
||||
}
|
||||
q += ` ORDER BY endpoint_id ASC LIMIT ?`
|
||||
args = append(args, limit)
|
||||
drows, err := a.db.Read.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = drows.Close() }()
|
||||
type item struct {
|
||||
EndpointID string `json:"endpoint_id"`
|
||||
State string `json:"state"`
|
||||
Reason string `json:"reason"`
|
||||
}
|
||||
items := make([]item, 0)
|
||||
var nextCursor string
|
||||
for drows.Next() {
|
||||
var it item
|
||||
if err := drows.Scan(&it.EndpointID, &it.State, &it.Reason); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, it)
|
||||
nextCursor = it.EndpointID
|
||||
}
|
||||
return map[string]any{
|
||||
"id": req.ID,
|
||||
"state": state,
|
||||
"reason": reason,
|
||||
"counts": c,
|
||||
"deliveries": items,
|
||||
"next_cursor": nextCursor,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ReceiptAck 确认回执已收下。
|
||||
func (a *App) ReceiptAck(ctx context.Context, endpointID string, req *protocol.ReceiptAck) error {
|
||||
if req == nil || req.ReceiptID == "" {
|
||||
return errCode(protocol.CodeBadRequest, "invalid receipt_ack")
|
||||
}
|
||||
rid, err := strconv.ParseInt(req.ReceiptID, 10, 64)
|
||||
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
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user