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

This commit is contained in:
Nixevol
2026-09-30 07:33:31 +08:00
parent bdd1d9e9f4
commit a78ab0d547
14 changed files with 2777 additions and 65 deletions
+353
View File
@@ -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)
+176
View File
@@ -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)
}
}