diff --git a/dcfs/app/sftp/handler.py b/dcfs/app/sftp/handler.py index 5f6b64b..674e051 100644 --- a/dcfs/app/sftp/handler.py +++ b/dcfs/app/sftp/handler.py @@ -322,6 +322,9 @@ async def fstat(self) -> asyncssh.SFTPAttrs: class DCFSSFTPBufferedFile(DCFSSFTPFileBase): + MAX_FORWARD_SKIP = 2 * 1024 * 1024 # 2 MB forward skip + MAX_BACKWARD_RETAIN = 4 * 1024 * 1024 # 4 MB backward retain + def __init__(self, ops: Ops, path: str, mode: str, client_name: str): self.ops = ops self.path = path @@ -336,6 +339,7 @@ def __init__(self, ops: Ops, path: str, mode: str, client_name: str): self._read_buf = bytearray() self._read_lock = asyncio.Lock() self._cached_attrs: Optional[asyncssh.SFTPAttrs] = None + self._highest_offset = 0 # Prefetch state self._prefetch_queue: Optional[asyncio.Queue[Optional[Any]]] = None @@ -376,43 +380,41 @@ async def _stop_prefetch(self) -> None: self._prefetch_queue = None self._prefetch_eof = False + async def _start_prefetch(self, offset: int) -> None: + self._read_buf = bytearray() + self._buf_offset = offset + self._highest_offset = offset + self._read_stream = await self.ops.download( + self.path, + offset, + -1, + os.path.basename(self.path), + validate=False, + ) + self._prefetch_queue = asyncio.Queue(maxsize=64) + self._prefetch_eof = False + self._prefetch_task = asyncio.create_task( + self._run_prefetch(self._read_stream, self._prefetch_queue) + ) + async def read(self, offset: int, size: int) -> bytes: if "r" not in self.mode: raise asyncssh.SFTPPermissionDenied("File not open for reading") async with self._read_lock: - buf_len = len(self._read_buf) - buf_end = self._buf_offset + buf_len + buf_end = self._buf_offset + len(self._read_buf) - in_buffer = ( + can_reuse_stream = ( self._read_stream is not None - and self._buf_offset <= offset <= buf_end + and self._buf_offset <= offset <= buf_end + self.MAX_FORWARD_SKIP ) - if not in_buffer: + if not can_reuse_stream: await self._stop_prefetch() + await self._start_prefetch(offset) - self._read_buf = bytearray() - self._buf_offset = offset - self._read_stream = await self.ops.download( - self.path, - offset, - -1, - os.path.basename(self.path), - validate=False, - ) - self._prefetch_queue = asyncio.Queue(maxsize=64) - self._prefetch_eof = False - self._prefetch_task = asyncio.create_task( - self._run_prefetch(self._read_stream, self._prefetch_queue) - ) - else: - discard = offset - self._buf_offset - if discard > 0: - self._read_buf = self._read_buf[discard:] - self._buf_offset = offset - - while len(self._read_buf) < size and not self._prefetch_eof: + target_end = offset + size + while (self._buf_offset + len(self._read_buf) < target_end) and not self._prefetch_eof: if self._prefetch_queue is None: break item = await self._prefetch_queue.get() @@ -424,10 +426,39 @@ async def read(self, offset: int, size: int) -> bytes: raise item self._read_buf.extend(item) - data = self._read_buf[:size] - self._read_buf = self._read_buf[size:] - self._buf_offset += len(data) - return bytes(data) + # If stream reached EOF before reaching offset, restart stream at offset + if self._prefetch_eof and (self._buf_offset + len(self._read_buf) <= offset) and size > 0: + await self._stop_prefetch() + await self._start_prefetch(offset) + while (self._buf_offset + len(self._read_buf) < target_end) and not self._prefetch_eof: + if self._prefetch_queue is None: + break + item = await self._prefetch_queue.get() + if item is None: + self._prefetch_eof = True + break + if isinstance(item, Exception): + self._prefetch_eof = True + raise item + self._read_buf.extend(item) + + rel_offset = offset - self._buf_offset + if rel_offset >= 0 and rel_offset < len(self._read_buf): + data = bytes(self._read_buf[rel_offset : rel_offset + size]) + else: + data = b"" + + self._highest_offset = max(self._highest_offset, offset + len(data)) + + # Prune buffer behind prune_target to keep memory bounded + prune_target = self._highest_offset - self.MAX_BACKWARD_RETAIN + if prune_target > self._buf_offset: + discard = min(prune_target - self._buf_offset, len(self._read_buf)) + if discard > 0: + self._read_buf = self._read_buf[discard:] + self._buf_offset += discard + + return data async def write(self, offset: int, data: bytes) -> int: if "w" not in self.mode: diff --git a/dcfs/core/api/message/__init__.py b/dcfs/core/api/message/__init__.py index f355151..da59b5f 100644 --- a/dcfs/core/api/message/__init__.py +++ b/dcfs/core/api/message/__init__.py @@ -1,6 +1,6 @@ import asyncio import logging -from typing import Iterable, Iterator, List +from typing import AsyncIterator, Iterable, Iterator, List from pyrate_limiter import Duration, InMemoryBucket, Limiter, Rate @@ -19,7 +19,6 @@ SendFileReq, SendTextReq, ) -from dcfs.utils.chained_async_iterator import ChainedAsyncIterator from dcfs.utils.others import exclude_none, is_big_file from .message_broker import MessageBroker @@ -180,7 +179,9 @@ async def download_file_parallel(self, message_id: int, begin: int, end: int): # Split the range into concurrent sub-range downloads so we can # utilise CDN bandwidth better for large single-part files. n = 4 - tasks = [ + sub_ranges = list(self.split_download_tasks(begin, end, n)) + + resps = await asyncio.gather(*[ self.discord_api.next_bot.download_file( DownloadFileReq( chat=self.private_file_channel, @@ -190,12 +191,50 @@ async def download_file_parallel(self, message_id: int, begin: int, end: int): end=e, ) ) - for b, e in self.split_download_tasks(begin, end, n) + for b, e in sub_ranges + ]) + + queues: list[asyncio.Queue[object]] = [ + asyncio.Queue(maxsize=32) for _ in range(n) ] - res = [t.chunks for t in await asyncio.gather(*tasks)] + async def _producer( + chunks_iter: Iterator[bytes] | AsyncIterator[bytes], + q: asyncio.Queue[object], + ) -> None: + try: + if hasattr(chunks_iter, "__anext__"): + async for chunk in chunks_iter: # type: ignore[union-attr] + await q.put(chunk) + else: + for chunk in chunks_iter: # type: ignore[union-attr] + await q.put(chunk) + await q.put(None) + except Exception as ex: + await q.put(ex) + + producer_tasks = [ + asyncio.create_task(_producer(resp.chunks, q)) + for resp, q in zip(resps, queues) + ] + + async def _parallel_chunks(): + try: + for q in queues: + while True: + item = await q.get() + if item is None: + break + if isinstance(item, Exception): + raise item + yield item + finally: + for task in producer_tasks: + if not task.done(): + task.cancel() + return DownloadFileResp( - chunks=ChainedAsyncIterator(res), size=self._size(begin, end) + chunks=_parallel_chunks(), size=self._size(begin, end) ) async def download_file( diff --git a/tests/dcfs/core/api/message/test_parallel.py b/tests/dcfs/core/api/message/test_parallel.py new file mode 100644 index 0000000..d4aa504 --- /dev/null +++ b/tests/dcfs/core/api/message/test_parallel.py @@ -0,0 +1,42 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from dcfs.core.api.message import MessageApi +from dcfs.reqres import DownloadFileResp + + +async def mock_chunks(data): + for chunk in data: + yield chunk + + +@pytest.mark.asyncio +async def test_download_file_parallel(): + discord_api = MagicMock() + bot = AsyncMock() + discord_api.next_bot = bot + message_api = MessageApi(discord_api, private_file_channel=123) + + # Return different chunk data for each sub-range request + bot.download_file.side_effect = [ + DownloadFileResp(chunks=mock_chunks([b"part1_chunk1", b"part1_chunk2"]), size=10), + DownloadFileResp(chunks=mock_chunks([b"part2_chunk1"]), size=10), + DownloadFileResp(chunks=mock_chunks([b"part3_chunk1"]), size=10), + DownloadFileResp(chunks=mock_chunks([b"part4_chunk1"]), size=10), + ] + + resp = await message_api.download_file_parallel(message_id=999, begin=0, end=39) + + chunks = [] + async for chunk in resp.chunks: + chunks.append(chunk) + + assert chunks == [ + b"part1_chunk1", + b"part1_chunk2", + b"part2_chunk1", + b"part3_chunk1", + b"part4_chunk1", + ] + assert bot.download_file.call_count == 4 diff --git a/tests/test_sftp_handler.py b/tests/test_sftp_handler.py index dbed880..3365cfb 100644 --- a/tests/test_sftp_handler.py +++ b/tests/test_sftp_handler.py @@ -106,3 +106,32 @@ async def mock_error_gen(): await file_handle.read(0, 100) await file_handle.close() + + +@pytest.mark.asyncio +async def test_sftp_buffered_file_pipelined_out_of_order_reads(): + mock_ops = MagicMock() + chunk1 = b"A" * 65536 # bytes 0..65535 + chunk2 = b"B" * 65536 # bytes 65536..131071 + + mock_ops.download = AsyncMock(return_value=mock_download_gen([chunk1, chunk2])) + + file_handle = DCFSSFTPBufferedFile(mock_ops, "/pipelined.txt", "r", "client1") + + # Read offset 0 first so download stream starts at 0 + data1 = await file_handle.read(0, 32768) + assert data1 == b"A" * 32768 + assert mock_ops.download.call_count == 1 + + # Pipelined request for offset 65536 arrives BEFORE request for offset 32768 + data3 = await file_handle.read(65536, 32768) + assert data3 == b"B" * 32768 + assert mock_ops.download.call_count == 1 + + # Request for offset 32768 arrives (out of order, behind current 65536 offset) + data2 = await file_handle.read(32768, 32768) + assert data2 == b"A" * 32768 + # Should STILL be 1 download call because byte range was retained in buffer! + assert mock_ops.download.call_count == 1 + + await file_handle.close()