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
104 changes: 86 additions & 18 deletions packages/orchestrator/pkg/sandbox/nbd/dispatch.go
Original file line number Diff line number Diff line change
Expand Up @@ -111,15 +111,21 @@ type Dispatch struct {
shuttingDown bool
shuttingDownLock sync.Mutex
fatal chan error
// asyncWriteZeroes controls whether WRITE_ZEROES/TRIM commands are handled
// in a goroutine (true) instead of inline on the read loop (false). Handling
// them inline lets a blocked reply writer holding writeLock stall the whole
// read loop; see cmdWriteZeroes.
asyncWriteZeroes bool
}

func NewDispatch(fp io.ReadWriter, prov Provider) *Dispatch {
func NewDispatch(fp io.ReadWriter, prov Provider, asyncWriteZeroes bool) *Dispatch {
d := &Dispatch{
responseHeader: make([]byte, 16),
fp: fp,
prov: prov,
provName: fmt.Sprintf("%T", prov),
fatal: make(chan error, 1),
responseHeader: make([]byte, 16),
fp: fp,
prov: prov,
provName: fmt.Sprintf("%T", prov),
fatal: make(chan error, 1),
asyncWriteZeroes: asyncWriteZeroes,
}

binary.BigEndian.PutUint32(d.responseHeader, NBDResponseMagic)
Expand Down Expand Up @@ -432,19 +438,81 @@ func (d *Dispatch) cmdWrite(ctx context.Context, cmdHandle uint64, cmdFrom uint6
return nil
}

// cmdWriteZeroes runs synchronously since WriteZeroesAt is cheap (mmap + bitmap).
// cmdWriteZeroes handles NBD WRITE_ZEROES and TRIM. The backend WriteZeroesAt
// call is cheap (mmap + bitmap), but the following writeResponse takes the
// shared writeLock and writes to the socket. Running this inline on the read
// loop (asyncWriteZeroes == false) means that if a concurrent reply writer is
// blocked inside writeResponse while holding writeLock (e.g. its socket write
// is blocked on a full send buffer), this command blocks the read loop on
// writeLock, the loop stops draining the socket, and the kernel eventually
// times out the connection (EIO). When asyncWriteZeroes is true the work runs
// in a goroutine like cmdRead/cmdWrite, so the read loop is never blocked.
func (d *Dispatch) cmdWriteZeroes(ctx context.Context, cmdHandle uint64, cmdFrom uint64, cmdLength int64) error {
var respErr uint32
if _, err := d.prov.WriteZeroesAt(int64(cmdFrom), cmdLength); err != nil {
respErr = 1
logger.L().Error(ctx, "nbd backend write-zeroes failed",
zap.Error(err),
zap.String("nbd_provider", d.provName),
zap.Uint64("nbd_handle", cmdHandle),
zap.Uint64("nbd_offset", cmdFrom),
zap.Int64("nbd_length", cmdLength),
)
performWriteZeroes := func() error {
// Run the backend call in a goroutine and select on ctx, mirroring
// cmdRead/cmdWrite, so a WriteZeroesAt that blocks cannot hang this
// goroutine (and the pendingResponses drain) during shutdown. The
// channel is buffered so the goroutine never leaks on the ctx.Done path.
errchan := make(chan error, 1)
go func() {
_, err := d.prov.WriteZeroesAt(int64(cmdFrom), cmdLength)
errchan <- err
}()

var zeroErr error
select {
case <-ctx.Done():
zeroErr = ctx.Err()
case zeroErr = <-errchan:
}

var respErr uint32
if zeroErr != nil {
respErr = 1
logger.L().Error(ctx, "nbd backend write-zeroes failed",
zap.Error(zeroErr),
zap.String("nbd_provider", d.provName),
zap.Uint64("nbd_handle", cmdHandle),
zap.Uint64("nbd_offset", cmdFrom),
zap.Int64("nbd_length", cmdLength),
)
}

return d.writeResponse(respErr, cmdHandle, nil)
}
Comment on lines +451 to +482

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.

high

Unlike cmdRead and cmdWrite, cmdWriteZeroes does not handle context cancellation when executing the backend operation. If the backend WriteZeroesAt call blocks indefinitely during shutdown, the goroutine will hang and prevent the orchestrator from draining pending responses cleanly, leading to a shutdown deadlock. Spawning a goroutine for the backend call and selecting on ctx.Done allows the dispatcher to unblock and exit gracefully.

	performWriteZeroes := func() error {
		errchan := make(chan error, 1)
		go func() {
			_, err := d.prov.WriteZeroesAt(int64(cmdFrom), cmdLength)
			errchan <- err
		}()

		var writeErr error
		select {
		case <-ctx.Done():
			writeErr = ctx.Err()
		case err := <-errchan:
			writeErr = err
		}

		var respErr uint32
		if writeErr != nil {
			respErr = 1
			logger.L().Error(ctx, "nbd backend write-zeroes failed",
				zap.Error(writeErr),
				zap.String("nbd_provider", d.provName),
				zap.Uint64("nbd_handle", cmdHandle),
				zap.Uint64("nbd_offset", cmdFrom),
				zap.Int64("nbd_length", cmdLength),
			)
		}

		return d.writeResponse(respErr, cmdHandle, nil)
	}

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Good catch — fixed. cmdWriteZeroes now runs WriteZeroesAt in a goroutine and selects on ctx, mirroring cmdRead/cmdWrite, so a blocked backend call can't hang the goroutine or the pending-response drain on shutdown. Folded into the refactor commit (performWriteZeroes).


return d.writeResponse(respErr, cmdHandle, nil)
if !d.asyncWriteZeroes {
return performWriteZeroes()
}

d.shuttingDownLock.Lock()
if d.shuttingDown {
d.shuttingDownLock.Unlock()

return ErrShuttingDown
}

d.pendingResponses.Add(1)
d.shuttingDownLock.Unlock()

go func() {
if err := performWriteZeroes(); err != nil {
select {
case d.fatal <- err:
default:
logger.L().Error(ctx, "nbd error cmd write-zeroes",
zap.Error(err),
zap.String("nbd_op", "write-zeroes"),
zap.String("nbd_provider", d.provName),
zap.Uint64("nbd_handle", cmdHandle),
zap.Uint64("nbd_offset", cmdFrom),
zap.Int64("nbd_length", cmdLength),
)
}
}

d.pendingResponses.Done()
}()

return nil
}
195 changes: 195 additions & 0 deletions packages/orchestrator/pkg/sandbox/nbd/dispatch_writezeroes_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
//go:build linux

package nbd

import (
"context"
"encoding/binary"
"io"
"sync"
"testing"
"time"
)

// ctrlConn is a fake NBD socket driving the real Dispatch.Handle loop.
//
// - Read serves queued request bytes (one queued slice per Read), blocking
// until a request is queued (or reqCh is closed -> io.EOF).
// - Write fires firstWrite once, then blocks until gate is closed, simulating
// a full socket send buffer (the kernel not draining replies).
type ctrlConn struct {
reqCh chan []byte
rbuf []byte
gate chan struct{}
firstWrite chan struct{}
fwOnce sync.Once
}

func (c *ctrlConn) Read(p []byte) (int, error) {
for len(c.rbuf) == 0 {
b, ok := <-c.reqCh
if !ok {
return 0, io.EOF
}
c.rbuf = b
}
n := copy(p, c.rbuf)
c.rbuf = c.rbuf[n:]

return n, nil
}

func (c *ctrlConn) Write(p []byte) (int, error) {
c.fwOnce.Do(func() { close(c.firstWrite) })
<-c.gate

return len(p), nil
}

// stallProv is a minimal Provider. It records the offsets passed to ReadAt so a
// test can tell whether the read loop served a particular request, and signals
// the first WriteZeroesAt call.
type stallProv struct {
mu sync.Mutex
seen map[int64]bool
wz chan struct{}
wzOnce sync.Once
}

func (m *stallProv) ReadAt(_ context.Context, p []byte, off int64) (int, error) {
m.mu.Lock()
m.seen[off] = true
m.mu.Unlock()

return len(p), nil
}

func (m *stallProv) sawRead(off int64) bool {
m.mu.Lock()
defer m.mu.Unlock()

return m.seen[off]
}

func (m *stallProv) Size(_ context.Context) (int64, error) { return 1 << 40, nil }
func (m *stallProv) WriteAt(p []byte, _ int64) (int, error) { return len(p), nil }

func (m *stallProv) WriteZeroesAt(_, length int64) (int, error) {
m.wzOnce.Do(func() { close(m.wz) })

return int(length), nil
}

func nbdRequest(typ uint16, handle, from uint64, length uint32) []byte {
b := make([]byte, 28)
binary.BigEndian.PutUint32(b[0:], NBDRequestMagic)
binary.BigEndian.PutUint16(b[4:], 0) // flags
binary.BigEndian.PutUint16(b[6:], typ)
binary.BigEndian.PutUint64(b[8:], handle)
binary.BigEndian.PutUint64(b[16:], from)
binary.BigEndian.PutUint32(b[24:], length)

return b
}

func waitForRead(p *stallProv, off int64, attempts int) bool {
for range attempts {
if p.sawRead(off) {
return true
}
time.Sleep(10 * time.Millisecond)
}

return p.sawRead(off)
}

// TestDispatchWriteZeroesReadLoopStall pins the head-of-line stall the
// nbd-async-write-zeroes flag fixes.
//
// A READ reply writer is made to block inside writeResponse while holding
// writeLock (simulating a full socket send buffer). A WRITE_ZEROES is then
// dispatched:
// - inline (asyncWriteZeroes=false): cmdWriteZeroes runs on the read loop and
// blocks acquiring writeLock, so the loop stops serving new requests.
// - async (asyncWriteZeroes=true): cmdWriteZeroes runs in a goroutine, so the
// read loop keeps serving.
//
// In both modes, once the blocked reply drains the loop must make progress
// again (matches the transient, self-clearing incidents this fix targets).
func TestDispatchWriteZeroesReadLoopStall(t *testing.T) {
t.Parallel()

for _, tc := range []struct {
name string
asyncWriteZeroes bool
wantStall bool
}{
{name: "inline_stalls_read_loop", asyncWriteZeroes: false, wantStall: true},
{name: "async_keeps_read_loop_alive", asyncWriteZeroes: true, wantStall: false},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

conn := &ctrlConn{
reqCh: make(chan []byte, 8),
gate: make(chan struct{}),
firstWrite: make(chan struct{}),
}
prov := &stallProv{seen: map[int64]bool{}, wz: make(chan struct{})}
d := NewDispatch(conn, prov, tc.asyncWriteZeroes)

done := make(chan struct{})
go func() {
_ = d.Handle(t.Context())
close(done)
}()
// Wind the loop down deterministically at the end: unblock any
// blocked reply Write, then close the request stream so Handle's
// Read returns io.EOF and the goroutine exits.
t.Cleanup(func() {
select {
case <-conn.gate:
default:
close(conn.gate)
}
close(conn.reqCh)
<-done
})

const offProbe = int64(0xABCDEF000)

// 1) READ -> its async reply writer blocks in Write holding writeLock.
conn.reqCh <- nbdRequest(NBDCmdRead, 1, 0, 512)
select {
case <-conn.firstWrite:
case <-time.After(2 * time.Second):
t.Fatal("read reply writer never reached the blocked socket Write")
}

// 2) WRITE_ZEROES -> inline blocks the loop on writeLock; async does not.
conn.reqCh <- nbdRequest(NBDCmdWriteZeroes, 2, 4096, 4096)
select {
case <-prov.wz:
case <-time.After(2 * time.Second):
t.Fatal("WriteZeroesAt was never called")
}

// 3) Probe with a fresh READ; a live loop serves it (records offProbe).
conn.reqCh <- nbdRequest(NBDCmdRead, 3, uint64(offProbe), 512)

served := waitForRead(prov, offProbe, 50) // ~500ms
switch {
case tc.wantStall && served:
t.Fatal("expected the read loop to STALL while a reply write is blocked, but it served a new request")
case !tc.wantStall && !served:
t.Fatal("expected the read loop to keep serving with async WRITE_ZEROES, but it stalled")
}

// Unblock replies; in both modes the loop must then serve the probe.
close(conn.gate)
if !waitForRead(prov, offProbe, 300) { // ~3s
t.Fatal("read loop did not serve the probe after replies drained")
}
})
}
}
3 changes: 2 additions & 1 deletion packages/orchestrator/pkg/sandbox/nbd/path_direct.go
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,7 @@ func (d *DirectPathMount) Open(ctx context.Context) (retDeviceIndex uint32, err
d.dispatchers = make([]*Dispatch, 0)

connections := d.featureFlags.IntFlag(ctx, featureflags.NBDConnectionsPerDevice)
asyncWriteZeroes := d.featureFlags.BoolFlag(ctx, featureflags.NBDAsyncWriteZeroesFlag)

for i := range connections {
// Create the socket pairs
Expand All @@ -151,7 +152,7 @@ func (d *DirectPathMount) Open(ctx context.Context) (retDeviceIndex uint32, err
}
server.Close()

dispatch := NewDispatch(serverc, d.Backend)
dispatch := NewDispatch(serverc, d.Backend, asyncWriteZeroes)
// Capture deviceIndex for the goroutine closure — it's reassigned on
// each retry iteration of the outer for-loop (not a range loop, so
// Go 1.22+ loop variable fix doesn't apply).
Expand Down
8 changes: 8 additions & 0 deletions packages/shared/pkg/featureflags/flags.go
Original file line number Diff line number Diff line change
Expand Up @@ -266,6 +266,14 @@ var (
// NBDConnectionsPerDevice the number of NBD socket connections per device
NBDConnectionsPerDevice = NewIntFlag("nbd-connections-per-device", 1)

// NBDAsyncWriteZeroesFlag, when enabled, handles NBD WRITE_ZEROES/TRIM
// commands in a goroutine instead of inline on the dispatch read loop.
// Inline handling can stall the read loop via head-of-line blocking on the
// shared write lock (when a reply writer is blocked on a full socket send
// buffer), which makes the kernel time out the NBD connection and surfaces
// as guest I/O errors. Disabled by default.
NBDAsyncWriteZeroesFlag = NewBoolFlag("nbd-async-write-zeroes", false)

// MemoryPrefetchMaxFetchWorkers is the maximum number of parallel fetch workers per sandbox for memory prefetching.
// Fetching is I/O bound so we can have more parallelism.
MemoryPrefetchMaxFetchWorkers = NewIntFlag("memory-prefetch-max-fetch-workers", 16)
Expand Down
Loading