mirror of
https://github.com/rookiestar28/ComfyUI-OpenClaw.git
synced 2026-08-14 00:48:07 +00:00
646 lines
23 KiB
Python
646 lines
23 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import contextlib
|
|
import json
|
|
import logging
|
|
from typing import Any, Dict, Optional
|
|
|
|
try:
|
|
from ..services.access_control import require_admin_token
|
|
from ..services.aiohttp_compat import import_aiohttp_web
|
|
from ..services.async_utils import run_in_thread
|
|
from ..services.automation_composer import AutomationComposerService
|
|
from ..services.planner import PlannerService
|
|
from ..services.planner_registry import get_planner_registry
|
|
from ..services.rate_limit import build_rate_limit_response, check_rate_limit
|
|
from ..services.reasoning_redaction import (
|
|
audit_reasoning_reveal,
|
|
resolve_reasoning_reveal,
|
|
sanitize_operator_payload,
|
|
)
|
|
from ..services.refiner import RefinerService
|
|
except ImportError:
|
|
# Fallback for ComfyUI's non-package loader or ad-hoc imports.
|
|
from services.access_control import require_admin_token
|
|
from services.aiohttp_compat import import_aiohttp_web
|
|
from services.async_utils import run_in_thread
|
|
from services.automation_composer import AutomationComposerService
|
|
from services.planner import PlannerService
|
|
from services.planner_registry import get_planner_registry
|
|
from services.rate_limit import build_rate_limit_response, check_rate_limit
|
|
from services.reasoning_redaction import (
|
|
audit_reasoning_reveal,
|
|
resolve_reasoning_reveal,
|
|
sanitize_operator_payload,
|
|
)
|
|
from services.refiner import RefinerService
|
|
|
|
# R98: Endpoint Metadata
|
|
if __package__ and "." in __package__:
|
|
from ..services.endpoint_manifest import (
|
|
AuthTier,
|
|
RiskTier,
|
|
RoutePlane,
|
|
endpoint_metadata,
|
|
)
|
|
else:
|
|
from services.endpoint_manifest import (
|
|
AuthTier,
|
|
RiskTier,
|
|
RoutePlane,
|
|
endpoint_metadata,
|
|
)
|
|
|
|
logger = logging.getLogger("ComfyUI-OpenClaw.api.assist")
|
|
web = import_aiohttp_web()
|
|
|
|
# Payload size limits (character count for strings, base64 length for images)
|
|
MAX_REQUIREMENTS_LEN = 8000
|
|
MAX_STYLE_LEN = 2000
|
|
MAX_IMAGE_B64_LEN = 5 * 1024 * 1024 # ~5MB base64 string length
|
|
MAX_STREAM_DELTA_CHARS = 256
|
|
MAX_STREAM_PREVIEW_CHARS = 16_000
|
|
STREAM_KEEPALIVE_SEC = 1.0
|
|
|
|
|
|
def _planner_profiles_payload() -> Dict[str, Any]:
|
|
registry = get_planner_registry()
|
|
return {
|
|
"profiles": [
|
|
{
|
|
"id": profile.id,
|
|
"label": profile.label,
|
|
"description": profile.description,
|
|
"version": profile.version,
|
|
}
|
|
for profile in registry.list_profiles()
|
|
],
|
|
"default_profile": registry.get_default_profile_id(),
|
|
}
|
|
|
|
|
|
class AssistHandlers:
|
|
def __init__(self):
|
|
self.planner = PlannerService()
|
|
self.refiner = RefinerService()
|
|
self.composer = AutomationComposerService()
|
|
|
|
async def _require_admin_and_rate_limit(
|
|
self, request: web.Request
|
|
) -> Optional[web.Response]:
|
|
authorized, _err_msg = require_admin_token(request)
|
|
if not authorized:
|
|
return web.json_response({"error": "Unauthorized"}, status=401)
|
|
if not check_rate_limit(request, "admin"):
|
|
return build_rate_limit_response(
|
|
request,
|
|
"admin",
|
|
web_module=web,
|
|
error="Rate limit exceeded",
|
|
include_ok=False,
|
|
)
|
|
return None
|
|
|
|
async def _parse_json_body(
|
|
self, request: web.Request
|
|
) -> tuple[Optional[dict], Optional[web.Response]]:
|
|
try:
|
|
data = await request.json()
|
|
except Exception:
|
|
return None, web.json_response({"error": "Invalid JSON"}, status=400)
|
|
if not isinstance(data, dict):
|
|
return None, web.json_response(
|
|
{"error": "JSON object required"}, status=400
|
|
)
|
|
return data, None
|
|
|
|
def _validate_planner_payload(
|
|
self, data: dict
|
|
) -> tuple[Optional[dict], Optional[web.Response]]:
|
|
registry = get_planner_registry()
|
|
profile = data.get("profile", registry.get_default_profile_id())
|
|
requirements = data.get("requirements", "")
|
|
style = data.get("style_directives", "")
|
|
seed = data.get("seed", 0)
|
|
|
|
if not isinstance(profile, str):
|
|
return None, web.json_response(
|
|
{"error": "profile must be string"}, status=400
|
|
)
|
|
if not registry.get_profile(profile):
|
|
return None, web.json_response(
|
|
{"error": f"Unknown profile: {profile}"}, status=400
|
|
)
|
|
if not isinstance(requirements, str):
|
|
return None, web.json_response(
|
|
{"error": "requirements must be string"}, status=400
|
|
)
|
|
if not isinstance(style, str):
|
|
return None, web.json_response(
|
|
{"error": "style_directives must be string"}, status=400
|
|
)
|
|
if len(requirements) > MAX_REQUIREMENTS_LEN:
|
|
return None, web.json_response(
|
|
{"error": f"requirements exceeds {MAX_REQUIREMENTS_LEN} chars"},
|
|
status=400,
|
|
)
|
|
if len(style) > MAX_STYLE_LEN:
|
|
return None, web.json_response(
|
|
{"error": f"style_directives exceeds {MAX_STYLE_LEN} chars"}, status=400
|
|
)
|
|
try:
|
|
seed = int(seed)
|
|
except Exception:
|
|
seed = 0
|
|
return {
|
|
"profile": profile,
|
|
"requirements": requirements,
|
|
"style_directives": style,
|
|
"seed": seed,
|
|
}, None
|
|
|
|
def _validate_refiner_payload(
|
|
self, data: dict
|
|
) -> tuple[Optional[dict], Optional[web.Response]]:
|
|
image_b64 = data.get("image_b64", "")
|
|
orig_pos = data.get("orig_positive", "")
|
|
orig_neg = data.get("orig_negative", "")
|
|
issue = data.get("issue", "Fix issues")
|
|
params_json = data.get("params_json", "{}")
|
|
goal = data.get("goal", "Fix issues")
|
|
|
|
if not isinstance(image_b64, str) or not image_b64:
|
|
return None, web.json_response({"error": "image_b64 required"}, status=400)
|
|
if len(image_b64) > MAX_IMAGE_B64_LEN:
|
|
return None, web.json_response(
|
|
{"error": f"image_b64 exceeds {MAX_IMAGE_B64_LEN // 1024 // 1024}MB"},
|
|
status=400,
|
|
)
|
|
for key, value in (
|
|
("orig_positive", orig_pos),
|
|
("orig_negative", orig_neg),
|
|
("issue", issue),
|
|
("params_json", params_json),
|
|
("goal", goal),
|
|
):
|
|
if not isinstance(value, str):
|
|
return None, web.json_response(
|
|
{"error": f"{key} must be string"}, status=400
|
|
)
|
|
if len(orig_pos) > MAX_REQUIREMENTS_LEN or len(orig_neg) > MAX_REQUIREMENTS_LEN:
|
|
return None, web.json_response({"error": "Prompt too long"}, status=400)
|
|
|
|
return {
|
|
"image_b64": image_b64,
|
|
"orig_positive": orig_pos,
|
|
"orig_negative": orig_neg,
|
|
"issue": issue,
|
|
"params_json": params_json,
|
|
"goal": goal,
|
|
}, None
|
|
|
|
def _finalize_operator_payload(
|
|
self,
|
|
request: web.Request,
|
|
*,
|
|
target: str,
|
|
payload: Dict[str, Any],
|
|
service: Any,
|
|
) -> Dict[str, Any]:
|
|
admin_allowed, _ = require_admin_token(request)
|
|
reveal = resolve_reasoning_reveal(request, admin_authorized=admin_allowed)
|
|
audit_reasoning_reveal(request, target=target, decision=reveal)
|
|
|
|
final_payload = sanitize_operator_payload(payload)
|
|
consume_debug = getattr(service, "consume_last_reasoning_debug", None)
|
|
reasoning_debug = consume_debug() if callable(consume_debug) else None
|
|
if reveal["allowed"] and reasoning_debug not in (None, {}, []):
|
|
final_payload = dict(final_payload)
|
|
final_payload["debug"] = {"reasoning": reasoning_debug}
|
|
return final_payload
|
|
|
|
@staticmethod
|
|
def _sse_frame(event: str, payload: Dict[str, Any]) -> bytes:
|
|
return (
|
|
f"event: {event}\n"
|
|
f"data: {json.dumps(payload, ensure_ascii=False, separators=(',', ':'))}\n\n"
|
|
).encode("utf-8")
|
|
|
|
async def _write_sse_event(
|
|
self, response: web.StreamResponse, event: str, payload: Dict[str, Any]
|
|
) -> bool:
|
|
try:
|
|
await response.write(self._sse_frame(event, payload))
|
|
return True
|
|
except (ConnectionError, RuntimeError):
|
|
return False
|
|
|
|
async def _assist_stream_session(
|
|
self,
|
|
request: web.Request,
|
|
*,
|
|
kind: str,
|
|
worker_fn,
|
|
worker_kwargs: Dict[str, Any],
|
|
) -> web.StreamResponse:
|
|
response = web.StreamResponse(
|
|
status=200,
|
|
headers={
|
|
"Content-Type": "text/event-stream",
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
},
|
|
)
|
|
await response.prepare(request)
|
|
|
|
loop = asyncio.get_running_loop()
|
|
queue: asyncio.Queue = asyncio.Queue()
|
|
preview_chars = 0
|
|
|
|
def emit(event: str, payload: Dict[str, Any]) -> None:
|
|
try:
|
|
loop.call_soon_threadsafe(queue.put_nowait, (event, payload))
|
|
except RuntimeError:
|
|
pass
|
|
|
|
def on_text_delta(delta: str) -> None:
|
|
nonlocal preview_chars
|
|
if not isinstance(delta, str) or not delta:
|
|
return
|
|
remaining = MAX_STREAM_PREVIEW_CHARS - preview_chars
|
|
if remaining <= 0:
|
|
return
|
|
clipped = delta[: min(remaining, MAX_STREAM_DELTA_CHARS)]
|
|
if not clipped:
|
|
return
|
|
preview_chars += len(clipped)
|
|
emit("delta", {"text": clipped, "preview_chars": preview_chars})
|
|
|
|
async def runner() -> None:
|
|
emit("ready", {"ok": True, "kind": kind, "mode": "sse"})
|
|
emit(
|
|
"stage", {"phase": "dispatch", "message": "Dispatching assist request"}
|
|
)
|
|
try:
|
|
call_kwargs = dict(worker_kwargs)
|
|
call_kwargs["on_text_delta"] = on_text_delta
|
|
result = await run_in_thread(worker_fn, **call_kwargs)
|
|
emit(
|
|
"stage",
|
|
{"phase": "finalize", "message": "Parsing and validating output"},
|
|
)
|
|
if kind == "planner":
|
|
pos, neg, params = result
|
|
final_payload = {
|
|
"positive": pos,
|
|
"negative": neg,
|
|
"params": params,
|
|
}
|
|
elif kind == "refiner":
|
|
new_pos, new_neg, patch, rationale = result
|
|
final_payload = {
|
|
"refined_positive": new_pos,
|
|
"refined_negative": new_neg,
|
|
"param_patch": patch,
|
|
"rationale": rationale,
|
|
}
|
|
else:
|
|
final_payload = {"result": result}
|
|
service = (
|
|
self.planner
|
|
if kind == "planner"
|
|
else self.refiner if kind == "refiner" else self.composer
|
|
)
|
|
final_payload = self._finalize_operator_payload(
|
|
request,
|
|
target=f"assist.{kind}.stream",
|
|
payload=final_payload,
|
|
service=service,
|
|
)
|
|
emit(
|
|
"final",
|
|
{
|
|
"ok": True,
|
|
"kind": kind,
|
|
"result": final_payload,
|
|
"streaming": {
|
|
"preview_chars": preview_chars,
|
|
"preview_truncated": preview_chars
|
|
>= MAX_STREAM_PREVIEW_CHARS,
|
|
},
|
|
},
|
|
)
|
|
except Exception as e:
|
|
logger.exception("Assist streaming API failed (%s)", kind)
|
|
emit(
|
|
"error",
|
|
{"ok": False, "kind": kind, "error": "Internal server error"},
|
|
)
|
|
finally:
|
|
emit("__done__", {})
|
|
|
|
runner_task = asyncio.create_task(runner())
|
|
|
|
try:
|
|
while True:
|
|
try:
|
|
event, payload = await asyncio.wait_for(
|
|
queue.get(), timeout=STREAM_KEEPALIVE_SEC
|
|
)
|
|
except asyncio.TimeoutError:
|
|
if runner_task.done():
|
|
break
|
|
if not await self._write_sse_event(
|
|
response, "keepalive", {"ok": True}
|
|
):
|
|
break
|
|
continue
|
|
|
|
if event == "__done__":
|
|
break
|
|
if not await self._write_sse_event(response, event, payload):
|
|
break
|
|
finally:
|
|
if not runner_task.done():
|
|
runner_task.cancel()
|
|
with contextlib.suppress(BaseException):
|
|
await runner_task
|
|
return response
|
|
|
|
@endpoint_metadata(
|
|
auth=AuthTier.ADMIN,
|
|
risk=RiskTier.LOW,
|
|
summary="List planner profiles",
|
|
description="Returns Prompt Planner profiles from the active registry.",
|
|
audit="assist.planner_profiles",
|
|
plane=RoutePlane.ADMIN,
|
|
)
|
|
async def planner_profiles_handler(self, request):
|
|
auth_resp = await self._require_admin_and_rate_limit(request)
|
|
if auth_resp:
|
|
return auth_resp
|
|
return web.json_response(_planner_profiles_payload())
|
|
|
|
@endpoint_metadata(
|
|
auth=AuthTier.ADMIN,
|
|
risk=RiskTier.MEDIUM,
|
|
summary="Run planner",
|
|
description="Generate prompts from requirements via LLM.",
|
|
audit="assist.planner",
|
|
plane=RoutePlane.ADMIN,
|
|
)
|
|
async def planner_handler(self, request):
|
|
"""
|
|
POST /openclaw/assist/planner (legacy: /moltbot/assist/planner)
|
|
JSON: { profile, requirements, style_directives, seed }
|
|
"""
|
|
# Security: Admin Token required
|
|
auth_resp = await self._require_admin_and_rate_limit(request)
|
|
if auth_resp:
|
|
return auth_resp
|
|
data, error_resp = await self._parse_json_body(request)
|
|
if error_resp:
|
|
return error_resp
|
|
assert data is not None
|
|
payload, payload_err = self._validate_planner_payload(data)
|
|
if payload_err:
|
|
return payload_err
|
|
assert payload is not None
|
|
|
|
try:
|
|
# Run sync LLM call in thread pool to avoid blocking event loop
|
|
pos, neg, params = await run_in_thread(
|
|
self.planner.plan_generation,
|
|
payload["profile"],
|
|
payload["requirements"],
|
|
payload["style_directives"],
|
|
payload["seed"],
|
|
)
|
|
|
|
return web.json_response(
|
|
self._finalize_operator_payload(
|
|
request,
|
|
target="assist.planner",
|
|
payload={"positive": pos, "negative": neg, "params": params},
|
|
service=self.planner,
|
|
)
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.exception("Planner API failed")
|
|
return web.json_response({"error": "Internal server error"}, status=500)
|
|
|
|
@endpoint_metadata(
|
|
auth=AuthTier.ADMIN,
|
|
risk=RiskTier.MEDIUM,
|
|
summary="Run refiner",
|
|
description="Refine prompt/parameters based on feedback.",
|
|
audit="assist.refiner",
|
|
plane=RoutePlane.ADMIN,
|
|
)
|
|
async def refiner_handler(self, request):
|
|
"""
|
|
POST /openclaw/assist/refiner (legacy: /moltbot/assist/refiner)
|
|
JSON: { image_b64, orig_positive, orig_negative, issue, params_json, goal }
|
|
"""
|
|
# Security checks
|
|
auth_resp = await self._require_admin_and_rate_limit(request)
|
|
if auth_resp:
|
|
return auth_resp
|
|
data, error_resp = await self._parse_json_body(request)
|
|
if error_resp:
|
|
return error_resp
|
|
assert data is not None
|
|
payload, payload_err = self._validate_refiner_payload(data)
|
|
if payload_err:
|
|
return payload_err
|
|
assert payload is not None
|
|
|
|
try:
|
|
# Run sync LLM call in thread pool
|
|
new_pos, new_neg, patch, rationale = await run_in_thread(
|
|
self.refiner.refine_prompt,
|
|
**payload,
|
|
)
|
|
|
|
return web.json_response(
|
|
self._finalize_operator_payload(
|
|
request,
|
|
target="assist.refiner",
|
|
payload={
|
|
"refined_positive": new_pos,
|
|
"refined_negative": new_neg,
|
|
"param_patch": patch,
|
|
"rationale": rationale,
|
|
},
|
|
service=self.refiner,
|
|
)
|
|
)
|
|
except Exception as e:
|
|
logger.exception("Refiner API failed")
|
|
return web.json_response({"error": "Internal server error"}, status=500)
|
|
|
|
@endpoint_metadata(
|
|
auth=AuthTier.ADMIN,
|
|
risk=RiskTier.MEDIUM,
|
|
summary="Run planner (streaming)",
|
|
description="Generate prompts from requirements via LLM with SSE-style incremental updates.",
|
|
audit="assist.planner.stream",
|
|
plane=RoutePlane.ADMIN,
|
|
)
|
|
async def planner_stream_handler(self, request):
|
|
auth_resp = await self._require_admin_and_rate_limit(request)
|
|
if auth_resp:
|
|
return auth_resp
|
|
data, error_resp = await self._parse_json_body(request)
|
|
if error_resp:
|
|
return error_resp
|
|
assert data is not None
|
|
payload, payload_err = self._validate_planner_payload(data)
|
|
if payload_err:
|
|
return payload_err
|
|
assert payload is not None
|
|
|
|
return await self._assist_stream_session(
|
|
request,
|
|
kind="planner",
|
|
worker_fn=self.planner.plan_generation,
|
|
worker_kwargs=payload,
|
|
)
|
|
|
|
@endpoint_metadata(
|
|
auth=AuthTier.ADMIN,
|
|
risk=RiskTier.MEDIUM,
|
|
summary="Run refiner (streaming)",
|
|
description="Refine prompt/parameters with SSE-style incremental updates.",
|
|
audit="assist.refiner.stream",
|
|
plane=RoutePlane.ADMIN,
|
|
)
|
|
async def refiner_stream_handler(self, request):
|
|
auth_resp = await self._require_admin_and_rate_limit(request)
|
|
if auth_resp:
|
|
return auth_resp
|
|
data, error_resp = await self._parse_json_body(request)
|
|
if error_resp:
|
|
return error_resp
|
|
assert data is not None
|
|
payload, payload_err = self._validate_refiner_payload(data)
|
|
if payload_err:
|
|
return payload_err
|
|
assert payload is not None
|
|
|
|
return await self._assist_stream_session(
|
|
request,
|
|
kind="refiner",
|
|
worker_fn=self.refiner.refine_prompt,
|
|
worker_kwargs=payload,
|
|
)
|
|
|
|
@endpoint_metadata(
|
|
auth=AuthTier.ADMIN,
|
|
risk=RiskTier.MEDIUM,
|
|
summary="Compose automation payload",
|
|
description="Generate-only automation payload draft for trigger/webhook endpoints.",
|
|
audit="assist.compose",
|
|
plane=RoutePlane.ADMIN,
|
|
)
|
|
async def compose_handler(self, request):
|
|
"""
|
|
POST /openclaw/assist/automation/compose (legacy: /moltbot/assist/automation/compose)
|
|
JSON:
|
|
{
|
|
kind: "trigger" | "webhook",
|
|
template_id: str,
|
|
intent: str,
|
|
inputs_hint?: object,
|
|
profile_id?: str,
|
|
require_approval?: bool,
|
|
trace_id?: str,
|
|
callback?: object
|
|
}
|
|
"""
|
|
authorized, err_msg = require_admin_token(request)
|
|
if not authorized:
|
|
return web.json_response({"error": "Unauthorized"}, status=401)
|
|
|
|
if not check_rate_limit(request, "admin"):
|
|
return build_rate_limit_response(
|
|
request,
|
|
"admin",
|
|
web_module=web,
|
|
error="Rate limit exceeded",
|
|
include_ok=False,
|
|
)
|
|
|
|
try:
|
|
data = await request.json()
|
|
except Exception:
|
|
return web.json_response({"error": "Invalid JSON"}, status=400)
|
|
|
|
kind = data.get("kind")
|
|
template_id = data.get("template_id")
|
|
intent = data.get("intent")
|
|
inputs_hint = data.get("inputs_hint", {})
|
|
profile_id = data.get("profile_id")
|
|
require_approval = data.get("require_approval")
|
|
trace_id = data.get("trace_id")
|
|
callback = data.get("callback")
|
|
|
|
if not isinstance(kind, str) or kind.strip().lower() not in {
|
|
"trigger",
|
|
"webhook",
|
|
}:
|
|
return web.json_response(
|
|
{"error": "kind must be 'trigger' or 'webhook'"}, status=400
|
|
)
|
|
if not isinstance(template_id, str) or not template_id.strip():
|
|
return web.json_response({"error": "template_id is required"}, status=400)
|
|
if not isinstance(intent, str) or not intent.strip():
|
|
return web.json_response({"error": "intent is required"}, status=400)
|
|
if len(intent) > MAX_REQUIREMENTS_LEN:
|
|
return web.json_response(
|
|
{"error": f"intent exceeds {MAX_REQUIREMENTS_LEN} chars"}, status=400
|
|
)
|
|
if not isinstance(inputs_hint, dict):
|
|
return web.json_response(
|
|
{"error": "inputs_hint must be object"}, status=400
|
|
)
|
|
if profile_id is not None and not isinstance(profile_id, str):
|
|
return web.json_response({"error": "profile_id must be string"}, status=400)
|
|
if require_approval is not None and not isinstance(require_approval, bool):
|
|
return web.json_response(
|
|
{"error": "require_approval must be boolean"}, status=400
|
|
)
|
|
if trace_id is not None and not isinstance(trace_id, str):
|
|
return web.json_response({"error": "trace_id must be string"}, status=400)
|
|
if callback is not None and not isinstance(callback, dict):
|
|
return web.json_response({"error": "callback must be object"}, status=400)
|
|
|
|
try:
|
|
result = await run_in_thread(
|
|
self.composer.compose_payload,
|
|
kind=kind,
|
|
template_id=template_id,
|
|
intent=intent,
|
|
inputs_hint=inputs_hint,
|
|
profile_id=profile_id,
|
|
require_approval=require_approval,
|
|
trace_id=trace_id,
|
|
callback=callback,
|
|
)
|
|
return web.json_response(
|
|
self._finalize_operator_payload(
|
|
request,
|
|
target="assist.compose",
|
|
payload={"ok": True, **result},
|
|
service=self.composer,
|
|
)
|
|
)
|
|
except ValueError as e:
|
|
return web.json_response({"error": str(e)}, status=400)
|
|
except Exception:
|
|
logger.exception("Automation compose API failed")
|
|
return web.json_response({"error": "Internal server error"}, status=500)
|