413 lines
13 KiB
Go
413 lines
13 KiB
Go
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)
|
||
}
|
||
}
|
||
}
|