diff --git a/server/cmd/api/main.go b/server/cmd/api/main.go index b3384c6f5..30670e8a7 100644 --- a/server/cmd/api/main.go +++ b/server/cmd/api/main.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "log/slog" + "net" "net/http" "net/url" "os" @@ -377,6 +378,7 @@ func main() { metrics.NewChromeCollector(upstreamMgr), metrics.NewGPUCollector(), metrics.NewSystemCollector(), + metrics.NewResponseDrainCollector(scaletozero.ResponseDrainOutcomeCounts, scaletozero.ActiveResponseHolds, scaletozero.FailClosedResponseHolds), } if otlpMetrics != nil { metricsCollectors = append(metricsCollectors, metrics.NewOTLPCollector(otlpMetrics)) @@ -387,29 +389,22 @@ func main() { Handler: rMetrics, } - go func() { - slogger.Info("http server starting", "addr", srv.Addr) - if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed { - slogger.Error("http server failed", "err", err) - stop() - } - }() - - go func() { - slogger.Info("devtools websocket proxy starting", "addr", srvDevtools.Addr) - if err := srvDevtools.ListenAndServe(); err != nil && err != http.ErrServerClosed { - slogger.Error("devtools websocket proxy failed", "err", err) - stop() - } - }() - - go func() { - slogger.Info("chromedriver proxy starting", "addr", srvChromeDriver.Addr) - if err := srvChromeDriver.ListenAndServe(); err != nil && err != http.ErrServerClosed { - slogger.Error("chromedriver proxy failed", "err", err) - stop() - } - }() + serveHTTP := func(name string, server *http.Server) { + go func() { + slogger.Info(name+" starting", "addr", server.Addr) + listener, err := net.Listen("tcp", server.Addr) + if err == nil { + err = scaletozero.Serve(server, listener) + } + if err != nil && err != http.ErrServerClosed { + slogger.Error(name+" failed", "err", err) + stop() + } + }() + } + serveHTTP("http server", srv) + serveHTTP("devtools websocket proxy", srvDevtools) + serveHTTP("chromedriver proxy", srvChromeDriver) go func() { slogger.Info("metrics server starting", "addr", srvMetrics.Addr) diff --git a/server/lib/metrics/scaletozero.go b/server/lib/metrics/scaletozero.go new file mode 100644 index 000000000..5a636f1b0 --- /dev/null +++ b/server/lib/metrics/scaletozero.go @@ -0,0 +1,40 @@ +package metrics + +import ( + "context" + "sort" +) + +type ResponseDrainSource func() map[string]uint64 +type ResponseDrainGauge func() int64 + +type ResponseDrainCollector struct { + snapshot ResponseDrainSource + active ResponseDrainGauge + failClosed ResponseDrainGauge +} + +func NewResponseDrainCollector(snapshot ResponseDrainSource, active, failClosed ResponseDrainGauge) *ResponseDrainCollector { + return &ResponseDrainCollector{snapshot: snapshot, active: active, failClosed: failClosed} +} + +func (c *ResponseDrainCollector) Name() string { return "scale-to-zero response drain" } + +func (c *ResponseDrainCollector) Collect(_ context.Context, w *Writer) error { + w.Metric("kernel_scale_to_zero_response_drain_total", "HTTP response drain events before scale-to-zero is re-enabled.", "counter") + counts := c.snapshot() + outcomes := make([]string, 0, len(counts)) + for outcome := range counts { + outcomes = append(outcomes, outcome) + } + sort.Strings(outcomes) + for _, outcome := range outcomes { + w.Sample("kernel_scale_to_zero_response_drain_total", []Label{{Name: "outcome", Value: outcome}}, float64(counts[outcome])) + } + + w.Metric("kernel_scale_to_zero_response_holds", "HTTP response scale-to-zero holds currently active.", "gauge") + w.Sample("kernel_scale_to_zero_response_holds", nil, float64(c.active())) + w.Metric("kernel_scale_to_zero_response_fail_closed_holds", "HTTP response holds awaiting terminal connection recovery or guest termination.", "gauge") + w.Sample("kernel_scale_to_zero_response_fail_closed_holds", nil, float64(c.failClosed())) + return nil +} diff --git a/server/lib/metrics/scaletozero_test.go b/server/lib/metrics/scaletozero_test.go new file mode 100644 index 000000000..d62d14865 --- /dev/null +++ b/server/lib/metrics/scaletozero_test.go @@ -0,0 +1,33 @@ +package metrics + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestResponseDrainCollector(t *testing.T) { + collector := NewResponseDrainCollector(func() map[string]uint64 { + return map[string]uint64{ + "timeout": 2, + "drained": 7, + } + }, func() int64 { return 3 }, func() int64 { return 1 }) + writer := &Writer{} + + require.NoError(t, collector.Collect(context.Background(), writer)) + + assert.Equal(t, `# HELP kernel_scale_to_zero_response_drain_total HTTP response drain events before scale-to-zero is re-enabled. +# TYPE kernel_scale_to_zero_response_drain_total counter +kernel_scale_to_zero_response_drain_total{outcome="drained"} 7 +kernel_scale_to_zero_response_drain_total{outcome="timeout"} 2 +# HELP kernel_scale_to_zero_response_holds HTTP response scale-to-zero holds currently active. +# TYPE kernel_scale_to_zero_response_holds gauge +kernel_scale_to_zero_response_holds 3 +# HELP kernel_scale_to_zero_response_fail_closed_holds HTTP response holds awaiting terminal connection recovery or guest termination. +# TYPE kernel_scale_to_zero_response_fail_closed_holds gauge +kernel_scale_to_zero_response_fail_closed_holds 1 +`, string(writer.Bytes())) +} diff --git a/server/lib/scaletozero/connection.go b/server/lib/scaletozero/connection.go new file mode 100644 index 000000000..73f7fc4a5 --- /dev/null +++ b/server/lib/scaletozero/connection.go @@ -0,0 +1,534 @@ +package scaletozero + +import ( + "context" + "errors" + "io" + "log/slog" + "net" + "net/http" + "sync" + "time" + + "golang.org/x/sys/unix" +) + +const responseReadFromChunkSize = 1 << 20 + +type drainListener struct { + net.Listener +} + +type drainConn struct { + *net.TCPConn + mu sync.Mutex + terminalMu sync.Mutex + pending []*requestDrain + state http.ConnState + generation uint64 + drainCancel context.CancelFunc + closed bool + closing bool + hijackedConn bool + writeTimeout time.Duration + setDeadline func(*net.TCPConn, time.Time) error + abortConnection func(*net.TCPConn) error + responseWriteErr error +} + +// Serve adds TCP response tracking to server before serving listener. +func Serve(server *http.Server, listener net.Listener) error { + server.ConnContext = connectionContext + server.ConnState = connectionState + return server.Serve(&drainListener{Listener: listener}) +} + +func (l *drainListener) Accept() (net.Conn, error) { + conn, err := l.Listener.Accept() + if err != nil { + return nil, err + } + tcp, ok := conn.(*net.TCPConn) + if !ok { + _ = conn.Close() + return nil, errors.New("scale-to-zero response drain requires a TCP listener") + } + return &drainConn{TCPConn: tcp, state: http.StateNew}, nil +} + +func (c *drainConn) configure(config responseDrainConfig) { + c.mu.Lock() + defer c.mu.Unlock() + c.writeTimeout = config.timeout + c.setDeadline = config.setDeadline + c.abortConnection = config.abort + c.responseWriteErr = nil +} + +func (c *drainConn) Write(p []byte) (int, error) { + c.mu.Lock() + if c.closed || c.closing { + c.mu.Unlock() + return 0, net.ErrClosed + } + timeout := c.writeTimeout + setDeadline := c.setDeadline + abort := c.abortConnection + c.mu.Unlock() + + if timeout > 0 { + if err := setDeadline(c.TCPConn, time.Now().Add(timeout)); err != nil { + c.recordWriteError(err) + c.abortNow(abort) + return 0, err + } + } + + n, err := c.TCPConn.Write(p) + if err != nil { + c.recordWriteError(err) + c.abortNow(abort) + } + return n, err +} + +func (c *drainConn) ReadFrom(r io.Reader) (int64, error) { + var total int64 + for { + c.mu.Lock() + if c.closed || c.closing { + c.mu.Unlock() + return total, net.ErrClosed + } + timeout := c.writeTimeout + setDeadline := c.setDeadline + abort := c.abortConnection + c.mu.Unlock() + + if timeout > 0 { + if err := setDeadline(c.TCPConn, time.Now().Add(timeout)); err != nil { + c.recordWriteError(err) + c.abortNow(abort) + return total, err + } + } + + limited := &io.LimitedReader{R: r, N: responseReadFromChunkSize} + n, err := c.TCPConn.ReadFrom(limited) + total += n + if errors.Is(err, io.EOF) { + return total, nil + } + if err != nil { + c.recordWriteError(err) + c.abortNow(abort) + return total, err + } + if limited.N > 0 { + return total, nil + } + } +} + +func (c *drainConn) recordWriteError(err error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.responseWriteErr == nil { + c.responseWriteErr = err + } +} + +func (c *drainConn) writeError() error { + c.mu.Lock() + defer c.mu.Unlock() + return c.responseWriteErr +} + +func (c *drainConn) addDrain(drain *requestDrain) bool { + c.mu.Lock() + switch { + case c.hijackedConn: + c.mu.Unlock() + drain.complete(responseDrainConnectionHijacked, 0, nil) + return false + case c.closed || c.closing: + c.mu.Unlock() + drain.complete(responseDrainConnectionClosed, 0, nil) + return false + default: + previous := c.pending + c.pending = []*requestDrain{drain} + c.mu.Unlock() + completeDrains(previous, responseDrainConnectionReused, 0, nil) + return true + } +} + +func (c *drainConn) setState(state http.ConnState) { + switch state { + case http.StateActive: + c.activate() + case http.StateIdle: + c.startIdleDrain() + case http.StateHijacked: + c.hijack() + } +} + +func (c *drainConn) activate() { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed || c.hijackedConn { + return + } + c.state = http.StateActive + c.generation++ + if c.drainCancel != nil { + c.drainCancel() + c.drainCancel = nil + } +} + +func (c *drainConn) startIdleDrain() { + c.mu.Lock() + if c.closed || c.closing || c.hijackedConn || len(c.pending) == 0 { + c.mu.Unlock() + return + } + c.state = http.StateIdle + c.generation++ + generation := c.generation + if c.drainCancel != nil { + c.drainCancel() + } + ctx, cancel := context.WithCancel(context.Background()) + c.drainCancel = cancel + config := c.pending[len(c.pending)-1].config + c.mu.Unlock() + + go c.runIdleDrain(ctx, generation, config) +} + +func (c *drainConn) runIdleDrain(ctx context.Context, generation uint64, config responseDrainConfig) { + outcome, queued, err := waitForResponseDrain(ctx, c.TCPConn, config, time.Now().Add(config.timeout)) + if errors.Is(err, context.Canceled) { + return + } + + c.mu.Lock() + if c.closed || c.hijackedConn || c.generation != generation || c.state != http.StateIdle { + c.mu.Unlock() + return + } + if outcome == responseDrainComplete { + if clearErr := config.setDeadline(c.TCPConn, time.Time{}); clearErr == nil { + drains := c.takePendingLocked() + c.drainCancel = nil + c.mu.Unlock() + completeDrains(drains, outcome, queued, nil) + return + } else { + outcome = responseDrainDeadlineClearError + err = clearErr + } + } + + c.closed = true + c.closing = true + c.drainCancel = nil + drains := c.takePendingLocked() + c.mu.Unlock() + go c.recoverResponseConnection(drains, config, outcome, queued, err) +} + +func (c *drainConn) abortNow(abort func(*net.TCPConn) error) { + config := defaultResponseDrainConfig() + if abort != nil { + config.abort = abort + } + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return + } + if len(c.pending) > 0 { + config = c.pending[len(c.pending)-1].config + if abort != nil { + config.abort = abort + } + } + c.closed = true + c.closing = true + drains := c.takePendingLocked() + c.mu.Unlock() + + c.terminalMu.Lock() + err := config.abort(c.TCPConn) + c.terminalMu.Unlock() + if err == nil || isTerminalConnectionError(err) { + completeDrains(drains, responseDrainConnectionClosed, 0, nil) + return + } + go c.recoverResponseConnection(drains, config, responseDrainWriteError, 0, err) +} + +func (c *drainConn) hijack() { + c.mu.Lock() + if c.closed || c.hijackedConn { + c.mu.Unlock() + return + } + c.hijackedConn = true + c.state = http.StateHijacked + c.generation++ + if c.drainCancel != nil { + c.drainCancel() + c.drainCancel = nil + } + c.writeTimeout = 0 + setDeadline := c.setDeadline + drains := c.takePendingLocked() + c.mu.Unlock() + + if setDeadline != nil { + if err := setDeadline(c.TCPConn, time.Time{}); err != nil { + recordResponseDrainOutcome(responseDrainDeadlineClearError) + if log := firstDrainLog(drains); log != nil { + log.Warn("failed to clear hijacked connection deadline", "outcome", responseDrainDeadlineClearError, "error", err) + } + } + } + completeDrains(drains, responseDrainConnectionHijacked, 0, nil) +} + +func (c *drainConn) Close() error { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return nil + } + c.closed = true + c.generation++ + if c.drainCancel != nil { + c.drainCancel() + c.drainCancel = nil + } + drains := c.takePendingLocked() + hijacked := c.hijackedConn + config := responseDrainConfig{} + if len(drains) > 0 { + config = drains[len(drains)-1].config + } + c.mu.Unlock() + + c.terminalMu.Lock() + defer c.terminalMu.Unlock() + if hijacked || len(drains) == 0 { + return c.TCPConn.Close() + } + + fd, err := config.duplicate(c.TCPConn) + if err == nil { + closeErr := c.TCPConn.Close() + go monitorClosedResponse(fd, drains, config, false, nil) + return closeErr + } + if isTerminalConnectionError(err) { + _ = c.TCPConn.Close() + completeDrains(drains, responseDrainConnectionClosed, 0, nil) + return nil + } + + recordResponseDrainOutcome(responseDrainAbortError) + abortErr := config.abort(c.TCPConn) + if abortErr == nil || isTerminalConnectionError(abortErr) { + completeDrains(drains, responseDrainConnectionClosed, 0, err) + return nil + } + go c.recoverResponseConnection(drains, config, responseDrainConnectionClosed, 0, errors.Join(err, abortErr)) + return errors.Join(err, abortErr) +} + +func (c *drainConn) takePendingLocked() []*requestDrain { + drains := c.pending + c.pending = nil + return drains +} + +func completeDrains(drains []*requestDrain, outcome responseDrainOutcome, queued int, err error) { + for _, drain := range drains { + drain.complete(outcome, queued, err) + } +} + +func firstDrainLog(drains []*requestDrain) *slog.Logger { + if len(drains) == 0 { + return nil + } + return drains[0].log +} + +func connectionContext(ctx context.Context, conn net.Conn) context.Context { + tracked, ok := conn.(*drainConn) + if !ok { + return ctx + } + return context.WithValue(ctx, connectionContextKey{}, tracked) +} + +func connectionState(conn net.Conn, state http.ConnState) { + if tracked, ok := conn.(*drainConn); ok { + tracked.setState(state) + } +} + +func duplicateAndShutdownWrite(conn *net.TCPConn) (int, error) { + raw, err := conn.SyscallConn() + if err != nil { + return -1, err + } + fd := -1 + var opErr error + if err := raw.Control(func(rawFD uintptr) { + fd, opErr = unix.FcntlInt(rawFD, unix.F_DUPFD_CLOEXEC, 0) + }); err != nil { + return -1, err + } + if opErr != nil { + return -1, opErr + } + if err := unix.Shutdown(fd, unix.SHUT_WR); err != nil { + _ = unix.Close(fd) + return -1, err + } + return fd, nil +} + +func (c *drainConn) recoverResponseConnection(drains []*requestDrain, config responseDrainConfig, outcome responseDrainOutcome, queued int, cause error) { + failedCount := len(drains) + if failedCount > 0 { + failClosedResponseHolds.Add(int64(failedCount)) + } + retryInterval := config.abortRetryInterval + deadline := time.Now().Add(config.terminalRecoveryTimeout) + for { + c.terminalMu.Lock() + abortErr := config.abort(c.TCPConn) + if abortErr == nil || isTerminalConnectionError(abortErr) { + c.terminalMu.Unlock() + if failedCount > 0 { + failClosedResponseHolds.Add(-int64(failedCount)) + } + completeDrains(drains, outcome, queued, cause) + return + } + + fd, duplicateErr := config.duplicate(c.TCPConn) + if duplicateErr == nil { + closeErr := c.TCPConn.Close() + c.terminalMu.Unlock() + go monitorClosedResponse(fd, drains, config, true, errors.Join(cause, abortErr, closeErr)) + return + } + c.terminalMu.Unlock() + if isTerminalConnectionError(duplicateErr) { + if failedCount > 0 { + failClosedResponseHolds.Add(-int64(failedCount)) + } + completeDrains(drains, responseDrainConnectionClosed, queued, cause) + return + } + + recordResponseDrainOutcome(responseDrainAbortError) + terminalErr := errors.Join(cause, abortErr, duplicateErr) + if time.Now().After(deadline) { + terminateAfterResponseFailure(drains, config, terminalErr) + } + if log := firstDrainLog(drains); log != nil { + log.Error("failed to terminate response connection; retrying while scale-to-zero remains held", "outcome", responseDrainAbortError, "error", terminalErr) + } + time.Sleep(min(retryInterval, time.Until(deadline))) + retryInterval = min(retryInterval*2, responseAbortMaxRetryInterval) + } +} + +func monitorClosedResponse(fd int, drains []*requestDrain, config responseDrainConfig, failed bool, cause error) { + userTimeoutErr := config.setUserTimeout(fd, config.terminalRecoveryTimeout) + outcome, queued, err := waitForClosedResponse(fd, config, time.Now().Add(config.timeout)) + cause = errors.Join(cause, userTimeoutErr, err) + if outcome == responseDrainCloseAcknowledged { + _ = config.closeFD(fd) + if failed { + failClosedResponseHolds.Add(-int64(len(drains))) + } + completeDrains(drains, outcome, queued, cause) + return + } + + if !failed { + failClosedResponseHolds.Add(int64(len(drains))) + } + retryInterval := config.abortRetryInterval + deadline := time.Now().Add(config.terminalRecoveryTimeout) + for { + abortErr := config.abortFD(fd) + if abortErr == nil || isTerminalConnectionError(abortErr) { + failClosedResponseHolds.Add(-int64(len(drains))) + completeDrains(drains, outcome, queued, cause) + return + } + + currentOutcome, currentQueued, inspectErr := waitForClosedResponse(fd, config, time.Now()) + if currentOutcome == responseDrainCloseAcknowledged || isTerminalConnectionError(inspectErr) { + _ = config.closeFD(fd) + failClosedResponseHolds.Add(-int64(len(drains))) + completeDrains(drains, responseDrainCloseAcknowledged, currentQueued, cause) + return + } + userTimeoutErr = config.setUserTimeout(fd, config.terminalRecoveryTimeout) + recordResponseDrainOutcome(responseDrainAbortError) + terminalErr := errors.Join(cause, abortErr, inspectErr, userTimeoutErr) + if time.Now().After(deadline) { + terminateAfterResponseFailure(drains, config, terminalErr) + } + if log := firstDrainLog(drains); log != nil { + log.Error("failed to abort closed response; retrying while scale-to-zero remains held", "outcome", responseDrainAbortError, "error", terminalErr) + } + time.Sleep(min(retryInterval, time.Until(deadline))) + retryInterval = min(retryInterval*2, responseAbortMaxRetryInterval) + } +} + +func terminateAfterResponseFailure(drains []*requestDrain, config responseDrainConfig, err error) { + recordResponseDrainOutcome(responseDrainGuestTermination) + if log := firstDrainLog(drains); log != nil { + log.Error("response connection could not be terminated; terminating guest", "outcome", responseDrainGuestTermination, "error", err) + } + config.terminateGuest() + panic("scale-to-zero guest termination returned") +} + +func waitForClosedResponse(fd int, config responseDrainConfig, deadline time.Time) (responseDrainOutcome, int, error) { + interval := config.initialPollInterval + queued := 0 + for { + var err error + queued, err = config.outboundFD(fd) + if err != nil { + return responseDrainIOError, queued, err + } + acked, err := config.closeAcknowledged(fd) + if err != nil { + return responseDrainIOError, queued, err + } + if queued == 0 && acked { + return responseDrainCloseAcknowledged, 0, nil + } + remaining := time.Until(deadline) + if remaining <= 0 { + return responseDrainTimeoutHit, queued, nil + } + time.Sleep(min(interval, remaining)) + interval = min(interval*2, config.maxPollInterval) + } +} diff --git a/server/lib/scaletozero/connection_test.go b/server/lib/scaletozero/connection_test.go new file mode 100644 index 000000000..e3dbfb912 --- /dev/null +++ b/server/lib/scaletozero/connection_test.go @@ -0,0 +1,484 @@ +package scaletozero + +import ( + "bufio" + "bytes" + "context" + "fmt" + "io" + "log/slog" + "net" + "net/http" + "os" + "path/filepath" + "runtime" + "sync/atomic" + "testing" + "time" + + "github.com/coder/websocket" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestHijackedWebSocketCancellationAndCloseReturnPromptly(t *testing.T) { + config := testDrainConfig(2*time.Second, func(*net.TCPConn) (int, error) { return 1, nil }) + accepted := make(chan struct{}) + startWrite := make(chan struct{}) + writeResult := make(chan struct { + duration time.Duration + err error + }, 1) + closeDuration := make(chan time.Duration, 1) + handler := middleware(NewNoopController(), config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + tracked := r.Context().Value(connectionContextKey{}).(*drainConn) + _ = tracked.SetWriteBuffer(4 << 10) + conn, err := websocket.Accept(w, r, nil) + if err != nil { + writeResult <- struct { + duration time.Duration + err error + }{err: err} + return + } + close(accepted) + <-startWrite + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + started := time.Now() + err = conn.Write(ctx, websocket.MessageBinary, bytes.Repeat([]byte("x"), 32<<20)) + cancel() + writeResult <- struct { + duration time.Duration + err error + }{duration: time.Since(started), err: err} + started = time.Now() + conn.CloseNow() + closeDuration <- time.Since(started) + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + conn, _, err := websocket.Dial(context.Background(), "ws://"+listener.Addr().String(), nil) + require.NoError(t, err) + defer conn.CloseNow() + <-accepted + close(startWrite) + + select { + case result := <-writeResult: + require.Error(t, result.err) + assert.Less(t, result.duration, 500*time.Millisecond) + case <-time.After(time.Second): + t.Fatal("WebSocket write cancellation blocked on response draining") + } + select { + case elapsed := <-closeDuration: + assert.Less(t, elapsed, 200*time.Millisecond) + case <-time.After(time.Second): + t.Fatal("WebSocket close blocked on response draining") + } +} + +func TestShutdownDoesNotBlockOnIdleResponseDrain(t *testing.T) { + started := make(chan struct{}) + config := testDrainConfig(5*time.Second, func(*net.TCPConn) (int, error) { + select { + case <-started: + default: + close(started) + } + return 1, nil + }) + handler := middleware(NewNoopController(), config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Length", "2") + _, _ = io.WriteString(w, "ok") + })) + + listener, serverDone, server := serveTestHandler(t, handler) + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + _, err = fmt.Fprint(conn, "GET / HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodGet}) + require.NoError(t, err) + _, err = io.Copy(io.Discard, response.Body) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + <-started + + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + before := time.Now() + require.NoError(t, server.Shutdown(ctx)) + assert.Less(t, time.Since(before), 200*time.Millisecond) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) +} + +func TestDisableAndAbortFailureNeverReturnsSuccess(t *testing.T) { + ctrl := &mockScaleToZeroer{disableErr: assert.AnError} + config := testDrainConfig(time.Second, outboundQueue) + config.abort = func(*net.TCPConn) error { return assert.AnError } + called := make(chan struct{}, 1) + handler := middleware(ctrl, config)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + called <- struct{}{} + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + _, err = fmt.Fprint(conn, "POST /mutate HTTP/1.1\r\nHost: test\r\nContent-Length: 0\r\n\r\n") + require.NoError(t, err) + response, readErr := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodPost}) + if readErr == nil { + defer response.Body.Close() + assert.GreaterOrEqual(t, response.StatusCode, http.StatusInternalServerError) + } + select { + case <-called: + t.Fatal("application handler ran after scale-to-zero disable failed") + default: + } +} + +func TestHTTP10CloseDelimitedResponseWaitsForCloseAcknowledgement(t *testing.T) { + const bodySize = 4 << 20 + base := newSignalController() + ctrl := NewDebouncedController(base) + outcome := make(chan responseDrainOutcome, 1) + config := testDrainConfig(5*time.Second, outboundQueue) + config.onComplete = func(value responseDrainOutcome) { outcome <- value } + handler := middleware(ctrl, config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + flusher := w.(http.Flusher) + for remaining := bodySize; remaining > 0; remaining -= 32 << 10 { + _, _ = w.Write(bytes.Repeat([]byte("x"), 32<<10)) + flusher.Flush() + } + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + require.NoError(t, conn.(*net.TCPConn).SetReadBuffer(32<<10)) + _, err = fmt.Fprint(conn, "GET /stream HTTP/1.0\r\nHost: test\r\n\r\n") + require.NoError(t, err) + response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodGet}) + require.NoError(t, err) + assert.Equal(t, int64(-1), response.ContentLength) + assert.True(t, response.Close) + written, err := io.Copy(io.Discard, response.Body) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + assert.Equal(t, int64(bodySize), written) + + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero was not enabled after close acknowledgement") + } + assert.Equal(t, responseDrainCloseAcknowledged, <-outcome) +} + +func TestResponseDrainControlsScaleToZeroFile(t *testing.T) { + const bodySize = 2 << 20 + scaleFile := filepath.Join(t.TempDir(), "scale_to_zero_disable") + require.NoError(t, os.WriteFile(scaleFile, []byte("-"), 0o600)) + ctrl := NewDebouncedController(&unikraftCloudController{path: scaleFile}) + handler := middleware(ctrl, testDrainConfig(5*time.Second, outboundQueue))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Length", fmt.Sprint(bodySize)) + _, _ = io.Copy(w, bytes.NewReader(make([]byte, bodySize))) + })) + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + require.NoError(t, conn.(*net.TCPConn).SetReadBuffer(32<<10)) + _, err = fmt.Fprint(conn, "GET /large HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodGet}) + require.NoError(t, err) + + require.Eventually(t, func() bool { + value, readErr := os.ReadFile(scaleFile) + return readErr == nil && string(value) == "+" + }, time.Second, time.Millisecond) + _, err = io.Copy(io.Discard, response.Body) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + require.Eventually(t, func() bool { + value, readErr := os.ReadFile(scaleFile) + return readErr == nil && string(value) == "-" + }, time.Second, time.Millisecond) +} + +func TestPersistentConnectionRecoveryTerminatesGuestWithinBound(t *testing.T) { + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + + activeBefore := ActiveResponseHolds() + failedBefore := FailClosedResponseHolds() + released := make(chan struct{}) + hold := newResponseHold(func() { close(released) }) + hold.finishHandler() + config := testDrainConfig(20*time.Millisecond, outboundQueue) + config.abort = func(*net.TCPConn) error { return assert.AnError } + config.duplicate = func(*net.TCPConn) (int, error) { return -1, assert.AnError } + terminated := make(chan struct{}) + config.terminateGuest = func() { + close(terminated) + runtime.Goexit() + } + drains := []*requestDrain{{hold: hold, config: config, log: slog.Default()}} + + go tracked.recoverResponseConnection(drains, config, responseDrainIOError, 1, assert.AnError) + select { + case <-terminated: + case <-time.After(time.Second): + t.Fatal("persistent connection recovery did not terminate the guest") + } + select { + case <-released: + t.Fatal("scale-to-zero hold released without a drained or aborted connection") + default: + } + assert.Equal(t, activeBefore+1, ActiveResponseHolds()) + assert.Equal(t, failedBefore+1, FailClosedResponseHolds()) + + failClosedResponseHolds.Add(-1) + completeDrains(drains, responseDrainConnectionClosed, 0, nil) + assert.Equal(t, activeBefore, ActiveResponseHolds()) + assert.Equal(t, failedBefore, FailClosedResponseHolds()) +} + +func TestPersistentClosedSocketRecoveryTerminatesGuestWithinBound(t *testing.T) { + activeBefore := ActiveResponseHolds() + failedBefore := FailClosedResponseHolds() + released := make(chan struct{}) + hold := newResponseHold(func() { close(released) }) + hold.finishHandler() + config := testDrainConfig(20*time.Millisecond, outboundQueue) + config.outboundFD = func(int) (int, error) { return 1, assert.AnError } + config.closeAcknowledged = func(int) (bool, error) { return false, assert.AnError } + config.abortFD = func(int) error { return assert.AnError } + config.setUserTimeout = func(int, time.Duration) error { return assert.AnError } + terminated := make(chan struct{}) + config.terminateGuest = func() { + close(terminated) + runtime.Goexit() + } + drains := []*requestDrain{{hold: hold, config: config, log: slog.Default()}} + + go monitorClosedResponse(-1, drains, config, false, assert.AnError) + select { + case <-terminated: + case <-time.After(time.Second): + t.Fatal("persistent closed-socket recovery did not terminate the guest") + } + select { + case <-released: + t.Fatal("scale-to-zero hold released without a drained or aborted connection") + default: + } + assert.Equal(t, activeBefore+1, ActiveResponseHolds()) + assert.Equal(t, failedBefore+1, FailClosedResponseHolds()) + + failClosedResponseHolds.Add(-1) + completeDrains(drains, responseDrainConnectionClosed, 0, nil) + assert.Equal(t, activeBefore, ActiveResponseHolds()) + assert.Equal(t, failedBefore, FailClosedResponseHolds()) +} + +func TestNewRequestCancelsPreviousIdleDrain(t *testing.T) { + var requests atomic.Int32 + config := testDrainConfig(50*time.Millisecond, func(*net.TCPConn) (int, error) { + if requests.Load() < 2 { + return 1, nil + } + return 0, nil + }) + handler := middleware(NewNoopController(), config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + request := requests.Add(1) + if request == 2 { + time.Sleep(100 * time.Millisecond) + } + w.Header().Set("Content-Length", "2") + _, _ = io.WriteString(w, "ok") + })) + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + transport := &http.Transport{MaxConnsPerHost: 1} + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport} + for i := 0; i < 2; i++ { + response, err := client.Get("http://" + listener.Addr().String()) + require.NoError(t, err) + body, err := io.ReadAll(response.Body) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + assert.Equal(t, "ok", string(body)) + } +} + +func TestReadFromRefreshesDeadlineDuringProgress(t *testing.T) { + const bodySize = 32 << 20 + const timeout = 200 * time.Millisecond + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + require.NoError(t, tracked.SetWriteBuffer(64<<10)) + + file, err := os.Create(filepath.Join(t.TempDir(), "response.bin")) + require.NoError(t, err) + defer file.Close() + require.NoError(t, file.Truncate(bodySize)) + + config := testDrainConfig(timeout, outboundQueue) + var deadlines atomic.Int32 + config.setDeadline = func(conn *net.TCPConn, deadline time.Time) error { + deadlines.Add(1) + return conn.SetWriteDeadline(deadline) + } + tracked.configure(config) + + readDone := make(chan error, 1) + go func() { + remaining := bodySize + buffer := make([]byte, 64<<10) + for remaining > 0 { + n, readErr := client.Read(buffer) + remaining -= n + if readErr != nil { + readDone <- readErr + return + } + time.Sleep(time.Millisecond) + } + readDone <- nil + }() + + started := time.Now() + written, err := tracked.ReadFrom(file) + require.NoError(t, err) + assert.Equal(t, int64(bodySize), written) + assert.Greater(t, time.Since(started), timeout) + require.NoError(t, <-readDone) + assert.GreaterOrEqual(t, deadlines.Load(), int32(bodySize/responseReadFromChunkSize)) +} + +func TestPipelinedRequestsKeepOnePendingHold(t *testing.T) { + const requestCount = 1000 + activeBefore := ActiveResponseHolds() + base := &mockScaleToZeroer{} + ctrl := NewDebouncedController(base) + var allowDrain atomic.Bool + config := testDrainConfig(time.Second, func(*net.TCPConn) (int, error) { + if allowDrain.Load() { + return 0, nil + } + return 1, nil + }) + trackedConn := make(chan *drainConn, 1) + handler := middleware(ctrl, config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case trackedConn <- r.Context().Value(connectionContextKey{}).(*drainConn): + default: + } + w.Header().Set("Content-Length", "1") + _, _ = io.WriteString(w, "x") + })) + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + requests := bytes.NewBuffer(make([]byte, 0, requestCount*32)) + for range requestCount { + _, _ = fmt.Fprint(requests, "GET / HTTP/1.1\r\nHost: test\r\n\r\n") + } + _, err = conn.Write(requests.Bytes()) + require.NoError(t, err) + reader := bufio.NewReader(conn) + for range requestCount { + response, readErr := http.ReadResponse(reader, &http.Request{Method: http.MethodGet}) + require.NoError(t, readErr) + body, readErr := io.ReadAll(response.Body) + require.NoError(t, readErr) + require.NoError(t, response.Body.Close()) + assert.Equal(t, "x", string(body)) + } + + tracked := <-trackedConn + tracked.mu.Lock() + pending := len(tracked.pending) + tracked.mu.Unlock() + assert.Equal(t, 1, pending) + assert.Equal(t, activeBefore+1, ActiveResponseHolds()) + base.mu.Lock() + assert.Equal(t, 1, base.disableCalls) + assert.Equal(t, 0, base.enableCalls) + base.mu.Unlock() + + allowDrain.Store(true) + require.Eventually(t, func() bool { + base.mu.Lock() + defer base.mu.Unlock() + return base.enableCalls == 1 && ActiveResponseHolds() == activeBefore + }, time.Second, time.Millisecond) + tracked.mu.Lock() + assert.Empty(t, tracked.pending) + tracked.mu.Unlock() +} + +func TestKeepAliveResponsesDoNotWaitForMaximumPollInterval(t *testing.T) { + config := testDrainConfig(time.Second, outboundQueue) + config.maxPollInterval = responseDrainMaxPollInterval + handler := middleware(NewNoopController(), config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Length", "2") + _, _ = io.WriteString(w, "ok") + })) + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + transport := &http.Transport{MaxConnsPerHost: 1} + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport} + started := time.Now() + for i := 0; i < 100; i++ { + response, err := client.Get("http://" + listener.Addr().String()) + require.NoError(t, err) + _, err = io.Copy(io.Discard, response.Body) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + } + assert.Less(t, time.Since(started), 5*time.Second) +} diff --git a/server/lib/scaletozero/middleware.go b/server/lib/scaletozero/middleware.go index b5452c06c..e6f47206d 100644 --- a/server/lib/scaletozero/middleware.go +++ b/server/lib/scaletozero/middleware.go @@ -2,17 +2,193 @@ package scaletozero import ( "context" + "errors" + "log/slog" "net" "net/http" + "sync" + "sync/atomic" + "time" "github.com/kernel/kernel-images/server/lib/logger" + "golang.org/x/sys/unix" ) -// Middleware returns a standard net/http middleware that disables scale-to-zero -// at the start of each request and re-enables it after the handler completes. -// Connections from loopback addresses are ignored and do not affect the -// scale-to-zero state. +const ( + responseDrainInitialPollInterval = time.Millisecond + responseDrainMaxPollInterval = 100 * time.Millisecond + responseDrainTimeout = 5 * time.Minute + responseAbortRetryInterval = time.Second + responseAbortMaxRetryInterval = 30 * time.Second + responseTerminalRecoveryTimeout = 5 * time.Minute +) + +type connectionContextKey struct{} + +type responseDrainConfig struct { + initialPollInterval time.Duration + maxPollInterval time.Duration + timeout time.Duration + outbound func(*net.TCPConn) (int, error) + setDeadline func(*net.TCPConn, time.Time) error + abort func(*net.TCPConn) error + duplicate func(*net.TCPConn) (int, error) + setUserTimeout func(int, time.Duration) error + abortFD func(int) error + closeFD func(int) error + outboundFD func(int) (int, error) + closeAcknowledged func(int) (bool, error) + terminateGuest func() + abortRetryInterval time.Duration + terminalRecoveryTimeout time.Duration + onComplete func(responseDrainOutcome) +} + +type responseDrainOutcome string + +const ( + responseDrainComplete responseDrainOutcome = "drained" + responseDrainCloseAcknowledged responseDrainOutcome = "close_acknowledged" + responseDrainTimeoutHit responseDrainOutcome = "timeout" + responseDrainIOError responseDrainOutcome = "ioctl_error" + responseDrainNonTCP responseDrainOutcome = "untracked_connection" + responseDrainWriteError responseDrainOutcome = "write_error" + responseDrainDeadlineClearError responseDrainOutcome = "deadline_clear_error" + responseDrainConnectionClosed responseDrainOutcome = "connection_closed" + responseDrainConnectionHijacked responseDrainOutcome = "connection_hijacked" + responseDrainConnectionReused responseDrainOutcome = "connection_reused" + responseDrainAbortError responseDrainOutcome = "abort_error" + responseDrainGuestTermination responseDrainOutcome = "guest_termination" +) + +var responseDrainCounters = map[responseDrainOutcome]*atomic.Uint64{ + responseDrainComplete: {}, + responseDrainCloseAcknowledged: {}, + responseDrainTimeoutHit: {}, + responseDrainIOError: {}, + responseDrainNonTCP: {}, + responseDrainWriteError: {}, + responseDrainDeadlineClearError: {}, + responseDrainConnectionClosed: {}, + responseDrainConnectionHijacked: {}, + responseDrainConnectionReused: {}, + responseDrainAbortError: {}, + responseDrainGuestTermination: {}, +} + +var activeResponseHolds atomic.Int64 +var failClosedResponseHolds atomic.Int64 + +// ResponseDrainOutcomeCounts returns process-lifetime response drain counters. +func ResponseDrainOutcomeCounts() map[string]uint64 { + counts := make(map[string]uint64, len(responseDrainCounters)) + for outcome, counter := range responseDrainCounters { + counts[string(outcome)] = counter.Load() + } + return counts +} + +func ActiveResponseHolds() int64 { return activeResponseHolds.Load() } + +func FailClosedResponseHolds() int64 { return failClosedResponseHolds.Load() } + +func recordResponseDrainOutcome(outcome responseDrainOutcome) { + responseDrainCounters[outcome].Add(1) +} + +type responseHold struct { + mu sync.Mutex + handlerDone bool + connectionDone bool + released bool + release func() +} + +type requestDrain struct { + hold *responseHold + config responseDrainConfig + log *slog.Logger +} + +func newResponseHold(release func()) *responseHold { + activeResponseHolds.Add(1) + return &responseHold{release: release} +} + +func (h *responseHold) finishHandler() { + h.finish(true) +} + +func (h *responseHold) finishConnection() { + h.finish(false) +} + +func (h *responseHold) finish(handler bool) { + h.mu.Lock() + if handler { + h.handlerDone = true + } else { + h.connectionDone = true + } + release := !h.released && h.handlerDone && h.connectionDone + if release { + h.released = true + } + h.mu.Unlock() + if release { + activeResponseHolds.Add(-1) + h.release() + } +} + +func (d *requestDrain) complete(outcome responseDrainOutcome, queued int, err error) { + recordResponseDrainOutcome(outcome) + if d.config.onComplete != nil { + d.config.onComplete(outcome) + } + attrs := []any{"outcome", outcome} + if queued > 0 { + attrs = append(attrs, "queued_bytes", queued) + } + if err != nil { + attrs = append(attrs, "error", err) + } + switch outcome { + case responseDrainComplete, responseDrainCloseAcknowledged, responseDrainConnectionClosed, responseDrainConnectionHijacked, responseDrainConnectionReused: + d.log.Debug("response drain finished", attrs...) + default: + d.log.Warn("response drain finished", attrs...) + } + d.hold.finishConnection() +} + +// Middleware holds scale-to-zero disabled until each non-loopback HTTP response +// is finalized and its TCP send queue drains or the connection terminates. func Middleware(ctrl Controller) func(http.Handler) http.Handler { + return middleware(ctrl, defaultResponseDrainConfig()) +} + +func defaultResponseDrainConfig() responseDrainConfig { + return responseDrainConfig{ + initialPollInterval: responseDrainInitialPollInterval, + maxPollInterval: responseDrainMaxPollInterval, + timeout: responseDrainTimeout, + outbound: outboundQueue, + setDeadline: setWriteDeadline, + abort: abortConnection, + duplicate: duplicateAndShutdownWrite, + setUserTimeout: setTCPUserTimeout, + abortFD: abortSocket, + closeFD: closeSocket, + outboundFD: outboundQueueFD, + closeAcknowledged: closeAcknowledged, + terminateGuest: terminateGuest, + abortRetryInterval: responseAbortRetryInterval, + terminalRecoveryTimeout: responseTerminalRecoveryTimeout, + } +} + +func middleware(ctrl Controller, config responseDrainConfig) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if isLoopbackAddr(r.RemoteAddr) { @@ -20,18 +196,117 @@ func Middleware(ctrl Controller) func(http.Handler) http.Handler { return } + ctx := context.WithoutCancel(r.Context()) + log := logger.FromContext(ctx) + conn, ok := r.Context().Value(connectionContextKey{}).(*drainConn) + if !ok { + recordResponseDrainOutcome(responseDrainNonTCP) + log.Error("response drain unavailable", "outcome", responseDrainNonTCP) + panic(http.ErrAbortHandler) + } if err := ctrl.Disable(r.Context()); err != nil { - logger.FromContext(r.Context()).Error("failed to disable scale-to-zero", "error", err) - http.Error(w, "failed to disable scale-to-zero", http.StatusInternalServerError) + log.Error("failed to disable scale-to-zero", "error", err) + conn.abortNow(config.abort) + panic(http.ErrAbortHandler) + } + + hold := newResponseHold(func() { + if err := ctrl.Enable(ctx); err != nil { + log.Error("failed to release response scale-to-zero hold", "error", err) + } + }) + drain := &requestDrain{hold: hold, config: config, log: log} + conn.configure(config) + registered := conn.addDrain(drain) + defer hold.finishHandler() + if !registered { return } - defer ctrl.Enable(context.WithoutCancel(r.Context())) + defer func() { + if err := conn.writeError(); err != nil { + recordResponseDrainOutcome(responseDrainWriteError) + log.Warn("response write failed", "outcome", responseDrainWriteError, "error", err) + } + }() next.ServeHTTP(w, r) }) } } +func waitForResponseDrain(ctx context.Context, conn *net.TCPConn, config responseDrainConfig, deadline time.Time) (responseDrainOutcome, int, error) { + queued, err := config.outbound(conn) + if err != nil { + return responseDrainIOError, queued, err + } + if queued == 0 { + return responseDrainComplete, 0, nil + } + + interval := config.initialPollInterval + for { + remaining := time.Until(deadline) + if remaining <= 0 { + return responseDrainTimeoutHit, queued, nil + } + if interval > remaining { + interval = remaining + } + timer := time.NewTimer(interval) + select { + case <-ctx.Done(): + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + return "", queued, ctx.Err() + case <-timer.C: + } + + queued, err = config.outbound(conn) + if err != nil { + return responseDrainIOError, queued, err + } + if queued == 0 { + return responseDrainComplete, 0, nil + } + interval = min(interval*2, config.maxPollInterval) + } +} + +func outboundQueue(conn *net.TCPConn) (int, error) { + raw, err := conn.SyscallConn() + if err != nil { + return 0, err + } + + var queued int + var ioctlErr error + if err := raw.Control(func(fd uintptr) { + queued, ioctlErr = unix.IoctlGetInt(int(fd), unix.TIOCOUTQ) + }); err != nil { + return 0, err + } + return queued, ioctlErr +} + +func setWriteDeadline(conn *net.TCPConn, deadline time.Time) error { + return conn.SetWriteDeadline(deadline) +} + +func abortConnection(conn *net.TCPConn) error { + if err := conn.SetLinger(0); err != nil { + return err + } + return conn.Close() +} + +func isTerminalConnectionError(err error) bool { + return errors.Is(err, net.ErrClosed) || errors.Is(err, unix.EBADF) || errors.Is(err, unix.ENOTCONN) +} + // isLoopbackAddr reports whether addr is a loopback address. // addr may be an "ip:port" pair or a bare IP. func isLoopbackAddr(addr string) bool { diff --git a/server/lib/scaletozero/middleware_test.go b/server/lib/scaletozero/middleware_test.go index c48b61226..2f4306af2 100644 --- a/server/lib/scaletozero/middleware_test.go +++ b/server/lib/scaletozero/middleware_test.go @@ -1,45 +1,57 @@ package scaletozero import ( + "bufio" + "bytes" + "context" + "fmt" + "io" + "net" "net/http" "net/http/httptest" + "sync" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func TestMiddlewareDisablesAndEnablesForExternalAddr(t *testing.T) { - t.Parallel() - mock := &mockScaleToZeroer{} - handler := Middleware(mock)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { +func TestMiddlewareDisablesUntilResponseFinalization(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + + handler := middleware(ctrl, testDrainConfig(time.Second, func(*net.TCPConn) (int, error) { + return 0, nil + }))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) + req := externalRequest(http.MethodGet, "/json/version", tracked) - req := httptest.NewRequest(http.MethodGet, "/", nil) - req.RemoteAddr = "203.0.113.50:12345" - rec := httptest.NewRecorder() - - handler.ServeHTTP(rec, req) + handler.ServeHTTP(httptest.NewRecorder(), req) + select { + case <-base.enabled: + t.Fatal("scale-to-zero enabled before response finalization") + default: + } - assert.Equal(t, http.StatusOK, rec.Code) - assert.Equal(t, 1, mock.disableCalls) - assert.Equal(t, 1, mock.enableCalls) + tracked.startIdleDrain() + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero was not enabled after response drain") + } } func TestMiddlewareSkipsLoopbackAddrs(t *testing.T) { t.Parallel() - loopbackAddrs := []struct { - name string - addr string - }{ - {"loopback-v4", "127.0.0.1:8080"}, - {"loopback-v6", "[::1]:8080"}, - } - - for _, tc := range loopbackAddrs { - t.Run(tc.name, func(t *testing.T) { + loopbackAddrs := []string{"127.0.0.1:8080", "[::1]:8080"} + for _, addr := range loopbackAddrs { + t.Run(addr, func(t *testing.T) { t.Parallel() mock := &mockScaleToZeroer{} var called bool @@ -49,38 +61,426 @@ func TestMiddlewareSkipsLoopbackAddrs(t *testing.T) { })) req := httptest.NewRequest(http.MethodGet, "/", nil) - req.RemoteAddr = tc.addr - rec := httptest.NewRecorder() + req.RemoteAddr = addr + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) - handler.ServeHTTP(rec, req) - - assert.True(t, called, "handler should still be called") - assert.Equal(t, http.StatusOK, rec.Code) - assert.Equal(t, 0, mock.disableCalls, "should not disable for loopback addr") - assert.Equal(t, 0, mock.enableCalls, "should not enable for loopback addr") + assert.True(t, called) + assert.Equal(t, http.StatusOK, recorder.Code) + assert.Equal(t, 0, mock.disableCalls) + assert.Equal(t, 0, mock.enableCalls) }) } } -func TestMiddlewareDisableError(t *testing.T) { +func TestMiddlewareDisableErrorAbortsConnection(t *testing.T) { t.Parallel() mock := &mockScaleToZeroer{disableErr: assert.AnError} + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() var called bool handler := Middleware(mock)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })) - req := httptest.NewRequest(http.MethodGet, "/", nil) - req.RemoteAddr = "203.0.113.50:12345" - rec := httptest.NewRecorder() - - handler.ServeHTTP(rec, req) + assert.PanicsWithValue(t, http.ErrAbortHandler, func() { + handler.ServeHTTP(httptest.NewRecorder(), externalRequest(http.MethodGet, "/", tracked)) + }) - assert.False(t, called, "handler should not be called on disable error") - assert.Equal(t, http.StatusInternalServerError, rec.Code) + assert.False(t, called) + assert.True(t, tracked.closed) assert.Equal(t, 0, mock.enableCalls) } +func TestMiddlewareRejectsUntrackedConnection(t *testing.T) { + ctrl := &mockScaleToZeroer{} + var called bool + handler := Middleware(ctrl)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + called = true + })) + req := httptest.NewRequest(http.MethodGet, "/fs/read_file", nil) + req.RemoteAddr = "192.0.2.1:1234" + + assert.PanicsWithValue(t, http.ErrAbortHandler, func() { + handler.ServeHTTP(httptest.NewRecorder(), req) + }) + + assert.False(t, called) + assert.Equal(t, 0, ctrl.disableCalls) + assert.Equal(t, 0, ctrl.enableCalls) +} + +func TestMiddlewareWaitsForTCPResponseDrain(t *testing.T) { + const bodySize = 2 << 20 + + type handlerResult struct { + at time.Time + err error + } + + base := newSignalController() + ctrl := NewDebouncedController(base) + handlerDone := make(chan handlerResult, 1) + var sawQueued bool + var queueEmptyAt time.Time + var queueEmptyOnce sync.Once + drain := testDrainConfig(5*time.Second, func(conn *net.TCPConn) (int, error) { + queued, err := outboundQueue(conn) + if queued > 0 { + sawQueued = true + } else if err == nil { + queueEmptyOnce.Do(func() { queueEmptyAt = time.Now() }) + } + return queued, err + }) + handler := middleware(ctrl, drain)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Length", fmt.Sprint(bodySize)) + _, err := io.Copy(w, bytes.NewReader(make([]byte, bodySize))) + handlerDone <- handlerResult{at: time.Now(), err: err} + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + require.NoError(t, conn.(*net.TCPConn).SetReadBuffer(32<<10)) + _, err = fmt.Fprint(conn, "GET /fs/read_file HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + + response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodGet}) + require.NoError(t, err) + defer response.Body.Close() + + buffer := make([]byte, 4<<10) + read := 0 + for read < bodySize { + n, readErr := response.Body.Read(buffer) + read += n + if readErr != nil { + require.ErrorIs(t, readErr, io.EOF) + break + } + time.Sleep(time.Millisecond) + } + + enabledAt := <-base.enabled + result := <-handlerDone + require.NoError(t, result.err) + assert.True(t, sawQueued) + assert.False(t, queueEmptyAt.IsZero()) + assert.False(t, enabledAt.Before(queueEmptyAt)) +} + +func TestMiddlewareDrainsAfterChunkedResponseFinalization(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + handlerErr := make(chan error, 1) + handler := middleware(ctrl, testDrainConfig(time.Second, outboundQueue))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, err := w.Write(bytes.Repeat([]byte("x"), 64<<10)) + handlerErr <- err + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + _, err = fmt.Fprint(conn, "GET /stream HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n") + require.NoError(t, err) + + response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodGet}) + require.NoError(t, err) + body, err := io.ReadAll(response.Body) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + assert.Equal(t, []string{"chunked"}, response.TransferEncoding) + assert.Len(t, body, 64<<10) + require.NoError(t, <-handlerErr) + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero was not enabled after chunked response drained") + } +} + +func TestMiddlewareResponseDrainTimeoutAbortsConnection(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + const ceiling = 50 * time.Millisecond + var aborted bool + drain := testDrainConfig(ceiling, func(*net.TCPConn) (int, error) { return 1, nil }) + drain.abort = func(conn *net.TCPConn) error { + aborted = true + return abortConnection(conn) + } + handler := middleware(ctrl, drain)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Length", "2") + _, _ = io.WriteString(w, "ok") + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + started := time.Now() + _, err = fmt.Fprint(conn, "GET /anything HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + + select { + case enabledAt := <-base.enabled: + assert.GreaterOrEqual(t, enabledAt.Sub(started), ceiling) + assert.Less(t, enabledAt.Sub(started), 3*ceiling) + case <-time.After(time.Second): + t.Fatal("scale-to-zero hold was not released after abort") + } + assert.True(t, aborted) +} + +func TestMiddlewareIOErrorAbortsConnection(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + var aborted bool + drain := testDrainConfig(time.Second, func(*net.TCPConn) (int, error) { + return 1, assert.AnError + }) + drain.abort = func(conn *net.TCPConn) error { + aborted = true + return abortConnection(conn) + } + handler := middleware(ctrl, drain)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.WriteString(w, "response") + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + _, err = fmt.Fprint(conn, "GET /anything HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero hold was not released after abort") + } + assert.True(t, aborted) +} + +func TestMiddlewareFallsBackToClosedSocketMonitoringAfterAbortFailure(t *testing.T) { + failClosedBefore := FailClosedResponseHolds() + base := newSignalController() + ctrl := NewDebouncedController(base) + drain := testDrainConfig(time.Second, func(*net.TCPConn) (int, error) { + return 1, assert.AnError + }) + abortAttempted := make(chan struct{}) + drain.abort = func(*net.TCPConn) error { + close(abortAttempted) + return assert.AnError + } + handler := middleware(ctrl, drain)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.WriteString(w, "response") + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + _, err = fmt.Fprint(conn, "GET /anything HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + <-abortAttempted + select { + case <-base.enabled: + t.Fatal("scale-to-zero enabled before connection recovery") + default: + } + + _, _ = io.Copy(io.Discard, conn) + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero hold was not released after closed-socket monitoring") + } + assert.Eventually(t, func() bool { + return FailClosedResponseHolds() == failClosedBefore + }, time.Second, time.Millisecond) +} + +func TestMiddlewareWriteDeadlineCoversHandler(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + const ceiling = 50 * time.Millisecond + handlerDone := make(chan error, 1) + handler := middleware(ctrl, testDrainConfig(ceiling, outboundQueue))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, err := io.Copy(w, bytes.NewReader(make([]byte, 32<<20))) + handlerDone <- err + })) + + listener, serverDone, server := serveTestHandler(t, handler) + defer func() { + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }() + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + require.NoError(t, conn.(*net.TCPConn).SetReadBuffer(4<<10)) + _, err = fmt.Fprint(conn, "GET /large HTTP/1.1\r\nHost: test\r\n\r\n") + require.NoError(t, err) + + select { + case err := <-handlerDone: + require.Error(t, err) + case <-time.After(2 * time.Second): + t.Fatal("handler write was not bounded by the response deadline") + } + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero hold was not released after timed-out connection was aborted") + } +} + +func TestMiddlewareClearsWriteDeadlineAfterDrain(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + + var deadlines []time.Time + drain := testDrainConfig(time.Second, func(*net.TCPConn) (int, error) { return 0, nil }) + drain.setDeadline = func(_ *net.TCPConn, deadline time.Time) error { + deadlines = append(deadlines, deadline) + return nil + } + handler := middleware(ctrl, drain)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + _, err := tracked.Write([]byte("response")) + require.NoError(t, err) + })) + + handler.ServeHTTP(httptest.NewRecorder(), externalRequest(http.MethodGet, "/", tracked)) + tracked.startIdleDrain() + <-base.enabled + + require.Len(t, deadlines, 2) + assert.False(t, deadlines[0].IsZero()) + assert.True(t, deadlines[1].IsZero()) +} + +func TestDrainConnClearsWriteDeadlineWhenHijacked(t *testing.T) { + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + + var deadlines []time.Time + config := testDrainConfig(time.Second, outboundQueue) + config.setDeadline = func(_ *net.TCPConn, deadline time.Time) error { + deadlines = append(deadlines, deadline) + return nil + } + tracked.configure(config) + tracked.hijack() + _, err := tracked.Write([]byte("frame")) + require.NoError(t, err) + + require.Len(t, deadlines, 1) + assert.True(t, deadlines[0].IsZero()) +} + +func TestMiddlewareDoesNotReleaseBeforeClosedHandlerReturns(t *testing.T) { + base := newSignalController() + ctrl := NewDebouncedController(base) + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + handler := middleware(ctrl, testDrainConfig(time.Second, outboundQueue))(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + tracked.abortNow(abortConnection) + select { + case <-base.enabled: + t.Fatal("scale-to-zero enabled while handler was still running") + default: + } + })) + + handler.ServeHTTP(httptest.NewRecorder(), externalRequest(http.MethodGet, "/", tracked)) + select { + case <-base.enabled: + case <-time.After(time.Second): + t.Fatal("scale-to-zero was not enabled after handler returned") + } +} + +func TestMiddlewareWriteDeadlineFailureAbortsConnection(t *testing.T) { + ctrl := &mockScaleToZeroer{} + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + var aborted bool + drain := testDrainConfig(time.Second, outboundQueue) + drain.setDeadline = func(*net.TCPConn, time.Time) error { return assert.AnError } + drain.abort = func(conn *net.TCPConn) error { + aborted = true + return abortConnection(conn) + } + handler := middleware(ctrl, drain)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + _, _ = tracked.Write([]byte("response")) + })) + + handler.ServeHTTP(httptest.NewRecorder(), externalRequest(http.MethodGet, "/", tracked)) + + assert.True(t, aborted) + assert.Equal(t, 1, ctrl.disableCalls) + assert.Equal(t, 1, ctrl.enableCalls) +} + +func TestWaitForResponseDrainCompletesImmediatelyForEmptyQueue(t *testing.T) { + outcome, queued, err := waitForResponseDrain(context.Background(), nil, responseDrainConfig{ + initialPollInterval: time.Millisecond, + maxPollInterval: time.Millisecond, + outbound: func(*net.TCPConn) (int, error) { return 0, nil }, + }, time.Now().Add(time.Second)) + + assert.Equal(t, responseDrainComplete, outcome) + assert.Zero(t, queued) + assert.NoError(t, err) +} + +func TestWaitForResponseDrainReportsIOError(t *testing.T) { + wantErr := assert.AnError + outcome, _, err := waitForResponseDrain(context.Background(), nil, responseDrainConfig{ + initialPollInterval: time.Millisecond, + maxPollInterval: time.Millisecond, + outbound: func(*net.TCPConn) (int, error) { return 0, wantErr }, + }, time.Now().Add(time.Second)) + + assert.Equal(t, responseDrainIOError, outcome) + assert.ErrorIs(t, err, wantErr) +} + func TestIsLoopbackAddr(t *testing.T) { t.Parallel() @@ -88,19 +488,10 @@ func TestIsLoopbackAddr(t *testing.T) { addr string loopback bool }{ - // Loopback {"127.0.0.1:80", true}, - {"[::1]:80", true}, - {"127.0.0.1", true}, - {"::1", true}, - // Non-loopback - {"10.0.0.1:80", false}, - {"172.16.0.1:80", false}, - {"192.168.1.1:80", false}, + {"[::1]:8080", true}, {"203.0.113.50:80", false}, - {"8.8.8.8:53", false}, - {"[2001:db8::1]:80", false}, - // Unparseable + {"2001:db8::1", false}, {"not-an-ip:80", false}, {"", false}, } @@ -112,3 +503,69 @@ func TestIsLoopbackAddr(t *testing.T) { }) } } + +func testDrainConfig(timeout time.Duration, outbound func(*net.TCPConn) (int, error)) responseDrainConfig { + config := defaultResponseDrainConfig() + config.initialPollInterval = time.Millisecond + config.maxPollInterval = time.Millisecond + config.timeout = timeout + config.outbound = outbound + config.abortRetryInterval = time.Millisecond + config.terminalRecoveryTimeout = timeout + return config +} + +func externalRequest(method, path string, conn *drainConn) *http.Request { + req := httptest.NewRequest(method, path, nil) + req.RemoteAddr = "192.0.2.1:1234" + return req.WithContext(connectionContext(req.Context(), conn)) +} + +func newTCPPair(t *testing.T) (*drainConn, net.Conn) { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + accepted := make(chan *net.TCPConn, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr == nil { + accepted <- conn.(*net.TCPConn) + } + }() + client, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + server := <-accepted + require.NoError(t, listener.Close()) + return &drainConn{TCPConn: server, state: http.StateActive}, client +} + +func serveTestHandler(t *testing.T, handler http.Handler) (net.Listener, <-chan error, *http.Server) { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + server := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + r.RemoteAddr = "192.0.2.1:1234" + handler.ServeHTTP(w, r) + }), + } + serverDone := make(chan error, 1) + go func() { serverDone <- Serve(server, listener) }() + return listener, serverDone, server +} + +type signalController struct { + enabled chan time.Time + once sync.Once +} + +func newSignalController() *signalController { + return &signalController{enabled: make(chan time.Time, 1)} +} + +func (*signalController) Disable(context.Context) error { return nil } + +func (c *signalController) Enable(context.Context) error { + c.once.Do(func() { c.enabled <- time.Now() }) + return nil +} diff --git a/server/lib/scaletozero/socket_linux.go b/server/lib/scaletozero/socket_linux.go new file mode 100644 index 000000000..86ed304a3 --- /dev/null +++ b/server/lib/scaletozero/socket_linux.go @@ -0,0 +1,59 @@ +//go:build linux + +package scaletozero + +import ( + "fmt" + "time" + + "golang.org/x/sys/unix" +) + +func setTCPUserTimeout(fd int, timeout time.Duration) error { + return unix.SetsockoptInt(fd, unix.IPPROTO_TCP, unix.TCP_USER_TIMEOUT, int(timeout.Milliseconds())) +} + +func outboundQueueFD(fd int) (int, error) { + return unix.IoctlGetInt(fd, unix.TIOCOUTQ) +} + +func abortSocket(fd int) error { + if err := unix.SetsockoptLinger(fd, unix.SOL_SOCKET, unix.SO_LINGER, &unix.Linger{Onoff: 1}); err != nil { + return err + } + return unix.Close(fd) +} + +func closeSocket(fd int) error { + return unix.Close(fd) +} + +// PID 1 is the image wrapper; it exits the guest after stopping supervisord. +func terminateGuest() { + if err := unix.Kill(1, unix.SIGTERM); err != nil { + panic(fmt.Sprintf("failed to terminate guest: %v", err)) + } + time.Sleep(30 * time.Second) + if err := unix.Kill(1, unix.SIGKILL); err != nil { + panic(fmt.Sprintf("failed to force guest termination: %v", err)) + } + select {} +} + +func closeAcknowledged(fd int) (bool, error) { + info, err := unix.GetsockoptTCPInfo(fd, unix.IPPROTO_TCP, unix.TCP_INFO) + if err != nil { + return false, err + } + const ( + tcpFinWait2 = 5 + tcpTimeWait = 6 + tcpClose = 7 + ) + switch info.State { + case tcpFinWait2, tcpTimeWait, tcpClose: + return true, nil + default: + return false, nil + } +} diff --git a/server/lib/scaletozero/socket_other.go b/server/lib/scaletozero/socket_other.go new file mode 100644 index 000000000..0ceeaa0a3 --- /dev/null +++ b/server/lib/scaletozero/socket_other.go @@ -0,0 +1,32 @@ +//go:build !linux + +package scaletozero + +import ( + "errors" + "time" +) + +func setTCPUserTimeout(int, time.Duration) error { + return errors.New("TCP_USER_TIMEOUT is unavailable") +} + +func outboundQueueFD(int) (int, error) { + return 0, errors.New("TCP outbound queue inspection is unavailable") +} + +func abortSocket(int) error { + return errors.New("abortive socket close is unavailable") +} + +func closeSocket(int) error { + return errors.New("socket close is unavailable") +} + +func terminateGuest() { + panic("guest termination is unavailable") +} + +func closeAcknowledged(int) (bool, error) { + return false, errors.New("TCP close acknowledgement inspection is unavailable") +}