From 9856e2ed97f16620b7223b8b7a59bfc8749978f2 Mon Sep 17 00:00:00 2001 From: Manish Kumar Date: Wed, 26 Aug 2026 12:47:18 -0500 Subject: [PATCH] test: avoid blocking middleware test clients --- tests/test_policy_middleware.py | 40 ++++++++++++++++++++++++--------- tests/test_s2s_middleware.py | 34 ++++++++++++++++------------ 2 files changed, 50 insertions(+), 24 deletions(-) diff --git a/tests/test_policy_middleware.py b/tests/test_policy_middleware.py index 2351cdc..a574433 100644 --- a/tests/test_policy_middleware.py +++ b/tests/test_policy_middleware.py @@ -1,10 +1,11 @@ -"""TestClient-based integration tests for middleware/policy.py PolicyMiddleware.""" +"""Async integration tests for middleware/policy.py PolicyMiddleware.""" +import json import sys import types import importlib import pytest from unittest.mock import AsyncMock, MagicMock, patch -from starlette.testclient import TestClient +from starlette.requests import Request from starlette.applications import Starlette from starlette.responses import JSONResponse from starlette.routing import Route @@ -42,23 +43,42 @@ async def endpoint(request): return app -def test_policy_middleware_allows(): +@pytest.mark.asyncio +async def test_policy_middleware_allows(): mock_policy = MagicMock() mock_policy.evaluate = AsyncMock(return_value={"allow": True, "reason": "ok"}) app = _make_app(mock_policy) - client = TestClient(app, raise_server_exceptions=False) + middleware = PolicyMiddleware(app, policy=mock_policy) + request = Request({ + "type": "http", "method": "GET", "path": "/test", + "headers": [], "query_string": b"", + }) + + async def call_next(_request): + return JSONResponse({"ok": True}) + with patch("middleware.policy.get_user", return_value={"user_id": "u1", "roles": ["researcher"]}): - resp = client.get("/test") + resp = await middleware.dispatch(request, call_next) assert resp.status_code == 200 -def test_policy_middleware_denies(): +@pytest.mark.asyncio +async def test_policy_middleware_denies(): mock_policy = MagicMock() mock_policy.evaluate = AsyncMock(return_value={"allow": False, "reason": "forbidden"}) app = _make_app(mock_policy) - client = TestClient(app, raise_server_exceptions=False) + middleware = PolicyMiddleware(app, policy=mock_policy) + request = Request({ + "type": "http", "method": "GET", "path": "/test", + "headers": [], "query_string": b"", + }) + + async def call_next(_request): + return JSONResponse({"ok": True}) + with patch("middleware.policy.get_user", return_value={"user_id": "u2", "roles": []}): - resp = client.get("/test") + resp = await middleware.dispatch(request, call_next) + body = json.loads(resp.body) assert resp.status_code == 403 - assert resp.json()["error"] == "forbidden" - assert resp.json()["reason"] == "forbidden" + assert body["error"] == "forbidden" + assert body["reason"] == "forbidden" diff --git a/tests/test_s2s_middleware.py b/tests/test_s2s_middleware.py index 139b547..db63b48 100644 --- a/tests/test_s2s_middleware.py +++ b/tests/test_s2s_middleware.py @@ -1,10 +1,10 @@ -"""TestClient-based integration tests for middleware/s2s.py ServiceAuthMiddleware.""" +"""Async integration tests for middleware/s2s.py ServiceAuthMiddleware.""" +import httpx import sys import types import importlib import jwt import pytest -from starlette.testclient import TestClient from starlette.applications import Starlette from starlette.responses import JSONResponse from starlette.routing import Route @@ -50,31 +50,37 @@ def _token(service="tes", aud=None): return jwt.encode({"service": service, "aud": aud or [SERVICE]}, SECRET, algorithm="HS256") -def test_s2s_missing_token(): - client = TestClient(_make_app(), raise_server_exceptions=False) - resp = client.get("/test") +async def _get(path="/test", headers=None): + transport = httpx.ASGITransport(app=_make_app()) + async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as client: + return await client.get(path, headers=headers) + + +@pytest.mark.asyncio +async def test_s2s_missing_token(): + resp = await _get() assert resp.status_code == 401 assert "missing service token" in resp.json()["error"] -def test_s2s_invalid_token(): - client = TestClient(_make_app(), raise_server_exceptions=False) - resp = client.get("/test", headers={"X-Service-Token": "bad.token"}) +@pytest.mark.asyncio +async def test_s2s_invalid_token(): + resp = await _get(headers={"X-Service-Token": "bad.token"}) assert resp.status_code == 401 assert "invalid service token" in resp.json()["error"] -def test_s2s_wrong_audience(): +@pytest.mark.asyncio +async def test_s2s_wrong_audience(): token = _token(aud=["other-service"]) - client = TestClient(_make_app(), raise_server_exceptions=False) - resp = client.get("/test", headers={"X-Service-Token": token}) + resp = await _get(headers={"X-Service-Token": token}) assert resp.status_code == 403 assert "service not allowed" in resp.json()["error"] -def test_s2s_valid_token_passes(): +@pytest.mark.asyncio +async def test_s2s_valid_token_passes(): token = _token(service="tes", aud=[SERVICE]) - client = TestClient(_make_app(), raise_server_exceptions=False) - resp = client.get("/test", headers={"X-Service-Token": token}) + resp = await _get(headers={"X-Service-Token": token}) assert resp.status_code == 200 assert resp.json()["service"] == "tes"