endpoints.py 5.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127
  1. from __future__ import annotations
  2. import json
  3. from collections.abc import Callable, Generator
  4. from typing import Any, Literal
  5. from starlette import status
  6. from starlette._utils import is_async_callable
  7. from starlette.concurrency import run_in_threadpool
  8. from starlette.exceptions import HTTPException
  9. from starlette.requests import Request
  10. from starlette.responses import PlainTextResponse, Response
  11. from starlette.types import Message, Receive, Scope, Send
  12. from starlette.websockets import WebSocket
  13. class HTTPEndpoint:
  14. def __init__(self, scope: Scope, receive: Receive, send: Send) -> None:
  15. assert scope["type"] == "http"
  16. self.scope = scope
  17. self.receive = receive
  18. self.send = send
  19. self._allowed_methods = [
  20. method
  21. for method in ("GET", "HEAD", "POST", "PUT", "PATCH", "DELETE", "OPTIONS")
  22. if getattr(self, method.lower(), None) is not None
  23. ]
  24. def __await__(self) -> Generator[Any, None, None]:
  25. return self.dispatch().__await__()
  26. async def dispatch(self) -> None:
  27. request = Request(self.scope, receive=self.receive)
  28. handler_name = "get" if request.method == "HEAD" and not hasattr(self, "head") else request.method.lower()
  29. handler: Callable[[Request], Any]
  30. if request.method in self._allowed_methods or (request.method == "HEAD" and "GET" in self._allowed_methods):
  31. handler = getattr(self, handler_name)
  32. else:
  33. handler = self.method_not_allowed
  34. is_async = is_async_callable(handler)
  35. if is_async:
  36. response = await handler(request)
  37. else:
  38. response = await run_in_threadpool(handler, request)
  39. await response(self.scope, self.receive, self.send)
  40. async def method_not_allowed(self, request: Request) -> Response:
  41. # If we're running inside a starlette application then raise an
  42. # exception, so that the configurable exception handler can deal with
  43. # returning the response. For plain ASGI apps, just return the response.
  44. headers = {"Allow": ", ".join(self._allowed_methods)}
  45. if "app" in self.scope:
  46. raise HTTPException(status_code=405, headers=headers)
  47. return PlainTextResponse("Method Not Allowed", status_code=405, headers=headers)
  48. class WebSocketEndpoint:
  49. encoding: Literal["text", "bytes", "json"] | None = None
  50. def __init__(self, scope: Scope, receive: Receive, send: Send) -> None:
  51. assert scope["type"] == "websocket"
  52. self.scope = scope
  53. self.receive = receive
  54. self.send = send
  55. def __await__(self) -> Generator[Any, None, None]:
  56. return self.dispatch().__await__()
  57. async def dispatch(self) -> None:
  58. websocket = WebSocket(self.scope, receive=self.receive, send=self.send)
  59. await self.on_connect(websocket)
  60. close_code = status.WS_1000_NORMAL_CLOSURE
  61. try:
  62. while True:
  63. message = await websocket.receive()
  64. if message["type"] == "websocket.receive":
  65. data = await self.decode(websocket, message)
  66. await self.on_receive(websocket, data)
  67. elif message["type"] == "websocket.disconnect": # pragma: no branch
  68. close_code = int(message.get("code") or status.WS_1000_NORMAL_CLOSURE)
  69. break
  70. except Exception as exc:
  71. close_code = status.WS_1011_INTERNAL_ERROR
  72. raise exc
  73. finally:
  74. await self.on_disconnect(websocket, close_code)
  75. async def decode(self, websocket: WebSocket, message: Message) -> Any:
  76. if self.encoding == "text":
  77. if "text" not in message:
  78. await websocket.close(code=status.WS_1003_UNSUPPORTED_DATA)
  79. raise RuntimeError("Expected text websocket messages, but got bytes")
  80. return message["text"]
  81. elif self.encoding == "bytes":
  82. if "bytes" not in message:
  83. await websocket.close(code=status.WS_1003_UNSUPPORTED_DATA)
  84. raise RuntimeError("Expected bytes websocket messages, but got text")
  85. return message["bytes"]
  86. elif self.encoding == "json":
  87. if message.get("text") is not None:
  88. text = message["text"]
  89. else:
  90. text = message["bytes"].decode("utf-8")
  91. try:
  92. return json.loads(text)
  93. except json.decoder.JSONDecodeError:
  94. await websocket.close(code=status.WS_1003_UNSUPPORTED_DATA)
  95. raise RuntimeError("Malformed JSON data received.")
  96. assert self.encoding is None, f"Unsupported 'encoding' attribute {self.encoding}"
  97. return message["text"] if message.get("text") else message["bytes"]
  98. async def on_connect(self, websocket: WebSocket) -> None:
  99. """Override to handle an incoming websocket connection"""
  100. await websocket.accept()
  101. async def on_receive(self, websocket: WebSocket, data: Any) -> None:
  102. """Override to handle an incoming websocket message"""
  103. async def on_disconnect(self, websocket: WebSocket, close_code: int) -> None:
  104. """Override to handle a disconnecting websocket"""