_connect.py 7.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255
  1. """A minimal hand-rolled Connect-over-HTTP client.
  2. This replaces the ``connectrpc`` package (and with it its ``pyqwest``
  3. dependency) for the one thing chattolib actually uses: unary request/response
  4. calls over the Connect JSON/binary protocol.
  5. The Connect unary protocol, in full::
  6. POST {address}/{service_name}/{Method}
  7. Content-Type: application/proto
  8. Authorization: Bearer <token> (chattolib adds this)
  9. <binary protobuf request body>
  10. -> 200, body is the binary protobuf response
  11. -> non-200, body is a JSON error envelope::
  12. {"code": "NOT_FOUND", "message": "...", "details": [...]}
  13. Everything else ``connectrpc`` ships (compression, interceptors, the sync
  14. client, the ASGI/WSGI server apps, streaming) is not used by chattolib and is
  15. deliberately not implemented here.
  16. """
  17. from __future__ import annotations
  18. import asyncio
  19. import enum
  20. from collections.abc import Mapping
  21. from dataclasses import dataclass
  22. from typing import Any
  23. import httpx
  24. from google.protobuf.message import Message
  25. from chattolib.exceptions import ChattoError
  26. class Code(enum.Enum):
  27. """Connect/gRPC status codes (the subset the protocol can return)."""
  28. CANCELED = "canceled"
  29. UNKNOWN = "unknown"
  30. INVALID_ARGUMENT = "invalid_argument"
  31. DEADLINE_EXCEEDED = "deadline_exceeded"
  32. NOT_FOUND = "not_found"
  33. ALREADY_EXISTS = "already_exists"
  34. PERMISSION_DENIED = "permission_denied"
  35. RESOURCE_EXHAUSTED = "resource_exhausted"
  36. FAILED_PRECONDITION = "failed_precondition"
  37. ABORTED = "aborted"
  38. OUT_OF_RANGE = "out_of_range"
  39. UNIMPLEMENTED = "unimplemented"
  40. INTERNAL = "internal"
  41. UNAVAILABLE = "unavailable"
  42. DATA_LOSS = "data_loss"
  43. UNAUTHENTICATED = "unauthenticated"
  44. # HTTP status -> Connect code, for non-JSON error bodies. Mirrors the
  45. # mapping in the Connect protocol spec.
  46. _HTTP_STATUS_TO_CODE: dict[int, Code] = {
  47. 400: Code.INVALID_ARGUMENT,
  48. 401: Code.UNAUTHENTICATED,
  49. 403: Code.PERMISSION_DENIED,
  50. 404: Code.NOT_FOUND,
  51. 409: Code.ALREADY_EXISTS,
  52. 413: Code.RESOURCE_EXHAUSTED,
  53. 429: Code.RESOURCE_EXHAUSTED,
  54. 500: Code.INTERNAL,
  55. 501: Code.UNIMPLEMENTED,
  56. 503: Code.UNAVAILABLE,
  57. }
  58. class ConnectError(Exception):
  59. """A ConnectRPC call returned a protocol error.
  60. Mirrors the shape ``connectrpc.errors.ConnectError`` exposes (``.code``,
  61. ``.message``, ``.details``) so ``_transport.translate_connect_error`` can
  62. consume it unchanged.
  63. """
  64. def __init__(
  65. self,
  66. code: Code,
  67. message: str,
  68. details: list[Any] | None = None,
  69. *,
  70. status_code: int | None = None,
  71. ) -> None:
  72. super().__init__(message)
  73. self.code = code
  74. self.message = message
  75. self.details = list(details or [])
  76. self.status_code = status_code
  77. def __str__(self) -> str:
  78. # Match connectrpc's ConnectError: str() is the message only, so
  79. # callers that embed str(exc) in a larger message don't double the code.
  80. return self.message
  81. @dataclass(frozen=True)
  82. class MethodInfo:
  83. """Static description of one unary method (what the stubs pass in)."""
  84. name: str
  85. service_name: str
  86. input: type[Message]
  87. output: type[Message]
  88. idempotency_level: str = "UNKNOWN"
  89. class _BinaryCodec:
  90. """Encode/decode protobuf messages as raw bytes (the `proto` codec)."""
  91. def name(self) -> str:
  92. return "proto"
  93. def encode(self, message: Message) -> bytes:
  94. return message.SerializeToString()
  95. def decode(self, data: bytes, message_class: type[Message]) -> Message:
  96. return message_class.FromString(data)
  97. def google_protobuf_binary_codec() -> _BinaryCodec:
  98. """Return the binary protobuf codec (drop-in for connectrpc's)."""
  99. return _BinaryCodec()
  100. def _code_from_name(name: str) -> Code:
  101. try:
  102. return Code(name.lower())
  103. except ValueError:
  104. return Code.UNKNOWN
  105. class ConnectClient:
  106. """Async unary Connect client.
  107. Subclassed by each generated ``*ServiceClient``; those subclasses add one
  108. method per RPC that calls :meth:`execute_unary`.
  109. """
  110. def __init__(
  111. self,
  112. address: str,
  113. *,
  114. codec: _BinaryCodec | None = None,
  115. timeout_ms: int | None = None,
  116. ) -> None:
  117. self._address = address
  118. self._codec = codec or _BinaryCodec()
  119. self._timeout_ms = timeout_ms
  120. self._http = httpx.AsyncClient(
  121. timeout=httpx.Timeout((timeout_ms / 1000.0) if timeout_ms else None)
  122. )
  123. async def close(self) -> None:
  124. await self._http.aclose()
  125. async def execute_unary(
  126. self,
  127. *,
  128. request: Message,
  129. method: MethodInfo,
  130. headers: Mapping[str, str] | None = None,
  131. ) -> Message:
  132. url = f"{self._address}/{method.service_name}/{method.name}"
  133. body = self._codec.encode(request)
  134. req_headers = {"Content-Type": "application/proto", **(headers or {})}
  135. try:
  136. resp = await self._http.post(url, content=body, headers=req_headers)
  137. except httpx.TimeoutException as e:
  138. raise ConnectError(Code.DEADLINE_EXCEEDED, "Request timed out") from e
  139. except httpx.HTTPError as e:
  140. raise ConnectError(Code.UNAVAILABLE, str(e)) from e
  141. if resp.status_code == 200:
  142. return self._codec.decode(resp.content, method.output)
  143. raise self._error_from_response(resp)
  144. @staticmethod
  145. def _error_from_response(resp: httpx.Response) -> ConnectError:
  146. try:
  147. data = resp.json()
  148. except Exception:
  149. data = None
  150. if isinstance(data, dict) and data.get("code"):
  151. code = _code_from_name(str(data["code"]))
  152. message = str(data.get("message", ""))
  153. else:
  154. code = _HTTP_STATUS_TO_CODE.get(resp.status_code, Code.UNKNOWN)
  155. message = resp.reason_phrase or ""
  156. return ConnectError(code, message, status_code=resp.status_code)
  157. class ConnectClientSync:
  158. """Synchronous unary Connect client.
  159. Wraps the async :class:`ConnectClient`: it runs the *same*
  160. ``execute_unary`` on a private event loop, so the protocol logic (codec,
  161. error parsing) is implemented once. Use it from synchronous code; from
  162. inside a running event loop it raises a clear :class:`ChattoError` instead
  163. of a nested-loop ``RuntimeError``.
  164. """
  165. def __init__(
  166. self,
  167. address: str,
  168. *,
  169. codec: _BinaryCodec | None = None,
  170. timeout_ms: int | None = None,
  171. ) -> None:
  172. self._async = ConnectClient(address, codec=codec, timeout_ms=timeout_ms)
  173. self._loop = asyncio.new_event_loop()
  174. def execute_unary(
  175. self,
  176. *,
  177. request: Message,
  178. method: MethodInfo,
  179. headers: Mapping[str, str] | None = None,
  180. ) -> Message:
  181. self._guard_no_running_loop()
  182. return self._loop.run_until_complete(
  183. self._async.execute_unary(request=request, method=method, headers=headers)
  184. )
  185. def close(self) -> None:
  186. self._loop.run_until_complete(self._async.close())
  187. self._loop.close()
  188. @staticmethod
  189. def _guard_no_running_loop() -> None:
  190. try:
  191. asyncio.get_running_loop()
  192. except RuntimeError:
  193. return # no running loop: safe to run_until_complete
  194. raise ChattoError(
  195. "the sync client cannot be used from within a running event loop; "
  196. "use the async ConnectClient in async code"
  197. )
  198. __all__ = [
  199. "Code",
  200. "ConnectClient",
  201. "ConnectClientSync",
  202. "ConnectError",
  203. "MethodInfo",
  204. "google_protobuf_binary_codec",
  205. ]