Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/utilities/subscriptions.py: 21%

58 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 02:04 +0000

1import asyncio 

2import hmac 

3from asyncio import IncompleteReadError as IOError 

4from logging import Logger 

5from typing import Optional 

6 

7from fastapi import WebSocket 

8from starlette.status import WS_1002_PROTOCOL_ERROR, WS_1008_POLICY_VIOLATION 

9from starlette.websockets import WebSocketDisconnect 

10from websockets.exceptions import ConnectionClosed 

11 

12from prefect.logging import get_logger 

13from prefect.settings import get_current_settings 

14 

15NORMAL_DISCONNECT_EXCEPTIONS = (IOError, ConnectionClosed, WebSocketDisconnect) 

16 

17logger: Logger = get_logger("prefect.server.utilities.subscriptions") 

18 

19 

20async def accept_prefect_socket( 

21 websocket: WebSocket, 

22 *, 

23 require_prefect_subprotocol: bool = False, 

24 authentication_failed_reason: str | None = None, 

25) -> Optional[WebSocket]: 

26 subprotocols = [ 

27 subprotocol.strip() 

28 for header in websocket.headers.getlist("Sec-WebSocket-Protocol") 

29 for subprotocol in header.split(",") 

30 ] 

31 has_prefect_subprotocol = "prefect" in subprotocols 

32 

33 auth_setting = ( 

34 auth_setting_secret.get_secret_value() 

35 if (auth_setting_secret := get_current_settings().server.api.auth_string) 

36 else None 

37 ) 

38 

39 # If client doesn't send "prefect" subprotocol: 

40 # - Reject if auth is configured (security requirement) 

41 # - Accept in legacy mode if auth is not configured (backward compatibility) 

42 if not has_prefect_subprotocol: 

43 if auth_setting or require_prefect_subprotocol: 

44 reason = ( 

45 "'prefect' subprotocol required" 

46 if require_prefect_subprotocol 

47 else "'prefect' subprotocol required when auth is configured" 

48 ) 

49 logger.warning("WebSocket connection rejected: %s", reason) 

50 return await websocket.close( 

51 WS_1002_PROTOCOL_ERROR 

52 if not require_prefect_subprotocol 

53 else WS_1008_POLICY_VIOLATION, 

54 reason=authentication_failed_reason, 

55 ) 

56 else: 

57 # Legacy mode: accept without auth handshake for old clients 

58 logger.debug( 

59 "Accepting WebSocket in legacy mode (no 'prefect' subprotocol)" 

60 ) 

61 await websocket.accept() 

62 return websocket 

63 

64 # New protocol: client sent "prefect" subprotocol, perform auth handshake 

65 await websocket.accept(subprotocol="prefect") 

66 

67 try: 

68 # Websocket connections are authenticated via messages. The first 

69 # message is expected to be an auth message, and if any other type of 

70 # message is received then the connection will be closed. 

71 # 

72 # The protocol requires receiving an auth message for compatibility 

73 # with Prefect Cloud, even if server-side auth is not configured. 

74 message = await websocket.receive_json() 

75 logger.debug( 

76 f"PREFECT_SERVER_API_AUTH_STRING setting: {'*' * len(auth_setting) if auth_setting else 'Not set'}" 

77 ) 

78 

79 if message.get("type") != "auth": 

80 logger.warning( 

81 "WebSocket connection closed: Expected 'auth' message first." 

82 ) 

83 return await websocket.close( 

84 WS_1008_POLICY_VIOLATION, 

85 reason=authentication_failed_reason or "Expected 'auth' message", 

86 ) 

87 

88 # Check authentication if PREFECT_SERVER_API_AUTH_STRING is set 

89 if auth_setting: 

90 received_token = message.get("token") 

91 logger.debug( 

92 f"Auth required. Received token: {'*' * len(received_token) if received_token else 'None'}" 

93 ) 

94 if not received_token: 

95 logger.warning( 

96 "WebSocket connection closed: Auth required but no token received." 

97 ) 

98 return await websocket.close( 

99 WS_1008_POLICY_VIOLATION, 

100 reason=( 

101 authentication_failed_reason 

102 or "Auth required but no token provided" 

103 ), 

104 ) 

105 

106 if not hmac.compare_digest(received_token, auth_setting): 

107 logger.warning("WebSocket connection closed: Invalid token.") 

108 return await websocket.close( 

109 WS_1008_POLICY_VIOLATION, 

110 reason=authentication_failed_reason or "Invalid token", 

111 ) 

112 logger.debug("WebSocket token authentication successful.") 

113 else: 

114 logger.debug("No server auth string set, skipping token check.") 

115 

116 await websocket.send_json({"type": "auth_success"}) 

117 logger.debug("Sent auth_success to WebSocket.") 

118 return websocket 

119 

120 except NORMAL_DISCONNECT_EXCEPTIONS: 

121 # it's fine if a client disconnects either normally or abnormally 

122 return None 

123 

124 

125async def still_connected(websocket: WebSocket) -> bool: 

126 """Checks that a client websocket still seems to be connected during a period where 

127 the server is expected to be sending messages.""" 

128 try: 

129 await asyncio.wait_for(websocket.receive(), timeout=0.1) 

130 return True # this should never happen, but if it does, we're still connected 

131 except asyncio.TimeoutError: 

132 # The fact that we timed out rather than getting another kind of error 

133 # here means we're still connected to our client, so we can continue to send 

134 # events. 

135 return True 

136 except RuntimeError: 

137 # starlette raises this if we test a client that's disconnected 

138 return False 

139 except NORMAL_DISCONNECT_EXCEPTIONS: 

140 return False