fix: make reboot routes awaitable (#1045)
This commit is contained in:
+2
-2
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user