Files
NixMsg/internal/app/group/group_test.go
T

413 lines
13 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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.LimitsFromConfig(config.Default().Limits)
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)
}
// 同编号新建群
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)
}
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.StateDispatched {
t.Fatalf("state=%s", 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 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 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)
}
}
}