Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 32 additions & 0 deletions oxyroute/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,9 @@ def __init__(
self.state: SimpleNamespace = SimpleNamespace()
self._docs_ui: str | None = normalize_docs_ui(docs_ui)
self._docs_mounted: bool = False
# Tracks request-middleware count so `set_middleware` can warn before it silently
# clobbers anything registered via `add_middleware` (issue #209).
self._request_mw_count: int = 0
if (
openapi_description is not None
or openapi_contact is not None
Expand Down Expand Up @@ -331,8 +334,37 @@ def set_middleware(self, handler: Callable[..., Any] | None) -> None:
(e.g. :class:`oxyroute.Response` or a ``dict`` with ``status`` / ``body`` / ``headers``);
the response is sent and routing / body read is skipped. Runs **before** the request
body is read (e.g. for CORS preflight).

**Replaces** the entire request-middleware stack, including anything registered via
:meth:`add_middleware` (or :func:`oxyroute.cors.apply_cors` / :func:`oxyroute.csrf.apply_csrf`,
which use :meth:`add_middleware` internally). Prefer :meth:`add_middleware` to compose
multiple middlewares instead of replacing what is already registered (issue #209).
"""
if handler is not None and self._request_mw_count > 0:
import warnings

warnings.warn(
f"set_middleware() is replacing {self._request_mw_count} already-registered "
"request middleware(s) (e.g. from add_middleware/apply_cors/apply_csrf). "
"Use add_middleware() instead to compose middlewares without clobbering them.",
stacklevel=2,
)
self._app.set_middleware(handler)
self._request_mw_count = 1 if handler is not None else 0

def add_middleware(self, handler: Callable[..., Any], phase: str = "request") -> None:
"""
Append ``handler`` to the middleware stack instead of replacing it.

``phase`` is one of ``"request"`` (runs before routing, same contract as
:meth:`set_middleware`), ``"response"`` (runs on the outgoing response, before
CORS / security-header merging), or ``"both"``. Request middlewares run in
registration order; the first one to return a non-``None`` value short-circuits
the rest and is sent as the response.
"""
if phase in ("request", "both"):
self._request_mw_count += 1
self._app.add_middleware(handler, phase)

def set_cors(self, config: Any | None) -> None:
"""
Expand Down
11 changes: 7 additions & 4 deletions oxyroute/cors.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,12 +136,15 @@ def apply_cors(
) -> None:
"""
Register CORS: stores ``config`` for native response merging and installs preflight
handling via :meth:`oxyroute.app.App.set_middleware`. If you already use middleware for
handling via :meth:`oxyroute.app.App.add_middleware`. If you already use middleware for
other work, pass it as ``chain`` so it runs when the request is not a CORS preflight
(your handler runs after the CORS layer returns ``None`` for continuation).

**Order:** this replaces ``set_middleware`` with an internal function. To combine with
another pre-route callback, use ``apply_cors(..., chain=your_middleware)``.
**Order:** this *appends* a request middleware rather than replacing the stack, so it
composes with anything already registered via :meth:`oxyroute.app.App.add_middleware`
(including another :func:`apply_cors` / :func:`oxyroute.csrf.apply_csrf` call) instead of
silently deleting it (issue #209). To combine with another pre-route callback that must run
strictly after the CORS preflight check, use ``apply_cors(..., chain=your_middleware)``.
"""
app.set_cors(config)

Expand All @@ -153,4 +156,4 @@ def _cors_middleware(scope: Any, protocol: Any) -> Response | None:
return chain(scope, protocol)
return None

app.set_middleware(_cors_middleware)
app.add_middleware(_cors_middleware, phase="request")
11 changes: 6 additions & 5 deletions oxyroute/csrf.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,10 +128,11 @@ def apply_csrf(
chain: _Middleware | None = None,
) -> None:
"""
Installs **one** :meth:`oxyroute.app.App.set_middleware` that runs :meth:`CSRFConfig.guard`
first, then ``chain`` (if any), then continues routing. Replaces any previous
pre-route callback — combine manually or use :func:`csrf_layer` inside
:func:`oxyroute.cors.apply_cors`.
Appends (via :meth:`oxyroute.app.App.add_middleware`) a request middleware that runs
:meth:`CSRFConfig.guard` first, then ``chain`` (if any), then continues routing. Composes
with anything already registered via :meth:`oxyroute.app.App.add_middleware` instead of
replacing it (issue #209) — combine manually, or use :func:`csrf_layer` inside
:func:`oxyroute.cors.apply_cors` for strict ordering relative to the CORS preflight check.
"""

def _mw(scope: Any, protocol: Any) -> Response | None:
Expand All @@ -142,4 +143,4 @@ def _mw(scope: Any, protocol: Any) -> Response | None:
return chain(scope, protocol)
return None

app.set_middleware(_mw)
app.add_middleware(_mw, phase="request")
3 changes: 2 additions & 1 deletion tests/test_csrf_units.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,8 @@ class _AppStub:
def __init__(self) -> None:
self.middleware = None

def set_middleware(self, mw): # type: ignore[no-untyped-def]
def add_middleware(self, mw, phase="request"): # type: ignore[no-untyped-def]
assert phase == "request"
self.middleware = mw

app = _AppStub()
Expand Down
61 changes: 61 additions & 0 deletions tests/test_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,12 @@
from __future__ import annotations

import asyncio
import warnings

import httpx
from oxyroute import App, Response
from oxyroute.cors import CORSConfig, apply_cors
from oxyroute.csrf import CSRFConfig, apply_csrf
from oxyroute.testing import asgi_test_app


Expand Down Expand Up @@ -43,3 +46,61 @@ async def _run() -> None:
assert n == 0

asyncio.run(_run())


def test_add_middleware_then_apply_cors_does_not_clobber_auth() -> None:
"""issue #209: add_middleware(auth) followed by apply_cors(...) must not silently
delete the auth middleware."""
calls: list[str] = []

def auth_mw(scope, _protocol): # type: ignore[no-untyped-def]
calls.append("auth")
return None

app = App()
app.add_middleware(auth_mw)
apply_cors(app, CORSConfig())

@app.get("/x")
def _x() -> str:
return "ok"

async def _run() -> None:
transport = httpx.ASGITransport(app=asgi_test_app(app))
async with httpx.AsyncClient(transport=transport, base_url="http://test") as c:
r = await c.get("/x")
assert r.status_code == 200, r.text

asyncio.run(_run())
assert calls == ["auth"]


def test_apply_csrf_then_apply_cors_compose_instead_of_clobbering() -> None:
"""issue #209: calling apply_csrf then apply_cors must not delete the CSRF guard."""
app = App()
apply_csrf(app, CSRFConfig())
apply_cors(app, CORSConfig())

@app.post("/x")
def _x() -> str:
return "ok"

async def _run() -> None:
transport = httpx.ASGITransport(app=asgi_test_app(app))
async with httpx.AsyncClient(transport=transport, base_url="http://test") as c:
# unsafe method, no CSRF token -> the CSRF guard (still registered) must block it
r = await c.post("/x")
assert r.status_code == 403, r.text

asyncio.run(_run())


def test_set_middleware_warns_when_clobbering_add_middleware() -> None:
app = App()
app.add_middleware(lambda scope, protocol: None)

with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
app.set_middleware(lambda scope, protocol: None)

assert any("set_middleware" in str(x.message) for x in w)
Loading