fix: make reboot routes awaitable (#1045)

This commit is contained in:
yolain
2026-09-29 12:23:08 +08:00
parent 8a0f2fc412
commit 67afdc8204
2 changed files with 26 additions and 8 deletions
+2 -2
View File
@@ -68,14 +68,14 @@ def _same_origin_request(request):
@PromptServer.instance.routes.get("/easyuse/reboot-token")
def get_reboot_token(request):
async def get_reboot_token(request):
if not _same_origin_request(request):
return web.Response(status=403)
return web.json_response({"token": _reboot_token}, headers={"Cache-Control": "no-store"})
@PromptServer.instance.routes.post("/easyuse/reboot")
def reboot(request):
async def reboot(request):
token = request.headers.get("X-EasyUse-Reboot-Token", "")
if not _same_origin_request(request) or not hmac.compare_digest(token, _reboot_token):
return web.Response(status=403)
+24 -6
View File
@@ -1,6 +1,7 @@
import ast
import asyncio
import hashlib
import inspect
import json
import os
from pathlib import Path
@@ -14,6 +15,7 @@ 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
@@ -133,24 +135,40 @@ class SecurityRouteTests(unittest.TestCase):
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(self.handlers.get_reboot_token(request).status, 403)
self.assertEqual(asyncio.run(self.handlers.get_reboot_token(request)).status, 403)
request.headers = {"Sec-Fetch-Site": "same-origin"}
self.assertEqual(json.loads(self.handlers.get_reboot_token(request).text)["token"], "test-reboot-token")
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(self.handlers.reboot(request).status, 403)
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(self.handlers.reboot(request).status, 403)
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(self.handlers.reboot(request).status, 403)
self.assertEqual(asyncio.run(self.handlers.reboot(request)).status, 403)
request.headers["Sec-Fetch-Site"] = "same-origin"
self.assertEqual(self.handlers.reboot(request), "restarted")
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()