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 } _ = ctx 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 } // 异步且略推迟:必须让处理该端上行的 worker 先 PublishDown resp。 // 若与 resp 同时向本连接注入 group_event,会与 mochi InlineClient 互相等待。 ids := append([]string(nil), recipients...) go func() { time.Sleep(20 * time.Millisecond) seen := map[string]struct{}{} for _, id := range ids { if _, ok := seen[id]; ok { continue } seen[id] = struct{}{} _ = a.down.PublishDown(context.Background(), 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)