Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion cmd/vz-shim/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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 {
Expand Down
10 changes: 8 additions & 2 deletions lib/hypervisor/vz/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
198 changes: 198 additions & 0 deletions lib/hypervisor/vz/client_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}
50 changes: 50 additions & 0 deletions lib/instances/lifecycle_noop_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,18 +4,21 @@ import (
"context"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"sync"
"testing"
"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"
Expand All @@ -30,6 +33,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 {
Expand Down Expand Up @@ -281,6 +296,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)
Expand Down
8 changes: 7 additions & 1 deletion lib/instances/stop.go
Original file line number Diff line number Diff line change
Expand Up @@ -244,7 +244,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())

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The existing VZ stop/restart test won't catch removal of this eviction: its first failed post-restart ExecIntoInstance attempt calls guest.CloseConn on retryable errors, and the surrounding test retries the exec. Please add a regression assertion that proves the stop path evicts the old pooled connection, or otherwise verifies the first post-restart RPC does not reuse it.

}
// 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)
Expand Down
Loading