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
24 changes: 17 additions & 7 deletions dcfs/app/sftp/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -333,7 +333,7 @@ def __init__(self, ops: Ops, path: str, mode: str, client_name: str):
# Read streaming state
self._read_stream: Optional[AsyncIterator[bytes]] = None
self._read_iter: Optional[AsyncIterator[bytes]] = None
self._read_pos = 0
self._buf_offset = 0
self._read_buf = bytearray()
self._read_lock = asyncio.Lock()
self._cached_attrs: Optional[asyncssh.SFTPAttrs] = None
Expand All @@ -343,15 +343,22 @@ async def read(self, offset: int, size: int) -> bytes:
raise asyncssh.SFTPPermissionDenied("File not open for reading")

async with self._read_lock:
# If we don't have a stream or it's at the wrong position, start a new one.
if self._read_stream is None or offset != self._read_pos:
buf_len = len(self._read_buf)
buf_end = self._buf_offset + buf_len

in_buffer = (
self._read_stream is not None
and self._buf_offset <= offset <= buf_end
)

if not in_buffer:
if self._read_stream is not None:
await cast(AsyncGenerator[bytes, None], self._read_stream).aclose()
self._read_stream = None
self._read_iter = None

self._read_buf = bytearray()
self._read_pos = offset
self._buf_offset = offset
self._read_stream = await self.ops.download(
self.path,
offset,
Expand All @@ -360,9 +367,12 @@ async def read(self, offset: int, size: int) -> bytes:
validate=False,
)
self._read_iter = self._read_stream.__aiter__()
else:
discard = offset - self._buf_offset
if discard > 0:
self._read_buf = self._read_buf[discard:]
self._buf_offset = offset

# it is impossible for self._read_iter to be None here given the logic above,
# but we cast to satisfy mypy.
it = cast(AsyncIterator[bytes], self._read_iter)
while len(self._read_buf) < size:
try:
Expand All @@ -373,7 +383,7 @@ async def read(self, offset: int, size: int) -> bytes:

data = self._read_buf[:size]
self._read_buf = self._read_buf[size:]
self._read_pos += len(data)
self._buf_offset += len(data)
return bytes(data)

async def write(self, offset: int, data: bytes) -> int:
Expand Down
90 changes: 90 additions & 0 deletions tests/test_sftp_handler.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,90 @@
import pytest
from unittest.mock import AsyncMock, MagicMock

import asyncssh

from dcfs.app.sftp.handler import DCFSSFTPBufferedFile


async def mock_download_gen(data_chunks):
for chunk in data_chunks:
yield chunk


@pytest.mark.asyncio
async def test_sftp_buffered_file_sequential_reads():
mock_ops = MagicMock()
# 3 chunks of 64KB
chunk1 = b"A" * 65536
chunk2 = b"B" * 65536
chunk3 = b"C" * 65536

mock_ops.download = AsyncMock(return_value=mock_download_gen([chunk1, chunk2, chunk3]))

file_handle = DCFSSFTPBufferedFile(mock_ops, "/test.txt", "r", "client1")

# Read 32KB at offset 0
data1 = await file_handle.read(0, 32768)
assert data1 == b"A" * 32768
assert mock_ops.download.call_count == 1

# Read next 32KB at offset 32768 (should hit buffer and NOT call download again)
data2 = await file_handle.read(32768, 32768)
assert data2 == b"A" * 32768
assert mock_ops.download.call_count == 1

# Read next 64KB at offset 65536
data3 = await file_handle.read(65536, 65536)
assert data3 == b"B" * 65536
assert mock_ops.download.call_count == 1

await file_handle.close()


@pytest.mark.asyncio
async def test_sftp_buffered_file_seek_out_of_order():
mock_ops = MagicMock()
chunk1 = b"0123456789"
chunk2 = b"ABCDEFGHIJ"

mock_ops.download = AsyncMock(
side_effect=[
mock_download_gen([chunk1]),
mock_download_gen([chunk2]),
]
)

file_handle = DCFSSFTPBufferedFile(mock_ops, "/test.txt", "r", "client1")

# Read offset 0, size 5
d1 = await file_handle.read(0, 5)
assert d1 == b"01234"
assert mock_ops.download.call_count == 1

# Jump forward out of current buffer (e.g. offset 100)
d2 = await file_handle.read(100, 5)
assert d2 == b"ABCDE"
assert mock_ops.download.call_count == 2
mock_ops.download.assert_called_with(
"/test.txt", 100, -1, "test.txt", validate=False
)

await file_handle.close()


@pytest.mark.asyncio
async def test_sftp_buffered_file_eof_and_mode_checks():
mock_ops = MagicMock()
mock_ops.download = AsyncMock(return_value=mock_download_gen([b"short"]))

file_handle = DCFSSFTPBufferedFile(mock_ops, "/test.txt", "r", "client1")

# Read more bytes than available
data = await file_handle.read(0, 100)
assert data == b"short"

# Write attempt on read mode should fail
with pytest.raises(asyncssh.SFTPPermissionDenied):
await file_handle.write(0, b"data")

await file_handle.close()
Loading