tests/core/test_middleware_std.py¶
Source from this local checkout, regenerated when the reader rebuilds.
Line links use #L<number>; a GitHub line range opens its first line.
1 # Copyright 2025 Softwell S.r.l.2 #3 # Licensed under the Apache License, Version 2.0 (the "License");4 # you may not use this file except in compliance with the License.5 # You may obtain a copy of the License at6 #7 # https://www.apache.org/licenses/LICENSE-2.08 #9 # Unless required by applicable law or agreed to in writing, software10 # distributed under the License is distributed on an "AS IS" BASIS,11 # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12 # See the License for the specific language governing permissions and13 # limitations under the License.14 15 """Standard middlewares tests (Macro 2 Phase 3 + Macro 4 Phase 3 deferrals).16 17 Same ASGI-level driving style as ``tests/test_middleware.py`` (no uvicorn):18 a canned http scope, a recording ``send``, chain assembled through19 ``MiddlewareMixin`` composed over ``BaseServer``. The request-driving and20 message-reading helpers now live in ``tests/conftest.py`` as fixtures21 (``http_request``, ``response_status``, ``response_headers``, ``response_body``).22 23 The Macro 4 Phase 3 additions cover the ``ErrorMiddleware`` on the real24 ``Response`` class (wire equivalence, forwarded exception headers, the25 response-started guard) and the two previously-untested CORS branches26 (restricted origins rejecting a foreign origin).27 """28 29 from __future__ import annotations30 31 import json32 import logging33 from urllib.parse import quote34 35 import pytest36 37 from genro_asgi import (38 AsgiServer,39 BaseApplication,40 BaseServer,41 MemorySessionStore,42 MiddlewareMixin,43 )44 from genro_asgi.exceptions import HTTPNotFound, HTTPUnauthorized, Redirect45 from genro_asgi.types import Message, Receive, Scope, Send46 47 48 class MwServer(MiddlewareMixin, BaseServer):49 """Phase 2 composition: middleware capability over the base server."""50 51 52 class RoutedApp(BaseApplication):53 """Test app: a plain 200 for every path it is given."""54 55 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:56 await send(57 {58 "type": "http.response.start",59 "status": 200,60 "headers": [(b"content-type", b"text/plain; charset=utf-8")],61 }62 )63 await send({"type": "http.response.body", "body": f"ok:{scope['path']}".encode()})64 65 66 class TestWellKnownMiddleware:67 async def test_probe_path_returns_404(self, http_request, response_status) -> None:68 server = MwServer(applications=[RoutedApp(mount="")], middleware={"wellknown": True})69 sent = await http_request(server, "/.well-known/probe")70 assert response_status(sent) == 40471 72 async def test_ordinary_path_still_reaches_the_app(73 self, http_request, response_status, response_body74 ) -> None:75 server = MwServer(applications=[RoutedApp(mount="")], middleware={"wellknown": True})76 sent = await http_request(server, "/")77 assert response_status(sent) == 20078 assert response_body(sent) == b"ok:/"79 80 81 class TestCORSMiddleware:82 async def test_preflight_returns_cors_headers(83 self, http_request, response_status, response_headers84 ) -> None:85 server = MwServer(applications=[RoutedApp(mount="")], middleware={"cors": True})86 sent = await http_request(87 server, "/", method="OPTIONS", headers=[(b"origin", b"https://example.test")]88 )89 assert response_status(sent) in (200, 204)90 headers = response_headers(sent)91 assert headers[b"access-control-allow-origin"] == b"*"92 assert b"access-control-allow-methods" in headers93 94 async def test_simple_get_carries_allow_origin_header(95 self, http_request, response_status, response_headers96 ) -> None:97 server = MwServer(applications=[RoutedApp(mount="")], middleware={"cors": True})98 sent = await http_request(server, "/", headers=[(b"origin", b"https://example.test")])99 assert response_status(sent) == 200100 assert response_headers(sent)[b"access-control-allow-origin"] == b"*"101 102 async def test_credentialed_wildcard_echoes_origin_with_vary(103 self, http_request, response_headers104 ) -> None:105 server = MwServer(applications=[RoutedApp(mount="")], middleware={"cors": {"allow_credentials": True}})106 sent = await http_request(server, "/", headers=[(b"origin", b"https://example.test")])107 headers = response_headers(sent)108 assert headers[b"access-control-allow-origin"] == b"https://example.test"109 assert headers[b"vary"] == b"Origin"110 assert headers[b"access-control-allow-credentials"] == b"true"111 112 async def test_restricted_origins_reject_a_foreign_origin(113 self, http_request, response_status, response_headers114 ) -> None:115 server = MwServer(116 applications=[RoutedApp(mount="")], middleware={"cors": {"allow_origins": ["https://allowed.test"]}}117 )118 sent = await http_request(server, "/", headers=[(b"origin", b"https://foreign.test")])119 assert response_status(sent) == 200120 assert b"access-control-allow-origin" not in response_headers(sent)121 122 async def test_preflight_disallowed_origin_returns_400(123 self, http_request, response_status, response_headers124 ) -> None:125 server = MwServer(126 applications=[RoutedApp(mount="")], middleware={"cors": {"allow_origins": ["https://allowed.test"]}}127 )128 sent = await http_request(129 server, "/", method="OPTIONS", headers=[(b"origin", b"https://foreign.test")]130 )131 assert response_status(sent) == 400132 assert b"access-control-allow-origin" not in response_headers(sent)133 134 135 class RaisingApp(BaseApplication):136 """Test app raising the control-flow exceptions the errors middleware maps."""137 138 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:139 path = scope["path"]140 if path == "/missing":141 raise HTTPNotFound("nothing here")142 if path == "/old":143 raise Redirect("/new")144 if path == "/challenge":145 raise HTTPUnauthorized("no", headers=[(b"www-authenticate", b"Bearer")])146 raise RuntimeError("boom")147 148 149 class StartThenRaiseApp(BaseApplication):150 """Test app that starts the response, then raises — nothing more can be sent."""151 152 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:153 await send({"type": "http.response.start", "status": 200, "headers": []})154 raise RuntimeError("boom after start")155 156 157 class TestErrorMiddlewareOnResponse:158 async def test_http_exception_wire_shape(159 self, http_request, response_status, response_headers, response_body160 ) -> None:161 server = MwServer(applications=[RaisingApp(mount="")])162 sent = await http_request(server, "/missing")163 assert response_status(sent) == 404164 assert response_headers(sent)[b"content-type"] == b"text/plain; charset=utf-8"165 assert response_body(sent) == b"nothing here"166 167 async def test_plain_exception_maps_to_500(168 self, http_request, response_status, response_body169 ) -> None:170 server = MwServer(applications=[RaisingApp(mount="")])171 sent = await http_request(server, "/boom")172 assert response_status(sent) == 500173 assert response_body(sent) == b"Internal Server Error"174 175 async def test_redirect_sets_location_header(176 self, http_request, response_status, response_headers177 ) -> None:178 server = MwServer(applications=[RaisingApp(mount="")])179 sent = await http_request(server, "/old")180 assert response_status(sent) == 302181 assert response_headers(sent)[b"location"] == b"/new"182 183 async def test_exception_headers_are_forwarded(184 self, http_request, response_status, response_headers185 ) -> None:186 server = MwServer(applications=[RaisingApp(mount="")])187 sent = await http_request(server, "/challenge")188 assert response_status(sent) == 401189 assert response_headers(sent)[b"www-authenticate"] == b"Bearer"190 191 async def test_error_after_start_is_reraised_not_double_sent(self) -> None:192 server = MwServer(applications=[StartThenRaiseApp(mount="")])193 scope: Scope = {"type": "http", "method": "GET", "path": "/", "headers": []}194 sent: list[Message] = []195 196 async def receive() -> Message:197 return {"type": "http.request"}198 199 async def send(message: Message) -> None:200 sent.append(message)201 202 with pytest.raises(RuntimeError, match="after start"):203 await server(scope, receive, send)204 205 starts = [m for m in sent if m["type"] == "http.response.start"]206 assert len(starts) == 1207 assert starts[0]["status"] == 200208 209 210 class TestErrorContentNegotiation:211 """Macro 5a Phase 5: the error body follows the caller's ``Accept``."""212 213 async def test_json_accept_gets_error_document(214 self, http_request, response_status, response_headers, response_body215 ) -> None:216 server = MwServer(applications=[RaisingApp(mount="")])217 sent = await http_request(server, "/missing", headers=[(b"accept", b"application/json")])218 assert response_status(sent) == 404219 assert response_headers(sent)[b"content-type"] == b"application/json"220 assert json.loads(response_body(sent)) == {"error": "nothing here"}221 222 async def test_wildcard_accept_gets_error_document(223 self, http_request, response_headers, response_body224 ) -> None:225 server = MwServer(applications=[RaisingApp(mount="")])226 sent = await http_request(server, "/missing", headers=[(b"accept", b"*/*")])227 assert response_headers(sent)[b"content-type"] == b"application/json"228 assert json.loads(response_body(sent)) == {"error": "nothing here"}229 230 async def test_html_accept_keeps_text_plain(231 self, http_request, response_headers, response_body232 ) -> None:233 server = MwServer(applications=[RaisingApp(mount="")])234 sent = await http_request(server, "/missing", headers=[(b"accept", b"text/html")])235 assert response_headers(sent)[b"content-type"] == b"text/plain; charset=utf-8"236 assert response_body(sent) == b"nothing here"237 238 async def test_no_accept_defaults_text_plain(239 self, http_request, response_headers, response_body240 ) -> None:241 server = MwServer(applications=[RaisingApp(mount="")])242 sent = await http_request(server, "/missing")243 assert response_headers(sent)[b"content-type"] == b"text/plain; charset=utf-8"244 assert response_body(sent) == b"nothing here"245 246 async def test_generic_500_json_hides_internal_message(247 self, http_request, response_status, response_body248 ) -> None:249 server = MwServer(applications=[RaisingApp(mount="")])250 sent = await http_request(server, "/boom", headers=[(b"accept", b"application/json")])251 assert response_status(sent) == 500252 assert json.loads(response_body(sent)) == {"error": "Internal Server Error"}253 254 255 class TestChallengeNegotiation:256 """Macro 5a Phase 5: a 401 is negotiated when the server has a login surface."""257 258 def test_login_enabled_reflects_registered_method(self) -> None:259 server = AsgiServer(applications=[BaseApplication(mount="")])260 assert server.login_enabled is True261 262 async def test_browser_navigation_redirects_to_login_page(263 self, http_request, response_status, response_headers264 ) -> None:265 server = AsgiServer(applications=[RaisingApp(mount="")])266 sent = await http_request(server, "/challenge", headers=[(b"accept", b"text/html")])267 assert response_status(sent) == 302268 assert response_headers(sent)[b"location"] == b"/_server/login_page?next=%2Fchallenge"269 270 async def test_api_caller_gets_login_url_and_challenge_header(271 self, http_request, response_status, response_headers, response_body272 ) -> None:273 server = AsgiServer(applications=[RaisingApp(mount="")])274 sent = await http_request(server, "/challenge", headers=[(b"accept", b"application/json")])275 assert response_status(sent) == 401276 assert response_headers(sent)[b"www-authenticate"] == b"Bearer"277 assert json.loads(response_body(sent)) == {"login_url": "/_server/login_page"}278 279 async def test_login_disabled_leaves_401_unchanged(280 self, http_request, response_status, response_headers281 ) -> None:282 server = MwServer(applications=[RaisingApp(mount="")])283 sent = await http_request(server, "/challenge", headers=[(b"accept", b"text/html")])284 assert response_status(sent) == 401285 assert response_headers(sent)[b"www-authenticate"] == b"Bearer"286 287 async def test_browser_redirect_preserves_path_and_query_through_safe_next(self) -> None:288 server = AsgiServer(applications=[RaisingApp(mount="")])289 scope: Scope = {290 "type": "http",291 "method": "GET",292 "path": "/challenge",293 "query_string": b"a=1&b=2",294 "headers": [(b"accept", b"text/html")],295 }296 sent: list[Message] = []297 298 async def receive() -> Message:299 return {"type": "http.request"}300 301 async def send(message: Message) -> None:302 sent.append(message)303 304 await server(scope, receive, send)305 start = next(m for m in sent if m["type"] == "http.response.start")306 assert start["status"] == 302307 location = dict(start["headers"])[b"location"].decode()308 assert location == "/_server/login_page?next=" + quote("/challenge?a=1&b=2", safe="")309 310 311 class TestLoggingMiddleware:312 async def test_records_one_entry_per_request(self, http_request, response_status) -> None:313 records: list[str] = []314 315 class RecordingHandler(logging.Handler):316 def emit(self, record: logging.LogRecord) -> None:317 records.append(record.getMessage())318 319 server = MwServer(applications=[RoutedApp(mount="")], middleware={"logging": True})320 access_logger = logging.getLogger("genro_asgi.middleware.logging.LoggingMiddleware")321 handler = RecordingHandler()322 access_logger.addHandler(handler)323 access_logger.setLevel(logging.INFO)324 try:325 sent = await http_request(server, "/")326 finally:327 access_logger.removeHandler(handler)328 329 assert response_status(sent) == 200330 assert len(records) == 2331 assert records[0].startswith("<- GET /")332 assert records[1].startswith("-> GET / 200")333 334 335 class CountingSessionStore(MemorySessionStore):336 """A memory store that counts ``save`` calls (write-back assertions)."""337 338 def __init__(self) -> None:339 super().__init__()340 self.saves = 0341 342 def save(self, session) -> None:343 self.saves += 1344 super().save(session)345 346 347 class SessionMutatingApp(BaseApplication):348 """Mutates the scope session on ``/write``, reads it on any other path."""349 350 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:351 session = scope.get("session")352 if session is not None and scope["path"] == "/write":353 session.data["hit"] = "yes"354 session.mark_dirty()355 await send({"type": "http.response.start", "status": 200, "headers": []})356 await send({"type": "http.response.body", "body": b"ok"})357 358 359 class TestSessionWriteBack:360 def _server(self) -> tuple[AsgiServer, CountingSessionStore]:361 store = CountingSessionStore()362 return AsgiServer(applications=[SessionMutatingApp(mount="")], session_store=store), store363 364 async def test_read_only_request_does_not_save(self, http_request, response_status) -> None:365 server, store = self._server()366 sent = await http_request(server, "/read")367 assert response_status(sent) == 200368 assert store.saves == 0 # read-only stays zero-I/O369 370 async def test_mutating_request_saves_once(self, http_request, response_status) -> None:371 server, store = self._server()372 sent = await http_request(server, "/write")373 assert response_status(sent) == 200374 assert store.saves == 1 # dirty → one write-back375 376 async def test_write_back_clears_the_dirty_flag(self, http_request) -> None:377 server, store = self._server()378 await http_request(server, "/write")379 # the single live session in the store is clean again after the save380 session = next(iter(store.dump()))381 assert store.get(session).dirty is False382 383 384 class TestDisabledByDefault:385 async def test_standard_middlewares_absent_without_switches(386 self, http_request, response_status, response_headers387 ) -> None:388 server = MwServer(applications=[RoutedApp(mount="")])389 390 wellknown_sent = await http_request(server, "/.well-known/probe")391 assert response_status(wellknown_sent) == 200392 393 cors_sent = await http_request(server, "/", headers=[(b"origin", b"https://example.test")])394 assert b"access-control-allow-origin" not in response_headers(cors_sent)395 396 records: list[str] = []397 398 class RecordingHandler(logging.Handler):399 def emit(self, record: logging.LogRecord) -> None:400 records.append(record.getMessage())401 402 access_logger = logging.getLogger("genro_asgi.middleware.logging.LoggingMiddleware")403 handler = RecordingHandler()404 access_logger.addHandler(handler)405 try:406 await http_request(server, "/")407 finally:408 access_logger.removeHandler(handler)409 assert records == []