2 Commits
17 changed files with 3985 additions and 133 deletions
+92 -12
View File
@@ -276,37 +276,117 @@
- 备选方案:N2 暴露回调给 M 注册。 - 备选方案:N2 暴露回调给 M 注册。
- 影响:接线后 M 需订阅或包装该钩子;当前接口可后续加 `OnPublishDropped` 回调字段。 - 影响:接线后 M 需订阅或包装该钩子;当前接口可后续加 `OnPublishDropped` 回调字段。
### N3 2026-09-30
1. **仍未接线 `cmd/nixmsg`**
- 原条款:serve 最终应挂上真实 Authenticator / Session。
- 实际做法:交付 `broker.Login`、`broker.Session` 与 F02 测试;不改 `cmd/nixmsg`/`wire.go`。
- 原因:与总控/其他线并行改 wire 冲突;N1/N2 已约定合并时接线。
- 备选方案:本分支改 wire(与隔离指令冲突)。
- 影响:进程默认仍 RejectAuthenticator,需接线注入 `Login`+`Session`。
2. **`session_hash` 存十六进制文本**
- 原条款:库中存 SHA-256;列为 TEXT,未规定编码。
- 实际做法:存 32 字节哈希的小写 hex(与后台 API 令牌存法一致)。
- 原因:TEXT 列无法直接存原始字节;hex 便于排查。
- 备选方案:BLOB 列或 base64。
- 影响:其他线读写 `session_hash` 需按 hex 编解码。
3. **上下线通知走 `PresenceSink` + port 回调**
- 原条款:写 `online_since`/`offline_since` 并通知;通过现有 port 接口供身份线订阅。
- 实际做法:N3 自己写时间戳;可选注入 `PresenceSink`(对齐 `presence.Service.SetOnline/SetOffline`);并继续调用 `OnHandshakeComplete`/`OnDisconnect`。旧连接断开用连接代号判断,只有当时仍是 current 才标离线。
- 原因:I3 尚未合入,不能依赖具体 presence 实现;双通道便于接线。
- 备选方案:只靠 port、由 I 线写库(与「N3 写 online_since」字面不符)。
- 影响:接线时避免 I 线重复写同一时间戳即可。
4. **InlineClient 的 `OnPublish` 必须放行**
- 原条款:客户端上行 `OnPublish` 返回 `CodeSuccessIgnore`。
- 实际做法:`cl.Net.Inline` 时原样返回,不 Ignore,否则 `PublishDown` 无法送达订阅者。
- 原因:mochi `Publish` 经 InlineClient `InjectPacket` 再进 `OnPublish`。
- 备选方案:不用 InlineClient,改直接 `publishToClient`(偏离文档装配)。
- 影响:N2 既有 PublishDown 测试此前未读回包,此缺陷在 N3 才暴露并修复。
5. **管理员踢线类入口挂在 `Session`**
- 原条款:停用/删除/重置密码先 fatal 再断开;踢下线只断开。
- 实际做法:`Session.Disable`/`Deleted`/`ResetPassword`/`Kick` 可调用;管理 HTTP 未接。
- 原因:A2 管理接口尚未接线。
- 备选方案:放到 `internal/admin`(超出 N 目录)。
- 影响:A/I 接线时调用这些方法即可。
## 消息 M ## 消息 M
### 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
})
}
+61 -50
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 {
+296
View File
@@ -0,0 +1,296 @@
package broker
import (
"context"
"database/sql"
"encoding/hex"
"errors"
"strings"
"sync"
"time"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
// ErrEndpointNotFound 编号不存在(业务拒绝,非内部故障)。
var ErrEndpointNotFound = errors.New("broker: endpoint not found")
// ErrEndpointDisabled 端已停用(业务拒绝)。
var ErrEndpointDisabled = errors.New("broker: endpoint disabled")
// Login 实现 Authenticator:会话令牌或登录密码(含锁定)。
type Login struct {
DB *store.DB
Pool auth.HashPool
Tokens auth.SessionTokens
Locks auth.LoginLocks
IdleDays int
Now func() time.Time
usedMu sync.Mutex
// 内存中的 session_used_at(毫秒)与上次落库时间。
usedAt map[string]int64
lastFlush map[string]int64
}
// LoginOptions 装配 Login。
type LoginOptions struct {
DB *store.DB
Pool auth.HashPool
Tokens auth.SessionTokens
Locks auth.LoginLocks
IdleDays int
Now func() time.Time
}
// NewLogin 创建登录校验器。
func NewLogin(opts LoginOptions) *Login {
now := opts.Now
if now == nil {
now = time.Now
}
tokens := opts.Tokens
if tokens == nil {
tokens = auth.NewSessionTokens()
}
locks := opts.Locks
if locks == nil {
locks = auth.NewLoginLocks()
}
return &Login{
DB: opts.DB,
Pool: opts.Pool,
Tokens: tokens,
Locks: locks,
IdleDays: opts.IdleDays,
Now: now,
usedAt: make(map[string]int64),
lastFlush: make(map[string]int64),
}
}
// Authenticate 按 DEVELOPMENT 第 5 节校验;内部故障返回 error。
func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) {
if l == nil || l.DB == nil {
return AuthResult{}, errors.New("broker: login not configured")
}
if endpointID == "" {
return AuthResult{OK: false}, nil
}
row, err := l.loadEndpoint(ctx, endpointID)
if err != nil {
if errors.Is(err, ErrEndpointNotFound) || errors.Is(err, ErrEndpointDisabled) {
return AuthResult{OK: false}, nil
}
return AuthResult{}, err
}
cred := string(password)
if l.Tokens.LooksLikeSessionToken(cred) {
ok, authErr := l.authSession(ctx, endpointID, cred, row)
if authErr != nil {
return AuthResult{}, authErr
}
return AuthResult{OK: ok}, nil
}
ok, token, authErr := l.authPassword(ctx, endpointID, cred, remoteIP, row)
if authErr != nil {
return AuthResult{}, authErr
}
return AuthResult{OK: ok, SessionToken: token}, nil
}
type endpointAuthRow struct {
loginHash string
sessionHash []byte // 原始 32 字节;无令牌时 nil
sessionUsedAt int64 // 毫秒;无则 0
}
func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) {
var (
loginHash string
enabled int
sessHex sql.NullString
usedAt sql.NullInt64
)
err := l.DB.Read.QueryRowContext(ctx, `
SELECT login_hash, enabled, session_hash, session_used_at
FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return endpointAuthRow{}, ErrEndpointNotFound
}
return endpointAuthRow{}, err
}
if enabled == 0 {
return endpointAuthRow{}, ErrEndpointDisabled
}
row := endpointAuthRow{loginHash: loginHash}
if usedAt.Valid {
row.sessionUsedAt = usedAt.Int64
}
if sessHex.Valid && sessHex.String != "" {
raw, decErr := hex.DecodeString(sessHex.String)
if decErr != nil || len(raw) != 32 {
// 损坏的哈希视为无有效会话(令牌校验失败),不是内部故障
row.sessionHash = nil
} else {
row.sessionHash = raw
}
}
return row, nil
}
func (l *Login) authSession(ctx context.Context, endpointID, token string, row endpointAuthRow) (bool, error) {
if len(row.sessionHash) == 0 {
return false, nil
}
got := l.Tokens.HashToken(token)
if !auth.EqualHash(got, row.sessionHash) {
return false, nil
}
now := l.Now()
nowMs := now.UnixMilli()
usedAt := row.sessionUsedAt
l.usedMu.Lock()
if mem, ok := l.usedAt[endpointID]; ok && mem > usedAt {
usedAt = mem
}
l.usedMu.Unlock()
if l.IdleDays > 0 {
idle := time.Duration(l.IdleDays) * 24 * time.Hour
if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle {
return false, nil
}
}
if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil {
return false, err
}
return true, nil
}
func (l *Login) touchSessionUsed(ctx context.Context, endpointID string, nowMs int64) error {
const flushEvery = int64(time.Hour / time.Millisecond)
l.usedMu.Lock()
l.usedAt[endpointID] = nowMs
last := l.lastFlush[endpointID]
needFlush := last == 0 || nowMs-last >= flushEvery
if needFlush {
l.lastFlush[endpointID] = nowMs
}
l.usedMu.Unlock()
if !needFlush {
return nil
}
return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET session_used_at = ? WHERE id = ? AND session_hash IS NOT NULL AND session_hash != ''`,
nowMs, endpointID)
return err
})
}
func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP string, row endpointAuthRow) (ok bool, token string, err error) {
ipKey := auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}
epKey := auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}
if locked, _ := l.Locks.Check(ipKey); locked {
return false, "", nil
}
if locked, _ := l.Locks.Check(epKey); locked {
return false, "", nil
}
if l.Pool == nil {
return false, "", errors.New("broker: password pool not configured")
}
match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash)
if verErr != nil {
return false, "", verErr
}
if !match {
l.Locks.Fail(ipKey)
l.Locks.Fail(epKey)
return false, "", nil
}
tok, hash, issErr := l.Tokens.Issue(ctx)
if issErr != nil {
return false, "", issErr
}
nowMs := l.Now().UnixMilli()
hashHex := hex.EncodeToString(hash)
writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`
UPDATE endpoints
SET session_hash = ?, session_issued_at = ?, session_used_at = ?
WHERE id = ?`, hashHex, nowMs, nowMs, endpointID)
return e
})
if writeErr != nil {
return false, "", writeErr
}
l.usedMu.Lock()
l.usedAt[endpointID] = nowMs
l.lastFlush[endpointID] = nowMs
l.usedMu.Unlock()
return true, tok, nil
}
// ClearSession 清空会话令牌(logout / 停用 / 删除 / 重置密码)。
func (l *Login) ClearSession(ctx context.Context, endpointID string) error {
if l == nil || l.DB == nil {
return errors.New("broker: login not configured")
}
err := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`
UPDATE endpoints
SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL
WHERE id = ?`, endpointID)
return e
})
if err != nil {
return err
}
l.usedMu.Lock()
delete(l.usedAt, endpointID)
delete(l.lastFlush, endpointID)
l.usedMu.Unlock()
return nil
}
// SetOnlineSince 握手完成时写入 online_since。
func (l *Login) SetOnlineSince(ctx context.Context, endpointID string, atMs int64) error {
return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, atMs, endpointID)
return err
})
}
// SetOfflineSince 当前连接断开时写入 offline_since。
func (l *Login) SetOfflineSince(ctx context.Context, endpointID string, atMs int64) error {
return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, atMs, endpointID)
return err
})
}
// SessionHashOf 返回当前库中的会话哈希(测试用);无则 nil。
func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, error) {
var sessHex sql.NullString
err := l.DB.Read.QueryRowContext(ctx, `SELECT session_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&sessHex)
if err != nil {
return nil, err
}
if !sessHex.Valid || sessHex.String == "" {
return nil, nil
}
return hex.DecodeString(sessHex.String)
}
// LooksLikeSessionToken 暴露给测试。
func (l *Login) LooksLikeSessionToken(s string) bool {
return strings.HasPrefix(s, "nst_")
}
var _ Authenticator = (*Login)(nil)
+117 -13
View File
@@ -9,6 +9,7 @@ import (
"net" "net"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
mqtt "github.com/mochi-mqtt/server/v2" mqtt "github.com/mochi-mqtt/server/v2"
@@ -85,18 +86,22 @@ type Broker struct {
} }
type connState struct { type connState struct {
connID port.ConnID connID port.ConnID
endpointID string endpointID string
transport port.Transport transport port.Transport
remoteIP string remoteIP string
client *mqtt.Client client *mqtt.Client
maxPacketSize uint32 maxPacketSize uint32
maxRecvBytes int maxRecvBytes int
authOK bool authOK bool
authErr error authErr error
sessionToken string sessionToken string
largeHeld int handshook bool
mu sync.Mutex subscribedDown bool
largeHeld int
mu sync.Mutex
handshakeTimer *time.Timer
} }
// New 创建并 Serve mochi(无监听器)。 // New 创建并 Serve mochi(无监听器)。
@@ -265,7 +270,12 @@ func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.Co
case port.DisconnectKicked, port.DisconnectFatal: case port.DisconnectKicked, port.DisconnectFatal:
code = packets.ErrAdministrativeAction code = packets.ErrAdministrativeAction
} }
return b.server.DisconnectClient(st.client, code) err := b.server.DisconnectClient(st.client, code)
// mochi 对错误类原因码会把 Code 当作 error 返回,表示已按该原因断开,不算失败。
if _, ok := err.(packets.Code); ok {
return nil
}
return err
} }
func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState { func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
@@ -362,6 +372,100 @@ func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) {
}, true }, true
} }
// IsHandshook 当前连接是否已完成握手。
func (b *Broker) IsHandshook(endpointID string) bool {
b.connsMu.RLock()
st := b.current[endpointID]
b.connsMu.RUnlock()
if st == nil {
return false
}
st.mu.Lock()
defer st.mu.Unlock()
return st.handshook
}
// CurrentConnID 返回端的当前连接代号。
func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) {
b.connsMu.RLock()
st := b.current[endpointID]
b.connsMu.RUnlock()
if st == nil {
return "", false
}
return st.connID, true
}
func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
for _, st := range b.byClient {
if st.endpointID == endpointID && st.connID == connID {
return st
}
}
return nil
}
func (b *Broker) hasDownSub(st *connState) bool {
if st == nil {
return false
}
st.mu.Lock()
defer st.mu.Unlock()
if st.subscribedDown {
return true
}
// 回退:直接看 mochi 订阅表
if st.client != nil && st.client.State.Subscriptions != nil {
_, ok := st.client.State.Subscriptions.Get(downTopic(st.endpointID))
return ok
}
return false
}
func (b *Broker) startHandshakeDeadline(endpointID string, connID port.ConnID, d time.Duration) {
st := b.connStateOf(endpointID, connID)
if st == nil {
return
}
st.mu.Lock()
if st.handshook {
st.mu.Unlock()
return
}
if st.handshakeTimer != nil {
st.handshakeTimer.Stop()
}
st.handshakeTimer = time.AfterFunc(d, func() {
cur := b.connStateOf(endpointID, connID)
if cur == nil {
return
}
cur.mu.Lock()
done := cur.handshook
cur.mu.Unlock()
if done {
return
}
_ = b.Disconnect(context.Background(), endpointID, connID, port.DisconnectIdle)
})
st.mu.Unlock()
}
func (b *Broker) cancelHandshakeDeadline(endpointID string, connID port.ConnID) {
st := b.connStateOf(endpointID, connID)
if st == nil {
return
}
st.mu.Lock()
if st.handshakeTimer != nil {
st.handshakeTimer.Stop()
st.handshakeTimer = nil
}
st.mu.Unlock()
}
func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) { func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) {
b.queuesMu.Lock() b.queuesMu.Lock()
q, ok := b.queues[endpointID] q, ok := b.queues[endpointID]
+734
View File
@@ -0,0 +1,734 @@
package broker_test
import (
"bytes"
"context"
"database/sql"
"encoding/json"
"io"
"net"
"sync"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/broker"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
"github.com/mochi-mqtt/server/v2/packets"
)
type presenceRec struct {
mu sync.Mutex
online []string
offline []string
}
func (p *presenceRec) SetOnline(_ context.Context, endpointID string, _ port.ConnID, _ int64) error {
p.mu.Lock()
defer p.mu.Unlock()
p.online = append(p.online, endpointID)
return nil
}
func (p *presenceRec) SetOffline(_ context.Context, endpointID string, _ port.ConnID, _ int64) error {
p.mu.Lock()
defer p.mu.Unlock()
p.offline = append(p.offline, endpointID)
return nil
}
type uplinkRec struct {
port.StubUplinkHandler
mu sync.Mutex
handshakes int
disconnects []port.DisconnectReason
}
func (u *uplinkRec) OnHandshakeComplete(context.Context, port.HandshakeInfo) error {
u.mu.Lock()
defer u.mu.Unlock()
u.handshakes++
return nil
}
func (u *uplinkRec) OnDisconnect(_ context.Context, _ port.ConnInfo, reason port.DisconnectReason) {
u.mu.Lock()
defer u.mu.Unlock()
u.disconnects = append(u.disconnects, reason)
}
type testEnv struct {
t *testing.T
db *store.DB
login *broker.Login
sess *broker.Session
b *broker.Broker
presence *presenceRec
uplink *uplinkRec
pool auth.HashPool
dir string
}
func openEnv(t *testing.T, idleDays int) *testEnv {
t.Helper()
dir := t.TempDir()
db, err := store.Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
pool := auth.NewStubHashPool()
locks := auth.NewLoginLocks()
login := broker.NewLogin(broker.LoginOptions{
DB: db,
Pool: pool,
Tokens: auth.NewSessionTokens(),
Locks: locks,
IdleDays: idleDays,
})
pres := &presenceRec{}
up := &uplinkRec{}
sess := broker.NewSession(broker.SessionOptions{
Login: login,
Inner: up,
Presence: pres,
Limits: broker.HelloLimits{ServerVersion: "0.1.0-test"},
})
b, err := broker.New(broker.Options{Authenticator: login, Uplink: sess})
if err != nil {
t.Fatal(err)
}
sess.Attach(b)
t.Cleanup(func() {
_ = b.Close()
_ = db.Close()
})
return &testEnv{t: t, db: db, login: login, sess: sess, b: b, presence: pres, uplink: up, pool: pool, dir: dir}
}
func (e *testEnv) insertEndpoint(id, password string) {
e.t.Helper()
phc, err := e.pool.Hash(context.Background(), auth.PasswordLogin, password)
if err != nil {
e.t.Fatal(err)
}
now := time.Now().UnixMilli()
err = e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, execErr := tx.Exec(`
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
VALUES (?, '', ?, NULL, 0, 0, 1, ?)`, id, phc, now)
return execErr
})
if err != nil {
e.t.Fatal(err)
}
}
type pipeClient struct {
t *testing.T
conn net.Conn
done chan struct{}
packet uint16
}
func (e *testEnv) dial() *pipeClient {
e.t.Helper()
r, w := net.Pipe()
done := make(chan struct{})
go func() {
defer close(done)
_ = e.b.AttachTCP(r)
}()
return &pipeClient{t: e.t, conn: w, done: done, packet: 1}
}
func (c *pipeClient) close() {
_ = c.conn.Close()
select {
case <-c.done:
case <-time.After(3 * time.Second):
}
}
func (c *pipeClient) connect(endpoint, password string, maxPacket uint32) (connack byte, ok bool) {
c.t.Helper()
pk := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Connect},
ProtocolVersion: 5,
Connect: packets.ConnectParams{
ProtocolName: []byte("MQTT"),
Clean: true,
ClientIdentifier: endpoint,
Keepalive: 30,
UsernameFlag: true,
Username: []byte(endpoint),
PasswordFlag: true,
Password: []byte(password),
},
Properties: packets.Properties{MaximumPacketSize: maxPacket},
}
var buf bytes.Buffer
if err := pk.ConnectEncode(&buf); err != nil {
c.t.Fatal(err)
}
if _, err := c.conn.Write(buf.Bytes()); err != nil {
c.t.Fatal(err)
}
_ = c.conn.SetReadDeadline(time.Now().Add(3 * time.Second))
raw := make([]byte, 256)
n, err := io.ReadAtLeast(c.conn, raw, 2)
if err != nil {
return 0, false
}
if raw[0]>>4 != packets.Connack {
c.t.Fatalf("want connack got %x", raw[:n])
}
// MQTT5 CONNACK: type, remaining len, flags, reason
reason := byte(0)
if n >= 4 {
reason = raw[3]
}
return reason, reason == 0
}
func (c *pipeClient) expectNoConnack() {
c.t.Helper()
_ = c.conn.SetReadDeadline(time.Now().Add(400 * time.Millisecond))
buf := make([]byte, 64)
n, err := c.conn.Read(buf)
if err == nil && n > 0 && buf[0]>>4 == packets.Connack {
c.t.Fatalf("unexpected connack %x", buf[:n])
}
}
func (c *pipeClient) subscribe(endpoint string) {
c.t.Helper()
c.packet++
pk := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1},
ProtocolVersion: 5,
PacketID: c.packet,
Filters: packets.Subscriptions{
{Filter: "nix/c/" + endpoint + "/down", Qos: 1},
},
}
var buf bytes.Buffer
if err := pk.SubscribeEncode(&buf); err != nil {
c.t.Fatal(err)
}
if _, err := c.conn.Write(buf.Bytes()); err != nil {
c.t.Fatal(err)
}
_ = c.conn.SetReadDeadline(time.Now().Add(3 * time.Second))
raw := make([]byte, 256)
n, err := io.ReadAtLeast(c.conn, raw, 2)
if err != nil {
c.t.Fatal(err)
}
if raw[0]>>4 != packets.Suback {
c.t.Fatalf("want suback got %x", raw[:n])
}
}
func (c *pipeClient) publishUp(endpoint string, payload []byte) {
c.t.Helper()
c.packet++
pk := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1},
ProtocolVersion: 5,
TopicName: "nix/c/" + endpoint + "/up",
PacketID: c.packet,
Payload: payload,
}
var buf bytes.Buffer
if err := pk.PublishEncode(&buf); err != nil {
c.t.Fatal(err)
}
if _, err := c.conn.Write(buf.Bytes()); err != nil {
c.t.Fatal(err)
}
}
func (c *pipeClient) readDownJSON(timeout time.Duration) map[string]any {
c.t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
_ = c.conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond))
hdr := make([]byte, 1)
if _, err := io.ReadFull(c.conn, hdr); err != nil {
continue
}
rem, err := readRemainingLength(c.conn)
if err != nil {
continue
}
body := make([]byte, rem)
if _, err := io.ReadFull(c.conn, body); err != nil {
continue
}
typ := hdr[0] >> 4
switch typ {
case packets.Publish:
pk := new(packets.Packet)
pk.ProtocolVersion = 5
pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem}
fhQos := (hdr[0] >> 1) & 0x3
pk.FixedHeader.Qos = fhQos
if err := pk.PublishDecode(body); err != nil {
c.t.Fatalf("publish decode: %v", err)
}
if fhQos > 0 {
ack := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Puback},
ProtocolVersion: 5,
PacketID: pk.PacketID,
}
var ab bytes.Buffer
_ = ack.PubackEncode(&ab)
_, _ = c.conn.Write(ab.Bytes())
}
var m map[string]any
if err := json.Unmarshal(pk.Payload, &m); err != nil {
c.t.Fatalf("json: %v payload=%s", err, pk.Payload)
}
return m
case packets.Puback, packets.Pingresp, packets.Disconnect:
continue
default:
continue
}
}
c.t.Fatal("timeout waiting down json")
return nil
}
func readRemainingLength(r io.Reader) (int, error) {
var mul uint32 = 1
var value uint32
for i := 0; i < 4; i++ {
var b [1]byte
if _, err := io.ReadFull(r, b[:]); err != nil {
return 0, err
}
value += uint32(b[0]&127) * mul
if b[0]&128 == 0 {
return int(value), nil
}
mul *= 128
}
return 0, io.ErrUnexpectedEOF
}
func helloPayload(rid string) []byte {
b, _ := protocol.Marshal(protocol.Hello{
V: protocol.Version, Type: protocol.TypeHello, RID: rid,
})
return b
}
func waitHandshook(t *testing.T, b *broker.Broker, endpoint string) {
t.Helper()
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
if b.IsHandshook(endpoint) {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatal("not handshook")
}
func TestF02PasswordLoginReturnsTokenAndHandshake(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep1", "password1")
c := e.dial()
defer c.close()
reason, ok := c.connect("ep1", "password1", 0)
if !ok {
t.Fatalf("connack reason=%d", reason)
}
c.subscribe("ep1")
c.publishUp("ep1", helloPayload("1"))
m := c.readDownJSON(3 * time.Second)
if m["type"] != "resp" || m["ok"] != true {
t.Fatalf("hello resp=%v", m)
}
data, _ := m["data"].(map[string]any)
tok, _ := data["session_token"].(string)
if tok == "" || tok[:4] != "nst_" {
t.Fatalf("session_token=%v", data["session_token"])
}
waitHandshook(t, e.b, "ep1")
e.presence.mu.Lock()
nOnline := len(e.presence.online)
e.presence.mu.Unlock()
if nOnline < 1 {
t.Fatal("expected presence online")
}
}
func TestF02TakenOverByPasswordLogin(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep2", "password1")
a := e.dial()
defer a.close()
if _, ok := a.connect("ep2", "password1", 0); !ok {
t.Fatal("A connect")
}
a.subscribe("ep2")
a.publishUp("ep2", helloPayload("1"))
_ = a.readDownJSON(3 * time.Second)
waitHandshook(t, e.b, "ep2")
infoA, _ := e.b.ConnInfoOf("ep2")
// 后台排空 A,避免顶号写 DISCONNECT 时 pipe 阻塞
go func() {
buf := make([]byte, 512)
for {
_ = a.conn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err := a.conn.Read(buf)
if err != nil {
return
}
}
}()
b := e.dial()
defer b.close()
if _, ok := b.connect("ep2", "password1", 0); !ok {
t.Fatal("B connect")
}
b.subscribe("ep2")
b.publishUp("ep2", helloPayload("2"))
m := b.readDownJSON(3 * time.Second)
data, _ := m["data"].(map[string]any)
tokB, _ := data["session_token"].(string)
if tokB == "" {
t.Fatal("B should get new token")
}
waitHandshook(t, e.b, "ep2")
infoB, ok := e.b.ConnInfoOf("ep2")
if !ok || infoB.ConnID == infoA.ConnID {
t.Fatalf("current should be B, got %+v old=%s", infoB, infoA.ConnID)
}
}
func TestF02OldTokenRejectedAfterPasswordLogin(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep3", "password1")
a := e.dial()
if _, ok := a.connect("ep3", "password1", 0); !ok {
t.Fatal("A")
}
a.subscribe("ep3")
a.publishUp("ep3", helloPayload("1"))
m := a.readDownJSON(3 * time.Second)
data, _ := m["data"].(map[string]any)
oldTok, _ := data["session_token"].(string)
a.close()
// 另一处密码登录换令牌
b := e.dial()
if _, ok := b.connect("ep3", "password1", 0); !ok {
t.Fatal("B")
}
b.subscribe("ep3")
b.publishUp("ep3", helloPayload("2"))
_ = b.readDownJSON(3 * time.Second)
b.close()
c := e.dial()
defer c.close()
reason, ok := c.connect("ep3", oldTok, 0)
if ok {
t.Fatal("old token should fail")
}
if reason != 0x86 {
t.Fatalf("want 0x86 got %#x", reason)
}
}
func TestF02TokenReconnectDifferentIPKeepsToken(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep4", "password1")
a := e.dial()
if _, ok := a.connect("ep4", "password1", 0); !ok {
t.Fatal("A")
}
a.subscribe("ep4")
a.publishUp("ep4", helloPayload("1"))
m := a.readDownJSON(3 * time.Second)
data, _ := m["data"].(map[string]any)
tok, _ := data["session_token"].(string)
hash1, err := e.login.SessionHashOf(context.Background(), "ep4")
if err != nil || hash1 == nil {
t.Fatalf("hash1=%v err=%v", hash1, err)
}
a.close()
b := e.dial()
defer b.close()
if _, ok := b.connect("ep4", tok, 0); !ok {
t.Fatal("token reconnect")
}
b.subscribe("ep4")
b.publishUp("ep4", helloPayload("2"))
m2 := b.readDownJSON(3 * time.Second)
data2, _ := m2["data"].(map[string]any)
if _, has := data2["session_token"]; has {
t.Fatalf("token reconnect must not return session_token: %v", data2)
}
hash2, _ := e.login.SessionHashOf(context.Background(), "ep4")
if !auth.EqualHash(hash1, hash2) {
t.Fatal("session hash changed on token reconnect")
}
}
func TestF02IPLockDoesNotAffectOtherIP(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep5", "password1")
login := e.login
for i := 0; i < 10; i++ {
res, err := login.Authenticate(context.Background(), "ep5", []byte("wrong-pass"), "1.1.1.1")
if err != nil {
t.Fatal(err)
}
if res.OK {
t.Fatal("should fail")
}
}
res, err := login.Authenticate(context.Background(), "ep5", []byte("password1"), "1.1.1.1")
if err != nil || res.OK {
t.Fatalf("locked same IP ok=%v err=%v", res.OK, err)
}
res, err = login.Authenticate(context.Background(), "ep5", []byte("password1"), "2.2.2.2")
if err != nil || !res.OK {
t.Fatalf("other IP ok=%v err=%v", res.OK, err)
}
}
func TestF02EndpointLockAllowsTokenReconnect(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep6", "password1")
login := e.login
// 先拿到令牌
res, err := login.Authenticate(context.Background(), "ep6", []byte("password1"), "10.0.0.1")
if err != nil || !res.OK || res.SessionToken == "" {
t.Fatalf("login=%+v err=%v", res, err)
}
tok := res.SessionToken
// 多 IP 累计 50 次失败
for i := 0; i < 50; i++ {
ip := "203.0.113." + itoa(i%250+1)
r, e2 := login.Authenticate(context.Background(), "ep6", []byte("bad"), ip)
if e2 != nil {
t.Fatal(e2)
}
if r.OK {
t.Fatal("unexpected ok")
}
}
// 密码登录暂停
r, err := login.Authenticate(context.Background(), "ep6", []byte("password1"), "198.51.100.1")
if err != nil || r.OK {
t.Fatalf("password should be locked ok=%v err=%v", r.OK, err)
}
// 令牌仍可
r, err = login.Authenticate(context.Background(), "ep6", []byte(tok), "198.51.100.9")
if err != nil || !r.OK {
t.Fatalf("token should work ok=%v err=%v", r.OK, err)
}
if r.SessionToken != "" {
t.Fatal("token auth must not issue new token")
}
}
func TestF02DBErrorClosesWithout086(t *testing.T) {
dir := t.TempDir()
db, err := store.Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
pool := auth.NewStubHashPool()
login := broker.NewLogin(broker.LoginOptions{
DB: db, Pool: pool, Tokens: auth.NewSessionTokens(), Locks: auth.NewLoginLocks(), IdleDays: 30,
})
phc, _ := pool.Hash(context.Background(), auth.PasswordLogin, "password1")
_ = db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, e := tx.Exec(`INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
VALUES ('ep7', '', ?, NULL, 0, 0, 1, ?)`, phc, time.Now().UnixMilli())
return e
})
_ = db.Read.Close()
b, err := broker.New(broker.Options{Authenticator: login})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
r, w := net.Pipe()
errCh := make(chan error, 1)
go func() { errCh <- b.AttachTCP(r) }()
c := &pipeClient{t: t, conn: w, done: make(chan struct{}), packet: 1}
c.expectNoConnack()
_ = w.Close()
select {
case <-errCh:
case <-time.After(2 * time.Second):
}
_ = db.Close()
}
func TestF02NotReadyBeforeHello(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep8", "password1")
c := e.dial()
defer c.close()
if _, ok := c.connect("ep8", "password1", 0); !ok {
t.Fatal("connect")
}
c.subscribe("ep8")
payload, _ := protocol.Marshal(map[string]any{
"v": 1, "type": "self.get", "rid": "9",
})
c.publishUp("ep8", payload)
m := c.readDownJSON(3 * time.Second)
if m["ok"] != false {
t.Fatalf("want not_ready resp got %v", m)
}
errObj, _ := m["error"].(map[string]any)
if errObj["code"] != protocol.CodeNotReady {
t.Fatalf("code=%v", errObj)
}
}
func TestF02LogoutClearsToken(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep9", "password1")
c := e.dial()
defer c.close()
if _, ok := c.connect("ep9", "password1", 0); !ok {
t.Fatal("connect")
}
c.subscribe("ep9")
c.publishUp("ep9", helloPayload("1"))
m := c.readDownJSON(3 * time.Second)
data, _ := m["data"].(map[string]any)
tok, _ := data["session_token"].(string)
waitHandshook(t, e.b, "ep9")
logout, _ := protocol.Marshal(protocol.SelfLogout{V: protocol.Version, Type: protocol.TypeSelfLogout, RID: "24"})
c.publishUp("ep9", logout)
m2 := c.readDownJSON(3 * time.Second)
if m2["ok"] != true {
t.Fatalf("logout resp=%v", m2)
}
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
h, _ := e.login.SessionHashOf(context.Background(), "ep9")
if h == nil {
break
}
time.Sleep(20 * time.Millisecond)
}
h, _ := e.login.SessionHashOf(context.Background(), "ep9")
if h != nil {
t.Fatal("session should be cleared")
}
c2 := e.dial()
defer c2.close()
if _, ok := c2.connect("ep9", tok, 0); ok {
t.Fatal("token after logout should fail")
}
}
func TestF02AdminResetPasswordFatal(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep10", "password1")
c := e.dial()
defer c.close()
if _, ok := c.connect("ep10", "password1", 0); !ok {
t.Fatal("connect")
}
c.subscribe("ep10")
c.publishUp("ep10", helloPayload("1"))
m := c.readDownJSON(3 * time.Second)
data, _ := m["data"].(map[string]any)
tok, _ := data["session_token"].(string)
waitHandshook(t, e.b, "ep10")
if err := e.sess.ResetPassword(context.Background(), "ep10"); err != nil {
t.Fatal(err)
}
fatal := c.readDownJSON(3 * time.Second)
if fatal["type"] != "fatal" || fatal["reason"] != "password_reset" {
t.Fatalf("fatal=%v", fatal)
}
c2 := e.dial()
defer c2.close()
if _, ok := c2.connect("ep10", tok, 0); ok {
t.Fatal("token after reset should fail")
}
}
func TestF02KickKeepsToken(t *testing.T) {
e := openEnv(t, 30)
e.insertEndpoint("ep11", "password1")
c := e.dial()
defer c.close()
if _, ok := c.connect("ep11", "password1", 0); !ok {
t.Fatal("connect")
}
c.subscribe("ep11")
c.publishUp("ep11", helloPayload("1"))
m := c.readDownJSON(3 * time.Second)
data, _ := m["data"].(map[string]any)
tok, _ := data["session_token"].(string)
waitHandshook(t, e.b, "ep11")
go func() {
buf := make([]byte, 512)
for {
_ = c.conn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err := c.conn.Read(buf)
if err != nil {
return
}
}
}()
if err := e.sess.Kick(context.Background(), "ep11"); err != nil {
t.Fatal(err)
}
time.Sleep(100 * time.Millisecond)
c2 := e.dial()
defer c2.close()
if _, ok := c2.connect("ep11", tok, 0); !ok {
t.Fatal("token should still work after kick")
}
}
func itoa(n int) string {
if n == 0 {
return "0"
}
var b [16]byte
i := len(b)
for n > 0 {
i--
b[i] = byte('0' + n%10)
n /= 10
}
return string(b[i:])
}
+43 -1
View File
@@ -26,13 +26,15 @@ func (h *nixHook) Provides(b byte) bool {
mqtt.OnSessionEstablished, mqtt.OnSessionEstablished,
mqtt.OnDisconnect, mqtt.OnDisconnect,
mqtt.OnQosComplete, mqtt.OnQosComplete,
mqtt.OnSubscribed,
}, []byte{b}) }, []byte{b})
} }
func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
endpointID := string(pk.Connect.Username) endpointID := string(pk.Connect.Username)
clientID := pk.Connect.ClientIdentifier
if endpointID == "" { if endpointID == "" {
endpointID = pk.Connect.ClientIdentifier endpointID = clientID
} }
remoteIP := remoteIPOf(cl) remoteIP := remoteIPOf(cl)
@@ -45,6 +47,13 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
maxPacketSize: pk.Properties.MaximumPacketSize, maxPacketSize: pk.Properties.MaximumPacketSize,
} }
// ClientID、Username 都必须等于端编号
if clientID == "" || endpointID == "" || clientID != endpointID {
st.authOK = false
h.rememberPending(cl, st)
return nil
}
// 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive // 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive
ka := pk.Connect.Keepalive ka := pk.Connect.Keepalive
if ka < keepaliveMin || ka > keepaliveMax { if ka < keepaliveMin || ka > keepaliveMax {
@@ -103,6 +112,10 @@ func (h *nixHook) OnACLCheck(cl *mqtt.Client, topic string, write bool) bool {
} }
func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, error) { func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, error) {
// InlineClient 的 PublishDown 走 InjectPacket → OnPublish;必须放行才能分发给订阅者。
if cl != nil && cl.Net.Inline {
return pk, nil
}
h.b.connsMu.RLock() h.b.connsMu.RLock()
st := h.b.byClient[cl] st := h.b.byClient[cl]
h.b.connsMu.RUnlock() h.b.connsMu.RUnlock()
@@ -122,6 +135,28 @@ func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet,
return pk, packets.CodeSuccessIgnore return pk, packets.CodeSuccessIgnore
} }
func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes []byte) {
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if st == nil {
return
}
down := downTopic(st.endpointID)
for i, sub := range pk.Filters {
if sub.Filter != down {
continue
}
if i < len(reasonCodes) && reasonCodes[i] >= 0x80 {
continue
}
st.mu.Lock()
st.subscribedDown = true
st.mu.Unlock()
return
}
}
func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) {
h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload))
} }
@@ -151,14 +186,17 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
h.b.connsMu.Lock() h.b.connsMu.Lock()
st := h.b.byClient[cl] st := h.b.byClient[cl]
delete(h.b.byClient, cl) delete(h.b.byClient, cl)
isCurrent := false
if st != nil && h.b.current[st.endpointID] == st { if st != nil && h.b.current[st.endpointID] == st {
delete(h.b.current, st.endpointID) delete(h.b.current, st.endpointID)
isCurrent = true
} }
h.b.connsMu.Unlock() h.b.connsMu.Unlock()
if st == nil { if st == nil {
return return
} }
h.b.releaseAllLarge(st) h.b.releaseAllLarge(st)
h.b.cancelHandshakeDeadline(st.endpointID, st.connID)
reason := port.DisconnectNormal reason := port.DisconnectNormal
if err != nil { if err != nil {
@@ -179,6 +217,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
SessionToken: st.sessionToken, SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize, MaxPacketSize: st.maxPacketSize,
} }
if sess, ok := h.b.uplink.(*Session); ok {
sess.HandleDisconnect(context.Background(), info, reason, isCurrent)
return
}
h.b.uplink.OnDisconnect(context.Background(), info, reason) h.b.uplink.OnDisconnect(context.Background(), info, reason)
} }
+363
View File
@@ -0,0 +1,363 @@
package broker
import (
"context"
"encoding/json"
"log/slog"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
const handshakeTimeout = 30 * time.Second
// PresenceSink 供身份线订阅上下线(与 presence.Service 的 SetOnline/SetOffline 对齐)。
type PresenceSink interface {
SetOnline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error
SetOffline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error
}
// HelloLimits 握手响应里的服务器限制。
type HelloLimits struct {
MaxBodyBytes int
MaxMetaBytes int
MaxFrameBytes int
MaxTTLSeconds int64
MaxScheduleSeconds int64
AckTimeoutSeconds int64
ServerVersion string
}
// Session 处理握手、logout、上下线落库,并转发其余上行给 Inner。
type Session struct {
b *Broker
login *Login
inner port.UplinkHandler
presence PresenceSink
limits HelloLimits
log *slog.Logger
now func() time.Time
}
// SessionOptions 装配 Session。
type SessionOptions struct {
Login *Login
Inner port.UplinkHandler
Presence PresenceSink
Limits HelloLimits
Logger *slog.Logger
Now func() time.Time
}
// NewSession 创建会话层;调用 Attach 绑定 Broker 后再接连接。
func NewSession(opts SessionOptions) *Session {
inner := opts.Inner
if inner == nil {
inner = port.StubUplinkHandler{}
}
log := opts.Logger
if log == nil {
log = slog.Default()
}
now := opts.Now
if now == nil {
now = time.Now
}
lim := opts.Limits
if lim.ServerVersion == "" {
lim.ServerVersion = "0.1.0"
}
if lim.MaxBodyBytes == 0 {
lim.MaxBodyBytes = protocol.DefaultMaxBodyBytes
}
if lim.MaxMetaBytes == 0 {
lim.MaxMetaBytes = protocol.DefaultMaxMetaBytes
}
if lim.MaxFrameBytes == 0 {
lim.MaxFrameBytes = protocol.DefaultMaxFrameBytes
}
if lim.MaxTTLSeconds == 0 {
lim.MaxTTLSeconds = 2592000
}
if lim.MaxScheduleSeconds == 0 {
lim.MaxScheduleSeconds = 31536000
}
if lim.AckTimeoutSeconds == 0 {
lim.AckTimeoutSeconds = 300
}
return &Session{
login: opts.Login,
inner: inner,
presence: opts.Presence,
limits: lim,
log: log,
now: now,
}
}
// Attach 绑定 Broker(PublishDown / Disconnect / 连接表)。
func (s *Session) Attach(b *Broker) {
s.b = b
}
func (s *Session) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error {
if s.b != nil {
s.b.startHandshakeDeadline(conn.EndpointID, conn.ConnID, handshakeTimeout)
}
return s.inner.OnSessionEstablished(ctx, conn)
}
func (s *Session) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error {
return s.inner.OnHandshakeComplete(ctx, hs)
}
func (s *Session) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) {
// 正常路径由 hooks 调 HandleDisconnect(带 isCurrent)。
// 此方法满足 UplinkHandler;直接调用时按非当前处理,避免误标离线。
s.HandleDisconnect(ctx, conn, reason, false)
}
// HandleDisconnect 由 hooks 在确知 isCurrent 后调用(含落库与 presence)。
func (s *Session) HandleDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason, isCurrent bool) {
if s.b != nil {
s.b.cancelHandshakeDeadline(conn.EndpointID, conn.ConnID)
}
if isCurrent && s.login != nil {
atMs := s.now().UnixMilli()
if err := s.login.SetOfflineSince(ctx, conn.EndpointID, atMs); err != nil {
s.log.Error("set offline_since", "endpoint", conn.EndpointID, "err", err)
}
if s.presence != nil {
if err := s.presence.SetOffline(ctx, conn.EndpointID, conn.ConnID, atMs); err != nil {
s.log.Error("presence offline", "endpoint", conn.EndpointID, "err", err)
}
}
}
s.inner.OnDisconnect(ctx, conn, reason)
}
func (s *Session) HandleUplink(ctx context.Context, conn port.ConnInfo, payload []byte) error {
if s.b == nil {
return nil
}
st := s.b.connStateOf(conn.EndpointID, conn.ConnID)
if st == nil {
return nil
}
frame, err := protocol.Decode(payload)
if err != nil {
s.replyErr(ctx, conn, peekRID(payload), protocol.CodeBadRequest, err.Error())
return nil
}
st.mu.Lock()
ready := st.handshook
st.mu.Unlock()
switch f := frame.(type) {
case *protocol.Hello:
return s.handleHello(ctx, conn, st, f)
case *protocol.SelfLogout:
if !ready {
s.replyErr(ctx, conn, f.RID, protocol.CodeNotReady, "handshake required")
return nil
}
return s.handleLogout(ctx, conn, f)
default:
if !ready {
rid := peekRID(payload)
s.replyErr(ctx, conn, rid, protocol.CodeNotReady, "handshake required")
return nil
}
return s.inner.HandleUplink(ctx, conn, payload)
}
}
func (s *Session) handleHello(ctx context.Context, conn port.ConnInfo, st *connState, hello *protocol.Hello) error {
if err := hello.Validate(); err != nil {
code := protocol.CodeBadRequest
if pe, ok := err.(*protocol.Error); ok {
code = pe.Code
}
s.replyErr(ctx, conn, hello.RID, code, err.Error())
return nil
}
st.mu.Lock()
if st.handshook {
st.mu.Unlock()
s.replyErr(ctx, conn, hello.RID, protocol.CodeBadRequest, "already handshook")
return nil
}
st.mu.Unlock()
if !s.b.hasDownSub(st) {
go func() {
_ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectIdle)
}()
return nil
}
maxRecv := 0
if hello.MaxReceiveBytes != nil {
maxRecv = *hello.MaxReceiveBytes
}
s.b.SetMaxReceiveBytes(conn.EndpointID, conn.ConnID, maxRecv)
data := protocol.HelloData{
ServerTimeMs: s.now().UnixMilli(),
ServerVersion: s.limits.ServerVersion,
MaxBodyBytes: s.limits.MaxBodyBytes,
MaxMetaBytes: s.limits.MaxMetaBytes,
MaxFrameBytes: s.limits.MaxFrameBytes,
MaxTTLSeconds: s.limits.MaxTTLSeconds,
MaxScheduleSeconds: s.limits.MaxScheduleSeconds,
AckTimeoutSeconds: s.limits.AckTimeoutSeconds,
}
if conn.SessionToken != "" {
data.SessionToken = conn.SessionToken
}
raw, err := protocol.Marshal(data)
if err != nil {
return err
}
resp := protocol.Resp{
V: protocol.Version,
Type: protocol.TypeResp,
RID: hello.RID,
OK: true,
Data: raw,
}
if err := s.publishJSON(ctx, conn, resp, 1); err != nil {
return err
}
atMs := s.now().UnixMilli()
if s.login != nil {
if err := s.login.SetOnlineSince(ctx, conn.EndpointID, atMs); err != nil {
s.log.Error("set online_since", "endpoint", conn.EndpointID, "err", err)
}
}
if s.presence != nil {
if err := s.presence.SetOnline(ctx, conn.EndpointID, conn.ConnID, atMs); err != nil {
s.log.Error("presence online", "endpoint", conn.EndpointID, "err", err)
}
}
st.mu.Lock()
st.handshook = true
st.mu.Unlock()
s.b.cancelHandshakeDeadline(conn.EndpointID, conn.ConnID)
hs := port.HandshakeInfo{
ConnInfo: conn,
MaxReceiveBytes: maxRecv,
Client: hello.Client,
}
return s.inner.OnHandshakeComplete(ctx, hs)
}
func (s *Session) handleLogout(ctx context.Context, conn port.ConnInfo, req *protocol.SelfLogout) error {
if err := req.Validate(); err != nil {
code := protocol.CodeBadRequest
if pe, ok := err.(*protocol.Error); ok {
code = pe.Code
}
s.replyErr(ctx, conn, req.RID, code, err.Error())
return nil
}
if s.login != nil {
if err := s.login.ClearSession(ctx, conn.EndpointID); err != nil {
s.replyErr(ctx, conn, req.RID, protocol.CodeBusy, "clear session failed")
return nil
}
}
resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: req.RID, OK: true}
if err := s.publishJSON(ctx, conn, resp, 1); err != nil {
s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", err)
}
go func() {
// 稍等让 QoS1 resp 写入连接,再断开
time.Sleep(50 * time.Millisecond)
_ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal)
}()
return nil
}
// Kick 只断开当前连接,令牌不变。
func (s *Session) Kick(ctx context.Context, endpointID string) error {
if s.b == nil {
return ErrNoConnection
}
return s.b.Disconnect(ctx, endpointID, "", port.DisconnectKicked)
}
// Disable 清空令牌,发 fatal(disabled) 后断开。
func (s *Session) Disable(ctx context.Context, endpointID string) error {
return s.fatalKick(ctx, endpointID, "disabled")
}
// Deleted 清空令牌,发 fatal(deleted) 后断开。
func (s *Session) Deleted(ctx context.Context, endpointID string) error {
return s.fatalKick(ctx, endpointID, "deleted")
}
// ResetPassword 清空令牌,发 fatal(password_reset) 后断开。
func (s *Session) ResetPassword(ctx context.Context, endpointID string) error {
return s.fatalKick(ctx, endpointID, "password_reset")
}
func (s *Session) fatalKick(ctx context.Context, endpointID, reason string) error {
if s.login != nil {
if err := s.login.ClearSession(ctx, endpointID); err != nil {
return err
}
}
if s.b == nil {
return nil
}
info, ok := s.b.ConnInfoOf(endpointID)
if !ok {
return nil
}
fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason}
_ = s.publishJSON(ctx, info, fatal, 1)
go func() {
time.Sleep(20 * time.Millisecond)
_ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal)
}()
return nil
}
func (s *Session) replyErr(ctx context.Context, conn port.ConnInfo, rid, code, message string) {
if rid == "" {
rid = "0"
}
resp := protocol.Resp{
V: protocol.Version,
Type: protocol.TypeResp,
RID: rid,
OK: false,
Error: &protocol.ErrorBody{Code: code, Message: message},
}
_ = s.publishJSON(ctx, conn, resp, 1)
}
func (s *Session) publishJSON(ctx context.Context, conn port.ConnInfo, v any, qos byte) error {
b, err := protocol.Marshal(v)
if err != nil {
return err
}
return s.b.PublishDown(ctx, conn.EndpointID, conn.ConnID, b, port.PublishOpts{QoS: qos})
}
func peekRID(payload []byte) string {
var peek struct {
RID string `json:"rid"`
}
_ = json.Unmarshal(payload, &peek)
return peek.RID
}
var _ port.UplinkHandler = (*Session)(nil)