Skip to content

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