From d00f56388da9a8326fdd35cfd70c7e4b7eb36ba8 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 12 Aug 2026 15:57:53 -0700 Subject: [PATCH] Move origin middleware into module Amp-Thread-ID: https://ampcode.com/threads/T-019ff807-a342-7628-863e-847d1dad3d5d Co-authored-by: Amp --- middleware/origin_middleware.py | 73 +++++++++++++++++++++++++++++++++ server.py | 70 +------------------------------ 2 files changed, 74 insertions(+), 69 deletions(-) create mode 100644 middleware/origin_middleware.py diff --git a/middleware/origin_middleware.py b/middleware/origin_middleware.py new file mode 100644 index 000000000..11f06ff75 --- /dev/null +++ b/middleware/origin_middleware.py @@ -0,0 +1,73 @@ +import ipaddress +import logging +import socket +import urllib + +from aiohttp import web + + +def is_loopback(host): + if host is None: + return False + try: + if ipaddress.ip_address(host).is_loopback: + return True + else: + return False + except: + pass + + loopback = False + for family in (socket.AF_INET, socket.AF_INET6): + try: + r = socket.getaddrinfo(host, None, family, socket.SOCK_STREAM) + for family, _, _, _, sockaddr in r: + if not ipaddress.ip_address(sockaddr[0]).is_loopback: + return loopback + else: + loopback = True + except socket.gaierror: + pass + + return loopback + + +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': + 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 + #I know the proper fix would be to add a cookie but this should take care of the problem in the meantime + if 'Host' in request.headers and 'Origin' in request.headers: + host = request.headers['Host'] + origin = request.headers['Origin'] + host_domain = host.lower() + parsed = urllib.parse.urlparse(origin) + origin_domain = parsed.netloc.lower() + host_domain_parsed = urllib.parse.urlsplit('//' + host_domain) + + #limit the check to when the host domain is localhost, this makes it slightly less safe but should still prevent the exploit + loopback = is_loopback(host_domain_parsed.hostname) + + if parsed.port is None: #if origin doesn't have a port strip it from the host to handle weird browsers, same for host + host_domain = host_domain_parsed.hostname + if host_domain_parsed.port is None: + origin_domain = parsed.hostname + + if loopback and host_domain is not None and origin_domain is not None and len(host_domain) > 0 and len(origin_domain) > 0: + if host_domain != origin_domain: + logging.warning("WARNING: request with non matching host and origin {} != {}, returning 403".format(host_domain, origin_domain)) + return web.Response(status=403) + + if request.method == "OPTIONS": + response = web.Response() + else: + response = await handler(request) + + return response + + return origin_only_middleware diff --git a/server.py b/server.py index c9ffcaa0d..8fee71cf1 100644 --- a/server.py +++ b/server.py @@ -23,8 +23,6 @@ import json import glob import struct import ssl -import socket -import ipaddress from PIL import Image, ImageOps from PIL.PngImagePlugin import PngInfo from io import BytesIO @@ -61,6 +59,7 @@ from protocol import BinaryEventTypes # Import cache control middleware from middleware.cache_middleware import cache_control +from middleware.origin_middleware import create_origin_only_middleware if args.enable_manager: import comfyui_manager @@ -130,73 +129,6 @@ def create_cors_middleware(allowed_origin: str): return cors_middleware -def is_loopback(host): - if host is None: - return False - try: - if ipaddress.ip_address(host).is_loopback: - return True - else: - return False - except: - pass - - loopback = False - for family in (socket.AF_INET, socket.AF_INET6): - try: - r = socket.getaddrinfo(host, None, family, socket.SOCK_STREAM) - for family, _, _, _, sockaddr in r: - if not ipaddress.ip_address(sockaddr[0]).is_loopback: - return loopback - else: - loopback = True - except socket.gaierror: - pass - - return loopback - - -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': - 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 - #I know the proper fix would be to add a cookie but this should take care of the problem in the meantime - if 'Host' in request.headers and 'Origin' in request.headers: - host = request.headers['Host'] - origin = request.headers['Origin'] - host_domain = host.lower() - parsed = urllib.parse.urlparse(origin) - origin_domain = parsed.netloc.lower() - host_domain_parsed = urllib.parse.urlsplit('//' + host_domain) - - #limit the check to when the host domain is localhost, this makes it slightly less safe but should still prevent the exploit - loopback = is_loopback(host_domain_parsed.hostname) - - if parsed.port is None: #if origin doesn't have a port strip it from the host to handle weird browsers, same for host - host_domain = host_domain_parsed.hostname - if host_domain_parsed.port is None: - origin_domain = parsed.hostname - - if loopback and host_domain is not None and origin_domain is not None and len(host_domain) > 0 and len(origin_domain) > 0: - if host_domain != origin_domain: - logging.warning("WARNING: request with non matching host and origin {} != {}, returning 403".format(host_domain, origin_domain)) - return web.Response(status=403) - - if request.method == "OPTIONS": - response = web.Response() - else: - response = await handler(request) - - return response - - return origin_only_middleware - - def create_block_external_middleware(): @web.middleware async def block_external_middleware(request: web.Request, handler):