diff --git a/app/api/v2/security.py b/app/api/v2/security.py index 0f70fb3a8..a78994ca5 100644 --- a/app/api/v2/security.py +++ b/app/api/v2/security.py @@ -69,6 +69,25 @@ async def authentication_required_middleware(request, handler): return authentication_required_middleware +def docs_guard_middleware_factory(auth_svc): + """Middleware that requires authentication for /api/docs and /static/swagger paths. + + Note: this middleware must be registered on the **main** application (not a + subapp) because ``init_swagger_documentation`` registers the ``/api/docs`` + and ``/static/swagger`` routes on the main application via + ``aiohttp_apispec.setup_aiohttp_apispec``. Middleware on the main app + intercepts these routes before they reach a handler. + """ + @web.middleware + async def docs_guard_middleware(request, handler): + if request.path.startswith('/api/docs') or request.path.startswith('/static/swagger'): + if not auth_svc.request_has_valid_api_key(request): + if not await auth_svc.request_has_valid_user_session(request): + raise web.HTTPUnauthorized(text='Authentication required for API documentation') + return await handler(request) + return docs_guard_middleware + + @web.middleware async def pass_option_middleware(request, handler): """Allow all 'OPTIONS' request to the server to return 200 diff --git a/server.py b/server.py index 2c2d0702d..34edfd56e 100644 --- a/server.py +++ b/server.py @@ -18,7 +18,7 @@ from app.ascii_banner import ASCII_BANNER, no_color, print_rich_banner from app.api.rest_api import RestApi from app.api.v2.responses import apispec_request_validation_middleware -from app.api.v2.security import pass_option_middleware +from app.api.v2.security import pass_option_middleware, docs_guard_middleware_factory from app.objects.c_agent import Agent from app.objects.c_obfuscator import Obfuscator from app.objects.secondclass.c_executor import Executor @@ -263,7 +263,11 @@ def list_str(values): app_svc = AppService( application=web.Application( - client_max_size=5120**2, middlewares=[pass_option_middleware] + client_max_size=5120**2, + middlewares=[ + docs_guard_middleware_factory(auth_svc), + pass_option_middleware, + ] ) ) app_svc.register_subapp("/api/v2", app.api.v2.make_app(app_svc.get_services())) diff --git a/tests/security/test_swagger_auth.py b/tests/security/test_swagger_auth.py new file mode 100644 index 000000000..d3d43133a --- /dev/null +++ b/tests/security/test_swagger_auth.py @@ -0,0 +1,70 @@ +import pytest +from unittest.mock import AsyncMock, MagicMock +from aiohttp import web + +from app.api.v2.security import docs_guard_middleware_factory + + +class TestSwaggerDocsAuth: + """Tests for the docs_guard_middleware that protects /api/docs and /static/swagger.""" + + def _make_middleware(self, *, api_key_valid=False, session_valid=False): + auth_svc = MagicMock() + auth_svc.request_has_valid_api_key.return_value = api_key_valid + auth_svc.request_has_valid_user_session = AsyncMock(return_value=session_valid) + return docs_guard_middleware_factory(auth_svc) + + @pytest.mark.asyncio + async def test_docs_guard_blocks_unauthenticated(self): + """Test that /api/docs paths are blocked without auth.""" + middleware = self._make_middleware() + request = MagicMock() + request.path = '/api/docs' + handler = AsyncMock() + + with pytest.raises(web.HTTPUnauthorized): + await middleware(request, handler) + + @pytest.mark.asyncio + async def test_docs_guard_allows_api_key(self): + """Test that /api/docs paths are allowed with valid API key.""" + middleware = self._make_middleware(api_key_valid=True) + request = MagicMock() + request.path = '/api/docs' + handler = AsyncMock(return_value='ok') + + result = await middleware(request, handler) + assert result == 'ok' + + @pytest.mark.asyncio + async def test_docs_guard_ignores_other_paths(self): + """Test that non-docs paths pass through without auth check.""" + middleware = self._make_middleware() + request = MagicMock() + request.path = '/api/v2/agents' + handler = AsyncMock(return_value='ok') + + result = await middleware(request, handler) + assert result == 'ok' + + @pytest.mark.asyncio + async def test_docs_guard_allows_valid_user_session(self): + """API key invalid but valid user session should allow /api/docs access.""" + middleware = self._make_middleware(session_valid=True) + request = MagicMock() + request.path = '/api/docs' + handler = AsyncMock(return_value='ok') + + result = await middleware(request, handler) + assert result == 'ok' + + @pytest.mark.asyncio + async def test_docs_guard_allows_static_swagger_with_session(self): + """Valid user session should also permit access to /static/swagger paths.""" + middleware = self._make_middleware(session_valid=True) + request = MagicMock() + request.path = '/static/swagger/ui.js' + handler = AsyncMock(return_value='script') + + result = await middleware(request, handler) + assert result == 'script'