fix: 群写操作在事务内复核并去重建群成员
This commit is contained in:
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user