diff options
| author | Peter Stone <thepeterstone@gmail.com> | 2026-07-15 09:49:01 +0000 |
|---|---|---|
| committer | Peter Stone <thepeterstone@gmail.com> | 2026-07-16 02:53:12 +0000 |
| commit | 1a825cc178b95a93cf5c8b210eb204b2ebe87acb (patch) | |
| tree | 70456a5dc5808f4197d1f112b5be345cf2377783 /internal | |
| parent | 3d69922b0a9904be7e2261d4376ea051bb29afd4 (diff) | |
refactor(store): extract agent session/trust methods from sqlite.go
Pure move: all agent-session and agent CRUD/trust methods (~340 lines)
plus their full test suite and setupTestStoreWithAgents helper, out of
sqlite.go/sqlite_test.go into agents.go/agents_test.go. No behavior
change. Mirrors internal/handlers/agent.go's domain naming.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VTUSAEKfsPc6WGDq45yPHD
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/store/agents.go | 352 | ||||
| -rw-r--r-- | internal/store/agents_test.go | 371 | ||||
| -rw-r--r-- | internal/store/sqlite.go | 345 | ||||
| -rw-r--r-- | internal/store/sqlite_test.go | 359 |
4 files changed, 723 insertions, 704 deletions
diff --git a/internal/store/agents.go b/internal/store/agents.go new file mode 100644 index 0000000..130c40d --- /dev/null +++ b/internal/store/agents.go @@ -0,0 +1,352 @@ +package store + +import ( + "database/sql" + "errors" + "time" + + "task-dashboard/internal/models" +) + +// CreateAgentSession creates a new pending agent session +func (s *Store) CreateAgentSession(session *models.AgentSession) error { + result, err := s.db.Exec(` + INSERT INTO agent_sessions (request_token, agent_name, agent_id, status, expires_at) + VALUES (?, ?, ?, 'pending', ?) + `, session.RequestToken, session.AgentName, session.AgentID, session.ExpiresAt) + if err != nil { + return err + } + id, err := result.LastInsertId() + if err != nil { + return err + } + session.ID = id + return nil +} + +// GetAgentSessionByRequestToken retrieves a session by request token +func (s *Store) GetAgentSessionByRequestToken(token string) (*models.AgentSession, error) { + var session models.AgentSession + var sessionToken sql.NullString + var sessionExpiresAt sql.NullTime + + err := s.db.QueryRow(` + SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at, session_token, session_expires_at + FROM agent_sessions + WHERE request_token = ? + `, token).Scan( + &session.ID, + &session.RequestToken, + &session.AgentName, + &session.AgentID, + &session.Status, + &session.CreatedAt, + &session.ExpiresAt, + &sessionToken, + &sessionExpiresAt, + ) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + if sessionToken.Valid { + session.SessionToken = sessionToken.String + } + if sessionExpiresAt.Valid { + session.SessionExpiresAt = &sessionExpiresAt.Time + } + return &session, nil +} + +// GetPendingAgentSessionByAgentID retrieves an existing pending session for an agent +func (s *Store) GetPendingAgentSessionByAgentID(agentID string) (*models.AgentSession, error) { + var session models.AgentSession + + err := s.db.QueryRow(` + SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at + FROM agent_sessions + WHERE agent_id = ? AND status = 'pending' AND expires_at > datetime('now', 'localtime') + ORDER BY created_at DESC + LIMIT 1 + `, agentID).Scan( + &session.ID, + &session.RequestToken, + &session.AgentName, + &session.AgentID, + &session.Status, + &session.CreatedAt, + &session.ExpiresAt, + ) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return &session, nil +} + +// GetAgentSessionBySessionToken retrieves a session by session token +func (s *Store) GetAgentSessionBySessionToken(token string) (*models.AgentSession, error) { + var session models.AgentSession + var sessionToken sql.NullString + var sessionExpiresAt sql.NullTime + + err := s.db.QueryRow(` + SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at, session_token, session_expires_at + FROM agent_sessions + WHERE session_token = ? AND status = 'approved' + `, token).Scan( + &session.ID, + &session.RequestToken, + &session.AgentName, + &session.AgentID, + &session.Status, + &session.CreatedAt, + &session.ExpiresAt, + &sessionToken, + &sessionExpiresAt, + ) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + if sessionToken.Valid { + session.SessionToken = sessionToken.String + } + if sessionExpiresAt.Valid { + session.SessionExpiresAt = &sessionExpiresAt.Time + } + return &session, nil +} + +// ApproveAgentSession approves a pending session +func (s *Store) ApproveAgentSession(requestToken, sessionToken string, sessionExpiresAt time.Time) error { + result, err := s.db.Exec(` + UPDATE agent_sessions + SET status = 'approved', session_token = ?, session_expires_at = ? + WHERE request_token = ? AND status = 'pending' + `, sessionToken, sessionExpiresAt, requestToken) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return errors.New("session not found or already processed") + } + return nil +} + +// DenyAgentSession denies a pending session +func (s *Store) DenyAgentSession(requestToken string) error { + result, err := s.db.Exec(` + UPDATE agent_sessions + SET status = 'denied' + WHERE request_token = ? AND status = 'pending' + `, requestToken) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return errors.New("session not found or already processed") + } + return nil +} + +// GetPendingAgentSessions retrieves all unexpired pending sessions +func (s *Store) GetPendingAgentSessions() ([]models.AgentSession, error) { + rows, err := s.db.Query(` + SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at + FROM agent_sessions + WHERE status = 'pending' AND expires_at > datetime('now', 'localtime') + ORDER BY created_at DESC + `) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + var sessions []models.AgentSession + for rows.Next() { + var session models.AgentSession + if err := rows.Scan( + &session.ID, + &session.RequestToken, + &session.AgentName, + &session.AgentID, + &session.Status, + &session.CreatedAt, + &session.ExpiresAt, + ); err != nil { + return nil, err + } + sessions = append(sessions, session) + } + return sessions, rows.Err() +} + +// InvalidatePreviousAgentSessions marks previous sessions for an agent as expired +func (s *Store) InvalidatePreviousAgentSessions(agentID string) error { + _, err := s.db.Exec(` + UPDATE agent_sessions + SET status = 'expired' + WHERE agent_id = ? AND status IN ('pending', 'approved') + `, agentID) + return err +} + +// GetAgentByAgentID retrieves an agent by their agent_id (UUID) +func (s *Store) GetAgentByAgentID(agentID string) (*models.Agent, error) { + var agent models.Agent + var lastSeen sql.NullTime + + err := s.db.QueryRow(` + SELECT id, name, agent_id, created_at, last_seen, trusted + FROM agents + WHERE agent_id = ? + `, agentID).Scan( + &agent.ID, + &agent.Name, + &agent.AgentID, + &agent.CreatedAt, + &lastSeen, + &agent.Trusted, + ) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + if lastSeen.Valid { + agent.LastSeen = &lastSeen.Time + } + return &agent, nil +} + +// GetAgentByName retrieves an agent by name +func (s *Store) GetAgentByName(name string) (*models.Agent, error) { + var agent models.Agent + var lastSeen sql.NullTime + + err := s.db.QueryRow(` + SELECT id, name, agent_id, created_at, last_seen, trusted + FROM agents + WHERE name = ? + `, name).Scan( + &agent.ID, + &agent.Name, + &agent.AgentID, + &agent.CreatedAt, + &lastSeen, + &agent.Trusted, + ) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + if lastSeen.Valid { + agent.LastSeen = &lastSeen.Time + } + return &agent, nil +} + +// CreateOrUpdateAgent creates or updates an agent record +func (s *Store) CreateOrUpdateAgent(name, agentID string) error { + _, err := s.db.Exec(` + INSERT INTO agents (name, agent_id, last_seen, trusted) + VALUES (?, ?, datetime('now'), 1) + ON CONFLICT(agent_id) DO UPDATE SET + name = excluded.name, + last_seen = datetime('now') + `, name, agentID) + return err +} + +// UpdateAgentLastSeen updates the last_seen timestamp for an agent +func (s *Store) UpdateAgentLastSeen(agentID string) error { + _, err := s.db.Exec(` + UPDATE agents SET last_seen = datetime('now') + WHERE agent_id = ? + `, agentID) + return err +} + +// GetAllAgents retrieves all agents +func (s *Store) GetAllAgents() ([]models.Agent, error) { + rows, err := s.db.Query(` + SELECT id, name, agent_id, created_at, last_seen, trusted + FROM agents + ORDER BY last_seen DESC NULLS LAST + `) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + var agents []models.Agent + for rows.Next() { + var agent models.Agent + var lastSeen sql.NullTime + if err := rows.Scan( + &agent.ID, + &agent.Name, + &agent.AgentID, + &agent.CreatedAt, + &lastSeen, + &agent.Trusted, + ); err != nil { + return nil, err + } + if lastSeen.Valid { + agent.LastSeen = &lastSeen.Time + } + agents = append(agents, agent) + } + return agents, rows.Err() +} + +// RevokeAgent sets trusted=false for an agent +func (s *Store) RevokeAgent(agentID string) error { + _, err := s.db.Exec(`UPDATE agents SET trusted = 0 WHERE agent_id = ?`, agentID) + return err +} + +// CheckAgentTrust determines trust level for an agent request +func (s *Store) CheckAgentTrust(name, agentID string) (models.AgentTrustLevel, error) { + // Check if this exact agent_id is known + existingByID, err := s.GetAgentByAgentID(agentID) + if err != nil { + return "", err + } + + // Check if this name is known with a different ID + existingByName, err := s.GetAgentByName(name) + if err != nil { + return "", err + } + + if existingByID != nil && existingByID.Name == name && existingByID.Trusted { + return models.AgentTrustRecognized, nil + } + + if existingByName != nil && existingByName.AgentID != agentID { + return models.AgentTrustSuspicious, nil + } + + return models.AgentTrustNew, nil +} diff --git a/internal/store/agents_test.go b/internal/store/agents_test.go new file mode 100644 index 0000000..1285e55 --- /dev/null +++ b/internal/store/agents_test.go @@ -0,0 +1,371 @@ +package store + +import ( + "database/sql" + "path/filepath" + "testing" + "time" + + _ "github.com/mattn/go-sqlite3" + "task-dashboard/internal/models" +) + +// ============================================================================= +// Agent Session Tests +// ============================================================================= + +func setupTestStoreWithAgents(t *testing.T) *Store { + t.Helper() + + tempDir := t.TempDir() + dbPath := filepath.Join(tempDir, "test.db") + + db, err := sql.Open("sqlite3", dbPath) + if err != nil { + t.Fatalf("Failed to open test database: %v", err) + } + + db.SetMaxOpenConns(1) + store := &Store{db: db} + + schema := ` + CREATE TABLE IF NOT EXISTS agents ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + agent_id TEXT UNIQUE NOT NULL, + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + last_seen DATETIME, + trusted BOOLEAN DEFAULT 1 + ); + CREATE TABLE IF NOT EXISTS agent_sessions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + request_token TEXT UNIQUE NOT NULL, + agent_name TEXT NOT NULL, + agent_id TEXT NOT NULL, + status TEXT DEFAULT 'pending', + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + expires_at DATETIME NOT NULL, + session_token TEXT, + session_expires_at DATETIME + ); + ` + if _, err := db.Exec(schema); err != nil { + t.Fatalf("Failed to create schema: %v", err) + } + + return store +} + +func TestAgentSession_CreateAndRetrieve(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + expiresAt := time.Now().Add(5 * time.Minute) + session := &models.AgentSession{ + RequestToken: "req-token-123", + AgentName: "TestAgent", + AgentID: "agent-uuid-123", + ExpiresAt: expiresAt, + } + + // Create session + if err := store.CreateAgentSession(session); err != nil { + t.Fatalf("Failed to create session: %v", err) + } + if session.ID == 0 { + t.Error("Session ID should be set after create") + } + + // Get by request token + retrieved, err := store.GetAgentSessionByRequestToken("req-token-123") + if err != nil { + t.Fatalf("Failed to get session: %v", err) + } + if retrieved == nil { + t.Fatal("Expected session to exist") + } + if retrieved.AgentName != "TestAgent" { + t.Errorf("Expected name 'TestAgent', got '%s'", retrieved.AgentName) + } + if retrieved.Status != "pending" { + t.Errorf("Expected status 'pending', got '%s'", retrieved.Status) + } + + // Get pending by agent ID + pending, err := store.GetPendingAgentSessionByAgentID("agent-uuid-123") + if err != nil { + t.Fatalf("Failed to get pending session: %v", err) + } + if pending == nil { + t.Fatal("Expected pending session to exist") + } +} + +func TestAgentSession_ApproveAndDeny(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + // Create two sessions + session1 := &models.AgentSession{ + RequestToken: "approve-token", + AgentName: "Agent1", + AgentID: "agent-1", + ExpiresAt: time.Now().Add(5 * time.Minute), + } + session2 := &models.AgentSession{ + RequestToken: "deny-token", + AgentName: "Agent2", + AgentID: "agent-2", + ExpiresAt: time.Now().Add(5 * time.Minute), + } + _ = store.CreateAgentSession(session1) + _ = store.CreateAgentSession(session2) + + // Approve session1 + sessionExpiry := time.Now().Add(1 * time.Hour) + if err := store.ApproveAgentSession("approve-token", "session-token-abc", sessionExpiry); err != nil { + t.Fatalf("Failed to approve session: %v", err) + } + + // Verify approval + approved, _ := store.GetAgentSessionByRequestToken("approve-token") + if approved.Status != "approved" { + t.Errorf("Expected status 'approved', got '%s'", approved.Status) + } + if approved.SessionToken != "session-token-abc" { + t.Errorf("Expected session token 'session-token-abc', got '%s'", approved.SessionToken) + } + + // Deny session2 + if err := store.DenyAgentSession("deny-token"); err != nil { + t.Fatalf("Failed to deny session: %v", err) + } + + denied, _ := store.GetAgentSessionByRequestToken("deny-token") + if denied.Status != "denied" { + t.Errorf("Expected status 'denied', got '%s'", denied.Status) + } +} + +func TestAgentSession_GetBySessionToken(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + session := &models.AgentSession{ + RequestToken: "req-for-session", + AgentName: "SessionAgent", + AgentID: "session-agent", + ExpiresAt: time.Now().Add(5 * time.Minute), + } + _ = store.CreateAgentSession(session) + _ = store.ApproveAgentSession("req-for-session", "active-session", time.Now().Add(1*time.Hour)) + + // Get by session token + retrieved, err := store.GetAgentSessionBySessionToken("active-session") + if err != nil { + t.Fatalf("Failed to get by session token: %v", err) + } + if retrieved == nil { + t.Fatal("Expected session to exist") + } + if retrieved.AgentName != "SessionAgent" { + t.Errorf("Expected 'SessionAgent', got '%s'", retrieved.AgentName) + } +} + +func TestAgentSession_GetPending(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + // Create pending sessions + for i := 0; i < 3; i++ { + session := &models.AgentSession{ + RequestToken: "pending-" + string(rune('0'+i)), + AgentName: "Agent" + string(rune('0'+i)), + AgentID: "agent-" + string(rune('0'+i)), + ExpiresAt: time.Now().Add(5 * time.Minute), + } + _ = store.CreateAgentSession(session) + } + + // Get pending sessions + pending, err := store.GetPendingAgentSessions() + if err != nil { + t.Fatalf("Failed to get pending sessions: %v", err) + } + if len(pending) != 3 { + t.Errorf("Expected 3 pending sessions, got %d", len(pending)) + } +} + +func TestAgentSession_Invalidate(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + // Create sessions for same agent + for i := 0; i < 2; i++ { + session := &models.AgentSession{ + RequestToken: "inv-" + string(rune('0'+i)), + AgentName: "SameAgent", + AgentID: "same-agent", + ExpiresAt: time.Now().Add(5 * time.Minute), + } + _ = store.CreateAgentSession(session) + } + + // Invalidate all sessions for agent + if err := store.InvalidatePreviousAgentSessions("same-agent"); err != nil { + t.Fatalf("Failed to invalidate sessions: %v", err) + } + + // Verify no pending sessions + pending, _ := store.GetPendingAgentSessions() + for _, s := range pending { + if s.AgentID == "same-agent" { + t.Error("Session should be invalidated") + } + } +} + +// ============================================================================= +// Agent Tests +// ============================================================================= + +func TestAgent_CreateAndRetrieve(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + // Create agent + if err := store.CreateOrUpdateAgent("TestBot", "bot-uuid-123"); err != nil { + t.Fatalf("Failed to create agent: %v", err) + } + + // Get by agent ID + agent, err := store.GetAgentByAgentID("bot-uuid-123") + if err != nil { + t.Fatalf("Failed to get agent: %v", err) + } + if agent == nil { + t.Fatal("Expected agent to exist") + } + if agent.Name != "TestBot" { + t.Errorf("Expected name 'TestBot', got '%s'", agent.Name) + } + if !agent.Trusted { + t.Error("New agent should be trusted by default") + } + + // Get by name + byName, err := store.GetAgentByName("TestBot") + if err != nil { + t.Fatalf("Failed to get agent by name: %v", err) + } + if byName == nil { + t.Fatal("Expected agent to exist by name") + } + + // Get all agents + all, err := store.GetAllAgents() + if err != nil { + t.Fatalf("Failed to get all agents: %v", err) + } + if len(all) != 1 { + t.Errorf("Expected 1 agent, got %d", len(all)) + } +} + +func TestAgent_UpdateLastSeen(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + _ = store.CreateOrUpdateAgent("SeenBot", "seen-uuid") + + // Update last seen + if err := store.UpdateAgentLastSeen("seen-uuid"); err != nil { + t.Fatalf("Failed to update last seen: %v", err) + } + + agent, _ := store.GetAgentByAgentID("seen-uuid") + if agent.LastSeen == nil { + t.Error("LastSeen should be set after update") + } +} + +func TestAgent_TrustLevels(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + // Check trust for unknown agent (new) + trust, err := store.CheckAgentTrust("UnknownBot", "unknown-uuid") + if err != nil { + t.Fatalf("Failed to check trust: %v", err) + } + if trust != models.AgentTrustNew { + t.Errorf("Expected AgentTrustNew, got %v", trust) + } + + // Create agent + _ = store.CreateOrUpdateAgent("TrustBot", "trust-uuid") + + // Check trust for recognized agent + trust, _ = store.CheckAgentTrust("TrustBot", "trust-uuid") + if trust != models.AgentTrustRecognized { + t.Errorf("Expected AgentTrustRecognized, got %v", trust) + } + + // Check trust for suspicious agent (same name, different uuid) + trust, _ = store.CheckAgentTrust("TrustBot", "different-uuid") + if trust != models.AgentTrustSuspicious { + t.Errorf("Expected AgentTrustSuspicious, got %v", trust) + } +} + +func TestAgent_Revoke(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + _ = store.CreateOrUpdateAgent("RevokeBot", "revoke-uuid") + + // Verify agent exists + agent, _ := store.GetAgentByAgentID("revoke-uuid") + if agent == nil { + t.Fatal("Agent should exist") + } + + // Revoke agent + if err := store.RevokeAgent("revoke-uuid"); err != nil { + t.Fatalf("Failed to revoke agent: %v", err) + } + + // After revoke, agent should still exist but be in different state + // (revoke doesn't delete, just marks somehow - let's verify it doesn't error) +} + +func TestAgent_NonExistent(t *testing.T) { + store := setupTestStoreWithAgents(t) + defer func() { _ = store.Close() }() + + // Get non-existent agent + agent, err := store.GetAgentByAgentID("does-not-exist") + if err != nil { + t.Fatalf("Should not error for non-existent agent: %v", err) + } + if agent != nil { + t.Error("Agent should be nil for non-existent") + } + + // Get non-existent by name + byName, err := store.GetAgentByName("unknown-name") + if err != nil { + t.Fatalf("Should not error for non-existent name: %v", err) + } + if byName != nil { + t.Error("Agent should be nil for non-existent name") + } + + // Check trust for non-existent (should be new) + trust, _ := store.CheckAgentTrust("UnknownBot", "unknown-uuid") + if trust != models.AgentTrustNew { + t.Errorf("Expected AgentTrustNew for unknown, got %v", trust) + } +} diff --git a/internal/store/sqlite.go b/internal/store/sqlite.go index d26dd3e..d48cb09 100644 --- a/internal/store/sqlite.go +++ b/internal/store/sqlite.go @@ -432,351 +432,6 @@ func (s *Store) GetCardsByDateRange(start, end time.Time) ([]models.Card, error) // Calendar event operations // SaveCalendarEvents replaces all cached calendar events -// Agent operations - -// CreateAgentSession creates a new pending agent session -func (s *Store) CreateAgentSession(session *models.AgentSession) error { - result, err := s.db.Exec(` - INSERT INTO agent_sessions (request_token, agent_name, agent_id, status, expires_at) - VALUES (?, ?, ?, 'pending', ?) - `, session.RequestToken, session.AgentName, session.AgentID, session.ExpiresAt) - if err != nil { - return err - } - id, err := result.LastInsertId() - if err != nil { - return err - } - session.ID = id - return nil -} - -// GetAgentSessionByRequestToken retrieves a session by request token -func (s *Store) GetAgentSessionByRequestToken(token string) (*models.AgentSession, error) { - var session models.AgentSession - var sessionToken sql.NullString - var sessionExpiresAt sql.NullTime - - err := s.db.QueryRow(` - SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at, session_token, session_expires_at - FROM agent_sessions - WHERE request_token = ? - `, token).Scan( - &session.ID, - &session.RequestToken, - &session.AgentName, - &session.AgentID, - &session.Status, - &session.CreatedAt, - &session.ExpiresAt, - &sessionToken, - &sessionExpiresAt, - ) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - if sessionToken.Valid { - session.SessionToken = sessionToken.String - } - if sessionExpiresAt.Valid { - session.SessionExpiresAt = &sessionExpiresAt.Time - } - return &session, nil -} - -// GetPendingAgentSessionByAgentID retrieves an existing pending session for an agent -func (s *Store) GetPendingAgentSessionByAgentID(agentID string) (*models.AgentSession, error) { - var session models.AgentSession - - err := s.db.QueryRow(` - SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at - FROM agent_sessions - WHERE agent_id = ? AND status = 'pending' AND expires_at > datetime('now', 'localtime') - ORDER BY created_at DESC - LIMIT 1 - `, agentID).Scan( - &session.ID, - &session.RequestToken, - &session.AgentName, - &session.AgentID, - &session.Status, - &session.CreatedAt, - &session.ExpiresAt, - ) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &session, nil -} - -// GetAgentSessionBySessionToken retrieves a session by session token -func (s *Store) GetAgentSessionBySessionToken(token string) (*models.AgentSession, error) { - var session models.AgentSession - var sessionToken sql.NullString - var sessionExpiresAt sql.NullTime - - err := s.db.QueryRow(` - SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at, session_token, session_expires_at - FROM agent_sessions - WHERE session_token = ? AND status = 'approved' - `, token).Scan( - &session.ID, - &session.RequestToken, - &session.AgentName, - &session.AgentID, - &session.Status, - &session.CreatedAt, - &session.ExpiresAt, - &sessionToken, - &sessionExpiresAt, - ) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - if sessionToken.Valid { - session.SessionToken = sessionToken.String - } - if sessionExpiresAt.Valid { - session.SessionExpiresAt = &sessionExpiresAt.Time - } - return &session, nil -} - -// ApproveAgentSession approves a pending session -func (s *Store) ApproveAgentSession(requestToken, sessionToken string, sessionExpiresAt time.Time) error { - result, err := s.db.Exec(` - UPDATE agent_sessions - SET status = 'approved', session_token = ?, session_expires_at = ? - WHERE request_token = ? AND status = 'pending' - `, sessionToken, sessionExpiresAt, requestToken) - if err != nil { - return err - } - affected, err := result.RowsAffected() - if err != nil { - return err - } - if affected == 0 { - return errors.New("session not found or already processed") - } - return nil -} - -// DenyAgentSession denies a pending session -func (s *Store) DenyAgentSession(requestToken string) error { - result, err := s.db.Exec(` - UPDATE agent_sessions - SET status = 'denied' - WHERE request_token = ? AND status = 'pending' - `, requestToken) - if err != nil { - return err - } - affected, err := result.RowsAffected() - if err != nil { - return err - } - if affected == 0 { - return errors.New("session not found or already processed") - } - return nil -} - -// GetPendingAgentSessions retrieves all unexpired pending sessions -func (s *Store) GetPendingAgentSessions() ([]models.AgentSession, error) { - rows, err := s.db.Query(` - SELECT id, request_token, agent_name, agent_id, status, created_at, expires_at - FROM agent_sessions - WHERE status = 'pending' AND expires_at > datetime('now', 'localtime') - ORDER BY created_at DESC - `) - if err != nil { - return nil, err - } - defer func() { _ = rows.Close() }() - - var sessions []models.AgentSession - for rows.Next() { - var session models.AgentSession - if err := rows.Scan( - &session.ID, - &session.RequestToken, - &session.AgentName, - &session.AgentID, - &session.Status, - &session.CreatedAt, - &session.ExpiresAt, - ); err != nil { - return nil, err - } - sessions = append(sessions, session) - } - return sessions, rows.Err() -} - -// InvalidatePreviousAgentSessions marks previous sessions for an agent as expired -func (s *Store) InvalidatePreviousAgentSessions(agentID string) error { - _, err := s.db.Exec(` - UPDATE agent_sessions - SET status = 'expired' - WHERE agent_id = ? AND status IN ('pending', 'approved') - `, agentID) - return err -} - -// GetAgentByAgentID retrieves an agent by their agent_id (UUID) -func (s *Store) GetAgentByAgentID(agentID string) (*models.Agent, error) { - var agent models.Agent - var lastSeen sql.NullTime - - err := s.db.QueryRow(` - SELECT id, name, agent_id, created_at, last_seen, trusted - FROM agents - WHERE agent_id = ? - `, agentID).Scan( - &agent.ID, - &agent.Name, - &agent.AgentID, - &agent.CreatedAt, - &lastSeen, - &agent.Trusted, - ) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - if lastSeen.Valid { - agent.LastSeen = &lastSeen.Time - } - return &agent, nil -} - -// GetAgentByName retrieves an agent by name -func (s *Store) GetAgentByName(name string) (*models.Agent, error) { - var agent models.Agent - var lastSeen sql.NullTime - - err := s.db.QueryRow(` - SELECT id, name, agent_id, created_at, last_seen, trusted - FROM agents - WHERE name = ? - `, name).Scan( - &agent.ID, - &agent.Name, - &agent.AgentID, - &agent.CreatedAt, - &lastSeen, - &agent.Trusted, - ) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - if lastSeen.Valid { - agent.LastSeen = &lastSeen.Time - } - return &agent, nil -} - -// CreateOrUpdateAgent creates or updates an agent record -func (s *Store) CreateOrUpdateAgent(name, agentID string) error { - _, err := s.db.Exec(` - INSERT INTO agents (name, agent_id, last_seen, trusted) - VALUES (?, ?, datetime('now'), 1) - ON CONFLICT(agent_id) DO UPDATE SET - name = excluded.name, - last_seen = datetime('now') - `, name, agentID) - return err -} - -// UpdateAgentLastSeen updates the last_seen timestamp for an agent -func (s *Store) UpdateAgentLastSeen(agentID string) error { - _, err := s.db.Exec(` - UPDATE agents SET last_seen = datetime('now') - WHERE agent_id = ? - `, agentID) - return err -} - -// GetAllAgents retrieves all agents -func (s *Store) GetAllAgents() ([]models.Agent, error) { - rows, err := s.db.Query(` - SELECT id, name, agent_id, created_at, last_seen, trusted - FROM agents - ORDER BY last_seen DESC NULLS LAST - `) - if err != nil { - return nil, err - } - defer func() { _ = rows.Close() }() - - var agents []models.Agent - for rows.Next() { - var agent models.Agent - var lastSeen sql.NullTime - if err := rows.Scan( - &agent.ID, - &agent.Name, - &agent.AgentID, - &agent.CreatedAt, - &lastSeen, - &agent.Trusted, - ); err != nil { - return nil, err - } - if lastSeen.Valid { - agent.LastSeen = &lastSeen.Time - } - agents = append(agents, agent) - } - return agents, rows.Err() -} - -// RevokeAgent sets trusted=false for an agent -func (s *Store) RevokeAgent(agentID string) error { - _, err := s.db.Exec(`UPDATE agents SET trusted = 0 WHERE agent_id = ?`, agentID) - return err -} - -// CheckAgentTrust determines trust level for an agent request -func (s *Store) CheckAgentTrust(name, agentID string) (models.AgentTrustLevel, error) { - // Check if this exact agent_id is known - existingByID, err := s.GetAgentByAgentID(agentID) - if err != nil { - return "", err - } - - // Check if this name is known with a different ID - existingByName, err := s.GetAgentByName(name) - if err != nil { - return "", err - } - - if existingByID != nil && existingByID.Name == name && existingByID.Trusted { - return models.AgentTrustRecognized, nil - } - - if existingByName != nil && existingByName.AgentID != agentID { - return models.AgentTrustSuspicious, nil - } - - return models.AgentTrustNew, nil -} - // Completed tasks log // SaveCompletedTask logs a completed task diff --git a/internal/store/sqlite_test.go b/internal/store/sqlite_test.go index dc9bfad..3c5a7b6 100644 --- a/internal/store/sqlite_test.go +++ b/internal/store/sqlite_test.go @@ -1145,362 +1145,3 @@ func TestSyncTokens_SetGetClear(t *testing.T) { } } -// ============================================================================= -// Agent Session Tests -// ============================================================================= - -func setupTestStoreWithAgents(t *testing.T) *Store { - t.Helper() - - tempDir := t.TempDir() - dbPath := filepath.Join(tempDir, "test.db") - - db, err := sql.Open("sqlite3", dbPath) - if err != nil { - t.Fatalf("Failed to open test database: %v", err) - } - - db.SetMaxOpenConns(1) - store := &Store{db: db} - - schema := ` - CREATE TABLE IF NOT EXISTS agents ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT NOT NULL, - agent_id TEXT UNIQUE NOT NULL, - created_at DATETIME DEFAULT CURRENT_TIMESTAMP, - last_seen DATETIME, - trusted BOOLEAN DEFAULT 1 - ); - CREATE TABLE IF NOT EXISTS agent_sessions ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - request_token TEXT UNIQUE NOT NULL, - agent_name TEXT NOT NULL, - agent_id TEXT NOT NULL, - status TEXT DEFAULT 'pending', - created_at DATETIME DEFAULT CURRENT_TIMESTAMP, - expires_at DATETIME NOT NULL, - session_token TEXT, - session_expires_at DATETIME - ); - ` - if _, err := db.Exec(schema); err != nil { - t.Fatalf("Failed to create schema: %v", err) - } - - return store -} - -func TestAgentSession_CreateAndRetrieve(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - expiresAt := time.Now().Add(5 * time.Minute) - session := &models.AgentSession{ - RequestToken: "req-token-123", - AgentName: "TestAgent", - AgentID: "agent-uuid-123", - ExpiresAt: expiresAt, - } - - // Create session - if err := store.CreateAgentSession(session); err != nil { - t.Fatalf("Failed to create session: %v", err) - } - if session.ID == 0 { - t.Error("Session ID should be set after create") - } - - // Get by request token - retrieved, err := store.GetAgentSessionByRequestToken("req-token-123") - if err != nil { - t.Fatalf("Failed to get session: %v", err) - } - if retrieved == nil { - t.Fatal("Expected session to exist") - } - if retrieved.AgentName != "TestAgent" { - t.Errorf("Expected name 'TestAgent', got '%s'", retrieved.AgentName) - } - if retrieved.Status != "pending" { - t.Errorf("Expected status 'pending', got '%s'", retrieved.Status) - } - - // Get pending by agent ID - pending, err := store.GetPendingAgentSessionByAgentID("agent-uuid-123") - if err != nil { - t.Fatalf("Failed to get pending session: %v", err) - } - if pending == nil { - t.Fatal("Expected pending session to exist") - } -} - -func TestAgentSession_ApproveAndDeny(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - // Create two sessions - session1 := &models.AgentSession{ - RequestToken: "approve-token", - AgentName: "Agent1", - AgentID: "agent-1", - ExpiresAt: time.Now().Add(5 * time.Minute), - } - session2 := &models.AgentSession{ - RequestToken: "deny-token", - AgentName: "Agent2", - AgentID: "agent-2", - ExpiresAt: time.Now().Add(5 * time.Minute), - } - _ = store.CreateAgentSession(session1) - _ = store.CreateAgentSession(session2) - - // Approve session1 - sessionExpiry := time.Now().Add(1 * time.Hour) - if err := store.ApproveAgentSession("approve-token", "session-token-abc", sessionExpiry); err != nil { - t.Fatalf("Failed to approve session: %v", err) - } - - // Verify approval - approved, _ := store.GetAgentSessionByRequestToken("approve-token") - if approved.Status != "approved" { - t.Errorf("Expected status 'approved', got '%s'", approved.Status) - } - if approved.SessionToken != "session-token-abc" { - t.Errorf("Expected session token 'session-token-abc', got '%s'", approved.SessionToken) - } - - // Deny session2 - if err := store.DenyAgentSession("deny-token"); err != nil { - t.Fatalf("Failed to deny session: %v", err) - } - - denied, _ := store.GetAgentSessionByRequestToken("deny-token") - if denied.Status != "denied" { - t.Errorf("Expected status 'denied', got '%s'", denied.Status) - } -} - -func TestAgentSession_GetBySessionToken(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - session := &models.AgentSession{ - RequestToken: "req-for-session", - AgentName: "SessionAgent", - AgentID: "session-agent", - ExpiresAt: time.Now().Add(5 * time.Minute), - } - _ = store.CreateAgentSession(session) - _ = store.ApproveAgentSession("req-for-session", "active-session", time.Now().Add(1*time.Hour)) - - // Get by session token - retrieved, err := store.GetAgentSessionBySessionToken("active-session") - if err != nil { - t.Fatalf("Failed to get by session token: %v", err) - } - if retrieved == nil { - t.Fatal("Expected session to exist") - } - if retrieved.AgentName != "SessionAgent" { - t.Errorf("Expected 'SessionAgent', got '%s'", retrieved.AgentName) - } -} - -func TestAgentSession_GetPending(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - // Create pending sessions - for i := 0; i < 3; i++ { - session := &models.AgentSession{ - RequestToken: "pending-" + string(rune('0'+i)), - AgentName: "Agent" + string(rune('0'+i)), - AgentID: "agent-" + string(rune('0'+i)), - ExpiresAt: time.Now().Add(5 * time.Minute), - } - _ = store.CreateAgentSession(session) - } - - // Get pending sessions - pending, err := store.GetPendingAgentSessions() - if err != nil { - t.Fatalf("Failed to get pending sessions: %v", err) - } - if len(pending) != 3 { - t.Errorf("Expected 3 pending sessions, got %d", len(pending)) - } -} - -func TestAgentSession_Invalidate(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - // Create sessions for same agent - for i := 0; i < 2; i++ { - session := &models.AgentSession{ - RequestToken: "inv-" + string(rune('0'+i)), - AgentName: "SameAgent", - AgentID: "same-agent", - ExpiresAt: time.Now().Add(5 * time.Minute), - } - _ = store.CreateAgentSession(session) - } - - // Invalidate all sessions for agent - if err := store.InvalidatePreviousAgentSessions("same-agent"); err != nil { - t.Fatalf("Failed to invalidate sessions: %v", err) - } - - // Verify no pending sessions - pending, _ := store.GetPendingAgentSessions() - for _, s := range pending { - if s.AgentID == "same-agent" { - t.Error("Session should be invalidated") - } - } -} - -// ============================================================================= -// Agent Tests -// ============================================================================= - -func TestAgent_CreateAndRetrieve(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - // Create agent - if err := store.CreateOrUpdateAgent("TestBot", "bot-uuid-123"); err != nil { - t.Fatalf("Failed to create agent: %v", err) - } - - // Get by agent ID - agent, err := store.GetAgentByAgentID("bot-uuid-123") - if err != nil { - t.Fatalf("Failed to get agent: %v", err) - } - if agent == nil { - t.Fatal("Expected agent to exist") - } - if agent.Name != "TestBot" { - t.Errorf("Expected name 'TestBot', got '%s'", agent.Name) - } - if !agent.Trusted { - t.Error("New agent should be trusted by default") - } - - // Get by name - byName, err := store.GetAgentByName("TestBot") - if err != nil { - t.Fatalf("Failed to get agent by name: %v", err) - } - if byName == nil { - t.Fatal("Expected agent to exist by name") - } - - // Get all agents - all, err := store.GetAllAgents() - if err != nil { - t.Fatalf("Failed to get all agents: %v", err) - } - if len(all) != 1 { - t.Errorf("Expected 1 agent, got %d", len(all)) - } -} - -func TestAgent_UpdateLastSeen(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - _ = store.CreateOrUpdateAgent("SeenBot", "seen-uuid") - - // Update last seen - if err := store.UpdateAgentLastSeen("seen-uuid"); err != nil { - t.Fatalf("Failed to update last seen: %v", err) - } - - agent, _ := store.GetAgentByAgentID("seen-uuid") - if agent.LastSeen == nil { - t.Error("LastSeen should be set after update") - } -} - -func TestAgent_TrustLevels(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - // Check trust for unknown agent (new) - trust, err := store.CheckAgentTrust("UnknownBot", "unknown-uuid") - if err != nil { - t.Fatalf("Failed to check trust: %v", err) - } - if trust != models.AgentTrustNew { - t.Errorf("Expected AgentTrustNew, got %v", trust) - } - - // Create agent - _ = store.CreateOrUpdateAgent("TrustBot", "trust-uuid") - - // Check trust for recognized agent - trust, _ = store.CheckAgentTrust("TrustBot", "trust-uuid") - if trust != models.AgentTrustRecognized { - t.Errorf("Expected AgentTrustRecognized, got %v", trust) - } - - // Check trust for suspicious agent (same name, different uuid) - trust, _ = store.CheckAgentTrust("TrustBot", "different-uuid") - if trust != models.AgentTrustSuspicious { - t.Errorf("Expected AgentTrustSuspicious, got %v", trust) - } -} - -func TestAgent_Revoke(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - _ = store.CreateOrUpdateAgent("RevokeBot", "revoke-uuid") - - // Verify agent exists - agent, _ := store.GetAgentByAgentID("revoke-uuid") - if agent == nil { - t.Fatal("Agent should exist") - } - - // Revoke agent - if err := store.RevokeAgent("revoke-uuid"); err != nil { - t.Fatalf("Failed to revoke agent: %v", err) - } - - // After revoke, agent should still exist but be in different state - // (revoke doesn't delete, just marks somehow - let's verify it doesn't error) -} - -func TestAgent_NonExistent(t *testing.T) { - store := setupTestStoreWithAgents(t) - defer func() { _ = store.Close() }() - - // Get non-existent agent - agent, err := store.GetAgentByAgentID("does-not-exist") - if err != nil { - t.Fatalf("Should not error for non-existent agent: %v", err) - } - if agent != nil { - t.Error("Agent should be nil for non-existent") - } - - // Get non-existent by name - byName, err := store.GetAgentByName("unknown-name") - if err != nil { - t.Fatalf("Should not error for non-existent name: %v", err) - } - if byName != nil { - t.Error("Agent should be nil for non-existent name") - } - - // Check trust for non-existent (should be new) - trust, _ := store.CheckAgentTrust("UnknownBot", "unknown-uuid") - if trust != models.AgentTrustNew { - t.Errorf("Expected AgentTrustNew for unknown, got %v", trust) - } -} |
