diff --git a/.github/workflows/deploy-staging.yml b/.github/workflows/deploy-staging.yml index 909ca4db3..0c257b308 100644 --- a/.github/workflows/deploy-staging.yml +++ b/.github/workflows/deploy-staging.yml @@ -31,5 +31,4 @@ jobs: --ref main \ --field env=staging \ --field ref="$HYPEMAN_REF" \ - --field cli-version=latest \ - --field triggered_by="${{ github.actor }}" + --field cli-version=latest diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 2e2910379..261148f24 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -13,17 +13,6 @@ on: required: false type: string -# Every branch push triggers a full run and superseded runs used to keep -# running, holding self-hosted slots that every repo on the shared pool -# queues behind. Group by ref so a newer push cancels the older run for -# that branch. main and workflow_dispatch runs get run_id, a unique group -# each, because sharing a group is unsafe even with cancellation off: -# GitHub cancels an existing pending run when a newer one enters the -# group, which would drop a commit's only test signal. -concurrency: - group: test-${{ (github.event_name == 'push' && github.ref != 'refs/heads/main') && github.ref || github.run_id }} - cancel-in-progress: ${{ github.event_name == 'push' && github.ref != 'refs/heads/main' }} - # A slash-command dispatch supplies a fork repository and immutable commit SHA. # Normal push runs use the upstream repository and pushed ref. env: diff --git a/cmd/api/api/images.go b/cmd/api/api/images.go index 7ae20ebdc..e9f6defc8 100644 --- a/cmd/api/api/images.go +++ b/cmd/api/api/images.go @@ -94,10 +94,6 @@ func (s *ApiService) CreateImage(ctx context.Context, request oapi.CreateImageRe return oapi.CreateImage202JSONResponse(imageToOAPI(*img)), nil } -// TagImage handles POST /images/{name}/tag. -// Note: ResolveResource skips POST /images/{name}/tag, so the source is -// resolved by ImageManager.TagImage and a missing source gets the specific -// 404 body below. func (s *ApiService) TagImage(ctx context.Context, request oapi.TagImageRequestObject) (oapi.TagImageResponseObject, error) { if request.Body == nil { return oapi.TagImage400JSONResponse{ @@ -108,23 +104,33 @@ func (s *ApiService) TagImage(ctx context.Context, request oapi.TagImageRequestO img, err := s.ImageManager.TagImage(ctx, request.Name, request.Body.Target) if err != nil { - return tagImageErrorResponse(ctx, err, request.Name, request.Body.Target), nil + log := logger.FromContext(ctx) + switch { + case errors.Is(err, images.ErrInvalidName): + return oapi.TagImage400JSONResponse{ + Code: "invalid_name", + Message: err.Error(), + }, nil + case errors.Is(err, images.ErrNotFound): + return oapi.TagImage404JSONResponse{ + Code: "not_found", + Message: "source image not found", + }, nil + case errors.Is(err, images.ErrImageNotReady): + return oapi.TagImage409JSONResponse{ + Code: "image_not_ready", + Message: err.Error(), + }, nil + default: + log.ErrorContext(ctx, "failed to tag image", "error", err, "source", request.Name, "target", request.Body.Target) + return oapi.TagImage500JSONResponse{ + Code: "internal_error", + Message: "failed to tag image", + }, nil + } } - return oapi.TagImage200JSONResponse(imageToOAPI(*img)), nil -} -func tagImageErrorResponse(ctx context.Context, err error, source, target string) oapi.TagImageResponseObject { - switch { - case errors.Is(err, images.ErrInvalidName): - return oapi.TagImage400JSONResponse{Code: "invalid_name", Message: err.Error()} - case errors.Is(err, images.ErrNotFound): - return oapi.TagImage404JSONResponse{Code: "not_found", Message: "source image not found"} - case errors.Is(err, images.ErrImageNotReady): - return oapi.TagImage409JSONResponse{Code: "image_not_ready", Message: err.Error()} - default: - logger.FromContext(ctx).ErrorContext(ctx, "failed to tag image", "error", err, "source", source, "target", target) - return oapi.TagImage500JSONResponse{Code: "internal_error", Message: "failed to tag image"} - } + return oapi.TagImage200JSONResponse(imageToOAPI(*img)), nil } // GetImage gets image details by name diff --git a/cmd/api/api/images_test.go b/cmd/api/api/images_test.go index b6a1857b5..2a5d02367 100644 --- a/cmd/api/api/images_test.go +++ b/cmd/api/api/images_test.go @@ -2,12 +2,14 @@ package api import ( "context" + "encoding/json" "fmt" + "os" + "path/filepath" "testing" "time" "github.com/kernel/hypeman/lib/images" - "github.com/kernel/hypeman/lib/images/testutil" "github.com/kernel/hypeman/lib/oapi" "github.com/kernel/hypeman/lib/paths" "github.com/stretchr/testify/assert" @@ -513,36 +515,58 @@ func seedReadyDigestOnlyImage(t *testing.T, svc *ApiService, imageRef string, im require.NoError(t, err) require.True(t, ref.IsDigest(), "test helper expects a digest reference") - testutil.SeedReadyImage(t, paths.New(svc.Config.DataDir), testutil.Seed{ - Repository: ref.Repository(), - DigestHex: ref.DigestHex(), - Name: imageRef, - Tags: imageTags, - }) + p := paths.New(svc.Config.DataDir) + digestDir := p.ImageDigestDir(ref.Repository(), ref.DigestHex()) + require.NoError(t, os.MkdirAll(digestDir, 0o755)) + require.NoError(t, os.WriteFile(p.ImageDigestPath(ref.Repository(), ref.DigestHex()), []byte("rootfs"), 0o644)) + + meta := struct { + Name string `json:"name"` + Digest string `json:"digest"` + Status string `json:"status"` + SizeBytes int64 `json:"size_bytes"` + Tags map[string]string `json:"tags,omitempty"` + CreatedAt time.Time `json:"created_at"` + }{ + Name: imageRef, + Digest: "sha256:" + ref.DigestHex(), + Status: "ready", + SizeBytes: int64(len("rootfs")), + Tags: imageTags, + CreatedAt: time.Now().UTC(), + } + + data, err := json.Marshal(meta) + require.NoError(t, err) + require.NoError(t, os.WriteFile(p.ImageMetadata(ref.Repository(), ref.DigestHex()), data, 0o644)) } func TestTagImage_ErrorStatusMapping(t *testing.T) { t.Parallel() cases := []struct { - name string - err error - want oapi.TagImageResponseObject + name string + err error + wantType any + wantCode string }{ { - name: "invalid name -> 400", - err: fmt.Errorf("tag: %w", images.ErrInvalidName), - want: oapi.TagImage400JSONResponse{Code: "invalid_name", Message: "tag: invalid image name"}, + name: "invalid name -> 400", + err: fmt.Errorf("tag: %w", images.ErrInvalidName), + wantType: oapi.TagImage400JSONResponse{}, + wantCode: "invalid_name", }, { - name: "not found -> 404", - err: fmt.Errorf("tag: %w", images.ErrNotFound), - want: oapi.TagImage404JSONResponse{Code: "not_found", Message: "source image not found"}, + name: "not found -> 404", + err: fmt.Errorf("tag: %w", images.ErrNotFound), + wantType: oapi.TagImage404JSONResponse{}, + wantCode: "not_found", }, { - name: "not ready -> 409", - err: fmt.Errorf("tag: %w", images.ErrImageNotReady), - want: oapi.TagImage409JSONResponse{Code: "image_not_ready", Message: "tag: image is not ready"}, + name: "not ready -> 409", + err: fmt.Errorf("tag: %w", images.ErrImageNotReady), + wantType: oapi.TagImage409JSONResponse{}, + wantCode: "image_not_ready", }, } @@ -554,7 +578,8 @@ func TestTagImage_ErrorStatusMapping(t *testing.T) { Body: &oapi.TagImageRequest{Target: "docker.io/library/alpine:stable"}, }) require.NoError(t, err) - require.Equal(t, tc.want, resp) + require.IsType(t, tc.wantType, resp) + require.Equal(t, tc.wantCode, tagImageErrorCode(resp)) }) } } @@ -570,16 +595,53 @@ func TestTagImage_MissingBody(t *testing.T) { require.IsType(t, oapi.TagImage400JSONResponse{}, resp) } +func tagImageErrorCode(resp oapi.TagImageResponseObject) string { + switch r := resp.(type) { + case oapi.TagImage400JSONResponse: + return r.Code + case oapi.TagImage404JSONResponse: + return r.Code + case oapi.TagImage409JSONResponse: + return r.Code + case oapi.TagImage500JSONResponse: + return r.Code + default: + return "" + } +} + // seedReadyContentImage writes a ready image into the shared content layout // plus a repository tag reference, without pulling from a registry. func seedReadyContentImage(t *testing.T, svc *ApiService, repository, tag, digestHex string) { t.Helper() - testutil.SeedReadyImage(t, paths.New(svc.Config.DataDir), testutil.Seed{ - Repository: repository, - Tag: tag, - DigestHex: digestHex, - Content: true, - }) + + p := paths.New(svc.Config.DataDir) + contentDir := p.ImageContentDir(digestHex) + require.NoError(t, os.MkdirAll(contentDir, 0o755)) + require.NoError(t, os.WriteFile(p.ImageContentPath(digestHex), []byte("rootfs"), 0o644)) + + meta := struct { + Name string `json:"name"` + Digest string `json:"digest"` + Status string `json:"status"` + SizeBytes int64 `json:"size_bytes"` + CreatedAt time.Time `json:"created_at"` + }{ + Name: repository + ":" + tag, + Digest: "sha256:" + digestHex, + Status: "ready", + SizeBytes: int64(len("rootfs")), + CreatedAt: time.Now().UTC(), + } + data, err := json.Marshal(meta) + require.NoError(t, err) + require.NoError(t, os.WriteFile(p.ImageContentMetadata(digestHex), data, 0o644)) + + linkPath := p.ImageRepositoryTagSymlink(repository, tag) + target, err := filepath.Rel(filepath.Dir(linkPath), contentDir) + require.NoError(t, err) + require.NoError(t, os.MkdirAll(filepath.Dir(linkPath), 0o755)) + require.NoError(t, os.Symlink(target, linkPath)) } func TestTagImage_Success(t *testing.T) { diff --git a/lib/devices/mdev_darwin.go b/lib/devices/mdev_darwin.go index 1427a5095..22dd3435a 100644 --- a/lib/devices/mdev_darwin.go +++ b/lib/devices/mdev_darwin.go @@ -2,10 +2,7 @@ package devices -import ( - "context" - "fmt" -) +import "context" // SetGPUProfileCacheTTL is a no-op on macOS. func SetGPUProfileCacheTTL(ttl string) { @@ -33,10 +30,6 @@ func ListMdevDevices() ([]MdevDevice, error) { return []MdevDevice{}, nil } -func CreateVGPU(ctx context.Context, profileName, instanceID string) (*VGPUDevice, error) { - return nil, ErrVGPUNotSupportedOnMacOS -} - // CreateMdev returns an error on macOS as mdev is not supported. func CreateMdev(ctx context.Context, profileName, instanceID string) (*MdevDevice, error) { return nil, ErrVGPUNotSupportedOnMacOS @@ -52,13 +45,6 @@ func IsMdevInUse(mdevUUID string) bool { return false } -func DestroyVGPU(ctx context.Context, assignment VGPUAssignment) error { - if assignment.Framework != VGPUFrameworkNone && assignment.Framework != VGPUFrameworkMdev { - return fmt.Errorf("unknown vGPU framework %q", assignment.Framework) - } - return nil -} - // ReconcileMdevs is a no-op on macOS. func ReconcileMdevs(ctx context.Context, instanceInfos []MdevReconcileInfo) error { return nil diff --git a/lib/devices/mdev_linux.go b/lib/devices/mdev_linux.go index 1a398a418..21335ca23 100644 --- a/lib/devices/mdev_linux.go +++ b/lib/devices/mdev_linux.go @@ -126,7 +126,7 @@ func DiscoverVFs() ([]VirtualFunction, error) { vfs = append(vfs, VirtualFunction{ PCIAddress: vfAddr, ParentGPU: parentGPU, - Allocated: hasMdev, + HasMdev: hasMdev, }) } @@ -253,7 +253,7 @@ func countAvailableVFsForProfilesParallel(vfs []VirtualFunction, profiles []prof // Group free VFs by parent GPU (done once, shared by all goroutines) freeVFsByParent := make(map[string][]VirtualFunction) for _, vf := range vfs { - if vf.Allocated { + if vf.HasMdev { continue } freeVFsByParent[vf.ParentGPU] = append(freeVFsByParent[vf.ParentGPU], vf) @@ -453,7 +453,7 @@ func selectLeastLoadedVF(ctx context.Context, vfs []VirtualFunction, profileType allGPUs := make(map[string]bool) for _, vf := range vfs { allGPUs[vf.ParentGPU] = true - if !vf.Allocated { + if !vf.HasMdev { freeVFsByGPU[vf.ParentGPU] = append(freeVFsByGPU[vf.ParentGPU], vf) } } diff --git a/lib/devices/types.go b/lib/devices/types.go index 809d669fe..57d870592 100644 --- a/lib/devices/types.go +++ b/lib/devices/types.go @@ -60,12 +60,7 @@ func ValidateDeviceName(name string) bool { // GPUMode represents the host's GPU configuration mode type GPUMode string -type VGPUFramework string - const ( - VGPUFrameworkNone VGPUFramework = "" - VGPUFrameworkMdev VGPUFramework = "mdev" - // GPUModePassthrough indicates whole GPU VFIO passthrough GPUModePassthrough GPUMode = "passthrough" // GPUModeVGPU indicates SR-IOV + mdev based vGPU @@ -78,23 +73,7 @@ const ( type VirtualFunction struct { PCIAddress string `json:"pci_address"` // e.g., "0000:82:00.4" ParentGPU string `json:"parent_gpu"` // e.g., "0000:82:00.0" - Allocated bool `json:"allocated"` // true if a vGPU is assigned to this VF -} - -// VGPUAssignment identifies an existing vGPU assignment to release. -type VGPUAssignment struct { - Framework VGPUFramework - DevicePath string - MdevUUID string -} - -type VGPUDevice struct { - Framework VGPUFramework - VFAddress string - ProfileType string - ProfileName string - SysfsPath string - MdevUUID string + HasMdev bool `json:"has_mdev"` // true if an mdev is created on this VF } // MdevDevice represents an active mediated device (vGPU instance) diff --git a/lib/devices/vgpu_linux.go b/lib/devices/vgpu_linux.go deleted file mode 100644 index eaf210b42..000000000 --- a/lib/devices/vgpu_linux.go +++ /dev/null @@ -1,38 +0,0 @@ -//go:build linux - -package devices - -import ( - "context" - "fmt" - "path/filepath" -) - -func CreateVGPU(ctx context.Context, profileName, instanceID string) (*VGPUDevice, error) { - mdev, err := CreateMdev(ctx, profileName, instanceID) - if err != nil { - return nil, err - } - return &VGPUDevice{ - Framework: VGPUFrameworkMdev, - VFAddress: mdev.VFAddress, - ProfileType: mdev.ProfileType, - ProfileName: mdev.ProfileName, - SysfsPath: mdev.SysfsPath, - MdevUUID: mdev.UUID, - }, nil -} - -func DestroyVGPU(ctx context.Context, assignment VGPUAssignment) error { - if assignment.Framework != VGPUFrameworkNone && assignment.Framework != VGPUFrameworkMdev { - return fmt.Errorf("unknown vGPU framework %q", assignment.Framework) - } - mdevUUID := assignment.MdevUUID - if mdevUUID == "" { - if assignment.DevicePath == "" { - return nil - } - mdevUUID = filepath.Base(assignment.DevicePath) - } - return DestroyMdev(ctx, mdevUUID) -} diff --git a/lib/diskutilization/diskutilization.go b/lib/diskutilization/diskutilization.go index 27fef3c81..a5f55d677 100644 --- a/lib/diskutilization/diskutilization.go +++ b/lib/diskutilization/diskutilization.go @@ -4,6 +4,7 @@ import ( "io/fs" "os" "path/filepath" + "strings" "syscall" "github.com/kernel/hypeman/lib/paths" @@ -56,7 +57,7 @@ func Collect(p *paths.Paths) (Breakdown, error) { return false } name := entry.Name() - return name == "rootfs.erofs" || name == "rootfs.ext4" + return name == "rootfs.erofs" || name == "rootfs.ext4" || strings.HasPrefix(name, "layer.") }) if err != nil { return Breakdown{}, err @@ -178,6 +179,7 @@ func sumDirectChildFileAllocatedBytes(root string, childFile string) (int64, err func sumMatchingFilesAllocatedBytes(root string, match func(path string, entry fs.DirEntry) bool) (int64, error) { var total int64 + seen := make(map[fileIdentity]struct{}) err := filepath.WalkDir(root, func(path string, entry fs.DirEntry, err error) error { if err != nil { if os.IsNotExist(err) { @@ -185,9 +187,24 @@ func sumMatchingFilesAllocatedBytes(root string, match func(path string, entry f } return err } - if match(path, entry) { - total += allocatedBytesForPath(path) + if !match(path, entry) { + return nil } + info, statErr := os.Lstat(path) + if statErr != nil { + if os.IsNotExist(statErr) { + return nil + } + return statErr + } + if stat, ok := info.Sys().(*syscall.Stat_t); ok { + identity := fileIdentity{dev: uint64(stat.Dev), ino: uint64(stat.Ino)} + if _, exists := seen[identity]; exists { + return nil + } + seen[identity] = struct{}{} + } + total += allocatedBytesForPath(path) return nil }) if err != nil { @@ -244,6 +261,11 @@ func sumSnapshotTreeAllocatedBytes(root string, sharedExtents *sharedExtentTrack return privateTotal, sharedTotal, nil } +type fileIdentity struct { + dev uint64 + ino uint64 +} + func allocatedBytesForPath(path string) int64 { info, err := os.Lstat(path) if err != nil { diff --git a/lib/diskutilization/diskutilization_test.go b/lib/diskutilization/diskutilization_test.go index 50a787017..ff854438d 100644 --- a/lib/diskutilization/diskutilization_test.go +++ b/lib/diskutilization/diskutilization_test.go @@ -98,6 +98,22 @@ func TestCollect_UsesAllocatedBytesAndClassifiesSnapshots(t *testing.T) { require.Equal(t, otherTotal, utilization.SnapshotOther) } +func TestCollect_DeduplicatesHardLinkedImagesAndCountsLayers(t *testing.T) { + p := paths.New(t.TempDir()) + imagePath := filepath.Join(p.ImagesDir(), "repo", "digest", "rootfs.erofs") + require.NoError(t, createSparseTestFile(imagePath, 8192, []sparseWrite{{offset: 0, data: []byte("image")}})) + aliasPath := filepath.Join(p.ImagesDir(), "content", "digest", "rootfs.erofs") + require.NoError(t, os.MkdirAll(filepath.Dir(aliasPath), 0755)) + require.NoError(t, os.Link(imagePath, aliasPath)) + + layerPath := filepath.Join(p.ImageLayersDir(), "layer-digest", "layer.erofs") + require.NoError(t, createSparseTestFile(layerPath, 8192, []sparseWrite{{offset: 0, data: []byte("layer")}})) + + utilization, err := Collect(p) + require.NoError(t, err) + require.Equal(t, allocatedBytesForPath(imagePath)+allocatedBytesForPath(layerPath), utilization.Images) +} + func createSparseTestFile(path string, size int64, writes []sparseWrite) error { if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { return err diff --git a/lib/hypervisor/cloudhypervisor/config.go b/lib/hypervisor/cloudhypervisor/config.go index ca3d98a55..e9f91fe4a 100644 --- a/lib/hypervisor/cloudhypervisor/config.go +++ b/lib/hypervisor/cloudhypervisor/config.go @@ -126,20 +126,13 @@ func ToVMConfig(cfg hypervisor.VMConfig) vmm.VmConfig { // Device passthrough configuration var devices *[]vmm.DeviceConfig - deviceCount := len(cfg.PCIDevices) - if cfg.VGPUDevicePath != "" { - deviceCount++ - } - if deviceCount > 0 { - deviceConfigs := make([]vmm.DeviceConfig, 0, deviceCount) + if len(cfg.PCIDevices) > 0 { + deviceConfigs := make([]vmm.DeviceConfig, 0, len(cfg.PCIDevices)) for _, path := range cfg.PCIDevices { deviceConfigs = append(deviceConfigs, vmm.DeviceConfig{ Path: path, }) } - if cfg.VGPUDevicePath != "" { - deviceConfigs = append(deviceConfigs, vmm.DeviceConfig{Path: cfg.VGPUDevicePath}) - } devices = &deviceConfigs } diff --git a/lib/hypervisor/cloudhypervisor/config_test.go b/lib/hypervisor/cloudhypervisor/config_test.go index be39d13af..b5cdb96e9 100644 --- a/lib/hypervisor/cloudhypervisor/config_test.go +++ b/lib/hypervisor/cloudhypervisor/config_test.go @@ -8,16 +8,6 @@ import ( "github.com/stretchr/testify/require" ) -func TestToVMConfigIncludesVGPUDevice(t *testing.T) { - path := "/sys/bus/mdev/devices/aa618089-8b16-4d01-a136-25a0f3c73123" - - vmCfg := ToVMConfig(hypervisor.VMConfig{VGPUDevicePath: path}) - - require.NotNil(t, vmCfg.Devices) - require.Len(t, *vmCfg.Devices, 1) - assert.Equal(t, path, (*vmCfg.Devices)[0].Path) -} - func TestToVMConfig_GuestMemoryBalloon(t *testing.T) { cfg := hypervisor.VMConfig{ VCPUs: 1, diff --git a/lib/hypervisor/config.go b/lib/hypervisor/config.go index 456775868..2562868da 100644 --- a/lib/hypervisor/config.go +++ b/lib/hypervisor/config.go @@ -24,8 +24,7 @@ type VMConfig struct { VsockSocket string // PCI device passthrough (GPU, etc.) - PCIDevices []string - VGPUDevicePath string + PCIDevices []string // Boot configuration KernelPath string diff --git a/lib/hypervisor/qemu/config.go b/lib/hypervisor/qemu/config.go index 24401c0f2..648763c22 100644 --- a/lib/hypervisor/qemu/config.go +++ b/lib/hypervisor/qemu/config.go @@ -94,10 +94,13 @@ func buildArgs(cfg hypervisor.VMConfig, machine MachineType) []string { args = append(args, "-device", fmt.Sprintf("%s,guest-cid=%d", virtioDevice(microvm, "vhost-vsock"), cfg.VsockCID)) } - // Whole-device PCI passthrough (vGPU attaches via VGPUDevicePath below) + // PCI device passthrough (GPU, mdev vGPU, etc.) for _, devicePath := range cfg.PCIDevices { var deviceArg string - if strings.HasPrefix(devicePath, "/sys/bus/pci/devices/") { + if strings.HasPrefix(devicePath, "/sys/bus/mdev/devices/") { + // mdev device (vGPU) - use sysfsdev parameter + deviceArg = fmt.Sprintf("vfio-pci,sysfsdev=%s", devicePath) + } else if strings.HasPrefix(devicePath, "/sys/bus/pci/devices/") { // Full sysfs path for regular PCI device - extract the PCI address // Using filepath.Base is more robust than manual string splitting pciAddr := filepath.Base(strings.TrimSuffix(devicePath, "/")) @@ -109,10 +112,6 @@ func buildArgs(cfg hypervisor.VMConfig, machine MachineType) []string { args = append(args, "-device", deviceArg) } - if cfg.VGPUDevicePath != "" { - args = append(args, "-device", fmt.Sprintf("vfio-pci,sysfsdev=%s", cfg.VGPUDevicePath)) - } - // Serial console output to file. Use a chardev with append=on so QEMU // opens the file with O_APPEND. Without it, QEMU writes at its internal // fd offset; if the file is externally truncated (e.g. log rotation via diff --git a/lib/hypervisor/qemu/config_test.go b/lib/hypervisor/qemu/config_test.go index 0e08e50cf..c8d25584f 100644 --- a/lib/hypervisor/qemu/config_test.go +++ b/lib/hypervisor/qemu/config_test.go @@ -123,49 +123,6 @@ func TestBuildArgs_Vsock(t *testing.T) { assert.Contains(t, args, "vhost-vsock-pci,guest-cid=123") } -func TestBuildArgs_VGPU(t *testing.T) { - t.Parallel() - - for _, path := range []string{ - "/sys/bus/mdev/devices/aa618089-8b16-4d01-a136-25a0f3c73123", - "/sys/bus/pci/devices/0000:82:00.4", - } { - path := path - t.Run(path, func(t *testing.T) { - t.Parallel() - args := BuildArgs(hypervisor.VMConfig{ - VCPUs: 1, - MemoryBytes: 512 * 1024 * 1024, - VGPUDevicePath: path, - }) - assert.Contains(t, args, "vfio-pci,sysfsdev="+path) - }) - } -} - -func TestBuildArgs_VGPUAfterPCIDevices(t *testing.T) { - args := BuildArgs(hypervisor.VMConfig{ - VCPUs: 1, - MemoryBytes: 512 * 1024 * 1024, - PCIDevices: []string{"0000:01:00.0"}, - VGPUDevicePath: "/sys/bus/mdev/devices/aa618089-8b16-4d01-a136-25a0f3c73123", - }) - - pciDeviceIndex := -1 - vgpuDeviceIndex := -1 - for i, arg := range args { - switch arg { - case "vfio-pci,host=0000:01:00.0": - pciDeviceIndex = i - case "vfio-pci,sysfsdev=/sys/bus/mdev/devices/aa618089-8b16-4d01-a136-25a0f3c73123": - vgpuDeviceIndex = i - } - } - - assert.Greater(t, pciDeviceIndex, -1) - assert.Greater(t, vgpuDeviceIndex, pciDeviceIndex) -} - func TestBuildArgs_PCIPassthrough(t *testing.T) { cfg := hypervisor.VMConfig{ VCPUs: 1, diff --git a/lib/hypervisor/qemu/machine_test.go b/lib/hypervisor/qemu/machine_test.go index 9484cdcde..9998d9280 100644 --- a/lib/hypervisor/qemu/machine_test.go +++ b/lib/hypervisor/qemu/machine_test.go @@ -81,19 +81,6 @@ func TestQEMUCapabilitiesAdvertiseFork(t *testing.T) { assert.True(t, (MicroVMProfile{}).capabilities().SupportsFork) } -func TestMicroVMValidateConfigRejectsVFIODevices(t *testing.T) { - t.Parallel() - err := (MicroVMProfile{}).validateConfig(hypervisor.VMConfig{ - PCIDevices: []string{"0000:82:00.4"}, - }) - require.ErrorContains(t, err, "microvm does not support PCI devices") - - err = (MicroVMProfile{}).validateConfig(hypervisor.VMConfig{ - VGPUDevicePath: "/sys/bus/mdev/devices/aa618089-8b16-4d01-a136-25a0f3c73123", - }) - require.ErrorContains(t, err, "microvm does not support PCI devices") -} - func TestValidateConfigMicroVM(t *testing.T) { t.Parallel() if _, err := microVMMachineType(); err != nil { diff --git a/lib/hypervisor/qemu/profile.go b/lib/hypervisor/qemu/profile.go index 2d61cd6fe..2ac3a2cbf 100644 --- a/lib/hypervisor/qemu/profile.go +++ b/lib/hypervisor/qemu/profile.go @@ -46,7 +46,7 @@ func (MicroVMProfile) validateConfig(cfg hypervisor.VMConfig) error { if cfg.HotplugBytes > 0 { return fmt.Errorf("microvm does not support hotplug memory") } - if len(cfg.PCIDevices) > 0 || cfg.VGPUDevicePath != "" { + if len(cfg.PCIDevices) > 0 { return fmt.Errorf("microvm does not support PCI devices") } diff --git a/lib/hypervisor/socket_pid.go b/lib/hypervisor/socket_pid.go deleted file mode 100644 index 0d009f9ee..000000000 --- a/lib/hypervisor/socket_pid.go +++ /dev/null @@ -1,5 +0,0 @@ -package hypervisor - -import "errors" - -var ErrNoOwningProcess = errors.New("no owning process found") diff --git a/lib/hypervisor/socket_pid_linux.go b/lib/hypervisor/socket_pid_linux.go index 45b280edc..7f46ebfa3 100644 --- a/lib/hypervisor/socket_pid_linux.go +++ b/lib/hypervisor/socket_pid_linux.go @@ -4,77 +4,36 @@ package hypervisor import ( "bufio" - "errors" "fmt" - "io/fs" "os" "path/filepath" - "slices" "strconv" "strings" - "syscall" ) -var procDir = "/proc" - -// soAcceptcon marks a listening socket in /proc/net/unix (__SO_ACCEPTCON). -const soAcceptcon = 0x10000 - // ResolveProcessPID finds the process currently holding the listening Unix -// socket for the given hypervisor control path, via the socket inode in -// /proc/net/unix and each process's fd table. The fd scan requires the -// caller to hold CAP_SYS_PTRACE (or run as root) so no live owner is missed; -// an ErrNoOwningProcess result is proof the listener is gone. -func ResolveProcessPID(socketPath string) (pid int, err error) { - return resolveProcessPID(socketPath, 0) -} - -// ResolveProcessPIDForOwner resolves a socket while preferring an expected -// owner when the socket descriptor is temporarily shared with a child process. -func ResolveProcessPIDForOwner(socketPath string, ownerPID int) (pid int, err error) { - return resolveProcessPID(socketPath, ownerPID) -} - -func resolveProcessPID(socketPath string, ownerPID int) (pid int, err error) { +// socket for the given hypervisor control path. +func ResolveProcessPID(socketPath string) (int, error) { socketRef, err := socketRefForPath(socketPath) - if err != nil { - return 0, err - } - // Confirm the expected owner first so a live stored PID does not - // require scanning every process fd. - if ownerPID > 0 && processHoldsSocketRef(ownerPID, socketRef) { - return ownerPID, nil + if err == nil { + if pid, refErr := pidBySocketRef(socketRef); refErr == nil { + return pid, nil + } } - return pidBySocketRef(socketRef, ownerPID) -} -func processHoldsSocketRef(pid int, socketRef string) bool { - fdEntries, err := os.ReadDir(filepath.Join(procDir, strconv.Itoa(pid), "fd")) - if err != nil { - return false + if pid, cmdErr := pidByCmdline(socketPath); cmdErr == nil { + return pid, nil } - for _, fdEntry := range fdEntries { - target, err := os.Readlink(filepath.Join(procDir, strconv.Itoa(pid), "fd", fdEntry.Name())) - if err != nil { - // Skip fds that cannot be read, like the full scan does: an fd - // vanishing mid-scan must not hide a listener held by a later fd. - continue - } - if strings.TrimSpace(target) == socketRef { - return true - } - } - return false + + return 0, fmt.Errorf("resolve process pid for socket %s: no owning process found", socketPath) } -func pidBySocketRef(socketRef string, ownerPID int) (int, error) { - procEntries, err := os.ReadDir(procDir) +func pidBySocketRef(socketRef string) (int, error) { + procEntries, err := os.ReadDir("/proc") if err != nil { return 0, fmt.Errorf("read /proc: %w", err) } - var owners []int - var scanErr error for _, entry := range procEntries { if !entry.IsDir() { continue @@ -85,57 +44,62 @@ func pidBySocketRef(socketRef string, ownerPID int) (int, error) { continue } - fdEntries, err := os.ReadDir(filepath.Join(procDir, entry.Name(), "fd")) + fdEntries, err := os.ReadDir(filepath.Join("/proc", entry.Name(), "fd")) if err != nil { - if errors.Is(err, fs.ErrNotExist) || errors.Is(err, syscall.ESRCH) { - continue - } - scanErr = err continue } for _, fdEntry := range fdEntries { - target, err := os.Readlink(filepath.Join(procDir, entry.Name(), "fd", fdEntry.Name())) + target, err := os.Readlink(filepath.Join("/proc", entry.Name(), "fd", fdEntry.Name())) if err != nil { - if errors.Is(err, fs.ErrNotExist) || errors.Is(err, syscall.ESRCH) { - continue - } - scanErr = err continue } if strings.TrimSpace(target) == socketRef { - owners = append(owners, pid) - break + return pid, nil } } } - // The scan observed ownerPID holding the listener fd — the same evidence - // the fast path uses — so a child transiently sharing the inherited fd - // must not turn a proven owner into an error. - if ownerPID > 0 && slices.Contains(owners, ownerPID) { - return ownerPID, nil - } - if len(owners) == 1 { - return owners[0], nil - } - if len(owners) > 1 { - return 0, fmt.Errorf("resolve process pid for %s: multiple owning processes found: %v", socketRef, owners) + return 0, fmt.Errorf("resolve process pid for %s: no owning process found", socketRef) +} + +func pidByCmdline(socketPath string) (int, error) { + procEntries, err := os.ReadDir("/proc") + if err != nil { + return 0, fmt.Errorf("read /proc: %w", err) } - if scanErr != nil { - return 0, fmt.Errorf("resolve process pid for %s: inspect process fds: %w", socketRef, scanErr) + + for _, entry := range procEntries { + if !entry.IsDir() { + continue + } + + pid, err := strconv.Atoi(entry.Name()) + if err != nil { + continue + } + + cmdline, err := os.ReadFile(filepath.Join("/proc", entry.Name(), "cmdline")) + if err != nil || len(cmdline) == 0 { + continue + } + for _, arg := range strings.Split(string(cmdline), "\x00") { + if arg == socketPath { + return pid, nil + } + } } - return 0, fmt.Errorf("resolve process pid for %s: %w", socketRef, ErrNoOwningProcess) + + return 0, fmt.Errorf("resolve process pid for socket %s: no matching command line found", socketPath) } func socketRefForPath(socketPath string) (string, error) { - file, err := os.Open(filepath.Join(procDir, "net", "unix")) + file, err := os.Open("/proc/net/unix") if err != nil { return "", fmt.Errorf("open /proc/net/unix: %w", err) } defer file.Close() scanner := bufio.NewScanner(file) - var socketRef string for scanner.Scan() { fields := strings.Fields(scanner.Text()) if len(fields) < 7 { @@ -148,26 +112,14 @@ func socketRefForPath(socketPath string) (string, error) { if path != socketPath { continue } - // Accepted server-side sockets list the bound path too; only the - // listener identifies the owning process. - flags, parseErr := strconv.ParseUint(fields[3], 16, 32) - if parseErr != nil || flags&soAcceptcon == 0 { - continue - } inode := fields[6] if inode == "" { break } - if socketRef != "" { - return "", fmt.Errorf("resolve process pid for socket %s: multiple socket inodes found", socketPath) - } - socketRef = fmt.Sprintf("socket:[%s]", inode) + return fmt.Sprintf("socket:[%s]", inode), nil } if err := scanner.Err(); err != nil { return "", fmt.Errorf("scan /proc/net/unix: %w", err) } - if socketRef != "" { - return socketRef, nil - } - return "", fmt.Errorf("resolve process pid for socket %s: socket inode not found: %w", socketPath, ErrNoOwningProcess) + return "", fmt.Errorf("resolve process pid for socket %s: socket inode not found", socketPath) } diff --git a/lib/hypervisor/socket_pid_linux_test.go b/lib/hypervisor/socket_pid_linux_test.go index d47044e9c..270524532 100644 --- a/lib/hypervisor/socket_pid_linux_test.go +++ b/lib/hypervisor/socket_pid_linux_test.go @@ -3,14 +3,10 @@ package hypervisor import ( - "context" - "errors" "net" "os" - "os/exec" "path/filepath" "testing" - "time" "github.com/stretchr/testify/require" ) @@ -27,338 +23,3 @@ func TestResolveProcessPID(t *testing.T) { require.NoError(t, err) require.Equal(t, os.Getpid(), pid) } - -func TestResolveProcessPIDIgnoresConnectedSocketEntries(t *testing.T) { - tmpDir := t.TempDir() - socketPath := filepath.Join(tmpDir, "test.sock") - - listener, err := net.Listen("unix", socketPath) - require.NoError(t, err) - defer listener.Close() - - // Accepted server-side sockets share the listener's path in - // /proc/net/unix; they must not make the listener's inode ambiguous. - conn, err := net.Dial("unix", socketPath) - require.NoError(t, err) - defer conn.Close() - accepted, err := listener.Accept() - require.NoError(t, err) - defer accepted.Close() - - pid, err := ResolveProcessPID(socketPath) - require.NoError(t, err) - require.Equal(t, os.Getpid(), pid) -} - -func TestResolveProcessPIDFailsForDuplicateSocketPaths(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte( - "00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"+ - "00000000: 00000002 00000000 00010000 0001 01 67890 "+socketPath+"\n"), 0o644)) - - _, err := ResolveProcessPID(socketPath) - require.ErrorContains(t, err, "multiple socket inodes found") -} - -func TestResolveProcessPIDToleratesExitedProcess(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "100"), 0o755)) - fdDir := filepath.Join(procDir, "200", "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(fdDir, "3"))) - - pid, err := ResolveProcessPID(socketPath) - require.NoError(t, err) - require.Equal(t, 200, pid) -} - -func TestResolveProcessPIDForOwnerPrefersExpectedProcess(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - for _, pid := range []string{"100", "101"} { - fdDir := filepath.Join(procDir, pid, "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(fdDir, "3"))) - } - - _, err := ResolveProcessPID(socketPath) - require.ErrorContains(t, err, "multiple owning processes found") - - pid, err := ResolveProcessPIDForOwner(socketPath, 101) - require.NoError(t, err) - require.Equal(t, 101, pid) -} - -func TestPidBySocketRefPrefersExpectedOwnerAmongMultiple(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - // The full scan itself must prefer the expected owner when a child - // transiently shares the inherited listener fd: the fast path can miss on - // a transient fd-dir read failure, and the scan's observation of the - // owner holding the fd is the same evidence the fast path would have used. - for _, pid := range []string{"100", "101"} { - fdDir := filepath.Join(procDir, pid, "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(fdDir, "3"))) - } - - pid, err := pidBySocketRef("socket:[12345]", 101) - require.NoError(t, err) - require.Equal(t, 101, pid) - - _, err = pidBySocketRef("socket:[12345]", 0) - require.ErrorContains(t, err, "multiple owning processes found") - - _, err = pidBySocketRef("socket:[12345]", 999) - require.ErrorContains(t, err, "multiple owning processes found") -} - -func TestResolveProcessPIDReportsNoOwnerAfterExitedProcesses(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "100"), 0o755)) - - _, err := ResolveProcessPID(socketPath) - require.ErrorIs(t, err, ErrNoOwningProcess) - require.NotContains(t, err.Error(), "inspect process fds") -} - -func TestResolveProcessPIDReportsMissingSocket(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), nil, 0o644)) - - _, err := ResolveProcessPID("/tmp/missing.sock") - require.ErrorIs(t, err, ErrNoOwningProcess) -} - -func TestResolveProcessPIDFailsWhenFDIsUnreadable(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - fdDir := filepath.Join(procDir, "123", "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(fdDir, "3"), nil, 0o644)) - - _, err := ResolveProcessPID(socketPath) - require.Error(t, err) - require.ErrorContains(t, err, "inspect process fds") - require.False(t, errors.Is(err, ErrNoOwningProcess)) -} - -func TestResolveProcessPIDForOwnerConfirmsCandidateWithoutFullScan(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - - fdDir := filepath.Join(procDir, "100", "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(fdDir, "3"))) - - // An unreadable sibling fd must not block confirming the candidate. - siblingFDDir := filepath.Join(procDir, "123", "fd") - require.NoError(t, os.MkdirAll(siblingFDDir, 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(siblingFDDir, "3"), nil, 0o644)) - - pid, err := ResolveProcessPIDForOwner(socketPath, 100) - require.NoError(t, err) - require.Equal(t, 100, pid) -} - -func TestResolveProcessPIDForOwnerSkipsUnreadableCandidateFD(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - - // An unreadable fd before the listener fd must not abort the candidate - // check; the scan skips it and still finds the match. - fdDir := filepath.Join(procDir, "100", "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(fdDir, "1"), nil, 0o644)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(fdDir, "3"))) - - pid, err := ResolveProcessPIDForOwner(socketPath, 100) - require.NoError(t, err) - require.Equal(t, 100, pid) -} - -func TestResolveProcessPIDForOwnerFallsThroughWhenCandidateLacksSocket(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - - candidateFDDir := filepath.Join(procDir, "999", "fd") - require.NoError(t, os.MkdirAll(candidateFDDir, 0o755)) - require.NoError(t, os.Symlink("socket:[99999]", filepath.Join(candidateFDDir, "3"))) - - ownerFDDir := filepath.Join(procDir, "200", "fd") - require.NoError(t, os.MkdirAll(ownerFDDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(ownerFDDir, "3"))) - - pid, err := ResolveProcessPIDForOwner(socketPath, 999) - require.NoError(t, err) - require.Equal(t, 200, pid) -} - -func TestResolveProcessPIDForOwnerFallsThroughWhenCandidateIsGone(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"), 0o644)) - - ownerFDDir := filepath.Join(procDir, "200", "fd") - require.NoError(t, os.MkdirAll(ownerFDDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(ownerFDDir, "3"))) - - pid, err := ResolveProcessPIDForOwner(socketPath, 999) - require.NoError(t, err) - require.Equal(t, 200, pid) -} - -func TestResolveProcessPIDForOwnerReportsMissingSocket(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("00000000: 00000002 00000000 00010000 0001 01 12345 /tmp/other.sock\n"), 0o644)) - - _, err := ResolveProcessPIDForOwner("/tmp/missing.sock", 100) - require.ErrorIs(t, err, ErrNoOwningProcess) -} - -func TestResolveProcessPIDForOwnerReportsMissingSocketWithHeaderOnlyUnixTable(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte("Num RefCount Protocol Flags Type St Inode Path\n"), 0o644)) - - _, err := ResolveProcessPIDForOwner("/tmp/missing.sock", 100) - require.ErrorIs(t, err, ErrNoOwningProcess) -} - -func TestResolveProcessPIDForOwnerReportsDuplicateSocketInodes(t *testing.T) { - oldProcDir := procDir - procDir = t.TempDir() - t.Cleanup(func() { procDir = oldProcDir }) - - socketPath := "/tmp/test.sock" - require.NoError(t, os.MkdirAll(filepath.Join(procDir, "net"), 0o755)) - require.NoError(t, os.WriteFile(filepath.Join(procDir, "net", "unix"), []byte( - "00000000: 00000002 00000000 00010000 0001 01 12345 "+socketPath+"\n"+ - "00000000: 00000002 00000000 00010000 0001 01 67890 "+socketPath+"\n"), 0o644)) - - fdDir := filepath.Join(procDir, "100", "fd") - require.NoError(t, os.MkdirAll(fdDir, 0o755)) - require.NoError(t, os.Symlink("socket:[12345]", filepath.Join(fdDir, "3"))) - - _, err := ResolveProcessPIDForOwner(socketPath, 100) - require.ErrorContains(t, err, "multiple socket inodes found") -} - -func TestResolveProcessPIDForOwnerConfirmsLiveListener(t *testing.T) { - tmpDir := t.TempDir() - socketPath := filepath.Join(tmpDir, "test.sock") - - listener, err := net.Listen("unix", socketPath) - require.NoError(t, err) - defer listener.Close() - - pid, err := ResolveProcessPIDForOwner(socketPath, os.Getpid()) - require.NoError(t, err) - require.Equal(t, os.Getpid(), pid) -} - -func TestResolveProcessPIDIgnoresCommandLineBystander(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - require.NoError(t, os.WriteFile(socketPath, nil, 0o600)) - - // A process carrying the socket path in its command line (e.g. a debug - // client like ch-remote) without holding the listener must not resolve - // as the owner; a missing listener is proof the hypervisor is gone. - bystander := exec.Command("sh", "-c", "sleep 30", "sh", socketPath) - require.NoError(t, bystander.Start()) - t.Cleanup(func() { - _ = bystander.Process.Kill() - _ = bystander.Wait() - }) - - _, err := ResolveProcessPID(socketPath) - require.ErrorIs(t, err, ErrNoOwningProcess) -} - -func TestResolveProcessPIDDuringProcessChurn(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - listener, err := net.Listen("unix", socketPath) - require.NoError(t, err) - defer listener.Close() - - ctx, cancel := context.WithCancel(context.Background()) - done := make(chan struct{}) - go func() { - defer close(done) - for ctx.Err() == nil { - _ = exec.CommandContext(ctx, "/bin/true").Run() - } - }() - defer func() { - cancel() - <-done - }() - - // Resolve without an owner hint so every iteration runs the full /proc - // scan; the owner fast path never exercises the churn tolerance. - deadline := time.Now().Add(2 * time.Second) - for time.Now().Before(deadline) { - pid, err := ResolveProcessPID(socketPath) - require.NoError(t, err) - require.Equal(t, os.Getpid(), pid) - } -} diff --git a/lib/hypervisor/socket_pid_other.go b/lib/hypervisor/socket_pid_other.go index 1fb594653..75db657e6 100644 --- a/lib/hypervisor/socket_pid_other.go +++ b/lib/hypervisor/socket_pid_other.go @@ -9,8 +9,3 @@ import "fmt" func ResolveProcessPID(socketPath string) (int, error) { return 0, fmt.Errorf("resolve process pid for socket %s: not supported on this platform", socketPath) } - -// ResolveProcessPIDForOwner is only implemented on Linux. -func ResolveProcessPIDForOwner(socketPath string, _ int) (int, error) { - return ResolveProcessPID(socketPath) -} diff --git a/lib/images/compose.go b/lib/images/compose.go index 3567644e1..718fadd50 100644 --- a/lib/images/compose.go +++ b/lib/images/compose.go @@ -3,42 +3,93 @@ package images import ( "fmt" "os" + "path/filepath" ) -// composeRootfs validates the persisted model and merges its layers into dest -// in manifest order, reading each layer blob from the shared OCI cache. -// Whiteout and opaque-directory markers are interpreted as each layer is -// applied. -func (c *ociClient) composeRootfs(dest, layoutTag string, model *imageManifestModel) error { +// validateModelPairing mirrors validateConfigFileForUnpack: the image config +// must carry one diff id per manifest layer so composition never indexes past +// the end of the pairing. +func validateModelPairing(layoutTag string, model *imageManifestModel) error { if err := validateManifestModel(layoutTag, model); err != nil { - return fmt.Errorf("validate manifest model: %w", err) + return fmt.Errorf("unpack rootfs: %w", err) } - if len(model.Layers) == 0 { - return fmt.Errorf("image has no layers") + return nil +} + +// composeRootfs merges an image's layers into dest in manifest order, reading +// each layer blob from the shared OCI cache. The result is one complete rootfs +// tree that is exported to a single disk, matching the guest's contract: one +// read-only lower filesystem and one writable overlay upper. Whiteout and +// opaque-directory markers are interpreted as each layer is applied instead of +// being left in the tree, so the composed rootfs never relies on tar-level +// whiteouts composing on overlayfs. +func (c *ociClient) composeRootfs(dest string, layers []layerDescriptor) error { + trees, err := c.composeRootfsWithLayerTrees(dest, layers) + if err != nil { + return err + } + cleanupLayerTrees(trees) + return nil +} + +func (c *ociClient) composeRootfsWithLayerTrees(dest string, layers []layerDescriptor) (map[string]layerTree, error) { + if len(layers) == 0 { + return nil, fmt.Errorf("image has no layers") } if err := os.MkdirAll(dest, 0755); err != nil { - return fmt.Errorf("create compose directory: %w", err) + return nil, fmt.Errorf("create compose directory: %w", err) } - for i, desc := range model.Layers { - if err := c.applyLayerToDir(dest, desc); err != nil { - return fmt.Errorf("apply layer %d (%s): %w", i, desc.Digest, err) + trees := make(map[string]layerTree, len(layers)) + for i, desc := range layers { + if tree, ok := trees[desc.Digest]; ok { + if err := applyLayerTree(tree.path, dest); err != nil { + cleanupLayerTrees(trees) + return nil, fmt.Errorf("apply layer %d (%s): %w", i, desc.Digest, err) + } + continue + } + tree, err := c.extractLayerTree(desc) + if err != nil { + cleanupLayerTrees(trees) + return nil, fmt.Errorf("extract layer %d (%s): %w", i, desc.Digest, err) + } + if err := applyLayerTree(tree.path, dest); err != nil { + cleanupLayerTrees(trees) + cleanupLayerTree(tree) + return nil, fmt.Errorf("apply layer %d (%s): %w", i, desc.Digest, err) } + trees[desc.Digest] = tree } - return nil + return trees, nil } -func (c *ociClient) applyLayerToDir(dest string, desc layerDescriptor) error { - layerDir, err := os.MkdirTemp("", "hypeman-layer-*") +// extractLayerTree extracts one layer into a private staging directory so it +// can be applied to a composed rootfs and later materialized as an artifact. +func (c *ociClient) extractLayerTree(desc layerDescriptor) (layerTree, error) { + layerHex, err := layerDigestHex(desc) if err != nil { - return fmt.Errorf("create layer staging directory: %w", err) + return layerTree{}, err + } + blobPath := filepath.Join(c.cacheDir, "blobs", "sha256", layerHex) + if _, err := os.Stat(blobPath); err != nil { + if os.IsNotExist(err) { + return layerTree{}, fmt.Errorf("layer blob missing from oci cache: %s", desc.Digest) + } + return layerTree{}, fmt.Errorf("stat layer blob: %w", err) } - defer os.RemoveAll(layerDir) - if _, err := unpackCachedLayer(c.cacheDir, desc, layerDir); err != nil { - return err + layerDir, err := os.MkdirTemp("", "hypeman-layer-*") + if err != nil { + return layerTree{}, fmt.Errorf("create layer staging directory: %w", err) } - if err := applyLayerTree(layerDir, dest); err != nil { - return fmt.Errorf("apply layer tree: %w", err) + stats, err := unpackLayerBlob(blobPath, desc.MediaType, layerDir) + if err != nil { + cleanupLayerTree(layerTree{path: layerDir}) + return layerTree{}, err } - return nil + if desc.DiffID != "" && stats.diffID != desc.DiffID { + cleanupLayerTree(layerTree{path: layerDir}) + return layerTree{}, fmt.Errorf("layer %s diff id mismatch: got %s, want %s", desc.Digest, stats.diffID, desc.DiffID) + } + return layerTree{path: layerDir, stats: stats}, nil } diff --git a/lib/images/compose_test.go b/lib/images/compose_test.go index b67c8dee8..286fc8a93 100644 --- a/lib/images/compose_test.go +++ b/lib/images/compose_test.go @@ -1,20 +1,62 @@ package images import ( + "archive/tar" + "bytes" + "compress/gzip" "io" "os" "os/exec" "path/filepath" - "strings" "testing" gcr "github.com/google/go-containerregistry/pkg/v1" "github.com/google/go-containerregistry/pkg/v1/empty" "github.com/google/go-containerregistry/pkg/v1/mutate" + "github.com/google/go-containerregistry/pkg/v1/tarball" "github.com/kernel/hypeman/lib/paths" "github.com/stretchr/testify/require" ) +type tarEntrySpec struct { + name string + content string + isDir bool + mode int64 +} + +// specLayer builds a gzipped tar layer from entry specs in order. +func specLayer(t *testing.T, entries []tarEntrySpec) gcr.Layer { + t.Helper() + + var buf bytes.Buffer + gzw := gzip.NewWriter(&buf) + tw := tar.NewWriter(gzw) + for _, entry := range entries { + if entry.isDir { + require.NoError(t, tw.WriteHeader(&tar.Header{Name: entry.name, Typeflag: tar.TypeDir, Mode: entry.mode})) + continue + } + require.NoError(t, tw.WriteHeader(&tar.Header{ + Name: entry.name, + Typeflag: tar.TypeReg, + Mode: entry.mode, + Size: int64(len(entry.content)), + })) + _, err := tw.Write([]byte(entry.content)) + require.NoError(t, err) + } + require.NoError(t, tw.Close()) + require.NoError(t, gzw.Close()) + + data := buf.Bytes() + layer, err := tarball.LayerFromOpener(func() (io.ReadCloser, error) { + return io.NopCloser(bytes.NewReader(data)), nil + }) + require.NoError(t, err) + return layer +} + // composeTestImage builds the standard two-layer fixture: a base layer with // content the top layer deletes, masks, replaces, and extends. func composeTestImage(t *testing.T) gcr.Image { @@ -45,11 +87,8 @@ func composeTestImage(t *testing.T) gcr.Image { return img } -// composeFixture composes the standard fixture image into the shared OCI cache -// and returns a client plus its validated manifest model. -func composeFixture(t *testing.T, p *paths.Paths) (*ociClient, string, *imageManifestModel) { - t.Helper() - +func TestComposeRootfsWhiteoutsAndOrdering(t *testing.T) { + p := paths.New(t.TempDir()) img := composeTestImage(t) writeLayerTestLayout(t, p, img) @@ -57,22 +96,15 @@ func composeFixture(t *testing.T, p *paths.Paths) (*ociClient, string, *imageMan require.NoError(t, err) digest, err := img.Digest() require.NoError(t, err) - tag := digestToLayoutTag(digest.String()) - bundle, err := client.extractOCIImageBundle(tag) + model, err := client.extractManifestModel(digestToLayoutTag(digest.String())) require.NoError(t, err) - return client, tag, bundle.Model -} - -func TestComposeRootfsWhiteoutsAndOrdering(t *testing.T) { - p := paths.New(t.TempDir()) - client, tag, model := composeFixture(t, p) require.Len(t, model.Layers, 2) dest := filepath.Join(t.TempDir(), "rootfs") - require.NoError(t, client.composeRootfs(dest, tag, model)) + require.NoError(t, client.composeRootfs(dest, model.Layers)) // Whiteout removed the base entry. - _, err := os.Lstat(filepath.Join(dest, "etc", "config.txt")) + _, err = os.Lstat(filepath.Join(dest, "etc", "config.txt")) require.True(t, os.IsNotExist(err), "whiteout must delete the base entry") // Plain replacement. @@ -108,75 +140,25 @@ func TestComposeRootfsWhiteoutsAndOrdering(t *testing.T) { })) } -// zeroLayerModel returns a schema-valid manifest model with no layers. -func zeroLayerModel() *imageManifestModel { - return &imageManifestModel{ - SchemaVersion: manifestModelSchemaVersion, - Digest: "sha256:" + strings.Repeat("ab", 32), - Config: manifestConfigRef{Digest: "sha256:" + strings.Repeat("cd", 32)}, - Layers: make([]layerDescriptor, 0), - } -} - func TestComposeRootfsEmptyLayers(t *testing.T) { p := paths.New(t.TempDir()) client, err := newOCIClient(p.SystemOCICache()) require.NoError(t, err) - model := zeroLayerModel() - err = client.composeRootfs(t.TempDir(), model.Digest, model) + err = client.composeRootfs(t.TempDir(), nil) require.ErrorContains(t, err, "no layers") } -func TestComposeRootfsInvalidModel(t *testing.T) { - p := paths.New(t.TempDir()) - client, tag, model := composeFixture(t, p) - - model.Config.DiffIDs = model.Config.DiffIDs[:1] - err := client.composeRootfs(filepath.Join(t.TempDir(), "rootfs"), tag, model) - require.ErrorContains(t, err, "1 diff ids for 2 layers") -} - func TestComposeRootfsMissingBlob(t *testing.T) { p := paths.New(t.TempDir()) client, err := newOCIClient(p.SystemOCICache()) require.NoError(t, err) - - digestHex := "sha256:" + strings.Repeat("ab", 32) - model := &imageManifestModel{ - SchemaVersion: manifestModelSchemaVersion, - Digest: digestHex, - Config: manifestConfigRef{ - Digest: "sha256:" + strings.Repeat("cd", 32), - DiffIDs: []string{"sha256:" + strings.Repeat("ef", 32)}, - }, - Layers: []layerDescriptor{{ - Digest: "sha256:" + strings.Repeat("01", 32), - MediaType: "application/vnd.oci.image.layer.v1.tar+gzip", - DiffID: "sha256:" + strings.Repeat("ef", 32), - }}, - } - err = client.composeRootfs(t.TempDir(), digestHex, model) + err = client.composeRootfs(t.TempDir(), []layerDescriptor{{ + Digest: "sha256:abababababababababababababababababababababababababababababababab", + MediaType: "application/vnd.oci.image.layer.v1.tar+gzip", + }}) require.ErrorContains(t, err, "missing from oci cache") } -func TestComposeRootfsDiffIDMismatch(t *testing.T) { - p := paths.New(t.TempDir()) - client, tag, model := composeFixture(t, p) - - // Replace the top layer's cached blob with different content so the - // unpacked diff id no longer matches the descriptor. - other := specLayer(t, []tarEntrySpec{{name: "other.txt", content: "other", mode: 0644}}) - otherBlob, err := other.Compressed() - require.NoError(t, err) - data, err := io.ReadAll(otherBlob) - require.NoError(t, err) - topHex := strings.TrimPrefix(model.Layers[1].Digest, "sha256:") - require.NoError(t, os.WriteFile(p.OCICacheBlob(topHex), data, 0644)) - - err = client.composeRootfs(filepath.Join(t.TempDir(), "rootfs"), tag, model) - require.ErrorContains(t, err, "diff id mismatch") -} - // TestComposeRootfsExportsValidErofs composes the fixture image and exports it // to erofs, then verifies the filesystem is intact and its contents match the // composed tree. @@ -189,10 +171,18 @@ func TestComposeRootfsExportsValidErofs(t *testing.T) { } p := paths.New(t.TempDir()) - client, tag, model := composeFixture(t, p) + img := composeTestImage(t) + writeLayerTestLayout(t, p, img) + + client, err := newOCIClient(p.SystemOCICache()) + require.NoError(t, err) + digest, err := img.Digest() + require.NoError(t, err) + model, err := client.extractManifestModel(digestToLayoutTag(digest.String())) + require.NoError(t, err) staging := filepath.Join(t.TempDir(), "rootfs") - require.NoError(t, client.composeRootfs(staging, tag, model)) + require.NoError(t, client.composeRootfs(staging, model.Layers)) diskPath := filepath.Join(t.TempDir(), "rootfs.erofs") size, err := ExportRootfs(staging, diskPath, FormatErofs) diff --git a/lib/images/credentials_test.go b/lib/images/credentials_test.go index 7f680b980..87a72c39b 100644 --- a/lib/images/credentials_test.go +++ b/lib/images/credentials_test.go @@ -83,10 +83,8 @@ func TestCreateImageRequestCredentialsAreNotPersisted(t *testing.T) { } func TestInflightPullRejectsDifferentCredentials(t *testing.T) { - m := &manager{ - inflightPulls: make(map[string]*inflightImagePull), - borrowedCredentialsTimeout: time.Minute, - } + m := newTestManager(nil) + m.borrowedCredentialsTimeout = time.Minute const digest = "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" credentials := &authn.AuthConfig{Username: "AWS", Password: "token-a"} inflight := m.registerInflightPull(digest, credentials) @@ -98,10 +96,8 @@ func TestInflightPullRejectsDifferentCredentials(t *testing.T) { } func TestBorrowedCredentialsExpireWhileQueued(t *testing.T) { - m := &manager{ - inflightPulls: make(map[string]*inflightImagePull), - borrowedCredentialsTimeout: time.Millisecond, - } + m := newTestManager(nil) + m.borrowedCredentialsTimeout = time.Millisecond const digest = "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" inflight := m.registerInflightPull(digest, &authn.AuthConfig{Username: "AWS", Password: "secret"}) defer m.releaseInflightPull(digest, inflight)() @@ -119,7 +115,7 @@ func TestBorrowedCredentialsExpireWhileQueued(t *testing.T) { } func TestBorrowedAuthRejectsReplacedInflightPull(t *testing.T) { - m := &manager{inflightPulls: make(map[string]*inflightImagePull)} + m := newTestManager(nil) const digest = "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" first := m.registerInflightPull(digest, &authn.AuthConfig{Username: "first"}) second := m.registerInflightPull(digest, &authn.AuthConfig{Username: "second"}) @@ -174,12 +170,9 @@ func TestRecoverInterruptedCredentialedPullFailsForFreshRetry(t *testing.T) { p := paths.New(t.TempDir()) client, err := newOCIClient(p.SystemOCICache()) require.NoError(t, err) - m := &manager{ - paths: p, - ociClient: client, - queue: queue.New(1), - readySubscribers: make(map[string][]chan StatusEvent), - } + m := newTestManager(p) + m.ociClient = client + m.queue = queue.New(1) const repository = "registry.example/private/image" const digest = "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" diff --git a/lib/images/disk_usage.go b/lib/images/disk_usage.go index b65dba1ac..c1cffbb2c 100644 --- a/lib/images/disk_usage.go +++ b/lib/images/disk_usage.go @@ -108,57 +108,80 @@ func totalOCICacheBlobBytesFromFilesystem(blobDir string) (int64, error) { return total, nil } -func (m *manager) getDiskUsageTotals() (int64, int64, error) { +// diskUsageTotals caches the size of each disk component the manager tracks. +type diskUsageTotals struct { + readyImageBytes int64 + layerBytes int64 + ociCacheBytes int64 +} + +func (m *manager) getDiskUsageTotals() (diskUsageTotals, error) { m.diskUsageMu.RLock() if m.diskUsageLoaded { - readyImageBytes := m.readyImageBytes - ociCacheBytes := m.ociCacheBytes + totals := diskUsageTotals{ + readyImageBytes: m.readyImageBytes, + layerBytes: m.layerBytes, + ociCacheBytes: m.ociCacheBytes, + } m.diskUsageMu.RUnlock() - return readyImageBytes, ociCacheBytes, nil + return totals, nil } m.diskUsageMu.RUnlock() - readyImageBytes, ociCacheBytes, err := m.computeDiskUsageTotals() + computed, err := m.computeDiskUsageTotals() if err != nil { - return 0, 0, err + return diskUsageTotals{}, err } m.diskUsageMu.Lock() if !m.diskUsageLoaded { - m.readyImageBytes = readyImageBytes - m.ociCacheBytes = ociCacheBytes + m.readyImageBytes = computed.readyImageBytes + m.layerBytes = computed.layerBytes + m.ociCacheBytes = computed.ociCacheBytes m.diskUsageLoaded = true } - readyImageBytes = m.readyImageBytes - ociCacheBytes = m.ociCacheBytes + totals := diskUsageTotals{ + readyImageBytes: m.readyImageBytes, + layerBytes: m.layerBytes, + ociCacheBytes: m.ociCacheBytes, + } m.diskUsageMu.Unlock() - return readyImageBytes, ociCacheBytes, nil + return totals, nil } func (m *manager) refreshDiskUsageTotals() { - readyImageBytes, ociCacheBytes, err := m.computeDiskUsageTotals() + computed, err := m.computeDiskUsageTotals() if err != nil { return } m.diskUsageMu.Lock() - m.readyImageBytes = readyImageBytes - m.ociCacheBytes = ociCacheBytes + m.readyImageBytes = computed.readyImageBytes + m.layerBytes = computed.layerBytes + m.ociCacheBytes = computed.ociCacheBytes m.diskUsageLoaded = true m.diskUsageMu.Unlock() } -func (m *manager) computeDiskUsageTotals() (int64, int64, error) { +func (m *manager) computeDiskUsageTotals() (diskUsageTotals, error) { readyImageBytes, err := totalReadyImageBytesFromMetadata(m.paths.ImagesDir()) if err != nil { - return 0, 0, err + return diskUsageTotals{}, err + } + layerBytes, err := totalLayerArtifactBytes(m.paths.ImageLayersDir()) + if err != nil { + return diskUsageTotals{}, err } ociCacheBytes, err := totalOCICacheBlobBytesFromFilesystem(m.paths.OCICacheBlobDir()) if err != nil { - return 0, 0, err + return diskUsageTotals{}, err } - return readyImageBytes, ociCacheBytes, nil + return diskUsageTotals{ + readyImageBytes: readyImageBytes, + layerBytes: layerBytes, + ociCacheBytes: ociCacheBytes, + }, nil } func totalRootfsBytesInDigestDir(digestDir string) (int64, error) { diff --git a/lib/images/layer_artifact.go b/lib/images/layer_artifact.go index 4283f8410..db64cb4bf 100644 --- a/lib/images/layer_artifact.go +++ b/lib/images/layer_artifact.go @@ -2,7 +2,6 @@ package images import ( "archive/tar" - "bytes" "compress/gzip" "crypto/sha256" "encoding/json" @@ -34,37 +33,23 @@ const ( const layerRecordSchemaVersion = 1 // layerArtifact is the persisted record for one materialized layer artifact. -// The key is the compressed layer blob digest plus the artifact format, so the -// same layer can coexist in several materializations. +// The key is the compressed layer blob digest plus the artifact format and +// options, so the same layer can coexist in several materializations. type layerArtifact struct { - SchemaVersion int `json:"schema_version"` - Digest string `json:"digest"` // compressed layer blob digest, sha256:... - DiffID string `json:"diff_id,omitempty"` - Format string `json:"format"` - SizeBytes int64 `json:"size_bytes"` // artifact bytes on disk - UnpackedBytes int64 `json:"unpacked_bytes"` - Entries int `json:"entries"` - Whiteouts []whiteoutRecord `json:"whiteouts,omitempty"` - CreatedAt time.Time `json:"created_at"` -} - -// validate checks a record read back from disk. The format fully determines -// the artifact options (erofs is always lz4-compressed, ext4 uncompressed), -// so only the format is stored. -func (a *layerArtifact) validate() error { - if a.SchemaVersion != layerRecordSchemaVersion { - return fmt.Errorf("unsupported schema version: %d", a.SchemaVersion) - } - if a.Digest == "" { - return fmt.Errorf("missing digest") - } - if a.Format != layerFormatErofs && a.Format != layerFormatExt4 { - return fmt.Errorf("invalid format: %s", a.Format) - } - if a.SizeBytes < 0 || a.UnpackedBytes < 0 || a.Entries < 0 { - return fmt.Errorf("invalid size or entry counts") - } - return nil + SchemaVersion int `json:"schema_version"` + Digest string `json:"digest"` // compressed layer blob digest, sha256:... + DiffID string `json:"diff_id,omitempty"` + Format string `json:"format"` + Options layerArtifactOptions `json:"options,omitempty"` + SizeBytes int64 `json:"size_bytes"` // artifact bytes on disk + UnpackedBytes int64 `json:"unpacked_bytes"` + Entries int `json:"entries"` + Whiteouts []whiteoutRecord `json:"whiteouts,omitempty"` + CreatedAt time.Time `json:"created_at"` +} + +type layerArtifactOptions struct { + Compression string `json:"compression,omitempty"` } // whiteoutRecord describes one whiteout marker found in a layer. Dir is the @@ -124,67 +109,67 @@ func readLayerRecord(p *paths.Paths, layerHex string) (*layerArtifact, error) { return &record, nil } -// materializeLayerArtifact ensures a layer has a materialized artifact keyed -// by its blob digest, building it from the shared OCI cache blob when absent. -// unpackCachedLayer resolves a layer blob in the shared OCI cache and unpacks -// it into dest, verifying the unpacked diff id against the descriptor. -func unpackCachedLayer(cacheDir string, desc layerDescriptor, dest string) (*unpackStats, error) { - layerHex := strings.TrimPrefix(desc.Digest, "sha256:") - if err := paths.ValidatePathComponent(layerHex); err != nil { - return nil, fmt.Errorf("invalid layer digest: %s", desc.Digest) +func (a *layerArtifact) validate() error { + if a.SchemaVersion != layerRecordSchemaVersion { + return fmt.Errorf("unsupported schema version: %d", a.SchemaVersion) } - blobPath := filepath.Join(cacheDir, "blobs", "sha256", layerHex) - if _, err := os.Stat(blobPath); err != nil { - if os.IsNotExist(err) { - return nil, fmt.Errorf("layer blob missing from oci cache: %s", desc.Digest) - } - return nil, fmt.Errorf("stat layer blob: %w", err) + if a.Digest == "" || (a.Format != layerFormatErofs && a.Format != layerFormatExt4) { + return fmt.Errorf("invalid digest or format") } - - stats, err := unpackLayerBlob(blobPath, desc.MediaType, dest) - if err != nil { - return nil, fmt.Errorf("unpack layer %s: %w", desc.Digest, err) + if a.SizeBytes < 0 || a.UnpackedBytes < 0 || a.Entries < 0 { + return fmt.Errorf("invalid size or entry counts") } - if stats.diffID != desc.DiffID { - return nil, fmt.Errorf("layer %s diff id mismatch: got %s, want %s", desc.Digest, stats.diffID, desc.DiffID) + if a.Format == layerFormatExt4 && a.Options.Compression != "" { + return fmt.Errorf("ext4 artifact has compression options") } - return stats, nil + if a.Format == layerFormatErofs && a.Options.Compression != "lz4" { + return fmt.Errorf("erofs artifact has invalid compression options") + } + return nil } -// The layer is unpacked into an isolated temp directory, converted to erofs, -// and installed atomically; an interrupted build leaves only temp files that +// materializeLayerArtifact ensures a layer has a materialized artifact keyed +// by its blob digest, converting the layer's staging tree from the pull and +// installing it atomically; an interrupted build leaves only temp files that // the next attempt replaces. -func (m *manager) materializeLayerArtifact(desc layerDescriptor) (*layerArtifact, error) { - layerHex := strings.TrimPrefix(desc.Digest, "sha256:") - if err := paths.ValidatePathComponent(layerHex); err != nil { - return nil, fmt.Errorf("invalid layer digest: %s", desc.Digest) - } - - if record, err := readLayerRecord(m.paths, layerHex); err != nil { +func (m *manager) materializeLayerArtifact(desc layerDescriptor, tree layerTree) (*layerArtifact, error) { + layerHex, err := layerDigestHex(desc) + if err != nil { return nil, err - } else if record != nil && record.matches(desc) { - if _, statErr := os.Stat(layerArtifactPath(m.paths, layerHex)); statErr == nil { - return record, nil - } - // Record without artifact: rebuild below. + } + if desc.DiffID != "" && tree.stats.diffID != desc.DiffID { + return nil, fmt.Errorf("layer %s diff id mismatch: got %s, want %s", desc.Digest, tree.stats.diffID, desc.DiffID) } - layerDir := m.paths.ImageLayerDir(layerHex) - if err := os.MkdirAll(layerDir, 0755); err != nil { - return nil, fmt.Errorf("create layer directory: %w", err) + lock := m.layerDigestLock(layerHex) + lock.Lock() + defer lock.Unlock() + if record, err := m.existingLayerArtifact(layerHex, desc); record != nil || err != nil { + return record, err } - unpackDir, err := os.MkdirTemp(layerDir, ".unpack-*") - if err != nil { - return nil, fmt.Errorf("create unpack directory: %w", err) + return m.installLayerArtifact(desc, layerHex, tree.path, tree.stats) +} + +func layerDigestHex(desc layerDescriptor) (string, error) { + layerHex := strings.TrimPrefix(desc.Digest, "sha256:") + if err := paths.ValidatePathComponent(layerHex); err != nil { + return "", fmt.Errorf("invalid layer digest: %s", desc.Digest) } - defer os.RemoveAll(unpackDir) + return layerHex, nil +} - stats, err := unpackCachedLayer(m.paths.SystemOCICache(), desc, unpackDir) +func (m *manager) existingLayerArtifact(layerHex string, desc layerDescriptor) (*layerArtifact, error) { + record, err := readLayerRecord(m.paths, layerHex) if err != nil { return nil, err } - - return m.installLayerArtifact(desc, layerHex, unpackDir, stats) + if record == nil || !record.matches(desc) { + return nil, nil + } + if _, err := os.Stat(layerArtifactPath(m.paths, layerHex)); err != nil { + return nil, nil + } + return record, nil } func (m *manager) installLayerArtifact(desc layerDescriptor, layerHex, unpackDir string, stats *unpackStats) (*layerArtifact, error) { @@ -193,18 +178,22 @@ func (m *manager) installLayerArtifact(desc layerDescriptor, layerHex, unpackDir Digest: desc.Digest, DiffID: desc.DiffID, Format: layerArtifactFormat(), + Options: artifactOptions(layerArtifactFormat()), UnpackedBytes: stats.unpackedBytes, Entries: stats.entries, - // Nothing reads Whiteouts yet; composition re-derives whiteouts - // from the unpacked tree once it lands. - Whiteouts: stats.whiteouts, - CreatedAt: time.Now(), + Whiteouts: stats.whiteouts, + CreatedAt: time.Now(), } if err := installAtomically(layerArtifactPath(m.paths, layerHex), func(path string) error { - // The artifact intentionally retains .wh. marker files so - // composition can re-derive whiteouts from the tree itself. - size, convErr := ExportRootfs(unpackDir, path, DefaultImageFormat) + var size int64 + var convErr error + switch layerArtifactFormat() { + case layerFormatExt4: + size, convErr = convertToExt4(unpackDir, path) + default: + size, convErr = convertToErofs(unpackDir, path) + } if convErr != nil { return convErr } @@ -225,6 +214,18 @@ func (m *manager) installLayerArtifact(desc layerDescriptor, layerHex, unpackDir return record, nil } +func artifactOptions(format string) layerArtifactOptions { + if format == layerFormatErofs { + return layerArtifactOptions{Compression: "lz4"} + } + return layerArtifactOptions{} +} + +type layerTree struct { + path string + stats *unpackStats +} + type unpackStats struct { entries int unpackedBytes int64 @@ -232,6 +233,18 @@ type unpackStats struct { whiteouts []whiteoutRecord } +func cleanupLayerTree(tree layerTree) { + if tree.path != "" { + _ = os.RemoveAll(tree.path) + } +} + +func cleanupLayerTrees(trees map[string]layerTree) { + for _, tree := range trees { + cleanupLayerTree(tree) + } +} + // unpackLayerBlob extracts one compressed layer blob into dest, preserving // whiteout marker files and recording them. Paths are confined to dest. func unpackLayerBlob(blobPath, mediaType, dest string) (*unpackStats, error) { @@ -249,9 +262,6 @@ func unpackLayerBlob(blobPath, mediaType, dest string) (*unpackStats, error) { hash := sha256.New() stats := &unpackStats{whiteouts: make([]whiteoutRecord, 0)} - // Directory metadata is re-applied after extraction, once children - // exist, so tar directory mtimes are not overwritten by later writes. - var pendingDirs []pendingDir hashedReader := io.TeeReader(reader, hash) tr := tar.NewReader(hashedReader) for { @@ -283,9 +293,6 @@ func unpackLayerBlob(blobPath, mediaType, dest string) (*unpackStats, error) { stats.whiteouts = append(stats.whiteouts, whiteoutRecord{Dir: dir, Target: targetName}) } - if header.Typeflag == tar.TypeDir { - pendingDirs = append(pendingDirs, pendingDir{target: target, header: header}) - } if err := extractTarEntry(tr, header, dest, target); err != nil { return nil, fmt.Errorf("extract %s: %w", header.Name, err) } @@ -296,26 +303,14 @@ func unpackLayerBlob(blobPath, mediaType, dest string) (*unpackStats, error) { if _, err := io.Copy(io.Discard, hashedReader); err != nil { return nil, fmt.Errorf("drain layer: %w", err) } - for _, dir := range pendingDirs { - if err := applyTarMetadata(dir.target, dir.header); err != nil { - return nil, fmt.Errorf("restore dir metadata %s: %w", dir.target, err) - } - } stats.diffID = fmt.Sprintf("sha256:%x", hash.Sum(nil)) return stats, nil } -type pendingDir struct { - target string - header *tar.Header -} - -// decompressLayer wraps the blob in the reader for its layer media type. Both -// OCI-style suffixes (+gzip, +zstd) and docker-style media types (tar.gzip, -// tar.zstd) are matched so neither encoding falls through to the raw path. +// decompressLayer wraps the blob in the reader for its layer media type. func decompressLayer(blob *os.File, mediaType string) (io.Reader, io.Closer, error) { switch { - case strings.HasSuffix(mediaType, "+zstd"), strings.Contains(mediaType, "tar.zstd"): + case strings.HasSuffix(mediaType, "+zstd"): decoder, err := zstd.NewReader(blob) if err != nil { return nil, nil, fmt.Errorf("zstd reader: %w", err) @@ -344,12 +339,7 @@ func (c multiCloser) Close() error { return firstErr } -// safeJoin resolves a tar entry name inside root and rejects symlinked -// parents: extraction must never create an entry through a symlink an earlier -// tar entry planted, so existing parents are Lstat-walked and rejected rather -// than resolved. Symlink entries themselves may legitimately name paths that -// do not exist yet, so their targets are checked by resolve in -// validateSymlinkTarget instead. +// safeJoin resolves a tar entry name inside root and rejects symlinked parents. func safeJoin(root, name string) (string, error) { if filepath.IsAbs(name) { return "", fmt.Errorf("tar entry escapes root: %s", name) @@ -363,10 +353,7 @@ func safeJoin(root, name string) (string, error) { } root = filepath.Clean(root) target := filepath.Join(root, clean) - if target == root { - return target, nil - } - if !strings.HasPrefix(target, root+string(filepath.Separator)) { + if target != root && !strings.HasPrefix(target, root+string(filepath.Separator)) { return "", fmt.Errorf("tar entry escapes root: %s", name) } for parent := filepath.Dir(target); parent != root; parent = filepath.Dir(parent) { @@ -402,7 +389,7 @@ func validateSymlinkTarget(root, target, linkname string) error { func extractTarEntry(tr *tar.Reader, header *tar.Header, root, target string) error { switch header.Typeflag { case tar.TypeDir: - return extractTarDir(target) + return extractTarDir(target, header) case tar.TypeReg: return extractTarFile(tr, target, header) case tar.TypeSymlink: @@ -422,16 +409,19 @@ func prepareTarTarget(target string) error { if err := os.MkdirAll(filepath.Dir(target), 0755); err != nil { return err } - return removePath(target) + return clearExisting(target) } -func extractTarDir(target string) error { +func extractTarDir(target string, header *tar.Header) error { if info, err := os.Lstat(target); err == nil && !info.IsDir() { - if err := removePath(target); err != nil { + if err := clearExisting(target); err != nil { return err } } - return os.MkdirAll(target, 0755) + if err := os.MkdirAll(target, 0755); err != nil { + return err + } + return applyTarMetadata(target, header) } func extractTarFile(tr *tar.Reader, target string, header *tar.Header) error { @@ -538,8 +528,6 @@ func applyTarMetadata(path string, header *tar.Header) error { return nil } -// removePath removes whatever entry occupies path, including non-empty -// directories, and tolerates a missing path. func removePath(path string) error { if err := os.RemoveAll(path); err != nil && !os.IsNotExist(err) { return err @@ -547,12 +535,27 @@ func removePath(path string) error { return nil } +// clearExisting removes whatever entry occupies path, including non-empty +// directories, so a layer entry of a different type can replace it. +func clearExisting(path string) error { + info, err := os.Lstat(path) + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + if info.IsDir() { + return os.RemoveAll(path) + } + return os.Remove(path) +} + // applyLayerTree merges one unpacked layer directory into targetDir following // OCI whiteout semantics: whiteouts and opaque markers remove what lower layers // contributed, then the layer's own entries are copied on top. Raw tar // whiteout files are interpreted here rather than passed through, because -// overlayfs does not understand them. No production caller yet; composition -// lands in a later change. +// overlayfs does not understand them. func applyLayerTree(layerDir, targetDir string) error { if err := os.MkdirAll(targetDir, 0755); err != nil { return fmt.Errorf("create target directory: %w", err) @@ -579,6 +582,9 @@ func applyLayerTree(layerDir, targetDir string) error { return clearDirContents(targetParent) } hidden := strings.TrimPrefix(base, whiteoutPrefix) + if hidden == "" || hidden == "." || hidden == ".." { + return fmt.Errorf("invalid whiteout entry: %s", filepath.Join(filepath.Dir(rel), base)) + } target, err := safeJoin(targetDir, filepath.Join(filepath.Dir(rel), hidden)) if err != nil { return err @@ -590,9 +596,6 @@ func applyLayerTree(layerDir, targetDir string) error { } // Phase 2: copy the layer's own entries, skipping whiteout markers. - // Directory metadata is deferred until all children are copied, so tar - // directory mtimes survive the merge. - var pendingDirs []dirMeta hardlinks := make(map[hardlinkIdentity]string) err = filepath.WalkDir(layerDir, func(path string, entry fs.DirEntry, err error) error { if err != nil { @@ -615,31 +618,14 @@ func applyLayerTree(layerDir, targetDir string) error { if err != nil { return err } - if entry.IsDir() { - info, err := entry.Info() - if err != nil { - return err - } - pendingDirs = append(pendingDirs, dirMeta{src: path, dst: target, info: info}) - } return copyEntryInto(path, target, hardlinks) }) if err != nil { return fmt.Errorf("copy layer tree: %w", err) } - for _, dir := range pendingDirs { - if err := copyEntryMetadata(dir.src, dir.dst, dir.info); err != nil { - return fmt.Errorf("restore dir metadata %s: %w", dir.dst, err) - } - } return nil } -type dirMeta struct { - src, dst string - info os.FileInfo -} - // clearDirContents removes everything inside dir without removing dir itself, // and without following symlinks. func clearDirContents(dir string) error { @@ -716,7 +702,10 @@ func copyDirectoryEntry(src, dst string, info os.FileInfo) error { return err } } - return os.MkdirAll(dst, info.Mode().Perm()) + if err := os.MkdirAll(dst, info.Mode().Perm()); err != nil { + return err + } + return copyEntryMetadata(src, dst, info) } func copySymlinkEntry(src, dst string) error { @@ -743,19 +732,10 @@ func copySpecialEntry(src, dst string, info os.FileInfo) error { } mode, err := specialFileMode(info.Mode() & fs.ModeType) if err != nil { - return fmt.Errorf("unsupported entry type for %s: %w", src, err) + return fmt.Errorf("unsupported entry type for %s", src) } if err := unix.Mknod(dst, mode|uint32(info.Mode().Perm()), int(stat.Rdev)); err != nil { - if !errors.Is(err, unix.EPERM) { - return fmt.Errorf("mknod: %w", err) - } - file, openErr := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|syscall.O_NOFOLLOW, 0644) - if openErr != nil { - return fmt.Errorf("create rootless device placeholder: %w", openErr) - } - if closeErr := file.Close(); closeErr != nil { - return closeErr - } + return err } return copyEntryMetadata(src, dst, info) } @@ -810,11 +790,10 @@ func copyXattrs(src, dst string) error { } names = names[:n] } - for _, attr := range bytes.Split(bytes.TrimSuffix(names, []byte{0}), []byte{0}) { - if len(attr) == 0 { + for _, name := range strings.Split(strings.TrimSuffix(string(names), "\\x00"), "\\x00") { + if name == "" { continue } - name := string(attr) size, err := unix.Lgetxattr(src, name, nil) if err != nil { if errors.Is(err, unix.ENOTSUP) || errors.Is(err, unix.EPERM) || errors.Is(err, unix.ENODATA) { diff --git a/lib/images/layer_artifact_test.go b/lib/images/layer_artifact_test.go index d92940b7a..2edf284f9 100644 --- a/lib/images/layer_artifact_test.go +++ b/lib/images/layer_artifact_test.go @@ -5,18 +5,13 @@ import ( "bytes" "compress/gzip" "crypto/sha256" - "errors" "fmt" "io" "io/fs" "os" - "os/exec" "path/filepath" "strings" "testing" - "time" - - "golang.org/x/sys/unix" gcr "github.com/google/go-containerregistry/pkg/v1" "github.com/google/go-containerregistry/pkg/v1/empty" @@ -56,19 +51,21 @@ func layerDescFromImage(t *testing.T, img gcr.Image, index int) layerDescriptor } func TestMaterializeLayerArtifact(t *testing.T) { - if _, err := exec.LookPath("mkfs.erofs"); err != nil { - t.Skip("mkfs.erofs not available") - } - p := paths.New(t.TempDir()) img, err := mutate.AppendLayers(empty.Image, syntheticLayer(t, "base.txt", "base layer content")) require.NoError(t, err) writeLayerTestLayout(t, p, img) desc := layerDescFromImage(t, img, 0) - m := &manager{paths: p} + client, err := newOCIClient(p.SystemOCICache()) + require.NoError(t, err) + m := newTestManager(p) + + tree, err := client.extractLayerTree(desc) + require.NoError(t, err) + defer cleanupLayerTree(tree) - record, err := m.materializeLayerArtifact(desc) + record, err := m.materializeLayerArtifact(desc, tree) require.NoError(t, err) require.Equal(t, desc.Digest, record.Digest) require.Equal(t, desc.DiffID, record.DiffID) @@ -78,25 +75,27 @@ func TestMaterializeLayerArtifact(t *testing.T) { require.Greater(t, record.Entries, 0) layerHex := desc.Digest[len("sha256:"):] - _, err = os.Stat(p.ImageLayerArtifactForFormat(layerHex, layerArtifactFormat())) - require.NoError(t, err, "layer.erofs must be installed") + artifactPath := p.ImageLayerArtifactForFormat(layerHex, layerArtifactFormat()) + _, err = os.Stat(artifactPath) + require.NoError(t, err, "layer artifact must be installed") // A second materialization reuses the existing artifact. - artifactInfo, err := os.Stat(p.ImageLayerArtifactForFormat(layerHex, layerArtifactFormat())) + artifactInfo, err := os.Stat(artifactPath) require.NoError(t, err) - reused, err := m.materializeLayerArtifact(desc) + reused, err := m.materializeLayerArtifact(desc, tree) require.NoError(t, err) require.True(t, record.CreatedAt.Equal(reused.CreatedAt), "reuse must return the stored record") - artifactInfoAfter, err := os.Stat(p.ImageLayerArtifactForFormat(layerHex, layerArtifactFormat())) + artifactInfoAfter, err := os.Stat(artifactPath) require.NoError(t, err) require.Equal(t, artifactInfo.ModTime(), artifactInfoAfter.ModTime(), "reuse must not rebuild") } -func TestMaterializeLayerArtifactMissingBlob(t *testing.T) { +func TestExtractLayerTreeMissingBlob(t *testing.T) { p := paths.New(t.TempDir()) - m := &manager{paths: p} + client, err := newOCIClient(p.SystemOCICache()) + require.NoError(t, err) - _, err := m.materializeLayerArtifact(layerDescriptor{ + _, err = client.extractLayerTree(layerDescriptor{ Digest: "sha256:abababababababababababababababababababababababababababababababab", MediaType: "application/vnd.oci.image.layer.v1.tar+gzip", }) @@ -141,19 +140,21 @@ func whiteoutLayer(t *testing.T) gcr.Layer { } func TestMaterializeLayerRecordsWhiteouts(t *testing.T) { - if _, err := exec.LookPath("mkfs.erofs"); err != nil { - t.Skip("mkfs.erofs not available") - } - p := paths.New(t.TempDir()) img, err := mutate.AppendLayers(empty.Image, whiteoutLayer(t)) require.NoError(t, err) writeLayerTestLayout(t, p, img) desc := layerDescFromImage(t, img, 0) - m := &manager{paths: p} + client, err := newOCIClient(p.SystemOCICache()) + require.NoError(t, err) + m := newTestManager(p) + + tree, err := client.extractLayerTree(desc) + require.NoError(t, err) + defer cleanupLayerTree(tree) - record, err := m.materializeLayerArtifact(desc) + record, err := m.materializeLayerArtifact(desc, tree) require.NoError(t, err) require.Contains(t, record.Whiteouts, whiteoutRecord{Dir: "gone", Target: "deleted.txt"}) @@ -322,69 +323,3 @@ func TestApplyLayerTreeSymlinksAndHardlinks(t *testing.T) { require.NoError(t, err) require.Equal(t, "a.txt", linkTarget) } - -func TestCopyXattrs(t *testing.T) { - root := t.TempDir() - src := filepath.Join(root, "src") - dst := filepath.Join(root, "dst") - require.NoError(t, os.WriteFile(src, []byte("payload"), 0644)) - require.NoError(t, os.WriteFile(dst, []byte("payload"), 0644)) - - err := unix.Lsetxattr(src, "user.one", []byte("1"), 0) - if errors.Is(err, unix.ENOTSUP) || errors.Is(err, unix.EPERM) { - t.Skip("filesystem does not support user xattrs") - } - require.NoError(t, err) - require.NoError(t, unix.Lsetxattr(src, "user.two", []byte("22"), 0)) - - require.NoError(t, copyXattrs(src, dst)) - - for name, want := range map[string]string{"user.one": "1", "user.two": "22"} { - size, err := unix.Lgetxattr(dst, name, nil) - require.NoError(t, err, "xattr %s must be copied", name) - value := make([]byte, size) - n, err := unix.Lgetxattr(dst, name, value) - require.NoError(t, err) - require.Equal(t, want, string(value[:n])) - } -} - -func TestUnpackLayerBlobPreservesDirMtime(t *testing.T) { - root := t.TempDir() - blobPath := filepath.Join(root, "layer.tar") - dirTime := time.Now().Add(-time.Hour).Truncate(time.Second) - - var buf bytes.Buffer - tw := tar.NewWriter(&buf) - require.NoError(t, tw.WriteHeader(&tar.Header{Name: "d/", Typeflag: tar.TypeDir, Mode: 0755, ModTime: dirTime})) - require.NoError(t, tw.WriteHeader(&tar.Header{Name: "d/file.txt", Typeflag: tar.TypeReg, Mode: 0644, Size: 1})) - _, err := tw.Write([]byte("x")) - require.NoError(t, err) - require.NoError(t, tw.Close()) - require.NoError(t, os.WriteFile(blobPath, buf.Bytes(), 0644)) - - dest := filepath.Join(root, "dest") - _, err = unpackLayerBlob(blobPath, "application/vnd.oci.image.layer.v1.tar", dest) - require.NoError(t, err) - - info, err := os.Stat(filepath.Join(dest, "d")) - require.NoError(t, err) - require.True(t, info.ModTime().Equal(dirTime), "dir mtime must come from the tar header") -} - -func TestApplyLayerTreePreservesDirMtime(t *testing.T) { - root := t.TempDir() - targetDir := filepath.Join(root, "target") - layerDir := filepath.Join(root, "layer") - - old := time.Now().Add(-time.Hour).Truncate(time.Second) - require.NoError(t, os.MkdirAll(filepath.Join(layerDir, "d"), 0700)) - require.NoError(t, os.WriteFile(filepath.Join(layerDir, "d", "file.txt"), []byte("x"), 0644)) - require.NoError(t, os.Chtimes(filepath.Join(layerDir, "d"), old, old)) - - require.NoError(t, applyLayerTree(layerDir, targetDir)) - - info, err := os.Stat(filepath.Join(targetDir, "d")) - require.NoError(t, err) - require.True(t, info.ModTime().Equal(old), "dir mtime must survive the merge") -} diff --git a/lib/images/layer_gc.go b/lib/images/layer_gc.go new file mode 100644 index 000000000..aefdadfc0 --- /dev/null +++ b/lib/images/layer_gc.go @@ -0,0 +1,212 @@ +package images + +import ( + "context" + "fmt" + "io/fs" + "log/slog" + "os" + "path/filepath" + "strings" + "time" +) + +// layerEvictionGracePeriod keeps freshly written layer artifacts and temp +// directories out of cleanup so recovery and eviction never race builds that +// are still writing them. +const layerEvictionGracePeriod = 10 * time.Minute + +// referencedLayerDigests returns the set of layer blob digests referenced by +// the manifest models of every image in the content layout, plus the digests +// currently referenced by in-flight builds. Layer artifacts in this set are +// protected from eviction. Unreadable manifest models are skipped with a +// warning so one corrupt record cannot disable eviction entirely. +func (m *manager) referencedLayerDigests() map[string]struct{} { + refs := m.inflightLayerRefSnapshot() + contentRoot := filepath.Join(m.paths.ImagesDir(), "content") + err := filepath.WalkDir(contentRoot, func(path string, entry fs.DirEntry, err error) error { + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + if entry.IsDir() || entry.Name() != "manifest.json" { + return nil + } + digestHex := filepath.Base(filepath.Dir(path)) + model, readErr := readManifestModel(m.paths, digestHex) + if readErr != nil { + slog.Warn("skipping unreadable manifest model for layer eviction", "digest", digestHex, "error", readErr) + return nil + } + if model == nil { + return nil + } + for _, layer := range model.Layers { + refs[strings.TrimPrefix(layer.Digest, "sha256:")] = struct{}{} + } + return nil + }) + if err != nil && !os.IsNotExist(err) { + slog.Warn("failed to walk content manifests for layer eviction", "error", err) + } + return refs +} + +// inflightLayerRefSnapshot returns the layer digests currently retained by +// in-flight builds. +func (m *manager) inflightLayerRefSnapshot() map[string]struct{} { + m.layerRefMu.Lock() + defer m.layerRefMu.Unlock() + refs := make(map[string]struct{}, len(m.inflightLayerRefs)) + for digestHex := range m.inflightLayerRefs { + refs[digestHex] = struct{}{} + } + return refs +} + +// reconcileLayerStore evicts unreferenced layer artifacts and refreshes the +// cached disk usage totals so accounting reflects the removals. +func (m *manager) reconcileLayerStore() { + m.evictUnreferencedLayerArtifacts() + m.refreshDiskUsageTotals() +} + +// evictUnreferencedLayerArtifacts removes layer artifacts that no image +// manifest model references, deleting the digest directory entirely. Artifacts +// newer than the grace period are kept so in-flight builds never lose work. +func (m *manager) evictUnreferencedLayerArtifacts() { + refs := m.referencedLayerDigests() + + layersDir := m.paths.ImageLayersDir() + entries, err := os.ReadDir(layersDir) + if err != nil { + if !os.IsNotExist(err) { + slog.Warn("layer eviction failed to list layer store", "error", err) + } + return + } + + cutoff := time.Now().Add(-m.layerEvictionGrace) + evicted := 0 + var evictedBytes int64 + for _, entry := range entries { + if !entry.IsDir() { + continue + } + digestHex := entry.Name() + if _, referenced := refs[digestHex]; referenced { + continue + } + size, removed := m.tryEvictLayerArtifact(digestHex, filepath.Join(layersDir, digestHex), cutoff) + if !removed { + continue + } + evicted++ + evictedBytes += size + } + if evicted > 0 { + slog.Info("evicted unreferenced layer artifacts", "count", evicted, "bytes", evictedBytes) + if m.metrics != nil { + m.metrics.layerArtifactsEvicted.Add(context.Background(), int64(evicted)) + } + } +} + +// tryEvictLayerArtifact removes one unreferenced layer artifact if it is still +// stale and no build is materializing it. The per-digest lock is taken with +// TryLock so eviction never blocks behind an in-flight conversion. +func (m *manager) tryEvictLayerArtifact(digestHex, dirPath string, cutoff time.Time) (int64, bool) { + lock := m.layerDigestLock(digestHex) + if !lock.TryLock() { + return 0, false + } + defer lock.Unlock() + + // The candidate was selected outside the lock; re-check that a build has + // not retained the digest and the artifact has not been rewritten since. + if _, referenced := m.inflightLayerRefSnapshot()[digestHex]; referenced { + return 0, false + } + info, statErr := os.Stat(dirPath) + if statErr != nil || info.ModTime().After(cutoff) { + return 0, false + } + size, err := dirSize(dirPath) + if err != nil { + slog.Warn("failed to measure layer artifact size", "digest", digestHex, "error", err) + } + if err := os.RemoveAll(dirPath); err != nil { + slog.Warn("failed to evict unreferenced layer artifact", "digest", digestHex, "error", err) + return 0, false + } + return size, true +} + +// cleanStaleImageTempDirs removes temp directories left behind by builds that +// were interrupted mid-install, mid-materialization, or mid-tag promotion. +// Only directories older than the grace period are removed so live builds are +// never disturbed. +func (m *manager) cleanStaleImageTempDirs() { + roots := []string{ + m.paths.ImageLayersDir(), + filepath.Join(m.paths.ImagesDir(), "content"), + } + cutoff := time.Now().Add(-m.layerEvictionGrace) + for _, root := range roots { + err := filepath.WalkDir(root, func(path string, entry fs.DirEntry, err error) error { + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + if !entry.IsDir() { + return nil + } + name := entry.Name() + if !strings.HasPrefix(name, ".unpack-") && !strings.HasPrefix(name, ".install-") && !strings.HasPrefix(name, ".tag-stage-") { + return nil + } + info, statErr := os.Stat(path) + if statErr == nil && info.ModTime().Before(cutoff) { + _ = os.RemoveAll(path) + } + return fs.SkipDir + }) + if err != nil && !os.IsNotExist(err) { + slog.Warn("failed to clean stale image temp dirs", "root", root, "error", err) + } + } +} + +// totalLayerArtifactBytes sums the bytes held by materialized layer +// artifacts, matching what diskutilization.Collect counts for the same store. +func totalLayerArtifactBytes(layersDir string) (int64, error) { + var total int64 + err := filepath.WalkDir(layersDir, func(path string, entry fs.DirEntry, err error) error { + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + if entry.IsDir() { + return nil + } + if !strings.HasPrefix(entry.Name(), "layer.") { + return nil + } + info, statErr := entry.Info() + if statErr != nil { + return nil + } + total += info.Size() + return nil + }) + if err != nil && !os.IsNotExist(err) { + return 0, fmt.Errorf("walk layer artifacts: %w", err) + } + return total, nil +} diff --git a/lib/images/lifecycle_test.go b/lib/images/lifecycle_test.go new file mode 100644 index 000000000..1554f1501 --- /dev/null +++ b/lib/images/lifecycle_test.go @@ -0,0 +1,207 @@ +package images + +import ( + "context" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "time" + + gcr "github.com/google/go-containerregistry/pkg/v1" + "github.com/google/go-containerregistry/pkg/v1/empty" + "github.com/google/go-containerregistry/pkg/v1/layout" + "github.com/google/go-containerregistry/pkg/v1/mutate" + "github.com/kernel/hypeman/lib/paths" + "github.com/stretchr/testify/require" +) + +// writeSharedLayout writes several images into one OCI layout cache, each +// annotated with its own digest tag. +func writeSharedLayout(t *testing.T, p *paths.Paths, imgs ...gcr.Image) []string { + t.Helper() + + layoutPath, err := layout.Write(p.SystemOCICache(), empty.Index) + require.NoError(t, err) + + digests := make([]string, 0, len(imgs)) + for _, img := range imgs { + digest, err := img.Digest() + require.NoError(t, err) + require.NoError(t, layoutPath.AppendImage(img, layout.WithAnnotations(map[string]string{ + "org.opencontainers.image.ref.name": digestToLayoutTag(digest.String()), + }))) + digests = append(digests, digest.String()) + } + return digests +} + +func layerHexes(t *testing.T, p *paths.Paths) map[string]struct{} { + t.Helper() + entries, err := os.ReadDir(p.ImageLayersDir()) + require.NoError(t, err) + hexes := make(map[string]struct{}) + for _, entry := range entries { + if entry.IsDir() { + hexes[entry.Name()] = struct{}{} + } + } + return hexes +} + +// TestSharedLayersMaterializeOnceAndEvictWithReferences is the end-to-end +// lifecycle: two images share a base layer, the shared artifact is created +// once, survives the deletion of one image, and is evicted only when its last +// reference is gone. +func TestSharedLayersMaterializeOnceAndEvictWithReferences(t *testing.T) { + if _, err := exec.LookPath("mkfs.erofs"); err != nil { + t.Skip("mkfs.erofs not available") + } + dataDir := t.TempDir() + p := paths.New(dataDir) + mgr, err := NewManager(p, 1, nil) + require.NoError(t, err) + m := mgr.(*manager) + m.layerEvictionGrace = 0 + + base := syntheticLayer(t, "base.txt", "shared base content") + topA := syntheticLayer(t, "a.txt", "app A payload") + topB := syntheticLayer(t, "b.txt", "app B payload") + + imgA, err := mutate.AppendLayers(empty.Image, base, topA) + require.NoError(t, err) + imgB, err := mutate.AppendLayers(empty.Image, base, topB) + require.NoError(t, err) + + digests := writeSharedLayout(t, p, imgA, imgB) + digestA, digestB := digests[0], digests[1] + + baseManifest, err := imgA.Manifest() + require.NoError(t, err) + baseHex := baseManifest.Layers[0].Digest.Hex + topAHex := baseManifest.Layers[1].Digest.Hex + topBManifest, err := imgB.Manifest() + require.NoError(t, err) + topBHex := topBManifest.Layers[1].Digest.Hex + + ctx := context.Background() + const repoA = "kernel.local/apps/app-a" + const repoB = "kernel.local/apps/app-b" + + eventsA := make(chan StatusEvent, 2) + m.subscribeToReady(digestToLayoutTag(digestA), eventsA) + defer m.unsubscribeFromReady(digestToLayoutTag(digestA), eventsA) + _, err = m.ImportLocalImage(ctx, repoA, "v1", digestA) + require.NoError(t, err) + select { + case event := <-eventsA: + require.Equal(t, StatusReady, event.Status) + case <-time.After(30 * time.Second): + t.Fatal("image A did not become ready") + } + + eventsB := make(chan StatusEvent, 2) + m.subscribeToReady(digestToLayoutTag(digestB), eventsB) + defer m.unsubscribeFromReady(digestToLayoutTag(digestB), eventsB) + _, err = m.ImportLocalImage(ctx, repoB, "v1", digestB) + require.NoError(t, err) + select { + case event := <-eventsB: + require.Equal(t, StatusReady, event.Status) + case <-time.After(30 * time.Second): + t.Fatal("image B did not become ready") + } + + // The shared base layer materialized exactly once, alongside the two tops. + hexes := layerHexes(t, p) + require.Len(t, hexes, 3) + require.Contains(t, hexes, baseHex) + require.Contains(t, hexes, topAHex) + require.Contains(t, hexes, topBHex) + + // Deleting image A evicts only its unique layer; the shared base survives. + require.NoError(t, m.DeleteImage(ctx, repoA+"@"+digestA)) + hexes = layerHexes(t, p) + require.Len(t, hexes, 2) + require.Contains(t, hexes, baseHex, "shared base must survive while referenced") + require.Contains(t, hexes, topBHex) + require.NotContains(t, hexes, topAHex) + + // Deleting image B removes the last references: everything is evicted. + require.NoError(t, m.DeleteImage(ctx, repoB+"@"+digestB)) + hexes = layerHexes(t, p) + require.Empty(t, hexes, "unreferenced layer artifacts must be evicted") +} + +func TestTotalImageBytesIncludesLayerArtifacts(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTestManager(p) + + digestHex := "cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01cd01" + require.NoError(t, os.MkdirAll(p.ImageLayerDir(digestHex), 0o755)) + payload := make([]byte, 4096) + require.NoError(t, os.WriteFile(p.ImageLayerArtifact(digestHex), payload, 0o644)) + + totals, err := m.getDiskUsageTotals() + require.NoError(t, err) + require.GreaterOrEqual(t, totals.layerBytes, int64(len(payload))) + + totalBytes, err := m.TotalImageBytes(context.Background()) + require.NoError(t, err) + require.Equal(t, totals.readyImageBytes+totals.layerBytes, totalBytes) +} + +func TestCleanStaleImageTempDirsRemovesOnlyOldDirectories(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTestManager(p) + m.layerEvictionGrace = time.Hour + + layersDir := p.ImageLayersDir() + staleDir := filepath.Join(layersDir, "ab12", ".unpack-stale") + freshDir := filepath.Join(layersDir, "cd34", ".unpack-fresh") + require.NoError(t, os.MkdirAll(staleDir, 0o755)) + require.NoError(t, os.MkdirAll(freshDir, 0o755)) + old := time.Now().Add(-2 * time.Hour) + require.NoError(t, os.Chtimes(staleDir, old, old)) + + m.cleanStaleImageTempDirs() + + _, err := os.Stat(staleDir) + require.True(t, os.IsNotExist(err), "stale temp dir must be removed") + _, err = os.Stat(freshDir) + require.NoError(t, err, "fresh temp dir must survive cleanup") +} + +func TestEvictionKeepsReferencedAndFreshArtifacts(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTestManager(p) + m.layerEvictionGrace = time.Hour + + referencedHex := "ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01ef01" + orphanFreshHex := "ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23ab23" + + // A manifest model referencing one layer protects it regardless of age. + model := &imageManifestModel{ + SchemaVersion: manifestModelSchemaVersion, + Digest: "sha256:" + referencedHex, + Config: manifestConfigRef{ + Digest: "sha256:" + strings.Repeat("c", 64), + DiffIDs: []string{"sha256:" + referencedHex}, + }, + Layers: []layerDescriptor{{Digest: "sha256:" + referencedHex, DiffID: "sha256:" + referencedHex}}, + } + require.NoError(t, writeManifestModel(p, referencedHex, model)) + require.NoError(t, os.MkdirAll(p.ImageLayerDir(referencedHex), 0o755)) + require.NoError(t, os.WriteFile(p.ImageLayerArtifact(referencedHex), []byte("kept"), 0o644)) + + // An unreferenced but fresh artifact is protected by the grace period. + require.NoError(t, os.MkdirAll(p.ImageLayerDir(orphanFreshHex), 0o755)) + require.NoError(t, os.WriteFile(p.ImageLayerArtifact(orphanFreshHex), []byte("fresh"), 0o644)) + + m.reconcileLayerStore() + + hexes := layerHexes(t, p) + require.Contains(t, hexes, referencedHex) + require.Contains(t, hexes, orphanFreshHex) +} diff --git a/lib/images/manager.go b/lib/images/manager.go index f4c547886..fb480a458 100644 --- a/lib/images/manager.go +++ b/lib/images/manager.go @@ -76,14 +76,19 @@ type manager struct { queue *queue.Queue createMu sync.Mutex diskUsageMu sync.RWMutex + layerDigestMu sync.Mutex + layerDigestLocks map[string]*sync.Mutex tagGenerations map[string]uint64 - requestedTags map[string]string // newest pull's digest per requested tag + layerRefMu sync.Mutex + inflightLayerRefs map[string]int diskUsageLoaded bool readyImageBytes int64 + layerBytes int64 ociCacheBytes int64 metrics *Metrics inflightPulls map[string]*inflightImagePull // keyed by digest borrowedCredentialsTimeout time.Duration + layerEvictionGrace time.Duration readySubscribers map[string][]chan StatusEvent // keyed by digestHex subscriberMu sync.RWMutex } @@ -104,9 +109,11 @@ func NewManager(p *paths.Paths, maxConcurrentBuilds int, meter metric.Meter) (Ma queue: queue.New(maxConcurrentBuilds), inflightPulls: make(map[string]*inflightImagePull), borrowedCredentialsTimeout: DefaultBorrowedCredentialsTimeout, - readySubscribers: make(map[string][]chan StatusEvent), + layerEvictionGrace: layerEvictionGracePeriod, tagGenerations: make(map[string]uint64), - requestedTags: make(map[string]string), + layerDigestLocks: make(map[string]*sync.Mutex), + inflightLayerRefs: make(map[string]int), + readySubscribers: make(map[string][]chan StatusEvent), } // Initialize metrics if meter is provided @@ -126,11 +133,27 @@ func NewManager(p *paths.Paths, maxConcurrentBuilds int, meter metric.Meter) (Ma if err != nil { fmt.Fprintf(os.Stderr, "Warning: failed to scan legacy images for promotion: %v\n", err) } else { - go promoteLegacyImages(p, legacyRefs) + go promoteLegacyImageRefs(p, legacyRefs) } + m.cleanStaleImageTempDirs() + m.reconcileLayerStore() return m, nil } +// layerDigestLock returns the mutex serializing materialization and eviction +// for one layer digest, so concurrent builds only contend on the digests they +// actually share. +func (m *manager) layerDigestLock(digestHex string) *sync.Mutex { + m.layerDigestMu.Lock() + defer m.layerDigestMu.Unlock() + lock := m.layerDigestLocks[digestHex] + if lock == nil { + lock = &sync.Mutex{} + m.layerDigestLocks[digestHex] = lock + } + return lock +} + func credentialsPresent(credentials *authn.AuthConfig) bool { return credentials != nil && (credentials.Username != "" || credentials.Password != "" || credentials.Auth != "" || credentials.IdentityToken != "" || credentials.RegistryToken != "") } @@ -158,7 +181,7 @@ func (m *manager) ListImages(ctx context.Context) ([]Image, error) { images := make([]Image, 0, len(metas)) for _, meta := range metas { - images = append(images, *meta.toImageFor(meta.Name)) + images = append(images, *meta.toImage()) } return images, nil @@ -256,17 +279,27 @@ func (m *manager) ImportLocalImage(ctx context.Context, repo, reference, digest // Check if we already have this digest (deduplication) if meta, err := readMetadata(m.paths, ref.Repository(), ref.DigestHex()); err == nil { - // Don't cache failed builds - allow retry by falling through to - // re-queue the build. + // Don't cache failed builds - allow retry if meta.Status == StatusFailed { - if err := removeDigestIfUnreferenced(m.paths, ref.Repository(), ref.DigestHex(), false); err != nil { - return nil, fmt.Errorf("remove failed image: %w", err) + if err := m.discardFailedImage(ref.Repository(), ref.DigestHex()); err != nil { + return nil, err } + // Fall through to re-queue the build } else { - if err := m.claimTagForStatus(meta, ref); err != nil { - return nil, fmt.Errorf("create image tag: %w", err) + if ref.Tag() != "" { + var tagErr error + if meta.Status == StatusReady { + m.nextTagGeneration(ref.Repository(), ref.Tag()) + tagErr = createTagSymlink(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()) + } else { + tagErr = ensurePendingTag(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()) + } + if tagErr != nil { + return nil, fmt.Errorf("create image tag: %w", tagErr) + } } - img := meta.toImageFor(ref.String()) + img := meta.toImage() + img.Name = ref.String() if meta.Status == StatusPending { img.QueuePosition = m.queue.GetPosition(meta.Digest) } @@ -284,30 +317,81 @@ func (m *manager) reuseExistingImage(ref *ResolvedRef, credentials *authn.AuthCo return nil, false, nil } if meta.Status == StatusFailed { - if err := removeDigestIfUnreferenced(m.paths, ref.Repository(), ref.DigestHex(), false); err != nil { - return nil, true, fmt.Errorf("remove failed image: %w", err) + if err := m.discardFailedImage(ref.Repository(), ref.DigestHex()); err != nil { + return nil, true, err } return nil, false, nil } if ref.Tag() != "" { - if meta.Status != StatusReady { - // A pending pull with different credentials does not get to point - // the tag at its digest. - if !m.inflightCredentialsMatch(ref.Digest(), credentials) { - return nil, true, fmt.Errorf("%w: retry after the current pull completes", ErrCredentialConflict) - } + if meta.Status == StatusReady { + m.nextTagGeneration(ref.Repository(), ref.Tag()) + err = createTagSymlink(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()) + } else { + err = ensurePendingTag(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()) } - if err := m.claimTagForStatus(meta, ref); err != nil { + if err != nil { return nil, true, fmt.Errorf("create image tag: %w", err) } } - img := meta.toImageFor(ref.String()) + img := meta.toImage() + img.Name = ref.String() + if meta.Status == StatusReady { + return img, true, nil + } + if !m.inflightCredentialsMatch(ref.Digest(), credentials) { + return nil, true, fmt.Errorf("%w: retry after the current pull completes", ErrCredentialConflict) + } if meta.Status == StatusPending { img.QueuePosition = m.queue.GetPosition(meta.Digest) } return img, true, nil } +// discardFailedImage clears a failed build's tags and digest layout so a +// retry can start clean. +func (m *manager) discardFailedImage(repository, digestHex string) error { + _ = deleteTagsForDigest(m.paths, repository, digestHex) + if err := removeDigestIfUnreferenced(m.paths, repository, digestHex, false); err != nil { + return fmt.Errorf("remove failed image: %w", err) + } + return nil +} + +func tagGenerationKey(repository, tag string) string { + return repository + ":" + tag +} + +func (m *manager) nextTagGeneration(repository, tag string) uint64 { + key := tagGenerationKey(repository, tag) + m.tagGenerations[key]++ + return m.tagGenerations[key] +} + +func (m *manager) restoreTagGenerations(metas []*imageMetadata) { + m.createMu.Lock() + defer m.createMu.Unlock() + for _, meta := range metas { + if meta.RequestedTag == "" { + continue + } + ref, err := ParseNormalizedRef(meta.Name) + if err != nil { + continue + } + key := tagGenerationKey(ref.Repository(), meta.RequestedTag) + if meta.TagGeneration > m.tagGenerations[key] { + m.tagGenerations[key] = meta.TagGeneration + } + } +} + +func (m *manager) claimReadyTag(ref *ResolvedRef) error { + m.createMu.Lock() + defer m.createMu.Unlock() + m.nextTagGeneration(ref.Repository(), ref.Tag()) + return createTagSymlink(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()) +} + func (m *manager) inflightCredentialsMatch(digest string, credentials *authn.AuthConfig) bool { inflight := m.inflightPulls[digest] var existingFingerprint [32]byte @@ -318,9 +402,6 @@ func (m *manager) inflightCredentialsMatch(digest string, credentials *authn.Aut } func (m *manager) registerInflightPull(digest string, credentials *authn.AuthConfig) *inflightImagePull { - if m.inflightPulls == nil { - m.inflightPulls = make(map[string]*inflightImagePull) - } if previous := m.inflightPulls[digest]; previous != nil && previous.timer != nil { previous.timer.Stop() } @@ -396,10 +477,10 @@ func (m *manager) createAndQueueImage(ref *ResolvedRef, req CreateImageRequest, return nil, fmt.Errorf("resolve existing image tag: %w", err) } } + tagGeneration := uint64(0) if ref.Tag() != "" { tagGeneration = m.nextTagGeneration(ref.Repository(), ref.Tag()) - m.trackRequestedTag(ref.Repository(), ref.Tag(), ref.DigestHex()) } meta := &imageMetadata{ Name: ref.String(), @@ -418,14 +499,10 @@ func (m *manager) createAndQueueImage(ref *ResolvedRef, req CreateImageRequest, // Write initial metadata if err := writeMetadata(m.paths, ref.Repository(), ref.DigestHex(), meta); err != nil { - if ref.Tag() != "" { - m.revertTagGeneration(ref.Repository(), ref.Tag()) - } return nil, fmt.Errorf("write initial metadata: %w", err) } if ref.Tag() != "" && previousTagDigest == "" { if err := createTagSymlink(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()); err != nil { - m.revertTagGeneration(ref.Repository(), ref.Tag()) return nil, fmt.Errorf("create pending image tag: %w", err) } } @@ -452,13 +529,38 @@ func (m *manager) createAndQueueImage(ref *ResolvedRef, req CreateImageRequest, m.buildImage(ctx, ref, credentials, buildID) }, m.releaseInflightPull(ref.Digest(), inflight)) - img := meta.toImageFor(ref.String()) + img := meta.toImage() if queuePos > 0 { img.QueuePosition = &queuePos } return img, nil } +func (m *manager) retainLayerRefs(model *imageManifestModel) func() { + if model == nil { + return func() {} + } + m.layerRefMu.Lock() + digests := make([]string, 0, len(model.Layers)) + for _, layer := range model.Layers { + digestHex := strings.TrimPrefix(layer.Digest, "sha256:") + m.inflightLayerRefs[digestHex]++ + digests = append(digests, digestHex) + } + m.layerRefMu.Unlock() + + return func() { + m.layerRefMu.Lock() + defer m.layerRefMu.Unlock() + for _, digestHex := range digests { + m.inflightLayerRefs[digestHex]-- + if m.inflightLayerRefs[digestHex] == 0 { + delete(m.inflightLayerRefs, digestHex) + } + } + } +} + func (m *manager) buildImage(ctx context.Context, ref *ResolvedRef, credentials *authn.AuthConfig, buildID string) { buildStart := time.Now() buildStatus := "failed" @@ -497,15 +599,22 @@ func (m *manager) buildImage(ctx context.Context, ref *ResolvedRef, credentials } m.recordPullMetrics(ctx, "success") + // The pulled layers own fully unpacked staging trees in /tmp; release them + // when the build ends no matter which path it takes. + defer result.cleanup() + + releaseLayerRefs := m.retainLayerRefs(result.Manifest) + defer func() { + releaseLayerRefs() + m.reconcileLayerStore() + }() + // Check if this digest already exists and is ready (deduplication) if meta, err := readMetadata(m.paths, ref.Repository(), ref.DigestHex()); err == nil { if meta.Status == StatusReady { // Another build completed first; last-pull-wins repoints the tag. if ref.Tag() != "" { - m.createMu.Lock() - err := m.claimReadyTag(ref.Repository(), ref.Tag(), ref.DigestHex()) - m.createMu.Unlock() - if err != nil { + if err := m.claimReadyTag(ref); err != nil { slog.Warn("failed to claim ready image tag", "repository", ref.Repository(), "tag", ref.Tag(), "error", err) } } @@ -514,6 +623,9 @@ func (m *manager) buildImage(ctx context.Context, ref *ResolvedRef, credentials } } + // Materialization is best effort; the composed rootfs does not depend on it. + m.materializeLayerArtifacts(ctx, ref.Digest(), result) + m.updateStatusByDigest(ref, StatusConverting, nil, buildID) diskPath := resolveImageLayout(m.paths, ref.Repository(), ref.DigestHex()).disk @@ -544,23 +656,42 @@ func (m *manager) buildImage(ctx context.Context, ref *ResolvedRef, credentials buildStatus = "success" } +func (m *manager) materializeLayerArtifacts(ctx context.Context, digest string, result *pullResult) { + if result == nil || result.Manifest == nil { + return + } + start := time.Now() + var firstErr error + for _, desc := range result.Manifest.Layers { + if _, err := m.materializeLayerArtifact(desc, result.LayerTrees[desc.Digest]); err != nil { + slog.WarnContext(ctx, "failed to materialize layer artifact", "digest", desc.Digest, "error", err) + if firstErr == nil { + firstErr = err + } + } + } + cacheStatus := "miss" + if result.CacheHit { + cacheStatus = "hit" + } + m.recordImageBuildPhase(ctx, digest, "layer_materialization", time.Since(start), phaseStatus(firstErr), cacheStatus) +} + func (m *manager) finalizeImage(ref *ResolvedRef, result *pullResult, diskSize int64, buildID, diskTempPath string) error { + if diskTempPath != "" { + defer os.Remove(diskTempPath) + } + m.createMu.Lock() defer m.createMu.Unlock() - layout := resolveImageLayout(m.paths, ref.Repository(), ref.DigestHex()) - // Read current metadata to preserve request info and reject stale builds. - meta, err := readMetadataAt(layout) + meta, err := readMetadata(m.paths, ref.Repository(), ref.DigestHex()) if err != nil || meta.BuildID != buildID { return errStaleBuild } - if err := installAtomically(layout.disk, func(path string) error { - return os.Rename(diskTempPath, path) - }); err != nil { - return fmt.Errorf("install image disk: %w", err) - } + finalDiskPath := resolveImageLayout(m.paths, ref.Repository(), ref.DigestHex()).disk // The pulled image config is the source of truth for the platform. var requestedPlatform string @@ -575,12 +706,13 @@ func (m *manager) finalizeImage(ref *ResolvedRef, result *pullResult, diskSize i // Persist the manifest content model beside the shared content so later // stages can recompose the image from per-layer artifacts and GC can tell // which OCI blobs are still referenced. - if result.Manifest != nil { - model := *result.Manifest - model.Platform = actualPlatform.String() - if err := writeManifestModel(m.paths, ref.DigestHex(), &model); err != nil { - return fmt.Errorf("write manifest model: %w", err) - } + if result.Manifest == nil { + return fmt.Errorf("manifest model missing for new image") + } + model := *result.Manifest + model.Platform = actualPlatform.String() + if err := writeManifestModel(m.paths, ref.DigestHex(), &model); err != nil { + return fmt.Errorf("write manifest model: %w", err) } meta.Status = StatusReady @@ -593,18 +725,61 @@ func (m *manager) finalizeImage(ref *ResolvedRef, result *pullResult, diskSize i meta.Labels = result.Metadata.Labels meta.WorkingDir = result.Metadata.WorkingDir - if err := writeMetadataFile(layout.metadata, meta); err != nil { + if err := installAtomically(finalDiskPath, func(path string) error { + return os.Rename(diskTempPath, path) + }); err != nil { + _ = os.Remove(m.paths.ImageContentManifestModel(ref.DigestHex())) + return fmt.Errorf("install image disk: %w", err) + } + if err := writeMetadata(m.paths, ref.Repository(), ref.DigestHex(), meta); err != nil { + _ = os.Remove(finalDiskPath) + _ = os.Remove(m.paths.ImageContentManifestModel(ref.DigestHex())) return fmt.Errorf("write final metadata: %w", err) } m.notifyReady(ref.DigestHex(), StatusReady, nil) - if !m.claimRequestedTag(ref, meta) { + if !m.publishReadyTag(ref, meta) { m.cleanupUnclaimedImage(ref) } m.refreshDiskUsageTotals() return nil } +func (m *manager) cleanupUnclaimedImage(ref *ResolvedRef) { + if err := removeDigestIfUnreferenced(m.paths, ref.Repository(), ref.DigestHex(), true); err != nil { + slog.Warn("failed to collect stale image", "repository", ref.Repository(), "digest", ref.DigestHex(), "error", err) + } +} + +func (m *manager) publishReadyTag(ref *ResolvedRef, meta *imageMetadata) bool { + requestedTag := meta.RequestedTag + allowMissing := requestedTag == "" + if requestedTag == "" { + requestedTag = ref.Tag() + } + if requestedTag == "" { + return false + } + + currentDigest, err := resolveTag(m.paths, ref.Repository(), requestedTag) + generationMatches := m.tagGenerations[tagGenerationKey(ref.Repository(), requestedTag)] == meta.TagGeneration + if !generationMatches { + return false + } + if err != nil { + if !allowMissing || !errors.Is(err, ErrNotFound) { + return false + } + } else if currentDigest != ref.DigestHex() && currentDigest != meta.PreviousTagDigest { + return false + } + if err := createTagSymlink(m.paths, ref.Repository(), requestedTag, ref.DigestHex()); err != nil { + fmt.Fprintf(os.Stderr, "Warning: failed to create tag symlink: %v\n", err) + return false + } + return true +} + func phaseStatus(err error) string { if err != nil { return "failed" @@ -649,8 +824,7 @@ func (m *manager) updateStatusByDigest(ref *ResolvedRef, status string, err erro m.createMu.Lock() defer m.createMu.Unlock() - layout := resolveImageLayout(m.paths, ref.Repository(), ref.DigestHex()) - meta, readErr := readMetadataAt(layout) + meta, readErr := readMetadata(m.paths, ref.Repository(), ref.DigestHex()) if readErr != nil || meta.BuildID != buildID { return } @@ -661,18 +835,22 @@ func (m *manager) updateStatusByDigest(ref *ResolvedRef, status string, err erro meta.Error = &errorMsg } - writeMetadataFile(layout.metadata, meta) + if writeErr := writeMetadata(m.paths, ref.Repository(), ref.DigestHex(), meta); writeErr != nil { + if status == StatusReady || status == StatusFailed { + m.notifyReady(ref.DigestHex(), status, errors.Join(err, writeErr)) + } + return + } + + if status == StatusFailed { + m.refreshDiskUsageTotals() + } // Notify while holding createMu so a delete/recreate cannot race the // metadata write and receive a terminal event for the old build. if status == StatusReady || status == StatusFailed { m.notifyReady(ref.DigestHex(), status, err) } - // A failed pull releases its tag claim so an older in-flight pull of the - // same tag can still repoint it. - if status == StatusFailed && meta.RequestedTag != "" { - m.releaseTagGeneration(ref.Repository(), meta.RequestedTag, meta.TagGeneration) - } } func (m *manager) RecoverInterruptedBuilds() { @@ -680,12 +858,12 @@ func (m *manager) RecoverInterruptedBuilds() { if err != nil { return // Best effort } + m.restoreTagGenerations(metas) // Sort by created_at to maintain FIFO order sort.Slice(metas, func(i, j int) bool { return metas[i].CreatedAt.Before(metas[j].CreatedAt) }) - m.restoreTagState(metas) seenDigests := make(map[string]struct{}) for _, meta := range metas { @@ -722,12 +900,29 @@ func (m *manager) GetImage(ctx context.Context, name string) (*Image, error) { return nil, fmt.Errorf("%w: %s", ErrInvalidName, err.Error()) } - _, meta, err := resolveRefMetadata(m.paths, ref) + repository := ref.Repository() + + var digestHex string + if ref.IsDigest() { + // Direct digest lookup + digestHex = ref.DigestHex() + } else { + // Tag lookup - resolve symlink + tag := ref.Tag() + d, err := resolveTag(m.paths, repository, tag) + if err != nil { + return nil, err + } + digestHex = d + } + + meta, err := readMetadata(m.paths, repository, digestHex) if err != nil { return nil, err } - img := meta.toImageFor(ref.String()) + img := meta.toImage() + img.Name = ref.String() if meta.Status == StatusPending { img.QueuePosition = m.queue.GetPosition(meta.Digest) @@ -736,6 +931,91 @@ func (m *manager) GetImage(ctx context.Context, name string) (*Image, error) { return img, nil } +// TagImage creates or updates a local tag pointing at an existing ready +// image. The source may be a tag or a digest; the target must carry a tag. +// No pull or conversion happens: the target becomes another reference to the +// same shared content. Source and target may live in different repositories; +// a cross-repository target first promotes the digest into shared content so +// both references resolve to one rootfs copy. +func (m *manager) TagImage(ctx context.Context, source, target string) (*Image, error) { + sourceRef, err := ParseNormalizedRef(source) + if err != nil { + return nil, fmt.Errorf("%w: invalid source reference: %s", ErrInvalidName, err) + } + targetRef, err := ParseNormalizedRef(target) + if err != nil { + return nil, fmt.Errorf("%w: invalid target reference: %s", ErrInvalidName, err) + } + if targetRef.IsDigest() { + return nil, fmt.Errorf("%w: target must include a tag", ErrInvalidName) + } + + m.createMu.Lock() + defer m.createMu.Unlock() + m.nextTagGeneration(targetRef.Repository(), targetRef.Tag()) + + digestHex, err := m.resolveTagSource(sourceRef) + if err != nil { + return nil, err + } + meta, err := m.readyImageMetadata(sourceRef.Repository(), digestHex) + if err != nil { + return nil, err + } + + previousDigest := "" + if targetDigest, tagErr := resolveTag(m.paths, targetRef.Repository(), targetRef.Tag()); tagErr == nil { + previousDigest = targetDigest + } else if !errors.Is(tagErr, ErrNotFound) { + return nil, tagErr + } + + if err := m.installImageTag(sourceRef, targetRef, digestHex, meta); err != nil { + return nil, err + } + if previousDigest != "" && previousDigest != digestHex { + if err := removeDigestIfUnreferenced(m.paths, targetRef.Repository(), previousDigest, true); err != nil { + return nil, fmt.Errorf("remove replaced image: %w", err) + } + m.reconcileLayerStore() + } + + img := meta.toImage() + img.Name = targetRef.String() + return img, nil +} + +func (m *manager) installImageTag(sourceRef, targetRef *NormalizedRef, digestHex string, meta *imageMetadata) error { + if sourceRef.Repository() != targetRef.Repository() { + if err := promoteImageToContent(m.paths, sourceRef.Repository(), digestHex, meta, &tagTarget{repository: targetRef.Repository(), tag: targetRef.Tag()}); err != nil { + return fmt.Errorf("create image alias: %w", err) + } + return nil + } + if err := createTagSymlink(m.paths, targetRef.Repository(), targetRef.Tag(), digestHex); err != nil { + return fmt.Errorf("create image tag: %w", err) + } + return nil +} + +func (m *manager) resolveTagSource(ref *NormalizedRef) (string, error) { + if ref.IsDigest() { + return ref.DigestHex(), nil + } + return resolveTag(m.paths, ref.Repository(), ref.Tag()) +} + +func (m *manager) readyImageMetadata(repository, digestHex string) (*imageMetadata, error) { + meta, err := readMetadata(m.paths, repository, digestHex) + if err != nil { + return nil, err + } + if meta.Status != StatusReady { + return nil, fmt.Errorf("%w: %s", ErrImageNotReady, meta.Status) + } + return meta, nil +} + func (m *manager) DeleteImage(ctx context.Context, name string) error { // Parse and normalize the reference ref, err := ParseNormalizedRef(name) @@ -758,15 +1038,15 @@ func (m *manager) DeleteImage(ctx context.Context, name string) error { if err := deleteTagsForDigest(m.paths, repository, digestHex); err != nil { return err } - m.pruneTagGenerations() if err := removeDigestIfUnreferenced(m.paths, repository, digestHex, false); err != nil { return err } - m.refreshDiskUsageTotals() + m.reconcileLayerStore() return nil } tag := ref.Tag() + m.nextTagGeneration(repository, tag) // Resolve the tag to get the digest before deleting digestHex, err := resolveTag(m.paths, repository, tag) @@ -778,7 +1058,6 @@ func (m *manager) DeleteImage(ctx context.Context, name string) error { if err := deleteTag(m.paths, repository, tag); err != nil { return err } - m.pruneTagGenerations() // Check if the digest is now orphaned (no other tags reference it) count, err := countTagsForDigest(m.paths, repository, digestHex) @@ -790,7 +1069,7 @@ func (m *manager) DeleteImage(ctx context.Context, name string) error { if err := removeDigestIfUnreferenced(m.paths, repository, digestHex, true); err != nil { return fmt.Errorf("delete orphaned digest %s: %w", digestHex, err) } - m.refreshDiskUsageTotals() + m.reconcileLayerStore() } return nil @@ -798,24 +1077,50 @@ func (m *manager) DeleteImage(ctx context.Context, name string) error { // TotalImageBytes returns the total size of all ready images on disk. func (m *manager) TotalImageBytes(ctx context.Context) (int64, error) { - readyImageBytes, _, err := m.getDiskUsageTotals() + totals, err := m.getDiskUsageTotals() if err != nil { return 0, err } - return readyImageBytes, nil + return totals.readyImageBytes + totals.layerBytes, nil } // TotalOCICacheBytes returns the total size of the OCI layer cache. func (m *manager) TotalOCICacheBytes(ctx context.Context) (int64, error) { - _, ociCacheBytes, err := m.getDiskUsageTotals() + totals, err := m.getDiskUsageTotals() if err != nil { return 0, err } - return ociCacheBytes, nil + return totals.ociCacheBytes, nil +} + +func (m *manager) findRequestedTagImage(ref *NormalizedRef) *Image { + metas, err := listAllMetadata(m.paths) + if err != nil { + return nil + } + var newest *imageMetadata + for _, meta := range metas { + if meta.RequestedTag != ref.Tag() || !strings.HasPrefix(meta.Name, ref.Repository()+":") { + continue + } + if newest == nil || newest.CreatedAt.Before(meta.CreatedAt) { + newest = meta + } + } + if newest == nil { + return nil + } + image := newest.toImage() + image.Name = ref.String() + return image } // WaitForReady blocks until the image reaches a terminal state (ready or failed) // or the context is cancelled. +// +// The image may not exist yet when this is called (e.g., the registry's +// triggerConversion goroutine hasn't called ImportLocalImage yet), so we +// poll briefly for the image to appear before subscribing for notifications. func (m *manager) WaitForReady(ctx context.Context, name string) error { ref, err := ParseNormalizedRef(name) if err != nil { @@ -864,7 +1169,7 @@ func (m *manager) waitForImage(ctx context.Context, name string, ref *Normalized for { var img *Image if !ref.IsDigest() { - img = m.requestedTagImage(ref) + img = m.findRequestedTagImage(ref) } if img == nil { img, lastErr = m.GetImage(ctx, name) diff --git a/lib/images/manager_test.go b/lib/images/manager_test.go index 13a6835e2..36015dd5b 100644 --- a/lib/images/manager_test.go +++ b/lib/images/manager_test.go @@ -7,6 +7,7 @@ import ( "os" "path/filepath" "strings" + "sync" "testing" "time" @@ -17,6 +18,19 @@ import ( "github.com/stretchr/testify/require" ) +// newTestManager returns a manager with the maps NewManager initializes, so +// tests can construct one directly without nil-map guards in the manager. +func newTestManager(p *paths.Paths) *manager { + return &manager{ + paths: p, + tagGenerations: make(map[string]uint64), + layerDigestLocks: make(map[string]*sync.Mutex), + inflightLayerRefs: make(map[string]int), + inflightPulls: make(map[string]*inflightImagePull), + readySubscribers: make(map[string][]chan StatusEvent), + } +} + func TestConversionFailedErr(t *testing.T) { t.Run("without detail", func(t *testing.T) { assert.EqualError(t, conversionFailedErr(nil, nil), "image conversion failed") @@ -734,9 +748,9 @@ func TestDeleteAndRecreateDuringBuildTail(t *testing.T) { require.NoError(t, err) staleRef := NewResolvedRef(normalized, digestStr) m.updateStatusByDigest(staleRef, StatusFailed, errors.New("stale build"), firstMeta.BuildID) - staleBundle, err := m.ociClient.extractOCIImageBundle(digestHex) + staleResult, _, _, err := m.ociClient.extractOCIImageDetails(digestHex) require.NoError(t, err) - require.ErrorIs(t, m.finalizeImage(staleRef, &pullResult{Metadata: staleBundle.Meta}, 1, firstMeta.BuildID, ""), errStaleBuild) + require.ErrorIs(t, m.finalizeImage(staleRef, &pullResult{Metadata: staleResult}, 1, firstMeta.BuildID, ""), errStaleBuild) currentMeta, err = readMetadata(p, repo, digestHex) require.NoError(t, err) require.Equal(t, StatusPending, currentMeta.Status) diff --git a/lib/images/manifest_model.go b/lib/images/manifest_model.go index b3e7dc25b..0a8331fa9 100644 --- a/lib/images/manifest_model.go +++ b/lib/images/manifest_model.go @@ -66,6 +66,7 @@ func (m *imageManifestModel) blobReferences() []string { return refs } +// writeManifestModel persists the manifest model for a digest atomically. func validateManifestModel(digestHex string, model *imageManifestModel) error { if model == nil { return fmt.Errorf("manifest model is nil") @@ -157,10 +158,28 @@ func readManifestModel(p *paths.Paths, digestHex string) (*imageManifestModel, e // writeJSONAtomic writes data to path via a temp file in the same directory // followed by a rename, so readers never observe a partial document. func writeJSONAtomic(path string, data []byte) error { - if err := installAtomically(path, func(tempPath string) error { - return os.WriteFile(tempPath, data, 0o644) - }); err != nil { - return fmt.Errorf("write %s: %w", filepath.Base(path), err) + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + return fmt.Errorf("create directory: %w", err) + } + tempFile, err := os.CreateTemp(filepath.Dir(path), "."+filepath.Base(path)+".tmp-*") + if err != nil { + return fmt.Errorf("create temp file: %w", err) + } + tempPath := tempFile.Name() + defer os.Remove(tempPath) + if err := tempFile.Chmod(0644); err != nil { + _ = tempFile.Close() + return fmt.Errorf("chmod temp file: %w", err) + } + if _, err := tempFile.Write(data); err != nil { + _ = tempFile.Close() + return fmt.Errorf("write temp file: %w", err) + } + if err := tempFile.Close(); err != nil { + return fmt.Errorf("close temp file: %w", err) + } + if err := os.Rename(tempPath, path); err != nil { + return fmt.Errorf("rename into place: %w", err) } return nil } diff --git a/lib/images/manifest_model_test.go b/lib/images/manifest_model_test.go index 1137e3912..8e9e63777 100644 --- a/lib/images/manifest_model_test.go +++ b/lib/images/manifest_model_test.go @@ -74,9 +74,8 @@ func TestExtractManifestModel(t *testing.T) { client, layoutTag := writeSyntheticLayout(t, img) - bundle, err := client.extractOCIImageBundle(layoutTag) + model, err := client.extractManifestModel(layoutTag) require.NoError(t, err) - model := bundle.Model manifest, err := img.Manifest() require.NoError(t, err) @@ -117,9 +116,8 @@ func TestExtractManifestModelPlatform(t *testing.T) { client, layoutTag := writeSyntheticLayout(t, img) - bundle, err := client.extractOCIImageBundle(layoutTag) + model, err := client.extractManifestModel(layoutTag) require.NoError(t, err) - model := bundle.Model require.Equal(t, "linux/amd64", model.Platform) } diff --git a/lib/images/metrics.go b/lib/images/metrics.go index d860885b2..75164a6cd 100644 --- a/lib/images/metrics.go +++ b/lib/images/metrics.go @@ -11,11 +11,12 @@ import ( // Metrics holds the metrics instruments for image operations. type Metrics struct { - buildDuration metric.Float64Histogram - buildPhaseDuration metric.Float64Histogram - ociLayerCount metric.Int64Histogram - ociCompressedBytes metric.Int64Histogram - pullsTotal metric.Int64Counter + buildDuration metric.Float64Histogram + buildPhaseDuration metric.Float64Histogram + ociLayerCount metric.Int64Histogram + ociCompressedBytes metric.Int64Histogram + pullsTotal metric.Int64Counter + layerArtifactsEvicted metric.Int64Counter } // newMetrics creates and registers all image metrics. @@ -78,6 +79,14 @@ func newMetrics(meter metric.Meter, m *manager) (*Metrics, error) { return nil, err } + layerArtifactsEvicted, err := meter.Int64Counter( + "hypeman_images_layer_artifacts_evicted_total", + metric.WithDescription("Total number of shared layer artifacts evicted after their last reference was removed"), + ) + if err != nil { + return nil, err + } + // Register observable gauges for queue length and total images buildQueueLength, err := meter.Int64ObservableGauge( "hypeman_images_build_queue_length", @@ -123,11 +132,12 @@ func newMetrics(meter metric.Meter, m *manager) (*Metrics, error) { } return &Metrics{ - buildDuration: buildDuration, - buildPhaseDuration: buildPhaseDuration, - ociLayerCount: ociLayerCount, - ociCompressedBytes: ociCompressedBytes, - pullsTotal: pullsTotal, + buildDuration: buildDuration, + buildPhaseDuration: buildPhaseDuration, + ociLayerCount: ociLayerCount, + ociCompressedBytes: ociCompressedBytes, + pullsTotal: pullsTotal, + layerArtifactsEvicted: layerArtifactsEvicted, }, nil } diff --git a/lib/images/metrics_test.go b/lib/images/metrics_test.go index 93007b300..1f58ab48c 100644 --- a/lib/images/metrics_test.go +++ b/lib/images/metrics_test.go @@ -16,10 +16,8 @@ import ( func TestImageBuildPhaseMetrics(t *testing.T) { reader := otelmetric.NewManualReader() provider := otelmetric.NewMeterProvider(otelmetric.WithReader(reader)) - m := &manager{ - paths: paths.New(t.TempDir()), - queue: queue.New(1), - } + m := newTestManager(paths.New(t.TempDir())) + m.queue = queue.New(1) metrics, err := newMetrics(provider.Meter("test"), m) require.NoError(t, err) diff --git a/lib/images/oci.go b/lib/images/oci.go index 5b66eaa53..9517c24b0 100644 --- a/lib/images/oci.go +++ b/lib/images/oci.go @@ -195,6 +195,7 @@ func (c *ociClient) inspectDigestPlatformAuth(ctx context.Context, imageRef stri type pullResult struct { Metadata *containerMetadata Manifest *imageManifestModel + LayerTrees map[string]layerTree Digest string // sha256:abc123... CacheHit bool LayerCount int @@ -208,14 +209,6 @@ type imageBuildPhaseMeasurement struct { Status string } -// ociImageBundle is the extracted content of one OCI image in the cache. -type ociImageBundle struct { - Meta *containerMetadata - Model *imageManifestModel - LayerCount int - CompressedBytes int64 -} - func (r *pullResult) measure(phase string, operation func() error) error { start := time.Now() err := operation() @@ -231,6 +224,11 @@ func (r *pullResult) measure(phase string, operation func() error) error { return err } +// cleanup removes the transient layer staging trees the result owns. +func (r *pullResult) cleanup() { + cleanupLayerTrees(r.LayerTrees) +} + func (c *ociClient) pullAndExport(ctx context.Context, imageRef, digest, exportDir string) (*pullResult, error) { return c.pullAndExportWithAuth(ctx, imageRef, digest, exportDir, nil) } @@ -269,24 +267,36 @@ func (c *ociClient) pullAndExportWithPlatformAuth(ctx context.Context, imageRef, // If cached, we skip the pull entirely // Extract metadata (from cache or freshly pulled) - var bundle *ociImageBundle + var meta *containerMetadata + var layerCount int + var compressedBytes int64 + var model *imageManifestModel err := result.measure("metadata_extract", func() error { var err error - bundle, err = c.extractOCIImageBundle(layoutTag) + meta, model, layerCount, compressedBytes, err = c.extractOCIImageBundle(layoutTag) return err }) if err != nil { return result, fmt.Errorf("extract metadata: %w", err) } - result.Metadata = bundle.Meta - result.Manifest = bundle.Model - result.LayerCount = bundle.LayerCount - result.CompressedBytes = bundle.CompressedBytes + result.Metadata = meta + result.Manifest = model + result.LayerCount = layerCount + result.CompressedBytes = compressedBytes - // Compose the rootfs from the shared layer blobs in manifest order. - // composeRootfs validates the model and rejects zero-layer manifests. + // Compose the rootfs from the shared layer blobs. The manifest model + // carries the ordered layer descriptors; fall back to the umoci-based + // unpack only when no model could be extracted. if err := result.measure("layer_unpack", func() error { - return c.composeRootfs(exportDir, layoutTag, bundle.Model) + if model != nil && len(model.Layers) > 0 { + if err := validateModelPairing(layoutTag, model); err != nil { + return err + } + var err error + result.LayerTrees, err = c.composeRootfsWithLayerTrees(exportDir, model.Layers) + return err + } + return c.unpackLayers(ctx, layoutTag, exportDir) }); err != nil { return result, fmt.Errorf("unpack layers: %w", err) } @@ -404,29 +414,43 @@ func imageByAnnotation(path layout.Path, layoutTag string) (gcr.Image, error) { return nil, fmt.Errorf("no image found with tag %s", layoutTag) } -// extractOCIImageBundle reads metadata, the manifest content model, and layer -// stats from the cached OCI layout. Uses go-containerregistry which handles -// both Docker v2 and OCI v1 manifests. -func (c *ociClient) extractOCIImageBundle(layoutTag string) (*ociImageBundle, error) { +// extractOCIMetadata reads metadata from OCI layout config.json +// Uses go-containerregistry which handles both Docker v2 and OCI v1 manifests. +func (c *ociClient) extractOCIMetadata(layoutTag string) (*containerMetadata, error) { + meta, _, _, err := c.extractOCIImageDetails(layoutTag) + return meta, err +} + +func (c *ociClient) extractOCIImageDetails(layoutTag string) (*containerMetadata, int, int64, error) { + meta, _, layerCount, compressedBytes, err := c.extractOCIImageBundle(layoutTag) + return meta, layerCount, compressedBytes, err +} + +func (c *ociClient) extractManifestModel(layoutTag string) (*imageManifestModel, error) { + _, model, _, _, err := c.extractOCIImageBundle(layoutTag) + return model, err +} + +func (c *ociClient) extractOCIImageBundle(layoutTag string) (*containerMetadata, *imageManifestModel, int, int64, error) { path, err := layout.FromPath(c.cacheDir) if err != nil { - return nil, fmt.Errorf("open oci layout: %w", err) + return nil, nil, 0, 0, fmt.Errorf("open oci layout: %w", err) } img, err := imageByAnnotation(path, layoutTag) if err != nil { - return nil, fmt.Errorf("find image by tag %s: %w", layoutTag, err) + return nil, nil, 0, 0, fmt.Errorf("find image by tag %s: %w", layoutTag, err) } configFile, err := img.ConfigFile() if err != nil { - return nil, fmt.Errorf("get config file: %w", err) + return nil, nil, 0, 0, fmt.Errorf("get config file: %w", err) } manifest, err := img.Manifest() if err != nil { - return nil, fmt.Errorf("get manifest: %w", err) + return nil, nil, 0, 0, fmt.Errorf("get manifest: %w", err) } configDigest, err := img.ConfigName() if err != nil { - return nil, fmt.Errorf("get config digest: %w", err) + return nil, nil, 0, 0, fmt.Errorf("get config digest: %w", err) } meta := &containerMetadata{ @@ -439,15 +463,9 @@ func (c *ociClient) extractOCIImageBundle(layoutTag string) (*ociImageBundle, er Labels: make(map[string]string), WorkingDir: configFile.Config.WorkingDir, } - // Parse environment variables for _, env := range configFile.Config.Env { - for i := 0; i < len(env); i++ { - if env[i] == '=' { - key := env[:i] - val := env[i+1:] - meta.Env[key] = val - break - } + if key, value, ok := strings.Cut(env, "="); ok { + meta.Env[key] = value } } for key, value := range configFile.Config.Labels { @@ -487,12 +505,7 @@ func (c *ociClient) extractOCIImageBundle(layoutTag string) (*ociImageBundle, er } model.Layers = append(model.Layers, layer) } - return &ociImageBundle{ - Meta: meta, - Model: model, - LayerCount: len(manifest.Layers), - CompressedBytes: compressedBytes, - }, nil + return meta, model, len(manifest.Layers), compressedBytes, nil } // unpackLayers unpacks all OCI layers to a target directory using umoci diff --git a/lib/images/oci_public.go b/lib/images/oci_public.go index 7d336745b..554e1967f 100644 --- a/lib/images/oci_public.go +++ b/lib/images/oci_public.go @@ -47,7 +47,10 @@ func (c *OCIClient) InspectManifestForLinux(ctx context.Context, imageRef string // PullAndUnpack pulls an OCI image and unpacks it to a directory (public for system manager). // Always targets Linux platform since hypeman VMs are Linux guests. func (c *OCIClient) PullAndUnpack(ctx context.Context, imageRef, digest, exportDir string) error { - _, err := c.client.pullAndExport(ctx, imageRef, digest, exportDir) + result, err := c.client.pullAndExport(ctx, imageRef, digest, exportDir) + if result != nil { + defer result.cleanup() + } if err != nil { return fmt.Errorf("pull and unpack: %w", err) } diff --git a/lib/images/oci_test.go b/lib/images/oci_test.go index 2cb765cc5..e076156ff 100644 --- a/lib/images/oci_test.go +++ b/lib/images/oci_test.go @@ -89,11 +89,11 @@ func TestExtractMetadataSucceedsOnBuildKitCache(t *testing.T) { // This succeeds because go-containerregistry doesn't validate config mediatype // The failure only happens in unpackLayers when umoci validates the config - bundle, err := client.extractOCIImageBundle("test-cache") - require.NoError(t, err, "extractOCIImageBundle succeeds - go-containerregistry is lenient") + meta, err := client.extractOCIMetadata("test-cache") + require.NoError(t, err, "extractOCIMetadata succeeds - go-containerregistry is lenient") // But the metadata will be empty/invalid since it's not a real OCI config - t.Logf("Got metadata (likely empty): %+v", bundle.Meta) + t.Logf("Got metadata (likely empty): %+v", meta) } // createBuildKitCacheLayout creates an OCI layout that mimics what BuildKit @@ -337,10 +337,9 @@ func TestDockerSaveTarballToOCILayoutRoundtrip(t *testing.T) { require.NoError(t, err) assert.True(t, client.existsInLayout(layoutTag), "image should exist in layout after AppendImage") - // Step 6: Verify extractOCIImageBundle reads correct config - bundle, err := client.extractOCIImageBundle(layoutTag) + // Step 6: Verify extractOCIMetadata reads correct config + meta, err := client.extractOCIMetadata(layoutTag) require.NoError(t, err) - meta := bundle.Meta assert.Equal(t, []string{"/usr/local/bin/guest-agent"}, meta.Entrypoint) assert.Equal(t, "/app", meta.WorkingDir) assert.Contains(t, meta.Env, "PATH") diff --git a/lib/images/recovery_regression_test.go b/lib/images/recovery_regression_test.go index ad3398316..c3d7577a4 100644 --- a/lib/images/recovery_regression_test.go +++ b/lib/images/recovery_regression_test.go @@ -43,12 +43,9 @@ func TestRecoverInterruptedBuildsCapturedFixtureMarksBuildFailed(t *testing.T) { client, err := newOCIClient(p.SystemOCICache()) require.NoError(t, err) - m := &manager{ - paths: p, - ociClient: client, - queue: queue.New(1), - readySubscribers: make(map[string][]chan StatusEvent), - } + m := newTestManager(p) + m.ociClient = client + m.queue = queue.New(1) m.RecoverInterruptedBuilds() diff --git a/lib/images/storage.go b/lib/images/storage.go index 91ce646f7..d1598d821 100644 --- a/lib/images/storage.go +++ b/lib/images/storage.go @@ -35,12 +35,6 @@ type imageMetadata struct { TagGeneration uint64 `json:"tag_generation,omitempty"` } -func (m *imageMetadata) toImageFor(reference string) *Image { - img := m.toImage() - img.Name = reference - return img -} - func (m *imageMetadata) toImage() *Image { platform := m.Platform if platform == "" { @@ -104,7 +98,12 @@ func resolveImageLayout(p *paths.Paths, repository, digestHex string) imageLayou metadata: p.ImageMetadata(repository, digestHex), disk: p.ImageDigestPath(repository, digestHex), } - content := contentLayout(p, digestHex) + content := imageLayout{ + dir: p.ImageContentDir(digestHex), + metadata: p.ImageContentMetadata(digestHex), + disk: p.ImageContentPath(digestHex), + content: true, + } if legacyImageExists(p, repository, digestHex) { contentStatus, contentOK := metadataStatus(content.metadata) @@ -121,20 +120,13 @@ func resolveImageLayout(p *paths.Paths, repository, digestHex string) imageLayou return content } -func contentLayout(p *paths.Paths, digestHex string) imageLayout { - return imageLayout{ - dir: p.ImageContentDir(digestHex), - metadata: p.ImageContentMetadata(digestHex), - disk: p.ImageContentPath(digestHex), - content: true, - } -} - func pathExists(path string) bool { _, err := os.Stat(path) return err == nil } +// digestDir returns the directory for a specific digest, using the same layout +// selection as metadata and disk lookup. func digestDir(p *paths.Paths, repository, digestHex string) string { return resolveImageLayout(p, repository, digestHex).dir } @@ -179,6 +171,7 @@ func metadataPath(p *paths.Paths, repository, digestHex string) string { return resolveImageLayout(p, repository, digestHex).metadata } +// tagSymlinkPath returns the path to a tag symlink in the active layout. func tagSymlinkPath(p *paths.Paths, repository, tag string) string { newPath := p.ImageRepositoryTagSymlink(repository, tag) if _, err := os.Lstat(newPath); err == nil { @@ -187,6 +180,7 @@ func tagSymlinkPath(p *paths.Paths, repository, tag string) string { return p.ImageTagSymlink(repository, tag) } +// writeMetadata writes metadata for a digest. func writeMetadata(p *paths.Paths, repository, digestHex string, meta *imageMetadata) error { return writeMetadataFile(resolveImageLayout(p, repository, digestHex).metadata, meta) } @@ -203,26 +197,13 @@ func readMetadata(p *paths.Paths, repository, digestHex string) (*imageMetadata, return readMetadataAt(resolveImageLayout(p, repository, digestHex)) } -// resolveRefMetadata resolves a digest reference directly or a tag reference -// through its symlink, then reads the metadata it points at. -func resolveRefMetadata(p *paths.Paths, ref *NormalizedRef) (string, *imageMetadata, error) { - digestHex := ref.DigestHex() - if !ref.IsDigest() { - var err error - digestHex, err = resolveTag(p, ref.Repository(), ref.Tag()) - if err != nil { - return "", nil, err - } - } - meta, err := readMetadata(p, ref.Repository(), digestHex) - if err != nil { - return "", nil, err - } - return digestHex, meta, nil -} - func readContentMetadata(p *paths.Paths, digestHex string) (*imageMetadata, error) { - return readMetadataAt(contentLayout(p, digestHex)) + return readMetadataAt(imageLayout{ + dir: p.ImageContentDir(digestHex), + metadata: p.ImageContentMetadata(digestHex), + disk: p.ImageContentPath(digestHex), + content: true, + }) } func readMetadataAt(layout imageLayout) (*imageMetadata, error) { @@ -252,7 +233,23 @@ func readMetadataAt(layout imageLayout) (*imageMetadata, error) { return &meta, nil } -func promoteImageToContent(p *paths.Paths, sourceRepository, digestHex string, sourceMeta *imageMetadata) error { +// tagTarget is a cross-repository tag to install while promoting an image +// into shared content storage. +type tagTarget struct { + repository string + tag string +} + +func promoteImageToContent(p *paths.Paths, sourceRepository, digestHex string, sourceMeta *imageMetadata, target *tagTarget) (retErr error) { + var installedTarget *stagedTagSymlink + defer func() { + if retErr == nil || installedTarget == nil { + return + } + if restoreErr := restoreSymlinkState(installedTarget.linkPath, installedTarget.previous); restoreErr != nil { + retErr = errors.Join(retErr, fmt.Errorf("restore target tag: %w", restoreErr)) + } + }() contentReady := false if contentMeta, err := readContentMetadata(p, digestHex); err == nil { contentReady = contentMeta.Status == StatusReady @@ -281,48 +278,59 @@ func promoteImageToContent(p *paths.Paths, sourceRepository, digestHex string, s } } - return promoteLegacyTags(p, sourceRepository, digestHex) -} - -func promoteLegacyTags(p *paths.Paths, repository, digestHex string) error { - if _, err := os.Stat(p.ImageDigestDir(repository, digestHex)); err != nil { - return nil - } - tags, err := listTags(p, repository) - if err != nil { - return err - } - staged := make([]stagedTagSymlink, 0, len(tags)) - defer func() { - for _, ref := range staged { - _ = os.RemoveAll(ref.tempDir) - } - }() - for _, tag := range tags { - target, err := resolveTag(p, repository, tag) - if err != nil || target != digestHex { - continue + // Install a cross-repository target before touching source references. This + // makes target installation failures leave the source layout unchanged. + if target != nil { + staged, err := stageTagSymlink(p, target.repository, target.tag, digestHex) + if err != nil { + return fmt.Errorf("stage target tag: %w", err) } - // resolveTag validated that content-relative links resolve to this - // digest's content dir, so only legacy links (bare digest target) - // still need restaging. - if raw, err := os.Readlink(tagSymlinkPath(p, repository, tag)); err == nil && raw != digestHex { - continue + installedTarget = &staged + if err := installStagedTag(p, &staged); err != nil { + return fmt.Errorf("install target tag: %w", err) } - ref, err := stageTagSymlink(p, repository, tag, digestHex) + } + + // A legacy source may still have tags pointing at its repository-local + // digest directory. Move those references to the shared content before + // removing the duplicate legacy tree. + legacyDir := p.ImageDigestDir(sourceRepository, digestHex) + if _, err := os.Stat(legacyDir); err == nil { + tags, err := listTags(p, sourceRepository) if err != nil { - return fmt.Errorf("stage legacy tag %s: %w", tag, err) + return err } - staged = append(staged, ref) - } - for i, ref := range staged { - if err := os.Rename(ref.tempPath, ref.linkPath); err != nil { - if rollbackErr := rollbackTagSymlinks(staged[:i]); rollbackErr != nil { - return errors.Join(fmt.Errorf("promote legacy tag: %w", err), rollbackErr) + staged := make([]stagedTagSymlink, 0, len(tags)) + cleanupStaged := func() { + for _, ref := range staged { + _ = os.RemoveAll(ref.tempDir) } - return fmt.Errorf("promote legacy tag: %w", err) + } + defer cleanupStaged() + for _, tag := range tags { + target, err := resolveTag(p, sourceRepository, tag) + if err != nil || target != digestHex { + continue + } + ref, err := stageTagSymlink(p, sourceRepository, tag, digestHex) + if err != nil { + return fmt.Errorf("stage legacy tag %s: %w", tag, err) + } + staged = append(staged, ref) + } + for i := range staged { + if err := installStagedTag(p, &staged[i]); err != nil { + if rollbackErr := rollbackTagSymlinks(staged[:i+1]); rollbackErr != nil { + return errors.Join(fmt.Errorf("promote legacy tag: %w", err), rollbackErr) + } + return fmt.Errorf("promote legacy tag: %w", err) + } + } + if err := os.RemoveAll(legacyDir); err != nil { + fmt.Fprintf(os.Stderr, "Warning: failed to remove legacy digest directory %s: %v\n", digestHex, err) } } + return nil } @@ -394,6 +402,19 @@ func stageTagSymlink(p *paths.Paths, repository, tag, digestHex string) (stagedT }, nil } +// installStagedTag atomically installs a staged tag symlink, removes the +// stale legacy-layout link it replaces, and cleans up the staging directory. +// A failure leaves the link's previous state recorded on ref available for +// restoreSymlinkState. +func installStagedTag(p *paths.Paths, ref *stagedTagSymlink) error { + err := os.Rename(ref.tempPath, ref.linkPath) + _ = os.RemoveAll(ref.tempDir) + if err != nil { + return err + } + return removeStaleTagSymlink(p, ref) +} + func readSymlinkState(path string) (symlinkState, error) { target, err := os.Readlink(path) if err != nil { @@ -441,29 +462,31 @@ func removeStaleTagSymlink(p *paths.Paths, ref *stagedTagSymlink) error { return nil } +// createTagSymlink creates or updates a tag symlink to point to a digest (only +// if the digest dir exists and the build is ready). +// +// Tag ownership is Docker last-pull-wins: the most recent pull of a tag always +// owns the symlink, regardless of platform. An earlier gate only repointed for +// host-native pulls, which silently stranded emulated variants (e.g. +// `pull --platform linux/amd64 alpine:3.19` could never make `image get` report +// amd64) and was non-recoverable. Always repointing is symmetric and matches +// Docker; callers repoint unconditionally on a ready digest. func createTagSymlink(p *paths.Paths, repository, tag, digestHex string) error { ref, err := stageTagSymlink(p, repository, tag, digestHex) if err != nil { return fmt.Errorf("stage tag symlink: %w", err) } - if err := os.Rename(ref.tempPath, ref.linkPath); err != nil { - _ = os.RemoveAll(ref.tempDir) + if err := installStagedTag(p, &ref); err != nil { return fmt.Errorf("install tag symlink: %w", err) } - if err := removeStaleTagSymlink(p, &ref); err != nil { - fmt.Fprintf(os.Stderr, "Warning: failed to remove stale tag symlink %s: %v\n", tag, err) - } - _ = os.RemoveAll(ref.tempDir) return nil } -// errInvalidSymlinkTarget marks a tag symlink that does not resolve to a -// digest or the shared content directory. -var errInvalidSymlinkTarget = errors.New("invalid symlink target") - +// resolveTag follows a tag symlink to get the digest hex func resolveTag(p *paths.Paths, repository, tag string) (string, error) { linkPath := tagSymlinkPath(p, repository, tag) + // Read the symlink target, err := os.Readlink(linkPath) if err != nil { if os.IsNotExist(err) { @@ -475,16 +498,16 @@ func resolveTag(p *paths.Paths, repository, tag string) (string, error) { // Legacy links contain only the digest. New links point relatively into the // shared content directory; validate that they resolve to that digest only. if filepath.IsAbs(target) { - return "", fmt.Errorf("%w: %s", errInvalidSymlinkTarget, target) + return "", fmt.Errorf("invalid symlink target: %s", target) } digestHex := filepath.Base(target) if digestHex == "." || digestHex == string(filepath.Separator) { - return "", fmt.Errorf("%w: %s", errInvalidSymlinkTarget, target) + return "", fmt.Errorf("invalid symlink target: %s", target) } if target != digestHex { resolved := filepath.Clean(filepath.Join(filepath.Dir(linkPath), target)) if resolved != filepath.Clean(p.ImageContentDir(digestHex)) { - return "", fmt.Errorf("%w: %s", errInvalidSymlinkTarget, target) + return "", fmt.Errorf("invalid symlink target: %s", target) } } diff --git a/lib/images/storage_refs.go b/lib/images/storage_refs.go index bc25b5eda..7a73d0822 100644 --- a/lib/images/storage_refs.go +++ b/lib/images/storage_refs.go @@ -79,14 +79,23 @@ func collectLegacyImages(p *paths.Paths) ([]legacyRef, error) { return refs, nil } -func promoteLegacyImages(p *paths.Paths, refs []legacyRef) { +func promoteLegacyImages(p *paths.Paths) { + refs, err := collectLegacyImages(p) + if err != nil { + fmt.Fprintf(os.Stderr, "Warning: failed to scan legacy images for promotion: %v\n", err) + return + } + promoteLegacyImageRefs(p, refs) +} + +func promoteLegacyImageRefs(p *paths.Paths, refs []legacyRef) { for _, ref := range refs { layout := resolveImageLayout(p, ref.repository, ref.digestHex) meta, readErr := readMetadataAt(layout) if readErr != nil || meta.Status != StatusReady { continue } - if promoteErr := promoteImageToContent(p, ref.repository, ref.digestHex, meta); promoteErr != nil { + if promoteErr := promoteImageToContent(p, ref.repository, ref.digestHex, meta, nil); promoteErr != nil { fmt.Fprintf(os.Stderr, "Warning: failed to promote legacy image %s@%s: %v\n", ref.repository, ref.digestHex, promoteErr) } } diff --git a/lib/images/storage_test.go b/lib/images/storage_test.go index d7c8fd994..cce34ec6c 100644 --- a/lib/images/storage_test.go +++ b/lib/images/storage_test.go @@ -203,7 +203,7 @@ func TestImageMetadataToImage_ClonesMetadata(t *testing.T) { CreatedAt: createdAt, } - img := source.toImageFor(source.Name) + img := source.toImage() require.Equal(t, source.Name, img.Name) require.Equal(t, source.Digest, img.Digest) require.Equal(t, map[string]string{"team": "backend", "env": "staging"}, img.Tags) @@ -220,7 +220,7 @@ func TestImageMetadataToImage_EmptyMetadataOmitted(t *testing.T) { Digest: "sha256:abc", Status: StatusPending, CreatedAt: time.Now().UTC(), - }).toImageFor("docker.io/library/alpine:latest") + }).toImage() require.Nil(t, img.Tags) } @@ -257,7 +257,9 @@ func TestPromoteLegacyImagesMovesContentAndTags(t *testing.T) { require.NoError(t, err) require.Equal(t, "rootfs!", string(data)) - require.FileExists(t, p.ImageDigestPath(repository, digest)) + // Legacy digest tree is retired once content is installed and tags moved. + _, err = os.Stat(legacyDir) + require.True(t, os.IsNotExist(err), "legacy digest dir should be removed") // The tag now resolves through the shared content layout. resolved, err := resolveTag(p, repository, tag) diff --git a/lib/images/tag_manager.go b/lib/images/tag_manager.go deleted file mode 100644 index a700e59c1..000000000 --- a/lib/images/tag_manager.go +++ /dev/null @@ -1,272 +0,0 @@ -package images - -import ( - "context" - "errors" - "fmt" - "log/slog" - "os" - "strings" -) - -func tagGenerationKey(repository, tag string) string { - return repository + ":" + tag -} - -// nextTagGeneration bumps the mutation generation for a tag and returns the -// new value. Pulls record the value with their metadata and only repoint the -// tag on completion when no later operation mutated it. -func (m *manager) nextTagGeneration(repository, tag string) uint64 { - if m.tagGenerations == nil { - m.tagGenerations = make(map[string]uint64) - } - key := tagGenerationKey(repository, tag) - m.tagGenerations[key]++ - return m.tagGenerations[key] -} - -// revertTagGeneration undoes a nextTagGeneration bump for an operation that -// failed before mutating the tag, so an in-flight pull of the tag is not -// permanently blocked from repointing it. -func (m *manager) revertTagGeneration(repository, tag string) { - key := tagGenerationKey(repository, tag) - if m.tagGenerations[key] <= 1 { - delete(m.tagGenerations, key) - return - } - m.tagGenerations[key]-- -} - -// releaseTagGeneration drops a failed pull's claim on its tag so an older -// in-flight pull of the same tag can still repoint it. -func (m *manager) releaseTagGeneration(repository, tag string, generation uint64) { - if m.tagGenerations[tagGenerationKey(repository, tag)] == generation { - m.revertTagGeneration(repository, tag) - } -} - -// pruneTagGenerations drops entries for tags that no longer resolve, keeping -// the map bounded over the process lifetime. Deleting a tag still suppresses -// an in-flight pull's repoint: the pull's recorded generation can no longer -// match the missing entry. -func (m *manager) pruneTagGenerations() { - for key := range m.tagGenerations { - // tagGenerationKey is repo+":"+tag and tags never contain ":", but - // repositories may (host:port), so split on the last colon. - colon := strings.LastIndexByte(key, ':') - if colon < 0 { - continue - } - if _, err := resolveTag(m.paths, key[:colon], key[colon+1:]); err != nil { - delete(m.tagGenerations, key) - } - } -} - -// restoreTagState re-seeds the tag indexes from recovered metadata. metas must -// be sorted oldest first so the newest requested pull wins. -func (m *manager) restoreTagState(metas []*imageMetadata) { - if m.tagGenerations == nil { - m.tagGenerations = make(map[string]uint64) - } - if m.requestedTags == nil { - m.requestedTags = make(map[string]string) - } - for _, meta := range metas { - if meta.RequestedTag == "" || meta.Digest == "" { - continue - } - ref, err := ParseNormalizedRef(meta.Name) - if err != nil { - continue - } - key := tagGenerationKey(ref.Repository(), meta.RequestedTag) - if meta.TagGeneration > m.tagGenerations[key] { - m.tagGenerations[key] = meta.TagGeneration - } - m.requestedTags[key] = strings.TrimPrefix(meta.Digest, "sha256:") - } -} - -// claimTagForStatus repoints ref's tag at the digest when the image is ready, -// or records a pending tag when the build is still in flight. It is a no-op -// for digest-only references. -func (m *manager) claimTagForStatus(meta *imageMetadata, ref *ResolvedRef) error { - if ref.Tag() == "" { - return nil - } - if meta.Status == StatusReady { - return m.claimReadyTag(ref.Repository(), ref.Tag(), ref.DigestHex()) - } - return ensurePendingTag(m.paths, ref.Repository(), ref.Tag(), ref.DigestHex()) -} - -// claimReadyTag repoints an existing tag at a ready digest, last pull wins. -func (m *manager) claimReadyTag(repository, tag, digestHex string) error { - m.nextTagGeneration(repository, tag) - if err := createTagSymlink(m.paths, repository, tag, digestHex); err != nil { - m.revertTagGeneration(repository, tag) - return err - } - return nil -} - -// trackRequestedTag records the digest of the newest pull requested for a tag -// so readiness waits can find it without walking the metadata tree. -func (m *manager) trackRequestedTag(repository, tag, digestHex string) { - if m.requestedTags == nil { - m.requestedTags = make(map[string]string) - } - m.requestedTags[tagGenerationKey(repository, tag)] = digestHex -} - -// requestedTagImage returns the newest image requested for a tag, so a -// readiness wait tracks the latest pull rather than the digest the tag -// currently points at. -func (m *manager) requestedTagImage(ref *NormalizedRef) *Image { - m.createMu.Lock() - digestHex, ok := m.requestedTags[tagGenerationKey(ref.Repository(), ref.Tag())] - m.createMu.Unlock() - if !ok { - return nil - } - meta, err := readMetadata(m.paths, ref.Repository(), digestHex) - if err != nil { - return nil - } - return meta.toImageFor(ref.String()) -} - -// claimRequestedTag repoints the pull's requested tag at the finished digest. -// The tag is only repointed when its generation still matches the pull's -// recorded claim, so a later mutation of the tag wins over the in-flight -// pull. A recovered build whose metadata predates requested-tag tracking has -// no recorded tag: fall back to the reference's own tag and recreate a -// missing symlink, matching the pre-tracking behavior. Otherwise an -// already-missing tag stays missing so a concurrent delete wins. -func (m *manager) claimRequestedTag(ref *ResolvedRef, meta *imageMetadata) bool { - requestedTag := meta.RequestedTag - allowMissing := requestedTag == "" - if requestedTag == "" { - requestedTag = ref.Tag() - } - if requestedTag == "" || m.tagGenerations[tagGenerationKey(ref.Repository(), requestedTag)] != meta.TagGeneration { - return false - } - current, err := resolveTag(m.paths, ref.Repository(), requestedTag) - if err != nil { - if !allowMissing || !errors.Is(err, ErrNotFound) { - return false - } - } else if current != ref.DigestHex() && current != meta.PreviousTagDigest { - return false - } - if err := createTagSymlink(m.paths, ref.Repository(), requestedTag, ref.DigestHex()); err != nil { - fmt.Fprintf(os.Stderr, "Warning: failed to create tag symlink: %v\n", err) - return false - } - return true -} - -// TagImage creates a ready-image tag without pulling or converting content. -// Cross-repository tags promote legacy content into the shared layout. A -// failed call leaves no side effects: the target tag's generation is only -// bumped after the new tag is on disk, so pending pulls that claimed the -// target tag keep their claim. When the target previously pointed at -// different content, that digest is collected after the new tag is live; -// cleanup failures are logged and do not fail the call, since the tag is -// already installed at that point. -// -// Promotion deliberately runs before the symlink install, so a symlink -// failure after a cross-repo promotion leaves the content promoted with no -// target tag. That state is gc-consistent (unreferenced content is -// collected) and retry is idempotent (promoteImageToContent short-circuits -// on ready content), which beats rolling back a completed promotion. -func (m *manager) TagImage(ctx context.Context, source, target string) (*Image, error) { - sourceRef, targetRef, err := parseTagReferences(source, target) - if err != nil { - return nil, err - } - - m.createMu.Lock() - defer m.createMu.Unlock() - - // A dangling or malformed target symlink is treated like a missing tag so - // the retag self-heals; createTagSymlink replaces the link either way. - previousDigest, err := resolveTag(m.paths, targetRef.Repository(), targetRef.Tag()) - if err != nil && !errors.Is(err, ErrNotFound) && !errors.Is(err, errInvalidSymlinkTarget) { - return nil, fmt.Errorf("resolve existing target tag: %w", err) - } - - digestHex, meta, err := m.readyTagImage(sourceRef) - if err != nil { - return nil, err - } - if sourceRef.Repository() != targetRef.Repository() { - if err := promoteImageToContent(m.paths, sourceRef.Repository(), digestHex, meta); err != nil { - return nil, fmt.Errorf("promote image to content: %w", err) - } - } - if err := createTagSymlink(m.paths, targetRef.Repository(), targetRef.Tag(), digestHex); err != nil { - return nil, fmt.Errorf("create image tag: %w", err) - } - if err := writeMetadata(m.paths, targetRef.Repository(), digestHex, meta); err != nil { - return nil, fmt.Errorf("write tagged image metadata: %w", err) - } - // Unlike updateExistingReference, which bumps the generation before - // installing the symlink, the bump happens after the install here so a - // failed tag call cannot invalidate a pending pull's claim on the target - // tag (pinned by TestTagImageFailureLeavesNoSideEffects). - m.nextTagGeneration(targetRef.Repository(), targetRef.Tag()) - m.cleanupReplacedTag(targetRef, previousDigest, digestHex) - - return meta.toImageFor(targetRef.String()), nil -} - -func (m *manager) cleanupUnclaimedImage(ref *ResolvedRef) { - if err := removeDigestIfUnreferenced(m.paths, ref.Repository(), ref.DigestHex(), true); err != nil { - slog.Warn("failed to collect stale image", "repository", ref.Repository(), "digest", ref.DigestHex(), "error", err) - } -} - -func (m *manager) cleanupReplacedTag(ref *NormalizedRef, previousDigest, digestHex string) { - if previousDigest != "" && previousDigest != digestHex { - // Sibling tags in this repository may still reference the previous - // digest; only collect when this was the last reference. - count, err := countTagsForDigest(m.paths, ref.Repository(), previousDigest) - if err != nil { - slog.Warn("failed to count tags for replaced image", "repository", ref.Repository(), "digest", previousDigest, "error", err) - } else if count == 0 { - if err := removeDigestIfUnreferenced(m.paths, ref.Repository(), previousDigest, true); err != nil { - slog.Warn("failed to collect replaced image content", "repository", ref.Repository(), "digest", previousDigest, "error", err) - } - m.refreshDiskUsageTotals() - } - } -} - -func parseTagReferences(source, target string) (*NormalizedRef, *NormalizedRef, error) { - sourceRef, err := ParseNormalizedRef(source) - if err != nil { - return nil, nil, fmt.Errorf("%w: invalid source reference: %s", ErrInvalidName, err) - } - targetRef, err := ParseNormalizedRef(target) - if err != nil { - return nil, nil, fmt.Errorf("%w: invalid target reference: %s", ErrInvalidName, err) - } - if targetRef.IsDigest() { - return nil, nil, fmt.Errorf("%w: target must be a tag reference, not a digest", ErrInvalidName) - } - return sourceRef, targetRef, nil -} - -func (m *manager) readyTagImage(ref *NormalizedRef) (string, *imageMetadata, error) { - digestHex, meta, err := resolveRefMetadata(m.paths, ref) - if err != nil { - return "", nil, err - } - if meta.Status != StatusReady { - return "", nil, fmt.Errorf("%w: %s", ErrImageNotReady, meta.Status) - } - return digestHex, meta, nil -} diff --git a/lib/images/tag_test.go b/lib/images/tag_test.go index 75796ab2f..8c25f6c6b 100644 --- a/lib/images/tag_test.go +++ b/lib/images/tag_test.go @@ -3,7 +3,6 @@ package images import ( "context" "os" - "path/filepath" "testing" "time" @@ -11,248 +10,223 @@ import ( "github.com/stretchr/testify/require" ) -func seedContent(t *testing.T, p *paths.Paths, repository, tag, digestHex string) { +// seedReadyContentImage writes a ready image directly into the shared content +// layout and tags it, without going through a pull. +func seedReadyContentImage(t *testing.T, p *paths.Paths, repository, tag, digestHex string) { t.Helper() - seedImage(t, p, repository, tag, digestHex, true) + require.NoError(t, os.MkdirAll(p.ImageContentDir(digestHex), 0o755)) + require.NoError(t, writeMetadataFile(p.ImageContentMetadata(digestHex), &imageMetadata{ + Name: repository + ":" + tag, + Digest: "sha256:" + digestHex, + Status: StatusReady, + SizeBytes: 7, + CreatedAt: time.Now().UTC(), + })) + require.NoError(t, os.WriteFile(p.ImageContentPath(digestHex), []byte("rootfs!"), 0o644)) + require.NoError(t, createTagSymlink(p, repository, tag, digestHex)) } -func seedLegacy(t *testing.T, p *paths.Paths, repository, tag, digestHex string) { - t.Helper() - seedImage(t, p, repository, tag, digestHex, false) +func newTagTestManager(p *paths.Paths) *manager { + return newTestManager(p) } -func seedImage(t *testing.T, p *paths.Paths, repository, tag, digestHex string, content bool) { - t.Helper() - dir := p.ImageDigestDir(repository, digestHex) - disk := p.ImageDigestPath(repository, digestHex) - metadata := p.ImageMetadata(repository, digestHex) - linkPath := p.ImageTagSymlink(repository, tag) - target := digestHex - if content { - dir = p.ImageContentDir(digestHex) - disk = p.ImageContentPath(digestHex) - metadata = p.ImageContentMetadata(digestHex) - linkPath = p.ImageRepositoryTagSymlink(repository, tag) - rel, err := filepath.Rel(filepath.Dir(linkPath), p.ImageContentDir(digestHex)) - require.NoError(t, err) - target = rel - } - - require.NoError(t, os.MkdirAll(dir, 0o755)) - require.NoError(t, os.WriteFile(disk, []byte("rootfs!"), 0o644)) - require.NoError(t, writeMetadataFile(metadata, &imageMetadata{ - Name: repository + ":" + tag, Digest: "sha256:" + digestHex, - Status: StatusReady, SizeBytes: int64(len("rootfs!")), CreatedAt: time.Now().UTC(), - })) - require.NoError(t, os.MkdirAll(filepath.Dir(linkPath), 0o755)) - require.NoError(t, os.Symlink(target, linkPath)) -} - -func newTagTestCase(t *testing.T) (*paths.Paths, *manager, string) { - t.Helper() +func TestTagImageSameRepository(t *testing.T) { p := paths.New(t.TempDir()) - return p, &manager{paths: p, tagGenerations: make(map[string]uint64)}, "docker.io/library/alpine" -} + m := newTagTestManager(p) + repository := "docker.io/library/alpine" + digest := "a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1" + seedReadyContentImage(t, p, repository, "latest", digest) -func requireTagResolvesTo(t *testing.T, p *paths.Paths, repository, tag, digest string) { - t.Helper() - resolved, err := resolveTag(p, repository, tag) + img, err := m.TagImage(context.Background(), repository+":latest", repository+":stable") require.NoError(t, err) - require.Equal(t, digest, resolved) -} + require.Equal(t, repository+":stable", img.Name) + require.Equal(t, "sha256:"+digest, img.Digest) -func tagImage(t *testing.T, m *manager, source, target, digest string) { - t.Helper() - img, err := m.TagImage(context.Background(), source, target) + resolved, err := resolveTag(p, repository, "stable") require.NoError(t, err) - require.Equal(t, target, img.Name) - require.Equal(t, "sha256:"+digest, img.Digest) + require.Equal(t, digest, resolved) + + // Both tags keep the shared content alive. + require.NoError(t, m.DeleteImage(context.Background(), repository+":latest")) + _, err = os.Stat(p.ImageContentDir(digest)) + require.NoError(t, err, "content must survive while another tag references it") + + require.NoError(t, m.DeleteImage(context.Background(), repository+":stable")) + _, err = os.Stat(p.ImageContentDir(digest)) + require.True(t, os.IsNotExist(err), "content must be removed once unreferenced") } -func TestTagImageAliasesReadyImage(t *testing.T) { - cases := []struct { - name, source, target, digest string - }{ - {"same repository", "docker.io/library/alpine:latest", "docker.io/library/alpine:stable", "a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1"}, - {"digest source", "docker.io/library/alpine@sha256:b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2", "docker.io/library/alpine:pinned", "b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2"}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - p, m, repository := newTagTestCase(t) - seedContent(t, p, repository, "latest", tc.digest) - tagImage(t, m, tc.source, tc.target, tc.digest) - targetRef, err := ParseNormalizedRef(tc.target) - require.NoError(t, err) - requireTagResolvesTo(t, p, targetRef.Repository(), targetRef.Tag(), tc.digest) - }) +func TestStaleTagClaimCollectsContent(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" + digest := "f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2" + seedReadyContentImage(t, p, repository, "latest", digest) + + linkPath := p.ImageRepositoryTagSymlink(repository, "latest") + require.NoError(t, os.Remove(linkPath)) + require.NoError(t, os.Symlink("replaced", linkPath)) + + normalized, err := ParseNormalizedRef(repository + ":latest") + require.NoError(t, err) + m.tagGenerations[repository+":latest"] = 2 + ref := NewResolvedRef(normalized, "sha256:"+digest) + meta := &imageMetadata{ + Name: repository + ":latest", Digest: "sha256:" + digest, + Status: StatusReady, RequestedTag: "latest", TagGeneration: 1, } + + require.False(t, m.publishReadyTag(ref, meta)) + m.cleanupUnclaimedImage(ref) + _, err = os.Stat(p.ImageContentDir(digest)) + require.ErrorIs(t, err, os.ErrNotExist) } -func TestTagImageSameRepositoryDeletesContent(t *testing.T) { - p, m, repository := newTagTestCase(t) - digest := "a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1a1" - seedContent(t, p, repository, "latest", digest) +func TestRecoveredTagClaimRecreatesMissingTag(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" + digest := "e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3e3" + seedReadyContentImage(t, p, repository, "latest", digest) + require.NoError(t, os.Remove(p.ImageRepositoryTagSymlink(repository, "latest"))) - tagImage(t, m, repository+":latest", repository+":stable", digest) - requireTagResolvesTo(t, p, repository, "stable", digest) + normalized, err := ParseNormalizedRef(repository + ":latest") + require.NoError(t, err) + ref := NewResolvedRef(normalized, "sha256:"+digest) + meta, err := readContentMetadata(p, digest) + require.NoError(t, err) - require.NoError(t, m.DeleteImage(context.Background(), repository+":latest")) - _, err := os.Stat(p.ImageContentDir(digest)) + require.True(t, m.publishReadyTag(ref, meta)) + resolved, err := resolveTag(p, repository, "latest") require.NoError(t, err) - require.NoError(t, m.DeleteImage(context.Background(), repository+":stable")) - _, err = os.Stat(p.ImageContentDir(digest)) - require.True(t, os.IsNotExist(err), "content must be removed once unreferenced") + require.Equal(t, digest, resolved) } -func TestTagImageLegacyLayoutSource(t *testing.T) { - p, m, repository := newTagTestCase(t) - digest := "d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8d8" - seedLegacy(t, p, repository, "latest", digest) +func TestTagImageFromDigestSource(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" + digest := "b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2b2" + seedReadyContentImage(t, p, repository, "latest", digest) + + img, err := m.TagImage(context.Background(), repository+"@sha256:"+digest, repository+":pinned") + require.NoError(t, err) + require.Equal(t, repository+":pinned", img.Name) - tagImage(t, m, repository+":latest", repository+":v2", digest) - requireTagResolvesTo(t, p, repository, "v2", digest) + resolved, err := resolveTag(p, repository, "pinned") + require.NoError(t, err) + require.Equal(t, digest, resolved) } func TestTagImageCrossRepository(t *testing.T) { - p, m, sourceRepo := newTagTestCase(t) + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + sourceRepo := "docker.io/library/alpine" targetRepo := "registry.example/apps/alpine" digest := "c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3" - seedLegacy(t, p, sourceRepo, "latest", digest) - tagImage(t, m, sourceRepo+":latest", targetRepo+":v1", digest) + // Ready image still living in the legacy per-repository layout. + legacyDir := p.ImageDigestDir(sourceRepo, digest) + require.NoError(t, os.MkdirAll(legacyDir, 0o755)) + require.NoError(t, writeMetadataFile(p.ImageMetadata(sourceRepo, digest), &imageMetadata{ + Name: sourceRepo + ":latest", + Digest: "sha256:" + digest, + Status: StatusReady, + SizeBytes: 7, + CreatedAt: time.Now().UTC(), + })) + require.NoError(t, os.WriteFile(p.ImageDigestPath(sourceRepo, digest), []byte("rootfs!"), 0o644)) + require.NoError(t, createTagSymlink(p, sourceRepo, "latest", digest)) + + img, err := m.TagImage(context.Background(), sourceRepo+":latest", targetRepo+":v1") + require.NoError(t, err) + require.Equal(t, targetRepo+":v1", img.Name) + + // Promotion moved the digest into shared content and retired the legacy tree. data, err := os.ReadFile(p.ImageContentPath(digest)) require.NoError(t, err) require.Equal(t, "rootfs!", string(data)) - require.DirExists(t, p.ImageDigestDir(sourceRepo, digest)) + _, err = os.Stat(legacyDir) + require.True(t, os.IsNotExist(err), "legacy tree should be retired after promotion") + // Both repositories resolve to the same digest. for repo, tag := range map[string]string{sourceRepo: "latest", targetRepo: "v1"} { - requireTagResolvesTo(t, p, repo, tag, digest) + resolved, err := resolveTag(p, repo, tag) + require.NoError(t, err) + require.Equal(t, digest, resolved) } + + // Deleting the source tag leaves the cross-repository alias intact. require.NoError(t, m.DeleteImage(context.Background(), sourceRepo+":latest")) _, err = os.Stat(p.ImageContentDir(digest)) - require.NoError(t, err) + require.NoError(t, err, "cross-repository tag must keep content alive") + require.NoError(t, m.DeleteImage(context.Background(), targetRepo+":v1")) _, err = os.Stat(p.ImageContentDir(digest)) require.True(t, os.IsNotExist(err), "content must be removed once unreferenced") } -func TestStaleTagClaimCollectsContent(t *testing.T) { - p, m, repository := newTagTestCase(t) - digest := "f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2f2" - seedContent(t, p, repository, "latest", digest) - linkPath := p.ImageRepositoryTagSymlink(repository, "latest") - require.NoError(t, os.Remove(linkPath)) - require.NoError(t, os.Symlink("replaced", linkPath)) - - normalized, err := ParseNormalizedRef(repository + ":latest") - require.NoError(t, err) - m.tagGenerations[repository+":latest"] = 2 - ref := NewResolvedRef(normalized, "sha256:"+digest) - meta, err := readContentMetadata(p, digest) - require.NoError(t, err) - meta.RequestedTag = "latest" - meta.TagGeneration = 1 - - require.False(t, m.claimRequestedTag(ref, meta)) - m.cleanupUnclaimedImage(ref) - - _, err = os.Stat(p.ImageContentDir(digest)) - require.ErrorIs(t, err, os.ErrNotExist) -} - func TestTagImageRejectsNotReady(t *testing.T) { - p, m, repository := newTagTestCase(t) + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" digest := "d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4" + require.NoError(t, os.MkdirAll(p.ImageContentDir(digest), 0o755)) require.NoError(t, writeMetadataFile(p.ImageContentMetadata(digest), &imageMetadata{ - Name: repository + ":latest", Digest: "sha256:" + digest, - Status: StatusConverting, CreatedAt: time.Now().UTC(), + Name: repository + ":latest", + Digest: "sha256:" + digest, + Status: StatusConverting, + CreatedAt: time.Now().UTC(), })) require.NoError(t, createTagSymlink(p, repository, "latest", digest)) - targetRepository := "registry.example/apps/alpine" - _, err := m.TagImage(context.Background(), repository+":latest", targetRepository+":stable") + _, err := m.TagImage(context.Background(), repository+":latest", repository+":stable") require.ErrorIs(t, err, ErrImageNotReady) - _, err = resolveTag(p, targetRepository, "stable") - require.ErrorIs(t, err, ErrNotFound) -} -func TestTagImageSourceNotFound(t *testing.T) { - _, m, repository := newTagTestCase(t) - _, err := m.TagImage(context.Background(), repository+":missing", repository+":stable") - require.ErrorIs(t, err, ErrNotFound) + _, err = resolveTag(p, repository, "stable") + require.ErrorIs(t, err, ErrNotFound, "failed tag must not create a reference") } -func TestTagImageFailureLeavesNoSideEffects(t *testing.T) { - p, m, repository := newTagTestCase(t) - digest := "d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5d5" - seedContent(t, p, repository, "latest", digest) +func TestTagImageSourceNotFound(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" - // Failed calls must not create the target tag or bump its generation, - // which would invalidate a pending pull's claim on the target tag. _, err := m.TagImage(context.Background(), repository+":missing", repository+":stable") require.ErrorIs(t, err, ErrNotFound) - _, err = resolveTag(p, repository, "stable") - require.ErrorIs(t, err, ErrNotFound) - require.Equal(t, uint64(0), m.tagGenerations[repository+":stable"]) - - tagImage(t, m, repository+":latest", repository+":stable", digest) - requireTagResolvesTo(t, p, repository, "stable", digest) } func TestTagImageRejectsDigestTarget(t *testing.T) { - p, m, repository := newTagTestCase(t) + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" digest := "e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5e5" - seedContent(t, p, repository, "latest", digest) + seedReadyContentImage(t, p, repository, "latest", digest) + _, err := m.TagImage(context.Background(), repository+":latest", repository+"@sha256:"+digest) require.ErrorIs(t, err, ErrInvalidName) } -func TestTagImageReplacesExistingTagAndCollectsOldContent(t *testing.T) { - p, m, repository := newTagTestCase(t) +func TestTagImageReplacesExistingTag(t *testing.T) { + p := paths.New(t.TempDir()) + m := newTagTestManager(p) + repository := "docker.io/library/alpine" first := "f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6f6" second := "a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7a7" - seedContent(t, p, repository, "latest", first) - seedContent(t, p, repository, "stable", second) + seedReadyContentImage(t, p, repository, "latest", first) + seedReadyContentImage(t, p, repository, "stable", second) - readyBytes, err := m.TotalImageBytes(context.Background()) + // "stable" moves from the second digest to the first. + img, err := m.TagImage(context.Background(), repository+":latest", repository+":stable") require.NoError(t, err) - require.Equal(t, int64(14), readyBytes) - - tagImage(t, m, repository+":latest", repository+":stable", first) - requireTagResolvesTo(t, p, repository, "stable", first) - _, err = os.Stat(p.ImageContentDir(second)) - require.ErrorIs(t, err, os.ErrNotExist) + require.Equal(t, repository+":stable", img.Name) - readyBytes, err = m.TotalImageBytes(context.Background()) + resolved, err := resolveTag(p, repository, "stable") require.NoError(t, err) - require.Equal(t, int64(7), readyBytes) -} - -func TestTagImageKeepsLegacySiblingTags(t *testing.T) { - p, m, repository := newTagTestCase(t) - first := "c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3c3" - second := "d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4d4" - seedLegacy(t, p, repository, "latest", first) - seedLegacy(t, p, repository, "stable", first) - seedLegacy(t, p, repository, "v2", second) - - tagImage(t, m, repository+":v2", repository+":stable", second) - requireTagResolvesTo(t, p, repository, "stable", second) - requireTagResolvesTo(t, p, repository, "latest", first) - _, err := os.Stat(p.ImageDigestDir(repository, first)) - require.NoError(t, err, "legacy content for a sibling tag must not be collected") -} - -func TestTagImageSelfHealsDanglingTargetSymlink(t *testing.T) { - p, m, repository := newTagTestCase(t) - digest := "e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6e6" - seedContent(t, p, repository, "latest", digest) + require.Equal(t, first, resolved) - linkPath := p.ImageRepositoryTagSymlink(repository, "stable") - require.NoError(t, os.MkdirAll(filepath.Dir(linkPath), 0o755)) - require.NoError(t, os.Symlink("already-collected", linkPath)) - - tagImage(t, m, repository+":latest", repository+":stable", digest) - requireTagResolvesTo(t, p, repository, "stable", digest) + // Replacing the only tag removes the old digest once it is unreferenced. + _, err = os.Stat(p.ImageContentDir(second)) + require.ErrorIs(t, err, os.ErrNotExist) } diff --git a/lib/images/testlayers_test.go b/lib/images/testlayers_test.go deleted file mode 100644 index d56cf1015..000000000 --- a/lib/images/testlayers_test.go +++ /dev/null @@ -1,52 +0,0 @@ -package images - -import ( - "archive/tar" - "bytes" - "compress/gzip" - "io" - "testing" - - gcr "github.com/google/go-containerregistry/pkg/v1" - "github.com/google/go-containerregistry/pkg/v1/tarball" - "github.com/stretchr/testify/require" -) - -type tarEntrySpec struct { - name string - content string - isDir bool - mode int64 -} - -// specLayer builds a gzipped tar layer from entry specs in order. -func specLayer(t *testing.T, entries []tarEntrySpec) gcr.Layer { - t.Helper() - - var buf bytes.Buffer - gzw := gzip.NewWriter(&buf) - tw := tar.NewWriter(gzw) - for _, entry := range entries { - if entry.isDir { - require.NoError(t, tw.WriteHeader(&tar.Header{Name: entry.name, Typeflag: tar.TypeDir, Mode: entry.mode})) - continue - } - require.NoError(t, tw.WriteHeader(&tar.Header{ - Name: entry.name, - Typeflag: tar.TypeReg, - Mode: entry.mode, - Size: int64(len(entry.content)), - })) - _, err := tw.Write([]byte(entry.content)) - require.NoError(t, err) - } - require.NoError(t, tw.Close()) - require.NoError(t, gzw.Close()) - - data := buf.Bytes() - layer, err := tarball.LayerFromOpener(func() (io.ReadCloser, error) { - return io.NopCloser(bytes.NewReader(data)), nil - }) - require.NoError(t, err) - return layer -} diff --git a/lib/images/testutil/testutil.go b/lib/images/testutil/testutil.go deleted file mode 100644 index aa28cd1e4..000000000 --- a/lib/images/testutil/testutil.go +++ /dev/null @@ -1,107 +0,0 @@ -// Package testutil provides helpers for seeding on-disk image state in tests. -package testutil - -import ( - "encoding/json" - "fmt" - "os" - "path/filepath" - "testing" - "time" - - "github.com/kernel/hypeman/lib/images" - "github.com/kernel/hypeman/lib/paths" - "github.com/stretchr/testify/require" -) - -const seedImageContent = "rootfs!" - -// Seed describes a ready image to seed directly to disk, bypassing pulls, -// builds, and conversion. -type Seed struct { - Repository string - DigestHex string - // Tag optionally creates a tag symlink pointing at the image. - Tag string - // Name overrides the metadata "name" field; it defaults to - // repository:tag, or repository@sha256:digest when Tag is empty. - Name string - // Tags records resource tags in the metadata. - Tags map[string]string - // Content writes the image into the shared content layout instead of - // the legacy per-repository digest layout. - Content bool -} - -type imageMetadata struct { - Name string `json:"name"` - Digest string `json:"digest"` - Status string `json:"status"` - SizeBytes int64 `json:"size_bytes"` - Tags map[string]string `json:"tags,omitempty"` - CreatedAt time.Time `json:"created_at"` -} - -// SeedReadyImage writes a ready image to disk per s. -func SeedReadyImage(t testing.TB, p *paths.Paths, s Seed) { - t.Helper() - require.NoError(t, seedImage(p, s)) -} - -func seedImage(p *paths.Paths, s Seed) error { - dir := p.ImageDigestDir(s.Repository, s.DigestHex) - disk := p.ImageDigestPath(s.Repository, s.DigestHex) - metadata := p.ImageMetadata(s.Repository, s.DigestHex) - linkPath := p.ImageTagSymlink(s.Repository, s.Tag) - target := s.DigestHex - if s.Content { - dir = p.ImageContentDir(s.DigestHex) - disk = p.ImageContentPath(s.DigestHex) - metadata = p.ImageContentMetadata(s.DigestHex) - linkPath = p.ImageRepositoryTagSymlink(s.Repository, s.Tag) - rel, err := filepath.Rel(filepath.Dir(linkPath), p.ImageContentDir(s.DigestHex)) - if err != nil { - return fmt.Errorf("rel content symlink target: %w", err) - } - target = rel - } - - if err := os.MkdirAll(dir, 0o755); err != nil { - return fmt.Errorf("create image dir: %w", err) - } - if err := os.WriteFile(disk, []byte(seedImageContent), 0o644); err != nil { - return fmt.Errorf("write disk image: %w", err) - } - - name := s.Name - if name == "" && s.Tag != "" { - name = s.Repository + ":" + s.Tag - } else if name == "" { - name = s.Repository + "@sha256:" + s.DigestHex - } - data, err := json.Marshal(imageMetadata{ - Name: name, - Digest: "sha256:" + s.DigestHex, - Status: images.StatusReady, - SizeBytes: int64(len(seedImageContent)), - Tags: s.Tags, - CreatedAt: time.Now().UTC(), - }) - if err != nil { - return fmt.Errorf("marshal metadata: %w", err) - } - if err := os.WriteFile(metadata, data, 0o644); err != nil { - return fmt.Errorf("write metadata: %w", err) - } - - if s.Tag == "" { - return nil - } - if err := os.MkdirAll(filepath.Dir(linkPath), 0o755); err != nil { - return fmt.Errorf("create tag dir: %w", err) - } - if err := os.Symlink(target, linkPath); err != nil { - return fmt.Errorf("create tag symlink: %w", err) - } - return nil -} diff --git a/lib/instances/admission_allocations.go b/lib/instances/admission_allocations.go index 17d0f5e94..e07cc71bd 100644 --- a/lib/instances/admission_allocations.go +++ b/lib/instances/admission_allocations.go @@ -107,7 +107,7 @@ func (m *manager) rollbackAdmissionAllocationActive(stored *StoredMetadata) { // Failed post-boot/restore steps should not leave the cached visible // allocation marked active. Clear the in-memory PID first so any later sync // from this metadata view also treats the instance as inactive. - stored.HypervisorProcessIdentity.Clear() + stored.HypervisorPID = nil m.setAdmissionAllocationActive(stored, false) } diff --git a/lib/instances/admission_allocations_test.go b/lib/instances/admission_allocations_test.go index 17674e046..26c4252c4 100644 --- a/lib/instances/admission_allocations_test.go +++ b/lib/instances/admission_allocations_test.go @@ -20,11 +20,11 @@ func TestRollbackAdmissionAllocationActiveClearsVisibleAllocation(t *testing.T) } pid := 1234 stored := &StoredMetadata{ - Id: "inst-1", - Name: "test-instance", - Vcpus: 2, - Size: 1024, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &pid}, + Id: "inst-1", + Name: "test-instance", + Vcpus: 2, + Size: 1024, + HypervisorPID: &pid, } m.setAdmissionAllocationActive(stored, true) @@ -52,12 +52,12 @@ func TestReconcileAdmissionAllocationsMarksMissingSocketInactive(t *testing.T) { pid := 4321 stored := StoredMetadata{ - Id: "inst-2", - Name: "test-instance", - Vcpus: 2, - Size: 1024, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &pid}, - SocketPath: socketPath, + Id: "inst-2", + Name: "test-instance", + Vcpus: 2, + Size: 1024, + HypervisorPID: &pid, + SocketPath: socketPath, } m := &manager{ diff --git a/lib/instances/create.go b/lib/instances/create.go index add25e4e7..0d21e8e91 100644 --- a/lib/instances/create.go +++ b/lib/instances/create.go @@ -52,7 +52,7 @@ var systemDirectories = []string{ "/var", } -func wrapCreateVGPUErr(profile string, err error) error { +func wrapCreateMdevErr(profile string, err error) error { if errors.Is(err, devices.ErrVGPUNotSupportedOnMacOS) { return fmt.Errorf("%w: %w", ErrInvalidRequest, err) } @@ -272,10 +272,7 @@ func (m *manager) createInstance( // whatever devices have been attached when cleanup runs. var attachedDeviceIDs []string var resolvedDeviceIDs []string - var gpuDevice *devices.VGPUDevice var gpuProfile string - var gpuFramework devices.VGPUFramework - var gpuDevicePath string var gpuMdevUUID string // Setup cleanup stack early so device attachment errors trigger cleanup @@ -295,30 +292,23 @@ func (m *manager) createInstance( }) } - // Handle vGPU profile request + // Handle vGPU profile request - create mdev device if req.GPU != nil && req.GPU.Profile != "" { - log.InfoContext(ctx, "creating vGPU", "instance_id", id, "profile", req.GPU.Profile) - gpuDevice, err = devices.CreateVGPU(ctx, req.GPU.Profile, id) + log.InfoContext(ctx, "creating vGPU mdev", "instance_id", id, "profile", req.GPU.Profile) + mdev, err := devices.CreateMdev(ctx, req.GPU.Profile, id) if err != nil { - log.ErrorContext(ctx, "failed to create vGPU", "profile", req.GPU.Profile, "error", err) - return nil, wrapCreateVGPUErr(req.GPU.Profile, err) + log.ErrorContext(ctx, "failed to create mdev", "profile", req.GPU.Profile, "error", err) + return nil, wrapCreateMdevErr(req.GPU.Profile, err) } - gpuProfile = gpuDevice.ProfileName - gpuFramework = gpuDevice.Framework - gpuDevicePath = gpuDevice.SysfsPath - gpuMdevUUID = gpuDevice.MdevUUID - log.InfoContext(ctx, "created vGPU", "instance_id", id, "profile", gpuProfile, "uuid", gpuMdevUUID) + gpuProfile = req.GPU.Profile + gpuMdevUUID = mdev.UUID + log.InfoContext(ctx, "created vGPU mdev", "instance_id", id, "profile", gpuProfile, "uuid", gpuMdevUUID) - // Add vGPU cleanup to stack + // Add mdev cleanup to stack cu.Add(func() { - log.DebugContext(ctx, "destroying vGPU on cleanup", "instance_id", id, "uuid", gpuDevice.MdevUUID) - assignment := devices.VGPUAssignment{ - Framework: gpuDevice.Framework, - DevicePath: gpuDevice.SysfsPath, - MdevUUID: gpuDevice.MdevUUID, - } - if err := devices.DestroyVGPU(ctx, assignment); err != nil { - log.WarnContext(ctx, "failed to destroy vGPU on cleanup", "instance_id", id, "uuid", gpuDevice.MdevUUID, "error", err) + log.DebugContext(ctx, "destroying mdev on cleanup", "instance_id", id, "uuid", gpuMdevUUID) + if err := devices.DestroyMdev(ctx, gpuMdevUUID); err != nil { + log.WarnContext(ctx, "failed to destroy mdev on cleanup", "instance_id", id, "uuid", gpuMdevUUID, "error", err) } }) } @@ -392,8 +382,6 @@ func (m *manager) createInstance( VsockSocket: vsockSocket, Devices: resolvedDeviceIDs, GPUProfile: gpuProfile, - GPUFramework: gpuFramework, - GPUDevicePath: gpuDevicePath, GPUMdevUUID: gpuMdevUUID, Entrypoint: req.Entrypoint, Cmd: req.Cmd, @@ -830,8 +818,10 @@ func (m *manager) startAndBootVM( if err != nil { return fmt.Errorf("start vm: %w", err) } - // Store the PID identity for later cleanup. - pid = resolveRuntimeHypervisorPID(log, stored, pid) + pid = resolveRuntimeHypervisorPID(log, stored.SocketPath, pid) + + // Store the PID for later cleanup + stored.HypervisorPID = &pid log.DebugContext(ctx, "VM started", "instance_id", stored.Id, "pid", pid) // Optional: Expand memory to max if hotplug configured @@ -847,26 +837,15 @@ func (m *manager) startAndBootVM( return nil } -// resolveRuntimeHypervisorPID resolves the runtime PID of the hypervisor -// serving the instance socket and records its process identity. The -// boot-scoped identity token is minted only for a trustworthy PID — the -// direct child we spawned or the confirmed socket owner. -func resolveRuntimeHypervisorPID(log *slog.Logger, stored *StoredMetadata, fallbackPID int) int { - if ProcessExists(fallbackPID) { - stored.HypervisorProcessIdentity.Set(fallbackPID) +func resolveRuntimeHypervisorPID(log *slog.Logger, socketPath string, fallbackPID int) int { + if processExists(fallbackPID) { return fallbackPID } - pid, err := hypervisor.ResolveProcessPID(stored.SocketPath) + pid, err := hypervisor.ResolveProcessPID(socketPath) if err != nil { - // The fallback PID was just proven dead, so it gets no identity - // token: minting one would stamp the current boot ID (and, if the - // PID is recycled mid-call, a live start time) onto a process that - // is not the hypervisor. - log.Debug("using fallback hypervisor pid", "socket_path", stored.SocketPath, "pid", fallbackPID, "error", err) - stored.HypervisorProcessIdentity.SetUnconfirmed(fallbackPID) + log.Debug("using fallback hypervisor pid", "socket_path", socketPath, "pid", fallbackPID, "error", err) return fallbackPID } - stored.HypervisorProcessIdentity.Set(pid) return pid } @@ -964,6 +943,12 @@ func (m *manager) buildHypervisorConfig(ctx context.Context, inst *Instance, ima } } + // Add vGPU mdev device if configured + if inst.GPUMdevUUID != "" { + mdevPath := filepath.Join("/sys/bus/mdev/devices", inst.GPUMdevUUID) + pciDevices = append(pciDevices, mdevPath) + } + // Build topology if available var topology *hypervisor.CPUTopology if hostTopo := calculateGuestTopology(inst.Vcpus, m.hostTopology); hostTopo != nil { @@ -983,22 +968,21 @@ func (m *manager) buildHypervisorConfig(ctx context.Context, inst *Instance, ima } return hypervisor.VMConfig{ - VCPUs: inst.Vcpus, - MemoryBytes: inst.Size, - HotplugBytes: inst.HotplugSize, - Topology: topology, - GuestMemory: m.guestMemoryConfig(), - Disks: disks, - Networks: networks, - SerialLogPath: m.paths.InstanceAppLog(inst.Id), - VsockCID: inst.VsockCID, - VsockSocket: inst.VsockSocket, - PCIDevices: pciDevices, - VGPUDevicePath: storedVGPUDevicePath(&inst.StoredMetadata), - KernelPath: kernelPath, - InitrdPath: initrdPath, - KernelArgs: m.kernelArgs(inst.HypervisorType), - EnableRosetta: inst.EnableRosetta, + VCPUs: inst.Vcpus, + MemoryBytes: inst.Size, + HotplugBytes: inst.HotplugSize, + Topology: topology, + GuestMemory: m.guestMemoryConfig(), + Disks: disks, + Networks: networks, + SerialLogPath: m.paths.InstanceAppLog(inst.Id), + VsockCID: inst.VsockCID, + VsockSocket: inst.VsockSocket, + PCIDevices: pciDevices, + KernelPath: kernelPath, + InitrdPath: initrdPath, + KernelArgs: m.kernelArgs(inst.HypervisorType), + EnableRosetta: inst.EnableRosetta, }, nil } diff --git a/lib/instances/create_mdev_test.go b/lib/instances/create_mdev_test.go index e6e4f55c5..db546c153 100644 --- a/lib/instances/create_mdev_test.go +++ b/lib/instances/create_mdev_test.go @@ -45,7 +45,7 @@ func TestCreateInstanceRejectsUnsupportedVGPUBeforeResourceReservation(t *testin assert.Zero(t, validator.reserveCalls) } -func TestWrapCreateVGPUErr(t *testing.T) { +func TestWrapCreateMdevErr(t *testing.T) { t.Parallel() for _, tc := range []struct { @@ -61,7 +61,7 @@ func TestWrapCreateVGPUErr(t *testing.T) { wantInvalidRequest: true, }, { - name: "other vGPU error", + name: "other mdev error", err: errors.New("boom"), wantMessage: "create vGPU mdev for profile profile: boom", }, @@ -69,7 +69,7 @@ func TestWrapCreateVGPUErr(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - err := wrapCreateVGPUErr("profile", tc.err) + err := wrapCreateMdevErr("profile", tc.err) assert.ErrorIs(t, err, tc.err) if tc.wantInvalidRequest { diff --git a/lib/instances/delete.go b/lib/instances/delete.go index e897ff049..412cf2736 100644 --- a/lib/instances/delete.go +++ b/lib/instances/delete.go @@ -7,6 +7,7 @@ import ( "syscall" "time" + "github.com/kernel/hypeman/lib/devices" "github.com/kernel/hypeman/lib/guest" "github.com/kernel/hypeman/lib/hypervisor" "github.com/kernel/hypeman/lib/logger" @@ -83,21 +84,6 @@ func (m *manager) deleteInstanceWithOptions( guest.CloseConn(dialer.Key()) } - // 3b. Block the restart policy before any teardown. If the delete fails - // partway (e.g. the hypervisor cannot be confirmed dead) the metadata is - // retained with the VMM already stopped, and without this marker the - // restart policy controller would start the instance again. - if err := m.markRestartManualStopLocked(ctx, id); err != nil { - return fmt.Errorf("block restart policy before delete: %w", err) - } - // markRestartManualStopLocked persists through a separate metadata load. - // Reload it so later saves in this delete do not overwrite the block. - meta, err = m.loadMetadata(id) - if err != nil { - return fmt.Errorf("reload metadata after blocking restart policy: %w", err) - } - stored = &meta.StoredMetadata - // 4. If active, try graceful guest shutdown before force kill. gracefulShutdown := false if !options.SkipGracefulShutdown && (inst.State == StateRunning || inst.State == StateInitializing) { @@ -131,34 +117,13 @@ func (m *manager) deleteInstanceWithOptions( err := m.killHypervisor(killCtx, &inst) killSpanEnd(err) if err != nil { - // The hypervisor may still be running, so tearing down its vGPU, - // network, and devices is unsafe. The restart policy is already - // blocked and the metadata is retained, so a retried delete is safe. - log.ErrorContext(ctx, "failed to kill hypervisor; retaining instance metadata", "instance_id", id, "error", err) - return fmt.Errorf("kill hypervisor: %w", err) + // Log error but continue with cleanup + // Best effort to clean up even if hypervisor is unresponsive + log.WarnContext(ctx, "failed to kill hypervisor, continuing with cleanup", "instance_id", id, "error", err) } } m.closeFirecrackerUFFDSession(ctx, stored) - // 5b. Release the vGPU assignment if present, before any network, device, - // or volume teardown. Release failure is logged and the delete continues, - // matching the pre-refactor contract: the VMM is already confirmed dead, - // the guards inside the release never destroy a device they cannot prove - // is unowned, and a skipped release is recovered by startup - // reconciliation. - hadVGPUAssignment := storedVGPUDevicePath(stored) != "" - if hadVGPUAssignment { - log.InfoContext(ctx, "destroying vGPU", "instance_id", id, "uuid", stored.GPUMdevUUID) - } - if err := releaseStoredVGPU(ctx, stored); err != nil { - // Log error but continue with cleanup. - log.WarnContext(ctx, "failed to destroy vGPU, continuing with cleanup", "instance_id", id, "uuid", stored.GPUMdevUUID, "error", err) - } else if hadVGPUAssignment { - if err := m.saveMetadata(meta); err != nil { - log.WarnContext(ctx, "failed to save metadata after vGPU release", "instance_id", id, "error", err) - } - } - // 6. Release network allocation if inst.NetworkEnabled { m.unregisterEgressProxyInstance(ctx, id) @@ -204,6 +169,15 @@ func (m *manager) deleteInstanceWithOptions( } } + // 7c. Destroy vGPU mdev device if present + if inst.GPUMdevUUID != "" { + log.InfoContext(ctx, "destroying vGPU mdev", "instance_id", id, "uuid", inst.GPUMdevUUID) + if err := devices.DestroyMdev(ctx, inst.GPUMdevUUID); err != nil { + // Log error but continue with cleanup + log.WarnContext(ctx, "failed to destroy mdev, continuing with cleanup", "instance_id", id, "uuid", inst.GPUMdevUUID, "error", err) + } + } + // 8. Delete all instance data log.DebugContext(ctx, "deleting instance data", "instance_id", id) _, dataSpanEnd := m.startLifecycleStep(ctx, "delete_instance_data", @@ -221,31 +195,46 @@ func (m *manager) deleteInstanceWithOptions( return nil } -// killHypervisor force kills the hypervisor process without graceful shutdown. -// Used by delete and as stop's final fallback after graceful shutdown fails. -// It returns an error when the hypervisor may still be running: neither process -// identity nor socket ownership can be confirmed, SIGKILL fails with an error -// other than ESRCH, or the process does not exit after SIGKILL. Callers must not -// tear down instance resources in that case. +// killHypervisor force kills the hypervisor process without graceful shutdown +// Used only for delete operations where we're removing all data anyway. +// For operations that need graceful shutdown (like standby), use the hypervisor API directly. func (m *manager) killHypervisor(ctx context.Context, inst *Instance) error { log := logger.FromContext(ctx) - pid, err := resolveLiveHypervisorPID(inst.HypervisorProcessIdentity, inst.SocketPath) - if err != nil { - return err - } - if pid > 0 { - if inst.HypervisorPID != nil && pid != *inst.HypervisorPID { - log.WarnContext(ctx, "stored hypervisor PID does not own the instance socket, killing the socket owner", - "instance_id", inst.Id, "stored_pid", *inst.HypervisorPID, "owner_pid", pid) - } - log.DebugContext(ctx, "killing hypervisor process", "instance_id", inst.Id, "pid", pid) - if err := killProcessAndWait(pid); err != nil { - return err + // If we have a PID, kill the process immediately + if inst.HypervisorPID != nil { + pid := *inst.HypervisorPID + + // Check if process exists + if err := syscall.Kill(pid, 0); err == nil { + // Process exists - kill it immediately with SIGKILL + // No graceful shutdown needed since we're deleting all data + log.DebugContext(ctx, "killing hypervisor process", "instance_id", inst.Id, "pid", pid) + if err := syscall.Kill(pid, syscall.SIGKILL); err != nil { + log.WarnContext(ctx, "failed to kill hypervisor process", "instance_id", inst.Id, "pid", pid, "error", err) + } + + // Wait for process to die and reap it to prevent zombies + // SIGKILL should be instant, but give it a moment + for i := 0; i < 50; i++ { // 50 * 100ms = 5 seconds + var wstatus syscall.WaitStatus + wpid, err := syscall.Wait4(pid, &wstatus, syscall.WNOHANG, nil) + if err != nil || wpid == pid { + // Process reaped successfully or error (likely ECHILD if already reaped) + log.DebugContext(ctx, "hypervisor process killed and reaped", "instance_id", inst.Id, "pid", pid) + break + } + if i == 49 { + log.WarnContext(ctx, "hypervisor process did not exit in time", "instance_id", inst.Id, "pid", pid) + } + time.Sleep(100 * time.Millisecond) + } + } else { + log.DebugContext(ctx, "hypervisor process not running", "instance_id", inst.Id, "pid", pid) } } - // The hypervisor is confirmed gone; remove its stale socket. + // Clean up socket if it still exists os.Remove(inst.SocketPath) return nil @@ -271,12 +260,12 @@ func WaitForProcessExit(pid int, timeout time.Duration) bool { // Process still running (or wait status not yet available). case waitErr == syscall.ECHILD: // Not our child (or already reaped elsewhere). Fall back to existence check. - if !ProcessExists(pid) { + if err := syscall.Kill(pid, 0); err != nil { return true } default: // Best effort fallback on transient/unexpected wait errors. - if !ProcessExists(pid) { + if err := syscall.Kill(pid, 0); err != nil { return true } } diff --git a/lib/instances/delete_test.go b/lib/instances/delete_test.go index dc03b5d26..0ed8efb0b 100644 --- a/lib/instances/delete_test.go +++ b/lib/instances/delete_test.go @@ -2,7 +2,6 @@ package instances import ( "os/exec" - "syscall" "testing" "time" @@ -23,15 +22,6 @@ func TestWaitForProcessExit_ReapsZombieChild(t *testing.T) { assert.Less(t, elapsed, 250*time.Millisecond, "reaping should be quick") } -func TestWaitForProcessExit_EPERMProcessIsAlive(t *testing.T) { - t.Parallel() - if syscall.Kill(1, 0) == nil { - t.Skip("running as root") - } - - assert.False(t, WaitForProcessExit(1, 100*time.Millisecond)) -} - func TestWaitForProcessExit_TimesOutForRunningProcess(t *testing.T) { t.Parallel() cmd := exec.Command("sleep", "2") diff --git a/lib/instances/fork.go b/lib/instances/fork.go index e0c778860..d65d5131b 100644 --- a/lib/instances/fork.go +++ b/lib/instances/fork.go @@ -281,7 +281,7 @@ func (m *manager) forkInstanceFromStoppedOrStandby(ctx context.Context, id strin forkMeta.ExpiresAt = nil forkMeta.StartedAt = nil forkMeta.StoppedAt = nil - forkMeta.HypervisorProcessIdentity.Clear() + forkMeta.HypervisorPID = nil forkMeta.SocketPath = m.paths.InstanceSocket(forkID, starter.SocketName()) forkMeta.DataDir = dstDir forkMeta.VsockSocket = m.paths.InstanceSocket(forkID, hypervisor.VsockSocketNameForType(forkMeta.HypervisorType)) @@ -299,11 +299,6 @@ func (m *manager) forkInstanceFromStoppedOrStandby(ctx context.Context, id strin // phase (Standby for snapshot forks, Stopped for stopped forks) will be // recorded by the appropriate operation when the fork is acted on. forkMeta.Phases.Reset() - // A vGPU assignment is never shared with a fork: normally stop already - // released it, and an assignment retained by a failed release must stay - // with the source so only one instance retries it. The fork acquires its - // own vGPU on start from GPUProfile. - clearStoredVGPUDevice(&forkMeta) switch source.State { case StateStandby: forkMeta.Phases.Record(phasetracking.PhaseStandby, now) diff --git a/lib/instances/fork_test.go b/lib/instances/fork_test.go index 26bc6cfcb..32dc06ee6 100644 --- a/lib/instances/fork_test.go +++ b/lib/instances/fork_test.go @@ -17,7 +17,6 @@ import ( "time" "github.com/kernel/hypeman/lib/autostandby" - "github.com/kernel/hypeman/lib/devices" "github.com/kernel/hypeman/lib/guest" "github.com/kernel/hypeman/lib/healthcheck" "github.com/kernel/hypeman/lib/hypervisor" @@ -30,39 +29,6 @@ import ( "github.com/stretchr/testify/require" ) -func TestForkInstanceClearsVGPUAssignment(t *testing.T) { - manager, _ := setupTestManager(t) - ctx := context.Background() - hvType := hypervisor.Type("fork-vgpu-test") - hypervisor.RegisterCapabilities(hvType, hypervisor.Capabilities{SupportsConcurrentForkPrepare: true}) - manager.vmStarters[hvType] = concurrentForkPrepareTestStarter{} - - sourceID := "fork-vgpu-source" - createStoppedSnapshotSourceFixture(t, manager, sourceID, sourceID, hvType) - - // A retained assignment (release failed during stop) must stay with the - // source; the fork keeps only the profile and acquires its own vGPU on - // start. - meta, err := manager.loadMetadata(sourceID) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - meta.GPUMdevUUID = "retained-uuid" - require.NoError(t, manager.saveMetadata(meta)) - - forked, err := manager.ForkInstance(ctx, sourceID, ForkInstanceRequest{Name: "fork-vgpu-copy"}) - require.NoError(t, err) - assert.Equal(t, "NVIDIA L40S-2Q", forked.GPUProfile) - assert.Equal(t, devices.VGPUFrameworkNone, forked.GPUFramework) - assert.Empty(t, forked.GPUDevicePath) - assert.Empty(t, forked.GPUMdevUUID) - - source, err := manager.loadMetadata(sourceID) - require.NoError(t, err) - assert.Equal(t, "/sys/bus/pci/devices/0000:82:00.4", source.GPUDevicePath) -} - func TestForkInstance_VZStoppedSourceSupported(t *testing.T) { t.Parallel() manager, _ := setupTestManager(t) @@ -719,20 +685,20 @@ func TestCloneStoredMetadataForFork_DeepCopiesReferenceFields(t *testing.T) { pendingLevel := 3 src := StoredMetadata{ - Image: "docker.io/library/alpine:3.19", - ResolvedImage: "docker.io/library/alpine@sha256:amd64digest", - Platform: "linux/amd64", - Env: map[string]string{"A": "1"}, - Tags: map[string]string{"m": "x"}, - Volumes: []VolumeAttachment{{VolumeID: "vol-1", MountPath: "/data"}}, - Devices: []string{"0000:01:00.0"}, - Entrypoint: []string{"/bin/sh", "-c"}, - Cmd: []string{"echo", "hello"}, - ExpiresAt: &expiresAt, - StartedAt: &startedAt, - StoppedAt: &stoppedAt, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &pid}, - ExitCode: &exitCode, + Image: "docker.io/library/alpine:3.19", + ResolvedImage: "docker.io/library/alpine@sha256:amd64digest", + Platform: "linux/amd64", + Env: map[string]string{"A": "1"}, + Tags: map[string]string{"m": "x"}, + Volumes: []VolumeAttachment{{VolumeID: "vol-1", MountPath: "/data"}}, + Devices: []string{"0000:01:00.0"}, + Entrypoint: []string{"/bin/sh", "-c"}, + Cmd: []string{"echo", "hello"}, + ExpiresAt: &expiresAt, + StartedAt: &startedAt, + StoppedAt: &stoppedAt, + HypervisorPID: &pid, + ExitCode: &exitCode, AutoStandby: &autostandby.Policy{ Enabled: true, IdleTimeout: "5m", diff --git a/lib/instances/guestmemory_linux_test.go b/lib/instances/guestmemory_linux_test.go index 87a40992e..224a74cbd 100644 --- a/lib/instances/guestmemory_linux_test.go +++ b/lib/instances/guestmemory_linux_test.go @@ -211,7 +211,7 @@ func requireHypervisorPID(t *testing.T, ctx context.Context, mgr *manager, insta t.Helper() inst, err := mgr.GetInstance(ctx, instanceID) require.NoError(t, err) - if inst.HypervisorPID != nil && ProcessExists(*inst.HypervisorPID) { + if inst.HypervisorPID != nil && processExists(*inst.HypervisorPID) { return *inst.HypervisorPID } if pid, err := hypervisor.ResolveProcessPID(inst.SocketPath); err == nil { diff --git a/lib/instances/lifecycle_noop_test.go b/lib/instances/lifecycle_noop_test.go index a9b918dfb..5ca7515fa 100644 --- a/lib/instances/lifecycle_noop_test.go +++ b/lib/instances/lifecycle_noop_test.go @@ -9,10 +9,8 @@ import ( "testing" "time" - "github.com/kernel/hypeman/lib/devices" "github.com/kernel/hypeman/lib/hypervisor" "github.com/kernel/hypeman/lib/paths" - restartpolicy "github.com/kernel/hypeman/lib/restart-policy" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -149,157 +147,6 @@ func TestLifecycleNoopStandbyWithOptionsStillRejectsStandbyInstance(t *testing.T assertNoLifecycleEvent(t, events) } -func TestDeleteContinuesWhenVGPUReleaseFails(t *testing.T) { - m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) - meta, err := m.loadMetadata(id) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - require.NoError(t, m.saveMetadata(meta)) - - // A failed release is logged and the delete continues, matching the - // pre-refactor contract; the leaked assignment is recovered by startup - // reconciliation. - require.NoError(t, m.DeleteInstance(context.Background(), id)) - - _, err = m.loadMetadata(id) - require.Error(t, err, "instance data must be deleted despite the failed release") -} - -func TestDeletePersistsVGPUReleaseBeforeTeardown(t *testing.T) { - m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) - var persisted *metadata - deviceManager := &recordingDeviceManager{ - onMarkDetached: func() { - var err error - persisted, err = m.loadMetadata(id) - require.NoError(t, err) - }, - } - m.deviceManager = deviceManager - meta, err := m.loadMetadata(id) - require.NoError(t, err) - meta.RestartPolicy = &restartpolicy.Policy{Policy: restartpolicy.PolicyAlways} - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUDevicePath = "/sys/bus/mdev/devices/test-mdev" - meta.GPUMdevUUID = "test-mdev" - meta.Devices = []string{"dev-1"} - require.NoError(t, m.saveMetadata(meta)) - - require.NoError(t, m.DeleteInstance(context.Background(), id)) - require.NotNil(t, persisted) - assert.Empty(t, persisted.GPUDevicePath) - assert.Empty(t, persisted.GPUMdevUUID) - assert.Equal(t, "NVIDIA L40S-2Q", persisted.GPUProfile) - assert.Equal(t, restartpolicy.BlockedReasonManualStop, persisted.RestartStatus.BlockedReason) -} - -func TestDeleteContinuesTeardownAfterFailedVGPURelease(t *testing.T) { - m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) - deviceManager := &recordingDeviceManager{} - m.deviceManager = deviceManager - meta, err := m.loadMetadata(id) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - meta.Devices = []string{"dev-1"} - require.NoError(t, m.saveMetadata(meta)) - - // The failed release must not block the rest of the teardown: devices - // are detached and the instance is fully deleted. - require.NoError(t, m.DeleteInstance(context.Background(), id)) - assert.Equal(t, []string{"dev-1"}, deviceManager.detached) - - _, err = m.loadMetadata(id) - require.Error(t, err, "instance data must be deleted despite the failed release") -} - -// A stale release during start must be persisted immediately: if start fails -// later (here at vGPU recreation on a host without VFs), the on-disk metadata -// must no longer point at the already-released device. -func TestStartPersistsStaleVGPUReleaseImmediately(t *testing.T) { - m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) - m.imageManager = readyFixtureImageManager{name: "test-image"} - meta, err := m.loadMetadata(id) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.HypervisorType = hypervisor.TypeQEMU - meta.GPUFramework = devices.VGPUFrameworkNone - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - require.NoError(t, m.saveMetadata(meta)) - - _, err = m.StartInstance(context.Background(), id, StartInstanceRequest{}) - require.Error(t, err) - - stored, err := m.loadMetadata(id) - require.NoError(t, err) - assert.Empty(t, stored.GPUDevicePath, "released assignment should be persisted despite the failed start") - assert.Equal(t, "NVIDIA L40S-2Q", stored.GPUProfile, "profile is kept for the next start") -} - -func TestStopStoppedInstanceReleasesRetainedVGPU(t *testing.T) { - m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) - meta, err := m.loadMetadata(id) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFrameworkNone - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - require.NoError(t, m.saveMetadata(meta)) - - inst, err := m.StopInstance(context.Background(), id) - require.NoError(t, err) - require.NotNil(t, inst) - assert.Equal(t, StateStopped, inst.State) - - stored, err := m.loadMetadata(id) - require.NoError(t, err) - assert.Empty(t, stored.GPUDevicePath) -} - -func TestStopStoppedInstanceVGPUReleaseFailureRemainsNoop(t *testing.T) { - m, id := newLifecycleNoopManagerWithInstance(t, StateStopped, time.Now().UTC()) - meta, err := m.loadMetadata(id) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - require.NoError(t, m.saveMetadata(meta)) - - inst, err := m.StopInstance(context.Background(), id) - require.NoError(t, err) - require.NotNil(t, inst) - assert.Equal(t, StateStopped, inst.State) - - stored, err := m.loadMetadata(id) - require.NoError(t, err) - assert.Equal(t, devices.VGPUFramework("future-framework"), stored.GPUFramework) - assert.Equal(t, "/sys/bus/pci/devices/0000:82:00.4", stored.GPUDevicePath) -} - -// recordingDeviceManager is a devices.Manager stub that records passthrough -// teardown calls. Only the methods delete exercises are implemented. -type recordingDeviceManager struct { - devices.Manager - detached []string - unbound []string - onMarkDetached func() -} - -func (m *recordingDeviceManager) MarkDetached(ctx context.Context, deviceID string) error { - m.detached = append(m.detached, deviceID) - if m.onMarkDetached != nil { - m.onMarkDetached() - } - return nil -} - -func (m *recordingDeviceManager) UnbindFromVFIO(ctx context.Context, id string) error { - m.unbound = append(m.unbound, id) - return nil -} - func newLifecycleNoopManagerWithInstance(t *testing.T, state State, now time.Time) (*manager, string) { t.Helper() diff --git a/lib/instances/manager.go b/lib/instances/manager.go index bc23fdf74..2dd027617 100644 --- a/lib/instances/manager.go +++ b/lib/instances/manager.go @@ -195,7 +195,6 @@ type manager struct { nativeCodecPaths map[string]string imageUsageRecorder ImageUsageRecorder guestAgentReadyProbe func(context.Context, *StoredMetadata) bool - shutdownGuestFn func(context.Context, hypervisor.VsockDialer, int32) error // Shared lifecycle event subscriptions for internal consumers. lifecycleEvents *lifecycleSubscribers @@ -648,12 +647,6 @@ func (m *manager) StopInstance(ctx context.Context, id string) (*Instance, error if err := m.markRestartManualStopLocked(ctx, id); err != nil { return nil, err } - // A stopped instance can retain a vGPU assignment when the release - // failed during the original stop. Retry it here so the vGPU slot is - // not held until the next start, delete, or hypeman restart. A failed - // retry only logs, keeping stop's no-op contract for already-stopped - // instances. - m.releaseRetainedVGPULocked(ctx, id) updated, err := m.currentInstanceWithoutHydration(ctx, id) if err != nil { return nil, err diff --git a/lib/instances/manager_test.go b/lib/instances/manager_test.go index 5029a5fb7..d9b405d35 100644 --- a/lib/instances/manager_test.go +++ b/lib/instances/manager_test.go @@ -4,7 +4,6 @@ import ( "bytes" "context" "crypto/tls" - "errors" "fmt" "io" "net" @@ -165,26 +164,6 @@ func waitForInstanceState(ctx context.Context, mgr Manager, instanceID string, e return nil, fmt.Errorf("instance %s did not reach %s within %v (last state: %s)", instanceID, expected, timeout, lastState) } -// deleteInstanceEventually deletes an instance, retrying while the hypervisor -// finishes dying. Delete fails closed when the VMM has not exited within its -// short post-SIGKILL wait; on loaded CI hosts kernel-side teardown can outlast -// that wait, and the contract is that a retried delete converges. -func deleteInstanceEventually(t *testing.T, ctx context.Context, mgr Manager, instanceID string) { - t.Helper() - deadline := time.Now().Add(integrationTestTimeout(30 * time.Second)) - for { - err := mgr.DeleteInstance(ctx, instanceID) - if err == nil || errors.Is(err, ErrNotFound) { - return - } - if time.Now().After(deadline) { - t.Fatalf("delete instance %s did not converge: %v", instanceID, err) - } - t.Logf("delete instance %s not yet converged, retrying: %v", instanceID, err) - time.Sleep(time.Second) - } -} - func integrationTestTimeout(timeout time.Duration) time.Duration { if os.Getenv("CI") == "true" && timeout < 45*time.Second { return 45 * time.Second @@ -1667,7 +1646,8 @@ func TestStandbyAndRestore(t *testing.T) { // Cleanup (no sleep needed - DeleteInstance handles process cleanup) t.Log("Cleaning up...") - deleteInstanceEventually(t, ctx, manager, inst.Id) + err = manager.DeleteInstance(ctx, inst.Id) + require.NoError(t, err) t.Log("Standby/restore test complete!") } diff --git a/lib/instances/network_test.go b/lib/instances/network_test.go index 67cf7b0f9..0c25455f0 100644 --- a/lib/instances/network_test.go +++ b/lib/instances/network_test.go @@ -296,7 +296,8 @@ func TestCreateInstanceWithNetwork(t *testing.T) { // Cleanup t.Log("Cleaning up instance...") - deleteInstanceEventually(t, ctx, manager, inst.Id) + err = manager.DeleteInstance(ctx, inst.Id) + require.NoError(t, err) // Verify TAP deleted after instance cleanup t.Log("Verifying TAP deleted after cleanup...") diff --git a/lib/instances/process_identity.go b/lib/instances/process_identity.go deleted file mode 100644 index 1d885d0b5..000000000 --- a/lib/instances/process_identity.go +++ /dev/null @@ -1,257 +0,0 @@ -package instances - -import ( - "errors" - "fmt" - "os" - "path/filepath" - "runtime" - "strconv" - "strings" - "sync" - "syscall" - "time" - - "github.com/kernel/hypeman/lib/hypervisor" -) - -// linuxBootIDPath is the kernel-provided boot ID used to scope process -// identities to a single host boot. -const linuxBootIDPath = "/proc/sys/kernel/random/boot_id" - -// hypervisorSIGKILLWaitTimeout bounds how long stop and delete wait for the -// hypervisor to exit after SIGKILL before reporting it still alive. A process -// that survives SIGKILL is stuck in uninterruptible sleep, and waiting longer -// does not unstick it, so the wait is short to keep stop and delete fast. -const hypervisorSIGKILLWaitTimeout = 2 * time.Second - -// killProcessAndWait SIGKILLs pid and waits for it to exit. A process that -// survives the first wait gets its process group killed too (the hypervisor -// may have spawned children in its own group) and a short grace period. An -// error means the process may still be running, so callers must not tear down -// instance resources. -func killProcessAndWait(pid int) error { - if err := syscall.Kill(pid, syscall.SIGKILL); err != nil { - if err == syscall.ESRCH { - return nil - } - return fmt.Errorf("kill hypervisor process %d: %w", pid, err) - } - if WaitForProcessExit(pid, hypervisorSIGKILLWaitTimeout) { - return nil - } - // The process may have spawned children in its own process group. - _ = syscall.Kill(-pid, syscall.SIGKILL) - if !WaitForProcessExit(pid, hypervisorSIGKILLWaitTimeout) { - return fmt.Errorf("hypervisor pid %d did not exit after SIGKILL", pid) - } - return nil -} - -// HypervisorProcessIdentity identifies a specific hypervisor process across -// PID reuse and host reboots. The zero value means no recorded identity. -// It is embedded anonymously in StoredMetadata so the persisted JSON keys -// (and on-disk metadata format) are unchanged. -type HypervisorProcessIdentity struct { - HypervisorPID *int // Hypervisor process ID (may be stale after host restart) - HypervisorStartTime uint64 // Start time of HypervisorPID from /proc//stat (clock ticks since boot). 0 = unknown. - HypervisorBootID string // Linux boot ID recorded with HypervisorStartTime; scopes the process identity across host reboots. -} - -// Set records pid's boot-scoped identity as the instance's hypervisor. -func (h *HypervisorProcessIdentity) Set(pid int) { - h.HypervisorPID = &pid - h.HypervisorStartTime = processStartTime(pid) - h.HypervisorBootID = hostBootID() -} - -// SetUnconfirmed records a bare PID without the boot-scoped identity token. -// Used when the PID was just proven dead: minting a token would stamp the -// current boot ID (and, if the PID is recycled mid-call, a live start time) -// onto a process that is not the hypervisor. Destructive paths must confirm -// socket ownership before trusting an unconfirmed PID. -func (h *HypervisorProcessIdentity) SetUnconfirmed(pid int) { - h.HypervisorPID = &pid - h.HypervisorStartTime = 0 - h.HypervisorBootID = "" -} - -// Clear erases the recorded identity. Call whenever the recorded hypervisor -// is known gone or a snapshot/fork must not inherit it. -func (h *HypervisorProcessIdentity) Clear() { - *h = HypervisorProcessIdentity{} -} - -// refreshHypervisorPID refreshes the stored PID for display and other -// non-destructive callers. It trusts a live stored PID without confirming -// socket ownership: hydration runs on every list/get, and its answer never -// authorizes teardown — stop, delete, standby, and the vGPU release guards -// all re-resolve identity through resolveLiveHypervisorPID before acting. -func refreshHypervisorPID(stored *StoredMetadata, state State) { - if !state.RequiresVMM() && state != StateUnknown { - return - } - if stored.HypervisorPID != nil && ProcessExists(*stored.HypervisorPID) { - return - } - if stored.SocketPath == "" { - return - } - pid, err := hypervisor.ResolveProcessPID(stored.SocketPath) - if err != nil { - return - } - stored.HypervisorProcessIdentity.Set(pid) -} - -// resolveLiveHypervisorPID returns the PID of the live hypervisor that owns -// the instance socket, or 0 when no live hypervisor is found. A live stored PID -// whose recorded boot ID and start time match is returned without socket -// confirmation. It returns an error when socket ownership cannot be confirmed -// because the socket scan itself failed. -func resolveLiveHypervisorPID(id HypervisorProcessIdentity, socketPath string) (int, error) { - stored := 0 - if id.HypervisorPID != nil { - if ProcessExists(*id.HypervisorPID) { - stored = *id.HypervisorPID - } else { - // ProcessExists treats zombies as dead, so a direct-child VMM that - // exited on its own never reaches the Wait4 in WaitForProcessExit - // and would sit unreaped. Reap it here: WNOHANG leaves a live child - // untouched, and a recycled or non-child PID fails with ECHILD. - var status syscall.WaitStatus - _, _ = syscall.Wait4(*id.HypervisorPID, &status, syscall.WNOHANG, nil) - } - } - if runtime.GOOS != "linux" { - return stored, nil - } - bootID := hostBootID() - if stored != 0 && id.HypervisorBootID != "" && bootID != "" && id.HypervisorBootID != bootID { - // The recorded identity is scoped to a previous host boot, so whatever - // process wears the stored PID now is provably not the recorded - // hypervisor. Treat the stored PID as dead rather than failing closed. - stored = 0 - } - if stored != 0 && id.HypervisorStartTime != 0 && id.HypervisorBootID != "" && bootID != "" && id.HypervisorBootID == bootID { - if processStartTime(stored) == id.HypervisorStartTime { - return stored, nil - } - stored = 0 - } - if socketPath == "" { - if stored != 0 { - return 0, fmt.Errorf("cannot confirm stored hypervisor PID %d without a socket path", stored) - } - return 0, nil - } - var resolved int - var err error - if stored != 0 && id.HypervisorStartTime == 0 { - resolved, err = hypervisor.ResolveProcessPIDForOwner(socketPath, stored) - } else { - resolved, err = hypervisor.ResolveProcessPID(socketPath) - } - return classifyResolvedHypervisorOwner(socketPath, stored, resolved, err) -} - -// classifyResolvedHypervisorOwner interprets a socket resolution result for -// resolveLiveHypervisorPID: which live PID, if any, is the recorded -// hypervisor, and whether the ambiguity must fail closed. -func classifyResolvedHypervisorOwner(socketPath string, stored, resolved int, err error) (int, error) { - switch { - case err == nil && ProcessExists(resolved): - return resolved, nil - case err == nil: - return 0, nil - } - if errors.Is(err, hypervisor.ErrNoOwningProcess) { - // The socket-owner scan found no process holding the listener. A live - // hypervisor always holds its control-socket listener, so the recorded - // hypervisor is gone and a live stored PID is a recycled number — the - // same conclusion already drawn above when the stored PID is dead. - // Without this, legacy metadata carrying no boot-scoped identity - // wedges stop and delete forever once its PID is reused. - return 0, nil - } - if stored != 0 { - return 0, fmt.Errorf("cannot confirm ownership of socket %s for stored hypervisor PID %d: %w", socketPath, stored, err) - } - return 0, fmt.Errorf("cannot confirm ownership of socket %s: %w", socketPath, err) -} - -// ProcessExists reports whether pid belongs to a live, non-zombie process. -func ProcessExists(pid int) bool { - if pid <= 0 { - return false - } - err := syscall.Kill(pid, 0) - if err != nil && err != syscall.EPERM { - return false - } - if runtime.GOOS != "linux" { - return true - } - state, err := readLinuxProcessState(pid) - if err != nil { - return true - } - return state != "Z" -} - -func readLinuxProcessState(pid int) (string, error) { - statusPath := filepath.Join("/proc", strconv.Itoa(pid), "status") - data, err := os.ReadFile(statusPath) - if err != nil { - return "", err - } - for _, line := range strings.Split(string(data), "\n") { - if !strings.HasPrefix(line, "State:") { - continue - } - fields := strings.Fields(line) - if len(fields) < 2 { - return "", fmt.Errorf("malformed process state in %s", statusPath) - } - return fields[1], nil - } - return "", fmt.Errorf("process state missing from %s", statusPath) -} - -var hostBootID = sync.OnceValue(readHostBootID) - -func readHostBootID() string { - if runtime.GOOS != "linux" { - return "" - } - data, err := os.ReadFile(linuxBootIDPath) - if err != nil { - return "" - } - return strings.TrimSpace(string(data)) -} - -// processStartTime returns the start time (field 22 of /proc//stat, clock -// ticks since boot) of pid, or 0 when it cannot be read. -func processStartTime(pid int) uint64 { - if runtime.GOOS != "linux" || pid <= 0 { - return 0 - } - data, err := os.ReadFile(filepath.Join("/proc", strconv.Itoa(pid), "stat")) - if err != nil { - return 0 - } - closingParen := strings.LastIndexByte(string(data), ')') - if closingParen == -1 { - return 0 - } - fields := strings.Fields(string(data[closingParen+1:])) - if len(fields) <= 19 { - return 0 - } - startTime, err := strconv.ParseUint(fields[19], 10, 64) - if err != nil { - return 0 - } - return startTime -} diff --git a/lib/instances/process_identity_linux_test.go b/lib/instances/process_identity_linux_test.go deleted file mode 100644 index 7450ed112..000000000 --- a/lib/instances/process_identity_linux_test.go +++ /dev/null @@ -1,650 +0,0 @@ -//go:build linux - -package instances - -import ( - "bufio" - "context" - "errors" - "fmt" - "io" - "log/slog" - "net" - "os" - "os/exec" - "path/filepath" - "syscall" - "testing" - "time" - - "github.com/kernel/hypeman/lib/hypervisor" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestResolveLiveHypervisorPIDWithoutStoredPID(t *testing.T) { - t.Run("missing socket", func(t *testing.T) { - pid, err := resolveLiveHypervisorPID(HypervisorProcessIdentity{}, filepath.Join(t.TempDir(), "missing.sock")) - require.NoError(t, err) - assert.Zero(t, pid) - }) - - t.Run("live owner", func(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - listener, err := net.Listen("unix", socketPath) - require.NoError(t, err) - defer listener.Close() - - pid, err := resolveLiveHypervisorPID(HypervisorProcessIdentity{}, socketPath) - require.NoError(t, err) - assert.Equal(t, os.Getpid(), pid) - }) -} - -// TestResolveLiveHypervisorPIDReapsZombieChild guards against leaking one -// zombie per direct-child VMM that exits on its own: ProcessExists treats -// zombies as dead, so the confirmed-gone paths in stop, delete, and standby -// never reach the Wait4 in WaitForProcessExit. -func TestResolveLiveHypervisorPIDReapsZombieChild(t *testing.T) { - child := exec.Command("true") - require.NoError(t, child.Start()) - pid := child.Process.Pid - - require.Eventually(t, func() bool { - state, err := readLinuxProcessState(pid) - return err == nil && state == "Z" - }, 5*time.Second, 10*time.Millisecond, "child never became a zombie") - - resolved, err := resolveLiveHypervisorPID(HypervisorProcessIdentity{HypervisorPID: &pid}, filepath.Join(t.TempDir(), "missing.sock")) - require.NoError(t, err) - assert.Zero(t, resolved) - - var status syscall.WaitStatus - _, waitErr := syscall.Wait4(pid, &status, syscall.WNOHANG, nil) - assert.ErrorIs(t, waitErr, syscall.ECHILD, "zombie child was not reaped") -} - -func TestHostBootIDIsStable(t *testing.T) { - first := hostBootID() - require.NotEmpty(t, first) - assert.Equal(t, first, hostBootID()) -} - -func TestProcessStartTime(t *testing.T) { - assert.NotZero(t, processStartTime(os.Getpid())) - assert.Zero(t, processStartTime(0)) - assert.Zero(t, processStartTime(-1)) - - const nonexistentPID = 1<<22 - 1 - require.False(t, ProcessExists(nonexistentPID)) - assert.Zero(t, processStartTime(nonexistentPID)) -} - -func TestResolveLiveHypervisorPIDUsesMatchingStartTime(t *testing.T) { - process := exec.Command("sleep", "30") - require.NoError(t, process.Start()) - t.Cleanup(func() { - _ = process.Process.Kill() - _ = process.Wait() - }) - - pid := process.Process.Pid - startTime := processStartTime(pid) - require.NotZero(t, startTime) - - resolved, err := resolveLiveHypervisorPID(HypervisorProcessIdentity{HypervisorPID: &pid, HypervisorStartTime: startTime, HypervisorBootID: hostBootID()}, "") - require.NoError(t, err) - assert.Equal(t, pid, resolved) -} - -func TestKillHypervisorUsesMatchingStartTimeWhenSocketIsGone(t *testing.T) { - process := exec.Command("sleep", "30") - require.NoError(t, process.Start()) - t.Cleanup(func() { - _ = process.Process.Kill() - _ = process.Wait() - }) - - pid := process.Process.Pid - startTime := processStartTime(pid) - require.NotZero(t, startTime) - socketPath := filepath.Join(t.TempDir(), "missing.sock") - - m := &manager{} - require.NoError(t, m.killHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{ - Id: "kill-test", - HypervisorProcessIdentity: HypervisorProcessIdentity{ - HypervisorPID: &pid, - HypervisorStartTime: startTime, - HypervisorBootID: hostBootID(), - }, - SocketPath: socketPath, - }, - })) - - assert.ErrorIs(t, syscall.Kill(pid, 0), syscall.ESRCH) - _, statErr := os.Stat(socketPath) - assert.True(t, os.IsNotExist(statErr), "instance socket should be removed") -} - -// TestResolveLiveHypervisorPIDTreatsDisprovenIdentityAsDead covers the -// disproof branches: an identity token that cannot belong to the live PID -// holder, plus no socket owner, resolves to "provably dead" (0, nil) so -// stop/delete proceed without signaling the recycled PID. -func TestResolveLiveHypervisorPIDTreatsDisprovenIdentityAsDead(t *testing.T) { - process := exec.Command("sleep", "30") - require.NoError(t, process.Start()) - t.Cleanup(func() { - _ = process.Process.Kill() - _ = process.Wait() - }) - - pid := process.Process.Pid - startTime := processStartTime(pid) - require.NotZero(t, startTime) - - for name, id := range map[string]HypervisorProcessIdentity{ - "different boot": {HypervisorPID: &pid, HypervisorStartTime: startTime, HypervisorBootID: "different-boot"}, - "mismatched start time": {HypervisorPID: &pid, HypervisorStartTime: startTime + 1, HypervisorBootID: hostBootID()}, - } { - t.Run(name, func(t *testing.T) { - resolved, err := resolveLiveHypervisorPID(id, "") - require.NoError(t, err) - assert.Zero(t, resolved) - assert.NoError(t, syscall.Kill(pid, 0), "disproven identity must not signal the PID holder") - }) - } -} - -func TestResolveLiveHypervisorPIDFailsClosedWithoutSocketOrIdentity(t *testing.T) { - pid := os.Getpid() - resolved, err := resolveLiveHypervisorPID(HypervisorProcessIdentity{HypervisorPID: &pid}, "") - require.ErrorContains(t, err, "without a socket path") - assert.Zero(t, resolved) -} - -func TestSocketListenerHelper(t *testing.T) { - if os.Getenv("HYPERVISOR_SOCKET_HELPER") != "1" { - return - } - - listener, err := net.Listen("unix", os.Getenv("HYPERVISOR_SOCKET_PATH")) - if err != nil { - fmt.Fprintln(os.Stderr, err) - os.Exit(1) - } - defer listener.Close() - fmt.Fprintln(os.Stdout, "ready") - _, _ = os.Stdin.Read(make([]byte, 1)) - if os.Getenv("HYPERVISOR_SOCKET_CLOSE_BEFORE_EXIT") == "1" { - _ = listener.Close() - fmt.Fprintln(os.Stdout, "closed") - _, _ = os.Stdin.Read(make([]byte, 1)) - } -} - -func TestKillHypervisorSparesReusedPIDAndKillsSocketOwner(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - owner := exec.Command(os.Args[0], "-test.run=^TestSocketListenerHelper$") - owner.Env = append(os.Environ(), "HYPERVISOR_SOCKET_HELPER=1", "HYPERVISOR_SOCKET_PATH="+socketPath) - stdin, err := owner.StdinPipe() - require.NoError(t, err) - stdout, err := owner.StdoutPipe() - require.NoError(t, err) - require.NoError(t, owner.Start()) - t.Cleanup(func() { - _ = stdin.Close() - _ = owner.Process.Kill() - _ = owner.Wait() - }) - _, err = bufio.NewReader(stdout).ReadString('\n') - require.NoError(t, err) - - stale := exec.Command("sleep", "30") - require.NoError(t, stale.Start()) - t.Cleanup(func() { - _ = stale.Process.Kill() - _ = stale.Wait() - }) - - stalePID := stale.Process.Pid - m := &manager{} - require.NoError(t, m.killHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{Id: "kill-test", HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &stalePID}, SocketPath: socketPath}, - })) - - assert.NoError(t, syscall.Kill(stalePID, 0), "unrelated process holding the stale PID must survive delete") - assert.True(t, WaitForProcessExit(owner.Process.Pid, 5*time.Second), "socket owner should be killed") - _, statErr := os.Stat(socketPath) - assert.True(t, os.IsNotExist(statErr), "instance socket should be removed") -} - -func TestGracefulShutdownWaitsForSocketOwnerInsteadOfExitedStoredPID(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - owner := exec.Command(os.Args[0], "-test.run=^TestSocketListenerHelper$") - owner.Env = append(os.Environ(), "HYPERVISOR_SOCKET_HELPER=1", "HYPERVISOR_SOCKET_PATH="+socketPath) - stdin, err := owner.StdinPipe() - require.NoError(t, err) - stdout, err := owner.StdoutPipe() - require.NoError(t, err) - require.NoError(t, owner.Start()) - t.Cleanup(func() { - _ = stdin.Close() - _ = owner.Process.Kill() - _ = owner.Wait() - }) - _, err = bufio.NewReader(stdout).ReadString('\n') - require.NoError(t, err) - - stale := exec.Command("true") - require.NoError(t, stale.Run()) - stalePID := stale.Process.Pid - inst := &Instance{StoredMetadata: StoredMetadata{ - Id: "graceful-stale-pid", - HypervisorType: hypervisor.TypeCloudHypervisor, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &stalePID}, - SocketPath: socketPath, - VsockSocket: filepath.Join(t.TempDir(), "missing-vsock.sock"), - }} - - m := &manager{} - assert.False(t, m.tryGracefulGuestShutdown(context.Background(), inst, 1), - "stop and delete must fall back to the hardened kill path while the socket owner is alive") - assert.True(t, ProcessExists(owner.Process.Pid)) -} - -func TestGracefulShutdownFallbackKillsCapturedPIDAfterListenerCloses(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - owner := exec.Command(os.Args[0], "-test.run=^TestSocketListenerHelper$") - owner.Env = append(os.Environ(), - "HYPERVISOR_SOCKET_HELPER=1", - "HYPERVISOR_SOCKET_CLOSE_BEFORE_EXIT=1", - "HYPERVISOR_SOCKET_PATH="+socketPath, - ) - stdin, err := owner.StdinPipe() - require.NoError(t, err) - stdout, err := owner.StdoutPipe() - require.NoError(t, err) - require.NoError(t, owner.Start()) - t.Cleanup(func() { - _ = stdin.Close() - _ = owner.Process.Kill() - _ = owner.Wait() - }) - output := bufio.NewReader(stdout) - _, err = output.ReadString('\n') - require.NoError(t, err) - - m := &manager{shutdownGuestFn: func(context.Context, hypervisor.VsockDialer, int32) error { - if _, err := stdin.Write([]byte{1}); err != nil { - return err - } - _, err := output.ReadString('\n') - return err - }} - inst := &Instance{StoredMetadata: StoredMetadata{ - Id: "graceful-listener-close", - HypervisorType: hypervisor.TypeCloudHypervisor, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &owner.Process.Pid}, - SocketPath: socketPath, - VsockSocket: filepath.Join(t.TempDir(), "missing-vsock.sock"), - }} - - assert.False(t, m.tryGracefulGuestShutdown(context.Background(), inst, 0)) - require.NoError(t, m.killHypervisor(context.Background(), inst)) - assert.False(t, ProcessExists(owner.Process.Pid)) -} - -func TestShutdownHypervisorSparesReusedPIDWhenNoProcessOwnsSocket(t *testing.T) { - process := exec.Command("sleep", "30") - require.NoError(t, process.Start()) - t.Cleanup(func() { - _ = process.Process.Kill() - _ = process.Wait() - }) - - // No process owns or references the socket, so the live stored PID is a - // recycled number: shutdown must not signal it. Any returned error comes - // from the unreachable control socket, not from ownership resolution. - pid := process.Process.Pid - m := &manager{} - err := m.shutdownHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{ - Id: "shutdown-reused-pid", - HypervisorType: hypervisor.TypeCloudHypervisor, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &pid}, - SocketPath: filepath.Join(t.TempDir(), "missing.sock"), - }, - }) - require.NotContains(t, fmt.Sprint(err), "confirm hypervisor ownership") - assert.NoError(t, syscall.Kill(pid, 0), "process with a recycled PID must not be killed") -} - -func TestShutdownHypervisorRemovesStaleSocketWhenNoLiveOwner(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "stale.sock") - require.NoError(t, os.WriteFile(socketPath, nil, 0600)) - - // No process owns or references the socket and no client factory exists - // for this hypervisor type: the hypervisor is provably gone, so shutdown - // must report success and remove the stale socket file. - m := &manager{} - require.NoError(t, m.shutdownHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{ - Id: "shutdown-stale-socket", - HypervisorType: hypervisor.Type("unregistered-stale-socket-test"), - SocketPath: socketPath, - }, - })) - _, statErr := os.Stat(socketPath) - assert.True(t, os.IsNotExist(statErr), "stale socket should be removed once the hypervisor is provably gone") -} - -func TestClassifyResolvedHypervisorOwner(t *testing.T) { - const deadPID = 1<<22 - 1 - require.False(t, ProcessExists(deadPID)) - - live := exec.Command("sleep", "30") - require.NoError(t, live.Start()) - t.Cleanup(func() { - _ = live.Process.Kill() - _ = live.Wait() - }) - livePID := live.Process.Pid - - // A resolved owner that exited between the scan and the liveness check is - // the same provable-death conclusion as ErrNoOwningProcess, even with a - // live stored PID. Only a failed scan fails closed. - for _, tc := range []struct { - name string - stored, resolved int - err error - wantPID int - wantErr bool - }{ - {name: "live resolved owner", resolved: livePID, wantPID: livePID}, - {name: "dead resolved owner with live stored PID", stored: livePID, resolved: deadPID}, - {name: "dead resolved owner without stored PID", resolved: deadPID}, - {name: "no owning process with live stored PID", stored: livePID, err: fmt.Errorf("wrapped: %w", hypervisor.ErrNoOwningProcess)}, - {name: "scan failure with stored PID", stored: livePID, err: errors.New("inspect process fds: permission denied"), wantErr: true}, - {name: "scan failure without stored PID", err: errors.New("inspect process fds: permission denied"), wantErr: true}, - } { - t.Run(tc.name, func(t *testing.T) { - pid, err := classifyResolvedHypervisorOwner("/fake.sock", tc.stored, tc.resolved, tc.err) - if tc.wantErr { - require.Error(t, err) - return - } - require.NoError(t, err) - assert.Equal(t, tc.wantPID, pid) - }) - } -} - -func TestShutdownHypervisorKillsResolvedOwnerWhenClientUnavailable(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - owner := exec.Command(os.Args[0], "-test.run=^TestSocketListenerHelper$") - owner.Env = append(os.Environ(), "HYPERVISOR_SOCKET_HELPER=1", "HYPERVISOR_SOCKET_PATH="+socketPath) - stdin, err := owner.StdinPipe() - require.NoError(t, err) - stdout, err := owner.StdoutPipe() - require.NoError(t, err) - require.NoError(t, owner.Start()) - t.Cleanup(func() { - _ = stdin.Close() - _ = owner.Process.Kill() - _ = owner.Wait() - }) - _, err = bufio.NewReader(stdout).ReadString('\n') - require.NoError(t, err) - - // No client factory exists for this hypervisor type, so getHypervisor - // fails while the resolved socket owner is alive: shutdown must kill the - // owner instead of reporting a completed shutdown. - m := &manager{} - require.NoError(t, m.shutdownHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{ - Id: "shutdown-no-client", - HypervisorType: hypervisor.Type("unregistered-shutdown-test"), - SocketPath: socketPath, - }, - })) - assert.True(t, WaitForProcessExit(owner.Process.Pid, 5*time.Second), "resolved socket owner must be killed when the control client is unavailable") - _, statErr := os.Stat(socketPath) - assert.True(t, os.IsNotExist(statErr), "instance socket should be removed") -} - -func TestShutdownHypervisorIgnoresCommandLineBystander(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - require.NoError(t, os.WriteFile(socketPath, nil, 0600)) - match := exec.Command("sh", "-c", "sleep 30", "sh", socketPath) - require.NoError(t, match.Start()) - t.Cleanup(func() { - _ = match.Process.Kill() - _ = match.Wait() - }) - - stale := exec.Command("sleep", "30") - require.NoError(t, stale.Start()) - t.Cleanup(func() { - _ = stale.Process.Kill() - _ = stale.Wait() - }) - - // No process owns the socket listener; a debug client carrying the socket - // path in its command line must not be mistaken for the hypervisor or - // block shutdown. - stalePID := stale.Process.Pid - m := &manager{} - err := m.shutdownHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{ - Id: "shutdown-cmdline-bystander", - HypervisorType: hypervisor.TypeCloudHypervisor, - HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &stalePID}, - SocketPath: socketPath, - }, - }) - require.NotContains(t, fmt.Sprint(err), "confirm hypervisor ownership") - assert.NoError(t, syscall.Kill(stalePID, 0), "process with a recycled PID must not be killed") - assert.NoError(t, syscall.Kill(match.Process.Pid, 0), "command-line bystander must not be signaled") -} - -func TestKillProcessAndWaitIgnoresExitedProcess(t *testing.T) { - process := exec.Command("true") - require.NoError(t, process.Run()) - - require.NoError(t, killProcessAndWait(process.Process.Pid)) -} - -func TestRefreshHypervisorPIDTrustsLiveStoredPID(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - owner := exec.Command(os.Args[0], "-test.run=^TestSocketListenerHelper$") - owner.Env = append(os.Environ(), "HYPERVISOR_SOCKET_HELPER=1", "HYPERVISOR_SOCKET_PATH="+socketPath) - stdin, err := owner.StdinPipe() - require.NoError(t, err) - stdout, err := owner.StdoutPipe() - require.NoError(t, err) - require.NoError(t, owner.Start()) - t.Cleanup(func() { - _ = stdin.Close() - _ = owner.Process.Kill() - _ = owner.Wait() - }) - _, err = bufio.NewReader(stdout).ReadString('\n') - require.NoError(t, err) - - stale := exec.Command("sleep", "30") - require.NoError(t, stale.Start()) - t.Cleanup(func() { - _ = stale.Process.Kill() - _ = stale.Wait() - }) - - // Hydration is display-only and must stay cheap: a live stored PID is - // trusted without resolving the socket, even when another process owns - // it. Destructive paths re-resolve through resolveLiveHypervisorPID. - stalePID := stale.Process.Pid - stored := StoredMetadata{HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &stalePID}, SocketPath: socketPath} - refreshHypervisorPID(&stored, StateRunning) - require.NotNil(t, stored.HypervisorPID) - assert.Equal(t, stalePID, *stored.HypervisorPID) -} - -func TestRefreshHypervisorPIDResolvesSocketOwnerWhenStoredPIDIsDead(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - listener, err := net.Listen("unix", socketPath) - require.NoError(t, err) - defer listener.Close() - - const deadPID = 1<<22 - 1 - require.False(t, ProcessExists(deadPID)) - storedPID := deadPID - stored := StoredMetadata{HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &storedPID}, SocketPath: socketPath} - refreshHypervisorPID(&stored, StateRunning) - - require.NotNil(t, stored.HypervisorPID) - assert.Equal(t, os.Getpid(), *stored.HypervisorPID) - assert.NotZero(t, stored.HypervisorStartTime, "confirmed socket owner mints the identity token") - assert.Equal(t, hostBootID(), stored.HypervisorBootID) -} - -func TestKillHypervisorSurvivesConcurrentReaper(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - process := exec.Command(os.Args[0], "-test.run=^TestSocketListenerHelper$") - process.Env = append(os.Environ(), "HYPERVISOR_SOCKET_HELPER=1", "HYPERVISOR_SOCKET_PATH="+socketPath) - stdin, err := process.StdinPipe() - require.NoError(t, err) - stdout, err := process.StdoutPipe() - require.NoError(t, err) - require.NoError(t, process.Start()) - _, err = bufio.NewReader(stdout).ReadString('\n') - require.NoError(t, err) - - pid := process.Process.Pid - waitDone := make(chan error, 1) - go func() { waitDone <- process.Wait() }() - t.Cleanup(func() { - _ = stdin.Close() - _ = process.Process.Kill() - <-waitDone - }) - - m := &manager{} - require.NoError(t, m.killHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{Id: "kill-test", HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &pid}, SocketPath: socketPath}, - })) -} - -func TestKillHypervisorSucceedsOnReusedPIDWhenNoProcessOwnsSocket(t *testing.T) { - stale := exec.Command("sleep", "30") - require.NoError(t, stale.Start()) - t.Cleanup(func() { - _ = stale.Process.Kill() - _ = stale.Wait() - }) - - // Legacy metadata: live stored PID, no boot-scoped identity, and no - // process anywhere owns or references the socket. That disproves the - // recorded hypervisor is alive, so the kill must succeed as a no-op - // instead of wedging stop/delete, while the PID holder stays untouched. - stalePID := stale.Process.Pid - socketPath := filepath.Join(t.TempDir(), "missing.sock") - m := &manager{} - require.NoError(t, m.killHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{Id: "kill-test", HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &stalePID}, SocketPath: socketPath}, - })) - - assert.NoError(t, syscall.Kill(stalePID, 0), "process with a recycled PID must not be killed") -} - -func TestKillHypervisorIgnoresCommandLineBystander(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - // A process whose command line contains the socket path but that does not - // own a listening socket (e.g. a debug client like ch-remote) is invisible - // to resolution: no listener exists, so the recorded hypervisor is - // provably gone and the kill completes without signaling the bystander. - match := exec.Command("sh", "-c", "sleep 30", "sh", socketPath) - require.NoError(t, match.Start()) - t.Cleanup(func() { - _ = match.Process.Kill() - _ = match.Wait() - }) - - matchPID := match.Process.Pid - m := &manager{} - require.NoError(t, m.killHypervisor(context.Background(), &Instance{ - StoredMetadata: StoredMetadata{Id: "kill-test", HypervisorProcessIdentity: HypervisorProcessIdentity{HypervisorPID: &matchPID}, SocketPath: socketPath}, - })) - - assert.NoError(t, syscall.Kill(matchPID, 0), "command-line bystander must not be killed") -} - -func TestResolveRuntimeHypervisorPIDMintsIdentityOnlyWhenConfirmed(t *testing.T) { - log := slog.New(slog.NewTextHandler(io.Discard, nil)) - const deadPID = 1<<22 - 1 - require.False(t, ProcessExists(deadPID)) - - t.Run("live direct child", func(t *testing.T) { - child := exec.Command("sleep", "30") - require.NoError(t, child.Start()) - t.Cleanup(func() { - _ = child.Process.Kill() - _ = child.Wait() - }) - - stored := &StoredMetadata{SocketPath: filepath.Join(t.TempDir(), "missing.sock")} - pid := resolveRuntimeHypervisorPID(log, stored, child.Process.Pid) - - assert.Equal(t, child.Process.Pid, pid) - require.NotNil(t, stored.HypervisorPID) - assert.Equal(t, child.Process.Pid, *stored.HypervisorPID) - assert.NotZero(t, stored.HypervisorStartTime) - assert.NotEmpty(t, stored.HypervisorBootID) - }) - - t.Run("confirmed socket owner", func(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - listener, err := net.Listen("unix", socketPath) - require.NoError(t, err) - defer listener.Close() - - stored := &StoredMetadata{SocketPath: socketPath} - pid := resolveRuntimeHypervisorPID(log, stored, deadPID) - - assert.Equal(t, os.Getpid(), pid) - require.NotNil(t, stored.HypervisorPID) - assert.Equal(t, os.Getpid(), *stored.HypervisorPID) - assert.NotZero(t, stored.HypervisorStartTime) - assert.NotEmpty(t, stored.HypervisorBootID) - }) - - t.Run("command-line bystander is not adopted", func(t *testing.T) { - socketPath := filepath.Join(t.TempDir(), "test.sock") - match := exec.Command("sh", "-c", "sleep 30", "sh", socketPath) - require.NoError(t, match.Start()) - t.Cleanup(func() { - _ = match.Process.Kill() - _ = match.Wait() - }) - - stored := &StoredMetadata{SocketPath: socketPath} - pid := resolveRuntimeHypervisorPID(log, stored, deadPID) - - assert.Equal(t, deadPID, pid, "a process matching only by command line must not be resolved") - require.NotNil(t, stored.HypervisorPID) - assert.Equal(t, deadPID, *stored.HypervisorPID) - assert.Zero(t, stored.HypervisorStartTime, "a dead fallback must not mint the identity token") - assert.Empty(t, stored.HypervisorBootID, "a dead fallback must not mint the identity token") - }) - - t.Run("dead fallback with unresolvable socket stores bare PID", func(t *testing.T) { - stored := &StoredMetadata{SocketPath: filepath.Join(t.TempDir(), "missing.sock")} - pid := resolveRuntimeHypervisorPID(log, stored, deadPID) - - assert.Equal(t, deadPID, pid) - require.NotNil(t, stored.HypervisorPID) - assert.Equal(t, deadPID, *stored.HypervisorPID) - assert.Zero(t, stored.HypervisorStartTime, "a dead fallback must not mint the identity token") - assert.Empty(t, stored.HypervisorBootID, "a dead fallback must not mint the identity token") - }) -} diff --git a/lib/instances/process_identity_test.go b/lib/instances/process_identity_test.go deleted file mode 100644 index e6fdf1f46..000000000 --- a/lib/instances/process_identity_test.go +++ /dev/null @@ -1,41 +0,0 @@ -package instances - -import ( - "encoding/json" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -// TestHypervisorProcessIdentityJSONKeysStayFlat guards the on-disk metadata -// format: the identity struct is embedded anonymously so its fields keep the -// flat JSON keys metadata files were written with before the struct existed. -func TestHypervisorProcessIdentityJSONKeysStayFlat(t *testing.T) { - pid := 1234 - stored := StoredMetadata{ - Id: "inst-json", - HypervisorProcessIdentity: HypervisorProcessIdentity{ - HypervisorPID: &pid, - HypervisorStartTime: 42, - HypervisorBootID: "boot-id", - }, - } - - data, err := json.Marshal(stored) - require.NoError(t, err) - - var keys map[string]json.RawMessage - require.NoError(t, json.Unmarshal(data, &keys)) - assert.Contains(t, keys, "HypervisorPID") - assert.Contains(t, keys, "HypervisorStartTime") - assert.Contains(t, keys, "HypervisorBootID") - assert.NotContains(t, keys, "HypervisorProcessIdentity") - - var decoded StoredMetadata - require.NoError(t, json.Unmarshal(data, &decoded)) - require.NotNil(t, decoded.HypervisorPID) - assert.Equal(t, pid, *decoded.HypervisorPID) - assert.Equal(t, uint64(42), decoded.HypervisorStartTime) - assert.Equal(t, "boot-id", decoded.HypervisorBootID) -} diff --git a/lib/instances/qemu_lifecycle_test.go b/lib/instances/qemu_lifecycle_test.go index ec7da113a..114c76390 100644 --- a/lib/instances/qemu_lifecycle_test.go +++ b/lib/instances/qemu_lifecycle_test.go @@ -154,7 +154,8 @@ func runQEMUStandbyAndRestore(t *testing.T, hypervisorType hypervisor.Type, inst // Cleanup t.Log("Cleaning up...") - deleteInstanceEventually(t, ctx, manager, inst.Id) + err = manager.DeleteInstance(ctx, inst.Id) + require.NoError(t, err) // Verify cleanup assert.NoDirExists(t, p.InstanceDir(inst.Id)) diff --git a/lib/instances/query.go b/lib/instances/query.go index 8e3ed61f1..98c5359e0 100644 --- a/lib/instances/query.go +++ b/lib/instances/query.go @@ -7,9 +7,11 @@ import ( "io" "os" "path/filepath" + "runtime" "slices" "strconv" "strings" + "syscall" "time" "github.com/kernel/hypeman/lib/guest" @@ -569,6 +571,59 @@ func (m *manager) toInstanceWithStateDerivation(ctx context.Context, meta *metad return inst } +func refreshHypervisorPID(stored *StoredMetadata, state State) { + if !state.RequiresVMM() && state != StateUnknown { + return + } + if stored.HypervisorPID != nil && processExists(*stored.HypervisorPID) { + return + } + if stored.SocketPath == "" { + return + } + if pid, err := hypervisor.ResolveProcessPID(stored.SocketPath); err == nil { + stored.HypervisorPID = &pid + return + } +} + +func processExists(pid int) bool { + if pid <= 0 { + return false + } + err := syscall.Kill(pid, 0) + if err != nil && err != syscall.EPERM { + return false + } + if runtime.GOOS != "linux" { + return true + } + state, err := readLinuxProcessState(pid) + if err != nil { + return true + } + return state != "Z" +} + +func readLinuxProcessState(pid int) (string, error) { + statusPath := filepath.Join("/proc", strconv.Itoa(pid), "status") + data, err := os.ReadFile(statusPath) + if err != nil { + return "", err + } + for _, line := range strings.Split(string(data), "\n") { + if !strings.HasPrefix(line, "State:") { + continue + } + fields := strings.Fields(line) + if len(fields) < 2 { + return "", fmt.Errorf("malformed process state in %s", statusPath) + } + return fields[1], nil + } + return "", fmt.Errorf("process state missing from %s", statusPath) +} + // parseExitSentinel reads the last lines of the serial console log to find the // HYPEMAN-EXIT sentinel written by init before shutdown. // Returns the exit code, message, and whether a sentinel was found. diff --git a/lib/instances/restore.go b/lib/instances/restore.go index ab27903ba..c209274e2 100644 --- a/lib/instances/restore.go +++ b/lib/instances/restore.go @@ -298,8 +298,7 @@ func (m *manager) restoreInstance( attribute.String("operation", "restore_from_snapshot"), ) log.InfoContext(ctx, "restoring from snapshot", "instance_id", id, "snapshot_dir", snapshotDir, "hypervisor", stored.HypervisorType) - // restoreFromSnapshot records the hypervisor process identity on stored. - _, hv, err := m.restoreFromSnapshot(restoreCtx, stored, snapshotDir, restoreOptions) + pid, hv, err := m.restoreFromSnapshot(restoreCtx, stored, snapshotDir, restoreOptions) restoreSpanEnd(err) if err != nil { log.ErrorContext(ctx, "failed to restore from snapshot", "instance_id", id, "error", err) @@ -309,6 +308,9 @@ func (m *manager) restoreInstance( return nil, err } + // Store the PID for later cleanup + stored.HypervisorPID = &pid + // 6. Transition: Paused → Running (resume) resumeCtx, resumeSpanEnd := m.startLifecycleStep(ctx, "resume_vm", attribute.String("instance_id", id), @@ -446,7 +448,7 @@ func (m *manager) restoreFromSnapshot( if err != nil { return 0, nil, fmt.Errorf("restore vm: %w", err) } - pid = resolveRuntimeHypervisorPID(log, stored, pid) + pid = resolveRuntimeHypervisorPID(log, stored.SocketPath, pid) log.DebugContext(ctx, "VM restored from snapshot successfully", "instance_id", stored.Id, "pid", pid) return pid, hv, nil diff --git a/lib/instances/snapshot.go b/lib/instances/snapshot.go index 2d2676c72..eef7a6092 100644 --- a/lib/instances/snapshot.go +++ b/lib/instances/snapshot.go @@ -301,17 +301,11 @@ func (m *manager) restoreSnapshot(ctx context.Context, id string, snapshotID str restored.Name = sourceMeta.Name restored.ExpiresAt = sourceMeta.ExpiresAt restored.DataDir = m.paths.InstanceDir(id) - restored.HypervisorProcessIdentity.Clear() + restored.HypervisorPID = nil restored.StartedAt = nil restored.StoppedAt = nil restored.ExitCode = nil restored.ExitMessage = "" - // vGPU assignments are live host state, not snapshot payload: keep the - // instance's current assignment (possibly retained from a failed release) - // instead of resurrecting the one embedded in the snapshot. - restored.GPUFramework = sourceMeta.GPUFramework - restored.GPUDevicePath = sourceMeta.GPUDevicePath - restored.GPUMdevUUID = sourceMeta.GPUMdevUUID restored.HypervisorType = targetHypervisor restored.HypervisorVersion = targetHypervisorVersion restored.SocketPath = m.paths.InstanceSocket(id, starter.SocketName()) @@ -432,7 +426,7 @@ func (m *manager) forkSnapshot(ctx context.Context, snapshotID string, req ForkS forkMeta.ExpiresAt = nil forkMeta.StartedAt = nil forkMeta.StoppedAt = nil - forkMeta.HypervisorProcessIdentity.Clear() + forkMeta.HypervisorPID = nil forkMeta.DataDir = dstDir forkMeta.HypervisorType = targetHypervisor forkMeta.HypervisorVersion = targetHypervisorVersion @@ -441,7 +435,6 @@ func (m *manager) forkSnapshot(ctx context.Context, snapshotID string, req ForkS forkMeta.ExitCode = nil forkMeta.ExitMessage = "" forkMeta.RestartStatus = restartpolicy.Status{} - clearStoredVGPUDevice(&forkMeta) forkMeta.FirecrackerUFFDSessionID = "" forkMeta.FirecrackerUFFDPagerVersion = "" forkMeta.FirecrackerUseUFFDOnNextRestore = useFirecrackerUFFDOnNextRestore(targetHypervisor, rec.Snapshot.Kind == SnapshotKindStandby, targetState) diff --git a/lib/instances/snapshot_test.go b/lib/instances/snapshot_test.go index b23e364ee..917f88d46 100644 --- a/lib/instances/snapshot_test.go +++ b/lib/instances/snapshot_test.go @@ -8,7 +8,6 @@ import ( "testing" "time" - "github.com/kernel/hypeman/lib/devices" "github.com/kernel/hypeman/lib/hypervisor" "github.com/kernel/hypeman/lib/images" snapshotstore "github.com/kernel/hypeman/lib/snapshot" @@ -16,118 +15,6 @@ import ( "github.com/stretchr/testify/require" ) -func TestForkSnapshotClearsVGPUAssignment(t *testing.T) { - mgr, _ := setupTestManager(t) - ctx := context.Background() - - sourceID := "snapshot-vgpu-source" - createStoppedSnapshotSourceFixture(t, mgr, sourceID, sourceID, mgr.defaultHypervisor) - - meta, err := mgr.loadMetadata(sourceID) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - meta.GPUMdevUUID = "retained-uuid" - require.NoError(t, mgr.saveMetadata(meta)) - - snapshot, err := mgr.CreateSnapshot(ctx, sourceID, CreateSnapshotRequest{ - Kind: SnapshotKindStopped, - Name: "snapshot-vgpu", - }) - require.NoError(t, err) - - forked, err := mgr.ForkSnapshot(ctx, snapshot.Id, ForkSnapshotRequest{ - Name: "snapshot-vgpu-fork", - TargetState: StateStopped, - }) - require.NoError(t, err) - assert.Equal(t, "NVIDIA L40S-2Q", forked.GPUProfile) - assert.Equal(t, devices.VGPUFrameworkNone, forked.GPUFramework) - assert.Empty(t, forked.GPUDevicePath) - assert.Empty(t, forked.GPUMdevUUID) - - source, err := mgr.loadMetadata(sourceID) - require.NoError(t, err) - assert.Equal(t, "/sys/bus/pci/devices/0000:82:00.4", source.GPUDevicePath) -} - -func TestRestoreSnapshotDoesNotResurrectStaleVGPUAssignment(t *testing.T) { - mgr, _ := setupTestManager(t) - ctx := context.Background() - - sourceID := "snapshot-vgpu-restore-stale" - createStoppedSnapshotSourceFixture(t, mgr, sourceID, sourceID, mgr.defaultHypervisor) - - meta, err := mgr.loadMetadata(sourceID) - require.NoError(t, err) - meta.GPUProfile = "NVIDIA L40S-2Q" - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - meta.GPUMdevUUID = "retained-uuid" - require.NoError(t, mgr.saveMetadata(meta)) - - snapshot, err := mgr.CreateSnapshot(ctx, sourceID, CreateSnapshotRequest{ - Kind: SnapshotKindStopped, - Name: "snapshot-vgpu-restore-stale", - }) - require.NoError(t, err) - - // The retained assignment is released successfully after the snapshot - // was taken; a restore must not resurrect the snapshot's embedded copy. - meta, err = mgr.loadMetadata(sourceID) - require.NoError(t, err) - clearStoredVGPUDevice(&meta.StoredMetadata) - require.NoError(t, mgr.saveMetadata(meta)) - - _, err = mgr.RestoreSnapshot(ctx, sourceID, snapshot.Id, RestoreSnapshotRequest{ - TargetState: StateStopped, - TargetHypervisor: mgr.defaultHypervisor, - }) - require.NoError(t, err) - - restored, err := mgr.loadMetadata(sourceID) - require.NoError(t, err) - assert.Equal(t, devices.VGPUFrameworkNone, restored.GPUFramework) - assert.Empty(t, restored.GPUDevicePath) - assert.Empty(t, restored.GPUMdevUUID) -} - -func TestRestoreSnapshotKeepsCurrentVGPUAssignment(t *testing.T) { - mgr, _ := setupTestManager(t) - ctx := context.Background() - - sourceID := "snapshot-vgpu-restore-retained" - createStoppedSnapshotSourceFixture(t, mgr, sourceID, sourceID, mgr.defaultHypervisor) - - snapshot, err := mgr.CreateSnapshot(ctx, sourceID, CreateSnapshotRequest{ - Kind: SnapshotKindStopped, - Name: "snapshot-vgpu-restore-retained", - }) - require.NoError(t, err) - - // An assignment retained after the snapshot was taken (e.g. from a - // failed release on stop) must survive the restore for the next retry. - meta, err := mgr.loadMetadata(sourceID) - require.NoError(t, err) - meta.GPUFramework = devices.VGPUFramework("future-framework") - meta.GPUDevicePath = "/sys/bus/pci/devices/0000:82:00.4" - meta.GPUMdevUUID = "retained-uuid" - require.NoError(t, mgr.saveMetadata(meta)) - - _, err = mgr.RestoreSnapshot(ctx, sourceID, snapshot.Id, RestoreSnapshotRequest{ - TargetState: StateStopped, - TargetHypervisor: mgr.defaultHypervisor, - }) - require.NoError(t, err) - - restored, err := mgr.loadMetadata(sourceID) - require.NoError(t, err) - assert.Equal(t, devices.VGPUFramework("future-framework"), restored.GPUFramework) - assert.Equal(t, "/sys/bus/pci/devices/0000:82:00.4", restored.GPUDevicePath) - assert.Equal(t, "retained-uuid", restored.GPUMdevUUID) -} - func TestStoppedSnapshotLifecycleAndForkAfterSourceDeletion(t *testing.T) { t.Parallel() mgr, _ := setupTestManager(t) diff --git a/lib/instances/standby.go b/lib/instances/standby.go index 4c6fd7d33..6913a9895 100644 --- a/lib/instances/standby.go +++ b/lib/instances/standby.go @@ -6,6 +6,7 @@ import ( "fmt" "os" "path/filepath" + "syscall" "time" "github.com/kernel/hypeman/lib/guest" @@ -184,15 +185,11 @@ func (m *manager) standbyInstance( ) if err := m.shutdownHypervisor(shutdownCtx, &inst); err != nil { shutdownSpanEnd(err) - // The hypervisor may still be running: releasing its TAP or clearing - // its identity now would orphan a live VMM. The snapshot on disk is - // harmless and a retried standby redoes it. - if resumeErr := hv.Resume(ctx); resumeErr != nil { - log.ErrorContext(ctx, "failed to resume VM after shutdown error", "instance_id", id, "error", resumeErr) - } - return nil, fmt.Errorf("shutdown hypervisor: %w", err) + // Log but continue - snapshot was created successfully + log.WarnContext(ctx, "failed to shutdown hypervisor gracefully, snapshot still valid", "instance_id", id, "error", err) + } else { + shutdownSpanEnd(nil) } - shutdownSpanEnd(nil) // Firecracker vsock sockets can persist across standby/restore if the process // exits ungracefully. Remove stale sockets before restore attempts. @@ -232,7 +229,7 @@ func (m *manager) standbyInstance( // 10. Update timestamp and clear PID (hypervisor no longer running) now := time.Now().UTC() stored.StoppedAt = &now - stored.HypervisorProcessIdentity.Clear() + stored.HypervisorPID = nil stored.PendingStandbyCompression = nil clearFirecrackerUFFDRestoreState(stored) if err := m.refreshFirecrackerSnapshotCacheKey(stored, snapshotDir); err != nil { @@ -358,31 +355,16 @@ func restoreRetainedSnapshotBase(snapshotDir string, retainedBaseDir string) err // shutdownHypervisor gracefully shuts down the hypervisor process via API func (m *manager) shutdownHypervisor(ctx context.Context, inst *Instance) error { log := logger.FromContext(ctx) - - // Resolve the live owner before any teardown: the stored PID may be stale - // or recycled, and signaling it raw would bypass the ownership checks the - // kill paths enforce. Failing closed here also keeps the control socket in - // place as evidence for a later hardened kill. - pid, err := resolveLiveHypervisorPID(inst.HypervisorProcessIdentity, inst.SocketPath) - if err != nil { - return fmt.Errorf("confirm hypervisor ownership before shutdown: %w", err) - } + defer func() { + // Clean stale sockets even if graceful shutdown fails. + _ = os.Remove(inst.SocketPath) + }() // Try to connect to hypervisor hv, err := m.getHypervisor(inst.SocketPath, inst.HypervisorType) if err != nil { - if pid > 0 { - // The control client cannot be built but the resolved owner is - // alive; teardown is committed, so kill it rather than report a - // completed shutdown for a VMM that is still running. - log.WarnContext(ctx, "could not connect to hypervisor, force killing resolved owner", "instance_id", inst.Id, "pid", pid, "error", err) - if err := killProcessAndWait(pid); err != nil { - return err - } - } - // The hypervisor is confirmed gone (killed above or no live owner); - // remove its stale socket. - _ = os.Remove(inst.SocketPath) + // Can't connect - hypervisor might already be stopped + log.DebugContext(ctx, "could not connect to hypervisor, may already be stopped", "instance_id", inst.Id) return nil } @@ -397,37 +379,54 @@ func (m *manager) shutdownHypervisor(ctx context.Context, inst *Instance) error shutdownErr = hv.Shutdown(ctx) } + // Teardown is committed; prevent new control-socket clients while the + // hypervisor exits. The deferred remove remains as a fallback for early + // returns above. + _ = os.Remove(inst.SocketPath) + // Wait for process to exit - if pid > 0 { + if inst.HypervisorPID != nil { + pid := *inst.HypervisorPID shouldWaitForGracefulExit := caps.SupportsGracefulVMMShutdown && shutdownErr != hypervisor.ErrNotSupported if shouldWaitForGracefulExit { if WaitForProcessExit(pid, 2*time.Second) { log.DebugContext(ctx, "hypervisor shutdown gracefully", "instance_id", inst.Id, "pid", pid) } else { log.WarnContext(ctx, "hypervisor did not exit gracefully in time, force killing process", "instance_id", inst.Id, "pid", pid) - if err := killProcessAndWait(pid); err != nil { + if err := forceKillHypervisorPID(pid); err != nil { return err } } } else { log.DebugContext(ctx, "skipping graceful exit wait; force killing hypervisor process", "instance_id", inst.Id, "pid", pid) - if err := killProcessAndWait(pid); err != nil { + if err := forceKillHypervisorPID(pid); err != nil { return err } } } - // The hypervisor is confirmed gone (graceful exit, force kill, or no live - // owner), so its socket is stale now. Removing it any earlier would unlink - // the control socket of a VMM that survives the kill and gets resumed, - // leaving that VM unreachable for a later graceful standby or stop. - _ = os.Remove(inst.SocketPath) - - // A graceful-API error at this point is not a failure: an error from this - // function means the hypervisor may still be running. if shutdownErr != nil && shutdownErr != hypervisor.ErrNotSupported { - log.WarnContext(ctx, "graceful hypervisor shutdown failed, process force killed", "instance_id", inst.Id, "error", shutdownErr) + return fmt.Errorf("graceful hypervisor shutdown failed: %w", shutdownErr) + } + + return nil +} + +func forceKillHypervisorPID(pid int) error { + if err := syscall.Kill(pid, syscall.SIGKILL); err != nil { + if err == syscall.ESRCH { + return nil + } + return fmt.Errorf("force kill hypervisor pid %d: %w", pid, err) + } + if WaitForProcessExit(pid, 2*time.Second) { + return nil } + // The process may have spawned children in its own process group. + _ = syscall.Kill(-pid, syscall.SIGKILL) + if !WaitForProcessExit(pid, 2*time.Second) { + return fmt.Errorf("hypervisor pid %d did not exit after SIGKILL", pid) + } return nil } diff --git a/lib/instances/start.go b/lib/instances/start.go index 7e7855eac..032b71076 100644 --- a/lib/instances/start.go +++ b/lib/instances/start.go @@ -49,21 +49,6 @@ func (m *manager) startInstance( return nil, fmt.Errorf("%w: cannot start from state %s, must be Stopped", ErrInvalidState, inst.State) } - // Release any assignment retained by an earlier failed release and - // persist the cleared fields immediately, so a failure later in start - // cannot leave on-disk metadata pointing at a device that is already - // gone (matching releaseRetainedVGPULocked). - if storedVGPUDevicePath(stored) != "" { - if err := releaseStoredVGPU(ctx, stored); err != nil { - log.ErrorContext(ctx, "failed to release stale vGPU before start", "instance_id", id, "error", err) - return nil, fmt.Errorf("release stale vGPU before start: %w", err) - } - if err := m.saveMetadata(meta); err != nil { - log.ErrorContext(ctx, "failed to save metadata after stale vGPU release", "instance_id", id, "error", err) - return nil, fmt.Errorf("save metadata after stale vGPU release: %w", err) - } - } - // 2a. Clear stale exit info from previous run and apply command overrides stored.ExitCode = nil stored.ExitMessage = "" @@ -159,27 +144,22 @@ func (m *manager) startInstance( } } - // 4b. Recreate the vGPU if this instance had a GPU profile + // 4b. Recreate vGPU mdev if this instance had a GPU profile // Note: GPU availability was already validated in step 2b if stored.GPUProfile != "" { log.InfoContext(ctx, "creating vGPU mdev for start", "instance_id", id, "profile", stored.GPUProfile) - device, err := devices.CreateVGPU(ctx, stored.GPUProfile, id) + mdev, err := devices.CreateMdev(ctx, stored.GPUProfile, id) if err != nil { - log.ErrorContext(ctx, "failed to create vGPU", "instance_id", id, "profile", stored.GPUProfile, "error", err) + log.ErrorContext(ctx, "failed to create mdev", "instance_id", id, "profile", stored.GPUProfile, "error", err) return nil, fmt.Errorf("create vGPU mdev for profile %s: %w", stored.GPUProfile, err) } - setStoredVGPUDevice(stored, device) - log.InfoContext(ctx, "created vGPU", "instance_id", id, "profile", stored.GPUProfile, "uuid", device.MdevUUID) - // Add vGPU cleanup to stack + stored.GPUMdevUUID = mdev.UUID + log.InfoContext(ctx, "created vGPU mdev", "instance_id", id, "profile", stored.GPUProfile, "uuid", mdev.UUID) + // Add mdev cleanup to stack cu.Add(func() { - log.DebugContext(ctx, "destroying vGPU on cleanup", "instance_id", id, "uuid", device.MdevUUID) - assignment := devices.VGPUAssignment{ - Framework: device.Framework, - DevicePath: device.SysfsPath, - MdevUUID: device.MdevUUID, - } - if err := devices.DestroyVGPU(ctx, assignment); err != nil { - log.WarnContext(ctx, "failed to destroy vGPU on cleanup", "instance_id", id, "uuid", device.MdevUUID, "error", err) + log.DebugContext(ctx, "destroying mdev on cleanup", "instance_id", id, "uuid", mdev.UUID) + if err := devices.DestroyMdev(ctx, mdev.UUID); err != nil { + log.WarnContext(ctx, "failed to destroy mdev on cleanup", "instance_id", id, "uuid", mdev.UUID, "error", err) } }) } diff --git a/lib/instances/stop.go b/lib/instances/stop.go index 7eddeb034..fb60c68c9 100644 --- a/lib/instances/stop.go +++ b/lib/instances/stop.go @@ -6,8 +6,10 @@ import ( "fmt" "os" "path/filepath" + "syscall" "time" + "github.com/kernel/hypeman/lib/devices" "github.com/kernel/hypeman/lib/guest" "github.com/kernel/hypeman/lib/hypervisor" "github.com/kernel/hypeman/lib/instances/phasetracking" @@ -50,27 +52,10 @@ func (m *manager) tryGracefulGuestShutdown(ctx context.Context, inst *Instance, return false } - // Capture the socket owner before shutdown can close its listener. Legacy - // metadata has no process identity token, so the PID cannot be recovered - // safely from the stored value once the listener disappears. - pid, err := resolveLiveHypervisorPID(inst.HypervisorProcessIdentity, inst.SocketPath) - if err != nil { - log.WarnContext(ctx, "could not confirm hypervisor ownership before graceful shutdown", "instance_id", inst.Id, "error", err) - return false - } - if pid == 0 { - return true - } - inst.HypervisorProcessIdentity.Set(pid) - - shutdownGuest := guest.ShutdownInstance - if m.shutdownGuestFn != nil { - shutdownGuest = m.shutdownGuestFn - } sendShutdown := func() error { shutdownCtx, cancel := context.WithTimeout(ctx, shutdownRPCDeadline) defer cancel() - return shutdownGuest(shutdownCtx, dialer, 0) + return guest.ShutdownInstance(shutdownCtx, dialer, 0) } shutdownSent := false @@ -86,21 +71,71 @@ func (m *manager) tryGracefulGuestShutdown(ctx context.Context, inst *Instance, shutdownSent = true } - waitTimeout := time.Duration(stopTimeout) * time.Second - if !shutdownSent && waitTimeout > shutdownFailureFallbackWait { - // If we couldn't signal the guest, don't burn the full graceful timeout. - waitTimeout = shutdownFailureFallbackWait - } + // Wait for the hypervisor process to exit (init calls reboot(POWER_OFF)) + if inst.HypervisorPID != nil { + waitTimeout := time.Duration(stopTimeout) * time.Second + if !shutdownSent && waitTimeout > shutdownFailureFallbackWait { + // If we couldn't signal the guest, don't burn the full graceful timeout. + waitTimeout = shutdownFailureFallbackWait + } - if WaitForProcessExit(pid, waitTimeout) { - log.DebugContext(ctx, "VM shut down gracefully", "instance_id", inst.Id) - return true + if WaitForProcessExit(*inst.HypervisorPID, waitTimeout) { + log.DebugContext(ctx, "VM shut down gracefully", "instance_id", inst.Id) + return true + } + + log.WarnContext(ctx, "graceful shutdown timed out, falling back to hypervisor shutdown", "instance_id", inst.Id) + return false } - log.WarnContext(ctx, "graceful shutdown timed out, falling back to hypervisor shutdown", "instance_id", inst.Id) return false } +// forceKillHypervisorProcess sends SIGKILL to the hypervisor process if it's still running +// and waits briefly for it to exit. +func (m *manager) forceKillHypervisorProcess(ctx context.Context, inst *Instance) error { + log := logger.FromContext(ctx) + + if inst.HypervisorPID == nil { + return nil + } + + pid := *inst.HypervisorPID + if err := syscall.Kill(pid, 0); err != nil { + // Process is already gone (likely ESRCH). + return nil + } + + log.WarnContext(ctx, "hypervisor still running after shutdown fallback, sending SIGKILL", "instance_id", inst.Id, "pid", pid) + if err := syscall.Kill(pid, syscall.SIGKILL); err != nil { + return fmt.Errorf("sigkill hypervisor pid %d: %w", pid, err) + } + + // Wait for process to die and reap it to avoid zombie false positives. + reaped := false + for i := 0; i < 50; i++ { // 50 * 100ms = 5s + var wstatus syscall.WaitStatus + wpid, err := syscall.Wait4(pid, &wstatus, syscall.WNOHANG, nil) + if err != nil || wpid == pid { + // Process reaped, or not our child (ECHILD) and no longer trackable here. + reaped = true + break + } + time.Sleep(100 * time.Millisecond) + } + + if !reaped { + // Timed out waiting for reap; if process still exists, treat as failure. + if err := syscall.Kill(pid, 0); err == nil { + return fmt.Errorf("hypervisor pid %d still alive after SIGKILL", pid) + } + log.WarnContext(ctx, "timeout waiting to reap hypervisor process after SIGKILL", "instance_id", inst.Id, "pid", pid) + } + + log.DebugContext(ctx, "hypervisor process force-killed", "instance_id", inst.Id, "pid", pid) + return nil +} + // stopInstance gracefully stops an active instance. // Flow: send Shutdown RPC -> wait for VM to power off -> // fall back to hypervisor shutdown -> final SIGKILL if still alive. @@ -173,25 +208,24 @@ func (m *manager) stopInstance( ) if err := m.shutdownHypervisor(shutdownCtx, &inst); err != nil { shutdownSpanEnd(err) + // Continue to final SIGKILL fallback if graceful shutdown API fails. log.WarnContext(ctx, "failed to shutdown hypervisor", "instance_id", id, "error", err) - - // Final fallback: force-kill the process. A nil return from - // shutdownHypervisor already confirmed the hypervisor is gone, so - // this only runs when shutdown could not. - killCtx, killSpanEnd := m.startLifecycleStep(ctx, "force_kill_hypervisor", - attribute.String("instance_id", id), - attribute.String("hypervisor", string(stored.HypervisorType)), - attribute.String("operation", "force_kill_hypervisor"), - ) - if err := m.killHypervisor(killCtx, &inst); err != nil { - killSpanEnd(err) - log.ErrorContext(ctx, "failed to force-kill hypervisor process", "instance_id", id, "error", err) - return nil, err - } - killSpanEnd(nil) } else { shutdownSpanEnd(nil) } + + // Final fallback: force-kill the process if it's still alive. + killCtx, killSpanEnd := m.startLifecycleStep(ctx, "force_kill_hypervisor", + attribute.String("instance_id", id), + attribute.String("hypervisor", string(stored.HypervisorType)), + attribute.String("operation", "force_kill_hypervisor"), + ) + if err := m.forceKillHypervisorProcess(killCtx, &inst); err != nil { + killSpanEnd(err) + log.ErrorContext(ctx, "failed to force-kill hypervisor process", "instance_id", id, "error", err) + return nil, err + } + killSpanEnd(nil) } // 6. Release network allocation (delete TAP device) @@ -229,11 +263,12 @@ func (m *manager) stopInstance( } } - // 7. Release the vGPU assignment if present (frees the vGPU slot for other VMs). - if path := storedVGPUDevicePath(stored); path != "" { - log.InfoContext(ctx, "destroying vGPU on stop", "instance_id", id, "device_path", path) - if err := releaseStoredVGPU(ctx, stored); err != nil { - log.WarnContext(ctx, "failed to destroy vGPU on stop; retaining assignment metadata", "instance_id", id, "device_path", path, "error", err) + // 7. Destroy vGPU mdev device if present (frees vGPU slot for other VMs) + if inst.GPUMdevUUID != "" { + log.InfoContext(ctx, "destroying vGPU mdev on stop", "instance_id", id, "uuid", inst.GPUMdevUUID) + if err := devices.DestroyMdev(ctx, inst.GPUMdevUUID); err != nil { + // Log error but continue - mdev cleanup is best-effort + log.WarnContext(ctx, "failed to destroy mdev on stop", "instance_id", id, "uuid", inst.GPUMdevUUID, "error", err) } } @@ -263,10 +298,11 @@ func (m *manager) stopInstance( } } - // 10. Update metadata (clear PID, set StoppedAt) + // 10. Update metadata (clear PID, mdev UUID, set StoppedAt) now := time.Now().UTC() stored.StoppedAt = &now - stored.HypervisorProcessIdentity.Clear() + stored.HypervisorPID = nil + stored.GPUMdevUUID = "" // Clear mdev UUID since we destroyed it // Boot markers are per-boot-run and must not carry across stop/restore/start. stored.ProgramStartedAt = nil stored.GuestAgentReadyAt = nil diff --git a/lib/instances/types.go b/lib/instances/types.go index 6aac15985..a90ecd574 100644 --- a/lib/instances/types.go +++ b/lib/instances/types.go @@ -4,7 +4,6 @@ import ( "time" "github.com/kernel/hypeman/lib/autostandby" - "github.com/kernel/hypeman/lib/devices" "github.com/kernel/hypeman/lib/healthcheck" "github.com/kernel/hypeman/lib/hypervisor" "github.com/kernel/hypeman/lib/instances/phasetracking" @@ -130,8 +129,7 @@ type StoredMetadata struct { // Hypervisor configuration HypervisorType hypervisor.Type // Hypervisor type (e.g., "cloud-hypervisor") HypervisorVersion string // Hypervisor version (e.g., "v51.1") - // Embedded so its fields keep their flat JSON keys in persisted metadata. - HypervisorProcessIdentity + HypervisorPID *int // Hypervisor process ID (may be stale after host restart) // Firecracker UFFD snapshot restore metadata. FirecrackerSnapshotCacheKey string @@ -151,10 +149,8 @@ type StoredMetadata struct { Devices []string // Device IDs attached to this instance // GPU configuration (vGPU mode) - GPUProfile string // vGPU profile name (e.g., "L40S-1Q") - GPUFramework devices.VGPUFramework - GPUDevicePath string - GPUMdevUUID string // populated for mdev-backed vGPUs + GPUProfile string // vGPU profile name (e.g., "L40S-1Q") + GPUMdevUUID string // mdev device UUID // Command overrides (like docker run ) Entrypoint []string // Override image entrypoint (nil = use image default) diff --git a/lib/instances/version_upgrade_test.go b/lib/instances/version_upgrade_test.go index 3b93eb54f..9e078d533 100644 --- a/lib/instances/version_upgrade_test.go +++ b/lib/instances/version_upgrade_test.go @@ -134,8 +134,8 @@ func TestCloudHypervisorVersionUpgradeRestore(t *testing.T) { // Cleanup t.Log("Cleaning up...") - deleteInstanceEventually(t, ctx, mgr, inst.Id) - deleteInstanceEventually(t, ctx, mgr, inst2.Id) + require.NoError(t, mgr.DeleteInstance(ctx, inst.Id)) + require.NoError(t, mgr.DeleteInstance(ctx, inst2.Id)) t.Log("Version upgrade restore test complete!") } diff --git a/lib/instances/vgpu.go b/lib/instances/vgpu.go deleted file mode 100644 index cffe2ac1d..000000000 --- a/lib/instances/vgpu.go +++ /dev/null @@ -1,71 +0,0 @@ -package instances - -import ( - "context" - "path/filepath" - - "github.com/kernel/hypeman/lib/devices" - "github.com/kernel/hypeman/lib/logger" -) - -func setStoredVGPUDevice(stored *StoredMetadata, device *devices.VGPUDevice) { - stored.GPUFramework = device.Framework - stored.GPUDevicePath = device.SysfsPath - stored.GPUMdevUUID = device.MdevUUID -} - -func clearStoredVGPUDevice(stored *StoredMetadata) { - stored.GPUFramework = devices.VGPUFrameworkNone - stored.GPUDevicePath = "" - stored.GPUMdevUUID = "" -} - -func releaseStoredVGPU(ctx context.Context, stored *StoredMetadata) error { - path := storedVGPUDevicePath(stored) - if path != "" { - assignment := devices.VGPUAssignment{ - Framework: stored.GPUFramework, - DevicePath: path, - MdevUUID: stored.GPUMdevUUID, - } - if err := devices.DestroyVGPU(ctx, assignment); err != nil { - return err - } - } - clearStoredVGPUDevice(stored) - return nil -} - -// releaseRetainedVGPULocked releases a vGPU assignment retained on a stopped -// instance after a failed release during the original stop. It is a no-op -// when no assignment is retained, and a failed retry only logs so the -// metadata stays for the next retry. The caller must hold the instance lock. -func (m *manager) releaseRetainedVGPULocked(ctx context.Context, id string) { - log := logger.FromContext(ctx) - meta, err := m.loadMetadata(id) - if err != nil { - log.WarnContext(ctx, "failed to load metadata for retained vGPU release", "instance_id", id, "error", err) - return - } - stored := &meta.StoredMetadata - if storedVGPUDevicePath(stored) == "" { - return - } - if err := releaseStoredVGPU(ctx, stored); err != nil { - log.WarnContext(ctx, "failed to destroy retained vGPU; retaining assignment metadata", "instance_id", id, "error", err) - return - } - if err := m.saveMetadata(meta); err != nil { - log.WarnContext(ctx, "failed to save metadata after retained vGPU release", "instance_id", id, "error", err) - } -} - -func storedVGPUDevicePath(stored *StoredMetadata) string { - if stored.GPUDevicePath != "" { - return stored.GPUDevicePath - } - if stored.GPUMdevUUID != "" { - return filepath.Join("/sys/bus/mdev/devices", stored.GPUMdevUUID) - } - return "" -} diff --git a/lib/instances/vgpu_test.go b/lib/instances/vgpu_test.go deleted file mode 100644 index 6f2c46819..000000000 --- a/lib/instances/vgpu_test.go +++ /dev/null @@ -1,54 +0,0 @@ -package instances - -import ( - "context" - "testing" - - "github.com/kernel/hypeman/lib/devices" - "github.com/stretchr/testify/assert" -) - -func TestStoredVGPUDevicePath(t *testing.T) { - t.Parallel() - - assert.Equal(t, "/sys/bus/mdev/devices/new-uuid", storedVGPUDevicePath(&StoredMetadata{ - GPUDevicePath: "/sys/bus/mdev/devices/new-uuid", - GPUMdevUUID: "legacy-uuid", - })) - assert.Equal(t, "/sys/bus/mdev/devices/legacy-uuid", storedVGPUDevicePath(&StoredMetadata{ - GPUMdevUUID: "legacy-uuid", - })) - assert.Empty(t, storedVGPUDevicePath(&StoredMetadata{})) -} - -func TestReleaseStoredVGPURetainsMetadataOnFailure(t *testing.T) { - t.Parallel() - - stored := &StoredMetadata{ - GPUFramework: devices.VGPUFramework("future-framework"), - GPUDevicePath: "/sys/bus/pci/devices/0000:82:00.4", - } - err := releaseStoredVGPU(context.Background(), stored) - assert.Error(t, err) - assert.Equal(t, devices.VGPUFramework("future-framework"), stored.GPUFramework) - assert.Equal(t, "/sys/bus/pci/devices/0000:82:00.4", stored.GPUDevicePath) -} - -func TestSetAndClearStoredVGPUDevice(t *testing.T) { - t.Parallel() - - stored := &StoredMetadata{} - setStoredVGPUDevice(stored, &devices.VGPUDevice{ - Framework: devices.VGPUFrameworkMdev, - SysfsPath: "/sys/bus/mdev/devices/new-uuid", - MdevUUID: "new-uuid", - }) - assert.Equal(t, devices.VGPUFrameworkMdev, stored.GPUFramework) - assert.Equal(t, "/sys/bus/mdev/devices/new-uuid", stored.GPUDevicePath) - assert.Equal(t, "new-uuid", stored.GPUMdevUUID) - - clearStoredVGPUDevice(stored) - assert.Empty(t, stored.GPUFramework) - assert.Empty(t, stored.GPUDevicePath) - assert.Empty(t, stored.GPUMdevUUID) -} diff --git a/lib/instances/vm_config_validation.go b/lib/instances/vm_config_validation.go index 61639e9b4..fe1530b54 100644 --- a/lib/instances/vm_config_validation.go +++ b/lib/instances/vm_config_validation.go @@ -6,10 +6,7 @@ import ( "github.com/kernel/hypeman/lib/hypervisor" ) -const ( - baseInstanceDiskCount = 3 // rootfs, writable overlay, and config disk - plannedVGPUDevicePath = "planned-vgpu-device" -) +const baseInstanceDiskCount = 3 // rootfs, writable overlay, and config disk func instanceDiskCount(volumes []VolumeAttachment) int { count := baseInstanceDiskCount @@ -25,13 +22,15 @@ func instanceDiskCount(volumes []VolumeAttachment) int { // validateCreateVMConfig performs side-effect-free backend validation against // the complete device plan before image, PCI, network, or filesystem work. func (m *manager) validateCreateVMConfig(starter hypervisor.VMStarter, req CreateInstanceRequest, hvType hypervisor.Type) error { - hasVGPU := req.GPU != nil && req.GPU.Profile != "" + pciDeviceCount := len(req.Devices) + if req.GPU != nil && req.GPU.Profile != "" { + pciDeviceCount++ + } return validatePlannedVMConfig(starter, hvType, m.plannedVMConfig( req.HotplugSize, req.Volumes, req.NetworkEnabled, - len(req.Devices), - hasVGPU, + pciDeviceCount, )) } @@ -42,13 +41,15 @@ func (m *manager) validateStoredVMConfig(starter hypervisor.VMStarter, snapshotK } func (m *manager) plannedStoredVMConfig(snapshotKind SnapshotKind, meta StoredMetadata) hypervisor.VMConfig { - hasVGPU := storedVGPUDevicePath(&meta) != "" || meta.GPUProfile != "" + pciDeviceCount := len(meta.Devices) + if meta.GPUMdevUUID != "" || meta.GPUProfile != "" { + pciDeviceCount++ + } config := m.plannedVMConfig( meta.HotplugSize, meta.Volumes, meta.NetworkEnabled, - len(meta.Devices), - hasVGPU, + pciDeviceCount, ) if snapshotKind == SnapshotKindStandby { // Standby restore/fork reuses the frozen snapshot device model, so live @@ -63,7 +64,6 @@ func (m *manager) plannedVMConfig( volumes []VolumeAttachment, networkEnabled bool, pciDeviceCount int, - hasVGPU bool, ) hypervisor.VMConfig { diskCount := instanceDiskCount(volumes) @@ -77,9 +77,6 @@ func (m *manager) plannedVMConfig( if networkEnabled { config.Networks = []hypervisor.NetworkConfig{{}} } - if hasVGPU { - config.VGPUDevicePath = plannedVGPUDevicePath - } return config } diff --git a/lib/instances/vm_config_validation_test.go b/lib/instances/vm_config_validation_test.go index d25456722..c7037e6ac 100644 --- a/lib/instances/vm_config_validation_test.go +++ b/lib/instances/vm_config_validation_test.go @@ -17,34 +17,21 @@ func TestPlannedVMConfig(t *testing.T) { []VolumeAttachment{{Overlay: true}, {Overlay: false}}, true, 2, - true, ) assert.Equal(t, int64(1024), config.HotplugBytes) require.Len(t, config.Disks, 6, "three instance disks plus two overlay-volume disks plus one plain volume") require.Len(t, config.Networks, 1) require.Len(t, config.PCIDevices, 2) - assert.Equal(t, plannedVGPUDevicePath, config.VGPUDevicePath) assert.Equal(t, int64(3), config.VsockCID) } func TestPlannedVMConfigWithoutOptionalDevices(t *testing.T) { t.Parallel() - config := (&manager{}).plannedVMConfig(0, nil, false, 0, false) + config := (&manager{}).plannedVMConfig(0, nil, false, 0) require.Len(t, config.Disks, baseInstanceDiskCount) assert.Empty(t, config.Networks) assert.Empty(t, config.PCIDevices) - assert.Empty(t, config.VGPUDevicePath) -} - -func TestPlannedStoredVMConfigSeparatesVGPUFromPCIDevices(t *testing.T) { - t.Parallel() - config := (&manager{}).plannedStoredVMConfig(SnapshotKindStopped, StoredMetadata{ - Devices: []string{"pci-device"}, - GPUProfile: "gpu-profile", - }) - require.Len(t, config.PCIDevices, 1) - assert.Equal(t, plannedVGPUDevicePath, config.VGPUDevicePath) } func TestPlannedStoredVMConfigStandbyIgnoresLiveBalloonPolicy(t *testing.T) { diff --git a/lib/middleware/resolve_test.go b/lib/middleware/resolve_test.go index 607d8f5de..e9cfa477f 100644 --- a/lib/middleware/resolve_test.go +++ b/lib/middleware/resolve_test.go @@ -93,41 +93,6 @@ func TestResolveResource_URLDecodesImageName(t *testing.T) { } } -func TestResolveResource_ResolvesBuilderByID(t *testing.T) { - // Regression test: the path-dispatch switch must include a /builders/ - // case, otherwise the Builder resolver is never invoked and resolved - // builders are missing from request context (handlers then 500). - - resolver := &mockResolver{} - - errResponder := func(w http.ResponseWriter, err error, lookup string) { - w.WriteHeader(http.StatusNotFound) - } - - middleware := ResolveResource(Resolvers{ - Builder: resolver, - }, errResponder) - - r := chi.NewRouter() - r.With(middleware).Get("/builders/{id}", func(w http.ResponseWriter, r *http.Request) { - if GetResolvedBuilder[struct{}](r.Context()) == nil { - w.WriteHeader(http.StatusInternalServerError) - return - } - w.WriteHeader(http.StatusOK) - }) - - req := httptest.NewRequest(http.MethodGet, "/builders/bld_123", nil) - w := httptest.NewRecorder() - - r.ServeHTTP(w, req) - - require.Equal(t, http.StatusOK, w.Code, - "Expected resolved builder in context, got %d", w.Code) - assert.Equal(t, "bld_123", resolver.receivedName, - "Builder resolver was not invoked with the path ID") -} - func TestResolveResource_SkipsOnlyImageTagPosts(t *testing.T) { resolver := &mockResolver{} @@ -165,3 +130,38 @@ func TestResolveResource_SkipsOnlyImageTagPosts(t *testing.T) { "only the tag route should bypass image resolution") }) } + +func TestResolveResource_ResolvesBuilderByID(t *testing.T) { + // Regression test: the path-dispatch switch must include a /builders/ + // case, otherwise the Builder resolver is never invoked and resolved + // builders are missing from request context (handlers then 500). + + resolver := &mockResolver{} + + errResponder := func(w http.ResponseWriter, err error, lookup string) { + w.WriteHeader(http.StatusNotFound) + } + + middleware := ResolveResource(Resolvers{ + Builder: resolver, + }, errResponder) + + r := chi.NewRouter() + r.With(middleware).Get("/builders/{id}", func(w http.ResponseWriter, r *http.Request) { + if GetResolvedBuilder[struct{}](r.Context()) == nil { + w.WriteHeader(http.StatusInternalServerError) + return + } + w.WriteHeader(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/builders/bld_123", nil) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code, + "Expected resolved builder in context, got %d", w.Code) + assert.Equal(t, "bld_123", resolver.receivedName, + "Builder resolver was not invoked with the path ID") +} diff --git a/lib/paths/paths.go b/lib/paths/paths.go index ee68f531c..0d747b023 100644 --- a/lib/paths/paths.go +++ b/lib/paths/paths.go @@ -188,6 +188,11 @@ func (p *Paths) ImageLayerDir(layerHex string) string { return filepath.Join(p.ImageLayersDir(), layerHex) } +// ImageLayerArtifact returns the path to the default materialized layer artifact. +func (p *Paths) ImageLayerArtifact(layerHex string) string { + return p.ImageLayerArtifactForFormat(layerHex, "erofs") +} + // ImageLayerArtifactForFormat returns the path to a materialized layer artifact. func (p *Paths) ImageLayerArtifactForFormat(layerHex, format string) string { return filepath.Join(p.ImageLayerDir(layerHex), "layer."+format) diff --git a/lib/resources/gpu.go b/lib/resources/gpu.go index 78788412e..dc2437bbf 100644 --- a/lib/resources/gpu.go +++ b/lib/resources/gpu.go @@ -42,7 +42,7 @@ func getVGPUStatus() *GPUResourceStatus { // Count used VFs (those with mdevs) usedSlots := 0 for _, vf := range vfs { - if vf.Allocated { + if vf.HasMdev { usedSlots++ } } diff --git a/openapi.yaml b/openapi.yaml index 4d5830c2e..6b8d6c800 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -2693,12 +2693,6 @@ paths: application/json: schema: $ref: "#/components/schemas/Error" - 401: - description: Unauthorized - content: - application/json: - schema: - $ref: "#/components/schemas/Error" 500: description: Internal server error content: