From ad4f13193c4a41623fbb1fb01c6048e6fab43208 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:22:46 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E7=BE=A4=E5=86=99=E6=93=8D=E4=BD=9C?= =?UTF-8?q?=E5=9C=A8=E4=BA=8B=E5=8A=A1=E5=86=85=E5=A4=8D=E6=A0=B8=E5=B9=B6?= =?UTF-8?q?=E5=8E=BB=E9=87=8D=E5=BB=BA=E7=BE=A4=E6=88=90=E5=91=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 30 + internal/app/group/app.go | 550 +++++++++++++----- internal/app/group/group_test.go | 208 +++++++ .../migrations/0003_group_members_fk.sql | 19 + 4 files changed, 648 insertions(+), 159 deletions(-) create mode 100644 internal/store/migrations/0003_group_members_fk.sql diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 2808272..a5ca39d 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -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 diff --git a/internal/app/group/app.go b/internal/app/group/app.go index aab588b..73e60d4 100644 --- a/internal/app/group/app.go +++ b/internal/app/group/app.go @@ -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) { - continue - } - if len(members)+len(added) >= a.maxMem { - failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull}) + 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 } + 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}) + } + 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 { - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - for _, id := range added { - if _, e := tx.Exec(`INSERT OR IGNORE INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`, - req.GroupID, id, now); e != nil { - return e - } - } - return nil - }) - if err != nil { - return AddResult{}, err + 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 { + curOwner, curMembers, e := loadGroupTx(tx, req.GroupID) + if e != nil { + return e } - all := append(append([]string{}, members...), added...) + if curOwner != actorID { + return errCode(protocol.CodeForbidden, "not owner") + } + present := memberSet(curMembers) + count := len(curMembers) + inserted = inserted[:0] for _, id := range added { - a.emit(ctx, all, req.GroupID, eventMemberAdded, id, now) + 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 + } + for _, id := range inserted { + a.emit(ctx, notify, req.GroupID, eventMemberAdded, id, now) } return AddResult{Failed: failed}, nil } @@ -234,34 +292,37 @@ 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 - } - if owner != actorID { - return errCode(protocol.CodeForbidden, "not owner") - } - if req.EndpointID == owner { - return errCode(protocol.CodeBadRequest, "cannot remove owner") - } - if !contains(members, req.EndpointID) { - return errCode(protocol.CodeNotFound, "member not found") - } now := a.nowMs() var revokes []revokeItem - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + 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") + } + if req.EndpointID == owner { + return errCode(protocol.CodeBadRequest, "cannot remove owner") + } + if !contains(members, req.EndpointID) { + return errCode(protocol.CodeNotFound, "member not found") + } 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,31 +335,34 @@ 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 - } - if !contains(members, actorID) { - return errCode(protocol.CodeNotMember, "not a member") - } - if owner == actorID { - return errCode(protocol.CodeOwnerCannotLeave, "owner cannot leave") - } now := a.nowMs() var revokes []revokeItem - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { + 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") + } + if owner == actorID { + return errCode(protocol.CodeOwnerCannotLeave, "owner cannot leave") + } 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 - } - if owner != actorID { - return errCode(protocol.CodeForbidden, "not owner") - } - if !contains(members, req.EndpointID) { - return errCode(protocol.CodeNotFound, "member not found") - } now := a.nowMs() - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - _, e := tx.Exec(`UPDATE groups SET owner_id = ? WHERE id = ?`, req.EndpointID, req.GroupID) - return e + 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(cur, req.EndpointID) { + return errCode(protocol.CodeNotFound, "member not found") + } + 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 - } - if owner != actorID { - return errCode(protocol.CodeForbidden, "not owner") - } now := a.nowMs() - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - _, e := tx.Exec(`UPDATE groups SET name = ? WHERE id = ?`, req.Name, req.GroupID) - return e + 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 _, 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 - } - 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 { + 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 e := voidGroupAllTx(tx, req.GroupID, now, &revokes); e != nil { return e } if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ?`, req.GroupID); e != nil { return e } - _, e := tx.Exec(`DELETE FROM groups WHERE id = ?`, req.GroupID) - return e + if _, 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) - now := a.nowMs() - for _, id := range memberIDs { - if contains(members, id) { + uniq := dedupeIDs(memberIDs, "") + already := memberSet(members) + candidates := make([]string, 0, len(uniq)) + for _, id := range uniq { + if _, ok := already[id]; ok { continue } - var enabled int - e := a.db.Read.QueryRowContext(ctx, `SELECT enabled FROM endpoints WHERE id = ?`, id).Scan(&enabled) - if errors.Is(e, sql.ErrNoRows) { - failed = append(failed, MemberFail{ID: id, Code: protocol.CodeInvalidTarget}) - continue - } - if e != nil { - return AddResult{}, e - } - if enabled == 0 { - failed = append(failed, MemberFail{ID: id, Code: protocol.CodeEndpointDisabled}) - continue - } - if len(members)+len(added) >= a.maxMem { - failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull}) - continue - } - added = append(added, id) + candidates = append(candidates, id) } - if len(added) == 0 { + 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() + 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(?,?,?)`, - groupID, id, now); e != nil { - return e + _, curMembers, e := loadGroupTx(tx, groupID) + if e != nil { + return e + } + present := memberSet(curMembers) + count := len(curMembers) + inserted = inserted[:0] + for _, id := range toAdd { + 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, 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, 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) { diff --git a/internal/app/group/group_test.go b/internal/app/group/group_test.go index 8fc0d6f..0fc916d 100644 --- a/internal/app/group/group_test.go +++ b/internal/app/group/group_test.go @@ -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") + } +} diff --git a/internal/store/migrations/0003_group_members_fk.sql b/internal/store/migrations/0003_group_members_fk.sql new file mode 100644 index 0000000..ff34f69 --- /dev/null +++ b/internal/store/migrations/0003_group_members_fk.sql @@ -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);