Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/starlette/websockets.py: 25%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1from __future__ import annotations
3import enum
4import json
5from collections.abc import AsyncIterator, Iterable
6from typing import Any, cast
8from starlette.requests import HTTPConnection, StateT
9from starlette.responses import Response
10from starlette.types import Message, Receive, Scope, Send
13class WebSocketState(enum.Enum):
14 CONNECTING = 0
15 CONNECTED = 1
16 DISCONNECTED = 2
17 RESPONSE = 3
20class WebSocketDisconnect(Exception):
21 def __init__(self, code: int = 1000, reason: str | None = None) -> None:
22 self.code = code
23 self.reason = reason or ""
26class WebSocketDisconnected(RuntimeError):
27 """
28 Raised when attempting to use a disconnected WebSocket.
29 """
32class WebSocket(HTTPConnection[StateT]):
33 def __init__(self, scope: Scope, receive: Receive, send: Send) -> None:
34 super().__init__(scope)
35 assert scope["type"] == "websocket"
36 self._receive = receive
37 self._send = send
38 self.client_state = WebSocketState.CONNECTING
39 self.application_state = WebSocketState.CONNECTING
41 async def receive(self) -> Message:
42 """
43 Receive ASGI websocket messages, ensuring valid state transitions.
44 """
45 if self.client_state == WebSocketState.CONNECTING:
46 message = await self._receive()
47 message_type = message["type"]
48 if message_type != "websocket.connect":
49 raise RuntimeError(f'Expected ASGI message "websocket.connect", but got {message_type!r}')
50 self.client_state = WebSocketState.CONNECTED
51 return message
52 elif self.client_state == WebSocketState.CONNECTED:
53 message = await self._receive()
54 message_type = message["type"]
55 if message_type not in {"websocket.receive", "websocket.disconnect"}:
56 raise RuntimeError(
57 f'Expected ASGI message "websocket.receive" or "websocket.disconnect", but got {message_type!r}'
58 )
59 if message_type == "websocket.disconnect":
60 self.client_state = WebSocketState.DISCONNECTED
61 return message
62 else:
63 raise WebSocketDisconnected('Cannot call "receive" once a disconnect message has been received.')
65 async def send(self, message: Message) -> None:
66 """
67 Send ASGI websocket messages, ensuring valid state transitions.
68 """
69 if self.application_state == WebSocketState.CONNECTING:
70 message_type = message["type"]
71 if message_type not in {"websocket.accept", "websocket.close", "websocket.http.response.start"}:
72 raise RuntimeError(
73 'Expected ASGI message "websocket.accept", "websocket.close" or "websocket.http.response.start", '
74 f"but got {message_type!r}"
75 )
76 if message_type == "websocket.close":
77 self.application_state = WebSocketState.DISCONNECTED
78 elif message_type == "websocket.http.response.start":
79 self.application_state = WebSocketState.RESPONSE
80 else:
81 self.application_state = WebSocketState.CONNECTED
82 await self._send(message)
83 elif self.application_state == WebSocketState.CONNECTED:
84 message_type = message["type"]
85 if message_type not in {"websocket.send", "websocket.close"}:
86 raise RuntimeError(
87 f'Expected ASGI message "websocket.send" or "websocket.close", but got {message_type!r}'
88 )
89 if message_type == "websocket.close":
90 self.application_state = WebSocketState.DISCONNECTED
91 try:
92 await self._send(message)
93 except OSError:
94 self.application_state = WebSocketState.DISCONNECTED
95 raise WebSocketDisconnect(code=1006)
96 elif self.application_state == WebSocketState.RESPONSE:
97 message_type = message["type"]
98 if message_type != "websocket.http.response.body":
99 raise RuntimeError(f'Expected ASGI message "websocket.http.response.body", but got {message_type!r}')
100 if not message.get("more_body", False):
101 self.application_state = WebSocketState.DISCONNECTED
102 await self._send(message)
103 else:
104 raise WebSocketDisconnected('Cannot call "send" once a close message has been sent.')
106 async def accept(
107 self,
108 subprotocol: str | None = None,
109 headers: Iterable[tuple[bytes, bytes]] | None = None,
110 ) -> None:
111 headers = headers or []
113 if self.client_state == WebSocketState.CONNECTING: # pragma: no branch
114 # If we haven't yet seen the 'connect' message, then wait for it first.
115 await self.receive()
116 await self.send({"type": "websocket.accept", "subprotocol": subprotocol, "headers": headers})
118 def _raise_on_disconnect(self, message: Message) -> None:
119 if message["type"] == "websocket.disconnect":
120 raise WebSocketDisconnect(message["code"], message.get("reason"))
122 async def receive_text(self) -> str:
123 if self.application_state != WebSocketState.CONNECTED:
124 raise WebSocketDisconnected('WebSocket is not connected. Need to call "accept" first.')
125 message = await self.receive()
126 self._raise_on_disconnect(message)
127 return cast(str, message["text"])
129 async def receive_bytes(self) -> bytes:
130 if self.application_state != WebSocketState.CONNECTED:
131 raise WebSocketDisconnected('WebSocket is not connected. Need to call "accept" first.')
132 message = await self.receive()
133 self._raise_on_disconnect(message)
134 return cast(bytes, message["bytes"])
136 async def receive_json(self, mode: str = "text") -> Any:
137 if mode not in {"text", "binary"}:
138 raise RuntimeError('The "mode" argument should be "text" or "binary".')
139 if self.application_state != WebSocketState.CONNECTED:
140 raise WebSocketDisconnected('WebSocket is not connected. Need to call "accept" first.')
141 message = await self.receive()
142 self._raise_on_disconnect(message)
144 if mode == "text":
145 text = message["text"]
146 else:
147 text = message["bytes"].decode("utf-8")
148 return json.loads(text)
150 async def iter_text(self) -> AsyncIterator[str]:
151 try:
152 while True:
153 yield await self.receive_text()
154 except WebSocketDisconnect:
155 pass
157 async def iter_bytes(self) -> AsyncIterator[bytes]:
158 try:
159 while True:
160 yield await self.receive_bytes()
161 except WebSocketDisconnect:
162 pass
164 async def iter_json(self) -> AsyncIterator[Any]:
165 try:
166 while True:
167 yield await self.receive_json()
168 except WebSocketDisconnect:
169 pass
171 async def send_text(self, data: str) -> None:
172 await self.send({"type": "websocket.send", "text": data})
174 async def send_bytes(self, data: bytes) -> None:
175 await self.send({"type": "websocket.send", "bytes": data})
177 async def send_json(self, data: Any, mode: str = "text") -> None:
178 if mode not in {"text", "binary"}:
179 raise RuntimeError('The "mode" argument should be "text" or "binary".')
180 text = json.dumps(data, separators=(",", ":"), ensure_ascii=False)
181 if mode == "text":
182 await self.send({"type": "websocket.send", "text": text})
183 else:
184 await self.send({"type": "websocket.send", "bytes": text.encode("utf-8")})
186 async def close(self, code: int = 1000, reason: str | None = None) -> None:
187 await self.send({"type": "websocket.close", "code": code, "reason": reason or ""})
189 async def send_denial_response(self, response: Response) -> None:
190 if "websocket.http.response" in self.scope.get("extensions", {}):
191 await response(self.scope, self.receive, self.send)
192 else:
193 raise RuntimeError("The server doesn't support the Websocket Denial Response extension.")
196class WebSocketClose:
197 def __init__(self, code: int = 1000, reason: str | None = None) -> None:
198 self.code = code
199 self.reason = reason or ""
201 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
202 await send({"type": "websocket.close", "code": self.code, "reason": self.reason})