diff --git a/.gitignore b/.gitignore index dc3adbc..c5f4e4c 100644 --- a/.gitignore +++ b/.gitignore @@ -23,6 +23,7 @@ vendor/ Thumbs.db .idea/ .vscode/ +.kilo/ *.swp *.swo diff --git a/README.md b/README.md index 94005f7..6637164 100644 --- a/README.md +++ b/README.md @@ -37,27 +37,6 @@ Hooks run automatically on `git commit`. To run them manually: pre-commit run --all-files ``` -## Scripts - -The `scripts/` directory is part of the public release because it shows how the -agent is built, installed, tested, and removed. - -- `scripts/install_agent.sh`: VM install helper. Requires a gateway URL argument - or `GATEWAY_URL`; optionally sends `INFRAHUB_KEY` as a download authorization - header without persisting it. -- `scripts/install.sh`: advanced systemd installer for a local agent binary. - Requires `BINARY_SOURCE` and does not embed credentials. -- `scripts/uninstall_agent.sh`: removes the systemd service and agent runtime - files. -- `scripts/build.sh`: builds linux release artifacts and SHA-256 checksums. -- `scripts/serve.sh`: local Docker helper that serves `/download` and - `/version` for install/update tests. -- `scripts/e2e.sh`: local Docker end-to-end test helper. - -Do not commit `.env` files, real gateway URLs, private IPs, credentials, tokens, -or customer-specific metadata into this repository. Use placeholders such as -`gateway.example.com` in docs and test fixtures. - ## Configuration The agent is configured through environment variables: @@ -66,8 +45,7 @@ The agent is configured through environment variables: - `HYPERSTACK_INTERVAL`: Collection interval. Defaults to `15s`. - `HYPERSTACK_ENABLE_NODE`: Enable node metrics. Defaults to `true`. - `HYPERSTACK_ENABLE_GPU`: Enable GPU metrics. Defaults to `true`. -- `HYPERSTACK_HEALTH_ADDR`: Health and self-metrics bind address. Defaults to - `127.0.0.1:9100`. +- `HYPERSTACK_HEALTH_ADDR`: Health and self-metrics bind address. Defaults to `127.0.0.1:9100`. - `METADATA_URL`: Optional metadata service URL override. ## Task Targets diff --git a/cmd/agent/main.go b/cmd/agent/main.go index a88bf2d..8286143 100644 --- a/cmd/agent/main.go +++ b/cmd/agent/main.go @@ -2,9 +2,13 @@ package main import ( "context" + "errors" + "fmt" + "io" "log/slog" "os" "os/signal" + "strings" "syscall" "time" @@ -13,6 +17,7 @@ import ( "github.com/NexGenCloud/hyperstack-agent/internal/config" "github.com/NexGenCloud/hyperstack-agent/internal/probes" "github.com/NexGenCloud/hyperstack-agent/internal/system" + "github.com/NexGenCloud/hyperstack-agent/internal/update" ) var ( @@ -21,12 +26,20 @@ var ( ) const ( + autoUpdateInterval = time.Hour + autoUpdateValidationTimeout = 2 * time.Minute + metricsConfigSyncInterval = time.Minute + metricsConfigSyncTimeout = 10 * time.Second startupMetadataInitialBackoff = 500 * time.Millisecond startupMetadataMaxBackoff = 1 * time.Minute defaultHealthAddr = "127.0.0.1:9100" ) func main() { + if handled, code := runDiagnosticCommand(os.Args[1:], os.Stdout); handled { + os.Exit(code) + } + // Configure log level from environment (HYPERSTACK_LOG_LEVEL or LOG_LEVEL) logLevel := slog.LevelInfo for _, key := range []string{"HYPERSTACK_LOG_LEVEL", "LOG_LEVEL"} { @@ -50,8 +63,11 @@ func main() { slog.Info("Hyperstack agent starting", "version", version) - ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + ctx, cancel := context.WithCancel(context.Background()) defer cancel() + signalCh := make(chan os.Signal, 1) + signal.Notify(signalCh, os.Interrupt, syscall.SIGTERM) + defer signal.Stop(signalCh) // Load configuration from environment cfg := config.Load() @@ -60,7 +76,7 @@ func main() { slog.Warn("security configuration warning", "code", warning.Code, "message", warning.Message) } - // Load startup metadata (uuid, infrahub_key, vm name, etc.) from metadata service once + // Load startup metadata (uuid, hyperstack key, vm name, etc.) from metadata service once // Retry with exponential backoff if metadata service is unavailable var meta system.StartupMetadata var err error @@ -101,13 +117,13 @@ func main() { // Hub client with API key from startup metadata. A KeyRefresher is // installed so that gateway 401 responses (e.g. after key rotation in - // infrahub) trigger a one-shot re-fetch from the metadata service + // Hyperstack) trigger a one-shot re-fetch from the metadata service // rather than burning the agent's retry budget on a stale credential. hub := client.NewHubClient(cfg.Hub.URL). WithPath(config.AgentPushPath) - hub.KeyRefresher = system.FetchInfrahubKey - if meta.InfrahubKey != "" { - hub.SetInfrahubKey(meta.InfrahubKey) + hub.KeyRefresher = system.FetchHyperstackKey + if meta.HyperstackKey != "" { + hub.SetHyperstackKey(meta.HyperstackKey) } var scheduled []collectors.ScheduledCollector @@ -170,23 +186,229 @@ func main() { submitLoopDone := hub.StartSubmitLoop(ctx) mgr := &collectors.Manager{Scheduled: scheduled} - if err := mgr.Run(ctx); err != nil && err != context.Canceled { - slog.Error("manager exited", "error", err) - os.Exit(1) + if meta.UUID == "" { + slog.Warn("metrics enablement sync disabled; instance uuid unavailable") + } else { + // Fix 6: do NOT pre-disable collectors before the first sync. The old + // behaviour was to always collect; defaulting to enabled until we receive + // an explicit false from the gateway preserves that contract. Starting + // disabled is risky: any first-sync failure (network blip, missing field, + // gateway rollout) would leave collectors permanently off. + go runMetricsEnabledSyncLoop(ctx, hub, mgr, meta.UUID, len(scheduled), metricsConfigSyncInterval) + } + + managerErrCh := make(chan error, 1) + go func() { + managerErrCh <- mgr.Run(ctx) + }() + + updateReadyCh := make(chan *update.Release, 1) + var executablePath string + var updater *update.Manager + currentPath, err := os.Executable() + if err != nil { + slog.Warn("self-update disabled; unable to resolve current executable", "error", err) + } else { + executablePath = currentPath + updateCheckURL := strings.TrimRight(cfg.Hub.URL, "/") + "/download" + updater = update.NewManager(updateCheckURL, version) + go runSelfUpdateLoop(ctx, updater, currentPath, autoUpdateInterval, updateReadyCh) } - // Wait for submit loop to exit before draining (prevents double-submission on shutdown) - slog.Info("collectors stopped, waiting for submit loop to exit") - <-submitLoopDone + var restartRelease *update.Release + externalShutdown := false + for { + select { + case sig := <-signalCh: + externalShutdown = true + slog.Info("shutdown signal received", "signal", sig) + cancel() + case err := <-managerErrCh: + if err != nil && !errors.Is(err, context.Canceled) { + slog.Error("manager exited", "error", err) + os.Exit(1) + } + + // Wait for submit loop to exit before draining (prevents double-submission on shutdown) + slog.Info("collectors stopped, waiting for submit loop to exit") + <-submitLoopDone + + // Collectors have exited and submit loop has stopped, now drain any remaining metrics + slog.Info("submit loop exited, draining pending metrics") + drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second) + if err := hub.DrainPending(drainCtx); err != nil { + slog.Warn("drain error", "error", err) + } + drainCancel() - // Collectors have exited and submit loop has stopped, now drain any remaining metrics - slog.Info("submit loop exited, draining pending metrics") - drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second) - defer drainCancel() + select { + case sig := <-signalCh: + externalShutdown = true + slog.Info("shutdown signal received", "signal", sig) + default: + } - if err := hub.DrainPending(drainCtx); err != nil { - slog.Warn("drain error", "error", err) + if restartRelease != nil && updater != nil && !externalShutdown { + if err := updater.PromoteRelease(executablePath, restartRelease); err != nil { + slog.Error("self-update promote failed", "version", restartRelease.Version, "error", err) + os.Exit(1) + } + slog.Info("restarting agent after self-update", "version", restartRelease.Version) + if err := update.RestartProcess(executablePath); err != nil { + slog.Error("self-update restart failed", "error", err) + os.Exit(1) + } + } else if restartRelease != nil { + _ = os.Remove(restartRelease.StagedPath) + if externalShutdown { + slog.Info("self-update skipped because shutdown was requested", "version", restartRelease.Version) + } + } + slog.Info("Hyperstack agent shutdown complete") + return + case release := <-updateReadyCh: + if release == nil || restartRelease != nil { + continue + } + restartRelease = release + slog.Info("self-update prepared; stopping collectors for restart", "version", release.Version) + cancel() + } + } +} + +func runMetricsEnabledSyncLoop( + ctx context.Context, + hub *client.HubClient, + manager *collectors.Manager, + uuid string, + collectorCount int, + interval time.Duration, +) { + if interval <= 0 { + interval = metricsConfigSyncInterval } - slog.Info("Hyperstack agent shutdown complete") + lastEnabled := true + hasLastEnabled := false + sync := func() { + syncCtx, cancel := context.WithTimeout(ctx, metricsConfigSyncTimeout) + metadata, err := hub.GetMetadata(syncCtx, uuid) + cancel() + if err != nil { + if ctx.Err() != nil { + return + } + slog.Warn("metrics enablement sync failed", "uuid", uuid, "error", err) + return + } + + // Fix 7a: MetricsEnabled is *bool; nil means the field was absent from + // the response (gateway rollout, mismatch). Treat nil as default-enabled + // so a missing field never silently turns off all collectors. + enabled := metadata.MetricsEnabled == nil || *metadata.MetricsEnabled + manager.SetEnabled(enabled) + if enabled { + hub.SetCollectorsRunning(int64(collectorCount)) + } else { + hub.SetCollectorsRunning(0) + } + + if !hasLastEnabled || lastEnabled != enabled { + if enabled { + slog.Info("metrics enabled; collectors resumed", "uuid", uuid) + } else { + slog.Info("metrics disabled; collectors sleeping", "uuid", uuid) + } + lastEnabled = enabled + hasLastEnabled = true + } + } + + sync() + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + sync() + } + } +} + +func runDiagnosticCommand(args []string, stdout io.Writer) (bool, int) { + if len(args) == 0 { + return false, 0 + } + + switch args[0] { + case "version", "--version", "-version": + _, _ = fmt.Fprintln(stdout, version) + return true, 0 + case "diagnose": + if len(args) == 2 && args[1] == "status" { + _, _ = fmt.Fprintf(stdout, "status=ok version=%s date=%s\n", version, date) + return true, 0 + } + _, _ = fmt.Fprintln(stdout, "usage: hyperstack-agent diagnose status") + return true, 2 + default: + return false, 0 + } +} + +func runSelfUpdateLoop( + ctx context.Context, + updater *update.Manager, + currentPath string, + interval time.Duration, + updateReadyCh chan<- *update.Release, +) { + check := func() { + release, err := updater.Check(ctx) + if err != nil { + if ctx.Err() != nil { + return + } + slog.Warn("self-update check failed", "url", updater.CheckURL, "error", err) + return + } + if release == nil { + return + } + slog.Info("new agent version available", "current_version", updater.CurrentVersion, "available_version", release.Version) + validationCtx, validationCancel := context.WithTimeout(context.Background(), autoUpdateValidationTimeout) + err = updater.DownloadRelease(validationCtx, release, currentPath) + validationCancel() + if err != nil { + if ctx.Err() != nil { + return + } + slog.Warn("self-update download failed", "version", release.Version, "error", err) + return + } + if ctx.Err() != nil { + _ = os.Remove(release.StagedPath) + return + } + select { + case updateReadyCh <- release: + default: + _ = os.Remove(release.StagedPath) + } + } + + check() + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + check() + } + } } diff --git a/cmd/agent/main_test.go b/cmd/agent/main_test.go new file mode 100644 index 0000000..96d9f25 --- /dev/null +++ b/cmd/agent/main_test.go @@ -0,0 +1,76 @@ +package main + +import ( + "bytes" + "strings" + "testing" +) + +func TestRunDiagnosticCommandVersion(t *testing.T) { + oldVersion := version + version = "1.2.3" + defer func() { version = oldVersion }() + + var out bytes.Buffer + handled, code := runDiagnosticCommand([]string{"version"}, &out) + if !handled { + t.Fatal("runDiagnosticCommand handled = false, want true") + } + if code != 0 { + t.Fatalf("runDiagnosticCommand code = %d, want 0", code) + } + if strings.TrimSpace(out.String()) != "1.2.3" { + t.Fatalf("output = %q, want version", out.String()) + } +} + +func TestRunDiagnosticCommandStatus(t *testing.T) { + oldVersion := version + oldDate := date + version = "2.0.0" + date = "2026-06-16T00:00:00Z" + defer func() { + version = oldVersion + date = oldDate + }() + + var out bytes.Buffer + handled, code := runDiagnosticCommand([]string{"diagnose", "status"}, &out) + if !handled { + t.Fatal("runDiagnosticCommand handled = false, want true") + } + if code != 0 { + t.Fatalf("runDiagnosticCommand code = %d, want 0", code) + } + if got := out.String(); !strings.Contains(got, "status=ok") || !strings.Contains(got, "version=2.0.0") { + t.Fatalf("output = %q, want status and version", got) + } +} + +func TestRunDiagnosticCommandInvalidDiagnoseUsage(t *testing.T) { + var out bytes.Buffer + handled, code := runDiagnosticCommand([]string{"diagnose"}, &out) + if !handled { + t.Fatal("runDiagnosticCommand handled = false, want true") + } + if code != 2 { + t.Fatalf("runDiagnosticCommand code = %d, want 2", code) + } + if !strings.Contains(out.String(), "usage:") { + t.Fatalf("output = %q, want usage", out.String()) + } +} + +func TestRunDiagnosticCommandUnknownCommand(t *testing.T) { + var out bytes.Buffer + handled, code := runDiagnosticCommand([]string{"run"}, &out) + if handled { + t.Fatal("runDiagnosticCommand handled = true, want false") + } + if code != 0 { + t.Fatalf("runDiagnosticCommand code = %d, want 0", code) + } + if out.Len() != 0 { + t.Fatalf("output = %q, want empty", out.String()) + } +} diff --git a/internal/client/hub.go b/internal/client/hub.go index 9f82b42..2e50343 100644 --- a/internal/client/hub.go +++ b/internal/client/hub.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "log/slog" @@ -30,7 +31,9 @@ const ( maxErrorBodyBytes = 4096 ) -// KeyRefresher returns a fresh raw infrahub key (no "VM " prefix) by re-fetching +var ErrMetricsDisabled = errors.New("metrics disabled for virtual machine") + +// KeyRefresher returns a fresh raw Hyperstack key (no "VM " prefix) by re-fetching // it from the source of truth (typically the OpenStack metadata service). // Implementations should be safe to call concurrently. type KeyRefresher func(ctx context.Context) (string, error) @@ -46,9 +49,9 @@ type HubClient struct { apiKeyMu sync.RWMutex APIKey string - // rawInfrahubKey caches the un-prefixed key so we can detect when a + // rawHyperstackKey caches the un-prefixed key so we can detect when a // refresh returned the same (still-bad) credential and avoid hot loops. - rawInfrahubKey string + rawHyperstackKey string // KeyRefresher, when non-nil, is invoked on a 401 response. If it // returns a new key, the request is retried once without consuming a @@ -109,13 +112,13 @@ func (h *HubClient) WithPath(path string) *HubClient { return h } -// SetInfrahubKey stores the raw infrahub key and updates the header value +// SetHyperstackKey stores the raw Hyperstack key and updates the header value // (prefixed with "VM ") used for outgoing requests. Safe for concurrent use. // Passing an empty string clears the key. -func (h *HubClient) SetInfrahubKey(rawKey string) { +func (h *HubClient) SetHyperstackKey(rawKey string) { h.apiKeyMu.Lock() defer h.apiKeyMu.Unlock() - h.rawInfrahubKey = rawKey + h.rawHyperstackKey = rawKey if rawKey == "" { h.APIKey = "" return @@ -130,18 +133,18 @@ func (h *HubClient) getAPIKey() string { return h.APIKey } -// getRawInfrahubKey returns the most recently stored raw key. -func (h *HubClient) getRawInfrahubKey() string { +// getRawHyperstackKey returns the most recently stored raw key. +func (h *HubClient) getRawHyperstackKey() string { h.apiKeyMu.RLock() defer h.apiKeyMu.RUnlock() - return h.rawInfrahubKey + return h.rawHyperstackKey } -// refreshInfrahubKey invokes KeyRefresher with cooldown + single-flight +// refreshHyperstackKey invokes KeyRefresher with cooldown + single-flight // semantics. It returns true when the stored key changed as a result of the // call, false otherwise (no refresher configured, cooldown active, refresher // returned the same key, or refresher errored). -func (h *HubClient) refreshInfrahubKey(ctx context.Context, observedRaw string) bool { +func (h *HubClient) refreshHyperstackKey(ctx context.Context, observedRaw string) bool { if h.KeyRefresher == nil { return false } @@ -150,7 +153,7 @@ func (h *HubClient) refreshInfrahubKey(ctx context.Context, observedRaw string) // If another goroutine already refreshed since we observed the bad key, // adopt its result instead of calling the refresher again. - if current := h.getRawInfrahubKey(); current != "" && current != observedRaw { + if current := h.getRawHyperstackKey(); current != "" && current != observedRaw { return true } @@ -163,22 +166,54 @@ func (h *HubClient) refreshInfrahubKey(ctx context.Context, observedRaw string) newKey, err := h.KeyRefresher(ctx) h.lastRefreshAt = time.Now() if err != nil { - slog.Warn("infrahub key refresh failed", "error", err) + slog.Warn("Hyperstack key refresh failed", "error", err) return false } if newKey == "" { - slog.Warn("infrahub key refresh returned empty key") + slog.Warn("Hyperstack key refresh returned empty key") return false } if newKey == observedRaw { - slog.Warn("infrahub key refresh returned the same key; not retrying") + slog.Warn("Hyperstack key refresh returned the same key; not retrying") return false } - h.SetInfrahubKey(newKey) - slog.Info("infrahub key refreshed after 401") + h.SetHyperstackKey(newKey) + slog.Info("Hyperstack key refreshed after 401") return true } +// doWithKeyRefresh executes an HTTP request and, if the response is 401 and a +// key refresh succeeds, rebuilds and retries the request exactly once. +// buildReq is called up to twice so that POST bodies (which are not replayable +// from a one-shot reader) can be reconstructed fresh for the retry. +func (h *HubClient) doWithKeyRefresh( + ctx context.Context, + buildReq func() (*http.Request, error), + operation string, +) (*http.Response, error) { + req, err := buildReq() + if err != nil { + return nil, err + } + observedRawKey := h.getRawHyperstackKey() + resp, err := h.HTTP.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode == http.StatusUnauthorized && h.refreshHyperstackKey(ctx, observedRawKey) { + closeResponseBody(resp, operation+" 401") + req, err = buildReq() + if err != nil { + return nil, err + } + resp, err = h.HTTP.Do(req) + if err != nil { + return nil, err + } + } + return resp, nil +} + func (h *HubClient) SetCollectorsRunning(count int64) { atomic.StoreInt64(&h.collectorsRunning, count) } @@ -367,6 +402,15 @@ func (h *HubClient) runSubmitLoop(ctx context.Context) { } slog.Info("hub batch submitted", "series", batchSize, "duration_ms", int(submitDuration/time.Millisecond)) + } else if errors.Is(err, ErrMetricsDisabled) { + h.pendingMu.Lock() + h.pending = h.pending[batchSize:] + h.pendingMu.Unlock() + + successesSinceFail = 0 + retryAttempts = 0 + + slog.Info("hub batch dropped; metrics disabled for vm", "size", batchSize, "error", err) } else { // Failure: increment failed counter (once per batch, not per retry attempt), // calculate backoff, and retry next iteration. Batch stays in pending queue automatically. @@ -535,37 +579,21 @@ func (h *HubClient) submitBatch(ctx context.Context, batch []metrics.Measure) (t slog.Info("hub submit payload", "bytes", len(body)) } - url := fmt.Sprintf("%s/%s", h.BaseURL, h.Path) - - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) - if err != nil { - return 0, fmt.Errorf("create request: %w", err) - } - - httpReq.Header.Set("Content-Type", "application/json") - if apiKey := h.getAPIKey(); apiKey != "" { - httpReq.Header.Set("api-key", apiKey) - } + reqURL := fmt.Sprintf("%s/%s", h.BaseURL, h.Path) - observedRawKey := h.getRawInfrahubKey() - resp, err := h.HTTP.Do(httpReq) - if err != nil { - return 0, fmt.Errorf("http error: %w", err) - } - if resp.StatusCode == http.StatusUnauthorized && h.refreshInfrahubKey(ctx, observedRawKey) { - closeResponseBody(resp, "submit batch 401") - httpReq, err = http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + resp, err := h.doWithKeyRefresh(ctx, func() (*http.Request, error) { + r, err := http.NewRequestWithContext(ctx, http.MethodPost, reqURL, bytes.NewReader(body)) if err != nil { - return 0, fmt.Errorf("create request: %w", err) + return nil, fmt.Errorf("create request: %w", err) } - httpReq.Header.Set("Content-Type", "application/json") + r.Header.Set("Content-Type", "application/json") if apiKey := h.getAPIKey(); apiKey != "" { - httpReq.Header.Set("api-key", apiKey) - } - resp, err = h.HTTP.Do(httpReq) - if err != nil { - return 0, fmt.Errorf("http error after key refresh: %w", err) + r.Header.Set("api-key", apiKey) } + return r, nil + }, "submit batch") + if err != nil { + return 0, fmt.Errorf("http error: %w", err) } defer closeResponseBody(resp, "submit batch") @@ -575,14 +603,19 @@ func (h *HubClient) submitBatch(ctx context.Context, batch []metrics.Measure) (t return 0, nil } + respBody, readErr := readErrorBody(resp.Body) + + if resp.StatusCode == http.StatusPreconditionFailed { + return 0, fmt.Errorf("%w: %s", ErrMetricsDisabled, respBody) + } + // Failure: extract Retry-After if present retryAfterHeader := resp.Header.Get("Retry-After") retryAfter, _ := parseRetryAfter(retryAfterHeader) retryAfter = capRetryAfter(retryAfter) - respBody, err := readErrorBody(resp.Body) - if err != nil { - return retryAfter, fmt.Errorf("http %d; read response body: %w", resp.StatusCode, err) + if readErr != nil { + return retryAfter, fmt.Errorf("http %d; read response body: %w", resp.StatusCode, readErr) } return retryAfter, fmt.Errorf("http %d: %s", resp.StatusCode, respBody) } @@ -603,7 +636,7 @@ func (h *HubClient) Submit(ctx context.Context, path string, payload any) error usePath = h.Path } - url := fmt.Sprintf("%s/%s", h.BaseURL, usePath) + reqURL := fmt.Sprintf("%s/%s", h.BaseURL, usePath) // Use PATCH for metadata, POST for metrics method := http.MethodPost @@ -611,35 +644,19 @@ func (h *HubClient) Submit(ctx context.Context, path string, payload any) error method = http.MethodPatch } - httpReq, err := http.NewRequestWithContext(ctx, method, url, bytes.NewReader(b)) - if err != nil { - return err - } - - httpReq.Header.Set("Content-Type", "application/json") - if apiKey := h.getAPIKey(); apiKey != "" { - httpReq.Header.Set("api-key", apiKey) - } - - observedRawKey := h.getRawInfrahubKey() - resp, err := h.HTTP.Do(httpReq) - if err != nil { - return err - } - if resp.StatusCode == http.StatusUnauthorized && h.refreshInfrahubKey(ctx, observedRawKey) { - closeResponseBody(resp, "submit metadata 401") - httpReq, err = http.NewRequestWithContext(ctx, method, url, bytes.NewReader(b)) + resp, err := h.doWithKeyRefresh(ctx, func() (*http.Request, error) { + r, err := http.NewRequestWithContext(ctx, method, reqURL, bytes.NewReader(b)) if err != nil { - return err + return nil, err } - httpReq.Header.Set("Content-Type", "application/json") + r.Header.Set("Content-Type", "application/json") if apiKey := h.getAPIKey(); apiKey != "" { - httpReq.Header.Set("api-key", apiKey) - } - resp, err = h.HTTP.Do(httpReq) - if err != nil { - return err + r.Header.Set("api-key", apiKey) } + return r, nil + }, "submit metadata") + if err != nil { + return err } defer closeResponseBody(resp, "submit metadata") @@ -691,3 +708,52 @@ func closeResponseBody(resp *http.Response, operation string) { slog.Debug("response body close failed", "operation", operation, "error", err) } } + +// VMMetadata holds VM-level configuration returned by the gateway. +// +// Fix 7a: MetricsEnabled is a *bool rather than a plain bool so that a missing +// JSON field is distinguishable from an explicit false. A nil value means the +// gateway did not set the field (e.g. during a rollout or a response mismatch) +// and callers should treat it as "default-enabled". A plain bool would make a +// dropped field silently disable all collectors. +type VMMetadata struct { + MetricsEnabled *bool `json:"metrics_enabled"` +} + +func (h *HubClient) GetMetadata(ctx context.Context, uuid string) (*VMMetadata, error) { + escapedUUID := url.PathEscape(strings.TrimSpace(uuid)) + if escapedUUID == "" { + return nil, errors.New("metadata uuid is required") + } + + reqURL := fmt.Sprintf("%s/api/v1/metadata/%s", strings.TrimRight(h.BaseURL, "/"), escapedUUID) + + // Fix 7b: mirror the refresh-and-retry pattern used by SubmitBatch and + // Submit so that a rotated Hyperstack key does not leave the sync loop stuck + // on 401 and permanently unable to re-enable collection. + resp, err := h.doWithKeyRefresh(ctx, func() (*http.Request, error) { + r, err := http.NewRequestWithContext(ctx, http.MethodGet, reqURL, nil) + if err != nil { + return nil, err + } + if apiKey := h.getAPIKey(); apiKey != "" { + r.Header.Set("api-key", apiKey) + } + return r, nil + }, "get metadata") + if err != nil { + return nil, err + } + defer closeResponseBody(resp, "get metadata") + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + respBody, _ := readErrorBody(resp.Body) + return nil, fmt.Errorf("metadata status %d: %s", resp.StatusCode, respBody) + } + + var metadata VMMetadata + if err := json.NewDecoder(resp.Body).Decode(&metadata); err != nil { + return nil, err + } + return &metadata, nil +} diff --git a/internal/client/hub_test.go b/internal/client/hub_test.go index 2b20930..8f92b14 100644 --- a/internal/client/hub_test.go +++ b/internal/client/hub_test.go @@ -345,7 +345,7 @@ func TestSubmitBatch_SingleAttemptNoRetry(t *testing.T) { } } -func TestSubmitBatch_RefreshesInfrahubKeyOnUnauthorized(t *testing.T) { +func TestSubmitBatch_RefreshesHyperstackKeyOnUnauthorized(t *testing.T) { var attempts atomic.Int32 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { attempt := attempts.Add(1) @@ -373,7 +373,7 @@ func TestSubmitBatch_RefreshesInfrahubKeyOnUnauthorized(t *testing.T) { hc := NewHubClient(srv.URL) hc.VMName = "test-vm" hc.InstanceUUID = "test-uuid" - hc.SetInfrahubKey("stale-key") + hc.SetHyperstackKey("stale-key") hc.KeyRefresher = func(ctx context.Context) (string, error) { refreshes.Add(1) return "fresh-key", nil @@ -401,7 +401,7 @@ func TestSubmitBatch_RefreshesInfrahubKeyOnUnauthorized(t *testing.T) { } } -func TestSubmit_RefreshesInfrahubKeyOnUnauthorized(t *testing.T) { +func TestSubmit_RefreshesHyperstackKeyOnUnauthorized(t *testing.T) { var attempts atomic.Int32 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPatch { @@ -430,7 +430,7 @@ func TestSubmit_RefreshesInfrahubKeyOnUnauthorized(t *testing.T) { var refreshes atomic.Int32 hc := NewHubClient(srv.URL) - hc.SetInfrahubKey("stale-key") + hc.SetHyperstackKey("stale-key") hc.KeyRefresher = func(ctx context.Context) (string, error) { refreshes.Add(1) return "fresh-key", nil @@ -470,7 +470,7 @@ func TestSubmitBatch_StripsAPIKeyOnCrossHostRedirect(t *testing.T) { hc := NewHubClient(redirector.URL) hc.VMName = "test-vm" hc.InstanceUUID = "test-uuid" - hc.SetInfrahubKey("secret-key") + hc.SetHyperstackKey("secret-key") ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() @@ -940,3 +940,100 @@ func TestDeepCopyMeasures(t *testing.T) { t.Fatalf("copy didn't update: key=%q, want 'modified'", copied[0].Labels["key"]) } } + +// Fix 7a + 7b: GetMetadata tests — *bool semantics and 401 retry. + +// TestGetMetadata_ExplicitTrue verifies that metrics_enabled:true is decoded correctly. +func TestGetMetadata_ExplicitTrue(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintln(w, `{"metrics_enabled":true}`) + })) + defer server.Close() + + h := NewHubClient(server.URL) + meta, err := h.GetMetadata(context.Background(), "test-uuid") + if err != nil { + t.Fatalf("GetMetadata() error = %v", err) + } + if meta.MetricsEnabled == nil { + t.Fatal("MetricsEnabled = nil, want non-nil") + } + if !*meta.MetricsEnabled { + t.Fatal("MetricsEnabled = false, want true") + } +} + +// TestGetMetadata_ExplicitFalse verifies that metrics_enabled:false is decoded correctly. +func TestGetMetadata_ExplicitFalse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintln(w, `{"metrics_enabled":false}`) + })) + defer server.Close() + + h := NewHubClient(server.URL) + meta, err := h.GetMetadata(context.Background(), "test-uuid") + if err != nil { + t.Fatalf("GetMetadata() error = %v", err) + } + if meta.MetricsEnabled == nil { + t.Fatal("MetricsEnabled = nil, want non-nil") + } + if *meta.MetricsEnabled { + t.Fatal("MetricsEnabled = true, want false") + } +} + +// TestGetMetadata_MissingFieldIsNil verifies that a response without +// metrics_enabled leaves MetricsEnabled as nil (distinguishable from false). +func TestGetMetadata_MissingFieldIsNil(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintln(w, `{}`) + })) + defer server.Close() + + h := NewHubClient(server.URL) + meta, err := h.GetMetadata(context.Background(), "test-uuid") + if err != nil { + t.Fatalf("GetMetadata() error = %v", err) + } + if meta.MetricsEnabled != nil { + t.Fatalf("MetricsEnabled = %v, want nil for absent field", *meta.MetricsEnabled) + } +} + +// TestGetMetadata_RefreshesHyperstackKeyOnUnauthorized verifies that GetMetadata +// retries the request after a 401 triggers a successful key refresh +// (mirrors the same test pattern used for SubmitBatch and Submit). +func TestGetMetadata_RefreshesHyperstackKeyOnUnauthorized(t *testing.T) { + callCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + callCount++ + if callCount == 1 { + w.WriteHeader(http.StatusUnauthorized) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintln(w, `{"metrics_enabled":true}`) + })) + defer server.Close() + + h := NewHubClient(server.URL) + h.SetHyperstackKey("old-key") + h.KeyRefresher = func(_ context.Context) (string, error) { + return "new-key", nil + } + + meta, err := h.GetMetadata(context.Background(), "test-uuid") + if err != nil { + t.Fatalf("GetMetadata() error = %v, want success after key refresh", err) + } + if meta == nil || meta.MetricsEnabled == nil || !*meta.MetricsEnabled { + t.Fatal("GetMetadata() returned unexpected metadata after key refresh") + } + if callCount != 2 { + t.Fatalf("server callCount = %d, want 2 (one 401 + one retry)", callCount) + } +} diff --git a/internal/collectors/manager.go b/internal/collectors/manager.go index a73131a..a08a4cb 100644 --- a/internal/collectors/manager.go +++ b/internal/collectors/manager.go @@ -3,6 +3,7 @@ package collectors import ( "context" "log/slog" + "sync/atomic" "time" "github.com/NexGenCloud/hyperstack-agent/internal/jitter" @@ -21,6 +22,15 @@ type Manager struct { Scheduled []ScheduledCollector // JitterFraction, if >0, adds up to Interval*JitterFraction random delay per tick JitterFraction float64 + disabled atomic.Bool +} + +func (m *Manager) SetEnabled(enabled bool) { + m.disabled.Store(!enabled) +} + +func (m *Manager) IsEnabled() bool { + return !m.disabled.Load() } func (m *Manager) Run(ctx context.Context) error { @@ -68,6 +78,10 @@ func (m *Manager) Run(ctx context.Context) error { } start := time.Now() slog.Debug("collector tick", "collector", collectorName, "jitter_ms", int(perTickJitter/time.Millisecond)) + if !m.IsEnabled() { + slog.Debug("collector skipped; metrics disabled", "collector", collectorName) + continue + } if err := scLocal.Collector.Run(ctx); err != nil && ctx.Err() == nil { slog.Error("collector run error", "collector", collectorName, "error", err) } else { diff --git a/internal/collectors/manager_test.go b/internal/collectors/manager_test.go new file mode 100644 index 0000000..888f9fa --- /dev/null +++ b/internal/collectors/manager_test.go @@ -0,0 +1,78 @@ +package collectors + +import ( + "context" + "sync/atomic" + "testing" + "time" +) + +type countingCollector struct { + count atomic.Int64 +} + +func (c *countingCollector) Run(context.Context) error { + c.count.Add(1) + return nil +} + +func TestManagerSkipsCollectorsWhenDisabled(t *testing.T) { + collector := &countingCollector{} + manager := &Manager{ + Scheduled: []ScheduledCollector{ + {Collector: collector, Interval: 5 * time.Millisecond}, + }, + JitterFraction: 0.01, + } + manager.SetEnabled(false) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- manager.Run(ctx) + }() + + time.Sleep(40 * time.Millisecond) + cancel() + <-done + + if got := collector.count.Load(); got != 0 { + t.Fatalf("collector runs = %d, want 0", got) + } +} + +func TestManagerRunsCollectorsWhenReenabled(t *testing.T) { + collector := &countingCollector{} + manager := &Manager{ + Scheduled: []ScheduledCollector{ + {Collector: collector, Interval: 5 * time.Millisecond}, + }, + JitterFraction: 0.01, + } + manager.SetEnabled(false) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- manager.Run(ctx) + }() + + time.Sleep(20 * time.Millisecond) + if got := collector.count.Load(); got != 0 { + cancel() + <-done + t.Fatalf("collector runs while disabled = %d, want 0", got) + } + + manager.SetEnabled(true) + deadline := time.Now().Add(100 * time.Millisecond) + for time.Now().Before(deadline) && collector.count.Load() == 0 { + time.Sleep(5 * time.Millisecond) + } + cancel() + <-done + + if got := collector.count.Load(); got == 0 { + t.Fatal("collector runs = 0, want at least one run after re-enable") + } +} diff --git a/internal/system/metadata.go b/internal/system/metadata.go index e909bea..49dca4b 100644 --- a/internal/system/metadata.go +++ b/internal/system/metadata.go @@ -304,18 +304,18 @@ func metadataMetaString(data map[string]any, key string) string { // StartupMetadata holds essential identity and auth data loaded once at startup. type StartupMetadata struct { - UUID string - InfrahubKey string - VMName string - Cluster string - Role string + UUID string + HyperstackKey string + VMName string + Cluster string + Role string } -// FetchInfrahubKey re-fetches the metadata from HTTP (skipping the cloud-init file) +// FetchHyperstackKey re-fetches the metadata from HTTP (skipping the cloud-init file) // and returns the `meta.infrahub_key` value. Used to recover from gateway 401s when // the originally-cached key has been rotated. The returned key is the raw value // (no "VM " prefix). -func FetchInfrahubKey(ctx context.Context) (string, error) { +func FetchHyperstackKey(ctx context.Context) (string, error) { // Honor caller cancellation by running the (synchronous) fetch in a // goroutine and selecting on ctx.Done(). The underlying fetch already // imposes its own per-request timeout. @@ -352,11 +352,11 @@ func LoadStartupMetadata() (StartupMetadata, error) { } meta := StartupMetadata{ - UUID: metadataString(metadata, "uuid"), - InfrahubKey: metadataMetaString(metadata, "infrahub_key"), - VMName: metadataString(metadata, "name"), - Cluster: metadataMetaString(metadata, "cluster"), - Role: metadataMetaString(metadata, "role"), + UUID: metadataString(metadata, "uuid"), + HyperstackKey: metadataMetaString(metadata, "infrahub_key"), + VMName: metadataString(metadata, "name"), + Cluster: metadataMetaString(metadata, "cluster"), + Role: metadataMetaString(metadata, "role"), } if meta.VMName == "" { diff --git a/internal/update/manager.go b/internal/update/manager.go new file mode 100644 index 0000000..54e66ec --- /dev/null +++ b/internal/update/manager.go @@ -0,0 +1,445 @@ +package update + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "hash" + "io" + "log/slog" + "net/http" + "net/url" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "sync" + "syscall" + "time" +) + +const ( + VersionHeaderName = "Hyperstack-Agent-Version" + DigestHeaderName = "Hyperstack-Agent-Digest" + maxBinarySize = 256 * 1024 * 1024 +) + +type Release struct { + Version string + Digest string + DownloadURL string + StagedPath string +} + +type Manager struct { + CheckURL string + CurrentVersion string + Client *http.Client + downloadMu sync.Mutex +} + +func NewManager(checkURL, currentVersion string) *Manager { + return &Manager{ + CheckURL: checkURL, + CurrentVersion: currentVersion, + Client: &http.Client{ + Timeout: 15 * time.Second, + CheckRedirect: func(req *http.Request, via []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + } +} + +func (m *Manager) Check(ctx context.Context) (*Release, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodHead, m.CheckURL, nil) + if err != nil { + return nil, err + } + + resp, err := m.Client.Do(req) + if err != nil { + return nil, err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode < 200 || resp.StatusCode >= 400 { + return nil, fmt.Errorf("update check returned status %s", resp.Status) + } + + version := strings.TrimSpace(resp.Header.Get(VersionHeaderName)) + if version == "" { + return nil, errors.New("update check missing version header") + } + + // Fix 5: treat non-semver current version (e.g. "dev", git SHA) as unknown — + // skip the update silently rather than spamming parse errors every check cycle. + shouldUpdate, err := isHigherVersion(version, m.CurrentVersion) + if err != nil { + return nil, err + } + if !shouldUpdate { + return nil, nil + } + + downloadURL := strings.TrimSpace(resp.Header.Get("Location")) + if downloadURL == "" { + downloadURL = m.CheckURL + } + if resolved, err := resolveURL(m.CheckURL, downloadURL); err == nil { + downloadURL = resolved + } + + // Fix 1: require the digest header before returning a release — an empty + // digest must never be silently passed through, since DownloadRelease would + // then promote the binary without any integrity verification. + digest := strings.TrimSpace(resp.Header.Get(DigestHeaderName)) + if digest == "" { + return nil, errors.New("update check missing digest header") + } + + return &Release{ + Version: version, + Digest: digest, + DownloadURL: downloadURL, + }, nil +} + +func (m *Manager) DownloadRelease(ctx context.Context, release *Release, currentPath string) error { + if release == nil { + return errors.New("release is required") + } + m.downloadMu.Lock() + defer m.downloadMu.Unlock() + + release.StagedPath = "" + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, release.DownloadURL, nil) + if err != nil { + return err + } + + resp, err := m.Client.Do(req) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("binary download returned status %s", resp.Status) + } + + // Fix 3: stage in os.TempDir() rather than next to the current executable. + // The directory containing the installed binary is read-only under systemd + // (ProtectSystem=strict) and in the Docker runtime (non-root). os.TempDir() + // is always writable. atomicSwap handles the cross-device rename fallback. + // + // Use os.CreateTemp so the file is created atomically with an unguessable + // name, eliminating the Remove→create TOCTOU window that a predictable fixed + // name in shared /tmp would expose (symlink replacement attack). + f, err := os.CreateTemp(os.TempDir(), "."+filepath.Base(currentPath)+"-*.tmp") /* #nosec G302 -- binary must be world-readable/executable */ + if err != nil { + return err + } + tmpPath := f.Name() + + copyErr := func() error { + if err := copyWithLimit(f, resp.Body, maxBinarySize); err != nil { + return err + } + if err := f.Sync(); err != nil { + return err + } + return nil + }() + closeErr := f.Close() + if copyErr != nil { + _ = os.Remove(tmpPath) + return copyErr + } + if closeErr != nil { + _ = os.Remove(tmpPath) + return closeErr + } + + if err := os.Chmod(tmpPath, 0o755); err != nil { /* #nosec G302 -- binary must be world-readable/executable */ + _ = os.Remove(tmpPath) + return err + } + + if err := verifyDigest(tmpPath, release.Digest); err != nil { + _ = os.Remove(tmpPath) + return err + } + + if err := smokeTestBinary(tmpPath); err != nil { + _ = os.Remove(tmpPath) + return err + } + + release.StagedPath = tmpPath + return nil +} + +func copyWithLimit(dst io.Writer, src io.Reader, maxBytes int64) error { + n, err := io.Copy(dst, io.LimitReader(src, maxBytes+1)) + if err != nil { + return err + } + if n > maxBytes { + return fmt.Errorf("binary download exceeds max size %d bytes", maxBytes) + } + return nil +} + +func (m *Manager) PromoteRelease(currentPath string, release *Release) error { + if release == nil { + return errors.New("release is required") + } + if strings.TrimSpace(release.StagedPath) == "" { + return errors.New("release staged path is required") + } + + if err := atomicSwap(currentPath, release.StagedPath); err != nil { + return err + } + + m.CurrentVersion = release.Version + return nil +} + +// resolveExecPath resolves symlinks and verifies the result is an absolute path. +// Both RestartProcess and smokeTestBinary call this before any exec so that +// the executed path is a concrete, fully-resolved value rather than a raw +// variable — which satisfies gosec G204 without needing a suppression annotation. +func resolveExecPath(p string) (string, error) { + resolved, err := filepath.EvalSymlinks(p) + if err != nil { + return "", fmt.Errorf("resolve exec path: %w", err) + } + if !filepath.IsAbs(resolved) { + return "", fmt.Errorf("exec path is not absolute: %s", resolved) + } + return resolved, nil +} + +// RestartProcess replaces the current process image with the binary at +// currentPath using a clean exec(2). The path is resolved through symlinks +// before the exec so the kernel receives a concrete, absolute path. +func RestartProcess(currentPath string) error { + resolved, err := resolveExecPath(currentPath) + if err != nil { + return err + } + return syscall.Exec(resolved, os.Args, os.Environ()) /* #nosec G204 G702 -- resolved is the symlink-evaluated, absolute-asserted current executable path */ +} + +func atomicSwap(currentPath, newPath string) error { + backup := currentPath + ".bak" + _ = os.Remove(backup) + if err := copyFile(currentPath, backup); err != nil { + return err + } + // os.Rename is atomic on the same filesystem. When the staged binary lives + // in os.TempDir() and the install dir is on a different device (common in + // Docker / systemd setups), Rename returns EXDEV. Fall back to a copy+remove + // so the promote step still succeeds across filesystem boundaries. + if err := os.Rename(newPath, currentPath); err != nil { + // Cross-device rename (EXDEV) fallback: copyFile opens the destination + // with O_TRUNC, so currentPath is zeroed the moment the copy begins. + // If the copy fails, restore from the backup made above to avoid leaving + // the agent with a truncated (unlaunchable) binary. + if err2 := copyFile(newPath, currentPath); err2 != nil { + _ = copyFile(backup, currentPath) // best-effort restore + _ = os.Remove(newPath) + return err2 + } + _ = os.Remove(newPath) + } + return nil +} + +// Fix 2: verifyDigest now fails closed — an empty digest is treated as an +// error rather than silently skipping verification. This prevents a stripped +// or missing Hyperstack-Agent-Digest header from allowing an unverified binary +// through (e.g. a bad gateway response that drops the header). +func verifyDigest(path, digest string) error { + digest = strings.TrimSpace(digest) + if digest == "" { + return errors.New("digest is required for binary verification") + } + + algorithm, expected, ok := strings.Cut(digest, ":") + if !ok || strings.TrimSpace(expected) == "" { + return fmt.Errorf("invalid digest %q", digest) + } + if strings.ToLower(strings.TrimSpace(algorithm)) != "sha256" { + return fmt.Errorf("unsupported digest algorithm %q", algorithm) + } + + actual, err := fileDigest(path, sha256.New()) + if err != nil { + return err + } + if !strings.EqualFold(actual, strings.TrimSpace(expected)) { + return fmt.Errorf("digest mismatch: got sha256:%s, want %s", actual, digest) + } + return nil +} + +func fileDigest(path string, h hash.Hash) (string, error) { + root, err := os.OpenRoot(filepath.Dir(path)) + if err != nil { + return "", err + } + defer func() { _ = root.Close() }() + f, err := root.Open(filepath.Base(path)) + if err != nil { + return "", err + } + defer func() { _ = f.Close() }() + + if _, err := io.Copy(h, f); err != nil { + return "", err + } + return hex.EncodeToString(h.Sum(nil)), nil +} + +func smokeTestBinary(binaryPath string) error { + // Fix 4: resolve the path through symlinks before exec so gosec G204 sees a + // concrete, absolute path rather than a raw variable. + resolved, err := resolveExecPath(binaryPath) + if err != nil { + return err + } + + smokeCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + cmd := exec.CommandContext(smokeCtx, resolved, "diagnose", "status") /* #nosec G204 G702 -- resolved is the symlink-evaluated, absolute-asserted staged binary path */ + cmd.Env = os.Environ() + output, err := cmd.CombinedOutput() + if smokeCtx.Err() == context.DeadlineExceeded { + return errors.New("smoke test timed out") + } + if err == nil { + return nil + } + + var exitErr *exec.ExitError + if errors.As(err, &exitErr) { + return fmt.Errorf("smoke test exit code %d: %s", exitErr.ExitCode(), strings.TrimSpace(string(output))) + } + + return err +} + +func copyFile(src, dst string) error { + srcRoot, err := os.OpenRoot(filepath.Dir(src)) + if err != nil { + return err + } + defer func() { _ = srcRoot.Close() }() + in, err := srcRoot.Open(filepath.Base(src)) + if err != nil { + return err + } + defer func() { _ = in.Close() }() + + info, err := in.Stat() + if err != nil { + return err + } + + dstRoot, err := os.OpenRoot(filepath.Dir(dst)) + if err != nil { + return err + } + defer func() { _ = dstRoot.Close() }() + out, err := dstRoot.OpenFile(filepath.Base(dst), os.O_CREATE|os.O_TRUNC|os.O_WRONLY, info.Mode()) + if err != nil { + return err + } + defer func() { + if err := out.Close(); err != nil { + slog.Debug("copyFile: close destination failed", "dst", dst, "error", err) + } + }() + + if _, err := io.Copy(out, in); err != nil { + return err + } + return out.Sync() +} + +// Fix 5: isHigherVersion treats a non-parseable current version (e.g. "dev", +// a git SHA, or any non-semver string) as "unknown / dev build" and returns +// false without an error. This prevents the update check loop from spamming +// parse errors every cycle when the agent is built without a semver tag. +// The server-supplied next version must still be valid semver. +func isHigherVersion(next, current string) (bool, error) { + nextParts, err := parseVersion(next) + if err != nil { + return false, fmt.Errorf("invalid next version %q: %w", next, err) + } + currentParts, err := parseVersion(current) + if err != nil { + // Non-semver current version (dev build, git sha) — skip update silently. + return false, nil + } + + for i := 0; i < len(nextParts) || i < len(currentParts); i++ { + var a, b int + if i < len(nextParts) { + a = nextParts[i] + } + if i < len(currentParts) { + b = currentParts[i] + } + if a > b { + return true, nil + } + if a < b { + return false, nil + } + } + return false, nil +} + +func parseVersion(v string) ([]int, error) { + trimmed := strings.TrimSpace(strings.TrimPrefix(v, "v")) + if trimmed == "" { + return nil, errors.New("empty version") + } + + core := strings.SplitN(trimmed, "+", 2)[0] + core = strings.SplitN(core, "-", 2)[0] + parts := strings.Split(core, ".") + out := make([]int, 0, len(parts)) + for _, part := range parts { + if part == "" { + return nil, errors.New("empty version segment") + } + n, err := strconv.Atoi(part) + if err != nil { + return nil, err + } + out = append(out, n) + } + return out, nil +} + +func resolveURL(baseURL, target string) (string, error) { + base, err := url.Parse(baseURL) + if err != nil { + return "", err + } + parsed, err := url.Parse(target) + if err != nil { + return "", err + } + return base.ResolveReference(parsed).String(), nil +} diff --git a/internal/update/manager_test.go b/internal/update/manager_test.go new file mode 100644 index 0000000..6fbe994 --- /dev/null +++ b/internal/update/manager_test.go @@ -0,0 +1,312 @@ +package update + +import ( + "bytes" + "context" + "crypto/sha256" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" +) + +func TestManagerCheckNoUpdate(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set(VersionHeaderName, "1.2.3") + w.Header().Set("Location", "/binary") + w.WriteHeader(http.StatusTemporaryRedirect) + })) + defer server.Close() + + manager := NewManager(server.URL+"/download", "1.2.3") + release, err := manager.Check(context.Background()) + if err != nil { + t.Fatalf("Check() error = %v", err) + } + if release != nil { + t.Fatalf("Check() release = %+v, want nil", release) + } +} + +func TestManagerCheckIgnoresLowerVersion(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set(VersionHeaderName, "1.2.2") + w.Header().Set("Location", "/binary") + w.WriteHeader(http.StatusTemporaryRedirect) + })) + defer server.Close() + + manager := NewManager(server.URL+"/download", "1.2.3") + release, err := manager.Check(context.Background()) + if err != nil { + t.Fatalf("Check() error = %v", err) + } + if release != nil { + t.Fatalf("Check() release = %+v, want nil", release) + } +} + +func TestManagerCheckFindsUpdate(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set(VersionHeaderName, "2.0.0") + w.Header().Set(DigestHeaderName, "sha256:abc123") + w.Header().Set("Location", "/binary") + w.WriteHeader(http.StatusTemporaryRedirect) + })) + defer server.Close() + + manager := NewManager(server.URL+"/download", "1.2.3") + release, err := manager.Check(context.Background()) + if err != nil { + t.Fatalf("Check() error = %v", err) + } + if release == nil { + t.Fatal("Check() release = nil, want update") + } + if release.Version != "2.0.0" { + t.Fatalf("release.Version = %q, want 2.0.0", release.Version) + } + if release.Digest != "sha256:abc123" { + t.Fatalf("release.Digest = %q, want sha256:abc123", release.Digest) + } + if release.DownloadURL != server.URL+"/binary" { + t.Fatalf("release.DownloadURL = %q, want %q", release.DownloadURL, server.URL+"/binary") + } +} + +func TestManagerDownloadReleaseAndPromote(t *testing.T) { + binary := []byte("#!/bin/sh\nexit 0\n") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/binary": + _, _ = w.Write(binary) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + dir := t.TempDir() + currentPath := filepath.Join(dir, "hyperstack-agent") + if err := os.WriteFile(currentPath, []byte("old-binary"), 0o755); err != nil { + t.Fatalf("WriteFile(currentPath) error = %v", err) + } + + manager := NewManager(server.URL+"/download", "1.0.0") + release := &Release{ + Version: "2.0.0", + Digest: digestFor(binary), + DownloadURL: server.URL + "/binary", + } + + if err := manager.DownloadRelease(context.Background(), release, currentPath); err != nil { + t.Fatalf("DownloadRelease() error = %v", err) + } + if release.StagedPath == "" { + t.Fatal("release.StagedPath = empty, want staged path") + } + + if err := manager.PromoteRelease(currentPath, release); err != nil { + t.Fatalf("PromoteRelease() error = %v", err) + } + + got, err := os.ReadFile(currentPath) + if err != nil { + t.Fatalf("ReadFile(currentPath) error = %v", err) + } + if len(got) == 0 { + t.Fatal("current binary is empty") + } + + backup, err := os.ReadFile(currentPath + ".bak") + if err != nil { + t.Fatalf("ReadFile(backup) error = %v", err) + } + if string(backup) != "old-binary" { + t.Fatalf("backup binary = %q, want %q", string(backup), "old-binary") + } + + if manager.CurrentVersion != "2.0.0" { + t.Fatalf("CurrentVersion = %q, want 2.0.0", manager.CurrentVersion) + } +} + +func TestManagerDownloadReleaseRejectsFailedSmokeTest(t *testing.T) { + binary := []byte("#!/bin/sh\nexit 1\n") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write(binary) + })) + defer server.Close() + + dir := t.TempDir() + currentPath := filepath.Join(dir, "hyperstack-agent") + if err := os.WriteFile(currentPath, []byte("old-binary"), 0o755); err != nil { + t.Fatalf("WriteFile(currentPath) error = %v", err) + } + + manager := NewManager(server.URL+"/download", "1.0.0") + release := &Release{ + Version: "2.0.0", + Digest: digestFor(binary), + DownloadURL: server.URL, + } + + if err := manager.DownloadRelease(context.Background(), release, currentPath); err == nil { + t.Fatal("DownloadRelease() error = nil, want smoke test error") + } + if release.StagedPath != "" { + t.Fatalf("release.StagedPath = %q, want empty", release.StagedPath) + } +} + +func TestManagerDownloadReleaseRejectsDigestMismatch(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte("#!/bin/sh\nexit 1\n")) + })) + defer server.Close() + + dir := t.TempDir() + currentPath := filepath.Join(dir, "hyperstack-agent") + if err := os.WriteFile(currentPath, []byte("old-binary"), 0o755); err != nil { + t.Fatalf("WriteFile(currentPath) error = %v", err) + } + + manager := NewManager(server.URL+"/download", "1.0.0") + release := &Release{ + Version: "2.0.0", + Digest: "sha256:deadbeef", + DownloadURL: server.URL, + } + + if err := manager.DownloadRelease(context.Background(), release, currentPath); err == nil { + t.Fatal("DownloadRelease() error = nil, want digest mismatch") + } + if release.StagedPath != "" { + t.Fatalf("release.StagedPath = %q, want empty", release.StagedPath) + } +} + +func TestCopyWithLimitRejectsOversizedBinary(t *testing.T) { + var dst bytes.Buffer + err := copyWithLimit(&dst, bytes.NewBufferString("12345"), 4) + if err == nil { + t.Fatal("copyWithLimit() error = nil, want size error") + } + if dst.String() != "12345" { + t.Fatalf("copied data = %q, want max+1 bytes", dst.String()) + } +} + +func TestCopyWithLimitAcceptsExactLimit(t *testing.T) { + var dst bytes.Buffer + err := copyWithLimit(&dst, bytes.NewBufferString("1234"), 4) + if err != nil { + t.Fatalf("copyWithLimit() error = %v", err) + } + if dst.String() != "1234" { + t.Fatalf("copied data = %q, want full data", dst.String()) + } +} + +// Fix 1: Check() must reject a release whose digest header is missing. +func TestManagerCheckRejectsAbsentDigestHeader(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set(VersionHeaderName, "2.0.0") + // DigestHeaderName intentionally omitted + w.Header().Set("Location", "/binary") + w.WriteHeader(http.StatusTemporaryRedirect) + })) + defer server.Close() + + manager := NewManager(server.URL+"/download", "1.0.0") + release, err := manager.Check(context.Background()) + if err == nil { + t.Fatal("Check() error = nil, want missing-digest error") + } + if release != nil { + t.Fatalf("Check() release = %+v, want nil on error", release) + } +} + +// Fix 2: DownloadRelease must fail when the release carries an empty digest. +func TestManagerDownloadReleaseRejectsEmptyDigest(t *testing.T) { + binary := []byte("#!/bin/sh\nexit 0\n") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write(binary) + })) + defer server.Close() + + dir := t.TempDir() + currentPath := filepath.Join(dir, "hyperstack-agent") + if err := os.WriteFile(currentPath, []byte("old-binary"), 0o755); err != nil { + t.Fatalf("WriteFile error = %v", err) + } + + manager := NewManager(server.URL+"/download", "1.0.0") + release := &Release{ + Version: "2.0.0", + Digest: "", // explicitly empty — verifyDigest must reject this + DownloadURL: server.URL, + } + + if err := manager.DownloadRelease(context.Background(), release, currentPath); err == nil { + t.Fatal("DownloadRelease() error = nil, want digest-required error") + } + if release.StagedPath != "" { + t.Fatalf("release.StagedPath = %q, want empty on failure", release.StagedPath) + } +} + +// Fix 5: isHigherVersion must skip the update silently (return false, nil) +// when the current version is non-semver (e.g. "dev" or a git SHA). +func TestIsHigherVersionDevCurrentSkipsUpdate(t *testing.T) { + cases := []string{"dev", "abc1234", "HEAD", ""} + for _, current := range cases { + got, err := isHigherVersion("1.0.0", current) + if err != nil { + t.Errorf("isHigherVersion(1.0.0, %q) error = %v, want nil", current, err) + } + if got { + t.Errorf("isHigherVersion(1.0.0, %q) = true, want false (dev build should skip)", current) + } + } +} + +// Fix 5: an invalid next version from the server must still propagate as an error. +func TestIsHigherVersionInvalidNextVersionErrors(t *testing.T) { + _, err := isHigherVersion("not-semver", "1.0.0") + if err == nil { + t.Fatal("isHigherVersion(not-semver, 1.0.0) error = nil, want parse error") + } +} + +// Fix 4: resolveExecPath must reject a non-existent path. +func TestResolveExecPathRejectsNonExistent(t *testing.T) { + _, err := resolveExecPath("/nonexistent/path/that/does/not/exist") + if err == nil { + t.Fatal("resolveExecPath() error = nil, want error for missing path") + } +} + +// Fix 4: resolveExecPath must succeed for a real file and return an absolute path. +func TestResolveExecPathResolvesRealFile(t *testing.T) { + dir := t.TempDir() + p := filepath.Join(dir, "mybinary") + if err := os.WriteFile(p, []byte("data"), 0o755); err != nil { + t.Fatalf("WriteFile error = %v", err) + } + resolved, err := resolveExecPath(p) + if err != nil { + t.Fatalf("resolveExecPath() error = %v", err) + } + if !filepath.IsAbs(resolved) { + t.Fatalf("resolveExecPath() = %q, want absolute path", resolved) + } +} + +func digestFor(data []byte) string { + sum := sha256.Sum256(data) + return fmt.Sprintf("sha256:%x", sum) +} diff --git a/scripts/serve.sh b/scripts/serve.sh index 40a5c9c..d1de001 100755 --- a/scripts/serve.sh +++ b/scripts/serve.sh @@ -5,6 +5,7 @@ set -e tmpdir="$(mktemp -d)" cd "$tmpdir" +echo "ok" > healthz cp /usr/local/bin/hyperstack-agent download sha256=$(sha256sum download | cut -d' ' -f1)