diff --git a/server/cmd/api/main.go b/server/cmd/api/main.go index b3384c6f5..690c0631b 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,13 @@ func main() { metrics.NewChromeCollector(upstreamMgr), metrics.NewGPUCollector(), metrics.NewSystemCollector(), + metrics.NewResponseDrainCollector( + scaletozero.ResponseDrainOutcomeCounts, + scaletozero.ActiveResponseHolds, + scaletozero.FailClosedResponseHolds, + scaletozero.ActiveResponseCloseMonitors, + scaletozero.ResponseCloseMonitorRejections, + ), } if otlpMetrics != nil { metricsCollectors = append(metricsCollectors, metrics.NewOTLPCollector(otlpMetrics)) @@ -387,29 +395,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..f56ccb597 --- /dev/null +++ b/server/lib/metrics/scaletozero.go @@ -0,0 +1,56 @@ +package metrics + +import ( + "context" + "sort" +) + +type ResponseDrainSource func() map[string]uint64 +type ResponseDrainGauge func() int64 + +type ResponseDrainCollector struct { + snapshot ResponseDrainSource + active ResponseDrainGauge + failClosed ResponseDrainGauge + closeMonitors ResponseDrainGauge + monitorRejections func() uint64 +} + +func NewResponseDrainCollector( + snapshot ResponseDrainSource, + active, failClosed, closeMonitors ResponseDrainGauge, + monitorRejections func() uint64, +) *ResponseDrainCollector { + return &ResponseDrainCollector{ + snapshot: snapshot, + active: active, + failClosed: failClosed, + closeMonitors: closeMonitors, + monitorRejections: monitorRejections, + } +} + +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())) + w.Metric("kernel_scale_to_zero_response_close_monitors", "Duplicated sockets currently monitored for acknowledged response close.", "gauge") + w.Sample("kernel_scale_to_zero_response_close_monitors", nil, float64(c.closeMonitors())) + w.Metric("kernel_scale_to_zero_response_close_monitor_rejections_total", "Response close monitor requests rejected at the concurrency limit.", "counter") + w.Sample("kernel_scale_to_zero_response_close_monitor_rejections_total", nil, float64(c.monitorRejections())) + return nil +} diff --git a/server/lib/metrics/scaletozero_test.go b/server/lib/metrics/scaletozero_test.go new file mode 100644 index 000000000..0128e02d4 --- /dev/null +++ b/server/lib/metrics/scaletozero_test.go @@ -0,0 +1,39 @@ +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 }, func() int64 { return 2 }, func() uint64 { return 4 }) + 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 +# HELP kernel_scale_to_zero_response_close_monitors Duplicated sockets currently monitored for acknowledged response close. +# TYPE kernel_scale_to_zero_response_close_monitors gauge +kernel_scale_to_zero_response_close_monitors 2 +# HELP kernel_scale_to_zero_response_close_monitor_rejections_total Response close monitor requests rejected at the concurrency limit. +# TYPE kernel_scale_to_zero_response_close_monitor_rejections_total counter +kernel_scale_to_zero_response_close_monitor_rejections_total 4 +`, string(writer.Bytes())) +} diff --git a/server/lib/scaletozero/connection.go b/server/lib/scaletozero/connection.go new file mode 100644 index 000000000..2afe6a2be --- /dev/null +++ b/server/lib/scaletozero/connection.go @@ -0,0 +1,662 @@ +package scaletozero + +import ( + "context" + "errors" + "io" + "log/slog" + "net" + "net/http" + "sync" + "time" + + "golang.org/x/sys/unix" +) + +const responseWriteChunkSize = 1 << 20 + +var errResponseMonitorLimit = errors.New("response close monitor limit reached") + +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 +} + +// 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 +} + +func (c *drainConn) Write(p []byte) (int, error) { + total := 0 + for len(p) > 0 { + timeout, setDeadline, abort, hijacked, err := c.writeConfig() + if err != nil { + return total, err + } + if hijacked { + n, err := c.TCPConn.Write(p) + return total + n, err + } + if timeout > 0 { + if err := setDeadline(c.TCPConn, time.Now().Add(timeout)); err != nil { + c.abortNowWithOutcome(abort, responseDrainWriteError, err) + return total, err + } + } + + chunk := p + if len(chunk) > responseWriteChunkSize { + chunk = chunk[:responseWriteChunkSize] + } + n, err := c.TCPConn.Write(chunk) + total += n + if err == nil && n != len(chunk) { + err = io.ErrShortWrite + } + if err != nil { + c.abortNowWithOutcome(abort, responseDrainWriteError, err) + return total, err + } + p = p[n:] + } + return total, nil +} + +func (c *drainConn) ReadFrom(r io.Reader) (int64, error) { + var total int64 + for { + timeout, setDeadline, abort, hijacked, err := c.writeConfig() + if err != nil { + return total, err + } + if hijacked { + n, err := c.TCPConn.ReadFrom(r) + return total + n, err + } + if timeout > 0 { + if err := setDeadline(c.TCPConn, time.Now().Add(timeout)); err != nil { + c.abortNowWithOutcome(abort, responseDrainWriteError, err) + return total, err + } + } + + limited, parents, exhausted := nextReadFromChunk(r) + if exhausted { + return total, nil + } + n, err := c.TCPConn.ReadFrom(limited) + total += n + for _, parent := range parents { + parent.N -= n + } + if errors.Is(err, io.EOF) { + return total, nil + } + if err != nil { + if isReadFromWriteError(limited.R, err) { + c.abortNowWithOutcome(abort, responseDrainWriteError, err) + } else { + c.markCurrentFailure(responseDrainSourceReadError, err) + } + return total, err + } + if limited.N > 0 || anyLimitExhausted(parents) { + return total, nil + } + } +} + +func (c *drainConn) writeConfig() (time.Duration, func(*net.TCPConn, time.Time) error, func(*net.TCPConn) error, bool, error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed || c.closing { + return 0, nil, nil, false, net.ErrClosed + } + if c.hijackedConn { + return 0, nil, nil, true, nil + } + return c.writeTimeout, c.setDeadline, c.abortConnection, false, nil +} + +func nextReadFromChunk(r io.Reader) (*io.LimitedReader, []*io.LimitedReader, bool) { + limit := int64(responseWriteChunkSize) + parents := make([]*io.LimitedReader, 0, 1) + for { + parent, ok := r.(*io.LimitedReader) + if !ok { + break + } + if parent.N <= 0 { + return nil, parents, true + } + parents = append(parents, parent) + limit = min(limit, parent.N) + r = parent.R + } + return &io.LimitedReader{R: r, N: limit}, parents, false +} + +func anyLimitExhausted(parents []*io.LimitedReader) bool { + for _, parent := range parents { + if parent.N == 0 { + return true + } + } + return false +} + +func isReadFromWriteError(source io.Reader, err error) bool { + if _, ambiguous := source.(net.Conn); ambiguous { + return false + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + return errors.Is(err, net.ErrClosed) || + errors.Is(err, unix.EPIPE) || + errors.Is(err, unix.ECONNRESET) || + errors.Is(err, unix.ECONNABORTED) || + errors.Is(err, unix.ETIMEDOUT) || + errors.Is(err, unix.ENETDOWN) || + errors.Is(err, unix.ENETRESET) || + errors.Is(err, unix.ENETUNREACH) || + errors.Is(err, unix.EHOSTDOWN) || + errors.Is(err, unix.EHOSTUNREACH) +} + +func (c *drainConn) markCurrentFailure(outcome responseDrainOutcome, err error) { + c.mu.Lock() + var drain *requestDrain + if len(c.pending) > 0 { + drain = c.pending[len(c.pending)-1] + } + c.mu.Unlock() + if drain != nil { + drain.fail(outcome, err) + } +} + +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) { + c.abortNowWithOutcome(abort, responseDrainConnectionClosed, nil) +} + +func (c *drainConn) abortNowWithOutcome(abort func(*net.TCPConn) error, outcome responseDrainOutcome, cause 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 len(drains) == 0 { + if err != nil && !isTerminalConnectionError(err) { + _ = c.TCPConn.Close() + } + return + } + if err == nil || isTerminalConnectionError(err) { + completeDrains(drains, outcome, 0, cause) + return + } + go c.recoverResponseConnection(drains, config, outcome, 0, errors.Join(cause, 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() + } + if !config.acquireMonitor() { + responseCloseMonitorRejections.Add(1) + abortErr := config.abort(c.TCPConn) + if abortErr == nil || isTerminalConnectionError(abortErr) { + completeDrains(drains, responseDrainMonitorLimit, 0, errResponseMonitorLimit) + return nil + } + go c.recoverResponseConnection(drains, config, responseDrainMonitorLimit, 0, errors.Join(errResponseMonitorLimit, abortErr)) + return abortErr + } + + fd, err := config.duplicate(c.TCPConn) + if err == nil { + closeErr := c.TCPConn.Close() + go monitorClosedResponse(fd, drains, config, false, nil, "") + return closeErr + } + config.releaseMonitor() + 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 := -1 + duplicateErr := errResponseMonitorLimit + monitorAcquired := config.acquireMonitor() + if monitorAcquired { + fd, duplicateErr = config.duplicate(c.TCPConn) + } else { + responseCloseMonitorRejections.Add(1) + } + if duplicateErr == nil { + closeErr := c.TCPConn.Close() + c.terminalMu.Unlock() + go monitorClosedResponse(fd, drains, config, true, errors.Join(cause, abortErr, closeErr), outcome) + return + } + if monitorAcquired { + config.releaseMonitor() + } + if isTerminalConnectionError(duplicateErr) { + _ = c.TCPConn.Close() + c.terminalMu.Unlock() + if failedCount > 0 { + failClosedResponseHolds.Add(-int64(failedCount)) + } + completeDrains(drains, responseDrainConnectionClosed, queued, cause) + return + } + c.terminalMu.Unlock() + + 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 isSafeClosedResponseOutcome(outcome responseDrainOutcome) bool { + return outcome == responseDrainCloseAcknowledged || outcome == responseDrainConnectionClosed +} + +func monitorClosedResponse(fd int, drains []*requestDrain, config responseDrainConfig, failed bool, cause error, failureOutcome responseDrainOutcome) { + activeResponseCloseMonitors.Add(1) + defer activeResponseCloseMonitors.Add(-1) + defer config.releaseMonitor() + + userTimeoutErr := config.setUserTimeout(fd, config.terminalRecoveryTimeout) + outcome, queued, err := waitForClosedResponse(fd, config, time.Now().Add(config.timeout)) + cause = errors.Join(cause, userTimeoutErr, err) + if isSafeClosedResponseOutcome(outcome) { + _ = config.closeFD(fd) + if failed { + failClosedResponseHolds.Add(-int64(len(drains))) + } + if failureOutcome != "" { + outcome = failureOutcome + } + 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))) + if failureOutcome != "" { + outcome = failureOutcome + } + completeDrains(drains, outcome, queued, cause) + return + } + + currentOutcome, currentQueued, inspectErr := waitForClosedResponse(fd, config, time.Now()) + if isSafeClosedResponseOutcome(currentOutcome) || isTerminalConnectionError(inspectErr) { + _ = config.closeFD(fd) + failClosedResponseHolds.Add(-int64(len(drains))) + if !isSafeClosedResponseOutcome(currentOutcome) { + currentOutcome = responseDrainConnectionClosed + } + if failureOutcome != "" { + currentOutcome = failureOutcome + } + completeDrains(drains, currentOutcome, 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 + } + state, err := config.closeState(fd) + if err != nil { + return responseDrainIOError, queued, err + } + switch { + case state == responseCloseTerminated: + return responseDrainConnectionClosed, queued, nil + case queued == 0 && state == responseCloseAcknowledged: + 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..43b6c6c1e --- /dev/null +++ b/server/lib/scaletozero/connection_test.go @@ -0,0 +1,843 @@ +package scaletozero + +import ( + "bufio" + "bytes" + "context" + "errors" + "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" + "golang.org/x/sys/unix" +) + +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 TestUntrackedWriteFailuresDoNotStartResponseRecovery(t *testing.T) { + tests := []struct { + name string + hijack bool + wantAborts int32 + }{ + {name: "pre-registration", wantAborts: 1}, + {name: "hijacked", hijack: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + tracked, peer := newTCPPair(t) + defer tracked.TCPConn.Close() + + var aborts atomic.Int32 + var duplicates atomic.Int32 + terminated := make(chan struct{}, 1) + config := testDrainConfig(time.Second, outboundQueue) + config.abort = func(*net.TCPConn) error { + aborts.Add(1) + return assert.AnError + } + config.duplicate = func(*net.TCPConn) (int, error) { + duplicates.Add(1) + return -1, assert.AnError + } + config.abortRetryInterval = time.Millisecond + config.terminalRecoveryTimeout = 5 * time.Millisecond + config.terminateGuest = func() { + terminated <- struct{}{} + runtime.Goexit() + } + tracked.configure(config) + if tc.hijack { + tracked.hijack() + } + + require.NoError(t, peer.(*net.TCPConn).SetLinger(0)) + require.NoError(t, peer.Close()) + require.Eventually(t, func() bool { + _, err := tracked.Write([]byte("response")) + return err != nil + }, time.Second, time.Millisecond) + + assert.Equal(t, tc.wantAborts, aborts.Load()) + assert.Zero(t, duplicates.Load()) + select { + case <-terminated: + t.Fatal("untracked write failure terminated the guest") + case <-time.After(20 * time.Millisecond): + } + }) + } +} + +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 TestResetClosedResponseCompletesWithQueuedBytes(t *testing.T) { + if runtime.GOOS != "linux" { + t.Skip("TCP output queue inspection requires Linux") + } + tracked, clientConn := newTCPPair(t) + client := clientConn.(*net.TCPConn) + defer tracked.TCPConn.Close() + require.NoError(t, tracked.SetWriteBuffer(64<<10)) + require.NoError(t, client.SetReadBuffer(4<<10)) + raw, err := tracked.SyscallConn() + require.NoError(t, err) + payload := make([]byte, 32<<20) + written := 0 + var writeErr error + require.NoError(t, raw.Write(func(fd uintptr) bool { + for written < len(payload) { + n, err := unix.Write(int(fd), payload[written:]) + written += n + if errors.Is(err, unix.EAGAIN) { + return true + } + if err != nil { + writeErr = err + return true + } + } + return true + })) + require.NoError(t, writeErr) + require.Positive(t, written) + require.Less(t, written, len(payload)) + queuedBeforeReset, err := outboundQueue(tracked.TCPConn) + require.NoError(t, err) + require.Positive(t, queuedBeforeReset) + + fd, err := duplicateAndShutdownWrite(tracked.TCPConn) + require.NoError(t, err) + defer closeSocket(fd) + require.NoError(t, tracked.TCPConn.Close()) + require.NoError(t, client.SetLinger(0)) + require.NoError(t, client.Close()) + + config := testDrainConfig(time.Second, outboundQueue) + started := time.Now() + outcome, _, err := waitForClosedResponse(fd, config, time.Now().Add(config.timeout)) + require.NoError(t, err) + assert.Equal(t, responseDrainConnectionClosed, outcome) + assert.Less(t, time.Since(started), 250*time.Millisecond) +} + +func TestTerminatedClosedResponseIgnoresQueuedBytes(t *testing.T) { + config := testDrainConfig(time.Second, outboundQueue) + config.outboundFD = func(int) (int, error) { return 123, nil } + config.closeState = func(int) (responseCloseState, error) { return responseCloseTerminated, nil } + + outcome, queued, err := waitForClosedResponse(-1, config, time.Now().Add(config.timeout)) + require.NoError(t, err) + assert.Equal(t, responseDrainConnectionClosed, outcome) + assert.Equal(t, 123, queued) +} + +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 TestTerminalDuplicationFailureClosesOriginalConnection(t *testing.T) { + tracked, peer := newTCPPair(t) + defer peer.Close() + + activeBefore := ActiveResponseHolds() + failedBefore := FailClosedResponseHolds() + released := make(chan struct{}) + outcomes := make(chan responseDrainOutcome, 1) + hold := newResponseHold(func() { close(released) }) + hold.finishHandler() + config := testDrainConfig(time.Second, outboundQueue) + config.abort = func(*net.TCPConn) error { return assert.AnError } + config.duplicate = func(*net.TCPConn) (int, error) { return -1, unix.ENOTCONN } + var monitors atomic.Int32 + config.acquireMonitor = func() bool { + monitors.Add(1) + return true + } + config.releaseMonitor = func() { monitors.Add(-1) } + config.onComplete = func(outcome responseDrainOutcome) { outcomes <- outcome } + config.terminateGuest = func() { t.Fatal("terminal duplication failure terminated the guest") } + drains := []*requestDrain{{hold: hold, config: config, log: slog.New(slog.NewTextHandler(io.Discard, nil))}} + + tracked.recoverResponseConnection(drains, config, responseDrainIOError, 1, assert.AnError) + + assert.Equal(t, int32(0), monitors.Load()) + assert.Equal(t, responseDrainConnectionClosed, <-outcomes) + select { + case <-released: + default: + t.Fatal("response hold was not released") + } + _, err := tracked.TCPConn.Write([]byte("closed")) + assert.ErrorIs(t, err, net.ErrClosed) + 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.closeState = func(int) (responseCloseState, error) { return responseClosePending, 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()}} + + require.True(t, config.acquireMonitor()) + 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 TestWritesRefreshDeadlineDuringProgress(t *testing.T) { + const bodySize = 32 << 20 + const timeout = 200 * time.Millisecond + tests := []struct { + name string + transfer func(*testing.T, *drainConn) (int64, error) + }{ + { + name: "write", + transfer: func(_ *testing.T, conn *drainConn) (int64, error) { + n, err := conn.Write(make([]byte, bodySize)) + return int64(n), err + }, + }, + { + name: "read_from", + transfer: func(t *testing.T, conn *drainConn) (int64, error) { + file, err := os.Create(filepath.Join(t.TempDir(), "response.bin")) + require.NoError(t, err) + defer file.Close() + require.NoError(t, file.Truncate(bodySize)) + return conn.ReadFrom(file) + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + require.NoError(t, tracked.SetWriteBuffer(64<<10)) + + 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 := tc.transfer(t, tracked) + 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/responseWriteChunkSize)) + }) + } +} + +func TestReadFromFlattensExistingLimits(t *testing.T) { + source := bytes.NewReader(make([]byte, 2<<20)) + inner := &io.LimitedReader{R: source, N: 2 << 20} + outer := &io.LimitedReader{R: inner, N: 3 << 20} + chunk, parents, exhausted := nextReadFromChunk(outer) + + assert.False(t, exhausted) + assert.Same(t, source, chunk.R) + assert.Equal(t, int64(responseWriteChunkSize), chunk.N) + require.Len(t, parents, 2) + assert.Same(t, outer, parents[0]) + assert.Same(t, inner, parents[1]) +} + +func TestReadFromPreservesExistingLimit(t *testing.T) { + const bodySize = 4 << 20 + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + tracked.configure(testDrainConfig(time.Second, outboundQueue)) + file, err := os.Create(filepath.Join(t.TempDir(), "response.bin")) + require.NoError(t, err) + defer file.Close() + require.NoError(t, file.Truncate(bodySize)) + limited := &io.LimitedReader{R: file, N: bodySize} + readDone := make(chan error, 1) + go func() { + _, err := io.CopyN(io.Discard, client, bodySize) + readDone <- err + }() + + written, err := tracked.ReadFrom(limited) + require.NoError(t, err) + assert.Equal(t, int64(bodySize), written) + assert.Zero(t, limited.N) + require.NoError(t, <-readDone) +} + +type failingSourceReader struct { + remaining int + err error +} + +func (r *failingSourceReader) Read(p []byte) (int, error) { + if r.remaining == 0 { + return 0, r.err + } + if len(p) > r.remaining { + p = p[:r.remaining] + } + r.remaining -= len(p) + return len(p), nil +} + +func TestReadFromSourceErrorClosesGracefully(t *testing.T) { + const bodySize = 2 << 20 + sourceErr := errors.New("source failed") + activeBefore := ActiveResponseHolds() + tracked, client := newTCPPair(t) + defer client.Close() + defer tracked.TCPConn.Close() + + outcomes := make(chan responseDrainOutcome, 1) + config := testDrainConfig(time.Second, outboundQueue) + config.onComplete = func(outcome responseDrainOutcome) { outcomes <- outcome } + var aborts atomic.Int32 + config.abort = func(conn *net.TCPConn) error { + aborts.Add(1) + return abortConnection(conn) + } + tracked.configure(config) + hold := newResponseHold(func() {}) + hold.finishHandler() + log := slog.New(slog.NewTextHandler(io.Discard, nil)) + require.True(t, tracked.addDrain(&requestDrain{hold: hold, config: config, log: log})) + + readDone := make(chan struct { + body []byte + err error + }, 1) + go func() { + body, err := io.ReadAll(client) + readDone <- struct { + body []byte + err error + }{body: body, err: err} + }() + + written, err := tracked.ReadFrom(&failingSourceReader{remaining: bodySize, err: sourceErr}) + require.ErrorIs(t, err, sourceErr) + assert.Equal(t, int64(bodySize), written) + require.NoError(t, tracked.Close()) + result := <-readDone + require.NoError(t, result.err) + assert.Len(t, result.body, bodySize) + assert.Zero(t, aborts.Load()) + assert.Equal(t, responseDrainSourceReadError, <-outcomes) + assert.Eventually(t, func() bool { return ActiveResponseHolds() == activeBefore }, time.Second, time.Millisecond) +} + +func TestReadFromErrorClassification(t *testing.T) { + sendfileErr := &net.OpError{Op: "readfrom", Net: "tcp", Err: &os.SyscallError{Syscall: "sendfile", Err: unix.EIO}} + assert.False(t, isReadFromWriteError(nil, sendfileErr)) + assert.True(t, isReadFromWriteError(bytes.NewReader(nil), os.ErrDeadlineExceeded)) + + source, peer := net.Pipe() + defer source.Close() + defer peer.Close() + assert.False(t, isReadFromWriteError(source, os.ErrDeadlineExceeded)) +} + +func TestCloseMonitorLimitBoundsResources(t *testing.T) { + const monitorLimit = 4 + const connectionCount = 12 + activeHoldsBefore := ActiveResponseHolds() + activeMonitorsBefore := ActiveResponseCloseMonitors() + rejectionsBefore := ResponseCloseMonitorRejections() + slots := make(chan struct{}, monitorLimit) + config := testDrainConfig(100*time.Millisecond, outboundQueue) + config.acquireMonitor = func() bool { + select { + case slots <- struct{}{}: + return true + default: + return false + } + } + config.releaseMonitor = func() { <-slots } + config.setUserTimeout = func(int, time.Duration) error { return nil } + config.outboundFD = func(int) (int, error) { return 1, nil } + config.closeState = func(int) (responseCloseState, error) { return responseClosePending, nil } + outcomes := make(chan responseDrainOutcome, connectionCount) + config.onComplete = func(outcome responseDrainOutcome) { outcomes <- outcome } + var aborts atomic.Int32 + config.abort = func(conn *net.TCPConn) error { + aborts.Add(1) + return abortConnection(conn) + } + + type connectionPair struct { + tracked *drainConn + client net.Conn + } + pairs := make([]connectionPair, 0, connectionCount) + log := slog.New(slog.NewTextHandler(io.Discard, nil)) + for range connectionCount { + tracked, client := newTCPPair(t) + hold := newResponseHold(func() {}) + hold.finishHandler() + require.True(t, tracked.addDrain(&requestDrain{hold: hold, config: config, log: log})) + pairs = append(pairs, connectionPair{tracked: tracked, client: client}) + } + for _, pair := range pairs { + require.NoError(t, pair.tracked.Close()) + } + + require.Eventually(t, func() bool { + return ActiveResponseCloseMonitors() == activeMonitorsBefore+monitorLimit + }, time.Second, time.Millisecond) + assert.Equal(t, int32(connectionCount-monitorLimit), aborts.Load()) + assert.Equal(t, rejectionsBefore+connectionCount-monitorLimit, ResponseCloseMonitorRejections()) + assert.Equal(t, activeHoldsBefore+monitorLimit, ActiveResponseHolds()) + assert.Len(t, slots, monitorLimit) + + for _, pair := range pairs { + require.NoError(t, pair.client.Close()) + } + require.Eventually(t, func() bool { + return ActiveResponseCloseMonitors() == activeMonitorsBefore && ActiveResponseHolds() == activeHoldsBefore + }, time.Second, time.Millisecond) + assert.Empty(t, slots) + counts := make(map[responseDrainOutcome]int) + for range connectionCount { + counts[<-outcomes]++ + } + assert.Equal(t, connectionCount-monitorLimit, counts[responseDrainMonitorLimit]) + assert.Equal(t, monitorLimit, counts[responseDrainTimeoutHit]) +} + +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..37a9d8f15 100644 --- a/server/lib/scaletozero/middleware.go +++ b/server/lib/scaletozero/middleware.go @@ -2,17 +2,253 @@ 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 + maxResponseCloseMonitors = 256 +) + +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) + closeState func(int) (responseCloseState, error) + acquireMonitor func() bool + releaseMonitor func() + terminateGuest func() + abortRetryInterval time.Duration + terminalRecoveryTimeout time.Duration + onComplete func(responseDrainOutcome) +} + +type responseDrainOutcome string + +type responseCloseState uint8 + +const ( + responseClosePending responseCloseState = iota + responseCloseAcknowledged + responseCloseTerminated +) + +const ( + responseDrainComplete responseDrainOutcome = "drained" + responseDrainCloseAcknowledged responseDrainOutcome = "close_acknowledged" + responseDrainTimeoutHit responseDrainOutcome = "timeout" + responseDrainIOError responseDrainOutcome = "ioctl_error" + responseDrainNonTCP responseDrainOutcome = "untracked_connection" + responseDrainWriteError responseDrainOutcome = "write_error" + responseDrainSourceReadError responseDrainOutcome = "source_read_error" + responseDrainDeadlineClearError responseDrainOutcome = "deadline_clear_error" + responseDrainConnectionClosed responseDrainOutcome = "connection_closed" + responseDrainConnectionHijacked responseDrainOutcome = "connection_hijacked" + responseDrainConnectionReused responseDrainOutcome = "connection_reused" + responseDrainAbortError responseDrainOutcome = "abort_error" + responseDrainMonitorLimit responseDrainOutcome = "monitor_limit" + responseDrainGuestTermination responseDrainOutcome = "guest_termination" +) + +var responseDrainCounters = map[responseDrainOutcome]*atomic.Uint64{ + responseDrainComplete: {}, + responseDrainCloseAcknowledged: {}, + responseDrainTimeoutHit: {}, + responseDrainIOError: {}, + responseDrainNonTCP: {}, + responseDrainWriteError: {}, + responseDrainSourceReadError: {}, + responseDrainDeadlineClearError: {}, + responseDrainConnectionClosed: {}, + responseDrainConnectionHijacked: {}, + responseDrainConnectionReused: {}, + responseDrainAbortError: {}, + responseDrainMonitorLimit: {}, + responseDrainGuestTermination: {}, +} + +var activeResponseHolds atomic.Int64 +var failClosedResponseHolds atomic.Int64 +var activeResponseCloseMonitors atomic.Int64 +var responseCloseMonitorRejections atomic.Uint64 + +var responseCloseMonitorSlots = make(chan struct{}, maxResponseCloseMonitors) + +// 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 ActiveResponseCloseMonitors() int64 { return activeResponseCloseMonitors.Load() } + +func ResponseCloseMonitorRejections() uint64 { return responseCloseMonitorRejections.Load() } + +func acquireResponseCloseMonitor() bool { + select { + case responseCloseMonitorSlots <- struct{}{}: + return true + default: + return false + } +} + +func releaseResponseCloseMonitor() { <-responseCloseMonitorSlots } + +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 + mu sync.Mutex + failureOutcome responseDrainOutcome + failureErr error +} + +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) fail(outcome responseDrainOutcome, err error) { + d.mu.Lock() + defer d.mu.Unlock() + if d.failureOutcome == "" { + d.failureOutcome = outcome + d.failureErr = err + } +} + +func (d *requestDrain) complete(outcome responseDrainOutcome, queued int, err error) { + d.mu.Lock() + failureOutcome := d.failureOutcome + failureErr := d.failureErr + d.mu.Unlock() + terminalOutcome := outcome + if failureOutcome != "" { + outcome = failureOutcome + err = errors.Join(failureErr, err) + } + recordResponseDrainOutcome(outcome) + if d.config.onComplete != nil { + d.config.onComplete(outcome) + } + attrs := []any{"outcome", outcome} + if failureOutcome != "" { + attrs = append(attrs, "terminal_outcome", terminalOutcome) + } + 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, + closeState: inspectResponseClose, + acquireMonitor: acquireResponseCloseMonitor, + releaseMonitor: releaseResponseCloseMonitor, + 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 +256,111 @@ 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())) 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..499cc2675 100644 --- a/server/lib/scaletozero/middleware_test.go +++ b/server/lib/scaletozero/middleware_test.go @@ -1,45 +1,58 @@ package scaletozero import ( + "bufio" + "bytes" + "context" + "fmt" + "io" + "net" "net/http" "net/http/httptest" + "sync" + "sync/atomic" "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 +62,541 @@ func TestMiddlewareSkipsLoopbackAddrs(t *testing.T) { })) req := httptest.NewRequest(http.MethodGet, "/", nil) - req.RemoteAddr = tc.addr - rec := httptest.NewRecorder() - - handler.ServeHTTP(rec, req) + req.RemoteAddr = addr + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, 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() + assert.PanicsWithValue(t, http.ErrAbortHandler, func() { + handler.ServeHTTP(httptest.NewRecorder(), externalRequest(http.MethodGet, "/", tracked)) + }) - handler.ServeHTTP(rec, req) - - 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 + }) + outcome := make(chan responseDrainOutcome, 1) + drain.onComplete = func(value responseDrainOutcome) { outcome <- value } + 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) + assert.Equal(t, responseDrainIOError, <-outcome) +} + +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 TestMiddlewareWriterToResponseRefreshesDeadline(t *testing.T) { + const bodySize = 32 << 20 + const timeout = 200 * time.Millisecond + var deadlines atomic.Int32 + config := testDrainConfig(timeout, outboundQueue) + config.setDeadline = func(conn *net.TCPConn, deadline time.Time) error { + if !deadline.IsZero() { + deadlines.Add(1) + } + return conn.SetWriteDeadline(deadline) + } + handlerDone := make(chan struct { + written int64 + elapsed time.Duration + err error + }, 1) + handler := middleware(NewNoopController(), config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Length", fmt.Sprint(bodySize)) + body := io.NopCloser(bytes.NewReader(make([]byte, bodySize))) + defer body.Close() + started := time.Now() + written, err := io.Copy(w, body) + handlerDone <- struct { + written int64 + elapsed time.Duration + err error + }{written: written, elapsed: time.Since(started), 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(64<<10)) + _, err = fmt.Fprint(conn, "GET /large 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) + buffer := make([]byte, 64<<10) + var received int64 + for { + n, readErr := response.Body.Read(buffer) + received += int64(n) + if readErr == io.EOF { + break + } + require.NoError(t, readErr) + time.Sleep(time.Millisecond) + } + require.NoError(t, response.Body.Close()) + + result := <-handlerDone + require.NoError(t, result.err) + assert.Equal(t, int64(bodySize), result.written) + assert.Equal(t, int64(bodySize), received) + assert.Greater(t, result.elapsed, timeout) + assert.GreaterOrEqual(t, deadlines.Load(), int32(bodySize/responseWriteChunkSize)) +} + +func TestWriteFailuresProduceOneWriteErrorOutcome(t *testing.T) { + tests := []struct { + name string + responseLen int + wantWriteErr bool + }{ + {name: "finish_request_flush", responseLen: 32, wantWriteErr: false}, + {name: "handler_write", responseLen: 4 << 10, wantWriteErr: true}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + config := testDrainConfig(time.Second, outboundQueue) + config.setDeadline = func(*net.TCPConn, time.Time) error { return assert.AnError } + outcomes := make(chan responseDrainOutcome, 2) + config.onComplete = func(outcome responseDrainOutcome) { outcomes <- outcome } + handlerWrite := make(chan error, 1) + handler := middleware(NewNoopController(), config)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, err := w.Write(make([]byte, tc.responseLen)) + handlerWrite <- err + })) + + listener, serverDone, server := serveTestHandler(t, handler) + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + _, err = fmt.Fprint(conn, "GET / HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n") + require.NoError(t, err) + if tc.wantWriteErr { + require.Error(t, <-handlerWrite) + } else { + require.NoError(t, <-handlerWrite) + } + select { + case outcome := <-outcomes: + assert.Equal(t, responseDrainWriteError, outcome) + case <-time.After(time.Second): + t.Fatal("write failure did not complete the response hold") + } + select { + case outcome := <-outcomes: + t.Fatalf("write failure produced a second outcome: %s", outcome) + case <-time.After(20 * time.Millisecond): + } + require.NoError(t, conn.Close()) + require.NoError(t, server.Close()) + assert.ErrorIs(t, <-serverDone, http.ErrServerClosed) + }) + } +} + +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 +604,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 +619,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..42b023baf --- /dev/null +++ b/server/lib/scaletozero/socket_linux.go @@ -0,0 +1,73 @@ +//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 { + return abortSocketWith( + func() error { + return unix.SetsockoptLinger(fd, unix.SOL_SOCKET, unix.SO_LINGER, &unix.Linger{Onoff: 1}) + }, + func() error { return unix.Close(fd) }, + ) +} + +func abortSocketWith(setLinger, closeFD func() error) error { + if err := setLinger(); err != nil { + return err + } + // Linux releases the descriptor even when close returns an error such as + // EINTR. Never let recovery operate on a descriptor number that may be reused. + _ = closeFD() + return nil +} + +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 inspectResponseClose(fd int) (responseCloseState, error) { + info, err := unix.GetsockoptTCPInfo(fd, unix.IPPROTO_TCP, unix.TCP_INFO) + if err != nil { + return responseClosePending, err + } + const ( + tcpFinWait2 = 5 + tcpTimeWait = 6 + tcpClose = 7 + ) + switch info.State { + case tcpFinWait2, tcpTimeWait: + return responseCloseAcknowledged, nil + case tcpClose: + return responseCloseTerminated, nil + default: + return responseClosePending, nil + } +} diff --git a/server/lib/scaletozero/socket_linux_test.go b/server/lib/scaletozero/socket_linux_test.go new file mode 100644 index 000000000..2853970b4 --- /dev/null +++ b/server/lib/scaletozero/socket_linux_test.go @@ -0,0 +1,54 @@ +//go:build linux + +package scaletozero + +import ( + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sys/unix" +) + +func TestAbortSocketDoesNotReuseDescriptorAfterCloseError(t *testing.T) { + var lingerCalls atomic.Int32 + var closeCalls atomic.Int32 + var inspectCalls atomic.Int32 + var abortCalls atomic.Int32 + var releaseCalls atomic.Int32 + config := testDrainConfig(0, nil) + config.setUserTimeout = func(int, time.Duration) error { return nil } + config.outboundFD = func(int) (int, error) { + inspectCalls.Add(1) + return 1, nil + } + config.closeState = func(int) (responseCloseState, error) { + inspectCalls.Add(1) + return responseClosePending, nil + } + config.abortFD = func(int) error { + abortCalls.Add(1) + return abortSocketWith( + func() error { + lingerCalls.Add(1) + return nil + }, + func() error { + closeCalls.Add(1) + return unix.EINTR + }, + ) + } + config.releaseMonitor = func() { releaseCalls.Add(1) } + config.terminateGuest = func() { t.Fatal("close EINTR triggered guest termination") } + + monitorClosedResponse(-1, nil, config, false, nil, "") + + assert.Equal(t, int32(1), abortCalls.Load()) + assert.Equal(t, int32(1), lingerCalls.Load()) + assert.Equal(t, int32(1), closeCalls.Load()) + assert.Equal(t, int32(2), inspectCalls.Load()) + require.Equal(t, int32(1), releaseCalls.Load()) +} diff --git a/server/lib/scaletozero/socket_other.go b/server/lib/scaletozero/socket_other.go new file mode 100644 index 000000000..168c0e477 --- /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 inspectResponseClose(int) (responseCloseState, error) { + return responseClosePending, errors.New("TCP close acknowledgement inspection is unavailable") +}