mirror of
https://github.com/rookiestar28/ComfyUI-OpenClaw.git
synced 2026-08-14 00:48:07 +00:00
326 lines
11 KiB
Python
326 lines
11 KiB
Python
import json
|
|
import os
|
|
import sys
|
|
import unittest
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
# Check if aiohttp is available
|
|
try:
|
|
from aiohttp import web
|
|
|
|
AIOHTTP_AVAILABLE = True
|
|
except ModuleNotFoundError:
|
|
AIOHTTP_AVAILABLE = False
|
|
|
|
# Ensure we can import the module from current directory
|
|
sys.path.append(os.getcwd())
|
|
|
|
|
|
@unittest.skipIf(not AIOHTTP_AVAILABLE, "aiohttp not available")
|
|
class TestAssistAPI(unittest.IsolatedAsyncioTestCase):
|
|
"""Unit tests for Assist API endpoints (F8/F21)."""
|
|
|
|
async def asyncSetUp(self):
|
|
from api.assist import AssistHandlers
|
|
|
|
self.handler = AssistHandlers()
|
|
# Mock services to avoid LLM calls
|
|
self.handler.planner = MagicMock()
|
|
self.handler.refiner = MagicMock()
|
|
self.handler.composer = MagicMock()
|
|
|
|
async def test_planner_no_auth(self):
|
|
"""Test that planner rejects unauthenticated requests."""
|
|
request = AsyncMock()
|
|
request.headers = {}
|
|
|
|
with patch("api.assist.require_admin_token", return_value=(False, "Denied")):
|
|
resp = await self.handler.planner_handler(request)
|
|
self.assertEqual(resp.status, 401)
|
|
|
|
async def test_planner_success(self):
|
|
"""Test planner returns expected response on success."""
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"profile": "SDXL-v1",
|
|
"requirements": "cat",
|
|
"style_directives": "photorealistic",
|
|
"seed": 123,
|
|
}
|
|
)
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.run_in_thread") as mock_run_in_thread,
|
|
):
|
|
|
|
# Mock Service Return via run_in_thread
|
|
mock_run_in_thread.return_value = ("pos", "neg", {"width": 1024})
|
|
|
|
resp = await self.handler.planner_handler(request)
|
|
self.assertEqual(resp.status, 200)
|
|
body = json.loads(resp.body)
|
|
self.assertEqual(body["positive"], "pos")
|
|
self.assertEqual(body["params"]["width"], 1024)
|
|
|
|
async def test_planner_profiles_success(self):
|
|
request = AsyncMock()
|
|
|
|
class _Profile:
|
|
def __init__(self, profile_id, label):
|
|
self.id = profile_id
|
|
self.label = label
|
|
self.description = f"{label} desc"
|
|
self.version = "1.0"
|
|
|
|
registry = MagicMock()
|
|
registry.list_profiles.return_value = [_Profile("P1", "Profile One")]
|
|
registry.get_default_profile_id.return_value = "P1"
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.check_rate_limit", return_value=True),
|
|
patch("api.assist.get_planner_registry", return_value=registry),
|
|
):
|
|
resp = await self.handler.planner_profiles_handler(request)
|
|
|
|
self.assertEqual(resp.status, 200)
|
|
body = json.loads(resp.body)
|
|
self.assertEqual(body["default_profile"], "P1")
|
|
self.assertEqual(body["profiles"][0]["id"], "P1")
|
|
|
|
async def test_planner_rejects_unknown_profile(self):
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"profile": "missing",
|
|
"requirements": "cat",
|
|
"style_directives": "photorealistic",
|
|
}
|
|
)
|
|
registry = MagicMock()
|
|
registry.get_default_profile_id.return_value = "SDXL-v1"
|
|
registry.get_profile.return_value = None
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.get_planner_registry", return_value=registry),
|
|
):
|
|
resp = await self.handler.planner_handler(request)
|
|
|
|
self.assertEqual(resp.status, 400)
|
|
body = json.loads(resp.body)
|
|
self.assertEqual(body["error"], "Unknown profile: missing")
|
|
|
|
async def test_refiner_missing_image(self):
|
|
"""Test refiner rejects requests without image."""
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"orig_positive": "cat"
|
|
# No image_b64
|
|
}
|
|
)
|
|
|
|
with patch("api.assist.require_admin_token", return_value=(True, None)):
|
|
|
|
resp = await self.handler.refiner_handler(request)
|
|
self.assertEqual(resp.status, 400)
|
|
self.assertIn("error", json.loads(resp.body))
|
|
|
|
async def test_refiner_success(self):
|
|
"""Test refiner returns expected response on success."""
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"image_b64": "fakeBase64",
|
|
"orig_positive": "cat",
|
|
"orig_negative": "",
|
|
"issue": "bad hands",
|
|
"params_json": "{}",
|
|
"goal": "fix",
|
|
}
|
|
)
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.run_in_thread") as mock_run_in_thread,
|
|
):
|
|
|
|
# Mock Service
|
|
mock_run_in_thread.return_value = (
|
|
"new_pos",
|
|
"new_neg",
|
|
{"steps": 30},
|
|
"Fixed hands",
|
|
)
|
|
|
|
resp = await self.handler.refiner_handler(request)
|
|
self.assertEqual(resp.status, 200)
|
|
body = json.loads(resp.body)
|
|
self.assertEqual(body["refined_positive"], "new_pos")
|
|
self.assertEqual(body["rationale"], "Fixed hands")
|
|
|
|
async def test_planner_stream_success_emits_delta_and_final(self):
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"profile": "SDXL-v1",
|
|
"requirements": "cat",
|
|
"style_directives": "cinematic",
|
|
"seed": 123,
|
|
}
|
|
)
|
|
|
|
class FakeStreamResponse:
|
|
def __init__(self, status=200, headers=None):
|
|
self.status = status
|
|
self.headers = headers or {}
|
|
self.writes = []
|
|
|
|
async def prepare(self, _request):
|
|
return self
|
|
|
|
async def write(self, data):
|
|
self.writes.append(data)
|
|
return None
|
|
|
|
async def fake_run_in_thread(func, *args, **kwargs):
|
|
cb = kwargs.get("on_text_delta")
|
|
if callable(cb):
|
|
cb("partial-json ")
|
|
cb("preview")
|
|
return ("pos", "neg", {"width": 1024, "seed": 123})
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.check_rate_limit", return_value=True),
|
|
patch("api.assist.web.StreamResponse", FakeStreamResponse),
|
|
patch("api.assist.run_in_thread", side_effect=fake_run_in_thread),
|
|
):
|
|
resp = await self.handler.planner_stream_handler(request)
|
|
|
|
self.assertEqual(resp.status, 200)
|
|
text = b"".join(resp.writes).decode("utf-8", errors="replace")
|
|
self.assertIn("event: ready", text)
|
|
self.assertIn("event: delta", text)
|
|
self.assertIn("event: final", text)
|
|
self.assertIn('"positive":"pos"', text)
|
|
self.assertIn('"preview_chars"', text)
|
|
|
|
async def test_refiner_stream_unauthorized(self):
|
|
request = AsyncMock()
|
|
request.headers = {}
|
|
with patch("api.assist.require_admin_token", return_value=(False, "Denied")):
|
|
resp = await self.handler.refiner_stream_handler(request)
|
|
self.assertEqual(resp.status, 401)
|
|
|
|
async def test_planner_stream_internal_error_emits_error_event(self):
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"profile": "SDXL-v1",
|
|
"requirements": "cat",
|
|
"style_directives": "cinematic",
|
|
}
|
|
)
|
|
|
|
class FakeStreamResponse:
|
|
def __init__(self, status=200, headers=None):
|
|
self.status = status
|
|
self.headers = headers or {}
|
|
self.writes = []
|
|
|
|
async def prepare(self, _request):
|
|
return self
|
|
|
|
async def write(self, data):
|
|
self.writes.append(data)
|
|
return None
|
|
|
|
async def fake_run_in_thread(func, *args, **kwargs):
|
|
raise RuntimeError("boom")
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.check_rate_limit", return_value=True),
|
|
patch("api.assist.web.StreamResponse", FakeStreamResponse),
|
|
patch("api.assist.run_in_thread", side_effect=fake_run_in_thread),
|
|
):
|
|
resp = await self.handler.planner_stream_handler(request)
|
|
|
|
text = b"".join(resp.writes).decode("utf-8", errors="replace")
|
|
self.assertIn("event: error", text)
|
|
self.assertIn("Internal server error", text)
|
|
|
|
async def test_compose_no_auth(self):
|
|
"""Test compose rejects unauthenticated requests."""
|
|
request = AsyncMock()
|
|
request.headers = {}
|
|
|
|
with patch("api.assist.require_admin_token", return_value=(False, "Denied")):
|
|
resp = await self.handler.compose_handler(request)
|
|
self.assertEqual(resp.status, 401)
|
|
|
|
async def test_compose_invalid_kind(self):
|
|
"""Test compose validates kind field."""
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"kind": "unknown",
|
|
"template_id": "portrait_v1",
|
|
"intent": "make draft",
|
|
}
|
|
)
|
|
|
|
with patch("api.assist.require_admin_token", return_value=(True, None)):
|
|
resp = await self.handler.compose_handler(request)
|
|
self.assertEqual(resp.status, 400)
|
|
body = json.loads(resp.body)
|
|
self.assertIn("kind must be", body["error"])
|
|
|
|
async def test_compose_success(self):
|
|
"""Test compose returns draft payload on success."""
|
|
request = AsyncMock()
|
|
request.json = AsyncMock(
|
|
return_value={
|
|
"kind": "webhook",
|
|
"template_id": "portrait_v1",
|
|
"profile_id": "SDXL-v1",
|
|
"intent": "render portrait with soft light",
|
|
"inputs_hint": {"requirements": "portrait"},
|
|
"trace_id": "trace_123",
|
|
}
|
|
)
|
|
|
|
with (
|
|
patch("api.assist.require_admin_token", return_value=(True, None)),
|
|
patch("api.assist.run_in_thread") as mock_run_in_thread,
|
|
):
|
|
mock_run_in_thread.return_value = {
|
|
"kind": "webhook",
|
|
"payload": {
|
|
"version": 1,
|
|
"template_id": "portrait_v1",
|
|
"profile_id": "SDXL-v1",
|
|
"inputs": {"requirements": "portrait"},
|
|
"trace_id": "trace_123",
|
|
"job_id": None,
|
|
"callback": None,
|
|
},
|
|
"warnings": [],
|
|
"used_tool_calling": False,
|
|
}
|
|
|
|
resp = await self.handler.compose_handler(request)
|
|
self.assertEqual(resp.status, 200)
|
|
body = json.loads(resp.body)
|
|
self.assertTrue(body["ok"])
|
|
self.assertEqual(body["kind"], "webhook")
|
|
self.assertEqual(body["payload"]["template_id"], "portrait_v1")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|