diff --git a/README.md b/README.md index 3bca92a..275e92b 100644 --- a/README.md +++ b/README.md @@ -37,8 +37,11 @@ The existing `python server.py` launcher also works using the installed environm - `get_kepler_formal_info`: report installed versions, build revision, and the available Python API options/results. - `run_kepler_formal_yaml`: verify the pair specified by an existing YAML file. - `create_yaml_and_run_kepler_formal`: create a YAML file from two Verilog paths and optional common Liberty libraries, then verify the pair. +- `open_session`, `load_designs`, `verify_session`, `close_session`: load designs once and reuse one Python worker for repeated checks. +- `attach_session`: bind to a bridge in a Python interpreter that already owns NajaEDA designs. +- `set_session`, `list_sessions`: select and inspect reusable sessions. -The tools use the Python library in a separate Python worker for each call. This keeps native solver output out of the MCP stdio channel and allows `timeout_seconds` to stop a verification. For verification, the worker loads both designs with NajaEDA and passes their live design handles to `kepler_formal.verify_designs`. +The two file tools use a separate Python worker for each call. Session tools retain one worker or attach to a caller's interpreter, preserving loaded designs between checks. In both cases, native solver output is separate from MCP messages, and Kepler verifies the existing NajaEDA design handles. `get_kepler_formal_info` reuses the selected session when one exists. A minimal YAML file is: @@ -61,6 +64,8 @@ Supported verification settings are `verification` (`lec` or `sec`; `mode` is an All current `VerificationOptions` fields are exposed. See [Python API coverage and SEC usage](docs/python-api.md) for the mapping and supported choices. MCP tool schemas list the accepted modes, solvers, engines and encodings; `get_kepler_formal_info` reads them from the installed library. +See [persistent sessions and attaching to live NajaEDA designs](docs/sessions.md) to reuse loaded netlists across calls or verify designs already owned by another Python process. + ## Results The returned JSON separates tool execution from the verification verdict: diff --git a/docs/python-api.md b/docs/python-api.md index f549eeb..56fcc01 100644 --- a/docs/python-api.md +++ b/docs/python-api.md @@ -6,7 +6,8 @@ NajaEDA loads both designs inside the worker, and Kepler borrows their existing native handles for the check. All nine `VerificationOptions` fields are available through -`create_yaml_and_run_kepler_formal` and existing YAML configurations: +`create_yaml_and_run_kepler_formal`, `verify_session`, and existing YAML configurations +(attached sessions cannot request process-relative skipped-output report files): | Python option | MCP argument | YAML key | Accepted values | | --- | --- | --- | --- | @@ -66,11 +67,15 @@ Call `get_kepler_formal_info` without arguments for the installed Kepler version and git revision, NajaEDA version, supported enums/statuses, Python option defaults, and result field names. This exposes `version()` and `git_hash()` through MCP and reads the API metadata from the installed package. +With an active session it queries that interpreter; otherwise it uses a temporary worker. `NativeDesign` and `from_najaeda()` operate on live Python/C++ objects within one process. The worker uses this interface internally; raw pointers and handles from another process cannot be sent through MCP JSON. The `najaeda` export is a compatibility alias for the separate NajaEDA package, not a remote netlist API. +To reuse loaded designs or bind to an existing interpreter, use the +[session tools and Python bridge](sessions.md). Verification then runs in the +process that owns those designs. The tests compare MCP option names, enum choices, and returned fields against the installed Python API, so missing coverage is detected when the dependency diff --git a/docs/sessions.md b/docs/sessions.md new file mode 100644 index 0000000..d07c637 --- /dev/null +++ b/docs/sessions.md @@ -0,0 +1,84 @@ +# Persistent verification sessions + +Sessions retain loaded designs between MCP calls. Use a managed session to load +files once, or attach to a Python process that already owns NajaEDA designs. +The existing YAML tools still run each comparison in an independent worker. + +## Managed sessions + +Call these MCP tools in order: + +1. `open_session(allowed_output_dir="/absolute/path/results")` creates and selects + a session. Keep its returned `session_id`. +2. `load_designs(input_paths=["/absolute/reference.v", "/absolute/candidate.v"], + liberty_files=[])` loads the pair as `reference` and `candidate`. +3. `verify_session()` checks those loaded designs. Repeated calls reuse the same + Python process and native netlists; input files are not reread. +4. `close_session()` releases the session and its process. + +`verify_session` accepts the same verification settings as the file tools: +`verification`, `solver`, `max_k`, `sec_engine`, `sec_encoding`, +`allow_boundary_mismatch`, `report_skipped_outputs`, `log_level`, and +`log_file_name`. For example, use `verification="sec", max_k=10, +sec_engine="pdr", sec_encoding="binary"` for a bounded SEC run. + +Open more sessions for independent designs. `list_sessions()` shows them; +`set_session(session_id=...)` changes the active one. Pass `session_id` directly +to loading, verification, or closing to target a session without switching. +Custom `names=["golden", "revised"]` in `load_designs` can be selected with +`verify_session(design1="golden", design2="revised")`. + +## Attach to existing NajaEDA designs + +Run the bridge in the Python process that owns your designs, using the same +installed Kepler Formal and NajaEDA packages as the MCP server: + +```python +from pathlib import Path +from kepler_formal_mcp.session_bridge import SessionBridge + +# reference and candidate are your already loaded NajaEDA design objects. +# Raw SNLDesigns, NajaEDA Instances, and KF NativeDesign handles are accepted. +with SessionBridge(output_dir=Path("results").resolve()) as bridge: + bridge.register_design("reference", reference) + bridge.register_design("candidate", candidate) + print("Attach MCP using:", bridge.connection_file) + input("Keep this process running; press Enter when finished. ") +``` + +Then call `attach_session(connection_file="/path/printed/by/the/bridge")` from +MCP, followed by `verify_session()`. Registering captures the selected design: +subsequent NajaEDA top-design changes do not retarget the registered handle. +Verification runs inside the owning Python process on its existing pointers. +Only authenticated local commands and results cross the connection; netlists +are not serialized or copied between processes. + +The connection file contains a credential. Keep it private and attach only to +a bridge you intend to control. The bridge offers registered-design operations, +not arbitrary Python execution. + +Coordinate edits in the owning Python process with the bridge's lock: + +```python +with bridge.lock: + # Edit reference or candidate with NajaEDA here. + ... +``` + +The next verification sees those edits. Do not destroy registered designs or +their universe while the bridge uses them. `close_session()` detaches MCP from +an attached session; it leaves the bridge and caller's netlists alive. +`bridge.close()` stops the bridge and releases its borrowed references without +destroying the caller's universe. + +## Timeouts and outputs + +Both session types restrict logs to their selected output directory. Managed +sessions can produce skipped-output reports in their own working directory. +Attached sessions reject `report_skipped_outputs=True`, because native reports +would otherwise write into the owning application's working directory. + +A managed-session timeout terminates its process and invalidates that session; +open and load a new one to retry. An attached-session timeout stops waiting but +does not kill the owning application. Verification may still be running there; +overlapping operations are rejected until it finishes. diff --git a/kepler_formal_mcp/runner.py b/kepler_formal_mcp/runner.py index 4580e66..f55cacb 100644 --- a/kepler_formal_mcp/runner.py +++ b/kepler_formal_mcp/runner.py @@ -59,7 +59,8 @@ def _invoke_worker(request: dict, timeout_seconds: int, yaml_path: Path | None = try: completed = subprocess.run( [sys.executable, "-m", "kepler_formal_mcp.worker", str(request_path), str(result_path)], - cwd=work, capture_output=True, text=True, encoding="utf-8", errors="replace", + cwd=work, stdin=subprocess.DEVNULL, + capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=timeout_seconds, check=False, ) except subprocess.TimeoutExpired as error: diff --git a/kepler_formal_mcp/server.py b/kepler_formal_mcp/server.py index d598091..91b4997 100644 --- a/kepler_formal_mcp/server.py +++ b/kepler_formal_mcp/server.py @@ -12,12 +12,19 @@ from . import config, runner from .options import LogLevel, Mode, SecEncoding, SecEngine, Solver +from .tool_dispatch import threaded_tool +from . import session_tools +from .session_tools import ( + attach_session, close_session, list_sessions, load_designs, + open_session, set_session, verify_session, +) app = FastMCP("kepler-formal") +session_tools.register(app) -@app.tool() +@threaded_tool(app) def get_kepler_formal_info() -> str: """Report the installed versions, build revision, and supported Python API. @@ -26,13 +33,14 @@ def get_kepler_formal_info() -> str: handles are process-local Python objects and cannot be passed through MCP. """ try: - result = runner.get_info() + result = (session_tools.manager.call({"operation": "info"}, timeout_seconds=30) + if session_tools.manager.active_session_id is not None else runner.get_info()) except (OSError, ValueError) as error: result = runner.error_result(str(error)) return json.dumps(result, indent=2) -@app.tool() +@threaded_tool(app) def run_kepler_formal_yaml( yaml_file: str, timeout_seconds: int = 600, @@ -58,7 +66,7 @@ def run_kepler_formal_yaml( return json.dumps(result, indent=2) -@app.tool() +@threaded_tool(app) def create_yaml_and_run_kepler_formal( input_paths: list[str], liberty_files: list[str], @@ -120,4 +128,7 @@ def create_yaml_and_run_kepler_formal( def main() -> None: logging.basicConfig(level=logging.INFO, stream=sys.stderr, format="[kepler-mcp] [%(levelname)s] %(message)s") - app.run() + try: + app.run() + finally: + session_tools.manager.close_all() diff --git a/kepler_formal_mcp/session_backend.py b/kepler_formal_mcp/session_backend.py new file mode 100644 index 0000000..63c24cb --- /dev/null +++ b/kepler_formal_mcp/session_backend.py @@ -0,0 +1,188 @@ +"""Persistent design handles hosted beside their native NajaEDA universe.""" + +from __future__ import annotations + +import os +from pathlib import Path +from threading import RLock +from typing import Any +from uuid import uuid4 + +from .runner import REPORT_FILENAMES, read_reports, tail +from .verification import OPTION_NAMES, load_designs, verify_loaded + + +# Naja's universe and verification caches are shared throughout an interpreter, +# including when a caller starts more than one bridge in that interpreter. +_NATIVE_LOCK = RLock() + + +class DesignSession: + """Host managed netlists or borrow live caller designs without taking ownership. + + Attached callers must hold ``lock`` while changing or destroying registered + designs, so edits cannot race a request received by the bridge. + """ + + def __init__(self, output_dir: str | Path, owned: bool = False): + from najaeda import naja + + self.lock = _NATIVE_LOCK + self.session_id = str(uuid4()) + self.kind = "managed" if owned else "attached" + self.pid = os.getpid() + self.output_dir = Path(output_dir).expanduser().resolve() + self.output_dir.mkdir(parents=True, exist_ok=True) + self._owned = owned + self._closed = False + self._designs: dict[str, Any] = {} + self._databases: list[Any] = [] + self._input_files: set[Path] = set() + self._workspace = Path.cwd().resolve() + self._universe = None + if owned: + if naja.NLUniverse.get() is not None: + raise RuntimeError("A managed session requires a fresh Naja universe") + self._universe = naja.NLUniverse.create() + + def _require_open(self): + if self._closed: + raise RuntimeError("The design session is closed") + if self._owned: + from najaeda import naja + + if naja.NLUniverse.get() is not self._universe: + raise ReferenceError("The managed session's universe is no longer active") + + def _name(self, name: Any) -> str: + if not isinstance(name, str) or not name.strip(): + raise ValueError("Design names must be nonempty strings") + return name + + def _metadata(self, name: str) -> dict[str, Any]: + result = {"name": name} + try: + result.update(design_name=self._designs[name].najaeda_design.getName(), valid=True) + except (RuntimeError, ReferenceError) as error: + result.update(valid=False, reason=str(error)) + return result + + def _inspect(self) -> dict[str, Any]: + return { + "status": "success", "session_id": self.session_id, "kind": self.kind, + "pid": self.pid, "output_dir": str(self.output_dir), + "designs": [self._metadata(name) for name in self._designs], + } + + def register_design(self, name: str, design: Any) -> dict[str, Any]: + """Capture an existing raw design or NajaEDA Instance without copying it.""" + from kepler_formal import NativeDesign, from_najaeda + + with self.lock: + self._require_open() + name = self._name(name) + if name in self._designs: + raise ValueError(f"A design is already registered as {name!r}") + handle = design if isinstance(design, NativeDesign) else from_najaeda(design) + self._designs[name] = handle + return {"status": "success", "session_id": self.session_id, **self._metadata(name)} + + def _load(self, request: dict[str, Any]) -> dict[str, Any]: + if not self._owned: + raise ValueError("Attached sessions borrow caller designs; load them in NajaEDA and register them") + names = request.get("names", ["reference", "candidate"]) + if not isinstance(names, list) or len(names) != 2: + raise ValueError("names must contain exactly two design names") + names = [self._name(name) for name in names] + if names[0] == names[1] or any(name in self._designs for name in names): + raise ValueError("Design names must be distinct and not already registered") + inputs, libraries = request.get("input_paths"), request.get("liberty_files", []) + databases, designs = load_designs(self._universe, inputs, libraries) + try: + from kepler_formal import from_najaeda + + handles = [from_najaeda(design) for design in designs] + except Exception: + for database in reversed(databases): + database.destroy() + raise + self._designs.update(zip(names, handles)) + self._databases.extend(databases) + self._input_files.update(Path(path).resolve() for path in inputs + libraries) + return {**self._inspect(), "loaded": names} + + def _verify(self, request: dict[str, Any]) -> dict[str, Any]: + names = [self._name(request.get(key)) for key in ("design1", "design2")] + for name in names: + if name not in self._designs: + raise ValueError(f"Unknown registered design: {name!r}") + raw_options = request.get("options", {}) + if not isinstance(raw_options, dict): + raise ValueError("options must be a mapping") + unknown = set(raw_options) - set(OPTION_NAMES) + if unknown: + raise ValueError("Unsupported verification options: " + ", ".join(sorted(map(str, unknown)))) + options = dict(raw_options) + reports_requested = options.get("report_skipped_outputs", False) + if not isinstance(reports_requested, bool): + raise ValueError("report_skipped_outputs must be a boolean") + if reports_requested and not self._owned: + raise ValueError("report_skipped_outputs is unavailable for attached sessions because it writes in the caller's working directory") + + log_value = options.get("log_file") + if log_value is not None and (not isinstance(log_value, str) or not log_value.strip()): + raise ValueError("log_file must be a nonempty path string or null") + log_path = Path(log_value).expanduser() if log_value is not None else Path(f"verify-{uuid4().hex}.log") + log_path = (log_path if log_path.is_absolute() else self.output_dir / log_path).resolve() + if not log_path.is_relative_to(self.output_dir): + raise ValueError(f"Log path is outside the session output directory: {log_path}") + if log_path in self._input_files or log_path.is_dir(): + raise ValueError("The log file must not overwrite an input file or directory") + options["log_file"] = str(log_path) + if self._owned: + if Path.cwd().resolve() != self._workspace: + raise RuntimeError("The managed session's private working directory changed") + for name in REPORT_FILENAMES: + (self._workspace / name).unlink(missing_ok=True) + log_path.parent.mkdir(parents=True, exist_ok=True) + result = verify_loaded(self._designs[names[0]], self._designs[names[1]], options) + return { + "status": "error" if result["status"] == "error" else "success", + "session_id": self.session_id, "exit_code": result["exit_code"], + "verdict": result["status"], "verification_result": result, + "generated_log_file": str(log_path), + "log_tail": tail(log_path.read_text(encoding="utf-8", errors="replace")) if log_path.is_file() else "", + "reports": read_reports(self._workspace) if self._owned else {}, + "stdout_tail": "", "stderr_tail": "", + } + + def dispatch(self, request: dict[str, Any]) -> dict[str, Any]: + with self.lock: + self._require_open() + if not isinstance(request, dict): + raise ValueError("Session requests must be mappings") + operation = request.get("operation") + if operation == "info": + from .capabilities import get_capabilities + + return {**get_capabilities(), "session_id": self.session_id, "kind": self.kind, "pid": self.pid} + if operation == "inspect": + return self._inspect() + if operation == "load": + return self._load(request) + if operation == "verify": + return self._verify(request) + raise ValueError(f"Unknown session operation: {operation!r}") + + def close(self) -> dict[str, Any]: + with self.lock: + if not self._closed: + self._designs.clear() + self._databases.clear() + if self._owned: + from najaeda import naja + + if naja.NLUniverse.get() is self._universe: + self._universe.destroy() + self._closed = True + return {"status": "success", "session_id": self.session_id, "closed": True} diff --git a/kepler_formal_mcp/session_bridge.py b/kepler_formal_mcp/session_bridge.py new file mode 100644 index 0000000..5ddfb62 --- /dev/null +++ b/kepler_formal_mcp/session_bridge.py @@ -0,0 +1,229 @@ +"""Authenticated loopback access to designs in the current Python process. + +Native pointers never leave this process. Use ``bridge.lock`` while editing the +registered designs so edits and verification cannot access Naja concurrently. +Closing the bridge waits for active work and leaves caller-owned netlists alive. +""" + +from __future__ import annotations + +import hmac +import json +import os +from pathlib import Path +import secrets +import socketserver +import tempfile +import threading +from typing import Any + + +PROTOCOL = "kepler-formal-mcp-session-v1" +MAX_REQUEST_BYTES = 1024 * 1024 +SOCKET_TIMEOUT_SECONDS = 10 + + +def _error(reason: str) -> dict[str, Any]: + return {"status": "error", "reason": reason} + + +class _SessionServer(socketserver.ThreadingTCPServer): + daemon_threads = True + allow_reuse_address = False + + +class _SessionHandler(socketserver.StreamRequestHandler): + def handle(self) -> None: + self.connection.settimeout(SOCKET_TIMEOUT_SECONDS) + bridge = self.server.bridge + try: + line = self.rfile.readline(MAX_REQUEST_BYTES + 1) + if len(line) > MAX_REQUEST_BYTES: + raise ValueError("Session request exceeds the size limit") + if not line.endswith(b"\n"): + raise ValueError("Session requests must be newline-terminated JSON") + envelope = json.loads(line) + if not isinstance(envelope, dict): + raise ValueError("Session request must be a JSON object") + token = envelope.get("token") + if not isinstance(token, str) or not hmac.compare_digest( + token.encode("utf-8"), bridge._token.encode("ascii") + ): + response = _error("Session authentication failed") + elif self.client_address[0] != "127.0.0.1": + response = _error("Only loopback session connections are allowed") + else: + response = bridge._dispatch(envelope.get("request")) + except Exception as error: + response = _error(f"{type(error).__name__}: {error}") + try: + self.wfile.write(json.dumps(response).encode("utf-8") + b"\n") + except OSError: + # A client can disconnect or time out without cancelling native + # work or changing the lifetime of the caller's Python process. + pass + + +class SessionBridge: + """Expose a local design session through an authenticated TCP descriptor. + + ``owned=False`` borrows designs from the caller's Naja universe. Closing this + bridge only releases its own references. ``owned=True`` is reserved for the + managed worker, which creates and owns its universe. + + The connection file contains a secret token: share its path only with trusted + local clients. A default file is created in a private temporary directory. + """ + + def __init__(self, output_dir: str | Path, + connection_file: str | Path | None = None, *, owned: bool = False): + # Importing this module is safe in the MCP server; native imports happen + # only when a bridge is instantiated in the process owning the designs. + from .session_backend import DesignSession + + self._backend = DesignSession(output_dir, owned=owned) + self.lock = self._backend.lock + self._lifecycle_lock = threading.RLock() + self._closing = threading.Event() + self._closed = False + self._server: _SessionServer | None = None + self._thread: threading.Thread | None = None + self._descriptor: dict[str, Any] | None = None + self._descriptor_identity: tuple[int, int] | None = None + self._token = secrets.token_hex(32) + self._temporary: tempfile.TemporaryDirectory | None = None + if connection_file is None: + self._temporary = tempfile.TemporaryDirectory(prefix="kepler-session-") + self.connection_file = Path(self._temporary.name) / "connection.json" + else: + requested = Path(connection_file).expanduser().absolute() + self.connection_file = requested.parent.resolve() / requested.name + + @property + def descriptor(self) -> dict[str, Any]: + if self._descriptor is None or self._closed: + raise RuntimeError("Session bridge is not running") + return dict(self._descriptor) + + @property + def session_id(self) -> str: + return self._backend.session_id + + def register_design(self, name: str, design: Any) -> Any: + with self.lock: + if self._closing.is_set(): + raise RuntimeError("Session bridge is closed") + return self._backend.register_design(name, design) + + def _dispatch(self, request: Any) -> dict[str, Any]: + if not isinstance(request, dict): + return _error("Session request payload must be a JSON object") + if request.get("operation") not in {"inspect", "load", "verify", "info"}: + return _error("Session operation must be inspect, load, verify, or info") + if self._closing.is_set(): + return _error("Session bridge is closing") + if not self.lock.acquire(blocking=False): + return _error("Session is busy with verification or caller edits; try again when idle") + try: + if self._closing.is_set(): + return _error("Session bridge is closing") + return self._backend.dispatch(request) + finally: + self.lock.release() + + def _write_descriptor(self) -> None: + path = self.connection_file + path.parent.mkdir(parents=True, exist_ok=True) + if os.path.lexists(path): + raise FileExistsError(f"Session connection file already exists: {path}") + fd, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) + temporary = Path(temporary_name) + try: + with os.fdopen(fd, "w", encoding="utf-8") as stream: + if hasattr(os, "fchmod"): + os.fchmod(stream.fileno(), 0o600) + else: + os.chmod(temporary, 0o600) + json.dump(self._descriptor, stream) + stream.write("\n") + stream.flush() + os.fsync(stream.fileno()) + os.replace(temporary, path) + stat = path.stat() + self._descriptor_identity = (stat.st_dev, stat.st_ino) + finally: + temporary.unlink(missing_ok=True) + + def start(self) -> SessionBridge: + with self._lifecycle_lock: + if self._closed: + raise RuntimeError("Session bridge is closed") + if self._server is not None: + return self + server = _SessionServer(("127.0.0.1", 0), _SessionHandler) + server.bridge = self + self._descriptor = { + "protocol": PROTOCOL, + "host": "127.0.0.1", + "port": server.server_address[1], + "token": self._token, + "session_id": self.session_id, + } + thread = threading.Thread( + target=server.serve_forever, kwargs={"poll_interval": 0.05}, + name=f"kepler-session-{self.session_id}", daemon=True, + ) + thread.start() + try: + self._write_descriptor() + except Exception: + server.shutdown() + server.server_close() + thread.join() + self._descriptor = None + raise + self._server, self._thread = server, thread + return self + + def close(self) -> None: + """Stop serving and wait for active verification before releasing state. + + An attached bridge never terminates or resets the caller's universe. A + native call already in progress is allowed to complete, even if its MCP + client has disconnected or timed out. + """ + with self._lifecycle_lock: + if self._closed: + return + self._closing.set() + if self._server is not None: + self._server.shutdown() + self._server.server_close() + self._thread.join() + try: + with self.lock: + self._backend.close() + finally: + self._closed = True + self._server = None + self._thread = None + self._descriptor = None + if self._descriptor_identity is not None: + try: + stat = self.connection_file.lstat() + if (stat.st_dev, stat.st_ino) == self._descriptor_identity: + self.connection_file.unlink() + except FileNotFoundError: + pass + if self._temporary is not None: + self._temporary.cleanup() + + def __enter__(self) -> SessionBridge: + try: + return self.start() + except Exception: + self.close() + raise + + def __exit__(self, exc_type, exc_value, traceback) -> None: + self.close() diff --git a/kepler_formal_mcp/session_manager.py b/kepler_formal_mcp/session_manager.py new file mode 100644 index 0000000..942d131 --- /dev/null +++ b/kepler_formal_mcp/session_manager.py @@ -0,0 +1,255 @@ +"""Manage persistent workers and authenticated connections to caller interpreters.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +import json +from pathlib import Path +import socket +import subprocess +import sys +import tempfile +from threading import Lock, RLock +import time +from typing import Any +from uuid import UUID + +from .runner import error_result, tail + + +PROTOCOL = "kepler-formal-mcp-session-v1" + + +def _descriptor(path: Path) -> dict: + value = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(value, dict) or value.get("protocol") != PROTOCOL: + raise ValueError("Not a Kepler session connection file") + if value.get("host") != "127.0.0.1": + raise ValueError("Session bridges must use the IPv4 loopback address") + if type(value.get("port")) is not int or not 0 < value["port"] < 65536: + raise ValueError("Invalid session bridge port") + token = value.get("token") + if not isinstance(token, str) or len(token) < 32: + raise ValueError("Invalid session bridge token") + if not isinstance(value.get("session_id"), str): + raise ValueError("Invalid session ID") + UUID(value["session_id"]) + return value + + +def _request(descriptor: dict, request: dict, timeout_seconds: int) -> dict: + """Send only JSON to a local, authenticated bridge; never deserialize Python objects.""" + if type(timeout_seconds) is not int or timeout_seconds <= 0: + raise ValueError("timeout_seconds must be a positive integer") + deadline = time.monotonic() + timeout_seconds + with socket.create_connection((descriptor["host"], descriptor["port"]), timeout_seconds) as connection: + connection.sendall(json.dumps({"token": descriptor["token"], "request": request}).encode() + b"\n") + chunks = bytearray() + while b"\n" not in chunks: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError("Session request timed out") + connection.settimeout(remaining) + chunk = connection.recv(65536) + if not chunk: + raise ConnectionError("Session bridge closed before returning a result") + chunks.extend(chunk) + if len(chunks) > 64 * 1024 * 1024: + raise ValueError("Session response exceeds 64 MiB") + result = json.loads(chunks.split(b"\n", 1)[0]) + if not isinstance(result, dict): + raise ValueError("Invalid session response") + return result + + +@dataclass +class _Session: + descriptor: dict = field(repr=False) + metadata: dict + process: subprocess.Popen | None = None + workspace: Any = None + streams: list = field(default_factory=list, repr=False) + call_lock: Any = field(default_factory=Lock, repr=False) + + +class SessionManager: + def __init__(self): + self._sessions: dict[str, _Session] = {} + self._active: str | None = None + self._lock = RLock() + + @property + def active_session_id(self) -> str | None: + with self._lock: + return self._active + + def _remember(self, entry: _Session) -> dict: + session_id = entry.metadata["session_id"] + with self._lock: + # Re-attaching the same bridge selects the existing connection. + if session_id not in self._sessions: + self._sessions[session_id] = entry + self._active = session_id + return {**entry.metadata, "active": True} + + def open(self, output_dir: Path) -> dict: + output_dir = output_dir.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + workspace = tempfile.TemporaryDirectory(prefix="kepler-session-") + root = Path(workspace.name) + streams = [(root / name).open("wb") for name in ("stdout.txt", "stderr.txt")] + entry = _Session({}, {}, workspace=workspace, streams=streams) + try: + connection_file = root / "connection.json" + entry.process = subprocess.Popen( + [sys.executable, "-m", "kepler_formal_mcp.session_worker", + "--connection-file", str(connection_file), "--output-dir", str(output_dir)], + cwd=root, stdin=subprocess.PIPE, stdout=streams[0], stderr=streams[1], + ) + deadline = time.monotonic() + 30 + while not connection_file.is_file(): + if entry.process.poll() is not None: + details = (root / "stderr.txt").read_text(encoding="utf-8", errors="replace") + raise RuntimeError("Session worker failed to start: " + tail(details)) + if time.monotonic() >= deadline: + raise TimeoutError("Session worker startup timed out after 30 seconds") + time.sleep(0.02) + entry.descriptor = _descriptor(connection_file) + entry.metadata = self._inspect(entry.descriptor) + if entry.metadata["kind"] != "managed": + raise ValueError("Worker did not create a managed session") + return self._remember(entry) + except Exception: + self._dispose(entry) + raise + + def _inspect(self, descriptor: dict) -> dict: + metadata = _request(descriptor, {"operation": "inspect"}, 30) + if metadata.get("status") != "success": + raise ValueError(metadata.get("reason", metadata.get("stderr_tail", "Session bridge rejected attachment"))) + if (metadata.get("session_id") != descriptor["session_id"] + or metadata.get("kind") not in {"managed", "attached"} + or not isinstance(metadata.get("output_dir"), str)): + raise ValueError("Bridge identity does not match its connection file") + return metadata + + def attach(self, connection_file: Path) -> dict: + descriptor = _descriptor(connection_file) + metadata = self._inspect(descriptor) + # An attached connection never acquires process/universe ownership, + # even if the external bridge happens to manage its own universe. + metadata["kind"] = "attached" + return self._remember(_Session(descriptor, metadata)) + + def _get(self, session_id: str | None) -> tuple[str, _Session]: + with self._lock: + selected = session_id if session_id is not None else self._active + if selected not in self._sessions: + raise ValueError(f"Unknown or closed session: {selected!r}") + return selected, self._sessions[selected] + + def select(self, session_id: str) -> dict: + selected, entry = self._get(session_id) + response = self.call({"operation": "inspect"}, selected, 30) + if response["status"] == "success": + with self._lock: + if self._sessions.get(selected) is not entry: + raise ValueError(f"Session closed while being selected: {selected!r}") + self._active = selected + response["active"] = True + return response + + def list(self) -> dict: + with self._lock: + return {"status": "success", "active_session_id": self._active, + "sessions": [dict(entry.metadata) for entry in self._sessions.values()]} + + def call(self, request: dict, session_id: str | None = None, timeout_seconds: int = 600) -> dict: + if type(timeout_seconds) is not int or timeout_seconds <= 0: + raise ValueError("timeout_seconds must be a positive integer") + selected, entry = self._get(session_id) + if not entry.call_lock.acquire(blocking=False): + return {**error_result("Session is busy"), "session_id": selected, "busy": True} + try: + # A close can win the race between lookup and acquiring this lock. + self._get(selected) + offsets = [] + if entry.workspace is not None: + for name in ("stdout.txt", "stderr.txt"): + offsets.append((Path(entry.workspace.name) / name).stat().st_size) + try: + response = _request(entry.descriptor, request, timeout_seconds) + except TimeoutError: + managed = entry.process is not None + response = error_result(f"Session request timed out after {timeout_seconds} seconds") + response.update(session_id=selected, session_invalidated=managed, + may_still_be_running=not managed) + if managed: + self._forget(selected) + self._dispose(entry, force=True) + return response + except (OSError, ValueError) as error: + response = error_result(f"Session connection failed: {error}") + response.update(session_id=selected, session_invalidated=True) + self._forget(selected) + self._dispose(entry, force=True) + return response + if response.get("status") == "success" and "designs" in response: + entry.metadata.update(response) + entry.metadata["kind"] = "managed" if entry.process is not None else "attached" + response.update(session_id=selected, pid=entry.metadata["pid"]) + if entry.workspace is not None: + for name, key, offset in zip(("stdout.txt", "stderr.txt"), ("stdout_tail", "stderr_tail"), offsets): + with (Path(entry.workspace.name) / name).open("rb") as stream: + stream.seek(offset) + captured = tail(stream.read()) + if captured: + response[key] = captured + return response + finally: + entry.call_lock.release() + + def _forget(self, selected: str): + with self._lock: + self._sessions.pop(selected, None) + if self._active == selected: + self._active = None + + def _dispose(self, entry: _Session, force: bool = False): + process = entry.process + if process is not None: + if process.stdin is not None: + process.stdin.close() + if force and process.poll() is None: + process.terminate() + try: + process.wait(timeout=3) + except subprocess.TimeoutExpired: + process.kill() + process.wait(timeout=3) + for stream in entry.streams: + stream.close() + if entry.workspace is not None: + entry.workspace.cleanup() + + def close(self, session_id: str | None = None) -> dict: + selected, entry = self._get(session_id) + if not entry.call_lock.acquire(blocking=False): + return {**error_result("Session is busy; wait for its request to finish"), + "session_id": selected, "busy": True} + try: + self._forget(selected) + self._dispose(entry) + return {"status": "success", "session_id": selected, + "state": "closed" if entry.process is not None else "detached"} + finally: + entry.call_lock.release() + + def close_all(self): + with self._lock: + entries = list(self._sessions.items()) + self._sessions.clear() + self._active = None + for _, entry in entries: + # Called on server shutdown: terminate owned workers, never callers. + self._dispose(entry, force=True) diff --git a/kepler_formal_mcp/session_tools.py b/kepler_formal_mcp/session_tools.py new file mode 100644 index 0000000..9a531e1 --- /dev/null +++ b/kepler_formal_mcp/session_tools.py @@ -0,0 +1,112 @@ +"""MCP session tools; the native runtime stays in a worker or caller interpreter.""" + +from __future__ import annotations + +import atexit +from functools import wraps +import json +from pathlib import Path + +from . import config +from .options import LogLevel, Mode, SecEncoding, SecEngine, Solver +from .runner import error_result +from .session_manager import SessionManager +from .tool_dispatch import threaded_tool + + +manager = SessionManager() +atexit.register(manager.close_all) + + +def _json_result(function): + @wraps(function) + def invoke(*args, **kwargs): + try: + result = function(*args, **kwargs) + except (OSError, ValueError, RuntimeError) as error: + result = error_result(str(error)) + return json.dumps(result, indent=2) + return invoke + + +@_json_result +def open_session(allowed_output_dir: str | None = None) -> str: + """Start and select a persistent Python worker with its own NajaEDA universe.""" + return manager.open(config.output_root(allowed_output_dir)) + + +@_json_result +def attach_session(connection_file: str) -> str: + """Bind to SessionBridge running in an existing Python interpreter on this machine. + + The caller must register its designs in that bridge. Closing this MCP + session only detaches; it never deletes the caller's designs or interpreter. + """ + return manager.attach(config.resolve_path(connection_file, Path.cwd())) + + +@_json_result +def set_session(session_id: str) -> str: + """Select the default session used by subsequent load_designs/verify_session calls.""" + return manager.select(session_id) + + +@_json_result +def list_sessions() -> str: + """List this MCP server's sessions and the currently selected session ID.""" + return manager.list() + + +@_json_result +def load_designs(input_paths: list[str], liberty_files: list[str] | None = None, + names: list[str] | None = None, session_id: str | None = None, + timeout_seconds: int = 600) -> str: + """Load two Verilog designs once into a managed session for repeated verification. + + Names default to reference/candidate and must not already exist. Inputs + are relative to the MCP launch directory. Attached interpreters load and + register their own designs through SessionBridge.register_design instead. + """ + request = {"operation": "load", + "input_paths": [str(config.resolve_path(path, Path.cwd())) for path in input_paths], + "liberty_files": [str(config.resolve_path(path, Path.cwd())) for path in (liberty_files or [])]} + if names is not None: + request["names"] = names + return manager.call(request, session_id, timeout_seconds) + + +@_json_result +def verify_session(design1: str = "reference", design2: str = "candidate", + session_id: str | None = None, verification: Mode = "lec", + solver: Solver = "kissat", max_k: int | None = None, + sec_engine: SecEngine | None = None, sec_encoding: SecEncoding | None = None, + allow_boundary_mismatch: bool = False, report_skipped_outputs: bool = False, + log_file_name: str | None = None, log_level: LogLevel | None = "info", + timeout_seconds: int = 600) -> str: + """Verify registered designs in place without reopening Python or rereading files. + + Log paths are relative to the session's fixed output directory. Managed + timeouts terminate/invalidate that session. Attached timeouts leave the + caller alive and verification may still be running; retries can report busy. + Attached sessions do not support report_skipped_outputs because those + native reports write into the caller's process-wide working directory. + """ + return manager.call({"operation": "verify", "design1": design1, "design2": design2, + "options": {"mode": verification, "solver": solver, "max_k": max_k, + "sec_engine": sec_engine, "sec_encoding": sec_encoding, + "allow_boundary_mismatch": allow_boundary_mismatch, + "report_skipped_outputs": report_skipped_outputs, + "log_file": log_file_name, "log_level": log_level}}, + session_id, timeout_seconds) + + +@_json_result +def close_session(session_id: str | None = None) -> str: + """Close a managed worker, or detach from a caller-owned Python interpreter.""" + return manager.close(session_id) + + +def register(app): + for function in (open_session, attach_session, set_session, list_sessions, + load_designs, verify_session, close_session): + threaded_tool(app)(function) diff --git a/kepler_formal_mcp/session_worker.py b/kepler_formal_mcp/session_worker.py new file mode 100644 index 0000000..ac7d8b5 --- /dev/null +++ b/kepler_formal_mcp/session_worker.py @@ -0,0 +1,26 @@ +"""Own a persistent Naja session until the parent closes the stdin pipe.""" + +from __future__ import annotations + +import argparse +from pathlib import Path +import sys + +from .session_bridge import SessionBridge + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--connection-file", required=True, type=Path) + parser.add_argument("--output-dir", required=True, type=Path) + arguments = parser.parse_args(argv) + with SessionBridge(arguments.output_dir, arguments.connection_file, owned=True): + # Native stdout/stderr belong to this managed process and are captured + # by its parent. The TCP response channel carries only JSON messages. + while sys.stdin.buffer.read(8192): + pass + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/kepler_formal_mcp/tool_dispatch.py b/kepler_formal_mcp/tool_dispatch.py new file mode 100644 index 0000000..1cb6e18 --- /dev/null +++ b/kepler_formal_mcp/tool_dispatch.py @@ -0,0 +1,21 @@ +"""Keep blocking verification and process I/O off the MCP event loop.""" + +from functools import partial, wraps +from inspect import signature + +import anyio + + +def threaded_tool(app): + def register(function): + @wraps(function) + async def dispatch(**arguments): + return await anyio.to_thread.run_sync(partial(function, **arguments)) + + # Resolve annotations in the original module before handing this wrapper + # to FastMCP, so Literal choices and defaults remain in the tool schema. + dispatch.__signature__ = signature(function, eval_str=True) + app.tool()(dispatch) + return function + + return register diff --git a/kepler_formal_mcp/verification.py b/kepler_formal_mcp/verification.py new file mode 100644 index 0000000..3cf994c --- /dev/null +++ b/kepler_formal_mcp/verification.py @@ -0,0 +1,64 @@ +"""Shared NajaEDA loading and native verification for isolated or live sessions.""" + +from __future__ import annotations + +from dataclasses import asdict +from pathlib import Path +from typing import Any + + +OPTION_NAMES = ( + "mode", "solver", "max_k", "sec_engine", "sec_encoding", + "allow_boundary_mismatch", "report_skipped_outputs", "log_file", "log_level", +) + + +def load_designs(universe: Any, input_paths: list[str], liberty_files: list[str]): + """Load two independent databases, rolling back only these databases on error.""" + from najaeda import naja + + if not isinstance(input_paths, list) or len(input_paths) != 2: + raise ValueError("input_paths must contain exactly two Verilog files") + if not isinstance(liberty_files, list): + raise ValueError("liberty_files must be a list of paths") + for path in input_paths + liberty_files: + if not isinstance(path, str) or not Path(path).is_absolute() or not Path(path).is_file(): + raise ValueError(f"Input must be an existing absolute file path: {path!r}") + + databases, designs = [], [] + try: + for path in input_paths: + database = naja.NLDB.create(universe) + databases.append(database) + if liberty_files: + database.loadLibertyPrimitives(liberty_files) + database.loadVerilog([path]) + design = database.getTopDesign() + if design is None: + raise RuntimeError(f"NajaEDA did not find a top design in {path}") + designs.append(design) + return databases, designs + except Exception: + for database in reversed(databases): + database.destroy() + raise + + +def verify_loaded(design1: Any, design2: Any, options: dict[str, Any]) -> dict[str, Any]: + """Borrow existing designs and return only owning, JSON-compatible result data.""" + from kepler_formal import VerificationOptions, verify_designs + + if not isinstance(options, dict): + raise ValueError("options must be a mapping") + unknown = set(options) - set(OPTION_NAMES) + if unknown: + raise ValueError("Unsupported verification options: " + ", ".join(sorted(map(str, unknown)))) + result = verify_designs(design1, design2, options=VerificationOptions(**options)) + serialized = asdict(result) + serialized.update( + status=result.status.value, + equivalent=result.equivalent, + conclusive=result.conclusive, + coverage_percent=result.coverage_percent, + ) + return serialized diff --git a/kepler_formal_mcp/worker.py b/kepler_formal_mcp/worker.py index 6125c81..bf3819e 100644 --- a/kepler_formal_mcp/worker.py +++ b/kepler_formal_mcp/worker.py @@ -7,7 +7,6 @@ from __future__ import annotations import argparse -from dataclasses import asdict import json from pathlib import Path import sys @@ -19,49 +18,18 @@ def verify(request: dict[str, Any]) -> dict[str, Any]: """Load the requested designs in NajaEDA and borrow them for verification.""" # Import NajaEDA first: it arranges the shared native runtime used by KF. from najaeda import naja - from kepler_formal import VerificationOptions, verify_designs + from .verification import OPTION_NAMES, load_designs, verify_loaded if naja.NLUniverse.get() is not None: raise RuntimeError("The verification worker requires a fresh Naja universe") universe = naja.NLUniverse.create() try: - designs = [] - for path in request["input_paths"]: - # Independent databases allow both inputs to have the same module - # names, while retaining their pointers in one shared runtime. - database = naja.NLDB.create(universe) - if request["liberty_files"]: - database.loadLibertyPrimitives(request["liberty_files"]) - database.loadVerilog([path]) - design = database.getTopDesign() - if design is None: - raise RuntimeError(f"NajaEDA did not find a top design in {path}") - designs.append(design) - - result = verify_designs( - designs[0], - designs[1], - options=VerificationOptions( - mode=request["mode"], - solver=request["solver"], - max_k=request.get("max_k"), - sec_engine=request.get("sec_engine"), - sec_encoding=request.get("sec_encoding"), - allow_boundary_mismatch=request["allow_boundary_mismatch"], - report_skipped_outputs=request["report_skipped_outputs"], - log_file=request["log_file"], - log_level=request["log_level"], - ), - ) - serialized = asdict(result) - serialized.update( - status=result.status.value, - equivalent=result.equivalent, - conclusive=result.conclusive, - coverage_percent=result.coverage_percent, + _databases, designs = load_designs(universe, request["input_paths"], request["liberty_files"]) + return verify_loaded( + designs[0], designs[1], + {name: request[name] for name in OPTION_NAMES if name in request}, ) - return serialized finally: # This universe belongs solely to this short-lived worker. The MCP # server and any callers' in-memory netlists are never reset. diff --git a/tests/test_server.py b/tests/test_server.py index f900a4f..a4953f6 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -295,6 +295,33 @@ async def exercise(): # A second protocol operation also works after native solver logging. self.assertEqual(names, {tool.name for tool in (await session.list_tools()).tools}) + async def call(name, arguments): + response = await session.call_tool(name, arguments) + self.assertFalse(response.isError, response) + result = json.loads(next(item.text for item in response.content if item.type == "text")) + self.assertEqual(result["status"], "success", result) + return result + + opened = await call("open_session", {"allowed_output_dir": str(self.outputs)}) + try: + await call("load_designs", { + "input_paths": [str(self.reference), str(self.candidate)], + "timeout_seconds": 30, + }) + self.reference.unlink() + self.candidate.unlink() + for mode in ("lec", "sec"): + result = await call("verify_session", { + "verification": mode, "solver": "glucose", "timeout_seconds": 30, + }) + self.assert_verdict(result, "different") + self.assertEqual(result["pid"], opened["pid"]) + information = await call("get_kepler_formal_info", {}) + self.assertEqual(information["pid"], opened["pid"]) + self.assertEqual(information["session_id"], opened["session_id"]) + finally: + await call("close_session", {"session_id": opened["session_id"]}) + asyncio.run(exercise()) diff --git a/tests/test_sessions.py b/tests/test_sessions.py new file mode 100644 index 0000000..db50b8d --- /dev/null +++ b/tests/test_sessions.py @@ -0,0 +1,295 @@ +"""Verify persistent workers and bridges to caller-owned NajaEDA designs.""" + +from __future__ import annotations + +import asyncio +import json +import os +from pathlib import Path +import tempfile +import unittest +from unittest.mock import patch + +from kepler_formal_mcp import server +from kepler_formal_mcp.session_bridge import SessionBridge + + +PASS_THROUGH = "module top(input a, output y); assign y = a; endmodule\n" +CONSTANT_OUTPUT = "module top(input a, output y); assign y = 1'b0; endmodule\n" + + +class SessionFixture(unittest.TestCase): + def setUp(self): + directory = tempfile.TemporaryDirectory(prefix="kepler sessions ") + self.addCleanup(directory.cleanup) + self.root = Path(directory.name).resolve() + self.reference = self.root / "reference.v" + self.candidate = self.root / "candidate.v" + self.reference.write_text(PASS_THROUGH, encoding="utf-8") + self.candidate.write_text(CONSTANT_OUTPUT, encoding="utf-8") + + def call(self, tool, **arguments): + return json.loads(getattr(server, tool)(**arguments)) + + def open(self, name="results"): + session = self.call("open_session", allowed_output_dir=str(self.root / name)) + self.assert_success(session) + self.addCleanup(server.close_session, session_id=session["session_id"]) + return session + + def attach(self, bridge): + session = self.call("attach_session", connection_file=str(bridge.connection_file)) + self.assert_success(session) + self.addCleanup(server.close_session, session_id=session["session_id"]) + return session + + def load(self, session_id=None, **arguments): + return self.call( + "load_designs", + input_paths=[str(self.reference), str(self.candidate)], + liberty_files=[], + session_id=session_id, + timeout_seconds=30, + **arguments, + ) + + def verify(self, session_id=None, **arguments): + return self.call("verify_session", session_id=session_id, timeout_seconds=30, **arguments) + + def assert_success(self, result): + self.assertEqual(result["status"], "success", result) + + def assert_verdict(self, result, expected): + self.assert_success(result) + self.assertEqual(result["verdict"], expected, result) + self.assertEqual(result["verification_result"]["status"], expected, result) + + +class ManagedSessionTest(SessionFixture): + def test_loaded_designs_survive_source_deletion_and_repeated_verification(self): + session = self.open() + self.assertEqual(session["kind"], "managed") + self.assertNotEqual(session["pid"], os.getpid()) + self.assert_success(self.load(session["session_id"])) + self.reference.unlink() + self.candidate.unlink() + for options in ({}, {"solver": "cadical"}, { + "verification": "sec", "sec_engine": "pdr", "sec_encoding": "binary", "max_k": 2, + }): + with self.subTest(options=options): + result = self.verify(session["session_id"], **options) + self.assert_verdict(result, "different") + self.assertEqual(result["pid"], session["pid"]) + self.assertEqual(result["session_id"], session["session_id"]) + + def test_multiple_sessions_are_independent_and_can_be_selected(self): + different = self.open("different") + self.assert_success(self.load()) + self.candidate.write_text(PASS_THROUGH, encoding="utf-8") + equivalent = self.open("equivalent") + self.assert_success(self.load(names=["golden", "revised"])) + self.assertNotEqual(different["pid"], equivalent["pid"]) + self.assertNotEqual(different["session_id"], equivalent["session_id"]) + self.assert_verdict(self.verify(design1="golden", design2="revised"), "equivalent") + self.assert_success(self.call("set_session", session_id=different["session_id"])) + self.assert_verdict(self.verify(), "different") + self.assert_verdict(self.verify( + equivalent["session_id"], design1="golden", design2="revised", + ), "equivalent") + listing = self.call("list_sessions") + self.assertEqual(listing["active_session_id"], different["session_id"]) + self.assertTrue( + {different["session_id"], equivalent["session_id"]} + <= {item["session_id"] for item in listing["sessions"]} + ) + + def test_closed_and_unknown_sessions_are_rejected(self): + session = self.open() + closed = self.call("close_session", session_id=session["session_id"]) + self.assert_success(closed) + self.assertEqual(closed["state"], "closed") + listing = self.call("list_sessions") + self.assertNotIn(session["session_id"], {item["session_id"] for item in listing["sessions"]}) + for session_id in (session["session_id"], "not-a-session"): + with self.subTest(session_id=session_id): + self.assertEqual(self.verify(session_id)["status"], "error") + self.assertEqual(self.call("set_session", session_id=session_id)["status"], "error") + + def test_invalid_load_does_not_discard_previously_loaded_designs(self): + session = self.open() + self.assert_success(self.load()) + invalid = self.call("load_designs", input_paths=[str(self.reference)], session_id=session["session_id"]) + self.assertEqual(invalid["status"], "error", invalid) + self.assert_verdict(self.verify(), "different") + missing = self.verify(design1="not-registered") + self.assertEqual(missing["status"], "error", missing) + self.assert_verdict(self.verify(), "different") + + def test_select_cannot_reactivate_a_session_closed_during_inspection(self): + session = self.open() + manager = server.session_tools.manager + original_call = manager.call + + def close_after_inspection(*args, **kwargs): + response = original_call(*args, **kwargs) + manager.close(session["session_id"]) + return response + + with patch.object(manager, "call", side_effect=close_after_inspection): + result = self.call("set_session", session_id=session["session_id"]) + self.assertEqual(result["status"], "error", result) + listing = self.call("list_sessions") + self.assertIsNone(listing["active_session_id"]) + self.assertNotIn(session["session_id"], {item["session_id"] for item in listing["sessions"]}) + + def test_session_logs_stay_in_selected_directory(self): + self.open() + self.assert_success(self.load()) + rejected = self.verify(log_file_name=str(self.root / "outside.log")) + self.assertEqual(rejected["status"], "error", rejected) + self.assertFalse((self.root / "outside.log").exists()) + result = self.verify(log_file_name="inside.log") + self.assert_verdict(result, "different") + self.assertEqual(Path(result["generated_log_file"]), self.root / "results" / "inside.log") + + def test_timed_out_managed_session_is_invalidated(self): + session = self.open() + self.assert_success(self.load()) + with patch("kepler_formal_mcp.session_manager._request", side_effect=TimeoutError("timed out")): + result = self.verify() + self.assertEqual(result["status"], "error", result) + self.assertTrue(result["session_invalidated"], result) + listing = self.call("list_sessions") + self.assertNotIn(session["session_id"], {item["session_id"] for item in listing["sessions"]}) + self.assertEqual(self.verify(session["session_id"])["status"], "error") + + def test_session_tools_are_exposed_to_mcp_clients(self): + names = {tool.name for tool in asyncio.run(server.app.list_tools())} + self.assertTrue({ + "open_session", "set_session", "list_sessions", "load_designs", + "verify_session", "close_session", "attach_session", + } <= names) + + +class AttachedSessionTest(SessionFixture): + def setUp(self): + super().setUp() + from najaeda import naja + + self.assertIsNone(naja.NLUniverse.get(), "These tests require a fresh caller-owned universe") + self.universe = naja.NLUniverse.create() + self.addCleanup(self.universe.destroy) + self.designs = [] + for path in (self.reference, self.candidate): + database = naja.NLDB.create(self.universe) + database.loadVerilog([str(path)]) + self.designs.append(database.getTopDesign()) + + def bridge(self): + bridge = SessionBridge(output_dir=self.root / "results") + self.addCleanup(bridge.close) + return bridge.start() + + def test_attached_designs_are_live_and_detaching_preserves_caller_ownership(self): + from kepler_formal import VerificationOptions, from_najaeda, verify_designs + from najaeda import naja, netlist + + reference, candidate = self.designs + bridge = self.bridge() + self.universe.setTopDesign(reference) + bridge.register_design("reference", netlist.get_top()) + bridge.register_design("raw_reference", reference) + self.universe.setTopDesign(candidate) + bridge.register_design("candidate", from_najaeda(candidate)) + session = self.attach(bridge) + self.assertEqual(session["kind"], "attached") + self.assertEqual(session["pid"], os.getpid()) + self.reference.unlink() + self.candidate.unlink() + self.assert_verdict(self.verify(), "different") + self.assert_verdict(self.verify(design1="raw_reference"), "different") + + with bridge.lock: + candidate.getScalarTerm("y").setNet(candidate.getScalarTerm("a").getNet()) + result = self.verify() + self.assert_verdict(result, "equivalent") + self.assertEqual(result["pid"], os.getpid()) + detached = self.call("close_session", session_id=session["session_id"]) + self.assert_success(detached) + self.assertEqual(detached["state"], "detached") + self.assertIsNotNone(naja.NLUniverse.get()) + self.assertEqual(candidate.getName(), "top") + # Detaching did not shut down the caller's bridge: it can be reattached. + reattached = self.attach(bridge) + self.assert_verdict(self.verify(reattached["session_id"]), "equivalent") + self.call("close_session", session_id=reattached["session_id"]) + bridge.close() + self.assertIsNotNone(naja.NLUniverse.get()) + native = verify_designs(reference, candidate, options=VerificationOptions( + log_file=self.root / "still-owned.log", + )) + self.assertEqual(native.status.value, "equivalent") + + def test_attached_sessions_reject_process_relative_native_reports(self): + bridge = self.bridge() + bridge.register_design("reference", self.designs[0]) + bridge.register_design("candidate", self.designs[1]) + self.attach(bridge) + original_cwd = Path.cwd() + rejected = self.verify(report_skipped_outputs=True) + self.assertEqual(rejected["status"], "error", rejected) + self.assertEqual(Path.cwd(), original_cwd) + self.assert_verdict(self.verify(), "different") + + def test_attached_timeout_does_not_kill_or_invalidate_callers_process(self): + bridge = self.bridge() + bridge.register_design("reference", self.designs[0]) + bridge.register_design("candidate", self.designs[1]) + session = self.attach(bridge) + with patch("kepler_formal_mcp.session_manager._request", side_effect=TimeoutError("timed out")): + result = self.verify() + self.assertEqual(result["status"], "error", result) + self.assertFalse(result["session_invalidated"], result) + self.assertTrue(result["may_still_be_running"], result) + listing = self.call("list_sessions") + self.assertIn(session["session_id"], {item["session_id"] for item in listing["sessions"]}) + self.assert_verdict(self.verify(session["session_id"]), "different") + self.assertEqual(self.designs[0].getName(), "top") + + def test_wrong_token_cannot_attach_to_a_live_bridge(self): + bridge = self.bridge() + bridge.register_design("reference", self.designs[0]) + bridge.register_design("candidate", self.designs[1]) + descriptor = json.loads(Path(bridge.connection_file).read_text(encoding="utf-8")) + descriptor["token"] = "0" * 64 if descriptor["token"] != "0" * 64 else "1" * 64 + invalid = self.root / "wrong-token.json" + invalid.write_text(json.dumps(descriptor), encoding="utf-8") + invalid.chmod(0o600) + rejected = self.call("attach_session", connection_file=str(invalid)) + self.assertEqual(rejected["status"], "error", rejected) + self.assertIn("auth", json.dumps(rejected).lower()) + self.attach(bridge) + self.assert_verdict(self.verify(), "different") + + def test_caller_edit_lock_prevents_overlapping_verification(self): + bridge = self.bridge() + bridge.register_design("reference", self.designs[0]) + bridge.register_design("candidate", self.designs[1]) + self.attach(bridge) + with bridge.lock: + busy = self.verify() + self.assertEqual(busy["status"], "error", busy) + self.assertIn("busy", json.dumps(busy).lower()) + self.assert_verdict(self.verify(), "different") + + def test_bad_connection_descriptor_is_a_structured_error(self): + descriptor = self.root / "invalid.json" + for text in ("not json", "{}", "[]"): + with self.subTest(text=text): + descriptor.write_text(text, encoding="utf-8") + result = self.call("attach_session", connection_file=str(descriptor)) + self.assertEqual(result["status"], "error", result) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_tool_dispatch.py b/tests/test_tool_dispatch.py new file mode 100644 index 0000000..30ca3bd --- /dev/null +++ b/tests/test_tool_dispatch.py @@ -0,0 +1,46 @@ +"""Keep MCP transport responsive while synchronous workers are busy.""" + +import asyncio +import subprocess +import threading +import unittest +from unittest.mock import patch + +from kepler_formal_mcp import runner, server + + +class ToolDispatchTest(unittest.TestCase): + def test_info_worker_does_not_inherit_the_mcp_input_pipe(self): + with patch.object(runner.subprocess, "run", return_value=subprocess.CompletedProcess( + args=[], returncode=1, stdout="", stderr="test worker stopped" + )) as launch: + runner.get_info() + self.assertEqual(launch.call_args.kwargs["stdin"], subprocess.DEVNULL) + + def test_mcp_event_loop_keeps_running_while_a_tool_waits(self): + started, release = threading.Event(), threading.Event() + + def blocking_info(): + started.set() + if not release.wait(3): + raise RuntimeError("MCP event loop did not unblock the worker") + return {"status": "success"} + + async def exercise(): + task = asyncio.create_task(server.app.call_tool("get_kepler_formal_info", {})) + try: + self.assertTrue(await asyncio.to_thread(started.wait, 2)) + self.assertFalse(task.done()) + release.set() + await task + finally: + release.set() + await asyncio.gather(task, return_exceptions=True) + + with patch.object(server.session_tools.manager, "_active", None), \ + patch.object(runner, "get_info", side_effect=blocking_info): + asyncio.run(exercise()) + + +if __name__ == "__main__": + unittest.main()