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

145 statements  

1from __future__ import annotations 

2 

3import enum 

4import json 

5from collections.abc import AsyncIterator, Iterable 

6from typing import Any, cast 

7 

8from starlette.requests import HTTPConnection, StateT 

9from starlette.responses import Response 

10from starlette.types import Message, Receive, Scope, Send 

11 

12 

13class WebSocketState(enum.Enum): 

14 CONNECTING = 0 

15 CONNECTED = 1 

16 DISCONNECTED = 2 

17 RESPONSE = 3 

18 

19 

20class WebSocketDisconnect(Exception): 

21 def __init__(self, code: int = 1000, reason: str | None = None) -> None: 

22 self.code = code 

23 self.reason = reason or "" 

24 

25 

26class WebSocketDisconnected(RuntimeError): 

27 """ 

28 Raised when attempting to use a disconnected WebSocket. 

29 """ 

30 

31 

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 

40 

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.') 

64 

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.') 

105 

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 [] 

112 

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}) 

117 

118 def _raise_on_disconnect(self, message: Message) -> None: 

119 if message["type"] == "websocket.disconnect": 

120 raise WebSocketDisconnect(message["code"], message.get("reason")) 

121 

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"]) 

128 

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"]) 

135 

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) 

143 

144 if mode == "text": 

145 text = message["text"] 

146 else: 

147 text = message["bytes"].decode("utf-8") 

148 return json.loads(text) 

149 

150 async def iter_text(self) -> AsyncIterator[str]: 

151 try: 

152 while True: 

153 yield await self.receive_text() 

154 except WebSocketDisconnect: 

155 pass 

156 

157 async def iter_bytes(self) -> AsyncIterator[bytes]: 

158 try: 

159 while True: 

160 yield await self.receive_bytes() 

161 except WebSocketDisconnect: 

162 pass 

163 

164 async def iter_json(self) -> AsyncIterator[Any]: 

165 try: 

166 while True: 

167 yield await self.receive_json() 

168 except WebSocketDisconnect: 

169 pass 

170 

171 async def send_text(self, data: str) -> None: 

172 await self.send({"type": "websocket.send", "text": data}) 

173 

174 async def send_bytes(self, data: bytes) -> None: 

175 await self.send({"type": "websocket.send", "bytes": data}) 

176 

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")}) 

185 

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 ""}) 

188 

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.") 

194 

195 

196class WebSocketClose: 

197 def __init__(self, code: int = 1000, reason: str | None = None) -> None: 

198 self.code = code 

199 self.reason = reason or "" 

200 

201 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: 

202 await send({"type": "websocket.close", "code": self.code, "reason": self.reason})