feat: 实现身份资料、对话密码、在线目录与群管理
This commit is contained in:
@@ -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