|
33 | 33 |
|
34 | 34 | import asyncio |
35 | 35 | import sys |
| 36 | +from contextlib import AsyncExitStack |
36 | 37 | from types import MappingProxyType, TracebackType |
37 | 38 |
|
38 | 39 | import httpx |
|
54 | 55 | from typesense.http_backend import ( |
55 | 56 | ASYNC_CLIENT_TYPES, |
56 | 57 | AsyncClientType, |
| 58 | + ResponseType, |
57 | 59 | backend_errors, |
58 | 60 | verify_option, |
59 | 61 | ) |
@@ -648,38 +650,41 @@ async def _open_stream( |
648 | 650 | headers = request_kwargs.get("headers", {}) |
649 | 651 | headers["Accept"] = "text/event-stream" |
650 | 652 | timeout = self._client.timeout |
651 | | - request = self._client.build_request( |
652 | | - method, |
653 | | - url, |
654 | | - params=typing.cast( |
655 | | - typing.Optional[_QueryParams], |
656 | | - request_kwargs.get("params"), |
657 | | - ), |
658 | | - content=request_kwargs.get("content"), |
659 | | - headers=headers, |
660 | | - timeout=( |
661 | | - timeout.connect, |
662 | | - self.config.stream_read_timeout_seconds, |
663 | | - timeout.write, |
664 | | - timeout.pool, |
665 | | - ), |
| 653 | + # Annotated so httpx and httpx2 responses unify as ``ResponseType``. |
| 654 | + response_context: typing.AsyncContextManager[ResponseType] = ( |
| 655 | + self._client.stream( |
| 656 | + method, |
| 657 | + url, |
| 658 | + params=typing.cast( |
| 659 | + typing.Optional[_QueryParams], |
| 660 | + request_kwargs.get("params"), |
| 661 | + ), |
| 662 | + content=request_kwargs.get("content"), |
| 663 | + headers=headers, |
| 664 | + timeout=( |
| 665 | + timeout.connect, |
| 666 | + self.config.stream_read_timeout_seconds, |
| 667 | + timeout.write, |
| 668 | + timeout.pool, |
| 669 | + ), |
| 670 | + ) |
666 | 671 | ) |
667 | 672 |
|
| 673 | + # Owns the concurrency slot and the response until the stream is closed. |
| 674 | + exit_stack = AsyncExitStack() |
668 | 675 | await self._concurrency_limit.acquire() |
| 676 | + exit_stack.callback(self._concurrency_limit.release) |
669 | 677 | try: |
670 | | - response = await self._client.send(request, stream=True) |
| 678 | + response = await exit_stack.enter_async_context(response_context) |
671 | 679 | if response.status_code < 200 or response.status_code >= 300: |
672 | | - try: |
673 | | - await response.aread() |
674 | | - finally: |
675 | | - await response.aclose() |
| 680 | + await response.aread() |
676 | 681 | self.request_handler.raise_for_status(response) |
677 | 682 | except BaseException: |
678 | | - self._concurrency_limit.release() |
| 683 | + await exit_stack.aclose() |
679 | 684 | raise |
680 | 685 |
|
681 | 686 | self.node_manager.set_node_health(node, is_healthy=True) |
682 | | - return AsyncSearchStream(response, self._concurrency_limit.release) |
| 687 | + return AsyncSearchStream(response, exit_stack) |
683 | 688 |
|
684 | 689 | def _prepare_request_params( |
685 | 690 | self, |
|
0 commit comments