From 8425e278c51eaec75de310cf3d13cf1e6ae0a05e Mon Sep 17 00:00:00 2001 From: Jamie Liu Date: Tue, 17 Sep 2024 21:17:36 -0700 Subject: [PATCH] segment: add Set.Remove[Full]RangeWith() These support the relatively common use case of removing all segments in a given range (unconditionally) but doing something with them before they're removed. This is always more compact, and may be slightly faster in some cases (every replaced loop calls Isolate per iteration, while RemoveRangeWith avoids redundant split checks between segments), at the cost of a direct function call. Also slightly optimize Set.LowerBoundSegmentSplitBefore() and Set.UpperBoundSegmentSplitAfter() by inlining LowerBoundSegment and UpperBoundSegment respectively; in the cases where Find() returns a GapIterator, the segment that is returned doesn't need to be split since it doesn't contain min/max respectively. PiperOrigin-RevId: 675824581 --- pkg/segment/set.go | 66 ++++++++++++++----- .../devices/tpuproxy/accel/gasket_ioctl.go | 11 +--- pkg/sentry/devices/tpuproxy/vfio/vfio_fd.go | 8 +-- pkg/sentry/fsimpl/gofer/regular_file.go | 7 +- pkg/sentry/fsutil/file_range_set.go | 7 +- pkg/sentry/mm/vma.go | 11 +--- 6 files changed, 61 insertions(+), 49 deletions(-) diff --git a/pkg/segment/set.go b/pkg/segment/set.go index f494e0dba..afa3a2285 100644 --- a/pkg/segment/set.go +++ b/pkg/segment/set.go @@ -603,21 +603,52 @@ func (s *Set) RemoveAll() { // if the caller needs to do additional work before removing each segment, // iterate segments and call Remove in a loop instead. func (s *Set) RemoveRange(r Range) GapIterator { - seg, gap := s.Find(r.Start) - if seg.Ok() { - seg = s.Isolate(seg, r) - gap = s.Remove(seg) - } - for seg = gap.NextSegment(); seg.Ok() && seg.Start() < r.End; seg = gap.NextSegment() { - seg = s.SplitAfter(seg, r.End) - gap = s.Remove(seg) - } - return gap + return s.RemoveRangeWith(r, nil) } // RemoveFullRange is equivalent to RemoveRange, except that if any key in the // given range does not correspond to a segment, RemoveFullRange panics. func (s *Set) RemoveFullRange(r Range) GapIterator { + return s.RemoveFullRangeWith(r, nil) +} + +// RemoveRangeWith removes all segments in the given range. An iterator to the +// newly formed gap is returned, and all existing iterators are invalidated. +// +// The function f is applied to each segment immediately before it is removed, +// in order of ascending keys. Segments that lie partially outside r are split +// before f is called, such that f only observes segments entirely within r. +// Non-empty gaps between segments are skipped. +// +// RemoveRangeWith searches the set to find segments to remove. If the caller +// already has an iterator to either end of the range of segments to remove, or +// if the caller needs to do additional work before removing each segment, +// iterate segments and call Remove in a loop instead. +// +// N.B. f must not invalidate iterators into s. +func (s *Set) RemoveRangeWith(r Range, f func(seg Iterator)) GapIterator { + seg, gap := s.Find(r.Start) + if seg.Ok() { + seg = s.Isolate(seg, r) + if f != nil { + f(seg) + } + gap = s.Remove(seg) + } + for seg = gap.NextSegment(); seg.Ok() && seg.Start() < r.End; seg = gap.NextSegment() { + seg = s.SplitAfter(seg, r.End) + if f != nil { + f(seg) + } + gap = s.Remove(seg) + } + return gap +} + +// RemoveFullRangeWith is equivalent to RemoveRangeWith, except that if any key +// in the given range does not correspond to a segment, RemoveFullRangeWith +// panics. +func (s *Set) RemoveFullRangeWith(r Range, f func(seg Iterator)) GapIterator { seg := s.FindSegment(r.Start) if !seg.Ok() { panic(fmt.Sprintf("missing segment at %v", r.Start)) @@ -625,6 +656,9 @@ func (s *Set) RemoveFullRange(r Range) GapIterator { seg = s.SplitBefore(seg, r.Start) for { seg = s.SplitAfter(seg, r.End) + if f != nil { + f(seg) + } end := seg.End() gap := s.Remove(seg) if r.End <= end { @@ -891,11 +925,11 @@ func (s *Set) Isolate(seg Iterator, r Range) Iterator { // LowerBoundSegmentSplitBefore provides an iterator to the first segment to be // mutated, suitable as the initial value for a loop variable. func (s *Set) LowerBoundSegmentSplitBefore(min Key) Iterator { - seg := s.LowerBoundSegment(min) + seg, gap := s.Find(min) if seg.Ok() { - seg = s.SplitBefore(seg, min) + return s.SplitBefore(seg, min) } - return seg + return gap.NextSegment() } // UpperBoundSegmentSplitAfter combines UpperBoundSegment and SplitAfter. @@ -905,11 +939,11 @@ func (s *Set) LowerBoundSegmentSplitBefore(min Key) Iterator { // UpperBoundSegmentSplitAfter provides an iterator to the first segment to be // mutated, suitable as the initial value for a loop variable. func (s *Set) UpperBoundSegmentSplitAfter(max Key) Iterator { - seg := s.UpperBoundSegment(max) + seg, gap := s.Find(max) if seg.Ok() { - seg = s.SplitAfter(seg, max) + return s.SplitAfter(seg, max) } - return seg + return gap.PrevSegment() } // VisitRange applies the function f to all segments intersecting the range r, diff --git a/pkg/sentry/devices/tpuproxy/accel/gasket_ioctl.go b/pkg/sentry/devices/tpuproxy/accel/gasket_ioctl.go index c399fc413..42deb475e 100644 --- a/pkg/sentry/devices/tpuproxy/accel/gasket_ioctl.go +++ b/pkg/sentry/devices/tpuproxy/accel/gasket_ioctl.go @@ -159,14 +159,9 @@ func gasketUnmapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, f defer fd.device.mu.Unlock() s := &fd.device.devAddrSet r := DevAddrRange{userIoctlParams.DeviceAddress, userIoctlParams.DeviceAddress + userIoctlParams.Size} - seg := s.LowerBoundSegment(r.Start) - for seg.Ok() && seg.Start() < r.End { - seg = s.Isolate(seg, r) - v := seg.Value() - mm.Unpin([]mm.PinnedRange{v.pinnedRange}) - gap := s.Remove(seg) - seg = gap.NextSegment() - } + s.RemoveRangeWith(r, func(seg DevAddrIterator) { + mm.Unpin([]mm.PinnedRange{seg.ValuePtr().pinnedRange}) + }) return n, nil } diff --git a/pkg/sentry/devices/tpuproxy/vfio/vfio_fd.go b/pkg/sentry/devices/tpuproxy/vfio/vfio_fd.go index b2269f984..0174882b0 100644 --- a/pkg/sentry/devices/tpuproxy/vfio/vfio_fd.go +++ b/pkg/sentry/devices/tpuproxy/vfio/vfio_fd.go @@ -246,13 +246,9 @@ func (fd *vfioFD) iommuUnmapDma(ctx context.Context, t *kernel.Task, arg hostarc func (fd *vfioFD) unpinRange(r DevAddrRange) { fd.mu.Lock() defer fd.mu.Unlock() - seg := fd.devAddrSet.LowerBoundSegment(r.Start) - for seg.Ok() && seg.Start() < r.End { - seg = fd.devAddrSet.Isolate(seg, r) + fd.devAddrSet.RemoveRangeWith(r, func(seg DevAddrIterator) { mm.Unpin([]mm.PinnedRange{seg.Value()}) - gap := fd.devAddrSet.Remove(seg) - seg = gap.NextSegment() - } + }) } // VFIO extension. diff --git a/pkg/sentry/fsimpl/gofer/regular_file.go b/pkg/sentry/fsimpl/gofer/regular_file.go index cd2d2e3a1..71754099b 100644 --- a/pkg/sentry/fsimpl/gofer/regular_file.go +++ b/pkg/sentry/fsimpl/gofer/regular_file.go @@ -295,12 +295,9 @@ func (fd *regularFileFD) writeCache(ctx context.Context, d *dentry, offset int64 var freed []memmap.FileRange d.dataMu.Lock() - cseg := d.cache.LowerBoundSegment(mr.Start) - for cseg.Ok() && cseg.Start() < mr.End { - cseg = d.cache.Isolate(cseg, mr) + d.cache.RemoveRangeWith(mr, func(cseg fsutil.FileRangeIterator) { freed = append(freed, memmap.FileRange{cseg.Value(), cseg.Value() + cseg.Range().Length()}) - cseg = d.cache.Remove(cseg).NextSegment() - } + }) d.dataMu.Unlock() // Invalidate mappings of removed pages. diff --git a/pkg/sentry/fsutil/file_range_set.go b/pkg/sentry/fsutil/file_range_set.go index c79a834d2..5febdfa98 100644 --- a/pkg/sentry/fsutil/file_range_set.go +++ b/pkg/sentry/fsutil/file_range_set.go @@ -183,12 +183,9 @@ func (s *FileRangeSet) Fill(ctx context.Context, required, optional memmap.Mappa // // Preconditions: mr must be page-aligned. func (s *FileRangeSet) Drop(mr memmap.MappableRange, mf *pgalloc.MemoryFile) { - seg := s.LowerBoundSegment(mr.Start) - for seg.Ok() && seg.Start() < mr.End { - seg = s.Isolate(seg, mr) + s.RemoveRangeWith(mr, func(seg FileRangeIterator) { mf.DecRef(seg.FileRange()) - seg = s.Remove(seg).NextSegment() - } + }) } // DropAll removes all segments in mr, freeing the corresponding diff --git a/pkg/sentry/mm/vma.go b/pkg/sentry/mm/vma.go index 13ee8f40d..4811085b5 100644 --- a/pkg/sentry/mm/vma.go +++ b/pkg/sentry/mm/vma.go @@ -401,12 +401,7 @@ func (mm *MemoryManager) removeVMAsLocked(ctx context.Context, ar hostarch.AddrR panic(fmt.Sprintf("invalid ar: %v", ar)) } } - vseg, vgap := mm.vmas.Find(ar.Start) - if vgap.Ok() { - vseg = vgap.NextSegment() - } - for vseg.Ok() && vseg.Start() < ar.End { - vseg = mm.vmas.Isolate(vseg, ar) + vgap := mm.vmas.RemoveRangeWith(ar, func(vseg vmaIterator) { vmaAR := vseg.Range() vma := vseg.ValuePtr() if vma.mappable != nil { @@ -422,9 +417,7 @@ func (mm *MemoryManager) removeVMAsLocked(ctx context.Context, ar hostarch.AddrR if vma.mlockMode != memmap.MLockNone { mm.lockedAS -= uint64(vmaAR.Length()) } - vgap = mm.vmas.Remove(vseg) - vseg = vgap.NextSegment() - } + }) return vgap, droppedIDs }