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

This reverts commit a78ab0d547.
This commit is contained in:
Nixevol
2026-09-30 07:33:59 +08:00
parent a78ab0d547
commit 16ece09a97
14 changed files with 65 additions and 2777 deletions
-65
View File
@@ -427,71 +427,6 @@
- 备选方案:16/32 位。 - 备选方案: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 ## 后台接口 A
### A1 2026-09-30 ### A1 2026-09-30
-727
View File
@@ -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)
-412
View File
@@ -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)
}
}
}
-176
View File
@@ -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")
}
-131
View File
@@ -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)
-7
View File
@@ -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}
}
+59
View File
@@ -247,6 +247,65 @@ func (h *RegisterHandler) logResult(result, id, ip string) {
h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip) 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) { func setCORS(w http.ResponseWriter) {
w.Header().Set("Access-Control-Allow-Origin", "*") w.Header().Set("Access-Control-Allow-Origin", "*")
} }
-222
View File
@@ -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 }
-314
View File
@@ -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 }
+4 -7
View File
@@ -46,16 +46,13 @@ type Service interface {
SelfGet(ctx context.Context, endpointID string) (SelfInfo, error) SelfGet(ctx context.Context, endpointID string) (SelfInfo, error)
SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error
SelfSetTalkPassword(ctx context.Context, endpointID string, talkPassword string) error SelfSetTalkPassword(ctx context.Context, endpointID string, talkPassword string) error
// SelfChangeLoginPassword 校验旧密码后换新密码并签发新会话令牌;remoteIP 计入登录锁定。 SelfChangeLoginPassword(ctx context.Context, endpointID string, oldPassword, newPassword string) (sessionToken string, err error)
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
SelfLogout(ctx context.Context, endpointID string) error SelfLogout(ctx context.Context, endpointID string) error
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。 // UnlockTalk 校验并写入对话密码授权(第 6.6 节 unlock)。
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error UnlockTalk(ctx context.Context, senderID, targetID, talkPassword string) error
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。 // HasTalkGrant 查询发送方对目标是否有有效授权。
HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error)
// CheckTalkPasswordForJoin 加人时校验对话密码:已有单聊授权不能代替,必须当次带对。
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
// Disable 停用端并作废相关消息/令牌(第 7.6 节)。 // Disable 停用端并作废相关消息/令牌(第 7.6 节)。
Disable(ctx context.Context, endpointID string) error Disable(ctx context.Context, endpointID string) error
+2 -6
View File
@@ -27,13 +27,13 @@ func (s *Stub) SelfSetTalkPassword(context.Context, string, string) error {
return ErrNotImplemented 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 return "", ErrNotImplemented
} }
func (s *Stub) SelfLogout(context.Context, 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 return ErrNotImplemented
} }
@@ -41,10 +41,6 @@ func (s *Stub) HasTalkGrant(context.Context, string, string) (bool, error) {
return false, nil 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) Disable(context.Context, string) error { return ErrNotImplemented }
func (s *Stub) Enable(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 } func (s *Stub) Delete(context.Context, string) error { return ErrNotImplemented }
-181
View File
@@ -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)
})
}
-353
View File
@@ -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)
-176
View File
@@ -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)
}
}