@@ -1,353 +0,0 @@
|
||||
package presence
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
// ConnTable is an injectable connection table (N3). When nil, use SetOnline/SetOffline and DB columns.
|
||||
type ConnTable interface {
|
||||
IsOnline(endpointID string) bool
|
||||
CurrentConn(endpointID string) (port.ConnID, bool)
|
||||
}
|
||||
|
||||
// Config holds presence service dependencies.
|
||||
type Config struct {
|
||||
DB *store.DB
|
||||
Downlink port.Downlink // QoS 0 presence notifies; nil skips push
|
||||
Conns ConnTable // optional
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
type watchSub struct {
|
||||
endpointID string
|
||||
all bool
|
||||
ids map[string]struct{}
|
||||
}
|
||||
|
||||
type onlineEntry struct {
|
||||
connID port.ConnID
|
||||
atMs int64
|
||||
}
|
||||
|
||||
// App implements presence.Service.
|
||||
type App struct {
|
||||
db *store.DB
|
||||
down port.Downlink
|
||||
conns ConnTable
|
||||
nowFn func() time.Time
|
||||
mu sync.Mutex
|
||||
online map[string]onlineEntry
|
||||
watches map[port.ConnID]watchSub
|
||||
}
|
||||
|
||||
// New constructs the presence service.
|
||||
func New(cfg Config) *App {
|
||||
now := cfg.Now
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
return &App{
|
||||
db: cfg.DB,
|
||||
down: cfg.Downlink,
|
||||
conns: cfg.Conns,
|
||||
nowFn: now,
|
||||
online: make(map[string]onlineEntry),
|
||||
watches: make(map[port.ConnID]watchSub),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) nowMs() int64 { return a.nowFn().UnixMilli() }
|
||||
|
||||
// Get queries online status for up to 200 ids.
|
||||
func (a *App) Get(ctx context.Context, ids []string) ([]StatusItem, error) {
|
||||
if len(ids) > protocol.MaxPresenceGetIDs {
|
||||
return nil, &protocol.Error{Code: protocol.CodeBadRequest, Message: "too many ids"}
|
||||
}
|
||||
out := make([]StatusItem, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
item := StatusItem{ID: id}
|
||||
var onlineSince, offlineSince sql.NullInt64
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT online_since, offline_since FROM endpoints WHERE id = ?`, id).Scan(&onlineSince, &offlineSince)
|
||||
if err == sql.ErrNoRows {
|
||||
item.NotFound = true
|
||||
out = append(out, item)
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
on, since := a.resolveOnline(id, onlineSince, offlineSince)
|
||||
item.Online = on
|
||||
item.SinceMs = since
|
||||
out = append(out, item)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Directory lists endpoints with optional query and pagination.
|
||||
func (a *App) Directory(ctx context.Context, req *protocol.DirectoryList) ([]DirectoryItem, string, error) {
|
||||
if req == nil {
|
||||
return nil, "", &protocol.Error{Code: protocol.CodeBadRequest, Message: "nil request"}
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > protocol.MaxPageLimit {
|
||||
limit = protocol.MaxPageLimit
|
||||
}
|
||||
cursor := req.Cursor
|
||||
query := strings.TrimSpace(req.Query)
|
||||
|
||||
var rows *sql.Rows
|
||||
var err error
|
||||
if query == "" {
|
||||
rows, err = a.db.Read.QueryContext(ctx, `
|
||||
SELECT id, name, online_since, offline_since, talk_hash
|
||||
FROM endpoints
|
||||
WHERE id > ?
|
||||
ORDER BY id ASC
|
||||
LIMIT ?`, cursor, limit+1)
|
||||
} else {
|
||||
like := "%" + strings.ToLower(query) + "%"
|
||||
prefix := strings.ToLower(query) + "%"
|
||||
rows, err = a.db.Read.QueryContext(ctx, `
|
||||
SELECT id, name, online_since, offline_since, talk_hash
|
||||
FROM endpoints
|
||||
WHERE id > ?
|
||||
AND (lower(id) LIKE ? OR lower(name) LIKE ?)
|
||||
ORDER BY id ASC
|
||||
LIMIT ?`, cursor, prefix, like, limit+1)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
items := make([]DirectoryItem, 0, limit)
|
||||
for rows.Next() {
|
||||
var id, name string
|
||||
var onlineSince, offlineSince sql.NullInt64
|
||||
var talk sql.NullString
|
||||
if err := rows.Scan(&id, &name, &onlineSince, &offlineSince, &talk); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
on, _ := a.resolveOnline(id, onlineSince, offlineSince)
|
||||
item := DirectoryItem{
|
||||
ID: id,
|
||||
Name: name,
|
||||
Online: on,
|
||||
TalkPasswordSet: talk.Valid && talk.String != "",
|
||||
}
|
||||
if onlineSince.Valid {
|
||||
v := onlineSince.Int64
|
||||
item.OnlineSinceMs = &v
|
||||
}
|
||||
if offlineSince.Valid {
|
||||
v := offlineSince.Int64
|
||||
item.OfflineSinceMs = &v
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
next := ""
|
||||
if len(items) > limit {
|
||||
items = items[:limit]
|
||||
next = items[len(items)-1].ID
|
||||
}
|
||||
return items, next, nil
|
||||
}
|
||||
|
||||
// Watch replaces this connection's presence subscription.
|
||||
func (a *App) Watch(_ context.Context, connID port.ConnID, endpointID string, req *protocol.PresenceWatch) error {
|
||||
if req == nil {
|
||||
return &protocol.Error{Code: protocol.CodeBadRequest, Message: "nil request"}
|
||||
}
|
||||
if err := req.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
sub := watchSub{endpointID: endpointID, all: req.All}
|
||||
if !req.All {
|
||||
sub.ids = make(map[string]struct{}, len(req.IDs))
|
||||
for _, id := range req.IDs {
|
||||
sub.ids[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
a.mu.Lock()
|
||||
a.watches[connID] = sub
|
||||
a.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClearWatch clears subscription on disconnect.
|
||||
func (a *App) ClearWatch(connID port.ConnID) {
|
||||
a.mu.Lock()
|
||||
delete(a.watches, connID)
|
||||
a.mu.Unlock()
|
||||
}
|
||||
|
||||
// SetOnline marks handshake complete.
|
||||
func (a *App) SetOnline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error {
|
||||
if atMs == 0 {
|
||||
atMs = a.nowMs()
|
||||
}
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, atMs, endpointID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.mu.Lock()
|
||||
a.online[endpointID] = onlineEntry{connID: connID, atMs: atMs}
|
||||
a.mu.Unlock()
|
||||
a.notify(ctx, endpointID, true, atMs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetOffline marks disconnect for the current connection.
|
||||
func (a *App) SetOffline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error {
|
||||
if atMs == 0 {
|
||||
atMs = a.nowMs()
|
||||
}
|
||||
a.mu.Lock()
|
||||
cur, ok := a.online[endpointID]
|
||||
if ok && cur.connID == connID {
|
||||
delete(a.online, endpointID)
|
||||
} else if ok {
|
||||
a.mu.Unlock()
|
||||
a.ClearWatch(connID)
|
||||
return nil
|
||||
}
|
||||
a.mu.Unlock()
|
||||
a.ClearWatch(connID)
|
||||
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, atMs, endpointID)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.notify(ctx, endpointID, false, atMs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsOnline reports whether the endpoint is online.
|
||||
func (a *App) IsOnline(endpointID string) bool {
|
||||
if a.conns != nil && a.conns.IsOnline(endpointID) {
|
||||
return true
|
||||
}
|
||||
a.mu.Lock()
|
||||
_, ok := a.online[endpointID]
|
||||
a.mu.Unlock()
|
||||
if ok {
|
||||
return true
|
||||
}
|
||||
var onlineSince, offlineSince sql.NullInt64
|
||||
err := a.db.Read.QueryRow(`SELECT online_since, offline_since FROM endpoints WHERE id = ?`, endpointID).
|
||||
Scan(&onlineSince, &offlineSince)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
on, _ := dbOnline(onlineSince, offlineSince)
|
||||
return on
|
||||
}
|
||||
|
||||
// CurrentConn returns the current connection id if any.
|
||||
func (a *App) CurrentConn(endpointID string) (port.ConnID, bool) {
|
||||
if a.conns != nil {
|
||||
if c, ok := a.conns.CurrentConn(endpointID); ok {
|
||||
return c, true
|
||||
}
|
||||
}
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
e, ok := a.online[endpointID]
|
||||
return e.connID, ok
|
||||
}
|
||||
|
||||
func (a *App) resolveOnline(id string, onlineSince, offlineSince sql.NullInt64) (online bool, sinceMs int64) {
|
||||
if a.conns != nil && a.conns.IsOnline(id) {
|
||||
a.mu.Lock()
|
||||
e, ok := a.online[id]
|
||||
a.mu.Unlock()
|
||||
if ok {
|
||||
return true, e.atMs
|
||||
}
|
||||
if onlineSince.Valid {
|
||||
return true, onlineSince.Int64
|
||||
}
|
||||
return true, 0
|
||||
}
|
||||
a.mu.Lock()
|
||||
e, ok := a.online[id]
|
||||
a.mu.Unlock()
|
||||
if ok {
|
||||
return true, e.atMs
|
||||
}
|
||||
return dbOnline(onlineSince, offlineSince)
|
||||
}
|
||||
|
||||
func dbOnline(onlineSince, offlineSince sql.NullInt64) (bool, int64) {
|
||||
if !onlineSince.Valid {
|
||||
return false, 0
|
||||
}
|
||||
if !offlineSince.Valid || onlineSince.Int64 > offlineSince.Int64 {
|
||||
return true, onlineSince.Int64
|
||||
}
|
||||
return false, offlineSince.Int64
|
||||
}
|
||||
|
||||
func (a *App) notify(ctx context.Context, changedID string, online bool, atMs int64) {
|
||||
if a.down == nil {
|
||||
return
|
||||
}
|
||||
frame := protocol.Presence{
|
||||
V: protocol.Version, Type: protocol.TypePresence,
|
||||
ID: changedID, Online: online, AtMs: atMs,
|
||||
}
|
||||
payload, err := encodeFrame(frame)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
a.mu.Lock()
|
||||
subs := make([]watchSub, 0, len(a.watches))
|
||||
for _, s := range a.watches {
|
||||
subs = append(subs, s)
|
||||
}
|
||||
a.mu.Unlock()
|
||||
for _, s := range subs {
|
||||
if !s.all {
|
||||
if _, ok := s.ids[changedID]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
_ = a.down.PublishDown(ctx, s.endpointID, "", payload, port.PublishOpts{QoS: 0})
|
||||
}
|
||||
}
|
||||
|
||||
func encodeFrame(v any) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
if err := protocol.Encode(&buf, v); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
var _ Service = (*App)(nil)
|
||||
Reference in New Issue
Block a user