Files

175 lines
7.9 KiB
Python

import ast
import asyncio
import hashlib
import inspect
import json
import os
from pathlib import Path
import shutil
import sys
import tempfile
import unittest
from functools import lru_cache
from types import SimpleNamespace
from urllib.parse import urlsplit
from unittest.mock import patch
from aiohttp import web
from aiohttp.test_utils import TestClient, TestServer
from PIL import Image, UnidentifiedImageError
ROUTES_PATH = Path(__file__).parents[1] / "py" / "routes.py"
def load_handlers(folder_paths, get_metadata):
"""Load the actual handlers without importing ComfyUI's GPU dependencies."""
names = {
"_same_origin_request", "get_reboot_token", "reboot", "_model_sha256",
"load_metadata", "save_notes", "save_preview",
}
tree = ast.parse(ROUTES_PATH.read_text())
functions = []
for node in tree.body:
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name in names:
node.decorator_list = []
functions.append(node)
namespace = {
"os": os, "sys": sys, "hashlib": hashlib, "hmac": __import__("hmac"),
"json": json, "shutil": shutil, "tempfile": tempfile,
"lru_cache": lru_cache, "urlsplit": urlsplit, "web": web,
"Image": Image, "UnidentifiedImageError": UnidentifiedImageError,
"folder_paths": folder_paths, "getMetadata": get_metadata,
"_reboot_token": "test-reboot-token",
"_PREVIEW_FORMATS": {
".png": "PNG", ".jpg": "JPEG", ".jpeg": "JPEG",
".webp": "WEBP", ".gif": "GIF",
},
}
exec(compile(ast.Module(body=functions, type_ignores=[]), str(ROUTES_PATH), "exec"), namespace)
return SimpleNamespace(**namespace)
class SecurityRouteTests(unittest.TestCase):
def setUp(self):
self.workspace = tempfile.TemporaryDirectory()
self.addCleanup(self.workspace.cleanup)
self.root = Path(self.workspace.name)
self.model_dir = self.root / "models"
self.temp_dir = self.root / "temp"
self.model_dir.mkdir()
self.temp_dir.mkdir()
self.model_path = self.model_dir / "sample.safetensors"
self.model_path.write_bytes(b"model data")
paths = SimpleNamespace(
get_filename_list=lambda kind: [self.model_path.name],
get_full_path=lambda kind, name: str(self.model_path),
get_directory_by_type=lambda kind: str(self.temp_dir),
)
self.handlers = load_handlers(
paths,
lambda path: json.dumps({"__metadata__": {"easyuse.notes": "<img onerror=alert(1)>"}}),
)
def request(self, name="loras/sample.safetensors", filename="preview.png", **body):
payload = {"type": "temp", "filename": filename, **body}
return SimpleNamespace(
match_info={"name": name},
json=lambda: asyncio.sleep(0, result=payload),
headers={}, host="localhost:8188",
)
def test_save_rejects_script_and_custom_node_target(self):
(self.temp_dir / "payload.py").write_text("print('sentinel')")
response = asyncio.run(self.handlers.save_preview(self.request(filename="payload.py")))
self.assertEqual(response.status, 400)
response = asyncio.run(self.handlers.save_preview(
self.request(name="custom_nodes/package/__init__.py", filename="payload.py")
))
self.assertEqual(response.status, 400)
def test_save_accepts_real_image_and_rejects_disguised_script(self):
Image.new("RGB", (1, 1)).save(self.temp_dir / "preview.png")
response = asyncio.run(self.handlers.save_preview(self.request()))
self.assertEqual(response.status, 200)
with Image.open(self.model_dir / "sample.png") as saved:
self.assertEqual(saved.format, "PNG")
(self.temp_dir / "preview.png").write_text("print('sentinel')")
response = asyncio.run(self.handlers.save_preview(self.request()))
self.assertEqual(response.status, 400)
with Image.open(self.model_dir / "sample.png") as saved:
self.assertEqual(saved.format, "PNG")
@unittest.skipUnless(hasattr(os, "symlink"), "symlinks are unavailable")
def test_save_does_not_follow_preview_symlink(self):
Image.new("RGB", (1, 1)).save(self.temp_dir / "preview.png")
protected = self.root / "protected.txt"
protected.write_text("untouched")
os.symlink(protected, self.model_dir / "sample.png")
response = asyncio.run(self.handlers.save_preview(self.request()))
self.assertEqual(response.status, 400)
self.assertEqual(protected.read_text(), "untouched")
def test_metadata_ignores_forged_hash_sidecar(self):
(self.model_dir / "sample.sha256").write_text("0" * 64)
response = asyncio.run(self.handlers.load_metadata(self.request()))
self.assertEqual(
json.loads(response.text)["easyuse.sha256"],
hashlib.sha256(self.model_path.read_bytes()).hexdigest(),
)
@unittest.skipUnless(hasattr(os, "symlink"), "symlinks are unavailable")
def test_notes_reject_custom_nodes_and_do_not_follow_symlinks(self):
request = self.request(name="custom_nodes/package/__init__.py")
request.text = lambda: asyncio.sleep(0, result="new notes")
self.assertEqual(asyncio.run(self.handlers.save_notes(request)).status, 400)
protected = self.root / "protected.txt"
protected.write_text("untouched")
os.symlink(protected, self.model_dir / "sample.txt")
request = self.request()
request.text = lambda: asyncio.sleep(0, result="new notes")
self.assertEqual(asyncio.run(self.handlers.save_notes(request)).status, 200)
self.assertEqual(protected.read_text(), "untouched")
self.assertEqual((self.model_dir / "sample.txt").read_text(), "new notes")
def test_reboot_requires_token_and_same_origin(self):
self.assertTrue(inspect.iscoroutinefunction(self.handlers.get_reboot_token))
self.assertTrue(inspect.iscoroutinefunction(self.handlers.reboot))
request = self.request()
request.headers = {"Sec-Fetch-Site": "cross-site"}
self.assertEqual(asyncio.run(self.handlers.get_reboot_token(request)).status, 403)
request.headers = {"Sec-Fetch-Site": "same-origin"}
self.assertEqual(json.loads(asyncio.run(self.handlers.get_reboot_token(request)).text)["token"], "test-reboot-token")
request.headers = {}
with patch.object(self.handlers.os, "execv", return_value="restarted") as restart:
self.assertEqual(asyncio.run(self.handlers.reboot(request)).status, 403)
request.headers = {"X-EasyUse-Reboot-Token": "test-reboot-token", "Origin": "http://other.test"}
self.assertEqual(asyncio.run(self.handlers.reboot(request)).status, 403)
restart.assert_not_called()
request.headers["Origin"] = "http://localhost:8188"
request.headers["Sec-Fetch-Site"] = "same-site"
self.assertEqual(asyncio.run(self.handlers.reboot(request)).status, 403)
request.headers["Sec-Fetch-Site"] = "same-origin"
self.assertEqual(asyncio.run(self.handlers.reboot(request)), "restarted")
restart.assert_called_once()
def test_reboot_routes_return_http_responses(self):
async def exercise_routes():
app = web.Application()
app.router.add_get("/easyuse/reboot-token", self.handlers.get_reboot_token)
app.router.add_post("/easyuse/reboot", self.handlers.reboot)
async with TestClient(TestServer(app)) as client:
token_response = await client.get("/easyuse/reboot-token")
self.assertEqual(token_response.status, 200)
self.assertEqual((await token_response.json())["token"], "test-reboot-token")
reboot_response = await client.post("/easyuse/reboot")
self.assertEqual(reboot_response.status, 403)
asyncio.run(exercise_routes())
if __name__ == "__main__":
unittest.main()