From 4284138a25528534cdf41b450de2e0d56141720c Mon Sep 17 00:00:00 2001 From: Arjan Singh Date: Mon, 25 Aug 2025 14:05:23 -0700 Subject: [PATCH] test: test cache handler and fix found issues --- server.py | 8 +- tests-unit/server_test/test_cache_control.py | 205 +++++++++++++++++++ 2 files changed, 209 insertions(+), 4 deletions(-) create mode 100644 tests-unit/server_test/test_cache_control.py diff --git a/server.py b/server.py index ae49ba9f3..7727caa37 100644 --- a/server.py +++ b/server.py @@ -35,7 +35,7 @@ from comfy_api.internal import _ComfyNodeInternal from app.user_manager import UserManager from app.model_manager import ModelFileManager from app.custom_node_manager import CustomNodeManager -from typing import Optional, Union +from typing import Optional, Union, Callable, Awaitable from api_server.routes.internal.internal_routes import InternalRoutes from protocol import BinaryEventTypes @@ -51,14 +51,14 @@ async def send_socket_catch_exception(function, message): logging.warning("send error: {}".format(err)) @web.middleware -async def cache_control(request: web.Request, handler) -> web.Response: +async def cache_control(request: web.Request, handler: Callable[[web.Request], Awaitable[web.Response]]) -> web.Response: response: web.Response = await handler(request) if request.path.endswith('.js') or request.path.endswith('.css') or request.path.endswith('index.json'): response.headers.setdefault('Cache-Control', 'no-cache') elif request.path.lower().endswith(IMG_EXTENSIONS): - max_age = ONE_DAY if response.status == 200 else ONE_HOUR - response.headers['Cache-Control'] = f"public, max-age={max_age}" + max_age = ONE_DAY if 200 <= response.status < 300 else ONE_HOUR + response.headers.setdefault('Cache-Control', f"public, max-age={max_age}") return response diff --git a/tests-unit/server_test/test_cache_control.py b/tests-unit/server_test/test_cache_control.py new file mode 100644 index 000000000..aab75ab1e --- /dev/null +++ b/tests-unit/server_test/test_cache_control.py @@ -0,0 +1,205 @@ +"""Tests for server cache control middleware""" +import pytest +from aiohttp import web +from aiohttp.test_utils import make_mocked_request +from unittest.mock import patch + +pytestmark = pytest.mark.asyncio # Apply asyncio mark to all tests + +# Mock the problematic imports before importing server +with patch('app.frontend_management.FrontendManager'): + with patch('utils.install_util.get_missing_requirements_message'): + with patch('utils.install_util.requirements_path'): + from server import cache_control, ONE_HOUR, ONE_DAY, IMG_EXTENSIONS + + +class TestCacheControl: + """Test cache control middleware functionality""" + + @pytest.fixture + def mock_handler(self): + """Create a mock handler that returns a response with 200 status""" + async def handler(request): + return web.Response(status=200) + return handler + + @pytest.fixture + def mock_handler_404(self): + """Create a mock handler that returns a 404 response""" + async def handler(request): + return web.Response(status=404) + return handler + + async def test_image_extensions_200_status(self, mock_handler): + """Test that images with 200 status get 24-hour cache""" + for ext in IMG_EXTENSIONS: + request = make_mocked_request('GET', f'/test{ext}') + response = await cache_control(request, mock_handler) + + assert response.status == 200 + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}' + + async def test_image_extensions_404_status(self, mock_handler_404): + """Test that images with 404 status get 1-hour cache""" + request = make_mocked_request('GET', '/missing.jpg') + response = await cache_control(request, mock_handler_404) + + assert response.status == 404 + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_HOUR}' + + async def test_case_insensitive_image_extension(self, mock_handler): + """Test that image extensions are matched case-insensitively""" + test_paths = ['/image.JPG', '/photo.PNG', '/pic.JpEg'] + + for path in test_paths: + request = make_mocked_request('GET', path) + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}' + + async def test_js_files_no_cache(self, mock_handler): + """Test that .js files get no-cache header""" + request = make_mocked_request('GET', '/script.js') + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == 'no-cache' + + async def test_css_files_no_cache(self, mock_handler): + """Test that .css files get no-cache header""" + request = make_mocked_request('GET', '/styles.css') + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == 'no-cache' + + async def test_index_json_no_cache(self, mock_handler): + """Test that index.json gets no-cache header""" + request = make_mocked_request('GET', '/api/index.json') + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == 'no-cache' + + async def test_js_css_preserves_existing_headers(self): + """Test that .js/.css files preserve existing Cache-Control headers""" + async def handler_with_cache(request): + return web.Response(status=200, headers={'Cache-Control': 'max-age=3600'}) + + request = make_mocked_request('GET', '/script.js') + response = await cache_control(request, handler_with_cache) + + # setdefault should preserve existing header + assert response.headers['Cache-Control'] == 'max-age=3600' + + async def test_image_preserves_existing_headers(self): + """Test that image cache headers preserve existing Cache-Control""" + async def handler_with_cache(request): + return web.Response(status=200, headers={'Cache-Control': 'private, no-cache'}) + + request = make_mocked_request('GET', '/image.jpg') + response = await cache_control(request, handler_with_cache) + + # setdefault should preserve existing header + assert response.headers['Cache-Control'] == 'private, no-cache' + + async def test_non_matching_files_unchanged(self, mock_handler): + """Test that non-matching files don't get cache headers""" + test_paths = ['/index.html', '/data.txt', '/api/endpoint', '/file.pdf'] + + for path in test_paths: + request = make_mocked_request('GET', path) + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' not in response.headers + + async def test_query_strings_ignored(self, mock_handler): + """Test that query strings don't affect image detection""" + request = make_mocked_request('GET', '/image.jpg?v=123&size=large') + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}' + + async def test_multiple_dots_in_path(self, mock_handler): + """Test files with multiple dots still match correctly""" + request = make_mocked_request('GET', '/image.min.jpg') + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}' + + async def test_various_error_statuses(self): + """Test that various error statuses get 1-hour cache for images""" + error_statuses = [403, 404, 500, 502, 503] + + for status in error_statuses: + async def handler_error(request): + return web.Response(status=status) + + request = make_mocked_request('GET', '/error.png') + response = await cache_control(request, handler_error) + + assert response.status == status + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_HOUR}' + + async def test_2xx_success_statuses_get_long_cache(self): + """Test that all 2xx success statuses get 24-hour cache for images""" + success_statuses = [200, 201, 202, 204, 206] + + for status in success_statuses: + async def handler_success(request): + return web.Response(status=status) + + request = make_mocked_request('GET', '/success.jpg') + response = await cache_control(request, handler_success) + + assert response.status == status + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}' + + async def test_3xx_redirect_statuses_get_short_cache(self): + """Test that 3xx redirect statuses get 1-hour cache for images""" + redirect_statuses = [301, 302, 304] + + for status in redirect_statuses: + async def handler_redirect(request): + return web.Response(status=status) + + request = make_mocked_request('GET', '/redirect.png') + response = await cache_control(request, handler_redirect) + + assert response.status == status + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_HOUR}' + + async def test_all_image_extensions(self, mock_handler): + """Test all defined image extensions are handled""" + expected_extensions = ('.jpg', '.jpeg', '.png', '.ppm', '.bmp', '.pgm', '.tif', '.tiff', '.webp') + assert IMG_EXTENSIONS == expected_extensions + + for ext in IMG_EXTENSIONS: + request = make_mocked_request('GET', f'/image{ext}') + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}' + + async def test_nested_paths_with_images(self, mock_handler): + """Test that images in nested paths are handled correctly""" + test_paths = [ + '/static/images/photo.jpg', + '/assets/img/banner.png', + '/uploads/2024/12/image.webp' + ] + + for path in test_paths: + request = make_mocked_request('GET', path) + response = await cache_control(request, mock_handler) + + assert 'Cache-Control' in response.headers + assert response.headers['Cache-Control'] == f'public, max-age={ONE_DAY}'