354 lines
8.4 KiB
Go
354 lines
8.4 KiB
Go
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)
|