diff --git a/middleware/origin_middleware.py b/middleware/origin_middleware.py index 11f06ff75..b4a482d9c 100644 --- a/middleware/origin_middleware.py +++ b/middleware/origin_middleware.py @@ -35,9 +35,14 @@ def is_loopback(host): def create_origin_only_middleware(): @web.middleware async def origin_only_middleware(request: web.Request, handler): - if 'Sec-Fetch-Site' in request.headers: - sec_fetch_site = request.headers['Sec-Fetch-Site'] - if sec_fetch_site == 'cross-site': + if request.headers.get('Sec-Fetch-Site') == 'cross-site': + is_top_level_navigation = ( + request.method == 'GET' + and request.path == '/' + and request.headers.get('Sec-Fetch-Mode') == 'navigate' + and request.headers.get('Sec-Fetch-Dest') == 'document' + ) + if not is_top_level_navigation: return web.Response(status=403) #this code is used to prevent the case where a random website can queue comfy workflows by making a POST to 127.0.0.1 which browsers don't prevent for some dumb reason. #in that case the Host and Origin hostnames won't match diff --git a/tests-unit/server_test/test_origin_middleware.py b/tests-unit/server_test/test_origin_middleware.py new file mode 100644 index 000000000..09ead7f66 --- /dev/null +++ b/tests-unit/server_test/test_origin_middleware.py @@ -0,0 +1,117 @@ +import pytest +from aiohttp import web +from aiohttp.test_utils import make_mocked_request + +from middleware.origin_middleware import create_origin_only_middleware + + +pytestmark = pytest.mark.asyncio + + +async def ok_handler(request): + return web.Response(status=200) + + +async def test_allows_cross_site_top_level_get_navigation(): + request = make_mocked_request( + 'GET', + '/', + headers={ + 'Sec-Fetch-Site': 'cross-site', + 'Sec-Fetch-Mode': 'navigate', + 'Sec-Fetch-Dest': 'document', + 'Sec-Fetch-User': '?1', + }, + ) + + response = await create_origin_only_middleware()(request, ok_handler) + + assert response.status == 200 + + +async def test_blocks_cross_site_top_level_navigation_to_other_routes(): + request = make_mocked_request( + 'GET', + '/system_stats', + headers={ + 'Sec-Fetch-Site': 'cross-site', + 'Sec-Fetch-Mode': 'navigate', + 'Sec-Fetch-Dest': 'document', + }, + ) + + response = await create_origin_only_middleware()(request, ok_handler) + + assert response.status == 403 + + +@pytest.mark.parametrize( + ('method', 'headers'), + [ + ( + 'POST', + { + 'Sec-Fetch-Site': 'cross-site', + 'Sec-Fetch-Mode': 'navigate', + 'Sec-Fetch-Dest': 'document', + }, + ), + ( + 'GET', + { + 'Sec-Fetch-Site': 'cross-site', + 'Sec-Fetch-Mode': 'cors', + 'Sec-Fetch-Dest': 'empty', + }, + ), + ( + 'GET', + { + 'Sec-Fetch-Site': 'cross-site', + 'Sec-Fetch-Mode': 'websocket', + 'Sec-Fetch-Dest': 'empty', + }, + ), + ( + 'GET', + { + 'Sec-Fetch-Site': 'cross-site', + 'Sec-Fetch-Mode': 'navigate', + 'Sec-Fetch-Dest': 'iframe', + }, + ), + ('GET', {'Sec-Fetch-Site': 'cross-site'}), + ], + ids=['form-post', 'fetch', 'websocket', 'iframe', 'missing-context'], +) +async def test_blocks_other_cross_site_requests(method, headers): + request = make_mocked_request(method, '/', headers=headers) + + response = await create_origin_only_middleware()(request, ok_handler) + + assert response.status == 403 + + +async def test_preserves_loopback_origin_check_without_fetch_metadata(): + request = make_mocked_request( + 'POST', + '/prompt', + headers={ + 'Host': '127.0.0.1:8188', + 'Origin': 'https://evil.example', + }, + ) + + response = await create_origin_only_middleware()(request, ok_handler) + + assert response.status == 403 + + +@pytest.mark.parametrize('site', ['same-origin', 'same-site', 'none', None]) +async def test_preserves_allowed_request_sources(site): + headers = {'Sec-Fetch-Site': site} if site is not None else {} + request = make_mocked_request('GET', '/', headers=headers) + + response = await create_origin_only_middleware()(request, ok_handler) + + assert response.status == 200