diff --git a/README.md b/README.md index 2bdbb09d..73e34147 100644 --- a/README.md +++ b/README.md @@ -50,6 +50,38 @@ Each deployment takes over all the traffic from the previously deployed instance. As soon as Kamal Proxy determines that the new instance is healthy, it will route all new traffic to that instance. +### Opt-in scale to zero + +Services can stop their write containers after an idle period and wake them on +the next application request: + + kamal-proxy run --docker-socket /var/run/docker.sock + kamal-proxy deploy service1 --target web-1:3000 --idle-timeout 15m --idle-wake-timeout 30s + +`--idle-timeout` defaults to `0` (disabled). `--idle-wake-timeout` defaults to +`30s` and bounds how long each request waits for Docker start and a successful +configured health check. The target hostname (`web-1` above) must be the Docker +container name. `DOCKER_SOCKET` and `KAMAL_PROXY_DOCKER_SOCKET` are equivalents +of the run flag. + +Requests are held before their bodies are read, so POST bodies are forwarded +unchanged after a successful wake. Concurrent wake requests are coalesced. +Open streaming responses and WebSockets count as activity/in-flight work and +prevent sleeping until they close; a new stream or WebSocket is held during +wake like any other request. Health-check requests do not wake or reset an idle +service and receive success while it is stopping, sleeping, or waking. + +Mounting the Docker socket gives the proxy host-level container control. Only +enable this feature where that trust is acceptable; the lifecycle calls are +isolated behind the `ContainerLifecycle` interface so they can be moved to an +external service later. + +The Docker client negotiates the API version once from the daemon's unversioned +`/version` endpoint and caches it for start/stop calls. If that endpoint is +unavailable or returns a non-success status, it falls back to the legacy +`v1.41` paths for compatibility with restricted socket proxies; a successful +but malformed version response is rejected instead of guessing. + The `deploy` command also waits for traffic to drain from the old instance before returning. This means it's safe to remove the old instance as soon as `deploy` returns successfully, without interrupting any in-flight requests. diff --git a/internal/cmd/deploy.go b/internal/cmd/deploy.go index 82af0e18..cfc5e2da 100644 --- a/internal/cmd/deploy.go +++ b/internal/cmd/deploy.go @@ -52,6 +52,9 @@ func newDeployCommand() *deployCommand { deployCommand.cmd.Flags().DurationVar(&deployCommand.args.ServiceOptions.WriterAffinityTimeout, "writer-affinity-timeout", server.DefaultWriterAffinityTimeout, "Time after a write before read requests will be routed to readers") deployCommand.cmd.Flags().BoolVar(&deployCommand.args.ServiceOptions.ReadTargetsAcceptWebsockets, "read-target-websockets", false, "Route WebSocket traffic to read targets, when available") + deployCommand.cmd.Flags().DurationVar(&deployCommand.args.ServiceOptions.IdleTimeout, "idle-timeout", 0, "Stop container after this duration of inactivity (0 to disable)") + deployCommand.cmd.Flags().DurationVar(&deployCommand.args.ServiceOptions.IdleWakeTimeout, "idle-wake-timeout", server.DefaultIdleWakeTimeout, "Max time to hold request while waking container") + deployCommand.cmd.Flags().DurationVar(&deployCommand.args.TargetOptions.ResponseTimeout, "target-timeout", server.DefaultTargetTimeout, "Maximum time to wait for the target server to respond when serving requests") deployCommand.cmd.Flags().BoolVar(&deployCommand.args.TargetOptions.BufferRequests, "buffer-requests", false, "Buffer requests before forwarding to target") diff --git a/internal/cmd/run.go b/internal/cmd/run.go index 0a70655e..0831455b 100644 --- a/internal/cmd/run.go +++ b/internal/cmd/run.go @@ -29,6 +29,7 @@ func newRunCommand() *runCommand { runCommand.cmd.Flags().IntVar(&globalConfig.HttpsPort, "https-port", getEnvInt("HTTPS_PORT", server.DefaultHttpsPort), "Port to serve HTTPS traffic on") runCommand.cmd.Flags().IntVar(&globalConfig.MetricsPort, "metrics-port", getEnvInt("METRICS_PORT", 0), "Publish metrics on the specified port (default zero to disable)") runCommand.cmd.Flags().BoolVar(&globalConfig.HTTP3Enabled, "http3", false, "Enable HTTP/3") + runCommand.cmd.Flags().StringVar(&globalConfig.DockerSocketPath, "docker-socket", getEnvString("DOCKER_SOCKET", server.DefaultDockerSocketPath), "Path to Docker socket") return runCommand } @@ -36,8 +37,10 @@ func newRunCommand() *runCommand { func (c *runCommand) run(cmd *cobra.Command, args []string) error { c.setLogger() - router := server.NewRouter(globalConfig.StatePath()) - router.RestoreLastSavedState() + router := server.NewRouter(globalConfig.StatePath(), globalConfig.DockerSocketPath) + if err := router.RestoreLastSavedState(); err != nil { + return err + } s := server.NewServer(&globalConfig, router) err := s.Start() diff --git a/internal/cmd/run_test.go b/internal/cmd/run_test.go new file mode 100644 index 00000000..3269eea5 --- /dev/null +++ b/internal/cmd/run_test.go @@ -0,0 +1,20 @@ +package cmd + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestRunCommandReturnsStateRestoreError(t *testing.T) { + previous := globalConfig + t.Cleanup(func() { globalConfig = previous }) + globalConfig.AlternateConfigDir = t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(globalConfig.AlternateConfigDir, "kamal-proxy.state"), []byte("invalid"), 0o600)) + + err := newRunCommand().run(nil, nil) + + require.ErrorContains(t, err, "invalid character 'i'") +} diff --git a/internal/cmd/util.go b/internal/cmd/util.go index d28618c6..000a303b 100644 --- a/internal/cmd/util.go +++ b/internal/cmd/util.go @@ -47,6 +47,14 @@ func getEnvInt(key string, defaultValue int) int { return intValue } +func getEnvString(key, defaultValue string) string { + value, ok := findEnv(key) + if !ok { + return defaultValue + } + return value +} + func getEnvBool(key string, defaultValue bool) bool { value, ok := findEnv(key) if !ok { diff --git a/internal/server/config.go b/internal/server/config.go index 66f632a2..92621f06 100644 --- a/internal/server/config.go +++ b/internal/server/config.go @@ -8,8 +8,9 @@ import ( ) const ( - DefaultHttpPort = 80 - DefaultHttpsPort = 443 + DefaultHttpPort = 80 + DefaultHttpsPort = 443 + DefaultDockerSocketPath = "/var/run/docker.sock" ) type Config struct { @@ -20,6 +21,7 @@ type Config struct { HTTP3Enabled bool AlternateConfigDir string + DockerSocketPath string } func (c Config) SocketPath() string { diff --git a/internal/server/docker_client.go b/internal/server/docker_client.go new file mode 100644 index 00000000..fc2c9ff0 --- /dev/null +++ b/internal/server/docker_client.go @@ -0,0 +1,141 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "strings" + "sync" +) + +const ( + legacyDockerAPIVersion = "1.41" + maxDockerErrorBody = 4096 +) + +type DockerClient struct { + httpClient *http.Client + + versionMu sync.Mutex + versionSet bool + apiVersion string + versionErr error +} + +func NewDockerClient(socketPath string) *DockerClient { + return &DockerClient{ + httpClient: &http.Client{ + Transport: &http.Transport{ + DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { + return (&net.Dialer{}).DialContext(ctx, "unix", socketPath) + }, + }, + }, + } +} + +func (c *DockerClient) StopContainer(ctx context.Context, name string) error { + return c.containerAction(ctx, name, "stop") +} + +func (c *DockerClient) StartContainer(ctx context.Context, name string) error { + return c.containerAction(ctx, name, "start") +} + +func (c *DockerClient) containerAction(ctx context.Context, name, action string) error { + version, err := c.negotiatedVersion(ctx) + if err != nil { + return err + } + endpoint := fmt.Sprintf("http://localhost/v%s/containers/%s/%s", version, url.PathEscape(name), action) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, nil) + if err != nil { + return err + } + resp, err := c.httpClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusNoContent && resp.StatusCode != http.StatusNotModified { + return dockerResponseError(action, resp) + } + return nil +} + +func (c *DockerClient) negotiatedVersion(ctx context.Context) (string, error) { + c.versionMu.Lock() + defer c.versionMu.Unlock() + if c.versionSet { + return c.apiVersion, c.versionErr + } + + negotiationCtx, cancel := dockerNegotiationContext(ctx) + defer cancel() + version, err := c.queryVersion(negotiationCtx) + if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + c.apiVersion, c.versionErr, c.versionSet = version, err, true + } + return version, err +} + +func dockerNegotiationContext(ctx context.Context) (context.Context, context.CancelFunc) { + if deadline, ok := ctx.Deadline(); ok { + return context.WithDeadline(context.Background(), deadline) + } + return context.WithCancel(context.Background()) +} + +func (c *DockerClient) queryVersion(ctx context.Context) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, "http://localhost/version", nil) + if err != nil { + return "", err + } + resp, err := c.httpClient.Do(req) + if err != nil { + if ctx.Err() != nil { + return "", ctx.Err() + } + // Some compatible Docker proxies do not expose /version. Preserve the + // legacy behavior and let the versioned operation return the useful error. + return legacyDockerAPIVersion, nil + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return legacyDockerAPIVersion, nil + } + var version struct { + APIVersion string `json:"ApiVersion"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, maxDockerErrorBody+1)).Decode(&version); err != nil { + return "", fmt.Errorf("invalid docker /version response: %w", err) + } + if version.APIVersion == "" { + return "", errors.New("docker /version response has no ApiVersion") + } + return version.APIVersion, nil +} + +func dockerResponseError(action string, resp *http.Response) error { + body, err := io.ReadAll(io.LimitReader(resp.Body, maxDockerErrorBody+1)) + if err != nil { + return fmt.Errorf("docker %s returned status %d (reading error body: %w)", action, resp.StatusCode, err) + } + truncated := len(body) > maxDockerErrorBody + if truncated { + body = body[:maxDockerErrorBody] + } + message := strings.TrimSpace(string(body)) + if message == "" { + return fmt.Errorf("docker %s returned status %d", action, resp.StatusCode) + } + if truncated { + message += "…" + } + return fmt.Errorf("docker %s returned status %d: %s", action, resp.StatusCode, message) +} diff --git a/internal/server/docker_client_test.go b/internal/server/docker_client_test.go new file mode 100644 index 00000000..284da437 --- /dev/null +++ b/internal/server/docker_client_test.go @@ -0,0 +1,160 @@ +package server + +import ( + "context" + "fmt" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDockerClientNegotiatesNewDaemonAndUsesStartStopPaths(t *testing.T) { + var pathsMu sync.Mutex + var paths []string + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { + pathsMu.Lock() + paths = append(paths, r.URL.Path) + pathsMu.Unlock() + if r.URL.Path == "/version" { + fmt.Fprint(w, `{"ApiVersion":"1.52","MinAPIVersion":"1.44"}`) + return + } + w.WriteHeader(http.StatusNoContent) + }) + + require.NoError(t, client.StopContainer(context.Background(), "test-container")) + require.NoError(t, client.StartContainer(context.Background(), "test-container")) + assert.Equal(t, []string{ + "/version", + "/v1.52/containers/test-container/stop", + "/v1.52/containers/test-container/start", + }, paths) +} + +func TestDockerClientNegotiatesOldCompatibleDaemon(t *testing.T) { + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/version" { + fmt.Fprint(w, `{"ApiVersion":"1.41","MinAPIVersion":"1.24"}`) + return + } + assert.Equal(t, "/v1.41/containers/web/stop", r.URL.Path) + w.WriteHeader(http.StatusNoContent) + }) + require.NoError(t, client.StopContainer(context.Background(), "web")) +} + +func TestDockerClientFallsBackWhenVersionEndpointFails(t *testing.T) { + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/version" { + http.Error(w, "not exposed", http.StatusNotFound) + return + } + assert.Equal(t, "/v1.41/containers/web/start", r.URL.Path) + w.WriteHeader(http.StatusNoContent) + }) + require.NoError(t, client.StartContainer(context.Background(), "web")) +} + +func TestDockerClientRejectsMalformedVersionResponse(t *testing.T) { + for name, body := range map[string]string{ + "malformed json": `{`, + "missing version": `{}`, + } { + t.Run(name, func(t *testing.T) { + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, body) }) + err := client.StartContainer(context.Background(), "web") + require.Error(t, err) + assert.Contains(t, err.Error(), "version") + }) + } +} + +func TestDockerClientNegotiatesOnlyOnceWithConcurrentFirstUse(t *testing.T) { + var versionCalls atomic.Int32 + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/version" { + versionCalls.Add(1) + fmt.Fprint(w, `{"ApiVersion":"1.52","MinAPIVersion":"1.44"}`) + return + } + assert.True(t, strings.HasPrefix(r.URL.Path, "/v1.52/containers/")) + w.WriteHeader(http.StatusNoContent) + }) + + var wg sync.WaitGroup + for range 20 { + wg.Add(1) + go func() { + defer wg.Done() + require.NoError(t, client.StartContainer(context.Background(), "web")) + }() + } + wg.Wait() + assert.Equal(t, int32(1), versionCalls.Load()) +} + +func TestDockerClientNegotiationSurvivesCallerCancellation(t *testing.T) { + versionStarted := make(chan struct{}) + releaseVersion := make(chan struct{}) + var startedOnce sync.Once + var actionPath string + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/version" { + startedOnce.Do(func() { close(versionStarted) }) + select { + case <-releaseVersion: + fmt.Fprint(w, `{"ApiVersion":"1.52","MinAPIVersion":"1.44"}`) + case <-r.Context().Done(): + } + return + } + actionPath = r.URL.Path + w.WriteHeader(http.StatusNoContent) + }) + + ctx, cancel := context.WithCancel(context.Background()) + firstDone := make(chan error, 1) + go func() { firstDone <- client.StartContainer(ctx, "web") }() + <-versionStarted + cancel() + close(releaseVersion) + require.Error(t, <-firstDone) + + require.NoError(t, client.StartContainer(context.Background(), "web")) + assert.Equal(t, "/v1.52/containers/web/start", actionPath) +} + +func TestDockerClientIncludesBoundedDockerErrorBody(t *testing.T) { + client := testDockerClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/version" { + fmt.Fprint(w, `{"ApiVersion":"1.52","MinAPIVersion":"1.44"}`) + return + } + w.WriteHeader(http.StatusBadRequest) + fmt.Fprint(w, "minimum API 1.44: "+strings.Repeat("x", maxDockerErrorBody*2)) + }) + err := client.StopContainer(context.Background(), "web") + require.Error(t, err) + assert.Contains(t, err.Error(), "minimum API 1.44") + assert.LessOrEqual(t, len(err.Error()), maxDockerErrorBody+100) + assert.True(t, strings.HasSuffix(err.Error(), "…")) +} + +func testDockerClient(t *testing.T, handler http.HandlerFunc) *DockerClient { + t.Helper() + server := httptest.NewUnstartedServer(handler) + socketPath := t.TempDir() + "/docker.sock" + listener, err := net.Listen("unix", socketPath) + require.NoError(t, err) + server.Listener = listener + server.Start() + t.Cleanup(server.Close) + return NewDockerClient(socketPath) +} diff --git a/internal/server/health_check.go b/internal/server/health_check.go index 59eeba30..f8d48ead 100644 --- a/internal/server/health_check.go +++ b/internal/server/health_check.go @@ -8,6 +8,7 @@ import ( "log/slog" "net/http" "net/url" + "sync" "time" ) @@ -33,9 +34,16 @@ type HealthCheck struct { ctx context.Context cancel context.CancelFunc + start sync.Once } func NewHealthCheck(consumer HealthCheckConsumer, endpoint *url.URL, interval time.Duration, timeout time.Duration, host string) *HealthCheck { + hc := newHealthCheck(consumer, endpoint, interval, timeout, host) + hc.Start() + return hc +} + +func newHealthCheck(consumer HealthCheckConsumer, endpoint *url.URL, interval time.Duration, timeout time.Duration, host string) *HealthCheck { ctx, cancel := context.WithCancel(context.Background()) hc := &HealthCheck{ @@ -49,10 +57,13 @@ func NewHealthCheck(consumer HealthCheckConsumer, endpoint *url.URL, interval ti cancel: cancel, } - go hc.run() return hc } +func (hc *HealthCheck) Start() { + hc.start.Do(func() { go hc.run() }) +} + func (hc *HealthCheck) Close() { hc.cancel() } diff --git a/internal/server/idle_controller.go b/internal/server/idle_controller.go new file mode 100644 index 00000000..39f2e994 --- /dev/null +++ b/internal/server/idle_controller.go @@ -0,0 +1,349 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "log/slog" + "sync" + "time" +) + +type IdleState int + +const ( + IdleStateActive IdleState = iota + IdleStateStopping + IdleStateSleeping + IdleStateWaking +) + +func (s IdleState) String() string { + switch s { + case IdleStateActive: + return "active" + case IdleStateStopping: + return "stopping" + case IdleStateSleeping: + return "sleeping" + case IdleStateWaking: + return "waking" + default: + return "" + } +} + +var ErrIdleWakeTimeout = errors.New("idle container wake timed out") + +type ContainerLifecycle interface { + StartContainer(context.Context, string) error + StopContainer(context.Context, string) error +} + +type IdleController struct { + State IdleState `json:"state"` + IdleTimeout time.Duration `json:"idle_timeout"` + WakeTimeout time.Duration `json:"wake_timeout"` + ContainerNames []string `json:"container_names"` + + mu sync.Mutex + lifecycle ContainerLifecycle + ready func(time.Duration) error + inflight int + lastRequest time.Time + wakeDone chan struct{} + wakeErr error + lifecycleCancel context.CancelFunc + changed chan struct{} + closed chan struct{} + disabled bool + closeOnce sync.Once + persist func() +} + +func NewIdleController(idleTimeout, wakeTimeout time.Duration, names []string, lifecycle ContainerLifecycle, ready func(time.Duration) error) *IdleController { + c := &IdleController{State: IdleStateActive} + c.configure(idleTimeout, wakeTimeout, names, lifecycle, ready) + return c +} + +func (c *IdleController) MarshalJSON() ([]byte, error) { + type persisted struct { + State IdleState `json:"state"` + IdleTimeout time.Duration `json:"idle_timeout"` + WakeTimeout time.Duration `json:"wake_timeout"` + ContainerNames []string `json:"container_names"` + } + c.mu.Lock() + defer c.mu.Unlock() + return json.Marshal(persisted{ + State: c.State, + IdleTimeout: c.IdleTimeout, + WakeTimeout: c.WakeTimeout, + ContainerNames: c.ContainerNames, + }) +} + +func (c *IdleController) UnmarshalJSON(data []byte) error { + type persisted struct { + State IdleState `json:"state"` + IdleTimeout time.Duration `json:"idle_timeout"` + WakeTimeout time.Duration `json:"wake_timeout"` + ContainerNames []string `json:"container_names"` + } + var p persisted + if err := json.Unmarshal(data, &p); err != nil { + return err + } + c.State, c.IdleTimeout, c.WakeTimeout, c.ContainerNames = p.State, p.IdleTimeout, p.WakeTimeout, p.ContainerNames + if c.State == IdleStateWaking || c.State == IdleStateStopping { + c.State = IdleStateSleeping + } + return nil +} + +func (c *IdleController) configure(idleTimeout, wakeTimeout time.Duration, names []string, lifecycle ContainerLifecycle, ready func(time.Duration) error) { + c.mu.Lock() + defer c.mu.Unlock() + c.IdleTimeout, c.WakeTimeout = idleTimeout, wakeTimeout + c.ContainerNames = append([]string(nil), names...) + c.lifecycle, c.ready = lifecycle, ready + c.lastRequest = time.Now() + if c.changed == nil { + c.changed, c.closed = make(chan struct{}, 1), make(chan struct{}) + go c.run() + } + c.signal() +} + +func (c *IdleController) BeginRequest(ctx context.Context) error { + c.mu.Lock() + c.inflight++ + c.lastRequest = time.Now() + c.mu.Unlock() + c.signal() + for { + c.mu.Lock() + if c.State == IdleStateSleeping && !c.disabled { + c.startWakeLocked() + } + done, timeout, state := c.wakeDone, c.WakeTimeout, c.State + c.mu.Unlock() + if state != IdleStateWaking && state != IdleStateStopping { + return nil + } + timer := time.NewTimer(timeout) + select { + case <-done: + timer.Stop() + if state == IdleStateStopping { + continue + } + c.mu.Lock() + err := c.wakeErr + c.mu.Unlock() + return err + case <-timer.C: + return ErrIdleWakeTimeout + case <-ctx.Done(): + timer.Stop() + return ctx.Err() + } + } +} + +func (c *IdleController) EndRequest() { + c.mu.Lock() + c.inflight-- + c.lastRequest = time.Now() + c.mu.Unlock() + c.signal() +} + +func (c *IdleController) StateValue() IdleState { c.mu.Lock(); defer c.mu.Unlock(); return c.State } + +func (c *IdleController) SetPersist(fn func()) { c.mu.Lock(); c.persist = fn; c.mu.Unlock() } + +func (c *IdleController) notifyPersist() { + c.mu.Lock() + fn := c.persist + c.mu.Unlock() + if fn != nil { + go fn() + } +} + +func (c *IdleController) Disable() { + c.mu.Lock() + c.disabled = true + if c.State == IdleStateWaking || c.State == IdleStateStopping { + c.cancelLifecycleLocked() + c.State = IdleStateSleeping + } + c.mu.Unlock() + c.signal() +} +func (c *IdleController) Enable() { + c.mu.Lock() + c.disabled = false + c.lastRequest = time.Now() + c.mu.Unlock() + c.signal() +} + +func (c *IdleController) Reset(names []string, ready func(time.Duration) error) { + c.mu.Lock() + c.ContainerNames = append([]string(nil), names...) + c.ready = ready + c.State, c.wakeErr, c.lastRequest = IdleStateActive, nil, time.Now() + c.cancelLifecycleLocked() + c.mu.Unlock() + c.signal() +} + +func (c *IdleController) cancelLifecycleLocked() { + if c.lifecycleCancel != nil { + c.lifecycleCancel() + c.lifecycleCancel = nil + } + if c.wakeDone != nil { + select { + case <-c.wakeDone: + default: + close(c.wakeDone) + } + c.wakeDone = nil + } +} + +func (c *IdleController) Close() { + c.closeOnce.Do(func() { + c.mu.Lock() + c.cancelLifecycleLocked() + c.mu.Unlock() + if c.closed != nil { + close(c.closed) + } + }) +} +func (c *IdleController) signal() { + select { + case c.changed <- struct{}{}: + default: + } +} + +func (c *IdleController) run() { + for { + c.mu.Lock() + wait := c.IdleTimeout - time.Since(c.lastRequest) + eligible := !c.disabled && c.State == IdleStateActive && c.inflight == 0 && c.IdleTimeout > 0 + c.mu.Unlock() + if !eligible { + wait = time.Hour + } + if wait < 0 { + wait = 0 + } + timer := time.NewTimer(wait) + select { + case <-timer.C: + c.trySleep() + case <-c.changed: + if !timer.Stop() { + <-timer.C + } + case <-c.closed: + if !timer.Stop() { + <-timer.C + } + return + } + } +} + +func (c *IdleController) trySleep() { + c.mu.Lock() + if c.disabled || c.State != IdleStateActive || c.inflight != 0 || c.IdleTimeout <= 0 || time.Since(c.lastRequest) < c.IdleTimeout { + c.mu.Unlock() + return + } + names, lifecycle := append([]string(nil), c.ContainerNames...), c.lifecycle + c.State, c.wakeDone = IdleStateStopping, make(chan struct{}) + stopDone := c.wakeDone + ctx, cancel := context.WithTimeout(context.Background(), DefaultIdleLifecycleTimeout) + c.lifecycleCancel = cancel + c.mu.Unlock() + defer cancel() + for _, name := range names { + if err := lifecycle.StopContainer(ctx, name); err != nil { + slog.Error("Failed to stop idle container", "container", name, "error", err) + c.mu.Lock() + if c.wakeDone == stopDone { + c.State = IdleStateSleeping + c.lifecycleCancel = nil + close(stopDone) + } + c.mu.Unlock() + c.notifyPersist() + return + } + } + c.mu.Lock() + if c.wakeDone == stopDone { + c.State = IdleStateSleeping + c.lifecycleCancel = nil + close(stopDone) + } + c.mu.Unlock() + c.notifyPersist() + c.signal() +} + +func (c *IdleController) startWakeLocked() { + c.State, c.wakeErr, c.wakeDone = IdleStateWaking, nil, make(chan struct{}) + done, names, timeout, lifecycle, ready := c.wakeDone, append([]string(nil), c.ContainerNames...), c.WakeTimeout, c.lifecycle, c.ready + ctx, cancel := context.WithTimeout(context.Background(), timeout) + deadline, _ := ctx.Deadline() + c.lifecycleCancel = cancel + go func() { + defer cancel() + var err error + for _, name := range names { + if ctx.Err() != nil { + err = ctx.Err() + break + } + if err = lifecycle.StartContainer(ctx, name); err != nil { + break + } + } + if err == nil { + c.mu.Lock() + current := c.wakeDone == done + c.mu.Unlock() + if current { + remaining := time.Until(deadline) + if remaining <= 0 { + err = ErrIdleWakeTimeout + } else { + err = ready(remaining) + } + } + } + c.mu.Lock() + if c.wakeDone == done { + c.wakeErr = err + c.lifecycleCancel = nil + if err == nil { + c.State, c.lastRequest = IdleStateActive, time.Now() + } else { + c.State = IdleStateSleeping + } + close(done) + } + c.mu.Unlock() + c.notifyPersist() + c.signal() + }() +} diff --git a/internal/server/idle_controller_test.go b/internal/server/idle_controller_test.go new file mode 100644 index 00000000..c6556228 --- /dev/null +++ b/internal/server/idle_controller_test.go @@ -0,0 +1,209 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type fakeLifecycle struct { + starts, stops atomic.Int32 + startErr, stopErr error + startDelay time.Duration + stopErrAt int32 +} + +type blockingLifecycle struct { + starts atomic.Int32 + started chan struct{} + finished chan struct{} +} + +func (f *blockingLifecycle) StartContainer(ctx context.Context, _ string) error { + f.starts.Add(1) + close(f.started) + <-ctx.Done() + close(f.finished) + return ctx.Err() +} + +func (f *blockingLifecycle) StopContainer(context.Context, string) error { return nil } + +func (f *fakeLifecycle) StartContainer(ctx context.Context, _ string) error { + f.starts.Add(1) + if f.startDelay > 0 { + select { + case <-time.After(f.startDelay): + case <-ctx.Done(): + return ctx.Err() + } + } + return f.startErr +} +func (f *fakeLifecycle) StopContainer(context.Context, string) error { + stops := f.stops.Add(1) + if f.stopErrAt > 0 && stops != f.stopErrAt { + return nil + } + return f.stopErr +} + +func waitFor(t *testing.T, check func() bool) { + t.Helper() + require.Eventually(t, check, time.Second, time.Millisecond) +} + +func TestIdleControllerDoesNotStopInflightRequest(t *testing.T) { + lifecycle := &fakeLifecycle{} + c := NewIdleController(10*time.Millisecond, time.Second, []string{"web"}, lifecycle, func(time.Duration) error { return nil }) + defer c.Close() + require.NoError(t, c.BeginRequest(context.Background())) + time.Sleep(30 * time.Millisecond) + assert.Zero(t, lifecycle.stops.Load()) + c.EndRequest() + waitFor(t, func() bool { return lifecycle.stops.Load() == 1 }) +} + +func TestIdleControllerCoalescesConcurrentWake(t *testing.T) { + lifecycle := &fakeLifecycle{} + ready := make(chan struct{}) + c := NewIdleController(time.Millisecond, time.Second, []string{"web"}, lifecycle, func(time.Duration) error { <-ready; return nil }) + defer c.Close() + waitFor(t, func() bool { return c.StateValue() == IdleStateSleeping }) + var wg sync.WaitGroup + for range 10 { + wg.Add(1) + go func() { defer wg.Done(); require.NoError(t, c.BeginRequest(context.Background())); c.EndRequest() }() + } + waitFor(t, func() bool { return lifecycle.starts.Load() == 1 }) + close(ready) + wg.Wait() + assert.Equal(t, int32(1), lifecycle.starts.Load()) +} + +func TestIdleControllerWakeFailureAndTimeout(t *testing.T) { + lifecycle := &fakeLifecycle{startErr: errors.New("start failed")} + c := NewIdleController(time.Millisecond, time.Second, []string{"web"}, lifecycle, func(time.Duration) error { return nil }) + defer c.Close() + waitFor(t, func() bool { return c.StateValue() == IdleStateSleeping }) + assert.ErrorContains(t, c.BeginRequest(context.Background()), "start failed") + c.EndRequest() + + blocking := make(chan struct{}) + c.Reset([]string{"web"}, func(time.Duration) error { <-blocking; return nil }) + c.mu.Lock() + c.State, c.WakeTimeout = IdleStateSleeping, 10*time.Millisecond + c.mu.Unlock() + lifecycle.startErr = nil + assert.ErrorIs(t, c.BeginRequest(context.Background()), ErrIdleWakeTimeout) + c.EndRequest() + close(blocking) +} + +func TestIdleControllerRestoresSleepingAndWakes(t *testing.T) { + original := &IdleController{State: IdleStateWaking, IdleTimeout: time.Minute, WakeTimeout: time.Second, ContainerNames: []string{"web"}} + data, err := json.Marshal(original) + require.NoError(t, err) + var restored IdleController + require.NoError(t, json.Unmarshal(data, &restored)) + lifecycle := &fakeLifecycle{} + restored.configure(time.Minute, time.Second, []string{"web"}, lifecycle, func(time.Duration) error { return nil }) + defer restored.Close() + assert.Equal(t, IdleStateSleeping, restored.StateValue()) + require.NoError(t, restored.BeginRequest(context.Background())) + restored.EndRequest() + assert.Equal(t, int32(1), lifecycle.starts.Load()) +} + +func TestIdleControllerMarshalWaitsForStateLock(t *testing.T) { + c := &IdleController{State: IdleStateSleeping} + c.mu.Lock() + done := make(chan error, 1) + go func() { + _, err := json.Marshal(c) + done <- err + }() + + select { + case <-done: + t.Fatal("MarshalJSON read state without acquiring the lock") + case <-time.After(10 * time.Millisecond): + } + c.mu.Unlock() + require.NoError(t, <-done) +} + +func TestIdleControllerResetCancelsWake(t *testing.T) { + lifecycle := &blockingLifecycle{started: make(chan struct{}), finished: make(chan struct{})} + readyCalled := atomic.Bool{} + c := NewIdleController(time.Hour, time.Second, []string{"old"}, lifecycle, func(time.Duration) error { + readyCalled.Store(true) + return nil + }) + defer c.Close() + c.mu.Lock() + c.State = IdleStateSleeping + c.mu.Unlock() + + done := make(chan error, 1) + go func() { done <- c.BeginRequest(context.Background()) }() + <-lifecycle.started + c.Reset([]string{"new"}, func(time.Duration) error { return nil }) + + require.NoError(t, <-done) + require.Eventually(t, func() bool { + select { + case <-lifecycle.finished: + return true + default: + return false + } + }, 100*time.Millisecond, time.Millisecond) + assert.False(t, readyCalled.Load()) + assert.Equal(t, IdleStateActive, c.StateValue()) +} + +func TestIdleControllerWakeReadinessUsesRemainingTimeout(t *testing.T) { + wakeTimeout := 100 * time.Millisecond + lifecycle := &fakeLifecycle{startDelay: 10 * time.Millisecond} + readyTimeout := make(chan time.Duration, 1) + c := NewIdleController(time.Hour, wakeTimeout, []string{"web"}, lifecycle, func(timeout time.Duration) error { + readyTimeout <- timeout + return nil + }) + defer c.Close() + c.mu.Lock() + c.State = IdleStateSleeping + c.mu.Unlock() + + require.NoError(t, c.BeginRequest(context.Background())) + c.EndRequest() + remaining := <-readyTimeout + assert.Positive(t, remaining) + assert.Less(t, remaining, wakeTimeout) +} + +func TestIdleControllerRecoversFromPartialStopFailure(t *testing.T) { + lifecycle := &fakeLifecycle{stopErr: errors.New("stop failed"), stopErrAt: 2} + c := NewIdleController(time.Hour, time.Second, []string{"web-1", "web-2"}, lifecycle, func(time.Duration) error { return nil }) + defer c.Close() + c.mu.Lock() + c.lastRequest = time.Now().Add(-2 * time.Hour) + c.mu.Unlock() + + c.trySleep() + + assert.Equal(t, int32(2), lifecycle.stops.Load()) + assert.Equal(t, IdleStateSleeping, c.StateValue()) + require.NoError(t, c.BeginRequest(context.Background())) + c.EndRequest() + assert.Equal(t, int32(2), lifecycle.starts.Load()) + assert.Equal(t, IdleStateActive, c.StateValue()) +} diff --git a/internal/server/load_balancer.go b/internal/server/load_balancer.go index 6bec9aec..a4d20267 100644 --- a/internal/server/load_balancer.go +++ b/internal/server/load_balancer.go @@ -153,6 +153,18 @@ func (lb *LoadBalancer) MarkAllHealthy() { lb.updateHealthyTargets() } +func (lb *LoadBalancer) PrepareForWake() { + lb.lock.Lock() + lb.markHealthy() + lb.waitForHealthyContext, lb.markHealthy = context.WithCancel(context.Background()) + lb.writers, lb.readers = TargetList{}, TargetList{} + lb.lock.Unlock() + for _, target := range lb.all { + target.markUnhealthyForWake() + target.RestartHealthChecks() + } +} + func (lb *LoadBalancer) Dispose() { lb.all.StopHealthChecks() } diff --git a/internal/server/load_balancer_test.go b/internal/server/load_balancer_test.go index 33e892c7..e7615be0 100644 --- a/internal/server/load_balancer_test.go +++ b/internal/server/load_balancer_test.go @@ -4,6 +4,7 @@ import ( "net/http" "net/http/httptest" "strconv" + "sync/atomic" "testing" "time" @@ -51,6 +52,48 @@ func TestLoadBalancer_WaitUntilHealthy(t *testing.T) { require.NoError(t, lb.WaitUntilHealthy(time.Second)) } +func TestLoadBalancer_PrepareForWakeRestartsSingleTargetHealthCheck(t *testing.T) { + var healthy atomic.Bool + healthy.Store(true) + _, targetURL := testBackendWithHandler(t, func(w http.ResponseWriter, r *http.Request) { + if !healthy.Load() { + w.WriteHeader(http.StatusServiceUnavailable) + } + }) + options := defaultTargetOptions + options.HealthCheckConfig.Interval = time.Millisecond + target, err := NewTarget(targetURL, options) + require.NoError(t, err) + lb := NewLoadBalancer(TargetList{target}, DefaultWriterAffinityTimeout, false) + t.Cleanup(lb.Dispose) + require.NoError(t, lb.WaitUntilHealthy(time.Second)) + require.Eventually(t, func() bool { + target.inflightLock.Lock() + defer target.inflightLock.Unlock() + return target.healthcheck == nil + }, time.Second, time.Millisecond) + + healthy.Store(false) + lb.PrepareForWake() + time.AfterFunc(10*time.Millisecond, func() { healthy.Store(true) }) + require.NoError(t, lb.WaitUntilHealthy(time.Second)) + assert.Equal(t, TargetStateHealthy, target.State()) +} + +func TestLoadBalancer_PrepareForWakeReleasesPreviousWaiters(t *testing.T) { + lb := NewLoadBalancer(TargetList{}, DefaultWriterAffinityTimeout, false) + t.Cleanup(lb.Dispose) + previous := lb.waitForHealthyContext + + lb.PrepareForWake() + + select { + case <-previous.Done(): + case <-time.After(time.Second): + t.Fatal("previous health wait was not released") + } +} + func TestLoadBalancer_StartRequest(t *testing.T) { lb := testLoadBalancerWithHandlers(t, func(w http.ResponseWriter, r *http.Request) { diff --git a/internal/server/router.go b/internal/server/router.go index 3b891b6a..5158fb35 100644 --- a/internal/server/router.go +++ b/internal/server/router.go @@ -5,6 +5,7 @@ import ( "crypto/tls" "encoding/json" "errors" + "fmt" "log/slog" "net/http" "os" @@ -47,9 +48,11 @@ func RoutedTargetPath(r *http.Request) string { } type Router struct { - statePath string - services *ServiceMap - serviceLock sync.RWMutex + statePath string + dockerClient *DockerClient + services *ServiceMap + serviceLock sync.RWMutex + stateLock sync.Mutex } type ServiceDescription struct { @@ -62,10 +65,11 @@ type ServiceDescription struct { type ServiceDescriptionMap map[string]ServiceDescription -func NewRouter(statePath string) *Router { +func NewRouter(statePath string, dockerSocketPath string) *Router { return &Router{ - statePath: statePath, - services: NewServiceMap(), + statePath: statePath, + dockerClient: NewDockerClient(dockerSocketPath), + services: NewServiceMap(), } } @@ -88,17 +92,19 @@ func (r *Router) RestoreLastSavedState() error { return err } - r.withWriteLock(func() error { + return r.withWriteLock(func() error { r.services = NewServiceMap() for _, service := range services { + service.lifecycle = r.dockerClient + service.stateChanged = func() { _ = r.saveStateSnapshot() } + if err := service.initialize(service.options, service.targetOptions); err != nil { + return fmt.Errorf("initialize restored service %q: %w", service.name, err) + } r.services.Set(service) } - + slog.Info("Restored saved state", "path", r.statePath) return nil }) - - slog.Info("Restored saved state", "path", r.statePath) - return nil } func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) { @@ -259,12 +265,20 @@ func (r *Router) ListActiveServices() ServiceDescriptionMap { path := strings.Join(service.options.PathPrefixes, ",") target := strings.Join(service.active.Targets().Names(), ",") + state := service.pauseController.GetState().String() + if service.idleController != nil && service.pauseController.GetState() == PauseStateRunning { + icState := service.idleController.StateValue() + if icState != IdleStateActive { + state = icState.String() + } + } + result[name] = ServiceDescription{ Host: host, Path: path, Target: target, TLS: service.options.TLSEnabled, - State: service.pauseController.GetState().String(), + State: state, } } } @@ -305,9 +319,17 @@ func (r *Router) GetCertificate(hello *tls.ClientHelloInfo) (*tls.Certificate, e func (r *Router) createOrUpdateService(name string, options ServiceOptions, targetOptions TargetOptions) (*Service, error) { service := r.services.Get(name) if service == nil { - return NewService(name, options, targetOptions) + service, err := NewService(name, options, targetOptions, r.dockerClient) + if err == nil { + service.stateChanged = func() { _ = r.saveStateSnapshot() } + if service.idleController != nil { + service.idleController.SetPersist(service.stateChanged) + } + } + return service, err } + service.lifecycle = r.dockerClient err := service.UpdateOptions(options, targetOptions) return service, err } @@ -357,6 +379,8 @@ func (r *Router) installLoadBalancer(name string, slot TargetSlot, lb *LoadBalan } func (r *Router) saveStateSnapshot() error { + r.stateLock.Lock() + defer r.stateLock.Unlock() services := []*Service{} r.withReadLock(func() error { for _, service := range r.services.All() { @@ -367,13 +391,19 @@ func (r *Router) saveStateSnapshot() error { f, err := os.Create(r.statePath) if err != nil { + slog.Error("Unable to save state snapshot", "error", err, "path", r.statePath) return err } - err = json.NewEncoder(f).Encode(services) - if err != nil { - slog.Error("Unable to save state", "error", err, "path", r.statePath) - return err + encodeErr := json.NewEncoder(f).Encode(services) + closeErr := f.Close() + if encodeErr != nil { + slog.Error("Unable to save state", "error", encodeErr, "path", r.statePath) + return encodeErr + } + if closeErr != nil { + slog.Error("Unable to close saved state", "error", closeErr, "path", r.statePath) + return closeErr } slog.Debug("Saved state", "path", r.statePath) diff --git a/internal/server/router_test.go b/internal/server/router_test.go index 28184868..cae77339 100644 --- a/internal/server/router_test.go +++ b/internal/server/router_test.go @@ -1,9 +1,11 @@ package server import ( + "bytes" "context" "crypto/tls" "encoding/json" + "log/slog" "net/http" "net/http/httptest" "os" @@ -26,6 +28,18 @@ func TestRouter_Empty(t *testing.T) { assert.Equal(t, http.StatusNotFound, statusCode) } +func TestRouter_StateChangedLogsSaveErrors(t *testing.T) { + var logs bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + router := NewRouter(t.TempDir(), DefaultDockerSocketPath) + + require.Error(t, router.saveStateSnapshot()) + + assert.Contains(t, logs.String(), "Unable to save state snapshot") +} + func TestRouter_DeployService(t *testing.T) { router := testRouter(t) _, target := testBackend(t, "first", http.StatusOK) @@ -767,7 +781,7 @@ func TestRouter_RestoreLastSavedState(t *testing.T) { _, second := testBackend(t, "second", http.StatusOK) _, third := testBackend(t, "third", http.StatusOK) - router := NewRouter(statePath) + router := NewRouter(statePath, DefaultDockerSocketPath) require.NoError(t, router.DeployService("default", []string{first}, defaultEmptyReaders, defaultServiceOptions, defaultTargetOptions, defaultDeploymentOptions)) serviceOptions := defaultServiceOptions @@ -791,7 +805,7 @@ func TestRouter_RestoreLastSavedState(t *testing.T) { assert.Equal(t, http.StatusOK, statusCode) assert.Equal(t, "third", body) - router = NewRouter(statePath) + router = NewRouter(statePath, DefaultDockerSocketPath) router.RestoreLastSavedState() statusCode, body = sendGETRequest(router, "http://something.example.com/") @@ -824,10 +838,10 @@ func TestRouter_RestoreLastSavedState_TLSOnDemandURL(t *testing.T) { serviceOptions.TLSEnabled = true serviceOptions.TLSOnDemandURL = allowServer.URL - router := NewRouter(statePath) + router := NewRouter(statePath, DefaultDockerSocketPath) require.NoError(t, router.DeployService("ondemand", []string{target}, defaultEmptyReaders, serviceOptions, defaultTargetOptions, defaultDeploymentOptions)) - router = NewRouter(statePath) + router = NewRouter(statePath, DefaultDockerSocketPath) require.NoError(t, router.RestoreLastSavedState()) service := router.services.Get("ondemand") @@ -845,7 +859,7 @@ func TestRouter_RestoreLastSavedState_TLSOnDemandURL(t *testing.T) { func testRouter(t *testing.T) *Router { statePath := filepath.Join(t.TempDir(), "state.json") - return NewRouter(statePath) + return NewRouter(statePath, DefaultDockerSocketPath) } func sendGETRequest(router *Router, url string) (int, string) { diff --git a/internal/server/service.go b/internal/server/service.go index d3c88d6d..b5275bb5 100644 --- a/internal/server/service.go +++ b/internal/server/service.go @@ -49,6 +49,9 @@ const ( DefaultMaxRequestBodySize = 0 DefaultMaxResponseBodySize = 0 + DefaultIdleWakeTimeout = 30 * time.Second + DefaultIdleLifecycleTimeout = 30 * time.Second + DefaultStopMessage = "" ) @@ -111,6 +114,8 @@ type ServiceOptions struct { ReadTargetsAcceptWebsockets bool `json:"read_targets_accept_websockets"` ExcludeMetricsPaths []string `json:"exclude_metrics_paths"` ClientIPHeader string `json:"client_ip_header"` + IdleTimeout time.Duration `json:"idle_timeout"` + IdleWakeTimeout time.Duration `json:"idle_wake_timeout"` } func (so *ServiceOptions) ShouldExcludeMetrics(r *http.Request) bool { @@ -120,6 +125,9 @@ func (so *ServiceOptions) ShouldExcludeMetrics(r *http.Request) bool { func (so *ServiceOptions) Normalize() { so.Hosts = NormalizeHosts(so.Hosts) so.PathPrefixes = NormalizePathPrefixes(so.PathPrefixes) + if so.IdleTimeout > 0 && so.IdleWakeTimeout == 0 { + so.IdleWakeTimeout = DefaultIdleWakeTimeout + } } func (so ServiceOptions) Validate() error { @@ -160,6 +168,9 @@ func (so ServiceOptions) Validate() error { return fmt.Errorf("%w: canonical-host '%s' must be present in the hosts list: %v", ErrServiceOptionsInvalid, so.CanonicalHost, so.Hosts) } } + if so.IdleTimeout < 0 || so.IdleWakeTimeout < 0 { + return fmt.Errorf("%w: idle timeouts cannot be negative", ErrServiceOptionsInvalid) + } return nil } @@ -203,15 +214,20 @@ type Service struct { serviceLock sync.RWMutex pauseController *PauseController + idleController *IdleController rolloutController *RolloutController - certManager CertManager - middleware http.Handler + lifecycle ContainerLifecycle + + certManager CertManager + middleware http.Handler + stateChanged func() } -func NewService(name string, options ServiceOptions, targetOptions TargetOptions) (*Service, error) { +func NewService(name string, options ServiceOptions, targetOptions TargetOptions, lifecycle ContainerLifecycle) (*Service, error) { service := &Service{ name: name, + lifecycle: lifecycle, pauseController: NewPauseController(), } @@ -226,6 +242,10 @@ func (s *Service) UpdateOptions(options ServiceOptions, targetOptions TargetOpti } func (s *Service) Dispose() { + if s.idleController != nil { + s.idleController.Close() + } + s.active.Dispose() if s.rollout != nil { s.rollout.Dispose() @@ -241,9 +261,15 @@ func (s *Service) UpdateLoadBalancer(lb *LoadBalancer, slot TargetSlot) *LoadBal if slot == TargetSlotRollout { replaced = s.rollout s.rollout = lb + if s.idleController != nil { + s.idleController.configure(s.options.IdleTimeout, s.options.IdleWakeTimeout, s.idleContainerNames(), s.lifecycle, s.waitUntilIdleTargetsHealthy) + } } else { replaced = s.active s.active = lb + if s.idleController != nil { + s.idleController.Reset(s.idleContainerNames(), s.waitUntilIdleTargetsHealthy) + } } return replaced @@ -279,6 +305,16 @@ func (s *Service) ServeHTTP(w http.ResponseWriter, r *http.Request) { defer metrics.Tracker.SubtractInflightRequest(s.name) } + if s.idleController != nil && s.pauseController.GetState() == PauseStateRunning && !s.targetOptions.IsHealthCheckRequest(r) { + if err := s.idleController.BeginRequest(r.Context()); err != nil { + s.idleController.EndRequest() + templateArguments := struct{ Message string }{err.Error()} + SetErrorResponse(w, r, http.StatusServiceUnavailable, templateArguments) + return + } + defer s.idleController.EndRequest() + } + s.middleware.ServeHTTP(w, r) } @@ -291,6 +327,7 @@ type marshalledService struct { RolloutTargets []string `json:"rollout_targets"` RolloutReaders []string `json:"rollout_readers"` PauseController *PauseController `json:"pause_controller"` + IdleController *IdleController `json:"idle_controller"` RolloutController *RolloutController `json:"rollout_controller"` LegacyActiveTarget string `json:"active_target,omitempty"` @@ -316,6 +353,7 @@ func (s *Service) MarshalJSON() ([]byte, error) { Options: s.options, TargetOptions: s.targetOptions, PauseController: s.pauseController, + IdleController: s.idleController, RolloutController: s.rolloutController, }) } @@ -344,6 +382,7 @@ func (s *Service) UnmarshalJSON(data []byte) error { s.name = ms.Name s.pauseController = ms.PauseController + s.idleController = ms.IdleController s.rolloutController = ms.RolloutController activeTargets, err := NewTargetList(ms.ActiveTargets, ms.ActiveReaders, ms.TargetOptions) @@ -366,6 +405,10 @@ func (s *Service) UnmarshalJSON(data []byte) error { } func (s *Service) Stop(drainTimeout time.Duration, message string) error { + if s.idleController != nil { + s.idleController.Disable() + } + err := s.pauseController.Stop(message) if err != nil { return err @@ -379,6 +422,10 @@ func (s *Service) Stop(drainTimeout time.Duration, message string) error { } func (s *Service) Pause(drainTimeout time.Duration, pauseTimeout time.Duration) error { + if s.idleController != nil { + s.idleController.Disable() + } + err := s.pauseController.Pause(pauseTimeout) if err != nil { return err @@ -392,6 +439,10 @@ func (s *Service) Pause(drainTimeout time.Duration, pauseTimeout time.Duration) } func (s *Service) Resume() error { + if s.idleController != nil { + s.idleController.Enable() + } + err := s.pauseController.Resume() if err != nil { return err @@ -419,6 +470,65 @@ func (s *Service) initialize(options ServiceOptions, targetOptions TargetOptions s.certManager = certManager s.middleware = middleware + if s.options.IdleTimeout > 0 { + if s.idleController == nil { + s.idleController = NewIdleController(s.options.IdleTimeout, s.options.IdleWakeTimeout, s.idleContainerNames(), s.lifecycle, s.waitUntilIdleTargetsHealthy) + } else { + s.idleController.configure(s.options.IdleTimeout, s.options.IdleWakeTimeout, s.idleContainerNames(), s.lifecycle, s.waitUntilIdleTargetsHealthy) + } + s.idleController.SetPersist(s.stateChanged) + } else if s.idleController != nil { + s.idleController.Close() + s.idleController = nil + } + + return nil +} + +func (s *Service) activeContainerNames() []string { + if s.active == nil { + return nil + } + names := make([]string, 0, len(s.active.WriteTargets())) + for _, target := range s.active.WriteTargets() { + names = append(names, target.ContainerName()) + } + return names +} + +func (s *Service) idleContainerNames() []string { + names := s.activeContainerNames() + if s.rollout != nil { + for _, target := range s.rollout.WriteTargets() { + names = append(names, target.ContainerName()) + } + } + return names +} + +func (s *Service) waitUntilIdleTargetsHealthy(timeout time.Duration) error { + s.serviceLock.RLock() + loadBalancers := make([]*LoadBalancer, 0, 2) + if s.active != nil { + loadBalancers = append(loadBalancers, s.active) + } + if s.rollout != nil { + loadBalancers = append(loadBalancers, s.rollout) + } + s.serviceLock.RUnlock() + if len(loadBalancers) == 0 { + return ErrorNoHealthyTargets + } + errs := make(chan error, len(loadBalancers)) + for _, lb := range loadBalancers { + lb.PrepareForWake() + go func() { errs <- lb.WaitUntilHealthy(timeout) }() + } + for range loadBalancers { + if err := <-errs; err != nil { + return err + } + } return nil } @@ -535,12 +645,35 @@ func (s *Service) serviceRequestWithTarget(w http.ResponseWriter, r *http.Reques return } + if s.handleIdleHealthCheck(w, r) { + return + } + sendRequest := s.startLoadBalancerRequest(w, r) if sendRequest != nil { sendRequest() } } +func (s *Service) handleIdleHealthCheck(w http.ResponseWriter, r *http.Request) bool { + if s.idleController == nil { + return false + } + + if s.targetOptions.IsHealthCheckRequest(r) { + // Health checks should not wake the service. + // While it is stopping, sleeping, or waking, just return 200. + icState := s.idleController.StateValue() + if icState != IdleStateActive { + w.WriteHeader(http.StatusOK) + return true + } + return false + } + + return false +} + func (s *Service) startLoadBalancerRequest(w http.ResponseWriter, r *http.Request) func() { s.serviceLock.RLock() defer s.serviceLock.RUnlock() diff --git a/internal/server/service_test.go b/internal/server/service_test.go index 4a74219f..bb1d5016 100644 --- a/internal/server/service_test.go +++ b/internal/server/service_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "net" "net/http" "net/http/httptest" @@ -14,6 +15,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/basecamp/kamal-proxy/internal/pages" ) func TestService_ServeRequest(t *testing.T) { @@ -26,6 +29,90 @@ func TestService_ServeRequest(t *testing.T) { require.Equal(t, http.StatusOK, w.Result().StatusCode) } +func TestService_StoppedRequestDoesNotWakeIdleContainer(t *testing.T) { + options := defaultServiceOptions + options.IdleTimeout = time.Minute + options.IdleWakeTimeout = time.Second + service := testCreateService(t, options, defaultTargetOptions) + t.Cleanup(service.Dispose) + lifecycle := &fakeLifecycle{} + service.idleController.mu.Lock() + service.idleController.lifecycle = lifecycle + service.idleController.State = IdleStateSleeping + service.idleController.mu.Unlock() + require.NoError(t, service.pauseController.Stop(DefaultStopMessage)) + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + w := httptest.NewRecorder() + service.ServeHTTP(w, req) + + assert.Equal(t, http.StatusServiceUnavailable, w.Code) + assert.Zero(t, lifecycle.starts.Load()) +} + +func TestService_WakeFailureRendersErrorMessage(t *testing.T) { + options := defaultServiceOptions + options.IdleTimeout = time.Minute + options.IdleWakeTimeout = time.Second + service := testCreateService(t, options, defaultTargetOptions) + t.Cleanup(service.Dispose) + lifecycle := &fakeLifecycle{startErr: errors.New("start failed")} + service.idleController.mu.Lock() + service.idleController.lifecycle = lifecycle + service.idleController.State = IdleStateSleeping + service.idleController.mu.Unlock() + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + w := httptest.NewRecorder() + handler, err := WithErrorPageMiddleware(pages.DefaultErrorPages, true, service) + require.NoError(t, err) + handler.ServeHTTP(w, req) + + assert.Equal(t, http.StatusServiceUnavailable, w.Code) + assert.Contains(t, w.Body.String(), "start failed") +} + +func TestService_WakeStartsActiveAndRolloutTargets(t *testing.T) { + options := defaultServiceOptions + options.IdleTimeout = time.Minute + options.IdleWakeTimeout = time.Second + service := testCreateServiceWithHandler(t, options, defaultTargetOptions, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("active")) + })) + t.Cleanup(service.Dispose) + _, rolloutTarget := testBackend(t, "rollout", http.StatusOK) + target, err := NewTarget(rolloutTarget, defaultTargetOptions) + require.NoError(t, err) + rollout := NewLoadBalancer(TargetList{target}, DefaultWriterAffinityTimeout, false) + require.NoError(t, rollout.WaitUntilHealthy(time.Second)) + service.UpdateLoadBalancer(rollout, TargetSlotRollout) + require.NoError(t, service.SetRolloutSplit(100, nil)) + lifecycle := &fakeLifecycle{} + service.idleController.mu.Lock() + service.idleController.lifecycle = lifecycle + service.idleController.State = IdleStateSleeping + service.idleController.mu.Unlock() + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.AddCookie(&http.Cookie{Name: RolloutCookieName, Value: "1"}) + w := httptest.NewRecorder() + service.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "rollout", w.Body.String()) + assert.Equal(t, int32(2), lifecycle.starts.Load()) +} + +func TestService_WaitUntilIdleTargetsHealthyWithoutLoadBalancer(t *testing.T) { + options := defaultServiceOptions + options.IdleTimeout = time.Minute + service, err := NewService("test", options, defaultTargetOptions, &fakeLifecycle{}) + require.NoError(t, err) + t.Cleanup(service.idleController.Close) + + assert.ErrorIs(t, service.waitUntilIdleTargetsHealthy(time.Second), ErrorNoHealthyTargets) +} + func TestService_ClientIPHeaderRewritesXForwardedFor(t *testing.T) { var xForwardedFor, trueClientIP string @@ -347,6 +434,17 @@ func TestService_UnmarshallingStateFromLegacyFormat(t *testing.T) { assert.Equal(t, 3*time.Second, service.targetOptions.ResponseTimeout) } +func TestNewServiceWithIdleLifecycleBeforeActiveLoadBalancer(t *testing.T) { + service, err := NewService("test", ServiceOptions{ + IdleTimeout: time.Minute, + IdleWakeTimeout: time.Second, + }, defaultTargetOptions, &fakeLifecycle{}) + require.NoError(t, err) + t.Cleanup(service.idleController.Close) + + assert.Empty(t, service.idleController.ContainerNames) +} + func testCreateService(t *testing.T, options ServiceOptions, targetOptions TargetOptions) *Service { return testCreateServiceWithHandler( t, options, targetOptions, @@ -364,7 +462,7 @@ func testCreateServiceWithHandler(t *testing.T, options ServiceOptions, targetOp target, err := NewTarget(serverURL.Host, targetOptions) require.NoError(t, err) - service, err := NewService("test", options, targetOptions) + service, err := NewService("test", options, targetOptions, &fakeLifecycle{}) require.NoError(t, err) service.UpdateLoadBalancer(NewLoadBalancer(TargetList{target}, DefaultWriterAffinityTimeout, false), TargetSlotActive) diff --git a/internal/server/target.go b/internal/server/target.go index b1bd627b..7bcf60bd 100644 --- a/internal/server/target.go +++ b/internal/server/target.go @@ -143,6 +143,14 @@ func (t *Target) Address() string { return t.targetURL.Host } +func (t *Target) ContainerName() string { return t.targetURL.Hostname() } + +func (t *Target) markUnhealthyForWake() { + t.inflightLock.Lock() + t.state = TargetStateUnhealthy + t.inflightLock.Unlock() +} + func (t *Target) State() TargetState { t.inflightLock.Lock() defer t.inflightLock.Unlock() @@ -217,16 +225,23 @@ WAIT_FOR_REQUESTS_TO_COMPLETE: func (t *Target) BeginHealthChecks(stateConsumer TargetStateConsumer) { t.stateConsumer = stateConsumer + t.RestartHealthChecks() +} +func (t *Target) RestartHealthChecks() { t.withInflightLock(func() { + if t.healthcheck != nil { + t.healthcheck.Close() + } healthCheckURL := t.buildHealthCheckURL() - t.healthcheck = NewHealthCheck( + t.healthcheck = newHealthCheck( t, healthCheckURL, t.options.HealthCheckConfig.Interval, t.options.HealthCheckConfig.Timeout, t.options.HealthCheckConfig.Host, ) + t.healthcheck.Start() }) } diff --git a/internal/server/testing.go b/internal/server/testing.go index 61576bf1..d58b7194 100644 --- a/internal/server/testing.go +++ b/internal/server/testing.go @@ -85,7 +85,7 @@ func testServer(t testing.TB, http3Enabled bool) *Server { AlternateConfigDir: t.TempDir(), HTTP3Enabled: http3Enabled, } - router := NewRouter(config.StatePath()) + router := NewRouter(config.StatePath(), DefaultDockerSocketPath) server := NewServer(config, router) err := server.Start() require.NoError(t, err)