wsgi.py 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156
  1. from __future__ import annotations
  2. import io
  3. import math
  4. import sys
  5. import warnings
  6. from collections.abc import Callable, MutableMapping
  7. from typing import Any
  8. import anyio
  9. from anyio.abc import ObjectReceiveStream, ObjectSendStream
  10. from starlette._utils import create_collapsing_task_group
  11. from starlette.exceptions import StarletteDeprecationWarning
  12. from starlette.types import Receive, Scope, Send
  13. warnings.warn(
  14. "starlette.middleware.wsgi is deprecated and will be removed in a future release. "
  15. "Please refer to https://github.com/abersheeran/a2wsgi as a replacement.",
  16. StarletteDeprecationWarning,
  17. stacklevel=2,
  18. )
  19. def build_environ(scope: Scope, body: bytes) -> dict[str, Any]:
  20. """
  21. Builds a scope and request body into a WSGI environ object.
  22. """
  23. script_name = scope.get("root_path", "").encode("utf8").decode("latin1")
  24. path_info = scope["path"].encode("utf8").decode("latin1")
  25. if path_info.startswith(script_name):
  26. path_info = path_info[len(script_name) :]
  27. environ = {
  28. "REQUEST_METHOD": scope["method"],
  29. "SCRIPT_NAME": script_name,
  30. "PATH_INFO": path_info,
  31. "QUERY_STRING": scope["query_string"].decode("ascii"),
  32. "SERVER_PROTOCOL": f"HTTP/{scope['http_version']}",
  33. "wsgi.version": (1, 0),
  34. "wsgi.url_scheme": scope.get("scheme", "http"),
  35. "wsgi.input": io.BytesIO(body),
  36. "wsgi.errors": sys.stdout,
  37. "wsgi.multithread": True,
  38. "wsgi.multiprocess": True,
  39. "wsgi.run_once": False,
  40. }
  41. # Get server name and port - required in WSGI, not in ASGI
  42. server = scope.get("server") or ("localhost", 80)
  43. environ["SERVER_NAME"] = server[0]
  44. environ["SERVER_PORT"] = server[1]
  45. # Get client IP address
  46. if scope.get("client"):
  47. environ["REMOTE_ADDR"] = scope["client"][0]
  48. # Go through headers and make them into environ entries
  49. for name, value in scope.get("headers", []):
  50. name = name.decode("latin1")
  51. if name == "content-length":
  52. corrected_name = "CONTENT_LENGTH"
  53. elif name == "content-type":
  54. corrected_name = "CONTENT_TYPE"
  55. else:
  56. corrected_name = f"HTTP_{name}".upper().replace("-", "_")
  57. # HTTPbis say only ASCII chars are allowed in headers, but we latin1 just in
  58. # case
  59. value = value.decode("latin1")
  60. if corrected_name in environ:
  61. value = environ[corrected_name] + "," + value
  62. environ[corrected_name] = value
  63. return environ
  64. class WSGIMiddleware:
  65. def __init__(self, app: Callable[..., Any]) -> None:
  66. self.app = app
  67. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  68. assert scope["type"] == "http"
  69. responder = WSGIResponder(self.app, scope)
  70. await responder(receive, send)
  71. class WSGIResponder:
  72. stream_send: ObjectSendStream[MutableMapping[str, Any]]
  73. stream_receive: ObjectReceiveStream[MutableMapping[str, Any]]
  74. def __init__(self, app: Callable[..., Any], scope: Scope) -> None:
  75. self.app = app
  76. self.scope = scope
  77. self.status = None
  78. self.response_headers = None
  79. self.stream_send, self.stream_receive = anyio.create_memory_object_stream(math.inf)
  80. self.response_started = False
  81. self.exc_info: Any = None
  82. async def __call__(self, receive: Receive, send: Send) -> None:
  83. body = b""
  84. more_body = True
  85. while more_body:
  86. message = await receive()
  87. body += message.get("body", b"")
  88. more_body = message.get("more_body", False)
  89. environ = build_environ(self.scope, body)
  90. async with create_collapsing_task_group() as task_group:
  91. task_group.start_soon(self.sender, send)
  92. async with self.stream_send:
  93. await anyio.to_thread.run_sync(self.wsgi, environ, self.start_response)
  94. if self.exc_info is not None:
  95. raise self.exc_info[0].with_traceback(self.exc_info[1], self.exc_info[2])
  96. async def sender(self, send: Send) -> None:
  97. async with self.stream_receive:
  98. async for message in self.stream_receive:
  99. await send(message)
  100. def start_response(
  101. self,
  102. status: str,
  103. response_headers: list[tuple[str, str]],
  104. exc_info: Any = None,
  105. ) -> None:
  106. self.exc_info = exc_info
  107. if not self.response_started: # pragma: no branch
  108. self.response_started = True
  109. status_code_string, _ = status.split(" ", 1)
  110. status_code = int(status_code_string)
  111. headers = [
  112. (name.strip().encode("ascii").lower(), value.strip().encode("ascii"))
  113. for name, value in response_headers
  114. ]
  115. anyio.from_thread.run(
  116. self.stream_send.send,
  117. {
  118. "type": "http.response.start",
  119. "status": status_code,
  120. "headers": headers,
  121. },
  122. )
  123. def wsgi(
  124. self,
  125. environ: dict[str, Any],
  126. start_response: Callable[..., Any],
  127. ) -> None:
  128. for chunk in self.app(environ, start_response):
  129. anyio.from_thread.run(
  130. self.stream_send.send,
  131. {"type": "http.response.body", "body": chunk, "more_body": True},
  132. )
  133. anyio.from_thread.run(self.stream_send.send, {"type": "http.response.body", "body": b""})