diff --git a/pilot_agent/agent/acceptance_blueprints.py b/pilot_agent/agent/acceptance_blueprints.py new file mode 100644 index 0000000..43201aa --- /dev/null +++ b/pilot_agent/agent/acceptance_blueprints.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class AcceptanceBlueprint: + id: str + title: str + cadence: str + purpose: str + commands: tuple[str, ...] + required_credentials: tuple[str, ...] = () + artifacts: tuple[str, ...] = () + notes: tuple[str, ...] = () + + +BLUEPRINTS: tuple[AcceptanceBlueprint, ...] = ( + AcceptanceBlueprint( + id="nightly-checks", + title="Nightly local acceptance", + cadence="nightly after main changes", + purpose="Re-run the repository quality bar and surface regressions before users hit them.", + commands=( + "UV_CACHE_DIR=.uv-cache uv sync --all-groups --frozen", + "scripts/run_tests.sh", + "pilot-agent doctor --json", + ), + artifacts=( + "test transcript", + "doctor JSON", + ), + notes=( + "This blueprint validates the existing local MVP pipeline; it does not deploy.", + "Treat a red doctor check as an acceptance failure unless explicitly waived.", + ), + ), + AcceptanceBlueprint( + id="dependency-audit", + title="Dependency drift audit", + cadence="weekly or before a release tag", + purpose="Detect lockfile drift and dependency metadata issues without adding new services.", + commands=( + "uv lock --check", + "UV_CACHE_DIR=.uv-cache uv sync --all-groups --frozen", + "UV_CACHE_DIR=.uv-cache uv run python -m pip check", + ), + artifacts=( + "uv lock check output", + "pip check output", + ), + notes=( + "Security scanners can be wired later; this v1 blueprint sticks to installed tooling.", + ), + ), + AcceptanceBlueprint( + id="deploy-verification", + title="Deploy readiness verification", + cadence="before release or deploy phase", + purpose="Confirm local package and container build paths still work before handoff.", + commands=( + "UV_CACHE_DIR=.uv-cache uv build", + "docker compose build", + "pilot-agent doctor --json", + ), + required_credentials=( + "VERCEL_TOKEN when deploy phase is enabled", + ), + artifacts=( + "dist build output", + "docker build output", + "doctor JSON", + ), + notes=( + "This is a readiness blueprint; publishing remains intentionally separate.", + ), + ), +) + +_BY_ID = {blueprint.id: blueprint for blueprint in BLUEPRINTS} + + +def list_blueprints() -> tuple[AcceptanceBlueprint, ...]: + return BLUEPRINTS + + +def get_blueprint(blueprint_id: str) -> AcceptanceBlueprint: + try: + return _BY_ID[blueprint_id] + except KeyError as exc: + known = ", ".join(sorted(_BY_ID)) + raise ValueError(f"unknown acceptance blueprint {blueprint_id!r}; known: {known}") from exc + + +def render_blueprint(blueprint: AcceptanceBlueprint) -> str: + lines = [ + f"# {blueprint.title}", + f"id: {blueprint.id}", + f"cadence: {blueprint.cadence}", + "", + blueprint.purpose, + "", + "Commands:", + *[f"- {command}" for command in blueprint.commands], + ] + if blueprint.required_credentials: + lines.extend( + [ + "", + "Required credentials:", + *[f"- {item}" for item in blueprint.required_credentials], + ] + ) + if blueprint.artifacts: + lines.extend(["", "Artifacts:", *[f"- {item}" for item in blueprint.artifacts]]) + if blueprint.notes: + lines.extend(["", "Notes:", *[f"- {item}" for item in blueprint.notes]]) + return "\n".join(lines) diff --git a/pilot_agent/agent/context.py b/pilot_agent/agent/context.py index 123b06b..fd0b6be 100644 --- a/pilot_agent/agent/context.py +++ b/pilot_agent/agent/context.py @@ -6,7 +6,15 @@ from pathlib import Path from pilot_agent.agent.safety import redact_sensitive_text -from pilot_agent.agent.types import CompletionResponse, Message, Role, ToolResult, ToolSpec +from pilot_agent.agent.types import ( + CompletionResponse, + Message, + Role, + SessionEvent, + ToolResult, + ToolSpec, + to_json, +) from pilot_agent.providers.base import Provider SUMMARY_PROMPT = """Compress the agent work history into a state document. @@ -32,6 +40,10 @@ def __init__( self.summarizer = summarizer or provider self.session_log = session_log self._ineffective_compactions = 0 + session_anchor = str(session_log.resolve()) if session_log is not None else "memory" + self._root_session_id = session_anchor + self._current_session_id = f"{session_anchor}#0" + self._compaction_depth = 0 def prepare(self, system: str, history: list[Message]) -> list[Message]: prepared = copy.deepcopy(history) @@ -67,7 +79,7 @@ def replace_provider(self, provider: Provider) -> None: def _truncate_tool_results(self, system: str, history: list[Message]) -> list[Message]: cutoff = self._last_turn_start(history, turns=5) seen_tool_outputs: dict[str, str] = {} - for message in reversed(history[:cutoff]): + for message in history[:cutoff]: if message.role is not Role.TOOL: continue for result in message.tool_results: @@ -181,20 +193,31 @@ def _last_turn_start(history: list[Message], turns: int) -> int: def _write_compaction_event(self, before: int, after: int) -> None: if self.session_log is None: return + next_depth = self._compaction_depth + 1 + parent_session_id = self._current_session_id + current_session_id = f"{self._root_session_id}#{next_depth}" + event = SessionEvent( + event_type="compaction", + payload={ + "before_tokens": before, + "after_tokens": after, + "saved_tokens": max(0, before - after), + "ineffective_count": self._ineffective_compactions, + "provenance": { + "root_session_id": self._root_session_id, + "parent_session_id": parent_session_id, + "current_session_id": current_session_id, + "session_kind": "continuation", + "creator_kind": "compaction", + "compaction_depth": next_depth, + }, + }, + ) self.session_log.parent.mkdir(parents=True, exist_ok=True) with self.session_log.open("a", encoding="utf-8") as handle: - handle.write( - json.dumps( - { - "_type": "compaction", - "before_tokens": before, - "after_tokens": after, - "saved_tokens": max(0, before - after), - "ineffective_count": self._ineffective_compactions, - } - ) - + "\n" - ) + handle.write(to_json(event) + "\n") + self._current_session_id = current_session_id + self._compaction_depth = next_depth def build_system_prompt( diff --git a/pilot_agent/agent/hooks.py b/pilot_agent/agent/hooks.py new file mode 100644 index 0000000..959f0b5 --- /dev/null +++ b/pilot_agent/agent/hooks.py @@ -0,0 +1,287 @@ +from __future__ import annotations + +import logging +from collections.abc import Callable +from copy import deepcopy +from dataclasses import dataclass, field +from typing import Any + +logger = logging.getLogger(__name__) + +OBSERVER_SCHEMA_VERSION = "pilot.observer.v1" +MIDDLEWARE_SCHEMA_VERSION = "pilot.middleware.v1" + +VALID_HOOKS = frozenset( + { + "pre_tool_call", + "post_tool_call", + "transform_tool_result", + "pre_llm_call", + "post_llm_call", + "api_request_error", + "on_compaction", + "on_session_event", + } +) + +TOOL_REQUEST_MIDDLEWARE = "tool_request" +TOOL_EXECUTION_MIDDLEWARE = "tool_execution" +LLM_REQUEST_MIDDLEWARE = "llm_request" +LLM_EXECUTION_MIDDLEWARE = "llm_execution" +VALID_MIDDLEWARE = frozenset( + { + TOOL_REQUEST_MIDDLEWARE, + TOOL_EXECUTION_MIDDLEWARE, + LLM_REQUEST_MIDDLEWARE, + LLM_EXECUTION_MIDDLEWARE, + } +) + +HookCallback = Callable[..., Any] +ExecutionCallback = Callable[[Any], Any] + + +@dataclass +class RequestMiddlewareResult: + payload: Any + original_payload: Any + changed: bool = False + trace: list[dict[str, Any]] = field(default_factory=list) + + +def observer_payload(**kwargs: Any) -> dict[str, Any]: + kwargs.setdefault("telemetry_schema_version", OBSERVER_SCHEMA_VERSION) + return kwargs + + +def middleware_payload(**kwargs: Any) -> dict[str, Any]: + kwargs.setdefault("telemetry_schema_version", OBSERVER_SCHEMA_VERSION) + kwargs.setdefault("middleware_schema_version", MIDDLEWARE_SCHEMA_VERSION) + return kwargs + + +def _safe_copy(payload: Any) -> Any: + try: + return deepcopy(payload) + except Exception as exc: + logger.debug("deepcopy failed for middleware payload: %s", exc) + return dict(payload) if isinstance(payload, dict) else payload + + +def _trace_entry(result: dict[str, Any]) -> dict[str, Any]: + name = result.get("name") or result.get("middleware") or result.get("id") or "anonymous" + entry: dict[str, Any] = {"name": str(name)} + if result.get("changed") is not None: + entry["changed"] = bool(result.get("changed")) + if result.get("reason"): + entry["reason"] = str(result["reason"]) + return entry + + +class RuntimeHooks: + """In-process hook and middleware registry for the agent harness. + + This intentionally does not discover or import arbitrary plugins. It gives + the core loop a stable extension contract that tests and future first-party + integrations can use without adding conditionals to AgentLoop. + """ + + def __init__(self) -> None: + self._hooks: dict[str, list[HookCallback]] = {name: [] for name in VALID_HOOKS} + self._middleware: dict[str, list[HookCallback]] = { + name: [] for name in VALID_MIDDLEWARE + } + + def register_hook(self, name: str, callback: HookCallback) -> None: + if name not in VALID_HOOKS: + raise ValueError(f"unknown hook: {name}") + self._hooks[name].append(callback) + + def register_middleware(self, kind: str, callback: HookCallback) -> None: + if kind not in VALID_MIDDLEWARE: + raise ValueError(f"unknown middleware: {kind}") + self._middleware[kind].append(callback) + + def has_hook(self, name: str) -> bool: + return bool(self._hooks.get(name)) + + def has_middleware(self, kind: str) -> bool: + return bool(self._middleware.get(kind)) + + def invoke_hook(self, name: str, **kwargs: Any) -> list[Any]: + callbacks = self._hooks.get(name, []) + if not callbacks: + return [] + payload = observer_payload(**kwargs) + results: list[Any] = [] + for callback in list(callbacks): + try: + results.append(callback(**payload)) + except Exception as exc: + logger.debug("hook %s failed: %s", name, exc) + return results + + def get_tool_block_message( + self, + tool_name: str, + args: dict[str, Any], + **context: Any, + ) -> str | None: + for result in self.invoke_hook("pre_tool_call", tool_name=tool_name, args=args, **context): + if isinstance(result, str) and result: + return result + if not isinstance(result, dict): + continue + action = str(result.get("action") or result.get("decision") or "").lower() + if action in {"block", "deny"}: + message = result.get("message") or result.get("reason") or "blocked by hook" + return str(message) + return None + + def transform_tool_result( + self, + result: str, + *, + tool_name: str, + args: dict[str, Any], + **context: Any, + ) -> str: + for hook_result in self.invoke_hook( + "transform_tool_result", + tool_name=tool_name, + args=args, + result=result, + **context, + ): + if isinstance(hook_result, str): + return hook_result + return result + + def apply_tool_request_middleware( + self, + tool_name: str, + args: dict[str, Any], + **context: Any, + ) -> RequestMiddlewareResult: + return self._apply_request_middleware( + TOOL_REQUEST_MIDDLEWARE, + payload_key="args", + payload=args, + tool_name=tool_name, + **context, + ) + + def apply_llm_request_middleware( + self, + request: dict[str, Any], + **context: Any, + ) -> RequestMiddlewareResult: + return self._apply_request_middleware( + LLM_REQUEST_MIDDLEWARE, + payload_key="request", + payload=request, + **context, + ) + + def run_tool_execution_middleware( + self, + tool_name: str, + args: dict[str, Any], + next_call: ExecutionCallback, + **context: Any, + ) -> Any: + return self._run_execution_middleware( + TOOL_EXECUTION_MIDDLEWARE, + payload=args, + next_call=next_call, + tool_name=tool_name, + **context, + ) + + def run_llm_execution_middleware( + self, + request: dict[str, Any], + next_call: ExecutionCallback, + **context: Any, + ) -> Any: + return self._run_execution_middleware( + LLM_EXECUTION_MIDDLEWARE, + payload=request, + next_call=next_call, + **context, + ) + + def _apply_request_middleware( + self, + kind: str, + *, + payload_key: str, + payload: Any, + **context: Any, + ) -> RequestMiddlewareResult: + callbacks = self._middleware.get(kind, []) + if not callbacks: + return RequestMiddlewareResult(payload=payload, original_payload=payload) + + original = _safe_copy(payload) + current = _safe_copy(original) + trace: list[dict[str, Any]] = [] + for callback in list(callbacks): + try: + result = callback( + **middleware_payload( + **{payload_key: current, f"original_{payload_key}": original}, + **context, + ) + ) + except Exception as exc: + logger.debug("request middleware %s failed: %s", kind, exc) + continue + if not isinstance(result, dict): + continue + next_payload = result.get(payload_key) + if next_payload is None: + continue + current = _safe_copy(next_payload) + trace.append(_trace_entry(result)) + return RequestMiddlewareResult( + payload=current, + original_payload=original, + changed=bool(trace), + trace=trace, + ) + + def _run_execution_middleware( + self, + kind: str, + *, + payload: Any, + next_call: ExecutionCallback, + **context: Any, + ) -> Any: + callbacks = list(self._middleware.get(kind, [])) + if not callbacks: + return next_call(payload) + + def call_at(index: int, current_payload: Any) -> Any: + if index >= len(callbacks): + return next_call(current_payload) + callback = callbacks[index] + next_used = False + + def downstream(next_payload: Any | None = None) -> Any: + nonlocal next_used + if next_used: + raise RuntimeError("execution middleware called next_call more than once") + next_used = True + return call_at(index + 1, current_payload if next_payload is None else next_payload) + + try: + return callback( + **middleware_payload(payload=current_payload, next_call=downstream, **context) + ) + except Exception as exc: + logger.debug("execution middleware %s failed: %s", kind, exc) + return downstream(current_payload) + + return call_at(0, payload) diff --git a/pilot_agent/agent/iteration_budget.py b/pilot_agent/agent/iteration_budget.py new file mode 100644 index 0000000..4687fe5 --- /dev/null +++ b/pilot_agent/agent/iteration_budget.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import threading +from dataclasses import dataclass + + +@dataclass(frozen=True) +class IterationBudgetSnapshot: + limit: int + consumed: int + remaining: int + + @property + def exhausted(self) -> bool: + return self.remaining <= 0 + + +class IterationBudget: + """Thread-safe consume/refund budget for agent loop iterations.""" + + def __init__(self, limit: int): + if limit < 1: + raise ValueError("iteration budget limit must be at least 1") + self._limit = limit + self._consumed = 0 + self._lock = threading.Lock() + + def consume(self, amount: int = 1) -> bool: + if amount < 1: + raise ValueError("consume amount must be at least 1") + with self._lock: + if self._consumed + amount > self._limit: + return False + self._consumed += amount + return True + + def refund(self, amount: int = 1) -> None: + if amount < 1: + raise ValueError("refund amount must be at least 1") + with self._lock: + self._consumed = max(0, self._consumed - amount) + + def snapshot(self) -> IterationBudgetSnapshot: + with self._lock: + consumed = self._consumed + return IterationBudgetSnapshot( + limit=self._limit, + consumed=consumed, + remaining=max(0, self._limit - consumed), + ) + + @property + def remaining(self) -> int: + return self.snapshot().remaining + + @property + def consumed(self) -> int: + return self.snapshot().consumed + + @property + def limit(self) -> int: + return self._limit diff --git a/pilot_agent/agent/loop.py b/pilot_agent/agent/loop.py index b13c489..a32593e 100644 --- a/pilot_agent/agent/loop.py +++ b/pilot_agent/agent/loop.py @@ -1,17 +1,26 @@ from __future__ import annotations import json +import time from collections.abc import Callable from dataclasses import dataclass from pathlib import Path -from typing import Protocol, cast +from typing import Any, Protocol, cast from pilot_agent.agent.context import ContextManager, build_system_prompt +from pilot_agent.agent.iteration_budget import IterationBudget from pilot_agent.agent.memory import Memory, SkillOutcomeRecorder from pilot_agent.agent.phases import PHASES, Phase from pilot_agent.agent.session_lock import ProjectSessionLock from pilot_agent.agent.state import read_state, write_session_record -from pilot_agent.agent.types import Message, Role, ToolCall, ToolResult +from pilot_agent.agent.types import ( + CompletionResponse, + Message, + Role, + ToolCall, + ToolResult, + ToolSpec, +) from pilot_agent.agent.usage import SessionUsage, normalize_usage from pilot_agent.cli.ui import UI from pilot_agent.providers.base import Provider @@ -69,13 +78,25 @@ def __init__( def run(self, max_turns: int = 200) -> None: with ProjectSessionLock(self.project_root): - turns = 0 + budget = IterationBudget(max_turns) interrupts = 0 - while self.phase is not None and turns < max_turns: + while self.phase is not None: + if not budget.consume(): + write_session_record( + self.project_root, + { + "_type": "iteration_budget_exhausted", + "limit": budget.limit, + "consumed": budget.consumed, + }, + ) + self.ui.warning(f"iteration budget exhausted after {budget.consumed} turns") + break try: self.run_turn() interrupts = 0 except KeyboardInterrupt: + budget.refund() interrupts += 1 if interrupts >= 2: write_session_record( @@ -88,7 +109,6 @@ def run(self, max_turns: int = 200) -> None: "interrupted - enter a new instruction or /quit; " "a second Ctrl+C exits the session" ) - turns += 1 def run_turn(self) -> None: if self.phase is None: @@ -103,9 +123,10 @@ def run_turn(self) -> None: self.memory.relevant_lessons(self.stack), ) messages = self.ctx.prepare(system, self.history) + tools = self.registry.specs(phase.tools, context_window=self.provider.context_window) with self.ui.api_spinner() as progress: progress.add_task("api", total=None) - resp = self.provider.complete(system, messages, self.registry.specs(phase.tools)) + resp = self._complete_with_hooks(system, messages, tools, phase) usage = normalize_usage(resp.usage) self.usage.add( usage, @@ -118,21 +139,7 @@ def run_turn(self) -> None: write_session_record(self.project_root, resp.message) self.ui.render(resp.message) if resp.stop_reason == "tool_use": - results: list[ToolResult] = [] - pin_tool_message = False - for call in resp.message.tool_calls: - if call.name == "complete_phase": - results.append(self.advance_phase(call)) - continue - timer = self.ui.tool_timer(call) - result = self.registry.execute(call) - self.memory.observe(call, result) - self.ui.render_tool_result(call, result, elapsed_s=timer.elapsed()) - if call.name == "load_skill" and not result.is_error: - skill_name = str(call.arguments.get("name", "")) - self.loaded_skills.setdefault(phase.name, set()).add(skill_name) - pin_tool_message = True - results.append(result) + results, pin_tool_message = self._execute_tool_calls(phase, resp.message.tool_calls) current_phase = self.phase.name if self.phase is not None else None tool_msg = Message( role=Role.TOOL, @@ -151,6 +158,109 @@ def run_turn(self) -> None: self.history.append(user_msg) write_session_record(self.project_root, user_msg) + def _complete_with_hooks( + self, + system: str, + messages: list[Message], + tools: list[ToolSpec], + phase: Phase, + ) -> CompletionResponse: + context = self._runtime_context(phase) + request: dict[str, Any] = { + "system": system, + "messages": messages, + "tools": tools, + "max_tokens": 4096, + } + self.registry.hooks.invoke_hook("pre_llm_call", **request, **context) + request_mw = self.registry.hooks.apply_llm_request_middleware(request, **context) + payload = request_mw.payload if isinstance(request_mw.payload, dict) else request + started_at = time.monotonic() + + def invoke_provider(next_request: Any) -> CompletionResponse: + typed = _coerce_llm_request(next_request, fallback=request) + return self.provider.complete( + typed["system"], + typed["messages"], + typed["tools"], + max_tokens=typed["max_tokens"], + ) + + try: + response = self.registry.hooks.run_llm_execution_middleware( + payload, + invoke_provider, + **context, + middleware_trace=request_mw.trace, + ) + except Exception as exc: + self.registry.hooks.invoke_hook( + "api_request_error", + error=str(exc), + error_type=exc.__class__.__name__, + **context, + ) + raise + if not isinstance(response, CompletionResponse): + raise TypeError("llm execution middleware must return CompletionResponse") + self.registry.hooks.invoke_hook( + "post_llm_call", + stop_reason=response.stop_reason, + usage=response.usage, + tool_call_count=len(response.message.tool_calls), + duration_ms=int((time.monotonic() - started_at) * 1000), + **context, + ) + return response + + def _execute_tool_calls( + self, + phase: Phase, + calls: list[ToolCall], + ) -> tuple[list[ToolResult], bool]: + results: list[ToolResult] = [] + pin_tool_message = False + pending: list[ToolCall] = [] + + def flush_pending() -> None: + nonlocal pin_tool_message + if not pending: + return + executions = self.registry.execute_batch( + list(pending), + context=self._runtime_context(phase), + ) + pending.clear() + for execution in executions: + call = execution.call + result = execution.result + self.memory.observe(call, result) + self.ui.render_tool_result(call, result, elapsed_s=execution.elapsed_s) + if call.name == "load_skill" and not result.is_error: + skill_name = str(call.arguments.get("name", "")) + self.loaded_skills.setdefault(phase.name, set()).add(skill_name) + pin_tool_message = True + results.append(result) + + for call in calls: + if call.name == "complete_phase": + flush_pending() + results.append(self.advance_phase(call)) + continue + pending.append(call) + flush_pending() + return results, pin_tool_message + + def _runtime_context(self, phase: Phase) -> dict[str, Any]: + provider_name = self.provider.__class__.__name__.removesuffix("Provider").lower() + return { + "phase": phase.name, + "allowed_tools": list(phase.tools), + "project_root": str(self.project_root), + "provider": provider_name, + "model": self.provider.model, + } + def handle_slash_command(self, raw: str) -> bool: command, _, arg = raw.partition(" ") if command == "/help": @@ -255,6 +365,34 @@ def _summarize_lesson(self, prompt: str) -> str: return response.message.content +def _coerce_llm_request( + request: Any, + *, + fallback: dict[str, Any], +) -> dict[str, Any]: + raw = request if isinstance(request, dict) else fallback + system = raw.get("system", fallback["system"]) + messages = raw.get("messages", fallback["messages"]) + tools = raw.get("tools", fallback["tools"]) + max_tokens = raw.get("max_tokens", fallback["max_tokens"]) + if not isinstance(system, str): + system = fallback["system"] + if not isinstance(messages, list): + messages = fallback["messages"] + if not isinstance(tools, list): + tools = fallback["tools"] + try: + typed_max_tokens = int(max_tokens) + except (TypeError, ValueError): + typed_max_tokens = int(fallback["max_tokens"]) + return { + "system": system, + "messages": cast(list[Message], messages), + "tools": cast(list[ToolSpec], tools), + "max_tokens": typed_max_tokens, + } + + def restore_phase_from_session(project_root: Path) -> str | None: session = project_root / ".pilot-agent" / "session.jsonl" phase: str | None = "discovery" diff --git a/pilot_agent/agent/session_lock.py b/pilot_agent/agent/session_lock.py index cf6a801..333ca30 100644 --- a/pilot_agent/agent/session_lock.py +++ b/pilot_agent/agent/session_lock.py @@ -3,6 +3,7 @@ import os from pathlib import Path from types import TracebackType +from typing import BinaryIO from pilot_agent.agent.state import pilot_agent_dir @@ -13,7 +14,7 @@ class ProjectSessionLock: def __init__(self, project_root: Path): self.project_root = project_root.resolve() self.path = pilot_agent_dir(self.project_root) / "session.lock" - self._fh = None + self._fh: BinaryIO | None = None def __enter__(self) -> ProjectSessionLock: self.path.parent.mkdir(parents=True, exist_ok=True) @@ -23,28 +24,29 @@ def __enter__(self) -> ProjectSessionLock: "Another Pilot Agent session is already running for this project. " "Stop it or remove .pilot-agent/session.lock if the process is gone." ) - self._fh = self.path.open("a+b") + fh = self.path.open("a+b") + self._fh = fh try: if os.name == "nt": import msvcrt - self._fh.seek(0) - msvcrt.locking(self._fh.fileno(), msvcrt.LK_NBLCK, 1) + fh.seek(0) + msvcrt.locking(fh.fileno(), msvcrt.LK_NBLCK, 1) # type: ignore[attr-defined] else: import fcntl - fcntl.flock(self._fh.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) + fcntl.flock(fh.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) except OSError as exc: - self._fh.close() + fh.close() self._fh = None raise RuntimeError( "Another Pilot Agent session is already running for this project. " "Stop it or remove .pilot-agent/session.lock if the process is gone." ) from exc - self._fh.seek(0) - self._fh.truncate() - self._fh.write(str(os.getpid()).encode()) - self._fh.flush() + fh.seek(0) + fh.truncate() + fh.write(str(os.getpid()).encode()) + fh.flush() _ACTIVE_LOCKS.add(resolved_lock) return self @@ -54,19 +56,20 @@ def __exit__( exc: BaseException | None, tb: TracebackType | None, ) -> None: - if self._fh is None: + fh = self._fh + if fh is None: return try: if os.name == "nt": import msvcrt - self._fh.seek(0) - msvcrt.locking(self._fh.fileno(), msvcrt.LK_UNLCK, 1) + fh.seek(0) + msvcrt.locking(fh.fileno(), msvcrt.LK_UNLCK, 1) # type: ignore[attr-defined] else: import fcntl - fcntl.flock(self._fh.fileno(), fcntl.LOCK_UN) + fcntl.flock(fh.fileno(), fcntl.LOCK_UN) finally: - self._fh.close() + fh.close() self._fh = None _ACTIVE_LOCKS.discard(self.path.resolve()) diff --git a/pilot_agent/agent/state.py b/pilot_agent/agent/state.py index e0f9552..e90dfb7 100644 --- a/pilot_agent/agent/state.py +++ b/pilot_agent/agent/state.py @@ -6,7 +6,7 @@ from typing import Any from pilot_agent.agent.safety import redact_sensitive_text, sanitize_text -from pilot_agent.agent.types import Message, from_json, to_json +from pilot_agent.agent.types import Message, SessionEvent, from_json, to_json TASKS_HEADING = "TO" + "DO" STATE_TEMPLATE = """# Project: {name} @@ -65,11 +65,14 @@ def append_reentry_request(project_root: Path, *, kind: str, description: str) - path.write_text(_append_to_section(text, TASKS_HEADING, task), encoding="utf-8") -def write_session_record(project_root: Path, record: Message | dict[str, Any]) -> None: +def write_session_record( + project_root: Path, + record: Message | SessionEvent | dict[str, Any], +) -> None: pilot_agent_dir(project_root).mkdir(parents=True, exist_ok=True) line = ( to_json(record) - if isinstance(record, Message) + if isinstance(record, (Message, SessionEvent)) else json.dumps(record, ensure_ascii=False) ) line = redact_sensitive_text(sanitize_text(line)) diff --git a/pilot_agent/agent/types.py b/pilot_agent/agent/types.py index 2558e11..e1a768e 100644 --- a/pilot_agent/agent/types.py +++ b/pilot_agent/agent/types.py @@ -54,13 +54,20 @@ class CompletionResponse: usage: dict[str, int] -Jsonable = ToolCall | ToolResult | Message | ToolSpec | CompletionResponse +@dataclass +class SessionEvent: + event_type: str + payload: dict[str, Any] = field(default_factory=dict) + + +Jsonable = ToolCall | ToolResult | Message | ToolSpec | CompletionResponse | SessionEvent _TYPE_MAP = { "ToolCall": ToolCall, "ToolResult": ToolResult, "Message": Message, "ToolSpec": ToolSpec, "CompletionResponse": CompletionResponse, + "SessionEvent": SessionEvent, } @@ -140,6 +147,13 @@ def _decode_completion(data: dict[str, Any]) -> CompletionResponse: ) +def _decode_session_event(data: dict[str, Any]) -> SessionEvent: + return SessionEvent( + event_type=str(data["event_type"]), + payload=cast(dict[str, Any], data.get("payload") or {}), + ) + + def from_json(raw: str) -> Jsonable: """Deserialize a canonical dataclass serialized by `to_json`.""" @@ -159,4 +173,6 @@ def from_json(raw: str) -> Jsonable: return _decode_tool_spec(clean) if cls is CompletionResponse: return _decode_completion(clean) + if cls is SessionEvent: + return _decode_session_event(clean) raise AssertionError(f"unsupported serialized type: {cls.__name__}") diff --git a/pilot_agent/cli/__init__.py b/pilot_agent/cli/__init__.py index a51f7f8..d0c7cad 100644 --- a/pilot_agent/cli/__init__.py +++ b/pilot_agent/cli/__init__.py @@ -1,5 +1,15 @@ """CLI package for the Typer app, setup wizard, auth helpers, and doctor checks.""" -from pilot_agent.cli.main import app +from __future__ import annotations + +from typing import Any + + +def __getattr__(name: str) -> Any: + if name == "app": + from pilot_agent.cli.main import app + + return app + raise AttributeError(name) __all__ = ["app"] diff --git a/pilot_agent/cli/main.py b/pilot_agent/cli/main.py index 575abbb..f6fc65e 100644 --- a/pilot_agent/cli/main.py +++ b/pilot_agent/cli/main.py @@ -17,6 +17,11 @@ import typer +from pilot_agent.agent.acceptance_blueprints import ( + get_blueprint, + list_blueprints, + render_blueprint, +) from pilot_agent.agent.context import ContextManager from pilot_agent.agent.loop import AgentLoop, restore_phase_from_session from pilot_agent.agent.phases import PHASES @@ -39,8 +44,8 @@ from pilot_agent.cli.ui.input import PilotAgentInput from pilot_agent.config.credentials import ( credential_services, - credentials_permissions, credentials_path, + credentials_permissions, get_credential, mask_secret, remove_credential, @@ -62,7 +67,7 @@ from pilot_agent.providers.base import Provider, from_config from pilot_agent.skills.registry import SkillRegistry from pilot_agent.tools.ask_user import AskUserTool -from pilot_agent.tools.base import ToolRegistry +from pilot_agent.tools.base import ToolRegistry, ToolSearchSettings from pilot_agent.tools.bash import BashTool from pilot_agent.tools.file_ops import EditFileTool, ListFilesTool, ReadFileTool, WriteFileTool from pilot_agent.tools.phase_tools import CompletePhaseTool @@ -79,12 +84,14 @@ sessions_app = typer.Typer(help="Manage sessions.") auth_app = typer.Typer(help="Manage credentials.") sandbox_app = typer.Typer(help="Manage Docker sandbox image and containers.") +blueprints_app = typer.Typer(help="Inspect scheduled acceptance blueprints.") app.add_typer(skills_app, name="skills") app.add_typer(config_app, name="config") app.add_typer(lessons_app, name="lessons") app.add_typer(sessions_app, name="sessions") app.add_typer(auth_app, name="auth") app.add_typer(sandbox_app, name="sandbox") +app.add_typer(blueprints_app, name="blueprints") console = create_console() INIT_PATH_ARGUMENT = typer.Argument(Path("."), help="Project path to initialize.") GLOBAL_PROVIDER: str | None = None @@ -193,7 +200,11 @@ def build_tool_registry( tools.append(WebFetchTool()) for tool in tools: tool.timeout_s = cfg.tool_timeout_s - return ToolRegistry(tools, project_root) + return ToolRegistry( + tools, + project_root, + tool_search=ToolSearchSettings.from_raw(cfg.tools.tool_search), + ) def run_loop(cfg: PilotAgentConfig) -> None: @@ -571,15 +582,18 @@ def tools_command( render_tools_table(cfg) emit("Change with: pilot-agent tools --enable|--disable") return - if tool not in {"web_search", "web_fetch", "deploy"}: - emit("Error: unknown tool. Use web_search, web_fetch, or deploy") + if tool not in {"web_search", "web_fetch", "deploy", "tool_search"}: + emit("Error: unknown tool. Use web_search, web_fetch, deploy, or tool_search") raise typer.Exit(1) if enable and disable: emit("Error: choose only one of --enable or --disable") raise typer.Exit(1) updates: dict[str, object] = {} if enable or disable: - updates[f"tools.{tool}.enabled"] = enable + enabled_value: object = "on" if enable else "off" + if tool != "tool_search": + enabled_value = enable + updates[f"tools.{tool}.enabled"] = enabled_value if provider is not None: if tool != "web_search": emit("Error: --provider applies only to web_search") @@ -610,6 +624,11 @@ def render_tools_table(cfg: PilotAgentConfig) -> None: "✓ enabled" if cfg.tools.web_fetch.enabled else "✗ disabled", "SSRF checks", ) + table.add_row( + "tool_search", + f"{cfg.tools.tool_search.enabled}", + f"threshold {cfg.tools.tool_search.threshold_pct:g}%", + ) vercel_key = get_credential("vercel", home, env_name=cfg.phases.deploy.vercel_token_env) deploy_state = ( "✓ enabled" if cfg.tools.deploy.enabled or cfg.phases.deploy.enabled else "✗ disabled" @@ -622,6 +641,27 @@ def render_tools_table(cfg: PilotAgentConfig) -> None: emit(table) +@blueprints_app.callback(invoke_without_command=True) +def blueprints_root(ctx: typer.Context) -> None: + if ctx.invoked_subcommand is not None: + return + table = simple_table("id", "cadence", "purpose") + for blueprint in list_blueprints(): + table.add_row(blueprint.id, blueprint.cadence, blueprint.purpose) + emit(table) + emit("Show details with: pilot-agent blueprints show ") + + +@blueprints_app.command("show") +def blueprints_show(blueprint_id: str) -> None: + try: + blueprint = get_blueprint(blueprint_id) + except ValueError as exc: + emit(f"Error: {exc}") + raise typer.Exit(1) from None + emit(render_blueprint(blueprint)) + + @app.command("settings") def settings() -> None: cfg = load_config_or_exit() diff --git a/pilot_agent/cli/setup_wizard.py b/pilot_agent/cli/setup_wizard.py index 8151ca1..cbf6c8d 100644 --- a/pilot_agent/cli/setup_wizard.py +++ b/pilot_agent/cli/setup_wizard.py @@ -187,7 +187,7 @@ def _ensure_provider_key( return link = _key_url(provider) if link: - console.print(f"Get a key: {link}", style="pilot_agent.dim") + console.print(f"Get a key: {link}", style="pilot_agent.muted") for attempt in range(1, 4): key = prompt.prompt(f"{provider} API key", password=True, default="") if not key: @@ -238,7 +238,10 @@ def _ensure_vercel_token(prompt: PilotAgentInput, console: Console, home: Path) style="pilot_agent.ok", ) return - console.print("Create a token: https://vercel.com/account/tokens", style="pilot_agent.dim") + console.print( + "Create a token: https://vercel.com/account/tokens", + style="pilot_agent.muted", + ) token = prompt.prompt("Vercel token (Enter to skip)", password=True, default="") if token: set_credential("vercel", token, home) diff --git a/pilot_agent/config/defaults.yaml b/pilot_agent/config/defaults.yaml index 36ded52..ada2a07 100644 --- a/pilot_agent/config/defaults.yaml +++ b/pilot_agent/config/defaults.yaml @@ -26,6 +26,11 @@ tools: enabled: false deploy: enabled: true + tool_search: + enabled: auto + threshold_pct: 10.0 + search_default_limit: 5 + max_search_limit: 20 ui: color: auto show_token_counter: true diff --git a/pilot_agent/config/schema.py b/pilot_agent/config/schema.py index 10bcfb3..2f214e4 100644 --- a/pilot_agent/config/schema.py +++ b/pilot_agent/config/schema.py @@ -74,10 +74,18 @@ class DeployToolConfig(BaseModel): enabled: bool = True +class ToolSearchBridgeConfig(BaseModel): + enabled: Literal["auto", "on", "off"] = "auto" + threshold_pct: float = 10.0 + search_default_limit: int = 5 + max_search_limit: int = 20 + + class ToolsConfig(BaseModel): web_search: WebSearchConfig = Field(default_factory=WebSearchConfig) web_fetch: WebFetchConfig = Field(default_factory=WebFetchConfig) deploy: DeployToolConfig = Field(default_factory=DeployToolConfig) + tool_search: ToolSearchBridgeConfig = Field(default_factory=ToolSearchBridgeConfig) class ProviderConfig(BaseModel): @@ -199,6 +207,10 @@ def _env_map() -> dict[str, str]: "tools.web_search.searxng_url": "PILOT_AGENT_TOOLS_WEB_SEARCH_SEARXNG_URL", "tools.web_fetch.enabled": "PILOT_AGENT_TOOLS_WEB_FETCH_ENABLED", "tools.deploy.enabled": "PILOT_AGENT_TOOLS_DEPLOY_ENABLED", + "tools.tool_search.enabled": "PILOT_AGENT_TOOLS_TOOL_SEARCH_ENABLED", + "tools.tool_search.threshold_pct": "PILOT_AGENT_TOOLS_TOOL_SEARCH_THRESHOLD_PCT", + "tools.tool_search.search_default_limit": "PILOT_AGENT_TOOLS_TOOL_SEARCH_DEFAULT_LIMIT", + "tools.tool_search.max_search_limit": "PILOT_AGENT_TOOLS_TOOL_SEARCH_MAX_LIMIT", "ui.color": "PILOT_AGENT_UI_COLOR", "ui.show_token_counter": "PILOT_AGENT_UI_SHOW_TOKEN_COUNTER", } diff --git a/pilot_agent/providers/anthropic.py b/pilot_agent/providers/anthropic.py index b5900a6..b572f3f 100644 --- a/pilot_agent/providers/anthropic.py +++ b/pilot_agent/providers/anthropic.py @@ -9,6 +9,7 @@ from pilot_agent.agent.safety import sanitize_jsonable, sanitize_text from pilot_agent.agent.types import CompletionResponse, Message, Role, ToolCall, ToolSpec from pilot_agent.providers.base import Provider, register +from pilot_agent.providers.hardening import sanitize_tool_schema logger = logging.getLogger(__name__) @@ -40,7 +41,7 @@ def anthropic_tools(tools: Iterable[ToolSpec]) -> list[dict[str, Any]]: { "name": tool.name, "description": tool.description, - "input_schema": tool.parameters, + "input_schema": sanitize_tool_schema(tool.parameters), } for tool in tools ] diff --git a/pilot_agent/providers/base.py b/pilot_agent/providers/base.py index f010b24..6466815 100644 --- a/pilot_agent/providers/base.py +++ b/pilot_agent/providers/base.py @@ -3,13 +3,14 @@ import sys from abc import ABC, abstractmethod from collections.abc import Callable +from dataclasses import dataclass, replace from importlib import import_module from typing import Protocol, TypeVar from tenacity import RetryCallState, retry, retry_if_exception, stop_after_attempt, wait_exponential from pilot_agent.agent.types import CompletionResponse, Message, ToolSpec -from pilot_agent.providers.errors import format_provider_error +from pilot_agent.providers.errors import classify_provider_error, format_provider_error class ProviderConfigLike(Protocol): @@ -29,6 +30,14 @@ def resolve_key(self) -> str: ... } +@dataclass(frozen=True) +class ProviderRequest: + system: str + messages: list[Message] + tools: list[ToolSpec] + max_tokens: int = 4096 + + def _status_code(exc: BaseException) -> int | None: status = getattr(exc, "status_code", None) if isinstance(status, int): @@ -68,14 +77,51 @@ def complete( tools: list[ToolSpec], max_tokens: int = 4096, ) -> CompletionResponse: + request = self.prepare_request( + ProviderRequest(system=system, messages=messages, tools=tools, max_tokens=max_tokens) + ) try: - return self._complete_with_retry(system, messages, tools, max_tokens=max_tokens) + return self._complete_with_retry( + request.system, + request.messages, + request.tools, + max_tokens=request.max_tokens, + ) except Exception as exc: + try: + context_fallback = self._context_fallback(exc, request) + except Exception as fallback_exc: + exc = fallback_exc + else: + if context_fallback is not None: + return context_fallback provider_name = self.__class__.__name__.removesuffix("Provider").lower() + if provider_name not in _PROVIDER_MODULES: + raise raise RuntimeError( format_provider_error(exc, provider=provider_name, model=self.model) ) from exc + def prepare_request(self, request: ProviderRequest) -> ProviderRequest: + return request + + def _context_fallback( + self, + exc: Exception, + request: ProviderRequest, + ) -> CompletionResponse | None: + provider_name = self.__class__.__name__.removesuffix("Provider").lower() + info = classify_provider_error(exc, provider=provider_name, model=self.model) + if info.kind != "context" or request.max_tokens <= 1024: + return None + fallback = replace(request, max_tokens=1024) + return self._complete_with_retry( + fallback.system, + fallback.messages, + fallback.tools, + max_tokens=fallback.max_tokens, + ) + @retry( retry=retry_if_exception(_retryable), wait=wait_exponential(multiplier=1, min=1, max=20), @@ -128,4 +174,11 @@ def from_config(cfg: ProviderConfigLike) -> Provider: return provider_cls(model=cfg.model, api_key=cfg.resolve_key(), base_url=cfg.base_url) -__all__ = ["Provider", "_PROVIDER_MODULES", "_REGISTRY", "from_config", "register"] +__all__ = [ + "Provider", + "ProviderRequest", + "_PROVIDER_MODULES", + "_REGISTRY", + "from_config", + "register", +] diff --git a/pilot_agent/providers/hardening.py b/pilot_agent/providers/hardening.py new file mode 100644 index 0000000..6a6bbef --- /dev/null +++ b/pilot_agent/providers/hardening.py @@ -0,0 +1,149 @@ +from __future__ import annotations + +import json +import re +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +from pilot_agent.agent.safety import sanitize_jsonable, sanitize_text + +_JSON_SCHEMA_KEYS = { + "type", + "description", + "properties", + "required", + "additionalProperties", + "items", + "enum", + "default", + "minimum", + "maximum", + "minLength", + "maxLength", + "minItems", + "maxItems", + "pattern", + "format", + "anyOf", + "oneOf", + "allOf", +} +_JSON_TYPES = {"object", "array", "string", "number", "integer", "boolean", "null"} + + +@dataclass(frozen=True) +class ToolArgumentParse: + arguments: dict[str, Any] + error: str | None = None + + @property + def ok(self) -> bool: + return self.error is None + + +def sanitize_tool_schema(schema: Mapping[str, Any]) -> dict[str, Any]: + clean = _sanitize_schema_node(sanitize_jsonable(dict(schema))) + if clean.get("type") != "object": + clean["type"] = "object" + properties = clean.get("properties") + if not isinstance(properties, dict): + clean["properties"] = {} + required = clean.get("required") + if isinstance(required, list): + known = set(clean["properties"]) + clean["required"] = [str(item) for item in required if str(item) in known] + else: + clean.pop("required", None) + clean.setdefault("additionalProperties", False) + return clean + + +def parse_tool_arguments(raw_arguments: Any, *, tool_name: str) -> ToolArgumentParse: + raw_text = _strip_code_fence(sanitize_text(str(raw_arguments or "{}"))) + for candidate in _argument_candidates(raw_text): + try: + parsed = json.loads(candidate) + except json.JSONDecodeError: + continue + if isinstance(parsed, dict): + return ToolArgumentParse(arguments=parsed) + return ToolArgumentParse( + arguments={"_raw_arguments": raw_text}, + error=f"tool arguments for {tool_name} must decode to an object", + ) + try: + json.loads(raw_text) + except json.JSONDecodeError as exc: + detail = str(exc) + else: + detail = "tool arguments JSON must decode to an object" + return ToolArgumentParse( + arguments={"_raw_arguments": raw_text}, + error=f"invalid JSON arguments for {tool_name}: {detail}", + ) + + +def _sanitize_schema_node(value: Any) -> dict[str, Any]: + if not isinstance(value, dict): + return {"type": "object", "properties": {}, "additionalProperties": False} + clean: dict[str, Any] = {} + for key, item in value.items(): + if key not in _JSON_SCHEMA_KEYS: + continue + if key == "type": + typed = _sanitize_type(item) + if typed is not None: + clean[key] = typed + continue + if key == "properties": + clean[key] = _sanitize_properties(item) + continue + if key == "items": + clean[key] = _sanitize_schema_node(item) + continue + if key in {"anyOf", "oneOf", "allOf"}: + clean[key] = [ + _sanitize_schema_node(entry) for entry in item if isinstance(entry, dict) + ] if isinstance(item, list) else [] + continue + if key == "additionalProperties" and isinstance(item, dict): + clean[key] = _sanitize_schema_node(item) + continue + clean[key] = sanitize_jsonable(item) + return clean + + +def _sanitize_properties(value: Any) -> dict[str, Any]: + if not isinstance(value, dict): + return {} + return { + str(name): _sanitize_schema_node(item) + for name, item in value.items() + if isinstance(item, dict) + } + + +def _sanitize_type(value: Any) -> str | list[str] | None: + if isinstance(value, str): + return value if value in _JSON_TYPES else None + if isinstance(value, list): + typed = [str(item) for item in value if str(item) in _JSON_TYPES] + return typed or None + return None + + +def _argument_candidates(raw_text: str) -> list[str]: + candidates = [raw_text] + without_trailing_commas = re.sub(r",\s*([}\]])", r"\1", raw_text) + if without_trailing_commas not in candidates: + candidates.append(without_trailing_commas) + single_quoted = re.sub(r"'([^'\\]*(?:\\.[^'\\]*)*)'", r'"\1"', without_trailing_commas) + if single_quoted not in candidates: + candidates.append(single_quoted) + return candidates + + +def _strip_code_fence(value: str) -> str: + match = re.fullmatch(r"```(?:json)?\s*(.*?)\s*```", value, flags=re.S | re.I) + return match.group(1).strip() if match else value diff --git a/pilot_agent/providers/openai.py b/pilot_agent/providers/openai.py index 14c4b3d..280881b 100644 --- a/pilot_agent/providers/openai.py +++ b/pilot_agent/providers/openai.py @@ -18,6 +18,7 @@ ToolSpec, ) from pilot_agent.providers.base import Provider, register +from pilot_agent.providers.hardening import parse_tool_arguments, sanitize_tool_schema logger = logging.getLogger(__name__) @@ -38,7 +39,7 @@ def openai_tools(tools: Iterable[ToolSpec]) -> list[dict[str, Any]]: "function": { "name": tool.name, "description": tool.description, - "parameters": tool.parameters, + "parameters": sanitize_tool_schema(tool.parameters), }, } for tool in tools @@ -94,20 +95,16 @@ def parse_openai_message(raw_message: Any) -> Message: call_id = str(_get(raw_call, "id", "")) name = str(_get(function, "name", "")) raw_args = _get(function, "arguments", "{}") or "{}" - try: - args = json.loads(sanitize_text(str(raw_args))) - if not isinstance(args, dict): - raise ValueError("tool arguments JSON must decode to an object") - except (json.JSONDecodeError, ValueError) as exc: + parsed = parse_tool_arguments(raw_args, tool_name=name) + if not parsed.ok: errors.append( ToolResult( tool_call_id=call_id, - content=f"invalid JSON arguments for {name}: {exc}", + content=str(parsed.error), is_error=True, ) ) - args = {"_raw_arguments": str(raw_args)} - tool_calls.append(ToolCall(id=call_id, name=name, arguments=cast(dict[str, Any], args))) + tool_calls.append(ToolCall(id=call_id, name=name, arguments=parsed.arguments)) if errors: return Message(role=Role.TOOL, tool_results=errors) return Message(role=Role.ASSISTANT, content=sanitize_text(str(content)), tool_calls=tool_calls) diff --git a/pilot_agent/tools/base.py b/pilot_agent/tools/base.py index 630c3a7..2a9bff0 100644 --- a/pilot_agent/tools/base.py +++ b/pilot_agent/tools/base.py @@ -1,13 +1,22 @@ from __future__ import annotations import concurrent.futures +import json +import logging +import math +import re +import threading +import time import traceback from abc import ABC, abstractmethod +from collections.abc import Callable, Iterable, Mapping +from dataclasses import dataclass from pathlib import Path from typing import Any from jsonschema import ValidationError, validate +from pilot_agent.agent.hooks import RuntimeHooks from pilot_agent.agent.safety import redact_sensitive_text, sanitize_jsonable from pilot_agent.agent.tool_guardrails import ( ToolCallGuardrailController, @@ -15,12 +24,24 @@ ) from pilot_agent.agent.types import ToolCall, ToolResult, ToolSpec +logger = logging.getLogger(__name__) + +TOOL_SEARCH_NAME = "tool_search" +TOOL_DESCRIBE_NAME = "tool_describe" +TOOL_CALL_NAME = "tool_call" +BRIDGE_TOOL_NAMES = frozenset({TOOL_SEARCH_NAME, TOOL_DESCRIBE_NAME, TOOL_CALL_NAME}) +CHARS_PER_TOKEN = 4.0 + class Tool(ABC): name: str description: str parameters: dict[str, Any] timeout_s: int = 120 + toolset: str = "core" + deferrable: bool = False + parallel_safe: bool = False + path_scope_args: tuple[str, ...] = () @abstractmethod def execute(self, **kwargs: Any) -> str: ... @@ -28,33 +49,153 @@ def execute(self, **kwargs: Any) -> str: ... def spec(self) -> ToolSpec: return ToolSpec(name=self.name, description=self.description, parameters=self.parameters) + def available(self) -> bool: + return True + + +@dataclass(frozen=True) +class ToolSearchSettings: + enabled: str = "auto" + threshold_pct: float = 10.0 + search_default_limit: int = 5 + max_search_limit: int = 20 + + @classmethod + def from_raw(cls, raw: object) -> ToolSearchSettings: + enabled = str(getattr(raw, "enabled", "auto")).lower() + if enabled not in {"auto", "on", "off"}: + enabled = "auto" + return cls( + enabled=enabled, + threshold_pct=float(getattr(raw, "threshold_pct", 10.0)), + search_default_limit=int(getattr(raw, "search_default_limit", 5)), + max_search_limit=int(getattr(raw, "max_search_limit", 20)), + ) + + +@dataclass(frozen=True) +class ToolEntry: + tool: Tool + check_fn: Callable[[], bool] + + @property + def name(self) -> str: + return self.tool.name + + @property + def deferrable(self) -> bool: + return self.tool.deferrable + + @property + def parallel_safe(self) -> bool: + return self.tool.parallel_safe + + @property + def path_scope_args(self) -> tuple[str, ...]: + return self.tool.path_scope_args + + def spec(self) -> ToolSpec: + return self.tool.spec() + + +@dataclass(frozen=True) +class ToolExecution: + call: ToolCall + result: ToolResult + elapsed_s: float + class ToolRegistry: MAX_RESULT_CHARS = 8_000 + CHECK_TTL_S = 30.0 + MAX_SPEC_CACHE = 8 def __init__( self, tools: list[Tool], project_root: Path, guardrails: ToolCallGuardrailController | None = None, + hooks: RuntimeHooks | None = None, + tool_search: ToolSearchSettings | None = None, ): - self.tools = {tool.name: tool for tool in tools} + self.tools: dict[str, Tool] = {} + self.entries: dict[str, ToolEntry] = {} self.project_root = project_root.resolve() self.artifacts_dir = self.project_root / ".pilot-agent" / "artifacts" self.guardrails = guardrails or ToolCallGuardrailController() + self.hooks = hooks or RuntimeHooks() + self.tool_search = tool_search or ToolSearchSettings() + self._generation = 0 + self._check_cache: dict[str, tuple[float, bool]] = {} + self._spec_cache: dict[tuple[tuple[str, ...] | None, int, int], list[ToolSpec]] = {} + self._guardrail_lock = threading.Lock() + for tool in tools: + self.register(tool) + + def register(self, tool: Tool, check_fn: Callable[[], bool] | None = None) -> None: + if tool.name in BRIDGE_TOOL_NAMES: + raise ValueError(f"tool name is reserved for the tool-search bridge: {tool.name}") + self.tools[tool.name] = tool + self.entries[tool.name] = ToolEntry(tool=tool, check_fn=check_fn or tool.available) + self._check_cache.pop(tool.name, None) + self._generation += 1 + self._spec_cache.clear() + + def specs(self, allowed: list[str] | None = None, *, context_window: int = 0) -> list[ToolSpec]: + names_key = tuple(allowed) if allowed is not None else None + key = (names_key, self._generation, int(context_window)) + cached = self._spec_cache.get(key) + if cached is not None: + return list(cached) - def specs(self, allowed: list[str] | None = None) -> list[ToolSpec]: names = allowed or list(self.tools) - return [self.tools[name].spec() for name in names if name in self.tools] + specs = [ + self.entries[name].spec() + for name in names + if name in self.entries and self._entry_available(self.entries[name]) + ] + specs = self._assemble_tool_search_specs(specs, context_window=context_window) + if len(self._spec_cache) >= self.MAX_SPEC_CACHE: + self._spec_cache.pop(next(iter(self._spec_cache))) + self._spec_cache[key] = specs + return list(specs) - def execute(self, call: ToolCall) -> ToolResult: + def execute(self, call: ToolCall, context: Mapping[str, Any] | None = None) -> ToolResult: + context_data = dict(context or {}) + if call.name in BRIDGE_TOOL_NAMES: + return self._execute_bridge(call, context_data) tool = self.tools.get(call.name) if tool is None: return self._result(call.id, f"unknown tool: {call.name}", is_error=True) - arguments = sanitize_jsonable(call.arguments) - decision = self.guardrails.before_call(call.name, arguments) + + request_mw = self.hooks.apply_tool_request_middleware( + call.name, + dict(call.arguments), + **context_data, + tool_call_id=call.id, + ) + arguments = sanitize_jsonable(request_mw.payload) + if not isinstance(arguments, dict): + arguments = {} + block_message = self.hooks.get_tool_block_message( + call.name, + arguments, + **context_data, + tool_call_id=call.id, + middleware_trace=request_mw.trace, + ) + if block_message is not None: + result = self._result(call.id, block_message, is_error=True) + self._emit_post_tool(call, arguments, result, status="blocked", context=context_data) + return result + + with self._guardrail_lock: + decision = self.guardrails.before_call(call.name, arguments) if not decision.allows_execution: - return self._result(call.id, decision.message, is_error=True) + result = self._result(call.id, decision.message, is_error=True) + self._emit_post_tool(call, arguments, result, status="blocked", context=context_data) + return result + try: validate(instance=arguments, schema=tool.parameters) except ValidationError as exc: @@ -63,40 +204,114 @@ def execute(self, call: ToolCall) -> ToolResult: f"argument validation failed: {exc.message}", is_error=True, ) - after = self.guardrails.after_call( - call.name, - arguments, - result.content, - failed=True, - ) - if after.action == "warn": - result.content = append_guardrail_guidance(result.content, after) + self._record_guardrail_after(call.name, arguments, result, failed=True) + self._emit_post_tool(call, arguments, result, status="error", context=context_data) return result + executor: concurrent.futures.ThreadPoolExecutor | None = None + started_at = time.monotonic() try: executor = concurrent.futures.ThreadPoolExecutor(max_workers=1) - future = executor.submit(tool.execute, **arguments) + future = executor.submit( + self.hooks.run_tool_execution_middleware, + call.name, + arguments, + lambda next_args: tool.execute(**next_args), + **context_data, + tool_call_id=call.id, + ) output = future.result(timeout=tool.timeout_s) executor.shutdown(wait=True) - result = self._result(call.id, output, is_error=False) + observed = self._result(call.id, str(output), is_error=False) + self._emit_post_tool( + call, + arguments, + observed, + status="ok", + context=context_data, + duration_ms=int((time.monotonic() - started_at) * 1000), + ) + transformed = self.hooks.transform_tool_result( + observed.content, + tool_name=call.name, + args=arguments, + **context_data, + tool_call_id=call.id, + ) + result = observed if transformed == observed.content else self._result( + call.id, + transformed, + is_error=False, + ) except concurrent.futures.TimeoutError: if executor is not None: executor.shutdown(wait=False, cancel_futures=True) result = self._result(call.id, f"tool timed out after {tool.timeout_s}s", is_error=True) + self._emit_post_tool(call, arguments, result, status="error", context=context_data) except Exception: if executor is not None: executor.shutdown(wait=True, cancel_futures=True) - err = traceback.format_exc(limit=5) - result = self._result(call.id, err, is_error=True) - after = self.guardrails.after_call( - call.name, - arguments, - result.content, - failed=result.is_error, - ) + result = self._result(call.id, traceback.format_exc(limit=5), is_error=True) + self._emit_post_tool(call, arguments, result, status="error", context=context_data) + + self._record_guardrail_after(call.name, arguments, result, failed=result.is_error) + return result + + def execute_batch( + self, + calls: list[ToolCall], + context: Mapping[str, Any] | None = None, + ) -> list[ToolExecution]: + if not self._should_parallelize(calls): + return [self._execute_timed(call, context=context) for call in calls] + max_workers = min(4, len(calls)) + with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor: + futures = [ + executor.submit(self._execute_timed, call, context=context) for call in calls + ] + return [future.result() for future in futures] + + def _execute_timed( + self, + call: ToolCall, + context: Mapping[str, Any] | None = None, + ) -> ToolExecution: + started_at = time.monotonic() + result = self.execute(call, context=context) + return ToolExecution(call=call, result=result, elapsed_s=time.monotonic() - started_at) + + def _record_guardrail_after( + self, + tool_name: str, + arguments: dict[str, Any], + result: ToolResult, + *, + failed: bool, + ) -> None: + with self._guardrail_lock: + after = self.guardrails.after_call( + tool_name, + arguments, + result.content, + failed=failed, + ) if after.action == "warn": result.content = append_guardrail_guidance(result.content, after) - return result + + def _entry_available(self, entry: ToolEntry) -> bool: + cached = self._check_cache.get(entry.name) + now = time.monotonic() + if cached is not None: + created_at, value = cached + if now - created_at < self.CHECK_TTL_S: + return value + try: + value = bool(entry.check_fn()) + except Exception as exc: + logger.debug("tool availability check failed for %s: %s", entry.name, exc) + value = False + self._check_cache[entry.name] = (now, value) + return value def _artifact_path(self, call_id: str) -> Path: safe_id = call_id.replace("/", "_").replace(":", "_") @@ -124,3 +339,230 @@ def _result(self, call_id: str, output: str, *, is_error: bool) -> ToolResult: truncated=True, artifact_path=str(artifact), ) + + def _emit_post_tool( + self, + call: ToolCall, + args: dict[str, Any], + result: ToolResult, + *, + status: str, + context: dict[str, Any], + duration_ms: int = 0, + ) -> None: + self.hooks.invoke_hook( + "post_tool_call", + tool_name=call.name, + args=args, + result=result.content, + status=status, + is_error=result.is_error, + tool_call_id=call.id, + duration_ms=duration_ms, + artifact_path=result.artifact_path or "", + **context, + ) + + def _assemble_tool_search_specs( + self, + specs: list[ToolSpec], + *, + context_window: int, + ) -> list[ToolSpec]: + if self.tool_search.enabled == "off": + return specs + deferred = [ + spec + for spec in specs + if spec.name in self.entries and self.entries[spec.name].deferrable + ] + if not deferred: + return specs + deferred_tokens = _estimate_spec_tokens(deferred) + if self.tool_search.enabled == "auto": + threshold = int(context_window * (self.tool_search.threshold_pct / 100.0)) + if context_window > 0 and deferred_tokens < threshold: + return specs + if context_window <= 0 and deferred_tokens < 20_000: + return specs + deferred_names = {item.name for item in deferred} + visible = [spec for spec in specs if spec.name not in deferred_names] + return [*visible, *_bridge_specs(len(deferred))] + + def _execute_bridge(self, call: ToolCall, context: dict[str, Any]) -> ToolResult: + allowed = context.get("allowed_tools") + allowed_names = [str(item) for item in allowed] if isinstance(allowed, list) else None + deferred = self._deferred_specs(allowed_names) + if call.name == TOOL_SEARCH_NAME: + query = str(call.arguments.get("query") or "").strip() + if not query: + return self._result(call.id, "query is required", is_error=True) + limit = _safe_int(call.arguments.get("limit"), self.tool_search.search_default_limit) + limit = max(1, min(self.tool_search.max_search_limit, limit)) + matches = _search_specs(deferred, query, limit) + payload = { + "matches": [ + {"name": spec.name, "description": spec.description[:400]} for spec in matches + ] + } + return self._result(call.id, json.dumps(payload, ensure_ascii=False), is_error=False) + if call.name == TOOL_DESCRIBE_NAME: + name = str(call.arguments.get("name") or "") + spec = next((item for item in deferred if item.name == name), None) + if spec is None: + return self._result(call.id, f"unknown deferred tool: {name}", is_error=True) + return self._result( + call.id, + json.dumps(spec.__dict__, ensure_ascii=False), + is_error=False, + ) + if call.name == TOOL_CALL_NAME: + name = str(call.arguments.get("name") or "") + arguments = call.arguments.get("arguments") + if not isinstance(arguments, dict): + return self._result(call.id, "tool_call.arguments must be an object", is_error=True) + if name not in {spec.name for spec in deferred}: + return self._result( + call.id, + f"deferred tool unavailable in this scope: {name}", + is_error=True, + ) + return self.execute( + ToolCall(id=call.id, name=name, arguments=arguments), + context={**context, "bridge_tool": TOOL_CALL_NAME}, + ) + return self._result(call.id, f"unknown bridge tool: {call.name}", is_error=True) + + def _deferred_specs(self, allowed: list[str] | None) -> list[ToolSpec]: + names = allowed or list(self.tools) + return [ + self.entries[name].spec() + for name in names + if name in self.entries + and self.entries[name].deferrable + and self._entry_available(self.entries[name]) + ] + + def _should_parallelize(self, calls: list[ToolCall]) -> bool: + if len(calls) <= 1: + return False + entries: list[ToolEntry] = [] + for call in calls: + if call.name in BRIDGE_TOOL_NAMES: + return False + entry = self.entries.get(call.name) + if entry is None: + return False + entries.append(entry) + if all(entry.parallel_safe for entry in entries): + return True + reserved_paths: list[Path] = [] + for entry, call in zip(entries, calls, strict=True): + scoped = self._path_scope(entry, call.arguments) + if scoped is None: + return False + if any(_paths_overlap(scoped, existing) for existing in reserved_paths): + return False + reserved_paths.append(scoped) + return True + + def _path_scope(self, entry: ToolEntry, args: Mapping[str, Any]) -> Path | None: + for key in entry.path_scope_args: + value = args.get(key) + if not isinstance(value, str) or not value.strip(): + continue + candidate = Path(value).expanduser() + if not candidate.is_absolute(): + candidate = self.project_root / candidate + return candidate.absolute() + return None + + +def _estimate_spec_tokens(specs: Iterable[ToolSpec]) -> int: + chars = 0 + for spec in specs: + try: + chars += len(json.dumps(spec.__dict__, ensure_ascii=False, default=str)) + except TypeError: + chars += len(str(spec)) + return int(math.ceil(chars / CHARS_PER_TOKEN)) + + +def _bridge_specs(count: int) -> list[ToolSpec]: + return [ + ToolSpec( + name=TOOL_SEARCH_NAME, + description=( + f"Search {count} additional tools loaded on demand. Follow with " + f"{TOOL_DESCRIBE_NAME}, then {TOOL_CALL_NAME}." + ), + parameters={ + "type": "object", + "properties": { + "query": {"type": "string"}, + "limit": {"type": "integer", "minimum": 1}, + }, + "required": ["query"], + "additionalProperties": False, + }, + ), + ToolSpec( + name=TOOL_DESCRIBE_NAME, + description="Load the full JSON schema for one deferred tool.", + parameters={ + "type": "object", + "properties": {"name": {"type": "string"}}, + "required": ["name"], + "additionalProperties": False, + }, + ), + ToolSpec( + name=TOOL_CALL_NAME, + description="Invoke a deferred tool by exact name and arguments.", + parameters={ + "type": "object", + "properties": { + "name": {"type": "string"}, + "arguments": {"type": "object"}, + }, + "required": ["name", "arguments"], + "additionalProperties": False, + }, + ), + ] + + +def _search_specs(specs: list[ToolSpec], query: str, limit: int) -> list[ToolSpec]: + tokens = _tokenize(query) + if not tokens: + return [] + scored: list[tuple[int, ToolSpec]] = [] + for spec in specs: + text = f"{spec.name} {spec.name.replace('_', ' ')} {spec.description}" + text_tokens = set(_tokenize(text)) + score = sum(1 for token in tokens if token in text_tokens) + if score: + scored.append((score, spec)) + if not scored: + lowered = query.lower() + scored = [(1, spec) for spec in specs if lowered in spec.name.lower()] + scored.sort(key=lambda item: (-item[0], item[1].name)) + return [spec for _, spec in scored[:limit]] + + +def _tokenize(value: str) -> list[str]: + return [part.lower() for part in re.findall(r"[A-Za-z0-9]+", value)] + + +def _safe_int(value: Any, fallback: int) -> int: + try: + return int(value) + except (TypeError, ValueError): + return fallback + + +def _paths_overlap(left: Path, right: Path) -> bool: + left_parts = left.parts + right_parts = right.parts + common = min(len(left_parts), len(right_parts)) + return left_parts[:common] == right_parts[:common] diff --git a/pilot_agent/tools/file_ops.py b/pilot_agent/tools/file_ops.py index d928cd6..513a6ee 100644 --- a/pilot_agent/tools/file_ops.py +++ b/pilot_agent/tools/file_ops.py @@ -51,6 +51,7 @@ def _inside(path: Path, parent: Path) -> bool: class ReadFileTool(Tool): name = "read_file" + parallel_safe = True description = "Read a file with line numbers for precise edit references." parameters: dict[str, Any] = { "type": "object", @@ -80,6 +81,7 @@ def execute(self, **kwargs: Any) -> str: class WriteFileTool(Tool): name = "write_file" + path_scope_args = ("path",) description = "Create parent directories and overwrite a file inside the project root." parameters: dict[str, Any] = { "type": "object", @@ -102,6 +104,7 @@ def execute(self, **kwargs: Any) -> str: class EditFileTool(Tool): name = "edit_file" + path_scope_args = ("path",) description = "Replace a unique string in a project file; fails unless old_str occurs once." parameters: dict[str, Any] = { "type": "object", @@ -132,6 +135,7 @@ def execute(self, **kwargs: Any) -> str: class ListFilesTool(Tool): name = "list_files" + parallel_safe = True description = "List a project tree up to depth 3, ignoring common generated directories." parameters: dict[str, Any] = { "type": "object", diff --git a/pilot_agent/tools/skill_tools.py b/pilot_agent/tools/skill_tools.py index b30960d..552a4ee 100644 --- a/pilot_agent/tools/skill_tools.py +++ b/pilot_agent/tools/skill_tools.py @@ -13,6 +13,7 @@ def save(self, content: str) -> object: ... class LoadSkillTool(Tool): name = "load_skill" + parallel_safe = True description = "Load the full markdown for a named skill before using that procedure." parameters: dict[str, Any] = { "type": "object", diff --git a/pilot_agent/tools/web_fetch.py b/pilot_agent/tools/web_fetch.py index baa6998..3426b77 100644 --- a/pilot_agent/tools/web_fetch.py +++ b/pilot_agent/tools/web_fetch.py @@ -39,6 +39,7 @@ def text(self) -> str: class WebFetchTool(Tool): name = "web_fetch" + parallel_safe = True description = "Fetch an HTTP(S) page, reject private IP targets, and return extracted text." parameters: dict[str, Any] = { "type": "object", diff --git a/pilot_agent/tools/web_search.py b/pilot_agent/tools/web_search.py index 657eeb0..22d15d6 100644 --- a/pilot_agent/tools/web_search.py +++ b/pilot_agent/tools/web_search.py @@ -10,6 +10,7 @@ class WebSearchTool(Tool): name = "web_search" + parallel_safe = True description = "Search the web and return a numbered list of titles, URLs, and snippets." parameters: dict[str, Any] = { "type": "object", diff --git a/tests/test_acceptance_blueprints.py b/tests/test_acceptance_blueprints.py new file mode 100644 index 0000000..e94efbe --- /dev/null +++ b/tests/test_acceptance_blueprints.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +import pytest +from typer.testing import CliRunner + +from pilot_agent.agent.acceptance_blueprints import get_blueprint, list_blueprints, render_blueprint +from pilot_agent.cli import app + + +def test_acceptance_blueprint_catalog_has_product_checks() -> None: + ids = {blueprint.id for blueprint in list_blueprints()} + + assert {"nightly-checks", "dependency-audit", "deploy-verification"} <= ids + assert "scripts/run_tests.sh" in get_blueprint("nightly-checks").commands + assert "docker compose build" in get_blueprint("deploy-verification").commands + + +def test_acceptance_blueprint_render_and_unknown_id() -> None: + rendered = render_blueprint(get_blueprint("dependency-audit")) + + assert "# Dependency drift audit" in rendered + assert "uv lock --check" in rendered + with pytest.raises(ValueError, match="unknown acceptance blueprint"): + get_blueprint("morning-briefing") + + +def test_blueprints_cli_lists_and_shows_blueprints() -> None: + runner = CliRunner() + + list_result = runner.invoke(app, ["blueprints"]) + show_result = runner.invoke(app, ["blueprints", "show", "nightly-checks"]) + + assert list_result.exit_code == 0 + assert "nightly-checks" in list_result.output + assert show_result.exit_code == 0 + assert "Nightly local acceptance" in show_result.output + assert "scripts/run_tests.sh" in show_result.output diff --git a/tests/test_cli_loop.py b/tests/test_cli_loop.py index 2ffc862..898e1c6 100644 --- a/tests/test_cli_loop.py +++ b/tests/test_cli_loop.py @@ -1,11 +1,14 @@ from __future__ import annotations +import threading +import time from pathlib import Path from typing import Any from typer.testing import CliRunner from pilot_agent.agent.context import ContextManager, StaticSummaryProvider +from pilot_agent.agent.hooks import LLM_REQUEST_MIDDLEWARE, RuntimeHooks from pilot_agent.agent.loop import AgentLoop, restore_phase_from_session from pilot_agent.agent.phases import PHASES, PIPELINE from pilot_agent.agent.prompts import ( @@ -67,6 +70,29 @@ def execute(self, **kwargs: Any) -> str: return "noop" +class ParallelNoopTool(NoopTool): + name = "parallel_noop" + parallel_safe = True + active = 0 + max_active = 0 + lock = threading.Lock() + + @classmethod + def reset(cls) -> None: + with cls.lock: + cls.active = 0 + cls.max_active = 0 + + def execute(self, **kwargs: Any) -> str: + with self.lock: + self.__class__.active += 1 + self.__class__.max_active = max(self.__class__.max_active, self.__class__.active) + time.sleep(0.05) + with self.lock: + self.__class__.active -= 1 + return "noop" + + class SequenceProvider(CompleteOnlyProvider): def __init__(self, calls: list[ToolCall]): super().__init__() @@ -87,6 +113,26 @@ def _complete( ) +class BatchProvider(CompleteOnlyProvider): + def __init__(self, calls: list[ToolCall]): + super().__init__() + self.calls = calls + + def _complete( + self, + system: str, + messages: list[Message], + tools: list[ToolSpec], + max_tokens: int = 4096, + ) -> CompletionResponse: + self.system_seen = system + return CompletionResponse( + message=Message(role=Role.ASSISTANT, tool_calls=self.calls), + stop_reason="tool_use", + usage={"input_tokens": 1, "output_tokens": 1}, + ) + + class SkillBackend: def __init__(self) -> None: self.outcomes: list[tuple[str, bool]] = [] @@ -176,6 +222,82 @@ def test_loop_turn_complete_phase_logs_and_advances(tmp_path: Path) -> None: assert "phase_change" in (tmp_path / ".pilot-agent" / "session.jsonl").read_text() +def test_loop_applies_llm_hooks_and_middleware(tmp_path: Path) -> None: + init_project_state(tmp_path, "Demo") + provider = CompleteOnlyProvider() + hooks = RuntimeHooks() + events: list[str] = [] + + def mutate_request(request: dict[str, Any], **_: Any) -> dict[str, Any]: + updated = dict(request) + updated["system"] = str(updated["system"]) + "\nHOOKED" + return {"name": "append-hook-marker", "request": updated, "changed": True} + + hooks.register_hook("pre_llm_call", lambda **_: events.append("pre")) + hooks.register_middleware(LLM_REQUEST_MIDDLEWARE, mutate_request) + hooks.register_hook("post_llm_call", lambda **_: events.append("post")) + loop = AgentLoop( + project_root=tmp_path, + provider=provider, + registry=ToolRegistry([NoopTool()], tmp_path, hooks=hooks), + ctx=ContextManager( + StaticSummaryProvider(), + session_log=tmp_path / ".pilot-agent/session.jsonl", + ), + ) + + loop.run_turn() + + assert "HOOKED" in provider.system_seen + assert events == ["pre", "post"] + + +def test_loop_batches_parallel_safe_tool_calls(tmp_path: Path) -> None: + init_project_state(tmp_path, "Demo") + ParallelNoopTool.reset() + provider = BatchProvider( + [ + ToolCall("a", "parallel_noop", {}), + ToolCall("b", "parallel_noop", {}), + ] + ) + loop = AgentLoop( + project_root=tmp_path, + provider=provider, + registry=ToolRegistry([ParallelNoopTool()], tmp_path), + ctx=ContextManager( + StaticSummaryProvider(), + session_log=tmp_path / ".pilot-agent/session.jsonl", + ), + ) + + loop.run_turn() + + assert ParallelNoopTool.max_active >= 2 + assert loop.history[-1].role is Role.TOOL + assert len(loop.history[-1].tool_results) == 2 + + +def test_loop_records_iteration_budget_exhaustion(tmp_path: Path) -> None: + init_project_state(tmp_path, "Demo") + loop = AgentLoop( + project_root=tmp_path, + provider=CompleteOnlyProvider(), + registry=ToolRegistry([NoopTool()], tmp_path), + ctx=ContextManager( + StaticSummaryProvider(), + session_log=tmp_path / ".pilot-agent/session.jsonl", + ), + ) + + loop.run(max_turns=1) + + assert loop.phase is not None + assert loop.phase.name == "planning" + session_text = (tmp_path / ".pilot-agent" / "session.jsonl").read_text(encoding="utf-8") + assert "iteration_budget_exhausted" in session_text + + def test_loop_pins_loaded_skill_and_records_outcome(tmp_path: Path) -> None: init_project_state(tmp_path, "Demo") backend = SkillBackend() diff --git a/tests/test_context.py b/tests/test_context.py index 2f61bcc..59e4713 100644 --- a/tests/test_context.py +++ b/tests/test_context.py @@ -39,7 +39,7 @@ def test_prepare_truncates_copy_not_original_and_preserves_recent_turns(tmp_path assert after_tokens <= manager.threshold assert history[1].tool_results[0].content == "x" * 300 assert prepared[0].pinned is True - assert "[output 300 chars -> a0.txt]" in prepared[1].tool_results[0].content + assert prepared[1].tool_results[0].content.startswith("[output 300 chars -> a0.txt") assert prepared[-1].tool_results[0].content == "x" * 300 @@ -58,8 +58,11 @@ def test_prepare_summarizes_when_truncation_is_not_enough(tmp_path: Path) -> Non assert prepared[0].content == "[Compressed history]\nphase summary" assert any(message.content == "pinned" for message in prepared) assert len([m for m in prepared if m.role in {Role.ASSISTANT, Role.TOOL}]) == 6 - assert event["_type"] == "compaction" - assert event["before_tokens"] > event["after_tokens"] + assert event["_type"] == "SessionEvent" + assert event["event_type"] == "compaction" + assert event["payload"]["before_tokens"] > event["payload"]["after_tokens"] + assert event["payload"]["provenance"]["compaction_depth"] == 1 + assert event["payload"]["provenance"]["session_kind"] == "continuation" def test_compaction_fallback_preserves_paths_when_summarizer_fails(tmp_path: Path) -> None: diff --git a/tests/test_iteration_budget.py b/tests/test_iteration_budget.py new file mode 100644 index 0000000..ed049f0 --- /dev/null +++ b/tests/test_iteration_budget.py @@ -0,0 +1,28 @@ +from __future__ import annotations + +import pytest + +from pilot_agent.agent.iteration_budget import IterationBudget + + +def test_iteration_budget_consumes_refunds_and_exhausts() -> None: + budget = IterationBudget(2) + + assert budget.consume() is True + assert budget.snapshot().remaining == 1 + budget.refund() + assert budget.remaining == 2 + assert budget.consume(2) is True + assert budget.consume() is False + assert budget.snapshot().exhausted is True + + +def test_iteration_budget_rejects_invalid_values() -> None: + with pytest.raises(ValueError): + IterationBudget(0) + + budget = IterationBudget(1) + with pytest.raises(ValueError): + budget.consume(0) + with pytest.raises(ValueError): + budget.refund(0) diff --git a/tests/test_providers.py b/tests/test_providers.py index b4120ee..de4e49f 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -12,7 +12,9 @@ anthropic_tools, parse_anthropic_response, ) -from pilot_agent.providers.base import Provider, from_config, register +from pilot_agent.providers.base import Provider, ProviderRequest, from_config, register +from pilot_agent.providers.errors import format_provider_error +from pilot_agent.providers.hardening import sanitize_tool_schema from pilot_agent.providers.openai import ( OpenAIProvider, openai_messages, @@ -23,13 +25,29 @@ def test_anthropic_tool_spec_uses_input_schema() -> None: - tools = [ToolSpec("read_file", "Read a file", {"type": "object", "properties": {}})] + tools = [ + ToolSpec( + "read_file", + "Read a file", + { + "type": "object", + "properties": {"path": {"type": "string", "x-local": object()}}, + "required": ["path", "missing"], + "x-unsafe": object(), + }, + ) + ] assert anthropic_tools(tools) == [ { "name": "read_file", "description": "Read a file", - "input_schema": {"type": "object", "properties": {}}, + "input_schema": { + "type": "object", + "properties": {"path": {"type": "string"}}, + "required": ["path"], + "additionalProperties": False, + }, } ] @@ -142,7 +160,11 @@ def test_openai_tool_spec_and_message_conversion() -> None: "function": { "name": "bash", "description": "Run shell", - "parameters": {"type": "object"}, + "parameters": { + "type": "object", + "properties": {}, + "additionalProperties": False, + }, }, } ] @@ -213,6 +235,56 @@ def test_openai_broken_json_returns_tool_result_error() -> None: ] +def test_openai_recovers_common_malformed_json_arguments() -> None: + response = { + "choices": [ + { + "message": { + "tool_calls": [ + { + "id": "call_recover", + "function": { + "name": "read_file", + "arguments": '{"path": "README.md",}', + }, + } + ], + }, + "finish_reason": "tool_calls", + } + ], + "usage": {}, + } + + parsed = parse_openai_response(response) + + assert parsed.message.role is Role.ASSISTANT + assert parsed.message.tool_calls == [ + ToolCall("call_recover", "read_file", {"path": "README.md"}) + ] + + +def test_sanitize_tool_schema_keeps_provider_safe_json_schema() -> None: + schema = sanitize_tool_schema( + { + "type": "function", + "properties": { + "ok": {"type": "string", "description": "safe", "callable": object()}, + "bad": "not schema", + }, + "required": ["ok", "bad"], + "callable": object(), + } + ) + + assert schema == { + "type": "object", + "properties": {"ok": {"type": "string", "description": "safe"}}, + "required": ["ok"], + "additionalProperties": False, + } + + def test_anthropic_complete_uses_mocked_client(monkeypatch: pytest.MonkeyPatch) -> None: provider = AnthropicProvider(model="claude-sonnet-4-6", api_key="test") response = Response( @@ -326,6 +398,10 @@ class NonRetryableError(Exception): status_code = 400 +class ContextWindowError(Exception): + status_code = 400 + + @register("retry-test") class RetryProvider(Provider): def __init__(self, model: str, api_key: str, base_url: str | None = None): @@ -369,6 +445,56 @@ def _complete( raise NonRetryableError("bad request") +@register("context-fallback-test") +class ContextFallbackProvider(RetryProvider): + def __init__(self, model: str, api_key: str, base_url: str | None = None): + super().__init__(model, api_key, base_url) + self.max_tokens_seen: list[int] = [] + + def _complete( + self, + system: str, + messages: list[Message], + tools: list[ToolSpec], + max_tokens: int = 4096, + ): + self.calls += 1 + self.max_tokens_seen.append(max_tokens) + if self.calls == 1: + raise ContextWindowError("context length too long") + return type( + "Completion", + (), + {"message": Message(Role.ASSISTANT), "stop_reason": "end_turn", "usage": {}}, + )() + + +@register("request-hook-test") +class RequestHookProvider(RetryProvider): + def prepare_request(self, request: ProviderRequest) -> ProviderRequest: + return ProviderRequest( + system=f"{request.system}\nprepared", + messages=request.messages, + tools=request.tools, + max_tokens=request.max_tokens, + ) + + def _complete( + self, + system: str, + messages: list[Message], + tools: list[ToolSpec], + max_tokens: int = 4096, + ): + self.calls += 1 + assert system.endswith("prepared") + return type( + "Completion", + (), + {"message": Message(Role.ASSISTANT), "stop_reason": "end_turn", "usage": {}}, + )() + + def test_provider_factory_uses_registry_without_if_else() -> None: class Cfg: provider = "retry-test" @@ -390,3 +516,28 @@ def test_provider_base_retries_429_and_not_400() -> None: with pytest.raises(NonRetryableError): no_retry.complete("", [], []) assert no_retry.calls == 1 + + +def test_provider_base_context_fallback_reduces_output_tokens() -> None: + provider = ContextFallbackProvider("m", "k") + + provider.complete("", [], [], max_tokens=4096) + + assert provider.calls == 2 + assert provider.max_tokens_seen == [4096, 1024] + + +def test_provider_base_prepare_request_hook_runs_before_call() -> None: + provider = RequestHookProvider("m", "k") + + provider.complete("sys", [], []) + + assert provider.calls == 1 + + +def test_provider_error_formatting_remains_available() -> None: + assert "Run /compact" in format_provider_error( + ContextWindowError("context length too long"), + provider="openai", + model="gpt", + ) diff --git a/tests/test_tools.py b/tests/test_tools.py index 57f6dfb..063529b 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -2,16 +2,18 @@ import json import subprocess +import threading import time from pathlib import Path from typing import Any import pytest +from pilot_agent.agent.hooks import TOOL_REQUEST_MIDDLEWARE, RuntimeHooks from pilot_agent.agent.tool_guardrails import ToolCallGuardrailController from pilot_agent.agent.types import ToolCall from pilot_agent.tools.ask_user import AskUserTool -from pilot_agent.tools.base import Tool, ToolRegistry +from pilot_agent.tools.base import Tool, ToolRegistry, ToolSearchSettings from pilot_agent.tools.bash import BashTool from pilot_agent.tools.file_ops import EditFileTool, ListFilesTool, ReadFileTool, WriteFileTool from pilot_agent.tools.phase_tools import CompletePhaseTool @@ -43,6 +45,58 @@ def execute(self, **kwargs: Any) -> str: return "too late" +class DeferredEchoTool(EchoTool): + name = "deferred_echo" + description = "Echo text through a deferred tool." + deferrable = True + + +class ParallelProbeTool(Tool): + name = "parallel_probe" + description = "Probe parallel execution." + parallel_safe = True + parameters: dict[str, Any] = { + "type": "object", + "properties": {"label": {"type": "string"}}, + "required": ["label"], + "additionalProperties": False, + } + active = 0 + max_active = 0 + lock = threading.Lock() + + @classmethod + def reset(cls) -> None: + with cls.lock: + cls.active = 0 + cls.max_active = 0 + + def execute(self, **kwargs: Any) -> str: + with self.lock: + self.__class__.active += 1 + self.__class__.max_active = max(self.__class__.max_active, self.__class__.active) + time.sleep(0.05) + with self.lock: + self.__class__.active -= 1 + return str(kwargs["label"]) + + +class PathScopedProbeTool(ParallelProbeTool): + name = "path_probe" + description = "Probe path scoped parallel execution." + parallel_safe = False + path_scope_args = ("path",) + parameters: dict[str, Any] = { + "type": "object", + "properties": {"path": {"type": "string"}}, + "required": ["path"], + "additionalProperties": False, + } + + def execute(self, **kwargs: Any) -> str: + return super().execute(label=str(kwargs["path"])) + + def test_registry_validates_writes_artifact_and_truncates(tmp_path: Path) -> None: registry = ToolRegistry([EchoTool()], tmp_path) long_text = "a" * 9_000 @@ -194,6 +248,110 @@ def test_tool_guardrails_warn_on_repeated_failure(tmp_path: Path) -> None: assert "Tool loop warning" in second.content +def test_runtime_hooks_can_mutate_block_and_transform_tool_calls(tmp_path: Path) -> None: + hooks = RuntimeHooks() + observed: list[tuple[str, str]] = [] + + def normalize_args(args: dict[str, Any], **_: Any) -> dict[str, Any]: + return { + "name": "normalize_args", + "args": {"text": str(args["text"]).upper()}, + "changed": True, + } + + def pre_tool(args: dict[str, Any], **_: Any) -> dict[str, str] | None: + if args["text"] == "BLOCK": + return {"action": "block", "message": "blocked by test hook"} + observed.append(("pre", str(args["text"]))) + return None + + hooks.register_middleware(TOOL_REQUEST_MIDDLEWARE, normalize_args) + hooks.register_hook("pre_tool_call", pre_tool) + hooks.register_hook( + "post_tool_call", + lambda result, **_: observed.append(("post", str(result))), + ) + hooks.register_hook("transform_tool_result", lambda result, **_: f"{result}!") + registry = ToolRegistry([EchoTool()], tmp_path, hooks=hooks) + + ok = registry.execute(ToolCall("ok", "echo", {"text": "hello"})) + blocked = registry.execute(ToolCall("blocked", "echo", {"text": "block"})) + + assert ok.content == "HELLO!" + assert blocked.is_error is True + assert "blocked by test hook" in blocked.content + assert observed == [("pre", "HELLO"), ("post", "HELLO"), ("post", "blocked by test hook")] + + +def test_tool_search_bridge_defers_and_scopes_tools(tmp_path: Path) -> None: + registry = ToolRegistry( + [EchoTool(), DeferredEchoTool()], + tmp_path, + tool_search=ToolSearchSettings(enabled="on", search_default_limit=2), + ) + + names = [spec.name for spec in registry.specs()] + search = registry.execute(ToolCall("search", "tool_search", {"query": "deferred echo"})) + describe = registry.execute( + ToolCall("describe", "tool_describe", {"name": "deferred_echo"}) + ) + called = registry.execute( + ToolCall( + "call", + "tool_call", + {"name": "deferred_echo", "arguments": {"text": "hi"}}, + ) + ) + out_of_scope = registry.execute( + ToolCall( + "scope", + "tool_call", + {"name": "deferred_echo", "arguments": {"text": "hi"}}, + ), + context={"allowed_tools": ["echo"]}, + ) + + assert "echo" in names + assert "deferred_echo" not in names + assert {"tool_search", "tool_describe", "tool_call"}.issubset(names) + assert json.loads(search.content)["matches"][0]["name"] == "deferred_echo" + assert json.loads(describe.content)["name"] == "deferred_echo" + assert called.content == "hi" + assert out_of_scope.is_error is True + assert "unavailable in this scope" in out_of_scope.content + + +def test_execute_batch_parallelizes_only_safe_groups(tmp_path: Path) -> None: + ParallelProbeTool.reset() + registry = ToolRegistry([ParallelProbeTool(), PathScopedProbeTool()], tmp_path) + + registry.execute_batch( + [ + ToolCall("a", "parallel_probe", {"label": "a"}), + ToolCall("b", "parallel_probe", {"label": "b"}), + ] + ) + assert ParallelProbeTool.max_active >= 2 + + PathScopedProbeTool.reset() + registry.execute_batch( + [ + ToolCall("c", "path_probe", {"path": "a.txt"}), + ToolCall("d", "path_probe", {"path": "b.txt"}), + ] + ) + assert PathScopedProbeTool.max_active >= 2 + + PathScopedProbeTool.reset() + registry.execute_batch( + [ + ToolCall("e", "path_probe", {"path": "same.txt"}), + ToolCall("f", "path_probe", {"path": "same.txt"}), + ] + ) + assert PathScopedProbeTool.max_active == 1 + + def test_run_and_check_kills_sleep_process(tmp_path: Path) -> None: registry = ToolRegistry([RunAndCheckTool(tmp_path)], tmp_path) diff --git a/tests/test_types.py b/tests/test_types.py index 498abc6..06d9dee 100644 --- a/tests/test_types.py +++ b/tests/test_types.py @@ -4,6 +4,7 @@ CompletionResponse, Message, Role, + SessionEvent, ToolCall, ToolResult, from_json, @@ -52,3 +53,14 @@ def test_completion_response_round_trip() -> None: decoded = from_json(to_json(response)) assert decoded == response + + +def test_session_event_round_trip() -> None: + event = SessionEvent( + event_type="compaction", + payload={"before_tokens": 100, "provenance": {"compaction_depth": 1}}, + ) + + decoded = from_json(to_json(event)) + + assert decoded == event diff --git a/web/app/globals.css b/web/app/globals.css index b600824..46a7316 100644 --- a/web/app/globals.css +++ b/web/app/globals.css @@ -103,6 +103,7 @@ --mono: var(--font-mono), ui-monospace, SFMono-Regular, Menlo, monospace; --maxw: 780px; + --content-shift: 24px; --shadow: 0 1px 3px rgba(0, 0, 0, 0.08); } @@ -143,7 +144,7 @@ p { .wrap { max-width: var(--maxw); - margin: 0; + margin: 0 0 0 var(--content-shift); padding: 0 24px 0 56px; } @@ -200,7 +201,7 @@ header.nav.scrolled { gap: 16px; height: 60px; max-width: var(--maxw); - margin: 0; + margin: 0 0 0 var(--content-shift); padding: 0 24px 0 56px; } .brand { @@ -1042,6 +1043,7 @@ footer.site { text-align: left; user-select: none; pointer-events: none; + margin-left: var(--content-shift); padding: 0 24px 24px 56px; letter-spacing: -0.02em; overflow: hidden; @@ -1049,6 +1051,9 @@ footer.site { /* ============ RESPONSIVE ============ */ @media (max-width: 640px) { + :root { + --content-shift: 10px; + } .nav-links { display: none; } diff --git a/web/components/sections/cta.tsx b/web/components/sections/cta.tsx index f8ccfd9..92c696c 100644 --- a/web/components/sections/cta.tsx +++ b/web/components/sections/cta.tsx @@ -3,7 +3,7 @@ import { CopyButton } from '../copy-button' import { GitHub } from '../icons' const REPO = 'https://github.com/Hqzdev/pilot-agent' -const CURL = 'curl -fsSL https://raw.githubusercontent.com/Hqzdev/pilot-agent/main/install.sh | bash' +const CURL = 'curl -fsSL https://pilotagent.vercel.app/install.sh | bash' export function Cta() { return ( diff --git a/web/components/sections/install.tsx b/web/components/sections/install.tsx index c3f7798..6383dd6 100644 --- a/web/components/sections/install.tsx +++ b/web/components/sections/install.tsx @@ -2,9 +2,10 @@ import { Reveal } from '../reveal' import { CopyButton } from '../copy-button' import { Monitor, Bolt, Code, Play } from '../icons' -const CURL = 'curl -fsSL https://raw.githubusercontent.com/Hqzdev/pilot-agent/main/install.sh | bash' +const INSTALL_URL = 'https://pilotagent.vercel.app/install.sh' +const CURL = `curl -fsSL ${INSTALL_URL} | bash` const CURL_SKIP = - 'curl -fsSL https://raw.githubusercontent.com/Hqzdev/pilot-agent/main/install.sh | bash -s -- --skip-setup' + `curl -fsSL ${INSTALL_URL} | bash -s -- --skip-setup` const UV = 'uv tool install git+https://github.com/Hqzdev/pilot-agent' const QUICKSTART = `pilot-agent setup cd your-project @@ -35,7 +36,7 @@ export function Install() {
                     curl -fsSL{' '}
                     
-                      https://raw.githubusercontent.com/Hqzdev/pilot-agent/main/install.sh
+                      {INSTALL_URL}
                     {' '}
                     | bash
                   
@@ -53,7 +54,7 @@ export function Install() {
                     curl -fsSL{' '}
                     
-                      https://raw.githubusercontent.com/Hqzdev/pilot-agent/main/install.sh
+                      {INSTALL_URL}
                     {' '}
                     | bash{' '}
                     -s -- --skip-setup
diff --git a/web/public/install.sh b/web/public/install.sh
new file mode 100644
index 0000000..dc2fa38
--- /dev/null
+++ b/web/public/install.sh
@@ -0,0 +1,13 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+INSTALL_URL="https://raw.githubusercontent.com/Hqzdev/pilot-agent/main/install.sh"
+
+if command -v curl >/dev/null 2>&1; then
+  curl -fsSL "$INSTALL_URL" | bash -s -- "$@"
+elif command -v wget >/dev/null 2>&1; then
+  wget -qO- "$INSTALL_URL" | bash -s -- "$@"
+else
+  echo "pilot-agent installer requires curl or wget." >&2
+  exit 1
+fi