From c769c5e7284bd8bb7da752a035d474d8f1472c93 Mon Sep 17 00:00:00 2001 From: Mike Bannister Date: Tue, 15 Sep 2026 14:50:48 -0400 Subject: [PATCH 1/2] Close transient VZ control connections Disable HTTP keep-alives for one-shot VZ control clients and honor dial cancellation. Bound idle control sessions in the shim, evict the pooled guest connection on terminal stop, and cover successful, error, cancellation, and repeated client paths with Unix ConnState regressions. --- cmd/vz-shim/main.go | 7 +- lib/hypervisor/vz/client.go | 10 +- lib/hypervisor/vz/client_test.go | 198 +++++++++++++++++++++++++++++++ lib/instances/stop.go | 8 +- 4 files changed, 219 insertions(+), 4 deletions(-) create mode 100644 lib/hypervisor/vz/client_test.go diff --git a/cmd/vz-shim/main.go b/cmd/vz-shim/main.go index 127a9a819..5c9f72807 100644 --- a/cmd/vz-shim/main.go +++ b/cmd/vz-shim/main.go @@ -22,6 +22,8 @@ import ( "github.com/kernel/hypeman/lib/hypervisor/vz/shimconfig" ) +const controlIdleTimeout = 30 * time.Second + func main() { configJSON := flag.String("config", "", "VM configuration as JSON") flag.Parse() @@ -96,7 +98,10 @@ func main() { defer vsockListener.Close() // Start HTTP server for control API - httpServer := &http.Server{Handler: server.Handler()} + httpServer := &http.Server{ + Handler: server.Handler(), + IdleTimeout: controlIdleTimeout, + } go func() { slog.Info("control API listening", "socket", config.ControlSocket) if err := httpServer.Serve(controlListener); err != nil && err != http.ErrServerClosed { diff --git a/lib/hypervisor/vz/client.go b/lib/hypervisor/vz/client.go index 53e56f4ae..3036014e4 100644 --- a/lib/hypervisor/vz/client.go +++ b/lib/hypervisor/vz/client.go @@ -28,10 +28,16 @@ type Client struct { // NewClient creates a new vz shim client. func NewClient(socketPath string) (*Client, error) { + dialer := &net.Dialer{} transport := &http.Transport{ - DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { - return net.Dial("unix", socketPath) + DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) { + return dialer.DialContext(ctx, "unix", socketPath) }, + // Callers construct VZ clients for short-lived state and lifecycle + // operations. A private keep-alive pool would retain one idle Unix socket + // pair after the client becomes unreachable, so close each connection when + // its response completes instead of pooling it on this one-shot transport. + DisableKeepAlives: true, } httpClient := &http.Client{ Transport: transport, diff --git a/lib/hypervisor/vz/client_test.go b/lib/hypervisor/vz/client_test.go new file mode 100644 index 000000000..c3090540f --- /dev/null +++ b/lib/hypervisor/vz/client_test.go @@ -0,0 +1,198 @@ +//go:build darwin + +package vz + +import ( + "context" + "errors" + "net" + "net/http" + "os" + "path/filepath" + "sync" + "sync/atomic" + "testing" + "time" +) + +const connectionCloseTimeout = 2 * time.Second + +type unixHTTPFixture struct { + listener net.Listener + server *http.Server + + mu sync.Mutex + conns map[net.Conn]http.ConnState +} + +func newUnixHTTPFixture(t *testing.T, handler http.Handler) *unixHTTPFixture { + t.Helper() + + dir, err := os.MkdirTemp("/tmp", "hypeman-vz-test-") + if err != nil { + t.Fatalf("create temporary Unix socket directory: %v", err) + } + t.Cleanup(func() { _ = os.RemoveAll(dir) }) + + listener, err := net.Listen("unix", filepath.Join(dir, "vz.sock")) + if err != nil { + t.Fatalf("listen on temporary Unix socket: %v", err) + } + fixture := &unixHTTPFixture{ + listener: listener, + conns: make(map[net.Conn]http.ConnState), + } + fixture.server = &http.Server{ + Handler: handler, + ConnState: func(conn net.Conn, state http.ConnState) { + fixture.mu.Lock() + defer fixture.mu.Unlock() + if state == http.StateClosed || state == http.StateHijacked { + delete(fixture.conns, conn) + return + } + fixture.conns[conn] = state + }, + } + go func() { + _ = fixture.server.Serve(listener) + }() + t.Cleanup(func() { + _ = fixture.server.Close() + }) + return fixture +} + +func (f *unixHTTPFixture) socketPath() string { + return f.listener.Addr().String() +} + +func (f *unixHTTPFixture) waitForNoConnections(t *testing.T) { + t.Helper() + deadline := time.Now().Add(connectionCloseTimeout) + for { + f.mu.Lock() + count := len(f.conns) + f.mu.Unlock() + if count == 0 { + return + } + if time.Now().After(deadline) { + t.Fatalf("temporary VZ server retained %d connection(s)", count) + } + time.Sleep(time.Millisecond) + } +} + +func TestRepeatedClientsReleaseSuccessfulControlConnections(t *testing.T) { + fixture := newUnixHTTPFixture(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v1/vmm.ping": + _, _ = w.Write([]byte("OK")) + case "/api/v1/vm.info": + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"state":"Running"}`)) + default: + http.NotFound(w, r) + } + })) + + for range 128 { + client, err := NewClient(fixture.socketPath()) + if err != nil { + t.Fatalf("create VZ client: %v", err) + } + if _, err := client.GetVMInfo(context.Background()); err != nil { + t.Fatalf("get VM info: %v", err) + } + fixture.waitForNoConnections(t) + } +} + +func TestControlConnectionsCloseAfterResponseErrors(t *testing.T) { + t.Run("http error", func(t *testing.T) { + fixture := newUnixHTTPFixture(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/v1/vmm.ping" { + _, _ = w.Write([]byte("OK")) + return + } + http.Error(w, "pause failed", http.StatusInternalServerError) + })) + + client, err := NewClient(fixture.socketPath()) + if err != nil { + t.Fatalf("create VZ client: %v", err) + } + if err := client.Pause(context.Background()); err == nil { + t.Fatal("expected HTTP error") + } + fixture.waitForNoConnections(t) + }) + + t.Run("invalid json", func(t *testing.T) { + fixture := newUnixHTTPFixture(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/v1/vmm.ping" { + _, _ = w.Write([]byte("OK")) + return + } + _, _ = w.Write([]byte("not-json")) + })) + + client, err := NewClient(fixture.socketPath()) + if err != nil { + t.Fatalf("create VZ client: %v", err) + } + if _, err := client.GetVMInfo(context.Background()); err == nil { + t.Fatal("expected JSON decoding error") + } + fixture.waitForNoConnections(t) + }) +} + +func TestControlConnectionClosesAfterCancellation(t *testing.T) { + requestStarted := make(chan struct{}) + var startedOnce sync.Once + var sawCancellation atomic.Bool + fixture := newUnixHTTPFixture(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/v1/vmm.ping" { + _, _ = w.Write([]byte("OK")) + return + } + startedOnce.Do(func() { close(requestStarted) }) + <-r.Context().Done() + sawCancellation.Store(true) + })) + + client, err := NewClient(fixture.socketPath()) + if err != nil { + t.Fatalf("create VZ client: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := client.GetVMInfo(ctx) + result <- err + }() + select { + case <-requestStarted: + case <-time.After(time.Second): + t.Fatal("VM info request did not reach the server") + } + cancel() + select { + case err := <-result: + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context canceled, got %v", err) + } + case <-time.After(time.Second): + t.Fatal("VM info request did not return after cancellation") + } + fixture.waitForNoConnections(t) + deadline := time.Now().Add(time.Second) + for !sawCancellation.Load() && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if !sawCancellation.Load() { + t.Fatal("server did not observe request cancellation") + } +} diff --git a/lib/instances/stop.go b/lib/instances/stop.go index ebcf58708..669a4f8ef 100644 --- a/lib/instances/stop.go +++ b/lib/instances/stop.go @@ -237,7 +237,13 @@ func (m *manager) stopInstance( } } - // 8. Always remove stale runtime sockets after process exit. + // 8. Drop the guest-agent connection for this VM incarnation and remove + // stale runtime sockets after process exit. The hypervisor shutdown closes + // the live peer, but the keyed gRPC ClientConn must not survive into a later + // start that reuses the same per-instance socket path. + if dialer, err := hypervisor.NewVsockDialer(inst.HypervisorType, inst.VsockSocket, inst.VsockCID); err == nil { + guest.CloseConn(dialer.Key()) + } // If graceful guest shutdown exits before shutdownHypervisor() is called, these // files may still exist and cause state derivation as Unknown or bind conflicts. _ = os.Remove(inst.SocketPath) From fda0883cd770e88487a51309f326ffd59f740239 Mon Sep 17 00:00:00 2001 From: Mike Bannister Date: Fri, 2 Oct 2026 19:56:19 -0400 Subject: [PATCH 2/2] test(instances): verify stop retires pooled guest connections --- lib/instances/lifecycle_noop_test.go | 50 ++++++++++++++++++++++++++++ 1 file changed, 50 insertions(+) diff --git a/lib/instances/lifecycle_noop_test.go b/lib/instances/lifecycle_noop_test.go index 4a0f76dcf..ec13654bf 100644 --- a/lib/instances/lifecycle_noop_test.go +++ b/lib/instances/lifecycle_noop_test.go @@ -3,6 +3,7 @@ package instances import ( "context" "errors" + "net" "os" "path/filepath" "sync" @@ -10,11 +11,13 @@ import ( "time" "github.com/kernel/hypeman/lib/devices" + "github.com/kernel/hypeman/lib/guest" "github.com/kernel/hypeman/lib/hypervisor" "github.com/kernel/hypeman/lib/paths" restartpolicy "github.com/kernel/hypeman/lib/restart-policy" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "google.golang.org/grpc/connectivity" ) const lifecycleNoopHypervisorType hypervisor.Type = "lifecycle-noop-test" @@ -29,6 +32,18 @@ func init() { } return lifecycleNoopHypervisor{state: state.(hypervisor.VMState)}, nil }) + hypervisor.RegisterVsockDialerFactory(lifecycleNoopHypervisorType, func(socketPath string, _ int64) hypervisor.VsockDialer { + return lifecycleNoopVsockDialer{socketPath: socketPath} + }) +} + +type lifecycleNoopVsockDialer struct { + socketPath string +} + +func (d lifecycleNoopVsockDialer) Key() string { return "lifecycle-noop:" + d.socketPath } +func (d lifecycleNoopVsockDialer) DialVsock(context.Context, int) (net.Conn, error) { + return nil, errors.New("no guest-agent in lifecycle fixture") } type lifecycleNoopHypervisor struct { @@ -221,6 +236,41 @@ func TestStartPersistsStaleVGPUReleaseImmediately(t *testing.T) { assert.Equal(t, "NVIDIA L40S-2Q", stored.GPUProfile, "profile is kept for the next start") } +func TestStopInstanceRetiresPooledGuestConnection(t *testing.T) { + m, id := newLifecycleNoopManagerWithInstance(t, StateRunning, time.Now().UTC()) + meta, err := m.loadMetadata(id) + require.NoError(t, err) + // Bypass guest shutdown and readiness RPCs so their retry cleanup cannot + // evict the connection on behalf of terminal stop cleanup. + meta.SkipGuestAgent = true + meta.VsockSocket = m.paths.InstanceSocket(id, "noop.vsock") + require.NoError(t, m.saveMetadata(meta)) + + ctx := context.Background() + dialer, err := hypervisor.NewVsockDialer(meta.HypervisorType, meta.VsockSocket, meta.VsockCID) + require.NoError(t, err) + t.Cleanup(func() { guest.CloseConn(dialer.Key()) }) + oldConn, err := guest.GetOrCreateConn(ctx, dialer) + require.NoError(t, err) + t.Cleanup(func() { _ = oldConn.Close() }) + cachedConn, err := guest.GetOrCreateConn(ctx, dialer) + require.NoError(t, err) + require.Same(t, oldConn, cachedConn, "connection must be pooled before stop") + + inst, err := m.StopInstance(ctx, id) + require.NoError(t, err) + require.Equal(t, StateStopped, inst.State) + + // Check the same key before any restart or RPC can recover a stale + // connection. GetOrCreateConn alone does not evict failed connections. + newConn, err := guest.GetOrCreateConn(ctx, dialer) + require.NoError(t, err) + require.NotSame(t, oldConn, newConn, "stop must evict the previous VM's pooled connection") + require.Eventually(t, func() bool { + return oldConn.GetState() == connectivity.Shutdown + }, time.Second, time.Millisecond, "stop must close the retired connection") +} + func TestStopStoppedInstanceLeavesVGPUForReconcile(t *testing.T) { m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) meta, err := m.loadMetadata(id)