diff --git a/dcfs/app/sftp/handler.py b/dcfs/app/sftp/handler.py index 90ee2bc..52da962 100644 --- a/dcfs/app/sftp/handler.py +++ b/dcfs/app/sftp/handler.py @@ -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 @@ -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, @@ -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: @@ -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: diff --git a/tests/test_sftp_handler.py b/tests/test_sftp_handler.py new file mode 100644 index 0000000..9667416 --- /dev/null +++ b/tests/test_sftp_handler.py @@ -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()