Skip to content
Open
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
40 changes: 30 additions & 10 deletions tests/test_policy_middleware.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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"
34 changes: 20 additions & 14 deletions tests/test_s2s_middleware.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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"