Files

325 lines
9.1 KiB
Go

package message
import (
"database/sql"
"encoding/base64"
"encoding/json"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
// 投递状态(DEVELOPMENT 7.1)。
const (
DeliveryPending = "pending"
DeliveryAccepted = "accepted"
DeliveryRecalled = "recalled"
DeliveryExpired = "expired"
DeliveryDropped = "dropped"
DeliveryRejected = "rejected"
)
// 常见 reason。
const (
ReasonEndpointDisabled = "endpoint_disabled"
ReasonQueueFull = "queue_full"
ReasonOffline = "offline"
ReasonNotAcked = "not_acked"
ReasonTTL = "ttl"
ReasonTooLarge = "too_large"
ReasonSenderLeft = "sender_left"
ReasonNoRecipients = "no_recipients"
ReasonRecalled = "recalled"
ReasonGroupDissolved = "group_dissolved"
)
const (
packetOverheadBudget = 128
largeFrameBytes = 64 * 1024
maxLargeInflight = 64
defaultDeliveryWindow = 32
defaultReceiptWindow = 64
)
// dispatchFullTx 按 DEVELOPMENT 7.4 完整分发一条已到点的 scheduled 消息。
// claimed=false 表示别人已先改状态,state 为当前状态。
func (a *App) dispatchFullTx(tx *sql.Tx, seq int64, senderID, destKind, destID string, sendAt int64, keep int, ttlSeconds int64, wantReceipt bool, nowMs int64) (state string, claimed bool, err error) {
graceMs := a.lim.GraceSeconds * 1000
res, err := tx.Exec(`UPDATE messages SET state = ? WHERE seq = ? AND state = ?`, StateDispatched, seq, StateScheduled)
if err != nil {
return "", false, err
}
aff, _ := res.RowsAffected()
if aff == 0 {
var st string
if err := tx.QueryRow(`SELECT state FROM messages WHERE seq = ?`, seq).Scan(&st); err != nil {
return "", false, err
}
return st, false, nil
}
type recip struct {
id string
enabled int
}
var recipients []recip
var msgReason string
completeEarly := false
switch destKind {
case protocol.TargetEndpoint:
ep, err := loadEndpointTx(tx, destID)
if err != nil {
if err == sql.ErrNoRows {
recipients = []recip{{id: destID, enabled: 0}}
} else {
return "", true, err
}
} else {
recipients = []recip{{id: ep.ID, enabled: ep.Enabled}}
}
case protocol.TargetGroup:
var one int
err := tx.QueryRow(`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, destID, senderID).Scan(&one)
if err == sql.ErrNoRows {
completeEarly = true
msgReason = ReasonSenderLeft
} else if err != nil {
return "", true, err
} else {
rows, qErr := tx.Query(`
SELECT gm.endpoint_id, e.enabled
FROM group_members gm
JOIN endpoints e ON e.id = gm.endpoint_id
WHERE gm.group_id = ? AND gm.endpoint_id != ?`, destID, senderID)
if qErr != nil {
return "", true, qErr
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var r recip
if sErr := rows.Scan(&r.id, &r.enabled); sErr != nil {
return "", true, sErr
}
recipients = append(recipients, r)
}
if err := rows.Err(); err != nil {
return "", true, err
}
if len(recipients) == 0 {
completeEarly = true
msgReason = ReasonNoRecipients
}
}
default:
return "", true, errCode(protocol.CodeBadRequest, "invalid dest_kind")
}
if completeEarly {
if err := finalizeMessageTx(tx, seq, wantReceipt, senderID, "", msgReason, nowMs, a.lim.RecordRetentionDays); err != nil {
return "", true, err
}
return StateCompleted, true, nil
}
pendingAny := false
for _, r := range recipients {
dState := DeliveryPending
reason := ""
var expireAt sql.NullInt64
if r.enabled == 0 {
dState = DeliveryRejected
reason = ReasonEndpointDisabled
} else if a.lim.MaxPendingPerReceiver > 0 {
var n int
if err := tx.QueryRow(`
SELECT COUNT(*) FROM deliveries
WHERE endpoint_id = ? AND state = 'pending'`, r.id).Scan(&n); err != nil {
return "", true, err
}
if n >= a.lim.MaxPendingPerReceiver {
dState = DeliveryRejected
reason = ReasonQueueFull
}
}
if dState == DeliveryPending {
_, online := a.lookupConn(r.id)
keepBool := keep != 0
switch {
case online && keepBool:
expireAt = sql.NullInt64{Int64: nowMs + ttlSeconds*1000, Valid: true}
case online && !keepBool:
// expire_at 空
case !online && keepBool:
expireAt = sql.NullInt64{Int64: nowMs + ttlSeconds*1000, Valid: true}
default:
var offlineSince sql.NullInt64
_ = tx.QueryRow(`SELECT offline_since FROM endpoints WHERE id = ?`, r.id).Scan(&offlineSince)
if !offlineSince.Valid {
dState = DeliveryDropped
reason = ReasonOffline
} else if nowMs-offlineSince.Int64 > graceMs {
dState = DeliveryDropped
reason = ReasonOffline
} else {
expireAt = sql.NullInt64{Int64: offlineSince.Int64 + graceMs, Valid: true}
}
}
}
keepVal := keep
if _, err := tx.Exec(`
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, expire_at, pushed_conn, pushed_at, attempts, updated_at)
VALUES(?,?,?,?,?,?,?,NULL,NULL,0,?)`,
seq, r.id, sendAt, keepVal, dState, reason, nullInt(expireAt), nowMs,
); err != nil {
return "", true, err
}
if dState == DeliveryPending {
pendingAny = true
} else if wantReceipt {
if err := insertReceiptTx(tx, senderID, seq, r.id, dState, reason, nowMs); err != nil {
return "", true, err
}
}
}
if pendingAny {
if _, err := tx.Exec(`UPDATE messages SET state = ? WHERE seq = ?`, StateDispatched, seq); err != nil {
return "", true, err
}
return StateDispatched, true, nil
}
if err := finalizeMessageTx(tx, seq, wantReceipt, senderID, "", "", nowMs, a.lim.RecordRetentionDays); err != nil {
return "", true, err
}
return StateCompleted, true, nil
}
func nullInt(v sql.NullInt64) any {
if !v.Valid {
return nil
}
return v.Int64
}
func (a *App) lookupConn(endpointID string) (LiveConn, bool) {
if a.conns == nil {
return LiveConn{}, false
}
return a.conns.Current(endpointID)
}
// finalizeMessageTx 无 pending 时收尾:completed、删正文;记录天数 0 则删消息与投递。
// msgReason 非空时写入消息 reason(发送前结束);endpointID 为空表示消息级回执。
func finalizeMessageTx(tx *sql.Tx, seq int64, wantReceipt bool, senderID, endpointID, msgReason string, nowMs int64, recordDays int) error {
var msgID string
var receipt int
if err := tx.QueryRow(`SELECT id, receipt FROM messages WHERE seq = ?`, seq).Scan(&msgID, &receipt); err != nil {
return err
}
if msgReason != "" && wantReceipt && receipt != 0 {
if err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, msgReason, nowMs); err != nil {
return err
}
}
if _, err := tx.Exec(`UPDATE messages SET state = ?, reason = CASE WHEN ? != '' THEN ? ELSE reason END WHERE seq = ?`,
StateCompleted, msgReason, msgReason, seq); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq); err != nil {
return err
}
if recordDays == 0 {
if _, err := tx.Exec(`DELETE FROM deliveries WHERE seq = ?`, seq); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM messages WHERE seq = ?`, seq); err != nil {
return err
}
// 防重行保留(消息记录已删,小时清理按条件删)
}
return nil
}
// tryFinalizeTx 若无 pending 则收尾。
func tryFinalizeTx(tx *sql.Tx, seq int64, nowMs int64, recordDays int) error {
var n int
if err := tx.QueryRow(`SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state = 'pending'`, seq).Scan(&n); err != nil {
return err
}
if n > 0 {
return nil
}
var senderID string
var receipt int
if err := tx.QueryRow(`SELECT sender_id, receipt FROM messages WHERE seq = ?`, seq).Scan(&senderID, &receipt); err != nil {
if err == sql.ErrNoRows {
return nil
}
return err
}
return finalizeMessageTx(tx, seq, receipt != 0, senderID, "", "", nowMs, recordDays)
}
func insertReceiptTx(tx *sql.Tx, senderID string, seq int64, endpointID, state, reason string, nowMs int64) error {
var msgID string
var want int
if err := tx.QueryRow(`SELECT id, receipt FROM messages WHERE seq = ?`, seq).Scan(&msgID, &want); err != nil {
return err
}
if want == 0 {
return nil
}
// 发送方仍存在
var one int
err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, senderID).Scan(&one)
if err == sql.ErrNoRows {
return nil
}
if err != nil {
return err
}
_, err = tx.Exec(`
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
VALUES(?,?,?,?,?,?,0)`, senderID, msgID, endpointID, state, reason, nowMs)
return err
}
func encodeStoredBody(enc, contentType string, raw []byte) protocol.Body {
if enc == protocol.EncBase64 {
return protocol.Body{Enc: enc, ContentType: contentType, Data: base64.StdEncoding.EncodeToString(raw)}
}
return protocol.Body{Enc: protocol.EncUTF8, ContentType: contentType, Data: string(raw)}
}
func decodeMetaJSON(s string) map[string]any {
if s == "" || s == "{}" {
return nil
}
var m map[string]any
if err := json.Unmarshal([]byte(s), &m); err != nil {
return nil
}
return m
}
func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
limit := 0
if maxPacketSize > 0 {
if maxPacketSize > packetOverheadBudget {
limit = int(maxPacketSize) - packetOverheadBudget
} else {
limit = 0
}
}
if maxRecvBytes > 0 {
if limit == 0 || maxRecvBytes < limit {
limit = maxRecvBytes
}
}
return limit
}