From c99c36edc1bc789140ebfe2986aad3fd0a9705cb Mon Sep 17 00:00:00 2001 From: nabil1440 <52530910+nabil1440@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:19:13 +0600 Subject: [PATCH 1/3] feat(agent): add fly agent run, the base of the monitoring agent - Read FLY_AGENT_URL, FLY_AGENT_TOKEN, FLY_AGENT_SERVER_ID and STATE_DIRECTORY. An error names each key that is not set or not valid. - Accept only an https URL. Plain http is accepted only for a loopback host, for tests and local development. - Lock a file in the state directory, so that only one agent operates. - Work at server_id % 60 seconds past each minute, and report after each report interval (1 to 10 samples, saved in state.json). - Stop on SIGTERM. The agent does not need Docker and does not run as root. - Add a client for the control plane and internal/statefile for crash-safe JSON files. The next layers use them. Refs #26 --- agent_test.go | 156 +++++++++++++++++++++++++++ cmd/agent.go | 43 ++++++++ internal/agent/agent.go | 114 ++++++++++++++++++++ internal/agent/agent_test.go | 150 ++++++++++++++++++++++++++ internal/agent/client.go | 108 +++++++++++++++++++ internal/agent/client_test.go | 133 +++++++++++++++++++++++ internal/agent/config.go | 133 +++++++++++++++++++++++ internal/agent/config_test.go | 126 ++++++++++++++++++++++ internal/agent/lock.go | 30 ++++++ internal/agent/lock_test.go | 27 +++++ internal/statefile/statefile.go | 67 ++++++++++++ internal/statefile/statefile_test.go | 103 ++++++++++++++++++ main_test.go | 9 +- 13 files changed, 1198 insertions(+), 1 deletion(-) create mode 100644 agent_test.go create mode 100644 cmd/agent.go create mode 100644 internal/agent/agent.go create mode 100644 internal/agent/agent_test.go create mode 100644 internal/agent/client.go create mode 100644 internal/agent/client_test.go create mode 100644 internal/agent/config.go create mode 100644 internal/agent/config_test.go create mode 100644 internal/agent/lock.go create mode 100644 internal/agent/lock_test.go create mode 100644 internal/statefile/statefile.go create mode 100644 internal/statefile/statefile_test.go diff --git a/agent_test.go b/agent_test.go new file mode 100644 index 0000000..34c56d4 --- /dev/null +++ b/agent_test.go @@ -0,0 +1,156 @@ +package main + +// End-to-end tests of "fly agent run": the configuration errors, the lock and +// a clean stop. The loop itself is tested in internal/agent with a fake clock. + +import ( + "bytes" + "os" + "os/exec" + "strings" + "sync" + "syscall" + "testing" + "time" +) + +const testToken = "flyagt_0123456789abcdefghijABCDEFGHIJ" + +// agentEnv is the environment of a valid agent with its own state directory. +// PATH holds no docker command: the agent must not need Docker. +func agentEnv(t *testing.T) []string { + t.Helper() + + if os.Geteuid() == 0 { + t.Skip("fly refuses to run as root") + } + + return []string{ + "PATH=" + t.TempDir(), + "HOME=" + t.TempDir(), + "FLY_AGENT_URL=http://127.0.0.1:9", + "FLY_AGENT_TOKEN=" + testToken, + "FLY_AGENT_SERVER_ID=17", + "STATE_DIRECTORY=" + t.TempDir(), + } +} + +// without returns environ without the variable key. +func without(environ []string, key string) []string { + var out []string + for _, kv := range environ { + if !strings.HasPrefix(kv, key+"=") { + out = append(out, kv) + } + } + return out +} + +// lockedBuffer is a bytes.Buffer that a child process and the test can use +// at the same time. +type lockedBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *lockedBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *lockedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +// startAgent starts "fly agent run" and waits until it logs that it started. +func startAgent(t *testing.T, environ []string) (*exec.Cmd, *lockedBuffer) { + t.Helper() + + cmd := exec.Command(flyBin, "agent", "run") + cmd.Env = environ + stderr := &lockedBuffer{} + cmd.Stderr = stderr + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = cmd.Process.Kill() }) + + deadline := time.Now().Add(10 * time.Second) + for !strings.Contains(stderr.String(), "agent started") { + if time.Now().After(deadline) { + t.Fatalf("the agent did not start within 10s. stderr:\n%s", stderr) + } + time.Sleep(20 * time.Millisecond) + } + + return cmd, stderr +} + +func TestAgentConfigErrors(t *testing.T) { + tests := []struct { + name string + environ func([]string) []string + want string + }{ + {"no token", func(e []string) []string { return without(e, "FLY_AGENT_TOKEN") }, "FLY_AGENT_TOKEN is not set"}, + {"no state directory", func(e []string) []string { return without(e, "STATE_DIRECTORY") }, "STATE_DIRECTORY is not set"}, + {"plain http", func(e []string) []string { return append(e, "FLY_AGENT_URL=http://example.com") }, "must be an https URL"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + res := runFly(t, t.TempDir(), tt.environ(agentEnv(t)), "agent", "run") + + if res.code != 1 { + t.Errorf("exit code = %d, want 1", res.code) + } + if !strings.Contains(res.stderr, tt.want) { + t.Errorf("stderr = %q, want it to contain %q", res.stderr, tt.want) + } + }) + } +} + +func TestAgentRejectsArguments(t *testing.T) { + res := runFly(t, t.TempDir(), agentEnv(t), "agent", "run", "extra") + if res.code != 1 || !strings.Contains(res.stderr, "unknown command") && !strings.Contains(res.stderr, "accepts 0 arg") { + t.Errorf("fly agent run extra: exit %d, stderr %q; want a usage error", res.code, res.stderr) + } +} + +func TestAgentRunsWithoutDockerAndStopsOnSIGTERM(t *testing.T) { + environ := agentEnv(t) + cmd, stderr := startAgent(t, environ) + + // Only one agent can use the state directory. + second := runFly(t, t.TempDir(), environ, "agent", "run") + if second.code != 1 || !strings.Contains(second.stderr, "a different agent is running") { + t.Errorf("second agent: exit %d, stderr %q; want exit 1 and a lock error", second.code, second.stderr) + } + + if err := cmd.Process.Signal(syscall.SIGTERM); err != nil { + t.Fatal(err) + } + + done := make(chan error, 1) + go func() { done <- cmd.Wait() }() + select { + case err := <-done: + if err != nil { + t.Errorf("agent exit after SIGTERM: %v, want exit status 0. stderr:\n%s", err, stderr) + } + case <-time.After(10 * time.Second): + t.Fatal("the agent did not stop within 10s after SIGTERM") + } + + out := stderr.String() + if !strings.Contains(out, "agent stopped") { + t.Errorf("stderr = %q, want an \"agent stopped\" line", out) + } + if strings.Contains(out, testToken) { + t.Error("the agent log contains the token") + } +} diff --git a/cmd/agent.go b/cmd/agent.go new file mode 100644 index 0000000..b7e44d6 --- /dev/null +++ b/cmd/agent.go @@ -0,0 +1,43 @@ +package cmd + +import ( + "log/slog" + "os" + "os/signal" + "syscall" + + "github.com/flywp/server-cli/internal/agent" + "github.com/spf13/cobra" +) + +var agentCmd = &cobra.Command{ + Use: "agent", + Short: "Run the FlyWP monitoring agent", +} + +// agentRunCmd does not need Docker: the agent must also report when Docker +// is down. +var agentRunCmd = &cobra.Command{ + Use: "run", + Short: "Run the monitoring agent until it is stopped", + Long: `Run the FlyWP monitoring agent until it is stopped. systemd starts this +command (fly-agent.service). The agent reads FLY_AGENT_URL, FLY_AGENT_TOKEN, +FLY_AGENT_SERVER_ID and STATE_DIRECTORY from the environment.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, args []string) error { + cfg, err := agent.ConfigFromEnv(os.Getenv) + if err != nil { + return err + } + + ctx, stop := signal.NotifyContext(cmd.Context(), syscall.SIGTERM, os.Interrupt) + defer stop() + + return agent.Run(ctx, cfg, slog.New(slog.NewTextHandler(os.Stderr, nil))) + }, +} + +func init() { + agentCmd.AddCommand(agentRunCmd) + rootCmd.AddCommand(agentCmd) +} diff --git a/internal/agent/agent.go b/internal/agent/agent.go new file mode 100644 index 0000000..3348ea0 --- /dev/null +++ b/internal/agent/agent.go @@ -0,0 +1,114 @@ +package agent + +import ( + "context" + "errors" + "io/fs" + "log/slog" + "path/filepath" + "time" + + "github.com/flywp/server-cli/internal/statefile" + "github.com/flywp/server-cli/internal/version" +) + +// The report interval is the number of samples in one report. The control +// plane sets it in each reply. +const ( + minReportInterval = 1 + maxReportInterval = 10 +) + +// state is the part of the agent state that is not a queue. +type state struct { + ReportInterval int `json:"report_interval"` +} + +type agent struct { + cfg Config + log *slog.Logger + + // interval is the report interval, and pending is the number of samples + // since the last report. + interval int + pending int +} + +// Run runs the agent until ctx is done. Only one agent can run with the same +// state directory. +func Run(ctx context.Context, cfg Config, log *slog.Logger) error { + unlock, err := lock(cfg.StateDir) + if err != nil { + return err + } + defer unlock() + + a := &agent{cfg: cfg, log: log, interval: loadInterval(cfg.StateDir, log)} + log.Info("agent started", "version", version.Version, "offset", cfg.Offset(), "report_interval", a.interval) + a.loop(ctx) + log.Info("agent stopped") + + return nil +} + +// loop calls tick at the offset second of each minute until ctx is done. +func (a *agent) loop(ctx context.Context) { + for { + timer := time.NewTimer(time.Until(nextTick(time.Now(), a.cfg.Offset()))) + select { + case <-ctx.Done(): + timer.Stop() + return + case now := <-timer.C: + a.tick(ctx, now) + } + } +} + +// tick does the work of one minute: it takes a sample and, after each +// interval samples, sends a report. +func (a *agent) tick(ctx context.Context, now time.Time) { + a.log.Debug("tick", "at", now) + + a.pending++ + if a.pending >= a.interval { + a.pending = 0 + a.report(ctx) + } +} + +// report sends the queued data to the control plane. +func (a *agent) report(_ context.Context) { + a.log.Debug("report") +} + +// nextTick returns the first time after now that is offset past a full minute. +// It comes from the wall clock, so a slow tick moves to the next minute and +// never runs two times in one minute. +func nextTick(now time.Time, offset time.Duration) time.Time { + t := now.Truncate(time.Minute).Add(offset) + if !t.After(now) { + t = t.Add(time.Minute) + } + + return t +} + +// loadInterval returns the report interval of the last reply, or 1. +func loadInterval(dir string, log *slog.Logger) int { + var s state + err := statefile.Read(filepath.Join(dir, "state.json"), &s) + switch { + case errors.Is(err, fs.ErrNotExist): + return minReportInterval + case err != nil: + log.Warn("ignoring the saved state", "error", err) + return minReportInterval + } + + return clampInterval(s.ReportInterval) +} + +func clampInterval(n int) int { + return min(max(n, minReportInterval), maxReportInterval) +} diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go new file mode 100644 index 0000000..2f24eab --- /dev/null +++ b/internal/agent/agent_test.go @@ -0,0 +1,150 @@ +package agent + +import ( + "context" + "log/slog" + "path/filepath" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/flywp/server-cli/internal/statefile" +) + +// recorder is a slog handler that keeps the message and the time of each +// record, so that tests can see when the agent did its work. +type recorder struct { + mu sync.Mutex + records []slog.Record +} + +func (r *recorder) Enabled(context.Context, slog.Level) bool { return true } +func (r *recorder) WithAttrs([]slog.Attr) slog.Handler { return r } +func (r *recorder) WithGroup(string) slog.Handler { return r } + +func (r *recorder) Handle(_ context.Context, rec slog.Record) error { + r.mu.Lock() + defer r.mu.Unlock() + r.records = append(r.records, rec) + return nil +} + +// times returns the times of the records with message msg. +func (r *recorder) times(msg string) []time.Time { + r.mu.Lock() + defer r.mu.Unlock() + + var ts []time.Time + for _, rec := range r.records { + if rec.Message == msg { + ts = append(ts, rec.Time) + } + } + return ts +} + +func TestNextTick(t *testing.T) { + base := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + tests := []struct { + now time.Time + offset time.Duration + want time.Time + }{ + {base, 17 * time.Second, base.Add(17 * time.Second)}, + {base.Add(10 * time.Second), 17 * time.Second, base.Add(17 * time.Second)}, + // At the offset itself, the next tick is one minute later. + {base.Add(17 * time.Second), 17 * time.Second, base.Add(77 * time.Second)}, + {base.Add(30 * time.Second), 17 * time.Second, base.Add(77 * time.Second)}, + {base.Add(59*time.Second + 999*time.Millisecond), 0, base.Add(time.Minute)}, + } + + for _, tt := range tests { + if got := nextTick(tt.now, tt.offset); !got.Equal(tt.want) { + t.Errorf("nextTick(%s, %v) = %s, want %s", tt.now.Format(time.TimeOnly), tt.offset, got.Format(time.TimeOnly), tt.want.Format(time.TimeOnly)) + } + } +} + +func TestLoopTicksAtTheOffsetAndReportsEachInterval(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + if err := statefile.Write(filepath.Join(dir, "state.json"), state{ReportInterval: 2}); err != nil { + t.Fatal(err) + } + + rec := &recorder{} + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error) + go func() { + done <- Run(ctx, Config{ServerID: 17, StateDir: dir}, slog.New(rec)) + }() + + // The bubble starts at 2000-01-01 00:00:00 UTC. Let 4 minutes pass. + time.Sleep(4 * time.Minute) + cancel() + if err := <-done; err != nil { + t.Fatal(err) + } + + start := time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC) + ticks := rec.times("tick") + if len(ticks) != 4 { + t.Fatalf("ticks at %v, want 4 ticks", ticks) + } + for i, got := range ticks { + if want := start.Add(time.Duration(i)*time.Minute + 17*time.Second); !got.Equal(want) { + t.Errorf("tick %d at %s, want %s", i, got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + } + + // Report interval 2: a report after tick 2 and tick 4. + if reports := rec.times("report"); len(reports) != 2 || !reports[0].Equal(ticks[1]) || !reports[1].Equal(ticks[3]) { + t.Errorf("reports at %v, want at ticks 2 and 4 (%v)", reports, ticks) + } + if len(rec.times("agent stopped")) != 1 { + t.Error("no \"agent stopped\" log record") + } + }) +} + +func TestRunRefusesASecondAgent(t *testing.T) { + dir := t.TempDir() + unlock, err := lock(dir) + if err != nil { + t.Fatal(err) + } + defer unlock() + + if err := Run(context.Background(), Config{StateDir: dir}, slog.New(&recorder{})); err == nil { + t.Fatal("Run() = nil, want an error while a different agent holds the lock") + } +} + +func TestLoadInterval(t *testing.T) { + tests := []struct { + name string + saved *state + want int + }{ + {"no state", nil, 1}, + {"saved", &state{ReportInterval: 5}, 5}, + {"too low", &state{ReportInterval: 0}, 1}, + {"too high", &state{ReportInterval: 60}, 10}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + if tt.saved != nil { + if err := statefile.Write(filepath.Join(dir, "state.json"), tt.saved); err != nil { + t.Fatal(err) + } + } + + if got := loadInterval(dir, slog.New(&recorder{})); got != tt.want { + t.Errorf("loadInterval() = %d, want %d", got, tt.want) + } + }) + } +} diff --git a/internal/agent/client.go b/internal/agent/client.go new file mode 100644 index 0000000..893854b --- /dev/null +++ b/internal/agent/client.go @@ -0,0 +1,108 @@ +package agent + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strconv" + "time" + + "github.com/flywp/server-cli/internal/version" +) + +// requestTimeout limits each request to the control plane. +const requestTimeout = 20 * time.Second + +// maxReplySize limits the reply body that the agent reads. +const maxReplySize = 1 << 20 + +// Client sends requests to the control plane. +type Client struct { + base *url.URL + token string + http *http.Client +} + +// NewClient returns a client for the control plane of cfg. A nil hc uses a +// client with the default timeout. +func NewClient(cfg Config, hc *http.Client) *Client { + if hc == nil { + hc = &http.Client{Timeout: requestTimeout} + } + + return &Client{base: cfg.URL, token: cfg.Token, http: hc} +} + +// StatusError is a reply from the control plane that is not 200 OK. +type StatusError struct { + StatusCode int + // RetryAfter is the wait that the Retry-After header asks for, or 0. + RetryAfter time.Duration +} + +func (e *StatusError) Error() string { + return fmt.Sprintf("control plane replied %d %s", e.StatusCode, http.StatusText(e.StatusCode)) +} + +// do sends in as JSON (no body when in is nil) to path and decodes a 200 reply +// into out (no decoding when out is nil). Another status is a *StatusError. +func (c *Client) do(ctx context.Context, method, path string, in, out any) error { + var body io.Reader + if in != nil { + data, err := json.Marshal(in) + if err != nil { + return err + } + body = bytes.NewReader(data) + } + + req, err := http.NewRequestWithContext(ctx, method, c.base.JoinPath(path).String(), body) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+c.token) + req.Header.Set("User-Agent", "fly/"+version.Version) + req.Header.Set("Accept", "application/json") + if in != nil { + req.Header.Set("Content-Type", "application/json") + } + + resp, err := c.http.Do(req) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, maxReplySize)) + return &StatusError{StatusCode: resp.StatusCode, RetryAfter: retryAfter(resp.Header.Get("Retry-After"), time.Now())} + } + + if out == nil { + return nil + } + if err := json.NewDecoder(io.LimitReader(resp.Body, maxReplySize)).Decode(out); err != nil { + return fmt.Errorf("reading the reply of %s: %w", path, err) + } + + return nil +} + +// retryAfter reads a Retry-After value: a number of seconds or an HTTP date. +func retryAfter(v string, now time.Time) time.Duration { + if v == "" { + return 0 + } + if s, err := strconv.Atoi(v); err == nil { + return max(time.Duration(s)*time.Second, 0) + } + if t, err := http.ParseTime(v); err == nil { + return max(t.Sub(now), 0) + } + + return 0 +} diff --git a/internal/agent/client_test.go b/internal/agent/client_test.go new file mode 100644 index 0000000..143b99c --- /dev/null +++ b/internal/agent/client_test.go @@ -0,0 +1,133 @@ +package agent + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "net/url" + "testing" + "time" + + "github.com/flywp/server-cli/internal/version" +) + +func testClient(t *testing.T, handler http.HandlerFunc) *Client { + t.Helper() + + srv := httptest.NewTLSServer(handler) + t.Cleanup(srv.Close) + + u, err := url.Parse(srv.URL + "/base") + if err != nil { + t.Fatal(err) + } + + return NewClient(Config{URL: u, Token: testToken}, srv.Client()) +} + +func TestClientSendsTheContractHeaders(t *testing.T) { + var got *http.Request + var body map[string]int + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + got = r + _ = json.NewDecoder(r.Body).Decode(&body) + _, _ = w.Write([]byte(`{"report_interval": 5}`)) + }) + + var reply struct { + ReportInterval int `json:"report_interval"` + } + if err := c.do(context.Background(), http.MethodPost, "agent/v1/metrics", map[string]int{"n": 1}, &reply); err != nil { + t.Fatal(err) + } + + if got.URL.Path != "/base/agent/v1/metrics" { + t.Errorf("path = %q, want the contract path under the base URL", got.URL.Path) + } + checks := map[string]string{ + "Authorization": "Bearer " + testToken, + "User-Agent": "fly/" + version.Version, + "Accept": "application/json", + "Content-Type": "application/json", + } + for header, want := range checks { + if v := got.Header.Get(header); v != want { + t.Errorf("%s = %q, want %q", header, v, want) + } + } + if body["n"] != 1 { + t.Errorf("body = %v, want the JSON of the request", body) + } + if reply.ReportInterval != 5 { + t.Errorf("reply = %+v, want the decoded JSON", reply) + } +} + +func TestClientWithoutBody(t *testing.T) { + var got *http.Request + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + got = r + _, _ = w.Write([]byte(`{}`)) + }) + + if err := c.do(context.Background(), http.MethodGet, "agent/v1/commands", nil, nil); err != nil { + t.Fatal(err) + } + if got.Method != http.MethodGet || got.Header.Get("Content-Type") != "" { + t.Errorf("request = %s with Content-Type %q, want GET without a body", got.Method, got.Header.Get("Content-Type")) + } +} + +func TestClientStatusError(t *testing.T) { + tests := []struct { + name string + code int + retryAfter string + want time.Duration + }{ + {"unauthorized", http.StatusUnauthorized, "", 0}, + {"throttled", http.StatusTooManyRequests, "30", 30 * time.Second}, + {"unavailable", http.StatusServiceUnavailable, "", 0}, + {"bad gateway", http.StatusBadGateway, "", 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if tt.retryAfter != "" { + w.Header().Set("Retry-After", tt.retryAfter) + } + w.WriteHeader(tt.code) + }) + + err := c.do(context.Background(), http.MethodPost, "agent/v1/events", map[string]int{}, nil) + var statusErr *StatusError + if !errors.As(err, &statusErr) { + t.Fatalf("do() error = %v, want a *StatusError", err) + } + if statusErr.StatusCode != tt.code || statusErr.RetryAfter != tt.want { + t.Errorf("StatusError = %+v, want code %d and wait %v", statusErr, tt.code, tt.want) + } + }) + } +} + +func TestRetryAfter(t *testing.T) { + now := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + tests := map[string]time.Duration{ + "": 0, + "120": 2 * time.Minute, + "-5": 0, + "soon": 0, + "Tue, 22 Sep 2026 10:01:30 GMT": 90 * time.Second, + "Tue, 22 Sep 2026 09:00:00 GMT": 0, + } + + for v, want := range tests { + if got := retryAfter(v, now); got != want { + t.Errorf("retryAfter(%q) = %v, want %v", v, got, want) + } + } +} diff --git a/internal/agent/config.go b/internal/agent/config.go new file mode 100644 index 0000000..0c34bae --- /dev/null +++ b/internal/agent/config.go @@ -0,0 +1,133 @@ +// Package agent is the FlyWP monitoring agent: the long-running mode of fly +// that "fly agent run" starts. It follows the FlyWP monitoring agent +// contract v0.2.1. +package agent + +import ( + "errors" + "fmt" + "net" + "net/url" + "os" + "strconv" + "strings" + "time" + "unicode" +) + +// The environment keys that the FlyWP installer writes to /etc/fly/agent.env, +// and the key that systemd sets for StateDirectory=. +const ( + EnvURL = "FLY_AGENT_URL" + EnvToken = "FLY_AGENT_TOKEN" + EnvServerID = "FLY_AGENT_SERVER_ID" + EnvStateDir = "STATE_DIRECTORY" +) + +// Config is the configuration of the agent. +type Config struct { + // URL is the base URL of the control plane. The agent adds the paths + // of the contract to it. + URL *url.URL + // Token is the bearer token of the agent. It is opaque: the agent sends + // it and never parses it. Never log it. + Token string + // ServerID sets the second of the minute at which the agent works. The + // agent never sends it. + ServerID int64 + // StateDir keeps the files that must survive a restart. + StateDir string +} + +// Offset is the time after each full minute at which the agent works. It +// spreads the reports of the fleet over the minute, and it does not change +// between restarts. +func (c Config) Offset() time.Duration { + return time.Duration(c.ServerID%60) * time.Second +} + +// ConfigFromEnv reads the configuration from the environment. The error names +// each key that is not set or not valid. +func ConfigFromEnv(getenv func(string) string) (Config, error) { + var cfg Config + var errs []error + + u, err := parseURL(getenv(EnvURL)) + if err != nil { + errs = append(errs, err) + } + cfg.URL = u + + cfg.Token = getenv(EnvToken) + switch { + case cfg.Token == "": + errs = append(errs, fmt.Errorf("%s is not set", EnvToken)) + case strings.ContainsFunc(cfg.Token, func(r rune) bool { return unicode.IsSpace(r) || unicode.IsControl(r) }): + errs = append(errs, fmt.Errorf("%s contains spaces or control characters", EnvToken)) + } + + if v := getenv(EnvServerID); v == "" { + errs = append(errs, fmt.Errorf("%s is not set", EnvServerID)) + } else if id, err := strconv.ParseInt(v, 10, 64); err != nil || id < 0 { + errs = append(errs, fmt.Errorf("%s must be an integer of 0 or more, not %q", EnvServerID, v)) + } else { + cfg.ServerID = id + } + + dir, err := stateDir(getenv(EnvStateDir)) + if err != nil { + errs = append(errs, err) + } + cfg.StateDir = dir + + return cfg, errors.Join(errs...) +} + +// parseURL accepts an https URL. It also accepts http for a loopback host, +// for tests and local development: plain http never leaves the machine. +func parseURL(v string) (*url.URL, error) { + if v == "" { + return nil, fmt.Errorf("%s is not set", EnvURL) + } + + u, err := url.Parse(strings.TrimRight(v, "/")) + if err != nil || u.Host == "" { + return nil, fmt.Errorf("%s is not a valid URL: %q", EnvURL, v) + } + + switch { + case u.Scheme == "https": + case u.Scheme == "http" && isLoopback(u.Hostname()): + default: + return nil, fmt.Errorf("%s must be an https URL, not %q", EnvURL, v) + } + + return u, nil +} + +func isLoopback(host string) bool { + if host == "localhost" { + return true + } + ip := net.ParseIP(host) + return ip != nil && ip.IsLoopback() +} + +// stateDir returns the first directory of STATE_DIRECTORY. systemd separates +// the directories with ":" when a unit has more than one. +func stateDir(v string) (string, error) { + if v == "" { + return "", fmt.Errorf("%s is not set (systemd sets it for StateDirectory=)", EnvStateDir) + } + + dir, _, _ := strings.Cut(v, ":") + info, err := os.Stat(dir) + if err != nil { + return "", fmt.Errorf("%s: %w", EnvStateDir, err) + } + if !info.IsDir() { + return "", fmt.Errorf("%s is not a directory: %s", EnvStateDir, dir) + } + + return dir, nil +} diff --git a/internal/agent/config_test.go b/internal/agent/config_test.go new file mode 100644 index 0000000..4503bcf --- /dev/null +++ b/internal/agent/config_test.go @@ -0,0 +1,126 @@ +package agent + +import ( + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +const testToken = "flyagt_0123456789abcdefghijABCDEFGHIJ" + +func validEnv(t *testing.T) map[string]string { + t.Helper() + return map[string]string{ + EnvURL: "https://app.flywp.com", + EnvToken: testToken, + EnvServerID: "59", + EnvStateDir: t.TempDir(), + } +} + +func getenv(env map[string]string) func(string) string { + return func(k string) string { return env[k] } +} + +func TestConfigFromEnv(t *testing.T) { + env := validEnv(t) + env[EnvURL] = "https://app.flywp.com/" + + cfg, err := ConfigFromEnv(getenv(env)) + if err != nil { + t.Fatal(err) + } + + if got := cfg.URL.String(); got != "https://app.flywp.com" { + t.Errorf("URL = %q, want the trailing slash removed", got) + } + if cfg.Token != testToken || cfg.ServerID != 59 || cfg.StateDir != env[EnvStateDir] { + t.Errorf("ConfigFromEnv() = %+v", cfg) + } + if got := cfg.Offset(); got != 59*time.Second { + t.Errorf("Offset() = %v, want 59s", got) + } +} + +func TestConfigOffsetWraps(t *testing.T) { + if got := (Config{ServerID: 125}).Offset(); got != 5*time.Second { + t.Errorf("Offset() = %v, want 5s (125 %% 60)", got) + } +} + +func TestConfigStateDirectoryList(t *testing.T) { + env := validEnv(t) + first := env[EnvStateDir] + env[EnvStateDir] = first + ":" + t.TempDir() + + cfg, err := ConfigFromEnv(getenv(env)) + if err != nil { + t.Fatal(err) + } + if cfg.StateDir != first { + t.Errorf("StateDir = %q, want the first directory %q", cfg.StateDir, first) + } +} + +func TestConfigFromEnvErrors(t *testing.T) { + file := filepath.Join(t.TempDir(), "file") + if err := os.WriteFile(file, nil, 0o600); err != nil { + t.Fatal(err) + } + + tests := []struct { + name, key, value, want string + }{ + {"no URL", EnvURL, "", "FLY_AGENT_URL is not set"}, + {"plain http", EnvURL, "http://app.flywp.com", "must be an https URL"}, + {"other scheme", EnvURL, "ftp://app.flywp.com", "must be an https URL"}, + {"no host", EnvURL, "https://", "not a valid URL"}, + {"no token", EnvToken, "", "FLY_AGENT_TOKEN is not set"}, + {"token with newline", EnvToken, "flyagt_abc\nX-Other: 1", "FLY_AGENT_TOKEN contains spaces"}, + {"no server id", EnvServerID, "", "FLY_AGENT_SERVER_ID is not set"}, + {"negative server id", EnvServerID, "-1", "FLY_AGENT_SERVER_ID must be an integer"}, + {"text server id", EnvServerID, "abc", "FLY_AGENT_SERVER_ID must be an integer"}, + {"no state directory", EnvStateDir, "", "STATE_DIRECTORY is not set"}, + {"missing state directory", EnvStateDir, filepath.Join(t.TempDir(), "missing"), "STATE_DIRECTORY"}, + {"state directory is a file", EnvStateDir, file, "STATE_DIRECTORY is not a directory"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + env := validEnv(t) + env[tt.key] = tt.value + + _, err := ConfigFromEnv(getenv(env)) + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("ConfigFromEnv() error = %v, want it to contain %q", err, tt.want) + } + if strings.Contains(err.Error(), testToken) { + t.Errorf("error %q contains the token", err) + } + }) + } +} + +func TestConfigFromEnvNamesEachProblem(t *testing.T) { + _, err := ConfigFromEnv(getenv(map[string]string{})) + if err == nil { + t.Fatal("ConfigFromEnv() = nil, want an error") + } + for _, key := range []string{EnvURL, EnvToken, EnvServerID, EnvStateDir} { + if !strings.Contains(err.Error(), key) { + t.Errorf("error %q does not name %s", err, key) + } + } +} + +func TestConfigAcceptsLoopbackHTTP(t *testing.T) { + for _, u := range []string{"http://127.0.0.1:8080", "http://[::1]:8080", "http://localhost:8080"} { + env := validEnv(t) + env[EnvURL] = u + if _, err := ConfigFromEnv(getenv(env)); err != nil { + t.Errorf("ConfigFromEnv(%s) error = %v, want loopback http accepted", u, err) + } + } +} diff --git a/internal/agent/lock.go b/internal/agent/lock.go new file mode 100644 index 0000000..5978c20 --- /dev/null +++ b/internal/agent/lock.go @@ -0,0 +1,30 @@ +package agent + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "syscall" +) + +// lock takes an exclusive lock on the lock file in dir, so that only one agent +// uses the state files. The lock ends when unlock runs or the process exits. +func lock(dir string) (unlock func(), err error) { + path := filepath.Join(dir, "lock") + f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, err + } + + if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX|syscall.LOCK_NB); err != nil { + _ = f.Close() + if errors.Is(err, syscall.EWOULDBLOCK) { + return nil, fmt.Errorf("a different agent is running: %s is locked", path) + } + return nil, fmt.Errorf("locking %s: %w", path, err) + } + + // Closing the file releases the lock. + return func() { _ = f.Close() }, nil +} diff --git a/internal/agent/lock_test.go b/internal/agent/lock_test.go new file mode 100644 index 0000000..9aaac86 --- /dev/null +++ b/internal/agent/lock_test.go @@ -0,0 +1,27 @@ +package agent + +import ( + "strings" + "testing" +) + +func TestLockAllowsOneAgent(t *testing.T) { + dir := t.TempDir() + + unlock, err := lock(dir) + if err != nil { + t.Fatal(err) + } + + if _, err := lock(dir); err == nil || !strings.Contains(err.Error(), "a different agent is running") { + t.Fatalf("second lock() error = %v, want a different agent is running", err) + } + + unlock() + + unlock, err = lock(dir) + if err != nil { + t.Fatalf("lock() after unlock error = %v", err) + } + unlock() +} diff --git a/internal/statefile/statefile.go b/internal/statefile/statefile.go new file mode 100644 index 0000000..a89e756 --- /dev/null +++ b/internal/statefile/statefile.go @@ -0,0 +1,67 @@ +// Package statefile keeps small JSON files that must survive a crash. +package statefile + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" +) + +// Write stores v as JSON in path. It writes a temporary file in the same +// directory, syncs it and renames it over path, so path always holds a +// complete file, also after a crash or a power loss. +func Write(path string, v any) (err error) { + data, err := json.Marshal(v) + if err != nil { + return fmt.Errorf("encoding %s: %w", filepath.Base(path), err) + } + + dir := filepath.Dir(path) + tmp, err := os.CreateTemp(dir, "."+filepath.Base(path)+".tmp-*") + if err != nil { + return err + } + defer func() { + if err != nil { + _ = os.Remove(tmp.Name()) + } + }() + + if _, err = tmp.Write(data); err != nil { + _ = tmp.Close() + return err + } + if err = tmp.Sync(); err != nil { + _ = tmp.Close() + return err + } + if err = tmp.Close(); err != nil { + return err + } + if err = os.Rename(tmp.Name(), path); err != nil { + return err + } + + // Sync the directory too, so that the rename itself survives a power loss. + d, err := os.Open(dir) + if err != nil { + return err + } + defer func() { _ = d.Close() }() + return d.Sync() +} + +// Read decodes the JSON in path into v. When path does not exist, the error +// matches fs.ErrNotExist and v is not changed. +func Read(path string, v any) error { + data, err := os.ReadFile(path) + if err != nil { + return err + } + if err := json.Unmarshal(data, v); err != nil { + return fmt.Errorf("reading %s: %w", filepath.Base(path), err) + } + + return nil +} diff --git a/internal/statefile/statefile_test.go b/internal/statefile/statefile_test.go new file mode 100644 index 0000000..11bc355 --- /dev/null +++ b/internal/statefile/statefile_test.go @@ -0,0 +1,103 @@ +package statefile + +import ( + "errors" + "io/fs" + "os" + "path/filepath" + "testing" +) + +type state struct { + Interval int `json:"interval"` + Names []string `json:"names"` +} + +func TestWriteThenRead(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + + want := state{Interval: 5, Names: []string{"a", "b"}} + if err := Write(path, want); err != nil { + t.Fatal(err) + } + + var got state + if err := Read(path, &got); err != nil { + t.Fatal(err) + } + if got.Interval != want.Interval || len(got.Names) != 2 { + t.Errorf("Read() = %+v, want %+v", got, want) + } + + // Only the file itself is left: no temporary files. + entries, err := os.ReadDir(filepath.Dir(path)) + if err != nil { + t.Fatal(err) + } + if len(entries) != 1 { + t.Errorf("directory holds %d entries, want only state.json", len(entries)) + } +} + +func TestWriteReplaces(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + + for i := 1; i <= 3; i++ { + if err := Write(path, state{Interval: i}); err != nil { + t.Fatal(err) + } + } + + var got state + if err := Read(path, &got); err != nil { + t.Fatal(err) + } + if got.Interval != 3 { + t.Errorf("Interval = %d, want 3", got.Interval) + } +} + +func TestReadMissing(t *testing.T) { + got := state{Interval: 7} + err := Read(filepath.Join(t.TempDir(), "missing.json"), &got) + if !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("Read() error = %v, want fs.ErrNotExist", err) + } + if got.Interval != 7 { + t.Errorf("Read() changed v to %+v", got) + } +} + +func TestReadIgnoresTornTemporaryFile(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "state.json") + if err := Write(path, state{Interval: 2}); err != nil { + t.Fatal(err) + } + + // A crash during a write leaves a partial temporary file next to the + // complete one. Read must still see the complete file. + if err := os.WriteFile(filepath.Join(dir, ".state.json.tmp-123"), []byte(`{"interv`), 0o600); err != nil { + t.Fatal(err) + } + + var got state + if err := Read(path, &got); err != nil { + t.Fatal(err) + } + if got.Interval != 2 { + t.Errorf("Interval = %d, want 2", got.Interval) + } +} + +func TestReadCorrupt(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + if err := os.WriteFile(path, []byte("{"), 0o600); err != nil { + t.Fatal(err) + } + + var got state + if err := Read(path, &got); err == nil { + t.Fatal("Read() = nil, want an error for invalid JSON") + } +} diff --git a/main_test.go b/main_test.go index 921b916..e9de3ab 100644 --- a/main_test.go +++ b/main_test.go @@ -72,13 +72,20 @@ func newEnv(t *testing.T, services ...string) *env { // the test if fly does not finish within 10 seconds. func (e *env) run(t *testing.T, dir string, args ...string) result { t.Helper() + return runFly(t, dir, append(append(e.docker.Env(), "HOME="+e.home), e.vars...), args...) +} + +// runFly executes fly in dir with the environment environ. It kills fly and +// fails the test if fly does not finish within 10 seconds. +func runFly(t *testing.T, dir string, environ []string, args ...string) result { + t.Helper() ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() cmd := exec.CommandContext(ctx, flyBin, args...) cmd.Dir = dir - cmd.Env = append(append(e.docker.Env(), "HOME="+e.home), e.vars...) + cmd.Env = environ var stdout, stderr bytes.Buffer cmd.Stdout, cmd.Stderr = &stdout, &stderr From 8179a0666d09127e604206a3de5110a33a645fbf Mon Sep 17 00:00:00 2001 From: nabil1440 <52530910+nabil1440@users.noreply.github.com> Date: Tue, 22 Sep 2026 16:16:45 +0600 Subject: [PATCH 2/3] fix(agent): refuse redirects, cap Retry-After, and never tick two times Fixes from the adversarial review of this layer. - The control plane client does not follow redirects. A followed redirect replays a POST as a GET without the body, so a 200 to it would drop samples that were never stored, and a redirect to http would send the token in plain text. A 3xx is now a reply that keeps the data. - A Retry-After value is at most one hour and cannot overflow. - A wall clock that steps back cannot make the same tick run two times. - FLY_AGENT_URL must not hold a user, a password, a query or a fragment, and an error never shows a password. localhost is accepted in any case. - FLY_AGENT_TOKEN may hold only printable ASCII characters. - At start, remove the temporary files that a crash during a write left. Refs #26 --- internal/agent/agent.go | 27 ++++++++++++++++++++++--- internal/agent/agent_test.go | 18 +++++++++++++++++ internal/agent/client.go | 29 ++++++++++++++++++++++----- internal/agent/client_test.go | 30 ++++++++++++++++++++++++++++ internal/agent/config.go | 21 ++++++++++++------- internal/agent/config_test.go | 29 +++++++++++++++++++++++---- internal/statefile/statefile.go | 14 ++++++++++++- internal/statefile/statefile_test.go | 26 ++++++++++++++++++++++++ 8 files changed, 174 insertions(+), 20 deletions(-) diff --git a/internal/agent/agent.go b/internal/agent/agent.go index 3348ea0..11f53e9 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -32,6 +32,9 @@ type agent struct { // since the last report. interval int pending int + + // last is the time of the last tick. + last time.Time } // Run runs the agent until ctx is done. Only one agent can run with the same @@ -43,6 +46,9 @@ func Run(ctx context.Context, cfg Config, log *slog.Logger) error { } defer unlock() + // A crash during a write can leave a temporary file. + statefile.RemoveTemp(cfg.StateDir) + a := &agent{cfg: cfg, log: log, interval: loadInterval(cfg.StateDir, log)} log.Info("agent started", "version", version.Version, "offset", cfg.Offset(), "report_interval", a.interval) a.loop(ctx) @@ -54,17 +60,32 @@ func Run(ctx context.Context, cfg Config, log *slog.Logger) error { // loop calls tick at the offset second of each minute until ctx is done. func (a *agent) loop(ctx context.Context) { for { - timer := time.NewTimer(time.Until(nextTick(time.Now(), a.cfg.Offset()))) + next := nextAfter(time.Now(), a.last, a.cfg.Offset()) + timer := time.NewTimer(time.Until(next)) select { case <-ctx.Done(): timer.Stop() return - case now := <-timer.C: - a.tick(ctx, now) + case <-timer.C: + a.last = next + a.tick(ctx, next) } } } +// nextAfter returns the next tick after now, and never a tick at or before +// last. The timer runs on the monotonic clock, but the tick times come from +// the wall clock: when the wall clock steps back, the timer fires before the +// tick time, and without last the same tick would run two times. +func nextAfter(now, last time.Time, offset time.Duration) time.Time { + next := nextTick(now, offset) + if !last.IsZero() && !next.After(last) { + next = nextTick(last, offset) + } + + return next +} + // tick does the work of one minute: it takes a sample and, after each // interval samples, sends a report. func (a *agent) tick(ctx context.Context, now time.Time) { diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 2f24eab..0c0cdc8 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -66,6 +66,24 @@ func TestNextTick(t *testing.T) { } } +func TestNextAfterAClockStepBack(t *testing.T) { + base := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + last := base.Add(17 * time.Second) + + // The wall clock stepped back 2 s after the tick at :17, so the timer + // fired at :15 wall time. The next tick must be in the next minute. + if got, want := nextAfter(base.Add(15*time.Second), last, 17*time.Second), base.Add(77*time.Second); !got.Equal(want) { + t.Errorf("nextAfter() = %s, want %s: never the same tick two times", got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + // Without a step, the last tick does not change the result. + if got, want := nextAfter(base.Add(30*time.Second), last, 17*time.Second), base.Add(77*time.Second); !got.Equal(want) { + t.Errorf("nextAfter() = %s, want %s", got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + if got, want := nextAfter(base, time.Time{}, 17*time.Second), last; !got.Equal(want) { + t.Errorf("nextAfter() without a last tick = %s, want %s", got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } +} + func TestLoopTicksAtTheOffsetAndReportsEachInterval(t *testing.T) { synctest.Test(t, func(t *testing.T) { dir := t.TempDir() diff --git a/internal/agent/client.go b/internal/agent/client.go index 893854b..a7b8d29 100644 --- a/internal/agent/client.go +++ b/internal/agent/client.go @@ -20,6 +20,10 @@ const requestTimeout = 20 * time.Second // maxReplySize limits the reply body that the agent reads. const maxReplySize = 1 << 20 +// maxRetryAfter limits the wait that a Retry-After header can ask for, so +// that one bad reply cannot stop the sends for a long time. +const maxRetryAfter = time.Hour + // Client sends requests to the control plane. type Client struct { base *url.URL @@ -28,13 +32,24 @@ type Client struct { } // NewClient returns a client for the control plane of cfg. A nil hc uses a -// client with the default timeout. +// client with the default timeout. The client never follows a redirect: see +// noRedirects. func NewClient(cfg Config, hc *http.Client) *Client { if hc == nil { hc = &http.Client{Timeout: requestTimeout} } + c := *hc + c.CheckRedirect = noRedirects + + return &Client{base: cfg.URL, token: cfg.Token, http: &c} +} - return &Client{base: cfg.URL, token: cfg.Token, http: hc} +// noRedirects makes a redirect a reply like any other status that is not 200. +// A followed redirect replays a POST as a GET without the body, so a 200 to +// that GET would drop samples that were never stored. It can also send the +// token over plain http. +func noRedirects(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse } // StatusError is a reply from the control plane that is not 200 OK. @@ -93,15 +108,19 @@ func (c *Client) do(ctx context.Context, method, path string, in, out any) error } // retryAfter reads a Retry-After value: a number of seconds or an HTTP date. +// The wait is at most maxRetryAfter. func retryAfter(v string, now time.Time) time.Duration { if v == "" { return 0 } - if s, err := strconv.Atoi(v); err == nil { - return max(time.Duration(s)*time.Second, 0) + if s, err := strconv.ParseInt(v, 10, 64); err == nil { + if s <= 0 { + return 0 + } + return time.Duration(min(s, int64(maxRetryAfter/time.Second))) * time.Second } if t, err := http.ParseTime(v); err == nil { - return max(t.Sub(now), 0) + return min(max(t.Sub(now), 0), maxRetryAfter) } return 0 diff --git a/internal/agent/client_test.go b/internal/agent/client_test.go index 143b99c..2f91bcc 100644 --- a/internal/agent/client_test.go +++ b/internal/agent/client_test.go @@ -123,6 +123,11 @@ func TestRetryAfter(t *testing.T) { "soon": 0, "Tue, 22 Sep 2026 10:01:30 GMT": 90 * time.Second, "Tue, 22 Sep 2026 09:00:00 GMT": 0, + // One bad reply must not stop the sends for years. + "31536000": time.Hour, + "10000000000": time.Hour, + "20000000000": time.Hour, + "Tue, 22 Sep 2027 10:00:00 GMT": time.Hour, } for v, want := range tests { @@ -131,3 +136,28 @@ func TestRetryAfter(t *testing.T) { } } } + +func TestClientDoesNotFollowRedirects(t *testing.T) { + var targetHits int + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + targetHits++ + _, _ = w.Write([]byte(`{"report_interval": 3}`)) + })) + defer target.Close() + + for _, code := range []int{http.StatusMovedPermanently, http.StatusFound, http.StatusTemporaryRedirect, http.StatusPermanentRedirect} { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + // For example a proxy during a deploy, or a redirect to plain http. + http.Redirect(w, r, target.URL+r.URL.Path, code) + }) + + err := c.do(context.Background(), http.MethodPost, "agent/v1/metrics", map[string]int{"samples": 5}, nil) + var statusErr *StatusError + if !errors.As(err, &statusErr) || statusErr.StatusCode != code { + t.Errorf("do() after a %d = %v, want a *StatusError with %d: the samples must stay in the queue", code, err, code) + } + } + if targetHits != 0 { + t.Errorf("the redirect target got %d requests, want none (the token must not go there)", targetHits) + } +} diff --git a/internal/agent/config.go b/internal/agent/config.go index 0c34bae..cdd081f 100644 --- a/internal/agent/config.go +++ b/internal/agent/config.go @@ -12,7 +12,6 @@ import ( "strconv" "strings" "time" - "unicode" ) // The environment keys that the FlyWP installer writes to /etc/fly/agent.env, @@ -62,8 +61,10 @@ func ConfigFromEnv(getenv func(string) string) (Config, error) { switch { case cfg.Token == "": errs = append(errs, fmt.Errorf("%s is not set", EnvToken)) - case strings.ContainsFunc(cfg.Token, func(r rune) bool { return unicode.IsSpace(r) || unicode.IsControl(r) }): - errs = append(errs, fmt.Errorf("%s contains spaces or control characters", EnvToken)) + case strings.ContainsFunc(cfg.Token, func(r rune) bool { return r < '!' || r > '~' }): + // The token is opaque, but it goes in a header: only printable ASCII + // (the format is flyagt_ and base62) is safe there. + errs = append(errs, fmt.Errorf("%s may contain only printable ASCII characters, without spaces", EnvToken)) } if v := getenv(EnvServerID); v == "" { @@ -84,7 +85,9 @@ func ConfigFromEnv(getenv func(string) string) (Config, error) { } // parseURL accepts an https URL. It also accepts http for a loopback host, -// for tests and local development: plain http never leaves the machine. +// for tests and local development: plain http never leaves the machine. The +// URL is a base URL only: no user, password, query or fragment. An error +// never shows a password. func parseURL(v string) (*url.URL, error) { if v == "" { return nil, fmt.Errorf("%s is not set", EnvURL) @@ -92,21 +95,25 @@ func parseURL(v string) (*url.URL, error) { u, err := url.Parse(strings.TrimRight(v, "/")) if err != nil || u.Host == "" { - return nil, fmt.Errorf("%s is not a valid URL: %q", EnvURL, v) + return nil, fmt.Errorf("%s is not a valid URL", EnvURL) } switch { + case u.User != nil: + return nil, fmt.Errorf("%s must not contain a user or a password: %s", EnvURL, u.Redacted()) + case u.RawQuery != "" || u.ForceQuery || u.Fragment != "": + return nil, fmt.Errorf("%s must not contain a query or a fragment: %s", EnvURL, u.Redacted()) case u.Scheme == "https": case u.Scheme == "http" && isLoopback(u.Hostname()): default: - return nil, fmt.Errorf("%s must be an https URL, not %q", EnvURL, v) + return nil, fmt.Errorf("%s must be an https URL, not %s", EnvURL, u.Redacted()) } return u, nil } func isLoopback(host string) bool { - if host == "localhost" { + if strings.EqualFold(host, "localhost") { return true } ip := net.ParseIP(host) diff --git a/internal/agent/config_test.go b/internal/agent/config_test.go index 4503bcf..21e9a7d 100644 --- a/internal/agent/config_test.go +++ b/internal/agent/config_test.go @@ -78,7 +78,13 @@ func TestConfigFromEnvErrors(t *testing.T) { {"other scheme", EnvURL, "ftp://app.flywp.com", "must be an https URL"}, {"no host", EnvURL, "https://", "not a valid URL"}, {"no token", EnvToken, "", "FLY_AGENT_TOKEN is not set"}, - {"token with newline", EnvToken, "flyagt_abc\nX-Other: 1", "FLY_AGENT_TOKEN contains spaces"}, + {"token with newline", EnvToken, "flyagt_abc\nX-Other: 1", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token with a space", EnvToken, "flyagt_abc def", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token that is not ASCII", EnvToken, "flyagt_é", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token with invalid UTF-8", EnvToken, "flyagt_\x85", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token with a zero-width space", EnvToken, "flyagt_\u200b", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"URL with a query", EnvURL, "https://app.flywp.com?x=1", "must not contain a query"}, + {"URL with a fragment", EnvURL, "https://app.flywp.com#x", "must not contain a query or a fragment"}, {"no server id", EnvServerID, "", "FLY_AGENT_SERVER_ID is not set"}, {"negative server id", EnvServerID, "-1", "FLY_AGENT_SERVER_ID must be an integer"}, {"text server id", EnvServerID, "abc", "FLY_AGENT_SERVER_ID must be an integer"}, @@ -115,12 +121,27 @@ func TestConfigFromEnvNamesEachProblem(t *testing.T) { } } -func TestConfigAcceptsLoopbackHTTP(t *testing.T) { - for _, u := range []string{"http://127.0.0.1:8080", "http://[::1]:8080", "http://localhost:8080"} { +func TestConfigAcceptsHTTPSAndLoopbackHTTP(t *testing.T) { + for _, u := range []string{"http://127.0.0.1:8080", "http://[::1]:8080", "http://localhost:8080", "http://LOCALHOST:8080", "HTTPS://app.flywp.com"} { env := validEnv(t) env[EnvURL] = u if _, err := ConfigFromEnv(getenv(env)); err != nil { - t.Errorf("ConfigFromEnv(%s) error = %v, want loopback http accepted", u, err) + t.Errorf("ConfigFromEnv(%s) error = %v, want it accepted", u, err) + } + } +} + +func TestConfigURLWithPasswordDoesNotShowIt(t *testing.T) { + for _, u := range []string{"https://user:s3cret@app.flywp.com", "http://user:s3cret@example.com"} { + env := validEnv(t) + env[EnvURL] = u + + _, err := ConfigFromEnv(getenv(env)) + if err == nil || !strings.Contains(err.Error(), "must not contain a user or a password") { + t.Errorf("ConfigFromEnv(%s) error = %v, want a refusal", u, err) + } + if err != nil && strings.Contains(err.Error(), "s3cret") { + t.Errorf("error %q shows the password", err) } } } diff --git a/internal/statefile/statefile.go b/internal/statefile/statefile.go index a89e756..d3ccdac 100644 --- a/internal/statefile/statefile.go +++ b/internal/statefile/statefile.go @@ -8,6 +8,9 @@ import ( "path/filepath" ) +// tempSuffix marks the temporary files of Write. +const tempSuffix = ".tmp-" + // Write stores v as JSON in path. It writes a temporary file in the same // directory, syncs it and renames it over path, so path always holds a // complete file, also after a crash or a power loss. @@ -18,7 +21,7 @@ func Write(path string, v any) (err error) { } dir := filepath.Dir(path) - tmp, err := os.CreateTemp(dir, "."+filepath.Base(path)+".tmp-*") + tmp, err := os.CreateTemp(dir, "."+filepath.Base(path)+tempSuffix+"*") if err != nil { return err } @@ -65,3 +68,12 @@ func Read(path string, v any) error { return nil } + +// RemoveTemp removes the temporary files that Write leaves in dir after a +// crash. Call it before any Write in dir starts. +func RemoveTemp(dir string) { + matches, _ := filepath.Glob(filepath.Join(dir, ".*"+tempSuffix+"*")) + for _, m := range matches { + _ = os.Remove(m) + } +} diff --git a/internal/statefile/statefile_test.go b/internal/statefile/statefile_test.go index 11bc355..049c969 100644 --- a/internal/statefile/statefile_test.go +++ b/internal/statefile/statefile_test.go @@ -101,3 +101,29 @@ func TestReadCorrupt(t *testing.T) { t.Fatal("Read() = nil, want an error for invalid JSON") } } + +func TestRemoveTemp(t *testing.T) { + dir := t.TempDir() + if err := Write(filepath.Join(dir, "state.json"), state{Interval: 1}); err != nil { + t.Fatal(err) + } + for _, name := range []string{".state.json.tmp-123", ".samples.json.tmp-9"} { + if err := os.WriteFile(filepath.Join(dir, name), []byte("{"), 0o600); err != nil { + t.Fatal(err) + } + } + + RemoveTemp(dir) + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + if len(entries) != 1 || entries[0].Name() != "state.json" { + var names []string + for _, e := range entries { + names = append(names, e.Name()) + } + t.Errorf("directory holds %v, want only state.json", names) + } +} From bad49b1d47b018a27e6def417ab479cc3e54c1e6 Mon Sep 17 00:00:00 2001 From: nabil1440 <52530910+nabil1440@users.noreply.github.com> Date: Tue, 22 Sep 2026 16:32:38 +0600 Subject: [PATCH 3/3] test(docker): give the probe timeout test room under load Under go test ./... -race, the fake "docker compose version" could take more than the 1 s test timeout, so the test reported the Compose plugin instead of the daemon. Use 3 s. Refs #26 --- internal/docker/probe_test.go | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/internal/docker/probe_test.go b/internal/docker/probe_test.go index 7f6d386..fbe8406 100644 --- a/internal/docker/probe_test.go +++ b/internal/docker/probe_test.go @@ -99,8 +99,11 @@ func TestStatusWithoutDockerCLI(t *testing.T) { func TestCheckTimesOut(t *testing.T) { useFakeDocker(t, "daemon-hang") + // Under load (go test ./... -race), the fake "docker compose version" + // can take more than 1 s, and would then time out first. 3 s is short + // enough for the test and long enough for the fake. old := probeTimeout - probeTimeout = time.Second + probeTimeout = 3 * time.Second t.Cleanup(func() { probeTimeout = old }) start := time.Now() @@ -110,7 +113,7 @@ func TestCheckTimesOut(t *testing.T) { if !errors.As(err, &unavailable) || unavailable.Part != PartDaemon || !strings.Contains(unavailable.Detail, "no answer") { t.Errorf("Check() = %v, want the daemon to be reported as not answering", err) } - if elapsed := time.Since(start); elapsed > 5*time.Second { + if elapsed := time.Since(start); elapsed > 10*time.Second { t.Errorf("Check() took %s, want it to stop soon after the timeout", elapsed) } }