mirror of
https://github.com/rookiestar28/ComfyUI-OpenClaw.git
synced 2026-08-14 00:48:07 +00:00
349 lines
11 KiB
Python
349 lines
11 KiB
Python
"""
|
|
Common types and registry for constrained transforms (S35/F42).
|
|
Refactored to avoid circular imports.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
import logging
|
|
import os
|
|
import time
|
|
from dataclasses import asdict, dataclass, field
|
|
from enum import Enum
|
|
from pathlib import Path
|
|
from typing import Any, Dict, List, Optional, Set
|
|
|
|
logger = logging.getLogger("ComfyUI-OpenClaw.services.transform_common")
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Feature gate
|
|
# ---------------------------------------------------------------------------
|
|
|
|
_FEATURE_FLAG = "OPENCLAW_ENABLE_TRANSFORMS"
|
|
|
|
|
|
def is_transforms_enabled() -> bool:
|
|
"""Check if constrained transforms are enabled (default: OFF)."""
|
|
val = os.environ.get(_FEATURE_FLAG, "").strip().lower()
|
|
return val in ("1", "true", "yes", "on")
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Runtime limits
|
|
# ---------------------------------------------------------------------------
|
|
|
|
DEFAULT_TRANSFORM_TIMEOUT_SEC = 5
|
|
DEFAULT_MAX_OUTPUT_BYTES = 64 * 1024 # 64KB
|
|
DEFAULT_MAX_TRANSFORMS_PER_REQUEST = 5
|
|
MAX_TRANSFORM_MODULE_SIZE_BYTES = 50 * 1024 # 50KB — prevent loading huge scripts
|
|
|
|
|
|
@dataclass
|
|
class TransformLimits:
|
|
"""Runtime limits for transform execution."""
|
|
|
|
timeout_sec: float = DEFAULT_TRANSFORM_TIMEOUT_SEC
|
|
max_output_bytes: int = DEFAULT_MAX_OUTPUT_BYTES
|
|
max_transforms_per_request: int = DEFAULT_MAX_TRANSFORMS_PER_REQUEST
|
|
|
|
@classmethod
|
|
def from_env(cls) -> "TransformLimits":
|
|
"""Load limits from environment variables."""
|
|
|
|
def _env_int(key: str, default: int) -> int:
|
|
try:
|
|
return int(os.environ.get(key, str(default)))
|
|
except (ValueError, TypeError):
|
|
return default
|
|
|
|
def _env_float(key: str, default: float) -> float:
|
|
try:
|
|
return float(os.environ.get(key, str(default)))
|
|
except (ValueError, TypeError):
|
|
return default
|
|
|
|
return cls(
|
|
timeout_sec=_env_float(
|
|
"OPENCLAW_TRANSFORM_TIMEOUT", DEFAULT_TRANSFORM_TIMEOUT_SEC
|
|
),
|
|
max_output_bytes=_env_int(
|
|
"OPENCLAW_TRANSFORM_MAX_OUTPUT", DEFAULT_MAX_OUTPUT_BYTES
|
|
),
|
|
max_transforms_per_request=_env_int(
|
|
"OPENCLAW_TRANSFORM_MAX_PER_REQUEST", DEFAULT_MAX_TRANSFORMS_PER_REQUEST
|
|
),
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Transform result
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TransformStatus(str, Enum):
|
|
SUCCESS = "success"
|
|
ERROR = "error"
|
|
TIMEOUT = "timeout"
|
|
DENIED = "denied"
|
|
SKIPPED = "skipped"
|
|
|
|
|
|
@dataclass
|
|
class TransformResult:
|
|
"""Result of a single transform execution."""
|
|
|
|
transform_id: str
|
|
status: str # TransformStatus.value
|
|
output: Optional[Dict[str, Any]] = None
|
|
error: str = ""
|
|
duration_ms: float = 0.0
|
|
output_bytes: int = 0
|
|
audit: Dict[str, Any] = field(default_factory=dict)
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
d: Dict[str, Any] = {
|
|
"transform_id": self.transform_id,
|
|
"status": self.status,
|
|
"duration_ms": round(self.duration_ms, 2),
|
|
"output_bytes": self.output_bytes,
|
|
}
|
|
if self.output is not None:
|
|
d["output"] = self.output
|
|
if self.error:
|
|
d["error"] = self.error
|
|
if self.audit:
|
|
d["audit"] = self.audit
|
|
return d
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Transform registry (trusted modules)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@dataclass
|
|
class TrustedTransform:
|
|
"""A registered, integrity-pinned transform module."""
|
|
|
|
id: str
|
|
label: str
|
|
module_path: str # Absolute path to .py module
|
|
sha256: str # Integrity hash of the module file
|
|
description: str = ""
|
|
trusted_source: str = "" # Who published this transform
|
|
registered_at: float = 0.0
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return asdict(self)
|
|
|
|
|
|
class TransformRegistryError(Exception):
|
|
"""Error in transform registry operations."""
|
|
|
|
pass
|
|
|
|
|
|
class TransformRegistry:
|
|
"""
|
|
Manages trusted transform modules with integrity pinning.
|
|
|
|
Transforms can only be loaded from explicitly trusted directories.
|
|
Each module is pinned by its SHA256 hash at registration time.
|
|
"""
|
|
|
|
def __init__(self, state_dir: str, trusted_dirs: Optional[List[str]] = None):
|
|
self._state_dir = state_dir
|
|
self._registry_dir = os.path.join(state_dir, "transforms")
|
|
self._index_path = os.path.join(self._registry_dir, "registry.json")
|
|
self._transforms: Dict[str, TrustedTransform] = {}
|
|
|
|
# Trusted directories where transform modules can live
|
|
self._trusted_dirs: Set[str] = set()
|
|
if trusted_dirs:
|
|
for d in trusted_dirs:
|
|
resolved = str(Path(d).resolve())
|
|
self._trusted_dirs.add(resolved)
|
|
|
|
os.makedirs(self._registry_dir, exist_ok=True)
|
|
self._load()
|
|
|
|
def _load(self) -> None:
|
|
"""Load transform registry from disk."""
|
|
if not os.path.exists(self._index_path):
|
|
self._transforms = {}
|
|
return
|
|
try:
|
|
with open(self._index_path, "r", encoding="utf-8") as f:
|
|
data = json.load(f)
|
|
self._transforms = {}
|
|
for tid, tdata in data.items():
|
|
self._transforms[tid] = TrustedTransform(**tdata)
|
|
except Exception as e:
|
|
logger.error(f"Failed to load transform registry: {e}")
|
|
self._transforms = {}
|
|
|
|
def _save(self) -> None:
|
|
"""Persist transform registry to disk."""
|
|
try:
|
|
data = {k: v.to_dict() for k, v in self._transforms.items()}
|
|
tmp_path = self._index_path + ".tmp"
|
|
with open(tmp_path, "w", encoding="utf-8") as f:
|
|
json.dump(data, f, indent=2)
|
|
f.write("\n")
|
|
os.replace(tmp_path, self._index_path)
|
|
except Exception as e:
|
|
logger.error(f"Failed to save transform registry: {e}")
|
|
|
|
@staticmethod
|
|
def _compute_sha256(file_path: str) -> str:
|
|
"""Compute SHA256 hash of a file."""
|
|
h = hashlib.sha256()
|
|
with open(file_path, "rb") as f:
|
|
for chunk in iter(lambda: f.read(4096), b""):
|
|
h.update(chunk)
|
|
return h.hexdigest()
|
|
|
|
def _is_in_trusted_dir(self, module_path: str) -> bool:
|
|
"""Check if a module path is inside a trusted directory."""
|
|
resolved = str(Path(module_path).resolve())
|
|
for trusted in self._trusted_dirs:
|
|
if resolved.startswith(trusted + os.sep) or resolved == trusted:
|
|
return True
|
|
return False
|
|
|
|
def register_transform(
|
|
self,
|
|
transform_id: str,
|
|
module_path: str,
|
|
*,
|
|
label: str = "",
|
|
description: str = "",
|
|
trusted_source: str = "",
|
|
) -> TrustedTransform:
|
|
"""
|
|
Register a transform module with integrity pinning.
|
|
|
|
The module must be in a trusted directory and within size limits.
|
|
"""
|
|
if not is_transforms_enabled():
|
|
raise TransformRegistryError(
|
|
f"Transforms disabled. Set {_FEATURE_FLAG}=1 to enable."
|
|
)
|
|
|
|
abs_path = str(Path(module_path).resolve())
|
|
|
|
# Security: must be in trusted directory
|
|
if not self._is_in_trusted_dir(abs_path):
|
|
raise TransformRegistryError(
|
|
f"Module path is not in a trusted directory: {abs_path}"
|
|
)
|
|
|
|
if not os.path.isfile(abs_path):
|
|
raise TransformRegistryError(f"Module file not found: {abs_path}")
|
|
|
|
# Size check
|
|
file_size = os.path.getsize(abs_path)
|
|
if file_size > MAX_TRANSFORM_MODULE_SIZE_BYTES:
|
|
raise TransformRegistryError(
|
|
f"Module exceeds size limit ({file_size} > {MAX_TRANSFORM_MODULE_SIZE_BYTES})"
|
|
)
|
|
|
|
# Must be a .py file
|
|
if not abs_path.endswith(".py"):
|
|
raise TransformRegistryError("Only .py modules are allowed as transforms")
|
|
|
|
sha256 = self._compute_sha256(abs_path)
|
|
|
|
transform = TrustedTransform(
|
|
id=transform_id,
|
|
label=label or transform_id,
|
|
module_path=abs_path,
|
|
sha256=sha256,
|
|
description=description,
|
|
trusted_source=trusted_source,
|
|
registered_at=time.time(),
|
|
)
|
|
|
|
self._transforms[transform_id] = transform
|
|
self._save()
|
|
logger.info(f"F42: Registered transform '{transform_id}' from {abs_path}")
|
|
return transform
|
|
|
|
def unregister_transform(self, transform_id: str) -> bool:
|
|
"""Remove a transform from the registry."""
|
|
if not is_transforms_enabled():
|
|
raise TransformRegistryError(
|
|
f"Transforms disabled. Set {_FEATURE_FLAG}=1 to enable."
|
|
)
|
|
|
|
if transform_id not in self._transforms:
|
|
return False
|
|
|
|
del self._transforms[transform_id]
|
|
self._save()
|
|
logger.info(f"F42: Unregistered transform '{transform_id}'")
|
|
return True
|
|
|
|
def get_transform(self, transform_id: str) -> Optional[TrustedTransform]:
|
|
"""Get a registered transform by ID."""
|
|
return self._transforms.get(transform_id)
|
|
|
|
def list_transforms(self) -> List[TrustedTransform]:
|
|
"""List all registered transforms."""
|
|
return list(self._transforms.values())
|
|
|
|
def verify_integrity(self, transform_id: str) -> bool:
|
|
"""Verify that a registered transform's file hasn't been modified."""
|
|
transform = self._transforms.get(transform_id)
|
|
if not transform:
|
|
return False
|
|
|
|
if not os.path.isfile(transform.module_path):
|
|
return False
|
|
|
|
actual_hash = self._compute_sha256(transform.module_path)
|
|
return actual_hash == transform.sha256
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Module-level convenience
|
|
# ---------------------------------------------------------------------------
|
|
|
|
_registry: Optional[TransformRegistry] = None
|
|
|
|
|
|
def get_transform_registry() -> TransformRegistry:
|
|
"""Get or create the global transform registry."""
|
|
global _registry
|
|
if _registry is None:
|
|
try:
|
|
from .state_dir import get_state_dir
|
|
|
|
state_dir = get_state_dir()
|
|
except ImportError:
|
|
try:
|
|
from services.state_dir import get_state_dir
|
|
|
|
state_dir = get_state_dir()
|
|
except ImportError:
|
|
state_dir = os.path.join(
|
|
os.path.dirname(os.path.dirname(__file__)), "data"
|
|
)
|
|
|
|
# Default trusted directory: pack-local transforms dir
|
|
pack_root = Path(__file__).resolve().parent.parent
|
|
trusted_dirs = [str(pack_root / "data" / "transforms")]
|
|
|
|
# Allow additional trusted dirs from env
|
|
extra = os.environ.get("OPENCLAW_TRANSFORM_TRUSTED_DIRS", "")
|
|
if extra:
|
|
for d in extra.split(os.pathsep):
|
|
d = d.strip()
|
|
if d:
|
|
trusted_dirs.append(d)
|
|
|
|
_registry = TransformRegistry(state_dir, trusted_dirs=trusted_dirs)
|
|
return _registry
|