fix(connector): target job cancellation requests

This commit is contained in:
rookiestar28
2026-06-24 13:21:30 +08:00
parent ba3c64bdd1
commit 4f64294b85
4 changed files with 199 additions and 11 deletions
+13 -4
View File
@@ -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
View File
@@ -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"
+74
View File
@@ -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."""
+51 -1
View File
@@ -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')