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 }