From fd1b34358448e58bca19d8b67f33fa9e62372dff Mon Sep 17 00:00:00 2001 From: goingforstudying-ctrl Date: Sat, 6 Jun 2026 06:16:27 -0400 Subject: [PATCH 1/2] fix(async): apply socket_timeout per read in async parsers The async client was wrapping the entire parser.read_response() call in async_timeout, which meant socket_timeout applied to the whole response. The sync client applies it per socket recv() instead. Change the async Python parsers (RESP2/RESP3) and the Hiredis async parser to accept a timeout parameter and apply it around each individual stream read. Connection.read_response now passes the timeout down to the parser rather than wrapping the parser call. Fixes #3454 --- redis/_parsers/base.py | 80 +++-- redis/_parsers/hiredis.py | 25 +- redis/_parsers/resp2.py | 20 +- redis/_parsers/resp3.py | 51 +++- redis/asyncio/connection.py | 17 +- .../test_async_timeout_per_read.py | 278 ++++++++++++++++++ 6 files changed, 416 insertions(+), 55 deletions(-) create mode 100644 tests/test_asyncio/test_async_timeout_per_read.py diff --git a/redis/_parsers/base.py b/redis/_parsers/base.py index b4b21f63a7..9fc9c6210f 100644 --- a/redis/_parsers/base.py +++ b/redis/_parsers/base.py @@ -1,8 +1,14 @@ import logging +import sys from abc import ABC, abstractmethod from asyncio import IncompleteReadError, StreamReader from typing import Awaitable, Callable, List, Optional, Protocol, Union +if sys.version_info >= (3, 11, 3): + from asyncio import timeout as async_timeout +else: + from async_timeout import timeout as async_timeout + from redis.maint_notifications import ( MaintenanceNotification, NodeFailedOverNotification, @@ -13,7 +19,6 @@ OSSNodeMigratedNotification, OSSNodeMigratingNotification, ) -from redis.utils import deprecated_function, safe_str from ..exceptions import ( AskError, @@ -36,6 +41,7 @@ TryAgainError, ) from ..typing import EncodableT +from ..utils import SENTINEL, deprecated_function, safe_str from .encoders import Encoder from .socket import SERVER_CLOSED_CONNECTION_ERROR, SocketBuffer @@ -183,7 +189,10 @@ async def can_read(self) -> bool: pass async def read_response( - self, disable_decoding: bool = False + self, + disable_decoding: bool = False, + push_request: bool = False, + timeout: Union[float, object] = SENTINEL, ) -> Union[EncodableT, ResponseError, None, List[EncodableT]]: raise NotImplementedError() @@ -545,40 +554,65 @@ async def can_read(self) -> bool: # parser and fail loudly if the private buffer API changes. return bool(self._stream._buffer) or self._stream.at_eof() - async def _read(self, length: int) -> bytes: + async def _read_from_stream( + self, timeout: Union[float, object] = SENTINEL, max_bytes: int = 0 + ) -> bytes: + """ + Read the next chunk from the underlying stream with an optional + per-read timeout. This mirrors the sync client's per-recv timeout + semantics: each individual socket read gets its own timeout window. + + ``max_bytes`` limits how many bytes may be returned. When 0, the + parser's ``_read_size`` is used. + """ + stream = self._stream + if stream is None: + raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) + size = max_bytes if max_bytes > 0 else self._read_size + if timeout is not SENTINEL and isinstance(timeout, (int, float)): + async with async_timeout(timeout): + return await stream.read(size) + return await stream.read(size) + + async def _read( + self, length: int, timeout: Union[float, object] = SENTINEL + ) -> bytes: """ Read `length` bytes of data. These are assumed to be followed by a '\r\n' terminator which is subsequently discarded. """ want = length + 2 end = self._pos + want - if len(self._buffer) >= end: - result = self._buffer[self._pos : end - 2] - else: - tail = self._buffer[self._pos :] + while len(self._buffer) < end: + need = end - len(self._buffer) try: - data = await self._stream.readexactly(want - len(tail)) + chunk = await self._read_from_stream( + timeout=timeout, max_bytes=min(need, self._read_size) + ) except IncompleteReadError as error: raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) from error - result = (tail + data)[:-2] - self._chunks.append(data) - self._pos += want + if not chunk: + raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) + self._buffer += chunk + result = self._buffer[self._pos : end - 2] + self._pos = end return result - async def _readline(self) -> bytes: + async def _readline(self, timeout: Union[float, object] = SENTINEL) -> bytes: """ read an unknown number of bytes up to the next '\r\n' line separator, which is discarded. """ - found = self._buffer.find(b"\r\n", self._pos) - if found >= 0: - result = self._buffer[self._pos : found] - else: - tail = self._buffer[self._pos :] - data = await self._stream.readline() - if not data.endswith(b"\r\n"): + while True: + found = self._buffer.find(b"\r\n", self._pos) + if found >= 0: + result = self._buffer[self._pos : found] + self._pos = found + 2 + return result + try: + chunk = await self._read_from_stream(timeout=timeout) + except IncompleteReadError as error: + raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) from error + if not chunk: raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) - result = (tail + data)[:-2] - self._chunks.append(data) - self._pos += len(result) + 2 - return result + self._buffer += chunk diff --git a/redis/_parsers/hiredis.py b/redis/_parsers/hiredis.py index 6ad4db8789..a630e70d1e 100644 --- a/redis/_parsers/hiredis.py +++ b/redis/_parsers/hiredis.py @@ -1,9 +1,15 @@ import select import selectors import socket +import sys from logging import getLogger from typing import Callable, List, Optional, TypedDict, Union +if sys.version_info >= (3, 11, 3): + from asyncio import timeout as async_timeout +else: + from async_timeout import timeout as async_timeout + from ..exceptions import ConnectionError, InvalidResponse, RedisError, TimeoutError from ..typing import EncodableT from ..utils import HIREDIS_AVAILABLE, SENTINEL, deprecated_function @@ -265,8 +271,12 @@ async def can_read(self) -> bool: # with a real StreamReader guard this private buffer API in CI. return bool(self._stream._buffer) - async def read_from_socket(self): - buffer = await self._stream.read(self._read_size) + async def read_from_socket(self, timeout: Union[float, object] = SENTINEL): + if timeout is not SENTINEL: + async with async_timeout(timeout): + buffer = await self._stream.read(self._read_size) + else: + buffer = await self._stream.read(self._read_size) if not buffer or not isinstance(buffer, bytes): raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) from None self._reader.feed(buffer) @@ -275,7 +285,10 @@ async def read_from_socket(self): return True async def read_response( - self, disable_decoding: bool = False, push_request: bool = False + self, + disable_decoding: bool = False, + push_request: bool = False, + timeout: Union[float, object] = SENTINEL, ) -> Union[EncodableT, List[EncodableT]]: # If `on_disconnect()` has been called, prohibit any more reads # even if they could happen because data might be present. @@ -289,7 +302,7 @@ async def read_response( response = self._reader.gets() while response is NOT_ENOUGH_DATA: - await self.read_from_socket() + await self.read_from_socket(timeout=timeout) if disable_decoding: response = self._reader.gets(False) else: @@ -306,7 +319,9 @@ async def read_response( response = await self.handle_push_response(response) if not push_request: return await self.read_response( - disable_decoding=disable_decoding, push_request=push_request + disable_decoding=disable_decoding, + push_request=push_request, + timeout=timeout, ) else: return response diff --git a/redis/_parsers/resp2.py b/redis/_parsers/resp2.py index 26701157f0..2c1eab24b5 100644 --- a/redis/_parsers/resp2.py +++ b/redis/_parsers/resp2.py @@ -78,7 +78,9 @@ def _read_response( class _AsyncRESP2Parser(_AsyncRESPBase): """Async class for the RESP2 protocol""" - async def read_response(self, disable_decoding: bool = False): + async def read_response( + self, disable_decoding: bool = False, timeout: Union[float, object] = SENTINEL + ): if not self._connected: raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) if self._chunks: @@ -86,15 +88,17 @@ async def read_response(self, disable_decoding: bool = False): self._buffer += b"".join(self._chunks) self._chunks.clear() self._pos = 0 - response = await self._read_response(disable_decoding=disable_decoding) + response = await self._read_response( + disable_decoding=disable_decoding, timeout=timeout + ) # Successfully parsing a response allows us to clear our parsing buffer self._clear() return response async def _read_response( - self, disable_decoding: bool = False + self, disable_decoding: bool = False, timeout: Union[float, object] = SENTINEL ) -> Union[EncodableT, ResponseError, None]: - raw = await self._readline() + raw = await self._readline(timeout=timeout) response: Any byte, response = raw[:1], raw[1:] @@ -122,13 +126,17 @@ async def _read_response( elif byte == b"$" and response == b"-1": return None elif byte == b"$": - response = await self._read(int(response)) + response = await self._read(int(response), timeout=timeout) # multi-bulk response elif byte == b"*" and response == b"-1": return None elif byte == b"*": response = [ - (await self._read_response(disable_decoding)) + ( + await self._read_response( + disable_decoding=disable_decoding, timeout=timeout + ) + ) for _ in range(int(response)) # noqa ] else: diff --git a/redis/_parsers/resp3.py b/redis/_parsers/resp3.py index 8a9e41e27e..6479e32339 100644 --- a/redis/_parsers/resp3.py +++ b/redis/_parsers/resp3.py @@ -153,6 +153,7 @@ def _read_response( return self._read_response( disable_decoding=disable_decoding, push_request=push_request, + timeout=timeout, ) else: raise InvalidResponse(f"Protocol Error: {raw!r}") @@ -175,7 +176,10 @@ async def handle_pubsub_push_response(self, response): return response async def read_response( - self, disable_decoding: bool = False, push_request: bool = False + self, + disable_decoding: bool = False, + push_request: bool = False, + timeout: Union[float, object] = SENTINEL, ): if self._chunks: # augment parsing buffer with previously read data @@ -183,18 +187,23 @@ async def read_response( self._chunks.clear() self._pos = 0 response = await self._read_response( - disable_decoding=disable_decoding, push_request=push_request + disable_decoding=disable_decoding, + push_request=push_request, + timeout=timeout, ) # Successfully parsing a response allows us to clear our parsing buffer self._clear() return response async def _read_response( - self, disable_decoding: bool = False, push_request: bool = False + self, + disable_decoding: bool = False, + push_request: bool = False, + timeout: Union[float, object] = SENTINEL, ) -> Union[EncodableT, ResponseError, None]: if not self._stream or not self.encoder: raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) - raw = await self._readline() + raw = await self._readline(timeout=timeout) response: Any byte, response = raw[:1], raw[1:] @@ -204,7 +213,7 @@ async def _read_response( # server returned an error if byte in (b"-", b"!"): if byte == b"!": - response = await self._read(int(response)) + response = await self._read(int(response), timeout=timeout) response = response.decode("utf-8", errors="replace") error = self.parse_error(response) # if the error is a ConnectionError, raise immediately so the user @@ -234,14 +243,18 @@ async def _read_response( return response == b"t" # bulk response elif byte == b"$": - response = await self._read(int(response)) + response = await self._read(int(response), timeout=timeout) # verbatim string response elif byte == b"=": - response = (await self._read(int(response)))[4:] + response = (await self._read(int(response), timeout=timeout))[4:] # array response elif byte == b"*": response = [ - (await self._read_response(disable_decoding=disable_decoding)) + ( + await self._read_response( + disable_decoding=disable_decoding, timeout=timeout + ) + ) for _ in range(int(response)) ] # set response @@ -249,7 +262,11 @@ async def _read_response( # redis can return unhashable types (like dict) in a set, # so we always convert to a list, to have predictable return types response = [ - (await self._read_response(disable_decoding=disable_decoding)) + ( + await self._read_response( + disable_decoding=disable_decoding, timeout=timeout + ) + ) for _ in range(int(response)) ] # map response @@ -259,9 +276,13 @@ async def _read_response( # became defined to be left-right in version 3.8 resp_dict = {} for _ in range(int(response)): - key = await self._read_response(disable_decoding=disable_decoding) + key = await self._read_response( + disable_decoding=disable_decoding, timeout=timeout + ) resp_dict[key] = await self._read_response( - disable_decoding=disable_decoding, push_request=push_request + disable_decoding=disable_decoding, + push_request=push_request, + timeout=timeout, ) response = resp_dict # push response @@ -269,7 +290,9 @@ async def _read_response( response = [ ( await self._read_response( - disable_decoding=disable_decoding, push_request=push_request + disable_decoding=disable_decoding, + push_request=push_request, + timeout=timeout, ) ) for _ in range(int(response)) @@ -277,7 +300,9 @@ async def _read_response( response = await self.handle_push_response(response) if not push_request: return await self._read_response( - disable_decoding=disable_decoding, push_request=push_request + disable_decoding=disable_decoding, + push_request=push_request, + timeout=timeout, ) else: return response diff --git a/redis/asyncio/connection.py b/redis/asyncio/connection.py index 781239df9c..9cbeec8a5c 100644 --- a/redis/asyncio/connection.py +++ b/redis/asyncio/connection.py @@ -779,15 +779,16 @@ async def read_response( host_error = self._host_error() try: if read_timeout is not None and self.protocol in ["3", 3]: - async with async_timeout(read_timeout): - response = await self._parser.read_response( - disable_decoding=disable_decoding, push_request=push_request - ) + response = await self._parser.read_response( + disable_decoding=disable_decoding, + push_request=push_request, + timeout=read_timeout, + ) elif read_timeout is not None: - async with async_timeout(read_timeout): - response = await self._parser.read_response( - disable_decoding=disable_decoding - ) + response = await self._parser.read_response( + disable_decoding=disable_decoding, + timeout=read_timeout, + ) elif self.protocol in ["3", 3]: response = await self._parser.read_response( disable_decoding=disable_decoding, push_request=push_request diff --git a/tests/test_asyncio/test_async_timeout_per_read.py b/tests/test_asyncio/test_async_timeout_per_read.py new file mode 100644 index 0000000000..679e9b5cd2 --- /dev/null +++ b/tests/test_asyncio/test_async_timeout_per_read.py @@ -0,0 +1,278 @@ +"""Tests for per-read socket_timeout semantics on async connections. + +Issue: redis/redis-py#3454 — socket_timeout on async connection should apply +per individual socket read, matching the sync client behavior, rather than to +the entire response. +""" + +import asyncio + +import pytest + +from redis._parsers import _AsyncRESP2Parser, _AsyncRESP3Parser +from redis.asyncio.connection import Connection +from redis.utils import HIREDIS_AVAILABLE + +if HIREDIS_AVAILABLE: + from redis._parsers import _AsyncHiredisParser + from redis._parsers.hiredis import NOT_ENOUGH_DATA + + +class SlowChunkStream: + """Mock StreamReader that returns data one chunk at a time with delays.""" + + def __init__(self, chunks, delay_between_chunks): + self._chunks = list(chunks) + self._delay = delay_between_chunks + self._buffer = b"" + self._pos = 0 + self._chunk_index = 0 + + def at_eof(self): + return self._chunk_index >= len(self._chunks) and self._pos >= len(self._buffer) + + async def read(self, _want): + if self._pos >= len(self._buffer): + if self._chunk_index >= len(self._chunks): + return b"" + if self._delay: + await asyncio.sleep(self._delay) + self._buffer = self._chunks[self._chunk_index] + self._chunk_index += 1 + self._pos = 0 + result = self._buffer[self._pos :] + self._pos += len(result) + return result + + async def readline(self): + if self._pos >= len(self._buffer): + if self._chunk_index >= len(self._chunks): + return b"" + if self._delay: + await asyncio.sleep(self._delay) + self._buffer = self._chunks[self._chunk_index] + self._chunk_index += 1 + self._pos = 0 + nl = self._buffer.find(b"\n", self._pos) + if nl < 0: + result = self._buffer[self._pos :] + self._pos = len(self._buffer) + return result + result = self._buffer[self._pos : nl + 1] + self._pos = nl + 1 + return result + + async def readexactly(self, length): + result = bytearray() + while len(result) < length: + if self._pos >= len(self._buffer): + if self._chunk_index >= len(self._chunks): + raise asyncio.IncompleteReadError(bytes(result), length) + if self._delay: + await asyncio.sleep(self._delay) + self._buffer = self._chunks[self._chunk_index] + self._chunk_index += 1 + self._pos = 0 + take = min(length - len(result), len(self._buffer) - self._pos) + result.extend(self._buffer[self._pos : self._pos + take]) + self._pos += take + return bytes(result) + + +class _DummyEncoder: + decode_responses = False + encoding = "utf-8" + encoding_errors = "strict" + + def decode(self, value): + if isinstance(value, bytes): + return value.decode(self.encoding, self.encoding_errors) + if isinstance(value, list): + return [self.decode(v) for v in value] + return value + + +def _make_resp2_parser(stream, read_size=4096): + parser = _AsyncRESP2Parser(socket_read_size=read_size) + parser._stream = stream + parser._connected = True + parser.encoder = _DummyEncoder() + return parser + + +def _make_resp3_parser(stream, read_size=4096): + parser = _AsyncRESP3Parser(socket_read_size=read_size) + parser._stream = stream + parser._connected = True + parser.encoder = _DummyEncoder() + return parser + + +@pytest.mark.parametrize( + "factory", + [_make_resp2_parser, _make_resp3_parser], + ids=["AsyncRESP2Parser", "AsyncRESP3Parser"], +) +async def test_per_read_timeout_allows_slow_multi_chunk_response(factory): + """ + A response that takes longer than the timeout in total, but where each + individual socket read completes quickly, must succeed under per-read + timeout semantics. + """ + # Bulk string payload split across several chunks with 0.05s delay each. + payload = b"hello world this is a moderately large bulk string value" + chunks = [ + b"$" + str(len(payload)).encode() + b"\r\n", + payload[:10], + payload[10:25], + payload[25:40], + payload[40:] + b"\r\n", + ] + stream = SlowChunkStream(chunks, delay_between_chunks=0.05) + parser = factory(stream) + + # Total elapsed will be ~0.2s, but each read is only 0.05s. + # With per-read semantics a 0.1s timeout should allow it. + response = await parser.read_response(timeout=0.1) + assert response == payload.decode() + + +@pytest.mark.parametrize( + "factory", + [_make_resp2_parser, _make_resp3_parser], + ids=["AsyncRESP2Parser", "AsyncRESP3Parser"], +) +async def test_per_read_timeout_fails_when_single_read_exceeds_timeout(factory): + """ + If an individual socket read itself exceeds the timeout, the parser must + raise a timeout error. + """ + chunks = [b"$5\r\n", b"hello", b"\r\n"] + # 0.3s delay per chunk means the second read will exceed a 0.1s timeout. + stream = SlowChunkStream(chunks, delay_between_chunks=0.3) + parser = factory(stream) + + with pytest.raises(asyncio.TimeoutError): + await parser.read_response(timeout=0.1) + + +@pytest.mark.parametrize( + "factory", + [_make_resp2_parser, _make_resp3_parser], + ids=["AsyncRESP2Parser", "AsyncRESP3Parser"], +) +async def test_per_read_timeout_propagates_through_nested_arrays(factory): + """ + Nested RESP arrays must keep the per-read timeout on every recursive + _readline/_read call. + """ + # *2\r\n$5\r\nhello\r\n$5\r\nworld\r\n split into many chunks + chunks = [ + b"*2\r\n", + b"$5\r\n", + b"hello", + b"\r\n$5\r\n", + b"world", + b"\r\n", + ] + stream = SlowChunkStream(chunks, delay_between_chunks=0.04) + parser = factory(stream) + + # Total ~0.24s but each read 0.04s; 0.1s per-read timeout should pass. + response = await parser.read_response(timeout=0.1) + assert response == ["hello", "world"] + + +@pytest.mark.parametrize( + "factory", + [_make_resp2_parser, _make_resp3_parser], + ids=["AsyncRESP2Parser", "AsyncRESP3Parser"], +) +async def test_no_timeout_when_sentinel_default(factory): + """When no timeout is supplied (SENTINEL default), reads must not time out.""" + chunks = [b"+OK\r\n"] + stream = SlowChunkStream(chunks, delay_between_chunks=0.1) + parser = factory(stream) + + from redis.utils import SENTINEL + + response = await parser.read_response(timeout=SENTINEL) + assert response == "OK" + + +@pytest.mark.skipif(not HIREDIS_AVAILABLE, reason="hiredis is not installed") +async def test_hiredis_per_read_timeout_allows_slow_multi_chunk_response(): + """ + The hiredis async parser must also apply timeout per read_from_socket call, + not across the entire read_response loop. + """ + import hiredis + + payload = b"hello world this is a moderately large bulk string value" + chunks = [ + b"$" + str(len(payload)).encode() + b"\r\n", + payload[:10], + payload[10:25], + payload[25:40], + payload[40:] + b"\r\n", + ] + stream = SlowChunkStream(chunks, delay_between_chunks=0.05) + + parser = _AsyncHiredisParser(socket_read_size=4096) + parser._stream = stream + parser._connected = True + parser._reader = hiredis.Reader( + protocolError=Exception, + replyError=Exception, + notEnoughData=NOT_ENOUGH_DATA, + ) + + response = await parser.read_response(timeout=0.1) + assert response == payload + + +@pytest.mark.skipif(not HIREDIS_AVAILABLE, reason="hiredis is not installed") +async def test_hiredis_per_read_timeout_fails_when_chunk_too_slow(): + """Hiredis parser must raise when a single read_from_socket exceeds timeout.""" + import hiredis + + chunks = [b"$5\r\n", b"hello", b"\r\n"] + stream = SlowChunkStream(chunks, delay_between_chunks=0.3) + + parser = _AsyncHiredisParser(socket_read_size=4096) + parser._stream = stream + parser._connected = True + parser._reader = hiredis.Reader( + protocolError=Exception, + replyError=Exception, + notEnoughData=NOT_ENOUGH_DATA, + ) + + with pytest.raises(asyncio.TimeoutError): + await parser.read_response(timeout=0.1) + + +@pytest.mark.parametrize("protocol", [2, 3]) +async def test_connection_passes_timeout_to_parser(protocol): + """ + Connection.read_response must pass its timeout through to the parser + rather than wrapping the parser call in an outer timeout context. + """ + conn = Connection(protocol=protocol, socket_timeout=0.05) + # Ensure parser is the Python-backed one for this test. + from redis._parsers import _AsyncRESP2Parser, _AsyncRESP3Parser + + expected_class = _AsyncRESP2Parser if protocol == 2 else _AsyncRESP3Parser + assert isinstance(conn._parser, expected_class) + + # Patch the parser to record the timeout it receives. + recorded = {} + + async def fake_read_response(*args, **kwargs): + recorded["timeout"] = kwargs.get("timeout") + return "OK" + + conn._parser.read_response = fake_read_response + response = await conn.read_response(timeout=0.42) + assert response == "OK" + assert recorded["timeout"] == 0.42 From ceb0bd496a982f0bcc2c88fdf35578128e01e81a Mon Sep 17 00:00:00 2001 From: goingforstudying-ctrl Date: Mon, 29 Jun 2026 12:18:23 -0400 Subject: [PATCH 2/2] Preserve pipelined data across async parser clear cycles _clear() now retains unconsumed bytes beyond _pos so pipelined responses buffered during a readline() are not discarded after a successful parse. on_connect() performs a full buffer reset on (re)connect to avoid carrying stale data from a previous session. --- redis/_parsers/base.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/redis/_parsers/base.py b/redis/_parsers/base.py index 9fc9c6210f..7b373a5592 100644 --- a/redis/_parsers/base.py +++ b/redis/_parsers/base.py @@ -518,7 +518,9 @@ def __init__(self, socket_read_size: int): self._pos = 0 def _clear(self): - self._buffer = b"" + """Clear parsed data but preserve unconsumed pipelined bytes.""" + self._buffer = self._buffer[self._pos :] + self._pos = 0 self._chunks.clear() def on_connect(self, connection): @@ -527,7 +529,10 @@ def on_connect(self, connection): if self._stream is None: raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) self.encoder = connection.encoder - self._clear() + # Full reset on (re)connect — discard stale data from old connection + self._buffer = b"" + self._pos = 0 + self._chunks.clear() self._connected = True def on_disconnect(self):