1import select
2import selectors
3import socket
4from logging import getLogger
5from typing import Callable, List, Optional, TypedDict, Union
6
7from ..exceptions import ConnectionError, InvalidResponse, RedisError, TimeoutError
8from ..typing import EncodableT
9from ..utils import HIREDIS_AVAILABLE, SENTINEL, deprecated_function
10from .base import (
11 AsyncBaseParser,
12 AsyncPushNotificationsParser,
13 BaseParser,
14 PushNotificationsParser,
15)
16from .socket import (
17 NONBLOCKING_EXCEPTION_ERROR_NUMBERS,
18 NONBLOCKING_EXCEPTIONS,
19 SERVER_CLOSED_CONNECTION_ERROR,
20)
21
22# Used to signal that hiredis-py does not have enough data to parse.
23# Using `False` or `None` is not reliable, given that the parser can
24# return `False` or `None` for legitimate reasons from RESP payloads.
25NOT_ENOUGH_DATA = object()
26
27# select.poll() is unavailable on Windows; fall back to selectors there.
28_HAS_POLL = hasattr(select, "poll")
29
30# POLLRDHUP (Linux) reports a peer half-close (FIN) that POLLHUP does not cover.
31# It is absent on macOS/Windows, where 0 turns the masks that use it into no-ops.
32_POLLRDHUP = getattr(select, "POLLRDHUP", 0)
33
34
35def _socket_can_read(sock, timeout: float) -> bool:
36 # SSL sockets can have decrypted bytes buffered above the OS socket layer.
37 if hasattr(sock, "pending") and sock.pending():
38 return True
39 # timeout=0 must be a non-blocking readiness check only; both branches
40 # below are non-destructive and have no FD_SETSIZE limit (select.select
41 # raises ValueError for fds >= 1024).
42 if _HAS_POLL:
43 # Prefer poll() over selectors.DefaultSelector: epoll/kqueue selectors
44 # allocate a file descriptor per check and so fail with EMFILE under
45 # fd exhaustion - the very condition that pushes sockets onto high
46 # fds. poll() allocates nothing.
47 poller = select.poll()
48 poller.register(sock, select.POLLIN)
49 # poll() takes milliseconds (None blocks forever). POLLHUP/POLLERR/
50 # POLLNVAL are always reported regardless of the registered mask, so
51 # closed or errored sockets still count as readable, like select().
52 poll_timeout = None if timeout is None else timeout * 1000
53 return bool(poller.poll(poll_timeout))
54 with selectors.DefaultSelector() as selector:
55 selector.register(sock, selectors.EVENT_READ)
56 return bool(selector.select(timeout))
57
58
59def _socket_is_closed(sock) -> bool:
60 # A server-closed socket reads as ready (it yields EOF), so readiness alone
61 # cannot tell it apart from a socket holding pending data, and both checks
62 # here are non-destructive so pending push messages (e.g. cache
63 # invalidations) are left intact to be processed. Without poll() the two
64 # states are indistinguishable, so report not-closed.
65 if not _HAS_POLL:
66 return False
67 # Decrypted TLS bytes buffered above the OS socket layer must be processed
68 # before the connection can be treated as closed, like kernel-level data.
69 if hasattr(sock, "pending") and sock.pending():
70 return False
71 poller = select.poll()
72 poller.register(sock, select.POLLIN | _POLLRDHUP)
73 events = poller.poll(0)
74 if not events:
75 return False
76 _, revents = events[0]
77 # A readable socket holds either data or EOF, and the poll flags alone
78 # cannot always tell which: POLLHUP/POLLRDHUP can be reported while unread
79 # data is still buffered (the peer sent data, then closed), and a drained
80 # closed socket reports plain POLLIN where POLLRDHUP is unavailable
81 # (PyPy). A non-destructive MSG_PEEK settles it: EOF peeks as b"".
82 try:
83 return sock.recv(1, socket.MSG_PEEK) == b""
84 except NONBLOCKING_EXCEPTIONS:
85 # Transient would-block: data may still arrive, keep the connection.
86 return False
87 except ValueError:
88 # SSL sockets do not support recv() flags; their buffered plaintext is
89 # covered by the pending() check above, so fall back to the poll
90 # flags. POLLRDHUP must be in the register mask or poll() won't
91 # report it: on Linux a graceful FIN reports POLLIN|POLLRDHUP and
92 # never POLLHUP. macOS/Windows set POLLHUP instead (_POLLRDHUP is 0).
93 closed_flags = select.POLLHUP | select.POLLERR | select.POLLNVAL | _POLLRDHUP
94 return bool(revents & closed_flags)
95 except OSError:
96 # The socket is errored (POLLERR/POLLNVAL); nothing left to read.
97 return True
98
99
100class _HiredisReaderArgs(TypedDict, total=False):
101 protocolError: Callable[[str], Exception]
102 replyError: Callable[[str], Exception]
103 encoding: Optional[str]
104 errors: Optional[str]
105
106
107class _HiredisParser(BaseParser, PushNotificationsParser):
108 "Parser class for connections using Hiredis"
109
110 def __init__(self, socket_read_size):
111 if not HIREDIS_AVAILABLE:
112 raise RedisError("Hiredis is not installed")
113 self.socket_read_size = socket_read_size
114 self._buffer = bytearray(socket_read_size)
115 self.pubsub_push_handler_func = self.handle_pubsub_push_response
116 self.node_moving_push_handler_func = None
117 self.maintenance_push_handler_func = None
118 self.oss_cluster_maint_push_handler_func = None
119 self.invalidation_push_handler_func = None
120 self._hiredis_PushNotificationType = None
121
122 def __del__(self):
123 try:
124 self.on_disconnect()
125 except Exception:
126 pass
127
128 def handle_pubsub_push_response(self, response):
129 logger = getLogger("push_response")
130 logger.debug("Push response: %s", response)
131 return response
132
133 def on_connect(self, connection, **kwargs):
134 import hiredis
135
136 self._sock = connection._sock
137 self._socket_timeout = connection.socket_timeout
138 kwargs = {
139 "protocolError": InvalidResponse,
140 "replyError": self.parse_error,
141 "errors": connection.encoder.encoding_errors,
142 "notEnoughData": NOT_ENOUGH_DATA,
143 }
144
145 if connection.encoder.decode_responses:
146 kwargs["encoding"] = connection.encoder.encoding
147 self._reader = hiredis.Reader(**kwargs)
148
149 try:
150 self._hiredis_PushNotificationType = hiredis.PushNotification
151 except AttributeError:
152 # hiredis < 3.2
153 self._hiredis_PushNotificationType = None
154
155 def on_disconnect(self):
156 self._sock = None
157 self._reader = None
158
159 def can_read(self, timeout: float = 0) -> bool:
160 # TODO: Rename this API; it detects pending data or dirty/closed
161 # connection state, not only whether application data can be read.
162 reader = self._reader
163 sock = self._sock
164 # Another thread may disconnect this connection while we are here (e.g.
165 # the multi-database health check taking a database out of service);
166 # on_disconnect() sets both _sock and _reader to None. Bind them locally
167 # so the checks below cannot disagree about the connection, and fail with
168 # a descriptive, retryable ConnectionError instead of a TypeError from
169 # registering None with poll().
170 if reader is None or sock is None:
171 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR)
172
173 if reader.has_data():
174 return True
175 try:
176 if not _socket_can_read(sock, timeout):
177 return False
178 # the socket reports readable but the reader has no buffered data. a
179 # server-closed socket also reads as ready (it yields EOF), so tell the
180 # two apart with a non-destructive poll: a peer-closed socket must not be
181 # reused, while a readable-but-open socket may just hold a pending push.
182 # this mirrors how the pure-Python parser (recv -> b"") and the async
183 # parser (StreamReader.at_eof()) already signal a closed connection.
184 if _socket_is_closed(sock):
185 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR)
186 except ValueError:
187 # Same race as above, one step later: a concurrent disconnect closed
188 # the socket, so its file descriptor is -1 and neither poll() nor a
189 # selector can register it. The connection is gone either way.
190 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) from None
191 return True
192
193 def read_from_socket(self, timeout=SENTINEL, raise_on_timeout=True):
194 sock = self._sock
195 reader = self._reader
196 # Another thread may disconnect this connection while we are here (e.g.
197 # a shared client closed via `with redis:`); on_disconnect() sets both
198 # _sock and _reader to None. Bind them locally and fail with a
199 # descriptive, retryable ConnectionError instead of an AttributeError.
200 if sock is None or reader is None:
201 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR)
202 custom_timeout = timeout is not SENTINEL
203 try:
204 if custom_timeout:
205 sock.settimeout(timeout)
206 bufflen = sock.recv_into(self._buffer)
207 if bufflen == 0:
208 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR)
209 reader.feed(self._buffer, 0, bufflen)
210 # data was read from the socket and added to the buffer.
211 # return True to indicate that data was read.
212 return True
213 except socket.timeout:
214 if raise_on_timeout:
215 raise TimeoutError("Timeout reading from socket")
216 return False
217 except NONBLOCKING_EXCEPTIONS as ex:
218 # if we're in nonblocking mode and the recv raises a
219 # blocking error, simply return False indicating that
220 # there's no data to be read. otherwise raise the
221 # original exception.
222 allowed = NONBLOCKING_EXCEPTION_ERROR_NUMBERS.get(ex.__class__, -1)
223 if ex.errno == allowed:
224 if not raise_on_timeout:
225 return False
226 if timeout == 0:
227 raise TimeoutError("Timeout reading from socket")
228 raise ConnectionError(f"Error while reading from socket: {ex.args}")
229 finally:
230 if custom_timeout:
231 try:
232 sock.settimeout(self._socket_timeout)
233 except OSError:
234 # A disconnect from another thread closes the socket, so there
235 # is nothing left to restore the timeout on. An exception from
236 # here would replace the outcome the body reported - including
237 # the retryable ConnectionError above - with EBADF.
238 pass
239
240 def read_response(
241 self,
242 disable_decoding=False,
243 push_request=False,
244 timeout: Union[float, object] = SENTINEL,
245 ):
246 # Bind the reader locally so a concurrent disconnect that clears
247 # self._reader can't turn a later .gets() into an AttributeError;
248 # re-checking the attribute each time would still race.
249 reader = self._reader
250 if reader is None:
251 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR)
252
253 if disable_decoding:
254 response = reader.gets(False)
255 else:
256 response = reader.gets()
257
258 while response is NOT_ENOUGH_DATA:
259 self.read_from_socket(timeout=timeout)
260 if disable_decoding:
261 response = reader.gets(False)
262 else:
263 response = reader.gets()
264 # if the response is a ConnectionError or the response is a list and
265 # the first item is a ConnectionError, raise it as something bad
266 # happened
267 if isinstance(response, ConnectionError):
268 raise response
269 elif self._hiredis_PushNotificationType is not None and isinstance(
270 response, self._hiredis_PushNotificationType
271 ):
272 response = self.handle_push_response(response)
273 if push_request:
274 return response
275 return self.read_response(
276 disable_decoding=disable_decoding,
277 push_request=push_request,
278 timeout=timeout,
279 )
280
281 elif (
282 isinstance(response, list)
283 and response
284 and isinstance(response[0], ConnectionError)
285 ):
286 raise response[0]
287 return response
288
289
290class _AsyncHiredisParser(AsyncBaseParser, AsyncPushNotificationsParser):
291 """Async implementation of parser class for connections using Hiredis"""
292
293 __slots__ = ("_reader",)
294
295 def __init__(self, socket_read_size: int):
296 if not HIREDIS_AVAILABLE:
297 raise RedisError("Hiredis is not available.")
298 super().__init__(socket_read_size=socket_read_size)
299 self._reader = None
300 self.pubsub_push_handler_func = self.handle_pubsub_push_response
301 self.invalidation_push_handler_func = None
302 self._hiredis_PushNotificationType = None
303
304 async def handle_pubsub_push_response(self, response):
305 logger = getLogger("push_response")
306 logger.debug("Push response: %s", response)
307 return response
308
309 def on_connect(self, connection):
310 import hiredis
311
312 self._stream = connection._reader
313 kwargs: _HiredisReaderArgs = {
314 "protocolError": InvalidResponse,
315 "replyError": self.parse_error,
316 "notEnoughData": NOT_ENOUGH_DATA,
317 }
318 if connection.encoder.decode_responses:
319 kwargs["encoding"] = connection.encoder.encoding
320 kwargs["errors"] = connection.encoder.encoding_errors
321
322 self._reader = hiredis.Reader(**kwargs)
323 self._connected = True
324
325 try:
326 self._hiredis_PushNotificationType = getattr(
327 hiredis, "PushNotification", None
328 )
329 except AttributeError:
330 # hiredis < 3.2
331 self._hiredis_PushNotificationType = None
332
333 def on_disconnect(self):
334 self._connected = False
335
336 @deprecated_function(
337 version="8.0.0", reason="Use can_read() instead", name="can_read_destructive"
338 )
339 async def can_read_destructive(self) -> bool:
340 return await self.can_read()
341
342 async def can_read(self) -> bool:
343 # TODO: Rename this API; it detects pending data or dirty/closed
344 # connection state, not only whether application data can be read.
345 if not self._connected:
346 raise OSError("Buffer is closed.")
347 # buffered data wins over EOF, like the sync parser: a pending
348 # response or push notification must stay readable even if the
349 # server has since closed the connection.
350 if self._reader.has_data():
351 return True
352 if self._stream.at_eof():
353 # Raise like the sync parser does on a server-closed connection,
354 # so callers that tolerate pending data (push notifications)
355 # can't mistake EOF for a readable connection.
356 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR)
357 # asyncio.StreamReader has no public non-destructive API for checking
358 # buffered bytes. Preserve dirty-connection detection for hiredis; tests
359 # with a real StreamReader guard this private buffer API in CI.
360 return bool(self._stream._buffer)
361
362 async def read_from_socket(self):
363 buffer = await self._stream.read(self._read_size)
364 if not buffer or not isinstance(buffer, bytes):
365 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) from None
366 self._reader.feed(buffer)
367 # data was read from the socket and added to the buffer.
368 # return True to indicate that data was read.
369 return True
370
371 async def read_response(
372 self, disable_decoding: bool = False, push_request: bool = False
373 ) -> Union[EncodableT, List[EncodableT]]:
374 # If `on_disconnect()` has been called, prohibit any more reads
375 # even if they could happen because data might be present.
376 # We still allow reads in progress to finish
377 if not self._connected:
378 raise ConnectionError(SERVER_CLOSED_CONNECTION_ERROR) from None
379
380 if disable_decoding:
381 response = self._reader.gets(False)
382 else:
383 response = self._reader.gets()
384
385 while response is NOT_ENOUGH_DATA:
386 await self.read_from_socket()
387 if disable_decoding:
388 response = self._reader.gets(False)
389 else:
390 response = self._reader.gets()
391
392 # if the response is a ConnectionError or the response is a list and
393 # the first item is a ConnectionError, raise it as something bad
394 # happened
395 if isinstance(response, ConnectionError):
396 raise response
397 elif self._hiredis_PushNotificationType is not None and isinstance(
398 response, self._hiredis_PushNotificationType
399 ):
400 response = await self.handle_push_response(response)
401 if not push_request:
402 return await self.read_response(
403 disable_decoding=disable_decoding, push_request=push_request
404 )
405 else:
406 return response
407 elif (
408 isinstance(response, list)
409 and response
410 and isinstance(response[0], ConnectionError)
411 ):
412 raise response[0]
413 return response