diff --git a/integration/agentcompat/internal/agent/adversarial_test.go b/integration/agentcompat/internal/agent/adversarial_test.go new file mode 100644 index 00000000..e487b321 --- /dev/null +++ b/integration/agentcompat/internal/agent/adversarial_test.go @@ -0,0 +1,59 @@ +//go:build linux && agentcompat + +package agent + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestAgent_RejectsMalformedStartWithoutWorkspaceArtifact(t *testing.T) { + // Given + parent := t.TempDir() + t.Setenv("TMPDIR", parent) + + // When + _, err := Start(t.Context(), AgentStartConfig{SourceDir: filepath.Join(parent, "missing"), Endpoint: "127.0.0.1:1", UUID: "bad"}) + + // Then + require.Error(t, err) + entries, readErr := os.ReadDir(parent) + require.NoError(t, readErr) + require.Empty(t, entries) +} + +func TestAgent_ContextInterruptionCleansWorkspace(t *testing.T) { + // Given + processContext, interrupt := context.WithCancel(t.Context()) + dashboardInstance := startTestDashboard(t, false) + agentInstance, err := Start(processContext, AgentStartConfig{ + SourceDir: testAgentSourceDir(t), + Endpoint: dashboardInstance.Endpoint(), + Secret: dashboardInstance.AgentSecret(), + UUID: "00000000-0000-0000-0000-000000000085", + }) + require.NoError(t, err) + t.Cleanup(func() { + cleanupContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, agentInstance.Stop(cleanupContext)) + }) + root := agentInstance.WorkspaceRoot() + + // When + interrupt() + + // Then + select { + case <-agentInstance.CleanupDone(): + case <-time.After(30 * time.Second): + t.Fatal("agent cleanup did not complete") + } + _, err = os.Stat(root) + require.ErrorIs(t, err, os.ErrNotExist) +} diff --git a/integration/agentcompat/internal/agent/agent.go b/integration/agentcompat/internal/agent/agent.go new file mode 100644 index 00000000..d97e27ce --- /dev/null +++ b/integration/agentcompat/internal/agent/agent.go @@ -0,0 +1,239 @@ +//go:build linux + +package agent + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "syscall" + "time" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard" + processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process" + "github.com/nezhahq/nezha/integration/agentcompat/internal/workspace" +) + +const ( + agentSecret = "0123456789abcdef0123456789abcdef" + agentMaxLogBytes = 1 << 20 + agentStopTimeout = 5 * time.Second + agentKillTimeout = 5 * time.Second +) + +type AgentStartConfig struct { + SourceDir string + PreparedBinary *PreparedBinary + Endpoint string + Secret string + UUID string + TLS bool + Debug bool + CAFilePath string + FMObserverRunID string + Credential *syscall.Credential + newSupervisor func(context.Context, processharness.Spec) *processharness.Supervisor + trackPID func(int) error + trackProcessGroup func(int) error +} + +type AgentStartError struct { + cause error + agent *Agent +} + +func (err *AgentStartError) Error() string { return err.cause.Error() } +func (err *AgentStartError) Unwrap() error { return err.cause } +func (err *AgentStartError) Finalize(ctx context.Context) error { + return err.agent.Stop(ctx) +} + +type Agent struct { + workspace *workspace.Workspace + supervisor *processharness.Supervisor + clients dashboard.Clients + configPath string + logPath string + binaryPath string + caFilePath string + environment []string + secret string + uuid string + releaseBinary func() + releasePending bool + cleanupOnce sync.Once + cleanupDone chan struct{} + cleanupAttemptMu sync.Mutex + cleanupMu sync.Mutex + cleanupErr error + readinessMu sync.Mutex + lastStateReport time.Time + fmObserver *FMProducerObserver + fmObserverPath string + startConfig AgentStartConfig + processMu sync.Mutex + currentProcess *processGeneration + processes []*processGeneration + generation uint64 + closed bool + trackPID func(int) error + trackProcessGroup func(int) error +} + +func Start(ctx context.Context, config AgentStartConfig) (*Agent, error) { + if config.PreparedBinary == nil { + if err := validateSourceDir(config.SourceDir); err != nil { + return nil, err + } + } + if config.Endpoint == "" || config.UUID == "" { + return nil, errors.New("agent endpoint and UUID are required") + } + if config.Secret == "" { + config.Secret = agentSecret + } + workspaceRoot, err := workspace.New(context.WithoutCancel(ctx)) + if err != nil { + return nil, fmt.Errorf("create agent workspace: %w", err) + } + trackPID := workspaceRoot.TrackPID + if config.trackPID != nil { + trackPID = config.trackPID + } + trackProcessGroup := workspaceRoot.TrackProcessGroup + if config.trackProcessGroup != nil { + trackProcessGroup = config.trackProcessGroup + } + agent := &Agent{workspace: workspaceRoot, secret: config.Secret, uuid: config.UUID, cleanupDone: make(chan struct{}), startConfig: config, trackPID: trackPID, trackProcessGroup: trackProcessGroup} + if err := agent.prepareFixture(ctx, config); err != nil { + return nil, cleanupFailedStart(ctx, agent, err) + } + if _, err := agent.StartProcess(ctx); err != nil { + return nil, cleanupFailedStart(ctx, agent, err) + } + go agent.cleanupOnCancellation(ctx) + return agent, nil +} + +func validateSourceDir(sourceDir string) error { + if sourceDir == "" || !filepath.IsAbs(sourceDir) { + return errors.New("agent source directory must be absolute") + } + return nil +} + +func cleanupFailedStart(ctx context.Context, agent *Agent, cause error) error { + cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 15*time.Second) + defer cancel() + startError := errors.Join(cause, agent.Stop(cleanupContext)) + if agent.finalizationPending() { + // A failed rollback can leave the prepared binary leased until its process group exits. + return &AgentStartError{cause: startError, agent: agent} + } + return startError +} + +func filteredEnvironment() []string { + result := make([]string, 0, len(os.Environ())) + for _, value := range os.Environ() { + if strings.HasPrefix(value, "NZ_") || strings.HasPrefix(value, "SSL_CERT_FILE=") || strings.HasPrefix(value, "AGENTCOMPAT_FM_OBSERVER_") { + continue + } + result = append(result, value) + } + return result +} + +func (agent *Agent) Stop(ctx context.Context) error { + agent.cleanupOnce.Do(func() { go agent.cleanup(context.WithoutCancel(ctx)) }) + select { + case <-agent.cleanupDone: + agent.retryFinalization(ctx) + return agent.cleanupResult() + case <-ctx.Done(): + return ctx.Err() + } +} + +func (agent *Agent) cleanupOnCancellation(ctx context.Context) { + select { + case <-ctx.Done(): + agent.cleanupOnce.Do(func() { go agent.cleanup(context.WithoutCancel(ctx)) }) + case <-agent.cleanupDone: + } +} + +func (agent *Agent) cleanup(ctx context.Context) { + defer close(agent.cleanupDone) + agent.finishCleanup(ctx, true) +} + +func (agent *Agent) finishCleanup(ctx context.Context, closeObserver bool) { + agent.cleanupAttemptMu.Lock() + defer agent.cleanupAttemptMu.Unlock() + cleanupError := agent.closeProcesses(ctx) + if closeObserver && agent.fmObserver != nil { + cleanupError = errors.Join(cleanupError, agent.fmObserver.Close(), removeFMObserverSocket(agent.fmObserverPath)) + } + if err := agent.workspace.Close(); err != nil { + cleanupError = errors.Join(cleanupError, err) + } + // The prepared workspace must outlive every consumer process group, even when process cleanup reports an error. + if agent.releasePending && agent.processesQuiescent() { + agent.releaseBinary() + agent.releasePending = false + } + agent.cleanupMu.Lock() + if agent.cleanupErr == nil { + agent.cleanupErr = cleanupError + } + agent.cleanupMu.Unlock() +} + +func (agent *Agent) retryFinalization(ctx context.Context) { + if agent.finalizationPending() { + agent.finishCleanup(context.WithoutCancel(ctx), false) + } +} + +func (agent *Agent) finalizationPending() bool { + agent.cleanupAttemptMu.Lock() + defer agent.cleanupAttemptMu.Unlock() + return agent.releasePending +} + +func (agent *Agent) cleanupResult() error { + agent.cleanupMu.Lock() + defer agent.cleanupMu.Unlock() + return agent.cleanupErr +} + +func closeError(first, second error) error { return errors.Join(first, second) } +func (agent *Agent) UUID() string { return agent.uuid } +func (agent *Agent) PID() int { + agent.processMu.Lock() + defer agent.processMu.Unlock() + if agent.currentProcess == nil { + return 0 + } + return agent.currentProcess.identity.PID +} +func (agent *Agent) CleanupReceipt() processharness.CleanupReceipt { + agent.processMu.Lock() + defer agent.processMu.Unlock() + records := make([]processharness.CleanupRecord, 0, len(agent.processes)) + for _, process := range agent.processes { + records = append(records, process.record) + } + return processharness.NewCleanupReceipt(records) +} +func (agent *Agent) ConfigPath() string { return agent.configPath } +func (agent *Agent) BinaryPath() string { return agent.binaryPath } +func (agent *Agent) LogPath() string { return agent.logPath } +func (agent *Agent) WorkspaceRoot() string { return agent.workspace.Root() } +func (agent *Agent) CleanupDone() <-chan struct{} { return agent.cleanupDone } +func (agent *Agent) FMProducerObserver() *FMProducerObserver { return agent.fmObserver } diff --git a/integration/agentcompat/internal/agent/agent_test.go b/integration/agentcompat/internal/agent/agent_test.go new file mode 100644 index 00000000..562e3fe4 --- /dev/null +++ b/integration/agentcompat/internal/agent/agent_test.go @@ -0,0 +1,294 @@ +//go:build linux && agentcompat + +package agent + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/client" + "github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard" + "github.com/nezhahq/nezha/integration/agentcompat/internal/fixture" + "github.com/nezhahq/nezha/integration/agentcompat/internal/testpaths" +) + +func TestAgent_BecomesOnlineOverH2C(t *testing.T) { + // Given + dashboardInstance := startTestDashboardWithReceiptGate(t) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{UUID: "00000000-0000-0000-0000-000000000081"}) + receiptAccepted := make(chan error, 1) + go func() { receiptAccepted <- dashboardInstance.WaitForReceiptAccepted(t.Context()) }() + select { + case err := <-receiptAccepted: + require.NoError(t, err) + case <-time.After(10 * time.Second): + t.Fatal("timed out waiting for withheld state receipt") + } + serverBeforeRelease := requireOnlineServer(t, dashboardInstance, agentInstance.UUID()) + require.NotZero(t, serverBeforeRelease.LastActive) + stateGeneration := dashboardInstance.StateGeneration(serverBeforeRelease.ID, agentInstance.UUID()) + require.NotZero(t, stateGeneration) + require.NoError(t, dashboardInstance.WaitForStateGeneration(t.Context(), serverBeforeRelease.ID, agentInstance.UUID(), stateGeneration, 1)) + stateTwoBeforeRelease, cancelStateTwo := context.WithTimeout(t.Context(), 1500*time.Millisecond) + require.ErrorIs(t, dashboardInstance.WaitForStateGeneration(stateTwoBeforeRelease, serverBeforeRelease.ID, agentInstance.UUID(), stateGeneration, 2), context.DeadlineExceeded) + cancelStateTwo() + require.Equal(t, uint64(1), dashboardInstance.ReceiptAcceptedCount()) + require.NoError(t, dashboardInstance.ReleaseReceipt(t.Context())) + secondState := make(chan error, 1) + go func() { secondState <- dashboardInstance.WaitForSecondState(t.Context()) }() + + // When + readiness, err := agentInstance.WaitReady(t.Context(), dashboardInstance) + + // Then + require.NoError(t, err) + require.NotZero(t, readiness.ServerID) + require.Equal(t, serverBeforeRelease.ID, readiness.ServerID) + require.Equal(t, agentInstance.UUID(), readiness.UUID) + require.Equal(t, "v2.1.0", readiness.Version) + require.True(t, readiness.VersionObserved) + require.True(t, readiness.RequestTaskEstablished) + require.True(t, readiness.StateReceiptObserved) + require.NoError(t, <-secondState) + require.Equal(t, uint64(2), dashboardInstance.ReceiptAcceptedCount()) + require.NoError(t, dashboardInstance.WaitForInfo2(t.Context(), serverBeforeRelease.ID, agentInstance.UUID())) + require.NotNil(t, readiness.Host) + require.NotNil(t, readiness.State) + require.True(t, readiness.Online) +} + +func TestAgent_BecomesOnlineOverVerifiedTLS(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, true) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{ + UUID: "00000000-0000-0000-0000-000000000082", + TLS: true, + CAFilePath: dashboardInstance.TLSCACertificatePath(), + }) + + // When + readiness, err := agentInstance.WaitReady(t.Context(), dashboardInstance) + + // Then + require.NoError(t, err) + require.True(t, readiness.Online) + require.NotNil(t, readiness.Host) + require.NotNil(t, readiness.State) + require.WithinDuration(t, time.Now(), readiness.LastActive, 30*time.Second) + var host struct { + Platform string `json:"platform"` + Version string `json:"version"` + } + require.NoError(t, json.Unmarshal(readiness.Host, &host)) + require.Equal(t, "v2.1.0", host.Version) + require.NotEmpty(t, host.Platform) + var state map[string]json.RawMessage + require.NoError(t, json.Unmarshal(readiness.State, &state)) + require.Contains(t, state, "uptime") +} + +func TestAgent_RejectsUnknownTLSAuthority(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, true) + wrongCAPath := filepath.Join(t.TempDir(), "wrong-ca.crt") + wrongFixture, fixtureErr := fixture.NewLocalTLSFixture(time.Now()) + require.NoError(t, fixtureErr) + require.NoError(t, os.WriteFile(wrongCAPath, wrongFixture.CAPEM(), 0o600)) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{ + UUID: "00000000-0000-0000-0000-000000000086", + TLS: true, + Debug: true, + CAFilePath: wrongCAPath, + }) + + // When + readinessContext, cancel := context.WithTimeout(t.Context(), 8*time.Second) + defer cancel() + _, err := agentInstance.WaitReady(readinessContext, dashboardInstance) + + // Then + require.Error(t, err) + require.ErrorIs(t, err, context.DeadlineExceeded) + logData, logErr := os.ReadFile(agentInstance.LogPath()) + require.NoError(t, logErr) + require.Contains(t, string(logData), "x509: certificate signed by unknown authority") + config, readErr := os.ReadFile(agentInstance.ConfigPath()) + require.NoError(t, readErr) + stopContext, stopCancel := context.WithTimeout(context.Background(), 15*time.Second) + require.NoError(t, agentInstance.Stop(stopContext)) + stopCancel() + require.Contains(t, string(config), "insecure_tls: false") + require.NotContains(t, string(config), "insecure_tls: true") +} + +func TestAgent_AssertNeverOnlineFailsClosedWhenDashboardUnavailable(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, false) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{ + UUID: "00000000-0000-0000-0000-000000000087", + Secret: "wrong-agent-secret", + }) + require.NoError(t, dashboardInstance.Stop(context.Background())) + + // When + err := agentInstance.AssertNeverOnline(t.Context(), dashboardInstance, time.Second) + + // Then + require.Error(t, err) +} + +func TestAgent_RejectsInvalidSecret(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, false) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{ + UUID: "00000000-0000-0000-0000-000000000083", + Secret: "wrong-agent-secret", + }) + + // When + err := agentInstance.AssertNeverOnline(t.Context(), dashboardInstance, 2*time.Second) + + // Then + require.NoError(t, err) + stopContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, agentInstance.Stop(stopContext)) +} + +func TestAgent_StopsCleanly(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, false) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{UUID: "00000000-0000-0000-0000-000000000084"}) + require.NoError(t, waitForAgentReady(t, agentInstance, dashboardInstance)) + + // When + stopContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + err := agentInstance.Stop(stopContext) + + // Then + require.NoError(t, err) + require.True(t, agentInstance.CleanupReceipt().Passed) + require.False(t, agentInstance.CleanupReceipt().Forced) +} + +func TestAgent_StopRemovesAgentFromOnlineOnlyList(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, false) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{UUID: "00000000-0000-0000-0000-000000000088"}) + require.NoError(t, waitForAgentReady(t, agentInstance, dashboardInstance)) + stopContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, agentInstance.Stop(stopContext)) + + // When + list, err := client.CallTool[serverListArguments, serverListResult]( + t.Context(), dashboardInstance.Clients().MCP, + client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}}, + ) + + // Then + require.NoError(t, err) + foundOnline := false + for _, server := range list.StructuredContent.Servers { + if server.UUID == agentInstance.UUID() { + foundOnline = true + } + } + require.False(t, foundOnline) +} + +func startTestDashboard(t *testing.T, enableTLS bool) *dashboard.Dashboard { + t.Helper() + sourceDir, err := testpaths.NezhaSource(t.Name()) + require.NoError(t, err) + instance, err := dashboard.Start(t.Context(), dashboard.StartConfig{SourceDir: sourceDir, EnableTLS: enableTLS, ReadinessTimeout: readinessBudget}) + require.NoError(t, err) + t.Cleanup(func() { + cleanupContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, instance.Stop(cleanupContext)) + }) + return instance +} + +func startTestDashboardWithReceiptGate(t *testing.T) *dashboard.Dashboard { + t.Helper() + sourceDir, err := testpaths.NezhaSource(t.Name()) + require.NoError(t, err) + instance, err := dashboard.Start(t.Context(), dashboard.StartConfig{SourceDir: sourceDir, ReceiptGate: true, ReadinessTimeout: readinessBudget}) + require.NoError(t, err) + t.Cleanup(func() { + cleanupContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, instance.Stop(cleanupContext)) + }) + return instance +} + +func startTestAgent(t *testing.T, dashboardInstance *dashboard.Dashboard, config AgentStartConfig) *Agent { + t.Helper() + if config.Secret == "" { + config.Secret = dashboardInstance.AgentSecret() + } + if config.TLS { + config.Endpoint = dashboardInstance.TLSEndpoint() + } else { + config.Endpoint = dashboardInstance.Endpoint() + } + instance, err := Start(t.Context(), AgentStartConfig{ + SourceDir: testAgentSourceDir(t), + Endpoint: config.Endpoint, + Secret: config.Secret, + UUID: config.UUID, + TLS: config.TLS, + Debug: config.Debug, + CAFilePath: config.CAFilePath, + }) + require.NoError(t, err) + t.Cleanup(func() { + cleanupContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, instance.Stop(cleanupContext)) + }) + return instance +} + +func testAgentSourceDir(t *testing.T) string { + t.Helper() + if sourceDir := os.Getenv("AGENT_SOURCE"); sourceDir != "" { + return sourceDir + } + nezhaSource, err := testpaths.NezhaSource(t.Name()) + require.NoError(t, err) + agentSourceDir, err := testpaths.AgentSource(nezhaSource) + require.NoError(t, err) + return agentSourceDir +} + +func requireOnlineServer(t *testing.T, dashboardInstance *dashboard.Dashboard, uuid string) serverListItem { + t.Helper() + list, err := client.CallTool[serverListArguments, serverListResult](t.Context(), dashboardInstance.Clients().MCP, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}}) + require.NoError(t, err) + for _, server := range list.StructuredContent.Servers { + if server.UUID == uuid { + return server + } + } + t.Fatalf("server %q is not online", uuid) + return serverListItem{} +} + +func waitForAgentReady(t *testing.T, instance *Agent, dashboardInstance *dashboard.Dashboard) error { + t.Helper() + readinessContext, cancel := context.WithTimeout(t.Context(), readinessBudget) + defer cancel() + _, err := instance.WaitReady(readinessContext, dashboardInstance) + return err +} diff --git a/integration/agentcompat/internal/agent/failed_start_recovery_test.go b/integration/agentcompat/internal/agent/failed_start_recovery_test.go new file mode 100644 index 00000000..3ac1ca16 --- /dev/null +++ b/integration/agentcompat/internal/agent/failed_start_recovery_test.go @@ -0,0 +1,68 @@ +//go:build linux && agentcompat + +package agent + +import ( + "context" + "errors" + "syscall" + "testing" + "time" + + processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process" + "github.com/stretchr/testify/require" +) + +func TestAgent_StartExposesFinalizerWhenPreparedConsumerSurvivesRollback(t *testing.T) { + prepared, err := PrepareBinary(t.Context(), testAgentSourceDir(t)) + require.NoError(t, err) + trackingErr := errors.New("injected failed-start PID tracking error") + var supervisor *processharness.Supervisor + + instance, startErr := Start(t.Context(), AgentStartConfig{ + PreparedBinary: prepared, + Endpoint: "127.0.0.1:1", + UUID: "00000000-0000-0000-0000-000000000200", + newSupervisor: func(ctx context.Context, spec processharness.Spec) *processharness.Supervisor { + supervisor = processharness.NewSupervisor(ctx, spec) + cancelledContext, cancel := context.WithCancel(ctx) + cancel() + _ = supervisor.Stop(cancelledContext) + require.NoError(t, supervisor.Stop(t.Context())) + return supervisor + }, + trackPID: func(int) error { return trackingErr }, + }) + require.Nil(t, instance) + require.ErrorIs(t, startErr, trackingErr) + require.NotNil(t, supervisor) + pid := supervisor.PID() + processGroupID := supervisor.ProcessGroupID() + t.Cleanup(func() { + _ = syscall.Kill(-processGroupID, syscall.SIGKILL) + select { + case <-supervisor.Exited(): + case <-time.After(5 * time.Second): + } + _ = prepared.Close() + }) + require.NoError(t, syscall.Kill(-processGroupID, 0)) + + var startFailure *AgentStartError + require.ErrorAs(t, startErr, &startFailure) + closeErr := prepared.Close() + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, closeErr, &usageErr) + require.Equal(t, "has active consumers", usageErr.Reason) + + require.NoError(t, syscall.Kill(-processGroupID, syscall.SIGKILL)) + select { + case <-supervisor.Exited(): + case <-time.After(5 * time.Second): + t.Fatalf("failed-start consumer PID %d was not reaped", pid) + } + require.ErrorIs(t, syscall.Kill(-processGroupID, 0), syscall.ESRCH) + require.NoError(t, startFailure.Finalize(t.Context())) + require.NoError(t, startFailure.Finalize(t.Context())) + require.NoError(t, prepared.Close()) +} diff --git a/integration/agentcompat/internal/agent/fixture.go b/integration/agentcompat/internal/agent/fixture.go new file mode 100644 index 00000000..c413f2eb --- /dev/null +++ b/integration/agentcompat/internal/agent/fixture.go @@ -0,0 +1,119 @@ +//go:build linux + +package agent + +import ( + "context" + "fmt" + "os" + "path/filepath" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/workspace" +) + +func agentBuildSpec(sourceDir string) workspace.BuildSpec { + return workspace.BuildSpec{Name: "agent", SourceDir: sourceDir, Package: "./cmd/agent", Tags: []string{"agentcompat"}, Ldflags: []string{"-X", "github.com/nezhahq/agent/pkg/monitor.Version=v2.1.0"}} +} + +func (agent *Agent) prepareFixture(ctx context.Context, config AgentStartConfig) error { + if err := agent.prepareConfig(config); err != nil { + return err + } + if err := agent.prepareFMObserver(config); err != nil { + return err + } + if err := agent.prepareBinary(ctx, config); err != nil { + return err + } + if err := agent.grantWorkspaceOwnership(config); err != nil { + return err + } + agent.prepareEnvironment(config) + return nil +} + +func (agent *Agent) prepareConfig(config AgentStartConfig) error { + configPath, err := agent.workspace.PayloadPath("config.yml") + if err != nil { + return err + } + agent.configPath = configPath + content := fmt.Sprintf("server: %q\nclient_secret: %q\nuuid: %q\ndisable_auto_update: true\ndisable_command_execute: false\ndisable_nat: false\nreport_delay: 1\nip_report_period: 30\ntls: %t\ninsecure_tls: false\ndebug: %t\n", config.Endpoint, config.Secret, config.UUID, config.TLS, config.Debug) + if err := os.WriteFile(configPath, []byte(content), 0o600); err != nil { + return fmt.Errorf("write agent config: %w", err) + } + if !config.TLS || config.CAFilePath == "" { + return nil + } + ca, err := os.ReadFile(config.CAFilePath) + if err != nil { + return fmt.Errorf("read agent CA certificate: %w", err) + } + caPath, err := agent.workspace.PayloadPath("agent-ca.crt") + if err != nil { + return err + } + if err := os.WriteFile(caPath, ca, 0o600); err != nil { + return fmt.Errorf("write agent CA certificate: %w", err) + } + agent.caFilePath = caPath + return nil +} + +func (agent *Agent) prepareFMObserver(config AgentStartConfig) error { + if config.FMObserverRunID == "" { + return nil + } + agent.fmObserverPath = fmObserverSocketPath(agent.workspace.Root()) + observer, err := newFMProducerObserver(agent.fmObserverPath) + if err != nil { + return err + } + agent.fmObserver = observer + return nil +} + +func (agent *Agent) prepareBinary(ctx context.Context, config AgentStartConfig) error { + if config.PreparedBinary != nil { + binaryPath, release, err := config.PreparedBinary.acquire() + if err != nil { + return err + } + agent.binaryPath = binaryPath + agent.releaseBinary = release + agent.releasePending = true + return nil + } + binaryPath, err := agent.workspace.Build(ctx, agentBuildSpec(config.SourceDir)) + if err != nil { + return err + } + agent.binaryPath = binaryPath + return nil +} + +func (agent *Agent) grantWorkspaceOwnership(config AgentStartConfig) error { + if config.Credential == nil { + return nil + } + if err := filepath.WalkDir(agent.workspace.Root(), func(path string, _ os.DirEntry, walkErr error) error { + if walkErr != nil { + return walkErr + } + return os.Chown(path, int(config.Credential.Uid), int(config.Credential.Gid)) + }); err != nil { + return fmt.Errorf("grant agent workspace ownership: %w", err) + } + return nil +} + +func (agent *Agent) prepareEnvironment(config AgentStartConfig) { + environment := filteredEnvironment() + if config.FMObserverRunID != "" { + environment = append(environment, "AGENTCOMPAT_FM_OBSERVER_SOCKET="+agent.fmObserverPath, "AGENTCOMPAT_FM_OBSERVER_RUN_ID="+config.FMObserverRunID) + } + if config.TLS && agent.caFilePath != "" { + environment = append(environment, "SSL_CERT_FILE="+agent.caFilePath) + } + agent.environment = append([]string(nil), environment...) +} diff --git a/integration/agentcompat/internal/agent/fm_observer.go b/integration/agentcompat/internal/agent/fm_observer.go new file mode 100644 index 00000000..8ce5302a --- /dev/null +++ b/integration/agentcompat/internal/agent/fm_observer.go @@ -0,0 +1,92 @@ +//go:build linux + +package agent + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net" + "os" + "path/filepath" + "sync" +) + +type FMProducerSample struct { + RunID string `json:"run_id"` + AgentUUID string `json:"agent_uuid"` + SessionID string `json:"session_id"` + Phase string `json:"phase"` + Active int64 `json:"active"` +} + +type FMProducerObserver struct { + listener net.Listener + samples chan FMProducerSample + done chan struct{} + once sync.Once +} + +func newFMProducerObserver(socketPath string) (*FMProducerObserver, error) { + listener, err := net.Listen("unix", socketPath) + if err != nil { + return nil, fmt.Errorf("listen for FM producer observations: %w", err) + } + observer := &FMProducerObserver{listener: listener, samples: make(chan FMProducerSample, 16), done: make(chan struct{})} + go observer.accept() + return observer, nil +} + +func (observer *FMProducerObserver) accept() { + defer close(observer.done) + for { + connection, err := observer.listener.Accept() + if err != nil { + if errors.Is(err, net.ErrClosed) { + return + } + continue + } + var sample FMProducerSample + err = json.NewDecoder(connection).Decode(&sample) + _ = connection.Close() + if err == nil { + observer.samples <- sample + } + } +} + +func (observer *FMProducerObserver) Await(ctx context.Context, match func(FMProducerSample) bool) (FMProducerSample, error) { + for { + select { + case sample := <-observer.samples: + if match(sample) { + return sample, nil + } + case <-ctx.Done(): + return FMProducerSample{}, ctx.Err() + } + } +} + +func (observer *FMProducerObserver) Close() error { + var closeErr error + observer.once.Do(func() { + closeErr = observer.listener.Close() + <-observer.done + }) + return closeErr +} + +func fmObserverSocketPath(workspaceRoot string) string { + return filepath.Join(workspaceRoot, "fm-observer.sock") +} + +func removeFMObserverSocket(path string) error { + err := os.Remove(path) + if errors.Is(err, os.ErrNotExist) { + return nil + } + return err +} diff --git a/integration/agentcompat/internal/agent/prepared_binary.go b/integration/agentcompat/internal/agent/prepared_binary.go new file mode 100644 index 00000000..d499674b --- /dev/null +++ b/integration/agentcompat/internal/agent/prepared_binary.go @@ -0,0 +1,125 @@ +//go:build linux + +package agent + +import ( + "context" + "fmt" + "os" + "path/filepath" + "sync" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/workspace" +) + +type PreparedBinaryUsageError struct { + Operation string + Reason string +} + +func (err *PreparedBinaryUsageError) Error() string { + return fmt.Sprintf("prepared agent binary %s: %s", err.Operation, err.Reason) +} + +// PreparedBinary owns a build-only workspace. Each successful Start lease keeps +// that workspace alive; Agent workspaces own their own config, logs, and process tracking. +type PreparedBinary struct { + workspace *workspace.Workspace + binaryPath string + mu sync.Mutex + consumers int + closed bool +} + +func PrepareBinary(ctx context.Context, sourceDir string) (*PreparedBinary, error) { + if err := validateSourceDir(sourceDir); err != nil { + return nil, &PreparedBinaryUsageError{Operation: "prepare", Reason: err.Error()} + } + workspaceRoot, err := workspace.New(context.WithoutCancel(ctx)) + if err != nil { + return nil, fmt.Errorf("create prepared agent workspace: %w", err) + } + binaryPath, err := workspaceRoot.Build(ctx, agentBuildSpec(sourceDir)) + if err != nil { + return nil, closePreparedWorkspace(workspaceRoot, err) + } + if err := exposePreparedBinary(workspaceRoot.Root(), binaryPath); err != nil { + return nil, closePreparedWorkspace(workspaceRoot, err) + } + return &PreparedBinary{workspace: workspaceRoot, binaryPath: binaryPath}, nil +} + +func closePreparedWorkspace(workspaceRoot *workspace.Workspace, cause error) error { + return fmt.Errorf("prepare agent binary: %w", closeError(cause, workspaceRoot.Close())) +} + +func exposePreparedBinary(root, binaryPath string) error { + for _, path := range []string{root, filepath.Dir(binaryPath)} { + if err := os.Chmod(path, 0o755); err != nil { + return fmt.Errorf("make prepared agent binary executable: %w", err) + } + } + return nil +} + +func (prepared *PreparedBinary) BinaryPath() string { + if prepared == nil { + return "" + } + prepared.mu.Lock() + defer prepared.mu.Unlock() + return prepared.binaryPath +} + +func (prepared *PreparedBinary) WorkspaceRoot() string { + if prepared == nil { + return "" + } + prepared.mu.Lock() + defer prepared.mu.Unlock() + return prepared.workspace.Root() +} + +func (prepared *PreparedBinary) acquire() (string, func(), error) { + if prepared == nil { + return "", nil, &PreparedBinaryUsageError{Operation: "start", Reason: "is nil"} + } + prepared.mu.Lock() + defer prepared.mu.Unlock() + if prepared.workspace == nil || prepared.binaryPath == "" { + return "", nil, &PreparedBinaryUsageError{Operation: "start", Reason: "is uninitialized"} + } + if prepared.closed { + return "", nil, &PreparedBinaryUsageError{Operation: "start", Reason: "is closed"} + } + prepared.consumers++ + return prepared.binaryPath, prepared.release, nil +} + +func (prepared *PreparedBinary) release() { + prepared.mu.Lock() + defer prepared.mu.Unlock() + prepared.consumers-- +} + +func (prepared *PreparedBinary) Close() error { + if prepared == nil { + return &PreparedBinaryUsageError{Operation: "close", Reason: "is nil"} + } + prepared.mu.Lock() + defer prepared.mu.Unlock() + if prepared.workspace == nil || prepared.binaryPath == "" { + return &PreparedBinaryUsageError{Operation: "close", Reason: "is uninitialized"} + } + if prepared.closed { + return nil + } + if prepared.consumers != 0 { + return &PreparedBinaryUsageError{Operation: "close", Reason: "has active consumers"} + } + if err := prepared.workspace.Close(); err != nil { + return fmt.Errorf("close prepared agent workspace: %w", err) + } + prepared.closed = true + return nil +} diff --git a/integration/agentcompat/internal/agent/prepared_binary_cleanup_failure_test.go b/integration/agentcompat/internal/agent/prepared_binary_cleanup_failure_test.go new file mode 100644 index 00000000..ebe1ed4f --- /dev/null +++ b/integration/agentcompat/internal/agent/prepared_binary_cleanup_failure_test.go @@ -0,0 +1,101 @@ +//go:build linux && agentcompat + +package agent + +import ( + "context" + "errors" + "net" + "os" + "syscall" + "testing" + "time" + + processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process" + "github.com/nezhahq/nezha/integration/agentcompat/internal/workspace" + "github.com/stretchr/testify/require" +) + +func TestPreparedBinary_ReleasesLeaseWhenOnlyWorkspaceCleanupFails(t *testing.T) { + prepared, err := PrepareBinary(t.Context(), testAgentSourceDir(t)) + require.NoError(t, err) + _, release, err := prepared.acquire() + require.NoError(t, err) + workspaceRoot, err := workspace.New(context.WithoutCancel(t.Context())) + require.NoError(t, err) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + ownedListener, err := workspaceRoot.AdoptListener(listener) + require.NoError(t, err) + heldListener, err := ownedListener.ExtraFile() + require.NoError(t, err) + + agent := &Agent{workspace: workspaceRoot, releaseBinary: release, releasePending: true, cleanupDone: make(chan struct{})} + stopErr := agent.Stop(t.Context()) + require.Error(t, stopErr) + require.NoError(t, prepared.Close()) + require.NoError(t, heldListener.Close()) + require.NoError(t, workspaceRoot.Close()) + require.NoDirExists(t, workspaceRoot.Root()) +} + +func TestPreparedBinary_RetainsLeaseWhileConsumerProcessGroupLivesAfterCleanupFailure(t *testing.T) { + prepared, err := PrepareBinary(t.Context(), testAgentSourceDir(t)) + require.NoError(t, err) + binaryPath, release, err := prepared.acquire() + require.NoError(t, err) + workspaceRoot, err := workspace.New(context.WithoutCancel(t.Context())) + require.NoError(t, err) + + supervisor := processharness.NewSupervisor(t.Context(), processharness.Spec{ + Name: "prepared-binary-lingering-consumer", Path: "/bin/sh", Args: []string{"-c", "exec tail -f /dev/null"}, + MaxLogBytes: 1024, TerminateTimeout: time.Second, KillTimeout: time.Second, + }) + cancelledContext, cancel := context.WithCancel(t.Context()) + cancel() + _ = supervisor.Stop(cancelledContext) + require.NoError(t, supervisor.Stop(t.Context())) + require.NoError(t, supervisor.Start()) + pid := supervisor.PID() + pgid := supervisor.ProcessGroupID() + require.NoError(t, workspaceRoot.TrackPID(pid)) + require.NoError(t, workspaceRoot.TrackProcessGroup(pgid)) + + agent := &Agent{ + workspace: workspaceRoot, binaryPath: binaryPath, releaseBinary: release, releasePending: true, + cleanupDone: make(chan struct{}), processes: []*processGeneration{{supervisor: supervisor, identity: ProcessIdentity{PID: pid, ProcessGroupID: pgid}}}, + } + t.Cleanup(func() { + _ = syscall.Kill(-pgid, syscall.SIGKILL) + select { + case <-supervisor.Exited(): + case <-time.After(5 * time.Second): + } + _ = workspaceRoot.Close() + _ = prepared.Close() + }) + + stopErr := agent.Stop(t.Context()) + require.Error(t, stopErr) + firstStopError := stopErr.Error() + require.NoError(t, syscall.Kill(-pgid, 0)) + require.FileExists(t, binaryPath) + closeErr := prepared.Close() + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, closeErr, &usageErr) + require.Equal(t, "has active consumers", usageErr.Reason) + + require.NoError(t, syscall.Kill(-pgid, syscall.SIGKILL)) + select { + case <-supervisor.Exited(): + case <-time.After(5 * time.Second): + t.Fatal("lingering consumer process was not reaped") + } + require.True(t, errors.Is(syscall.Kill(-pgid, 0), syscall.ESRCH)) + recoveryErr := agent.Stop(t.Context()) + require.Error(t, recoveryErr) + require.Equal(t, firstStopError, recoveryErr.Error()) + require.NoError(t, prepared.Close()) + _, statErr := os.Stat(prepared.WorkspaceRoot()) + require.ErrorIs(t, statErr, os.ErrNotExist) +} diff --git a/integration/agentcompat/internal/agent/prepared_binary_concurrency_test.go b/integration/agentcompat/internal/agent/prepared_binary_concurrency_test.go new file mode 100644 index 00000000..f3bbcb70 --- /dev/null +++ b/integration/agentcompat/internal/agent/prepared_binary_concurrency_test.go @@ -0,0 +1,57 @@ +//go:build linux && agentcompat + +package agent + +import ( + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestPreparedBinary_ConcurrentConsumersSharePathAndBlockClose(t *testing.T) { + prepared, err := PrepareBinary(t.Context(), testAgentSourceDir(t)) + require.NoError(t, err) + t.Cleanup(func() { _ = prepared.Close() }) + + releases := make(chan func(), 8) + errorsChannel := make(chan error, 8) + var acquireGroup sync.WaitGroup + for range 8 { + acquireGroup.Add(1) + go func() { + defer acquireGroup.Done() + path, release, acquireErr := prepared.acquire() + if acquireErr == nil && path != prepared.BinaryPath() { + acquireErr = fmt.Errorf("acquired path %q differs from prepared path", path) + } + errorsChannel <- acquireErr + if acquireErr == nil { + releases <- release + } + }() + } + acquireGroup.Wait() + close(errorsChannel) + close(releases) + for acquireErr := range errorsChannel { + require.NoError(t, acquireErr) + } + + closeErr := prepared.Close() + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, closeErr, &usageErr) + require.Equal(t, "has active consumers", usageErr.Reason) + + var releaseGroup sync.WaitGroup + for release := range releases { + releaseGroup.Add(1) + go func() { + defer releaseGroup.Done() + release() + }() + } + releaseGroup.Wait() + require.NoError(t, prepared.Close()) +} diff --git a/integration/agentcompat/internal/agent/prepared_binary_test.go b/integration/agentcompat/internal/agent/prepared_binary_test.go new file mode 100644 index 00000000..4ab3c233 --- /dev/null +++ b/integration/agentcompat/internal/agent/prepared_binary_test.go @@ -0,0 +1,190 @@ +//go:build linux && agentcompat + +package agent + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestPreparedBinary_EightIndependentAgentsShareBinaryAndCleanUp(t *testing.T) { + prepared, err := PrepareBinary(t.Context(), testAgentSourceDir(t)) + require.NoError(t, err) + preparedRoot := prepared.WorkspaceRoot() + binaryPath := prepared.BinaryPath() + require.FileExists(t, binaryPath) + initialBinaryInfo, err := os.Stat(binaryPath) + require.NoError(t, err) + initialBinaryStat, ok := initialBinaryInfo.Sys().(*syscall.Stat_t) + require.True(t, ok) + instances := make([]*Agent, 0, 8) + t.Cleanup(func() { + for _, instance := range instances { + cleanupContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + _ = instance.Stop(cleanupContext) + cancel() + } + _ = prepared.Close() + }) + dashboardInstance := startTestDashboard(t, true) + configPaths := make(map[string]struct{}, 8) + logPaths := make(map[string]struct{}, 8) + for index := range 8 { + config := AgentStartConfig{ + PreparedBinary: prepared, + Endpoint: dashboardInstance.TLSEndpoint(), + Secret: dashboardInstance.AgentSecret(), + UUID: preparedBinaryUUID(index), + TLS: true, + CAFilePath: dashboardInstance.TLSCACertificatePath(), + Debug: index%2 == 0, + } + if index == 6 { + config.FMObserverRunID = "prepared-binary-observer" + } + if index == 7 { + config.Credential = &syscall.Credential{Uid: 65534, Gid: 65534} + } + instance, startErr := Start(t.Context(), config) + require.NoError(t, startErr) + require.Equal(t, binaryPath, instance.BinaryPath()) + processBinaryInfo, statErr := os.Stat(fmt.Sprintf("/proc/%d/exe", instance.PID())) + if errors.Is(statErr, syscall.EACCES) || errors.Is(statErr, syscall.EPERM) { + t.Logf("kernel denied /proc/%d/exe metadata for agent %d", instance.PID(), index) + } else { + require.NoError(t, statErr) + processBinaryStat, statOK := processBinaryInfo.Sys().(*syscall.Stat_t) + require.True(t, statOK) + require.Equal(t, initialBinaryStat.Dev, processBinaryStat.Dev) + require.Equal(t, initialBinaryStat.Ino, processBinaryStat.Ino) + } + require.NotEqual(t, preparedRoot, instance.WorkspaceRoot()) + require.NoFileExists(t, filepath.Join(instance.WorkspaceRoot(), "bin", "agent")) + configPaths[instance.ConfigPath()] = struct{}{} + logPaths[instance.LogPath()] = struct{}{} + require.NoError(t, waitForAgentReady(t, instance, dashboardInstance)) + if index == 6 { + require.NotNil(t, instance.FMProducerObserver()) + require.FileExists(t, instance.fmObserverPath) + } + if index == 7 { + status, statusErr := os.ReadFile(fmt.Sprintf("/proc/%d/status", instance.PID())) + require.NoError(t, statusErr) + require.Contains(t, string(status), "Uid:\t65534\t65534\t65534\t65534") + } + instances = append(instances, instance) + } + finalBinaryInfo, err := os.Stat(binaryPath) + require.NoError(t, err) + require.True(t, os.SameFile(initialBinaryInfo, finalBinaryInfo)) + require.Len(t, configPaths, 8) + require.Len(t, logPaths, 8) + closeErr := prepared.Close() + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, closeErr, &usageErr) + require.Equal(t, "close", usageErr.Operation) + require.Equal(t, "has active consumers", usageErr.Reason) + for index, instance := range instances { + stopContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + require.NoError(t, instance.Stop(stopContext), "agent %d", index) + cancel() + require.True(t, instance.CleanupReceipt().Passed) + require.False(t, instance.CleanupReceipt().Forced) + if index < len(instances)-1 { + require.DirExists(t, preparedRoot) + require.FileExists(t, binaryPath) + } + } + require.NoError(t, prepared.Close()) + require.NoError(t, prepared.Close()) + _, statErr := os.Stat(preparedRoot) + require.ErrorIs(t, statErr, os.ErrNotExist) +} + +func TestAgent_StartBuildsAndOwnsItsBinary(t *testing.T) { + instance, err := Start(t.Context(), AgentStartConfig{ + SourceDir: testAgentSourceDir(t), + Endpoint: "127.0.0.1:1", + UUID: "00000000-0000-0000-0000-000000000197", + }) + require.NoError(t, err) + workspaceRoot := instance.WorkspaceRoot() + require.Equal(t, filepath.Join(workspaceRoot, "bin", "agent"), instance.BinaryPath()) + require.FileExists(t, instance.BinaryPath()) + stopContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, instance.Stop(stopContext)) + _, statErr := os.Stat(workspaceRoot) + require.ErrorIs(t, statErr, os.ErrNotExist) +} + +func TestPreparedBinary_RejectsConsumerAfterClose(t *testing.T) { + prepared, err := PrepareBinary(t.Context(), testAgentSourceDir(t)) + require.NoError(t, err) + require.NoError(t, prepared.Close()) + _, err = Start(t.Context(), AgentStartConfig{ + PreparedBinary: prepared, + Endpoint: "127.0.0.1:1", + UUID: "00000000-0000-0000-0000-000000000198", + }) + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, err, &usageErr) + require.Equal(t, "start", usageErr.Operation) + require.Equal(t, "is closed", usageErr.Reason) +} + +func TestPreparedBinary_RejectsInvalidSourceDirectory(t *testing.T) { + _, err := PrepareBinary(t.Context(), "relative") + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, err, &usageErr) + require.Equal(t, "prepare", usageErr.Operation) + require.True(t, strings.Contains(usageErr.Reason, "must be absolute")) +} + +func TestPreparedBinary_CancelledBuildRemovesWorkspace(t *testing.T) { + temporaryRoot := t.TempDir() + t.Setenv("TMPDIR", temporaryRoot) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err := PrepareBinary(ctx, testAgentSourceDir(t)) + require.Error(t, err) + entries, readErr := os.ReadDir(temporaryRoot) + require.NoError(t, readErr) + require.Empty(t, entries) +} + +func TestPreparedBinary_RejectsUninitializedValue(t *testing.T) { + _, err := Start(t.Context(), AgentStartConfig{ + PreparedBinary: &PreparedBinary{}, + Endpoint: "127.0.0.1:1", + UUID: "00000000-0000-0000-0000-000000000199", + }) + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, err, &usageErr) + require.Equal(t, "start", usageErr.Operation) + require.Equal(t, "is uninitialized", usageErr.Reason) +} + +func TestPreparedBinary_CloseRejectsNilAndUninitializedValues(t *testing.T) { + var nilPrepared *PreparedBinary + for _, prepared := range []*PreparedBinary{nilPrepared, &PreparedBinary{}} { + err := prepared.Close() + var usageErr *PreparedBinaryUsageError + require.ErrorAs(t, err, &usageErr) + require.Equal(t, "close", usageErr.Operation) + } +} + +func preparedBinaryUUID(index int) string { + return fmt.Sprintf("00000000-0000-0000-0000-%012d", 190+index) +} diff --git a/integration/agentcompat/internal/agent/process_lifecycle.go b/integration/agentcompat/internal/agent/process_lifecycle.go new file mode 100644 index 00000000..2082173c --- /dev/null +++ b/integration/agentcompat/internal/agent/process_lifecycle.go @@ -0,0 +1,170 @@ +//go:build linux + +package agent + +import ( + "context" + "errors" + "fmt" + "strings" + "syscall" + "time" + + processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process" +) + +type ProcessIdentity struct { + Generation uint64 + PID int + ProcessGroupID int +} + +type ProcessTransition struct { + Previous ProcessIdentity + Current ProcessIdentity +} + +type processGeneration struct { + supervisor *processharness.Supervisor + identity ProcessIdentity + record processharness.CleanupRecord +} + +func (agent *Agent) RuntimeIdentity() ProcessIdentity { + agent.processMu.Lock() + defer agent.processMu.Unlock() + if agent.currentProcess == nil { + return ProcessIdentity{} + } + return agent.currentProcess.identity +} + +func (agent *Agent) StartProcess(ctx context.Context) (ProcessTransition, error) { + agent.processMu.Lock() + defer agent.processMu.Unlock() + if agent.closed { + return ProcessTransition{}, errors.New("agent is closed") + } + if agent.currentProcess != nil { + return ProcessTransition{}, errors.New("agent process is already running") + } + agent.generation++ + logFile, err := agent.workspace.Log(fmt.Sprintf("agent-%s-generation-%d", strings.ReplaceAll(agent.uuid, "-", ""), agent.generation)) + if err != nil { + return ProcessTransition{}, err + } + agent.logPath = logFile.Name() + newSupervisor := processharness.NewSupervisor + if agent.startConfig.newSupervisor != nil { + newSupervisor = agent.startConfig.newSupervisor + } + supervisor := newSupervisor(ctx, processharness.Spec{ + Name: "agent", Path: agent.binaryPath, Args: []string{"-c", agent.configPath}, Env: agent.environment, + Stdout: logFile, Stderr: logFile, MaxLogBytes: agentMaxLogBytes, + TerminateTimeout: agentStopTimeout, KillTimeout: agentKillTimeout, + Credential: agent.startConfig.Credential, + }) + if err := supervisor.Start(); err != nil { + return ProcessTransition{}, err + } + identity := ProcessIdentity{Generation: agent.generation, PID: supervisor.PID(), ProcessGroupID: supervisor.ProcessGroupID()} + generation := &processGeneration{supervisor: supervisor, identity: identity, record: supervisor.CleanupRecord()} + // Register the started generation before post-start setup so failures remain cleanup-owned. + agent.currentProcess = generation + agent.supervisor = supervisor + agent.processes = append(agent.processes, generation) + if err := agent.trackPID(identity.PID); err != nil { + return agent.rollbackStartedProcess(ctx, generation, err) + } + if err := agent.trackProcessGroup(identity.ProcessGroupID); err != nil { + return agent.rollbackStartedProcess(ctx, generation, err) + } + previous := ProcessIdentity{} + return ProcessTransition{Previous: previous, Current: identity}, nil +} + +func (agent *Agent) rollbackStartedProcess(ctx context.Context, generation *processGeneration, trackingErr error) (ProcessTransition, error) { + agent.currentProcess = nil + agent.supervisor = nil + rollbackContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 15*time.Second) + defer cancel() + rollbackErr := generation.supervisor.Stop(rollbackContext) + generation.record = generation.supervisor.CleanupRecord() + return ProcessTransition{}, errors.Join(trackingErr, rollbackErr) +} + +func (agent *Agent) StopProcess(ctx context.Context) (ProcessTransition, error) { + agent.processMu.Lock() + process := agent.currentProcess + if process == nil { + agent.processMu.Unlock() + return ProcessTransition{}, errors.New("agent process is not running") + } + agent.currentProcess = nil + agent.supervisor = nil + agent.processMu.Unlock() + if err := process.supervisor.Stop(ctx); err != nil { + return ProcessTransition{Previous: process.identity}, fmt.Errorf("stop agent process: %w", err) + } + process.record = process.supervisor.CleanupRecord() + return ProcessTransition{Previous: process.identity}, nil +} + +func (agent *Agent) RestartProcess(ctx context.Context) (ProcessTransition, error) { + stopped, err := agent.StopProcess(ctx) + if err != nil { + return stopped, err + } + started, err := agent.StartProcess(ctx) + if err != nil { + return ProcessTransition{Previous: stopped.Previous}, err + } + return ProcessTransition{Previous: stopped.Previous, Current: started.Current}, nil +} + +func (agent *Agent) Restart(ctx context.Context) error { + _, err := agent.RestartProcess(ctx) + return err +} + +func (agent *Agent) Close(ctx context.Context) error { + agent.processMu.Lock() + agent.closed = true + agent.processMu.Unlock() + return agent.Stop(ctx) +} + +func (agent *Agent) closeProcesses(ctx context.Context) error { + agent.processMu.Lock() + processes := append([]*processGeneration(nil), agent.processes...) + agent.processMu.Unlock() + var cleanupError error + for _, process := range processes { + stopContext, cancel := context.WithTimeout(ctx, 15*time.Second) + cleanupError = errors.Join(cleanupError, process.supervisor.Stop(stopContext)) + cancel() + process.record = process.supervisor.CleanupRecord() + } + return cleanupError +} + +func (agent *Agent) processesQuiescent() bool { + agent.processMu.Lock() + processes := append([]*processGeneration(nil), agent.processes...) + agent.processMu.Unlock() + for _, process := range processes { + select { + case <-process.supervisor.Exited(): + default: + return false + } + err := syscall.Kill(-process.identity.ProcessGroupID, 0) + if err == nil || errors.Is(err, syscall.EPERM) { + return false + } + if !errors.Is(err, syscall.ESRCH) { + return false + } + } + return true +} diff --git a/integration/agentcompat/internal/agent/process_tracking_rollback_test.go b/integration/agentcompat/internal/agent/process_tracking_rollback_test.go new file mode 100644 index 00000000..ed1bcd9a --- /dev/null +++ b/integration/agentcompat/internal/agent/process_tracking_rollback_test.go @@ -0,0 +1,142 @@ +//go:build linux && agentcompat + +package agent + +import ( + "context" + "errors" + "os" + "path/filepath" + "strconv" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/workspace" +) + +func TestAgent_StartProcess_tracksStartedGenerationForStop(t *testing.T) { + agent := newUnstartedTestAgent(t) + transition, err := agent.StartProcess(t.Context()) + require.NoError(t, err) + require.NotZero(t, transition.Current.PID) + require.Equal(t, transition.Current, agent.RuntimeIdentity()) + require.NoError(t, agent.Stop(t.Context())) +} + +func TestAgent_StartProcess_rollsBackStartedGenerationWhenPIDTrackingFails(t *testing.T) { + agent := newUnstartedTestAgent(t) + trackingErr := errors.New("injected PID tracking failure") + var startedPID int + agent.trackPID = func(pid int) error { + startedPID = pid + return trackingErr + } + _, err := agent.StartProcess(t.Context()) + require.ErrorIs(t, err, trackingErr) + require.Empty(t, agent.RuntimeIdentity()) + receipt := agent.CleanupReceipt() + require.Len(t, receipt.Processes, 1) + t.Logf("started_pid=%d started_pgid=%d injected_failure=%q runtime_identity=%+v cleanup_record=%+v forced=%t", startedPID, startedPID, trackingErr, agent.RuntimeIdentity(), receipt.Processes[0], receipt.Forced) + require.Equal(t, "agent", receipt.Processes[0].Name) + require.NotZero(t, receipt.Processes[0].PID) + require.False(t, receipt.Processes[0].Forced) + requireProcessAndGroupGone(t, receipt.Processes[0].PID, receipt.Processes[0].PID) + require.FileExists(t, agent.ConfigPath()) + require.NoError(t, agent.Stop(t.Context())) + require.NoDirExists(t, agent.WorkspaceRoot()) +} + +func TestAgent_StartProcess_rollsBackStartedGenerationWhenProcessGroupTrackingFails(t *testing.T) { + agent := newUnstartedTestAgent(t) + trackingErr := errors.New("injected process group tracking failure") + var startedPID, startedProcessGroupID int + agent.trackPID = func(pid int) error { + startedPID = pid + return nil + } + agent.trackProcessGroup = func(processGroupID int) error { + startedProcessGroupID = processGroupID + return trackingErr + } + _, err := agent.StartProcess(t.Context()) + require.ErrorIs(t, err, trackingErr) + require.Empty(t, agent.RuntimeIdentity()) + receipt := agent.CleanupReceipt() + require.Len(t, receipt.Processes, 1) + t.Logf("started_pid=%d started_pgid=%d injected_failure=%q runtime_identity=%+v cleanup_record=%+v forced=%t", startedPID, startedProcessGroupID, trackingErr, agent.RuntimeIdentity(), receipt.Processes[0], receipt.Forced) + require.Equal(t, "agent", receipt.Processes[0].Name) + require.NotZero(t, receipt.Processes[0].PID) + require.False(t, receipt.Processes[0].Forced) + requireProcessAndGroupGone(t, receipt.Processes[0].PID, startedProcessGroupID) + require.FileExists(t, agent.ConfigPath()) + require.NoError(t, agent.Stop(t.Context())) + require.NoDirExists(t, agent.WorkspaceRoot()) +} + +func newUnstartedTestAgent(t *testing.T) *Agent { + workspaceRoot, err := workspace.New(t.Context()) + require.NoError(t, err) + agent := &Agent{workspace: workspaceRoot, uuid: "00000000-0000-0000-0000-000000000091", cleanupDone: make(chan struct{}), startConfig: AgentStartConfig{SourceDir: testAgentSourceDir(t), Endpoint: "127.0.0.1:1", Secret: agentSecret}, trackPID: workspaceRoot.TrackPID, trackProcessGroup: workspaceRoot.TrackProcessGroup} + require.NoError(t, agent.prepareFixture(t.Context(), agent.startConfig)) + t.Cleanup(func() { + cleanupContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, agent.Stop(cleanupContext)) + }) + return agent +} + +func requireProcessAndGroupGone(t *testing.T, pid, processGroupID int) { + _, err := os.Stat(filepath.Join("/proc", strconv.Itoa(pid))) + require.ErrorIs(t, err, os.ErrNotExist) + require.ErrorIs(t, syscall.Kill(-processGroupID, 0), syscall.ESRCH) +} + +func TestAgent_StopProcessPreservesConfig(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, false) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{UUID: "00000000-0000-0000-0000-000000000089"}) + require.NoError(t, waitForAgentReady(t, agentInstance, dashboardInstance)) + configBefore, err := os.ReadFile(agentInstance.ConfigPath()) + require.NoError(t, err) + pidBefore := agentInstance.PID() + + // When + stopContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + _, err = agentInstance.StopProcess(stopContext) + + // Then + require.NoError(t, err) + configAfter, readErr := os.ReadFile(agentInstance.ConfigPath()) + require.NoError(t, readErr) + require.Equal(t, configBefore, configAfter) + require.Equal(t, "00000000-0000-0000-0000-000000000089", agentInstance.UUID()) + require.NotZero(t, pidBefore) +} + +func TestAgent_RestartProcessPreservesConfigBytesAndUUID(t *testing.T) { + // Given + dashboardInstance := startTestDashboard(t, false) + agentInstance := startTestAgent(t, dashboardInstance, AgentStartConfig{UUID: "00000000-0000-0000-0000-000000000090"}) + require.NoError(t, waitForAgentReady(t, agentInstance, dashboardInstance)) + configBefore, err := os.ReadFile(agentInstance.ConfigPath()) + require.NoError(t, err) + pidBefore := agentInstance.PID() + + // When + restartContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + _, err = agentInstance.RestartProcess(restartContext) + + // Then + require.NoError(t, err) + require.NotEqual(t, pidBefore, agentInstance.PID()) + configAfter, readErr := os.ReadFile(agentInstance.ConfigPath()) + require.NoError(t, readErr) + require.Equal(t, configBefore, configAfter) + require.Equal(t, "00000000-0000-0000-0000-000000000090", agentInstance.UUID()) +} diff --git a/integration/agentcompat/internal/agent/readiness.go b/integration/agentcompat/internal/agent/readiness.go new file mode 100644 index 00000000..eaa0c34f --- /dev/null +++ b/integration/agentcompat/internal/agent/readiness.go @@ -0,0 +1,241 @@ +//go:build linux + +package agent + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "time" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/client" + "github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard" +) + +const readinessBudget = 45 * time.Second + +type Readiness struct { + ServerID uint64 + UUID string + Version string + Online bool + LastActive time.Time + VersionObserved bool + RequestTaskEstablished bool + StateReceiptObserved bool + Host json.RawMessage + State json.RawMessage +} + +type serverListArguments struct { + OnlineOnly bool `json:"online_only"` +} + +type serverListResult struct { + Servers []serverListItem `json:"servers"` + Count int `json:"count"` +} + +type serverListItem struct { + ID uint64 `json:"id"` + UUID string `json:"uuid"` + Online bool `json:"online"` + Platform string `json:"platform"` + Arch string `json:"arch"` + LastActive time.Time `json:"last_active"` +} + +type serverGetArguments struct { + ServerID uint64 `json:"server_id"` +} + +type serverGetResult struct { + ID uint64 `json:"id"` + UUID string `json:"uuid"` + Host json.RawMessage `json:"host"` + State json.RawMessage `json:"state"` + LastActive time.Time `json:"last_active"` +} + +type execArguments struct { + ServerID uint64 `json:"server_id"` + Cmd string `json:"cmd"` + Args []string `json:"args"` +} + +type execResult struct { + ExitCode int `json:"exit_code"` + Stdout string `json:"stdout"` + Error string `json:"error"` +} + +func (agent *Agent) WaitReady(ctx context.Context, dashboardInstance *dashboard.Dashboard) (Readiness, error) { + deadline, cancel := context.WithTimeout(ctx, readinessBudget) + defer cancel() + ticker := time.NewTicker(500 * time.Millisecond) + defer ticker.Stop() + for { + readiness, err := agent.probeReadiness(deadline, dashboardInstance) + if err == nil { + return readiness, nil + } + select { + case <-agent.supervisor.Exited(): + return Readiness{}, fmt.Errorf("agent exited before readiness: %w", err) + case <-deadline.Done(): + return Readiness{}, fmt.Errorf("agent readiness: %w", errors.Join(err, deadline.Err())) + case <-ticker.C: + } + } +} + +func (agent *Agent) WaitReadyEventDriven(ctx context.Context, dashboardInstance *dashboard.Dashboard) (Readiness, error) { + serverID, err := dashboardInstance.WaitForInfo2UUID(ctx, agent.uuid) + if err != nil { + return Readiness{}, fmt.Errorf("agent info2 readiness: %w", err) + } + readiness, err := agent.probeReadinessForServer(ctx, dashboardInstance.Clients().MCP, serverID) + if err != nil { + return Readiness{}, err + } + return readiness, nil +} + +func (agent *Agent) WaitReadyEventDrivenWithClient(ctx context.Context, dashboardInstance *dashboard.Dashboard, mcpClient *client.Client) (Readiness, error) { + serverID, err := dashboardInstance.WaitForInfo2UUID(ctx, agent.uuid) + if err != nil { + return Readiness{}, fmt.Errorf("agent info2 readiness: %w", err) + } + return agent.probeReadinessForServer(ctx, mcpClient, serverID) +} + +func (agent *Agent) probeReadinessForServer(ctx context.Context, mcpClient *client.Client, serverID uint64) (Readiness, error) { + serverResponse, err := client.CallTool[serverGetArguments, serverGetResult](ctx, mcpClient, client.ToolCall[serverGetArguments]{Name: "server.get", Arguments: serverGetArguments{ServerID: serverID}}) + if err != nil { + return Readiness{}, err + } + server := serverListItem{ID: serverID, UUID: agent.uuid, Online: true} + if err := verifyServerGetResult(server, serverResponse.StructuredContent); err != nil { + return Readiness{}, err + } + execResponse, err := client.CallTool[execArguments, execResult](ctx, mcpClient, client.ToolCall[execArguments]{Name: "server.exec", Arguments: execArguments{ServerID: serverID, Cmd: "sh", Args: []string{"-c", "printf agentcompat-ready"}}}) + if err != nil { + return Readiness{}, fmt.Errorf("live RequestTask probe: %w", err) + } + if execResponse.StructuredContent.ExitCode != 0 || execResponse.StructuredContent.Stdout != "agentcompat-ready" { + return Readiness{}, errors.New("live RequestTask probe returned unexpected result") + } + version, versionObserved, err := decodeHostVersionEvidence(serverResponse.StructuredContent.Host) + if err != nil { + return Readiness{}, err + } + return Readiness{ServerID: serverID, UUID: agent.uuid, Version: version, Online: true, LastActive: serverResponse.StructuredContent.LastActive, VersionObserved: versionObserved, RequestTaskEstablished: true, StateReceiptObserved: true, Host: serverResponse.StructuredContent.Host, State: serverResponse.StructuredContent.State}, nil +} + +func (agent *Agent) probeReadiness(ctx context.Context, dashboardInstance *dashboard.Dashboard) (Readiness, error) { + call := dashboardInstance.Clients().MCP + list, err := client.CallTool[serverListArguments, serverListResult](ctx, call, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}}) + if err != nil { + return Readiness{}, err + } + var server serverListItem + for _, candidate := range list.StructuredContent.Servers { + if candidate.UUID == agent.uuid { + server = candidate + break + } + } + if server.ID == 0 || server.UUID != agent.uuid || !server.Online { + return Readiness{}, errors.New("agent UUID is not online in dashboard server.list") + } + execResponse, err := client.CallTool[execArguments, execResult](ctx, call, client.ToolCall[execArguments]{Name: "server.exec", Arguments: execArguments{ServerID: server.ID, Cmd: "sh", Args: []string{"-c", "printf agentcompat-ready"}}}) + if err != nil { + return Readiness{}, fmt.Errorf("live RequestTask probe: %w", err) + } + if execResponse.StructuredContent.ExitCode != 0 || execResponse.StructuredContent.Stdout != "agentcompat-ready" { + return Readiness{}, errors.New("live RequestTask probe returned unexpected result") + } + serverResponse, err := client.CallTool[serverGetArguments, serverGetResult](ctx, call, client.ToolCall[serverGetArguments]{Name: "server.get", Arguments: serverGetArguments{ServerID: server.ID}}) + if err != nil { + return Readiness{}, err + } + if err := verifyServerGetResult(server, serverResponse.StructuredContent); err != nil { + return Readiness{}, err + } + version, versionObserved, err := decodeHostVersionEvidence(serverResponse.StructuredContent.Host) + if err != nil { + return Readiness{}, err + } + stateReceiptObserved := dashboardInstance.ReceiptAccepted() + if !dashboardInstance.ReceiptGateEnabled() { + stateReceiptObserved = agent.observeStateReceipt(serverResponse.StructuredContent.LastActive) + } + if !stateReceiptObserved { + return Readiness{}, errors.New("waiting for a second state report after receipt") + } + return Readiness{ + ServerID: server.ID, UUID: agent.uuid, Version: version, Online: true, LastActive: serverResponse.StructuredContent.LastActive, VersionObserved: versionObserved, + RequestTaskEstablished: true, StateReceiptObserved: stateReceiptObserved, + Host: serverResponse.StructuredContent.Host, State: serverResponse.StructuredContent.State, + }, nil +} + +func decodeHostVersionEvidence(raw json.RawMessage) (string, bool, error) { + var host struct { + Version string `json:"version"` + } + if err := json.Unmarshal(raw, &host); err != nil { + return "", false, fmt.Errorf("decode dashboard Host: %w", err) + } + // A decoded Host object is not version evidence unless the Agent reported a value. + return host.Version, host.Version != "", nil +} + +func verifyServerGetResult(server serverListItem, result serverGetResult) error { + if server.ID == 0 || result.ID != server.ID || result.UUID != server.UUID { + return errors.New("dashboard server.get identity does not match server.list") + } + if len(result.Host) == 0 || len(result.State) == 0 || string(result.Host) == "null" || string(result.State) == "null" { + return errors.New("dashboard server.get omitted Host or State") + } + return nil +} + +func (agent *Agent) observeStateReceipt(lastActive time.Time) bool { + agent.readinessMu.Lock() + defer agent.readinessMu.Unlock() + observed := !agent.lastStateReport.IsZero() && lastActive.After(agent.lastStateReport) + if lastActive.After(agent.lastStateReport) { + agent.lastStateReport = lastActive + } + return observed +} + +func (agent *Agent) AssertNeverOnline(ctx context.Context, dashboardInstance *dashboard.Dashboard, duration time.Duration) error { + deadline, cancel := context.WithTimeout(ctx, duration) + defer cancel() + ticker := time.NewTicker(500 * time.Millisecond) + defer ticker.Stop() + var lastError error + for { + list, err := client.CallTool[serverListArguments, serverListResult](deadline, dashboardInstance.Clients().MCP, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}}) + if err == nil { + for _, server := range list.StructuredContent.Servers { + if server.UUID == agent.uuid { + return errors.New("invalid-secret agent became online") + } + } + } else if deadline.Err() == nil { + lastError = err + } + select { + case <-deadline.Done(): + if lastError != nil { + return fmt.Errorf("server.list unavailable while asserting agent stayed offline: %w", lastError) + } + return nil + case <-ticker.C: + } + } +} diff --git a/integration/agentcompat/internal/agent/readiness_server_id_test.go b/integration/agentcompat/internal/agent/readiness_server_id_test.go new file mode 100644 index 00000000..c2aecce3 --- /dev/null +++ b/integration/agentcompat/internal/agent/readiness_server_id_test.go @@ -0,0 +1,47 @@ +//go:build linux && agentcompat + +package agent + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestServerGetResult_VerifiesListedServerIdentity_whenUUIDOrIDMismatch(t *testing.T) { + listed := serverListItem{ID: 81, UUID: "00000000-0000-0000-0000-000000000081", Online: true} + tests := []struct { + name string + result serverGetResult + }{ + { + name: "UUID differs", + result: serverGetResult{ID: listed.ID, UUID: "00000000-0000-0000-0000-000000000082", Host: json.RawMessage(`{}`), State: json.RawMessage(`{}`)}, + }, + { + name: "ID differs", + result: serverGetResult{ID: 82, UUID: listed.UUID, Host: json.RawMessage(`{}`), State: json.RawMessage(`{}`)}, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + result := test.result + err := verifyServerGetResult(listed, result) + require.Error(t, err) + }) + } +} + +func TestServerGetResult_RejectsZeroListedServerID_whenReturnedIdentityMatches(t *testing.T) { + // Given + listed := serverListItem{UUID: "00000000-0000-0000-0000-000000000081", Online: true} + result := serverGetResult{UUID: listed.UUID, Host: json.RawMessage(`{}`), State: json.RawMessage(`{}`)} + + // When + err := verifyServerGetResult(listed, result) + + // Then + require.Error(t, err) + require.EqualError(t, err, "dashboard server.get identity does not match server.list") +} diff --git a/integration/agentcompat/internal/agent/readiness_version_test.go b/integration/agentcompat/internal/agent/readiness_version_test.go new file mode 100644 index 00000000..804393c3 --- /dev/null +++ b/integration/agentcompat/internal/agent/readiness_version_test.go @@ -0,0 +1,36 @@ +//go:build linux + +package agent + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestHostVersionEvidence_IsNotObserved_whenReportedVersionIsEmpty(t *testing.T) { + // Given + host := json.RawMessage(`{"version":""}`) + + // When + version, observed, err := decodeHostVersionEvidence(host) + + // Then + require.NoError(t, err) + require.Empty(t, version) + require.False(t, observed) +} + +func TestHostVersionEvidence_IsObserved_whenReportedVersionIsNonempty(t *testing.T) { + // Given + host := json.RawMessage(`{"version":"v2.1.0"}`) + + // When + version, observed, err := decodeHostVersionEvidence(host) + + // Then + require.NoError(t, err) + require.Equal(t, "v2.1.0", version) + require.True(t, observed) +}