package group_test import ( "context" "database/sql" "errors" "path/filepath" "sync" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/app/group" "git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/app/message" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/config" "git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/store" ) type memDown struct { mu sync.Mutex msgs []struct { to string qos byte raw []byte } } func (d *memDown) PublishDown(_ context.Context, endpointID string, _ port.ConnID, payload []byte, opts port.PublishOpts) error { d.mu.Lock() defer d.mu.Unlock() d.msgs = append(d.msgs, struct { to string qos byte raw []byte }{endpointID, opts.QoS, append([]byte(nil), payload...)}) return nil } func setup(t *testing.T) (*group.App, *identity.App, *message.App, *store.DB, *memDown) { t.Helper() db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) fixed := time.UnixMilli(1_700_000_000_000) locks := auth.NewLoginLocks() idApp := identity.New(identity.Config{ DB: db, Hash: auth.NewStubHashPool(), Locks: locks, Sessions: auth.NewSessionTokens(), Now: func() time.Time { return fixed }, }) down := &memDown{} gApp := group.New(group.Config{ DB: db, Talk: idApp, Downlink: down, MaxGroupMembers: 1000, Now: func() time.Time { return fixed }, DefaultRemoteIP: "1.1.1.1", }) lim := message.LimitsFromFullConfig(config.Default()) lim.RequestsPerSecond = 0 msgApp := message.New(db, lim, auth.NewStubHashPool(), message.WithNow(func() time.Time { return fixed }), message.WithLocks(locks), ) return gApp, idApp, msgApp, db, down } func insertEP(t *testing.T, db *store.DB, id string, enabled int) { t.Helper() err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { _, e := tx.Exec(` INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) VALUES(?,?,?,?,0,0,?,?)`, id, id, "stub$login", nil, enabled, 1_700_000_000_000) return e }) if err != nil { t.Fatal(err) } } func protoCode(err error) string { var pe *protocol.Error if errors.As(err, &pe) { return pe.Code } return "" } func TestF16CreateAddPasswordAndOwner(t *testing.T) { t.Parallel() gApp, idApp, _, db, _ := setup(t) ctx := context.Background() insertEP(t, db, "alice", 1) insertEP(t, db, "bob", 1) insertEP(t, db, "carol", 1) insertEP(t, db, "dave", 0) // disabled if err := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil { t.Fatal(err) } // 非群主加人:先建群 created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", ID: "g_test01", Name: "一组", Members: []protocol.GroupMemberIn{ {ID: "bob", TalkPassword: "wrong"}, {ID: "carol"}, {ID: "dave"}, {ID: "nobody"}, }, }) if err != nil { t.Fatal(err) } if created.OwnerID != "alice" || created.ID != "g_test01" { t.Fatalf("%+v", created) } // bob 密码错、dave 停用、nobody 无效 → failed;carol 加入 codes := map[string]string{} for _, f := range created.Failed { codes[f.ID] = f.Code } if codes["bob"] != protocol.CodeTalkPasswordInvalid { t.Fatalf("bob fail %+v", created.Failed) } if codes["dave"] != protocol.CodeEndpointDisabled { t.Fatalf("dave fail %+v", created.Failed) } if codes["nobody"] != protocol.CodeInvalidTarget { t.Fatalf("nobody fail %+v", created.Failed) } var n int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=?`, "g_test01").Scan(&n) if n != 2 { // alice + carol t.Fatalf("members=%d", n) } var bobN int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id=? AND endpoint_id=?`, "g_test01", "bob").Scan(&bobN) if bobN != 0 { t.Fatal("bob should not be member") } // 非群主加人失败 _, err = gApp.Add(ctx, "carol", &protocol.GroupAdd{ V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "2", GroupID: "g_test01", Members: []protocol.GroupMemberIn{{ID: "bob", TalkPassword: "secret"}}, }) if protoCode(err) != protocol.CodeForbidden { t.Fatalf("got %v", err) } // 群主带对密码加人;已有单聊授权也不能省略 if err = idApp.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil { t.Fatal(err) } addRes, err := gApp.Add(ctx, "alice", &protocol.GroupAdd{ V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "3", GroupID: "g_test01", Members: []protocol.GroupMemberIn{{ID: "bob"}}, }) if err != nil { t.Fatal(err) } if len(addRes.Failed) != 1 || addRes.Failed[0].Code != protocol.CodeTalkPasswordRequired { t.Fatalf("%+v", addRes.Failed) } addRes, err = gApp.Add(ctx, "alice", &protocol.GroupAdd{ V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "4", GroupID: "g_test01", Members: []protocol.GroupMemberIn{{ID: "bob", TalkPassword: "secret"}}, }) if err != nil || len(addRes.Failed) != 0 { t.Fatalf("err=%v failed=%+v", err, addRes.Failed) } } func TestF16LeaveRemoveDissolve(t *testing.T) { t.Parallel() gApp, _, msgApp, db, down := setup(t) ctx := context.Background() insertEP(t, db, "alice", 1) insertEP(t, db, "bob", 1) insertEP(t, db, "carol", 1) created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", ID: "g_leave1", Name: "L", Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "carol"}}, }) if err != nil { t.Fatal(err) } // 群主不能直接退出 err = gApp.Leave(ctx, "alice", &protocol.GroupLeave{ V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "2", GroupID: created.ID, }) if protoCode(err) != protocol.CodeOwnerCannotLeave { t.Fatalf("got %v", err) } // 提交延迟群消息,再让 bob 退出 → pending 未推送改 left_group delay := int64(60_000) _, err = msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "s1", ID: "gm1", To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"}, DelayMs: &delay, }) if err != nil { t.Fatal(err) } // 到点最小分发:手动把消息改成 dispatched + pending deliveries(模拟已分发未推送) err = db.Queue.Do(ctx, func(tx *sql.Tx) error { var seq int64 if e := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id=? AND id=?`, "alice", "gm1").Scan(&seq); e != nil { return e } if _, e := tx.Exec(`UPDATE messages SET state='dispatched', send_at=? WHERE seq=?`, 1_700_000_000_000, seq); e != nil { return e } for _, ep := range []string{"bob", "carol"} { if _, e := tx.Exec(` INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at) VALUES(?,?,?,0,'pending','',?)`, seq, ep, 1_700_000_000_000, 1_700_000_000_000); e != nil { return e } } return nil }) if err != nil { t.Fatal(err) } if err = gApp.Leave(ctx, "bob", &protocol.GroupLeave{ V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "3", GroupID: created.ID, }); err != nil { t.Fatal(err) } var reason string err = db.Read.QueryRow(` SELECT d.reason FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='gm1' AND d.endpoint_id='bob'`).Scan(&reason) if err != nil || reason != "left_group" { t.Fatalf("reason=%q err=%v", reason, err) } // 已推送的踢人发 revoked err = db.Queue.Do(ctx, func(tx *sql.Tx) error { _, e := tx.Exec(`UPDATE deliveries SET pushed_at=?, pushed_conn='c' WHERE endpoint_id='carol' AND state='pending'`, 1_700_000_000_000) return e }) if err != nil { t.Fatal(err) } down.mu.Lock() down.msgs = nil down.mu.Unlock() if err = gApp.Remove(ctx, "alice", &protocol.GroupRemove{ V: protocol.Version, Type: protocol.TypeGroupRemove, RID: "4", GroupID: created.ID, EndpointID: "carol", }); err != nil { t.Fatal(err) } down.mu.Lock() nRev := 0 for _, m := range down.msgs { if m.qos == 1 { nRev++ } } down.mu.Unlock() if nRev < 1 { t.Fatal("expected revoked for pushed delivery") } // 解散:scheduled 作废,编号可复用 _, err = msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "s2", ID: "gm2", To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "later"}, DelayMs: &delay, }) if err != nil { t.Fatal(err) } // 重新加 carol 以便解散时有成员(alice 仍是群主) _, _ = gApp.AdminAddMembers(ctx, created.ID, []string{"carol"}) if err = gApp.Dissolve(ctx, "alice", &protocol.GroupDissolve{ V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "5", GroupID: created.ID, }); err != nil { t.Fatal(err) } var state, mreason string err = db.Read.QueryRow(`SELECT state, reason FROM messages WHERE id='gm2'`).Scan(&state, &mreason) if err != nil || state != "completed" || mreason != "group_dissolved" { t.Fatalf("state=%s reason=%s err=%v", state, mreason, err) } // 消息级回执 state 须为协议枚举 rejected,不得写成消息状态 completed(issue #6 / DEVELOPMENT 6.4) var rState, rReason, rEndpoint string err = db.Read.QueryRow(` SELECT state, reason, endpoint_id FROM receipts WHERE sender_id='alice' AND msg_id='gm2'`).Scan(&rState, &rReason, &rEndpoint) if err != nil { t.Fatalf("receipt for dissolved scheduled: %v", err) } if rState != "rejected" || rReason != "group_dissolved" || rEndpoint != "" { t.Fatalf("receipt state=%q reason=%q endpoint=%q want rejected/group_dissolved/empty", rState, rReason, rEndpoint) } // 同编号新建群 created2, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "6", ID: created.ID, Name: "新群", Members: nil, }) if err != nil || created2.ID != created.ID { t.Fatalf("reuse id err=%v %+v", err, created2) } // 旧 scheduled 不应再存在为 scheduled _ = db.Read.QueryRow(`SELECT state FROM messages WHERE id='gm2'`).Scan(&state) if state == "scheduled" { t.Fatal("old scheduled should stay completed") } } func TestF06GroupSendMembership(t *testing.T) { t.Parallel() gApp, _, msgApp, db, _ := setup(t) ctx := context.Background() insertEP(t, db, "alice", 1) insertEP(t, db, "bob", 1) insertEP(t, db, "carol", 1) created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", Name: "G", Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "carol"}}, }) if err != nil { t.Fatal(err) } // 默认不保留 + 成员离线(无连接、无 offline_since):投递立即 dropped,消息 completed(DEVELOPMENT 7.4/7.6) res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "s", ID: "m1", To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "broadcast"}, }) if err != nil { t.Fatal(err) } if res.State != message.StateCompleted { t.Fatalf("state=%s want completed (offline, keep=false)", res.State) } var cnt int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1'`).Scan(&cnt) if cnt != 2 { t.Fatalf("deliveries=%d want 2 (not sender)", cnt) } var pending, dropped int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.state='pending'`).Scan(&pending) _ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.state='dropped'`).Scan(&dropped) if pending != 0 || dropped != 2 { t.Fatalf("pending=%d dropped=%d want pending=0 dropped=2", pending, dropped) } var self int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.endpoint_id='alice'`).Scan(&self) if self != 0 { t.Fatal("sender should not receive") } // 发送后入群收不到旧消息:新成员 dave 入群后不应有该投递 insertEP(t, db, "dave", 1) _, err = gApp.Add(ctx, "alice", &protocol.GroupAdd{ V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "2", GroupID: created.ID, Members: []protocol.GroupMemberIn{{ID: "dave"}}, }) if err != nil { t.Fatal(err) } var daveN int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.endpoint_id='dave'`).Scan(&daveN) if daveN != 0 { t.Fatal("late joiner should not get old delivery") } } func TestF06GroupSendKeepOfflinePending(t *testing.T) { t.Parallel() gApp, _, msgApp, db, _ := setup(t) ctx := context.Background() insertEP(t, db, "alice", 1) insertEP(t, db, "bob", 1) insertEP(t, db, "carol", 1) created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", Name: "GKeep", Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "carol"}}, }) if err != nil { t.Fatal(err) } ttl := int64(3600) res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "s", ID: "mkeep", To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "kept"}, Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}, }) if err != nil { t.Fatal(err) } if res.State != message.StateDispatched { t.Fatalf("state=%s want dispatched (offline keep)", res.State) } var pending int _ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='mkeep' AND d.state='pending'`).Scan(&pending) if pending != 2 { t.Fatalf("pending=%d want 2", pending) } } func TestGroupTransferRenameListGet(t *testing.T) { t.Parallel() gApp, _, _, db, _ := setup(t) ctx := context.Background() insertEP(t, db, "alice", 1) insertEP(t, db, "bob", 1) created, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", Name: "N", Members: []protocol.GroupMemberIn{{ID: "bob"}}, }) if err != nil { t.Fatal(err) } if err = gApp.Rename(ctx, "alice", &protocol.GroupRename{ V: protocol.Version, Type: protocol.TypeGroupRename, RID: "2", GroupID: created.ID, Name: "新名", }); err != nil { t.Fatal(err) } if err = gApp.Transfer(ctx, "alice", &protocol.GroupTransfer{ V: protocol.Version, Type: protocol.TypeGroupTransfer, RID: "3", GroupID: created.ID, EndpointID: "bob", }); err != nil { t.Fatal(err) } items, _, err := gApp.List(ctx, "alice", &protocol.GroupList{ V: protocol.Version, Type: protocol.TypeGroupList, RID: "4", Limit: 10, }) if err != nil || len(items) != 1 || items[0].OwnerID != "bob" { t.Fatalf("%+v err=%v", items, err) } got, err := gApp.Get(ctx, "alice", &protocol.GroupGet{ V: protocol.Version, Type: protocol.TypeGroupGet, RID: "5", GroupID: created.ID, Limit: 10, }) if err != nil || got.Name != "新名" || len(got.Members) != 2 { t.Fatalf("%+v err=%v", got, err) } _, err = gApp.Get(ctx, "nobody", &protocol.GroupGet{ V: protocol.Version, Type: protocol.TypeGroupGet, RID: "6", GroupID: created.ID, }) if protoCode(err) != protocol.CodeNotMember && protoCode(err) != protocol.CodeNotFound { // nobody 不是端也不是成员 if protoCode(err) == "" { t.Fatalf("got %v", err) } } } 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) } } // raceTalk 在对话密码校验成功后执行一次 mutate,模拟事务外密码窗口里创建者被停用/删除。 type raceTalk struct { inner group.TalkGate mutate func() once sync.Once } func (r *raceTalk) CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error { err := r.inner.CheckTalkPasswordForJoin(ctx, actorID, targetID, talkPassword, remoteIP) if err == nil && r.mutate != nil { r.once.Do(r.mutate) } return err } func TestR307CreateRejectsActorDisabledOrDeletedBeforeWrite(t *testing.T) { t.Parallel() cases := []struct { name string code string mutate func(ctx context.Context, t *testing.T, idApp *identity.App, db *store.DB) }{ { name: "disabled", code: protocol.CodeEndpointDisabled, mutate: func(ctx context.Context, t *testing.T, idApp *identity.App, _ *store.DB) { t.Helper() if err := idApp.Disable(ctx, "alice"); err != nil { t.Fatal(err) } }, }, { name: "deleted", code: protocol.CodeInvalidTarget, mutate: func(ctx context.Context, t *testing.T, idApp *identity.App, _ *store.DB) { t.Helper() if err := idApp.Delete(ctx, "alice"); err != nil { t.Fatal(err) } }, }, } for _, tc := range cases { tc := tc t.Run(tc.name, func(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 }, }) gate := &raceTalk{inner: idApp} gApp := group.New(group.Config{ DB: db, Talk: gate, MaxGroupMembers: 1000, Now: func() time.Time { return fixed }, DefaultRemoteIP: "1.1.1.1", }) ctx := context.Background() insertEP(t, db, "alice", 1) insertEP(t, db, "bob", 1) if err := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil { t.Fatal(err) } gate.mutate = func() { tc.mutate(ctx, t, idApp, db) } var before int if qErr := db.Read.QueryRow(`SELECT COUNT(*) FROM groups`).Scan(&before); qErr != nil { t.Fatal(qErr) } _, err = gApp.Create(ctx, "alice", &protocol.GroupCreate{ V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1", ID: "g_r307_" + tc.name, Name: "竞态群", Members: []protocol.GroupMemberIn{{ID: "bob", TalkPassword: "secret"}}, }) if protoCode(err) != tc.code { t.Fatalf("want %s got %v", tc.code, err) } var after int if qErr := db.Read.QueryRow(`SELECT COUNT(*) FROM groups`).Scan(&after); qErr != nil { t.Fatal(qErr) } if after != before { t.Fatalf("groups leaked: before=%d after=%d", before, after) } var exists int qErr := db.Read.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, "g_r307_"+tc.name).Scan(&exists) if !errors.Is(qErr, sql.ErrNoRows) { t.Fatalf("expected no group row, got exists=%d err=%v", exists, qErr) } }) } } 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") } }