Allow cross-site UI navigation

Amp-Thread-ID: https://ampcode.com/threads/T-019ff807-a342-7628-863e-847d1dad3d5d
Co-authored-by: Amp <amp@ampcode.com>
This commit is contained in:
Jedrzej Kosinski
2026-08-12 15:59:23 -07:00
parent d00f56388d
commit f8d63fb01f
2 changed files with 125 additions and 3 deletions

View File

@@ -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

View File

@@ -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