_client_async.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467
  1. from __future__ import annotations
  2. import asyncio
  3. import functools
  4. import sys
  5. from asyncio import CancelledError, sleep, wait_for
  6. from typing import TYPE_CHECKING, Any, Protocol, TypeVar
  7. from urllib.parse import urlencode
  8. from pyqwest import Client as HTTPClient
  9. from pyqwest import FullResponse, Response
  10. from pyqwest import Headers as HTTPHeaders
  11. from . import _client_shared
  12. from ._codec import proto_binary_codec
  13. from ._compression import IdentityCompression, _gzip, resolve_compressions
  14. from ._interceptor_async import (
  15. BidiStreamInterceptor,
  16. ClientStreamInterceptor,
  17. Interceptor,
  18. ServerStreamInterceptor,
  19. UnaryInterceptor,
  20. resolve_interceptors,
  21. )
  22. from ._protocol import ConnectWireError
  23. from ._protocol_connect import ConnectClientProtocol, ConnectEnvelopeWriter
  24. from ._protocol_grpc import GRPCClientProtocol, GRPCWebClientProtocol
  25. from ._response_metadata import handle_response_headers
  26. from .code import Code
  27. from .errors import ConnectError
  28. from .protocol import ProtocolType
  29. if sys.version_info >= (3, 11):
  30. from asyncio import timeout as asyncio_timeout
  31. else:
  32. from ._asyncio_timeout import timeout as asyncio_timeout
  33. if TYPE_CHECKING:
  34. from collections.abc import AsyncIterator, Iterable, Mapping
  35. from types import TracebackType
  36. from ._envelope import EnvelopeReader
  37. from .codec import Codec
  38. from .compression import Compression
  39. from .method import MethodInfo
  40. from .request import Headers, RequestContext
  41. if sys.version_info >= (3, 11):
  42. from typing import Self
  43. else:
  44. from typing_extensions import Self
  45. else:
  46. Self = "Self"
  47. REQ = TypeVar("REQ")
  48. RES = TypeVar("RES")
  49. class _ExecuteUnary(Protocol[REQ, RES]):
  50. async def __call__(self, request: REQ, ctx: RequestContext[REQ, RES]) -> RES: ...
  51. class _ExecuteClientStream(Protocol[REQ, RES]):
  52. async def __call__(
  53. self, request: AsyncIterator[REQ], ctx: RequestContext[REQ, RES]
  54. ) -> RES: ...
  55. class _ExecuteServerStream(Protocol[REQ, RES]):
  56. def __call__(
  57. self, request: REQ, ctx: RequestContext[REQ, RES]
  58. ) -> AsyncIterator[RES]: ...
  59. class _ExecuteBidiStream(Protocol[REQ, RES]):
  60. def __call__(
  61. self, request: AsyncIterator[REQ], ctx: RequestContext[REQ, RES]
  62. ) -> AsyncIterator[RES]: ...
  63. class ConnectClient:
  64. """An asynchronous client for the Connect protocol."""
  65. _execute_unary: _ExecuteUnary
  66. _execute_client_stream: _ExecuteClientStream
  67. _execute_server_stream: _ExecuteServerStream
  68. _execute_bidi_stream: _ExecuteBidiStream
  69. def __init__(
  70. self,
  71. address: str,
  72. *,
  73. codec: Codec | None = None,
  74. protocol: ProtocolType = ProtocolType.CONNECT,
  75. accept_compression: Iterable[Compression] | None = None,
  76. send_compression: Compression | None = _gzip,
  77. timeout_ms: int | None = None,
  78. read_max_bytes: int | None = None,
  79. interceptors: Iterable[Interceptor] = (),
  80. http_client: HTTPClient | None = None,
  81. ) -> None:
  82. """Creates a new asynchronous Connect client.
  83. When providing an HTTP client, for example to configure TLS settings,
  84. it is the caller's responsibility to close it.
  85. Examples:
  86. ```python
  87. from pyqwest import Client
  88. from my_service import MyServiceClient
  89. async with (
  90. Client() as http_client,
  91. MyServiceClient("http://localhost:8000", http_client=http_client) as client,
  92. ):
  93. # Use the client!
  94. ```
  95. Args:
  96. address: The address of the server to connect to, including scheme.
  97. codec: The [Codec][] to use for requests. If unset, defaults to binary protobuf.
  98. For JSON encoding, use [proto_json_codec][connectrpc.codec.proto_json_codec].
  99. protocol: The [ProtocolType][] to use for requests.
  100. accept_compression: Compression algorithms to accept from the server. If unset,
  101. defaults to gzip. If set to empty, disables response compression.
  102. send_compression: Compression algorithm to use for sending requests. If unset,
  103. defaults to gzip. If set to None, disables request compression.
  104. timeout_ms: The timeout for requests in milliseconds.
  105. read_max_bytes: The maximum number of bytes to read from the response.
  106. interceptors: A list of interceptors to apply to requests.
  107. http_client: A pyqwest Client to use for requests.
  108. """
  109. self._address = address
  110. self._codec = codec or proto_binary_codec()
  111. self._response_compressions = resolve_compressions(accept_compression)
  112. self._accept_compression_header = ",".join(self._response_compressions.keys())
  113. self._send_compression = send_compression or IdentityCompression()
  114. self._timeout_ms = timeout_ms
  115. self._read_max_bytes = read_max_bytes
  116. if http_client:
  117. self._http_client = http_client
  118. else:
  119. # Use shared default transport if not specified
  120. self._http_client = HTTPClient()
  121. self._closed = False
  122. match protocol:
  123. case ProtocolType.CONNECT:
  124. self._protocol = ConnectClientProtocol()
  125. case ProtocolType.GRPC:
  126. self._protocol = GRPCClientProtocol()
  127. case ProtocolType.GRPC_WEB:
  128. self._protocol = GRPCWebClientProtocol()
  129. interceptors = resolve_interceptors(interceptors)
  130. execute_unary = self._send_request_unary
  131. for interceptor in (
  132. i for i in reversed(interceptors) if isinstance(i, UnaryInterceptor)
  133. ):
  134. execute_unary = functools.partial(
  135. interceptor.intercept_unary, execute_unary
  136. )
  137. self._execute_unary = execute_unary
  138. execute_client_stream = self._send_request_client_stream
  139. for interceptor in (
  140. i for i in reversed(interceptors) if isinstance(i, ClientStreamInterceptor)
  141. ):
  142. execute_client_stream = functools.partial(
  143. interceptor.intercept_client_stream, execute_client_stream
  144. )
  145. self._execute_client_stream = execute_client_stream
  146. execute_server_stream: _ExecuteServerStream = self._send_request_server_stream
  147. for interceptor in (
  148. i for i in reversed(interceptors) if isinstance(i, ServerStreamInterceptor)
  149. ):
  150. execute_server_stream = functools.partial(
  151. interceptor.intercept_server_stream, execute_server_stream
  152. )
  153. self._execute_server_stream = execute_server_stream
  154. execute_bidi_stream = self._send_request_bidi_stream
  155. for interceptor in (
  156. i for i in reversed(interceptors) if isinstance(i, BidiStreamInterceptor)
  157. ):
  158. execute_bidi_stream = functools.partial(
  159. interceptor.intercept_bidi_stream, execute_bidi_stream
  160. )
  161. self._execute_bidi_stream = execute_bidi_stream
  162. async def close(self) -> None:
  163. """Close the client. After closing, the client cannot be used to make requests."""
  164. if not self._closed:
  165. self._closed = True
  166. async def __aenter__(self) -> Self:
  167. return self
  168. async def __aexit__(
  169. self,
  170. _exc_type: type[BaseException] | None,
  171. _exc_value: BaseException | None,
  172. _traceback: TracebackType | None,
  173. ) -> None:
  174. await self.close()
  175. async def execute_unary(
  176. self,
  177. *,
  178. request: REQ,
  179. method: MethodInfo[REQ, RES],
  180. headers: Headers | Mapping[str, str] | None = None,
  181. timeout_ms: int | None = None,
  182. use_get: bool = False,
  183. ) -> RES:
  184. ctx = self._protocol.create_request_context(
  185. method=method,
  186. url=self._address,
  187. http_method="GET" if use_get else "POST",
  188. user_headers=headers,
  189. timeout_ms=timeout_ms or self._timeout_ms,
  190. codec=self._codec,
  191. stream=False,
  192. accept_compression=self._accept_compression_header,
  193. send_compression=self._send_compression,
  194. )
  195. return await self._execute_unary(request, ctx)
  196. async def execute_client_stream(
  197. self,
  198. *,
  199. request: AsyncIterator[REQ],
  200. method: MethodInfo[REQ, RES],
  201. headers: Headers | Mapping[str, str] | None = None,
  202. timeout_ms: int | None = None,
  203. ) -> RES:
  204. ctx = self._protocol.create_request_context(
  205. method=method,
  206. url=self._address,
  207. http_method="POST",
  208. user_headers=headers,
  209. timeout_ms=timeout_ms or self._timeout_ms,
  210. codec=self._codec,
  211. stream=True,
  212. accept_compression=self._accept_compression_header,
  213. send_compression=self._send_compression,
  214. )
  215. return await self._execute_client_stream(request, ctx)
  216. def execute_server_stream(
  217. self,
  218. *,
  219. request: REQ,
  220. method: MethodInfo[REQ, RES],
  221. headers: Headers | Mapping[str, str] | None = None,
  222. timeout_ms: int | None = None,
  223. ) -> AsyncIterator[RES]:
  224. ctx = self._protocol.create_request_context(
  225. method=method,
  226. url=self._address,
  227. http_method="POST",
  228. user_headers=headers,
  229. timeout_ms=timeout_ms or self._timeout_ms,
  230. codec=self._codec,
  231. stream=True,
  232. accept_compression=self._accept_compression_header,
  233. send_compression=self._send_compression,
  234. )
  235. return self._execute_server_stream(request, ctx)
  236. def execute_bidi_stream(
  237. self,
  238. *,
  239. request: AsyncIterator[REQ],
  240. method: MethodInfo[REQ, RES],
  241. headers: Headers | Mapping[str, str] | None = None,
  242. timeout_ms: int | None = None,
  243. ) -> AsyncIterator[RES]:
  244. ctx = self._protocol.create_request_context(
  245. method=method,
  246. url=self._address,
  247. http_method="POST",
  248. user_headers=headers,
  249. timeout_ms=timeout_ms or self._timeout_ms,
  250. codec=self._codec,
  251. stream=True,
  252. accept_compression=self._accept_compression_header,
  253. send_compression=self._send_compression,
  254. )
  255. return self._execute_bidi_stream(request, ctx)
  256. async def _send_request_unary(
  257. self, request: REQ, ctx: RequestContext[REQ, RES]
  258. ) -> RES:
  259. if isinstance(self._protocol, GRPCClientProtocol):
  260. return await _consume_single_response(
  261. self._send_request_bidi_stream(_yield_single_message(request), ctx)
  262. )
  263. request_headers = HTTPHeaders(ctx.request_headers.allitems())
  264. url = f"{self._address}/{ctx.method.service_name}/{ctx.method.name}"
  265. if (timeout_ms := ctx.timeout_ms) is not None:
  266. timeout_s = timeout_ms / 1000.0
  267. else:
  268. timeout_s = None
  269. try:
  270. request_data = self._codec.encode(request)
  271. if self._send_compression:
  272. request_data = self._send_compression.compress(request_data)
  273. if ctx.http_method == "GET":
  274. params = _client_shared.prepare_get_params(
  275. self._codec, request_data, request_headers
  276. )
  277. params_str = urlencode(params)
  278. url = f"{url}?{params_str}"
  279. request_headers.pop("content-type", None)
  280. resp = await wait_for(
  281. self._http_client.get(url=url, headers=request_headers), timeout_s
  282. )
  283. else:
  284. resp = await wait_for(
  285. self._http_client.post(
  286. url=url, headers=request_headers, content=request_data
  287. ),
  288. timeout_s,
  289. )
  290. self._protocol.validate_response(
  291. self._codec.name(), resp.status, resp.headers.get("content-type", "")
  292. )
  293. # Decompression itself is handled by pyqwest, but we validate it
  294. # by resolving it.
  295. self._protocol.handle_response_compression(
  296. resp.headers, self._response_compressions, stream=False
  297. )
  298. handle_response_headers(resp.headers)
  299. if resp.status == 200:
  300. if (
  301. self._read_max_bytes is not None
  302. and len(resp.content) > self._read_max_bytes
  303. ):
  304. raise ConnectError(
  305. Code.RESOURCE_EXHAUSTED,
  306. f"message is larger than configured max {self._read_max_bytes}",
  307. )
  308. return self._codec.decode(resp.content, ctx.method.output)
  309. raise ConnectWireError.from_response(resp).to_exception()
  310. except (TimeoutError, asyncio.TimeoutError) as e:
  311. raise ConnectError(Code.DEADLINE_EXCEEDED, "Request timed out") from e
  312. except ConnectError:
  313. raise
  314. except CancelledError as e:
  315. raise ConnectError(Code.CANCELED, "Request was cancelled") from e
  316. except Exception as e:
  317. raise ConnectError(Code.UNAVAILABLE, str(e)) from e
  318. async def _send_request_client_stream(
  319. self, request: AsyncIterator[REQ], ctx: RequestContext[REQ, RES]
  320. ) -> RES:
  321. return await _consume_single_response(
  322. self._send_request_bidi_stream(request, ctx)
  323. )
  324. def _send_request_server_stream(
  325. self, request: REQ, ctx: RequestContext[REQ, RES]
  326. ) -> AsyncIterator[RES]:
  327. return self._send_request_bidi_stream(_yield_single_message(request), ctx)
  328. async def _send_request_bidi_stream(
  329. self, request: AsyncIterator[REQ], ctx: RequestContext[REQ, RES]
  330. ) -> AsyncIterator[RES]:
  331. request_headers = HTTPHeaders(ctx.request_headers.allitems())
  332. url = f"{self._address}/{ctx.method.service_name}/{ctx.method.name}"
  333. if (timeout_ms := ctx.timeout_ms) is not None:
  334. timeout_s = timeout_ms / 1000.0
  335. else:
  336. timeout_s = None
  337. reader: EnvelopeReader | None = None
  338. resp: Response | None = None
  339. try:
  340. request_data = _streaming_request_content(
  341. request, self._codec, self._send_compression
  342. )
  343. async with (
  344. asyncio_timeout(timeout_s),
  345. self._http_client.stream(
  346. "POST", url, headers=request_headers, content=request_data
  347. ) as resp,
  348. ):
  349. handle_response_headers(resp.headers)
  350. if resp.status == 200:
  351. self._protocol.validate_stream_response(
  352. self._codec.name(), resp.headers.get("content-type", "")
  353. )
  354. compression = self._protocol.handle_response_compression(
  355. resp.headers, self._response_compressions, stream=True
  356. )
  357. reader = self._protocol.create_envelope_reader(
  358. ctx.method.output,
  359. self._codec,
  360. compression,
  361. self._read_max_bytes,
  362. )
  363. async for chunk in resp.content:
  364. for message in reader.feed(bytes(chunk)):
  365. yield message
  366. # Check for cancellation each message. While this seems heavyweight,
  367. # conformance tests require it.
  368. await sleep(0)
  369. reader.handle_response_complete(resp)
  370. else:
  371. content = bytearray()
  372. async for chunk in resp.content:
  373. content.extend(chunk)
  374. fres = FullResponse(
  375. status=resp.status,
  376. headers=resp.headers,
  377. content=bytes(content),
  378. trailers=resp.trailers,
  379. )
  380. raise ConnectWireError.from_response(fres).to_exception()
  381. except (TimeoutError, asyncio.TimeoutError) as e:
  382. raise ConnectError(Code.DEADLINE_EXCEEDED, "Request timed out") from e
  383. except ConnectError:
  384. raise
  385. except CancelledError as e:
  386. raise ConnectError(Code.CANCELED, "Request was cancelled") from e
  387. except Exception as e:
  388. if rst_err := _client_shared.maybe_map_stream_reset(e, ctx):
  389. # It is possible for a reset to come with trailers which should
  390. # be used.
  391. if reader and resp:
  392. reader.handle_response_complete(resp, rst_err)
  393. raise rst_err from e
  394. raise ConnectError(Code.UNAVAILABLE, str(e)) from e
  395. async def _streaming_request_content(
  396. msgs: AsyncIterator[Any], codec: Codec, compression: Compression | None
  397. ) -> AsyncIterator[bytes]:
  398. writer = ConnectEnvelopeWriter(codec, compression)
  399. async for msg in msgs:
  400. yield writer.write(msg)
  401. async def _yield_single_message(message: REQ) -> AsyncIterator[REQ]:
  402. yield message
  403. async def _consume_single_response(stream: AsyncIterator[RES]) -> RES:
  404. res = None
  405. async for message in stream:
  406. if res is not None:
  407. raise ConnectError(
  408. Code.UNIMPLEMENTED, "unary response has multiple messages"
  409. )
  410. res = message
  411. if res is None:
  412. raise ConnectError(Code.UNIMPLEMENTED, "unary response has zero messages")
  413. return res