diff --git a/.env.example b/.env.example index a7b78029..109b90cd 100644 --- a/.env.example +++ b/.env.example @@ -53,6 +53,18 @@ CCF_PLAYBACK_TIMEOUT="5s" CCF_PLAYBACK_MAX_BYTES=2097152 CCF_PLAYBACK_MAX_CONCURRENT=4 +# Agent remote configuration: instance freshness, retention, pruning (River job; needs +# CCF_WORKER_ENABLED) and the per-agent instance cap. +CCF_AGENT_INSTANCE_STALE_AFTER="10m" +CCF_AGENT_INSTANCE_RETENTION="720h" +CCF_AGENT_INSTANCE_ONESHOT_RETENTION="24h" +CCF_AGENT_INSTANCE_PRUNE_ENABLED=true +# The prune job is deduplicated per hour: it runs at most hourly, whatever the schedule. +CCF_AGENT_INSTANCE_PRUNE_SCHEDULE="0 17 * * * *" +# Cap on non-prunable instances per agent. When full, a new instance replaces the oldest +# stale one (not seen within CCF_AGENT_INSTANCE_STALE_AFTER); 409 only when all are fresh. +CCF_AGENT_MAX_INSTANCES=500 + # Policy evaluation artifacts: POST /api/agent/artifacts (see docs/artifacts.md) CCF_ARTIFACT_MAX_BYTES=16777216 CCF_ARTIFACT_MAX_CONCURRENT=8 diff --git a/cmd/root.go b/cmd/root.go index 701fea9d..defd1427 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -99,6 +99,12 @@ func bindEnvironmentVariables() { viper.MustBindEnv("playback_timeout") viper.MustBindEnv("playback_max_bytes") viper.MustBindEnv("playback_max_concurrent") + viper.MustBindEnv("agent_instance_stale_after") + viper.MustBindEnv("agent_instance_retention") + viper.MustBindEnv("agent_instance_oneshot_retention") + viper.MustBindEnv("agent_instance_prune_enabled") + viper.MustBindEnv("agent_instance_prune_schedule") + viper.MustBindEnv("agent_max_instances") viper.MustBindEnv("artifact_max_bytes") viper.MustBindEnv("artifact_max_concurrent") } diff --git a/internal/config/agents.go b/internal/config/agents.go new file mode 100644 index 00000000..29350a03 --- /dev/null +++ b/internal/config/agents.go @@ -0,0 +1,72 @@ +package config + +import ( + "time" + + "github.com/spf13/viper" +) + +// AgentsConfig tunes agent remote configuration: instance freshness, retention, pruning and +// the per-agent instance cap (R14, R37). +type AgentsConfig struct { + // InstanceStaleAfter: an instance is fresh when seen (heartbeat or report) within this + // window. CCF_AGENT_INSTANCE_STALE_AFTER, default 10m. + InstanceStaleAfter time.Duration `json:"instanceStaleAfter"` + // InstanceRetention: daemon (or unknown) instances not seen for this long are pruned. + // CCF_AGENT_INSTANCE_RETENTION, default 720h. + InstanceRetention time.Duration `json:"instanceRetention"` + // OneShotInstanceRetention: daemon=false instances not seen for this long are pruned. + // CCF_AGENT_INSTANCE_ONESHOT_RETENTION, default 24h. + OneShotInstanceRetention time.Duration `json:"oneShotInstanceRetention"` + // InstancePruneEnabled schedules the prune job (needs the worker service). + // CCF_AGENT_INSTANCE_PRUNE_ENABLED, default true. + InstancePruneEnabled bool `json:"instancePruneEnabled"` + // InstancePruneSchedule is the River (6-field, seconds first) cron of the prune job. + // CCF_AGENT_INSTANCE_PRUNE_SCHEDULE, default "0 17 * * * *" (hourly). The job is + // deduplicated per hour, so it runs at most hourly: a more frequent schedule is not + // honored. + InstancePruneSchedule string `json:"instancePruneSchedule"` + // MaxInstancesPerAgent caps the non-prunable instances of one agent. When the cap is + // reached, a new instance replaces the oldest stale one (not seen within + // InstanceStaleAfter); only when every counted instance is fresh does a report from a new + // instance get 409 (and its heartbeats are not recorded). CCF_AGENT_MAX_INSTANCES, + // default 500. + MaxInstancesPerAgent int `json:"maxInstancesPerAgent"` +} + +// DefaultAgentsConfig returns the defaults. +func DefaultAgentsConfig() *AgentsConfig { + return &AgentsConfig{ + InstanceStaleAfter: 10 * time.Minute, + InstanceRetention: 720 * time.Hour, + OneShotInstanceRetention: 24 * time.Hour, + InstancePruneEnabled: true, + InstancePruneSchedule: "0 17 * * * *", + MaxInstancesPerAgent: 500, + } +} + +// LoadAgentsConfig reads the CCF_AGENT_* settings, falling back to the default for any value +// that is unset or not positive. +func LoadAgentsConfig() *AgentsConfig { + cfg := DefaultAgentsConfig() + if d := viper.GetDuration("agent_instance_stale_after"); d > 0 { + cfg.InstanceStaleAfter = d + } + if d := viper.GetDuration("agent_instance_retention"); d > 0 { + cfg.InstanceRetention = d + } + if d := viper.GetDuration("agent_instance_oneshot_retention"); d > 0 { + cfg.OneShotInstanceRetention = d + } + if viper.IsSet("agent_instance_prune_enabled") { + cfg.InstancePruneEnabled = viper.GetBool("agent_instance_prune_enabled") + } + if s := viper.GetString("agent_instance_prune_schedule"); s != "" { + cfg.InstancePruneSchedule = s + } + if n := viper.GetInt("agent_max_instances"); n > 0 { + cfg.MaxInstancesPerAgent = n + } + return cfg +} diff --git a/internal/config/agents_test.go b/internal/config/agents_test.go new file mode 100644 index 00000000..4e5dfa37 --- /dev/null +++ b/internal/config/agents_test.go @@ -0,0 +1,74 @@ +package config + +import ( + "testing" + "time" + + "github.com/spf13/viper" + "github.com/stretchr/testify/assert" +) + +func TestLoadAgentsConfigDefaults(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + + cfg := LoadAgentsConfig() + + assert.Equal(t, 10*time.Minute, cfg.InstanceStaleAfter) + assert.Equal(t, 720*time.Hour, cfg.InstanceRetention) + assert.Equal(t, 24*time.Hour, cfg.OneShotInstanceRetention) + assert.True(t, cfg.InstancePruneEnabled) + assert.Equal(t, "0 17 * * * *", cfg.InstancePruneSchedule) + assert.Equal(t, 500, cfg.MaxInstancesPerAgent) + assert.Equal(t, DefaultAgentsConfig(), cfg) +} + +func TestLoadAgentsConfigOverrides(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + viper.Set("agent_instance_stale_after", "5m") + viper.Set("agent_instance_retention", "48h") + viper.Set("agent_instance_oneshot_retention", "2h") + viper.Set("agent_instance_prune_enabled", false) + viper.Set("agent_instance_prune_schedule", "0 0 * * * *") + viper.Set("agent_max_instances", 20) + + cfg := LoadAgentsConfig() + + assert.Equal(t, 5*time.Minute, cfg.InstanceStaleAfter) + assert.Equal(t, 48*time.Hour, cfg.InstanceRetention) + assert.Equal(t, 2*time.Hour, cfg.OneShotInstanceRetention) + assert.False(t, cfg.InstancePruneEnabled) + assert.Equal(t, "0 0 * * * *", cfg.InstancePruneSchedule) + assert.Equal(t, 20, cfg.MaxInstancesPerAgent) +} + +func TestLoadAgentsConfigFromEnv(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + viper.SetEnvPrefix("ccf") + viper.AutomaticEnv() + t.Setenv("CCF_AGENT_INSTANCE_PRUNE_ENABLED", "false") + t.Setenv("CCF_AGENT_MAX_INSTANCES", "3") + t.Setenv("CCF_AGENT_INSTANCE_STALE_AFTER", "90s") + + cfg := LoadAgentsConfig() + + assert.False(t, cfg.InstancePruneEnabled) + assert.Equal(t, 3, cfg.MaxInstancesPerAgent) + assert.Equal(t, 90*time.Second, cfg.InstanceStaleAfter) +} + +func TestLoadAgentsConfigIgnoresNonPositiveValues(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + viper.Set("agent_instance_stale_after", "0s") + viper.Set("agent_instance_retention", "-1h") + viper.Set("agent_instance_oneshot_retention", "0") + viper.Set("agent_instance_prune_schedule", "") + viper.Set("agent_max_instances", 0) + + cfg := LoadAgentsConfig() + + assert.Equal(t, DefaultAgentsConfig(), cfg) +} diff --git a/internal/config/config.go b/internal/config/config.go index fa6f63d6..c1cd0b79 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -47,6 +47,7 @@ type Config struct { StrictDisablePublicAgentEndpoints bool Authz *AuthzConfig Playback *PlaybackConfig + Agents *AgentsConfig Artifact *ArtifactConfig EvidenceSubjects *EvidenceSubjectConfig } @@ -284,6 +285,7 @@ func NewConfig(logger *zap.SugaredLogger) *Config { StrictDisablePublicAgentEndpoints: strictDisablePublicAgentEndpoints, Authz: authzConfig, Playback: LoadPlaybackConfig(), + Agents: LoadAgentsConfig(), Artifact: LoadArtifactConfig(), } diff --git a/internal/service/migrator.go b/internal/service/migrator.go index d7d8ac1e..060e1a2e 100644 --- a/internal/service/migrator.go +++ b/internal/service/migrator.go @@ -184,6 +184,7 @@ func MigrateUpWithConfig(db *gorm.DB, cfg *config.Config) error { &relational.Agent{}, &relational.AgentServiceAccountKey{}, &relational.AgentAuthEvent{}, + &relational.AgentConfigRevision{}, &relational.UserNotificationSubscription{}, &relational.SystemNotificationDestination{}, &Heartbeat{}, @@ -1159,6 +1160,7 @@ func MigrateDown(db *gorm.DB) error { &poamrel.PoamItemMilestone{}, &poamrel.PoamItem{}, + &relational.AgentConfigRevision{}, &relational.AgentAuthEvent{}, &relational.AgentServiceAccountKey{}, &relational.Agent{}, diff --git a/internal/service/relational/agent_config.go b/internal/service/relational/agent_config.go new file mode 100644 index 00000000..5cd1fd59 --- /dev/null +++ b/internal/service/relational/agent_config.go @@ -0,0 +1,37 @@ +package relational + +import ( + "errors" + "time" + + "github.com/google/uuid" + "gorm.io/datatypes" + "gorm.io/gorm" +) + +// ErrAgentConfigRevisionAppendOnly is returned when code tries to update or delete a +// configuration revision. +var ErrAgentConfigRevisionAppendOnly = errors.New("agent config revisions are append-only") + +// AgentConfigRevision is one immutable revision of an agent's configuration overlay (D11). +// The desired revision of an agent is MAX(revision). The row id feeds the agent-facing +// opaque ETag (R7), so a DB reset never produces a false 304. +type AgentConfigRevision struct { + UUIDModel + CreatedAt time.Time `json:"createdAt"` + AgentID uuid.UUID `json:"agentId" gorm:"type:uuid;not null;uniqueIndex:idx_agent_config_rev,priority:1"` + Revision int64 `json:"revision" gorm:"not null;uniqueIndex:idx_agent_config_rev,priority:2"` + Overlay datatypes.JSON `json:"overlay" gorm:"type:jsonb;not null"` + Comment *string `json:"comment,omitempty" gorm:"type:text"` + CreatedBy string `json:"createdBy" gorm:"type:text;not null"` // user subject id (email) + CreatedByID *uuid.UUID `json:"createdById,omitempty" gorm:"type:uuid"` // user_uuid claim when present + RevertOf *int64 `json:"revertOf,omitempty"` +} + +func (AgentConfigRevision) TableName() string { return "ccf_agent_config_revisions" } + +// BeforeUpdate keeps revisions append-only. +func (*AgentConfigRevision) BeforeUpdate(*gorm.DB) error { return ErrAgentConfigRevisionAppendOnly } + +// BeforeDelete keeps revisions append-only. +func (*AgentConfigRevision) BeforeDelete(*gorm.DB) error { return ErrAgentConfigRevisionAppendOnly } diff --git a/internal/service/relational/agentcfg/service.go b/internal/service/relational/agentcfg/service.go new file mode 100644 index 00000000..2f93baae --- /dev/null +++ b/internal/service/relational/agentcfg/service.go @@ -0,0 +1,263 @@ +// Package agentcfg stores and serves agent remote-configuration revisions (append-only +// overlays per agent) and the agent instances that report against them. +package agentcfg + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "time" + + "github.com/compliance-framework/api/internal/config" + "github.com/compliance-framework/api/internal/service" + "github.com/compliance-framework/api/internal/service/relational" + "github.com/google/uuid" + "go.uber.org/zap" + "gorm.io/datatypes" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +// Settings tunes instance freshness, retention and the per-agent instance cap. +type Settings struct { + InstanceStaleAfter time.Duration // fresh <=> last_seen_at >= now - InstanceStaleAfter + InstanceRetention time.Duration // daemon (or unknown) instances + OneShotInstanceRetention time.Duration // daemon=false instances (R37) + MaxInstancesPerAgent int // non-prunable instances per agent; the oldest stale one is replaced when full +} + +// WithDefaults fills zero or negative values with config.DefaultAgentsConfig (R37, R14), the +// single source of the defaults. +func (s Settings) WithDefaults() Settings { + d := config.DefaultAgentsConfig() + if s.InstanceStaleAfter <= 0 { + s.InstanceStaleAfter = d.InstanceStaleAfter + } + if s.InstanceRetention <= 0 { + s.InstanceRetention = d.InstanceRetention + } + if s.OneShotInstanceRetention <= 0 { + s.OneShotInstanceRetention = d.OneShotInstanceRetention + } + if s.MaxInstancesPerAgent <= 0 { + s.MaxInstancesPerAgent = d.MaxInstancesPerAgent + } + return s +} + +// SettingsFromConfig maps the CCF_AGENT_* config onto Settings (nil config or nil Agents => +// defaults). +func SettingsFromConfig(c *config.Config) Settings { + if c == nil || c.Agents == nil { + return Settings{}.WithDefaults() + } + cfg := c.Agents + return Settings{ + InstanceStaleAfter: cfg.InstanceStaleAfter, + InstanceRetention: cfg.InstanceRetention, + OneShotInstanceRetention: cfg.OneShotInstanceRetention, + MaxInstancesPerAgent: cfg.MaxInstancesPerAgent, + }.WithDefaults() +} + +var ( + // ErrNotFound is returned when the agent, revision or instance does not exist. + ErrNotFound = errors.New("not found") + // ErrRevisionConflict is returned (wrapped in *RevisionConflictError) when the expected + // revision is not the current one. + ErrRevisionConflict = errors.New("configuration revision conflict") +) + +// RevisionConflictError carries the current revision of a failed CreateRevision. +type RevisionConflictError struct { + Current int64 +} + +func (e *RevisionConflictError) Error() string { + return fmt.Sprintf("%s: current revision is %d", ErrRevisionConflict.Error(), e.Current) +} + +// Is makes errors.Is(err, ErrRevisionConflict) true. +func (e *RevisionConflictError) Is(target error) bool { return target == ErrRevisionConflict } + +// Service is the agent remote-configuration store. +type Service struct { + db *gorm.DB + settings Settings + logger *zap.SugaredLogger + now func() time.Time +} + +// NewService builds a Service. Zero settings take the defaults; a nil logger is a no-op. +func NewService(db *gorm.DB, s Settings, logger *zap.SugaredLogger) *Service { + if logger == nil { + logger = zap.NewNop().Sugar() + } + return &Service{ + db: db, + settings: s.WithDefaults(), + logger: logger, + now: func() time.Time { return time.Now().UTC() }, + } +} + +// Settings returns the effective settings. +func (s *Service) Settings() Settings { return s.settings } + +// Now returns the service clock (UTC). +func (s *Service) Now() time.Time { return s.now() } + +// SetClock overrides the service clock (tests). +func (s *Service) SetClock(now func() time.Time) { s.now = now } + +// ---- Revisions ---- + +// Current returns the latest revision of an agent, or nil (revision 0) when none exists. +func (s *Service) Current(ctx context.Context, agentID uuid.UUID) (*relational.AgentConfigRevision, error) { + var rev relational.AgentConfigRevision + err := s.db.WithContext(ctx). + Where("agent_id = ?", agentID). + Order("revision DESC"). + Limit(1). + Take(&rev).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &rev, nil +} + +// CurrentRevisionNumber returns the latest revision number, 0 when none exists. +func (s *Service) CurrentRevisionNumber(ctx context.Context, agentID uuid.UUID) (int64, error) { + var cur int64 + err := s.db.WithContext(ctx).Model(&relational.AgentConfigRevision{}). + Where("agent_id = ?", agentID). + Select("COALESCE(MAX(revision), 0)"). + Scan(&cur).Error + return cur, err +} + +// GetRevision returns one revision or ErrNotFound. +func (s *Service) GetRevision(ctx context.Context, agentID uuid.UUID, rev int64) (*relational.AgentConfigRevision, error) { + var out relational.AgentConfigRevision + err := s.db.WithContext(ctx).Where("agent_id = ? AND revision = ?", agentID, rev).Take(&out).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrNotFound + } + if err != nil { + return nil, err + } + return &out, nil +} + +// RevisionMeta is a revision without its overlay (list views, R12). +type RevisionMeta struct { + ID uuid.UUID + AgentID uuid.UUID + Revision int64 + CreatedAt time.Time + Comment *string + CreatedBy string + CreatedByID *uuid.UUID + RevertOf *int64 + OverlaySize int // bytes of the stored overlay's text form +} + +// ListRevisions returns one page of revisions, newest first, and the total count. +func (s *Service) ListRevisions(ctx context.Context, agentID uuid.UUID, p service.PaginationParams) ([]RevisionMeta, int64, error) { + db := s.db.WithContext(ctx) + var total int64 + if err := db.Model(&relational.AgentConfigRevision{}).Where("agent_id = ?", agentID).Count(&total).Error; err != nil { + return nil, 0, err + } + var rows []RevisionMeta + err := db.Model(&relational.AgentConfigRevision{}). + Select("id, agent_id, revision, created_at, comment, created_by, created_by_id, revert_of, octet_length(overlay::text) AS overlay_size"). + Where("agent_id = ?", agentID). + Order("revision DESC"). + Limit(p.Limit). + Offset(p.Offset). + Scan(&rows).Error + if err != nil { + return nil, 0, err + } + return rows, total, nil +} + +// CreateRevisionParams are the inputs of CreateRevision. +type CreateRevisionParams struct { + AgentID uuid.UUID + ExpectedRevision int64 + Overlay json.RawMessage + Comment *string + CreatedBy string + CreatedByID *uuid.UUID + RevertOf *int64 +} + +// CreateRevision appends revision ExpectedRevision+1 in one transaction: it locks the agent +// row, re-reads the current revision and fails with *RevisionConflictError when it is not +// ExpectedRevision. The agent row lock serializes writers, so the insert cannot hit the +// unique (agent_id, revision) index; any insert error is returned as is. The returned row +// is read back after the insert, so its Overlay bytes are the stored form GetRevision and +// ListRevisions (overlay size) see. +func (s *Service) CreateRevision(ctx context.Context, p CreateRevisionParams) (*relational.AgentConfigRevision, error) { + var created *relational.AgentConfigRevision + err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var agent relational.Agent + if err := tx.Clauses(clause.Locking{Strength: clause.LockingStrengthUpdate}). + Select("id"). + Where("id = ?", p.AgentID). + Take(&agent).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return ErrNotFound + } + return err + } + var cur int64 + if err := tx.Model(&relational.AgentConfigRevision{}). + Where("agent_id = ?", p.AgentID). + Select("COALESCE(MAX(revision), 0)"). + Scan(&cur).Error; err != nil { + return err + } + if cur != p.ExpectedRevision { + return &RevisionConflictError{Current: cur} + } + rev := &relational.AgentConfigRevision{ + AgentID: p.AgentID, + Revision: cur + 1, + Overlay: datatypes.JSON(p.Overlay), + Comment: p.Comment, + CreatedBy: p.CreatedBy, + CreatedByID: p.CreatedByID, + RevertOf: p.RevertOf, + CreatedAt: s.now(), + } + if err := tx.Create(rev).Error; err != nil { + return err + } + var stored relational.AgentConfigRevision + if err := tx.Where("agent_id = ? AND revision = ?", p.AgentID, rev.Revision).Take(&stored).Error; err != nil { + return err + } + created = &stored + return nil + }) + if err != nil { + return nil, err + } + return created, nil +} + +// DeleteRevisionsForAgent removes every configuration revision of an agent (agent deletion). +// It is the purge path for an overlay that held a secret, so it bypasses the append-only +// BeforeDelete hook, which still blocks every other delete of a revision. +func DeleteRevisionsForAgent(tx *gorm.DB, agentID uuid.UUID) error { + return tx.Session(&gorm.Session{SkipHooks: true}). + Where("agent_id = ?", agentID). + Delete(&relational.AgentConfigRevision{}).Error +} diff --git a/internal/service/relational/agentcfg/service_integration_test.go b/internal/service/relational/agentcfg/service_integration_test.go new file mode 100644 index 00000000..171a89e4 --- /dev/null +++ b/internal/service/relational/agentcfg/service_integration_test.go @@ -0,0 +1,296 @@ +//go:build integration + +package agentcfg_test + +import ( + "context" + "encoding/json" + "errors" + "sync" + "testing" + "time" + + "github.com/compliance-framework/api/internal/service" + "github.com/compliance-framework/api/internal/service/relational" + "github.com/compliance-framework/api/internal/service/relational/agentcfg" + "github.com/compliance-framework/api/internal/tests" + "github.com/google/uuid" + "github.com/stretchr/testify/suite" + "gorm.io/datatypes" + "gorm.io/gorm" +) + +type AgentCfgServiceIntegrationSuite struct { + tests.IntegrationTestSuite + ctx context.Context + now time.Time + svc *agentcfg.Service +} + +func TestAgentCfgServiceIntegration(t *testing.T) { + suite.Run(t, new(AgentCfgServiceIntegrationSuite)) +} + +func (s *AgentCfgServiceIntegrationSuite) SetupTest() { + s.Require().NoError(s.Migrator.Refresh()) + s.ctx = context.Background() + s.now = time.Date(2026, 3, 1, 12, 0, 0, 0, time.UTC) + s.svc = s.newService(agentcfg.Settings{}) +} + +// newService builds a service whose clock follows s.now. +func (s *AgentCfgServiceIntegrationSuite) newService(settings agentcfg.Settings) *agentcfg.Service { + svc := agentcfg.NewService(s.DB, settings, nil) + svc.SetClock(func() time.Time { return s.now }) + return svc +} + +func (s *AgentCfgServiceIntegrationSuite) newAgent(name string) uuid.UUID { + agent, err := s.CreateAgent(name) + s.Require().NoError(err) + return *agent.ID +} + +func (s *AgentCfgServiceIntegrationSuite) createRevision(agentID uuid.UUID, expected int64, overlay string) *relational.AgentConfigRevision { + rev, err := s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: agentID, + ExpectedRevision: expected, + Overlay: json.RawMessage(overlay), + CreatedBy: "alice@example.com", + }) + s.Require().NoError(err) + return rev +} + +func ptr[T any](v T) *T { return &v } + +// ---- Revisions ---- + +func (s *AgentCfgServiceIntegrationSuite) TestRevisionHooksAreAppendOnly() { + agentID := s.newAgent("hooks") + rev := s.createRevision(agentID, 0, `{"verbosity":1}`) + + err := s.DB.Model(rev).Update("comment", "changed").Error + s.ErrorIs(err, relational.ErrAgentConfigRevisionAppendOnly) + + rev.CreatedBy = "mallory@example.com" + s.ErrorIs(s.DB.Save(rev).Error, relational.ErrAgentConfigRevisionAppendOnly) + + s.ErrorIs(s.DB.Delete(rev).Error, relational.ErrAgentConfigRevisionAppendOnly) + s.ErrorIs(s.DB.Where("agent_id = ?", agentID).Delete(&relational.AgentConfigRevision{}).Error, + relational.ErrAgentConfigRevisionAppendOnly) + + stored, err := s.svc.GetRevision(s.ctx, agentID, 1) + s.Require().NoError(err) + s.Nil(stored.Comment) + s.Equal("alice@example.com", stored.CreatedBy) +} + +func (s *AgentCfgServiceIntegrationSuite) TestRevisionUniqueIndex() { + agentA := s.newAgent("unique-a") + agentB := s.newAgent("unique-b") + mk := func(agentID uuid.UUID, rev int64) *relational.AgentConfigRevision { + return &relational.AgentConfigRevision{ + AgentID: agentID, Revision: rev, Overlay: datatypes.JSON(`{}`), CreatedBy: "a", CreatedAt: s.now, + } + } + s.Require().NoError(s.DB.Create(mk(agentA, 1)).Error) + s.Error(s.DB.Create(mk(agentA, 1)).Error, "(agent_id, revision) must be unique") + s.NoError(s.DB.Create(mk(agentB, 1)).Error, "the same revision number is fine for another agent") + s.NoError(s.DB.Create(mk(agentA, 2)).Error) +} + +func (s *AgentCfgServiceIntegrationSuite) TestCreateRevisionSequenceAndConflicts() { + agentID := s.newAgent("seq") + + cur, err := s.svc.Current(s.ctx, agentID) + s.Require().NoError(err) + s.Nil(cur, "no revision yet means revision 0") + n, err := s.svc.CurrentRevisionNumber(s.ctx, agentID) + s.Require().NoError(err) + s.Equal(int64(0), n) + + userID := uuid.New() + rev1, err := s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: agentID, ExpectedRevision: 0, Overlay: json.RawMessage(`{"verbosity":1}`), + Comment: ptr("first"), CreatedBy: "alice@example.com", CreatedByID: &userID, + }) + s.Require().NoError(err) + s.Equal(int64(1), rev1.Revision) + s.NotNil(rev1.ID) + s.True(rev1.CreatedAt.Equal(s.now)) + + s.now = s.now.Add(time.Minute) + rev2, err := s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: agentID, ExpectedRevision: 1, Overlay: json.RawMessage(`{"verbosity":2}`), + CreatedBy: "bob@example.com", RevertOf: ptr(int64(1)), + }) + s.Require().NoError(err) + s.Equal(int64(2), rev2.Revision) + s.NotEqual(*rev1.ID, *rev2.ID) + + for _, expected := range []int64{0, 1, 5} { + _, err = s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: agentID, ExpectedRevision: expected, Overlay: json.RawMessage(`{}`), CreatedBy: "x", + }) + s.Require().Error(err) + s.True(errors.Is(err, agentcfg.ErrRevisionConflict), "expected %d", expected) + var conflict *agentcfg.RevisionConflictError + s.Require().True(errors.As(err, &conflict)) + s.Equal(int64(2), conflict.Current) + } + + _, err = s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: uuid.New(), ExpectedRevision: 0, Overlay: json.RawMessage(`{}`), CreatedBy: "x", + }) + s.ErrorIs(err, agentcfg.ErrNotFound) + + cur, err = s.svc.Current(s.ctx, agentID) + s.Require().NoError(err) + s.Require().NotNil(cur) + s.Equal(int64(2), cur.Revision) + s.JSONEq(`{"verbosity":2}`, string(cur.Overlay)) + s.Equal("bob@example.com", cur.CreatedBy) + s.Require().NotNil(cur.RevertOf) + s.Equal(int64(1), *cur.RevertOf) + + n, err = s.svc.CurrentRevisionNumber(s.ctx, agentID) + s.Require().NoError(err) + s.Equal(int64(2), n) + + got, err := s.svc.GetRevision(s.ctx, agentID, 1) + s.Require().NoError(err) + s.JSONEq(`{"verbosity":1}`, string(got.Overlay)) + s.Equal("first", *got.Comment) + s.Equal(userID, *got.CreatedByID) + s.True(got.CreatedAt.Equal(rev1.CreatedAt)) + + _, err = s.svc.GetRevision(s.ctx, agentID, 3) + s.ErrorIs(err, agentcfg.ErrNotFound) + other := s.newAgent("seq-other") + _, err = s.svc.GetRevision(s.ctx, other, 1) + s.ErrorIs(err, agentcfg.ErrNotFound) +} + +func (s *AgentCfgServiceIntegrationSuite) TestCreateRevisionConcurrentWritersExactlyOneWins() { + for round := 0; round < 5; round++ { + agentID := s.newAgent("concurrent") + start := make(chan struct{}) + errs := make([]error, 2) + var wg sync.WaitGroup + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + <-start + _, errs[i] = s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: agentID, ExpectedRevision: 0, Overlay: json.RawMessage(`{"verbosity":1}`), CreatedBy: "x", + }) + }(i) + } + close(start) + wg.Wait() + + var ok, conflicts int + for _, err := range errs { + switch { + case err == nil: + ok++ + case errors.Is(err, agentcfg.ErrRevisionConflict): + conflicts++ + var conflict *agentcfg.RevisionConflictError + s.Require().True(errors.As(err, &conflict)) + s.Equal(int64(1), conflict.Current) + default: + s.Failf("unexpected error", "%v", err) + } + } + s.Equal(1, ok, "round %d", round) + s.Equal(1, conflicts, "round %d", round) + n, err := s.svc.CurrentRevisionNumber(s.ctx, agentID) + s.Require().NoError(err) + s.Equal(int64(1), n) + } +} + +func (s *AgentCfgServiceIntegrationSuite) TestListRevisionsPagesNewestFirst() { + agentID := s.newAgent("list") + other := s.newAgent("list-other") + for i := int64(0); i < 5; i++ { + s.now = s.now.Add(time.Minute) + _, err := s.svc.CreateRevision(s.ctx, agentcfg.CreateRevisionParams{ + AgentID: agentID, ExpectedRevision: i, Overlay: json.RawMessage(`{"verbosity":1,"plugins":{"p":{"source":"x"}}}`), + Comment: ptr("c"), CreatedBy: "alice@example.com", RevertOf: ptr(int64(1)), + }) + s.Require().NoError(err) + } + s.createRevision(other, 0, `{}`) + + page, total, err := s.svc.ListRevisions(s.ctx, agentID, service.PaginationParams{Page: 1, Limit: 2, Offset: 0}) + s.Require().NoError(err) + s.Equal(int64(5), total) + s.Require().Len(page, 2) + s.Equal(int64(5), page[0].Revision) + s.Equal(int64(4), page[1].Revision) + for _, m := range page { + s.NotEqual(uuid.Nil, m.ID) + s.Equal(agentID, m.AgentID) + s.Greater(m.OverlaySize, 0) + s.Equal("alice@example.com", m.CreatedBy) + s.Require().NotNil(m.Comment) + s.Equal("c", *m.Comment) + s.Require().NotNil(m.RevertOf) + s.False(m.CreatedAt.IsZero()) + } + s.True(page[0].CreatedAt.After(page[1].CreatedAt)) + + page, total, err = s.svc.ListRevisions(s.ctx, agentID, service.PaginationParams{Page: 3, Limit: 2, Offset: 4}) + s.Require().NoError(err) + s.Equal(int64(5), total) + s.Require().Len(page, 1) + s.Equal(int64(1), page[0].Revision) + + page, total, err = s.svc.ListRevisions(s.ctx, s.newAgent("list-empty"), service.PaginationParams{Page: 1, Limit: 10}) + s.Require().NoError(err) + s.Equal(int64(0), total) + s.Empty(page) +} + +// ---- Instances ---- + +func (s *AgentCfgServiceIntegrationSuite) TestDeleteRevisionsForAgent() { + agentA := s.newAgent("purge-a") + agentB := s.newAgent("purge-b") + s.createRevision(agentA, 0, `{"plugins":{"p":{"config":{"password":"hunter2"}}}}`) + s.createRevision(agentA, 1, `{"verbosity":2}`) + s.createRevision(agentB, 0, `{"verbosity":1}`) + + s.Require().NoError(s.DB.Transaction(func(tx *gorm.DB) error { + return agentcfg.DeleteRevisionsForAgent(tx, agentA) + })) + + var n int64 + s.Require().NoError(s.DB.Model(&relational.AgentConfigRevision{}).Where("agent_id = ?", agentA).Count(&n).Error) + s.Zero(n, "every revision of the agent is purged") + cur, err := s.svc.CurrentRevisionNumber(s.ctx, agentB) + s.Require().NoError(err) + s.Equal(int64(1), cur, "other agents keep theirs") + + // Every other delete path stays append-only. + s.ErrorIs(s.DB.Where("agent_id = ?", agentB).Delete(&relational.AgentConfigRevision{}).Error, + relational.ErrAgentConfigRevisionAppendOnly) +} + +func (s *AgentCfgServiceIntegrationSuite) TestCreateRevisionReturnsTheStoredRow() { + agentID := s.newAgent("stored") + rev := s.createRevision(agentID, 0, `{"verbosity":1,"plugins":{"b":{"enabled":false},"a":{"enabled":true}}}`) + got, err := s.svc.GetRevision(s.ctx, agentID, rev.Revision) + s.Require().NoError(err) + s.Equal(string(got.Overlay), string(rev.Overlay), "same bytes as a later read") + s.Equal(*got.ID, *rev.ID) + + metas, _, err := s.svc.ListRevisions(s.ctx, agentID, service.PaginationParams{Page: 1, Limit: 10}) + s.Require().NoError(err) + s.Require().Len(metas, 1) + s.Equal(len(rev.Overlay), metas[0].OverlaySize) +} diff --git a/internal/tests/migrate.go b/internal/tests/migrate.go index 01449d38..b6152090 100644 --- a/internal/tests/migrate.go +++ b/internal/tests/migrate.go @@ -173,6 +173,7 @@ func (t *TestMigrator) Up() error { &relational.Agent{}, &relational.AgentServiceAccountKey{}, &relational.AgentAuthEvent{}, + &relational.AgentConfigRevision{}, &relational.SSOUserLink{}, &relational.SlackLinkAttempt{}, &relational.SlackUserLink{}, @@ -550,6 +551,7 @@ func (t *TestMigrator) Down() error { "poam_findings", "poam_risks", + &relational.AgentConfigRevision{}, &relational.AgentAuthEvent{}, &relational.AgentServiceAccountKey{}, &relational.Agent{},