From 16ece09a97f0a82ffc8aa68814db1bca1ec005bb Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 07:33:59 +0800 Subject: [PATCH] =?UTF-8?q?Revert=20"feat:=20=E5=AE=9E=E7=8E=B0=E8=BA=AB?= =?UTF-8?q?=E4=BB=BD=E8=B5=84=E6=96=99=E3=80=81=E5=AF=B9=E8=AF=9D=E5=AF=86?= =?UTF-8?q?=E7=A0=81=E3=80=81=E5=9C=A8=E7=BA=BF=E7=9B=AE=E5=BD=95=E4=B8=8E?= =?UTF-8?q?=E7=BE=A4=E7=AE=A1=E7=90=86"?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit a78ab0d547b32439eff6fe159a8b51e00b62f74d. --- docs/DEVIATIONS.md | 65 --- internal/app/group/app.go | 727 ------------------------ internal/app/group/group_test.go | 412 -------------- internal/app/group/void.go | 176 ------ internal/app/identity/app.go | 131 ----- internal/app/identity/errors.go | 7 - internal/app/identity/register.go | 59 ++ internal/app/identity/self.go | 222 -------- internal/app/identity/self_talk_test.go | 314 ---------- internal/app/identity/service.go | 11 +- internal/app/identity/stub.go | 8 +- internal/app/identity/talk.go | 181 ------ internal/app/presence/app.go | 353 ------------ internal/app/presence/presence_test.go | 176 ------ 14 files changed, 65 insertions(+), 2777 deletions(-) delete mode 100644 internal/app/group/app.go delete mode 100644 internal/app/group/group_test.go delete mode 100644 internal/app/group/void.go delete mode 100644 internal/app/identity/app.go delete mode 100644 internal/app/identity/errors.go delete mode 100644 internal/app/identity/self.go delete mode 100644 internal/app/identity/self_talk_test.go delete mode 100644 internal/app/identity/talk.go delete mode 100644 internal/app/presence/app.go delete mode 100644 internal/app/presence/presence_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index f29f66c..f7c254f 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -427,71 +427,6 @@ - 备选方案:16/32 位。 - 影响:无产品行为差异。 -### I2 / I3 / I4 2026-09-30 - -1. **服务方法可直接调用,未接 MQTT 分发** - - 原条款:端协议帧经 broker 上行分发到 app。 - - 实际做法:`identity.App` / `presence.App` / `group.App` 实现 Service 方法;单测直接调用,不经 mochi。`cmd/nixmsg`/`wire.go` 未改。 - - 原因:任务要求可被协议分发调用的服务方法 + 直接调用验证;N3 连接事件与总控接线另波。 - - 备选方案:本分支顺带改 wire(与隔离冲突)。 - - 影响:合入后需总控/N 接线 HandleUplink → 各 Service;MQTT 帧路径未测。 - -2. **在线状态:库字段 + 可注入连接表 + 本包握手表** - - 原条款:在线以真实连接为准;依赖 N3 连接事件。 - - 实际做法:`presence.ConnTable` 可注入;另用 `SetOnline`/`SetOffline` 维护内存表并写 `endpoints.online_since`/`offline_since`;查询优先 ConnTable,其次内存表,再回退库字段(`online_since` 晚于 `offline_since` 或后者为空)。 - - 原因:N3 可能尚未合入,不阻塞 I3。 - - 备选方案:阻塞等 N3。 - - 影响:未接线时须调用 SetOnline/SetOffline 或写库字段;拔网线心跳超时属 N 线,本波单测不覆盖。 - -3. **presence / group_event 经 Downlink QoS 0,可 nil** - - 原条款:上下线与群事件尽力推送、不落库。 - - 实际做法:注入 `port.Downlink` 时编码帧并 `PublishDown`(presence/group_event QoS 0;退群已推送投递的 `revoked` 用 QoS 1);Downlink 为 nil 时跳过推送,业务库操作仍完成。 - - 原因:无 MQTT 时仍可测库逻辑。 - - 备选方案:强制假 Downlink。 - - 影响:接线后必须注入真实 Downlink 才有通知。 - -4. **self.logout 踢线可选** - - 原条款:回 resp 后断开连接。 - - 实际做法:清 `session_hash`;若注入 `port.ConnControl` 则 `Disconnect`,否则仅清令牌。 - - 原因:未接 broker。 - - 备选方案:无。 - - 影响:接线方应注入 ConnControl。 - -5. **session_hash 存 SHA-256 十六进制** - - 原条款:库中存会话令牌 SHA-256。 - - 实际做法:`hex.EncodeToString(hash)` 写入 TEXT 列。 - - 原因:文档未规定编码;十六进制便于调试与比对。 - - 备选方案:BLOB/Base64。 - - 影响:N3 校验须用同一编码。 - -6. **self.update 空 name 不写库** - - 原条款:可更新 name。 - - 实际做法:`name` 非空才 UPDATE name;仅改 `default_delay_ms` 时不碰 name(JSON omitempty 无法区分省略与空串)。 - - 原因:避免误清空名称。 - - 备选方案:用指针字段区分。 - - 影响:端无法通过协议把名称改成空字符串(可用空格等)。 - -7. **进群密码校验不产生单聊授权** - - 原条款:拉人须当次带密码;已有授权不能代替。 - - 实际做法:`CheckTalkPasswordForJoin` 只校验,不写 `talk_grants`。 - - 原因:与 F15「进群仍要密码」一致,避免进群副作用放宽单聊。 - - 备选方案:校验成功顺带写 password 授权。 - - 影响:仅进群成功后,单聊仍须 unlock/发送带密。 - -8. **群作废在 group 包内写 deliveries/messages** - - 原条款:退群/踢人/解散的投递作废属 7.6,消息线亦相关。 - - 实际做法:I4 在 `group` 写操作里直接改 `pending→rejected`、`scheduled→completed/group_dissolved`,删正文行,需要时插消息级回执,已推送则经 Downlink 发 `revoked`。 - - 原因:I4 验收依赖作废规则;M 线完整推送循环可能未合入。 - - 备选方案:只调 message 钩子(接口尚未暴露)。 - - 影响:与后续 M 作废路径需保持同语义,避免重复作废。 - -9. **UnlockTalk / SelfChangeLoginPassword / CheckTalkPasswordForJoin 增加 remoteIP** - - 原条款:锁定按发送方+对方 / 编号+IP。 - - 实际做法:Service 方法增加 `remoteIP` 参数供锁定计数;T0.4 Stub 同步改签名。 - - 原因:无 ConnInfo 的直接调用测试需要显式 IP。 - - 备选方案:塞进 context。 - - 影响:协议分发接线时从 `ConnInfo.RemoteIP` 传入。 - ## 后台接口 A ### A1 2026-09-30 diff --git a/internal/app/group/app.go b/internal/app/group/app.go deleted file mode 100644 index 4de8138..0000000 --- a/internal/app/group/app.go +++ /dev/null @@ -1,727 +0,0 @@ -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) diff --git a/internal/app/group/group_test.go b/internal/app/group/group_test.go deleted file mode 100644 index 44c2231..0000000 --- a/internal/app/group/group_test.go +++ /dev/null @@ -1,412 +0,0 @@ -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) - } - } -} diff --git a/internal/app/group/void.go b/internal/app/group/void.go deleted file mode 100644 index 7411ad0..0000000 --- a/internal/app/group/void.go +++ /dev/null @@ -1,176 +0,0 @@ -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") -} diff --git a/internal/app/identity/app.go b/internal/app/identity/app.go deleted file mode 100644 index c22779a..0000000 --- a/internal/app/identity/app.go +++ /dev/null @@ -1,131 +0,0 @@ -package identity - -import ( - "context" - "encoding/hex" - "log/slog" - "net/http" - "time" - - "git.asio.asia/nixevol/NixMsg/internal/app/port" - "git.asio.asia/nixevol/NixMsg/internal/auth" - "git.asio.asia/nixevol/NixMsg/internal/protocol" - "git.asio.asia/nixevol/NixMsg/internal/store" -) - -const ( - grantKindPassword = "password" - grantKindReply = "reply" -) - -// Config 是身份服务依赖(注册 + self + 对话密码)。 -type Config struct { - DB *store.DB - Hash auth.HashPool - Locks auth.LoginLocks - Logger *slog.Logger - Now func() time.Time - // ClientIP 仅注册 HTTP 用。 - ClientIP func(*http.Request) string - - // Sessions 签发会话令牌;改登录密码必填。 - Sessions auth.SessionTokens - // MaxScheduleSeconds 限制 self.update 的 default_delay_ms。 - MaxScheduleSeconds int64 - // ConnControl 可选:logout 后踢线;未接线时为 nil。 - ConnControl port.ConnControl -} - -// App 实现 identity.Service(含 I1 注册与 I2 self/对话密码)。 -type App struct { - handler *RegisterHandler - db *store.DB - hash auth.HashPool - locks auth.LoginLocks - sessions auth.SessionTokens - maxScheduleSeconds int64 - connCtrl port.ConnControl - nowFn func() time.Time -} - -// New 构造完整身份服务。 -func New(cfg Config) *App { - if cfg.Logger == nil { - cfg.Logger = slog.Default() - } - if cfg.Now == nil { - cfg.Now = time.Now - } - if cfg.ClientIP == nil { - cfg.ClientIP = clientIPFromRemoteAddr - } - if cfg.Sessions == nil { - cfg.Sessions = auth.NewStubSessionTokens() - } - if cfg.Locks == nil { - cfg.Locks = auth.NewStubLoginLocks() - } - h := NewRegisterHandler(RegisterConfig{ - DB: cfg.DB, - Hash: cfg.Hash, - Locks: cfg.Locks, - Logger: cfg.Logger, - Now: cfg.Now, - ClientIP: cfg.ClientIP, - }) - return &App{ - handler: h, - db: cfg.DB, - hash: cfg.Hash, - locks: cfg.Locks, - sessions: cfg.Sessions, - maxScheduleSeconds: cfg.MaxScheduleSeconds, - connCtrl: cfg.ConnControl, - nowFn: cfg.Now, - } -} - -// NewServer 兼容 I1:用注册配置构造 Service(会话令牌用 Stub)。 -func NewServer(cfg RegisterConfig) *App { - return New(Config{ - DB: cfg.DB, - Hash: cfg.Hash, - Locks: cfg.Locks, - Logger: cfg.Logger, - Now: cfg.Now, - ClientIP: cfg.ClientIP, - Sessions: auth.NewStubSessionTokens(), - }) -} - -func (a *App) now() time.Time { return a.nowFn() } - -// Handler 返回可挂载的注册 HTTP 处理器。 -func (a *App) Handler() http.Handler { return a.handler } - -// Register 实现自助注册。 -func (a *App) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) { - preq := &protocol.RegisterRequest{ - RegistrationCode: req.RegistrationCode, - ID: req.ID, - LoginPassword: req.LoginPassword, - Name: req.Name, - TalkPassword: req.TalkPassword, - } - ip := req.RemoteIP - if ip == "" { - ip = "0.0.0.0" - } - result, apiErr := a.handler.register(ctx, preq, ip) - if apiErr != nil { - return RegisterResult{}, apiErr - } - return result, nil -} - -func encodeSessionHash(hash []byte) string { - return hex.EncodeToString(hash) -} - -var _ Service = (*App)(nil) -var _ http.Handler = (*RegisterHandler)(nil) diff --git a/internal/app/identity/errors.go b/internal/app/identity/errors.go deleted file mode 100644 index ab0466f..0000000 --- a/internal/app/identity/errors.go +++ /dev/null @@ -1,7 +0,0 @@ -package identity - -import "git.asio.asia/nixevol/NixMsg/internal/protocol" - -func errCode(code, msg string) *protocol.Error { - return &protocol.Error{Code: code, Message: msg} -} diff --git a/internal/app/identity/register.go b/internal/app/identity/register.go index 5f27b52..020a500 100644 --- a/internal/app/identity/register.go +++ b/internal/app/identity/register.go @@ -247,6 +247,65 @@ func (h *RegisterHandler) logResult(result, id, ip string) { h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip) } +// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。 +type Server struct { + handler *RegisterHandler +} + +// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。 +func NewServer(cfg RegisterConfig) *Server { + return &Server{handler: NewRegisterHandler(cfg)} +} + +// Handler 返回可挂载的注册 HTTP 处理器。 +func (s *Server) Handler() http.Handler { return s.handler } + +// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。 +func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) { + preq := &protocol.RegisterRequest{ + RegistrationCode: req.RegistrationCode, + ID: req.ID, + LoginPassword: req.LoginPassword, + Name: req.Name, + TalkPassword: req.TalkPassword, + } + ip := req.RemoteIP + if ip == "" { + ip = "0.0.0.0" + } + result, err := s.handler.register(ctx, preq, ip) + if err != nil { + return RegisterResult{}, err + } + return result, nil +} + +func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) { + return SelfInfo{}, ErrNotImplemented +} +func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error { + return ErrNotImplemented +} +func (s *Server) SelfSetTalkPassword(context.Context, string, string) error { + return ErrNotImplemented +} +func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) { + return "", ErrNotImplemented +} +func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented } +func (s *Server) UnlockTalk(context.Context, string, string, string) error { + return ErrNotImplemented +} +func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) { + return false, nil +} +func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented } +func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented } +func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented } + +var _ Service = (*Server)(nil) +var _ http.Handler = (*RegisterHandler)(nil) + func setCORS(w http.ResponseWriter) { w.Header().Set("Access-Control-Allow-Origin", "*") } diff --git a/internal/app/identity/self.go b/internal/app/identity/self.go deleted file mode 100644 index 912039d..0000000 --- a/internal/app/identity/self.go +++ /dev/null @@ -1,222 +0,0 @@ -package identity - -import ( - "context" - "database/sql" - "errors" - - "git.asio.asia/nixevol/NixMsg/internal/app/port" - "git.asio.asia/nixevol/NixMsg/internal/auth" - "git.asio.asia/nixevol/NixMsg/internal/protocol" -) - -// SelfGet 返回自己的资料。 -func (a *App) SelfGet(ctx context.Context, endpointID string) (SelfInfo, error) { - if !protocol.ValidEndpointID(endpointID) { - return SelfInfo{}, errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - var info SelfInfo - var talk sql.NullString - err := a.db.Read.QueryRowContext(ctx, ` -SELECT id, name, default_delay_ms, talk_hash -FROM endpoints WHERE id = ?`, endpointID).Scan(&info.ID, &info.Name, &info.DefaultDelayMs, &talk) - if errors.Is(err, sql.ErrNoRows) { - return SelfInfo{}, errCode(protocol.CodeNotFound, "endpoint not found") - } - if err != nil { - return SelfInfo{}, err - } - info.TalkPasswordSet = talk.Valid && talk.String != "" - return info, nil -} - -// SelfUpdate 更新名称与默认延迟。 -func (a *App) SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error { - if !protocol.ValidEndpointID(endpointID) { - return errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - if req == nil { - return errCode(protocol.CodeBadRequest, "nil request") - } - if err := req.Validate(); err != nil { - return err - } - if req.Name == "" && req.DefaultDelayMs == nil { - return errCode(protocol.CodeBadRequest, "nothing to update") - } - if req.DefaultDelayMs != nil && a.maxScheduleSeconds > 0 { - maxMs := a.maxScheduleSeconds * 1000 - if *req.DefaultDelayMs > maxMs { - return errCode(protocol.CodeBadRequest, "default_delay_ms exceeds max_schedule_seconds") - } - } - - return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - var exists int - if err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, endpointID).Scan(&exists); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return errCode(protocol.CodeNotFound, "endpoint not found") - } - return err - } - if req.Name != "" { - if _, err := tx.Exec(`UPDATE endpoints SET name = ? WHERE id = ?`, req.Name, endpointID); err != nil { - return err - } - } - if req.DefaultDelayMs != nil { - if _, err := tx.Exec(`UPDATE endpoints SET default_delay_ms = ? WHERE id = ?`, *req.DefaultDelayMs, endpointID); err != nil { - return err - } - } - return nil - }) -} - -// SelfSetTalkPassword 设置或清除对话密码,并增加 talk_version;改密清零对方锁定计数。 -func (a *App) SelfSetTalkPassword(ctx context.Context, endpointID, talkPassword string) error { - if !protocol.ValidEndpointID(endpointID) { - return errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - if !protocol.ValidTalkPassword(talkPassword) { - return errCode(protocol.CodeBadRequest, "invalid talk_password") - } - - var talkHash any - if talkPassword != "" { - if a.hash == nil { - return errors.New("identity: hash pool required") - } - phc, err := a.hash.Hash(ctx, auth.PasswordTalk, talkPassword) - if err != nil { - return err - } - talkHash = phc - } - - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - res, err := tx.Exec(` -UPDATE endpoints -SET talk_hash = ?, talk_version = talk_version + 1 -WHERE id = ?`, talkHash, endpointID) - if err != nil { - return err - } - n, _ := res.RowsAffected() - if n == 0 { - return errCode(protocol.CodeNotFound, "endpoint not found") - } - return nil - }) - if err != nil { - return err - } - // D24:改密清零按对方计的对话密码失败计数。 - a.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: endpointID}) - return nil -} - -// SelfChangeLoginPassword 要求旧密码;成功返回新 session_token,旧令牌作废,当前连接由调用方保留。 -func (a *App) SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (string, error) { - if !protocol.ValidEndpointID(endpointID) { - return "", errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - if protocol.LoginPasswordForbiddenPrefix(newPassword) { - return "", errCode(protocol.CodeBadRequest, "login password must not start with nst_") - } - if !protocol.ValidLoginPassword(newPassword) || newPassword == "" { - return "", errCode(protocol.CodeBadRequest, "invalid new_password") - } - if a.hash == nil || a.sessions == nil { - return "", errors.New("identity: hash/sessions required") - } - - var loginHash string - err := a.db.Read.QueryRowContext(ctx, `SELECT login_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&loginHash) - if errors.Is(err, sql.ErrNoRows) { - return "", errCode(protocol.CodeNotFound, "endpoint not found") - } - if err != nil { - return "", err - } - - if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}); locked { - return "", errCode(protocol.CodeRateLimited, "login locked") - } - if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}); locked { - return "", errCode(protocol.CodeRateLimited, "login locked") - } - - ok, err := a.hash.Verify(ctx, auth.PasswordLogin, oldPassword, loginHash) - if err != nil { - return "", err - } - if !ok { - a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}) - a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}) - return "", errCode(protocol.CodeUnauthorized, "old password invalid") - } - - newHash, err := a.hash.Hash(ctx, auth.PasswordLogin, newPassword) - if err != nil { - return "", err - } - token, tokenHash, err := a.sessions.Issue(ctx) - if err != nil { - return "", err - } - nowMs := a.now().UnixMilli() - hashHex := encodeSessionHash(tokenHash) - - err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - res, e := tx.Exec(` -UPDATE endpoints -SET login_hash = ?, session_hash = ?, session_issued_at = ?, session_used_at = ? -WHERE id = ?`, newHash, hashHex, nowMs, nowMs, endpointID) - if e != nil { - return e - } - n, _ := res.RowsAffected() - if n == 0 { - return errCode(protocol.CodeNotFound, "endpoint not found") - } - return nil - }) - if err != nil { - return "", err - } - return token, nil -} - -// SelfLogout 清空会话令牌;若注入了 ConnControl 则断开当前连接。 -func (a *App) SelfLogout(ctx context.Context, endpointID string) error { - if !protocol.ValidEndpointID(endpointID) { - return errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - res, e := tx.Exec(` -UPDATE endpoints -SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL -WHERE id = ?`, endpointID) - if e != nil { - return e - } - n, _ := res.RowsAffected() - if n == 0 { - return errCode(protocol.CodeNotFound, "endpoint not found") - } - return nil - }) - if err != nil { - return err - } - if a.connCtrl != nil { - _ = a.connCtrl.Disconnect(ctx, endpointID, "", port.DisconnectNormal) - } - return nil -} - -// Disable / Enable / Delete 属 I5,此处保留未实现。 -func (a *App) Disable(context.Context, string) error { return ErrNotImplemented } -func (a *App) Enable(context.Context, string) error { return ErrNotImplemented } -func (a *App) Delete(context.Context, string) error { return ErrNotImplemented } diff --git a/internal/app/identity/self_talk_test.go b/internal/app/identity/self_talk_test.go deleted file mode 100644 index 4a20087..0000000 --- a/internal/app/identity/self_talk_test.go +++ /dev/null @@ -1,314 +0,0 @@ -package identity_test - -import ( - "context" - "database/sql" - "errors" - "path/filepath" - "testing" - "time" - - "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" -) - -func openIdentity(t *testing.T) (*identity.App, *store.DB, *auth.MemoryLocks) { - t.Helper() - dir := t.TempDir() - db, err := store.Open(filepath.Join(dir, "data"), "FULL") - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { _ = db.Close() }) - fixed := time.UnixMilli(1_700_000_000_000) - locks := auth.NewLoginLocks() - locks.SetClock(func() time.Time { return fixed }) - app := identity.New(identity.Config{ - DB: db, - Hash: auth.NewStubHashPool(), - Locks: locks, - Sessions: auth.NewSessionTokens(), - MaxScheduleSeconds: int64(config.Default().Limits.MaxScheduleSeconds), - Now: func() time.Time { return fixed }, - }) - return app, db, locks -} - -func insertEP(t *testing.T, db *store.DB, id, loginPW string) { - 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,1,?)`, id, id, "stub$"+loginPW, nil, 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 TestF15TalkPasswordAuth(t *testing.T) { - t.Parallel() - app, db, _ := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "alice", "password1") - insertEP(t, db, "bob", "password1") - - if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil { - t.Fatal(err) - } - - // A 不带密码失败 - if err := app.UnlockTalk(ctx, "alice", "bob", "", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordRequired { - t.Fatalf("want talk_password_required got %v", err) - } - ok, err := app.HasTalkGrant(ctx, "alice", "bob") - if err != nil || ok { - t.Fatalf("grant=%v err=%v", ok, err) - } - - // 带对后成功,之后无密码也有授权 - if err = app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil { - t.Fatal(err) - } - ok, err = app.HasTalkGrant(ctx, "alice", "bob") - if err != nil || !ok { - t.Fatalf("grant=%v err=%v", ok, err) - } - - // B 改密后旧授权失效 - if err = app.SelfSetTalkPassword(ctx, "bob", "newsecret"); err != nil { - t.Fatal(err) - } - ok, err = app.HasTalkGrant(ctx, "alice", "bob") - if err != nil || ok { - t.Fatalf("after change grant=%v err=%v", ok, err) - } - if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordInvalid { - t.Fatalf("want invalid got %v", err) - } - if err := app.UnlockTalk(ctx, "alice", "bob", "newsecret", "1.1.1.1"); err != nil { - t.Fatal(err) - } -} - -func TestF15ReplyGrantAndChange(t *testing.T) { - t.Parallel() - app, db, _ := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "alice", "password1") - insertEP(t, db, "bob", "password1") - if err := app.SelfSetTalkPassword(ctx, "alice", "alice-pw"); err != nil { - t.Fatal(err) - } - // B 先给 A 发 → 写入 reply 授权(A 可回 B 免密;这里记的是 bob→alice 的授权给「alice 作为接收方」... - // 回复授权:对方曾成功提交发给我的单聊 → 我对对方有 reply 权。 - // 即 B 发给 A 后,A 对 B 有授权(sender=alice, target=bob)。 - if err := app.RecordReplyGrant(ctx, "alice", "bob"); err != nil { - t.Fatal(err) - } - // 但 bob 还没设密码,alice→bob 本就不需要。给 bob 设密后验证 reply: - if err := app.SelfSetTalkPassword(ctx, "bob", "bob-pw"); err != nil { - t.Fatal(err) - } - // 重新记 reply:B 发给 A 成功后 A 获得对 B 的回复权 - if err := app.RecordReplyGrant(ctx, "alice", "bob"); err != nil { - t.Fatal(err) - } - ok, err := app.HasTalkGrant(ctx, "alice", "bob") - if err != nil || !ok { - t.Fatalf("reply grant=%v err=%v", ok, err) - } - if err = app.SelfSetTalkPassword(ctx, "bob", "bob-pw2"); err != nil { - t.Fatal(err) - } - ok, err = app.HasTalkGrant(ctx, "alice", "bob") - if err != nil || ok { - t.Fatalf("after change reply should die grant=%v", ok) - } -} - -func TestF15JoinNeedsPasswordDespiteGrant(t *testing.T) { - t.Parallel() - app, db, _ := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "alice", "password1") - insertEP(t, db, "bob", "password1") - if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil { - t.Fatal(err) - } - if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil { - t.Fatal(err) - } - // 已有单聊授权,进群仍要密码 - if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordRequired { - t.Fatalf("want required got %v", err) - } - if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "wrong", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordInvalid { - t.Fatalf("want invalid got %v", err) - } - if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil { - t.Fatal(err) - } -} - -func TestF15SubmittedMessageUnaffectedByPasswordChange(t *testing.T) { - t.Parallel() - idApp, db, locks := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "alice", "password1") - insertEP(t, db, "bob", "password1") - if err := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil { - t.Fatal(err) - } - lim := message.LimitsFromConfig(config.Default().Limits) - lim.RequestsPerSecond = 0 - msgApp := message.New(db, lim, auth.NewStubHashPool(), - message.WithNow(func() time.Time { return time.UnixMilli(1_700_000_000_000) }), - message.WithLocks(locks), - ) - req := &protocol.Send{ - V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "m1", - To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"}, - Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"}, - DelayMs: ptrInt64(60_000), - TalkPassword: "secret", - } - res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "1.1.1.1"}, req) - if err != nil { - t.Fatal(err) - } - if res.State != message.StateScheduled { - t.Fatalf("state=%s", res.State) - } - if err := idApp.SelfSetTalkPassword(ctx, "bob", "changed"); err != nil { - t.Fatal(err) - } - // 已提交消息行仍在且状态不变 - var state string - if err := db.Read.QueryRow(`SELECT state FROM messages WHERE sender_id=? AND id=?`, "alice", "m1").Scan(&state); err != nil { - t.Fatal(err) - } - if state != message.StateScheduled { - t.Fatalf("message state changed to %s", state) - } -} - -func TestF15TargetLockAndExistingGrant(t *testing.T) { - t.Parallel() - app, db, locks := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "bob", "password1") - insertEP(t, db, "authd", "password1") - if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil { - t.Fatal(err) - } - if err := app.UnlockTalk(ctx, "authd", "bob", "secret", "9.9.9.9"); err != nil { - t.Fatal(err) - } - // 50 次错误触发对方总数锁 - for i := 0; i < 50; i++ { - id := "u" + string(rune('0'+i/100)) + string(rune('0'+(i/10)%10)) + string(rune('0'+i%10)) - insertEP(t, db, id, "password1") - _ = app.UnlockTalk(ctx, id, "bob", "wrong", "2.2.2.2") - } - locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}) - if !locked { - t.Fatal("expected talk target lock") - } - // 正确密码也暂时无法解锁 - insertEP(t, db, "newbie", "password1") - if err := app.UnlockTalk(ctx, "newbie", "bob", "secret", "3.3.3.3"); protoCode(err) != protocol.CodeRateLimited { - t.Fatalf("want rate_limited got %v", err) - } - // 已有授权仍可用 - ok, err := app.HasTalkGrant(ctx, "authd", "bob") - if err != nil || !ok { - t.Fatalf("authd grant=%v err=%v", ok, err) - } - // 改密清零 - if err := app.SelfSetTalkPassword(ctx, "bob", "secret2"); err != nil { - t.Fatal(err) - } - locked, _ = locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}) - if locked { - t.Fatal("lock should clear on password change") - } -} - -func TestSelfLoginPasswordAndLogout(t *testing.T) { - t.Parallel() - app, db, locks := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "alice", "oldpass12") - tok, err := app.SelfChangeLoginPassword(ctx, "alice", "wrongpass", "newpass12", "1.1.1.1") - if protoCode(err) != protocol.CodeUnauthorized || tok != "" { - t.Fatalf("got tok=%q err=%v", tok, err) - } - locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: "alice", IP: "1.1.1.1"}) // 确保 Fail 路径可调用 - - tok, err = app.SelfChangeLoginPassword(ctx, "alice", "oldpass12", "newpass12", "1.1.1.1") - if err != nil || tok == "" || !protocol.ValidEndpointID("alice") { - t.Fatalf("tok=%q err=%v", tok, err) - } - if !auth.NewSessionTokens().LooksLikeSessionToken(tok) { - t.Fatalf("token prefix %q", tok) - } - var hash sql.NullString - if err = db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id=?`, "alice").Scan(&hash); err != nil || !hash.Valid { - t.Fatal(err) - } - info, err := app.SelfGet(ctx, "alice") - if err != nil || info.ID != "alice" { - t.Fatal(err) - } - name := "门口" - delay := int64(10000) - if err := app.SelfUpdate(ctx, "alice", &protocol.SelfUpdate{ - V: protocol.Version, Type: protocol.TypeSelfUpdate, RID: "1", Name: name, DefaultDelayMs: &delay, - }); err != nil { - t.Fatal(err) - } - info, _ = app.SelfGet(ctx, "alice") - if info.Name != name || info.DefaultDelayMs != delay { - t.Fatalf("%+v", info) - } - if err := app.SelfLogout(ctx, "alice"); err != nil { - t.Fatal(err) - } - if err := db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id=?`, "alice").Scan(&hash); err != nil { - t.Fatal(err) - } - if hash.Valid { - t.Fatal("session should be cleared") - } -} - -func TestUnlockSelfAndNoPassword(t *testing.T) { - t.Parallel() - app, db, _ := openIdentity(t) - ctx := context.Background() - insertEP(t, db, "alice", "password1") - insertEP(t, db, "bob", "password1") - if err := app.UnlockTalk(ctx, "alice", "alice", "", ""); err != nil { - t.Fatal(err) - } - if err := app.UnlockTalk(ctx, "alice", "bob", "anything", ""); err != nil { - t.Fatal(err) - } -} - -func ptrInt64(v int64) *int64 { return &v } diff --git a/internal/app/identity/service.go b/internal/app/identity/service.go index 741c555..e5ce538 100644 --- a/internal/app/identity/service.go +++ b/internal/app/identity/service.go @@ -46,16 +46,13 @@ type Service interface { SelfGet(ctx context.Context, endpointID string) (SelfInfo, error) SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error SelfSetTalkPassword(ctx context.Context, endpointID string, talkPassword string) error - // SelfChangeLoginPassword 校验旧密码后换新密码并签发新会话令牌;remoteIP 计入登录锁定。 - SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error) + SelfChangeLoginPassword(ctx context.Context, endpointID string, oldPassword, newPassword string) (sessionToken string, err error) SelfLogout(ctx context.Context, endpointID string) error - // UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。 - UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error - // HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。 + // UnlockTalk 校验并写入对话密码授权(第 6.6 节 unlock)。 + UnlockTalk(ctx context.Context, senderID, targetID, talkPassword string) error + // HasTalkGrant 查询发送方对目标是否有有效授权。 HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) - // CheckTalkPasswordForJoin 加人时校验对话密码:已有单聊授权不能代替,必须当次带对。 - CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error // Disable 停用端并作废相关消息/令牌(第 7.6 节)。 Disable(ctx context.Context, endpointID string) error diff --git a/internal/app/identity/stub.go b/internal/app/identity/stub.go index 307b73d..41738a3 100644 --- a/internal/app/identity/stub.go +++ b/internal/app/identity/stub.go @@ -27,13 +27,13 @@ func (s *Stub) SelfSetTalkPassword(context.Context, string, string) error { return ErrNotImplemented } -func (s *Stub) SelfChangeLoginPassword(context.Context, string, string, string, string) (string, error) { +func (s *Stub) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) { return "", ErrNotImplemented } func (s *Stub) SelfLogout(context.Context, string) error { return ErrNotImplemented } -func (s *Stub) UnlockTalk(context.Context, string, string, string, string) error { +func (s *Stub) UnlockTalk(context.Context, string, string, string) error { return ErrNotImplemented } @@ -41,10 +41,6 @@ func (s *Stub) HasTalkGrant(context.Context, string, string) (bool, error) { return false, nil } -func (s *Stub) CheckTalkPasswordForJoin(context.Context, string, string, string, string) error { - return ErrNotImplemented -} - func (s *Stub) Disable(context.Context, string) error { return ErrNotImplemented } func (s *Stub) Enable(context.Context, string) error { return ErrNotImplemented } func (s *Stub) Delete(context.Context, string) error { return ErrNotImplemented } diff --git a/internal/app/identity/talk.go b/internal/app/identity/talk.go deleted file mode 100644 index cdc8eaa..0000000 --- a/internal/app/identity/talk.go +++ /dev/null @@ -1,181 +0,0 @@ -package identity - -import ( - "context" - "database/sql" - "errors" - - "git.asio.asia/nixevol/NixMsg/internal/auth" - "git.asio.asia/nixevol/NixMsg/internal/protocol" -) - -// UnlockTalk 校验对话密码并写入 password 授权;未设密码则直接成功。 -func (a *App) UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error { - if !protocol.ValidEndpointID(senderID) || !protocol.ValidEndpointID(targetID) { - return errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - if senderID == targetID { - return nil - } - - var talkHash sql.NullString - var talkVer int64 - err := a.db.Read.QueryRowContext(ctx, ` -SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer) - if errors.Is(err, sql.ErrNoRows) { - return errCode(protocol.CodeInvalidTarget, "target not found") - } - if err != nil { - return err - } - if !talkHash.Valid || talkHash.String == "" { - return nil - } - - if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}); locked { - return errCode(protocol.CodeRateLimited, "talk password locked") - } - if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked { - return errCode(protocol.CodeRateLimited, "talk password locked") - } - if talkPassword == "" { - return errCode(protocol.CodeTalkPasswordRequired, "talk password required") - } - if a.hash == nil { - return errors.New("identity: hash pool required") - } - ok, err := a.hash.Verify(ctx, auth.PasswordTalk, talkPassword, talkHash.String) - if err != nil { - return err - } - if !ok { - a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) - a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}) - return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") - } - a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) - - nowMs := a.now().UnixMilli() - return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - return upsertGrantTx(tx, senderID, targetID, talkVer, grantKindPassword, nowMs) - }) -} - -// HasTalkGrant 发给自己、对方未设密码、或存在匹配版本授权时为 true。 -func (a *App) HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) { - if senderID == targetID { - return true, nil - } - if !protocol.ValidEndpointID(senderID) || !protocol.ValidEndpointID(targetID) { - return false, errCode(protocol.CodeBadRequest, "invalid endpoint id") - } - var talkHash sql.NullString - var talkVer int64 - err := a.db.Read.QueryRowContext(ctx, ` -SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer) - if errors.Is(err, sql.ErrNoRows) { - return false, errCode(protocol.CodeInvalidTarget, "target not found") - } - if err != nil { - return false, err - } - if !talkHash.Valid || talkHash.String == "" { - return true, nil - } - var n int - err = a.db.Read.QueryRowContext(ctx, ` -SELECT 1 FROM talk_grants -WHERE sender_id = ? AND target_id = ? AND target_talk_version = ? -LIMIT 1`, senderID, targetID, talkVer).Scan(&n) - if errors.Is(err, sql.ErrNoRows) { - return false, nil - } - return err == nil, err -} - -// CheckTalkPasswordForJoin 加人时必须当次带对密码;已有单聊授权不能代替。 -func (a *App) CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error { - if !protocol.ValidEndpointID(targetID) { - return errCode(protocol.CodeInvalidTarget, "invalid target") - } - var talkHash sql.NullString - var talkVer int64 - var enabled int - err := a.db.Read.QueryRowContext(ctx, ` -SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer, &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") - } - if !talkHash.Valid || talkHash.String == "" { - return nil - } - - if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}); locked { - return errCode(protocol.CodeRateLimited, "talk password locked") - } - if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked { - return errCode(protocol.CodeRateLimited, "talk password locked") - } - if talkPassword == "" { - return errCode(protocol.CodeTalkPasswordRequired, "talk password required") - } - if a.hash == nil { - return errors.New("identity: hash pool required") - } - ok, err := a.hash.Verify(ctx, auth.PasswordTalk, talkPassword, talkHash.String) - if err != nil { - return err - } - if !ok { - a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) - a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}) - return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") - } - a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) - // 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。 - _ = talkVer - return nil -} - -func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error { - _, err := tx.Exec(` -INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at) -VALUES(?,?,?,?,?) -ON CONFLICT(sender_id, target_id) DO UPDATE SET - target_talk_version = excluded.target_talk_version, - kind = excluded.kind, - created_at = excluded.created_at`, - senderID, targetID, talkVersion, kind, nowMs, - ) - return err -} - -// RecordReplyGrant 在对方成功提交单聊后写入 reply 授权(供消息线或测试调用)。 -func (a *App) RecordReplyGrant(ctx context.Context, senderID, targetID string) error { - if senderID == targetID { - return nil - } - var talkHash sql.NullString - var talkVer int64 - err := a.db.Read.QueryRowContext(ctx, ` -SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer) - if errors.Is(err, sql.ErrNoRows) { - return errCode(protocol.CodeInvalidTarget, "target not found") - } - if err != nil { - return err - } - if !talkHash.Valid || talkHash.String == "" { - return nil - } - nowMs := a.now().UnixMilli() - return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - return upsertGrantTx(tx, senderID, targetID, talkVer, grantKindReply, nowMs) - }) -} diff --git a/internal/app/presence/app.go b/internal/app/presence/app.go deleted file mode 100644 index f75966f..0000000 --- a/internal/app/presence/app.go +++ /dev/null @@ -1,353 +0,0 @@ -package presence - -import ( - "bytes" - "context" - "database/sql" - "strings" - "sync" - "time" - - "git.asio.asia/nixevol/NixMsg/internal/app/port" - "git.asio.asia/nixevol/NixMsg/internal/protocol" - "git.asio.asia/nixevol/NixMsg/internal/store" -) - -// ConnTable is an injectable connection table (N3). When nil, use SetOnline/SetOffline and DB columns. -type ConnTable interface { - IsOnline(endpointID string) bool - CurrentConn(endpointID string) (port.ConnID, bool) -} - -// Config holds presence service dependencies. -type Config struct { - DB *store.DB - Downlink port.Downlink // QoS 0 presence notifies; nil skips push - Conns ConnTable // optional - Now func() time.Time -} - -type watchSub struct { - endpointID string - all bool - ids map[string]struct{} -} - -type onlineEntry struct { - connID port.ConnID - atMs int64 -} - -// App implements presence.Service. -type App struct { - db *store.DB - down port.Downlink - conns ConnTable - nowFn func() time.Time - mu sync.Mutex - online map[string]onlineEntry - watches map[port.ConnID]watchSub -} - -// New constructs the presence service. -func New(cfg Config) *App { - now := cfg.Now - if now == nil { - now = time.Now - } - return &App{ - db: cfg.DB, - down: cfg.Downlink, - conns: cfg.Conns, - nowFn: now, - online: make(map[string]onlineEntry), - watches: make(map[port.ConnID]watchSub), - } -} - -func (a *App) nowMs() int64 { return a.nowFn().UnixMilli() } - -// Get queries online status for up to 200 ids. -func (a *App) Get(ctx context.Context, ids []string) ([]StatusItem, error) { - if len(ids) > protocol.MaxPresenceGetIDs { - return nil, &protocol.Error{Code: protocol.CodeBadRequest, Message: "too many ids"} - } - out := make([]StatusItem, 0, len(ids)) - for _, id := range ids { - item := StatusItem{ID: id} - var onlineSince, offlineSince sql.NullInt64 - err := a.db.Read.QueryRowContext(ctx, ` -SELECT online_since, offline_since FROM endpoints WHERE id = ?`, id).Scan(&onlineSince, &offlineSince) - if err == sql.ErrNoRows { - item.NotFound = true - out = append(out, item) - continue - } - if err != nil { - return nil, err - } - on, since := a.resolveOnline(id, onlineSince, offlineSince) - item.Online = on - item.SinceMs = since - out = append(out, item) - } - return out, nil -} - -// Directory lists endpoints with optional query and pagination. -func (a *App) Directory(ctx context.Context, req *protocol.DirectoryList) ([]DirectoryItem, string, error) { - if req == nil { - return nil, "", &protocol.Error{Code: protocol.CodeBadRequest, Message: "nil request"} - } - if err := req.Validate(); err != nil { - return nil, "", err - } - limit := req.Limit - if limit <= 0 { - limit = 100 - } - if limit > protocol.MaxPageLimit { - limit = protocol.MaxPageLimit - } - cursor := req.Cursor - query := strings.TrimSpace(req.Query) - - var rows *sql.Rows - var err error - if query == "" { - rows, err = a.db.Read.QueryContext(ctx, ` -SELECT id, name, online_since, offline_since, talk_hash -FROM endpoints -WHERE id > ? -ORDER BY id ASC -LIMIT ?`, cursor, limit+1) - } else { - like := "%" + strings.ToLower(query) + "%" - prefix := strings.ToLower(query) + "%" - rows, err = a.db.Read.QueryContext(ctx, ` -SELECT id, name, online_since, offline_since, talk_hash -FROM endpoints -WHERE id > ? - AND (lower(id) LIKE ? OR lower(name) LIKE ?) -ORDER BY id ASC -LIMIT ?`, cursor, prefix, like, limit+1) - } - if err != nil { - return nil, "", err - } - defer func() { _ = rows.Close() }() - - items := make([]DirectoryItem, 0, limit) - for rows.Next() { - var id, name string - var onlineSince, offlineSince sql.NullInt64 - var talk sql.NullString - if err := rows.Scan(&id, &name, &onlineSince, &offlineSince, &talk); err != nil { - return nil, "", err - } - on, _ := a.resolveOnline(id, onlineSince, offlineSince) - item := DirectoryItem{ - ID: id, - Name: name, - Online: on, - TalkPasswordSet: talk.Valid && talk.String != "", - } - if onlineSince.Valid { - v := onlineSince.Int64 - item.OnlineSinceMs = &v - } - if offlineSince.Valid { - v := offlineSince.Int64 - item.OfflineSinceMs = &v - } - items = append(items, item) - } - if err := rows.Err(); err != nil { - return nil, "", err - } - next := "" - if len(items) > limit { - items = items[:limit] - next = items[len(items)-1].ID - } - return items, next, nil -} - -// Watch replaces this connection's presence subscription. -func (a *App) Watch(_ context.Context, connID port.ConnID, endpointID string, req *protocol.PresenceWatch) error { - if req == nil { - return &protocol.Error{Code: protocol.CodeBadRequest, Message: "nil request"} - } - if err := req.Validate(); err != nil { - return err - } - sub := watchSub{endpointID: endpointID, all: req.All} - if !req.All { - sub.ids = make(map[string]struct{}, len(req.IDs)) - for _, id := range req.IDs { - sub.ids[id] = struct{}{} - } - } - a.mu.Lock() - a.watches[connID] = sub - a.mu.Unlock() - return nil -} - -// ClearWatch clears subscription on disconnect. -func (a *App) ClearWatch(connID port.ConnID) { - a.mu.Lock() - delete(a.watches, connID) - a.mu.Unlock() -} - -// SetOnline marks handshake complete. -func (a *App) SetOnline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error { - if atMs == 0 { - atMs = a.nowMs() - } - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - _, e := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, atMs, endpointID) - return e - }) - if err != nil { - return err - } - a.mu.Lock() - a.online[endpointID] = onlineEntry{connID: connID, atMs: atMs} - a.mu.Unlock() - a.notify(ctx, endpointID, true, atMs) - return nil -} - -// SetOffline marks disconnect for the current connection. -func (a *App) SetOffline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error { - if atMs == 0 { - atMs = a.nowMs() - } - a.mu.Lock() - cur, ok := a.online[endpointID] - if ok && cur.connID == connID { - delete(a.online, endpointID) - } else if ok { - a.mu.Unlock() - a.ClearWatch(connID) - return nil - } - a.mu.Unlock() - a.ClearWatch(connID) - - err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { - _, e := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, atMs, endpointID) - return e - }) - if err != nil { - return err - } - a.notify(ctx, endpointID, false, atMs) - return nil -} - -// IsOnline reports whether the endpoint is online. -func (a *App) IsOnline(endpointID string) bool { - if a.conns != nil && a.conns.IsOnline(endpointID) { - return true - } - a.mu.Lock() - _, ok := a.online[endpointID] - a.mu.Unlock() - if ok { - return true - } - var onlineSince, offlineSince sql.NullInt64 - err := a.db.Read.QueryRow(`SELECT online_since, offline_since FROM endpoints WHERE id = ?`, endpointID). - Scan(&onlineSince, &offlineSince) - if err != nil { - return false - } - on, _ := dbOnline(onlineSince, offlineSince) - return on -} - -// CurrentConn returns the current connection id if any. -func (a *App) CurrentConn(endpointID string) (port.ConnID, bool) { - if a.conns != nil { - if c, ok := a.conns.CurrentConn(endpointID); ok { - return c, true - } - } - a.mu.Lock() - defer a.mu.Unlock() - e, ok := a.online[endpointID] - return e.connID, ok -} - -func (a *App) resolveOnline(id string, onlineSince, offlineSince sql.NullInt64) (online bool, sinceMs int64) { - if a.conns != nil && a.conns.IsOnline(id) { - a.mu.Lock() - e, ok := a.online[id] - a.mu.Unlock() - if ok { - return true, e.atMs - } - if onlineSince.Valid { - return true, onlineSince.Int64 - } - return true, 0 - } - a.mu.Lock() - e, ok := a.online[id] - a.mu.Unlock() - if ok { - return true, e.atMs - } - return dbOnline(onlineSince, offlineSince) -} - -func dbOnline(onlineSince, offlineSince sql.NullInt64) (bool, int64) { - if !onlineSince.Valid { - return false, 0 - } - if !offlineSince.Valid || onlineSince.Int64 > offlineSince.Int64 { - return true, onlineSince.Int64 - } - return false, offlineSince.Int64 -} - -func (a *App) notify(ctx context.Context, changedID string, online bool, atMs int64) { - if a.down == nil { - return - } - frame := protocol.Presence{ - V: protocol.Version, Type: protocol.TypePresence, - ID: changedID, Online: online, AtMs: atMs, - } - payload, err := encodeFrame(frame) - if err != nil { - return - } - a.mu.Lock() - subs := make([]watchSub, 0, len(a.watches)) - for _, s := range a.watches { - subs = append(subs, s) - } - a.mu.Unlock() - for _, s := range subs { - if !s.all { - if _, ok := s.ids[changedID]; !ok { - continue - } - } - _ = a.down.PublishDown(ctx, s.endpointID, "", 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 -} - -var _ Service = (*App)(nil) diff --git a/internal/app/presence/presence_test.go b/internal/app/presence/presence_test.go deleted file mode 100644 index a46e5e0..0000000 --- a/internal/app/presence/presence_test.go +++ /dev/null @@ -1,176 +0,0 @@ -package presence_test - -import ( - "context" - "database/sql" - "path/filepath" - "sync" - "testing" - "time" - - "git.asio.asia/nixevol/NixMsg/internal/app/port" - "git.asio.asia/nixevol/NixMsg/internal/app/presence" - "git.asio.asia/nixevol/NixMsg/internal/protocol" - "git.asio.asia/nixevol/NixMsg/internal/store" -) - -type memDownlink struct { - mu sync.Mutex - msgs []downMsg -} - -type downMsg struct { - endpointID string - payload []byte - qos byte -} - -func (d *memDownlink) PublishDown(_ context.Context, endpointID string, _ port.ConnID, payload []byte, opts port.PublishOpts) error { - d.mu.Lock() - defer d.mu.Unlock() - cp := append([]byte(nil), payload...) - d.msgs = append(d.msgs, downMsg{endpointID: endpointID, payload: cp, qos: opts.QoS}) - return nil -} - -func (d *memDownlink) take() []downMsg { - d.mu.Lock() - defer d.mu.Unlock() - out := d.msgs - d.msgs = nil - return out -} - -func openPresence(t *testing.T) (*presence.App, *store.DB, *memDownlink) { - t.Helper() - db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL") - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { _ = db.Close() }) - down := &memDownlink{} - fixed := time.UnixMilli(1_700_000_000_000) - app := presence.New(presence.Config{ - DB: db, Downlink: down, Now: func() time.Time { return fixed }, - }) - return app, db, down -} - -func insertEP(t *testing.T, db *store.DB, id, name string) { - 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,1,?)`, id, name, "stub$login", nil, 1_700_000_000_000) - return e - }) - if err != nil { - t.Fatal(err) - } -} - -func TestF03PresenceAndDirectory(t *testing.T) { - t.Parallel() - app, db, _ := openPresence(t) - ctx := context.Background() - insertEP(t, db, "alice", "Alice") - insertEP(t, db, "bob", "Bob") - insertEP(t, db, "carol", "Carol") - - items, err := app.Get(ctx, []string{"alice", "nobody"}) - if err != nil { - t.Fatal(err) - } - if len(items) != 2 || items[0].Online || !items[1].NotFound { - t.Fatalf("%+v", items) - } - - if err = app.SetOnline(ctx, "alice", "c1", 1_700_000_000_100); err != nil { - t.Fatal(err) - } - items, _ = app.Get(ctx, []string{"alice"}) - if !items[0].Online || items[0].SinceMs != 1_700_000_000_100 { - t.Fatalf("%+v", items[0]) - } - - // 正常断开后查为离线 - if err = app.SetOffline(ctx, "alice", "c1", 1_700_000_000_200); err != nil { - t.Fatal(err) - } - items, _ = app.Get(ctx, []string{"alice"}) - if items[0].Online { - t.Fatal("should be offline") - } - - dir, next, err := app.Directory(ctx, &protocol.DirectoryList{ - V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "1", Limit: 2, - }) - if err != nil || len(dir) != 2 || next == "" { - t.Fatalf("dir=%+v next=%q err=%v", dir, next, err) - } - dir2, next2, err := app.Directory(ctx, &protocol.DirectoryList{ - V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "2", Cursor: next, Limit: 10, - }) - if err != nil || len(dir2) != 1 || next2 != "" { - t.Fatalf("dir2=%+v next=%q", dir2, next2) - } - - // query:编号前缀 - q, _, err := app.Directory(ctx, &protocol.DirectoryList{ - V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "3", Query: "bo", Limit: 10, - }) - if err != nil || len(q) != 1 || q[0].ID != "bob" { - t.Fatalf("%+v err=%v", q, err) - } - // 名称包含不区分大小写 - q, _, err = app.Directory(ctx, &protocol.DirectoryList{ - V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "4", Query: "car", Limit: 10, - }) - if err != nil || len(q) != 1 || q[0].ID != "carol" { - t.Fatalf("%+v", q) - } -} - -func TestF04PresenceWatch(t *testing.T) { - t.Parallel() - app, db, down := openPresence(t) - ctx := context.Background() - insertEP(t, db, "alice", "A") - insertEP(t, db, "bob", "B") - insertEP(t, db, "carol", "C") - - if err := app.Watch(ctx, "conn-sub", "watcher", &protocol.PresenceWatch{ - V: protocol.Version, Type: protocol.TypePresenceWatch, RID: "1", IDs: []string{"alice"}, - }); err != nil { - t.Fatal(err) - } - _ = app.SetOnline(ctx, "alice", "ca", 100) - _ = app.SetOffline(ctx, "alice", "ca", 200) - _ = app.SetOnline(ctx, "bob", "cb", 300) // 未订阅 - msgs := down.take() - if len(msgs) != 2 { - t.Fatalf("want 2 presence for alice got %d", len(msgs)) - } - for _, m := range msgs { - if m.endpointID != "watcher" || m.qos != 0 { - t.Fatalf("%+v", m) - } - } - - // 再次 Watch 覆盖;断线清空 - if err := app.Watch(ctx, "conn-sub", "watcher", &protocol.PresenceWatch{ - V: protocol.Version, Type: protocol.TypePresenceWatch, RID: "2", All: true, - }); err != nil { - t.Fatal(err) - } - down.take() - _ = app.SetOnline(ctx, "carol", "cc", 400) - if n := len(down.take()); n != 1 { - t.Fatalf("all watch got %d", n) - } - app.ClearWatch("conn-sub") - _ = app.SetOffline(ctx, "carol", "cc", 500) - if n := len(down.take()); n != 0 { - t.Fatalf("after clear got %d", n) - } -}