mirror of
https://github.com/Comfy-Org/ComfyUI.git
synced 2026-08-18 15:30:46 +08:00
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:
@@ -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
|
||||
|
||||
117
tests-unit/server_test/test_origin_middleware.py
Normal file
117
tests-unit/server_test/test_origin_middleware.py
Normal 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
|
||||
Reference in New Issue
Block a user