main
py 276 lines 7.83 KB
Raw
1 import sys
2 from pathlib import Path
3 from unittest.mock import patch
4
5 import pytest
6
7 PROJECT_ROOT = Path(__file__).resolve().parents[1]
8 if str(PROJECT_ROOT) not in sys.path:
9 sys.path.insert(0, str(PROJECT_ROOT))
10
11 from helpers.ws import WsHandler, _SecurityContext, _check_security
12
13
14 # Handler variants with different security flag combinations
15
16 class _OpenHandler(WsHandler):
17 @classmethod
18 def requires_auth(cls) -> bool:
19 return False
20
21 @classmethod
22 def requires_csrf(cls) -> bool:
23 return False
24
25 async def process(self, event, data, sid):
26 return None
27
28
29 class _AuthOnlyHandler(WsHandler):
30 @classmethod
31 def requires_auth(cls) -> bool:
32 return True
33
34 @classmethod
35 def requires_csrf(cls) -> bool:
36 return False
37
38 async def process(self, event, data, sid):
39 return None
40
41
42 class _CsrfHandler(WsHandler):
43 """Default: requires_auth=True, requires_csrf=True (via default)."""
44
45 async def process(self, event, data, sid):
46 return None
47
48
49 class _LoopbackHandler(WsHandler):
50 @classmethod
51 def requires_loopback(cls) -> bool:
52 return True
53
54 @classmethod
55 def requires_auth(cls) -> bool:
56 return False
57
58 @classmethod
59 def requires_csrf(cls) -> bool:
60 return False
61
62 async def process(self, event, data, sid):
63 return None
64
65
66 class _ApiKeyHandler(WsHandler):
67 @classmethod
68 def requires_api_key(cls) -> bool:
69 return True
70
71 @classmethod
72 def requires_auth(cls) -> bool:
73 return False
74
75 @classmethod
76 def requires_csrf(cls) -> bool:
77 return False
78
79 async def process(self, event, data, sid):
80 return None
81
82
83 # Helper to build _SecurityContext quickly
84
85 def _ctx(
86 *,
87 auth_hash=None,
88 csrf_token="tok",
89 client_csrf_token="tok",
90 csrf_cookie="tok",
91 remote_addr="127.0.0.1",
92 api_key=None,
93 ) -> _SecurityContext:
94 return _SecurityContext(
95 auth_hash=auth_hash,
96 csrf_token=csrf_token,
97 client_csrf_token=client_csrf_token,
98 csrf_cookie=csrf_cookie,
99 remote_addr=remote_addr,
100 api_key=api_key,
101 )
102
103
104 # Open handler (no security)
105
106 def test_open_handler_always_passes():
107 assert _check_security(_OpenHandler, _ctx()) is None
108
109
110 def test_open_handler_passes_even_without_tokens():
111 assert _check_security(_OpenHandler, _ctx(csrf_token=None, client_csrf_token=None, csrf_cookie=None)) is None
112
113
114 # Loopback
115
116 def test_loopback_allows_127_0_0_1():
117 assert _check_security(_LoopbackHandler, _ctx(remote_addr="127.0.0.1")) is None
118
119
120 def test_loopback_allows_ipv6():
121 assert _check_security(_LoopbackHandler, _ctx(remote_addr="::1")) is None
122
123
124 def test_loopback_rejects_remote():
125 result = _check_security(_LoopbackHandler, _ctx(remote_addr="192.168.1.50"))
126 assert result is not None
127 assert result["code"] == "FORBIDDEN"
128
129
130 def test_loopback_rejects_none():
131 result = _check_security(_LoopbackHandler, _ctx(remote_addr=None))
132 assert result is not None
133 assert result["code"] == "FORBIDDEN"
134
135
136 # Auth
137
138 @patch("helpers.login.get_credentials_hash", return_value="hashed123")
139 def test_auth_passes_with_matching_hash(_mock):
140 result = _check_security(_AuthOnlyHandler, _ctx(auth_hash="hashed123"))
141 assert result is None
142
143
144 @patch("helpers.login.get_credentials_hash", return_value="hashed123")
145 def test_auth_rejects_wrong_hash(_mock):
146 result = _check_security(_AuthOnlyHandler, _ctx(auth_hash="wrong"))
147 assert result is not None
148 assert result["code"] == "AUTH_REQUIRED"
149
150
151 @patch("helpers.login.get_credentials_hash", return_value="hashed123")
152 def test_auth_rejects_missing_hash(_mock):
153 result = _check_security(_AuthOnlyHandler, _ctx(auth_hash=None))
154 assert result is not None
155 assert result["code"] == "AUTH_REQUIRED"
156
157
158 @patch("helpers.login.get_credentials_hash", return_value=None)
159 def test_auth_passes_when_no_credentials_configured(_mock):
160 """When no password is set (get_credentials_hash returns None/empty),
161 auth check should pass regardless of the client hash."""
162 result = _check_security(_AuthOnlyHandler, _ctx(auth_hash=None))
163 assert result is None
164
165
166 # CSRF
167
168 @patch("helpers.login.get_credentials_hash", return_value=None)
169 def test_csrf_passes_with_all_tokens_matching(_mock):
170 result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token="abc", csrf_cookie="abc"))
171 assert result is None
172
173
174 @patch("helpers.login.get_credentials_hash", return_value=None)
175 def test_csrf_rejects_missing_server_token(_mock):
176 result = _check_security(_CsrfHandler, _ctx(csrf_token=None, client_csrf_token="abc", csrf_cookie="abc"))
177 assert result is not None
178 assert result["code"] == "CSRF_MISSING"
179
180
181 @patch("helpers.login.get_credentials_hash", return_value=None)
182 def test_csrf_rejects_missing_client_token(_mock):
183 result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token=None, csrf_cookie="abc"))
184 assert result is not None
185 assert result["code"] == "CSRF_INVALID"
186
187
188 @patch("helpers.login.get_credentials_hash", return_value=None)
189 def test_csrf_rejects_mismatched_client_token(_mock):
190 result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token="xyz", csrf_cookie="abc"))
191 assert result is not None
192 assert result["code"] == "CSRF_INVALID"
193
194
195 @patch("helpers.login.get_credentials_hash", return_value=None)
196 def test_csrf_rejects_mismatched_cookie(_mock):
197 result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token="abc", csrf_cookie="wrong"))
198 assert result is not None
199 assert result["code"] == "CSRF_COOKIE"
200
201
202 # API Key
203
204 @patch("helpers.settings.get_settings", return_value={"mcp_server_token": "secret-key-123"})
205 def test_api_key_passes_with_correct_key(_mock):
206 result = _check_security(_ApiKeyHandler, _ctx(api_key="secret-key-123"))
207 assert result is None
208
209
210 @patch("helpers.settings.get_settings", return_value={"mcp_server_token": "secret-key-123"})
211 def test_api_key_rejects_wrong_key(_mock):
212 result = _check_security(_ApiKeyHandler, _ctx(api_key="wrong-key"))
213 assert result is not None
214 assert result["code"] == "API_KEY_REQUIRED"
215
216
217 @patch("helpers.settings.get_settings", return_value={"mcp_server_token": "secret-key-123"})
218 def test_api_key_rejects_missing_key(_mock):
219 result = _check_security(_ApiKeyHandler, _ctx(api_key=None))
220 assert result is not None
221 assert result["code"] == "API_KEY_REQUIRED"
222
223
224 # Combined flags
225
226 class _FullSecurityHandler(WsHandler):
227 @classmethod
228 def requires_loopback(cls) -> bool:
229 return True
230
231 @classmethod
232 def requires_auth(cls) -> bool:
233 return True
234
235 @classmethod
236 def requires_api_key(cls) -> bool:
237 return True
238
239 async def process(self, event, data, sid):
240 return None
241
242
243 @patch("helpers.login.get_credentials_hash", return_value="hash")
244 @patch("helpers.settings.get_settings", return_value={"mcp_server_token": "key"})
245 def test_full_security_passes_when_all_match(_mock_settings, _mock_login):
246 result = _check_security(
247 _FullSecurityHandler,
248 _ctx(
249 remote_addr="127.0.0.1",
250 auth_hash="hash",
251 csrf_token="tok",
252 client_csrf_token="tok",
253 csrf_cookie="tok",
254 api_key="key",
255 ),
256 )
257 assert result is None
258
259
260 @patch("helpers.login.get_credentials_hash", return_value="hash")
261 @patch("helpers.settings.get_settings", return_value={"mcp_server_token": "key"})
262 def test_full_security_fails_at_first_check_loopback(_mock_settings, _mock_login):
263 """Loopback check runs first; if it fails, later checks don't matter."""
264 result = _check_security(
265 _FullSecurityHandler,
266 _ctx(
267 remote_addr="10.0.0.1",
268 auth_hash="hash",
269 csrf_token="tok",
270 client_csrf_token="tok",
271 csrf_cookie="tok",
272 api_key="key",
273 ),
274 )
275 assert result is not None
276 assert result["code"] == "FORBIDDEN"