feat: 实现身份资料、对话密码、在线目录与群管理

This commit is contained in:
Nixevol
2026-09-30 07:33:31 +08:00
parent bdd1d9e9f4
commit a78ab0d547
14 changed files with 2777 additions and 65 deletions
+727
View File
@@ -0,0 +1,727 @@
package group
import (
"bytes"
"context"
"crypto/rand"
"database/sql"
"errors"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
const (
eventMemberAdded = "member_added"
eventMemberRemoved = "member_removed"
eventLeft = "left"
eventOwnerChanged = "owner_changed"
eventRenamed = "renamed"
eventDissolved = "dissolved"
reasonLeftGroup = "left_group"
reasonGroupDissolved = "group_dissolved"
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
)
// TalkGate checks talk password when adding members (implemented by identity).
type TalkGate interface {
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
}
// OnlineLookup reports member online status.
type OnlineLookup interface {
IsOnline(endpointID string) bool
}
// Config holds group service dependencies.
type Config struct {
DB *store.DB
Talk TalkGate
Online OnlineLookup
Downlink port.Downlink
MaxGroupMembers int
Now func() time.Time
DefaultRemoteIP string
}
// App implements group.Service.
type App struct {
db *store.DB
talk TalkGate
online OnlineLookup
down port.Downlink
maxMem int
nowFn func() time.Time
remoteIP string
}
// New constructs the group service.
func New(cfg Config) *App {
now := cfg.Now
if now == nil {
now = time.Now
}
max := cfg.MaxGroupMembers
if max <= 0 {
max = 1000
}
return &App{
db: cfg.DB,
talk: cfg.Talk,
online: cfg.Online,
down: cfg.Downlink,
maxMem: max,
nowFn: now,
remoteIP: cfg.DefaultRemoteIP,
}
}
func (a *App) nowMs() int64 { return a.nowFn().UnixMilli() }
func errCode(code, msg string) *protocol.Error {
return &protocol.Error{Code: code, Message: msg}
}
func protoCode(err error) string {
var pe *protocol.Error
if errors.As(err, &pe) {
return pe.Code
}
return ""
}
// Create creates a group; creator becomes owner and member. Partial member failures still create the group.
func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCreate) (CreateResult, error) {
if req == nil {
return CreateResult{}, errCode(protocol.CodeBadRequest, "nil request")
}
if err := req.Validate(); err != nil {
return CreateResult{}, err
}
if !protocol.ValidEndpointID(actorID) {
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid actor")
}
gid := req.ID
if gid == "" {
var genErr error
gid, genErr = generateGroupID()
if genErr != nil {
return CreateResult{}, genErr
}
}
now := a.nowMs()
failed := make([]MemberFail, 0)
added := make([]string, 0, len(req.Members))
for _, m := range req.Members {
if m.ID == actorID {
continue
}
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
continue
}
added = append(added, m.ID)
}
if 1+len(added) > a.maxMem {
return CreateResult{}, errCode(protocol.CodeGroupFull, "group full")
}
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)
if qErr == nil {
return errCode(protocol.CodeIDTaken, "group id taken")
}
if !errors.Is(qErr, sql.ErrNoRows) {
return qErr
}
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
gid, req.Name, actorID, now); e != nil {
if isUnique(e) {
return errCode(protocol.CodeIDTaken, "group id taken")
}
return e
}
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
gid, actorID, 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 {
return e
}
}
return nil
})
if err != nil {
return CreateResult{}, err
}
for _, id := range added {
a.emit(ctx, append([]string{actorID}, added...), gid, eventMemberAdded, id, now)
}
return CreateResult{ID: gid, Name: req.Name, OwnerID: actorID, Failed: failed}, nil
}
// Add adds members (owner only).
func (a *App) Add(ctx context.Context, actorID string, req *protocol.GroupAdd) (AddResult, error) {
if req == nil {
return AddResult{}, errCode(protocol.CodeBadRequest, "nil request")
}
if err := req.Validate(); err != nil {
return AddResult{}, err
}
owner, members, err := a.loadGroup(ctx, req.GroupID)
if err != nil {
return AddResult{}, err
}
if owner != actorID {
return AddResult{}, errCode(protocol.CodeForbidden, "not owner")
}
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})
continue
}
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
continue
}
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
}
all := append(append([]string{}, members...), added...)
for _, id := range added {
a.emit(ctx, all, req.GroupID, eventMemberAdded, id, now)
}
}
return AddResult{Failed: failed}, nil
}
// Remove kicks a member (owner only).
func (a *App) Remove(ctx context.Context, actorID string, req *protocol.GroupRemove) error {
if req == nil {
return errCode(protocol.CodeBadRequest, "nil request")
}
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 {
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 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
}
// Leave lets a non-owner member leave.
func (a *App) Leave(ctx context.Context, actorID string, req *protocol.GroupLeave) error {
if req == nil {
return errCode(protocol.CodeBadRequest, "nil request")
}
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 {
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 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
}
// Transfer transfers ownership.
func (a *App) Transfer(ctx context.Context, actorID string, req *protocol.GroupTransfer) error {
if req == nil {
return errCode(protocol.CodeBadRequest, "nil request")
}
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
})
if err != nil {
return err
}
a.emit(ctx, members, req.GroupID, eventOwnerChanged, req.EndpointID, now)
return nil
}
// Rename renames the group.
func (a *App) Rename(ctx context.Context, actorID string, req *protocol.GroupRename) error {
if req == nil {
return errCode(protocol.CodeBadRequest, "nil request")
}
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
})
if err != nil {
return err
}
a.emit(ctx, members, req.GroupID, eventRenamed, "", now)
return nil
}
// Dissolve dissolves the group and voids unfinished deliveries / scheduled messages.
func (a *App) Dissolve(ctx context.Context, actorID string, req *protocol.GroupDissolve) error {
if req == nil {
return errCode(protocol.CodeBadRequest, "nil request")
}
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 {
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 err != nil {
return err
}
a.publishRevokes(ctx, revokes)
a.emit(ctx, members, req.GroupID, eventDissolved, "", now)
return nil
}
// List lists groups the actor belongs to.
func (a *App) List(ctx context.Context, actorID string, req *protocol.GroupList) ([]ListItem, string, error) {
if req == nil {
req = &protocol.GroupList{V: protocol.Version, Type: protocol.TypeGroupList, RID: "x", Limit: 100}
}
if err := req.Validate(); err != nil {
return nil, "", err
}
limit := req.Limit
if limit <= 0 {
limit = 100
}
if limit > protocol.MaxPageLimit {
limit = protocol.MaxPageLimit
}
rows, err := a.db.Read.QueryContext(ctx, `
SELECT g.id, g.name, g.owner_id,
(SELECT COUNT(*) FROM group_members gm2 WHERE gm2.group_id = g.id) AS cnt
FROM groups g
JOIN group_members gm ON gm.group_id = g.id AND gm.endpoint_id = ?
WHERE g.id > ?
ORDER BY g.id ASC
LIMIT ?`, actorID, req.Cursor, limit+1)
if err != nil {
return nil, "", err
}
defer func() { _ = rows.Close() }()
items := make([]ListItem, 0, limit)
for rows.Next() {
var it ListItem
if scanErr := rows.Scan(&it.ID, &it.Name, &it.OwnerID, &it.MemberCount); scanErr != nil {
return nil, "", scanErr
}
items = append(items, it)
}
next := ""
if len(items) > limit {
items = items[:limit]
next = items[len(items)-1].ID
}
return items, next, rows.Err()
}
// Get returns group details; caller must be a member.
func (a *App) Get(ctx context.Context, actorID string, req *protocol.GroupGet) (GetResult, error) {
if req == nil {
return GetResult{}, errCode(protocol.CodeBadRequest, "nil request")
}
if err := req.Validate(); err != nil {
return GetResult{}, err
}
var name, owner string
err := a.db.Read.QueryRowContext(ctx, `SELECT name, owner_id FROM groups WHERE id = ?`, req.GroupID).
Scan(&name, &owner)
if errors.Is(err, sql.ErrNoRows) {
return GetResult{}, errCode(protocol.CodeNotFound, "group not found")
}
if err != nil {
return GetResult{}, err
}
var one int
err = a.db.Read.QueryRowContext(ctx,
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, req.GroupID, actorID).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return GetResult{}, errCode(protocol.CodeNotMember, "not a member")
}
if err != nil {
return GetResult{}, err
}
limit := req.Limit
if limit <= 0 {
limit = 100
}
if limit > protocol.MaxPageLimit {
limit = protocol.MaxPageLimit
}
rows, err := a.db.Read.QueryContext(ctx, `
SELECT gm.endpoint_id, e.name
FROM group_members gm
JOIN endpoints e ON e.id = gm.endpoint_id
WHERE gm.group_id = ? AND gm.endpoint_id > ?
ORDER BY gm.endpoint_id ASC
LIMIT ?`, req.GroupID, req.Cursor, limit+1)
if err != nil {
return GetResult{}, err
}
defer func() { _ = rows.Close() }()
members := make([]MemberItem, 0, limit)
for rows.Next() {
var m MemberItem
if scanErr := rows.Scan(&m.ID, &m.Name); scanErr != nil {
return GetResult{}, scanErr
}
if a.online != nil {
m.Online = a.online.IsOnline(m.ID)
}
members = append(members, m)
}
next := ""
if len(members) > limit {
members = members[:limit]
next = members[len(members)-1].ID
}
return GetResult{ID: req.GroupID, Name: name, OwnerID: owner, Members: members, NextCursor: next}, rows.Err()
}
// AdminCreate creates a group without talk-password checks.
func (a *App) AdminCreate(ctx context.Context, name, ownerID string, memberIDs []string) (CreateResult, error) {
members := make([]protocol.GroupMemberIn, 0, len(memberIDs))
for _, id := range memberIDs {
members = append(members, protocol.GroupMemberIn{ID: id})
}
return a.createAdmin(ctx, ownerID, name, "", members)
}
// AdminAddMembers adds members without talk-password checks.
func (a *App) AdminAddMembers(ctx context.Context, groupID string, memberIDs []string) (AddResult, error) {
_, 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) {
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)
}
if len(added) == 0 {
return AddResult{Failed: failed}, nil
}
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
}
}
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)
}
return AddResult{Failed: failed}, nil
}
func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, members []protocol.GroupMemberIn) (CreateResult, error) {
if gid == "" {
var genErr error
gid, genErr = generateGroupID()
if genErr != nil {
return CreateResult{}, genErr
}
}
if !protocol.ValidName(name) || name == "" {
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid name")
}
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")
}
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
gid, name, ownerID, now); e != nil {
if isUnique(e) {
return errCode(protocol.CodeIDTaken, "group id taken")
}
return e
}
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
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 {
return e
}
}
return nil
})
if err != nil {
return CreateResult{}, err
}
return CreateResult{ID: gid, Name: name, OwnerID: ownerID, Failed: failed}, nil
}
func (a *App) checkAddMember(ctx context.Context, actorID, targetID, talkPassword string) error {
if a.talk == nil {
return errCode(protocol.CodeBusy, "talk gate not configured")
}
return a.talk.CheckTalkPasswordForJoin(ctx, actorID, targetID, talkPassword, a.remoteIP)
}
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) {
return "", nil, errCode(protocol.CodeNotFound, "group not found")
}
if err != nil {
return "", nil, err
}
rows, qErr := a.db.Read.QueryContext(ctx, `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 (a *App) emit(ctx context.Context, recipients []string, groupID, event, endpointID string, atMs int64) {
if a.down == nil {
return
}
frame := protocol.GroupEvent{
V: protocol.Version, Type: protocol.TypeGroupEvent,
GroupID: groupID, Event: event, EndpointID: endpointID, AtMs: atMs,
}
payload, encErr := encodeFrame(frame)
if encErr != nil {
return
}
seen := map[string]struct{}{}
for _, id := range recipients {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
_ = a.down.PublishDown(ctx, id, "", payload, port.PublishOpts{QoS: 0})
}
}
func encodeFrame(v any) ([]byte, error) {
var buf bytes.Buffer
if err := protocol.Encode(&buf, v); err != nil {
return nil, err
}
return buf.Bytes(), nil
}
func generateGroupID() (string, error) {
b := make([]byte, 8)
if _, err := rand.Read(b); err != nil {
return "", err
}
out := make([]byte, 8)
for i := range b {
out[i] = idAlphabet[int(b[i])%len(idAlphabet)]
}
return "g_" + string(out), nil
}
func failCode(err error) string {
if c := protoCode(err); c != "" {
return c
}
return protocol.CodeBusy
}
func contains(ss []string, x string) bool {
for _, s := range ss {
if s == x {
return true
}
}
return false
}
func without(ss []string, x string) []string {
out := make([]string, 0, len(ss))
for _, s := range ss {
if s != x {
out = append(out, s)
}
}
return out
}
var _ Service = (*App)(nil)
+412
View File
@@ -0,0 +1,412 @@
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)
}
}
}
+176
View File
@@ -0,0 +1,176 @@
package group
import (
"context"
"database/sql"
"strings"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
type revokeItem struct {
endpointID string
msgID string
fromID string
reason string
}
// voidMemberDeliveriesTx rejects pending deliveries for a leaving member; records revokes for pushed ones.
func voidMemberDeliveriesTx(tx *sql.Tx, groupID, endpointID, reason string, nowMs int64, revokes *[]revokeItem) error {
rows, err := tx.Query(`
SELECT d.seq, d.pushed_at, m.id, m.sender_id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.endpoint_id = ? AND d.state = 'pending'
AND m.dest_kind = 'group' AND m.dest_id = ?`, endpointID, groupID)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
type row struct {
seq int64
pushed sql.NullInt64
msgID string
senderID string
}
var list []row
for rows.Next() {
var r row
if scanErr := rows.Scan(&r.seq, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
return scanErr
}
list = append(list, r)
}
if err = rows.Err(); err != nil {
return err
}
for _, r := range list {
if _, execErr := tx.Exec(`
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
reason, nowMs, r.seq, endpointID); execErr != nil {
return execErr
}
if r.pushed.Valid && revokes != nil {
*revokes = append(*revokes, revokeItem{
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason,
})
}
}
return nil
}
// voidGroupAllTx rejects all pending group deliveries and completes scheduled messages.
func voidGroupAllTx(tx *sql.Tx, groupID string, nowMs int64, revokes *[]revokeItem) error {
rows, err := tx.Query(`
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
if err != nil {
return err
}
type drow struct {
seq int64
endpointID string
pushed sql.NullInt64
msgID string
senderID string
}
var dlist []drow
for rows.Next() {
var r drow
if scanErr := rows.Scan(&r.seq, &r.endpointID, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
_ = rows.Close()
return scanErr
}
dlist = append(dlist, r)
}
_ = rows.Close()
if err = rows.Err(); err != nil {
return err
}
for _, r := range dlist {
if _, execErr := tx.Exec(`
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
reasonGroupDissolved, nowMs, r.seq, r.endpointID); execErr != nil {
return execErr
}
if r.pushed.Valid && revokes != nil {
*revokes = append(*revokes, revokeItem{
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved,
})
}
}
srows, err := tx.Query(`
SELECT seq, id, sender_id, receipt FROM messages
WHERE dest_kind = 'group' AND dest_id = ? AND state = 'scheduled'`, groupID)
if err != nil {
return err
}
type srow struct {
seq int64
msgID string
senderID string
receipt int
}
var slist []srow
for srows.Next() {
var r srow
if scanErr := srows.Scan(&r.seq, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
_ = srows.Close()
return scanErr
}
slist = append(slist, r)
}
_ = srows.Close()
if err = srows.Err(); err != nil {
return err
}
for _, r := range slist {
if _, execErr := tx.Exec(`
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
reasonGroupDissolved, r.seq); execErr != nil {
return execErr
}
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, r.seq); execErr != nil {
return execErr
}
if r.receipt != 0 {
if _, execErr := tx.Exec(`
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
VALUES(?,?,?,?,?,?,0)`,
r.senderID, r.msgID, "", "completed", reasonGroupDissolved, nowMs); execErr != nil {
return execErr
}
}
}
return nil
}
func (a *App) publishRevokes(ctx context.Context, items []revokeItem) {
if a.down == nil || len(items) == 0 {
return
}
for _, it := range items {
frame := protocol.Revoked{
V: protocol.Version, Type: protocol.TypeRevoked,
ID: it.msgID, From: it.fromID, Reason: it.reason,
}
payload, encErr := encodeFrame(frame)
if encErr != nil {
continue
}
_ = a.down.PublishDown(ctx, it.endpointID, "", payload, port.PublishOpts{QoS: 1})
}
}
func isUnique(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "unique") || strings.Contains(msg, "constraint failed")
}