feat: add sqlite session store for remote server
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -0,0 +1,253 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/longbin/agent-notify/internal/remote"
|
||||
)
|
||||
|
||||
const offlineThreshold = 5 * time.Minute
|
||||
|
||||
type Store struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
type SessionRow struct {
|
||||
SessionKey string `json:"session_key"`
|
||||
Hostname string `json:"hostname"`
|
||||
IPs []string `json:"ips"`
|
||||
Agent string `json:"agent"`
|
||||
CWD string `json:"cwd"`
|
||||
Status string `json:"status"`
|
||||
Event string `json:"event,omitempty"`
|
||||
ConversationID string `json:"conversation_id,omitempty"`
|
||||
LastUser string `json:"last_user,omitempty"`
|
||||
LastAgent string `json:"last_agent,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
ReceivedAt time.Time `json:"received_at"`
|
||||
}
|
||||
|
||||
type ListFilters struct {
|
||||
Host string
|
||||
CWD string
|
||||
Agent string
|
||||
Status string
|
||||
}
|
||||
|
||||
type MetaResult struct {
|
||||
Hosts []string `json:"hosts"`
|
||||
CWDs []string `json:"cwds"`
|
||||
Agents []string `json:"agents"`
|
||||
}
|
||||
|
||||
func Open(path string) (*Store, error) {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db, err := sql.Open("sqlite", path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s := &Store{db: db}
|
||||
if err := s.initSchema(); err != nil {
|
||||
db.Close()
|
||||
return nil, err
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (s *Store) Close() error {
|
||||
return s.db.Close()
|
||||
}
|
||||
|
||||
func (s *Store) initSchema() error {
|
||||
stmts := []string{
|
||||
`CREATE TABLE IF NOT EXISTS sessions (
|
||||
session_key TEXT PRIMARY KEY,
|
||||
hostname TEXT NOT NULL,
|
||||
ips TEXT NOT NULL,
|
||||
agent TEXT NOT NULL,
|
||||
cwd TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
event TEXT NOT NULL DEFAULT '',
|
||||
conversation_id TEXT NOT NULL DEFAULT '',
|
||||
last_user TEXT NOT NULL DEFAULT '',
|
||||
last_agent TEXT NOT NULL DEFAULT '',
|
||||
updated_at TEXT NOT NULL,
|
||||
received_at TEXT NOT NULL
|
||||
)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_sessions_hostname ON sessions(hostname)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_sessions_cwd ON sessions(cwd)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_sessions_agent ON sessions(agent)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_sessions_updated_at ON sessions(updated_at)`,
|
||||
}
|
||||
for _, stmt := range stmts {
|
||||
if _, err := s.db.Exec(stmt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) Upsert(report remote.StatusReport) error {
|
||||
key := remote.SessionKey(report.Hostname, remote.PrimaryIP(report.IPs), report.CWD, report.Agent)
|
||||
ipsJSON, err := json.Marshal(report.IPs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
receivedAt := time.Now().UTC()
|
||||
updatedAt := report.UpdatedAt.UTC()
|
||||
if updatedAt.IsZero() {
|
||||
updatedAt = receivedAt
|
||||
}
|
||||
|
||||
_, err = s.db.Exec(`
|
||||
INSERT INTO sessions (
|
||||
session_key, hostname, ips, agent, cwd, status, event,
|
||||
conversation_id, last_user, last_agent, updated_at, received_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(session_key) DO UPDATE SET
|
||||
hostname = excluded.hostname,
|
||||
ips = excluded.ips,
|
||||
agent = excluded.agent,
|
||||
cwd = excluded.cwd,
|
||||
status = excluded.status,
|
||||
event = excluded.event,
|
||||
conversation_id = excluded.conversation_id,
|
||||
last_user = CASE WHEN excluded.last_user = '' THEN sessions.last_user ELSE excluded.last_user END,
|
||||
last_agent = CASE WHEN excluded.last_agent = '' THEN sessions.last_agent ELSE excluded.last_agent END,
|
||||
updated_at = excluded.updated_at,
|
||||
received_at = excluded.received_at
|
||||
`, key, report.Hostname, string(ipsJSON), report.Agent, report.CWD, report.Status,
|
||||
report.Event, report.ConversationID, report.LastUser, report.LastAgent,
|
||||
updatedAt.Format(time.RFC3339Nano), receivedAt.Format(time.RFC3339Nano))
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) List(filters ListFilters) ([]SessionRow, error) {
|
||||
query := `SELECT session_key, hostname, ips, agent, cwd, status, event,
|
||||
conversation_id, last_user, last_agent, updated_at, received_at
|
||||
FROM sessions WHERE 1=1`
|
||||
var args []any
|
||||
|
||||
if filters.Host != "" {
|
||||
query += ` AND hostname = ?`
|
||||
args = append(args, filters.Host)
|
||||
}
|
||||
if filters.CWD != "" {
|
||||
query += ` AND cwd = ?`
|
||||
args = append(args, filters.CWD)
|
||||
}
|
||||
if filters.Agent != "" {
|
||||
query += ` AND agent = ?`
|
||||
args = append(args, filters.Agent)
|
||||
}
|
||||
query += ` ORDER BY updated_at DESC`
|
||||
|
||||
rows, err := s.db.Query(query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []SessionRow
|
||||
for rows.Next() {
|
||||
row, err := scanSessionRow(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
row.Status = displayStatus(row.Status, row.UpdatedAt)
|
||||
if filters.Status != "" && row.Status != filters.Status {
|
||||
continue
|
||||
}
|
||||
result = append(result, row)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func displayStatus(stored string, updatedAt time.Time) string {
|
||||
if time.Since(updatedAt) > offlineThreshold {
|
||||
return "offline"
|
||||
}
|
||||
return stored
|
||||
}
|
||||
|
||||
func (s *Store) Meta() (MetaResult, error) {
|
||||
var meta MetaResult
|
||||
for _, q := range []struct {
|
||||
sql string
|
||||
dest *[]string
|
||||
}{
|
||||
{`SELECT DISTINCT hostname FROM sessions ORDER BY hostname`, &meta.Hosts},
|
||||
{`SELECT DISTINCT cwd FROM sessions ORDER BY cwd`, &meta.CWDs},
|
||||
{`SELECT DISTINCT agent FROM sessions ORDER BY agent`, &meta.Agents},
|
||||
} {
|
||||
rows, err := s.db.Query(q.sql)
|
||||
if err != nil {
|
||||
return MetaResult{}, err
|
||||
}
|
||||
for rows.Next() {
|
||||
var v string
|
||||
if err := rows.Scan(&v); err != nil {
|
||||
rows.Close()
|
||||
return MetaResult{}, err
|
||||
}
|
||||
*q.dest = append(*q.dest, v)
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return MetaResult{}, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return MetaResult{}, err
|
||||
}
|
||||
}
|
||||
return meta, nil
|
||||
}
|
||||
|
||||
type rowScanner interface {
|
||||
Scan(dest ...any) error
|
||||
}
|
||||
|
||||
func scanSessionRow(rows rowScanner) (SessionRow, error) {
|
||||
var row SessionRow
|
||||
var ipsJSON, updatedAtStr, receivedAtStr string
|
||||
if err := rows.Scan(
|
||||
&row.SessionKey, &row.Hostname, &ipsJSON, &row.Agent, &row.CWD, &row.Status,
|
||||
&row.Event, &row.ConversationID, &row.LastUser, &row.LastAgent,
|
||||
&updatedAtStr, &receivedAtStr,
|
||||
); err != nil {
|
||||
return SessionRow{}, err
|
||||
}
|
||||
if err := json.Unmarshal([]byte(ipsJSON), &row.IPs); err != nil {
|
||||
return SessionRow{}, fmt.Errorf("decode ips: %w", err)
|
||||
}
|
||||
if row.IPs == nil {
|
||||
row.IPs = []string{}
|
||||
}
|
||||
var err error
|
||||
row.UpdatedAt, err = parseTime(updatedAtStr)
|
||||
if err != nil {
|
||||
return SessionRow{}, fmt.Errorf("updated_at: %w", err)
|
||||
}
|
||||
row.ReceivedAt, err = parseTime(receivedAtStr)
|
||||
if err != nil {
|
||||
return SessionRow{}, fmt.Errorf("received_at: %w", err)
|
||||
}
|
||||
return row, nil
|
||||
}
|
||||
|
||||
func parseTime(s string) (time.Time, error) {
|
||||
for _, layout := range []string{time.RFC3339Nano, time.RFC3339} {
|
||||
if t, err := time.Parse(layout, s); err == nil {
|
||||
return t.UTC(), nil
|
||||
}
|
||||
}
|
||||
return time.Time{}, fmt.Errorf("invalid time %q", s)
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/longbin/agent-notify/internal/remote"
|
||||
)
|
||||
|
||||
func TestStoreUpsertAndList(t *testing.T) {
|
||||
db := filepath.Join(t.TempDir(), "test.db")
|
||||
s, err := Open(db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
now := time.Now().UTC()
|
||||
base := remote.StatusReport{
|
||||
Hostname: "host-a",
|
||||
IPs: []string{"10.0.0.1"},
|
||||
Agent: "Cursor",
|
||||
CWD: "/proj",
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
r1 := base
|
||||
r1.Status = "waiting"
|
||||
if err := s.Upsert(r1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
r2 := base
|
||||
r2.Status = "running"
|
||||
if err := s.Upsert(r2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
rows, err := s.List(ListFilters{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected 1 row, got %d", len(rows))
|
||||
}
|
||||
if rows[0].Status != "running" {
|
||||
t.Fatalf("expected latest status running, got %q", rows[0].Status)
|
||||
}
|
||||
key := remote.SessionKey("host-a", "10.0.0.1", "/proj", "Cursor")
|
||||
if rows[0].SessionKey != key {
|
||||
t.Fatalf("session_key: got %q want %q", rows[0].SessionKey, key)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStoreOfflineAfter5Min(t *testing.T) {
|
||||
db := filepath.Join(t.TempDir(), "test.db")
|
||||
s, err := Open(db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
stale := time.Now().UTC().Add(-10 * time.Minute)
|
||||
report := remote.StatusReport{
|
||||
Hostname: "host-b",
|
||||
IPs: []string{"192.168.1.5"},
|
||||
Agent: "Claude",
|
||||
CWD: "/work",
|
||||
Status: "waiting",
|
||||
UpdatedAt: stale,
|
||||
}
|
||||
if err := s.Upsert(report); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
rows, err := s.List(ListFilters{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected 1 row, got %d", len(rows))
|
||||
}
|
||||
if rows[0].Status != "offline" {
|
||||
t.Fatalf("expected offline status, got %q", rows[0].Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStoreMergePreservesLastUser(t *testing.T) {
|
||||
db := filepath.Join(t.TempDir(), "test.db")
|
||||
s, err := Open(db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
now := time.Now().UTC()
|
||||
base := remote.StatusReport{
|
||||
Hostname: "host-c",
|
||||
IPs: []string{"10.0.0.2"},
|
||||
Agent: "Cursor",
|
||||
CWD: "/proj",
|
||||
Status: "waiting",
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
r1 := base
|
||||
r1.LastUser = "keep me"
|
||||
if err := s.Upsert(r1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
r2 := base
|
||||
r2.LastUser = ""
|
||||
r2.Status = "running"
|
||||
if err := s.Upsert(r2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
rows, err := s.List(ListFilters{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected 1 row, got %d", len(rows))
|
||||
}
|
||||
if rows[0].LastUser != "keep me" {
|
||||
t.Fatalf("last_user: got %q want %q", rows[0].LastUser, "keep me")
|
||||
}
|
||||
if rows[0].Status != "running" {
|
||||
t.Fatalf("expected status running, got %q", rows[0].Status)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user