| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443 |
- from __future__ import annotations
- import struct
- import sys
- import urllib.parse
- from base64 import b64decode, b64encode
- from http import HTTPStatus
- from typing import TYPE_CHECKING, Any, TypeVar
- from pyqwest import Headers as HTTPHeaders
- from ._compression import IdentityCompression, negotiate_compression
- from ._envelope import EnvelopeReader, EnvelopeWriter
- from ._gen.google.rpc.status_pb import Status
- from ._protocol import (
- ConnectWireError,
- HTTPException,
- host_to_server_address,
- url_to_server_address,
- )
- from ._response_metadata import handle_response_trailers
- from ._version import __version__
- from .code import Code
- from .errors import ConnectError
- from .request import Headers, RequestContext
- if TYPE_CHECKING:
- from collections.abc import Mapping
- from pyqwest import Response, SyncResponse
- from ._codec import Codec
- from ._compression import Compression
- from .method import MethodInfo
- REQ = TypeVar("REQ")
- RES = TypeVar("RES")
- GRPC_CONTENT_TYPE_DEFAULT = "application/grpc"
- GRPC_CONTENT_TYPE_PREFIX = f"{GRPC_CONTENT_TYPE_DEFAULT}+"
- GRPC_WEB_CONTENT_TYPE_DEFAULT = "application/grpc-web"
- GRPC_WEB_CONTENT_TYPE_PREFIX = f"{GRPC_WEB_CONTENT_TYPE_DEFAULT}+"
- GRPC_HEADER_TIMEOUT = "grpc-timeout"
- GRPC_HEADER_COMPRESSION = "grpc-encoding"
- GRPC_HEADER_ACCEPT_COMPRESSION = "grpc-accept-encoding"
- _DEFAULT_GRPC_USER_AGENT = f"grpc-python-connect/{__version__} ({sys.version})"
- class GRPCServerProtocol:
- def create_request_context(
- self,
- method: MethodInfo[REQ, RES],
- http_method: str,
- http_scheme: str,
- headers: Headers,
- client_address: str | None = None,
- ) -> RequestContext[REQ, RES]:
- if http_method != "POST":
- raise HTTPException(HTTPStatus.METHOD_NOT_ALLOWED, [("allow", "POST")])
- timeout_header = headers.get(GRPC_HEADER_TIMEOUT)
- timeout_ms = _parse_timeout(timeout_header) if timeout_header else None
- server_address = host_to_server_address(headers.get("host"), http_scheme)
- return RequestContext(
- method=method,
- http_method=http_method,
- request_headers=headers,
- timeout_ms=timeout_ms,
- server_address=server_address,
- client_address=client_address,
- )
- def create_envelope_writer(
- self, codec: Codec[RES, Any], compression: Compression | None
- ) -> EnvelopeWriter[RES]:
- return GRPCEnvelopeWriter(codec, compression)
- def uses_trailers(self) -> bool:
- return True
- def content_type(self, codec: Codec) -> str:
- return f"{GRPC_CONTENT_TYPE_PREFIX}{codec.name()}"
- def compression_header_name(self) -> str:
- return GRPC_HEADER_COMPRESSION
- def codec_name_from_content_type(self, content_type: str, *, stream: bool) -> str:
- if content_type.startswith(GRPC_CONTENT_TYPE_PREFIX):
- return content_type[len(GRPC_CONTENT_TYPE_PREFIX) :]
- return "proto"
- def negotiate_stream_compression(
- self, headers: Headers, compressions: dict[str, Compression]
- ) -> tuple[Compression | None, Compression]:
- req_compression_name = headers.get(GRPC_HEADER_COMPRESSION, "identity")
- req_compression = compressions.get(req_compression_name)
- accept_compression = headers.get(GRPC_HEADER_ACCEPT_COMPRESSION, "")
- resp_compression = negotiate_compression(accept_compression, compressions)
- return req_compression, resp_compression
- class GRPCWebServerProtocol(GRPCServerProtocol):
- def uses_trailers(self) -> bool:
- return False
- def content_type(self, codec: Codec) -> str:
- return f"{GRPC_WEB_CONTENT_TYPE_PREFIX}{codec.name()}"
- def codec_name_from_content_type(self, content_type: str, *, stream: bool) -> str:
- if content_type.startswith(GRPC_WEB_CONTENT_TYPE_PREFIX):
- return content_type[len(GRPC_WEB_CONTENT_TYPE_PREFIX) :]
- return "proto"
- def create_envelope_writer(
- self, codec: Codec[RES, Any], compression: Compression | None
- ) -> EnvelopeWriter[RES]:
- return GRPCWebEnvelopeWriter(codec, compression)
- def _parse_timeout(timeout: str) -> int:
- # We normalize to int milliseconds matching connect's timeout header.
- value_to_ms = _lookup_timeout_unit(timeout[-1])
- try:
- value = int(timeout[:-1])
- except ValueError as e:
- msg = f"protocol error: invalid timeout '{timeout}'"
- raise ValueError(msg) from e
- # timeout must be ASCII string of at most 8 digits
- if value > 99999999:
- msg = f"protocol error: timeout '{timeout}' is too long"
- raise ValueError(msg)
- return int(value * value_to_ms)
- def _lookup_timeout_unit(unit: str) -> float:
- match unit:
- case "H":
- return 60 * 60 * 1000
- case "M":
- return 60 * 1000
- case "S":
- return 1 * 1000
- case "m":
- return 1
- case "u":
- return 1 / 1000
- case "n":
- return 1 / 1000 / 1000
- case _:
- msg = f"protocol error: timeout has invalid unit '{unit}'"
- raise ValueError(msg)
- class GRPCEnvelopeWriter(EnvelopeWriter):
- def end(self, user_trailers: Headers, error: ConnectWireError | None) -> Headers:
- trailers = Headers(list(user_trailers.allitems()))
- if error:
- status = _connect_status_to_grpc[error.code]
- trailers["grpc-status"] = status
- message = error.message
- if message:
- message = urllib.parse.quote(message, safe="")
- trailers["grpc-message"] = message
- if error.details:
- grpc_status = Status(
- code=int(status),
- message=error.message,
- details=[d._any for d in error.details], # noqa: SLF001
- )
- grpc_status_bin = (
- b64encode(grpc_status.to_binary()).decode().rstrip("=")
- )
- trailers["grpc-status-details-bin"] = grpc_status_bin
- else:
- trailers["grpc-status"] = "0"
- return trailers
- class GRPCWebEnvelopeWriter(GRPCEnvelopeWriter):
- def end(self, user_trailers: Headers, error: ConnectWireError | None) -> bytes: # ty: ignore[invalid-method-override]
- trailers = super().end(user_trailers, error)
- data = "".join(f"{k}: {v}\r\n" for k, v in trailers.allitems()).encode()
- if self._compression:
- data = self._compression.compress(data)
- prefix = self._prefix | 0b10000000
- return struct.pack(">BI", prefix, len(data)) + data
- class GRPCClientProtocol:
- def __init__(self) -> None:
- self._content_type = GRPC_CONTENT_TYPE_DEFAULT
- self._content_type_prefix = GRPC_CONTENT_TYPE_PREFIX
- def create_request_context(
- self,
- *,
- method: MethodInfo[REQ, RES],
- url: str,
- http_method: str,
- user_headers: Headers | Mapping[str, str] | None,
- timeout_ms: int | None,
- codec: Codec,
- stream: bool,
- accept_compression: str,
- send_compression: Compression | None,
- ) -> RequestContext[REQ, RES]:
- match user_headers:
- case Headers():
- # Copy to prevent modification if user keeps reference
- # TODO: Optimize
- headers = Headers(tuple(user_headers.allitems()))
- case None:
- headers = Headers()
- case _:
- headers = Headers(user_headers)
- headers["te"] = "trailers"
- if "user-agent" not in headers:
- headers["user-agent"] = _DEFAULT_GRPC_USER_AGENT
- headers[GRPC_HEADER_ACCEPT_COMPRESSION] = accept_compression
- if send_compression is not None:
- headers[GRPC_HEADER_COMPRESSION] = send_compression.name()
- else:
- headers.pop(GRPC_HEADER_COMPRESSION, None)
- headers["content-type"] = f"{self._content_type_prefix}{codec.name()}"
- if timeout_ms is not None:
- headers[GRPC_HEADER_TIMEOUT] = _serialize_timeout(timeout_ms)
- server_address = url_to_server_address(url)
- return RequestContext(
- method=method,
- http_method=http_method,
- request_headers=headers,
- timeout_ms=timeout_ms,
- server_address=server_address,
- )
- def validate_response(
- self, request_codec_name: str, status_code: int, response_content_type: str
- ) -> None:
- raise NotImplementedError
- def validate_stream_response(
- self, request_codec_name: str, response_content_type: str
- ) -> None:
- if not (
- response_content_type == self._content_type
- or response_content_type.startswith(self._content_type_prefix)
- ):
- raise ConnectError(
- Code.UNKNOWN,
- f"invalid content-type: '{response_content_type}'; expecting '{self._content_type_prefix}{request_codec_name}'",
- )
- if response_content_type.startswith(self._content_type_prefix):
- response_codec_name = response_content_type[
- len(self._content_type_prefix) :
- ]
- else:
- response_codec_name = "proto"
- if response_codec_name != request_codec_name:
- raise ConnectError(
- Code.INTERNAL,
- f"invalid content-type: '{response_content_type}'; expecting '{self._content_type_prefix}{request_codec_name}'",
- )
- def handle_response_compression(
- self, headers: HTTPHeaders, compressions: dict[str, Compression], stream: bool
- ) -> Compression:
- encoding = headers.get(GRPC_HEADER_COMPRESSION)
- if not encoding:
- return IdentityCompression()
- res = compressions.get(encoding)
- if not res:
- raise ConnectError(
- Code.INTERNAL,
- f"unknown encoding '{encoding}'; accepted encodings are {', '.join(compressions.keys())}",
- )
- return res
- def create_envelope_reader(
- self,
- message_class: type[RES],
- codec: Codec,
- compression: Compression,
- read_max_bytes: int | None,
- ) -> EnvelopeReader[RES]:
- return GRPCEnvelopeReader(message_class, codec, compression, read_max_bytes)
- class GRPCWebClientProtocol(GRPCClientProtocol):
- def __init__(self) -> None:
- self._content_type = GRPC_WEB_CONTENT_TYPE_DEFAULT
- self._content_type_prefix = GRPC_WEB_CONTENT_TYPE_PREFIX
- def create_envelope_reader(
- self,
- message_class: type[RES],
- codec: Codec,
- compression: Compression,
- read_max_bytes: int | None,
- ) -> EnvelopeReader[RES]:
- return GRPCWebEnvelopeReader(message_class, codec, compression, read_max_bytes)
- class GRPCEnvelopeReader(EnvelopeReader[RES]):
- def __init__(
- self,
- message_class: type[RES],
- codec: Codec,
- compression: Compression,
- read_max_bytes: int | None,
- ) -> None:
- super().__init__(message_class, codec, compression, read_max_bytes)
- self._read_message = False
- def handle_end_message(
- self, prefix_byte: int, message_data: bytes | bytearray
- ) -> bool:
- # It's coincidence that this method is called when there is a body and not
- # when there isn't. Somewhat hacky, but easiest way to handle the case
- # where there is a body and no trailers, which conformance tests verify.
- self._read_message = True
- return False
- def get_response_trailers(self, response: Response | SyncResponse) -> HTTPHeaders:
- return response.trailers
- def handle_response_complete(
- self, response: Response | SyncResponse, e: ConnectError | None = None
- ) -> None:
- # Get the actual HTTP trailers
- trailers = self.get_response_trailers(response)
- # gRPC trailers are either the HTTP trailers if there was body present
- # or HTTP headers if there was no body.
- grpc_status = trailers.get("grpc-status")
- if grpc_status is None:
- # If there was a body message, we do not read response headers
- if self._read_message:
- raise e or ConnectError(Code.INTERNAL, "missing grpc-status trailer")
- trailers = response.headers
- handle_response_trailers(trailers)
- grpc_status = trailers.get("grpc-status")
- if grpc_status is None:
- raise e or ConnectError(Code.INTERNAL, "missing grpc-status trailer")
- # e is present for RST_STREAM. We prioritize its code while reading message and details
- # from trailers when available.
- code = e.code if e else None
- if grpc_status != "0":
- message = trailers.get("grpc-message", "")
- if grpc_status_details := trailers.get("grpc-status-details-bin"):
- status = Status.from_binary(b64decode(grpc_status_details + "==="))
- connect_code = code or _grpc_status_to_connect.get(
- str(status.code), Code.UNKNOWN
- )
- raise ConnectError(connect_code, status.message, details=status.details)
- connect_code = code or _grpc_status_to_connect.get(
- grpc_status, Code.UNKNOWN
- )
- raise ConnectError(connect_code, urllib.parse.unquote(message))
- class GRPCWebEnvelopeReader(GRPCEnvelopeReader):
- def __init__(
- self,
- message_class: type[RES],
- codec: Codec,
- compression: Compression,
- read_max_bytes: int | None,
- ) -> None:
- super().__init__(message_class, codec, compression, read_max_bytes)
- self._trailers = HTTPHeaders()
- def handle_end_message(
- self, prefix_byte: int, message_data: bytes | bytearray
- ) -> bool:
- super().handle_end_message(prefix_byte, message_data)
- end_stream = prefix_byte & 0b10000000 != 0
- if not end_stream:
- return False
- for line in message_data.split(b"\r\n"):
- if not line:
- continue
- key, value = line.split(b":", 1)
- self._trailers.add(key.decode().strip(), value.decode().strip())
- return True
- def get_response_trailers(self, response: Response | SyncResponse) -> HTTPHeaders:
- return self._trailers
- _GRPC_TIMEOUT_MAX_VALUE = 1e8
- def _serialize_timeout(timeout_ms: int) -> str:
- if timeout_ms <= 0:
- return "0n"
- # The gRPC protocol limits timeouts to 8 characters (not counting the unit),
- # so timeouts must be strictly less than 1e8 of the appropriate unit.
- if timeout_ms < _GRPC_TIMEOUT_MAX_VALUE:
- size, unit = 1, "m"
- elif timeout_ms < _GRPC_TIMEOUT_MAX_VALUE * 1000:
- size, unit = 1 * 1000, "S"
- elif timeout_ms < _GRPC_TIMEOUT_MAX_VALUE * 60 * 1000:
- size, unit = 60 * 1000, "M"
- else:
- size, unit = 60 * 60 * 1000, "H"
- return f"{timeout_ms // size}{unit}"
- _connect_status_to_grpc = {
- Code.CANCELED: "1",
- Code.UNKNOWN: "2",
- Code.INVALID_ARGUMENT: "3",
- Code.DEADLINE_EXCEEDED: "4",
- Code.NOT_FOUND: "5",
- Code.ALREADY_EXISTS: "6",
- Code.PERMISSION_DENIED: "7",
- Code.RESOURCE_EXHAUSTED: "8",
- Code.FAILED_PRECONDITION: "9",
- Code.ABORTED: "10",
- Code.OUT_OF_RANGE: "11",
- Code.UNIMPLEMENTED: "12",
- Code.INTERNAL: "13",
- Code.UNAVAILABLE: "14",
- Code.DATA_LOSS: "15",
- Code.UNAUTHENTICATED: "16",
- }
- _grpc_status_to_connect = {v: k for k, v in _connect_status_to_grpc.items()}
|