Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5aa45b44e0 | ||
|
|
786fe3590b | ||
|
|
e95062d9eb | ||
|
|
af8278d2a9 | ||
|
|
ac90495137 |
@@ -67,6 +67,10 @@ func (u *appUplink) HandleUplink(ctx context.Context, conn port.ConnInfo, payloa
|
||||
u.replyErr(ctx, conn, peekRID(payload), protocol.CodeBadRequest, err.Error())
|
||||
return nil
|
||||
}
|
||||
if !uplinkRateExempt(frame) && u.msg != nil && !u.msg.AllowRequest(conn.EndpointID) {
|
||||
u.replyErr(ctx, conn, peekRID(payload), protocol.CodeRateLimited, "request rate exceeded")
|
||||
return nil
|
||||
}
|
||||
|
||||
rid, data, callErr := u.dispatch(ctx, conn, frame)
|
||||
if callErr != nil {
|
||||
@@ -77,6 +81,15 @@ func (u *appUplink) HandleUplink(ctx context.Context, conn port.ConnInfo, payloa
|
||||
return nil
|
||||
}
|
||||
|
||||
func uplinkRateExempt(frame any) bool {
|
||||
switch frame.(type) {
|
||||
case *protocol.Ack, *protocol.ReceiptAck:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (u *appUplink) dispatch(ctx context.Context, conn port.ConnInfo, frame any) (rid string, data any, err error) {
|
||||
switch f := frame.(type) {
|
||||
case *protocol.Send:
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"io"
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"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 TestHandleUplinkRateLimitStatusAndAckExempt(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
nowMs := int64(1_700_000_000_000)
|
||||
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('alice','alice','stub$login',NULL,0,0,1,?)`, nowMs)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
lim := message.LimitsFromFullConfig(config.Default())
|
||||
lim.RequestsPerSecond = 50
|
||||
lim.RequestBurst = 100
|
||||
app := message.New(db, lim, auth.NewStubHashPool(),
|
||||
message.WithNow(func() time.Time { return time.UnixMilli(nowMs) }),
|
||||
)
|
||||
down := &message.RecordingDownlink{}
|
||||
conns := message.NewMemoryConns()
|
||||
conns.Set("alice", message.LiveConn{ConnID: "c1"})
|
||||
u := &appUplink{msg: app, conns: conns, down: down, log: slog.New(slog.NewTextHandler(io.Discard, nil))}
|
||||
conn := port.ConnInfo{EndpointID: "alice", ConnID: "c1"}
|
||||
ctx := context.Background()
|
||||
|
||||
ackPayload, err := protocol.Marshal(&protocol.Ack{
|
||||
V: protocol.Version, Type: protocol.TypeAck, RID: "a", From: "alice", ID: "missing",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := 0; i < 150; i++ {
|
||||
if e := u.HandleUplink(ctx, conn, ackPayload); e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
}
|
||||
if n := countRespCode(down, protocol.CodeRateLimited); n != 0 {
|
||||
t.Fatalf("ack should not count, rate_limited=%d", n)
|
||||
}
|
||||
|
||||
statusPayload, err := protocol.Marshal(&protocol.Status{
|
||||
V: protocol.Version, Type: protocol.TypeStatus, RID: "s", ID: "no-such",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := 0; i < 150; i++ {
|
||||
if e := u.HandleUplink(ctx, conn, statusPayload); e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
}
|
||||
limited := countRespCode(down, protocol.CodeRateLimited)
|
||||
if limited != 50 {
|
||||
t.Fatalf("status rate_limited=%d want 50 (burst 100 of 150)", limited)
|
||||
}
|
||||
}
|
||||
|
||||
func countRespCode(down *message.RecordingDownlink, code string) int {
|
||||
n := 0
|
||||
for _, p := range down.Snapshots() {
|
||||
var resp protocol.Resp
|
||||
if err := protocol.Unmarshal(p.Payload, &resp); err != nil {
|
||||
continue
|
||||
}
|
||||
if !resp.OK && resp.Error != nil && resp.Error.Code == code {
|
||||
n++
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
+69
-3
@@ -385,10 +385,10 @@
|
||||
|
||||
2. **请求频率突发容量写死为 100**
|
||||
- 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。
|
||||
- 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 不计入桶(与 6.10 一致)。
|
||||
- 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。`message.App.AllowRequest` 导出同一令牌桶;`HandleUplink` 在分发前对 ack/receipt_ack 以外的帧调用。`Submit` 不再单独扣桶,避免 send 计两次。
|
||||
- 原因:配置无独立 burst 字段。
|
||||
- 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。
|
||||
- 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。
|
||||
- 备选方案:配置增加 `request_burst`。
|
||||
- 影响:改 `requests_per_second` 不改突发;非 send 请求也受同一桶限制。
|
||||
|
||||
3. **未接线 `cmd/nixmsg`**
|
||||
- 原条款:可替换 T0.4 假实现。
|
||||
@@ -448,6 +448,42 @@
|
||||
- 备选方案:总控在 `protocol` 增类型。
|
||||
- 影响:接线编码 `resp.data` 时直接 Marshal 该 map 即可。
|
||||
|
||||
### 复审修复 C-04
|
||||
|
||||
1. **退群/踢人/解散/停用/删除作废投递走统一终态函数**
|
||||
- 原条款:DEVELOPMENT 7.6 投递进入 rejected 时写回执;没有 pending 时收尾 completed、删正文;保留 0 天同一事务删行。PRD F14/F18。
|
||||
- 实际做法:message 导出 `RejectPendingTx` / `TryFinalizeTx` / `FinalizeMessageTx`。group `void.go` 与 identity `lifecycle.go` 的作废/收尾改为调用它们,去掉复制 SQL。`sender_disabled`/`sender_deleted` 仍不写回执。`CleanupOnce` 分批收尾「dispatched 且无 pending」的卡住消息。作废路径未接线 `record_retention_days` 时按默认 7 天收尾(不在同一事务删行);保留 0 天由 message 自己的 finalize 覆盖。不改 group `emit`。
|
||||
- 原因:原先 group 只改投递状态,identity 收尾但不写回执,最后一个 pending 被作废后消息永远停在 dispatched。
|
||||
- 备选方案:在 group/identity 各自补写回执与收尾(继续分叉)。
|
||||
- 影响:退群/解散/停用后发送方可收到 rejected 回执,配额释放,正文删除。
|
||||
|
||||
### 复审修复 C-05
|
||||
|
||||
1. **每端请求限速覆盖非 send 帧**
|
||||
- 原条款:PRD F05 / DEVELOPMENT 6.10:除 ack、receipt_ack 外共用一个桶,默认每秒 50、突发 100。
|
||||
- 实际做法:message 导出 `AllowRequest`;`cmd/nixmsg/uplink.go` 的 `HandleUplink` 解码后、分发前检查;超限回 `rate_limited`。去掉 `Submit` 内扣桶。不改 uplink 生命周期与 `publishResp`。
|
||||
- 原因:原先只有 send 限速,unlock/status/目录/群等可打满哈希池与读库。
|
||||
- 备选方案:把桶挪到 broker 层(B-09 范围)。
|
||||
- 影响:开放注册后的非 send 请求也计入配额;直接调 `Submit` 的单测不再覆盖限速。
|
||||
|
||||
### 复审修复 C-06
|
||||
|
||||
1. **推送 meta 数字用 UseNumber 解码**
|
||||
- 原条款:PRD F07 / D11 自定义键值送达应与提交一致。
|
||||
- 实际做法:`decodeMetaJSON` 改用 `protocol.Unmarshal`(`UseNumber`),超过 2^53 的整数以 `json.Number` 保留原文再编码进推送帧。不改 `Msg.Meta` 类型与协议包。
|
||||
- 原因:标准 `json.Unmarshal` 把数字变成 float64,雪花 ID 会被改掉。
|
||||
- 备选方案:`Meta` 改为 `json.RawMessage` 原样输出(需改 protocol,牵动 SDK)。
|
||||
- 影响:仅推送路径;入库仍是提交时的规范 JSON。
|
||||
|
||||
### 复审修复 C-07
|
||||
|
||||
1. **提交校验、停用检查、入群时间过滤、保留期按完成时刻**
|
||||
- 原条款:DEVELOPMENT 6.2 ttl/定时上限;PRD F01 停用后不能再发;F06 发送时刻之后入群的端收不到;F18 记录保留从完成起算。
|
||||
- 实际做法:`keep` 且 `ttl_seconds<=0` 回 `bad_request`;`delay_ms` 先与 `max_schedule_seconds*1000` 比较再加法。`Send.Validate` 同步(`protocol.Limits` 增加可选 MaxTTL/MaxSchedule,0 表示不查上限)。写事务内检查发送方 `enabled`,停用回 `unauthorized`。群分发 `joined_at <= send_at`。未做 C-03 的 `completed_at` 列,清理暂用 `MAX(deliveries.updated_at)` 否则 `send_at` 近似完成时刻。不改对话密码锁键语义。
|
||||
- 原因:ttl=0/负数、delay 溢出、停用窗口内仍能提交、晚入群仍能收到、按 created_at 清理会误删长定时/长保留消息。
|
||||
- 备选方案:等 C-03 迁移后改用 `completed_at`;发送方停用改用 `endpoint_disabled`(与目标停用混用)。
|
||||
- 影响:发送方停用错误码为 `unauthorized`;保留期口径在 C-03 合入前对无投递的 scheduled 作废行用 `send_at` 近似。
|
||||
|
||||
## 身份 I
|
||||
|
||||
### I1 2026-09-30
|
||||
@@ -582,6 +618,36 @@
|
||||
- 备选方案:仅按 joined_at。
|
||||
- 影响:同毫秒加入时编号小者优先。
|
||||
|
||||
### 复审修复 U-02
|
||||
|
||||
1. **群写操作在同一写事务内复核**
|
||||
- 原条款:PRD F16 群主同时是成员、停用端不能加入、成员上限、新群收不到旧群消息;issue #40。
|
||||
- 实际做法:加人/踢人/退群/转让/改名/解散在 `Queue.Do` 内重读群主、成员关系和成员数;加人再复核目标端 `enabled`。对话密码(argon2)仍在事务外,事务里只做廉价 SQL。`INSERT OR IGNORE` 改为先复核再 `INSERT`;外键失败按 `not_found`。不改 `emit`,不改 message `RejectPendingTx`。
|
||||
- 原因:读后写会在解散后留下孤儿成员、并发加人超过上限、转让后群主不在成员里。
|
||||
- 备选方案:只靠外键、事务外校验(否决,无法给出原错误码)。
|
||||
- 影响:加人与解散并发时整次加人返回 `not_found`,不写孤儿行。
|
||||
|
||||
2. **建群/加人先去重再截断,单请求成员数设上限**
|
||||
- 原条款:部分失败仍建群;成员上限。
|
||||
- 实际做法:先去掉自己和重复编号,再按剩余名额截断,超出记 `group_full`,然后才做密码校验。整表请求成员数超过 `2*max_group_members`(至少 256)回 `bad_request`。原先「校验通过人数加群主超上限则整次建群失败」改为截断后仍建群。
|
||||
- 原因:重复编号会校验两次并在插入时主键冲突,客户端按 `busy` 一直重试;一个请求可带上万个成员打满哈希池。
|
||||
- 备选方案:协议层去重(禁止改 protocol)。
|
||||
- 影响:带重复成员的建群会成功且只留一条;超上限的多余成员在 `failed` 里而不是整次失败。
|
||||
|
||||
3. **后台建群校验群主并补推 `member_added`**
|
||||
- 原条款:群主必须是已启用的端。
|
||||
- 实际做法:`createAdmin` 校验群主编号格式、存在且 `enabled`;成员去重;建成后按与客户端建群相同方式 `emit` `member_added`。群主不存在 `invalid_target`,已停用 `endpoint_disabled`,格式非法 `bad_request`。
|
||||
- 原因:原先可不存在/已停用的编号当群主,成员也不去重,也不推事件。
|
||||
- 备选方案:由 admin HTTP 层预校验(仍会与写路径竞态)。
|
||||
- 影响:后台建群失败码与加人目标错误码对齐。
|
||||
|
||||
4. **可选迁移 `0003_group_members_fk.sql`**
|
||||
- 原条款:TASKS 4.2 改表加新文件,rebase 时取当时最大号加一;issue 写「排在 C-03 的 0003 之后」。
|
||||
- 实际做法:本分支基于 C-04,当时最大号 0002,按 TASKS 4.2 用 0003:重建 `group_members` 并 `REFERENCES groups(id) ON DELETE CASCADE`。C-03 尚未合入。
|
||||
- 原因:无外键时同编号新建群会继承旧孤儿成员。
|
||||
- 备选方案:等 C-03 占用 0003 后再用 0004(rebase 时改号)。
|
||||
- 影响:若 C-03 先合入并占用 0003,本文件 rebase 时改号。
|
||||
|
||||
## 后台接口 A
|
||||
|
||||
### A1 2026-09-30
|
||||
|
||||
+352
-120
@@ -6,6 +6,7 @@ import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
@@ -27,6 +28,14 @@ const (
|
||||
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||
)
|
||||
|
||||
func (a *App) memberRequestCap() int {
|
||||
n := a.maxMem * 2
|
||||
if n < 256 {
|
||||
n = 256
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// TalkGate checks talk password when adding members (implemented by identity).
|
||||
type TalkGate interface {
|
||||
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
|
||||
@@ -105,6 +114,9 @@ func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCre
|
||||
if !protocol.ValidEndpointID(actorID) {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid actor")
|
||||
}
|
||||
if err := a.rejectOversizedMemberList(len(req.Members)); err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
gid := req.ID
|
||||
if gid == "" {
|
||||
var genErr error
|
||||
@@ -115,12 +127,17 @@ func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCre
|
||||
}
|
||||
now := a.nowMs()
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0, len(req.Members))
|
||||
|
||||
for _, m := range req.Members {
|
||||
if m.ID == actorID {
|
||||
continue
|
||||
uniq := dedupeMemberIns(req.Members, actorID)
|
||||
room := a.maxMem - 1
|
||||
if room < 0 {
|
||||
room = 0
|
||||
}
|
||||
toCheck, overflow := splitMemberIns(uniq, room)
|
||||
for _, m := range overflow {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
||||
}
|
||||
added := make([]string, 0, len(toCheck))
|
||||
for _, m := range toCheck {
|
||||
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
@@ -128,10 +145,7 @@ func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCre
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
|
||||
if 1+len(added) > a.maxMem {
|
||||
return CreateResult{}, errCode(protocol.CodeGroupFull, "group full")
|
||||
}
|
||||
|
||||
var inserted []string
|
||||
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)
|
||||
@@ -148,15 +162,19 @@ func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCre
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, actorID, now); e != nil {
|
||||
if e := insertMemberTx(tx, gid, actorID, now); e != nil {
|
||||
return e
|
||||
}
|
||||
inserted = inserted[:0]
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, id, now); e != nil {
|
||||
if e := endpointCheckTx(tx, id); e != nil {
|
||||
failed = append(failed, MemberFail{ID: id, Code: failCode(e)})
|
||||
continue
|
||||
}
|
||||
if e := insertMemberTx(tx, gid, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
inserted = append(inserted, id)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
@@ -164,8 +182,9 @@ func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCre
|
||||
return CreateResult{}, err
|
||||
}
|
||||
|
||||
for _, id := range added {
|
||||
a.emit(ctx, append([]string{actorID}, added...), gid, eventMemberAdded, id, now)
|
||||
notify := append([]string{actorID}, inserted...)
|
||||
for _, id := range inserted {
|
||||
a.emit(ctx, notify, gid, eventMemberAdded, id, now)
|
||||
}
|
||||
return CreateResult{ID: gid, Name: req.Name, OwnerID: actorID, Failed: failed}, nil
|
||||
}
|
||||
@@ -178,6 +197,9 @@ func (a *App) Add(ctx context.Context, actorID string, req *protocol.GroupAdd) (
|
||||
if err := req.Validate(); err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
if err := a.rejectOversizedMemberList(len(req.Members)); err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
@@ -187,17 +209,23 @@ func (a *App) Add(ctx context.Context, actorID string, req *protocol.GroupAdd) (
|
||||
}
|
||||
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
now := a.nowMs()
|
||||
|
||||
for _, m := range req.Members {
|
||||
if contains(members, m.ID) {
|
||||
uniq := dedupeMemberIns(req.Members, actorID)
|
||||
already := memberSet(members)
|
||||
candidates := make([]protocol.GroupMemberIn, 0, len(uniq))
|
||||
for _, m := range uniq {
|
||||
if _, ok := already[m.ID]; ok {
|
||||
continue
|
||||
}
|
||||
if len(members)+len(added) >= a.maxMem {
|
||||
candidates = append(candidates, m)
|
||||
}
|
||||
room := a.maxMem - len(members)
|
||||
toCheck, overflow := splitMemberIns(candidates, room)
|
||||
for _, m := range overflow {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
added := make([]string, 0, len(toCheck))
|
||||
for _, m := range toCheck {
|
||||
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
@@ -205,23 +233,53 @@ func (a *App) Add(ctx context.Context, actorID string, req *protocol.GroupAdd) (
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
|
||||
if len(added) > 0 {
|
||||
if len(added) == 0 {
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
|
||||
var inserted []string
|
||||
var notify []string
|
||||
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 {
|
||||
curOwner, curMembers, e := loadGroupTx(tx, req.GroupID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if curOwner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
present := memberSet(curMembers)
|
||||
count := len(curMembers)
|
||||
inserted = inserted[:0]
|
||||
for _, id := range added {
|
||||
if _, ok := present[id]; ok {
|
||||
continue
|
||||
}
|
||||
if count >= a.maxMem {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
if checkErr := endpointCheckTx(tx, id); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: id, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
if insErr := insertMemberTx(tx, req.GroupID, id, now); insErr != nil {
|
||||
return insErr
|
||||
}
|
||||
present[id] = struct{}{}
|
||||
count++
|
||||
inserted = append(inserted, id)
|
||||
}
|
||||
notify = make([]string, 0, count)
|
||||
for id := range present {
|
||||
notify = append(notify, id)
|
||||
}
|
||||
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)
|
||||
}
|
||||
for _, id := range inserted {
|
||||
a.emit(ctx, notify, req.GroupID, eventMemberAdded, id, now)
|
||||
}
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
@@ -234,9 +292,13 @@ func (a *App) Remove(ctx context.Context, actorID string, req *protocol.GroupRem
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
var notify []string
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
owner, members, e := loadGroupTx(tx, req.GroupID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
@@ -247,21 +309,20 @@ func (a *App) Remove(ctx context.Context, actorID string, req *protocol.GroupRem
|
||||
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 e := voidMemberDeliveriesTx(tx, req.GroupID, req.EndpointID, reasonLeftGroup, now, &revokes); e != nil {
|
||||
return e
|
||||
}
|
||||
notify = append(without(members, req.EndpointID), req.EndpointID)
|
||||
return nil
|
||||
})
|
||||
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
|
||||
}
|
||||
@@ -274,9 +335,13 @@ func (a *App) Leave(ctx context.Context, actorID string, req *protocol.GroupLeav
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
var notify []string
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
owner, members, e := loadGroupTx(tx, req.GroupID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if !contains(members, actorID) {
|
||||
return errCode(protocol.CodeNotMember, "not a member")
|
||||
@@ -284,21 +349,20 @@ func (a *App) Leave(ctx context.Context, actorID string, req *protocol.GroupLeav
|
||||
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 e := voidMemberDeliveriesTx(tx, req.GroupID, actorID, reasonLeftGroup, now, &revokes); e != nil {
|
||||
return e
|
||||
}
|
||||
notify = append(without(members, actorID), actorID)
|
||||
return nil
|
||||
})
|
||||
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
|
||||
}
|
||||
@@ -311,20 +375,24 @@ func (a *App) Transfer(ctx context.Context, actorID string, req *protocol.GroupT
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
now := a.nowMs()
|
||||
var members []string
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
owner, cur, e := loadGroupTx(tx, req.GroupID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
if !contains(members, req.EndpointID) {
|
||||
if !contains(cur, 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)
|
||||
if _, e := tx.Exec(`UPDATE groups SET owner_id = ? WHERE id = ?`, req.EndpointID, req.GroupID); e != nil {
|
||||
return e
|
||||
}
|
||||
members = cur
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -341,17 +409,21 @@ func (a *App) Rename(ctx context.Context, actorID string, req *protocol.GroupRen
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
now := a.nowMs()
|
||||
var members []string
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
owner, cur, e := loadGroupTx(tx, req.GroupID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
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)
|
||||
if _, e := tx.Exec(`UPDATE groups SET name = ? WHERE id = ?`, req.Name, req.GroupID); e != nil {
|
||||
return e
|
||||
}
|
||||
members = cur
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -368,24 +440,28 @@ func (a *App) Dissolve(ctx context.Context, actorID string, req *protocol.GroupD
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
var members []string
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
owner, cur, e := loadGroupTx(tx, req.GroupID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
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)
|
||||
if _, e := tx.Exec(`DELETE FROM groups WHERE id = ?`, req.GroupID); e != nil {
|
||||
return e
|
||||
}
|
||||
members = cur
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -512,59 +588,80 @@ func (a *App) AdminCreate(ctx context.Context, name, ownerID string, memberIDs [
|
||||
|
||||
// AdminAddMembers adds members without talk-password checks.
|
||||
func (a *App) AdminAddMembers(ctx context.Context, groupID string, memberIDs []string) (AddResult, error) {
|
||||
if err := a.rejectOversizedMemberList(len(memberIDs)); err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
_, members, err := a.loadGroup(ctx, groupID)
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
uniq := dedupeIDs(memberIDs, "")
|
||||
already := memberSet(members)
|
||||
candidates := make([]string, 0, len(uniq))
|
||||
for _, id := range uniq {
|
||||
if _, ok := already[id]; ok {
|
||||
continue
|
||||
}
|
||||
candidates = append(candidates, id)
|
||||
}
|
||||
room := a.maxMem - len(members)
|
||||
toAdd, overflow := splitIDs(candidates, room)
|
||||
for _, id := range overflow {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
||||
}
|
||||
if len(toAdd) == 0 {
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
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
|
||||
}
|
||||
var inserted []string
|
||||
var notify []string
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, curMembers, e := loadGroupTx(tx, groupID)
|
||||
if e != nil {
|
||||
return AddResult{}, e
|
||||
return e
|
||||
}
|
||||
if enabled == 0 {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeEndpointDisabled})
|
||||
present := memberSet(curMembers)
|
||||
count := len(curMembers)
|
||||
inserted = inserted[:0]
|
||||
for _, id := range toAdd {
|
||||
if _, ok := present[id]; ok {
|
||||
continue
|
||||
}
|
||||
if len(members)+len(added) >= a.maxMem {
|
||||
if count >= a.maxMem {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
added = append(added, id)
|
||||
if checkErr := endpointCheckTx(tx, id); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: id, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
if len(added) == 0 {
|
||||
return AddResult{Failed: failed}, nil
|
||||
if insErr := insertMemberTx(tx, groupID, id, now); insErr != nil {
|
||||
return insErr
|
||||
}
|
||||
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
|
||||
present[id] = struct{}{}
|
||||
count++
|
||||
inserted = append(inserted, id)
|
||||
}
|
||||
notify = make([]string, 0, count)
|
||||
for id := range present {
|
||||
notify = append(notify, id)
|
||||
}
|
||||
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)
|
||||
for _, id := range inserted {
|
||||
a.emit(ctx, notify, 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 !protocol.ValidEndpointID(ownerID) {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid owner")
|
||||
}
|
||||
if gid == "" {
|
||||
var genErr error
|
||||
gid, genErr = generateGroupID()
|
||||
@@ -575,32 +672,29 @@ func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, member
|
||||
if !protocol.ValidName(name) || name == "" {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid name")
|
||||
}
|
||||
if err := a.rejectOversizedMemberList(len(members)); err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
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")
|
||||
uniq := dedupeMemberIns(members, ownerID)
|
||||
room := a.maxMem - 1
|
||||
toAdd, overflow := splitMemberIns(uniq, room)
|
||||
for _, m := range overflow {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
||||
}
|
||||
|
||||
var inserted []string
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if e := endpointCheckTx(tx, ownerID); e != nil {
|
||||
if protoCode(e) == protocol.CodeInvalidTarget {
|
||||
return errCode(protocol.CodeInvalidTarget, "owner not found")
|
||||
}
|
||||
if protoCode(e) == protocol.CodeEndpointDisabled {
|
||||
return errCode(protocol.CodeEndpointDisabled, "owner disabled")
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
gid, name, ownerID, now); e != nil {
|
||||
if isUnique(e) {
|
||||
@@ -608,21 +702,29 @@ func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, member
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, ownerID, now); e != nil {
|
||||
if e := insertMemberTx(tx, 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 {
|
||||
inserted = inserted[:0]
|
||||
for _, m := range toAdd {
|
||||
if checkErr := endpointCheckTx(tx, m.ID); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
if e := insertMemberTx(tx, gid, m.ID, now); e != nil {
|
||||
return e
|
||||
}
|
||||
inserted = append(inserted, m.ID)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
notify := append([]string{ownerID}, inserted...)
|
||||
for _, id := range inserted {
|
||||
a.emit(ctx, notify, gid, eventMemberAdded, id, now)
|
||||
}
|
||||
return CreateResult{ID: gid, Name: name, OwnerID: ownerID, Failed: failed}, nil
|
||||
}
|
||||
|
||||
@@ -633,6 +735,136 @@ func (a *App) checkAddMember(ctx context.Context, actorID, targetID, talkPasswor
|
||||
return a.talk.CheckTalkPasswordForJoin(ctx, actorID, targetID, talkPassword, a.remoteIP)
|
||||
}
|
||||
|
||||
func (a *App) rejectOversizedMemberList(n int) error {
|
||||
if n > a.memberRequestCap() {
|
||||
return errCode(protocol.CodeBadRequest, "too many members")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadGroupTx(tx *sql.Tx, groupID string) (owner string, members []string, err error) {
|
||||
err = tx.QueryRow(`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 := tx.Query(`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 endpointCheckTx(tx *sql.Tx, id string) error {
|
||||
var enabled int
|
||||
err := tx.QueryRow(`SELECT enabled FROM endpoints WHERE id = ?`, id).Scan(&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")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func insertMemberTx(tx *sql.Tx, groupID, endpointID string, now int64) error {
|
||||
_, err := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
groupID, endpointID, now)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if isForeignKey(err) {
|
||||
return errCode(protocol.CodeNotFound, "group not found")
|
||||
}
|
||||
if isUnique(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func isForeignKey(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "foreign key")
|
||||
}
|
||||
|
||||
func dedupeMemberIns(members []protocol.GroupMemberIn, skipID string) []protocol.GroupMemberIn {
|
||||
seen := make(map[string]struct{}, len(members)+1)
|
||||
if skipID != "" {
|
||||
seen[skipID] = struct{}{}
|
||||
}
|
||||
out := make([]protocol.GroupMemberIn, 0, len(members))
|
||||
for _, m := range members {
|
||||
if _, ok := seen[m.ID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[m.ID] = struct{}{}
|
||||
out = append(out, m)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func splitMemberIns(members []protocol.GroupMemberIn, room int) (keep, overflow []protocol.GroupMemberIn) {
|
||||
if room < 0 {
|
||||
room = 0
|
||||
}
|
||||
if len(members) <= room {
|
||||
return members, nil
|
||||
}
|
||||
return members[:room], members[room:]
|
||||
}
|
||||
|
||||
func dedupeIDs(ids []string, skipID string) []string {
|
||||
seen := make(map[string]struct{}, len(ids)+1)
|
||||
if skipID != "" {
|
||||
seen[skipID] = struct{}{}
|
||||
}
|
||||
out := make([]string, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
out = append(out, id)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func splitIDs(ids []string, room int) (keep, overflow []string) {
|
||||
if room < 0 {
|
||||
room = 0
|
||||
}
|
||||
if len(ids) <= room {
|
||||
return ids, nil
|
||||
}
|
||||
return ids[:room], ids[room:]
|
||||
}
|
||||
|
||||
func memberSet(ss []string) map[string]struct{} {
|
||||
m := make(map[string]struct{}, len(ss))
|
||||
for _, s := range ss {
|
||||
m[s] = struct{}{}
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
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) {
|
||||
|
||||
@@ -461,3 +461,270 @@ func TestGroupTransferRenameListGet(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLeaveLastPendingFinalizesAndReceipt(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, _, msgApp, 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: "OnlyBob", Members: []protocol.GroupMemberIn{{ID: "bob"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ttl := int64(3600)
|
||||
_, err = msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "s", ID: "keep1",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||
Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = gApp.Leave(ctx, "bob", &protocol.GroupLeave{
|
||||
V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "2", GroupID: created.ID,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var state, reason string
|
||||
if err = db.Read.QueryRow(`SELECT state, reason FROM messages WHERE id='keep1'`).Scan(&state, &reason); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if state != message.StateCompleted {
|
||||
t.Fatalf("state=%s want completed", state)
|
||||
}
|
||||
var bodies int
|
||||
if err = db.Read.QueryRow(`SELECT COUNT(*) FROM message_bodies b JOIN messages m ON m.seq=b.seq WHERE m.id='keep1'`).Scan(&bodies); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bodies != 0 {
|
||||
t.Fatalf("body still present: %d", bodies)
|
||||
}
|
||||
var rState, rReason, rEP string
|
||||
if err = db.Read.QueryRow(`
|
||||
SELECT state, reason, endpoint_id FROM receipts WHERE sender_id='alice' AND msg_id='keep1'`).Scan(&rState, &rReason, &rEP); err != nil {
|
||||
t.Fatalf("receipt: %v", err)
|
||||
}
|
||||
if rState != "rejected" || rReason != "left_group" || rEP != "bob" {
|
||||
t.Fatalf("receipt state=%q reason=%q ep=%q", rState, rReason, rEP)
|
||||
}
|
||||
var pending int
|
||||
if err = db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id='alice' AND state IN ('scheduled','dispatched')`).Scan(&pending); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if pending != 0 {
|
||||
t.Fatalf("sender pending count=%d", pending)
|
||||
}
|
||||
}
|
||||
|
||||
type dissolveOnJoin struct {
|
||||
app *group.App
|
||||
gid string
|
||||
owner string
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (d *dissolveOnJoin) CheckTalkPasswordForJoin(ctx context.Context, _, _, _, _ string) error {
|
||||
d.once.Do(func() {
|
||||
if d.app == nil || d.gid == "" {
|
||||
return
|
||||
}
|
||||
_ = d.app.Dissolve(ctx, d.owner, &protocol.GroupDissolve{
|
||||
V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "hook", GroupID: d.gid,
|
||||
})
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestU02AddAfterTalkGateDissolvesReturnsNotFound(t *testing.T) {
|
||||
t.Parallel()
|
||||
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)
|
||||
hook := &dissolveOnJoin{owner: "alice"}
|
||||
gApp := group.New(group.Config{
|
||||
DB: db, Talk: hook, MaxGroupMembers: 1000,
|
||||
Now: func() time.Time { return fixed }, DefaultRemoteIP: "1.1.1.1",
|
||||
})
|
||||
hook.app = gApp
|
||||
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: "G",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
hook.gid = created.ID
|
||||
_, err = gApp.Add(ctx, "alice", &protocol.GroupAdd{
|
||||
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "2", GroupID: created.ID,
|
||||
Members: []protocol.GroupMemberIn{{ID: "bob"}},
|
||||
})
|
||||
if protoCode(err) != protocol.CodeNotFound {
|
||||
t.Fatalf("got %v want not_found", err)
|
||||
}
|
||||
var n int
|
||||
if qErr := db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=?`, created.ID).Scan(&n); qErr != nil {
|
||||
t.Fatal(qErr)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("orphan members=%d", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestU02CreateDedupesMembers(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: "G",
|
||||
Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "bob"}, {ID: "alice"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(created.Failed) != 0 {
|
||||
t.Fatalf("failed=%+v", created.Failed)
|
||||
}
|
||||
var n, bobN int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=?`, created.ID).Scan(&n)
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=? AND endpoint_id=?`, created.ID, "bob").Scan(&bobN)
|
||||
if n != 2 || bobN != 1 {
|
||||
t.Fatalf("members=%d bob=%d", n, bobN)
|
||||
}
|
||||
}
|
||||
|
||||
func TestU02AdminCreateOwnerMustExistAndEnabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, _, _, db, down := setup(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
insertEP(t, db, "bob", 1)
|
||||
insertEP(t, db, "dave", 0)
|
||||
|
||||
_, err := gApp.AdminCreate(ctx, "G", "nobody", nil)
|
||||
if protoCode(err) != protocol.CodeInvalidTarget {
|
||||
t.Fatalf("missing owner got %v", err)
|
||||
}
|
||||
_, err = gApp.AdminCreate(ctx, "G", "dave", nil)
|
||||
if protoCode(err) != protocol.CodeEndpointDisabled {
|
||||
t.Fatalf("disabled owner got %v", err)
|
||||
}
|
||||
_, err = gApp.AdminCreate(ctx, "G", "Alice", nil)
|
||||
if protoCode(err) != protocol.CodeBadRequest {
|
||||
t.Fatalf("invalid owner format got %v", err)
|
||||
}
|
||||
|
||||
down.mu.Lock()
|
||||
down.msgs = nil
|
||||
down.mu.Unlock()
|
||||
created, err := gApp.AdminCreate(ctx, "一组", "alice", []string{"bob", "bob"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n, bobN int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=?`, created.ID).Scan(&n)
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=? AND endpoint_id=?`, created.ID, "bob").Scan(&bobN)
|
||||
if n != 2 || bobN != 1 {
|
||||
t.Fatalf("members=%d bob=%d", n, bobN)
|
||||
}
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
down.mu.Lock()
|
||||
got := len(down.msgs)
|
||||
down.mu.Unlock()
|
||||
if got >= 2 {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
down.mu.Lock()
|
||||
defer down.mu.Unlock()
|
||||
t.Fatalf("expected member_added downlink, got %d msgs", len(down.msgs))
|
||||
}
|
||||
|
||||
func TestU02ConcurrentAddRespectsLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
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 },
|
||||
})
|
||||
gApp := group.New(group.Config{
|
||||
DB: db, Talk: idApp, MaxGroupMembers: 3,
|
||||
Now: func() time.Time { return fixed }, DefaultRemoteIP: "1.1.1.1",
|
||||
})
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
insertEP(t, db, "bob", 1)
|
||||
insertEP(t, db, "carol", 1)
|
||||
insertEP(t, db, "dave", 1)
|
||||
created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", Name: "G",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var wg sync.WaitGroup
|
||||
for _, id := range []string{"bob", "carol", "dave"} {
|
||||
wg.Add(1)
|
||||
go func(id string) {
|
||||
defer wg.Done()
|
||||
_, _ = gApp.Add(ctx, "alice", &protocol.GroupAdd{
|
||||
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "a" + id,
|
||||
GroupID: created.ID, Members: []protocol.GroupMemberIn{{ID: id}},
|
||||
})
|
||||
}(id)
|
||||
}
|
||||
wg.Wait()
|
||||
var n int
|
||||
if qErr := db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=?`, created.ID).Scan(&n); qErr != nil {
|
||||
t.Fatal(qErr)
|
||||
}
|
||||
if n != 3 {
|
||||
t.Fatalf("members=%d want 3", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestU02GroupMembersFKRejectsOrphan(t *testing.T) {
|
||||
t.Parallel()
|
||||
gApp, _, _, db, _ := setup(t)
|
||||
ctx := context.Background()
|
||||
insertEP(t, db, "alice", 1)
|
||||
created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", Name: "G",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = gApp.Dissolve(ctx, "alice", &protocol.GroupDissolve{
|
||||
V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "2", GroupID: created.ID,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
created.ID, "alice", 1_700_000_000_000)
|
||||
return e
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected foreign key failure")
|
||||
}
|
||||
}
|
||||
|
||||
+20
-33
@@ -5,6 +5,7 @@ import (
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/message"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
@@ -19,7 +20,7 @@ type revokeItem struct {
|
||||
// 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
|
||||
SELECT d.seq, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
@@ -30,14 +31,13 @@ WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
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 {
|
||||
if scanErr := rows.Scan(&r.seq, &r.msgID, &r.senderID); scanErr != nil {
|
||||
return scanErr
|
||||
}
|
||||
list = append(list, r)
|
||||
@@ -46,16 +46,18 @@ WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
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 {
|
||||
pushed, execErr := message.RejectPendingTx(tx, r.seq, endpointID, reason, nowMs)
|
||||
if execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
if pushed && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason,
|
||||
})
|
||||
}
|
||||
if e := message.TryFinalizeTx(tx, r.seq, nowMs, message.DefaultVoidRetentionDays); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -63,7 +65,7 @@ UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ? WHERE seq =
|
||||
// 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
|
||||
SELECT d.seq, d.endpoint_id, 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)
|
||||
@@ -73,14 +75,13 @@ WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
|
||||
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 {
|
||||
if scanErr := rows.Scan(&r.seq, &r.endpointID, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -91,35 +92,35 @@ WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
|
||||
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 {
|
||||
pushed, execErr := message.RejectPendingTx(tx, r.seq, r.endpointID, reasonGroupDissolved, nowMs)
|
||||
if execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
if pushed && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved,
|
||||
})
|
||||
}
|
||||
if e := message.TryFinalizeTx(tx, r.seq, nowMs, message.DefaultVoidRetentionDays); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
|
||||
srows, err := tx.Query(`
|
||||
SELECT seq, id, sender_id, receipt FROM messages
|
||||
SELECT seq, 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 {
|
||||
if scanErr := srows.Scan(&r.seq, &r.senderID, &r.receipt); scanErr != nil {
|
||||
_ = srows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -130,23 +131,9 @@ WHERE dest_kind = 'group' AND dest_id = ? AND state = 'scheduled'`, groupID)
|
||||
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 {
|
||||
if execErr := message.FinalizeMessageTx(tx, r.seq, r.receipt != 0, r.senderID, "", reasonGroupDissolved, nowMs, message.DefaultVoidRetentionDays); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, r.seq); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.receipt != 0 {
|
||||
// 消息级作废回执:endpoint_id 空,state=rejected(DEVELOPMENT 6.4);消息行仍为 completed
|
||||
if _, execErr := tx.Exec(`
|
||||
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
|
||||
VALUES(?,?,?,?,?,?,0)`,
|
||||
r.senderID, r.msgID, "", "rejected", reasonGroupDissolved, nowMs); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/message"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
@@ -140,9 +141,10 @@ WHERE id = ?`, endpointID); e != nil {
|
||||
}
|
||||
|
||||
func voidEndpointMessagesTx(tx *sql.Tx, endpointID, recvReason, sendReason string, nowMs int64, revokes *[]revokeItem) error {
|
||||
days := message.DefaultVoidRetentionDays
|
||||
// 发给 X 的 pending → rejected
|
||||
rows, err := tx.Query(`
|
||||
SELECT d.seq, d.pushed_at, m.id, m.sender_id, m.receipt
|
||||
SELECT d.seq, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.endpoint_id = ? AND d.state = 'pending'`, endpointID)
|
||||
@@ -151,15 +153,13 @@ WHERE d.endpoint_id = ? AND d.state = 'pending'`, endpointID)
|
||||
}
|
||||
type pendRow struct {
|
||||
seq int64
|
||||
pushed sql.NullInt64
|
||||
msgID string
|
||||
senderID string
|
||||
receipt int
|
||||
}
|
||||
var pending []pendRow
|
||||
for rows.Next() {
|
||||
var r pendRow
|
||||
if scanErr := rows.Scan(&r.seq, &r.pushed, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
|
||||
if scanErr := rows.Scan(&r.seq, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -171,18 +171,11 @@ WHERE d.endpoint_id = ? AND d.state = 'pending'`, endpointID)
|
||||
}
|
||||
finalSeqs := map[int64]struct{}{}
|
||||
for _, r := range pending {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
recvReason, nowMs, r.seq, endpointID); execErr != nil {
|
||||
pushed, execErr := message.RejectPendingTx(tx, r.seq, endpointID, recvReason, nowMs)
|
||||
if execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.receipt != 0 {
|
||||
if e := insertReceiptIfWantedTx(tx, r.senderID, r.msgID, endpointID, "rejected", recvReason, nowMs, true); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
if pushed && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: recvReason,
|
||||
})
|
||||
@@ -192,21 +185,20 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
|
||||
// 发给 X 的 scheduled 单聊 → completed,要回执则写
|
||||
srows, err := tx.Query(`
|
||||
SELECT seq, id, sender_id, receipt FROM messages
|
||||
SELECT seq, sender_id, receipt FROM messages
|
||||
WHERE dest_kind = 'endpoint' AND dest_id = ? AND state = 'scheduled'`, endpointID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type schedRow struct {
|
||||
seq int64
|
||||
msgID string
|
||||
senderID string
|
||||
receipt int
|
||||
}
|
||||
var scheduledTo []schedRow
|
||||
for srows.Next() {
|
||||
var r schedRow
|
||||
if scanErr := srows.Scan(&r.seq, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
|
||||
if scanErr := srows.Scan(&r.seq, &r.senderID, &r.receipt); scanErr != nil {
|
||||
_ = srows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -217,21 +209,10 @@ WHERE dest_kind = 'endpoint' AND dest_id = ? AND state = 'scheduled'`, endpointI
|
||||
return err
|
||||
}
|
||||
for _, r := range scheduledTo {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
|
||||
recvReason, 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 {
|
||||
// 消息级作废:endpoint_id 空,state=rejected(DEVELOPMENT 6.4)
|
||||
if e := insertReceiptIfWantedTx(tx, r.senderID, r.msgID, "", "rejected", recvReason, nowMs, true); e != nil {
|
||||
if e := message.FinalizeMessageTx(tx, r.seq, r.receipt != 0, r.senderID, "", recvReason, nowMs, days); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// X 发出的 scheduled → completed(sender_*),不写回执
|
||||
outSched, err := tx.Query(`SELECT seq FROM messages WHERE sender_id = ? AND state = 'scheduled'`, endpointID)
|
||||
@@ -252,19 +233,14 @@ UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 's
|
||||
return err
|
||||
}
|
||||
for _, seq := range outSeqs {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
|
||||
sendReason, seq); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq); execErr != nil {
|
||||
return execErr
|
||||
if e := message.FinalizeMessageTx(tx, seq, false, endpointID, "", sendReason, nowMs, days); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
|
||||
// X 发出的消息的 pending 投递 → rejected(sender_*),不写回执
|
||||
drows, err := tx.Query(`
|
||||
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
|
||||
SELECT d.seq, d.endpoint_id, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE m.sender_id = ? AND d.state = 'pending'`, endpointID)
|
||||
@@ -274,14 +250,13 @@ WHERE m.sender_id = ? AND d.state = 'pending'`, endpointID)
|
||||
type outPend struct {
|
||||
seq int64
|
||||
endpointID string
|
||||
pushed sql.NullInt64
|
||||
msgID string
|
||||
senderID string
|
||||
}
|
||||
var outPending []outPend
|
||||
for drows.Next() {
|
||||
var r outPend
|
||||
if scanErr := drows.Scan(&r.seq, &r.endpointID, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
|
||||
if scanErr := drows.Scan(&r.seq, &r.endpointID, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = drows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -292,13 +267,11 @@ WHERE m.sender_id = ? AND d.state = 'pending'`, endpointID)
|
||||
return err
|
||||
}
|
||||
for _, r := range outPending {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
sendReason, nowMs, r.seq, r.endpointID); execErr != nil {
|
||||
pushed, execErr := message.RejectPendingTx(tx, r.seq, r.endpointID, sendReason, nowMs)
|
||||
if execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
if pushed && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: sendReason,
|
||||
})
|
||||
@@ -307,7 +280,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
}
|
||||
|
||||
for seq := range finalSeqs {
|
||||
if e := tryFinalizeTx(tx, seq); e != nil {
|
||||
if e := message.TryFinalizeTx(tx, seq, nowMs, days); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
@@ -445,7 +418,7 @@ func withoutMember(ids []string, drop string) []string {
|
||||
// voidMemberDeliveriesTx 与 group 包同语义:退群成员的 pending 群投递改 rejected。
|
||||
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
|
||||
SELECT d.seq, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
@@ -455,14 +428,13 @@ WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
}
|
||||
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 {
|
||||
if scanErr := rows.Scan(&r.seq, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -472,19 +444,18 @@ WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
if err = rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
days := message.DefaultVoidRetentionDays
|
||||
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 {
|
||||
pushed, execErr := message.RejectPendingTx(tx, r.seq, endpointID, reason, nowMs)
|
||||
if execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
if pushed && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason,
|
||||
})
|
||||
}
|
||||
if e := tryFinalizeTx(tx, r.seq); e != nil {
|
||||
if e := message.TryFinalizeTx(tx, r.seq, nowMs, days); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
@@ -492,8 +463,9 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
}
|
||||
|
||||
func voidGroupAllTx(tx *sql.Tx, groupID string, nowMs int64, revokes *[]revokeItem) error {
|
||||
days := message.DefaultVoidRetentionDays
|
||||
rows, err := tx.Query(`
|
||||
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
|
||||
SELECT d.seq, d.endpoint_id, 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)
|
||||
@@ -503,14 +475,13 @@ WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
|
||||
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 {
|
||||
if scanErr := rows.Scan(&r.seq, &r.endpointID, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -521,38 +492,35 @@ WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
|
||||
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 {
|
||||
pushed, execErr := message.RejectPendingTx(tx, r.seq, r.endpointID, reasonGroupDissolved, nowMs)
|
||||
if execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
if pushed && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved,
|
||||
})
|
||||
}
|
||||
if e := tryFinalizeTx(tx, r.seq); e != nil {
|
||||
if e := message.TryFinalizeTx(tx, r.seq, nowMs, days); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
|
||||
srows, err := tx.Query(`
|
||||
SELECT seq, id, sender_id, receipt FROM messages
|
||||
SELECT seq, 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 {
|
||||
if scanErr := srows.Scan(&r.seq, &r.senderID, &r.receipt); scanErr != nil {
|
||||
_ = srows.Close()
|
||||
return scanErr
|
||||
}
|
||||
@@ -563,66 +531,13 @@ WHERE dest_kind = 'group' AND dest_id = ? AND state = 'scheduled'`, groupID)
|
||||
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 e := insertReceiptIfWantedTx(tx, r.senderID, r.msgID, "", "rejected", reasonGroupDissolved, nowMs, true); e != nil {
|
||||
if e := message.FinalizeMessageTx(tx, r.seq, r.receipt != 0, r.senderID, "", reasonGroupDissolved, nowMs, days); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func insertReceiptIfWantedTx(tx *sql.Tx, senderID, msgID, endpointID, state, reason string, nowMs int64, alreadyWanted bool) error {
|
||||
if !alreadyWanted {
|
||||
return nil
|
||||
}
|
||||
var one int
|
||||
err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, senderID).Scan(&one)
|
||||
if errors.Is(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 tryFinalizeTx(tx *sql.Tx, seq int64) 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 state string
|
||||
if err := tx.QueryRow(`SELECT state FROM messages WHERE seq = ?`, seq).Scan(&state); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if state == "completed" {
|
||||
return nil
|
||||
}
|
||||
if _, err := tx.Exec(`UPDATE messages SET state = 'completed' WHERE seq = ?`, seq); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq)
|
||||
return err
|
||||
}
|
||||
|
||||
func (a *App) publishRevokes(ctx context.Context, items []revokeItem) {
|
||||
if a.down == nil || len(items) == 0 {
|
||||
return
|
||||
|
||||
@@ -228,6 +228,54 @@ func TestF01DisableVoidsScheduledAndRejectsNew(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisableLastPendingFinalizesAndReceipt(t *testing.T) {
|
||||
t.Parallel()
|
||||
idApp, msgApp, db := openLifecycle(t)
|
||||
ctx := context.Background()
|
||||
insertEPFull(t, db, "alice")
|
||||
insertEPFull(t, db, "bob")
|
||||
ttl := int64(3600)
|
||||
if _, err := msgApp.Submit(ctx, "alice", port.ConnInfo{EndpointID: "alice"}, &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "keep-bob",
|
||||
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||
Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := idApp.Disable(ctx, "bob"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var state, reason string
|
||||
if err := db.Read.QueryRow(`SELECT state, reason FROM messages WHERE id='keep-bob'`).Scan(&state, &reason); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if state != "completed" {
|
||||
t.Fatalf("state=%s want completed", state)
|
||||
}
|
||||
var bodies int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM message_bodies b JOIN messages m ON m.seq=b.seq WHERE m.id='keep-bob'`).Scan(&bodies); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bodies != 0 {
|
||||
t.Fatalf("body still present: %d", bodies)
|
||||
}
|
||||
var rState, rReason string
|
||||
if err := db.Read.QueryRow(`SELECT state, reason FROM receipts WHERE sender_id='alice' AND msg_id='keep-bob'`).Scan(&rState, &rReason); err != nil {
|
||||
t.Fatalf("receipt: %v", err)
|
||||
}
|
||||
if rState != "rejected" || rReason != "endpoint_disabled" {
|
||||
t.Fatalf("receipt state=%q reason=%q", rState, rReason)
|
||||
}
|
||||
var pending int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id='alice' AND state IN ('scheduled','dispatched')`).Scan(&pending); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if pending != 0 {
|
||||
t.Fatalf("sender pending=%d", pending)
|
||||
}
|
||||
}
|
||||
|
||||
func TestF01DeleteOwnerTransfersEarliest(t *testing.T) {
|
||||
t.Parallel()
|
||||
idApp, _, db := openLifecycle(t)
|
||||
|
||||
@@ -48,7 +48,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
if e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs); e != nil {
|
||||
return e
|
||||
}
|
||||
return tryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
}
|
||||
var state string
|
||||
err = tx.QueryRow(`
|
||||
@@ -175,7 +175,7 @@ SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state IN ('expired','dropped',
|
||||
default:
|
||||
data.Result = "failed"
|
||||
}
|
||||
return tryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
})
|
||||
if err != nil {
|
||||
return data, err
|
||||
|
||||
@@ -165,6 +165,8 @@ func (a *App) protocolLimits() protocol.Limits {
|
||||
MaxBodyBytes: a.lim.MaxBodyBytes,
|
||||
MaxMetaBytes: a.lim.MaxMetaBytes,
|
||||
MaxFrameBytes: a.lim.MaxFrameBytes,
|
||||
MaxTTLSeconds: a.lim.MaxTTLSeconds,
|
||||
MaxScheduleSeconds: a.lim.MaxScheduleSeconds,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
@@ -742,3 +743,37 @@ func TestPushRevokedOnRecallAfterPush(t *testing.T) {
|
||||
t.Fatal("expected revoked frame")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPushPreservesLargeMetaInteger(t *testing.T) {
|
||||
t.Parallel()
|
||||
raw := `{"id":12345678901234567890}`
|
||||
decoded := decodeMetaJSON(raw)
|
||||
n, ok := decoded["id"].(json.Number)
|
||||
if !ok || n.String() != "12345678901234567890" {
|
||||
t.Fatalf("decode meta=%v", decoded)
|
||||
}
|
||||
|
||||
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("meta-big", "bob")
|
||||
req.Meta = map[string]any{"id": json.Number("12345678901234567890")}
|
||||
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)
|
||||
}
|
||||
found := false
|
||||
for _, p := range e.down.Snapshots() {
|
||||
if bytes.Contains(p.Payload, []byte("12345678901234567890")) {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("push payloads missing large int: %v", e.down.Snapshots())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,11 +3,13 @@ package message
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
// DefaultVoidRetentionDays 是 group/identity 作废路径未接线配置时的记录保留天数。
|
||||
const DefaultVoidRetentionDays = 7
|
||||
|
||||
// 投递状态(DEVELOPMENT 7.1)。
|
||||
const (
|
||||
DeliveryPending = "pending"
|
||||
@@ -91,7 +93,7 @@ func (a *App) dispatchFullTx(tx *sql.Tx, seq int64, senderID, destKind, destID s
|
||||
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)
|
||||
WHERE gm.group_id = ? AND gm.endpoint_id != ? AND gm.joined_at <= ?`, destID, senderID, sendAt)
|
||||
if qErr != nil {
|
||||
return "", true, qErr
|
||||
}
|
||||
@@ -116,7 +118,7 @@ WHERE gm.group_id = ? AND gm.endpoint_id != ?`, destID, senderID)
|
||||
}
|
||||
|
||||
if completeEarly {
|
||||
if err := finalizeMessageTx(tx, seq, wantReceipt, senderID, "", msgReason, nowMs, a.lim.RecordRetentionDays); err != nil {
|
||||
if err := FinalizeMessageTx(tx, seq, wantReceipt, senderID, "", msgReason, nowMs, a.lim.RecordRetentionDays); err != nil {
|
||||
return "", true, err
|
||||
}
|
||||
return StateCompleted, true, nil
|
||||
@@ -192,7 +194,7 @@ VALUES(?,?,?,?,?,?,?,NULL,NULL,0,?)`,
|
||||
}
|
||||
return StateDispatched, true, nil
|
||||
}
|
||||
if err := finalizeMessageTx(tx, seq, wantReceipt, senderID, "", "", nowMs, a.lim.RecordRetentionDays); err != nil {
|
||||
if err := FinalizeMessageTx(tx, seq, wantReceipt, senderID, "", "", nowMs, a.lim.RecordRetentionDays); err != nil {
|
||||
return "", true, err
|
||||
}
|
||||
return StateCompleted, true, nil
|
||||
@@ -212,9 +214,9 @@ func (a *App) lookupConn(endpointID string) (LiveConn, bool) {
|
||||
return a.conns.Current(endpointID)
|
||||
}
|
||||
|
||||
// finalizeMessageTx 无 pending 时收尾:completed、删正文;记录天数 0 则删消息与投递。
|
||||
// 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 {
|
||||
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 {
|
||||
@@ -244,8 +246,8 @@ func finalizeMessageTx(tx *sql.Tx, seq int64, wantReceipt bool, senderID, endpoi
|
||||
return nil
|
||||
}
|
||||
|
||||
// tryFinalizeTx 若无 pending 则收尾。
|
||||
func tryFinalizeTx(tx *sql.Tx, seq int64, nowMs int64, recordDays int) error {
|
||||
// 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
|
||||
@@ -261,7 +263,46 @@ func tryFinalizeTx(tx *sql.Tx, seq int64, nowMs int64, recordDays int) error {
|
||||
}
|
||||
return err
|
||||
}
|
||||
return finalizeMessageTx(tx, seq, receipt != 0, senderID, "", "", nowMs, recordDays)
|
||||
return FinalizeMessageTx(tx, seq, receipt != 0, senderID, "", "", nowMs, recordDays)
|
||||
}
|
||||
|
||||
func skipVoidReceipt(reason string) bool {
|
||||
return reason == "sender_disabled" || reason == "sender_deleted"
|
||||
}
|
||||
|
||||
// RejectPendingTx 把一条 pending 投递改为 rejected;消息要求回执且发送方存在时写回执。
|
||||
// 停用/删除发送方(sender_disabled / sender_deleted)不写回执(DEVELOPMENT 7.6)。
|
||||
// 返回该投递是否曾推送,供调用方发 revoked。
|
||||
func RejectPendingTx(tx *sql.Tx, seq int64, endpointID, reason string, nowMs int64) (pushed bool, err error) {
|
||||
var pushedAt sql.NullInt64
|
||||
err = tx.QueryRow(`SELECT pushed_at FROM deliveries WHERE seq = ? AND endpoint_id = ?`, seq, endpointID).Scan(&pushedAt)
|
||||
if err == sql.ErrNoRows {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
res, err := tx.Exec(`
|
||||
UPDATE deliveries SET state = ?, reason = ?, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = ?`,
|
||||
DeliveryRejected, reason, nowMs, seq, endpointID, DeliveryPending)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
aff, _ := res.RowsAffected()
|
||||
if aff == 0 {
|
||||
return false, nil
|
||||
}
|
||||
if !skipVoidReceipt(reason) {
|
||||
var senderID string
|
||||
if err := tx.QueryRow(`SELECT sender_id FROM messages WHERE seq = ?`, seq).Scan(&senderID); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, reason, nowMs); err != nil {
|
||||
return false, err
|
||||
}
|
||||
}
|
||||
return pushedAt.Valid, nil
|
||||
}
|
||||
|
||||
func insertReceiptTx(tx *sql.Tx, senderID string, seq int64, endpointID, state, reason string, nowMs int64) error {
|
||||
@@ -300,7 +341,7 @@ func decodeMetaJSON(s string) map[string]any {
|
||||
return nil
|
||||
}
|
||||
var m map[string]any
|
||||
if err := json.Unmarshal([]byte(s), &m); err != nil {
|
||||
if err := protocol.Unmarshal([]byte(s), &m); err != nil {
|
||||
return nil
|
||||
}
|
||||
return m
|
||||
|
||||
@@ -259,7 +259,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL`
|
||||
if err := insertReceiptTx(tx, senderID, seq, endpointID, DeliveryRejected, ReasonTooLarge, nowMs); err != nil {
|
||||
return err
|
||||
}
|
||||
return tryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -361,7 +361,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := tryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays); err != nil {
|
||||
if err := TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays); err != nil {
|
||||
return err
|
||||
}
|
||||
if sendRevoked {
|
||||
|
||||
@@ -57,3 +57,11 @@ func (r *rateLimiter) allow(endpointID string, now time.Time) bool {
|
||||
b.tokens--
|
||||
return true
|
||||
}
|
||||
|
||||
// AllowRequest 消耗该端 1 个请求令牌;允许则 true。rps<=0 时不限速。
|
||||
func (a *App) AllowRequest(endpointID string) bool {
|
||||
if a == nil {
|
||||
return true
|
||||
}
|
||||
return a.rates.allow(endpointID, a.now())
|
||||
}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAllowRequestBurstAndAckExemptBucket(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
lim.RequestsPerSecond = 50
|
||||
lim.RequestBurst = 100
|
||||
app, _ := openTestApp(t, lim)
|
||||
allowed := 0
|
||||
for i := 0; i < 150; i++ {
|
||||
if app.AllowRequest("alice") {
|
||||
allowed++
|
||||
}
|
||||
}
|
||||
if allowed != 100 {
|
||||
t.Fatalf("allowed=%d want 100 (burst)", allowed)
|
||||
}
|
||||
}
|
||||
@@ -78,11 +78,23 @@ WHERE d.state = 'pending' AND d.pushed_conn IS NULL
|
||||
}
|
||||
}
|
||||
|
||||
if err := finalizeStuckDispatchedTx(tx, nowMs, a.lim.RecordRetentionDays); 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
|
||||
SELECT seq FROM (
|
||||
SELECT m.seq FROM messages m
|
||||
WHERE m.state = 'completed'
|
||||
AND COALESCE(
|
||||
(SELECT MAX(d.updated_at) FROM deliveries d WHERE d.seq = m.seq),
|
||||
m.send_at
|
||||
) < ?
|
||||
LIMIT 5000
|
||||
)
|
||||
)`, cutoff); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -120,3 +132,36 @@ DELETE FROM send_keys WHERE rowid IN (
|
||||
a.flushRevokes(ctx)
|
||||
return nil
|
||||
}
|
||||
|
||||
// finalizeStuckDispatchedTx 收尾「dispatched 且已无 pending 投递」的消息(C-04 兜底,修复已卡住的数据)。
|
||||
func finalizeStuckDispatchedTx(tx *sql.Tx, nowMs int64, recordDays int) error {
|
||||
rows, err := tx.Query(`
|
||||
SELECT seq FROM messages
|
||||
WHERE state = ?
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM deliveries d WHERE d.seq = messages.seq AND d.state = 'pending'
|
||||
)
|
||||
LIMIT 500`, StateDispatched)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var seqs []int64
|
||||
for rows.Next() {
|
||||
var seq int64
|
||||
if err := rows.Scan(&seq); err != nil {
|
||||
_ = rows.Close()
|
||||
return err
|
||||
}
|
||||
seqs = append(seqs, seq)
|
||||
}
|
||||
_ = rows.Close()
|
||||
if err := rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, seq := range seqs {
|
||||
if err := TryFinalizeTx(tx, seq, nowMs, recordDays); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -21,9 +21,6 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid sender")
|
||||
}
|
||||
now := a.now()
|
||||
if !a.rates.allow(senderID, now) {
|
||||
return SubmitResult{}, errCode(protocol.CodeRateLimited, "request rate exceeded")
|
||||
}
|
||||
|
||||
if err := req.Validate(a.protocolLimits()); err != nil {
|
||||
return SubmitResult{}, err
|
||||
@@ -49,6 +46,9 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
keep := protocol.EffectiveOfflineKeep(req)
|
||||
ttl := protocol.EffectiveOfflineTTL(req)
|
||||
receipt := protocol.EffectiveReceipt(req)
|
||||
if keep && ttl <= 0 {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds must be > 0")
|
||||
}
|
||||
if keep && a.lim.MaxTTLSeconds > 0 && ttl > a.lim.MaxTTLSeconds {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds exceeds max_ttl_seconds")
|
||||
}
|
||||
@@ -69,6 +69,9 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
}
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
if sender.Enabled == 0 {
|
||||
return SubmitResult{}, errCode(protocol.CodeUnauthorized, "sender disabled")
|
||||
}
|
||||
|
||||
sendAt, err := a.computeSendAt(req, sender.DefaultDelayMs, nowMs)
|
||||
if err != nil {
|
||||
@@ -158,6 +161,17 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
return e
|
||||
}
|
||||
|
||||
snd, se := loadEndpointTx(tx, senderID)
|
||||
if se != nil {
|
||||
if errors.Is(se, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeInvalidTarget, "sender not found")
|
||||
}
|
||||
return se
|
||||
}
|
||||
if snd.Enabled == 0 {
|
||||
return errCode(protocol.CodeUnauthorized, "sender disabled")
|
||||
}
|
||||
|
||||
// 写事务内再确认目标与授权(防并发停用/退群)。
|
||||
switch req.To.Kind {
|
||||
case protocol.TargetEndpoint:
|
||||
@@ -200,10 +214,6 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
|
||||
}
|
||||
}
|
||||
// 发送方设了对话密码且发给别人的单聊:给对方写回复授权。
|
||||
snd, se := loadEndpointTx(tx, senderID)
|
||||
if se != nil {
|
||||
return se
|
||||
}
|
||||
if senderID != req.To.ID && snd.TalkHash != nil && *snd.TalkHash != "" {
|
||||
if ge := upsertGrantTx(tx, req.To.ID, senderID, snd.TalkVersion, GrantKindReply, nowMs); ge != nil {
|
||||
return ge
|
||||
@@ -298,6 +308,12 @@ func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (in
|
||||
if *req.DelayMs < 0 {
|
||||
return 0, errCode(protocol.CodeBadRequest, "delay_ms negative")
|
||||
}
|
||||
if a.lim.MaxScheduleSeconds > 0 {
|
||||
maxDelay := a.lim.MaxScheduleSeconds * 1000
|
||||
if *req.DelayMs > maxDelay {
|
||||
return 0, errCode(protocol.CodeBadRequest, "send time exceeds max_schedule_seconds")
|
||||
}
|
||||
}
|
||||
sendAt = nowMs + *req.DelayMs
|
||||
default:
|
||||
if defaultDelayMs < 0 {
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"math"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -363,6 +364,141 @@ SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ttl_zero_rejected", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ttl := int64(0)
|
||||
req := baseSend("ttl0", "bob")
|
||||
req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}
|
||||
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req)
|
||||
if protoCode(err) != protocol.CodeBadRequest {
|
||||
t.Fatalf("ttl=0 want bad_request got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ttl_negative_rejected", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ttl := int64(-1)
|
||||
req := baseSend("ttlneg", "bob")
|
||||
req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}
|
||||
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req)
|
||||
if protoCode(err) != protocol.CodeBadRequest {
|
||||
t.Fatalf("ttl=-1 want bad_request got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delay_maxint64_rejected", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
delay := int64(math.MaxInt64)
|
||||
req := baseSend("delaymax", "bob")
|
||||
req.DelayMs = &delay
|
||||
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req)
|
||||
if protoCode(err) != protocol.CodeBadRequest {
|
||||
t.Fatalf("delay=MaxInt64 want bad_request got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("sender_disabled_unauthorized", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 0, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
_, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, baseSend("from-off", "bob"))
|
||||
if protoCode(err) != protocol.CodeUnauthorized {
|
||||
t.Fatalf("want unauthorized got %v", err)
|
||||
}
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("deliveries=%d", n)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("group_late_joiner_skipped", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
insertEndpoint(t, db, "dave", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
nowMs := int64(1_700_000_000_000)
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
"g-late", "g", "alice", nowMs); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, m := range []string{"alice", "bob"} {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
"g-late", m, nowMs); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
delay := int64(10_000)
|
||||
req := &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "late-1",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g-late"},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||
DelayMs: &delay,
|
||||
Offline: keepTrue(),
|
||||
}
|
||||
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != StateScheduled {
|
||||
t.Fatalf("state=%s", res.State)
|
||||
}
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
"g-late", "dave", res.SendAtMs+500)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := app.DispatchDue(ctx, res.SendAtMs, 10); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var daveN int
|
||||
if err := db.Read.QueryRow(`
|
||||
SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq
|
||||
WHERE m.id='late-1' AND d.endpoint_id='dave'`).Scan(&daveN); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if daveN != 0 {
|
||||
t.Fatalf("late joiner deliveries=%d", daveN)
|
||||
}
|
||||
var bobN int
|
||||
if err := db.Read.QueryRow(`
|
||||
SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq
|
||||
WHERE m.id='late-1' AND d.endpoint_id='bob'`).Scan(&bobN); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bobN != 1 {
|
||||
t.Fatalf("bob deliveries=%d", bobN)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("rate_limited", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
@@ -371,16 +507,15 @@ SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice")
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
if !app.AllowRequest("alice") || !app.AllowRequest("alice") {
|
||||
t.Fatal("burst should allow first two")
|
||||
}
|
||||
if app.AllowRequest("alice") {
|
||||
t.Fatal("third request should be rate limited")
|
||||
}
|
||||
ctx := context.Background()
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r2", "bob")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r3", "bob"))
|
||||
if protoCode(err) != protocol.CodeRateLimited {
|
||||
t.Fatalf("want rate_limited got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
)
|
||||
|
||||
func TestRejectPendingAndFinalizeRetentionZero(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
lim.RecordRetentionDays = 0
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
req := baseSend("z1", "bob")
|
||||
req.Offline = keepTrue()
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var seq int64
|
||||
if e := tx.QueryRow(`SELECT seq FROM messages WHERE id='z1'`).Scan(&seq); e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e := RejectPendingTx(tx, seq, "bob", ReasonEndpointDisabled, 1_700_000_000_000); e != nil {
|
||||
return e
|
||||
}
|
||||
return TryFinalizeTx(tx, seq, 1_700_000_000_000, 0)
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE id='z1'`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("message row should be deleted when retention=0, n=%d", n)
|
||||
}
|
||||
var receipts int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM receipts WHERE msg_id='z1' AND state='rejected'`).Scan(&receipts); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if receipts != 1 {
|
||||
t.Fatalf("receipts=%d", receipts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupStuckDispatched(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
nowMs := int64(1_700_000_000_000)
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, e := tx.Exec(`
|
||||
INSERT INTO messages(id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
|
||||
send_at, keep, ttl_seconds, receipt, state, reason, created_at)
|
||||
VALUES('stuck','alice','endpoint','bob','{}','text/plain; charset=utf-8','utf8',
|
||||
?,0,0,1,'dispatched','',?)`, nowMs, nowMs)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
seq, e := res.LastInsertId()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e = tx.Exec(`INSERT INTO message_bodies(seq, body) VALUES(?, ?)`, seq, []byte("x")); e != nil {
|
||||
return e
|
||||
}
|
||||
_, e = tx.Exec(`
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at)
|
||||
VALUES(?,?,?,0,'rejected','left_group',?)`, seq, "bob", nowMs, nowMs)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.CleanupOnce(ctx, nowMs); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var state string
|
||||
if err := db.Read.QueryRow(`SELECT state FROM messages WHERE id='stuck'`).Scan(&state); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if state != StateCompleted {
|
||||
t.Fatalf("state=%s want completed", state)
|
||||
}
|
||||
var bodies int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM message_bodies b JOIN messages m ON m.seq=b.seq WHERE m.id='stuck'`).Scan(&bodies); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bodies != 0 {
|
||||
t.Fatalf("body still present: %d", bodies)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupKeepsRecentlyCompletedOldCreated(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
lim.RecordRetentionDays = 7
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
nowMs := int64(1_700_000_000_000)
|
||||
created := nowMs - int64(30)*24*3600*1000
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, e := tx.Exec(`
|
||||
INSERT INTO messages(id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
|
||||
send_at, keep, ttl_seconds, receipt, state, reason, created_at)
|
||||
VALUES('old-created','alice','endpoint','bob','{}','text/plain; charset=utf-8','utf8',
|
||||
?,0,0,1,'completed','',?)`, nowMs, created)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
seq, e := res.LastInsertId()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
_, e = tx.Exec(`
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at)
|
||||
VALUES(?,?,?,0,'accepted','',?)`, seq, "bob", nowMs, nowMs)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.CleanupOnce(ctx, nowMs); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE id='old-created'`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("recently completed message should remain, n=%d", n)
|
||||
}
|
||||
}
|
||||
@@ -26,6 +26,8 @@ type Limits struct {
|
||||
MaxBodyBytes int
|
||||
MaxMetaBytes int
|
||||
MaxFrameBytes int
|
||||
MaxTTLSeconds int64
|
||||
MaxScheduleSeconds int64
|
||||
}
|
||||
|
||||
// DefaultLimits 返回 DEVELOPMENT 示例中的默认上限。
|
||||
|
||||
@@ -167,6 +167,23 @@ func (s *Send) Validate(lim Limits) error {
|
||||
if s.SendAtMs != nil && s.DelayMs != nil {
|
||||
return badRequest("send_at_ms and delay_ms are mutually exclusive")
|
||||
}
|
||||
if EffectiveOfflineKeep(s) {
|
||||
if s.Offline != nil && s.Offline.TTLSeconds != nil && *s.Offline.TTLSeconds <= 0 {
|
||||
return badRequest("ttl_seconds must be > 0")
|
||||
}
|
||||
ttl := EffectiveOfflineTTL(s)
|
||||
if lim.MaxTTLSeconds > 0 && ttl > lim.MaxTTLSeconds {
|
||||
return badRequest("ttl_seconds exceeds max_ttl_seconds")
|
||||
}
|
||||
}
|
||||
if s.DelayMs != nil {
|
||||
if *s.DelayMs < 0 {
|
||||
return badRequest("delay_ms negative")
|
||||
}
|
||||
if lim.MaxScheduleSeconds > 0 && *s.DelayMs > lim.MaxScheduleSeconds*1000 {
|
||||
return badRequest("delay_ms exceeds max_schedule_seconds")
|
||||
}
|
||||
}
|
||||
n, err := FrameBytes(s)
|
||||
if err != nil {
|
||||
return badRequest("cannot encode frame")
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
-- U-02: group_members 增加指向 groups 的外键,避免解散后残留孤儿行。
|
||||
-- 本分支基于 C-04 时最大迁移号为 0002,按 TASKS 4.2 取 0003。
|
||||
DELETE FROM group_members WHERE group_id NOT IN (SELECT id FROM groups);
|
||||
|
||||
CREATE TABLE group_members_new (
|
||||
group_id TEXT NOT NULL REFERENCES groups(id) ON DELETE CASCADE,
|
||||
endpoint_id TEXT NOT NULL,
|
||||
joined_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (group_id, endpoint_id)
|
||||
);
|
||||
|
||||
INSERT INTO group_members_new (group_id, endpoint_id, joined_at)
|
||||
SELECT group_id, endpoint_id, joined_at FROM group_members;
|
||||
|
||||
DROP TABLE group_members;
|
||||
|
||||
ALTER TABLE group_members_new RENAME TO group_members;
|
||||
|
||||
CREATE INDEX idx_group_members_endpoint ON group_members(endpoint_id);
|
||||
Reference in New Issue
Block a user