Compare commits
2
Commits
16ece09a97
...
653b0d1866
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
653b0d1866 | ||
|
|
14e2e65a8c |
@@ -427,6 +427,78 @@
|
||||
- 备选方案: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` 传入。
|
||||
|
||||
10. **F06 群发断言与 M2 分发语义对齐(2026-09-30)**
|
||||
- 原条款:F06 验收「发到当时成员」;早期单测在仅写 `pending` 的桩分发下断言 Submit 返回 `dispatched`。
|
||||
- 实际做法:成员离线且默认不保留时,投递立即 `dropped`,消息为 `completed`(DEVELOPMENT 7.4/7.6);`TestF06GroupSendMembership` 改断言 `completed`;另加 `TestF06GroupSendKeepOfflinePending`:选离线保留时期望 `dispatched` 且接收者有 `pending`。不改消息分发实现。
|
||||
- 原因:合入含 M2 完整分发的 main 后,旧断言与产品规则冲突;永远 `dispatched` 才是错的。
|
||||
- 备选方案:测试里注入在线连接表使默认不保留也走 pending(与「离线不保留」场景重复覆盖)。
|
||||
- 影响:仅测试期望;产品行为不变。
|
||||
|
||||
## 后台接口 A
|
||||
|
||||
### A1 2026-09-30
|
||||
|
||||
@@ -0,0 +1,727 @@
|
||||
package group
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
eventMemberAdded = "member_added"
|
||||
eventMemberRemoved = "member_removed"
|
||||
eventLeft = "left"
|
||||
eventOwnerChanged = "owner_changed"
|
||||
eventRenamed = "renamed"
|
||||
eventDissolved = "dissolved"
|
||||
|
||||
reasonLeftGroup = "left_group"
|
||||
reasonGroupDissolved = "group_dissolved"
|
||||
|
||||
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||
)
|
||||
|
||||
// TalkGate checks talk password when adding members (implemented by identity).
|
||||
type TalkGate interface {
|
||||
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
|
||||
}
|
||||
|
||||
// OnlineLookup reports member online status.
|
||||
type OnlineLookup interface {
|
||||
IsOnline(endpointID string) bool
|
||||
}
|
||||
|
||||
// Config holds group service dependencies.
|
||||
type Config struct {
|
||||
DB *store.DB
|
||||
Talk TalkGate
|
||||
Online OnlineLookup
|
||||
Downlink port.Downlink
|
||||
MaxGroupMembers int
|
||||
Now func() time.Time
|
||||
DefaultRemoteIP string
|
||||
}
|
||||
|
||||
// App implements group.Service.
|
||||
type App struct {
|
||||
db *store.DB
|
||||
talk TalkGate
|
||||
online OnlineLookup
|
||||
down port.Downlink
|
||||
maxMem int
|
||||
nowFn func() time.Time
|
||||
remoteIP string
|
||||
}
|
||||
|
||||
// New constructs the group service.
|
||||
func New(cfg Config) *App {
|
||||
now := cfg.Now
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
max := cfg.MaxGroupMembers
|
||||
if max <= 0 {
|
||||
max = 1000
|
||||
}
|
||||
return &App{
|
||||
db: cfg.DB,
|
||||
talk: cfg.Talk,
|
||||
online: cfg.Online,
|
||||
down: cfg.Downlink,
|
||||
maxMem: max,
|
||||
nowFn: now,
|
||||
remoteIP: cfg.DefaultRemoteIP,
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) nowMs() int64 { return a.nowFn().UnixMilli() }
|
||||
|
||||
func errCode(code, msg string) *protocol.Error {
|
||||
return &protocol.Error{Code: code, Message: msg}
|
||||
}
|
||||
|
||||
func protoCode(err error) string {
|
||||
var pe *protocol.Error
|
||||
if errors.As(err, &pe) {
|
||||
return pe.Code
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// Create creates a group; creator becomes owner and member. Partial member failures still create the group.
|
||||
func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCreate) (CreateResult, error) {
|
||||
if req == nil {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
if !protocol.ValidEndpointID(actorID) {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid actor")
|
||||
}
|
||||
gid := req.ID
|
||||
if gid == "" {
|
||||
var genErr error
|
||||
gid, genErr = generateGroupID()
|
||||
if genErr != nil {
|
||||
return CreateResult{}, genErr
|
||||
}
|
||||
}
|
||||
now := a.nowMs()
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0, len(req.Members))
|
||||
|
||||
for _, m := range req.Members {
|
||||
if m.ID == actorID {
|
||||
continue
|
||||
}
|
||||
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
|
||||
if 1+len(added) > a.maxMem {
|
||||
return CreateResult{}, errCode(protocol.CodeGroupFull, "group full")
|
||||
}
|
||||
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var exists int
|
||||
qErr := tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, gid).Scan(&exists)
|
||||
if qErr == nil {
|
||||
return errCode(protocol.CodeIDTaken, "group id taken")
|
||||
}
|
||||
if !errors.Is(qErr, sql.ErrNoRows) {
|
||||
return qErr
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
gid, req.Name, actorID, now); e != nil {
|
||||
if isUnique(e) {
|
||||
return errCode(protocol.CodeIDTaken, "group id taken")
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, actorID, now); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
|
||||
for _, id := range added {
|
||||
a.emit(ctx, append([]string{actorID}, added...), gid, eventMemberAdded, id, now)
|
||||
}
|
||||
return CreateResult{ID: gid, Name: req.Name, OwnerID: actorID, Failed: failed}, nil
|
||||
}
|
||||
|
||||
// Add adds members (owner only).
|
||||
func (a *App) Add(ctx context.Context, actorID string, req *protocol.GroupAdd) (AddResult, error) {
|
||||
if req == nil {
|
||||
return AddResult{}, errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
if owner != actorID {
|
||||
return AddResult{}, errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
now := a.nowMs()
|
||||
|
||||
for _, m := range req.Members {
|
||||
if contains(members, m.ID) {
|
||||
continue
|
||||
}
|
||||
if len(members)+len(added) >= a.maxMem {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
if checkErr := a.checkAddMember(ctx, actorID, m.ID, m.TalkPassword); checkErr != nil {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: failCode(checkErr)})
|
||||
continue
|
||||
}
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
|
||||
if len(added) > 0 {
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT OR IGNORE INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
req.GroupID, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
all := append(append([]string{}, members...), added...)
|
||||
for _, id := range added {
|
||||
a.emit(ctx, all, req.GroupID, eventMemberAdded, id, now)
|
||||
}
|
||||
}
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
|
||||
// Remove kicks a member (owner only).
|
||||
func (a *App) Remove(ctx context.Context, actorID string, req *protocol.GroupRemove) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
if req.EndpointID == owner {
|
||||
return errCode(protocol.CodeBadRequest, "cannot remove owner")
|
||||
}
|
||||
if !contains(members, req.EndpointID) {
|
||||
return errCode(protocol.CodeNotFound, "member not found")
|
||||
}
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`,
|
||||
req.GroupID, req.EndpointID); e != nil {
|
||||
return e
|
||||
}
|
||||
return voidMemberDeliveriesTx(tx, req.GroupID, req.EndpointID, reasonLeftGroup, now, &revokes)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.publishRevokes(ctx, revokes)
|
||||
left := without(members, req.EndpointID)
|
||||
notify := append(left, req.EndpointID)
|
||||
a.emit(ctx, notify, req.GroupID, eventMemberRemoved, req.EndpointID, now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Leave lets a non-owner member leave.
|
||||
func (a *App) Leave(ctx context.Context, actorID string, req *protocol.GroupLeave) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !contains(members, actorID) {
|
||||
return errCode(protocol.CodeNotMember, "not a member")
|
||||
}
|
||||
if owner == actorID {
|
||||
return errCode(protocol.CodeOwnerCannotLeave, "owner cannot leave")
|
||||
}
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`,
|
||||
req.GroupID, actorID); e != nil {
|
||||
return e
|
||||
}
|
||||
return voidMemberDeliveriesTx(tx, req.GroupID, actorID, reasonLeftGroup, now, &revokes)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.publishRevokes(ctx, revokes)
|
||||
left := without(members, actorID)
|
||||
notify := append(left, actorID)
|
||||
a.emit(ctx, notify, req.GroupID, eventLeft, actorID, now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Transfer transfers ownership.
|
||||
func (a *App) Transfer(ctx context.Context, actorID string, req *protocol.GroupTransfer) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
if !contains(members, req.EndpointID) {
|
||||
return errCode(protocol.CodeNotFound, "member not found")
|
||||
}
|
||||
now := a.nowMs()
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE groups SET owner_id = ? WHERE id = ?`, req.EndpointID, req.GroupID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.emit(ctx, members, req.GroupID, eventOwnerChanged, req.EndpointID, now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Rename renames the group.
|
||||
func (a *App) Rename(ctx context.Context, actorID string, req *protocol.GroupRename) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
now := a.nowMs()
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE groups SET name = ? WHERE id = ?`, req.Name, req.GroupID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.emit(ctx, members, req.GroupID, eventRenamed, "", now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Dissolve dissolves the group and voids unfinished deliveries / scheduled messages.
|
||||
func (a *App) Dissolve(ctx context.Context, actorID string, req *protocol.GroupDissolve) error {
|
||||
if req == nil {
|
||||
return errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
owner, members, err := a.loadGroup(ctx, req.GroupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if owner != actorID {
|
||||
return errCode(protocol.CodeForbidden, "not owner")
|
||||
}
|
||||
now := a.nowMs()
|
||||
var revokes []revokeItem
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if e := voidGroupAllTx(tx, req.GroupID, now, &revokes); e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ?`, req.GroupID); e != nil {
|
||||
return e
|
||||
}
|
||||
_, e := tx.Exec(`DELETE FROM groups WHERE id = ?`, req.GroupID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.publishRevokes(ctx, revokes)
|
||||
a.emit(ctx, members, req.GroupID, eventDissolved, "", now)
|
||||
return nil
|
||||
}
|
||||
|
||||
// List lists groups the actor belongs to.
|
||||
func (a *App) List(ctx context.Context, actorID string, req *protocol.GroupList) ([]ListItem, string, error) {
|
||||
if req == nil {
|
||||
req = &protocol.GroupList{V: protocol.Version, Type: protocol.TypeGroupList, RID: "x", Limit: 100}
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > protocol.MaxPageLimit {
|
||||
limit = protocol.MaxPageLimit
|
||||
}
|
||||
rows, err := a.db.Read.QueryContext(ctx, `
|
||||
SELECT g.id, g.name, g.owner_id,
|
||||
(SELECT COUNT(*) FROM group_members gm2 WHERE gm2.group_id = g.id) AS cnt
|
||||
FROM groups g
|
||||
JOIN group_members gm ON gm.group_id = g.id AND gm.endpoint_id = ?
|
||||
WHERE g.id > ?
|
||||
ORDER BY g.id ASC
|
||||
LIMIT ?`, actorID, req.Cursor, limit+1)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
items := make([]ListItem, 0, limit)
|
||||
for rows.Next() {
|
||||
var it ListItem
|
||||
if scanErr := rows.Scan(&it.ID, &it.Name, &it.OwnerID, &it.MemberCount); scanErr != nil {
|
||||
return nil, "", scanErr
|
||||
}
|
||||
items = append(items, it)
|
||||
}
|
||||
next := ""
|
||||
if len(items) > limit {
|
||||
items = items[:limit]
|
||||
next = items[len(items)-1].ID
|
||||
}
|
||||
return items, next, rows.Err()
|
||||
}
|
||||
|
||||
// Get returns group details; caller must be a member.
|
||||
func (a *App) Get(ctx context.Context, actorID string, req *protocol.GroupGet) (GetResult, error) {
|
||||
if req == nil {
|
||||
return GetResult{}, errCode(protocol.CodeBadRequest, "nil request")
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
var name, owner string
|
||||
err := a.db.Read.QueryRowContext(ctx, `SELECT name, owner_id FROM groups WHERE id = ?`, req.GroupID).
|
||||
Scan(&name, &owner)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return GetResult{}, errCode(protocol.CodeNotFound, "group not found")
|
||||
}
|
||||
if err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
var one int
|
||||
err = a.db.Read.QueryRowContext(ctx,
|
||||
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, req.GroupID, actorID).Scan(&one)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return GetResult{}, errCode(protocol.CodeNotMember, "not a member")
|
||||
}
|
||||
if err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > protocol.MaxPageLimit {
|
||||
limit = protocol.MaxPageLimit
|
||||
}
|
||||
rows, err := a.db.Read.QueryContext(ctx, `
|
||||
SELECT gm.endpoint_id, e.name
|
||||
FROM group_members gm
|
||||
JOIN endpoints e ON e.id = gm.endpoint_id
|
||||
WHERE gm.group_id = ? AND gm.endpoint_id > ?
|
||||
ORDER BY gm.endpoint_id ASC
|
||||
LIMIT ?`, req.GroupID, req.Cursor, limit+1)
|
||||
if err != nil {
|
||||
return GetResult{}, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
members := make([]MemberItem, 0, limit)
|
||||
for rows.Next() {
|
||||
var m MemberItem
|
||||
if scanErr := rows.Scan(&m.ID, &m.Name); scanErr != nil {
|
||||
return GetResult{}, scanErr
|
||||
}
|
||||
if a.online != nil {
|
||||
m.Online = a.online.IsOnline(m.ID)
|
||||
}
|
||||
members = append(members, m)
|
||||
}
|
||||
next := ""
|
||||
if len(members) > limit {
|
||||
members = members[:limit]
|
||||
next = members[len(members)-1].ID
|
||||
}
|
||||
return GetResult{ID: req.GroupID, Name: name, OwnerID: owner, Members: members, NextCursor: next}, rows.Err()
|
||||
}
|
||||
|
||||
// AdminCreate creates a group without talk-password checks.
|
||||
func (a *App) AdminCreate(ctx context.Context, name, ownerID string, memberIDs []string) (CreateResult, error) {
|
||||
members := make([]protocol.GroupMemberIn, 0, len(memberIDs))
|
||||
for _, id := range memberIDs {
|
||||
members = append(members, protocol.GroupMemberIn{ID: id})
|
||||
}
|
||||
return a.createAdmin(ctx, ownerID, name, "", members)
|
||||
}
|
||||
|
||||
// AdminAddMembers adds members without talk-password checks.
|
||||
func (a *App) AdminAddMembers(ctx context.Context, groupID string, memberIDs []string) (AddResult, error) {
|
||||
_, members, err := a.loadGroup(ctx, groupID)
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
now := a.nowMs()
|
||||
for _, id := range memberIDs {
|
||||
if contains(members, id) {
|
||||
continue
|
||||
}
|
||||
var enabled int
|
||||
e := a.db.Read.QueryRowContext(ctx, `SELECT enabled FROM endpoints WHERE id = ?`, id).Scan(&enabled)
|
||||
if errors.Is(e, sql.ErrNoRows) {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeInvalidTarget})
|
||||
continue
|
||||
}
|
||||
if e != nil {
|
||||
return AddResult{}, e
|
||||
}
|
||||
if enabled == 0 {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeEndpointDisabled})
|
||||
continue
|
||||
}
|
||||
if len(members)+len(added) >= a.maxMem {
|
||||
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
|
||||
continue
|
||||
}
|
||||
added = append(added, id)
|
||||
}
|
||||
if len(added) == 0 {
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT OR IGNORE INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
groupID, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return AddResult{}, err
|
||||
}
|
||||
all := append(append([]string{}, members...), added...)
|
||||
for _, id := range added {
|
||||
a.emit(ctx, all, groupID, eventMemberAdded, id, now)
|
||||
}
|
||||
return AddResult{Failed: failed}, nil
|
||||
}
|
||||
|
||||
func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, members []protocol.GroupMemberIn) (CreateResult, error) {
|
||||
if gid == "" {
|
||||
var genErr error
|
||||
gid, genErr = generateGroupID()
|
||||
if genErr != nil {
|
||||
return CreateResult{}, genErr
|
||||
}
|
||||
}
|
||||
if !protocol.ValidName(name) || name == "" {
|
||||
return CreateResult{}, errCode(protocol.CodeBadRequest, "invalid name")
|
||||
}
|
||||
now := a.nowMs()
|
||||
failed := make([]MemberFail, 0)
|
||||
added := make([]string, 0)
|
||||
for _, m := range members {
|
||||
if m.ID == ownerID {
|
||||
continue
|
||||
}
|
||||
var enabled int
|
||||
e := a.db.Read.QueryRowContext(ctx, `SELECT enabled FROM endpoints WHERE id = ?`, m.ID).Scan(&enabled)
|
||||
if errors.Is(e, sql.ErrNoRows) {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeInvalidTarget})
|
||||
continue
|
||||
}
|
||||
if e != nil {
|
||||
return CreateResult{}, e
|
||||
}
|
||||
if enabled == 0 {
|
||||
failed = append(failed, MemberFail{ID: m.ID, Code: protocol.CodeEndpointDisabled})
|
||||
continue
|
||||
}
|
||||
added = append(added, m.ID)
|
||||
}
|
||||
if 1+len(added) > a.maxMem {
|
||||
return CreateResult{}, errCode(protocol.CodeGroupFull, "group full")
|
||||
}
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
gid, name, ownerID, now); e != nil {
|
||||
if isUnique(e) {
|
||||
return errCode(protocol.CodeIDTaken, "group id taken")
|
||||
}
|
||||
return e
|
||||
}
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, ownerID, now); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, id := range added {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
gid, id, now); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return CreateResult{}, err
|
||||
}
|
||||
return CreateResult{ID: gid, Name: name, OwnerID: ownerID, Failed: failed}, nil
|
||||
}
|
||||
|
||||
func (a *App) checkAddMember(ctx context.Context, actorID, targetID, talkPassword string) error {
|
||||
if a.talk == nil {
|
||||
return errCode(protocol.CodeBusy, "talk gate not configured")
|
||||
}
|
||||
return a.talk.CheckTalkPasswordForJoin(ctx, actorID, targetID, talkPassword, a.remoteIP)
|
||||
}
|
||||
|
||||
func (a *App) loadGroup(ctx context.Context, groupID string) (owner string, members []string, err error) {
|
||||
err = a.db.Read.QueryRowContext(ctx, `SELECT owner_id FROM groups WHERE id = ?`, groupID).Scan(&owner)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return "", nil, errCode(protocol.CodeNotFound, "group not found")
|
||||
}
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
rows, qErr := a.db.Read.QueryContext(ctx, `SELECT endpoint_id FROM group_members WHERE group_id = ?`, groupID)
|
||||
if qErr != nil {
|
||||
return "", nil, qErr
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var id string
|
||||
if scanErr := rows.Scan(&id); scanErr != nil {
|
||||
return "", nil, scanErr
|
||||
}
|
||||
members = append(members, id)
|
||||
}
|
||||
return owner, members, rows.Err()
|
||||
}
|
||||
|
||||
func (a *App) emit(ctx context.Context, recipients []string, groupID, event, endpointID string, atMs int64) {
|
||||
if a.down == nil {
|
||||
return
|
||||
}
|
||||
frame := protocol.GroupEvent{
|
||||
V: protocol.Version, Type: protocol.TypeGroupEvent,
|
||||
GroupID: groupID, Event: event, EndpointID: endpointID, AtMs: atMs,
|
||||
}
|
||||
payload, encErr := encodeFrame(frame)
|
||||
if encErr != nil {
|
||||
return
|
||||
}
|
||||
seen := map[string]struct{}{}
|
||||
for _, id := range recipients {
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
_ = a.down.PublishDown(ctx, id, "", payload, port.PublishOpts{QoS: 0})
|
||||
}
|
||||
}
|
||||
|
||||
func encodeFrame(v any) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
if err := protocol.Encode(&buf, v); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func generateGroupID() (string, error) {
|
||||
b := make([]byte, 8)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, 8)
|
||||
for i := range b {
|
||||
out[i] = idAlphabet[int(b[i])%len(idAlphabet)]
|
||||
}
|
||||
return "g_" + string(out), nil
|
||||
}
|
||||
|
||||
func failCode(err error) string {
|
||||
if c := protoCode(err); c != "" {
|
||||
return c
|
||||
}
|
||||
return protocol.CodeBusy
|
||||
}
|
||||
|
||||
func contains(ss []string, x string) bool {
|
||||
for _, s := range ss {
|
||||
if s == x {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func without(ss []string, x string) []string {
|
||||
out := make([]string, 0, len(ss))
|
||||
for _, s := range ss {
|
||||
if s != x {
|
||||
out = append(out, s)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
var _ Service = (*App)(nil)
|
||||
@@ -0,0 +1,453 @@
|
||||
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.LimitsFromFullConfig(config.Default())
|
||||
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)
|
||||
}
|
||||
// 默认不保留 + 成员离线(无连接、无 offline_since):投递立即 dropped,消息 completed(DEVELOPMENT 7.4/7.6)
|
||||
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.StateCompleted {
|
||||
t.Fatalf("state=%s want completed (offline, keep=false)", 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 pending, dropped int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.state='pending'`).Scan(&pending)
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='m1' AND d.state='dropped'`).Scan(&dropped)
|
||||
if pending != 0 || dropped != 2 {
|
||||
t.Fatalf("pending=%d dropped=%d want pending=0 dropped=2", pending, dropped)
|
||||
}
|
||||
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 TestF06GroupSendKeepOfflinePending(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: "GKeep", Members: []protocol.GroupMemberIn{{ID: "bob"}, {ID: "carol"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ttl := int64(3600)
|
||||
res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{}, &protocol.Send{
|
||||
V: protocol.Version, Type: protocol.TypeSend, RID: "s", ID: "mkeep",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: created.ID},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "kept"},
|
||||
Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != message.StateDispatched {
|
||||
t.Fatalf("state=%s want dispatched (offline keep)", res.State)
|
||||
}
|
||||
var pending int
|
||||
_ = db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='mkeep' AND d.state='pending'`).Scan(&pending)
|
||||
if pending != 2 {
|
||||
t.Fatalf("pending=%d want 2", pending)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package group
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
type revokeItem struct {
|
||||
endpointID string
|
||||
msgID string
|
||||
fromID string
|
||||
reason string
|
||||
}
|
||||
|
||||
// voidMemberDeliveriesTx rejects pending deliveries for a leaving member; records revokes for pushed ones.
|
||||
func voidMemberDeliveriesTx(tx *sql.Tx, groupID, endpointID, reason string, nowMs int64, revokes *[]revokeItem) error {
|
||||
rows, err := tx.Query(`
|
||||
SELECT d.seq, d.pushed_at, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.endpoint_id = ? AND d.state = 'pending'
|
||||
AND m.dest_kind = 'group' AND m.dest_id = ?`, endpointID, groupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
type row struct {
|
||||
seq int64
|
||||
pushed sql.NullInt64
|
||||
msgID string
|
||||
senderID string
|
||||
}
|
||||
var list []row
|
||||
for rows.Next() {
|
||||
var r row
|
||||
if scanErr := rows.Scan(&r.seq, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
|
||||
return scanErr
|
||||
}
|
||||
list = append(list, r)
|
||||
}
|
||||
if err = rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range list {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
reason, nowMs, r.seq, endpointID); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason,
|
||||
})
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// voidGroupAllTx rejects all pending group deliveries and completes scheduled messages.
|
||||
func voidGroupAllTx(tx *sql.Tx, groupID string, nowMs int64, revokes *[]revokeItem) error {
|
||||
rows, err := tx.Query(`
|
||||
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
|
||||
FROM deliveries d
|
||||
JOIN messages m ON m.seq = d.seq
|
||||
WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type drow struct {
|
||||
seq int64
|
||||
endpointID string
|
||||
pushed sql.NullInt64
|
||||
msgID string
|
||||
senderID string
|
||||
}
|
||||
var dlist []drow
|
||||
for rows.Next() {
|
||||
var r drow
|
||||
if scanErr := rows.Scan(&r.seq, &r.endpointID, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
|
||||
_ = rows.Close()
|
||||
return scanErr
|
||||
}
|
||||
dlist = append(dlist, r)
|
||||
}
|
||||
_ = rows.Close()
|
||||
if err = rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range dlist {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
reasonGroupDissolved, nowMs, r.seq, r.endpointID); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.pushed.Valid && revokes != nil {
|
||||
*revokes = append(*revokes, revokeItem{
|
||||
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
srows, err := tx.Query(`
|
||||
SELECT seq, id, sender_id, receipt FROM messages
|
||||
WHERE dest_kind = 'group' AND dest_id = ? AND state = 'scheduled'`, groupID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type srow struct {
|
||||
seq int64
|
||||
msgID string
|
||||
senderID string
|
||||
receipt int
|
||||
}
|
||||
var slist []srow
|
||||
for srows.Next() {
|
||||
var r srow
|
||||
if scanErr := srows.Scan(&r.seq, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
|
||||
_ = srows.Close()
|
||||
return scanErr
|
||||
}
|
||||
slist = append(slist, r)
|
||||
}
|
||||
_ = srows.Close()
|
||||
if err = srows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range slist {
|
||||
if _, execErr := tx.Exec(`
|
||||
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
|
||||
reasonGroupDissolved, r.seq); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, r.seq); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
if r.receipt != 0 {
|
||||
if _, execErr := tx.Exec(`
|
||||
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
|
||||
VALUES(?,?,?,?,?,?,0)`,
|
||||
r.senderID, r.msgID, "", "completed", reasonGroupDissolved, nowMs); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *App) publishRevokes(ctx context.Context, items []revokeItem) {
|
||||
if a.down == nil || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
for _, it := range items {
|
||||
frame := protocol.Revoked{
|
||||
V: protocol.Version, Type: protocol.TypeRevoked,
|
||||
ID: it.msgID, From: it.fromID, Reason: it.reason,
|
||||
}
|
||||
payload, encErr := encodeFrame(frame)
|
||||
if encErr != nil {
|
||||
continue
|
||||
}
|
||||
_ = a.down.PublishDown(ctx, it.endpointID, "", payload, port.PublishOpts{QoS: 1})
|
||||
}
|
||||
}
|
||||
|
||||
func isUnique(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
return strings.Contains(msg, "unique") || strings.Contains(msg, "constraint failed")
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
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)
|
||||
@@ -0,0 +1,7 @@
|
||||
package identity
|
||||
|
||||
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
|
||||
func errCode(code, msg string) *protocol.Error {
|
||||
return &protocol.Error{Code: code, Message: msg}
|
||||
}
|
||||
@@ -247,65 +247,6 @@ 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", "*")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
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 }
|
||||
@@ -0,0 +1,314 @@
|
||||
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 }
|
||||
@@ -46,13 +46,16 @@ 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(ctx context.Context, endpointID string, oldPassword, newPassword string) (sessionToken string, err error)
|
||||
// SelfChangeLoginPassword 校验旧密码后换新密码并签发新会话令牌;remoteIP 计入登录锁定。
|
||||
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
|
||||
SelfLogout(ctx context.Context, endpointID string) error
|
||||
|
||||
// UnlockTalk 校验并写入对话密码授权(第 6.6 节 unlock)。
|
||||
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword string) error
|
||||
// HasTalkGrant 查询发送方对目标是否有有效授权。
|
||||
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。
|
||||
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP 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
|
||||
|
||||
@@ -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, error) {
|
||||
func (s *Stub) SelfChangeLoginPassword(context.Context, string, 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) error {
|
||||
func (s *Stub) UnlockTalk(context.Context, string, string, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
|
||||
@@ -41,6 +41,10 @@ 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 }
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,353 @@
|
||||
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)
|
||||
@@ -0,0 +1,176 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user