"""Long-lived Pi RPC client for gateway services. Extracted and generalised from the memo-inbox gateway (``pi-workspaces/memo-inbox/telegram-gateway/gateway.py``, class ``PiRPC``), which has run this pattern in production since 2026-07. Compared with the original it adds the four gaps that both existing scenarios shared: * an explicit minimal ``env`` so provider and backend API keys never reach the node process (``docs/gateway-patterns.md`` pattern 8); * ``start_new_session`` plus process-group termination so a timeout cannot leave orphaned node children burning provider quota (pattern 7); * the loading-isolation flags from ``docs/isolation-baseline.md`` are part of the launch contract rather than something each caller remembers; * a per-turn deadline enforced with RPC ``abort`` instead of killing the process. It deliberately keeps the original's proven mechanics: strict ``\\n``-only JSONL framing, request/response correlation by ``id``, ``agent_settled`` as the terminal event, and deterministic receipts harvested from ``tool_execution_end`` rather than from model prose (pattern 3). Stdlib only. Requires Python 3.11+ (``asyncio.timeout``). Protocol reference: ``docs/pi-runtime-notes.md`` §15. """ from __future__ import annotations import asyncio import contextlib import json import logging import os import signal import time from dataclasses import dataclass, field, replace from pathlib import Path from typing import Any, Awaitable, Callable, Iterable, Sequence LOG = logging.getLogger("pi_rpc") __all__ = [ "PiLaunchConfig", "PiRpcClient", "PiRpcError", "PiRpcRejected", "PiTurnAborted", "TurnResult", "Usage", ] # Environment variables a pi subprocess legitimately needs. Everything else -- # in particular every ``*_API_KEY`` and ``*_TOKEN`` belonging to the host # service -- is withheld. The provider credential is read by pi itself from # ~/.pi/agent/models.json and must not be passed through the environment. DEFAULT_ENV_ALLOWLIST: tuple[str, ...] = ( "PATH", "HOME", "LANG", "LC_ALL", "LC_CTYPE", "TZ", "TERM", "NODE_OPTIONS", "SSL_CERT_FILE", "SSL_CERT_DIR", "NO_PROXY", "no_proxy", ) class PiRpcError(RuntimeError): """Transport or lifecycle failure.""" class PiRpcRejected(PiRpcError): """Pi returned ``success: false`` for a command.""" class PiTurnAborted(PiRpcError): """The turn exceeded its deadline and was aborted.""" @dataclass(frozen=True) class Usage: """Token and cost accounting for one turn. Populated from ``message_update.usage`` deltas. See ``docs/pi-runtime-notes.md`` §16 -- observability needs no extra instrumentation because the stream already carries this. """ input: int = 0 output: int = 0 cache_read: int = 0 cache_write: int = 0 total_tokens: int = 0 cost_total: float = 0.0 def merged(self, raw: dict[str, Any]) -> "Usage": cost = raw.get("cost") or {} return Usage( input=self.input + int(raw.get("input") or 0), output=self.output + int(raw.get("output") or 0), cache_read=self.cache_read + int(raw.get("cacheRead") or 0), cache_write=self.cache_write + int(raw.get("cacheWrite") or 0), total_tokens=self.total_tokens + int(raw.get("totalTokens") or 0), cost_total=self.cost_total + float(cost.get("total") or 0.0), ) @dataclass class ToolCallRecord: """One completed tool execution.""" tool_name: str args: dict[str, Any] text: str details: Any is_error: bool @dataclass class TurnResult: """Everything one ``prompt`` produced. ``receipts`` holds the text of tool results whose tool is named in ``PiLaunchConfig.receipt_tools``. Callers must render user-visible outcomes of state changes from ``receipts`` (or from ``tool_calls``), never from ``replies`` -- see ``docs/gateway-patterns.md`` pattern 3 for the production failure this prevents. """ replies: list[str] = field(default_factory=list) receipts: list[str] = field(default_factory=list) tool_calls: list[ToolCallRecord] = field(default_factory=list) usage: Usage = field(default_factory=Usage) latency_seconds: float = 0.0 model: str = "" thinking: str = "" aborted: bool = False extension_errors: list[dict[str, Any]] = field(default_factory=list) @property def text(self) -> str: """The final assistant message, or an empty string.""" return self.replies[-1] if self.replies else "" @dataclass(frozen=True) class PiLaunchConfig: """Launch contract for a dedicated Pi agent. The isolation defaults implement ``docs/isolation-baseline.md``. Overriding them widens what the agent loads or can call, so each override should be justified in the scenario's ``profile.toml``. """ pi_bin: str workspace: Path session_dir: Path provider: str model: str thinking: str = "medium" display_name: str = "Pi Agent" # --- session identity ------------------------------------------------- # A stable prefix; the client appends a rotation counter so that history # stays auditable instead of being summarised away. Empty means --no-session. session_id_prefix: str = "" # --- personality ------------------------------------------------------ # system_prompt REPLACES pi's coding-assistant prompt. Note that pi then # omits the tool list and guidelines entirely: the customPrompt branch of # dist/core/system-prompt.js returns before they are assembled, so # promptSnippet and promptGuidelines never reach the model and the prompt # file must enumerate the tools itself. # See docs/pi-runtime-notes.md section 22. system_prompt: Path | None = None append_system_prompt: Path | None = None # --- layer 1: loading isolation -------------------------------------- extensions: tuple[Path, ...] = () skills: tuple[Path, ...] = () no_extensions: bool = True no_skills: bool = True no_prompt_templates: bool = True no_themes: bool = True no_context_files: bool = False approve: bool = True # --- layer 3: capability --------------------------------------------- # Prefer no_builtin_tools over an explicit ``tools`` allowlist: --tools is a # registry-level filter, which prevents an extension from registering tools # at runtime. See docs/pi-runtime-notes.md section 4. no_builtin_tools: bool = True tools: tuple[str, ...] = () exclude_tools: tuple[str, ...] = () # --- budgets ---------------------------------------------------------- turn_deadline_seconds: float = 180.0 startup_timeout_seconds: float = 60.0 stop_timeout_seconds: float = 10.0 # --- rotation --------------------------------------------------------- rotate_after_prompts: int = 24 rotate_after_messages: int = 60 # --- receipts --------------------------------------------------------- receipt_tools: frozenset[str] = frozenset() # Set by a scenario whose extension registers a tool named ``read``. Only # affects whether loading skills is warned about; see _read_reachable. extension_registers_read: bool = False # --- environment ------------------------------------------------------ env_allowlist: tuple[str, ...] = DEFAULT_ENV_ALLOWLIST extra_env: tuple[tuple[str, str], ...] = () def __post_init__(self) -> None: if not self.no_builtin_tools and not self.tools: LOG.warning( "PiLaunchConfig for %r enables built-in tools with no allowlist; " "bash/edit/write will be active", self.display_name, ) if self.skills and not self._read_reachable(): LOG.warning( "PiLaunchConfig for %r loads %d skill(s) but no tool named 'read' " "will be active, so pi omits the skills section entirely and the " "skills have no effect whatsoever. Either set " "extension_registers_read=True if the extension provides one, or " "put the content in the system prompt instead " "(see docs/pi-runtime-notes.md sections 1 and 23)", self.display_name, len(self.skills), ) def _read_reachable(self) -> bool: """Whether a tool named ``read`` will actually be active. Pi emits the skills section only when ``read`` is active. This used to assume that ``no_builtin_tools`` implied an extension supplying a restricted ``read`` override, which is how memo-inbox is built -- but an agent whose tools are all domain-specific has no ``read`` at all, and then every ``--skill`` argument is silently discarded. Measured: with tools [query_library, lookup_online, counts] the prompt contained no skills section and no skill names, with and without --system-prompt. So this no longer guesses. An extension that registers ``read`` must say so. """ if self.tools: return "read" in self.tools if not self.no_builtin_tools: return True # the built-in read is active return self.extension_registers_read def build_env(self) -> dict[str, str]: env = {k: os.environ[k] for k in self.env_allowlist if k in os.environ} env.setdefault("HOME", str(Path.home())) env.update(dict(self.extra_env)) return env def build_args(self, session_id: str | None) -> list[str]: args: list[str] = [self.pi_bin, "--mode", "rpc"] if session_id: args += ["--session-id", session_id, "--session-dir", str(self.session_dir)] else: args += ["--no-session"] if self.no_builtin_tools: args += ["--no-builtin-tools"] if self.tools: args += ["--tools", ",".join(self.tools)] if self.exclude_tools: args += ["--exclude-tools", ",".join(self.exclude_tools)] if self.no_extensions: args += ["--no-extensions"] for path in self.extensions: args += ["-e", str(path)] if self.no_skills: args += ["--no-skills"] for path in self.skills: args += ["--skill", str(path)] if self.no_prompt_templates: args += ["--no-prompt-templates"] if self.no_themes: args += ["--no-themes"] if self.no_context_files: args += ["--no-context-files"] if self.approve: args += ["--approve"] if self.system_prompt: args += ["--system-prompt", str(self.system_prompt)] if self.append_system_prompt: args += ["--append-system-prompt", str(self.append_system_prompt)] args += [ "--provider", self.provider, "--model", self.model, "--thinking", self.thinking, "--name", self.display_name, ] return args class PiRpcClient: """A single long-lived ``pi --mode rpc`` process. One client owns one conversation. Instantiate one per Telegram chat, per ticket, or per whatever your unit of continuity is; do not multiplex unrelated conversations through one client, because they would share a session and contaminate each other. Not safe for concurrent ``prompt`` calls -- an internal lock serialises commands, so callers queue rather than interleave. """ def __init__( self, config: PiLaunchConfig, *, on_event: Callable[[dict[str, Any]], Awaitable[None] | None] | None = None, ) -> None: self.config = config self._on_event = on_event self._process: asyncio.subprocess.Process | None = None self._events: asyncio.Queue[dict[str, Any]] = asyncio.Queue() self._lock = asyncio.Lock() self._stdout_task: asyncio.Task[None] | None = None self._stderr_task: asyncio.Task[None] | None = None self._stderr_tail: list[str] = [] self._prompt_count = 0 self._rotation = 0 # -- lifecycle --------------------------------------------------------- @property def running(self) -> bool: return self._process is not None and self._process.returncode is None @property def pid(self) -> int | None: return self._process.pid if self._process else None def _session_id(self) -> str | None: if not self.config.session_id_prefix: return None return f"{self.config.session_id_prefix}-{self._rotation:04d}" async def start(self) -> None: if self.running: return cfg = self.config cfg.session_dir.mkdir(parents=True, exist_ok=True) args = cfg.build_args(self._session_id()) LOG.info("Starting pi: %s", " ".join(args)) self._process = await asyncio.create_subprocess_exec( *args, cwd=str(cfg.workspace), env=cfg.build_env(), stdin=asyncio.subprocess.PIPE, stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE, # Own process group so a stuck node tree can be killed wholesale. start_new_session=True, ) self._stdout_task = asyncio.create_task(self._read_stdout(), name="pi-rpc-stdout") self._stderr_task = asyncio.create_task(self._read_stderr(), name="pi-rpc-stderr") try: async with asyncio.timeout(cfg.startup_timeout_seconds): await self._command({"type": "get_state"}, wait_for_settled=False, use_lock=False) except TimeoutError as exc: tail = " | ".join(self._stderr_tail[-5:]) await self.stop() raise PiRpcError(f"pi did not answer get_state within " f"{cfg.startup_timeout_seconds}s: {tail}") from exc LOG.info("pi RPC ready pid=%s session=%s", self.pid, self._session_id()) async def stop(self) -> None: process, self._process = self._process, None for task in (self._stdout_task, self._stderr_task): if task: task.cancel() for task in (self._stdout_task, self._stderr_task): if task: with contextlib.suppress(asyncio.CancelledError, Exception): await task self._stdout_task = self._stderr_task = None if process is None or process.returncode is not None: return # Terminate the whole group: pi is a node CLI and may have children. with contextlib.suppress(ProcessLookupError, PermissionError): os.killpg(os.getpgid(process.pid), signal.SIGTERM) try: async with asyncio.timeout(self.config.stop_timeout_seconds): await process.wait() except TimeoutError: LOG.warning("pi pid=%s ignored SIGTERM, sending SIGKILL to group", process.pid) with contextlib.suppress(ProcessLookupError, PermissionError): os.killpg(os.getpgid(process.pid), signal.SIGKILL) with contextlib.suppress(Exception): await process.wait() async def __aenter__(self) -> "PiRpcClient": await self.start() return self async def __aexit__(self, *_exc: object) -> None: await self.stop() # -- stream plumbing --------------------------------------------------- async def _read_stdout(self) -> None: assert self._process and self._process.stdout buffer = b"" # Strict JSONL: split on \n only. A generic line reader is not # protocol-compliant because U+2028/U+2029 are legal inside JSON # strings. See docs/pi-runtime-notes.md section 15. while chunk := await self._process.stdout.read(65536): buffer += chunk while b"\n" in buffer: line, buffer = buffer.split(b"\n", 1) line = line.rstrip(b"\r") if not line: continue try: event = json.loads(line) except json.JSONDecodeError: LOG.warning("Discarding malformed pi RPC record (%d bytes)", len(line)) continue await self._events.put(event) if self._on_event is not None: try: result = self._on_event(event) if asyncio.iscoroutine(result): await result except Exception: LOG.exception("on_event callback failed") # EOF: the child died or closed its pipe. Without a sentinel, _consume # would sit on _events.get() until the turn deadline -- a process that # exited at second 3 would be reported as "no reply" at second 180, and a # fallback that launches a fresh pi would hide the death entirely. Put a # sentinel so the waiter fails immediately instead. LOG.warning("pi RPC stdout closed pid=%s rc=%s", self.pid, self._process.returncode) await self._events.put({"type": "_process_exited"}) async def _read_stderr(self) -> None: assert self._process and self._process.stderr while line := await self._process.stderr.readline(): text = line.decode(errors="replace").rstrip() if not text: continue # Kept even on success: the memo/curator scenarios both discarded # stderr and then could not explain model failures afterwards. self._stderr_tail.append(text) del self._stderr_tail[:-50] LOG.info("pi stderr: %s", text) async def _send(self, payload: dict[str, Any]) -> None: await self.start() assert self._process and self._process.stdin blob = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode() self._process.stdin.write(blob + b"\n") await self._process.stdin.drain() # -- commands ---------------------------------------------------------- async def command(self, payload: dict[str, Any], *, wait_for_settled: bool = False, deadline: float | None = None) -> tuple[dict[str, Any], TurnResult]: async with self._lock: return await self._command(payload, wait_for_settled=wait_for_settled, use_lock=False, deadline=deadline) async def _command( self, payload: dict[str, Any], *, wait_for_settled: bool, use_lock: bool = True, deadline: float | None = None, ) -> tuple[dict[str, Any], TurnResult]: if use_lock: async with self._lock: return await self._command(payload, wait_for_settled=wait_for_settled, use_lock=False, deadline=deadline) request_id = payload.setdefault("id", f"req-{time.monotonic_ns():x}") # Drop anything left over from an aborted turn so correlation is clean. while not self._events.empty(): self._events.get_nowait() started = time.monotonic() result = TurnResult(model=self.config.model, thinking=self.config.thinking) await self._send(payload) budget = deadline if deadline is not None else self.config.turn_deadline_seconds response: dict[str, Any] | None = None try: async with asyncio.timeout(budget): response = await self._consume(request_id, result, wait_for_settled) except TimeoutError: result.aborted = True result.latency_seconds = time.monotonic() - started LOG.warning("pi turn exceeded %.0fs; aborting", budget) with contextlib.suppress(Exception): await self._send({"id": f"{request_id}-abort", "type": "abort"}) raise PiTurnAborted(f"pi turn exceeded {budget:.0f}s") from None result.latency_seconds = time.monotonic() - started assert response is not None return response, result async def _consume(self, request_id: str, result: TurnResult, wait_for_settled: bool) -> dict[str, Any]: """Collect events until the command is complete. ``success: true`` only means *accepted*; post-acceptance failures arrive as events. So for a prompt we wait for ``agent_settled`` as well. """ response: dict[str, Any] | None = None pending_args: dict[str, dict[str, Any]] = {} while True: event = await self._events.get() kind = event.get("type") if kind == "_process_exited": rc = self._process.returncode if self._process else None raise PiRpcError( f"pi process exited (pid={self.pid}, returncode={rc}) while " f"awaiting {request_id}; stderr tail: {' | '.join(self._stderr_tail[-3:])}" ) if kind == "response" and event.get("id") == request_id: response = event if not wait_for_settled or not event.get("success"): return response elif kind == "message_update": raw = event.get("usage") if isinstance(raw, dict): result.usage = result.usage.merged(raw) elif kind == "message_end": text = _assistant_text(event.get("message") or {}) if text: result.replies.append(text) elif kind == "tool_execution_start": call_id = str(event.get("toolCallId") or "") if call_id: pending_args[call_id] = event.get("args") or {} elif kind == "tool_execution_end": name = str(event.get("toolName") or "") payload = event.get("result") or {} text = _result_text(payload) record = ToolCallRecord( tool_name=name, args=pending_args.pop(str(event.get("toolCallId") or ""), {}), text=text, details=payload.get("details"), is_error=bool(event.get("isError")), ) result.tool_calls.append(record) # Deterministic receipts, not model prose. Pattern 3. if name in self.config.receipt_tools and text and not record.is_error: result.receipts.append(text) elif kind == "extension_error": LOG.error("pi extension error: %s", event) result.extension_errors.append(event) elif kind == "agent_settled" and response is not None: return response # -- high level -------------------------------------------------------- async def prompt(self, message: str, *, deadline: float | None = None, images: Sequence[dict[str, Any]] | None = None) -> TurnResult: """Send one user turn and return everything it produced. Rotates the session first when either rotation bound is reached: pi has no session TTL and auto-compaction does not fire on large-context models (``docs/pi-runtime-notes.md`` sections 17-18). """ await self._maybe_rotate() payload: dict[str, Any] = {"type": "prompt", "message": message} if images: payload["images"] = list(images) response, result = await self.command(payload, wait_for_settled=True, deadline=deadline) if not response.get("success"): raise PiRpcRejected(str(response.get("error") or "pi rejected the prompt")) self._prompt_count += 1 return result async def _maybe_rotate(self) -> None: cfg = self.config if not cfg.session_id_prefix: return if cfg.rotate_after_prompts and self._prompt_count >= cfg.rotate_after_prompts: LOG.info("Rotating pi session: prompt_count=%d", self._prompt_count) await self.rotate_session() return if cfg.rotate_after_messages: state = await self.get_state() count = int(state.get("messageCount") or 0) if count >= cfg.rotate_after_messages: LOG.info("Rotating pi session: message_count=%d", count) await self.rotate_session() async def rotate_session(self) -> None: """Start a fresh session, preserving the previous file for audit. Restarting the process with the next ``--session-id`` is preferred over the ``new_session`` command because the session id then encodes the rotation, so history remains greppable on disk. """ self._rotation += 1 self._prompt_count = 0 await self.stop() await self.start() async def get_state(self) -> dict[str, Any]: response, _ = await self.command({"type": "get_state"}) return response.get("data") or {} async def get_session_stats(self) -> dict[str, Any]: response, _ = await self.command({"type": "get_session_stats"}) return response.get("data") or {} async def set_model(self, model: str) -> None: response, _ = await self.command({"type": "set_model", "model": model}) if not response.get("success"): raise PiRpcRejected(str(response.get("error") or f"cannot select model {model}")) self.config = replace(self.config, model=model) async def set_thinking_level(self, level: str) -> None: response, _ = await self.command({"type": "set_thinking_level", "level": level}) if not response.get("success"): raise PiRpcRejected(str(response.get("error") or f"cannot set thinking {level}")) self.config = replace(self.config, thinking=level) async def abort(self) -> None: await self.command({"type": "abort"}) async def compact(self, instructions: str | None = None) -> None: payload: dict[str, Any] = {"type": "compact"} if instructions: payload["customInstructions"] = instructions await self.command(payload, wait_for_settled=False) @property def stderr_tail(self) -> list[str]: """Recent stderr lines, retained even on success for diagnostics.""" return list(self._stderr_tail) # -- helpers --------------------------------------------------------------- def _assistant_text(message: dict[str, Any]) -> str: """Concatenate the text blocks of an assistant message. Thinking blocks are excluded: they are reasoning, not output. """ content = message.get("content") if isinstance(content, str): return content.strip() if not isinstance(content, Iterable): return "" parts = [ block["text"] for block in content if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) ] return "\n".join(parts).strip() def _result_text(result: dict[str, Any]) -> str: content = result.get("content") or [] if not isinstance(content, Iterable): return "" parts = [ block["text"] for block in content if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) ] return "\n".join(parts).strip()