Source code for nvflare.fuel.utils.buffer_list

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.


[docs] class ConsumableBufferList(list): """Buffer list whose owner permits readers to release consumed entries."""
[docs] class BufferList: """A buffer list that can be treated as a single buffer""" def __init__(self, buf_list: list): self.buf_list = buf_list self.start_offset = 0
[docs] def get_size(self): if self.buf_list: size = sum(len(buf) for buf in self.buf_list) else: size = 0 return size
[docs] def get_list(self): return self.buf_list
[docs] def append(self, buf: bytes): if not self.buf_list: self.buf_list = [] self.buf_list.append(buf)
[docs] def read(self, start: int, end: int): buffer = None view_start = 0 pos = 0 for view in self.buf_list: view_end = view_start + len(view) if view_start <= start < view_end and end <= view_end: return view[start - view_start : end - view_start] buf_start = start + pos if buf_start < view_end: if not buffer: buffer = bytearray(end - start) remaining = min(end, view_end) - buf_start view_pos = buf_start - view_start buffer[pos : pos + remaining] = view[view_pos : view_pos + remaining] pos = pos + remaining if view_end >= end: break view_start = view_end return buffer
[docs] def read_bytes(self, start: int, end: int) -> bytes: """Read a range into one immutable bytes allocation. Unlike ``read()``, this method does not first assemble a multi-buffer range into a bytearray. This matters for large streamed FOBS sections: callers that require bytes would otherwise copy the complete range once into a bytearray and a second time into bytes. """ if start < 0: raise ValueError(f"start must be non-negative, got {start}") if start < self.start_offset: raise ValueError(f"start {start} precedes discarded data at offset {self.start_offset}") if end < start: raise ValueError(f"end {end} must not be less than start {start}") available_end = self.start_offset + self.get_size() if end > available_end: raise ValueError(f"end {end} exceeds available data ending at {available_end}") if start == end: return b"" parts = [] view_start = self.start_offset for buffer in self.buf_list or []: view_end = view_start + len(buffer) if view_end <= start: view_start = view_end continue if view_start >= end: break part_start = max(start, view_start) - view_start part_end = min(end, view_end) - view_start parts.append(memoryview(buffer)[part_start:part_end]) view_start = view_end if not parts: return b"" if len(parts) == 1: return parts[0].tobytes() return b"".join(parts)
[docs] def discard_before(self, offset: int) -> None: """Release complete buffers ending at or before an absolute offset.""" if offset < self.start_offset: raise ValueError(f"offset {offset} precedes discarded data at {self.start_offset}") discard_count = 0 discarded_size = 0 for buffer in self.buf_list or []: buffer_end = self.start_offset + discarded_size + len(buffer) if buffer_end > offset: break discard_count += 1 discarded_size += len(buffer) if discard_count: del self.buf_list[:discard_count] self.start_offset += discarded_size
[docs] def flatten(self): size = self.get_size() if not size: return None result = bytearray(size) start = 0 for b in self.buf_list: size = len(b) result[start : start + size] = b start += size return result