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 { q := strings.ToLower(query) esc := escapeLikePattern(q) like := "%" + esc + "%" prefix := esc + "%" rows, err = a.db.Read.QueryContext(ctx, ` SELECT id, name, online_since, offline_since, talk_hash FROM endpoints WHERE id > ? AND (lower(id) LIKE ? ESCAPE '\' OR lower(name) LIKE ? ESCAPE '\') 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)