diff --git a/elasticsearch/_async/transport.py b/elasticsearch/_async/transport.py index f75b79fe..3fd63733 100644 --- a/elasticsearch/_async/transport.py +++ b/elasticsearch/_async/transport.py @@ -84,6 +84,7 @@ class AsyncTransport(Transport): self.sniffing_task = None self.loop = None self._async_init_called = False + self._sniff_on_start_event = None # type: asyncio.Event super(AsyncTransport, self).__init__( *args, hosts=[], sniff_on_start=False, **kwargs @@ -112,14 +113,35 @@ class AsyncTransport(Transport): self.loop = get_running_loop() self.kwargs["loop"] = self.loop - # Now that we have a loop we can create all our HTTP connections + # Now that we have a loop we can create all our HTTP connections... self.set_connections(self.hosts) self.seed_connections = list(self.connection_pool.connections[:]) # ... and we can start sniffing in the background. if self.sniffing_task is None and self.sniff_on_start: - self.last_sniff = self.loop.time() - self.create_sniff_task(initial=True) + + # Create an asyncio.Event for future calls to block on + # until the initial sniffing task completes. + self._sniff_on_start_event = asyncio.Event() + + try: + self.last_sniff = self.loop.time() + self.create_sniff_task(initial=True) + + # Since this is the first one we wait for it to complete + # in case there's an error it'll get raised here. + await self.sniffing_task + + # If the task gets cancelled here it likely means the + # transport got closed. + except asyncio.CancelledError: + pass + + # Once we exit this section we want to unblock any _async_calls() + # that are blocking on our initial sniff attempt regardless of it + # was successful or not. + finally: + self._sniff_on_start_event.set() async def _async_call(self): """This method is called within any async method of AsyncTransport @@ -130,6 +152,14 @@ class AsyncTransport(Transport): self._async_init_called = True await self._async_init() + # If the initial sniff_on_start hasn't returned yet + # then we need to wait for node information to come back + # or for the task to be cancelled via AsyncTransport.close() + if self._sniff_on_start_event and not self._sniff_on_start_event.is_set(): + # This is already a no-op if the event is set but we try to + # avoid an 'await' by checking 'not event.is_set()' above first. + await self._sniff_on_start_event.wait() + if self.sniffer_timeout: if self.loop.time() >= self.last_sniff + self.sniffer_timeout: self.create_sniff_task() @@ -187,6 +217,12 @@ class AsyncTransport(Transport): for t in done: try: _, headers, node_info = t.result() + + # Lowercase all the header names for consistency in accessing them. + headers = { + header.lower(): value for header, value in headers.items() + } + node_info = self.deserializer.loads( node_info, headers.get("content-type") ) @@ -212,6 +248,8 @@ class AsyncTransport(Transport): """ # Without a loop we can't do anything. if not self.loop: + if initial: + raise RuntimeError("Event loop not running on initial sniffing task") return node_info = await self._get_sniff_data(initial) @@ -293,7 +331,7 @@ class AsyncTransport(Transport): connection = self.get_connection() try: - status, headers, data = await connection.perform_request( + status, headers_response, data = await connection.perform_request( method, url, params, @@ -302,6 +340,11 @@ class AsyncTransport(Transport): ignore=ignore, timeout=timeout, ) + + # Lowercase all the header names for consistency in accessing them. + headers_response = { + header.lower(): value for header, value in headers_response.items() + } except TransportError as e: if method == "HEAD" and e.status_code == 404: return False @@ -336,7 +379,9 @@ class AsyncTransport(Transport): return 200 <= status < 300 if data: - data = self.deserializer.loads(data, headers.get("content-type")) + data = self.deserializer.loads( + data, headers_response.get("content-type") + ) return data async def close(self): @@ -350,5 +395,6 @@ class AsyncTransport(Transport): except asyncio.CancelledError: pass self.sniffing_task = None + for connection in self.connection_pool.connections: await connection.close() diff --git a/elasticsearch/transport.py b/elasticsearch/transport.py index b9e07a99..cf46e9b2 100644 --- a/elasticsearch/transport.py +++ b/elasticsearch/transport.py @@ -278,6 +278,12 @@ class Transport(object): "/_nodes/_all/http", timeout=self.sniff_timeout if not initial else None, ) + + # Lowercase all the header names for consistency in accessing them. + headers = { + header.lower(): value for header, value in headers.items() + } + node_info = self.deserializer.loads( node_info, headers.get("content-type") ) @@ -388,6 +394,11 @@ class Transport(object): timeout=timeout, ) + # Lowercase all the header names for consistency in accessing them. + headers_response = { + header.lower(): value for header, value in headers_response.items() + } + except TransportError as e: if method == "HEAD" and e.status_code == 404: return False diff --git a/test_elasticsearch/test_async/test_transport.py b/test_elasticsearch/test_async/test_transport.py index c7ea0c34..13bc492f 100644 --- a/test_elasticsearch/test_async/test_transport.py +++ b/test_elasticsearch/test_async/test_transport.py @@ -494,3 +494,83 @@ class TestTransport: assert not any([conn.closed for conn in t.connection_pool.connections]) await t.close() assert all([conn.closed for conn in t.connection_pool.connections]) + + async def test_sniff_on_start_error_if_no_sniffed_hosts(self, event_loop): + t = AsyncTransport( + [ + {"data": ""}, + {"data": ""}, + {"data": ""}, + ], + connection_class=DummyConnection, + sniff_on_start=True, + ) + + # If our initial sniffing attempt comes back + # empty then we raise an error. + with pytest.raises(TransportError) as e: + await t._async_call() + assert str(e.value) == "TransportError(N/A, 'Unable to sniff hosts.')" + + async def test_sniff_on_start_waits_for_sniff_to_complete(self, event_loop): + t = AsyncTransport( + [ + {"delay": 1, "data": ""}, + {"delay": 1, "data": ""}, + {"delay": 1, "data": CLUSTER_NODES}, + ], + connection_class=DummyConnection, + sniff_on_start=True, + ) + + # Start the timer right before the first task + # and have a bunch of tasks come in immediately. + tasks = [] + start_time = event_loop.time() + for _ in range(5): + tasks.append(event_loop.create_task(t._async_call())) + await asyncio.sleep(0) # Yield to the loop + + assert t.sniffing_task is not None + + # Tasks streaming in later. + for _ in range(5): + tasks.append(event_loop.create_task(t._async_call())) + await asyncio.sleep(0.1) + + # Now that all the API calls have come in we wait for + # them all to resolve before + await asyncio.gather(*tasks) + end_time = event_loop.time() + duration = end_time - start_time + + # All the tasks blocked on the sniff of each node + # and then resolved immediately after. + assert 1 <= duration < 2 + + async def test_sniff_on_start_close_unlocks_async_calls(self, event_loop): + t = AsyncTransport( + [ + {"delay": 10, "data": CLUSTER_NODES}, + ], + connection_class=DummyConnection, + sniff_on_start=True, + ) + + # Start making _async_calls() before we cancel + tasks = [] + start_time = event_loop.time() + for _ in range(3): + tasks.append(event_loop.create_task(t._async_call())) + await asyncio.sleep(0) + + # Close the transport while the sniffing task is active! :( + await t.close() + + # Now we start waiting on all those _async_calls() + await asyncio.gather(*tasks) + end_time = event_loop.time() + duration = end_time - start_time + + # A lot quicker than 10 seconds defined in 'delay' + assert duration < 1