967 lines
25 KiB
Go
967 lines
25 KiB
Go
package group
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/rand"
|
|
"database/sql"
|
|
"errors"
|
|
"strings"
|
|
"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"
|
|
)
|
|
|
|
func (a *App) memberRequestCap() int {
|
|
n := a.maxMem * 2
|
|
if n < 256 {
|
|
n = 256
|
|
}
|
|
return n
|
|
}
|
|
|
|
// 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")
|
|
}
|
|
if err := a.rejectOversizedMemberList(len(req.Members)); err != nil {
|
|
return CreateResult{}, err
|
|
}
|
|
gid := req.ID
|
|
if gid == "" {
|
|
var genErr error
|
|
gid, genErr = generateGroupID()
|
|
if genErr != nil {
|
|
return CreateResult{}, genErr
|
|
}
|
|
}
|
|
now := a.nowMs()
|
|
failed := make([]MemberFail, 0)
|
|
uniq := dedupeMemberIns(req.Members, actorID)
|
|
room := a.maxMem - 1
|
|
if room < 0 {
|
|
room = 0
|
|
}
|
|
toCheck, overflow := splitMemberIns(uniq, room)
|
|
for _, m := range overflow {
|
|
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
|
}
|
|
added := make([]string, 0, len(toCheck))
|
|
for _, m := range toCheck {
|
|
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)
|
|
}
|
|
|
|
var inserted []string
|
|
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 := insertMemberTx(tx, gid, actorID, now); e != nil {
|
|
return e
|
|
}
|
|
inserted = inserted[:0]
|
|
for _, id := range added {
|
|
if e := endpointCheckTx(tx, id); e != nil {
|
|
failed = append(failed, MemberFail{ID: id, Code: failCode(e)})
|
|
continue
|
|
}
|
|
if e := insertMemberTx(tx, gid, id, now); e != nil {
|
|
return e
|
|
}
|
|
inserted = append(inserted, id)
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return CreateResult{}, err
|
|
}
|
|
|
|
notify := append([]string{actorID}, inserted...)
|
|
for _, id := range inserted {
|
|
a.emit(ctx, notify, 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
|
|
}
|
|
if err := a.rejectOversizedMemberList(len(req.Members)); 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)
|
|
now := a.nowMs()
|
|
uniq := dedupeMemberIns(req.Members, actorID)
|
|
already := memberSet(members)
|
|
candidates := make([]protocol.GroupMemberIn, 0, len(uniq))
|
|
for _, m := range uniq {
|
|
if _, ok := already[m.ID]; ok {
|
|
continue
|
|
}
|
|
candidates = append(candidates, m)
|
|
}
|
|
room := a.maxMem - len(members)
|
|
toCheck, overflow := splitMemberIns(candidates, room)
|
|
for _, m := range overflow {
|
|
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
|
}
|
|
added := make([]string, 0, len(toCheck))
|
|
for _, m := range toCheck {
|
|
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 {
|
|
return AddResult{Failed: failed}, nil
|
|
}
|
|
|
|
var inserted []string
|
|
var notify []string
|
|
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
curOwner, curMembers, e := loadGroupTx(tx, req.GroupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
if curOwner != actorID {
|
|
return errCode(protocol.CodeForbidden, "not owner")
|
|
}
|
|
present := memberSet(curMembers)
|
|
count := len(curMembers)
|
|
inserted = inserted[:0]
|
|
for _, id := range added {
|
|
if _, ok := present[id]; ok {
|
|
continue
|
|
}
|
|
if count >= a.maxMem {
|
|
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
|
continue
|
|
}
|
|
if checkErr := endpointCheckTx(tx, id); checkErr != nil {
|
|
failed = append(failed, MemberFail{ID: id, Code: failCode(checkErr)})
|
|
continue
|
|
}
|
|
if insErr := insertMemberTx(tx, req.GroupID, id, now); insErr != nil {
|
|
return insErr
|
|
}
|
|
present[id] = struct{}{}
|
|
count++
|
|
inserted = append(inserted, id)
|
|
}
|
|
notify = make([]string, 0, count)
|
|
for id := range present {
|
|
notify = append(notify, id)
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return AddResult{}, err
|
|
}
|
|
for _, id := range inserted {
|
|
a.emit(ctx, notify, 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
|
|
}
|
|
now := a.nowMs()
|
|
var revokes []revokeItem
|
|
var notify []string
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
owner, members, e := loadGroupTx(tx, req.GroupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
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")
|
|
}
|
|
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`,
|
|
req.GroupID, req.EndpointID); e != nil {
|
|
return e
|
|
}
|
|
if e := voidMemberDeliveriesTx(tx, req.GroupID, req.EndpointID, reasonLeftGroup, now, &revokes); e != nil {
|
|
return e
|
|
}
|
|
notify = append(without(members, req.EndpointID), req.EndpointID)
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
a.publishRevokes(ctx, revokes)
|
|
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
|
|
}
|
|
now := a.nowMs()
|
|
var revokes []revokeItem
|
|
var notify []string
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
owner, members, e := loadGroupTx(tx, req.GroupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
if !contains(members, actorID) {
|
|
return errCode(protocol.CodeNotMember, "not a member")
|
|
}
|
|
if owner == actorID {
|
|
return errCode(protocol.CodeOwnerCannotLeave, "owner cannot leave")
|
|
}
|
|
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`,
|
|
req.GroupID, actorID); e != nil {
|
|
return e
|
|
}
|
|
if e := voidMemberDeliveriesTx(tx, req.GroupID, actorID, reasonLeftGroup, now, &revokes); e != nil {
|
|
return e
|
|
}
|
|
notify = append(without(members, actorID), actorID)
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
a.publishRevokes(ctx, revokes)
|
|
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
|
|
}
|
|
now := a.nowMs()
|
|
var members []string
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
owner, cur, e := loadGroupTx(tx, req.GroupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
if owner != actorID {
|
|
return errCode(protocol.CodeForbidden, "not owner")
|
|
}
|
|
if !contains(cur, req.EndpointID) {
|
|
return errCode(protocol.CodeNotFound, "member not found")
|
|
}
|
|
if _, e := tx.Exec(`UPDATE groups SET owner_id = ? WHERE id = ?`, req.EndpointID, req.GroupID); e != nil {
|
|
return e
|
|
}
|
|
members = cur
|
|
return nil
|
|
})
|
|
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
|
|
}
|
|
now := a.nowMs()
|
|
var members []string
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
owner, cur, e := loadGroupTx(tx, req.GroupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
if owner != actorID {
|
|
return errCode(protocol.CodeForbidden, "not owner")
|
|
}
|
|
if _, e := tx.Exec(`UPDATE groups SET name = ? WHERE id = ?`, req.Name, req.GroupID); e != nil {
|
|
return e
|
|
}
|
|
members = cur
|
|
return nil
|
|
})
|
|
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
|
|
}
|
|
now := a.nowMs()
|
|
var revokes []revokeItem
|
|
var members []string
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
owner, cur, e := loadGroupTx(tx, req.GroupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
if owner != actorID {
|
|
return errCode(protocol.CodeForbidden, "not owner")
|
|
}
|
|
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
|
|
}
|
|
if _, e := tx.Exec(`DELETE FROM groups WHERE id = ?`, req.GroupID); e != nil {
|
|
return e
|
|
}
|
|
members = cur
|
|
return nil
|
|
})
|
|
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) {
|
|
if err := a.rejectOversizedMemberList(len(memberIDs)); err != nil {
|
|
return AddResult{}, err
|
|
}
|
|
_, members, err := a.loadGroup(ctx, groupID)
|
|
if err != nil {
|
|
return AddResult{}, err
|
|
}
|
|
failed := make([]MemberFail, 0)
|
|
uniq := dedupeIDs(memberIDs, "")
|
|
already := memberSet(members)
|
|
candidates := make([]string, 0, len(uniq))
|
|
for _, id := range uniq {
|
|
if _, ok := already[id]; ok {
|
|
continue
|
|
}
|
|
candidates = append(candidates, id)
|
|
}
|
|
room := a.maxMem - len(members)
|
|
toAdd, overflow := splitIDs(candidates, room)
|
|
for _, id := range overflow {
|
|
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
|
}
|
|
if len(toAdd) == 0 {
|
|
return AddResult{Failed: failed}, nil
|
|
}
|
|
now := a.nowMs()
|
|
var inserted []string
|
|
var notify []string
|
|
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
_, curMembers, e := loadGroupTx(tx, groupID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
present := memberSet(curMembers)
|
|
count := len(curMembers)
|
|
inserted = inserted[:0]
|
|
for _, id := range toAdd {
|
|
if _, ok := present[id]; ok {
|
|
continue
|
|
}
|
|
if count >= a.maxMem {
|
|
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
|
continue
|
|
}
|
|
if checkErr := endpointCheckTx(tx, id); checkErr != nil {
|
|
failed = append(failed, MemberFail{ID: id, Code: failCode(checkErr)})
|
|
continue
|
|
}
|
|
if insErr := insertMemberTx(tx, groupID, id, now); insErr != nil {
|
|
return insErr
|
|
}
|
|
present[id] = struct{}{}
|
|
count++
|
|
inserted = append(inserted, id)
|
|
}
|
|
notify = make([]string, 0, count)
|
|
for id := range present {
|
|
notify = append(notify, id)
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return AddResult{}, err
|
|
}
|
|
for _, id := range inserted {
|
|
a.emit(ctx, notify, 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 !protocol.ValidEndpointID(ownerID) {
|
|
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid owner")
|
|
}
|
|
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")
|
|
}
|
|
if err := a.rejectOversizedMemberList(len(members)); err != nil {
|
|
return CreateResult{}, err
|
|
}
|
|
now := a.nowMs()
|
|
failed := make([]MemberFail, 0)
|
|
uniq := dedupeMemberIns(members, ownerID)
|
|
room := a.maxMem - 1
|
|
toAdd, overflow := splitMemberIns(uniq, room)
|
|
for _, m := range overflow {
|
|
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
|
}
|
|
|
|
var inserted []string
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
if e := endpointCheckTx(tx, ownerID); e != nil {
|
|
if protoCode(e) == protocol.CodeInvalidTarget {
|
|
return errCode(protocol.CodeInvalidTarget, "owner not found")
|
|
}
|
|
if protoCode(e) == protocol.CodeEndpointDisabled {
|
|
return errCode(protocol.CodeEndpointDisabled, "owner disabled")
|
|
}
|
|
return e
|
|
}
|
|
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 := insertMemberTx(tx, gid, ownerID, now); e != nil {
|
|
return e
|
|
}
|
|
inserted = inserted[:0]
|
|
for _, m := range toAdd {
|
|
if checkErr := endpointCheckTx(tx, m.ID); checkErr != nil {
|
|
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
|
continue
|
|
}
|
|
if e := insertMemberTx(tx, gid, m.ID, now); e != nil {
|
|
return e
|
|
}
|
|
inserted = append(inserted, m.ID)
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return CreateResult{}, err
|
|
}
|
|
notify := append([]string{ownerID}, inserted...)
|
|
for _, id := range inserted {
|
|
a.emit(ctx, notify, gid, eventMemberAdded, id, now)
|
|
}
|
|
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) rejectOversizedMemberList(n int) error {
|
|
if n > a.memberRequestCap() {
|
|
return errCode(protocol.CodeBadRequest, "too many members")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func loadGroupTx(tx *sql.Tx, groupID string) (owner string, members []string, err error) {
|
|
err = tx.QueryRow(`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 := tx.Query(`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 endpointCheckTx(tx *sql.Tx, id string) error {
|
|
var enabled int
|
|
err := tx.QueryRow(`SELECT enabled FROM endpoints WHERE id = ?`, id).Scan(&enabled)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return errCode(protocol.CodeInvalidTarget, "target not found")
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if enabled == 0 {
|
|
return errCode(protocol.CodeEndpointDisabled, "endpoint disabled")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func insertMemberTx(tx *sql.Tx, groupID, endpointID string, now int64) error {
|
|
_, err := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
|
groupID, endpointID, now)
|
|
if err == nil {
|
|
return nil
|
|
}
|
|
if isForeignKey(err) {
|
|
return errCode(protocol.CodeNotFound, "group not found")
|
|
}
|
|
if isUnique(err) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
|
|
func isForeignKey(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
return strings.Contains(strings.ToLower(err.Error()), "foreign key")
|
|
}
|
|
|
|
func dedupeMemberIns(members []protocol.GroupMemberIn, skipID string) []protocol.GroupMemberIn {
|
|
seen := make(map[string]struct{}, len(members)+1)
|
|
if skipID != "" {
|
|
seen[skipID] = struct{}{}
|
|
}
|
|
out := make([]protocol.GroupMemberIn, 0, len(members))
|
|
for _, m := range members {
|
|
if _, ok := seen[m.ID]; ok {
|
|
continue
|
|
}
|
|
seen[m.ID] = struct{}{}
|
|
out = append(out, m)
|
|
}
|
|
return out
|
|
}
|
|
|
|
func splitMemberIns(members []protocol.GroupMemberIn, room int) (keep, overflow []protocol.GroupMemberIn) {
|
|
if room < 0 {
|
|
room = 0
|
|
}
|
|
if len(members) <= room {
|
|
return members, nil
|
|
}
|
|
return members[:room], members[room:]
|
|
}
|
|
|
|
func dedupeIDs(ids []string, skipID string) []string {
|
|
seen := make(map[string]struct{}, len(ids)+1)
|
|
if skipID != "" {
|
|
seen[skipID] = struct{}{}
|
|
}
|
|
out := make([]string, 0, len(ids))
|
|
for _, id := range ids {
|
|
if id == "" {
|
|
continue
|
|
}
|
|
if _, ok := seen[id]; ok {
|
|
continue
|
|
}
|
|
seen[id] = struct{}{}
|
|
out = append(out, id)
|
|
}
|
|
return out
|
|
}
|
|
|
|
func splitIDs(ids []string, room int) (keep, overflow []string) {
|
|
if room < 0 {
|
|
room = 0
|
|
}
|
|
if len(ids) <= room {
|
|
return ids, nil
|
|
}
|
|
return ids[:room], ids[room:]
|
|
}
|
|
|
|
func memberSet(ss []string) map[string]struct{} {
|
|
m := make(map[string]struct{}, len(ss))
|
|
for _, s := range ss {
|
|
m[s] = struct{}{}
|
|
}
|
|
return m
|
|
}
|
|
|
|
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)
|