Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 61 additions & 30 deletions dcfs/app/sftp/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand All @@ -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:
Expand Down
51 changes: 45 additions & 6 deletions dcfs/core/api/message/__init__.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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(
Expand Down
42 changes: 42 additions & 0 deletions tests/dcfs/core/api/message/test_parallel.py
Original file line number Diff line number Diff line change
@@ -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
29 changes: 29 additions & 0 deletions tests/test_sftp_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading