diff --git a/cmd/guestbd/guestbd-main.go b/cmd/guestbd/guestbd-main.go new file mode 100644 index 0000000..92de675 --- /dev/null +++ b/cmd/guestbd/guestbd-main.go @@ -0,0 +1,81 @@ +// Command guestbd runs a guestbd NBD server. It listens for NBD client +// connections over TCP and serves a backing file (raw or qcow2). By default, +// each connection gets its own ephemeral writable snapshot that is discarded +// on disconnect. With --shared-snapshot, all connections share a single +// writable snapshot so that reconnections see previous writes. +// It also runs a debug HTTP server exposing expvar metrics and pprof endpoints. +package main + +import ( + "context" + "flag" + "log" + "net" + "net/http" + "os" + "os/signal" + + "github.com/tailscale/tb/guestbd" + "tailscale.com/tsweb" +) + +var ( + flagListen = flag.String("listen", ":10809", "NBD listen address") + flagFile = flag.String("file", "", "path to the backing file to serve; files are treated as raw files, unless filename ends in .qcow2") + flagPageSize = flag.Int("page-size", 4096, "page size in bytes (must be a power of two)") + flagMaxMem = flag.Int64("max-mem", 1<<30, "maximum memory for page cache in bytes") + flagDebug = flag.String("debug-addr", ":8080", "debug HTTP listen address") + flagSharedSnapshot = flag.Bool("shared-snapshot", false, "use a single shared writable snapshot for all connections instead of one per connection") +) + +func main() { + flag.Parse() + + if *flagFile == "" { + log.Fatal("--file is required") + } + if *flagPageSize <= 0 || (*flagPageSize&(*flagPageSize-1)) != 0 { + log.Fatal("--page-size must be a positive power of two") + } + + opts := []guestbd.ServerOption{ + guestbd.WithPageSize(*flagPageSize), + guestbd.WithMaxMem(*flagMaxMem), + } + if *flagSharedSnapshot { + opts = append(opts, guestbd.WithSharedSnapshot()) + } + + srv := guestbd.NewServer(guestbd.FileSource(*flagFile), opts...) + defer srv.Close() + srv.InitExpvar() + + // Debug HTTP server with tsweb. + debugMux := http.NewServeMux() + tsweb.Debugger(debugMux) + go func() { + log.Printf("debug HTTP server listening on %s", *flagDebug) + if err := http.ListenAndServe(*flagDebug, debugMux); err != nil { + log.Fatalf("debug HTTP: %v", err) + } + }() + + // NBD TCP listener. + ln, err := net.Listen("tcp", *flagListen) + if err != nil { + log.Fatalf("listen: %v", err) + } + log.Printf("NBD server listening on %s, serving %s", *flagListen, *flagFile) + + ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt) + defer cancel() + + go func() { + <-ctx.Done() + ln.Close() + }() + + if err := srv.Serve(ln); err != nil && ctx.Err() == nil { + log.Fatalf("serve: %v", err) + } +} diff --git a/go.mod b/go.mod index a0e51da..4d00203 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,7 @@ go 1.27.1 require ( github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900 + github.com/bradfitz/qcow2 v0.0.0-20260303185237-93afc730382b github.com/cespare/xxhash/v2 v2.3.0 github.com/go-jose/go-jose/v4 v4.1.3 github.com/golang-jwt/jwt/v5 v5.3.1 @@ -28,6 +29,7 @@ require ( github.com/google/uuid v1.6.0 // indirect github.com/hdevalence/ed25519consensus v0.2.0 // indirect github.com/jsimonetti/rtnetlink v1.4.2 // indirect + github.com/klauspost/compress v1.20.0 // indirect github.com/mattn/go-isatty v0.0.24 // indirect github.com/mdlayher/netlink v1.11.2 // indirect github.com/mdlayher/socket v0.7.0 // indirect diff --git a/go.sum b/go.sum index b0cc92a..507d011 100644 --- a/go.sum +++ b/go.sum @@ -38,6 +38,8 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900 h1:YTPrKVaBJlca4xkN+YU2VG3WHIDFme0vItEKuj0wrBQ= github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900/go.mod h1:nmAuQ8iUcGbqQsa847JnlFq6PQMMxpEeK0QCPAPYf3w= +github.com/bradfitz/qcow2 v0.0.0-20260303185237-93afc730382b h1:D6BX3KA9oZ7WGn9D7fKsDNyQjrHb1kRsQTsyc7nv35s= +github.com/bradfitz/qcow2 v0.0.0-20260303185237-93afc730382b/go.mod h1:829+KZfDIY07C9uUedqANoUnClohrBIXLTQdhr3gC9w= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cilium/ebpf v0.22.0 h1:v2ktp0roffpMOj2MMf3idtCQZOsAoC4BJbAJN+ke2bY= diff --git a/guestbd/.gitignore b/guestbd/.gitignore new file mode 100644 index 0000000..9301f2f --- /dev/null +++ b/guestbd/.gitignore @@ -0,0 +1,2 @@ +./guestbd +*~ diff --git a/guestbd/README.md b/guestbd/README.md new file mode 100644 index 0000000..06d8d04 --- /dev/null +++ b/guestbd/README.md @@ -0,0 +1,105 @@ +# guestbd + +guestbd is a userspace NBD server that gives each TCP connection its own virtual +read/write namespace on top of a given file on disk that's only read-only. Then +the TCP connection breaks, any writes that were made by that client are lost. + +Both raw disk images and qcow2 images are supported. Files ending in `.qcow2` +are automatically opened as qcow2 (with support for deflate and zstd +compression). qcow2 support is provided by +[github.com/bradfitz/qcow2](https://pkg.go.dev/github.com/bradfitz/qcow2). + +This is meant for the being the root block device for short-lived ephemeral CI +workload VMs, where we prefer speed over any sort of durability. + +Any written data exists only in memory (for the first configured N gigabytes) +and only spills to disk as needed as a cache. + +Everything is content-addressable and de-duped. + +When a new TCP connection is accepted, the named file (given by a flag to the +binary) is opened, and its *os.File is compared against existing open TCP +connections to see if it maps to the same inode on disk. That means all +connections to same base file image share the same cache. + +Then, each page's 4KB data is stored just once for the whole process as a +function of its hash. This means if the file on disk is replaced and gets a new +inode, but there are clients still open on the old version, and the new version +has 80% of the same contents overall (e.g. the ext4 filesystem was rebuilt with +mostly identical contents, but different inode dentry metadata), then the cache +will be mostly shared. + +That hash=>contents is stored in a configurably sized LRU cache tracking hit +rates, and used for all connections. No reference counting is done on it; things +simply age out if they no longer exist in the base image or any old versions of +the base image. + +Likewise, all writes update the underlying block device in 4KB units, so a small +1KB write from a client ends up reading the existing 4KB, mutating the 1KB in +the middle of it, and then hashing the whole 4KB. + +It's expected that many connections will end up doing writes with the same 4KB +page contents (at different offsets in the block device), so we also share that +memory. But because written data needs to always be re-readable later, we need +to guarantee it either exists in memory or on disk. To start simple, we don't +try to be clever with how data is lazily written to disk. Each connection +maintains a table: + + pageNum => *struct{ hash pageHash, dirtyPage int } + +Where dirtyPage is the 4KB page number on disk (a per TCP connection temp file +that's open for read/write and then unlinked immediately) where the page was +written. Whenever a new dirty page is written, it gets a monotonically +increasing dirtyPageNum. All writes to the connDirtyFile are in 4KB units, or +whatever the flag-configured page size is (which can be limited to a power of +two) + +### Observability + +The server uses tailscale.com/tsweb (including its DebugHandler) to +expose pprof, expvar, and a Prometheus-compatible `/debug/varz` endpoint +on the `--debug-addr` (default `:8080`). + +Metrics use normal expvar metrics (which tsweb Prometheus-ifies) and +tailscale.com/metrics's LabelMap and Histogram types. + +#### Counters + +| Metric | Description | +|--------|-------------| +| `guestbd_total_conns` | Total TCP connections accepted | +| `guestbd_nbd_ops{type=read\|write\|disconnect\|flush\|trim}` | NBD operations by type | +| `guestbd_read_bytes` | Total bytes read by clients | +| `guestbd_read_pages` | Total page reads (dirty + base layer) | +| `guestbd_read_path{type=base_mem\|base_disk_cold\|base_disk_miss\|from_write}` | Read source breakdown | +| `guestbd_write_bytes` | Total bytes written by clients | +| `guestbd_write_pages` | Total dirty pages written | +| `guestbd_cache{path=hits\|misses\|evictions}` | Page cache operations | + +#### Gauges + +| Metric | Description | +|--------|-------------| +| `guestbd_active_conns` | Currently connected clients | +| `guestbd_cache_entries` | Pages in the LRU cache | +| `guestbd_cache_bytes` | Bytes used by the LRU cache | +| `guestbd_base_images_active` | readonlyFile entries with active connections | +| `guestbd_base_images_cached` | readonlyFile entries in memory (including idle) | +| `guestbd_page_size` | Configured page size | +| `guestbd_max_dirty_bytes` | Dirty page bytes of the connection with the most dirty pages | + +#### Histograms + +| Metric | Description | +|--------|-------------| +| `guestbd_read_size_bytes` | Distribution of NBD read request sizes | +| `guestbd_write_size_bytes` | Distribution of NBD write request sizes | +| `guestbd_read_latency_seconds` | Distribution of NBD read latencies (snapshot ReadAt only, excludes wire I/O) | +| `guestbd_write_latency_seconds` | Distribution of NBD write latencies (snapshot WriteAt only, excludes wire I/O) | + +The `read_path` metric is particularly useful for diagnosing cache effectiveness: + +- **`base_mem`** — page hash was known and data was in the LRU cache (ideal) +- **`base_disk_cold`** — page hash was unknown (first read ever, or readonlyFile was recreated due to inode change); had to read from disk +- **`base_disk_miss`** — page hash was known but data was evicted from the LRU cache; had to re-read from disk (indicates cache is too small) +- **`from_write`** — read of a page that was previously written on this connection diff --git a/guestbd/cache.go b/guestbd/cache.go new file mode 100644 index 0000000..7ec1c25 --- /dev/null +++ b/guestbd/cache.go @@ -0,0 +1,105 @@ +package guestbd + +import ( + "bytes" + "container/list" + "crypto/sha256" + "expvar" + "sync" + + "tailscale.com/metrics" +) + +// pageHash is the sha256 hash of a page's contents. +// A zero value means the page has not been read from disk yet. +type pageHash [sha256.Size]byte + +// hashPage returns the SHA-256 hash of data. +func hashPage(data []byte) pageHash { + return sha256.Sum256(data) +} + +// pageCache is a content-addressable LRU cache of page data, +// shared across all connections. Pages are keyed by their sha256 hash. +type pageCache struct { + mu sync.Mutex + maxPages int + pageSize int + + items map[pageHash]*list.Element + lru *list.List // front = most recently used + + path metrics.LabelMap // counter_guestbd_cache{path="hits|misses|evictions"} + entries expvar.Int // gauge_guestbd_cache_entries + bytes expvar.Int // gauge_guestbd_cache_bytes +} + +// cacheEntry is a single element in the LRU list, pairing a page hash with +// its data. +type cacheEntry struct { + hash pageHash + data []byte +} + +// newPageCache returns a new page cache that holds at most maxPages pages of +// the given pageSize. +func newPageCache(maxPages, pageSize int) *pageCache { + return &pageCache{ + maxPages: maxPages, + pageSize: pageSize, + items: make(map[pageHash]*list.Element), + lru: list.New(), + path: metrics.LabelMap{Label: "path"}, + } +} + +// Get returns the page data for the given hash, if present in the cache. +func (c *pageCache) Get(h pageHash) ([]byte, bool) { + c.mu.Lock() + defer c.mu.Unlock() + + if elem, ok := c.items[h]; ok { + c.lru.MoveToFront(elem) + c.path.Add("hits", 1) + return elem.Value.(*cacheEntry).data, true + } + c.path.Add("misses", 1) + return nil, false +} + +// Put adds page data to the cache, keyed by its hash. +// The data slice is cloned; the caller retains ownership of the original. +// If the hash already exists, it's moved to the front of the LRU. +func (c *pageCache) Put(h pageHash, data []byte) { + c.mu.Lock() + defer c.mu.Unlock() + + if elem, ok := c.items[h]; ok { + c.lru.MoveToFront(elem) + return + } + + entry := &cacheEntry{hash: h, data: bytes.Clone(data)} + elem := c.lru.PushFront(entry) + c.items[h] = elem + c.entries.Add(1) + c.bytes.Add(int64(c.pageSize)) + + for c.lru.Len() > c.maxPages { + c.evict() + } +} + +// evict removes the least recently used entry from the cache. +func (c *pageCache) evict() { + elem := c.lru.Back() + if elem == nil { + return + } + c.lru.Remove(elem) + entry := elem.Value.(*cacheEntry) + delete(c.items, entry.hash) + c.path.Add("evictions", 1) + c.entries.Add(-1) + c.bytes.Add(-int64(c.pageSize)) +} diff --git a/guestbd/doc.go b/guestbd/doc.go new file mode 100644 index 0000000..c7c974c --- /dev/null +++ b/guestbd/doc.go @@ -0,0 +1,18 @@ +// Package guestbd implements a userspace NBD (Network Block Device) server +// designed for ephemeral CI workload VMs. +// +// Each Snapshot provides a read/write namespace layered on top of a shared +// read-only base image. By default each TCP connection gets its own Snapshot +// whose writes are discarded on disconnect, but a Snapshot can also be shared +// across multiple connections so that reconnecting clients see previous writes. +// +// The base image is provided as a [BaseImageSource] — a function returning a +// [BaseImage] — so callers can serve images from files, object stores, +// or memory. [BaseImage.BaseImageKey] enables equivalence-keyed caching of +// idle baseImageStates so that reconnecting clients reuse the page hash table. +// The [FileSource] helper provides the common file-based workflow, with +// automatic qcow2 detection by file extension. +// +// Pages are content-addressed by their SHA-256 hash and de-duplicated in a +// global LRU cache shared across all snapshots. +package guestbd diff --git a/guestbd/fileid_other.go b/guestbd/fileid_other.go new file mode 100644 index 0000000..e640b45 --- /dev/null +++ b/guestbd/fileid_other.go @@ -0,0 +1,11 @@ +//go:build !unix && !windows + +package guestbd + +import "os" + +// fileIdentity returns nil, as files have no device and inode identity +// here. Base images opened from files are then never coalesced. +func fileIdentity(f *os.File, fi os.FileInfo) any { + return nil +} diff --git a/guestbd/fileid_unix.go b/guestbd/fileid_unix.go new file mode 100644 index 0000000..4f4d62d --- /dev/null +++ b/guestbd/fileid_unix.go @@ -0,0 +1,15 @@ +//go:build unix + +package guestbd + +import ( + "os" + "syscall" +) + +// fileIdentity returns a BaseImageKey identifying the open file f, +// described by fi, by its device and inode. +func fileIdentity(f *os.File, fi os.FileInfo) any { + st := fi.Sys().(*syscall.Stat_t) + return fileIdentityKey{dev: uint64(st.Dev), ino: st.Ino} // Dev is int32 on darwin +} diff --git a/guestbd/fileid_windows.go b/guestbd/fileid_windows.go new file mode 100644 index 0000000..95b886d --- /dev/null +++ b/guestbd/fileid_windows.go @@ -0,0 +1,20 @@ +package guestbd + +import ( + "os" + "syscall" +) + +// fileIdentity returns a BaseImageKey identifying the open file f by its +// volume serial number and file index, Windows' equivalent of a device +// and inode. It returns nil, disabling coalescing, if they can't be read. +func fileIdentity(f *os.File, fi os.FileInfo) any { + var d syscall.ByHandleFileInformation + if err := syscall.GetFileInformationByHandle(syscall.Handle(f.Fd()), &d); err != nil { + return nil + } + return fileIdentityKey{ + dev: uint64(d.VolumeSerialNumber), + ino: uint64(d.FileIndexHigh)<<32 | uint64(d.FileIndexLow), + } +} diff --git a/guestbd/guestbd_linux_test.go b/guestbd/guestbd_linux_test.go new file mode 100644 index 0000000..33457d6 --- /dev/null +++ b/guestbd/guestbd_linux_test.go @@ -0,0 +1,322 @@ +package guestbd + +import ( + "bytes" + "crypto/rand" + "fmt" + "net" + "os" + "os/exec" + "testing" + "time" +) + +// skipUnlessCanNBD skips the test unless we're running on Linux +// as root or with passwordless sudo, nbd-client is installed, +// and the nbd kernel module can be loaded. It returns the sudo +// prefix ("sudo" or ""). +func skipUnlessCanNBD(t *testing.T) string { + t.Helper() + + var sudo string + if os.Getuid() != 0 { + out, err := exec.Command("sudo", "-n", "true").CombinedOutput() + if err != nil { + t.Skipf("skipping: not root and no passwordless sudo: %s", out) + } + sudo = "sudo" + } + + if _, err := exec.LookPath("nbd-client"); err != nil { + t.Skip("skipping: nbd-client not found in PATH") + } + + out, err := sudoExec(sudo, "modprobe", "nbd") + if err != nil { + t.Skipf("skipping: cannot load nbd module: %s %v", out, err) + } + + return sudo +} + +// sudoExec runs a command, optionally prefixed with sudo. +func sudoExec(sudo string, args ...string) ([]byte, error) { + if sudo != "" { + args = append([]string{sudo}, args...) + } + return exec.Command(args[0], args[1:]...).CombinedOutput() +} + +// findFreeNBD finds an unused /dev/nbdN device. +func findFreeNBD(t *testing.T, sudo string) string { + t.Helper() + for i := 0; i < 16; i++ { + dev := fmt.Sprintf("/dev/nbd%d", i) + if _, err := os.Stat(dev); err != nil { + continue + } + // nbd-client -c exits non-zero if the device is not connected. + var cmd *exec.Cmd + if sudo != "" { + cmd = exec.Command(sudo, "nbd-client", "-c", dev) + } else { + cmd = exec.Command("nbd-client", "-c", dev) + } + if err := cmd.Run(); err != nil { + return dev + } + } + t.Skip("skipping: no free /dev/nbdX device found") + return "" +} + +// nbdClientConnect connects the kernel NBD device to our server. +func nbdClientConnect(t *testing.T, sudo, addr, dev string, blockSize int) { + t.Helper() + host, port, err := net.SplitHostPort(addr) + if err != nil { + t.Fatalf("bad addr %q: %v", addr, err) + } + out, err := sudoExec(sudo, + "nbd-client", "-N", "", "-b", fmt.Sprint(blockSize), + host, port, dev) + if err != nil { + t.Fatalf("nbd-client connect %s to %s: %s %v", dev, addr, out, err) + } +} + +// nbdClientDisconnect disconnects the kernel NBD device. +func nbdClientDisconnect(sudo, dev string) { + sudoExec(sudo, "nbd-client", "-d", dev) +} + +// devRead reads length bytes at the given byte offset from a block device +// using dd, bypassing any page cache with iflag=direct. +func devRead(t *testing.T, sudo, dev string, offset, length int) []byte { + t.Helper() + args := []string{ + "dd", + "if=" + dev, + fmt.Sprintf("bs=%d", length), + "count=1", + fmt.Sprintf("skip=%d", offset), + "iflag=skip_bytes,direct", + "status=none", + } + if sudo != "" { + args = append([]string{sudo}, args...) + } + cmd := exec.Command(args[0], args[1:]...) + out, err := cmd.Output() + if err != nil { + ee, _ := err.(*exec.ExitError) + t.Fatalf("dd read at offset %d: %v (stderr: %s)", offset, err, ee.Stderr) + } + return out +} + +// devWrite writes data at the given byte offset to a block device +// using dd, bypassing any page cache with oflag=direct. +func devWrite(t *testing.T, sudo, dev string, offset int, data []byte) { + t.Helper() + args := []string{ + "dd", + "of=" + dev, + fmt.Sprintf("bs=%d", len(data)), + "count=1", + fmt.Sprintf("seek=%d", offset), + "oflag=seek_bytes,direct", + "conv=notrunc", + "status=none", + } + if sudo != "" { + args = append([]string{sudo}, args...) + } + cmd := exec.Command(args[0], args[1:]...) + cmd.Stdin = bytes.NewReader(data) + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("dd write at offset %d: %s %v", offset, out, err) + } +} + +func TestLinuxNBDRead(t *testing.T) { + sudo := skipUnlessCanNBD(t) + dev := findFreeNBD(t, sudo) + + const pageSize = 4096 + data := make([]byte, pageSize*10) + rand.Read(data) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + nbdClientConnect(t, sudo, addr, dev, pageSize) + t.Cleanup(func() { nbdClientDisconnect(sudo, dev) }) + time.Sleep(200 * time.Millisecond) + + // Read first page. + got := devRead(t, sudo, dev, 0, pageSize) + if !bytes.Equal(got, data[:pageSize]) { + t.Fatalf("first page mismatch: got %x... want %x...", got[:16], data[:16]) + } + + // Read a page in the middle. + off := 5 * pageSize + got = devRead(t, sudo, dev, off, pageSize) + if !bytes.Equal(got, data[off:off+pageSize]) { + t.Fatalf("middle page mismatch at offset %d", off) + } + + // Read last page. + off = 9 * pageSize + got = devRead(t, sudo, dev, off, pageSize) + if !bytes.Equal(got, data[off:off+pageSize]) { + t.Fatalf("last page mismatch at offset %d", off) + } +} + +func TestLinuxNBDWriteAndReadBack(t *testing.T) { + sudo := skipUnlessCanNBD(t) + dev := findFreeNBD(t, sudo) + + const pageSize = 4096 + original := make([]byte, pageSize*4) + rand.Read(original) + + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + + nbdClientConnect(t, sudo, addr, dev, pageSize) + t.Cleanup(func() { nbdClientDisconnect(sudo, dev) }) + time.Sleep(200 * time.Millisecond) + + // Write a full page of 0xBE. + writeData := bytes.Repeat([]byte{0xBE}, pageSize) + devWrite(t, sudo, dev, 0, writeData) + + // Read it back. + got := devRead(t, sudo, dev, 0, pageSize) + if !bytes.Equal(got, writeData) { + t.Fatalf("written page read-back mismatch: got %x... want %x...", got[:16], writeData[:16]) + } + + // Verify an unmodified page is still the original data. + got = devRead(t, sudo, dev, pageSize, pageSize) + if !bytes.Equal(got, original[pageSize:2*pageSize]) { + t.Fatal("unmodified page should match original") + } + + // Write a second distinct pattern to page 2. + writeData2 := bytes.Repeat([]byte{0xEF}, pageSize) + devWrite(t, sudo, dev, 2*pageSize, writeData2) + + got = devRead(t, sudo, dev, 2*pageSize, pageSize) + if !bytes.Equal(got, writeData2) { + t.Fatal("second write read-back mismatch") + } +} + +func TestLinuxNBDWritesLostOnReconnect(t *testing.T) { + sudo := skipUnlessCanNBD(t) + dev := findFreeNBD(t, sudo) + + const pageSize = 4096 + original := make([]byte, pageSize*4) + rand.Read(original) + + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + + // First connection: write data via the kernel. + nbdClientConnect(t, sudo, addr, dev, pageSize) + time.Sleep(200 * time.Millisecond) + + writeData := bytes.Repeat([]byte{0xAA}, pageSize) + devWrite(t, sudo, dev, 0, writeData) + + // Confirm the write is visible. + got := devRead(t, sudo, dev, 0, pageSize) + if !bytes.Equal(got, writeData) { + t.Fatalf("write not visible: got %x... want %x...", got[:16], writeData[:16]) + } + + // Disconnect. + nbdClientDisconnect(sudo, dev) + time.Sleep(200 * time.Millisecond) + + // Reconnect: the server's per-connection writes should be gone. + nbdClientConnect(t, sudo, addr, dev, pageSize) + t.Cleanup(func() { nbdClientDisconnect(sudo, dev) }) + time.Sleep(200 * time.Millisecond) + + got = devRead(t, sudo, dev, 0, pageSize) + if !bytes.Equal(got, original[:pageSize]) { + t.Fatalf("writes should be lost after reconnect: got %x... want %x...", + got[:16], original[:16]) + } +} + +func TestLinuxNBDMultipleConnections(t *testing.T) { + sudo := skipUnlessCanNBD(t) + + // Find two free devices. + dev1 := findFreeNBD(t, sudo) + dev2 := "" + for i := 0; i < 16; i++ { + d := fmt.Sprintf("/dev/nbd%d", i) + if d == dev1 { + continue + } + if _, err := os.Stat(d); err != nil { + continue + } + var cmd *exec.Cmd + if sudo != "" { + cmd = exec.Command(sudo, "nbd-client", "-c", d) + } else { + cmd = exec.Command("nbd-client", "-c", d) + } + if err := cmd.Run(); err != nil { + dev2 = d + break + } + } + if dev2 == "" { + t.Skip("skipping: need two free /dev/nbdX devices") + } + + const pageSize = 4096 + original := make([]byte, pageSize*4) + rand.Read(original) + + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + + nbdClientConnect(t, sudo, addr, dev1, pageSize) + t.Cleanup(func() { nbdClientDisconnect(sudo, dev1) }) + nbdClientConnect(t, sudo, addr, dev2, pageSize) + t.Cleanup(func() { nbdClientDisconnect(sudo, dev2) }) + time.Sleep(200 * time.Millisecond) + + // Both should see the same original data. + got1 := devRead(t, sudo, dev1, 0, pageSize) + got2 := devRead(t, sudo, dev2, 0, pageSize) + if !bytes.Equal(got1, original[:pageSize]) || !bytes.Equal(got2, original[:pageSize]) { + t.Fatal("both connections should see original data") + } + + // Write on dev1 should not be visible on dev2. + writeData := bytes.Repeat([]byte{0xCC}, pageSize) + devWrite(t, sudo, dev1, 0, writeData) + + got1 = devRead(t, sudo, dev1, 0, pageSize) + if !bytes.Equal(got1, writeData) { + t.Fatal("dev1 should see its own write") + } + + got2 = devRead(t, sudo, dev2, 0, pageSize) + if !bytes.Equal(got2, original[:pageSize]) { + t.Fatalf("dev2 should still see original: got %x... want %x...", + got2[:16], original[:16]) + } +} diff --git a/guestbd/guestbd_test.go b/guestbd/guestbd_test.go new file mode 100644 index 0000000..b773183 --- /dev/null +++ b/guestbd/guestbd_test.go @@ -0,0 +1,1486 @@ +package guestbd + +import ( + "bytes" + "crypto/rand" + "encoding/binary" + "fmt" + "io" + "math" + "net" + "os" + "os/exec" + "path/filepath" + "runtime" + "slices" + "sync" + "testing" + "time" +) + +func TestLatencyBuckets(t *testing.T) { + got := latencyBuckets() + want := []float64{ + 10e-6, 20e-6, 40e-6, 80e-6, + 160e-6, 320e-6, 640e-6, 1280e-6, + 2560e-6, 5120e-6, 10240e-6, 20480e-6, + 40960e-6, 81920e-6, 163840e-6, 327680e-6, + 655360e-6, 1310720e-6, 2621440e-6, 5242880e-6, + 10485760e-6, + } + if len(got) != len(want) { + t.Fatalf("got %d buckets, want %d: %v", len(got), len(want), got) + } + for i, g := range got { + if math.Abs(g-want[i])/want[i] > 1e-9 { + t.Errorf("bucket %d: got %g, want %g", i, g, want[i]) + } + } + t.Logf("%d buckets, first=%v, last=%v", + len(got), time.Duration(got[0]*float64(time.Second)), + time.Duration(got[len(got)-1]*float64(time.Second))) +} + +// testNBDClient is a pure Go NBD client for testing. +type testNBDClient struct { + t *testing.T + conn net.Conn + exportSize uint64 + nextHandle uint64 +} + +func newTestClient(t *testing.T, addr string) *testNBDClient { + t.Helper() + conn, err := net.Dial("tcp", addr) + if err != nil { + t.Fatalf("dial: %v", err) + } + c := &testNBDClient{t: t, conn: conn} + c.handshake() + return c +} + +func (c *testNBDClient) handshake() { + c.t.Helper() + + // Read server handshake: magic(8) + opts_magic(8) + flags(2) + var handshake [18]byte + if _, err := io.ReadFull(c.conn, handshake[:]); err != nil { + c.t.Fatalf("read handshake: %v", err) + } + magic := binary.BigEndian.Uint64(handshake[0:8]) + if magic != nbdMagic { + c.t.Fatalf("bad magic: %#x", magic) + } + optsMagic := binary.BigEndian.Uint64(handshake[8:16]) + if optsMagic != nbdOptsMagic { + c.t.Fatalf("bad opts magic: %#x", optsMagic) + } + + // Send client flags (fixed newstyle + no zeroes). + var clientFlags [4]byte + binary.BigEndian.PutUint32(clientFlags[:], uint32(nbdFlagCFixedNewstyle|nbdFlagCNoZeroes)) + if _, err := c.conn.Write(clientFlags[:]); err != nil { + c.t.Fatalf("write client flags: %v", err) + } + + c.infoOpt(nbdOptGo) +} + +// infoOpt sends NBD_OPT_GO or NBD_OPT_INFO with an empty export name and +// the given info requests, and reads replies until the ACK. It returns the +// info types the server sent. +func (c *testNBDClient) infoOpt(optCode uint32, infoReqs ...uint16) (infoTypes []uint16) { + c.t.Helper() + optBuf := make([]byte, 22+2*len(infoReqs)) + binary.BigEndian.PutUint64(optBuf[0:8], nbdOptsMagic) + binary.BigEndian.PutUint32(optBuf[8:12], optCode) + binary.BigEndian.PutUint32(optBuf[12:16], uint32(6+2*len(infoReqs))) // name_len(4) + name(0) + info_count(2) + requests + binary.BigEndian.PutUint32(optBuf[16:20], 0) // name length + binary.BigEndian.PutUint16(optBuf[20:22], uint16(len(infoReqs))) // number of info requests + for i, r := range infoReqs { + binary.BigEndian.PutUint16(optBuf[22+2*i:], r) + } + if _, err := c.conn.Write(optBuf); err != nil { + c.t.Fatalf("write opt %d: %v", optCode, err) + } + + // Read option replies until ACK. + for { + var replyHeader [20]byte + if _, err := io.ReadFull(c.conn, replyHeader[:]); err != nil { + c.t.Fatalf("read opt reply: %v", err) + } + replyMagic := binary.BigEndian.Uint64(replyHeader[0:8]) + if replyMagic != nbdOptReplyMagic { + c.t.Fatalf("bad opt reply magic: %#x", replyMagic) + } + replyType := binary.BigEndian.Uint32(replyHeader[12:16]) + replyLen := binary.BigEndian.Uint32(replyHeader[16:20]) + + replyData := make([]byte, replyLen) + if replyLen > 0 { + if _, err := io.ReadFull(c.conn, replyData); err != nil { + c.t.Fatalf("read opt reply data: %v", err) + } + } + + if replyType == nbdRepInfo && replyLen >= 2 { + infoType := binary.BigEndian.Uint16(replyData[0:2]) + infoTypes = append(infoTypes, infoType) + if infoType == nbdInfoExport && replyLen >= 12 { + c.exportSize = binary.BigEndian.Uint64(replyData[2:10]) + } + } + + if replyType == nbdRepAck { + return infoTypes + } + if replyType&(1<<31) != 0 { + c.t.Fatalf("opt reply error: type=%#x", replyType) + } + } +} + +// TestOptInfoThenGo tests the option sequence Apple's Virtualization.framework +// NBD client uses: NBD_OPT_INFO asking for the block size, then NBD_OPT_GO. +func TestOptInfoThenGo(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*2) + for i := range data { + data[i] = byte(i % 251) + } + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + conn, err := net.Dial("tcp", addr) + if err != nil { + t.Fatal(err) + } + c := &testNBDClient{t: t, conn: conn} + defer c.disconnect() + var greeting [18]byte + if _, err := io.ReadFull(conn, greeting[:]); err != nil { + t.Fatal(err) + } + // Fixed newstyle only, as Apple's client sends. + var clientFlags [4]byte + binary.BigEndian.PutUint32(clientFlags[:], nbdFlagCFixedNewstyle) + if _, err := conn.Write(clientFlags[:]); err != nil { + t.Fatal(err) + } + + infos := c.infoOpt(nbdOptInfo, nbdInfoBlockSize) + if !slices.Contains(infos, nbdInfoExport) || !slices.Contains(infos, nbdInfoBlockSize) { + t.Fatalf("NBD_OPT_INFO replies = %v; want export and block size info", infos) + } + if c.exportSize != uint64(len(data)) { + t.Fatalf("export size = %d, want %d", c.exportSize, len(data)) + } + c.infoOpt(nbdOptGo, nbdInfoBlockSize) + if got := c.read(pageSize, 100); !bytes.Equal(got, data[pageSize:pageSize+100]) { + t.Fatal("read after NBD_OPT_INFO + NBD_OPT_GO mismatch") + } +} + +func (c *testNBDClient) read(offset uint64, length uint32) []byte { + c.t.Helper() + handle := c.nextHandle + c.nextHandle++ + + var req [28]byte + binary.BigEndian.PutUint32(req[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(req[4:6], 0) + binary.BigEndian.PutUint16(req[6:8], uint16(nbdCmdRead)) + binary.BigEndian.PutUint64(req[8:16], handle) + binary.BigEndian.PutUint64(req[16:24], offset) + binary.BigEndian.PutUint32(req[24:28], length) + if _, err := c.conn.Write(req[:]); err != nil { + c.t.Fatalf("write read request: %v", err) + } + + var reply [16]byte + if _, err := io.ReadFull(c.conn, reply[:]); err != nil { + c.t.Fatalf("read reply: %v", err) + } + replyMagic := binary.BigEndian.Uint32(reply[0:4]) + if replyMagic != nbdReplyMagic { + c.t.Fatalf("bad reply magic: %#x", replyMagic) + } + errCode := binary.BigEndian.Uint32(reply[4:8]) + if errCode != 0 { + c.t.Fatalf("read error: %d", errCode) + } + replyHandle := binary.BigEndian.Uint64(reply[8:16]) + if replyHandle != handle { + c.t.Fatalf("handle mismatch: got %d want %d", replyHandle, handle) + } + + data := make([]byte, length) + if _, err := io.ReadFull(c.conn, data); err != nil { + c.t.Fatalf("read data: %v", err) + } + return data +} + +func (c *testNBDClient) write(offset uint64, data []byte) { + c.t.Helper() + handle := c.nextHandle + c.nextHandle++ + + var req [28]byte + binary.BigEndian.PutUint32(req[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(req[4:6], 0) + binary.BigEndian.PutUint16(req[6:8], uint16(nbdCmdWrite)) + binary.BigEndian.PutUint64(req[8:16], handle) + binary.BigEndian.PutUint64(req[16:24], offset) + binary.BigEndian.PutUint32(req[24:28], uint32(len(data))) + if _, err := c.conn.Write(req[:]); err != nil { + c.t.Fatalf("write request: %v", err) + } + if _, err := c.conn.Write(data); err != nil { + c.t.Fatalf("write data: %v", err) + } + + var reply [16]byte + if _, err := io.ReadFull(c.conn, reply[:]); err != nil { + c.t.Fatalf("read write reply: %v", err) + } + replyMagic := binary.BigEndian.Uint32(reply[0:4]) + if replyMagic != nbdReplyMagic { + c.t.Fatalf("bad reply magic: %#x", replyMagic) + } + errCode := binary.BigEndian.Uint32(reply[4:8]) + if errCode != 0 { + c.t.Fatalf("write error: %d", errCode) + } +} + +func (c *testNBDClient) trim(offset uint64, length uint32) { + c.t.Helper() + handle := c.nextHandle + c.nextHandle++ + + var req [28]byte + binary.BigEndian.PutUint32(req[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(req[4:6], 0) + binary.BigEndian.PutUint16(req[6:8], uint16(nbdCmdTrim)) + binary.BigEndian.PutUint64(req[8:16], handle) + binary.BigEndian.PutUint64(req[16:24], offset) + binary.BigEndian.PutUint32(req[24:28], length) + if _, err := c.conn.Write(req[:]); err != nil { + c.t.Fatalf("write trim request: %v", err) + } + + var reply [16]byte + if _, err := io.ReadFull(c.conn, reply[:]); err != nil { + c.t.Fatalf("read trim reply: %v", err) + } + errCode := binary.BigEndian.Uint32(reply[4:8]) + if errCode != 0 { + c.t.Fatalf("trim error: %d", errCode) + } +} + +func (c *testNBDClient) flush() { + c.t.Helper() + handle := c.nextHandle + c.nextHandle++ + + var req [28]byte + binary.BigEndian.PutUint32(req[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(req[4:6], 0) + binary.BigEndian.PutUint16(req[6:8], uint16(nbdCmdFlush)) + binary.BigEndian.PutUint64(req[8:16], handle) + if _, err := c.conn.Write(req[:]); err != nil { + c.t.Fatalf("write flush request: %v", err) + } + + var reply [16]byte + if _, err := io.ReadFull(c.conn, reply[:]); err != nil { + c.t.Fatalf("read flush reply: %v", err) + } + errCode := binary.BigEndian.Uint32(reply[4:8]) + if errCode != 0 { + c.t.Fatalf("flush error: %d", errCode) + } +} + +func (c *testNBDClient) disconnect() { + handle := c.nextHandle + c.nextHandle++ + + var req [28]byte + binary.BigEndian.PutUint32(req[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(req[4:6], 0) + binary.BigEndian.PutUint16(req[6:8], uint16(nbdCmdDisc)) + binary.BigEndian.PutUint64(req[8:16], handle) + c.conn.Write(req[:]) + c.conn.Close() +} + +// startTestServer creates a temp file, starts a server, and returns +// the address, server, and a cleanup function. +func startTestServer(t *testing.T, fileData []byte, pageSize int) (addr string, srv *Server, cleanup func()) { + t.Helper() + + tmpFile, err := os.CreateTemp("", "guestbd-test-*") + if err != nil { + t.Fatalf("create temp file: %v", err) + } + if _, err := tmpFile.Write(fileData); err != nil { + t.Fatalf("write temp file: %v", err) + } + tmpFile.Close() + + return startTestServerFile(t, tmpFile.Name(), pageSize) +} + +func startTestServerFile(t *testing.T, filePath string, pageSize int) (addr string, srv *Server, cleanup func()) { + t.Helper() + + srv = NewServer(FileSource(filePath), WithPageSize(pageSize), WithMaxMem(int64(pageSize)*256)) + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go srv.HandleConn(conn) + } + }() + + return ln.Addr().String(), srv, func() { + ln.Close() + os.Remove(filePath) + } +} + +func TestReadBasic(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*4) + for i := range data { + data[i] = byte(i % 251) + } + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + if c.exportSize != uint64(len(data)) { + t.Fatalf("export size: got %d, want %d", c.exportSize, len(data)) + } + + // Read full first page. + got := c.read(0, pageSize) + if !bytes.Equal(got, data[:pageSize]) { + t.Fatal("first page mismatch") + } + + // Read across page boundary. + off := uint64(pageSize - 100) + got = c.read(off, 200) + if !bytes.Equal(got, data[off:off+200]) { + t.Fatal("cross-page read mismatch") + } + + // Read last page. + off = uint64(pageSize * 3) + got = c.read(off, pageSize) + if !bytes.Equal(got, data[off:off+pageSize]) { + t.Fatal("last page mismatch") + } +} + +func TestReadAtAlignment(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*6) + for i := range data { + data[i] = byte(i % 251) + } + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + tests := []struct { + name string + off uint64 + length uint32 + }{ + // Within a single page, not at offset 0. + {"mid-page-short-read", 100, 50}, + // Page-aligned start, sub-page length. + {"aligned-start-sub-page", 0, 100}, + // Starts mid-page, ends at exact page boundary. + {"mid-to-page-boundary", 100, pageSize - 100}, + // Starts at page boundary, ends mid-next-page. + {"page-boundary-into-next", pageSize, pageSize + 200}, + // Unaligned start and end spanning 3 pages: + // partial first + full middle + partial last. + {"unaligned-multi-page", 100, pageSize*2 + 200}, + // Multiple full aligned pages. + {"multi-full-pages", 0, pageSize * 3}, + // Single byte at page boundary. + {"one-byte-at-boundary", pageSize, 1}, + // Single byte just before page boundary. + {"one-byte-before-boundary", pageSize - 1, 1}, + // Single byte just after page boundary. + {"one-byte-after-boundary", pageSize + 1, 1}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := c.read(tt.off, tt.length) + want := data[tt.off : tt.off+uint64(tt.length)] + if !bytes.Equal(got, want) { + t.Fatalf("off=%d len=%d: got %x... want %x...", + tt.off, tt.length, got[:min(len(got), 16)], want[:min(len(want), 16)]) + } + }) + } +} + +func TestWriteAndReadBack(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*4) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Write a full page. + writeData := make([]byte, pageSize) + for i := range writeData { + writeData[i] = 0xAB + } + c.write(0, writeData) + + // Read it back. + got := c.read(0, pageSize) + if !bytes.Equal(got, writeData) { + t.Fatal("written page mismatch") + } + + // Second page should still be zeros. + got = c.read(pageSize, pageSize) + if !bytes.Equal(got, make([]byte, pageSize)) { + t.Fatal("unwritten page should be zeros") + } +} + +func TestSubPageWrite(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize) + for i := range data { + data[i] = byte(i % 256) + } + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Write 100 bytes in the middle of the page. + patch := bytes.Repeat([]byte{0xFF}, 100) + c.write(200, patch) + + // Read back full page. + got := c.read(0, pageSize) + + // Build expected result. + expected := make([]byte, pageSize) + copy(expected, data) + copy(expected[200:], patch) + if !bytes.Equal(got, expected) { + t.Fatal("sub-page write mismatch") + } +} + +func TestCrossPageWrite(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*2) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Write data spanning page boundary. + writeData := make([]byte, 200) + for i := range writeData { + writeData[i] = byte(i) + } + c.write(uint64(pageSize-100), writeData) + + // Read back the region. + got := c.read(uint64(pageSize-100), 200) + if !bytes.Equal(got, writeData) { + t.Fatal("cross-page write mismatch") + } +} + +func TestWritesLostOnDisconnect(t *testing.T) { + const pageSize = 4096 + original := make([]byte, pageSize) + for i := range original { + original[i] = byte(i) + } + + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + + // First connection: write some data. + c1 := newTestClient(t, addr) + patch := bytes.Repeat([]byte{0xDE}, pageSize) + c1.write(0, patch) + got := c1.read(0, pageSize) + if !bytes.Equal(got, patch) { + t.Fatal("write not reflected") + } + c1.disconnect() + + // Second connection: should see original data. + c2 := newTestClient(t, addr) + defer c2.disconnect() + got = c2.read(0, pageSize) + if !bytes.Equal(got, original) { + t.Fatal("writes should be lost after disconnect") + } +} + +func TestTrim(t *testing.T) { + const pageSize = 4096 + original := make([]byte, pageSize) + for i := range original { + original[i] = byte(i % 256) + } + + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Write new data. + writeData := bytes.Repeat([]byte{0xFF}, pageSize) + c.write(0, writeData) + got := c.read(0, pageSize) + if !bytes.Equal(got, writeData) { + t.Fatal("write not reflected") + } + + // Trim the page. + c.trim(0, pageSize) + + // Should revert to original data. + got = c.read(0, pageSize) + if !bytes.Equal(got, original) { + t.Fatal("trim should revert to original") + } +} + +func TestTrimAlignment(t *testing.T) { + const pageSize = 4096 + // Non-zero original so we can distinguish "reverted to base" from "zeroed". + original := make([]byte, pageSize*4) + for i := range original { + original[i] = byte(i % 251) + } + + t.Run("mid-page", func(t *testing.T) { + // Trim within a single dirty page: trimmed bytes become zero, + // rest of page stays dirty. + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + c := newTestClient(t, addr) + defer c.disconnect() + + c.write(0, bytes.Repeat([]byte{0xFF}, pageSize)) + c.trim(100, 200) // zero bytes 100..299 + + got := c.read(0, pageSize) + expected := bytes.Repeat([]byte{0xFF}, pageSize) + for i := 100; i < 300; i++ { + expected[i] = 0 + } + if !bytes.Equal(got, expected) { + t.Fatal("mid-page trim mismatch") + } + }) + + t.Run("cross-boundary-partial", func(t *testing.T) { + // Trim crossing page boundary, not fully covering either page. + // Both partial pages get zero-writes. + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + c := newTestClient(t, addr) + defer c.disconnect() + + c.write(0, bytes.Repeat([]byte{0xFF}, pageSize*2)) + c.trim(uint64(pageSize-100), 200) // zero last 100 bytes of page 0, first 100 of page 1 + + got := c.read(0, pageSize*2) + expected := bytes.Repeat([]byte{0xFF}, pageSize*2) + for i := pageSize - 100; i < pageSize+100; i++ { + expected[i] = 0 + } + if !bytes.Equal(got, expected) { + t.Fatal("cross-boundary trim mismatch") + } + }) + + t.Run("partial-full-partial", func(t *testing.T) { + // Trim partially covers first and last page, fully covers middle. + // Partial pages get zero-writes; full middle page reverts to original. + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + c := newTestClient(t, addr) + defer c.disconnect() + + c.write(0, bytes.Repeat([]byte{0xFF}, pageSize*3)) + c.trim(uint64(pageSize-100), pageSize+200) // partial page 0, full page 1, partial page 2 + + got := c.read(0, pageSize*3) + expected := bytes.Repeat([]byte{0xFF}, pageSize*3) + // Partial end of page 0: zeroed. + for i := pageSize - 100; i < pageSize; i++ { + expected[i] = 0 + } + // Full page 1: reverts to original. + copy(expected[pageSize:pageSize*2], original[pageSize:pageSize*2]) + // Partial start of page 2: zeroed. + for i := pageSize * 2; i < pageSize*2+100; i++ { + expected[i] = 0 + } + if !bytes.Equal(got, expected) { + t.Fatal("partial-full-partial trim mismatch") + } + }) + + t.Run("exact-full-page", func(t *testing.T) { + // Trim exactly covers one page: reverts to base image. + addr, _, cleanup := startTestServer(t, original, pageSize) + defer cleanup() + c := newTestClient(t, addr) + defer c.disconnect() + + c.write(0, bytes.Repeat([]byte{0xFF}, pageSize*3)) + c.trim(uint64(pageSize), uint32(pageSize)) + + got := c.read(0, pageSize*3) + expected := bytes.Repeat([]byte{0xFF}, pageSize*3) + // Page 1 reverts to original. + copy(expected[pageSize:pageSize*2], original[pageSize:pageSize*2]) + if !bytes.Equal(got, expected) { + t.Fatal("exact-full-page trim mismatch") + } + }) +} + +func TestFlush(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + c.write(0, bytes.Repeat([]byte{1}, pageSize)) + c.flush() // should not error +} + +func TestConcurrentConnections(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*2) + for i := range data { + data[i] = byte(i % 256) + } + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + // Multiple concurrent connections, each with independent writes. + var wg sync.WaitGroup + for i := 0; i < 5; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + c := newTestClient(t, addr) + defer c.disconnect() + + // Each connection writes a different byte pattern. + pattern := bytes.Repeat([]byte{byte(idx)}, pageSize) + c.write(0, pattern) + + got := c.read(0, pageSize) + if !bytes.Equal(got, pattern) { + t.Errorf("conn %d: write/read mismatch", idx) + } + + // Second page should still be original. + got = c.read(pageSize, pageSize) + if !bytes.Equal(got, data[pageSize:]) { + t.Errorf("conn %d: second page should be original", idx) + } + }(i) + } + wg.Wait() +} + +func TestContentDedup(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*4) + + addr, srv, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Write the same content to two different pages. + pattern := make([]byte, pageSize) + rand.Read(pattern) + + c.write(0, pattern) + c.write(pageSize, pattern) + + // Both pages should read back correctly. + got1 := c.read(0, pageSize) + got2 := c.read(pageSize, pageSize) + if !bytes.Equal(got1, pattern) || !bytes.Equal(got2, pattern) { + t.Fatal("dedup read mismatch") + } + + // The cache should have only one entry for this content + // (plus potentially zero-page entries for the other pages). + h := hashPage(pattern) + if _, ok := srv.cache.Get(h); !ok { + t.Fatal("expected pattern to be in cache") + } +} + +func TestZeroPage(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + got := c.read(0, pageSize) + if !bytes.Equal(got, make([]byte, pageSize)) { + t.Fatal("zero page should be all zeros") + } +} + +func TestLargeFile(t *testing.T) { + const pageSize = 4096 + const numPages = 100 + data := make([]byte, pageSize*numPages) + rand.Read(data) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + if c.exportSize != uint64(len(data)) { + t.Fatalf("export size: got %d, want %d", c.exportSize, len(data)) + } + + // Read a few random pages and verify. + for _, pageIdx := range []int{0, 1, 50, 99} { + off := uint64(pageIdx * pageSize) + got := c.read(off, pageSize) + if !bytes.Equal(got, data[off:off+pageSize]) { + t.Fatalf("page %d mismatch", pageIdx) + } + } +} + +func TestNonPageAlignedFile(t *testing.T) { + const pageSize = 4096 + // File size not a multiple of page size. + data := make([]byte, pageSize+500) + for i := range data { + data[i] = byte(i % 256) + } + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Read past the file's actual data (within the last page). + // The remainder should be zero-filled. + got := c.read(pageSize, pageSize) + expected := make([]byte, pageSize) + copy(expected, data[pageSize:]) + if !bytes.Equal(got, expected) { + t.Fatal("non-aligned last page mismatch") + } +} + +func TestHandshakeExportName(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize) + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + // Manual handshake using NBD_OPT_EXPORT_NAME instead of NBD_OPT_GO. + conn, err := net.Dial("tcp", addr) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + + var handshake [18]byte + if _, err := io.ReadFull(conn, handshake[:]); err != nil { + t.Fatalf("read handshake: %v", err) + } + + // Send client flags (fixed newstyle, no zeroes). + var clientFlags [4]byte + binary.BigEndian.PutUint32(clientFlags[:], uint32(nbdFlagCFixedNewstyle|nbdFlagCNoZeroes)) + if _, err := conn.Write(clientFlags[:]); err != nil { + t.Fatalf("write client flags: %v", err) + } + + // Send NBD_OPT_EXPORT_NAME. + var optBuf [16]byte + binary.BigEndian.PutUint64(optBuf[0:8], nbdOptsMagic) + binary.BigEndian.PutUint32(optBuf[8:12], nbdOptExportName) + binary.BigEndian.PutUint32(optBuf[12:16], 0) // empty name + if _, err := conn.Write(optBuf[:]); err != nil { + t.Fatalf("write opt: %v", err) + } + + // Read reply: size(8) + flags(2), no zeroes since we set the flag. + var reply [10]byte + if _, err := io.ReadFull(conn, reply[:]); err != nil { + t.Fatalf("read export reply: %v", err) + } + exportSize := binary.BigEndian.Uint64(reply[0:8]) + if exportSize != uint64(pageSize) { + t.Fatalf("export size: got %d, want %d", exportSize, pageSize) + } +} + +func TestInodeSharing(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*2) + rand.Read(data) + + addr, srv, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + // Two connections to the same file should share the same baseImageState. + c1 := newTestClient(t, addr) + defer c1.disconnect() + c2 := newTestClient(t, addr) + defer c2.disconnect() + + // Read same page from both; second should hit cache. + c1.read(0, pageSize) + c2.read(0, pageSize) + + srv.mu.Lock() + hasRO := len(srv.roFiles) > 0 + srv.mu.Unlock() + + if !hasRO { + t.Fatal("expected roFiles to be non-empty") + } +} + +func TestReconnectHitsCache(t *testing.T) { + const pageSize = 4096 + const numPages = 20 + data := make([]byte, pageSize*numPages) + rand.Read(data) + + addr, srv, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + // First connection: read all pages, populating the hash table and cache. + c1 := newTestClient(t, addr) + for i := 0; i < numPages; i++ { + c1.read(uint64(i*pageSize), pageSize) + } + c1.disconnect() + + // Snapshot read path counters after first connection. + diskColdBefore := srv.readPath.Get("base_disk_cold").Value() + diskMissBefore := srv.readPath.Get("base_disk_miss").Value() + baseMemBefore := srv.readPath.Get("base_mem").Value() + + // Second connection after full disconnect: FileSource provides identity + // so the idle baseImageState is matched and reused, preserving the page + // hash table. All reads should come from cache (0 cold reads). + c2 := newTestClient(t, addr) + defer c2.disconnect() + for i := 0; i < numPages; i++ { + got := c2.read(uint64(i*pageSize), pageSize) + want := data[i*pageSize : (i+1)*pageSize] + if !bytes.Equal(got, want) { + t.Fatalf("page %d mismatch on reconnect", i) + } + } + + diskColdAfter := srv.readPath.Get("base_disk_cold").Value() + diskMissAfter := srv.readPath.Get("base_disk_miss").Value() + baseMemAfter := srv.readPath.Get("base_mem").Value() + + newCold := diskColdAfter - diskColdBefore + newMiss := diskMissAfter - diskMissBefore + newMem := baseMemAfter - baseMemBefore + + if newCold != 0 { + t.Errorf("expected 0 base_disk_cold reads, got %d", newCold) + } + if newMiss != 0 { + t.Errorf("expected 0 base_disk_miss reads, got %d", newMiss) + } + if newMem != int64(numPages) { + t.Errorf("expected %d base_mem reads, got %d", numPages, newMem) + } +} + +func TestConcurrentSharesHashTable(t *testing.T) { + const pageSize = 4096 + const numPages = 20 + data := make([]byte, pageSize*numPages) + rand.Read(data) + + addr, srv, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + // First connection: read all pages, populating the hash table. + c1 := newTestClient(t, addr) + for i := 0; i < numPages; i++ { + c1.read(uint64(i*pageSize), pageSize) + } + + // Snapshot read path counters while c1 is still connected. + diskColdBefore := srv.readPath.Get("base_disk_cold").Value() + diskMissBefore := srv.readPath.Get("base_disk_miss").Value() + baseMemBefore := srv.readPath.Get("base_mem").Value() + + // Second connection while c1 is alive: shares the baseImageState + // and its page hash table, so all reads come from cache. + c2 := newTestClient(t, addr) + for i := 0; i < numPages; i++ { + got := c2.read(uint64(i*pageSize), pageSize) + want := data[i*pageSize : (i+1)*pageSize] + if !bytes.Equal(got, want) { + t.Fatalf("page %d mismatch", i) + } + } + c2.disconnect() + c1.disconnect() + + diskColdAfter := srv.readPath.Get("base_disk_cold").Value() + diskMissAfter := srv.readPath.Get("base_disk_miss").Value() + baseMemAfter := srv.readPath.Get("base_mem").Value() + + newCold := diskColdAfter - diskColdBefore + newMiss := diskMissAfter - diskMissBefore + newMem := baseMemAfter - baseMemBefore + + if newCold != 0 { + t.Errorf("expected 0 base_disk_cold reads, got %d", newCold) + } + if newMiss != 0 { + t.Errorf("expected 0 base_disk_miss reads, got %d", newMiss) + } + if newMem != int64(numPages) { + t.Errorf("expected %d base_mem reads, got %d", numPages, newMem) + } +} + +// noKeyBaseImage wraps a BaseImage and returns nil from BaseImageKey, +// disabling equivalence-keyed caching. +type noKeyBaseImage struct { + BaseImage +} + +func (n *noKeyBaseImage) BaseImageKey() any { return nil } + +// waitBaseImagesIdle waits for srv to have no active base images, as +// happens once it has finished handling all disconnected clients. +func waitBaseImagesIdle(t *testing.T, srv *Server) { + t.Helper() + deadline := time.Now().Add(10 * time.Second) + for srv.baseImagesActive.Value() != 0 { + if time.Now().After(deadline) { + t.Fatalf("timeout waiting for base images to go idle; %d active", srv.baseImagesActive.Value()) + } + time.Sleep(time.Millisecond) + } +} + +func TestReconnectNoIdentity(t *testing.T) { + const pageSize = 4096 + const numPages = 20 + data := make([]byte, pageSize*numPages) + rand.Read(data) + + tmpFile, err := os.CreateTemp("", "guestbd-test-*") + if err != nil { + t.Fatalf("create temp file: %v", err) + } + if _, err := tmpFile.Write(data); err != nil { + t.Fatalf("write temp file: %v", err) + } + tmpFile.Close() + defer os.Remove(tmpFile.Name()) + + // Use a BaseImageSource that wraps FileSource's output with nil key. + inner := FileSource(tmpFile.Name()) + noKeySource := func() (BaseImage, error) { + base, err := inner() + if err != nil { + return nil, err + } + return &noKeyBaseImage{base}, nil + } + + srv := NewServer(noKeySource, WithPageSize(pageSize), WithMaxMem(int64(pageSize)*256)) + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go srv.HandleConn(conn) + } + }() + defer ln.Close() + addr := ln.Addr().String() + + // First connection: read all pages. + c1 := newTestClient(t, addr) + for i := 0; i < numPages; i++ { + c1.read(uint64(i*pageSize), pageSize) + } + c1.disconnect() + + // Wait for the server to release the first connection's base image. + // Until then it is still active and would be shared, not replaced. + waitBaseImagesIdle(t, srv) + + // Snapshot counters. + diskColdBefore := srv.readPath.Get("base_disk_cold").Value() + + // Second connection: nil key means the old idle roFile is replaced, + // so all reads should be cold. + c2 := newTestClient(t, addr) + defer c2.disconnect() + for i := 0; i < numPages; i++ { + got := c2.read(uint64(i*pageSize), pageSize) + want := data[i*pageSize : (i+1)*pageSize] + if !bytes.Equal(got, want) { + t.Fatalf("page %d mismatch on reconnect", i) + } + } + + diskColdAfter := srv.readPath.Get("base_disk_cold").Value() + newCold := diskColdAfter - diskColdBefore + if newCold != int64(numPages) { + t.Errorf("expected %d base_disk_cold reads (no identity), got %d", numPages, newCold) + } +} + +func TestBaseImageReplaced(t *testing.T) { + if runtime.GOOS == "windows" { + // The server keeps the idle base image open for reuse, and + // Windows can't rename over a file that's open. + t.Skip("can't replace an open file by rename on Windows") + } + const pageSize = 4096 + const numPages = 4 + data1 := make([]byte, pageSize*numPages) + rand.Read(data1) + + dir := t.TempDir() + filePath := filepath.Join(dir, "disk.raw") + if err := os.WriteFile(filePath, data1, 0644); err != nil { + t.Fatal(err) + } + + addr, _, cleanup := startTestServerFile(t, filePath, pageSize) + defer cleanup() + + // First connection: read all pages. + c1 := newTestClient(t, addr) + for i := 0; i < numPages; i++ { + got := c1.read(uint64(i*pageSize), pageSize) + if !bytes.Equal(got, data1[i*pageSize:(i+1)*pageSize]) { + t.Fatalf("page %d mismatch (first conn)", i) + } + } + c1.disconnect() + + // Replace the file: write new content to a temp file and rename + // (new inode → different identity key). + data2 := make([]byte, pageSize*numPages) + rand.Read(data2) + tmpPath := filepath.Join(dir, "disk.raw.new") + if err := os.WriteFile(tmpPath, data2, 0644); err != nil { + t.Fatal(err) + } + if err := os.Rename(tmpPath, filePath); err != nil { + t.Fatal(err) + } + + // Second connection: should see the new data. + c2 := newTestClient(t, addr) + defer c2.disconnect() + for i := 0; i < numPages; i++ { + got := c2.read(uint64(i*pageSize), pageSize) + if !bytes.Equal(got, data2[i*pageSize:(i+1)*pageSize]) { + t.Fatalf("page %d mismatch after file replacement: got %x... want %x...", + i, got[:8], data2[i*pageSize:i*pageSize+8]) + } + } +} + +func BenchmarkRead(b *testing.B) { + const pageSize = 4096 + data := make([]byte, pageSize*100) + rand.Read(data) + + tmpFile, err := os.CreateTemp("", "guestbd-bench-*") + if err != nil { + b.Fatal(err) + } + tmpFile.Write(data) + tmpFile.Close() + defer os.Remove(tmpFile.Name()) + + srv := NewServer(FileSource(tmpFile.Name()), WithPageSize(pageSize), WithMaxMem(int64(pageSize)*256)) + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + b.Fatal(err) + } + defer ln.Close() + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go srv.HandleConn(conn) + } + }() + + conn, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + b.Fatal(err) + } + c := &testNBDClient{t: nil, conn: conn} + c.t = &testing.T{} // unused, but handshake needs it + // Skip handshake in benchmark... actually let's do it properly + // Just inline the handshake since testNBDClient needs *testing.T + + // For benchmark, do the handshake manually. + var handshake [18]byte + io.ReadFull(conn, handshake[:]) + var clientFlags [4]byte + binary.BigEndian.PutUint32(clientFlags[:], uint32(nbdFlagCFixedNewstyle|nbdFlagCNoZeroes)) + conn.Write(clientFlags[:]) + var optBuf [16]byte + binary.BigEndian.PutUint64(optBuf[0:8], nbdOptsMagic) + binary.BigEndian.PutUint32(optBuf[8:12], nbdOptExportName) + binary.BigEndian.PutUint32(optBuf[12:16], 0) + conn.Write(optBuf[:]) + var exportReply [10]byte + io.ReadFull(conn, exportReply[:]) + + b.ResetTimer() + for i := 0; i < b.N; i++ { + handle := uint64(i) + var req [28]byte + binary.BigEndian.PutUint32(req[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(req[6:8], uint16(nbdCmdRead)) + binary.BigEndian.PutUint64(req[8:16], handle) + offset := uint64((i % 100) * pageSize) + binary.BigEndian.PutUint64(req[16:24], offset) + binary.BigEndian.PutUint32(req[24:28], pageSize) + if _, err := conn.Write(req[:]); err != nil { + b.Fatal(err) + } + + var reply [16]byte + if _, err := io.ReadFull(conn, reply[:]); err != nil { + b.Fatal(err) + } + replyData := make([]byte, pageSize) + if _, err := io.ReadFull(conn, replyData); err != nil { + b.Fatal(err) + } + } + b.SetBytes(pageSize) + + // Send disconnect. + var disc [28]byte + binary.BigEndian.PutUint32(disc[0:4], nbdRequestMagic) + binary.BigEndian.PutUint16(disc[6:8], uint16(nbdCmdDisc)) + conn.Write(disc[:]) + conn.Close() +} + +func TestPageCacheLRU(t *testing.T) { + cache := newPageCache(3, 4096) + + pages := make([][]byte, 5) + hashes := make([]pageHash, 5) + for i := range pages { + pages[i] = make([]byte, 4096) + pages[i][0] = byte(i + 1) + hashes[i] = hashPage(pages[i]) + } + + // Fill cache. + cache.Put(hashes[0], pages[0]) + cache.Put(hashes[1], pages[1]) + cache.Put(hashes[2], pages[2]) + + // All three should be present. + for i := 0; i < 3; i++ { + if _, ok := cache.Get(hashes[i]); !ok { + t.Fatalf("page %d should be in cache", i) + } + } + + // Add a 4th; oldest (page 0) should be evicted. + // But first access page 0 to make it recent. + cache.Get(hashes[0]) + cache.Put(hashes[3], pages[3]) + // Now page 1 should be evicted (it's the LRU). + if _, ok := cache.Get(hashes[1]); ok { + t.Fatal("page 1 should have been evicted") + } + if _, ok := cache.Get(hashes[0]); !ok { + t.Fatal("page 0 should still be in cache (was recently accessed)") + } +} + +func TestNoCacheMode(t *testing.T) { + const pageSize = 4096 + data := make([]byte, pageSize*4) + for i := range data { + data[i] = byte(i % 251) + } + + tmpFile, err := os.CreateTemp("", "guestbd-test-*") + if err != nil { + t.Fatal(err) + } + if _, err := tmpFile.Write(data); err != nil { + t.Fatal(err) + } + tmpFile.Close() + defer os.Remove(tmpFile.Name()) + + srv := NewServer(FileSource(tmpFile.Name()), WithPageSize(pageSize), WithMaxMem(0)) + if srv.cache != nil { + t.Fatal("expected nil cache with WithMaxMem(0)") + } + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go srv.HandleConn(conn) + } + }() + addr := ln.Addr().String() + + c := newTestClient(t, addr) + defer c.disconnect() + + // Read base image data. + got := c.read(0, pageSize) + if !bytes.Equal(got, data[:pageSize]) { + t.Fatal("first page mismatch") + } + + // Full page write and read back. + writeData := make([]byte, pageSize) + for i := range writeData { + writeData[i] = 0xAB + } + c.write(0, writeData) + got = c.read(0, pageSize) + if !bytes.Equal(got, writeData) { + t.Fatal("written page mismatch") + } + + // Sub-page write (read-modify-write without cache). + patch := bytes.Repeat([]byte{0xFF}, 100) + c.write(pageSize+200, patch) + got = c.read(pageSize, pageSize) + expected := make([]byte, pageSize) + copy(expected, data[pageSize:2*pageSize]) + copy(expected[200:], patch) + if !bytes.Equal(got, expected) { + t.Fatal("sub-page write mismatch") + } + + // Trim reverts to base image. + c.trim(0, pageSize) + got = c.read(0, pageSize) + if !bytes.Equal(got, data[:pageSize]) { + t.Fatal("trim should revert to original") + } + + // Unwritten page still reads from base. + got = c.read(uint64(pageSize*3), pageSize) + if !bytes.Equal(got, data[pageSize*3:]) { + t.Fatal("unwritten page mismatch") + } +} + +func TestMultiplePageSizes(t *testing.T) { + for _, pageSize := range []int{512, 1024, 4096, 8192} { + t.Run(fmt.Sprintf("pageSize=%d", pageSize), func(t *testing.T) { + data := make([]byte, pageSize*4) + for i := range data { + data[i] = byte(i % 251) + } + + addr, _, cleanup := startTestServer(t, data, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + got := c.read(0, uint32(pageSize)) + if !bytes.Equal(got, data[:pageSize]) { + t.Fatal("page mismatch") + } + + // Sub-page write. + patch := bytes.Repeat([]byte{0xCC}, 16) + c.write(100, patch) + got = c.read(0, uint32(pageSize)) + expected := make([]byte, pageSize) + copy(expected, data[:pageSize]) + copy(expected[100:], patch) + if !bytes.Equal(got, expected) { + t.Fatal("sub-page write mismatch") + } + }) + } +} + +func TestQcow2BaseImage(t *testing.T) { + qemuImg, err := exec.LookPath("qemu-img") + if err != nil { + t.Skip("qemu-img not found") + } + + const pageSize = 4096 + const diskSize = pageSize * 4 // 16KB virtual disk + + // Create a raw image with known data. + dir := t.TempDir() + rawPath := filepath.Join(dir, "test.raw") + data := make([]byte, diskSize) + for i := range data { + data[i] = byte(i % 251) + } + if err := os.WriteFile(rawPath, data, 0644); err != nil { + t.Fatal(err) + } + + // Convert to qcow2. + qcow2Path := filepath.Join(dir, "test.qcow2") + out, err := exec.Command(qemuImg, "convert", "-f", "raw", "-O", "qcow2", rawPath, qcow2Path).CombinedOutput() + if err != nil { + t.Fatalf("qemu-img convert: %v\n%s", err, out) + } + + addr, _, cleanup := startTestServerFile(t, qcow2Path, pageSize) + defer cleanup() + + c := newTestClient(t, addr) + defer c.disconnect() + + if c.exportSize != uint64(diskSize) { + t.Fatalf("export size: got %d, want %d", c.exportSize, diskSize) + } + + // Read full first page. + got := c.read(0, pageSize) + if !bytes.Equal(got, data[:pageSize]) { + t.Fatal("first page mismatch") + } + + // Read across page boundary. + off := uint64(pageSize - 100) + got = c.read(off, 200) + if !bytes.Equal(got, data[off:off+200]) { + t.Fatal("cross-page read mismatch") + } + + // Read last page. + off = uint64(pageSize * 3) + got = c.read(off, pageSize) + if !bytes.Equal(got, data[off:off+pageSize]) { + t.Fatal("last page mismatch") + } + + // Write and read back (ephemeral write on top of qcow2 base). + writeData := make([]byte, pageSize) + for i := range writeData { + writeData[i] = 0xAB + } + c.write(0, writeData) + got = c.read(0, pageSize) + if !bytes.Equal(got, writeData) { + t.Fatal("write on qcow2 mismatch") + } + + // Second page should still be original. + got = c.read(pageSize, pageSize) + if !bytes.Equal(got, data[pageSize:2*pageSize]) { + t.Fatal("unwritten page should still be original qcow2 data") + } +} diff --git a/guestbd/nbd.go b/guestbd/nbd.go new file mode 100644 index 0000000..bc7b487 --- /dev/null +++ b/guestbd/nbd.go @@ -0,0 +1,67 @@ +package guestbd + +// NBD protocol constants. +// See https://github.com/NetworkBlockDevice/nbd/blob/master/doc/proto.md + +const ( + // Handshake magic values + nbdMagic uint64 = 0x4e42444d41474943 // "NBDMAGIC" + nbdOptsMagic uint64 = 0x49484156454F5054 // "IHAVEOPT" + nbdRequestMagic uint32 = 0x25609513 + nbdReplyMagic uint32 = 0x67446698 + nbdOptReplyMagic uint64 = 0x3e889045565a9 + + // Handshake flags (server to client) + nbdFlagFixedNewstyle uint16 = 1 << 0 + nbdFlagNoZeroes uint16 = 1 << 1 + + // Client flags + nbdFlagCFixedNewstyle uint32 = 1 << 0 + nbdFlagCNoZeroes uint32 = 1 << 1 + + // Transmission flags + nbdFlagHasFlags uint16 = 1 << 0 + nbdFlagSendFlush uint16 = 1 << 2 + nbdFlagSendTrim uint16 = 1 << 5 + + // Option codes + nbdOptExportName uint32 = 1 + nbdOptAbort uint32 = 2 + nbdOptList uint32 = 3 + nbdOptInfo uint32 = 6 + nbdOptGo uint32 = 7 + + // Option reply types + nbdRepAck uint32 = 1 + nbdRepServer uint32 = 2 + nbdRepInfo uint32 = 3 + nbdRepErrUnsup uint32 = (1 << 31) + 1 + + // Info types + nbdInfoExport uint16 = 0 + nbdInfoBlockSize uint16 = 3 + + // Command types + nbdCmdRead uint16 = 0 + nbdCmdWrite uint16 = 1 + nbdCmdDisc uint16 = 2 + nbdCmdFlush uint16 = 3 + nbdCmdTrim uint16 = 4 + + // Maximum payload size per the NBD spec. + nbdMaxPayload uint32 = 32 * 1024 * 1024 // 32 MB + + // Error codes (Linux errno values) + nbdEIO uint32 = 5 + nbdEINVAL uint32 = 22 +) + +// nbdRequest is the wire format for an NBD transmission request. +type nbdRequest struct { + Magic uint32 + Flags uint16 + Type uint16 + Handle uint64 + Offset uint64 + Length uint32 +} diff --git a/guestbd/persist_test.go b/guestbd/persist_test.go new file mode 100644 index 0000000..9747769 --- /dev/null +++ b/guestbd/persist_test.go @@ -0,0 +1,77 @@ +package guestbd + +import ( + "bytes" + "net" + "os" + "path/filepath" + "testing" +) + +// TestWriteDirtyTo checks that writing a shared snapshot's dirty pages onto +// a copy of the base image reproduces what clients see. +func TestWriteDirtyTo(t *testing.T) { + const pageSize = 4096 + base := make([]byte, 3*pageSize+100) // not a page multiple + for i := range base { + base[i] = byte(i % 251) + } + dir := t.TempDir() + basePath := filepath.Join(dir, "base.img") + if err := os.WriteFile(basePath, base, 0o644); err != nil { + t.Fatal(err) + } + srv := NewServer(FileSource(basePath), WithPageSize(pageSize), WithSharedSnapshot(), WithMaxMem(0)) + defer srv.Close() + if srv.SharedSnapshot() != nil { + t.Fatal("SharedSnapshot non-nil before any connection") + } + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + go srv.Serve(ln) + + c := newTestClient(t, ln.Addr().String()) + c.write(10, bytes.Repeat([]byte{0xaa}, 20)) // partial first page + c.write(pageSize-5, bytes.Repeat([]byte{0xbb}, 10)) // straddles pages 0 and 1 + c.write(3*pageSize+50, bytes.Repeat([]byte{0xcc}, 50)) // partial last page, to EOF + c.write(2*pageSize, bytes.Repeat([]byte{0xdd}, pageSize)) // whole page 2... + c.trim(2*pageSize, pageSize) // ...then trimmed back to base + want := c.read(0, uint32(len(base))) + c.disconnect() + + snap := srv.SharedSnapshot() + if snap == nil { + t.Fatal("SharedSnapshot nil after a connection") + } + outPath := filepath.Join(dir, "persisted.img") + if err := os.WriteFile(outPath, base, 0o644); err != nil { + t.Fatal(err) + } + f, err := os.OpenFile(outPath, os.O_WRONLY, 0) + if err != nil { + t.Fatal(err) + } + pages, err := snap.WriteDirtyTo(f) + if err != nil { + t.Fatal(err) + } + if err := f.Close(); err != nil { + t.Fatal(err) + } + if pages != 3 { + t.Errorf("wrote %d pages, want 3 (pages 0, 1 and 3; page 2 was trimmed)", pages) + } + got, err := os.ReadFile(outPath) + if err != nil { + t.Fatal(err) + } + if len(got) != len(base) { + t.Fatalf("persisted image is %d bytes, want %d", len(got), len(base)) + } + if !bytes.Equal(got, want) { + t.Fatal("persisted image differs from what the client read") + } +} diff --git a/guestbd/readonly.go b/guestbd/readonly.go new file mode 100644 index 0000000..41c3974 --- /dev/null +++ b/guestbd/readonly.go @@ -0,0 +1,116 @@ +package guestbd + +import ( + "io" + "sync" +) + +// baseImageState represents a shared read-only backing image. +// Multiple connections share one baseImageState. +// +// Locking: fields are protected by either mu or Server.mu as noted. +// base and size are immutable after construction and need no lock. +// When both locks are needed, Server.mu must be acquired first; +// mu is never held when acquiring Server.mu. +type baseImageState struct { + srv *Server // immutable + base BaseImage // immutable; used for page reads + size int64 // immutable; cached from base.Size() + identityKey any // immutable; non-nil if base returned a key; used as roFiles map key + + idleSince int64 // protected by Server.mu; monotonic seq for LRU eviction; 0 while active + + mu sync.Mutex + refcount int32 // protected by mu + // pageHashes is lazily computed per page. + // Absent from the map means the page has not been read from disk yet. + // A value equal to Server.zeroPageHash means the page is all zeros. + pageHashes map[int64]pageHash // protected by mu +} + +// newBaseImageState creates a new baseImageState for the given base image. +// The refcount starts at 1. +func newBaseImageState(srv *Server, base BaseImage, key any) *baseImageState { + bs := &baseImageState{ + srv: srv, + base: base, + size: base.Size(), + identityKey: key, + refcount: 1, + } + if srv.cache != nil { + bs.pageHashes = make(map[int64]pageHash) + } + return bs +} + +// readResult describes where a page read was served from. +type readResult int + +const ( + readFromCache readResult = iota // hash known, data in LRU cache + readFromDiskCold // hash unknown (first read of this page), read from disk + readFromDiskMiss // hash known but evicted from LRU cache, re-read from disk +) + +// readPage reads page n from the backing file into buf and returns its hash +// and how the read was served (cache hit, cold disk read, or cache miss disk read). +// buf must be at least pageSize bytes long or readPage panics. +// +// When the server has no page cache (WithMaxMem(0)), readPage skips hashing +// and cache operations and reads directly from the base image. +func (bs *baseImageState) readPage(buf []byte, n int64) (hash pageHash, result readResult, err error) { + cache := bs.srv.cache + pageSize := bs.srv.pageSize + _ = buf[pageSize-1] // bounds check hint; panics if too small + + if cache == nil { + // No caching; read straight from the base image. + offset := n * int64(pageSize) + nr, readErr := bs.base.ReadAt(buf[:pageSize], offset) + if readErr != nil && readErr != io.EOF { + return pageHash{}, 0, readErr + } + for i := nr; i < pageSize; i++ { + buf[i] = 0 + } + return pageHash{}, readFromDiskCold, nil + } + + bs.mu.Lock() + h, hashKnown := bs.pageHashes[n] + bs.mu.Unlock() + + if hashKnown { + // Already have the hash; try cache. + if d, ok := cache.Get(h); ok { + copy(buf, d) + return h, readFromCache, nil + } + } + + // Need to read from disk. + offset := n * int64(pageSize) + nr, readErr := bs.base.ReadAt(buf[:pageSize], offset) + if readErr != nil && readErr != io.EOF { + return pageHash{}, 0, readErr + } + // Zero-fill remainder (last page may be short). + for i := nr; i < pageSize; i++ { + buf[i] = 0 + } + + h = hashPage(buf[:pageSize]) + + bs.mu.Lock() + if _, ok := bs.pageHashes[n]; !ok { + bs.pageHashes[n] = h + } + bs.mu.Unlock() + + cache.Put(h, buf[:pageSize]) + if hashKnown { + return h, readFromDiskMiss, nil + } + return h, readFromDiskCold, nil +} diff --git a/guestbd/server.go b/guestbd/server.go new file mode 100644 index 0000000..928dcb4 --- /dev/null +++ b/guestbd/server.go @@ -0,0 +1,836 @@ +package guestbd + +import ( + "bufio" + "encoding/binary" + "expvar" + "fmt" + "io" + "log" + "net" + "os" + "slices" + "strings" + "sync" + "time" + + "github.com/bradfitz/qcow2" + "tailscale.com/metrics" + "tailscale.com/util/set" +) + +// BaseImage is a read-only random-access image with a known size. +// The server calls Close when the image is no longer needed. +// Implementations that return a non-nil key from BaseImageKey enable +// equivalence-keyed caching: idle baseImageStates whose key matches a +// newly opened image are reused, preserving the page hash table. +type BaseImage interface { + io.ReaderAt + io.Closer + Size() int64 + BaseImageKey() any // nil = no coalescing; non-nil must be comparable +} + +// BaseImageSource is a function that opens and returns the base image. +// It is called lazily when the first snapshot is created. +type BaseImageSource func() (BaseImage, error) + +// NewBaseImage returns a BaseImage wrapping an io.ReaderAt with a fixed size. +// The returned BaseImage has no identity key (BaseImageKey returns nil), +// so idle baseImageStates using it will not be reused across connections. +func NewBaseImage(r io.ReaderAt, size int64) BaseImage { + return &plainBaseImage{r: r, size: size} +} + +type plainBaseImage struct { + r io.ReaderAt + size int64 +} + +func (p *plainBaseImage) ReadAt(b []byte, off int64) (int, error) { return p.r.ReadAt(b, off) } +func (p *plainBaseImage) Size() int64 { return p.size } +func (p *plainBaseImage) Close() error { return nil } +func (p *plainBaseImage) BaseImageKey() any { return nil } + +// fileIdentityKey uniquely identifies an on-disk file by device and inode. +// See fileIdentity. +type fileIdentityKey struct { + dev uint64 + ino uint64 +} + +// fileSizeReaderAt wraps an *os.File as a BaseImage. +type fileSizeReaderAt struct { + f *os.File + size int64 + key any // from fileIdentity +} + +func (r *fileSizeReaderAt) ReadAt(p []byte, off int64) (int, error) { return r.f.ReadAt(p, off) } +func (r *fileSizeReaderAt) Size() int64 { return r.size } +func (r *fileSizeReaderAt) Close() error { return r.f.Close() } +func (r *fileSizeReaderAt) BaseImageKey() any { return r.key } + +// qcow2SizeReaderAt wraps a *qcow2.Image as a BaseImage, holding the +// underlying *os.File so it can be closed. +type qcow2SizeReaderAt struct { + img *qcow2.Image + f *os.File + key any // from fileIdentity +} + +func (r *qcow2SizeReaderAt) ReadAt(p []byte, off int64) (int, error) { return r.img.ReadAt(p, off) } +func (r *qcow2SizeReaderAt) Size() int64 { return r.img.Size() } +func (r *qcow2SizeReaderAt) Close() error { return r.f.Close() } +func (r *qcow2SizeReaderAt) BaseImageKey() any { return r.key } + +// FileSource returns a BaseImageSource that opens the file at path. +// Files ending in ".qcow2" are opened as qcow2 images; all others are +// treated as raw disk images. +func FileSource(path string) BaseImageSource { + return func() (BaseImage, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + fi, err := f.Stat() + if err != nil { + f.Close() + return nil, err + } + key := fileIdentity(f, fi) + if strings.HasSuffix(path, ".qcow2") { + img, err := qcow2.Open(f) + if err != nil { + f.Close() + return nil, fmt.Errorf("opening qcow2 image: %w", err) + } + return &qcow2SizeReaderAt{img, f, key}, nil + } + return &fileSizeReaderAt{f, fi.Size(), key}, nil + } +} + +// ServerOption configures optional Server parameters. +type ServerOption func(*serverConfig) + +type serverConfig struct { + pageSize int + maxMem int64 + sharedSnapshot bool + maxIdleBase int +} + +// WithPageSize sets the page size in bytes. It must be a positive power of two. +// The default is 4096. +func WithPageSize(n int) ServerOption { + return func(c *serverConfig) { c.pageSize = n } +} + +// WithMaxMem sets the maximum memory used by the shared page cache. +// The default is 1GB. A value of 0 disables page hashing and caching +// entirely; reads of non-dirty pages go directly to the base image. +func WithMaxMem(n int64) ServerOption { + return func(c *serverConfig) { c.maxMem = n } +} + +// WithSharedSnapshot makes all connections share a single writable Snapshot +// instead of each getting its own. This allows reconnecting clients to see +// writes from previous connections. +func WithSharedSnapshot() ServerOption { + return func(c *serverConfig) { c.sharedSnapshot = true } +} + +// WithMaxIdleBase sets the maximum number of idle base images to keep cached +// for equivalence-keyed reuse. When a baseImageState becomes idle (refcount 0) +// and its BaseImageKey is non-nil, it is kept in a cache of up to n entries. +// The default is 1. +func WithMaxIdleBase(n int) ServerOption { + return func(c *serverConfig) { c.maxIdleBase = n } +} + +// Server is an NBD server that serves a single backing image to multiple +// concurrent clients. Each client connection may get an independent +// copy-on-write Snapshot, or multiple connections may share a single Snapshot +// to allow reconnection to a previous writable state. The snapshot policy is +// configured via WithSharedSnapshot at construction time. +type Server struct { + // Immutable after construction: + getBase BaseImageSource // immutable + pageSize int // immutable + zeroPageHash pageHash // immutable; sha256 of an all-zero page + cache *pageCache // immutable; nil when WithMaxMem(0) disables caching + useSharedSnapshot bool // immutable; set by WithSharedSnapshot + maxIdleBase int // immutable; max idle entries in roFiles + pageBufPool sync.Pool // immutable; pool of *[]byte, each pageSize long + + mu sync.Mutex + sharedSnap *Snapshot // lazily created when useSharedSnapshot is true + roFiles map[any]*baseImageState // identity-keyed entries (active + idle) + roFileNoKey *baseImageState // single entry for nil-key sources + nextIdleSeq int64 // monotonic counter for idle LRU ordering + snapshots set.Set[*Snapshot] + + // Metrics (atomic, no mutex needed): + totalConns expvar.Int // counter_guestbd_total_conns + activeConns expvar.Int // gauge_guestbd_active_conns + ops metrics.LabelMap // counter_guestbd_nbd_ops{type=...} + readPath metrics.LabelMap // counter_guestbd_read_path{type=...} + readBytes expvar.Int // counter_guestbd_read_bytes + readPages expvar.Int // counter_guestbd_read_pages + writeBytes expvar.Int // counter_guestbd_write_bytes + writePages expvar.Int // counter_guestbd_write_pages + readSizeHist *metrics.Histogram // histogram_guestbd_read_size_bytes + writeSizeHist *metrics.Histogram // histogram_guestbd_write_size_bytes + readLatencyHist *metrics.Histogram // histogram_guestbd_read_latency_seconds + writeLatencyHist *metrics.Histogram // histogram_guestbd_write_latency_seconds + baseImagesActive expvar.Int // gauge_guestbd_base_images_active + baseImagesCached expvar.Int // gauge_guestbd_base_images_cached +} + +// latencyBuckets returns the histogram bucket boundaries (in seconds) +// for NBD read/write latency. The buckets are powers of two from 10µs +// through ~10s, covering cache-hit reads (microseconds), cold disk +// reads (milliseconds), and pathological slow ops. +func latencyBuckets() []float64 { + var b []float64 + for v := 10e-6; v < 16; v *= 2 { + b = append(b, v) + } + return b +} + +// NewServer creates a new Server that serves the base image returned by getBase. +func NewServer(getBase BaseImageSource, opts ...ServerOption) *Server { + cfg := serverConfig{ + pageSize: 4096, + maxMem: 1 << 30, // 1GB default + } + for _, o := range opts { + o(&cfg) + } + + pageSize := cfg.pageSize + + var cache *pageCache + if cfg.maxMem > 0 { + maxPages := int(cfg.maxMem) / pageSize + if maxPages < 1 { + maxPages = 1 + } + cache = newPageCache(maxPages, pageSize) + } + + // Powers of two from 1 to 4MB. + var sizeBuckets []float64 + for v := float64(1); v <= 4*1024*1024; v *= 2 { + sizeBuckets = append(sizeBuckets, v) + } + + maxIdleBase := cfg.maxIdleBase + if maxIdleBase == 0 { + maxIdleBase = 1 + } + + srv := &Server{ + getBase: getBase, + pageSize: pageSize, + useSharedSnapshot: cfg.sharedSnapshot, + cache: cache, + pageBufPool: sync.Pool{New: func() any { b := make([]byte, pageSize); return &b }}, + roFiles: make(map[any]*baseImageState), + maxIdleBase: maxIdleBase, + snapshots: make(set.Set[*Snapshot]), + ops: metrics.LabelMap{Label: "type"}, + readPath: metrics.LabelMap{Label: "type"}, + readSizeHist: metrics.NewHistogram(sizeBuckets), + writeSizeHist: metrics.NewHistogram(sizeBuckets), + readLatencyHist: metrics.NewHistogram(latencyBuckets()), + writeLatencyHist: metrics.NewHistogram(latencyBuckets()), + } + if cache != nil { + srv.zeroPageHash = hashPage(make([]byte, pageSize)) + } + + return srv +} + +// InitExpvar publishes the server's metrics to expvar. It should be called +// once before serving. It is not called in tests to avoid duplicate +// registration panics. +func (s *Server) InitExpvar() { + expvar.Publish("counter_guestbd_total_conns", &s.totalConns) + expvar.Publish("gauge_guestbd_active_conns", &s.activeConns) + expvar.Publish("counter_guestbd_nbd_ops", &s.ops) + expvar.Publish("counter_guestbd_read_path", &s.readPath) + expvar.Publish("counter_guestbd_read_bytes", &s.readBytes) + expvar.Publish("counter_guestbd_read_pages", &s.readPages) + expvar.Publish("counter_guestbd_write_bytes", &s.writeBytes) + expvar.Publish("counter_guestbd_write_pages", &s.writePages) + expvar.Publish("histogram_guestbd_read_size_bytes", s.readSizeHist) + expvar.Publish("histogram_guestbd_write_size_bytes", s.writeSizeHist) + expvar.Publish("histogram_guestbd_read_latency_seconds", s.readLatencyHist) + expvar.Publish("histogram_guestbd_write_latency_seconds", s.writeLatencyHist) + if s.cache != nil { + expvar.Publish("counter_guestbd_cache", &s.cache.path) + expvar.Publish("gauge_guestbd_cache_entries", &s.cache.entries) + expvar.Publish("gauge_guestbd_cache_bytes", &s.cache.bytes) + } + + expvar.Publish("gauge_guestbd_base_images_active", &s.baseImagesActive) + expvar.Publish("gauge_guestbd_base_images_cached", &s.baseImagesCached) + + ps := new(expvar.Int) + ps.Set(int64(s.pageSize)) + expvar.Publish("gauge_guestbd_page_size", ps) + + expvar.Publish("gauge_guestbd_max_dirty_bytes", expvar.Func(func() any { + s.mu.Lock() + defer s.mu.Unlock() + + var max int64 + for c := range s.snapshots { + c.mu.Lock() + n := int64(len(c.dirtyPages)) * int64(s.pageSize) + c.mu.Unlock() + if n > max { + max = n + } + } + return max + })) +} + +// getBaseImageState returns a baseImageState for the backing image. +// getBase is called on every invocation. If there is already an active or +// idle baseImageState with a matching BaseImageKey, the newly opened base is +// closed and the existing baseImageState is reused (preserving its page hash +// table). Sources returning nil keys share while active but are replaced +// when idle. +func (s *Server) getBaseImageState() (*baseImageState, error) { + base, err := s.getBase() + if err != nil { + return nil, err + } + + key := base.BaseImageKey() + + s.mu.Lock() + defer s.mu.Unlock() + + if key != nil { + if ro, ok := s.roFiles[key]; ok { + // Found a matching entry (active or idle). Close the new base + // and bump refcount. + base.Close() + ro.mu.Lock() + wasIdle := ro.refcount == 0 + ro.refcount++ + ro.mu.Unlock() + if wasIdle { + s.baseImagesActive.Add(1) + ro.idleSince = 0 + } + return ro, nil + } + // Not found — create new and insert. + ro := newBaseImageState(s, base, key) + s.roFiles[key] = ro + s.baseImagesActive.Add(1) + s.baseImagesCached.Add(1) + s.evictIdleLocked() + return ro, nil + } + + // key == nil + if ro := s.roFileNoKey; ro != nil { + ro.mu.Lock() + rc := ro.refcount + ro.mu.Unlock() + if rc > 0 { + // Active — share. + base.Close() + ro.mu.Lock() + ro.refcount++ + ro.mu.Unlock() + return ro, nil + } + // Idle — close old base and replace. + ro.base.Close() + s.baseImagesCached.Add(-1) + } + + ro := newBaseImageState(s, base, nil) + s.roFileNoKey = ro + s.baseImagesActive.Add(1) + s.baseImagesCached.Add(1) + return ro, nil +} + +// releaseBaseImageState decrements the refcount. When it reaches zero the +// baseImagesActive metric is updated. For keyed entries, the baseImageState +// stays in the roFiles map as an idle entry (subject to LRU eviction). +// For nil-key entries, it stays in roFileNoKey and is replaced next call. +func (s *Server) releaseBaseImageState(ro *baseImageState) { + ro.mu.Lock() + ro.refcount-- + rc := ro.refcount + ro.mu.Unlock() + + if rc > 0 { + return + } + + s.baseImagesActive.Add(-1) + + s.mu.Lock() + defer s.mu.Unlock() + + if ro.identityKey != nil { + s.nextIdleSeq++ + ro.idleSince = s.nextIdleSeq + s.evictIdleLocked() + } +} + +// evictIdleLocked removes the oldest idle entries from s.roFiles until +// the number of idle entries is at most s.maxIdleBase. +func (s *Server) evictIdleLocked() { + for { + var ( + idleCount int + oldestKey any + oldestSeq int64 + oldestEntry *baseImageState + ) + for k, ro := range s.roFiles { + ro.mu.Lock() + rc := ro.refcount + ro.mu.Unlock() + if rc == 0 { + idleCount++ + if oldestEntry == nil || ro.idleSince < oldestSeq { + oldestKey = k + oldestSeq = ro.idleSince + oldestEntry = ro + } + } + } + if idleCount <= s.maxIdleBase { + return + } + // Evict oldest idle. + oldestEntry.base.Close() + delete(s.roFiles, oldestKey) + s.baseImagesCached.Add(-1) + } +} + +// NewSnapshot creates a new writable Snapshot backed by the server's base +// image. The caller is responsible for calling Close when the snapshot is no +// longer needed. +func (s *Server) NewSnapshot() (*Snapshot, error) { + ro, err := s.getBaseImageState() + if err != nil { + return nil, fmt.Errorf("opening base image: %w", err) + } + snap, err := newSnapshot(s, ro) + if err != nil { + s.releaseBaseImageState(ro) + return nil, err + } + s.mu.Lock() + s.snapshots.Add(snap) + s.mu.Unlock() + return snap, nil +} + +// Close cleans up server resources, including any shared Snapshot created +// via WithSharedSnapshot and all cached baseImageStates. +func (s *Server) Close() error { + s.mu.Lock() + snap := s.sharedSnap + s.sharedSnap = nil + roFiles := s.roFiles + s.roFiles = nil + roNoKey := s.roFileNoKey + s.roFileNoKey = nil + s.mu.Unlock() + if snap != nil { + snap.Close() + } + for _, ro := range roFiles { + ro.base.Close() + } + if roNoKey != nil { + roNoKey.base.Close() + } + return nil +} + +// getOrCreateSnapshot returns a snapshot for a new connection. If the server +// is configured with WithSharedSnapshot, it lazily creates a single shared +// snapshot and returns it with owned=false. Otherwise it creates a fresh +// per-connection snapshot with owned=true. The caller must close the snapshot +// when owned is true. +func (s *Server) getOrCreateSnapshot() (snap *Snapshot, owned bool, err error) { + if !s.useSharedSnapshot { + snap, err = s.NewSnapshot() + return snap, true, err + } + + s.mu.Lock() + snap = s.sharedSnap + s.mu.Unlock() + if snap != nil { + return snap, false, nil + } + + // Lazy creation; must not hold s.mu across NewSnapshot. + snap, err = s.NewSnapshot() + if err != nil { + return nil, false, err + } + + s.mu.Lock() + if s.sharedSnap != nil { + // Another goroutine created it first. + s.mu.Unlock() + snap.Close() + return s.sharedSnap, false, nil + } + s.sharedSnap = snap + s.mu.Unlock() + return snap, false, nil +} + +// SharedSnapshot returns the Snapshot shared by all connections of a server +// created with WithSharedSnapshot, or nil if no client has connected yet +// (or the server doesn't share one). +func (s *Server) SharedSnapshot() *Snapshot { + s.mu.Lock() + defer s.mu.Unlock() + return s.sharedSnap +} + +// Serve accepts incoming connections on the listener ln, handling +// each one on a new goroutine. Serve blocks until the listener +// returns an error. The caller is responsible for closing ln. +func (s *Server) Serve(ln net.Listener) error { + for { + conn, err := ln.Accept() + if err != nil { + return err + } + go s.HandleConn(conn) + } +} + +// HandleConn handles a new TCP connection, performing the NBD +// handshake and entering the transmission phase. The snapshot policy +// (per-connection or shared) is determined by the server's configuration. +func (s *Server) HandleConn(nc net.Conn) { + s.totalConns.Add(1) + s.activeConns.Add(1) + defer s.activeConns.Add(-1) + + snap, owned, err := s.getOrCreateSnapshot() + if err != nil { + log.Printf("create snapshot for %v: %v", nc.RemoteAddr(), err) + nc.Close() + return + } + + defer func() { + if owned { + snap.Close() + } + nc.Close() + }() + + if err := s.serveNBD(snap, nc); err != nil { + log.Printf("nbd %v: %v", nc.RemoteAddr(), err) + } +} + +// serveNBD runs the NBD protocol on the given connection. +func (s *Server) serveNBD(snap *Snapshot, nc net.Conn) error { + r := nc + bw := bufio.NewWriterSize(nc, 1<<20) // 1MB + + // Phase 1: Newstyle handshake. + var handshake [18]byte + binary.BigEndian.PutUint64(handshake[0:8], nbdMagic) + binary.BigEndian.PutUint64(handshake[8:16], nbdOptsMagic) + binary.BigEndian.PutUint16(handshake[16:18], uint16(nbdFlagFixedNewstyle)) + if _, err := bw.Write(handshake[:]); err != nil { + return fmt.Errorf("writing handshake: %w", err) + } + if err := bw.Flush(); err != nil { + return fmt.Errorf("flushing handshake: %w", err) + } + + var clientFlagsBuf [4]byte + if _, err := io.ReadFull(r, clientFlagsBuf[:]); err != nil { + return fmt.Errorf("reading client flags: %w", err) + } + cflags := binary.BigEndian.Uint32(clientFlagsBuf[:]) + noZeroes := cflags&nbdFlagCNoZeroes != 0 + + // Phase 2: Option haggling. + exportSize := uint64(snap.roFile.size) + txFlags := nbdFlagHasFlags | nbdFlagSendFlush | nbdFlagSendTrim + + for { + var optHeader [16]byte + if _, err := io.ReadFull(r, optHeader[:]); err != nil { + return fmt.Errorf("reading option: %w", err) + } + optMagic := binary.BigEndian.Uint64(optHeader[0:8]) + if optMagic != nbdOptsMagic { + return fmt.Errorf("bad option magic: %#x", optMagic) + } + optCode := binary.BigEndian.Uint32(optHeader[8:12]) + optLen := binary.BigEndian.Uint32(optHeader[12:16]) + + optData := make([]byte, optLen) + if optLen > 0 { + if _, err := io.ReadFull(r, optData); err != nil { + return fmt.Errorf("reading option data: %w", err) + } + } + + switch optCode { + case nbdOptExportName: + var reply [10]byte + binary.BigEndian.PutUint64(reply[0:8], exportSize) + binary.BigEndian.PutUint16(reply[8:10], txFlags) + if _, err := bw.Write(reply[:]); err != nil { + return err + } + if !noZeroes { + var zeros [124]byte + if _, err := bw.Write(zeros[:]); err != nil { + return err + } + } + if err := bw.Flush(); err != nil { + return err + } + return s.serveTransmission(snap, r, bw) + + case nbdOptAbort: + if err := s.sendOptReply(bw, optCode, nbdRepAck, nil); err != nil { + return err + } + return bw.Flush() + + case nbdOptInfo, nbdOptGo: + // NBD_OPT_INFO gets the same replies as NBD_OPT_GO but stays in + // option haggling. Apple's Virtualization.framework client + // sends NBD_OPT_INFO before NBD_OPT_GO and gives up if it is + // unsupported. + // + // Send NBD_INFO_EXPORT. + var infoData [12]byte + binary.BigEndian.PutUint16(infoData[0:2], nbdInfoExport) + binary.BigEndian.PutUint64(infoData[2:10], exportSize) + binary.BigEndian.PutUint16(infoData[10:12], txFlags) + if err := s.sendOptReply(bw, optCode, nbdRepInfo, infoData[:]); err != nil { + return err + } + // Send NBD_INFO_BLOCK_SIZE. + var bsData [14]byte + binary.BigEndian.PutUint16(bsData[0:2], nbdInfoBlockSize) + binary.BigEndian.PutUint32(bsData[2:6], 1) // minimum + binary.BigEndian.PutUint32(bsData[6:10], uint32(s.pageSize)) // preferred + binary.BigEndian.PutUint32(bsData[10:14], 32*1024*1024) // maximum + if err := s.sendOptReply(bw, optCode, nbdRepInfo, bsData[:]); err != nil { + return err + } + // ACK + if err := s.sendOptReply(bw, optCode, nbdRepAck, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + if optCode == nbdOptInfo { + continue + } + return s.serveTransmission(snap, r, bw) + + case nbdOptList: + // One export with empty name (default). + var nameData [4]byte + if err := s.sendOptReply(bw, optCode, nbdRepServer, nameData[:]); err != nil { + return err + } + if err := s.sendOptReply(bw, optCode, nbdRepAck, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + + default: + if err := s.sendOptReply(bw, optCode, nbdRepErrUnsup, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + } + } +} + +// serveTransmission handles the NBD transmission phase. +func (s *Server) serveTransmission(snap *Snapshot, r io.Reader, bw *bufio.Writer) error { + var buf []byte + + for { + var req nbdRequest + if err := binary.Read(r, binary.BigEndian, &req); err != nil { + if err == io.EOF || err == io.ErrUnexpectedEOF { + return nil + } + return fmt.Errorf("reading request: %w", err) + } + if req.Magic != nbdRequestMagic { + return fmt.Errorf("bad request magic: %#x", req.Magic) + } + + if (req.Type == nbdCmdRead || req.Type == nbdCmdWrite) && req.Length > nbdMaxPayload { + return fmt.Errorf("request length %d exceeds NBD max payload %d", req.Length, nbdMaxPayload) + } + + switch req.Type { + case nbdCmdRead: + s.ops.Add("read", 1) + s.readSizeHist.Observe(float64(req.Length)) + buf = slices.Grow(buf[:0], int(req.Length))[:req.Length] + t0 := time.Now() + _, err := snap.ReadAt(buf, int64(req.Offset)) + s.readLatencyHist.Observe(time.Since(t0).Seconds()) + if err != nil { + if werr := s.sendReply(bw, req.Handle, nbdEIO, nil); werr != nil { + return werr + } + if werr := bw.Flush(); werr != nil { + return werr + } + continue + } + if err := s.sendReply(bw, req.Handle, 0, buf); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + + case nbdCmdWrite: + s.ops.Add("write", 1) + s.writeSizeHist.Observe(float64(req.Length)) + buf = slices.Grow(buf[:0], int(req.Length))[:req.Length] + if _, err := io.ReadFull(r, buf); err != nil { + return fmt.Errorf("reading write data: %w", err) + } + t0 := time.Now() + _, err := snap.WriteAt(buf, int64(req.Offset)) + s.writeLatencyHist.Observe(time.Since(t0).Seconds()) + if err != nil { + if werr := s.sendReply(bw, req.Handle, nbdEIO, nil); werr != nil { + return werr + } + if werr := bw.Flush(); werr != nil { + return werr + } + continue + } + if err := s.sendReply(bw, req.Handle, 0, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + + case nbdCmdDisc: + s.ops.Add("disconnect", 1) + return nil + + case nbdCmdFlush: + s.ops.Add("flush", 1) + // No-op: dirty data is ephemeral and lost on disconnect, + // so there's no point in syncing to disk for guest VM performance. + if err := s.sendReply(bw, req.Handle, 0, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + + case nbdCmdTrim: + s.ops.Add("trim", 1) + if err := snap.handleTrim(req.Offset, uint64(req.Length)); err != nil { + if werr := s.sendReply(bw, req.Handle, nbdEIO, nil); werr != nil { + return werr + } + if werr := bw.Flush(); werr != nil { + return werr + } + continue + } + if err := s.sendReply(bw, req.Handle, 0, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + + default: + if err := s.sendReply(bw, req.Handle, nbdEINVAL, nil); err != nil { + return err + } + if err := bw.Flush(); err != nil { + return err + } + } + } +} + +// sendOptReply writes an NBD option reply with the given option code, reply +// type, and optional payload data. +func (s *Server) sendOptReply(w io.Writer, optCode uint32, replyType uint32, data []byte) error { + var header [20]byte + binary.BigEndian.PutUint64(header[0:8], nbdOptReplyMagic) + binary.BigEndian.PutUint32(header[8:12], optCode) + binary.BigEndian.PutUint32(header[12:16], replyType) + binary.BigEndian.PutUint32(header[16:20], uint32(len(data))) + if _, err := w.Write(header[:]); err != nil { + return err + } + if len(data) > 0 { + if _, err := w.Write(data); err != nil { + return err + } + } + return nil +} + +// sendReply writes an NBD transmission-phase reply with the given handle, +// error code, and optional payload data. +func (s *Server) sendReply(w io.Writer, handle uint64, errCode uint32, data []byte) error { + var header [16]byte + binary.BigEndian.PutUint32(header[0:4], nbdReplyMagic) + binary.BigEndian.PutUint32(header[4:8], errCode) + binary.BigEndian.PutUint64(header[8:16], handle) + if _, err := w.Write(header[:]); err != nil { + return err + } + if len(data) > 0 { + if _, err := w.Write(data); err != nil { + return err + } + } + return nil +} diff --git a/guestbd/snapshot.go b/guestbd/snapshot.go new file mode 100644 index 0000000..c26633b --- /dev/null +++ b/guestbd/snapshot.go @@ -0,0 +1,304 @@ +package guestbd + +import ( + "fmt" + "io" + "maps" + "os" + "slices" + "sync" +) + +// Compile-time interface checks. +var ( + _ io.ReaderAt = (*Snapshot)(nil) + _ io.WriterAt = (*Snapshot)(nil) +) + +// Snapshot is a writable layer on top of a read-only base image. +// Reads fall through to the shared base image, while writes are tracked +// in the snapshot. A Snapshot can outlive individual TCP connections, +// allowing reconnections to reattach to the same writable state. +// +// Snapshot implements io.ReaderAt and io.WriterAt. +type Snapshot struct { + server *Server + roFile *baseImageState + + mu sync.Mutex + dirtyPages map[int64]dirtyPageInfo // pageNum => dirty info + dirtyFile *os.File // temp file for dirty page data + nextDirtyNum int // monotonically increasing +} + +// dirtyPageInfo tracks a written page in a snapshot. +type dirtyPageInfo struct { + hash pageHash + dirtyNum int // index into the snapshot's dirtyFile +} + +// newSnapshot creates a new Snapshot for the given read-only backing file. +// It creates an unlinked temporary file for dirty page storage. +func newSnapshot(srv *Server, ro *baseImageState) (*Snapshot, error) { + tmpFile, err := os.CreateTemp("", "guestbd-dirty-*") + if err != nil { + return nil, fmt.Errorf("creating temp file: %w", err) + } + // Unlink immediately so it's cleaned up when the fd closes. + os.Remove(tmpFile.Name()) + + return &Snapshot{ + server: srv, + roFile: ro, + dirtyPages: make(map[int64]dirtyPageInfo), + dirtyFile: tmpFile, + }, nil +} + +// Close releases the snapshot's dirty page storage and unregisters it from +// the server. It does not close any TCP connection using the snapshot. +func (c *Snapshot) Close() error { + c.mu.Lock() + if c.dirtyFile != nil { + c.dirtyFile.Close() + c.dirtyFile = nil + } + c.mu.Unlock() + + c.server.mu.Lock() + c.server.snapshots.Delete(c) + c.server.mu.Unlock() + + c.server.releaseBaseImageState(c.roFile) + return nil +} + +// readPageData reads a full page into buf, checking dirty pages first, +// then falling back to the readonly file. +// buf must be at least pageSize bytes long or readPageData panics. +func (c *Snapshot) readPageData(buf []byte, pageNum int64) error { + pageSize := c.server.pageSize + _ = buf[pageSize-1] // bounds check hint; panics if too small + + c.mu.Lock() + dp, dirty := c.dirtyPages[pageNum] + c.mu.Unlock() + + c.server.readPages.Add(1) + + if dirty { + c.server.readPath.Add("from_write", 1) + cache := c.server.cache + if cache != nil { + // Try the global cache first. + if data, ok := cache.Get(dp.hash); ok { + copy(buf, data) + return nil + } + } + // Read from the snapshot's dirty file. + _, err := c.dirtyFile.ReadAt(buf[:pageSize], int64(dp.dirtyNum)*int64(pageSize)) + if err != nil { + return fmt.Errorf("reading dirty page: %w", err) + } + if cache != nil { + cache.Put(dp.hash, buf[:pageSize]) + } + return nil + } + + _, result, err := c.roFile.readPage(buf, pageNum) + if err != nil { + return err + } + switch result { + case readFromCache: + c.server.readPath.Add("base_mem", 1) + case readFromDiskCold: + c.server.readPath.Add("base_disk_cold", 1) + case readFromDiskMiss: + c.server.readPath.Add("base_disk_miss", 1) + } + return nil +} + +// ReadAt reads len(p) bytes from the snapshot starting at byte offset off. +// It implements io.ReaderAt. +func (c *Snapshot) ReadAt(p []byte, off int64) (int, error) { + length := uint64(len(p)) + c.server.readBytes.Add(int64(length)) + + pageSize := uint64(c.server.pageSize) + + pos := uint64(0) + for pos < length { + absOffset := uint64(off) + pos + pageNum := int64(absOffset / pageSize) + pageOffset := absOffset % pageSize + + // Read up to the end of the page or the end of the requested length, + // whichever is smaller. Usually this will be pageSize, except for the + // first and last pages if the ReadAt call isn't page-aligned. + n := min(pageSize-pageOffset, length-pos) + + if pageOffset == 0 && n == pageSize { + // Full page; read directly into the caller's buffer. + if err := c.readPageData(p[pos:pos+pageSize], pageNum); err != nil { + return int(pos), err + } + } else { + // Partial page; need a scratch buffer from the pool. + bufp := c.server.pageBufPool.Get().(*[]byte) + if err := c.readPageData(*bufp, pageNum); err != nil { + return int(pos), err + } + copy(p[pos:pos+n], (*bufp)[pageOffset:pageOffset+n]) + c.server.pageBufPool.Put(bufp) + } + pos += n + } + return int(pos), nil +} + +// WriteAt writes len(p) bytes to the snapshot starting at byte offset off. +// Sub-page writes trigger a read-modify-write cycle. +// It implements io.WriterAt. +func (c *Snapshot) WriteAt(p []byte, off int64) (int, error) { + length := uint64(len(p)) + pageSize := uint64(c.server.pageSize) + + c.server.writeBytes.Add(int64(length)) + + offset := uint64(off) + pos := uint64(0) + for pos < length { + absOffset := offset + pos + pageNum := int64(absOffset / pageSize) + pageOffset := absOffset % pageSize + + n := pageSize - pageOffset + if n > length-pos { + n = length - pos + } + + pageData := make([]byte, pageSize) + if pageOffset == 0 && n == pageSize { + // Full page write. + copy(pageData, p[pos:pos+n]) + } else { + // Partial page write: read-modify-write. + if err := c.readPageData(pageData, pageNum); err != nil { + return int(pos), err + } + copy(pageData[pageOffset:], p[pos:pos+n]) + } + + cache := c.server.cache + var h pageHash + if cache != nil { + h = hashPage(pageData) + } + + c.mu.Lock() + dirtyNum := c.nextDirtyNum + c.nextDirtyNum++ + + _, err := c.dirtyFile.WriteAt(pageData, int64(dirtyNum)*int64(pageSize)) + if err != nil { + c.mu.Unlock() + return int(pos), fmt.Errorf("writing dirty page: %w", err) + } + + c.dirtyPages[pageNum] = dirtyPageInfo{ + hash: h, + dirtyNum: dirtyNum, + } + c.mu.Unlock() + + if cache != nil { + cache.Put(h, pageData) + } + c.server.writePages.Add(1) + pos += n + } + return int(pos), nil +} + +// WriteDirtyTo writes every page written to the snapshot, at its offset, +// to w, and returns the number of pages written. Writing them to a copy of +// the base image makes that copy equal to the snapshot's contents, which +// persists a snapshot (for example, a VM image warmed up by booting it over +// NBD). The final page is truncated to the image size. +// +// The caller should ensure nothing writes to the snapshot meanwhile; +// concurrent writes may or may not be included. +func (c *Snapshot) WriteDirtyTo(w io.WriterAt) (pages int, err error) { + c.mu.Lock() + pageNums := slices.Sorted(maps.Keys(c.dirtyPages)) + c.mu.Unlock() + + pageSize := int64(c.server.pageSize) + size := c.roFile.size + bufp := c.server.pageBufPool.Get().(*[]byte) + defer c.server.pageBufPool.Put(bufp) + buf := *bufp + for _, p := range pageNums { + off := p * pageSize + if off >= size { + continue + } + if err := c.readPageData(buf, p); err != nil { + return pages, err + } + n := min(pageSize, size-off) + if _, err := w.WriteAt(buf[:n], off); err != nil { + return pages, err + } + pages++ + } + return pages, nil +} + +// handleTrim zeros the given byte range and reverts any fully-covered +// dirty pages to the base image. Partial pages at the start and end of +// the range are zero-filled via WriteAt. +func (c *Snapshot) handleTrim(offset, length uint64) error { + end := offset + length + pageSize := uint64(c.server.pageSize) + + // Round up to the first fully-covered page. + firstFullPage := int64((offset + pageSize - 1) / pageSize) + // Round down to the last fully-covered page (exclusive). + lastFullPageExcl := int64(end / pageSize) + + // Zero partial start page. + if startRem := offset % pageSize; startRem != 0 { + n := min(pageSize-startRem, length) + bufp := c.server.pageBufPool.Get().(*[]byte) + clear((*bufp)[:n]) + _, err := c.WriteAt((*bufp)[:n], int64(offset)) + c.server.pageBufPool.Put(bufp) + if err != nil { + return err + } + } + + // Zero partial end page, if on a different page than the start. + if endRem := end % pageSize; endRem != 0 && int64(end/pageSize) >= firstFullPage { + bufp := c.server.pageBufPool.Get().(*[]byte) + clear((*bufp)[:endRem]) + _, err := c.WriteAt((*bufp)[:endRem], int64(lastFullPageExcl)*int64(pageSize)) + c.server.pageBufPool.Put(bufp) + if err != nil { + return err + } + } + + // Delete fully-covered pages, reverting them to the base image. + c.mu.Lock() + defer c.mu.Unlock() + for p := firstFullPage; p < lastFullPageExcl; p++ { + delete(c.dirtyPages, p) + } + return nil +}