mirror of
https://github.com/rookiestar28/ComfyUI-OpenClaw.git
synced 2026-08-14 00:48:07 +00:00
fix(connector): target job cancellation requests
This commit is contained in:
@@ -7,6 +7,7 @@ import json
|
||||
import logging
|
||||
import uuid
|
||||
from typing import Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
from .config import ConnectorConfig
|
||||
|
||||
@@ -68,7 +69,7 @@ class OpenClawClient:
|
||||
async with session.request(
|
||||
method, url, headers=self.headers, json=json_data, timeout=timeout
|
||||
) as resp:
|
||||
result = {"ok": resp.status in (200, 201, 202)}
|
||||
result = {"ok": resp.status in (200, 201, 202), "status": resp.status}
|
||||
|
||||
try:
|
||||
data = await resp.json()
|
||||
@@ -166,9 +167,17 @@ class OpenClawClient:
|
||||
}
|
||||
return await self._request("POST", "/openclaw/triggers/fire", data)
|
||||
|
||||
async def interrupt_output(self) -> dict:
|
||||
# Remediation: Cancel -> Interrupt (Global)
|
||||
return await self._request("POST", "/api/interrupt", {})
|
||||
async def cancel_job(self, job_id: str) -> dict:
|
||||
encoded_job_id = quote(str(job_id), safe="")
|
||||
return await self._request("POST", f"/api/jobs/{encoded_job_id}/cancel", {})
|
||||
|
||||
async def cancel_jobs(self, job_ids: list[str]) -> dict:
|
||||
return await self._request("POST", "/api/jobs/cancel", {"job_ids": job_ids})
|
||||
|
||||
async def interrupt_output(self, prompt_id: Optional[str] = None) -> dict:
|
||||
# No prompt_id means explicit global interrupt. A prompt_id is targeted.
|
||||
payload = {"prompt_id": str(prompt_id)} if prompt_id else {}
|
||||
return await self._request("POST", "/api/interrupt", payload)
|
||||
|
||||
async def get_view(
|
||||
self, filename: str, subfolder: str = "", type: str = "output"
|
||||
|
||||
+61
-6
@@ -498,13 +498,68 @@ class CommandRouter:
|
||||
if err := self._require_admin_token_configured():
|
||||
return err
|
||||
|
||||
# Remediation: Global Interrupt
|
||||
res = await self.client.interrupt_output()
|
||||
if res.get("ok"):
|
||||
return CommandResponse(text="[Stop] Global Interrupt sent to ComfyUI.")
|
||||
else:
|
||||
targets = self._parse_stop_targets(args)
|
||||
if not targets:
|
||||
res = await self.client.interrupt_output()
|
||||
if res.get("ok"):
|
||||
return CommandResponse(text="[Stop] Global Interrupt sent to ComfyUI.")
|
||||
return CommandResponse(text=f"[Stop Failed] {res.get('error')}")
|
||||
|
||||
if len(targets) == 1:
|
||||
job_id = targets[0]
|
||||
res = await self.client.cancel_job(job_id)
|
||||
if res.get("ok"):
|
||||
return CommandResponse(
|
||||
text=f"[Stop] Cancellation requested for job {job_id}."
|
||||
)
|
||||
|
||||
# IMPORTANT: Targeted stops must never degrade to no-payload global
|
||||
# interrupt. Older-host fallback is allowed only with prompt_id set.
|
||||
if self._jobs_cancel_unsupported(res):
|
||||
fallback = await self.client.interrupt_output(prompt_id=job_id)
|
||||
if fallback.get("ok"):
|
||||
return CommandResponse(
|
||||
text=(
|
||||
f"[Stop] Targeted interrupt sent for job {job_id} "
|
||||
"(jobs cancel unsupported)."
|
||||
)
|
||||
)
|
||||
return CommandResponse(text=f"[Stop Failed] {fallback.get('error')}")
|
||||
|
||||
return CommandResponse(text=f"[Stop Failed] {res.get('error')}")
|
||||
|
||||
res = await self.client.cancel_jobs(targets)
|
||||
if res.get("ok"):
|
||||
return CommandResponse(
|
||||
text=f"[Stop] Cancellation requested for {len(targets)} jobs."
|
||||
)
|
||||
return CommandResponse(text=f"[Stop Failed] {res.get('error')}")
|
||||
|
||||
@staticmethod
|
||||
def _parse_stop_targets(args: List[str]) -> List[str]:
|
||||
targets: List[str] = []
|
||||
for arg in args:
|
||||
for part in str(arg).split(","):
|
||||
target = part.strip()
|
||||
if target:
|
||||
targets.append(target)
|
||||
return targets
|
||||
|
||||
@staticmethod
|
||||
def _jobs_cancel_unsupported(res: Dict[str, Any]) -> bool:
|
||||
status = res.get("status")
|
||||
if status in (404, 405, 501):
|
||||
return True
|
||||
error = str(res.get("error", "")).lower()
|
||||
unsupported_markers = (
|
||||
"404",
|
||||
"not found",
|
||||
"method not allowed",
|
||||
"unsupported",
|
||||
"not implemented",
|
||||
)
|
||||
return any(marker in error for marker in unsupported_markers)
|
||||
|
||||
async def _handle_approvals_list(
|
||||
self, req: CommandRequest, args: List[str]
|
||||
) -> CommandResponse:
|
||||
@@ -679,7 +734,7 @@ class CommandRouter:
|
||||
"OpenClaw Connector\n"
|
||||
"/status - Check system health and queue\n"
|
||||
"/run <template> [prompt] [k=v] - Run a generation (trusted users auto-exec; others require approval)\n"
|
||||
"/stop - Global Interrupt (Admin)\n"
|
||||
"/stop [job_id ...] - Cancel jobs by id; no args sends Global Interrupt (Admin)\n"
|
||||
"/history <id> - Job details\n"
|
||||
"/jobs - Queue summary\n"
|
||||
"Admin Only:\n"
|
||||
|
||||
@@ -112,8 +112,82 @@ class TestOpenClawClient(unittest.TestCase):
|
||||
asyncio.run(self.client.interrupt_output())
|
||||
|
||||
method, url = mock_session.request.call_args[0]
|
||||
kwargs = mock_session.request.call_args[1]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertTrue(url.endswith("/api/interrupt"))
|
||||
self.assertEqual(kwargs["json"], {})
|
||||
|
||||
def test_targeted_interrupt_output(self):
|
||||
"""Verify targeted interrupt carries prompt_id and does not use global payload."""
|
||||
with patch("connector.openclaw_client._create_session") as create_session:
|
||||
mock_session = MagicMock()
|
||||
mock_session.close = AsyncMock()
|
||||
create_session.return_value = mock_session
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status = 200
|
||||
mock_resp.json = AsyncMock(return_value={})
|
||||
|
||||
mock_ctx = MagicMock()
|
||||
mock_ctx.__aenter__.return_value = mock_resp
|
||||
mock_ctx.__aexit__.return_value = None
|
||||
mock_session.request.return_value = mock_ctx
|
||||
|
||||
asyncio.run(self.client.interrupt_output(prompt_id="p-1"))
|
||||
|
||||
method, url = mock_session.request.call_args[0]
|
||||
kwargs = mock_session.request.call_args[1]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertTrue(url.endswith("/api/interrupt"))
|
||||
self.assertEqual(kwargs["json"], {"prompt_id": "p-1"})
|
||||
|
||||
def test_cancel_job_uses_jobs_namespace(self):
|
||||
"""Verify single job cancel calls /api/jobs/{job_id}/cancel."""
|
||||
with patch("connector.openclaw_client._create_session") as create_session:
|
||||
mock_session = MagicMock()
|
||||
mock_session.close = AsyncMock()
|
||||
create_session.return_value = mock_session
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status = 200
|
||||
mock_resp.json = AsyncMock(return_value={})
|
||||
|
||||
mock_ctx = MagicMock()
|
||||
mock_ctx.__aenter__.return_value = mock_resp
|
||||
mock_ctx.__aexit__.return_value = None
|
||||
mock_session.request.return_value = mock_ctx
|
||||
|
||||
asyncio.run(self.client.cancel_job("job/one"))
|
||||
|
||||
method, url = mock_session.request.call_args[0]
|
||||
kwargs = mock_session.request.call_args[1]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertTrue(url.endswith("/api/jobs/job%2Fone/cancel"))
|
||||
self.assertEqual(kwargs["json"], {})
|
||||
|
||||
def test_cancel_jobs_uses_batch_jobs_namespace(self):
|
||||
"""Verify batch job cancel calls /api/jobs/cancel with job_ids."""
|
||||
with patch("connector.openclaw_client._create_session") as create_session:
|
||||
mock_session = MagicMock()
|
||||
mock_session.close = AsyncMock()
|
||||
create_session.return_value = mock_session
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status = 200
|
||||
mock_resp.json = AsyncMock(return_value={})
|
||||
|
||||
mock_ctx = MagicMock()
|
||||
mock_ctx.__aenter__.return_value = mock_resp
|
||||
mock_ctx.__aexit__.return_value = None
|
||||
mock_session.request.return_value = mock_ctx
|
||||
|
||||
asyncio.run(self.client.cancel_jobs(["p-1", "p-2"]))
|
||||
|
||||
method, url = mock_session.request.call_args[0]
|
||||
kwargs = mock_session.request.call_args[1]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertTrue(url.endswith("/api/jobs/cancel"))
|
||||
self.assertEqual(kwargs["json"], {"job_ids": ["p-1", "p-2"]})
|
||||
|
||||
def test_get_approvals_query(self):
|
||||
"""Verify get_approvals uses query param and parses nested response."""
|
||||
|
||||
@@ -69,6 +69,8 @@ class TestCommandRouterPhase2(unittest.TestCase):
|
||||
}
|
||||
)
|
||||
self.client.interrupt_output = AsyncMock(return_value={"ok": True})
|
||||
self.client.cancel_job = AsyncMock(return_value={"ok": True})
|
||||
self.client.cancel_jobs = AsyncMock(return_value={"ok": True})
|
||||
|
||||
# Phase 4: Approve result
|
||||
self.client.approve_request = AsyncMock(
|
||||
@@ -163,13 +165,61 @@ class TestCommandRouterPhase2(unittest.TestCase):
|
||||
req = self._req("/stop", sender="999")
|
||||
resp = asyncio.run(self.router.handle(req))
|
||||
self.assertIn("Global Interrupt sent", resp.text)
|
||||
self.client.interrupt_output.assert_called_once()
|
||||
self.client.interrupt_output.assert_called_once_with()
|
||||
self.client.cancel_job.assert_not_called()
|
||||
self.client.cancel_jobs.assert_not_called()
|
||||
|
||||
# Deny non-admin
|
||||
req = self._req("/stop", sender="123")
|
||||
resp = asyncio.run(self.router.handle(req))
|
||||
self.assertIn("Access Denied", resp.text)
|
||||
|
||||
def test_interrupt_single_job_uses_jobs_cancel(self):
|
||||
req = self._req("/stop p-123", sender="999")
|
||||
resp = asyncio.run(self.router.handle(req))
|
||||
|
||||
self.assertIn("Cancellation requested for job p-123", resp.text)
|
||||
self.client.cancel_job.assert_called_once_with("p-123")
|
||||
self.client.cancel_jobs.assert_not_called()
|
||||
self.client.interrupt_output.assert_not_called()
|
||||
|
||||
def test_interrupt_multiple_jobs_uses_batch_cancel(self):
|
||||
req = self._req("/cancel p-1,p-2 p-3", sender="999")
|
||||
resp = asyncio.run(self.router.handle(req))
|
||||
|
||||
self.assertIn("Cancellation requested for 3 jobs", resp.text)
|
||||
self.client.cancel_jobs.assert_called_once_with(["p-1", "p-2", "p-3"])
|
||||
self.client.cancel_job.assert_not_called()
|
||||
self.client.interrupt_output.assert_not_called()
|
||||
|
||||
def test_interrupt_single_job_falls_back_to_targeted_interrupt_only(self):
|
||||
self.client.cancel_job.return_value = {
|
||||
"ok": False,
|
||||
"status": 404,
|
||||
"error": "HTTP 404",
|
||||
}
|
||||
|
||||
req = self._req("/interrupt p-legacy", sender="999")
|
||||
resp = asyncio.run(self.router.handle(req))
|
||||
|
||||
self.assertIn("Targeted interrupt sent for job p-legacy", resp.text)
|
||||
self.client.cancel_job.assert_called_once_with("p-legacy")
|
||||
self.client.interrupt_output.assert_called_once_with(prompt_id="p-legacy")
|
||||
|
||||
def test_interrupt_multiple_jobs_does_not_global_fallback(self):
|
||||
self.client.cancel_jobs.return_value = {
|
||||
"ok": False,
|
||||
"status": 404,
|
||||
"error": "HTTP 404",
|
||||
}
|
||||
|
||||
req = self._req("/stop p-1 p-2", sender="999")
|
||||
resp = asyncio.run(self.router.handle(req))
|
||||
|
||||
self.assertIn("[Stop Failed]", resp.text)
|
||||
self.client.cancel_jobs.assert_called_once_with(["p-1", "p-2"])
|
||||
self.client.interrupt_output.assert_not_called()
|
||||
|
||||
def test_complex_quotes(self):
|
||||
# Unbalanced
|
||||
req = self._req('/run "oops')
|
||||
|
||||
Reference in New Issue
Block a user