Compare commits
3
Commits
ba6cc9f6e0
...
a78ab0d547
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a78ab0d547 | ||
|
|
bdd1d9e9f4 | ||
|
|
34e7c2827f |
+157
-12
@@ -276,37 +276,117 @@
|
||||
- 备选方案:N2 暴露回调给 M 注册。
|
||||
- 影响:接线后 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
|
||||
|
||||
### M1 2026-09-30
|
||||
|
||||
1. **提交时分发做成最小正确版**
|
||||
- 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发(停用拒绝、`queue_full`、`expire_at`/宽限、无接收者 `completed`、回执等)。
|
||||
- 实际做法:单聊只插一条 `pending`;群按当时 `group_members` 去掉发送者各插 `pending`;消息改为 `dispatched`。不设 `expire_at`,不检查接收端配额/在线/停用,不因无接收者改为 `completed`,不写回执,不唤醒推送循环。
|
||||
- 原因:M1 范围是提交;完整分发与推送属 M2。
|
||||
- 备选方案:M1 直接实现完整 7.4(抢 M2)。
|
||||
- 影响:到点消息已有投递行,但停用成员仍会有 `pending`;无成员群仍为 `dispatched` 且无投递;推送需等 M2。
|
||||
1. **提交时分发做成最小正确版**(已被 M2 取代)
|
||||
- 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发。
|
||||
- 实际做法(M1):单聊/群插 `pending` 后改 `dispatched`,不做停用/配额/宽限。
|
||||
- 现状(M2):`Submit` 到点与 `DispatchDue` 均走完整 `dispatchFullTx`(7.4)。
|
||||
- 原因 / 备选 / 影响:见 M2。
|
||||
|
||||
2. **请求频率突发容量写死为 100**
|
||||
- 原条款: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 字段。
|
||||
- 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。
|
||||
- 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。
|
||||
|
||||
3. **未接线 `cmd/nixmsg`**
|
||||
- 原条款:可替换 T0.4 假实现。
|
||||
- 实际做法:新增 `message.App` 实现 `Submit`;保留 `Stub`;按任务隔离要求未改 `cmd/nixmsg`/`wire.go`。
|
||||
- 原因:本任务禁止改 `cmd/nixmsg`;总控接线或后续任务再换。
|
||||
- 实际做法:`message.App` 实现 Service;保留 `Stub`;未改 `cmd/nixmsg`/`wire.go`。
|
||||
- 原因:本任务隔离;总控接线。
|
||||
- 备选方案:本任务直接改 `wire.go`。
|
||||
- 影响:进程内仍用 Stub,需显式构造 `message.New` 才能用真实提交。
|
||||
- 影响:进程内仍用 Stub,需显式 `message.New` 并注入 `Downlink`/`ConnRegistry`。
|
||||
|
||||
4. **防重键在、消息行已删时返回 `not_found`**
|
||||
- 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为(状态查询为 `not_found`)。
|
||||
- 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为。
|
||||
- 实际做法:`send_keys` 指纹相同但 `messages` 无行时返回 `not_found`。
|
||||
- 原因:无法构造 `send_at`/`state`。
|
||||
- 备选方案:在 `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
|
||||
|
||||
@@ -347,6 +427,71 @@
|
||||
- 备选方案:16/32 位。
|
||||
- 影响:无产品行为差异。
|
||||
|
||||
### I2 / I3 / I4 2026-09-30
|
||||
|
||||
1. **服务方法可直接调用,未接 MQTT 分发**
|
||||
- 原条款:端协议帧经 broker 上行分发到 app。
|
||||
- 实际做法:`identity.App` / `presence.App` / `group.App` 实现 Service 方法;单测直接调用,不经 mochi。`cmd/nixmsg`/`wire.go` 未改。
|
||||
- 原因:任务要求可被协议分发调用的服务方法 + 直接调用验证;N3 连接事件与总控接线另波。
|
||||
- 备选方案:本分支顺带改 wire(与隔离冲突)。
|
||||
- 影响:合入后需总控/N 接线 HandleUplink → 各 Service;MQTT 帧路径未测。
|
||||
|
||||
2. **在线状态:库字段 + 可注入连接表 + 本包握手表**
|
||||
- 原条款:在线以真实连接为准;依赖 N3 连接事件。
|
||||
- 实际做法:`presence.ConnTable` 可注入;另用 `SetOnline`/`SetOffline` 维护内存表并写 `endpoints.online_since`/`offline_since`;查询优先 ConnTable,其次内存表,再回退库字段(`online_since` 晚于 `offline_since` 或后者为空)。
|
||||
- 原因:N3 可能尚未合入,不阻塞 I3。
|
||||
- 备选方案:阻塞等 N3。
|
||||
- 影响:未接线时须调用 SetOnline/SetOffline 或写库字段;拔网线心跳超时属 N 线,本波单测不覆盖。
|
||||
|
||||
3. **presence / group_event 经 Downlink QoS 0,可 nil**
|
||||
- 原条款:上下线与群事件尽力推送、不落库。
|
||||
- 实际做法:注入 `port.Downlink` 时编码帧并 `PublishDown`(presence/group_event QoS 0;退群已推送投递的 `revoked` 用 QoS 1);Downlink 为 nil 时跳过推送,业务库操作仍完成。
|
||||
- 原因:无 MQTT 时仍可测库逻辑。
|
||||
- 备选方案:强制假 Downlink。
|
||||
- 影响:接线后必须注入真实 Downlink 才有通知。
|
||||
|
||||
4. **self.logout 踢线可选**
|
||||
- 原条款:回 resp 后断开连接。
|
||||
- 实际做法:清 `session_hash`;若注入 `port.ConnControl` 则 `Disconnect`,否则仅清令牌。
|
||||
- 原因:未接 broker。
|
||||
- 备选方案:无。
|
||||
- 影响:接线方应注入 ConnControl。
|
||||
|
||||
5. **session_hash 存 SHA-256 十六进制**
|
||||
- 原条款:库中存会话令牌 SHA-256。
|
||||
- 实际做法:`hex.EncodeToString(hash)` 写入 TEXT 列。
|
||||
- 原因:文档未规定编码;十六进制便于调试与比对。
|
||||
- 备选方案:BLOB/Base64。
|
||||
- 影响:N3 校验须用同一编码。
|
||||
|
||||
6. **self.update 空 name 不写库**
|
||||
- 原条款:可更新 name。
|
||||
- 实际做法:`name` 非空才 UPDATE name;仅改 `default_delay_ms` 时不碰 name(JSON omitempty 无法区分省略与空串)。
|
||||
- 原因:避免误清空名称。
|
||||
- 备选方案:用指针字段区分。
|
||||
- 影响:端无法通过协议把名称改成空字符串(可用空格等)。
|
||||
|
||||
7. **进群密码校验不产生单聊授权**
|
||||
- 原条款:拉人须当次带密码;已有授权不能代替。
|
||||
- 实际做法:`CheckTalkPasswordForJoin` 只校验,不写 `talk_grants`。
|
||||
- 原因:与 F15「进群仍要密码」一致,避免进群副作用放宽单聊。
|
||||
- 备选方案:校验成功顺带写 password 授权。
|
||||
- 影响:仅进群成功后,单聊仍须 unlock/发送带密。
|
||||
|
||||
8. **群作废在 group 包内写 deliveries/messages**
|
||||
- 原条款:退群/踢人/解散的投递作废属 7.6,消息线亦相关。
|
||||
- 实际做法:I4 在 `group` 写操作里直接改 `pending→rejected`、`scheduled→completed/group_dissolved`,删正文行,需要时插消息级回执,已推送则经 Downlink 发 `revoked`。
|
||||
- 原因:I4 验收依赖作废规则;M 线完整推送循环可能未合入。
|
||||
- 备选方案:只调 message 钩子(接口尚未暴露)。
|
||||
- 影响:与后续 M 作废路径需保持同语义,避免重复作废。
|
||||
|
||||
9. **UnlockTalk / SelfChangeLoginPassword / CheckTalkPasswordForJoin 增加 remoteIP**
|
||||
- 原条款:锁定按发送方+对方 / 编号+IP。
|
||||
- 实际做法:Service 方法增加 `remoteIP` 参数供锁定计数;T0.4 Stub 同步改签名。
|
||||
- 原因:无 ConnInfo 的直接调用测试需要显式 IP。
|
||||
- 备选方案:塞进 context。
|
||||
- 影响:协议分发接线时从 `ConnInfo.RemoteIP` 传入。
|
||||
|
||||
## 后台接口 A
|
||||
|
||||
### A1 2026-09-30
|
||||
|
||||
@@ -0,0 +1,727 @@
|
||||
package group
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
eventMemberAdded = "member_added"
|
||||
eventMemberRemoved = "member_removed"
|
||||
eventLeft = "left"
|
||||
eventOwnerChanged = "owner_changed"
|
||||
eventRenamed = "renamed"
|
||||
eventDissolved = "dissolved"
|
||||
|
||||
reasonLeftGroup = "left_group"
|
||||
reasonGroupDissolved = "group_dissolved"
|
||||
|
||||
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||
)
|
||||
|
||||
// TalkGate checks talk password when adding members (implemented by identity).
|
||||
type TalkGate interface {
|
||||
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
|
||||
}
|
||||
|
||||
// OnlineLookup reports member online status.
|
||||
type OnlineLookup interface {
|
||||
IsOnline(endpointID string) bool
|
||||
}
|
||||
|
||||
// Config holds group service dependencies.
|
||||
type Config struct {
|
||||
DB *store.DB
|
||||
Talk TalkGate
|
||||
Online OnlineLookup
|
||||
Downlink port.Downlink
|
||||
MaxGroupMembers int
|
||||
Now func() time.Time
|
||||
DefaultRemoteIP string
|
||||
}
|
||||
|
||||
// App implements group.Service.
|
||||
type App struct {
|
||||
db *store.DB
|
||||
talk TalkGate
|
||||
online OnlineLookup
|
||||
down port.Downlink
|
||||
maxMem int
|
||||
nowFn func() time.Time
|
||||
remoteIP string
|
||||
}
|
||||
|
||||
// New constructs the group service.
|
||||
func New(cfg Config) *App {
|
||||
now := cfg.Now
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
max := cfg.MaxGroupMembers
|
||||
if max <= 0 {
|
||||
max = 1000
|
||||
}
|
||||
return &App{
|
||||
db: cfg.DB,
|
||||
talk: cfg.Talk,
|
||||
online: cfg.Online,
|
||||
down: cfg.Downlink,
|
||||
maxMem: max,
|
||||
nowFn: now,
|
||||
remoteIP: cfg.DefaultRemoteIP,
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) nowMs() int64 { return a.nowFn().UnixMilli() }
|
||||
|
||||
func errCode(code, msg string) *protocol.Error {
|
||||
return &protocol.Error{Code: code, Message: msg}
|
||||
}
|
||||
|
||||
func protoCode(err error) string {
|
||||
var pe *protocol.Error
|
||||
if errors.As(err, &pe) {
|
||||
return pe.Code
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// Create creates a group; creator becomes owner and member. Partial member failures still create the group.
|
||||
func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCreate) (CreateResult, error) {
|
||||
if req == nil {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
if !protocol.ValidEndpointID(actorID) {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid actor")
|
||||
}
|
||||
gid := req.ID
|
||||
if gid == "" {
|
||||
var genErr error
|
||||
gid, genErr = generateGroupID()
|
||||
if genErr != nil {
|
||||
return CreateResult{}, genErr
|
||||
}
|
||||
}
|
||||
now := a.nowMs()
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0, len(req.Members))
|
||||
|
||||
for _, m := range req.Members {
|
||||
if m.ID == actorID {
|
||||
continue
|
||||
}
|
||||
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
|
||||
if 1+len(added) > a.maxMem {
|
||||
return CreateResult{}, errCode(protocol.CodeGroupFull, "group full")
|
||||
}
|
||||
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var exists int
|
||||
qErr := tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, gid).Scan(&exists)
|
||||
if qErr == nil {
|
||||
return errCode(protocol.CodeIDTaken, "group id taken")
|
||||
}
|
||||
if !errors.Is(qErr, sql.ErrNoRows) {
|
||||
return qErr
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
gid, req.Name, actorID, now); e != nil {
|
||||
if isUnique(e) {
|
||||
return errCode(protocol.CodeIDTaken, "group id taken")
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, actorID, now); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
|
||||
for _, id := range added {
|
||||
a.emit(ctx, append([]string{actorID}, added...), gid, eventMemberAdded, id, now)
|
||||
}
|
||||
return CreateResult{ID: gid, Name: req.Name, OwnerID: actorID, Failed: failed}, nil
|
||||
}
|
||||
|
||||
// Add adds members (owner only).
|
||||
func (a *App) Add(ctx context.Context, actorID string, req *protocol.GroupAdd) (AddResult, error) {
|
||||
if req == nil {
|
||||
return AddResult{}, errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
if owner != actorID {
|
||||
return AddResult{}, errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
now := a.nowMs()
|
||||
|
||||
for _, m := range req.Members {
|
||||
if contains(members, m.ID) {
|
||||
continue
|
||||
}
|
||||
if len(members)+len(added) >= a.maxMem {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
|
||||
if len(added) > 0 {
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT OR IGNORE INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
req.GroupID, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
all := append(append([]string{}, members...), added...)
|
||||
for _, id := range added {
|
||||
a.emit(ctx, all, req.GroupID, eventMemberAdded, id, now)
|
||||
}
|
||||
}
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
|
||||
// Remove kicks a member (owner only).
|
||||
func (a *App) Remove(ctx context.Context, actorID string, req *protocol.GroupRemove) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
if req.EndpointID == owner {
|
||||
return errCode(protocol.CodeBadRequest, "cannot remove owner")
|
||||
}
|
||||
if !contains(members, req.EndpointID) {
|
||||
return errCode(protocol.CodeNotFound, "member not found")
|
||||
}
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`,
|
||||
req.GroupID, req.EndpointID); e != nil {
|
||||
return e
|
||||
}
|
||||
return voidMemberDeliveriesTx(tx, req.GroupID, req.EndpointID, reasonLeftGroup, now, &revokes)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.publishRevokes(ctx, revokes)
|
||||
left := without(members, req.EndpointID)
|
||||
notify := append(left, req.EndpointID)
|
||||
a.emit(ctx, notify, req.GroupID, eventMemberRemoved, req.EndpointID, now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Leave lets a non-owner member leave.
|
||||
func (a *App) Leave(ctx context.Context, actorID string, req *protocol.GroupLeave) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !contains(members, actorID) {
|
||||
return errCode(protocol.CodeNotMember, "not a member")
|
||||
}
|
||||
if owner == actorID {
|
||||
return errCode(protocol.CodeOwnerCannotLeave, "owner cannot leave")
|
||||
}
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`,
|
||||
req.GroupID, actorID); e != nil {
|
||||
return e
|
||||
}
|
||||
return voidMemberDeliveriesTx(tx, req.GroupID, actorID, reasonLeftGroup, now, &revokes)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.publishRevokes(ctx, revokes)
|
||||
left := without(members, actorID)
|
||||
notify := append(left, actorID)
|
||||
a.emit(ctx, notify, req.GroupID, eventLeft, actorID, now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Transfer transfers ownership.
|
||||
func (a *App) Transfer(ctx context.Context, actorID string, req *protocol.GroupTransfer) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
if !contains(members, req.EndpointID) {
|
||||
return errCode(protocol.CodeNotFound, "member not found")
|
||||
}
|
||||
now := a.nowMs()
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE groups SET owner_id = ? WHERE id = ?`, req.EndpointID, req.GroupID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.emit(ctx, members, req.GroupID, eventOwnerChanged, req.EndpointID, now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Rename renames the group.
|
||||
func (a *App) Rename(ctx context.Context, actorID string, req *protocol.GroupRename) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
now := a.nowMs()
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE groups SET name = ? WHERE id = ?`, req.Name, req.GroupID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.emit(ctx, members, req.GroupID, eventRenamed, "", now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Dissolve dissolves the group and voids unfinished deliveries / scheduled messages.
|
||||
func (a *App) Dissolve(ctx context.Context, actorID string, req *protocol.GroupDissolve) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if e := voidGroupAllTx(tx, req.GroupID, now, &revokes); e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ?`, req.GroupID); e != nil {
|
||||
return e
|
||||
}
|
||||
_, e := tx.Exec(`DELETE FROM groups WHERE id = ?`, req.GroupID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.publishRevokes(ctx, revokes)
|
||||
a.emit(ctx, members, req.GroupID, eventDissolved, "", now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// List lists groups the actor belongs to.
|
||||
func (a *App) List(ctx context.Context, actorID string, req *protocol.GroupList) ([]ListItem, string, error) {
|
||||
if req == nil {
|
||||
req = &protocol.GroupList{V: protocol.Version, Type: protocol.TypeGroupList, RID: "x", Limit: 100}
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > protocol.MaxPageLimit {
|
||||
limit = protocol.MaxPageLimit
|
||||
}
|
||||
rows, err := a.db.Read.QueryContext(ctx, `
|
||||
SELECT g.id, g.name, g.owner_id,
|
||||
(SELECT COUNT(*) FROM group_members gm2 WHERE gm2.group_id = g.id) AS cnt
|
||||
FROM groups g
|
||||
JOIN group_members gm ON gm.group_id = g.id AND gm.endpoint_id = ?
|
||||
WHERE g.id > ?
|
||||
ORDER BY g.id ASC
|
||||
LIMIT ?`, actorID, req.Cursor, limit+1)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
items := make([]ListItem, 0, limit)
|
||||
for rows.Next() {
|
||||
var it ListItem
|
||||
if scanErr := rows.Scan(&it.ID, &it.Name, &it.OwnerID, &it.MemberCount); scanErr != nil {
|
||||
return nil, "", scanErr
|
||||
}
|
||||
items = append(items, it)
|
||||
}
|
||||
next := ""
|
||||
if len(items) > limit {
|
||||
items = items[:limit]
|
||||
next = items[len(items)-1].ID
|
||||
}
|
||||
return items, next, rows.Err()
|
||||
}
|
||||
|
||||
// Get returns group details; caller must be a member.
|
||||
func (a *App) Get(ctx context.Context, actorID string, req *protocol.GroupGet) (GetResult, error) {
|
||||
if req == nil {
|
||||
return GetResult{}, errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
var name, owner string
|
||||
err := a.db.Read.QueryRowContext(ctx, `SELECT name, owner_id FROM groups WHERE id = ?`, req.GroupID).
|
||||
Scan(&name, &owner)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return GetResult{}, errCode(protocol.CodeNotFound, "group not found")
|
||||
}
|
||||
if err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
var one int
|
||||
err = a.db.Read.QueryRowContext(ctx,
|
||||
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, req.GroupID, actorID).Scan(&one)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return GetResult{}, errCode(protocol.CodeNotMember, "not a member")
|
||||
}
|
||||
if err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > protocol.MaxPageLimit {
|
||||
limit = protocol.MaxPageLimit
|
||||
}
|
||||
rows, err := a.db.Read.QueryContext(ctx, `
|
||||
SELECT gm.endpoint_id, e.name
|
||||
FROM group_members gm
|
||||
JOIN endpoints e ON e.id = gm.endpoint_id
|
||||
WHERE gm.group_id = ? AND gm.endpoint_id > ?
|
||||
ORDER BY gm.endpoint_id ASC
|
||||
LIMIT ?`, req.GroupID, req.Cursor, limit+1)
|
||||
if err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
members := make([]MemberItem, 0, limit)
|
||||
for rows.Next() {
|
||||
var m MemberItem
|
||||
if scanErr := rows.Scan(&m.ID, &m.Name); scanErr != nil {
|
||||
return GetResult{}, scanErr
|
||||
}
|
||||
if a.online != nil {
|
||||
m.Online = a.online.IsOnline(m.ID)
|
||||
}
|
||||
members = append(members, m)
|
||||
}
|
||||
next := ""
|
||||
if len(members) > limit {
|
||||
members = members[:limit]
|
||||
next = members[len(members)-1].ID
|
||||
}
|
||||
return GetResult{ID: req.GroupID, Name: name, OwnerID: owner, Members: members, NextCursor: next}, rows.Err()
|
||||
}
|
||||
|
||||
// AdminCreate creates a group without talk-password checks.
|
||||
func (a *App) AdminCreate(ctx context.Context, name, ownerID string, memberIDs []string) (CreateResult, error) {
|
||||
members := make([]protocol.GroupMemberIn, 0, len(memberIDs))
|
||||
for _, id := range memberIDs {
|
||||
members = append(members, protocol.GroupMemberIn{ID: id})
|
||||
}
|
||||
return a.createAdmin(ctx, ownerID, name, "", members)
|
||||
}
|
||||
|
||||
// AdminAddMembers adds members without talk-password checks.
|
||||
func (a *App) AdminAddMembers(ctx context.Context, groupID string, memberIDs []string) (AddResult, error) {
|
||||
_, members, err := a.loadGroup(ctx, groupID)
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
now := a.nowMs()
|
||||
for _, id := range memberIDs {
|
||||
if contains(members, id) {
|
||||
continue
|
||||
}
|
||||
var enabled int
|
||||
e := a.db.Read.QueryRowContext(ctx, `SELECT enabled FROM endpoints WHERE id = ?`, id).Scan(&enabled)
|
||||
if errors.Is(e, sql.ErrNoRows) {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeInvalidTarget})
|
||||
continue
|
||||
}
|
||||
if e != nil {
|
||||
return AddResult{}, e
|
||||
}
|
||||
if enabled == 0 {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeEndpointDisabled})
|
||||
continue
|
||||
}
|
||||
if len(members)+len(added) >= a.maxMem {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
added = append(added, id)
|
||||
}
|
||||
if len(added) == 0 {
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT OR IGNORE INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
groupID, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
all := append(append([]string{}, members...), added...)
|
||||
for _, id := range added {
|
||||
a.emit(ctx, all, groupID, eventMemberAdded, id, now)
|
||||
}
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
|
||||
func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, members []protocol.GroupMemberIn) (CreateResult, error) {
|
||||
if gid == "" {
|
||||
var genErr error
|
||||
gid, genErr = generateGroupID()
|
||||
if genErr != nil {
|
||||
return CreateResult{}, genErr
|
||||
}
|
||||
}
|
||||
if !protocol.ValidName(name) || name == "" {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid name")
|
||||
}
|
||||
now := a.nowMs()
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
for _, m := range members {
|
||||
if m.ID == ownerID {
|
||||
continue
|
||||
}
|
||||
var enabled int
|
||||
e := a.db.Read.QueryRowContext(ctx, `SELECT enabled FROM endpoints WHERE id = ?`, m.ID).Scan(&enabled)
|
||||
if errors.Is(e, sql.ErrNoRows) {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeInvalidTarget})
|
||||
continue
|
||||
}
|
||||
if e != nil {
|
||||
return CreateResult{}, e
|
||||
}
|
||||
if enabled == 0 {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeEndpointDisabled})
|
||||
continue
|
||||
}
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
if 1+len(added) > a.maxMem {
|
||||
return CreateResult{}, errCode(protocol.CodeGroupFull, "group full")
|
||||
}
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
gid, name, ownerID, now); e != nil {
|
||||
if isUnique(e) {
|
||||
return errCode(protocol.CodeIDTaken, "group id taken")
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, ownerID, now); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
return CreateResult{ID: gid, Name: name, OwnerID: ownerID, Failed: failed}, nil
|
||||
}
|
||||
|
||||
func (a *App) checkAddMember(ctx context.Context, actorID, targetID, talkPassword string) error {
|
||||
if a.talk == nil {
|
||||
return errCode(protocol.CodeBusy, "talk gate not configured")
|
||||
}
|
||||
return a.talk.CheckTalkPasswordForJoin(ctx, actorID, targetID, talkPassword, a.remoteIP)
|
||||
}
|
||||
|
||||
func (a *App) loadGroup(ctx context.Context, groupID string) (owner string, members []string, err error) {
|
||||
err = a.db.Read.QueryRowContext(ctx, `SELECT owner_id FROM groups WHERE id = ?`, groupID).Scan(&owner)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return "", nil, errCode(protocol.CodeNotFound, "group not found")
|
||||
}
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
rows, qErr := a.db.Read.QueryContext(ctx, `SELECT endpoint_id FROM group_members WHERE group_id = ?`, groupID)
|
||||
if qErr != nil {
|
||||
return "", nil, qErr
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var id string
|
||||
if scanErr := rows.Scan(&id); scanErr != nil {
|
||||
return "", nil, scanErr
|
||||
}
|
||||
members = append(members, id)
|
||||
}
|
||||
return owner, members, rows.Err()
|
||||
}
|
||||
|
||||
func (a *App) emit(ctx context.Context, recipients []string, groupID, event, endpointID string, atMs int64) {
|
||||
if a.down == nil {
|
||||
return
|
||||
}
|
||||
frame := protocol.GroupEvent{
|
||||
V: protocol.Version, Type: protocol.TypeGroupEvent,
|
||||
GroupID: groupID, Event: event, EndpointID: endpointID, AtMs: atMs,
|
||||
}
|
||||
payload, encErr := encodeFrame(frame)
|
||||
if encErr != nil {
|
||||
return
|
||||
}
|
||||
seen := map[string]struct{}{}
|
||||
for _, id := range recipients {
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
_ = a.down.PublishDown(ctx, id, "", payload, port.PublishOpts{QoS: 0})
|
||||
}
|
||||
}
|
||||
|
||||
func encodeFrame(v any) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
if err := protocol.Encode(&buf, v); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func generateGroupID() (string, error) {
|
||||
b := make([]byte, 8)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, 8)
|
||||
for i := range b {
|
||||
out[i] = idAlphabet[int(b[i])%len(idAlphabet)]
|
||||
}
|
||||
return "g_" + string(out), nil
|
||||
}
|
||||
|
||||
func failCode(err error) string {
|
||||
if c := protoCode(err); c != "" {
|
||||
return c
|
||||
}
|
||||
return protocol.CodeBusy
|
||||
}
|
||||
|
||||
func contains(ss []string, x string) bool {
|
||||
for _, s := range ss {
|
||||
if s == x {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func without(ss []string, x string) []string {
|
||||
out := make([]string, 0, len(ss))
|
||||
for _, s := range ss {
|
||||
if s != x {
|
||||
out = append(out, s)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
var _ Service = (*App)(nil)
|
||||
@@ -0,0 +1,412 @@
|
||||
package group_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/group"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/message"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
type memDown struct {
|
||||
mu sync.Mutex
|
||||
msgs []struct {
|
||||
to string
|
||||
qos byte
|
||||
raw []byte
|
||||
}
|
||||
}
|
||||
|
||||
func (d *memDown) PublishDown(_ context.Context, endpointID string, _ port.ConnID, payload []byte, opts port.PublishOpts) error {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
d.msgs = append(d.msgs, struct {
|
||||
to string
|
||||
qos byte
|
||||
raw []byte
|
||||
}{endpointID, opts.QoS, append([]byte(nil), payload...)})
|
||||
return nil
|
||||
}
|
||||
|
||||
func setup(t *testing.T) (*group.App, *identity.App, *message.App, *store.DB, *memDown) {
|
||||
t.Helper()
|
||||
db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
fixed := time.UnixMilli(1_700_000_000_000)
|
||||
locks := auth.NewLoginLocks()
|
||||
idApp := identity.New(identity.Config{
|
||||
DB: db, Hash: auth.NewStubHashPool(), Locks: locks,
|
||||
Sessions: auth.NewSessionTokens(),
|
||||
Now: func() time.Time { return fixed },
|
||||
})
|
||||
down := &memDown{}
|
||||
gApp := group.New(group.Config{
|
||||
DB: db, Talk: idApp, Downlink: down, MaxGroupMembers: 1000,
|
||||
Now: func() time.Time { return fixed }, DefaultRemoteIP: "1.1.1.1",
|
||||
})
|
||||
lim := message.LimitsFromConfig(config.Default().Limits)
|
||||
lim.RequestsPerSecond = 0
|
||||
msgApp := message.New(db, lim, auth.NewStubHashPool(),
|
||||
message.WithNow(func() time.Time { return fixed }),
|
||||
message.WithLocks(locks),
|
||||
)
|
||||
return gApp, idApp, msgApp, db, down
|
||||
}
|
||||
|
||||
func insertEP(t *testing.T, db *store.DB, id string, enabled int) {
|
||||
t.Helper()
|
||||
err := 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(?,?,?,?,0,0,?,?)`, id, id, "stub$login", nil, enabled, 1_700_000_000_000)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func protoCode(err error) string {
|
||||
var pe *protocol.Error
|
||||
if errors.As(err, &pe) {
|
||||
return pe.Code
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func TestF16CreateAddPasswordAndOwner(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, idApp, _, db, _ := setup(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
insertEP(t, db, "bob", 1)
|
||||
insertEP(t, db, "carol", 1)
|
||||
insertEP(t, db, "dave", 0) // disabled
|
||||
if err := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// 非群主加人:先建群
|
||||
created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1",
|
||||
ID: "g_test01", Name: "一组",
|
||||
Members: []protocol.GroupMemberIn{
|
||||
{ID: "bob", TalkPassword: "wrong"},
|
||||
{ID: "carol"},
|
||||
{ID: "dave"},
|
||||
{ID: "nobody"},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if created.OwnerID != "alice" || created.ID != "g_test01" {
|
||||
t.Fatalf("%+v", created)
|
||||
}
|
||||
// bob 密码错、dave 停用、nobody 无效 → failed;carol 加入
|
||||
codes := map[string]string{}
|
||||
for _, f := range created.Failed {
|
||||
codes[f.ID] = f.Code
|
||||
}
|
||||
if codes["bob"] != protocol.CodeTalkPasswordInvalid {
|
||||
t.Fatalf("bob fail %+v", created.Failed)
|
||||
}
|
||||
if codes["dave"] != protocol.CodeEndpointDisabled {
|
||||
t.Fatalf("dave fail %+v", created.Failed)
|
||||
}
|
||||
if codes["nobody"] != protocol.CodeInvalidTarget {
|
||||
t.Fatalf("nobody fail %+v", created.Failed)
|
||||
}
|
||||
var n int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=?`, "g_test01").Scan(&n)
|
||||
if n != 2 { // alice + carol
|
||||
t.Fatalf("members=%d", n)
|
||||
}
|
||||
var bobN int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=? AND endpoint_id=?`, "g_test01", "bob").Scan(&bobN)
|
||||
if bobN != 0 {
|
||||
t.Fatal("bob should not be member")
|
||||
}
|
||||
|
||||
// 非群主加人失败
|
||||
_, err = gApp.Add(ctx, "carol", &protocol.GroupAdd{
|
||||
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "2", GroupID: "g_test01",
|
||||
Members: []protocol.GroupMemberIn{{ID: "bob", TalkPassword: "secret"}},
|
||||
})
|
||||
if protoCode(err) != protocol.CodeForbidden {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
// 群主带对密码加人;已有单聊授权也不能省略
|
||||
if err = idApp.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
addRes, err := gApp.Add(ctx, "alice", &protocol.GroupAdd{
|
||||
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "3", GroupID: "g_test01",
|
||||
Members: []protocol.GroupMemberIn{{ID: "bob"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(addRes.Failed) != 1 || addRes.Failed[0].Code != protocol.CodeTalkPasswordRequired {
|
||||
t.Fatalf("%+v", addRes.Failed)
|
||||
}
|
||||
addRes, err = gApp.Add(ctx, "alice", &protocol.GroupAdd{
|
||||
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "4", GroupID: "g_test01",
|
||||
Members: []protocol.GroupMemberIn{{ID: "bob", TalkPassword: "secret"}},
|
||||
})
|
||||
if err != nil || len(addRes.Failed) != 0 {
|
||||
t.Fatalf("err=%v failed=%+v", err, addRes.Failed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF16LeaveRemoveDissolve(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, _, msgApp, db, down := setup(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
insertEP(t, db, "bob", 1)
|
||||
insertEP(t, db, "carol", 1)
|
||||
created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1",
|
||||
ID: "g_leave1", Name: "L",
|
||||
Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "carol"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// 群主不能直接退出
|
||||
err = gApp.Leave(ctx, "alice", &protocol.GroupLeave{
|
||||
V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "2", GroupID: created.ID,
|
||||
})
|
||||
if protoCode(err) != protocol.CodeOwnerCannotLeave {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
// 提交延迟群消息,再让 bob 退出 → pending 未推送改 left_group
|
||||
delay := int64(60_000)
|
||||
_, err = msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "s1", ID: "gm1",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||
DelayMs: &delay,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 到点最小分发:手动把消息改成 dispatched + pending deliveries(模拟已分发未推送)
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var seq int64
|
||||
if e := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id=? AND id=?`, "alice", "gm1").Scan(&seq); e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`UPDATE messages SET state='dispatched', send_at=? WHERE seq=?`, 1_700_000_000_000, seq); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, ep := range []string{"bob", "carol"} {
|
||||
if _, e := tx.Exec(`
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at)
|
||||
VALUES(?,?,?,0,'pending','',?)`, seq, ep, 1_700_000_000_000, 1_700_000_000_000); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err = gApp.Leave(ctx, "bob", &protocol.GroupLeave{
|
||||
V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "3", GroupID: created.ID,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var reason string
|
||||
err = db.Read.QueryRow(`
|
||||
SELECT d.reason FROM deliveries d
|
||||
JOIN messages m ON m.seq=d.seq
|
||||
WHERE m.id='gm1' AND d.endpoint_id='bob'`).Scan(&reason)
|
||||
if err != nil || reason != "left_group" {
|
||||
t.Fatalf("reason=%q err=%v", reason, err)
|
||||
}
|
||||
|
||||
// 已推送的踢人发 revoked
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE deliveries SET pushed_at=?, pushed_conn='c' WHERE endpoint_id='carol' AND state='pending'`, 1_700_000_000_000)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
down.mu.Lock()
|
||||
down.msgs = nil
|
||||
down.mu.Unlock()
|
||||
if err = gApp.Remove(ctx, "alice", &protocol.GroupRemove{
|
||||
V: protocol.Version, Type: protocol.TypeGroupRemove, RID: "4",
|
||||
GroupID: created.ID, EndpointID: "carol",
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
down.mu.Lock()
|
||||
nRev := 0
|
||||
for _, m := range down.msgs {
|
||||
if m.qos == 1 {
|
||||
nRev++
|
||||
}
|
||||
}
|
||||
down.mu.Unlock()
|
||||
if nRev < 1 {
|
||||
t.Fatal("expected revoked for pushed delivery")
|
||||
}
|
||||
|
||||
// 解散:scheduled 作废,编号可复用
|
||||
_, err = msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "s2", ID: "gm2",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "later"},
|
||||
DelayMs: &delay,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 重新加 carol 以便解散时有成员(alice 仍是群主)
|
||||
_, _ = gApp.AdminAddMembers(ctx, created.ID, []string{"carol"})
|
||||
if err = gApp.Dissolve(ctx, "alice", &protocol.GroupDissolve{
|
||||
V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "5", GroupID: created.ID,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var state, mreason string
|
||||
err = db.Read.QueryRow(`SELECT state, reason FROM messages WHERE id='gm2'`).Scan(&state, &mreason)
|
||||
if err != nil || state != "completed" || mreason != "group_dissolved" {
|
||||
t.Fatalf("state=%s reason=%s err=%v", state, mreason, err)
|
||||
}
|
||||
// 同编号新建群
|
||||
created2, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "6",
|
||||
ID: created.ID, Name: "新群", Members: nil,
|
||||
})
|
||||
if err != nil || created2.ID != created.ID {
|
||||
t.Fatalf("reuse id err=%v %+v", err, created2)
|
||||
}
|
||||
// 旧 scheduled 不应再存在为 scheduled
|
||||
_ = db.Read.QueryRow(`SELECT state FROM messages WHERE id='gm2'`).Scan(&state)
|
||||
if state == "scheduled" {
|
||||
t.Fatal("old scheduled should stay completed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestF06GroupSendMembership(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, _, msgApp, db, _ := setup(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
insertEP(t, db, "bob", 1)
|
||||
insertEP(t, db, "carol", 1)
|
||||
created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1",
|
||||
Name: "G", Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "carol"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "s", ID: "m1",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "broadcast"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != message.StateDispatched {
|
||||
t.Fatalf("state=%s", res.State)
|
||||
}
|
||||
var cnt int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1'`).Scan(&cnt)
|
||||
if cnt != 2 {
|
||||
t.Fatalf("deliveries=%d want 2 (not sender)", cnt)
|
||||
}
|
||||
var self int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.endpoint_id='alice'`).Scan(&self)
|
||||
if self != 0 {
|
||||
t.Fatal("sender should not receive")
|
||||
}
|
||||
|
||||
// 发送后入群收不到旧消息:新成员 dave 入群后不应有该投递
|
||||
insertEP(t, db, "dave", 1)
|
||||
_, err = gApp.Add(ctx, "alice", &protocol.GroupAdd{
|
||||
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "2", GroupID: created.ID,
|
||||
Members: []protocol.GroupMemberIn{{ID: "dave"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var daveN int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.endpoint_id='dave'`).Scan(&daveN)
|
||||
if daveN != 0 {
|
||||
t.Fatal("late joiner should not get old delivery")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGroupTransferRenameListGet(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, _, _, db, _ := setup(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
insertEP(t, db, "bob", 1)
|
||||
created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1",
|
||||
Name: "N", Members: []protocol.GroupMemberIn{{ID: "bob"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = gApp.Rename(ctx, "alice", &protocol.GroupRename{
|
||||
V: protocol.Version, Type: protocol.TypeGroupRename, RID: "2",
|
||||
GroupID: created.ID, Name: "新名",
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = gApp.Transfer(ctx, "alice", &protocol.GroupTransfer{
|
||||
V: protocol.Version, Type: protocol.TypeGroupTransfer, RID: "3",
|
||||
GroupID: created.ID, EndpointID: "bob",
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
items, _, err := gApp.List(ctx, "alice", &protocol.GroupList{
|
||||
V: protocol.Version, Type: protocol.TypeGroupList, RID: "4", Limit: 10,
|
||||
})
|
||||
if err != nil || len(items) != 1 || items[0].OwnerID != "bob" {
|
||||
t.Fatalf("%+v err=%v", items, err)
|
||||
}
|
||||
got, err := gApp.Get(ctx, "alice", &protocol.GroupGet{
|
||||
V: protocol.Version, Type: protocol.TypeGroupGet, RID: "5", GroupID: created.ID, Limit: 10,
|
||||
})
|
||||
if err != nil || got.Name != "新名" || len(got.Members) != 2 {
|
||||
t.Fatalf("%+v err=%v", got, err)
|
||||
}
|
||||
_, err = gApp.Get(ctx, "nobody", &protocol.GroupGet{
|
||||
V: protocol.Version, Type: protocol.TypeGroupGet, RID: "6", GroupID: created.ID,
|
||||
})
|
||||
if protoCode(err) != protocol.CodeNotMember && protoCode(err) != protocol.CodeNotFound {
|
||||
// nobody 不是端也不是成员
|
||||
if protoCode(err) == "" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package group
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
type revokeItem struct {
|
||||
endpointID string
|
||||
msgID string
|
||||
fromID string
|
||||
reason string
|
||||
}
|
||||
|
||||
// voidMemberDeliveriesTx rejects pending deliveries for a leaving member; records revokes for pushed ones.
|
||||
func voidMemberDeliveriesTx(tx *sql.Tx, groupID, endpointID, reason string, nowMs int64, revokes *[]revokeItem) error {
|
||||
rows, err := tx.Query(`
|
||||
SELECT d.seq, d.pushed_at, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
AND m.dest_kind = 'group' AND m.dest_id = ?`, endpointID, groupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
type row struct {
|
||||
seq int64
|
||||
pushed sql.NullInt64
|
||||
msgID string
|
||||
senderID string
|
||||
}
|
||||
var list []row
|
||||
for rows.Next() {
|
||||
var r row
|
||||
if scanErr := rows.Scan(&r.seq, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
|
||||
return scanErr
|
||||
}
|
||||
list = append(list, r)
|
||||
}
|
||||
if err = rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range list {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
reason, nowMs, r.seq, endpointID); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason,
|
||||
})
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// voidGroupAllTx rejects all pending group deliveries and completes scheduled messages.
|
||||
func voidGroupAllTx(tx *sql.Tx, groupID string, nowMs int64, revokes *[]revokeItem) error {
|
||||
rows, err := tx.Query(`
|
||||
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type drow struct {
|
||||
seq int64
|
||||
endpointID string
|
||||
pushed sql.NullInt64
|
||||
msgID string
|
||||
senderID string
|
||||
}
|
||||
var dlist []drow
|
||||
for rows.Next() {
|
||||
var r drow
|
||||
if scanErr := rows.Scan(&r.seq, &r.endpointID, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return scanErr
|
||||
}
|
||||
dlist = append(dlist, r)
|
||||
}
|
||||
_ = rows.Close()
|
||||
if err = rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range dlist {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
reasonGroupDissolved, nowMs, r.seq, r.endpointID); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
srows, err := tx.Query(`
|
||||
SELECT seq, id, sender_id, receipt FROM messages
|
||||
WHERE dest_kind = 'group' AND dest_id = ? AND state = 'scheduled'`, groupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type srow struct {
|
||||
seq int64
|
||||
msgID string
|
||||
senderID string
|
||||
receipt int
|
||||
}
|
||||
var slist []srow
|
||||
for srows.Next() {
|
||||
var r srow
|
||||
if scanErr := srows.Scan(&r.seq, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
|
||||
_ = srows.Close()
|
||||
return scanErr
|
||||
}
|
||||
slist = append(slist, r)
|
||||
}
|
||||
_ = srows.Close()
|
||||
if err = srows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range slist {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
|
||||
reasonGroupDissolved, r.seq); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, r.seq); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.receipt != 0 {
|
||||
if _, execErr := tx.Exec(`
|
||||
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
|
||||
VALUES(?,?,?,?,?,?,0)`,
|
||||
r.senderID, r.msgID, "", "completed", reasonGroupDissolved, nowMs); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *App) publishRevokes(ctx context.Context, items []revokeItem) {
|
||||
if a.down == nil || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
for _, it := range items {
|
||||
frame := protocol.Revoked{
|
||||
V: protocol.Version, Type: protocol.TypeRevoked,
|
||||
ID: it.msgID, From: it.fromID, Reason: it.reason,
|
||||
}
|
||||
payload, encErr := encodeFrame(frame)
|
||||
if encErr != nil {
|
||||
continue
|
||||
}
|
||||
_ = a.down.PublishDown(ctx, it.endpointID, "", payload, port.PublishOpts{QoS: 1})
|
||||
}
|
||||
}
|
||||
|
||||
func isUnique(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
return strings.Contains(msg, "unique") || strings.Contains(msg, "constraint failed")
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"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"
|
||||
)
|
||||
|
||||
const (
|
||||
grantKindPassword = "password"
|
||||
grantKindReply = "reply"
|
||||
)
|
||||
|
||||
// Config 是身份服务依赖(注册 + self + 对话密码)。
|
||||
type Config struct {
|
||||
DB *store.DB
|
||||
Hash auth.HashPool
|
||||
Locks auth.LoginLocks
|
||||
Logger *slog.Logger
|
||||
Now func() time.Time
|
||||
// ClientIP 仅注册 HTTP 用。
|
||||
ClientIP func(*http.Request) string
|
||||
|
||||
// Sessions 签发会话令牌;改登录密码必填。
|
||||
Sessions auth.SessionTokens
|
||||
// MaxScheduleSeconds 限制 self.update 的 default_delay_ms。
|
||||
MaxScheduleSeconds int64
|
||||
// ConnControl 可选:logout 后踢线;未接线时为 nil。
|
||||
ConnControl port.ConnControl
|
||||
}
|
||||
|
||||
// App 实现 identity.Service(含 I1 注册与 I2 self/对话密码)。
|
||||
type App struct {
|
||||
handler *RegisterHandler
|
||||
db *store.DB
|
||||
hash auth.HashPool
|
||||
locks auth.LoginLocks
|
||||
sessions auth.SessionTokens
|
||||
maxScheduleSeconds int64
|
||||
connCtrl port.ConnControl
|
||||
nowFn func() time.Time
|
||||
}
|
||||
|
||||
// New 构造完整身份服务。
|
||||
func New(cfg Config) *App {
|
||||
if cfg.Logger == nil {
|
||||
cfg.Logger = slog.Default()
|
||||
}
|
||||
if cfg.Now == nil {
|
||||
cfg.Now = time.Now
|
||||
}
|
||||
if cfg.ClientIP == nil {
|
||||
cfg.ClientIP = clientIPFromRemoteAddr
|
||||
}
|
||||
if cfg.Sessions == nil {
|
||||
cfg.Sessions = auth.NewStubSessionTokens()
|
||||
}
|
||||
if cfg.Locks == nil {
|
||||
cfg.Locks = auth.NewStubLoginLocks()
|
||||
}
|
||||
h := NewRegisterHandler(RegisterConfig{
|
||||
DB: cfg.DB,
|
||||
Hash: cfg.Hash,
|
||||
Locks: cfg.Locks,
|
||||
Logger: cfg.Logger,
|
||||
Now: cfg.Now,
|
||||
ClientIP: cfg.ClientIP,
|
||||
})
|
||||
return &App{
|
||||
handler: h,
|
||||
db: cfg.DB,
|
||||
hash: cfg.Hash,
|
||||
locks: cfg.Locks,
|
||||
sessions: cfg.Sessions,
|
||||
maxScheduleSeconds: cfg.MaxScheduleSeconds,
|
||||
connCtrl: cfg.ConnControl,
|
||||
nowFn: cfg.Now,
|
||||
}
|
||||
}
|
||||
|
||||
// NewServer 兼容 I1:用注册配置构造 Service(会话令牌用 Stub)。
|
||||
func NewServer(cfg RegisterConfig) *App {
|
||||
return New(Config{
|
||||
DB: cfg.DB,
|
||||
Hash: cfg.Hash,
|
||||
Locks: cfg.Locks,
|
||||
Logger: cfg.Logger,
|
||||
Now: cfg.Now,
|
||||
ClientIP: cfg.ClientIP,
|
||||
Sessions: auth.NewStubSessionTokens(),
|
||||
})
|
||||
}
|
||||
|
||||
func (a *App) now() time.Time { return a.nowFn() }
|
||||
|
||||
// Handler 返回可挂载的注册 HTTP 处理器。
|
||||
func (a *App) Handler() http.Handler { return a.handler }
|
||||
|
||||
// Register 实现自助注册。
|
||||
func (a *App) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
|
||||
preq := &protocol.RegisterRequest{
|
||||
RegistrationCode: req.RegistrationCode,
|
||||
ID: req.ID,
|
||||
LoginPassword: req.LoginPassword,
|
||||
Name: req.Name,
|
||||
TalkPassword: req.TalkPassword,
|
||||
}
|
||||
ip := req.RemoteIP
|
||||
if ip == "" {
|
||||
ip = "0.0.0.0"
|
||||
}
|
||||
result, apiErr := a.handler.register(ctx, preq, ip)
|
||||
if apiErr != nil {
|
||||
return RegisterResult{}, apiErr
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func encodeSessionHash(hash []byte) string {
|
||||
return hex.EncodeToString(hash)
|
||||
}
|
||||
|
||||
var _ Service = (*App)(nil)
|
||||
var _ http.Handler = (*RegisterHandler)(nil)
|
||||
@@ -0,0 +1,7 @@
|
||||
package identity
|
||||
|
||||
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
|
||||
func errCode(code, msg string) *protocol.Error {
|
||||
return &protocol.Error{Code: code, Message: msg}
|
||||
}
|
||||
@@ -247,65 +247,6 @@ func (h *RegisterHandler) logResult(result, id, ip string) {
|
||||
h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip)
|
||||
}
|
||||
|
||||
// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。
|
||||
type Server struct {
|
||||
handler *RegisterHandler
|
||||
}
|
||||
|
||||
// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。
|
||||
func NewServer(cfg RegisterConfig) *Server {
|
||||
return &Server{handler: NewRegisterHandler(cfg)}
|
||||
}
|
||||
|
||||
// Handler 返回可挂载的注册 HTTP 处理器。
|
||||
func (s *Server) Handler() http.Handler { return s.handler }
|
||||
|
||||
// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。
|
||||
func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
|
||||
preq := &protocol.RegisterRequest{
|
||||
RegistrationCode: req.RegistrationCode,
|
||||
ID: req.ID,
|
||||
LoginPassword: req.LoginPassword,
|
||||
Name: req.Name,
|
||||
TalkPassword: req.TalkPassword,
|
||||
}
|
||||
ip := req.RemoteIP
|
||||
if ip == "" {
|
||||
ip = "0.0.0.0"
|
||||
}
|
||||
result, err := s.handler.register(ctx, preq, ip)
|
||||
if err != nil {
|
||||
return RegisterResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) {
|
||||
return SelfInfo{}, ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfSetTalkPassword(context.Context, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
|
||||
return "", ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Server) UnlockTalk(context.Context, string, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented }
|
||||
|
||||
var _ Service = (*Server)(nil)
|
||||
var _ http.Handler = (*RegisterHandler)(nil)
|
||||
|
||||
func setCORS(w http.ResponseWriter) {
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
// SelfGet 返回自己的资料。
|
||||
func (a *App) SelfGet(ctx context.Context, endpointID string) (SelfInfo, error) {
|
||||
if !protocol.ValidEndpointID(endpointID) {
|
||||
return SelfInfo{}, errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
var info SelfInfo
|
||||
var talk sql.NullString
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT id, name, default_delay_ms, talk_hash
|
||||
FROM endpoints WHERE id = ?`, endpointID).Scan(&info.ID, &info.Name, &info.DefaultDelayMs, &talk)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SelfInfo{}, errCode(protocol.CodeNotFound, "endpoint not found")
|
||||
}
|
||||
if err != nil {
|
||||
return SelfInfo{}, err
|
||||
}
|
||||
info.TalkPasswordSet = talk.Valid && talk.String != ""
|
||||
return info, nil
|
||||
}
|
||||
|
||||
// SelfUpdate 更新名称与默认延迟。
|
||||
func (a *App) SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error {
|
||||
if !protocol.ValidEndpointID(endpointID) {
|
||||
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
if req.Name == "" && req.DefaultDelayMs == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nothing to update")
|
||||
}
|
||||
if req.DefaultDelayMs != nil && a.maxScheduleSeconds > 0 {
|
||||
maxMs := a.maxScheduleSeconds * 1000
|
||||
if *req.DefaultDelayMs > maxMs {
|
||||
return errCode(protocol.CodeBadRequest, "default_delay_ms exceeds max_schedule_seconds")
|
||||
}
|
||||
}
|
||||
|
||||
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var exists int
|
||||
if err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, endpointID).Scan(&exists); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeNotFound, "endpoint not found")
|
||||
}
|
||||
return err
|
||||
}
|
||||
if req.Name != "" {
|
||||
if _, err := tx.Exec(`UPDATE endpoints SET name = ? WHERE id = ?`, req.Name, endpointID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if req.DefaultDelayMs != nil {
|
||||
if _, err := tx.Exec(`UPDATE endpoints SET default_delay_ms = ? WHERE id = ?`, *req.DefaultDelayMs, endpointID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// SelfSetTalkPassword 设置或清除对话密码,并增加 talk_version;改密清零对方锁定计数。
|
||||
func (a *App) SelfSetTalkPassword(ctx context.Context, endpointID, talkPassword string) error {
|
||||
if !protocol.ValidEndpointID(endpointID) {
|
||||
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
if !protocol.ValidTalkPassword(talkPassword) {
|
||||
return errCode(protocol.CodeBadRequest, "invalid talk_password")
|
||||
}
|
||||
|
||||
var talkHash any
|
||||
if talkPassword != "" {
|
||||
if a.hash == nil {
|
||||
return errors.New("identity: hash pool required")
|
||||
}
|
||||
phc, err := a.hash.Hash(ctx, auth.PasswordTalk, talkPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
talkHash = phc
|
||||
}
|
||||
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, err := tx.Exec(`
|
||||
UPDATE endpoints
|
||||
SET talk_hash = ?, talk_version = talk_version + 1
|
||||
WHERE id = ?`, talkHash, endpointID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return errCode(protocol.CodeNotFound, "endpoint not found")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// D24:改密清零按对方计的对话密码失败计数。
|
||||
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: endpointID})
|
||||
return nil
|
||||
}
|
||||
|
||||
// SelfChangeLoginPassword 要求旧密码;成功返回新 session_token,旧令牌作废,当前连接由调用方保留。
|
||||
func (a *App) SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (string, error) {
|
||||
if !protocol.ValidEndpointID(endpointID) {
|
||||
return "", errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
if protocol.LoginPasswordForbiddenPrefix(newPassword) {
|
||||
return "", errCode(protocol.CodeBadRequest, "login password must not start with nst_")
|
||||
}
|
||||
if !protocol.ValidLoginPassword(newPassword) || newPassword == "" {
|
||||
return "", errCode(protocol.CodeBadRequest, "invalid new_password")
|
||||
}
|
||||
if a.hash == nil || a.sessions == nil {
|
||||
return "", errors.New("identity: hash/sessions required")
|
||||
}
|
||||
|
||||
var loginHash string
|
||||
err := a.db.Read.QueryRowContext(ctx, `SELECT login_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&loginHash)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return "", errCode(protocol.CodeNotFound, "endpoint not found")
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}); locked {
|
||||
return "", errCode(protocol.CodeRateLimited, "login locked")
|
||||
}
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}); locked {
|
||||
return "", errCode(protocol.CodeRateLimited, "login locked")
|
||||
}
|
||||
|
||||
ok, err := a.hash.Verify(ctx, auth.PasswordLogin, oldPassword, loginHash)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !ok {
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP})
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID})
|
||||
return "", errCode(protocol.CodeUnauthorized, "old password invalid")
|
||||
}
|
||||
|
||||
newHash, err := a.hash.Hash(ctx, auth.PasswordLogin, newPassword)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
token, tokenHash, err := a.sessions.Issue(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
nowMs := a.now().UnixMilli()
|
||||
hashHex := encodeSessionHash(tokenHash)
|
||||
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, e := tx.Exec(`
|
||||
UPDATE endpoints
|
||||
SET login_hash = ?, session_hash = ?, session_issued_at = ?, session_used_at = ?
|
||||
WHERE id = ?`, newHash, hashHex, nowMs, nowMs, endpointID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return errCode(protocol.CodeNotFound, "endpoint not found")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// SelfLogout 清空会话令牌;若注入了 ConnControl 则断开当前连接。
|
||||
func (a *App) SelfLogout(ctx context.Context, endpointID string) error {
|
||||
if !protocol.ValidEndpointID(endpointID) {
|
||||
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, e := tx.Exec(`
|
||||
UPDATE endpoints
|
||||
SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL
|
||||
WHERE id = ?`, endpointID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return errCode(protocol.CodeNotFound, "endpoint not found")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if a.connCtrl != nil {
|
||||
_ = a.connCtrl.Disconnect(ctx, endpointID, "", port.DisconnectNormal)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Disable / Enable / Delete 属 I5,此处保留未实现。
|
||||
func (a *App) Disable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (a *App) Enable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (a *App) Delete(context.Context, string) error { return ErrNotImplemented }
|
||||
@@ -0,0 +1,314 @@
|
||||
package identity_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/message"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
func openIdentity(t *testing.T) (*identity.App, *store.DB, *auth.MemoryLocks) {
|
||||
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() })
|
||||
fixed := time.UnixMilli(1_700_000_000_000)
|
||||
locks := auth.NewLoginLocks()
|
||||
locks.SetClock(func() time.Time { return fixed })
|
||||
app := identity.New(identity.Config{
|
||||
DB: db,
|
||||
Hash: auth.NewStubHashPool(),
|
||||
Locks: locks,
|
||||
Sessions: auth.NewSessionTokens(),
|
||||
MaxScheduleSeconds: int64(config.Default().Limits.MaxScheduleSeconds),
|
||||
Now: func() time.Time { return fixed },
|
||||
})
|
||||
return app, db, locks
|
||||
}
|
||||
|
||||
func insertEP(t *testing.T, db *store.DB, id, loginPW string) {
|
||||
t.Helper()
|
||||
err := 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(?,?,?,?,0,0,1,?)`, id, id, "stub$"+loginPW, nil, 1_700_000_000_000)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func protoCode(err error) string {
|
||||
var pe *protocol.Error
|
||||
if errors.As(err, &pe) {
|
||||
return pe.Code
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func TestF15TalkPasswordAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, _ := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "password1")
|
||||
insertEP(t, db, "bob", "password1")
|
||||
|
||||
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// A 不带密码失败
|
||||
if err := app.UnlockTalk(ctx, "alice", "bob", "", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordRequired {
|
||||
t.Fatalf("want talk_password_required got %v", err)
|
||||
}
|
||||
ok, err := app.HasTalkGrant(ctx, "alice", "bob")
|
||||
if err != nil || ok {
|
||||
t.Fatalf("grant=%v err=%v", ok, err)
|
||||
}
|
||||
|
||||
// 带对后成功,之后无密码也有授权
|
||||
if err = app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ok, err = app.HasTalkGrant(ctx, "alice", "bob")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("grant=%v err=%v", ok, err)
|
||||
}
|
||||
|
||||
// B 改密后旧授权失效
|
||||
if err = app.SelfSetTalkPassword(ctx, "bob", "newsecret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ok, err = app.HasTalkGrant(ctx, "alice", "bob")
|
||||
if err != nil || ok {
|
||||
t.Fatalf("after change grant=%v err=%v", ok, err)
|
||||
}
|
||||
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordInvalid {
|
||||
t.Fatalf("want invalid got %v", err)
|
||||
}
|
||||
if err := app.UnlockTalk(ctx, "alice", "bob", "newsecret", "1.1.1.1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF15ReplyGrantAndChange(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, _ := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "password1")
|
||||
insertEP(t, db, "bob", "password1")
|
||||
if err := app.SelfSetTalkPassword(ctx, "alice", "alice-pw"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// B 先给 A 发 → 写入 reply 授权(A 可回 B 免密;这里记的是 bob→alice 的授权给「alice 作为接收方」...
|
||||
// 回复授权:对方曾成功提交发给我的单聊 → 我对对方有 reply 权。
|
||||
// 即 B 发给 A 后,A 对 B 有授权(sender=alice, target=bob)。
|
||||
if err := app.RecordReplyGrant(ctx, "alice", "bob"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 但 bob 还没设密码,alice→bob 本就不需要。给 bob 设密后验证 reply:
|
||||
if err := app.SelfSetTalkPassword(ctx, "bob", "bob-pw"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 重新记 reply:B 发给 A 成功后 A 获得对 B 的回复权
|
||||
if err := app.RecordReplyGrant(ctx, "alice", "bob"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ok, err := app.HasTalkGrant(ctx, "alice", "bob")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("reply grant=%v err=%v", ok, err)
|
||||
}
|
||||
if err = app.SelfSetTalkPassword(ctx, "bob", "bob-pw2"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ok, err = app.HasTalkGrant(ctx, "alice", "bob")
|
||||
if err != nil || ok {
|
||||
t.Fatalf("after change reply should die grant=%v", ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF15JoinNeedsPasswordDespiteGrant(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, _ := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "password1")
|
||||
insertEP(t, db, "bob", "password1")
|
||||
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 已有单聊授权,进群仍要密码
|
||||
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordRequired {
|
||||
t.Fatalf("want required got %v", err)
|
||||
}
|
||||
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "wrong", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordInvalid {
|
||||
t.Fatalf("want invalid got %v", err)
|
||||
}
|
||||
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF15SubmittedMessageUnaffectedByPasswordChange(t *testing.T) {
|
||||
t.Parallel()
|
||||
idApp, db, locks := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "password1")
|
||||
insertEP(t, db, "bob", "password1")
|
||||
if err := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
lim := message.LimitsFromConfig(config.Default().Limits)
|
||||
lim.RequestsPerSecond = 0
|
||||
msgApp := message.New(db, lim, auth.NewStubHashPool(),
|
||||
message.WithNow(func() time.Time { return time.UnixMilli(1_700_000_000_000) }),
|
||||
message.WithLocks(locks),
|
||||
)
|
||||
req := &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "m1",
|
||||
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||
DelayMs: ptrInt64(60_000),
|
||||
TalkPassword: "secret",
|
||||
}
|
||||
res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "1.1.1.1"}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != message.StateScheduled {
|
||||
t.Fatalf("state=%s", res.State)
|
||||
}
|
||||
if err := idApp.SelfSetTalkPassword(ctx, "bob", "changed"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 已提交消息行仍在且状态不变
|
||||
var state string
|
||||
if err := db.Read.QueryRow(`SELECT state FROM messages WHERE sender_id=? AND id=?`, "alice", "m1").Scan(&state); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if state != message.StateScheduled {
|
||||
t.Fatalf("message state changed to %s", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF15TargetLockAndExistingGrant(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, locks := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "bob", "password1")
|
||||
insertEP(t, db, "authd", "password1")
|
||||
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.UnlockTalk(ctx, "authd", "bob", "secret", "9.9.9.9"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 50 次错误触发对方总数锁
|
||||
for i := 0; i < 50; i++ {
|
||||
id := "u" + string(rune('0'+i/100)) + string(rune('0'+(i/10)%10)) + string(rune('0'+i%10))
|
||||
insertEP(t, db, id, "password1")
|
||||
_ = app.UnlockTalk(ctx, id, "bob", "wrong", "2.2.2.2")
|
||||
}
|
||||
locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"})
|
||||
if !locked {
|
||||
t.Fatal("expected talk target lock")
|
||||
}
|
||||
// 正确密码也暂时无法解锁
|
||||
insertEP(t, db, "newbie", "password1")
|
||||
if err := app.UnlockTalk(ctx, "newbie", "bob", "secret", "3.3.3.3"); protoCode(err) != protocol.CodeRateLimited {
|
||||
t.Fatalf("want rate_limited got %v", err)
|
||||
}
|
||||
// 已有授权仍可用
|
||||
ok, err := app.HasTalkGrant(ctx, "authd", "bob")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("authd grant=%v err=%v", ok, err)
|
||||
}
|
||||
// 改密清零
|
||||
if err := app.SelfSetTalkPassword(ctx, "bob", "secret2"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
locked, _ = locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"})
|
||||
if locked {
|
||||
t.Fatal("lock should clear on password change")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelfLoginPasswordAndLogout(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, locks := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "oldpass12")
|
||||
tok, err := app.SelfChangeLoginPassword(ctx, "alice", "wrongpass", "newpass12", "1.1.1.1")
|
||||
if protoCode(err) != protocol.CodeUnauthorized || tok != "" {
|
||||
t.Fatalf("got tok=%q err=%v", tok, err)
|
||||
}
|
||||
locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: "alice", IP: "1.1.1.1"}) // 确保 Fail 路径可调用
|
||||
|
||||
tok, err = app.SelfChangeLoginPassword(ctx, "alice", "oldpass12", "newpass12", "1.1.1.1")
|
||||
if err != nil || tok == "" || !protocol.ValidEndpointID("alice") {
|
||||
t.Fatalf("tok=%q err=%v", tok, err)
|
||||
}
|
||||
if !auth.NewSessionTokens().LooksLikeSessionToken(tok) {
|
||||
t.Fatalf("token prefix %q", tok)
|
||||
}
|
||||
var hash sql.NullString
|
||||
if err = db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id=?`, "alice").Scan(&hash); err != nil || !hash.Valid {
|
||||
t.Fatal(err)
|
||||
}
|
||||
info, err := app.SelfGet(ctx, "alice")
|
||||
if err != nil || info.ID != "alice" {
|
||||
t.Fatal(err)
|
||||
}
|
||||
name := "门口"
|
||||
delay := int64(10000)
|
||||
if err := app.SelfUpdate(ctx, "alice", &protocol.SelfUpdate{
|
||||
V: protocol.Version, Type: protocol.TypeSelfUpdate, RID: "1", Name: name, DefaultDelayMs: &delay,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
info, _ = app.SelfGet(ctx, "alice")
|
||||
if info.Name != name || info.DefaultDelayMs != delay {
|
||||
t.Fatalf("%+v", info)
|
||||
}
|
||||
if err := app.SelfLogout(ctx, "alice"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id=?`, "alice").Scan(&hash); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hash.Valid {
|
||||
t.Fatal("session should be cleared")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnlockSelfAndNoPassword(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, _ := openIdentity(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "password1")
|
||||
insertEP(t, db, "bob", "password1")
|
||||
if err := app.UnlockTalk(ctx, "alice", "alice", "", ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.UnlockTalk(ctx, "alice", "bob", "anything", ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func ptrInt64(v int64) *int64 { return &v }
|
||||
@@ -46,13 +46,16 @@ type Service interface {
|
||||
SelfGet(ctx context.Context, endpointID string) (SelfInfo, error)
|
||||
SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error
|
||||
SelfSetTalkPassword(ctx context.Context, endpointID string, talkPassword string) error
|
||||
SelfChangeLoginPassword(ctx context.Context, endpointID string, oldPassword, newPassword string) (sessionToken string, err error)
|
||||
// SelfChangeLoginPassword 校验旧密码后换新密码并签发新会话令牌;remoteIP 计入登录锁定。
|
||||
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
|
||||
SelfLogout(ctx context.Context, endpointID string) error
|
||||
|
||||
// UnlockTalk 校验并写入对话密码授权(第 6.6 节 unlock)。
|
||||
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword string) error
|
||||
// HasTalkGrant 查询发送方对目标是否有有效授权。
|
||||
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。
|
||||
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error
|
||||
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。
|
||||
HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error)
|
||||
// CheckTalkPasswordForJoin 加人时校验对话密码:已有单聊授权不能代替,必须当次带对。
|
||||
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
|
||||
|
||||
// Disable 停用端并作废相关消息/令牌(第 7.6 节)。
|
||||
Disable(ctx context.Context, endpointID string) error
|
||||
|
||||
@@ -27,13 +27,13 @@ func (s *Stub) SelfSetTalkPassword(context.Context, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
|
||||
func (s *Stub) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
|
||||
func (s *Stub) SelfChangeLoginPassword(context.Context, string, string, string, string) (string, error) {
|
||||
return "", ErrNotImplemented
|
||||
}
|
||||
|
||||
func (s *Stub) SelfLogout(context.Context, string) error { return ErrNotImplemented }
|
||||
|
||||
func (s *Stub) UnlockTalk(context.Context, string, string, string) error {
|
||||
func (s *Stub) UnlockTalk(context.Context, string, string, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
|
||||
@@ -41,6 +41,10 @@ func (s *Stub) HasTalkGrant(context.Context, string, string) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (s *Stub) CheckTalkPasswordForJoin(context.Context, string, string, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
|
||||
func (s *Stub) Disable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Stub) Enable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Stub) Delete(context.Context, string) error { return ErrNotImplemented }
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
// UnlockTalk 校验对话密码并写入 password 授权;未设密码则直接成功。
|
||||
func (a *App) UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error {
|
||||
if !protocol.ValidEndpointID(senderID) || !protocol.ValidEndpointID(targetID) {
|
||||
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
if senderID == targetID {
|
||||
return nil
|
||||
}
|
||||
|
||||
var talkHash sql.NullString
|
||||
var talkVer int64
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeInvalidTarget, "target not found")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !talkHash.Valid || talkHash.String == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}); locked {
|
||||
return errCode(protocol.CodeRateLimited, "talk password locked")
|
||||
}
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
|
||||
return errCode(protocol.CodeRateLimited, "talk password locked")
|
||||
}
|
||||
if talkPassword == "" {
|
||||
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
|
||||
}
|
||||
if a.hash == nil {
|
||||
return errors.New("identity: hash pool required")
|
||||
}
|
||||
ok, err := a.hash.Verify(ctx, auth.PasswordTalk, talkPassword, talkHash.String)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP})
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
|
||||
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
|
||||
}
|
||||
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP})
|
||||
|
||||
nowMs := a.now().UnixMilli()
|
||||
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
return upsertGrantTx(tx, senderID, targetID, talkVer, grantKindPassword, nowMs)
|
||||
})
|
||||
}
|
||||
|
||||
// HasTalkGrant 发给自己、对方未设密码、或存在匹配版本授权时为 true。
|
||||
func (a *App) HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) {
|
||||
if senderID == targetID {
|
||||
return true, nil
|
||||
}
|
||||
if !protocol.ValidEndpointID(senderID) || !protocol.ValidEndpointID(targetID) {
|
||||
return false, errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
||||
}
|
||||
var talkHash sql.NullString
|
||||
var talkVer int64
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, errCode(protocol.CodeInvalidTarget, "target not found")
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !talkHash.Valid || talkHash.String == "" {
|
||||
return true, nil
|
||||
}
|
||||
var n int
|
||||
err = a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT 1 FROM talk_grants
|
||||
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
|
||||
LIMIT 1`, senderID, targetID, talkVer).Scan(&n)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, nil
|
||||
}
|
||||
return err == nil, err
|
||||
}
|
||||
|
||||
// CheckTalkPasswordForJoin 加人时必须当次带对密码;已有单聊授权不能代替。
|
||||
func (a *App) CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error {
|
||||
if !protocol.ValidEndpointID(targetID) {
|
||||
return errCode(protocol.CodeInvalidTarget, "invalid target")
|
||||
}
|
||||
var talkHash sql.NullString
|
||||
var talkVer int64
|
||||
var enabled int
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer, &enabled)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeInvalidTarget, "target not found")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if enabled == 0 {
|
||||
return errCode(protocol.CodeEndpointDisabled, "endpoint disabled")
|
||||
}
|
||||
if !talkHash.Valid || talkHash.String == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}); locked {
|
||||
return errCode(protocol.CodeRateLimited, "talk password locked")
|
||||
}
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
|
||||
return errCode(protocol.CodeRateLimited, "talk password locked")
|
||||
}
|
||||
if talkPassword == "" {
|
||||
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
|
||||
}
|
||||
if a.hash == nil {
|
||||
return errors.New("identity: hash pool required")
|
||||
}
|
||||
ok, err := a.hash.Verify(ctx, auth.PasswordTalk, talkPassword, talkHash.String)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP})
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
|
||||
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
|
||||
}
|
||||
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP})
|
||||
// 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。
|
||||
_ = talkVer
|
||||
return nil
|
||||
}
|
||||
|
||||
func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error {
|
||||
_, err := tx.Exec(`
|
||||
INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at)
|
||||
VALUES(?,?,?,?,?)
|
||||
ON CONFLICT(sender_id, target_id) DO UPDATE SET
|
||||
target_talk_version = excluded.target_talk_version,
|
||||
kind = excluded.kind,
|
||||
created_at = excluded.created_at`,
|
||||
senderID, targetID, talkVersion, kind, nowMs,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// RecordReplyGrant 在对方成功提交单聊后写入 reply 授权(供消息线或测试调用)。
|
||||
func (a *App) RecordReplyGrant(ctx context.Context, senderID, targetID string) error {
|
||||
if senderID == targetID {
|
||||
return nil
|
||||
}
|
||||
var talkHash sql.NullString
|
||||
var talkVer int64
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeInvalidTarget, "target not found")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !talkHash.Valid || talkHash.String == "" {
|
||||
return nil
|
||||
}
|
||||
nowMs := a.now().UnixMilli()
|
||||
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
return upsertGrantTx(tx, senderID, targetID, talkVer, grantKindReply, nowMs)
|
||||
})
|
||||
}
|
||||
@@ -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
@@ -1,7 +1,7 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
@@ -18,11 +18,6 @@ const (
|
||||
StateCompleted = "completed"
|
||||
)
|
||||
|
||||
// 投递状态。
|
||||
const (
|
||||
DeliveryPending = "pending"
|
||||
)
|
||||
|
||||
// talk_grants.kind。
|
||||
const (
|
||||
GrantKindPassword = "password"
|
||||
@@ -32,7 +27,7 @@ const (
|
||||
// 请求频率桶默认突发容量(DEVELOPMENT 6.10;配置无单独字段)。
|
||||
const defaultRequestBurst = 100
|
||||
|
||||
// Limits 是提交所需的配置上限(来自 config.LimitsConfig)。
|
||||
// Limits 是消息子系统所需配置上限。
|
||||
type Limits struct {
|
||||
MaxBodyBytes int
|
||||
MaxMetaBytes int
|
||||
@@ -44,11 +39,16 @@ type Limits struct {
|
||||
MaxPendingPerSender int
|
||||
MaxPendingPerReceiver int
|
||||
GraceSeconds int64
|
||||
AckTimeoutSeconds int64
|
||||
DeliveryWindow int
|
||||
ReceiptWindow int
|
||||
RecordRetentionDays int
|
||||
ReceiptRetentionDays int
|
||||
IdempotencyHours int
|
||||
}
|
||||
|
||||
// LimitsFromConfig 从平台配置构造 Limits。
|
||||
func LimitsFromConfig(c config.LimitsConfig) Limits {
|
||||
burst := defaultRequestBurst
|
||||
return Limits{
|
||||
MaxBodyBytes: c.MaxBodyBytes,
|
||||
MaxMetaBytes: c.MaxMetaBytes,
|
||||
@@ -56,14 +56,26 @@ func LimitsFromConfig(c config.LimitsConfig) Limits {
|
||||
MaxTTLSeconds: int64(c.MaxTTLSeconds),
|
||||
MaxScheduleSeconds: int64(c.MaxScheduleSeconds),
|
||||
RequestsPerSecond: float64(c.RequestsPerSecond),
|
||||
RequestBurst: burst,
|
||||
RequestBurst: defaultRequestBurst,
|
||||
MaxPendingPerSender: c.MaxPendingPerSender,
|
||||
MaxPendingPerReceiver: c.MaxPendingPerReceiver,
|
||||
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 {
|
||||
db *store.DB
|
||||
lim Limits
|
||||
@@ -71,6 +83,15 @@ type App struct {
|
||||
locks auth.LoginLocks
|
||||
nowFn func() time.Time
|
||||
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。
|
||||
@@ -86,17 +107,41 @@ func WithLocks(locks auth.LoginLocks) Option {
|
||||
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 {
|
||||
if lim.RequestBurst <= 0 {
|
||||
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{
|
||||
db: db,
|
||||
lim: lim,
|
||||
hash: hash,
|
||||
nowFn: time.Now,
|
||||
rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst),
|
||||
db: db,
|
||||
lim: lim,
|
||||
hash: hash,
|
||||
nowFn: time.Now,
|
||||
rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst),
|
||||
largeSem: make(chan struct{}, maxLargeInflight),
|
||||
largeHeld: make(map[string]bool),
|
||||
}
|
||||
for _, opt := range opts {
|
||||
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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"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) {
|
||||
if req == nil {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "nil send")
|
||||
@@ -250,10 +250,12 @@ INSERT INTO messages(
|
||||
|
||||
finalState := state
|
||||
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 {
|
||||
return e
|
||||
}
|
||||
_ = claimed
|
||||
}
|
||||
result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState}
|
||||
return nil
|
||||
@@ -261,9 +263,29 @@ INSERT INTO messages(
|
||||
if err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
if result.State == StateDispatched {
|
||||
a.wakeReceivers(ctx, result.ID, senderID)
|
||||
}
|
||||
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) {
|
||||
if req.SendAtMs != nil && req.DelayMs != nil {
|
||||
return 0, errCode(protocol.CodeBadRequest, "send_at_ms and delay_ms are mutually exclusive")
|
||||
|
||||
@@ -50,6 +50,9 @@ func defaultTestLimits() Limits {
|
||||
cfg := config.Default().Limits
|
||||
lim := LimitsFromConfig(cfg)
|
||||
lim.RequestsPerSecond = 0 // 测试默认不限速
|
||||
lim.RecordRetentionDays = 7
|
||||
lim.ReceiptRetentionDays = 7
|
||||
lim.IdempotencyHours = 24
|
||||
return lim
|
||||
}
|
||||
|
||||
@@ -62,11 +65,12 @@ func insertEndpoint(t *testing.T, db *store.DB, id string, talkPassword string,
|
||||
talk = "stub$" + talkPassword
|
||||
talkVer = 1
|
||||
}
|
||||
nowMs := int64(1_700_000_000_000)
|
||||
err := db.Queue.Do(ctx, 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(?,?,?,?,?,?,?,?)`,
|
||||
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, 1_700_000_000_000)
|
||||
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at, offline_since)
|
||||
VALUES(?,?,?,?,?,?,?,?,?)`,
|
||||
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, nowMs, nowMs)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,353 @@
|
||||
package presence
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
// ConnTable is an injectable connection table (N3). When nil, use SetOnline/SetOffline and DB columns.
|
||||
type ConnTable interface {
|
||||
IsOnline(endpointID string) bool
|
||||
CurrentConn(endpointID string) (port.ConnID, bool)
|
||||
}
|
||||
|
||||
// Config holds presence service dependencies.
|
||||
type Config struct {
|
||||
DB *store.DB
|
||||
Downlink port.Downlink // QoS 0 presence notifies; nil skips push
|
||||
Conns ConnTable // optional
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
type watchSub struct {
|
||||
endpointID string
|
||||
all bool
|
||||
ids map[string]struct{}
|
||||
}
|
||||
|
||||
type onlineEntry struct {
|
||||
connID port.ConnID
|
||||
atMs int64
|
||||
}
|
||||
|
||||
// App implements presence.Service.
|
||||
type App struct {
|
||||
db *store.DB
|
||||
down port.Downlink
|
||||
conns ConnTable
|
||||
nowFn func() time.Time
|
||||
mu sync.Mutex
|
||||
online map[string]onlineEntry
|
||||
watches map[port.ConnID]watchSub
|
||||
}
|
||||
|
||||
// New constructs the presence service.
|
||||
func New(cfg Config) *App {
|
||||
now := cfg.Now
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
return &App{
|
||||
db: cfg.DB,
|
||||
down: cfg.Downlink,
|
||||
conns: cfg.Conns,
|
||||
nowFn: now,
|
||||
online: make(map[string]onlineEntry),
|
||||
watches: make(map[port.ConnID]watchSub),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) nowMs() int64 { return a.nowFn().UnixMilli() }
|
||||
|
||||
// Get queries online status for up to 200 ids.
|
||||
func (a *App) Get(ctx context.Context, ids []string) ([]StatusItem, error) {
|
||||
if len(ids) > protocol.MaxPresenceGetIDs {
|
||||
return nil, &protocol.Error{Code: protocol.CodeBadRequest, Message: "too many ids"}
|
||||
}
|
||||
out := make([]StatusItem, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
item := StatusItem{ID: id}
|
||||
var onlineSince, offlineSince sql.NullInt64
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT online_since, offline_since FROM endpoints WHERE id = ?`, id).Scan(&onlineSince, &offlineSince)
|
||||
if err == sql.ErrNoRows {
|
||||
item.NotFound = true
|
||||
out = append(out, item)
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
on, since := a.resolveOnline(id, onlineSince, offlineSince)
|
||||
item.Online = on
|
||||
item.SinceMs = since
|
||||
out = append(out, item)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Directory lists endpoints with optional query and pagination.
|
||||
func (a *App) Directory(ctx context.Context, req *protocol.DirectoryList) ([]DirectoryItem, string, error) {
|
||||
if req == nil {
|
||||
return nil, "", &protocol.Error{Code: protocol.CodeBadRequest, Message: "nil request"}
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > protocol.MaxPageLimit {
|
||||
limit = protocol.MaxPageLimit
|
||||
}
|
||||
cursor := req.Cursor
|
||||
query := strings.TrimSpace(req.Query)
|
||||
|
||||
var rows *sql.Rows
|
||||
var err error
|
||||
if query == "" {
|
||||
rows, err = a.db.Read.QueryContext(ctx, `
|
||||
SELECT id, name, online_since, offline_since, talk_hash
|
||||
FROM endpoints
|
||||
WHERE id > ?
|
||||
ORDER BY id ASC
|
||||
LIMIT ?`, cursor, limit+1)
|
||||
} else {
|
||||
like := "%" + strings.ToLower(query) + "%"
|
||||
prefix := strings.ToLower(query) + "%"
|
||||
rows, err = a.db.Read.QueryContext(ctx, `
|
||||
SELECT id, name, online_since, offline_since, talk_hash
|
||||
FROM endpoints
|
||||
WHERE id > ?
|
||||
AND (lower(id) LIKE ? OR lower(name) LIKE ?)
|
||||
ORDER BY id ASC
|
||||
LIMIT ?`, cursor, prefix, like, limit+1)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
items := make([]DirectoryItem, 0, limit)
|
||||
for rows.Next() {
|
||||
var id, name string
|
||||
var onlineSince, offlineSince sql.NullInt64
|
||||
var talk sql.NullString
|
||||
if err := rows.Scan(&id, &name, &onlineSince, &offlineSince, &talk); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
on, _ := a.resolveOnline(id, onlineSince, offlineSince)
|
||||
item := DirectoryItem{
|
||||
ID: id,
|
||||
Name: name,
|
||||
Online: on,
|
||||
TalkPasswordSet: talk.Valid && talk.String != "",
|
||||
}
|
||||
if onlineSince.Valid {
|
||||
v := onlineSince.Int64
|
||||
item.OnlineSinceMs = &v
|
||||
}
|
||||
if offlineSince.Valid {
|
||||
v := offlineSince.Int64
|
||||
item.OfflineSinceMs = &v
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
next := ""
|
||||
if len(items) > limit {
|
||||
items = items[:limit]
|
||||
next = items[len(items)-1].ID
|
||||
}
|
||||
return items, next, nil
|
||||
}
|
||||
|
||||
// Watch replaces this connection's presence subscription.
|
||||
func (a *App) Watch(_ context.Context, connID port.ConnID, endpointID string, req *protocol.PresenceWatch) error {
|
||||
if req == nil {
|
||||
return &protocol.Error{Code: protocol.CodeBadRequest, Message: "nil request"}
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
sub := watchSub{endpointID: endpointID, all: req.All}
|
||||
if !req.All {
|
||||
sub.ids = make(map[string]struct{}, len(req.IDs))
|
||||
for _, id := range req.IDs {
|
||||
sub.ids[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
a.mu.Lock()
|
||||
a.watches[connID] = sub
|
||||
a.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClearWatch clears subscription on disconnect.
|
||||
func (a *App) ClearWatch(connID port.ConnID) {
|
||||
a.mu.Lock()
|
||||
delete(a.watches, connID)
|
||||
a.mu.Unlock()
|
||||
}
|
||||
|
||||
// SetOnline marks handshake complete.
|
||||
func (a *App) SetOnline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error {
|
||||
if atMs == 0 {
|
||||
atMs = a.nowMs()
|
||||
}
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, atMs, endpointID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.mu.Lock()
|
||||
a.online[endpointID] = onlineEntry{connID: connID, atMs: atMs}
|
||||
a.mu.Unlock()
|
||||
a.notify(ctx, endpointID, true, atMs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetOffline marks disconnect for the current connection.
|
||||
func (a *App) SetOffline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error {
|
||||
if atMs == 0 {
|
||||
atMs = a.nowMs()
|
||||
}
|
||||
a.mu.Lock()
|
||||
cur, ok := a.online[endpointID]
|
||||
if ok && cur.connID == connID {
|
||||
delete(a.online, endpointID)
|
||||
} else if ok {
|
||||
a.mu.Unlock()
|
||||
a.ClearWatch(connID)
|
||||
return nil
|
||||
}
|
||||
a.mu.Unlock()
|
||||
a.ClearWatch(connID)
|
||||
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, atMs, endpointID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.notify(ctx, endpointID, false, atMs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsOnline reports whether the endpoint is online.
|
||||
func (a *App) IsOnline(endpointID string) bool {
|
||||
if a.conns != nil && a.conns.IsOnline(endpointID) {
|
||||
return true
|
||||
}
|
||||
a.mu.Lock()
|
||||
_, ok := a.online[endpointID]
|
||||
a.mu.Unlock()
|
||||
if ok {
|
||||
return true
|
||||
}
|
||||
var onlineSince, offlineSince sql.NullInt64
|
||||
err := a.db.Read.QueryRow(`SELECT online_since, offline_since FROM endpoints WHERE id = ?`, endpointID).
|
||||
Scan(&onlineSince, &offlineSince)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
on, _ := dbOnline(onlineSince, offlineSince)
|
||||
return on
|
||||
}
|
||||
|
||||
// CurrentConn returns the current connection id if any.
|
||||
func (a *App) CurrentConn(endpointID string) (port.ConnID, bool) {
|
||||
if a.conns != nil {
|
||||
if c, ok := a.conns.CurrentConn(endpointID); ok {
|
||||
return c, true
|
||||
}
|
||||
}
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
e, ok := a.online[endpointID]
|
||||
return e.connID, ok
|
||||
}
|
||||
|
||||
func (a *App) resolveOnline(id string, onlineSince, offlineSince sql.NullInt64) (online bool, sinceMs int64) {
|
||||
if a.conns != nil && a.conns.IsOnline(id) {
|
||||
a.mu.Lock()
|
||||
e, ok := a.online[id]
|
||||
a.mu.Unlock()
|
||||
if ok {
|
||||
return true, e.atMs
|
||||
}
|
||||
if onlineSince.Valid {
|
||||
return true, onlineSince.Int64
|
||||
}
|
||||
return true, 0
|
||||
}
|
||||
a.mu.Lock()
|
||||
e, ok := a.online[id]
|
||||
a.mu.Unlock()
|
||||
if ok {
|
||||
return true, e.atMs
|
||||
}
|
||||
return dbOnline(onlineSince, offlineSince)
|
||||
}
|
||||
|
||||
func dbOnline(onlineSince, offlineSince sql.NullInt64) (bool, int64) {
|
||||
if !onlineSince.Valid {
|
||||
return false, 0
|
||||
}
|
||||
if !offlineSince.Valid || onlineSince.Int64 > offlineSince.Int64 {
|
||||
return true, onlineSince.Int64
|
||||
}
|
||||
return false, offlineSince.Int64
|
||||
}
|
||||
|
||||
func (a *App) notify(ctx context.Context, changedID string, online bool, atMs int64) {
|
||||
if a.down == nil {
|
||||
return
|
||||
}
|
||||
frame := protocol.Presence{
|
||||
V: protocol.Version, Type: protocol.TypePresence,
|
||||
ID: changedID, Online: online, AtMs: atMs,
|
||||
}
|
||||
payload, err := encodeFrame(frame)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
a.mu.Lock()
|
||||
subs := make([]watchSub, 0, len(a.watches))
|
||||
for _, s := range a.watches {
|
||||
subs = append(subs, s)
|
||||
}
|
||||
a.mu.Unlock()
|
||||
for _, s := range subs {
|
||||
if !s.all {
|
||||
if _, ok := s.ids[changedID]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
_ = a.down.PublishDown(ctx, s.endpointID, "", payload, port.PublishOpts{QoS: 0})
|
||||
}
|
||||
}
|
||||
|
||||
func encodeFrame(v any) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
if err := protocol.Encode(&buf, v); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
var _ Service = (*App)(nil)
|
||||
@@ -0,0 +1,176 @@
|
||||
package presence_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/presence"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
type memDownlink struct {
|
||||
mu sync.Mutex
|
||||
msgs []downMsg
|
||||
}
|
||||
|
||||
type downMsg struct {
|
||||
endpointID string
|
||||
payload []byte
|
||||
qos byte
|
||||
}
|
||||
|
||||
func (d *memDownlink) PublishDown(_ context.Context, endpointID string, _ port.ConnID, payload []byte, opts port.PublishOpts) error {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
cp := append([]byte(nil), payload...)
|
||||
d.msgs = append(d.msgs, downMsg{endpointID: endpointID, payload: cp, qos: opts.QoS})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *memDownlink) take() []downMsg {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
out := d.msgs
|
||||
d.msgs = nil
|
||||
return out
|
||||
}
|
||||
|
||||
func openPresence(t *testing.T) (*presence.App, *store.DB, *memDownlink) {
|
||||
t.Helper()
|
||||
db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
down := &memDownlink{}
|
||||
fixed := time.UnixMilli(1_700_000_000_000)
|
||||
app := presence.New(presence.Config{
|
||||
DB: db, Downlink: down, Now: func() time.Time { return fixed },
|
||||
})
|
||||
return app, db, down
|
||||
}
|
||||
|
||||
func insertEP(t *testing.T, db *store.DB, id, name string) {
|
||||
t.Helper()
|
||||
err := 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(?,?,?,?,0,0,1,?)`, id, name, "stub$login", nil, 1_700_000_000_000)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF03PresenceAndDirectory(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, _ := openPresence(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "Alice")
|
||||
insertEP(t, db, "bob", "Bob")
|
||||
insertEP(t, db, "carol", "Carol")
|
||||
|
||||
items, err := app.Get(ctx, []string{"alice", "nobody"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(items) != 2 || items[0].Online || !items[1].NotFound {
|
||||
t.Fatalf("%+v", items)
|
||||
}
|
||||
|
||||
if err = app.SetOnline(ctx, "alice", "c1", 1_700_000_000_100); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
items, _ = app.Get(ctx, []string{"alice"})
|
||||
if !items[0].Online || items[0].SinceMs != 1_700_000_000_100 {
|
||||
t.Fatalf("%+v", items[0])
|
||||
}
|
||||
|
||||
// 正常断开后查为离线
|
||||
if err = app.SetOffline(ctx, "alice", "c1", 1_700_000_000_200); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
items, _ = app.Get(ctx, []string{"alice"})
|
||||
if items[0].Online {
|
||||
t.Fatal("should be offline")
|
||||
}
|
||||
|
||||
dir, next, err := app.Directory(ctx, &protocol.DirectoryList{
|
||||
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "1", Limit: 2,
|
||||
})
|
||||
if err != nil || len(dir) != 2 || next == "" {
|
||||
t.Fatalf("dir=%+v next=%q err=%v", dir, next, err)
|
||||
}
|
||||
dir2, next2, err := app.Directory(ctx, &protocol.DirectoryList{
|
||||
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "2", Cursor: next, Limit: 10,
|
||||
})
|
||||
if err != nil || len(dir2) != 1 || next2 != "" {
|
||||
t.Fatalf("dir2=%+v next=%q", dir2, next2)
|
||||
}
|
||||
|
||||
// query:编号前缀
|
||||
q, _, err := app.Directory(ctx, &protocol.DirectoryList{
|
||||
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "3", Query: "bo", Limit: 10,
|
||||
})
|
||||
if err != nil || len(q) != 1 || q[0].ID != "bob" {
|
||||
t.Fatalf("%+v err=%v", q, err)
|
||||
}
|
||||
// 名称包含不区分大小写
|
||||
q, _, err = app.Directory(ctx, &protocol.DirectoryList{
|
||||
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "4", Query: "car", Limit: 10,
|
||||
})
|
||||
if err != nil || len(q) != 1 || q[0].ID != "carol" {
|
||||
t.Fatalf("%+v", q)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF04PresenceWatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
app, db, down := openPresence(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", "A")
|
||||
insertEP(t, db, "bob", "B")
|
||||
insertEP(t, db, "carol", "C")
|
||||
|
||||
if err := app.Watch(ctx, "conn-sub", "watcher", &protocol.PresenceWatch{
|
||||
V: protocol.Version, Type: protocol.TypePresenceWatch, RID: "1", IDs: []string{"alice"},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = app.SetOnline(ctx, "alice", "ca", 100)
|
||||
_ = app.SetOffline(ctx, "alice", "ca", 200)
|
||||
_ = app.SetOnline(ctx, "bob", "cb", 300) // 未订阅
|
||||
msgs := down.take()
|
||||
if len(msgs) != 2 {
|
||||
t.Fatalf("want 2 presence for alice got %d", len(msgs))
|
||||
}
|
||||
for _, m := range msgs {
|
||||
if m.endpointID != "watcher" || m.qos != 0 {
|
||||
t.Fatalf("%+v", m)
|
||||
}
|
||||
}
|
||||
|
||||
// 再次 Watch 覆盖;断线清空
|
||||
if err := app.Watch(ctx, "conn-sub", "watcher", &protocol.PresenceWatch{
|
||||
V: protocol.Version, Type: protocol.TypePresenceWatch, RID: "2", All: true,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
down.take()
|
||||
_ = app.SetOnline(ctx, "carol", "cc", 400)
|
||||
if n := len(down.take()); n != 1 {
|
||||
t.Fatalf("all watch got %d", n)
|
||||
}
|
||||
app.ClearWatch("conn-sub")
|
||||
_ = app.SetOffline(ctx, "carol", "cc", 500)
|
||||
if n := len(down.take()); n != 0 {
|
||||
t.Fatalf("after clear got %d", n)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -9,6 +9,7 @@ import (
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
mqtt "github.com/mochi-mqtt/server/v2"
|
||||
@@ -85,18 +86,22 @@ type Broker struct {
|
||||
}
|
||||
|
||||
type connState struct {
|
||||
connID port.ConnID
|
||||
endpointID string
|
||||
transport port.Transport
|
||||
remoteIP string
|
||||
client *mqtt.Client
|
||||
maxPacketSize uint32
|
||||
maxRecvBytes int
|
||||
authOK bool
|
||||
authErr error
|
||||
sessionToken string
|
||||
largeHeld int
|
||||
mu sync.Mutex
|
||||
connID port.ConnID
|
||||
endpointID string
|
||||
transport port.Transport
|
||||
remoteIP string
|
||||
client *mqtt.Client
|
||||
maxPacketSize uint32
|
||||
maxRecvBytes int
|
||||
authOK bool
|
||||
authErr error
|
||||
sessionToken string
|
||||
handshook bool
|
||||
subscribedDown bool
|
||||
largeHeld int
|
||||
mu sync.Mutex
|
||||
|
||||
handshakeTimer *time.Timer
|
||||
}
|
||||
|
||||
// New 创建并 Serve mochi(无监听器)。
|
||||
@@ -265,7 +270,12 @@ func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.Co
|
||||
case port.DisconnectKicked, port.DisconnectFatal:
|
||||
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 {
|
||||
@@ -362,6 +372,100 @@ func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) {
|
||||
}, 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) {
|
||||
b.queuesMu.Lock()
|
||||
q, ok := b.queues[endpointID]
|
||||
|
||||
@@ -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:])
|
||||
}
|
||||
@@ -26,13 +26,15 @@ func (h *nixHook) Provides(b byte) bool {
|
||||
mqtt.OnSessionEstablished,
|
||||
mqtt.OnDisconnect,
|
||||
mqtt.OnQosComplete,
|
||||
mqtt.OnSubscribed,
|
||||
}, []byte{b})
|
||||
}
|
||||
|
||||
func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
||||
endpointID := string(pk.Connect.Username)
|
||||
clientID := pk.Connect.ClientIdentifier
|
||||
if endpointID == "" {
|
||||
endpointID = pk.Connect.ClientIdentifier
|
||||
endpointID = clientID
|
||||
}
|
||||
remoteIP := remoteIPOf(cl)
|
||||
|
||||
@@ -45,6 +47,13 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
||||
maxPacketSize: pk.Properties.MaximumPacketSize,
|
||||
}
|
||||
|
||||
// ClientID、Username 都必须等于端编号
|
||||
if clientID == "" || endpointID == "" || clientID != endpointID {
|
||||
st.authOK = false
|
||||
h.rememberPending(cl, st)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive
|
||||
ka := pk.Connect.Keepalive
|
||||
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) {
|
||||
// InlineClient 的 PublishDown 走 InjectPacket → OnPublish;必须放行才能分发给订阅者。
|
||||
if cl != nil && cl.Net.Inline {
|
||||
return pk, nil
|
||||
}
|
||||
h.b.connsMu.RLock()
|
||||
st := h.b.byClient[cl]
|
||||
h.b.connsMu.RUnlock()
|
||||
@@ -122,6 +135,28 @@ func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet,
|
||||
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) {
|
||||
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()
|
||||
st := h.b.byClient[cl]
|
||||
delete(h.b.byClient, cl)
|
||||
isCurrent := false
|
||||
if st != nil && h.b.current[st.endpointID] == st {
|
||||
delete(h.b.current, st.endpointID)
|
||||
isCurrent = true
|
||||
}
|
||||
h.b.connsMu.Unlock()
|
||||
if st == nil {
|
||||
return
|
||||
}
|
||||
h.b.releaseAllLarge(st)
|
||||
h.b.cancelHandshakeDeadline(st.endpointID, st.connID)
|
||||
|
||||
reason := port.DisconnectNormal
|
||||
if err != nil {
|
||||
@@ -179,6 +217,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
||||
SessionToken: st.sessionToken,
|
||||
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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user