fix: 群写操作在事务内复核并去重建群成员
This commit is contained in:
@@ -716,6 +716,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) {
|
||||
|
||||
@@ -520,3 +520,211 @@ SELECT state, reason, endpoint_id FROM receipts WHERE sender_id='alice' AND msg_
|
||||
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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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