_server_async.py 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665
  1. from __future__ import annotations
  2. import base64
  3. import contextlib
  4. import functools
  5. import inspect
  6. from abc import ABC, abstractmethod
  7. from asyncio import CancelledError, Event, create_task, sleep
  8. from dataclasses import replace
  9. from http import HTTPStatus
  10. from typing import TYPE_CHECKING, Generic, TypeVar, cast
  11. from urllib.parse import parse_qs
  12. from ._codec import Codec, get_default_codecs
  13. from ._compression import negotiate_compression, resolve_compressions
  14. from ._envelope import EnvelopeReader
  15. from ._interceptor_async import (
  16. BidiStreamInterceptor,
  17. ClientStreamInterceptor,
  18. Interceptor,
  19. ServerStreamInterceptor,
  20. UnaryInterceptor,
  21. resolve_interceptors,
  22. )
  23. from ._protocol import ConnectWireError, HTTPException, ServerProtocol
  24. from ._protocol_connect import CONNECT_UNARY_CONTENT_TYPE_PREFIX, ConnectServerProtocol
  25. from ._protocol_server import negotiate_server_protocol
  26. from ._server_shared import (
  27. EndpointBidiStream,
  28. EndpointClientStream,
  29. EndpointServerStream,
  30. EndpointUnary,
  31. )
  32. from .code import Code
  33. from .errors import ConnectError
  34. from .request import Headers, RequestContext
  35. if TYPE_CHECKING:
  36. # We don't use asgiref code so only import from it for type checking
  37. from collections.abc import (
  38. AsyncGenerator,
  39. AsyncIterator,
  40. Callable,
  41. Iterable,
  42. Mapping,
  43. Sequence,
  44. )
  45. from asgiref.typing import ASGIReceiveCallable, ASGISendCallable, HTTPScope, Scope
  46. from . import _server_shared
  47. from .compression import Compression
  48. else:
  49. ASGIReceiveCallable = "asgiref.typing.ASGIReceiveCallable"
  50. ASGISendCallable = "asgiref.typing.ASGISendCallable"
  51. HTTPScope = "asgiref.typing.HTTPScope"
  52. Scope = "asgiref.typing.Scope"
  53. _SVC = TypeVar("_SVC")
  54. _REQ = TypeVar("_REQ")
  55. _RES = TypeVar("_RES")
  56. # We don't mutate query params so use a singleton for when they're not set.
  57. _UNSET_QUERY_PARAMS: dict[str, list[str]] = {}
  58. # While _server_shared.Endpoint is a closed type, we can't indicate that to Python so define
  59. # a more precise type here.
  60. Endpoint = (
  61. EndpointBidiStream[_REQ, _RES]
  62. | EndpointClientStream[_REQ, _RES]
  63. | EndpointServerStream[_REQ, _RES]
  64. | EndpointUnary[_REQ, _RES]
  65. )
  66. class ConnectASGIApplication(ABC, Generic[_SVC]):
  67. """An ASGI application for the Connect protocol."""
  68. _resolved_endpoints: Mapping[str, Endpoint] | None
  69. @property
  70. @abstractmethod
  71. def path(self) -> str: ...
  72. def __init__(
  73. self,
  74. *,
  75. service: _SVC | AsyncGenerator[_SVC],
  76. endpoints: Callable[[_SVC], Mapping[str, Endpoint]],
  77. interceptors: Iterable[Interceptor] = (),
  78. read_max_bytes: int | None = None,
  79. compressions: Iterable[Compression] | None = None,
  80. codecs: Iterable[Codec] | None = None,
  81. ) -> None:
  82. """Initialize the ASGI application.
  83. Args:
  84. service: The service instance or async generator that yields the service during lifespan.
  85. endpoints: A callable that takes the service instance and returns a mapping of URL
  86. paths to endpoints. Typically provided directly by generated code from the
  87. Connect Python plugin.
  88. interceptors: A sequence of interceptors to apply to the endpoints.
  89. read_max_bytes: Maximum size of request messages.
  90. compressions: Supported compression algorithms. If unset, defaults to gzip.
  91. If set to empty, disables compression.
  92. codecs: The codecs supported by the server. If unset, defaults to Protocol Buffers
  93. binary and JSON codecs.
  94. """
  95. super().__init__()
  96. self._service = service
  97. self._endpoints = endpoints
  98. self._interceptors = interceptors
  99. self._resolved_endpoints = None
  100. self._read_max_bytes = read_max_bytes
  101. self._compressions = resolve_compressions(compressions)
  102. codecs = codecs if codecs is not None else get_default_codecs()
  103. self._codecs = {codec.name(): codec for codec in codecs}
  104. async def __call__(
  105. self, scope: Scope, receive: ASGIReceiveCallable, send: ASGISendCallable
  106. ) -> None:
  107. if scope["type"] == "websocket":
  108. msg = "connect does not support websockets"
  109. raise RuntimeError(msg)
  110. if scope["type"] == "lifespan":
  111. service_iter = None
  112. while True:
  113. msg = await receive()
  114. match msg["type"]:
  115. case "lifespan.startup":
  116. # Need to cast since type checking doesn't seem to narrow well with isasyncgen
  117. if inspect.isasyncgen(self._service):
  118. service_iter = cast(
  119. "AsyncGenerator[_SVC, None]", self._service
  120. )
  121. try:
  122. service = await anext(service_iter)
  123. except Exception as e:
  124. await send(
  125. {
  126. "type": "lifespan.startup.failed",
  127. "message": str(e),
  128. }
  129. )
  130. return None
  131. else:
  132. service = cast("_SVC", self._service)
  133. self._resolved_endpoints = self._resolve_endpoints(service)
  134. await send({"type": "lifespan.startup.complete"})
  135. case "lifespan.shutdown":
  136. if service_iter is not None:
  137. try:
  138. await service_iter.aclose()
  139. except Exception as e:
  140. await send(
  141. {
  142. "type": "lifespan.shutdown.failed",
  143. "message": str(e),
  144. }
  145. )
  146. return None
  147. await send({"type": "lifespan.shutdown.complete"})
  148. return None
  149. if not self._resolved_endpoints:
  150. if inspect.isasyncgen(self._service):
  151. msg = "ASGI server does not support lifespan but async generator passed for service. Enable lifespan support."
  152. raise RuntimeError(msg)
  153. self._resolved_endpoints = self._resolve_endpoints(
  154. cast("_SVC", self._service)
  155. )
  156. endpoints = self._resolved_endpoints
  157. ctx: RequestContext | None = None
  158. try:
  159. path = scope["path"]
  160. endpoint = endpoints.get(path)
  161. if not endpoint and scope["root_path"]:
  162. # The application was mounted at some root so try stripping the prefix.
  163. path = path.removeprefix(scope["root_path"])
  164. endpoint = endpoints.get(path)
  165. if not endpoint:
  166. raise HTTPException(HTTPStatus.NOT_FOUND, [])
  167. http_method = scope["method"]
  168. http_scheme = scope.get("scheme", "http")
  169. headers = _process_headers(scope.get("headers", ()))
  170. client_address = f"{ca[0]}:{ca[1]}" if (ca := scope.get("client")) else None
  171. content_type = headers.get("content-type", "")
  172. protocol = negotiate_server_protocol(content_type)
  173. if protocol.uses_trailers() and "http.response.trailers" not in cast(
  174. "dict", scope.get("extensions", {})
  175. ):
  176. msg = f"ASGI server does not support ASGI trailers extension but protocol for content-type '{content_type}' requires trailers"
  177. raise RuntimeError(msg)
  178. ctx = protocol.create_request_context(
  179. endpoint.method, http_method, http_scheme, headers, client_address
  180. )
  181. is_unary = isinstance(endpoint, EndpointUnary)
  182. if http_method == "GET":
  183. query_string = scope.get("query_string", b"").decode("utf-8")
  184. query_params = parse_qs(query_string, keep_blank_values=True)
  185. codec_name = query_params.get("encoding", ("",))[0]
  186. else:
  187. query_params = _UNSET_QUERY_PARAMS
  188. codec_name = protocol.codec_name_from_content_type(
  189. headers.get("content-type", ""), stream=not is_unary
  190. )
  191. codec = self._codecs.get(codec_name)
  192. if not codec:
  193. raise HTTPException(
  194. HTTPStatus.UNSUPPORTED_MEDIA_TYPE,
  195. [("Accept-Post", "application/json, application/proto")],
  196. )
  197. if is_unary and isinstance(protocol, ConnectServerProtocol):
  198. return await self._handle_unary_connect(
  199. http_method,
  200. headers,
  201. codec,
  202. query_params,
  203. endpoint,
  204. receive,
  205. send,
  206. ctx,
  207. )
  208. except Exception as e:
  209. await self._handle_error(e, ctx, send)
  210. if not isinstance(e, (ConnectError, HTTPException)):
  211. raise
  212. return None
  213. # Streams have their own error handling so move out of the try block.
  214. return await self._handle_stream(
  215. receive, send, protocol, endpoint, codec, headers, ctx
  216. )
  217. async def _handle_unary_connect(
  218. self,
  219. http_method: str,
  220. headers: Headers,
  221. codec: Codec,
  222. query_params: dict[str, list[str]],
  223. endpoint: EndpointUnary[_REQ, _RES],
  224. receive: ASGIReceiveCallable,
  225. send: ASGISendCallable,
  226. ctx: RequestContext,
  227. ) -> None:
  228. accept_encoding = headers.get("accept-encoding", "")
  229. compression = negotiate_compression(accept_encoding, self._compressions)
  230. if http_method == "GET":
  231. request = await self._read_get_request(endpoint, codec, query_params)
  232. else:
  233. request = await self._read_post_request(endpoint, receive, codec, headers)
  234. response_data = await endpoint.function(request, ctx)
  235. res_bytes = codec.encode(response_data)
  236. response_headers: list[tuple[bytes, bytes]] = [
  237. (
  238. b"content-type",
  239. f"{CONNECT_UNARY_CONTENT_TYPE_PREFIX}{codec.name()}".encode(),
  240. )
  241. ]
  242. res_bytes = compression.compress(res_bytes)
  243. response_headers.append((b"content-encoding", compression.name().encode()))
  244. response_headers.append((b"vary", b"Accept-Encoding"))
  245. _add_context_headers(response_headers, ctx)
  246. await send(
  247. {
  248. "type": "http.response.start",
  249. "status": 200,
  250. "headers": response_headers,
  251. "trailers": False,
  252. }
  253. )
  254. await send(
  255. {"type": "http.response.body", "body": res_bytes, "more_body": False}
  256. )
  257. async def _read_get_request(
  258. self,
  259. endpoint: EndpointUnary[_REQ, _RES],
  260. codec: Codec,
  261. params: dict[str, list[str]],
  262. ) -> _REQ:
  263. """Handle GET request with query parameters."""
  264. # Validation
  265. if "message" not in params:
  266. raise ConnectError(
  267. Code.INVALID_ARGUMENT,
  268. "'message' parameter is required for GET requests",
  269. )
  270. # Get and decode message
  271. message = params["message"][0]
  272. is_base64 = "base64" in params and params["base64"][0] == "1"
  273. if is_base64:
  274. try:
  275. message = base64.urlsafe_b64decode(message + "===")
  276. except Exception as e:
  277. raise ConnectError(
  278. Code.INVALID_ARGUMENT, "Invalid base64 encoding"
  279. ) from e
  280. else:
  281. message = message.encode("utf-8")
  282. # Handle compression
  283. compression_name = params.get("compression", ["identity"])[0]
  284. compression = self._compressions.get(compression_name)
  285. if not compression:
  286. raise ConnectError(
  287. Code.UNIMPLEMENTED,
  288. f"unknown compression: '{compression_name}': supported encodings are {', '.join(self._compressions.keys())}",
  289. )
  290. # Decompress and decode message
  291. if message: # Don't decompress empty messages
  292. message = compression.decompress(message)
  293. # Get the appropriate decoder for the endpoint
  294. return codec.decode(message, endpoint.method.input)
  295. async def _read_post_request(
  296. self,
  297. endpoint: Endpoint[_REQ, _RES],
  298. receive: ASGIReceiveCallable,
  299. codec: Codec,
  300. headers: Headers,
  301. ) -> _REQ:
  302. """Handle POST request with body."""
  303. # Get request body
  304. chunks: list[bytes] = [chunk async for chunk in _read_body(receive)]
  305. req_body = b"".join(chunks)
  306. # Handle compression if specified
  307. compression_name = headers.get("content-encoding", "identity").lower()
  308. compression = self._compressions.get(compression_name)
  309. if not compression:
  310. raise ConnectError(
  311. Code.UNIMPLEMENTED,
  312. f"unknown compression: '{compression_name}': supported encodings are {', '.join(self._compressions.keys())}",
  313. )
  314. if req_body: # Don't decompress empty body
  315. req_body = compression.decompress(req_body)
  316. if self._read_max_bytes is not None and len(req_body) > self._read_max_bytes:
  317. raise ConnectError(
  318. Code.RESOURCE_EXHAUSTED,
  319. f"message is larger than configured max {self._read_max_bytes}",
  320. )
  321. return codec.decode(req_body, endpoint.method.input)
  322. async def _handle_stream(
  323. self,
  324. receive: ASGIReceiveCallable,
  325. send: ASGISendCallable,
  326. protocol: ServerProtocol,
  327. endpoint: Endpoint[_REQ, _RES],
  328. codec: Codec,
  329. headers: Headers,
  330. ctx: _server_shared.RequestContext,
  331. ) -> None:
  332. req_compression, resp_compression = protocol.negotiate_stream_compression(
  333. headers, self._compressions
  334. )
  335. writer = protocol.create_envelope_writer(codec, resp_compression)
  336. error: Exception | None = None
  337. sent_headers = False
  338. try:
  339. if not req_compression:
  340. raise ConnectError(
  341. Code.UNIMPLEMENTED, "Unrecognized request compression"
  342. )
  343. request_stream = _request_stream(
  344. receive,
  345. endpoint.method.input,
  346. codec,
  347. req_compression,
  348. self._read_max_bytes,
  349. )
  350. disconnect_detected: Event | None = None
  351. monitor_task = None
  352. match endpoint:
  353. case EndpointUnary():
  354. request = await _consume_single_request(request_stream)
  355. response = await endpoint.function(request, ctx)
  356. response_stream = _yield_single_response(response)
  357. case EndpointClientStream():
  358. response = await endpoint.function(request_stream, ctx)
  359. response_stream = _yield_single_response(response)
  360. case EndpointServerStream():
  361. request = await _consume_single_request(request_stream)
  362. response_stream = endpoint.function(request, ctx)
  363. # The request has been fully consumed; monitor receive() for a
  364. # client disconnect so we can stop streaming promptly.
  365. disconnect_detected = Event()
  366. async def _watch_for_disconnect() -> None:
  367. while True:
  368. msg = await receive()
  369. if msg["type"] == "http.disconnect":
  370. disconnect_detected.set()
  371. return
  372. monitor_task = create_task(_watch_for_disconnect())
  373. case EndpointBidiStream():
  374. response_stream = endpoint.function(request_stream, ctx)
  375. try:
  376. async for message in response_stream:
  377. if disconnect_detected is not None and disconnect_detected.is_set():
  378. raise ConnectError(Code.CANCELED, "Client disconnected")
  379. # Don't send headers until the first message to allow logic a chance to add
  380. # response headers.
  381. if not sent_headers:
  382. await _send_stream_response_headers(
  383. send, protocol, codec, resp_compression.name(), ctx
  384. )
  385. sent_headers = True
  386. body = writer.write(message)
  387. await send(
  388. {"type": "http.response.body", "body": body, "more_body": True}
  389. )
  390. finally:
  391. # Cancel the monitor first so a throwing generator finally-block
  392. # doesn't leak the task.
  393. if monitor_task is not None:
  394. monitor_task.cancel()
  395. with contextlib.suppress(CancelledError):
  396. await monitor_task
  397. # Explicitly close the stream so that any generator finally-blocks
  398. # run promptly (Python defers async-generator cleanup to GC otherwise).
  399. aclose = getattr(response_stream, "aclose", None)
  400. if aclose is not None:
  401. await aclose()
  402. except CancelledError as e:
  403. raise ConnectError(Code.CANCELED, "Request was cancelled") from e
  404. except Exception as e:
  405. error = e
  406. finally:
  407. end_message = writer.end(
  408. ctx.response_trailers,
  409. ConnectWireError.from_exception(error) if error else None,
  410. )
  411. if not sent_headers:
  412. # Exception before any response message is returned
  413. await _send_stream_response_headers(
  414. send, protocol, codec, resp_compression.name(), ctx
  415. )
  416. if isinstance(end_message, bytes):
  417. await send(
  418. {
  419. "type": "http.response.body",
  420. "body": end_message,
  421. "more_body": False,
  422. }
  423. )
  424. else:
  425. await send(
  426. {"type": "http.response.body", "body": b"", "more_body": False}
  427. )
  428. await send(
  429. {
  430. "type": "http.response.trailers",
  431. "headers": [
  432. (k.encode(), v.encode()) for k, v in end_message.allitems()
  433. ],
  434. "more_trailers": False,
  435. }
  436. )
  437. if error and not isinstance(error, ConnectError):
  438. raise error
  439. async def _handle_error(
  440. self, exc: Exception, ctx: RequestContext | None, send: ASGISendCallable
  441. ) -> None:
  442. """Handle errors that occur during request processing."""
  443. headers: list[tuple[bytes, bytes]]
  444. body: bytes
  445. status: int
  446. if isinstance(exc, HTTPException):
  447. status = exc.status.value
  448. headers = [(k.encode("utf-8"), v.encode("utf-8")) for k, v in exc.headers]
  449. body = b""
  450. else:
  451. wire_error = ConnectWireError.from_exception(exc)
  452. status = wire_error.to_http_status().code
  453. headers = [(b"content-type", b"application/json")]
  454. body = wire_error.to_json_bytes()
  455. if ctx:
  456. _add_context_headers(headers, ctx)
  457. await send(
  458. {
  459. "type": "http.response.start",
  460. "status": status,
  461. "headers": headers,
  462. "trailers": False,
  463. }
  464. )
  465. await send({"type": "http.response.body", "body": body, "more_body": False})
  466. def _resolve_endpoints(self, service: _SVC) -> Mapping[str, Endpoint]:
  467. resolved_endpoints = self._endpoints(service)
  468. if self._interceptors:
  469. resolved_endpoints = {
  470. path: _apply_interceptors(
  471. endpoint, resolve_interceptors(self._interceptors)
  472. )
  473. for path, endpoint in resolved_endpoints.items()
  474. }
  475. return resolved_endpoints
  476. async def _send_stream_response_headers(
  477. send: ASGISendCallable,
  478. protocol: ServerProtocol,
  479. codec: Codec,
  480. compression_name: str,
  481. ctx: RequestContext,
  482. ) -> None:
  483. response_headers = [
  484. (b"content-type", protocol.content_type(codec).encode()),
  485. (protocol.compression_header_name().encode(), compression_name.encode()),
  486. ]
  487. response_headers.extend(
  488. (key.encode(), value.encode()) for key, value in ctx.response_headers.allitems()
  489. )
  490. await send(
  491. {
  492. "type": "http.response.start",
  493. "status": 200,
  494. "headers": response_headers,
  495. "trailers": protocol.uses_trailers(),
  496. }
  497. )
  498. async def _request_stream(
  499. receive: ASGIReceiveCallable,
  500. request_class: type[_REQ],
  501. codec: Codec,
  502. compression: Compression,
  503. read_max_bytes: int | None = None,
  504. ) -> AsyncIterator[_REQ]:
  505. reader = EnvelopeReader(request_class, codec, compression, read_max_bytes)
  506. try:
  507. async for chunk in _read_body(receive):
  508. for message in reader.feed(chunk):
  509. yield message
  510. # Check for cancellation each message. While this seems heavyweight,
  511. # conformance tests require it.
  512. await sleep(0)
  513. except CancelledError as e:
  514. raise ConnectError(Code.CANCELED, "Request was cancelled") from e
  515. async def _read_body(receive: ASGIReceiveCallable) -> AsyncIterator[bytes]:
  516. """Read the body of the request."""
  517. while True:
  518. message = await receive()
  519. match message["type"]:
  520. case "http.request":
  521. body = message.get("body", b"")
  522. yield body
  523. if not message.get("more_body", False):
  524. return
  525. case "http.disconnect":
  526. raise ConnectError(
  527. Code.CANCELED, "Client disconnected before request completion"
  528. )
  529. case _:
  530. raise ConnectError(Code.UNKNOWN, "Unexpected message type")
  531. async def _consume_single_request(stream: AsyncIterator[_REQ]) -> _REQ:
  532. req = None
  533. async for message in stream:
  534. if req is not None:
  535. raise ConnectError(
  536. Code.UNIMPLEMENTED, "unary request has multiple messages"
  537. )
  538. req = message
  539. if req is None:
  540. raise ConnectError(Code.UNIMPLEMENTED, "unary request has zero messages")
  541. return req
  542. async def _yield_single_response(response: _RES) -> AsyncIterator[_RES]:
  543. yield response
  544. def _process_headers(iterable: Iterable[tuple[bytes, bytes]]) -> Headers:
  545. result = Headers()
  546. for key, value in iterable:
  547. result.add(key.decode(), value.decode())
  548. return result
  549. def _add_context_headers(
  550. headers: list[tuple[bytes, bytes]], ctx: RequestContext
  551. ) -> None:
  552. headers.extend(
  553. (key.encode(), value.encode()) for key, value in ctx.response_headers.allitems()
  554. )
  555. headers.extend(
  556. (f"trailer-{key}".encode(), value.encode())
  557. for key, value in ctx.response_trailers.allitems()
  558. )
  559. def _apply_interceptors(
  560. endpoint: Endpoint[_REQ, _RES], interceptors: Sequence[Interceptor]
  561. ) -> Endpoint[_REQ, _RES]:
  562. match endpoint:
  563. case EndpointUnary():
  564. func = endpoint.function
  565. for interceptor in reversed(interceptors):
  566. if not isinstance(interceptor, UnaryInterceptor):
  567. continue
  568. func = functools.partial(interceptor.intercept_unary, func)
  569. return replace(endpoint, function=func)
  570. case EndpointClientStream():
  571. func = endpoint.function
  572. for interceptor in reversed(interceptors):
  573. if not isinstance(interceptor, ClientStreamInterceptor):
  574. continue
  575. func = functools.partial(interceptor.intercept_client_stream, func)
  576. return replace(endpoint, function=func)
  577. case EndpointServerStream():
  578. func = endpoint.function
  579. for interceptor in reversed(interceptors):
  580. if not isinstance(interceptor, ServerStreamInterceptor):
  581. continue
  582. func = functools.partial(interceptor.intercept_server_stream, func)
  583. return replace(endpoint, function=func)
  584. case EndpointBidiStream():
  585. func = endpoint.function
  586. for interceptor in reversed(interceptors):
  587. if not isinstance(interceptor, BidiStreamInterceptor):
  588. continue
  589. func = functools.partial(interceptor.intercept_bidi_stream, func)
  590. return replace(endpoint, function=func)