# 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