feat: 实现消息分发推送确认撤回与启动恢复

This commit is contained in:
Nixevol
2026-09-30 07:32:53 +08:00
parent 34e7c2827f
commit bdd1d9e9f4
12 changed files with 2395 additions and 119 deletions
+55 -12
View File
@@ -317,33 +317,76 @@
### M1 2026-09-30 ### M1 2026-09-30
1. **提交时分发做成最小正确版** 1. **提交时分发做成最小正确版**(已被 M2 取代)
- 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发(停用拒绝、`queue_full`、`expire_at`/宽限、无接收者 `completed`、回执等)。 - 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发。
- 实际做法:单聊只插一条 `pending`;群按当时 `group_members` 去掉发送者各插 `pending`;消息改为 `dispatched`。不设 `expire_at`,不检查接收端配额/在线/停用,不因无接收者改为 `completed`,不写回执,不唤醒推送循环。 - 实际做法(M1):单聊/群插 `pending` 后改 `dispatched`,不做停用/配额/宽限。
- 原因:M1 范围是提交;完整分发与推送属 M2。 - 现状(M2):`Submit` 到点与 `DispatchDue` 均走完整 `dispatchFullTx`(7.4)。
- 备选方案:M1 直接实现完整 7.4(抢 M2)。 - 原因 / 备选 / 影响:见 M2。
- 影响:到点消息已有投递行,但停用成员仍会有 `pending`;无成员群仍为 `dispatched` 且无投递;推送需等 M2。
2. **请求频率突发容量写死为 100** 2. **请求频率突发容量写死为 100**
- 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。 - 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。
- 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 尚未实现故未接桶。 - 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 不计入桶(与 6.10 一致)。
- 原因:配置无独立 burst 字段。 - 原因:配置无独立 burst 字段。
- 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。 - 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。
- 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。 - 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。
3. **未接线 `cmd/nixmsg`** 3. **未接线 `cmd/nixmsg`**
- 原条款:可替换 T0.4 假实现。 - 原条款:可替换 T0.4 假实现。
- 实际做法:新增 `message.App` 实现 `Submit`;保留 `Stub`;按任务隔离要求未改 `cmd/nixmsg`/`wire.go`。 - 实际做法:`message.App` 实现 Service;保留 `Stub`;未改 `cmd/nixmsg`/`wire.go`。
- 原因:本任务禁止改 `cmd/nixmsg`;总控接线或后续任务再换。 - 原因:本任务隔离;总控接线。
- 备选方案:本任务直接改 `wire.go`。 - 备选方案:本任务直接改 `wire.go`。
- 影响:进程内仍用 Stub,需显式构造 `message.New` 才能用真实提交。 - 影响:进程内仍用 Stub,需显式 `message.New` 并注入 `Downlink`/`ConnRegistry`。
4. **防重键在、消息行已删时返回 `not_found`** 4. **防重键在、消息行已删时返回 `not_found`**
- 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为(状态查询为 `not_found`)。 - 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为。
- 实际做法:`send_keys` 指纹相同但 `messages` 无行时返回 `not_found`。 - 实际做法:`send_keys` 指纹相同但 `messages` 无行时返回 `not_found`。
- 原因:无法构造 `send_at`/`state`。 - 原因:无法构造 `send_at`/`state`。
- 备选方案:在 `send_keys` 冗余存结果快照。 - 备选方案:在 `send_keys` 冗余存结果快照。
- 影响:保留期过后的重试不再幂等成功。 - 影响:保留天数 0 完成后重试不再幂等成功(与 F18 防重「记录还在时」一致)。
### M2 / M3 / M4 2026-09-30
1. **下行与在线用可注入接口,测试用假实现**
- 原条款:推送经 broker `Downlink`;连接表在 N 线内存。
- 实际做法:`WithDownlink` / `WithConnRegistry`;测试用 `RecordingDownlink`、`MemoryConns`。状态机全在 `message` 包。未接真实 MQTT/mochi。
- 原因:N3 握手与 wire 本波未强制合入;任务允许假下行。
- 备选方案:直接依赖 `internal/broker.Broker`。
- 影响:接线方需在握手/断线时调用 `OnHandshakeComplete`/`OnDisconnect`,登记连接,并把 `OnPublishDropped` 转到 `App`。
2. **大帧并发名额在 message 包再管一份**
- 原条款:大于 64KiB 全局同时不超过 64(DEVELOPMENT 7.5);N2 broker 已有信号量。
- 实际做法:`App` 内另有容量 64 的 `largeSem`,发布前申请,确认/超时/清标记时释放。
- 原因:假 `Downlink` 不经 broker 时仍要满足上限。
- 备选方案:只依赖 broker,测试也走真实 PublishDown。
- 影响:接线真实 broker 后可能双重限流(更严,不破坏语义)。
3. **确认超时按库内 `pushed_at` 判定,不另开每连接计时器 goroutine**
- 原条款:推送循环在内存里计时。
- 实际做法:`PushPending` 开头扫描该连接已推且 `now - pushed_at >= ack_timeout` 的投递,再按 keep/expire 规则处理。
- 原因:与崩溃恢复一致、测试可拨钟;避免无调度器时泄漏计时器。
- 备选方案:每连接 `time.AfterFunc`。
- 影响:需周期性调用 `PushPending`(或 `WakePush`)才会触发超时。
4. **后台调度/清理循环未在 App 内自启**
- 原条款:调度按 `send_at` 唤醒;清理约每秒;推送每连接一循环。
- 实际做法:导出 `DispatchDue`、`PushPending`、`CleanupOnce`、`RecoverOnStart`、`WakePush`;由接线方起 goroutine。`WakePush` 在有连接时异步 `PushPending`。
- 原因:未改 `cmd/nixmsg`;避免无 context 的后台泄漏。
- 备选方案:`App.Start(ctx)` 内启三循环。
- 影响:未接线则定时消息不会自动到点,需外部调用 `DispatchDue`。
5. **回执推送窗口未单独记 inflight**
- 原条款:回执窗口默认 64,确认一笔再推下一笔。
- 实际做法:按 `acked=0` 取最多 `ReceiptWindow` 条尽力发布;不因未 `receipt_ack` 停推后续。
- 原因:简化;回执可重复、SDK 按 `receipt_id` 去重。
- 备选方案:内存记已推未确认回执数。
- 影响:发送方慢确认时可能多推几条回执(协议允许重复)。
6. **`Status` 返回自建 map,非独立协议类型**
- 原条款:6.4 状态响应字段。
- 实际做法:`map[string]any`(`state`/`reason`/`counts`/`deliveries`/`next_cursor`)。
- 原因:`protocol` 无 StatusData 结构且不可改共享协议包时取稳妥形状。
- 备选方案:总控在 `protocol` 增类型。
- 影响:接线编码 `resp.data` 时直接 Marshal 该 map 即可。
## 身份 I ## 身份 I
+288
View File
@@ -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
})
}
+56 -45
View File
@@ -1,7 +1,7 @@
package message package message
import ( import (
"context" "sync"
"time" "time"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
@@ -18,11 +18,6 @@ const (
StateCompleted = "completed" StateCompleted = "completed"
) )
// 投递状态。
const (
DeliveryPending = "pending"
)
// talk_grants.kind。 // talk_grants.kind。
const ( const (
GrantKindPassword = "password" GrantKindPassword = "password"
@@ -32,7 +27,7 @@ const (
// 请求频率桶默认突发容量(DEVELOPMENT 6.10;配置无单独字段)。 // 请求频率桶默认突发容量(DEVELOPMENT 6.10;配置无单独字段)。
const defaultRequestBurst = 100 const defaultRequestBurst = 100
// Limits 是提交所需的配置上限(来自 config.LimitsConfig)。 // Limits 是消息子系统所需配置上限。
type Limits struct { type Limits struct {
MaxBodyBytes int MaxBodyBytes int
MaxMetaBytes int MaxMetaBytes int
@@ -44,11 +39,16 @@ type Limits struct {
MaxPendingPerSender int MaxPendingPerSender int
MaxPendingPerReceiver int MaxPendingPerReceiver int
GraceSeconds int64 GraceSeconds int64
AckTimeoutSeconds int64
DeliveryWindow int
ReceiptWindow int
RecordRetentionDays int
ReceiptRetentionDays int
IdempotencyHours int
} }
// LimitsFromConfig 从平台配置构造 Limits。 // LimitsFromConfig 从平台配置构造 Limits。
func LimitsFromConfig(c config.LimitsConfig) Limits { func LimitsFromConfig(c config.LimitsConfig) Limits {
burst := defaultRequestBurst
return Limits{ return Limits{
MaxBodyBytes: c.MaxBodyBytes, MaxBodyBytes: c.MaxBodyBytes,
MaxMetaBytes: c.MaxMetaBytes, MaxMetaBytes: c.MaxMetaBytes,
@@ -56,14 +56,26 @@ func LimitsFromConfig(c config.LimitsConfig) Limits {
MaxTTLSeconds: int64(c.MaxTTLSeconds), MaxTTLSeconds: int64(c.MaxTTLSeconds),
MaxScheduleSeconds: int64(c.MaxScheduleSeconds), MaxScheduleSeconds: int64(c.MaxScheduleSeconds),
RequestsPerSecond: float64(c.RequestsPerSecond), RequestsPerSecond: float64(c.RequestsPerSecond),
RequestBurst: burst, RequestBurst: defaultRequestBurst,
MaxPendingPerSender: c.MaxPendingPerSender, MaxPendingPerSender: c.MaxPendingPerSender,
MaxPendingPerReceiver: c.MaxPendingPerReceiver, MaxPendingPerReceiver: c.MaxPendingPerReceiver,
GraceSeconds: int64(c.GraceSeconds), GraceSeconds: int64(c.GraceSeconds),
AckTimeoutSeconds: int64(c.AckTimeoutSeconds),
DeliveryWindow: c.DeliveryWindow,
ReceiptWindow: c.ReceiptWindow,
} }
} }
// App 实现 Service 的提交路径(M1);其余方法暂返回未实现或空操作。 // LimitsFromFullConfig 附带保留天数等顶层配置。
func LimitsFromFullConfig(cfg config.Config) Limits {
lim := LimitsFromConfig(cfg.Limits)
lim.RecordRetentionDays = cfg.RecordRetentionDays
lim.ReceiptRetentionDays = cfg.ReceiptRetentionDays
lim.IdempotencyHours = cfg.IdempotencyHours
return lim
}
// App 实现 Service:提交、分发、推送、确认、撤回、回执、清理与启动恢复。
type App struct { type App struct {
db *store.DB db *store.DB
lim Limits lim Limits
@@ -71,6 +83,15 @@ type App struct {
locks auth.LoginLocks locks auth.LoginLocks
nowFn func() time.Time nowFn func() time.Time
rates *rateLimiter rates *rateLimiter
down port.Downlink
conns ConnRegistry
mu sync.Mutex
largeSem chan struct{}
largeHeld map[string]bool
pendingRevoke []revokeJob
repushTimers map[string]*time.Timer
} }
// Option 配置 App。 // Option 配置 App。
@@ -86,17 +107,41 @@ func WithLocks(locks auth.LoginLocks) Option {
return func(a *App) { a.locks = locks } return func(a *App) { a.locks = locks }
} }
// New 创建消息服务实现。hash 用于校验对话密码;locks 可为 nil。 // WithDownlink 注入下行发布器(未接线时测试用 RecordingDownlink)。
func WithDownlink(d port.Downlink) Option {
return func(a *App) { a.down = d }
}
// WithConnRegistry 注入连接查询。
func WithConnRegistry(c ConnRegistry) Option {
return func(a *App) { a.conns = c }
}
// New 创建消息服务实现。
func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App { func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App {
if lim.RequestBurst <= 0 { if lim.RequestBurst <= 0 {
lim.RequestBurst = defaultRequestBurst lim.RequestBurst = defaultRequestBurst
} }
if lim.DeliveryWindow <= 0 {
lim.DeliveryWindow = defaultDeliveryWindow
}
if lim.ReceiptWindow <= 0 {
lim.ReceiptWindow = defaultReceiptWindow
}
if lim.AckTimeoutSeconds <= 0 {
lim.AckTimeoutSeconds = 300
}
if lim.GraceSeconds < 0 {
lim.GraceSeconds = 60
}
a := &App{ a := &App{
db: db, db: db,
lim: lim, lim: lim,
hash: hash, hash: hash,
nowFn: time.Now, nowFn: time.Now,
rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst), rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst),
largeSem: make(chan struct{}, maxLargeInflight),
largeHeld: make(map[string]bool),
} }
for _, opt := range opts { for _, opt := range opts {
opt(a) opt(a)
@@ -116,38 +161,4 @@ func (a *App) protocolLimits() protocol.Limits {
} }
} }
func (a *App) Ack(context.Context, string, *protocol.Ack) (AckResult, error) {
return AckResult{}, ErrNotImplemented
}
func (a *App) Recall(context.Context, string, *protocol.Recall) (protocol.RecallData, error) {
return protocol.RecallData{}, ErrNotImplemented
}
func (a *App) Status(context.Context, string, *protocol.Status) (any, error) {
return nil, ErrNotImplemented
}
func (a *App) ReceiptAck(context.Context, string, *protocol.ReceiptAck) error {
return ErrNotImplemented
}
func (a *App) DispatchDue(context.Context, int64, int) (int, error) {
return 0, nil
}
func (a *App) PushPending(context.Context, string, port.ConnID) error {
return nil
}
func (a *App) OnPublishDropped(context.Context, string, port.ConnID, []byte) error {
return nil
}
func (a *App) CleanupOnce(context.Context, int64) error { return nil }
func (a *App) RecoverOnStart(context.Context) error { return nil }
func (a *App) WakePush(string) {}
var _ Service = (*App)(nil) var _ Service = (*App)(nil)
+144
View File
@@ -0,0 +1,144 @@
package message
import (
"context"
"encoding/json"
"errors"
"sync"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
)
var (
errPayloadTooLarge = errors.New("message: payload too large")
errPublishFailed = errors.New("message: publish failed")
)
// LiveConn 是端当前连接的快照(含握手中)。
type LiveConn struct {
ConnID port.ConnID
MaxReceiveBytes int
MaxPacketSize uint32
}
// ConnRegistry 查询端是否有连接(由 N 线或测试假实现注入)。
// 有连接(含握手中)即视为在线,用于分发时计算 expire_at。
type ConnRegistry interface {
Current(endpointID string) (LiveConn, bool)
}
// MemoryConns 是测试用的内存连接表。
type MemoryConns struct {
mu sync.RWMutex
m map[string]LiveConn
}
// NewMemoryConns 创建空连接表。
func NewMemoryConns() *MemoryConns {
return &MemoryConns{m: make(map[string]LiveConn)}
}
// Set 登记或更新端的当前连接。
func (c *MemoryConns) Set(endpointID string, conn LiveConn) {
c.mu.Lock()
defer c.mu.Unlock()
c.m[endpointID] = conn
}
// Clear 移除端的当前连接;若代号不匹配则不动。
func (c *MemoryConns) Clear(endpointID string, connID port.ConnID) {
c.mu.Lock()
defer c.mu.Unlock()
cur, ok := c.m[endpointID]
if !ok {
return
}
if connID != "" && cur.ConnID != connID {
return
}
delete(c.m, endpointID)
}
// Current 实现 ConnRegistry。
func (c *MemoryConns) Current(endpointID string) (LiveConn, bool) {
c.mu.RLock()
defer c.mu.RUnlock()
v, ok := c.m[endpointID]
return v, ok
}
// RecordingDownlink 记录下行发布,供测试断言。
type RecordingDownlink struct {
mu sync.Mutex
Published []DownPublish
FailNext int // 接下来 N 次 PublishDown 返回错误
MaxSize int // >0 时超限返回错误
}
// DownPublish 是一次下行记录。
type DownPublish struct {
EndpointID string
ConnID port.ConnID
Payload []byte
QoS byte
}
// PublishDown 实现 port.Downlink。
func (d *RecordingDownlink) PublishDown(_ context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error {
d.mu.Lock()
defer d.mu.Unlock()
if d.MaxSize > 0 && len(payload) > d.MaxSize {
return errPayloadTooLarge
}
if d.FailNext > 0 {
d.FailNext--
return errPublishFailed
}
d.Published = append(d.Published, DownPublish{
EndpointID: endpointID,
ConnID: connID,
Payload: append([]byte(nil), payload...),
QoS: opts.QoS,
})
return nil
}
// Count 返回已发布条数。
func (d *RecordingDownlink) Count() int {
d.mu.Lock()
defer d.mu.Unlock()
return len(d.Published)
}
// Snapshots 返回发布副本。
func (d *RecordingDownlink) Snapshots() []DownPublish {
d.mu.Lock()
defer d.mu.Unlock()
out := make([]DownPublish, len(d.Published))
copy(out, d.Published)
return out
}
// FilterType 统计 type 字段匹配的发布次数。
func (d *RecordingDownlink) FilterType(typ string) int {
d.mu.Lock()
defer d.mu.Unlock()
n := 0
for _, p := range d.Published {
if payloadType(p.Payload) == typ {
n++
}
}
return n
}
func payloadType(payload []byte) string {
var head struct {
Type string `json:"type"`
}
_ = json.Unmarshal(payload, &head)
return head.Type
}
var _ port.Downlink = (*RecordingDownlink)(nil)
var _ ConnRegistry = (*MemoryConns)(nil)
+744
View File
@@ -0,0 +1,744 @@
package message
import (
"context"
"database/sql"
"encoding/json"
"path/filepath"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
type deliveryEnv struct {
t *testing.T
app *App
db *store.DB
conns *MemoryConns
down *RecordingDownlink
nowMs int64
}
func openDeliveryEnv(t *testing.T, mutate func(*Limits)) *deliveryEnv {
t.Helper()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
lim := defaultTestLimits()
lim.GraceSeconds = 60
lim.AckTimeoutSeconds = 300
lim.DeliveryWindow = 32
if mutate != nil {
mutate(&lim)
}
nowMs := int64(1_700_000_000_000)
conns := NewMemoryConns()
down := &RecordingDownlink{}
app := New(db, lim, auth.NewStubHashPool(),
WithNow(func() time.Time { return time.UnixMilli(nowMs) }),
WithLocks(auth.NewStubLoginLocks()),
WithConnRegistry(conns),
WithDownlink(down),
)
return &deliveryEnv{t: t, app: app, db: db, conns: conns, down: down, nowMs: nowMs}
}
func (e *deliveryEnv) setNow(ms int64) {
e.nowMs = ms
e.app.nowFn = func() time.Time { return time.UnixMilli(e.nowMs) }
}
func (e *deliveryEnv) online(id string, connID port.ConnID) {
e.conns.Set(id, LiveConn{ConnID: connID, MaxReceiveBytes: 0, MaxPacketSize: 0})
}
func (e *deliveryEnv) deliveryState(seq int64, endpointID string) (state, reason string) {
e.t.Helper()
err := e.db.Read.QueryRow(`SELECT state, reason FROM deliveries WHERE seq=? AND endpoint_id=?`, seq, endpointID).Scan(&state, &reason)
if err != nil {
e.t.Fatal(err)
}
return
}
func (e *deliveryEnv) msgState(sender, id string) (state, reason string) {
e.t.Helper()
err := e.db.Read.QueryRow(`SELECT state, reason FROM messages WHERE sender_id=? AND id=?`, sender, id).Scan(&state, &reason)
if err != nil {
e.t.Fatal(err)
}
return
}
func (e *deliveryEnv) seqOf(sender, id string) int64 {
e.t.Helper()
var seq int64
if err := e.db.Read.QueryRow(`SELECT seq FROM messages WHERE sender_id=? AND id=?`, sender, id).Scan(&seq); err != nil {
e.t.Fatal(err)
}
return seq
}
func keepTrue() *protocol.OfflineOpts {
ttl := int64(3600)
return &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}
}
func TestDeliveryStateMachine(t *testing.T) {
t.Parallel()
t.Run("F10_grace_within_keeps_pending", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
// offline_since = now → 宽限内
ctx := context.Background()
res, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("g1", "bob"))
if err != nil {
t.Fatal(err)
}
if res.State != StateDispatched {
t.Fatalf("state=%s", res.State)
}
seq := e.seqOf("alice", "g1")
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryPending || reason != "" {
t.Fatalf("got %s/%s", st, reason)
}
var exp sql.NullInt64
_ = e.db.Read.QueryRow(`SELECT expire_at FROM deliveries WHERE seq=?`, seq).Scan(&exp)
if !exp.Valid || exp.Int64 != e.nowMs+60_000 {
t.Fatalf("expire_at=%v want %d", exp, e.nowMs+60_000)
}
})
t.Run("F10_grace_exceeded_dropped_offline", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
_ = e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id=?`, e.nowMs-120_000, "bob")
return err
})
ctx := context.Background()
res, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("g2", "bob"))
if err != nil {
t.Fatal(err)
}
if res.State != StateCompleted {
t.Fatalf("state=%s", res.State)
}
seq := e.seqOf("alice", "g2")
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryDropped || reason != ReasonOffline {
t.Fatalf("got %s/%s", st, reason)
}
})
t.Run("F09_keep_ttl_expire_via_cleanup", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) { l.GraceSeconds = 60 })
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
req := baseSend("k1", "bob")
req.Offline = keepTrue()
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "k1")
st, _ := e.deliveryState(seq, "bob")
if st != DeliveryPending {
t.Fatalf("state=%s", st)
}
e.setNow(e.nowMs + 3600_000 + 1)
if err := e.app.CleanupOnce(ctx, e.nowMs); err != nil {
t.Fatal(err)
}
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryExpired || reason != ReasonTTL {
t.Fatalf("got %s/%s", st, reason)
}
})
t.Run("F08_ack_timeout_not_keep_dropped", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) { l.AckTimeoutSeconds = 10 })
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a1", "bob")); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
if e.down.FilterType(protocol.TypeMsg) < 1 {
t.Fatal("expected msg publish")
}
seq := e.seqOf("alice", "a1")
var pushed sql.NullString
_ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed)
if !pushed.Valid {
t.Fatal("expected pushed")
}
e.setNow(e.nowMs + 11_000)
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryDropped || reason != ReasonNotAcked {
t.Fatalf("got %s/%s", st, reason)
}
})
t.Run("F08_ack_timeout_keep_repush_then_expire", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) {
l.AckTimeoutSeconds = 10
})
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
req := baseSend("a2", "bob")
ttl := int64(30)
req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "a2")
// 未过 expire:清标记重推
e.setNow(e.nowMs + 11_000)
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
st, _ := e.deliveryState(seq, "bob")
if st != DeliveryPending {
t.Fatalf("want pending got %s", st)
}
var pushed sql.NullString
_ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed)
// 可能已重推或仍清空后待推
// 过 expire
e.setNow(e.nowMs + 30_000)
_ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE deliveries SET pushed_conn=?, pushed_at=? WHERE seq=?`, "c-bob", e.nowMs-11_000, seq)
return err
})
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryExpired || reason != ReasonTTL {
t.Fatalf("got %s/%s", st, reason)
}
})
t.Run("F12_scheduled_recall", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
delay := int64(10_000)
req := baseSend("r1", "bob")
req.DelayMs = &delay
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
data, err := e.app.Recall(ctx, "alice", &protocol.Recall{V: 1, Type: protocol.TypeRecall, RID: "1", ID: "r1"})
if err != nil {
t.Fatal(err)
}
if data.Result != "recalled" || data.Recalled != 0 {
t.Fatalf("%+v", data)
}
st, reason := e.msgState("alice", "r1")
if st != StateCompleted || reason != ReasonRecalled {
t.Fatalf("%s/%s", st, reason)
}
})
t.Run("F13_recall_race_partial", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
insertEndpoint(t, e.db, "carol", "", 1, 0)
ctx := context.Background()
_ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES('g1','g','alice',?)`, e.nowMs)
if err != nil {
return err
}
for _, m := range []string{"alice", "bob", "carol"} {
if _, err := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES('g1',?,?)`, m, e.nowMs); err != nil {
return err
}
}
return nil
})
req := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "grp1",
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
}
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "grp1")
// bob 先确认
if _, err := e.app.Ack(ctx, "bob", &protocol.Ack{V: 1, Type: protocol.TypeAck, RID: "a", From: "alice", ID: "grp1"}); err != nil {
t.Fatal(err)
}
data, err := e.app.Recall(ctx, "alice", &protocol.Recall{V: 1, Type: protocol.TypeRecall, RID: "2", ID: "grp1"})
if err != nil {
t.Fatal(err)
}
if data.Result != "partial" || data.Accepted != 1 || data.Recalled != 1 {
t.Fatalf("%+v", data)
}
st, _ := e.deliveryState(seq, "carol")
if st != DeliveryRecalled {
t.Fatalf("carol=%s", st)
}
st, _ = e.deliveryState(seq, "bob")
if st != DeliveryAccepted {
t.Fatalf("bob=%s", st)
}
})
t.Run("F13_recall_pending_unpushed", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
ctx := context.Background()
req := baseSend("r2", "bob")
req.Offline = keepTrue()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
data, err := e.app.Recall(ctx, "alice", &protocol.Recall{V: 1, Type: protocol.TypeRecall, RID: "1", ID: "r2"})
if err != nil {
t.Fatal(err)
}
if data.Result != "recalled" || data.Recalled != 1 {
t.Fatalf("%+v", data)
}
})
t.Run("F14_receipt_when_sender_offline", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("rc1", "bob")); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
if _, err := e.app.Ack(ctx, "bob", &protocol.Ack{V: 1, Type: protocol.TypeAck, RID: "a", From: "alice", ID: "rc1"}); err != nil {
t.Fatal(err)
}
var n int
if err := e.db.Read.QueryRow(`SELECT COUNT(*) FROM receipts WHERE sender_id=? AND msg_id=? AND state=?`, "alice", "rc1", DeliveryAccepted).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("receipts=%d", n)
}
// alice 上线后能推到回执
e.online("alice", "c-alice")
if err := e.app.PushPending(ctx, "alice", "c-alice"); err != nil {
t.Fatal(err)
}
if e.down.FilterType(protocol.TypeReceipt) < 1 {
t.Fatal("expected receipt push")
}
})
t.Run("F18_body_gone_record_zero_idempotent_kept", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) { l.RecordRetentionDays = 0 })
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
req := baseSend("f18", "bob")
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
if _, err := e.app.Ack(ctx, "bob", &protocol.Ack{V: 1, Type: protocol.TypeAck, RID: "a", From: "alice", ID: "f18"}); err != nil {
t.Fatal(err)
}
var bodies int
_ = e.db.Read.QueryRow(`SELECT COUNT(*) FROM message_bodies`).Scan(&bodies)
if bodies != 0 {
t.Fatalf("bodies=%d", bodies)
}
var msgs int
_ = e.db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE id=?`, "f18").Scan(&msgs)
if msgs != 0 {
t.Fatalf("messages=%d want 0", msgs)
}
var keys int
_ = e.db.Read.QueryRow(`SELECT COUNT(*) FROM send_keys WHERE msg_id=?`, "f18").Scan(&keys)
if keys != 1 {
t.Fatalf("send_keys=%d", keys)
}
// 防重:消息已删 → not_found
_, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req)
if protoCode(err) != protocol.CodeNotFound {
t.Fatalf("want not_found got %v", err)
}
})
t.Run("F11_dispatch_due_after_downtime", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
delay := int64(60_000)
req := baseSend("due1", "bob")
req.DelayMs = &delay
ctx := context.Background()
res, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if res.State != StateScheduled {
t.Fatalf("%s", res.State)
}
e.setNow(e.nowMs + 60_000)
n, err := e.app.DispatchDue(ctx, e.nowMs, 10)
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("dispatched=%d", n)
}
st, _ := e.msgState("alice", "due1")
if st != StateDispatched {
t.Fatalf("%s", st)
}
})
t.Run("F10_recover_clears_pushed_and_extends_grace", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
req := baseSend("rec1", "bob")
req.Offline = keepTrue()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "rec1")
// 模拟保留期在停机期间已过
_ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE deliveries SET expire_at=?, pushed_conn=? WHERE seq=?`, e.nowMs-1000, "old-conn", seq)
return err
})
_ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET online_since=?, offline_since=NULL WHERE id=?`, e.nowMs-5000, "bob")
return err
})
start := e.nowMs + 1_000
e.setNow(start)
if err := e.app.RecoverOnStart(ctx); err != nil {
t.Fatal(err)
}
var pushed sql.NullString
var exp sql.NullInt64
_ = e.db.Read.QueryRow(`SELECT pushed_conn, expire_at FROM deliveries WHERE seq=?`, seq).Scan(&pushed, &exp)
if pushed.Valid {
t.Fatalf("pushed_conn still set: %v", pushed.String)
}
wantMin := start + 60_000
if !exp.Valid || exp.Int64 < wantMin {
t.Fatalf("expire_at=%v want >= %d", exp, wantMin)
}
var offline sql.NullInt64
_ = e.db.Read.QueryRow(`SELECT offline_since FROM endpoints WHERE id=?`, "bob").Scan(&offline)
if !offline.Valid || offline.Int64 != start {
t.Fatalf("offline_since=%v want %d", offline, start)
}
})
t.Run("disabled_target_rejected", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
delay := int64(5_000)
req := baseSend("dis1", "bob")
req.DelayMs = &delay
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
_ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET enabled=0 WHERE id=?`, "bob")
return err
})
e.setNow(e.nowMs + 5_000)
if _, err := e.app.DispatchDue(ctx, e.nowMs, 10); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "dis1")
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryRejected || reason != ReasonEndpointDisabled {
t.Fatalf("%s/%s", st, reason)
}
})
t.Run("queue_full_rejected", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) { l.MaxPendingPerReceiver = 1 })
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
req1 := baseSend("qf1", "bob")
req1.Offline = keepTrue()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req1); err != nil {
t.Fatal(err)
}
req2 := baseSend("qf2", "bob")
req2.Offline = keepTrue()
res, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req2)
if err != nil {
t.Fatal(err)
}
if res.State != StateCompleted {
t.Fatalf("state=%s", res.State)
}
seq := e.seqOf("alice", "qf2")
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryRejected || reason != ReasonQueueFull {
t.Fatalf("%s/%s", st, reason)
}
})
t.Run("sender_left_no_recipients", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
ctx := context.Background()
_ = e.db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`INSERT INTO groups(id,name,owner_id,created_at) VALUES('g2','g','alice',?)`, e.nowMs)
if err != nil {
return err
}
_, err = tx.Exec(`INSERT INTO group_members(group_id,endpoint_id,joined_at) VALUES('g2','alice',?)`, e.nowMs)
return err
})
delay := int64(1000)
req := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "nr1",
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g2"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"}, DelayMs: &delay,
}
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
e.setNow(e.nowMs + 1000)
if _, err := e.app.DispatchDue(ctx, e.nowMs, 10); err != nil {
t.Fatal(err)
}
st, reason := e.msgState("alice", "nr1")
if st != StateCompleted || reason != ReasonNoRecipients {
t.Fatalf("%s/%s", st, reason)
}
})
t.Run("too_large_rejected", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.conns.Set("bob", LiveConn{ConnID: "c-bob", MaxReceiveBytes: 50})
ctx := context.Background()
req := baseSend("big1", "bob")
req.Body.Data = string(make([]byte, 200))
for i := range req.Body.Data {
// utf8 valid
_ = i
}
req.Body.Data = "{\"x\":\"" + string(make([]byte, 80)) + "\"}"
// simpler: long ascii
b := make([]byte, 80)
for i := range b {
b[i] = 'a'
}
req.Body.Data = string(b)
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "big1")
st, reason := e.deliveryState(seq, "bob")
if st != DeliveryRejected || reason != ReasonTooLarge {
t.Fatalf("%s/%s", st, reason)
}
})
t.Run("on_publish_dropped_clears_mark", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("drop1", "bob")); err != nil {
t.Fatal(err)
}
if err := e.app.PushPending(ctx, "bob", "c-bob"); err != nil {
t.Fatal(err)
}
payload := e.down.Snapshots()[0].Payload
if err := e.app.OnPublishDropped(ctx, "bob", "c-bob", payload); err != nil {
t.Fatal(err)
}
seq := e.seqOf("alice", "drop1")
var pushed sql.NullString
_ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=?`, seq).Scan(&pushed)
if pushed.Valid {
t.Fatalf("still pushed %s", pushed.String)
}
})
t.Run("ack_duplicate_and_terminal", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("ack1", "bob")); err != nil {
t.Fatal(err)
}
_ = e.app.PushPending(ctx, "bob", "c-bob")
r1, err := e.app.Ack(ctx, "bob", &protocol.Ack{V: 1, Type: protocol.TypeAck, RID: "1", From: "alice", ID: "ack1"})
if err != nil {
t.Fatal(err)
}
if r1.Result != DeliveryAccepted {
t.Fatalf("%+v", r1)
}
r2, err := e.app.Ack(ctx, "bob", &protocol.Ack{V: 1, Type: protocol.TypeAck, RID: "2", From: "alice", ID: "ack1"})
if err != nil {
t.Fatal(err)
}
if r2.Result != DeliveryAccepted {
t.Fatalf("dup %+v", r2)
}
})
t.Run("disconnect_extends_grace", func(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("dc1", "bob")); err != nil {
t.Fatal(err)
}
_ = e.app.PushPending(ctx, "bob", "c-bob")
if err := e.app.OnDisconnect(ctx, "bob", "c-bob", true); err != nil {
t.Fatal(err)
}
e.conns.Clear("bob", "c-bob")
seq := e.seqOf("alice", "dc1")
var pushed sql.NullString
var exp sql.NullInt64
_ = e.db.Read.QueryRow(`SELECT pushed_conn, expire_at FROM deliveries WHERE seq=?`, seq).Scan(&pushed, &exp)
if pushed.Valid {
t.Fatal("pushed should clear")
}
if !exp.Valid || exp.Int64 != e.nowMs+60_000 {
t.Fatalf("expire_at=%v", exp)
}
})
}
func TestRecallDoesNotWriteReceipt(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
ctx := context.Background()
req := baseSend("nr", "bob")
req.Offline = keepTrue()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
if _, err := e.app.Recall(ctx, "alice", &protocol.Recall{V: 1, Type: protocol.TypeRecall, RID: "1", ID: "nr"}); err != nil {
t.Fatal(err)
}
var n int
_ = e.db.Read.QueryRow(`SELECT COUNT(*) FROM receipts WHERE msg_id=?`, "nr").Scan(&n)
if n != 0 {
t.Fatalf("receipts=%d", n)
}
}
func TestPushRevokedOnRecallAfterPush(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
e.online("bob", "c-bob")
ctx := context.Background()
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("rv1", "bob")); err != nil {
t.Fatal(err)
}
_ = e.app.PushPending(ctx, "bob", "c-bob")
if _, err := e.app.Recall(ctx, "alice", &protocol.Recall{V: 1, Type: protocol.TypeRecall, RID: "1", ID: "rv1"}); err != nil {
t.Fatal(err)
}
found := false
for _, p := range e.down.Snapshots() {
var head struct {
Type string `json:"type"`
}
_ = json.Unmarshal(p.Payload, &head)
if head.Type == protocol.TypeRevoked {
found = true
}
}
if !found {
t.Fatal("expected revoked frame")
}
}
+324
View File
@@ -0,0 +1,324 @@
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
}
-52
View File
@@ -1,52 +0,0 @@
package message
import (
"database/sql"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
// dispatchMinimalTx 是 M1 最小分发:单聊插一条 pending;群按当前成员去掉发送者各插 pending;
// 消息改为 dispatched。完整 7.4 规则见 DEVIATIONS「消息 M」。
func dispatchMinimalTx(tx *sql.Tx, seq int64, senderID, destKind, destID string, sendAt int64, keep int, nowMs int64) (string, error) {
recipients := make([]string, 0, 8)
switch destKind {
case protocol.TargetEndpoint:
recipients = append(recipients, destID)
case protocol.TargetGroup:
rows, err := tx.Query(`
SELECT endpoint_id FROM group_members WHERE group_id = ? AND endpoint_id != ?`, destID, senderID)
if err != nil {
return "", err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return "", err
}
recipients = append(recipients, id)
}
if err := rows.Err(); err != nil {
return "", err
}
default:
return "", errCode(protocol.CodeBadRequest, "invalid dest_kind")
}
for _, ep := range recipients {
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,NULL,0,?)`,
seq, ep, sendAt, keep, DeliveryPending, "", nowMs,
); err != nil {
return "", err
}
}
state := StateDispatched
if _, err := tx.Exec(`UPDATE messages SET state = ? WHERE seq = ?`, state, seq); err != nil {
return "", err
}
return state, nil
}
+551
View File
@@ -0,0 +1,551 @@
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)
}()
}
+122
View File
@@ -0,0 +1,122 @@
package message
import (
"context"
"database/sql"
)
// RecoverOnStart 启动恢复(DEVELOPMENT 7.8)。
func (a *App) RecoverOnStart(ctx context.Context) error {
nowMs := a.now().UnixMilli()
graceMs := a.lim.GraceSeconds * 1000
minExpire := nowMs + graceMs
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, err := tx.Exec(`
UPDATE deliveries SET
pushed_conn = NULL,
expire_at = CASE
WHEN expire_at IS NULL OR expire_at < ? THEN ?
ELSE expire_at
END,
updated_at = ?
WHERE state = 'pending'`, minExpire, minExpire, nowMs); err != nil {
return err
}
// 停机前在线:online_since 晚于 offline_since,或 offline_since 空而 online_since 非空
_, err := tx.Exec(`
UPDATE endpoints SET offline_since = ?
WHERE online_since IS NOT NULL
AND (offline_since IS NULL OR online_since > offline_since)`, nowMs)
return err
})
if err != nil {
return err
}
// 停机期间到点的 scheduled 立即分发
_, err = a.DispatchDue(ctx, nowMs, 1000)
return err
}
// CleanupOnce 处理未推送且到期的 pending,并做记录/回执/防重清理。
func (a *App) CleanupOnce(ctx context.Context, nowMs int64) error {
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
rows, err := tx.Query(`
SELECT d.seq, d.endpoint_id, d.keep, m.sender_id, m.id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.state = 'pending' AND d.pushed_conn IS NULL
AND d.expire_at IS NOT NULL AND d.expire_at <= ?`, nowMs)
if err != nil {
return err
}
type item struct {
seq, keep int64
endpointID, senderID string
msgID string
}
var list []item
for rows.Next() {
var it item
if err := rows.Scan(&it.seq, &it.endpointID, &it.keep, &it.senderID, &it.msgID); err != nil {
_ = rows.Close()
return err
}
list = append(list, it)
}
_ = rows.Close()
for _, it := range list {
state := DeliveryDropped
reason := ReasonOffline
if it.keep != 0 {
state = DeliveryExpired
reason = ReasonTTL
}
if err := a.finishDeliveryTx(tx, it.seq, it.endpointID, it.senderID, it.msgID, state, reason, false, nowMs); err != nil {
return err
}
}
if a.lim.RecordRetentionDays > 0 {
cutoff := nowMs - int64(a.lim.RecordRetentionDays)*24*3600*1000
if _, err := tx.Exec(`
DELETE FROM messages WHERE seq IN (
SELECT seq FROM messages WHERE state = 'completed' AND created_at < ? LIMIT 5000
)`, cutoff); err != nil {
return err
}
}
if a.lim.ReceiptRetentionDays > 0 {
cutoff := nowMs - int64(a.lim.ReceiptRetentionDays)*24*3600*1000
if _, err := tx.Exec(`
DELETE FROM receipts WHERE receipt_id IN (
SELECT receipt_id FROM receipts WHERE created_at < ? LIMIT 5000
)`, cutoff); err != nil {
return err
}
}
if a.lim.IdempotencyHours > 0 {
cutoff := nowMs - int64(a.lim.IdempotencyHours)*3600*1000
if _, err := tx.Exec(`
DELETE FROM send_keys WHERE rowid IN (
SELECT sk.rowid FROM send_keys sk
WHERE sk.created_at < ?
AND NOT EXISTS (
SELECT 1 FROM messages m WHERE m.sender_id = sk.sender_id AND m.id = sk.msg_id
)
LIMIT 5000
)`, cutoff); err != nil {
return err
}
}
return nil
})
if err != nil {
return err
}
a.flushRevokes(ctx)
return nil
}
+75
View File
@@ -0,0 +1,75 @@
package message
import (
"context"
"database/sql"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
)
// OnHandshakeComplete 握手完成:写 online_since、清空不 keep 的 expire_at,并推送。
// 调用方须先把连接登记进 ConnRegistry(MemoryConns.Set)。
func (a *App) OnHandshakeComplete(ctx context.Context, endpointID string, conn LiveConn) error {
nowMs := a.now().UnixMilli()
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, err := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, nowMs, endpointID); err != nil {
return err
}
_, err := tx.Exec(`
UPDATE deliveries SET expire_at = NULL, updated_at = ?
WHERE endpoint_id = ? AND state = 'pending' AND keep = 0`, nowMs, endpointID)
return err
})
if err != nil {
return err
}
return a.PushPending(ctx, endpointID, conn.ConnID)
}
// OnDisconnect 连接断开:当前连接则延长宽限;按代号清 pushed_conn。
// 调用方负责从 ConnRegistry 移除连接。
func (a *App) OnDisconnect(ctx context.Context, endpointID string, connID port.ConnID, isCurrent bool) error {
nowMs := a.now().UnixMilli()
graceMs := a.lim.GraceSeconds * 1000
deadline := nowMs + graceMs
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if isCurrent {
if _, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, nowMs, endpointID); err != nil {
return err
}
if _, err := tx.Exec(`
UPDATE deliveries SET
expire_at = CASE
WHEN expire_at IS NOT NULL AND expire_at > ? THEN expire_at
ELSE ?
END,
updated_at = ?
WHERE endpoint_id = ? AND state = 'pending' AND keep = 0`, deadline, deadline, nowMs, endpointID); err != nil {
return err
}
if _, err := tx.Exec(`
UPDATE deliveries SET
expire_at = CASE
WHEN expire_at IS NULL OR expire_at < ? THEN ?
ELSE expire_at
END,
updated_at = ?
WHERE endpoint_id = ? AND state = 'pending' AND keep = 1 AND pushed_conn = ?`,
deadline, deadline, nowMs, endpointID, string(connID)); err != nil {
return err
}
}
_, err := tx.Exec(`
UPDATE deliveries SET pushed_conn = NULL, updated_at = ?
WHERE state = 'pending' AND pushed_conn = ?`, nowMs, string(connID))
return err
})
if err != nil {
return err
}
if live, ok := a.lookupConn(endpointID); ok && live.ConnID != connID {
a.WakePush(endpointID)
}
return nil
}
+24 -2
View File
@@ -12,7 +12,7 @@ import (
"git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/protocol"
) )
// Submit 处理发送提交(DEVELOPMENT 7.3):防重 → 校验/配额/授权 → 写入 → 到点则最小分发。 // Submit 处理发送提交(DEVELOPMENT 7.3):防重 → 校验/配额/授权 → 写入 → 到点则完整分发。
func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, req *protocol.Send) (SubmitResult, error) { func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, req *protocol.Send) (SubmitResult, error) {
if req == nil { if req == nil {
return SubmitResult{}, errCode(protocol.CodeBadRequest, "nil send") return SubmitResult{}, errCode(protocol.CodeBadRequest, "nil send")
@@ -250,10 +250,12 @@ INSERT INTO messages(
finalState := state finalState := state
if sendAt <= nowMs { if sendAt <= nowMs {
finalState, e = dispatchMinimalTx(tx, seq, senderID, req.To.Kind, req.To.ID, sendAt, keepInt, nowMs) var claimed bool
finalState, claimed, e = a.dispatchFullTx(tx, seq, senderID, req.To.Kind, req.To.ID, sendAt, keepInt, ttl, receipt, nowMs)
if e != nil { if e != nil {
return e return e
} }
_ = claimed
} }
result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState} result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState}
return nil return nil
@@ -261,9 +263,29 @@ INSERT INTO messages(
if err != nil { if err != nil {
return SubmitResult{}, err return SubmitResult{}, err
} }
if result.State == StateDispatched {
a.wakeReceivers(ctx, result.ID, senderID)
}
return result, nil return result, nil
} }
func (a *App) wakeReceivers(ctx context.Context, msgID, senderID string) {
rows, err := a.db.Read.QueryContext(ctx, `
SELECT d.endpoint_id FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE m.sender_id = ? AND m.id = ? AND d.state = 'pending'`, senderID, msgID)
if err != nil {
return
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var ep string
if rows.Scan(&ep) == nil {
a.WakePush(ep)
}
}
}
func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (int64, error) { func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (int64, error) {
if req.SendAtMs != nil && req.DelayMs != nil { if req.SendAtMs != nil && req.DelayMs != nil {
return 0, errCode(protocol.CodeBadRequest, "send_at_ms and delay_ms are mutually exclusive") return 0, errCode(protocol.CodeBadRequest, "send_at_ms and delay_ms are mutually exclusive")
+7 -3
View File
@@ -50,6 +50,9 @@ func defaultTestLimits() Limits {
cfg := config.Default().Limits cfg := config.Default().Limits
lim := LimitsFromConfig(cfg) lim := LimitsFromConfig(cfg)
lim.RequestsPerSecond = 0 // 测试默认不限速 lim.RequestsPerSecond = 0 // 测试默认不限速
lim.RecordRetentionDays = 7
lim.ReceiptRetentionDays = 7
lim.IdempotencyHours = 24
return lim return lim
} }
@@ -62,11 +65,12 @@ func insertEndpoint(t *testing.T, db *store.DB, id string, talkPassword string,
talk = "stub$" + talkPassword talk = "stub$" + talkPassword
talkVer = 1 talkVer = 1
} }
nowMs := int64(1_700_000_000_000)
err := db.Queue.Do(ctx, func(tx *sql.Tx) error { err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(` _, e := tx.Exec(`
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at, offline_since)
VALUES(?,?,?,?,?,?,?,?)`, VALUES(?,?,?,?,?,?,?,?,?)`,
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, 1_700_000_000_000) id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, nowMs, nowMs)
return e return e
}) })
if err != nil { if err != nil {