_protocol_grpc.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443
  1. from __future__ import annotations
  2. import struct
  3. import sys
  4. import urllib.parse
  5. from base64 import b64decode, b64encode
  6. from http import HTTPStatus
  7. from typing import TYPE_CHECKING, Any, TypeVar
  8. from pyqwest import Headers as HTTPHeaders
  9. from ._compression import IdentityCompression, negotiate_compression
  10. from ._envelope import EnvelopeReader, EnvelopeWriter
  11. from ._gen.google.rpc.status_pb import Status
  12. from ._protocol import (
  13. ConnectWireError,
  14. HTTPException,
  15. host_to_server_address,
  16. url_to_server_address,
  17. )
  18. from ._response_metadata import handle_response_trailers
  19. from ._version import __version__
  20. from .code import Code
  21. from .errors import ConnectError
  22. from .request import Headers, RequestContext
  23. if TYPE_CHECKING:
  24. from collections.abc import Mapping
  25. from pyqwest import Response, SyncResponse
  26. from ._codec import Codec
  27. from ._compression import Compression
  28. from .method import MethodInfo
  29. REQ = TypeVar("REQ")
  30. RES = TypeVar("RES")
  31. GRPC_CONTENT_TYPE_DEFAULT = "application/grpc"
  32. GRPC_CONTENT_TYPE_PREFIX = f"{GRPC_CONTENT_TYPE_DEFAULT}+"
  33. GRPC_WEB_CONTENT_TYPE_DEFAULT = "application/grpc-web"
  34. GRPC_WEB_CONTENT_TYPE_PREFIX = f"{GRPC_WEB_CONTENT_TYPE_DEFAULT}+"
  35. GRPC_HEADER_TIMEOUT = "grpc-timeout"
  36. GRPC_HEADER_COMPRESSION = "grpc-encoding"
  37. GRPC_HEADER_ACCEPT_COMPRESSION = "grpc-accept-encoding"
  38. _DEFAULT_GRPC_USER_AGENT = f"grpc-python-connect/{__version__} ({sys.version})"
  39. class GRPCServerProtocol:
  40. def create_request_context(
  41. self,
  42. method: MethodInfo[REQ, RES],
  43. http_method: str,
  44. http_scheme: str,
  45. headers: Headers,
  46. client_address: str | None = None,
  47. ) -> RequestContext[REQ, RES]:
  48. if http_method != "POST":
  49. raise HTTPException(HTTPStatus.METHOD_NOT_ALLOWED, [("allow", "POST")])
  50. timeout_header = headers.get(GRPC_HEADER_TIMEOUT)
  51. timeout_ms = _parse_timeout(timeout_header) if timeout_header else None
  52. server_address = host_to_server_address(headers.get("host"), http_scheme)
  53. return RequestContext(
  54. method=method,
  55. http_method=http_method,
  56. request_headers=headers,
  57. timeout_ms=timeout_ms,
  58. server_address=server_address,
  59. client_address=client_address,
  60. )
  61. def create_envelope_writer(
  62. self, codec: Codec[RES, Any], compression: Compression | None
  63. ) -> EnvelopeWriter[RES]:
  64. return GRPCEnvelopeWriter(codec, compression)
  65. def uses_trailers(self) -> bool:
  66. return True
  67. def content_type(self, codec: Codec) -> str:
  68. return f"{GRPC_CONTENT_TYPE_PREFIX}{codec.name()}"
  69. def compression_header_name(self) -> str:
  70. return GRPC_HEADER_COMPRESSION
  71. def codec_name_from_content_type(self, content_type: str, *, stream: bool) -> str:
  72. if content_type.startswith(GRPC_CONTENT_TYPE_PREFIX):
  73. return content_type[len(GRPC_CONTENT_TYPE_PREFIX) :]
  74. return "proto"
  75. def negotiate_stream_compression(
  76. self, headers: Headers, compressions: dict[str, Compression]
  77. ) -> tuple[Compression | None, Compression]:
  78. req_compression_name = headers.get(GRPC_HEADER_COMPRESSION, "identity")
  79. req_compression = compressions.get(req_compression_name)
  80. accept_compression = headers.get(GRPC_HEADER_ACCEPT_COMPRESSION, "")
  81. resp_compression = negotiate_compression(accept_compression, compressions)
  82. return req_compression, resp_compression
  83. class GRPCWebServerProtocol(GRPCServerProtocol):
  84. def uses_trailers(self) -> bool:
  85. return False
  86. def content_type(self, codec: Codec) -> str:
  87. return f"{GRPC_WEB_CONTENT_TYPE_PREFIX}{codec.name()}"
  88. def codec_name_from_content_type(self, content_type: str, *, stream: bool) -> str:
  89. if content_type.startswith(GRPC_WEB_CONTENT_TYPE_PREFIX):
  90. return content_type[len(GRPC_WEB_CONTENT_TYPE_PREFIX) :]
  91. return "proto"
  92. def create_envelope_writer(
  93. self, codec: Codec[RES, Any], compression: Compression | None
  94. ) -> EnvelopeWriter[RES]:
  95. return GRPCWebEnvelopeWriter(codec, compression)
  96. def _parse_timeout(timeout: str) -> int:
  97. # We normalize to int milliseconds matching connect's timeout header.
  98. value_to_ms = _lookup_timeout_unit(timeout[-1])
  99. try:
  100. value = int(timeout[:-1])
  101. except ValueError as e:
  102. msg = f"protocol error: invalid timeout '{timeout}'"
  103. raise ValueError(msg) from e
  104. # timeout must be ASCII string of at most 8 digits
  105. if value > 99999999:
  106. msg = f"protocol error: timeout '{timeout}' is too long"
  107. raise ValueError(msg)
  108. return int(value * value_to_ms)
  109. def _lookup_timeout_unit(unit: str) -> float:
  110. match unit:
  111. case "H":
  112. return 60 * 60 * 1000
  113. case "M":
  114. return 60 * 1000
  115. case "S":
  116. return 1 * 1000
  117. case "m":
  118. return 1
  119. case "u":
  120. return 1 / 1000
  121. case "n":
  122. return 1 / 1000 / 1000
  123. case _:
  124. msg = f"protocol error: timeout has invalid unit '{unit}'"
  125. raise ValueError(msg)
  126. class GRPCEnvelopeWriter(EnvelopeWriter):
  127. def end(self, user_trailers: Headers, error: ConnectWireError | None) -> Headers:
  128. trailers = Headers(list(user_trailers.allitems()))
  129. if error:
  130. status = _connect_status_to_grpc[error.code]
  131. trailers["grpc-status"] = status
  132. message = error.message
  133. if message:
  134. message = urllib.parse.quote(message, safe="")
  135. trailers["grpc-message"] = message
  136. if error.details:
  137. grpc_status = Status(
  138. code=int(status),
  139. message=error.message,
  140. details=[d._any for d in error.details], # noqa: SLF001
  141. )
  142. grpc_status_bin = (
  143. b64encode(grpc_status.to_binary()).decode().rstrip("=")
  144. )
  145. trailers["grpc-status-details-bin"] = grpc_status_bin
  146. else:
  147. trailers["grpc-status"] = "0"
  148. return trailers
  149. class GRPCWebEnvelopeWriter(GRPCEnvelopeWriter):
  150. def end(self, user_trailers: Headers, error: ConnectWireError | None) -> bytes: # ty: ignore[invalid-method-override]
  151. trailers = super().end(user_trailers, error)
  152. data = "".join(f"{k}: {v}\r\n" for k, v in trailers.allitems()).encode()
  153. if self._compression:
  154. data = self._compression.compress(data)
  155. prefix = self._prefix | 0b10000000
  156. return struct.pack(">BI", prefix, len(data)) + data
  157. class GRPCClientProtocol:
  158. def __init__(self) -> None:
  159. self._content_type = GRPC_CONTENT_TYPE_DEFAULT
  160. self._content_type_prefix = GRPC_CONTENT_TYPE_PREFIX
  161. def create_request_context(
  162. self,
  163. *,
  164. method: MethodInfo[REQ, RES],
  165. url: str,
  166. http_method: str,
  167. user_headers: Headers | Mapping[str, str] | None,
  168. timeout_ms: int | None,
  169. codec: Codec,
  170. stream: bool,
  171. accept_compression: str,
  172. send_compression: Compression | None,
  173. ) -> RequestContext[REQ, RES]:
  174. match user_headers:
  175. case Headers():
  176. # Copy to prevent modification if user keeps reference
  177. # TODO: Optimize
  178. headers = Headers(tuple(user_headers.allitems()))
  179. case None:
  180. headers = Headers()
  181. case _:
  182. headers = Headers(user_headers)
  183. headers["te"] = "trailers"
  184. if "user-agent" not in headers:
  185. headers["user-agent"] = _DEFAULT_GRPC_USER_AGENT
  186. headers[GRPC_HEADER_ACCEPT_COMPRESSION] = accept_compression
  187. if send_compression is not None:
  188. headers[GRPC_HEADER_COMPRESSION] = send_compression.name()
  189. else:
  190. headers.pop(GRPC_HEADER_COMPRESSION, None)
  191. headers["content-type"] = f"{self._content_type_prefix}{codec.name()}"
  192. if timeout_ms is not None:
  193. headers[GRPC_HEADER_TIMEOUT] = _serialize_timeout(timeout_ms)
  194. server_address = url_to_server_address(url)
  195. return RequestContext(
  196. method=method,
  197. http_method=http_method,
  198. request_headers=headers,
  199. timeout_ms=timeout_ms,
  200. server_address=server_address,
  201. )
  202. def validate_response(
  203. self, request_codec_name: str, status_code: int, response_content_type: str
  204. ) -> None:
  205. raise NotImplementedError
  206. def validate_stream_response(
  207. self, request_codec_name: str, response_content_type: str
  208. ) -> None:
  209. if not (
  210. response_content_type == self._content_type
  211. or response_content_type.startswith(self._content_type_prefix)
  212. ):
  213. raise ConnectError(
  214. Code.UNKNOWN,
  215. f"invalid content-type: '{response_content_type}'; expecting '{self._content_type_prefix}{request_codec_name}'",
  216. )
  217. if response_content_type.startswith(self._content_type_prefix):
  218. response_codec_name = response_content_type[
  219. len(self._content_type_prefix) :
  220. ]
  221. else:
  222. response_codec_name = "proto"
  223. if response_codec_name != request_codec_name:
  224. raise ConnectError(
  225. Code.INTERNAL,
  226. f"invalid content-type: '{response_content_type}'; expecting '{self._content_type_prefix}{request_codec_name}'",
  227. )
  228. def handle_response_compression(
  229. self, headers: HTTPHeaders, compressions: dict[str, Compression], stream: bool
  230. ) -> Compression:
  231. encoding = headers.get(GRPC_HEADER_COMPRESSION)
  232. if not encoding:
  233. return IdentityCompression()
  234. res = compressions.get(encoding)
  235. if not res:
  236. raise ConnectError(
  237. Code.INTERNAL,
  238. f"unknown encoding '{encoding}'; accepted encodings are {', '.join(compressions.keys())}",
  239. )
  240. return res
  241. def create_envelope_reader(
  242. self,
  243. message_class: type[RES],
  244. codec: Codec,
  245. compression: Compression,
  246. read_max_bytes: int | None,
  247. ) -> EnvelopeReader[RES]:
  248. return GRPCEnvelopeReader(message_class, codec, compression, read_max_bytes)
  249. class GRPCWebClientProtocol(GRPCClientProtocol):
  250. def __init__(self) -> None:
  251. self._content_type = GRPC_WEB_CONTENT_TYPE_DEFAULT
  252. self._content_type_prefix = GRPC_WEB_CONTENT_TYPE_PREFIX
  253. def create_envelope_reader(
  254. self,
  255. message_class: type[RES],
  256. codec: Codec,
  257. compression: Compression,
  258. read_max_bytes: int | None,
  259. ) -> EnvelopeReader[RES]:
  260. return GRPCWebEnvelopeReader(message_class, codec, compression, read_max_bytes)
  261. class GRPCEnvelopeReader(EnvelopeReader[RES]):
  262. def __init__(
  263. self,
  264. message_class: type[RES],
  265. codec: Codec,
  266. compression: Compression,
  267. read_max_bytes: int | None,
  268. ) -> None:
  269. super().__init__(message_class, codec, compression, read_max_bytes)
  270. self._read_message = False
  271. def handle_end_message(
  272. self, prefix_byte: int, message_data: bytes | bytearray
  273. ) -> bool:
  274. # It's coincidence that this method is called when there is a body and not
  275. # when there isn't. Somewhat hacky, but easiest way to handle the case
  276. # where there is a body and no trailers, which conformance tests verify.
  277. self._read_message = True
  278. return False
  279. def get_response_trailers(self, response: Response | SyncResponse) -> HTTPHeaders:
  280. return response.trailers
  281. def handle_response_complete(
  282. self, response: Response | SyncResponse, e: ConnectError | None = None
  283. ) -> None:
  284. # Get the actual HTTP trailers
  285. trailers = self.get_response_trailers(response)
  286. # gRPC trailers are either the HTTP trailers if there was body present
  287. # or HTTP headers if there was no body.
  288. grpc_status = trailers.get("grpc-status")
  289. if grpc_status is None:
  290. # If there was a body message, we do not read response headers
  291. if self._read_message:
  292. raise e or ConnectError(Code.INTERNAL, "missing grpc-status trailer")
  293. trailers = response.headers
  294. handle_response_trailers(trailers)
  295. grpc_status = trailers.get("grpc-status")
  296. if grpc_status is None:
  297. raise e or ConnectError(Code.INTERNAL, "missing grpc-status trailer")
  298. # e is present for RST_STREAM. We prioritize its code while reading message and details
  299. # from trailers when available.
  300. code = e.code if e else None
  301. if grpc_status != "0":
  302. message = trailers.get("grpc-message", "")
  303. if grpc_status_details := trailers.get("grpc-status-details-bin"):
  304. status = Status.from_binary(b64decode(grpc_status_details + "==="))
  305. connect_code = code or _grpc_status_to_connect.get(
  306. str(status.code), Code.UNKNOWN
  307. )
  308. raise ConnectError(connect_code, status.message, details=status.details)
  309. connect_code = code or _grpc_status_to_connect.get(
  310. grpc_status, Code.UNKNOWN
  311. )
  312. raise ConnectError(connect_code, urllib.parse.unquote(message))
  313. class GRPCWebEnvelopeReader(GRPCEnvelopeReader):
  314. def __init__(
  315. self,
  316. message_class: type[RES],
  317. codec: Codec,
  318. compression: Compression,
  319. read_max_bytes: int | None,
  320. ) -> None:
  321. super().__init__(message_class, codec, compression, read_max_bytes)
  322. self._trailers = HTTPHeaders()
  323. def handle_end_message(
  324. self, prefix_byte: int, message_data: bytes | bytearray
  325. ) -> bool:
  326. super().handle_end_message(prefix_byte, message_data)
  327. end_stream = prefix_byte & 0b10000000 != 0
  328. if not end_stream:
  329. return False
  330. for line in message_data.split(b"\r\n"):
  331. if not line:
  332. continue
  333. key, value = line.split(b":", 1)
  334. self._trailers.add(key.decode().strip(), value.decode().strip())
  335. return True
  336. def get_response_trailers(self, response: Response | SyncResponse) -> HTTPHeaders:
  337. return self._trailers
  338. _GRPC_TIMEOUT_MAX_VALUE = 1e8
  339. def _serialize_timeout(timeout_ms: int) -> str:
  340. if timeout_ms <= 0:
  341. return "0n"
  342. # The gRPC protocol limits timeouts to 8 characters (not counting the unit),
  343. # so timeouts must be strictly less than 1e8 of the appropriate unit.
  344. if timeout_ms < _GRPC_TIMEOUT_MAX_VALUE:
  345. size, unit = 1, "m"
  346. elif timeout_ms < _GRPC_TIMEOUT_MAX_VALUE * 1000:
  347. size, unit = 1 * 1000, "S"
  348. elif timeout_ms < _GRPC_TIMEOUT_MAX_VALUE * 60 * 1000:
  349. size, unit = 60 * 1000, "M"
  350. else:
  351. size, unit = 60 * 60 * 1000, "H"
  352. return f"{timeout_ms // size}{unit}"
  353. _connect_status_to_grpc = {
  354. Code.CANCELED: "1",
  355. Code.UNKNOWN: "2",
  356. Code.INVALID_ARGUMENT: "3",
  357. Code.DEADLINE_EXCEEDED: "4",
  358. Code.NOT_FOUND: "5",
  359. Code.ALREADY_EXISTS: "6",
  360. Code.PERMISSION_DENIED: "7",
  361. Code.RESOURCE_EXHAUSTED: "8",
  362. Code.FAILED_PRECONDITION: "9",
  363. Code.ABORTED: "10",
  364. Code.OUT_OF_RANGE: "11",
  365. Code.UNIMPLEMENTED: "12",
  366. Code.INTERNAL: "13",
  367. Code.UNAVAILABLE: "14",
  368. Code.DATA_LOSS: "15",
  369. Code.UNAUTHENTICATED: "16",
  370. }
  371. _grpc_status_to_connect = {v: k for k, v in _connect_status_to_grpc.items()}