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
« 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
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
12from prefect.logging import get_logger
13from prefect.settings import get_current_settings
15NORMAL_DISCONNECT_EXCEPTIONS = (IOError, ConnectionClosed, WebSocketDisconnect)
17logger: Logger = get_logger("prefect.server.utilities.subscriptions")
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
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 )
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
64 # New protocol: client sent "prefect" subprotocol, perform auth handshake
65 await websocket.accept(subprotocol="prefect")
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 )
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 )
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 )
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.")
116 await websocket.send_json({"type": "auth_success"})
117 logger.debug("Sent auth_success to WebSocket.")
118 return websocket
120 except NORMAL_DISCONNECT_EXCEPTIONS:
121 # it's fine if a client disconnects either normally or abnormally
122 return None
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