package cli import ( "bytes" "context" "errors" "fmt" "iter" "os" "path/filepath" "testing" "git.cloudinit.dev/coreci/orca/internal/drift" ) type mockDriftTransport struct { execOut []byte execErr error readOut []byte readErr error writes []mockDriftWrite execFn func(ctx context.Context, peer string, cmd string) ([]byte, error) execs []string } type mockDriftWrite struct { peer string path string content []byte mode os.FileMode } func (m *mockDriftTransport) Exec(ctx context.Context, peer string, cmd string) ([]byte, error) { if m.execFn != nil { return m.execFn(ctx, peer, cmd) } m.execs = append(m.execs, cmd) return m.execOut, m.execErr } func (m *mockDriftTransport) WriteFileIdempotent(ctx context.Context, peer string, path string, content []byte, mode os.FileMode) (bool, error) { m.writes = append(m.writes, mockDriftWrite{peer, path, content, mode}) return true, nil } func (m *mockDriftTransport) ReadFile(ctx context.Context, peer string, path string) ([]byte, error) { return m.readOut, m.readErr } // fakeDetector is a record-replay drift.Detector for CLI tests. type fakeDetector struct { aggEvents []drift.Event aggErr error remediateErr error ackWrites int remediateCalls []remediateCall } type remediateCall struct { peer string path string force bool } func (f *fakeDetector) Watch(ctx context.Context, paths []drift.PathSpec) iter.Seq2[drift.Event, error] { return func(yield func(drift.Event, error) bool) { for _, e := range f.aggEvents { if !yield(e, nil) { return } } <-ctx.Done() } } func (f *fakeDetector) Aggregate(ctx context.Context, leadPeer string) ([]drift.Event, error) { return f.aggEvents, f.aggErr } func (f *fakeDetector) Remediate(ctx context.Context, leadPeer, path string, force bool) error { f.remediateCalls = append(f.remediateCalls, remediateCall{leadPeer, path, force}) return f.remediateErr } func (f *fakeDetector) Acknowledge(ctx context.Context, leadPeer, path string) error { f.ackWrites++ return nil } func setupDriftCLITest(t *testing.T) string { t.Helper() home := t.TempDir() t.Setenv("ORCA_HOME", home) t.Setenv("ORCA_LEAD_STATE_DIR", filepath.Join(home, "state")) return home } func TestDriftCmdRegistered(t *testing.T) { for _, c := range rootCmd.Commands() { if c.Name() == "drift" { return } } t.Fatal("drift command not registered on root") } func TestDriftSubcommandsRegistered(t *testing.T) { for _, c := range rootCmd.Commands() { if c.Name() != "drift" { continue } want := map[string]bool{ "show": false, "watch": false, "acknowledge": false, "remediate": false, "config": false, } for _, sub := range c.Commands() { if _, ok := want[sub.Name()]; ok { want[sub.Name()] = true } } for name, found := range want { if !found { t.Errorf("drift subcommand %q not registered", name) } } return } t.Fatal("drift command not registered") } func TestDriftConfigSubcommands(t *testing.T) { for _, c := range rootCmd.Commands() { if c.Name() != "drift" { continue } for _, sub := range c.Commands() { if sub.Name() != "config" { continue } want := map[string]bool{"show": false, "validate": false} for _, s := range sub.Commands() { if _, ok := want[s.Name()]; ok { want[s.Name()] = true } } for name, found := range want { if !found { t.Errorf("drift config subcommand %q not registered", name) } } return } } t.Fatal("drift config not registered") } func TestJobRestartRegistered(t *testing.T) { for _, c := range rootCmd.Commands() { if c.Name() != "job" { continue } for _, sub := range c.Commands() { if sub.Name() == "restart" { return } } } t.Fatal("job restart not registered") } func TestDriftShowEmpty(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{aggEvents: nil} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "show"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift show: %v", err) } if !bytesContains(buf.String(), "No drift events") { t.Errorf("empty show output: %s", buf.String()) } } func TestDriftShowTable(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{aggEvents: []drift.Event{ {EventID: "E1", Host: "peer1", Path: "/etc/traefik/dynamic/orca.yml", Status: drift.StatusModified, DriftConfirmed: true}, }} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "show"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift show: %v", err) } out := buf.String() for _, want := range []string{"E1", "peer1", "modified"} { if !bytesContains(out, want) { t.Errorf("output missing %q: %s", want, out) } } } func TestDriftShowJSON(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{aggEvents: []drift.Event{ {EventID: "E1", Host: "peer1", Path: "/etc/x", Status: drift.StatusCreated, DriftConfirmed: true}, }} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) _ = rootCmd.PersistentFlags().Set("json", "true") var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "show"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift show --json: %v", err) } if !bytesContains(buf.String(), `"event_id"`) { t.Errorf("json output missing event_id: %s", buf.String()) } } func TestDriftShowPeerFilter(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{aggEvents: []drift.Event{ {EventID: "E1", Host: "peer1", Path: "/a", Status: drift.StatusModified, DriftConfirmed: true}, {EventID: "E2", Host: "peer2", Path: "/b", Status: drift.StatusModified, DriftConfirmed: true}, }} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "show", "--peer", "peer1"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift show --peer: %v", err) } out := buf.String() if !bytesContains(out, "E1") { t.Errorf("filtered output should have E1: %s", out) } if bytesContains(out, "E2") { t.Errorf("filtered output should NOT have E2: %s", out) } } func TestDriftWatchStreamsAndCancels(t *testing.T) { setupDriftCLITest(t) // Use a fakeDetector whose Watch yields one event then blocks on // ctx so the stream terminates when ctrl-c (signal.NotifyContext) // cancels. We simulate the cancel by constructing a fake that yields // then returns when the consumer stops pulling OR ctx is cancelled. fd := &drainingFakeDetector{events: []drift.Event{ {EventID: "E1", Host: "p", Path: "/etc/x", Status: drift.StatusModified, DriftConfirmed: true}, }} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "watch", "--interval", "10ms"}) // Inject a context that auto-cancels after the events drain so // the watch loop exits without polluting rootCmd's context (which // is shared across tests). We use PersistentPreRunE's context by // overriding it here and restoring after. origCtx := rootCmd.Context() ctx, cancel := context.WithCancel(origCtx) defer cancel() rootCmd.SetContext(ctx) fd.cancelAfterYield = cancel if err := rootCmd.Execute(); err != nil { t.Fatalf("drift watch: %v", err) } // Restore rootCmd context for subsequent tests. rootCmd.SetContext(origCtx) if !bytesContains(buf.String(), "E1") { t.Errorf("watch output missing E1: %s", buf.String()) } } // drainingFakeDetector yields the events then cancels the provided // cancel func (so the watch loop's signal.NotifyContext ctx is // cancelled and the stream terminates cleanly). type drainingFakeDetector struct { events []drift.Event cancelAfterYield context.CancelFunc remediateCalls []remediateCall ackWrites int } func (d *drainingFakeDetector) Watch(ctx context.Context, paths []drift.PathSpec) iter.Seq2[drift.Event, error] { return func(yield func(drift.Event, error) bool) { for _, e := range d.events { if !yield(e, nil) { return } } if d.cancelAfterYield != nil { d.cancelAfterYield() } <-ctx.Done() } } func (d *drainingFakeDetector) Aggregate(ctx context.Context, leadPeer string) ([]drift.Event, error) { return d.events, nil } func (d *drainingFakeDetector) Remediate(ctx context.Context, leadPeer, path string, force bool) error { d.remediateCalls = append(d.remediateCalls, remediateCall{leadPeer, path, force}) return nil } func (d *drainingFakeDetector) Acknowledge(ctx context.Context, leadPeer, path string) error { d.ackWrites++ return nil } func TestDriftAcknowledge(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "acknowledge", "peer1", "/etc/traefik/dynamic/orca.yml"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift acknowledge: %v", err) } if fd.ackWrites != 1 { t.Errorf("ackWrites = %d, want 1", fd.ackWrites) } if !bytesContains(buf.String(), "Acknowledged") { t.Errorf("output missing Acknowledged: %s", buf.String()) } } func TestDriftRemediate(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "remediate", "peer1", "/etc/x"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift remediate: %v", err) } if len(fd.remediateCalls) != 1 { t.Fatalf("remediateCalls = %d, want 1", len(fd.remediateCalls)) } if fd.remediateCalls[0].peer != "peer1" || fd.remediateCalls[0].path != "/etc/x" { t.Errorf("remediate call: %+v", fd.remediateCalls[0]) } if fd.remediateCalls[0].force { t.Errorf("force should be false without --force") } if !bytesContains(buf.String(), "Remediated") { t.Errorf("output missing Remediated: %s", buf.String()) } } func TestDriftRemediateForce(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "remediate", "peer1", "/etc/x", "--force"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift remediate --force: %v", err) } if len(fd.remediateCalls) != 1 { t.Fatalf("remediateCalls = %d, want 1", len(fd.remediateCalls)) } if !fd.remediateCalls[0].force { t.Errorf("force should be true with --force") } } func TestDriftRemediateCooldown(t *testing.T) { setupDriftCLITest(t) fd := &fakeDetector{remediateErr: drift.ErrCooldown} driftDetectorOverride = fd defer func() { driftDetectorOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "remediate", "peer1", "/etc/x"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift remediate cooldown: %v", err) } if !bytesContains(buf.String(), "cooldown") { t.Errorf("output should mention cooldown: %s", buf.String()) } } func TestDriftConfigShow(t *testing.T) { setupDriftCLITest(t) resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "config", "show"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift config show: %v", err) } out := buf.String() for _, want := range []string{"Polling:", "Critical paths:", "Standard paths:", "Excluded paths:", "Remediation:"} { if !bytesContains(out, want) { t.Errorf("config show missing %q: %s", want, out) } } } func TestDriftConfigValidate(t *testing.T) { setupDriftCLITest(t) resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "config", "validate"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("drift config validate: %v", err) } if !bytesContains(buf.String(), "valid") { t.Errorf("output missing 'valid': %s", buf.String()) } } func TestDriftConfigValidateFails(t *testing.T) { setupDriftCLITest(t) cfgPath := filepath.Join(t.TempDir(), "drift.json") bad := `{"polling":{"enabled":true,"default_interval":60000000000,"max_concurrent_peers":4},"paths":{"critical":[{"tier":"critical","pattern":"","interval":5000000000}]},"remediate":{"auto":true}}` if err := os.WriteFile(cfgPath, []byte(bad), 0o644); err != nil { t.Fatalf("write: %v", err) } resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"drift", "config", "validate", "--config", cfgPath}) err := rootCmd.Execute() if err == nil { t.Fatal("expected validate error") } } func TestJobRestartRequiresPeer(t *testing.T) { setupDriftCLITest(t) resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"job", "restart", "alloc1"}) err := rootCmd.Execute() if err == nil { t.Fatal("expected error for missing --peer") } } func TestJobRestartExecs(t *testing.T) { setupDriftCLITest(t) mt := &mockDriftTransport{execOut: []byte("restarted")} driftTransportOverride = mt defer func() { driftTransportOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"job", "restart", "alloc1", "--peer", "peer1:22"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("job restart: %v", err) } if len(mt.execs) != 1 { t.Fatalf("execs = %d, want 1", len(mt.execs)) } if !bytesContains(mt.execs[0], "orca-alloc-alloc1.service") { t.Errorf("restart cmd missing: %s", mt.execs[0]) } } func TestJobRestartTransientError(t *testing.T) { setupDriftCLITest(t) mt := &mockDriftTransport{execErr: fmt.Errorf("connection refused")} driftTransportOverride = mt defer func() { driftTransportOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"job", "restart", "alloc1", "--peer", "peer1:22"}) err := rootCmd.Execute() if err == nil { t.Fatal("expected error for exec failure") } if !errors.Is(err, err) { } } func TestPeerSetupCmdRegistered(t *testing.T) { for _, c := range rootCmd.Commands() { if c.Name() == "peer-setup" { return } } t.Fatal("peer-setup command not registered") } func TestPeerSetupCreatesUserAndDir(t *testing.T) { setupDriftCLITest(t) mt := &mockDriftTransport{execOut: []byte("ext4")} peerSetupTransportOverride = mt defer func() { peerSetupTransportOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"peer-setup", "peer1:22"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("peer-setup: %v", err) } if len(mt.execs) < 2 { t.Fatalf("execs = %d, want >= 2", len(mt.execs)) } useraddSeen := false mkdirSeen := false statSeen := false for _, c := range mt.execs { if bytesContains(c, "useradd -r orca") { useraddSeen = true } if bytesContains(c, "mkdir -p /etc/orca/state/drift-events") { mkdirSeen = true } if bytesContains(c, "stat -f") { statSeen = true } } if !useraddSeen { t.Errorf("useradd not run") } if !mkdirSeen { t.Errorf("mkdir drift-events not run") } if !statSeen { t.Errorf("stat (NFS detect) not run") } if !bytesContains(buf.String(), "nfs=false") { t.Errorf("output should report nfs=false: %s", buf.String()) } } func TestPeerSetupNoOrcaUser(t *testing.T) { setupDriftCLITest(t) mt := &mockDriftTransport{execOut: []byte("ext4")} peerSetupTransportOverride = mt defer func() { peerSetupTransportOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"peer-setup", "peer1:22", "--no-orca-user"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("peer-setup --no-orca-user: %v", err) } useraddSeen := false for _, c := range mt.execs { if bytesContains(c, "useradd -r orca") { useraddSeen = true } } if useraddSeen { t.Errorf("useradd should NOT run with --no-orca-user") } if !bytesContains(buf.String(), "user=false") { t.Errorf("output should report user=false: %s", buf.String()) } } func TestPeerSetupDetectsNFS(t *testing.T) { setupDriftCLITest(t) mt := &mockDriftTransport{execOut: []byte("nfs4")} peerSetupTransportOverride = mt defer func() { peerSetupTransportOverride = nil }() resetRootFlags(t) var buf bytes.Buffer rootCmd.SetOut(&buf) rootCmd.SetErr(&buf) rootCmd.SetArgs([]string{"peer-setup", "peer1:22"}) if err := rootCmd.Execute(); err != nil { t.Fatalf("peer-setup: %v", err) } if !bytesContains(buf.String(), "nfs=true") { t.Errorf("output should report nfs=true: %s", buf.String()) } }