diff --git a/gitpilot/agent/approvals.py b/gitpilot/agent/approvals.py index 6c4f675..93b32a1 100644 --- a/gitpilot/agent/approvals.py +++ b/gitpilot/agent/approvals.py @@ -30,6 +30,8 @@ from dataclasses import dataclass, field from typing import Any, Dict, Optional, Set +from ..toolkit.registry import Effect + logger = logging.getLogger(__name__) #: How long a request waits before it is treated as a refusal. Per-topology in @@ -40,6 +42,15 @@ SCOPE_ONCE = "once" SCOPE_SESSION = "session" +#: Side effects that must always be approved one operation at a time. A session +#: grant is convenient for local, reversible edits; it is too broad for actions +#: another person can observe or that change an external system. +_ONE_SHOT_EFFECTS = frozenset({ + Effect.GIT_REMOTE, + Effect.FORGE_WRITE, + Effect.EXTERNAL_WRITE, +}) + @dataclass class PendingApproval: @@ -52,6 +63,7 @@ class PendingApproval: risk: str reason: str command_class: str = "" + allow_session: bool = True #: Always set by :meth:`ApprovalRegistry.register`; optional only so the #: dataclass can be constructed in a test without an event loop. future: Optional["asyncio.Future[Dict[str, Any]]"] = field(repr=False, default=None) @@ -65,6 +77,8 @@ def to_dict(self) -> Dict[str, Any]: "risk": self.risk, "reason": self.reason, "command_class": self.command_class, + "allow_session": self.allow_session, + "allowed_scopes": [SCOPE_ONCE, SCOPE_SESSION] if self.allow_session else [SCOPE_ONCE], } @@ -92,12 +106,13 @@ def register( risk: str = "approval", reason: str = "", command_class: str = "", + allow_session: bool = True, ) -> PendingApproval: future: "asyncio.Future[Dict[str, Any]]" = asyncio.get_running_loop().create_future() pending = PendingApproval( request_id=request_id, session_id=session_id, tool=tool, arguments=dict(arguments or {}), risk=risk, reason=reason, - command_class=command_class, future=future, + command_class=command_class, allow_session=allow_session, future=future, ) self._pending[request_id] = pending self._by_session.setdefault(session_id, set()).add(request_id) @@ -133,12 +148,19 @@ def resolve(self, request_id: str, approved: bool, scope: str = SCOPE_ONCE) -> b """Answer a pending request. Returns whether there was one. Idempotent: a duplicate answer (the user double-clicks, or two transports - both deliver it) is dropped rather than raising. + both deliver it) is dropped rather than raising. A client cannot elevate + a one-shot external approval into a session-wide grant: unsupported + ``session`` scope is deterministically reduced to ``once``. """ pending = self._pending.get(request_id) if pending is None or pending.future is None or pending.future.done(): return False - pending.future.set_result({"approved": bool(approved), "scope": scope}) + effective_scope = ( + SCOPE_SESSION + if scope == SCOPE_SESSION and pending.allow_session + else SCOPE_ONCE + ) + pending.future.set_result({"approved": bool(approved), "scope": effective_scope}) return True def deny_session(self, session_id: str, *, reason: str = "session closed") -> int: @@ -159,7 +181,7 @@ def discard(self, request_id: str) -> None: if pending is not None: requests = self._by_session.get(pending.session_id) if requests is not None: - requests.discard(request_id) + requests.discard(pending.request_id) if not requests: self._by_session.pop(pending.session_id, None) @@ -198,6 +220,38 @@ def reset_registry() -> None: _registry = None +def _allows_session_scope(call: Any, ctx: Any) -> bool: + """Session grants are only for local/reversible effects. + + The registry is already attached to the execution context by the runner. If + anything about the spec cannot be established, fail closed to one-shot: a + missing registry must never widen an approval. + """ + tool_context = getattr(ctx, "tool_context", None) + extras = getattr(tool_context, "extras", {}) or {} + registry = extras.get("registry") if isinstance(extras, dict) else None + if registry is None: + return False + try: + spec = registry.spec(call.tool) + except Exception: # noqa: BLE001 - unknown ⇒ one-shot is the safe direction + return False + + if spec.effects & _ONE_SHOT_EFFECTS: + return False + + # ``fs.write``/``fs.edit`` share one spec between local and GitHub-backed + # workspaces. When the target is a remote repo rather than a local checkout, + # treat filesystem writes as external and keep them one-shot as well. + if Effect.WRITES_FS in spec.effects: + workspace = getattr(tool_context, "workspace", None) + repo = getattr(tool_context, "repo", None) + if workspace is None and repo is not None: + return False + + return True + + # --------------------------------------------------------------------------- # The loop's approver # --------------------------------------------------------------------------- @@ -241,6 +295,7 @@ async def request(self, call: Any, decision: Any, ctx: Any) -> bool: risk=getattr(decision, "risk", "approval"), reason=getattr(decision, "reason", ""), command_class=getattr(decision, "command_class", ""), + allow_session=_allows_session_scope(call, ctx), ) if self.gate is not None: # Let the gate's own pending map see it too, so a WS client calling diff --git a/gitpilot/agent_tools.py b/gitpilot/agent_tools.py index 019ce74..a539e17 100644 --- a/gitpilot/agent_tools.py +++ b/gitpilot/agent_tools.py @@ -12,6 +12,7 @@ from .glob_match import GLOB_DEFAULT_MAX_RESULTS as _GLOB_DEFAULT_MAX_RESULTS from .glob_match import GLOB_HARD_MAX_RESULTS as _GLOB_HARD_MAX_RESULTS from .glob_match import glob_match, glob_to_regex +from .idempotency import run_idempotent_mutation def _sanitize_tool_arg(value: Any, fallback_key: str = "description") -> str: @@ -403,14 +404,13 @@ def edit_file( old_string: Any, new_string: Any, commit_message: Any, + idempotency_key: str, expected_occurrences: Any = 1, ) -> str: - """Surgical edit — replace a small section of a file without - re-emitting the rest. Use this whenever you want to fix a bug, - rename a symbol, or insert a few lines into a file that already - exists. Never use ``Write or update a file`` to apply a small - change — that requires re-emitting the whole file and corrupts - long files on small-context models. + """Surgical approved edit with replay protection. + + idempotency_key is the stable approval/request id supplied by the runtime. + Reuse the same key for retries of the exact same edit. file_path: path relative to the repo root. Plain string. old_string: the exact text to find — including surrounding @@ -440,29 +440,52 @@ def edit_file( try: owner, repo, token, branch = get_repo_context() - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - current = loop.run_until_complete( - get_file(owner, repo, file_path, token=token, ref=branch) - ) - new_content, report = apply_edit( - current or "", - old_string=old_string_s, - new_string=new_string_s, - expected_occurrences=expected, - ) - result = loop.run_until_complete( - put_file(owner, repo, file_path, new_content, commit_message_s, token=token, branch=branch) - ) - finally: - loop.close() + args = { + "file_path": file_path, + "old_string": old_string_s, + "new_string": new_string_s, + "commit_message": commit_message_s, + "expected_occurrences": expected, + "branch": branch, + } + def _operation(): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + current = loop.run_until_complete( + get_file(owner, repo, file_path, token=token, ref=branch) + ) + new_content, report = apply_edit( + current or "", + old_string=old_string_s, + new_string=new_string_s, + expected_occurrences=expected, + ) + result = loop.run_until_complete( + put_file(owner, repo, file_path, new_content, commit_message_s, token=token, branch=branch) + ) + return { + "result": result, + "occurrences_replaced": report.occurrences_replaced, + "bytes_before": report.bytes_before, + "bytes_after": report.bytes_after, + } + finally: + loop.close() + + outcome = run_idempotent_mutation( + scope=f"github.file.edit:{owner}/{repo}:{file_path}", + idempotency_key=idempotency_key, + arguments=args, + operation=_operation, + ) + result = outcome["result"] sha = result.get("commit_sha", "") return ( f"File '{file_path}' edited " - f"({report.occurrences_replaced} occurrence(s) replaced, " - f"{report.bytes_before} → {report.bytes_after} bytes). " + f"({outcome['occurrences_replaced']} occurrence(s) replaced, " + f"{outcome['bytes_before']} → {outcome['bytes_after']} bytes). " f"Commit: {sha[:8]}" ) except EditError as e: @@ -478,11 +501,13 @@ def apply_patch_to_file( file_path: Any, diff: Any, commit_message: Any, + idempotency_key: str, ) -> str: - """Apply a unified-diff patch to a single file. Use this when the - change involves several non-contiguous edits inside one file and - a single ``Edit a section of a file`` call wouldn't capture all - of them cleanly. + """Apply one approved unified-diff patch to a single file. + + idempotency_key is the stable approval/request id supplied by the runtime. + Use this when the change involves several non-contiguous edits inside one + file and a single ``Edit a section of a file`` call would not capture them. file_path: path relative to the repo root. diff: a single-file unified diff with one or more @@-hunks. The @@ -502,24 +527,45 @@ def apply_patch_to_file( try: owner, repo, token, branch = get_repo_context() - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - current = loop.run_until_complete( - get_file(owner, repo, file_path, token=token, ref=branch) - ) - new_content, report = apply_unified_diff(current or "", diff_s) - result = loop.run_until_complete( - put_file(owner, repo, file_path, new_content, commit_message_s, token=token, branch=branch) - ) - finally: - loop.close() + args = { + "file_path": file_path, + "diff": diff_s, + "commit_message": commit_message_s, + "branch": branch, + } + def _operation(): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + current = loop.run_until_complete( + get_file(owner, repo, file_path, token=token, ref=branch) + ) + new_content, report = apply_unified_diff(current or "", diff_s) + result = loop.run_until_complete( + put_file(owner, repo, file_path, new_content, commit_message_s, token=token, branch=branch) + ) + return { + "result": result, + "occurrences_replaced": report.occurrences_replaced, + "bytes_before": report.bytes_before, + "bytes_after": report.bytes_after, + } + finally: + loop.close() + + outcome = run_idempotent_mutation( + scope=f"github.file.patch:{owner}/{repo}:{file_path}", + idempotency_key=idempotency_key, + arguments=args, + operation=_operation, + ) + result = outcome["result"] sha = result.get("commit_sha", "") return ( f"File '{file_path}' patched " - f"({report.occurrences_replaced} hunk(s) applied, " - f"{report.bytes_before} → {report.bytes_after} bytes). " + f"({outcome['occurrences_replaced']} hunk(s) applied, " + f"{outcome['bytes_before']} → {outcome['bytes_after']} bytes). " f"Commit: {sha[:8]}" ) except EditError as e: @@ -529,9 +575,10 @@ def apply_patch_to_file( @tool("Write or update a file in the repository") -def write_file(file_path: Any, content: Any, commit_message: Any) -> str: - """Create or update a file in the repository. +def write_file(file_path: Any, content: Any, commit_message: Any, idempotency_key: str) -> str: + """Create or update one approved file in the repository. + idempotency_key is the stable approval/request id supplied by the runtime. file_path: path relative to the repo root (plain string, e.g. ``"src/main.py"``). content: the full new file content (plain string). commit_message: a short imperative commit summary. Do @@ -544,15 +591,27 @@ def write_file(file_path: Any, content: Any, commit_message: Any) -> str: owner, repo, token, branch = get_repo_context() from .github_api import put_file - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - result = loop.run_until_complete( - put_file(owner, repo, file_path, content, commit_message, token=token, branch=branch) - ) - finally: - loop.close() - + def _operation(): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + return loop.run_until_complete( + put_file(owner, repo, file_path, content, commit_message, token=token, branch=branch) + ) + finally: + loop.close() + + result = run_idempotent_mutation( + scope=f"github.file.write:{owner}/{repo}:{file_path}", + idempotency_key=idempotency_key, + arguments={ + "file_path": file_path, + "content": content, + "commit_message": commit_message, + "branch": branch, + }, + operation=_operation, + ) sha = result.get("commit_sha", "") return f"File '{file_path}' written successfully. Commit: {sha[:8]}" except Exception as e: @@ -560,9 +619,10 @@ def write_file(file_path: Any, content: Any, commit_message: Any) -> str: @tool("Delete a file from the repository") -def delete_repo_file(file_path: Any, commit_message: Any) -> str: - """Delete a file from the repository. +def delete_repo_file(file_path: Any, commit_message: Any, idempotency_key: str) -> str: + """Delete one approved file from the repository. + idempotency_key is the stable approval/request id supplied by the runtime. file_path: the path relative to the repo root (plain string, e.g. ``"docs/old.md"``). commit_message: a short imperative commit summary. Both are plain strings — never wrap them in a schema dict. @@ -573,15 +633,26 @@ def delete_repo_file(file_path: Any, commit_message: Any) -> str: owner, repo, token, branch = get_repo_context() from .github_api import delete_file - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - result = loop.run_until_complete( - delete_file(owner, repo, file_path, commit_message, token=token, branch=branch) - ) - finally: - loop.close() - + def _operation(): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + return loop.run_until_complete( + delete_file(owner, repo, file_path, commit_message, token=token, branch=branch) + ) + finally: + loop.close() + + result = run_idempotent_mutation( + scope=f"github.file.delete:{owner}/{repo}:{file_path}", + idempotency_key=idempotency_key, + arguments={ + "file_path": file_path, + "commit_message": commit_message, + "branch": branch, + }, + operation=_operation, + ) sha = result.get("commit_sha", "") return f"File '{file_path}' deleted. Commit: {sha[:8]}" except Exception as e: @@ -589,22 +660,32 @@ def delete_repo_file(file_path: Any, commit_message: Any) -> str: @tool("Create a new branch in the repository") -def create_repo_branch(branch_name: str) -> str: - """Creates a new branch from the current HEAD.""" +def create_repo_branch(branch_name: str, idempotency_key: str) -> str: + """Create one approved branch from the current HEAD. + + idempotency_key is the stable approval/request id supplied by the runtime. + """ branch_name = _sanitize_tool_arg(branch_name) try: owner, repo, token, _branch = get_repo_context() from .github_api import create_branch - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - loop.run_until_complete( - create_branch(owner, repo, branch_name, from_ref="HEAD", token=token) - ) - finally: - loop.close() - + def _operation(): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + return loop.run_until_complete( + create_branch(owner, repo, branch_name, from_ref="HEAD", token=token) + ) + finally: + loop.close() + + run_idempotent_mutation( + scope=f"github.branch.create:{owner}/{repo}:{branch_name}", + idempotency_key=idempotency_key, + arguments={"branch_name": branch_name, "from_ref": "HEAD"}, + operation=_operation, + ) return f"Branch '{branch_name}' created successfully." except Exception as e: if "already exists" in str(e).lower() or "422" in str(e): @@ -821,4 +902,4 @@ def list_files_alias() -> str: write_file, delete_repo_file, create_repo_branch, -] +] \ No newline at end of file diff --git a/gitpilot/idempotency.py b/gitpilot/idempotency.py new file mode 100644 index 0000000..373f406 --- /dev/null +++ b/gitpilot/idempotency.py @@ -0,0 +1,374 @@ +"""Durable idempotency for mutating agent tools. + +Human approval answers *whether* a side effect may happen. Idempotency answers a +different production question: what happens when the approved call is delivered +again because a model, worker, network, or orchestrator retries it? + +GitPilot uses a tiny SQLite ledger because it is local, transactional, requires no +new service, and survives process restarts. Reads never touch this module. A +mutation reserves ``(scope, idempotency_key)`` before the side effect, binds the +key to a canonical hash of the exact arguments, and stores the successful result. +A retry with the same key and arguments replays the stored result without calling +the downstream service again. + +There is deliberately no automatic retry after a process dies or an exception is +raised while a mutation is in flight. Many downstream APIs (including GitHub's +create-issue / create-PR APIs) do not accept an idempotency header, so after a +lost response the outcome is unknowable. The safe industry pattern is to mark +that key indeterminate and require reconciliation/new approval rather than risk a +duplicate externally-visible action. +""" +from __future__ import annotations + +import hashlib +import json +import os +import sqlite3 +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Awaitable, Callable, Generic, Optional, TypeVar + +T = TypeVar("T") + +DEFAULT_DB_PATH = Path.home() / ".gitpilot" / "idempotency.sqlite3" +MAX_KEY_LENGTH = 256 +PENDING_FRESH_SECONDS = 300.0 + + +class IdempotencyError(RuntimeError): + """Base class for safe-to-surface idempotency failures.""" + + +class IdempotencyConflict(IdempotencyError): + """A key was reused for a different operation or argument set.""" + + +class IdempotencyInProgress(IdempotencyError): + """Another worker appears to be executing this key right now.""" + + +class IdempotencyIndeterminate(IdempotencyError): + """A previous execution may have committed but its result was lost.""" + + +@dataclass(frozen=True) +class Reservation(Generic[T]): + execute: bool + result: Optional[T] = None + + +def _json_default(value: Any) -> str: + """Stable fallback for values such as ``Path`` without storing secrets.""" + if isinstance(value, Path): + return str(value) + return repr(value) + + +def canonical_fingerprint(scope: str, arguments: Any) -> str: + """Hash a mutation's identity without persisting its raw arguments.""" + payload = json.dumps( + {"scope": scope, "arguments": arguments}, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + default=_json_default, + ).encode("utf-8") + return hashlib.sha256(payload).hexdigest() + + +def _validate_key(key: str) -> str: + value = str(key or "").strip() + if not value: + raise IdempotencyError( + "idempotency_key is required for mutating tools; use the stable " + "approval/request id and reuse it for retries" + ) + if len(value) > MAX_KEY_LENGTH: + raise IdempotencyError( + f"idempotency_key is too long ({len(value)} > {MAX_KEY_LENGTH})" + ) + return value + + +class IdempotencyStore: + """Small durable execution ledger backed by SQLite. + + Each method opens a short-lived connection. Mutations are rare compared with + reads, so this avoids shared-connection/thread hazards while WAL mode and + ``BEGIN IMMEDIATE`` make concurrent reservations deterministic. + """ + + def __init__(self, path: Optional[Path | str] = None) -> None: + configured = os.getenv("GITPILOT_IDEMPOTENCY_DB") + self.path = Path(path or configured or DEFAULT_DB_PATH).expanduser() + + def _connect(self) -> sqlite3.Connection: + self.path.parent.mkdir(parents=True, exist_ok=True) + conn = sqlite3.connect(str(self.path), timeout=5.0) + conn.row_factory = sqlite3.Row + conn.execute("PRAGMA journal_mode=WAL") + conn.execute("PRAGMA synchronous=FULL") + conn.execute("PRAGMA busy_timeout=5000") + conn.execute( + """ + CREATE TABLE IF NOT EXISTS mutation_idempotency ( + scope TEXT NOT NULL, + idempotency_key TEXT NOT NULL, + fingerprint TEXT NOT NULL, + status TEXT NOT NULL, + result_json TEXT, + error TEXT, + updated_at REAL NOT NULL, + PRIMARY KEY (scope, idempotency_key) + ) + """ + ) + return conn + + def reserve( + self, + *, + scope: str, + idempotency_key: str, + arguments: Any, + ) -> Reservation[Any]: + key = _validate_key(idempotency_key) + fingerprint = canonical_fingerprint(scope, arguments) + now = time.time() + + with self._connect() as conn: + conn.execute("BEGIN IMMEDIATE") + row = conn.execute( + "SELECT fingerprint, status, result_json, error, updated_at " + "FROM mutation_idempotency WHERE scope = ? AND idempotency_key = ?", + (scope, key), + ).fetchone() + + if row is None: + conn.execute( + "INSERT INTO mutation_idempotency " + "(scope, idempotency_key, fingerprint, status, updated_at) " + "VALUES (?, ?, ?, 'pending', ?)", + (scope, key, fingerprint, now), + ) + conn.commit() + return Reservation(execute=True) + + if row["fingerprint"] != fingerprint: + raise IdempotencyConflict( + "idempotency_key was already used for different arguments; " + "request a new approval instead of reusing the key" + ) + + status = str(row["status"]) + if status == "completed": + raw = row["result_json"] + return Reservation( + execute=False, + result=json.loads(raw) if raw is not None else None, + ) + + if status == "indeterminate": + detail = f" ({row['error']})" if row["error"] else "" + raise IdempotencyIndeterminate( + "a previous attempt may have committed but its response was lost" + f"{detail}; reconcile the target before requesting a new approval" + ) + + age = max(0.0, now - float(row["updated_at"] or now)) + if age <= PENDING_FRESH_SECONDS: + raise IdempotencyInProgress( + "this approved mutation is already executing; do not issue it again" + ) + raise IdempotencyIndeterminate( + "a previous execution started but never recorded a result; " + "reconcile the target before requesting a new approval" + ) + + def complete( + self, + *, + scope: str, + idempotency_key: str, + arguments: Any, + result: Any, + ) -> None: + key = _validate_key(idempotency_key) + fingerprint = canonical_fingerprint(scope, arguments) + encoded = json.dumps(result, ensure_ascii=False, default=_json_default) + with self._connect() as conn: + cursor = conn.execute( + "UPDATE mutation_idempotency SET status = 'completed', result_json = ?, " + "error = NULL, updated_at = ? WHERE scope = ? AND idempotency_key = ? " + "AND fingerprint = ?", + (encoded, time.time(), scope, key, fingerprint), + ) + if cursor.rowcount != 1: + raise IdempotencyConflict( + "could not complete idempotency record because its arguments changed" + ) + + def mark_indeterminate( + self, + *, + scope: str, + idempotency_key: str, + arguments: Any, + error: BaseException, + ) -> None: + key = _validate_key(idempotency_key) + fingerprint = canonical_fingerprint(scope, arguments) + message = f"{type(error).__name__}: {error}"[:500] + with self._connect() as conn: + conn.execute( + "UPDATE mutation_idempotency SET status = 'indeterminate', error = ?, " + "updated_at = ? WHERE scope = ? AND idempotency_key = ? AND fingerprint = ?", + (message, time.time(), scope, key, fingerprint), + ) + + def run_once( + self, + *, + scope: str, + idempotency_key: str, + arguments: Any, + operation: Callable[[], T], + ) -> T: + reservation = self.reserve( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + ) + if not reservation.execute: + return reservation.result # type: ignore[return-value] + + try: + result = operation() + except BaseException as exc: + # Do not silently retry an externally-visible side effect after an + # ambiguous failure. The caller surfaces the original error; a + # later retry of the same key gets the clearer indeterminate message. + self.mark_indeterminate( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + error=exc, + ) + raise + + self.complete( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + result=result, + ) + return result + + async def run_once_async( + self, + *, + scope: str, + idempotency_key: str, + arguments: Any, + operation: Callable[[], Awaitable[T]], + ) -> T: + """Async twin used by the production runtime registry.""" + reservation = self.reserve( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + ) + if not reservation.execute: + return reservation.result # type: ignore[return-value] + + try: + result = await operation() + except BaseException as exc: + self.mark_indeterminate( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + error=exc, + ) + raise + + self.complete( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + result=result, + ) + return result + + +_default_store: Optional[IdempotencyStore] = None + + +def get_idempotency_store() -> IdempotencyStore: + global _default_store + if _default_store is None: + _default_store = IdempotencyStore() + return _default_store + + +def run_idempotent_mutation( + *, + scope: str, + idempotency_key: str, + arguments: Any, + operation: Callable[[], T], +) -> T: + """Execute one approved mutation at most once for a stable request key.""" + return get_idempotency_store().run_once( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + operation=operation, + ) + + +def run_legacy_mutation( + *, + scope: str, + idempotency_key: str = "", + arguments: Any, + operation: Callable[[], T], +) -> T: + """Compatibility bridge for pre-V4 CrewAI mutators. + + Historical CrewAI tools predate the approval/request-id plumbing and are also + called directly by integrations and parity tests. Breaking those signatures + would force callers—or worse, small models—to fabricate operational IDs. + + When the legacy caller supplies the real approval/request key, use the same + durable ledger as the V4 runtime. When it does not, preserve the historical + one-shot behavior. The production V4 runtime never uses this fallback: its + :class:`RuntimeToolRegistry` always has the canonical call id and enforces + durable replay protection there. + """ + key = str(idempotency_key or "").strip() + if not key: + return operation() + return run_idempotent_mutation( + scope=scope, + idempotency_key=key, + arguments=arguments, + operation=operation, + ) + + +async def run_idempotent_mutation_async( + *, + scope: str, + idempotency_key: str, + arguments: Any, + operation: Callable[[], Awaitable[T]], +) -> T: + """Async execution path for canonical agent tools.""" + return await get_idempotency_store().run_once_async( + scope=scope, + idempotency_key=idempotency_key, + arguments=arguments, + operation=operation, + ) diff --git a/gitpilot/issue_tools.py b/gitpilot/issue_tools.py index 3850674..10fec6e 100644 --- a/gitpilot/issue_tools.py +++ b/gitpilot/issue_tools.py @@ -11,6 +11,7 @@ from .agent_tools import get_repo_context from . import github_issues as gi +from .idempotency import run_legacy_mutation def _run_async(coro): @@ -78,14 +79,34 @@ def create_issue( body: str = "", labels: str = "", assignees: str = "", + idempotency_key: str = "", ) -> str: - """Creates a new GitHub issue. labels and assignees are comma-separated strings.""" + """Creates a new GitHub issue. + + Existing callers keep the historical argument order. When supplied, + ``idempotency_key`` must be the stable approval/request id and makes retries + of this exact action replay-safe. + """ try: owner, repo, token, _branch = get_repo_context() label_list = [l.strip() for l in labels.split(",") if l.strip()] if labels else None assignee_list = [a.strip() for a in assignees.split(",") if a.strip()] if assignees else None - issue = _run_async( - gi.create_issue(owner, repo, title, body=body or None, labels=label_list, assignees=assignee_list, token=token) + args = { + "title": title, + "body": body, + "labels": label_list, + "assignees": assignee_list, + } + issue = run_legacy_mutation( + scope=f"github.issue.create:{owner}/{repo}", + idempotency_key=idempotency_key, + arguments=args, + operation=lambda: _run_async( + gi.create_issue( + owner, repo, title, body=body or None, + labels=label_list, assignees=assignee_list, token=token, + ) + ), ) return f"Created issue #{issue.get('number')}: {issue.get('title')}\nURL: {issue.get('html_url', '')}" except Exception as e: @@ -100,8 +121,14 @@ def update_issue( state: str = "", labels: str = "", assignees: str = "", + idempotency_key: str = "", ) -> str: - """Updates an existing issue. Only non-empty fields are changed. labels/assignees are comma-separated.""" + """Updates an existing issue. + + Existing callers keep the historical argument order. When supplied, + ``idempotency_key`` is the stable approval/request id. Only non-empty fields + are changed; labels/assignees are comma-separated. + """ try: owner, repo, token, _branch = get_repo_context() kwargs: dict = {} @@ -115,18 +142,39 @@ def update_issue( kwargs["labels"] = [l.strip() for l in labels.split(",") if l.strip()] if assignees: kwargs["assignees"] = [a.strip() for a in assignees.split(",") if a.strip()] - issue = _run_async(gi.update_issue(owner, repo, issue_number, token=token, **kwargs)) + issue = run_legacy_mutation( + scope=f"github.issue.update:{owner}/{repo}:{issue_number}", + idempotency_key=idempotency_key, + arguments=kwargs, + operation=lambda: _run_async( + gi.update_issue(owner, repo, issue_number, token=token, **kwargs) + ), + ) return f"Updated issue #{issue.get('number')}: {issue.get('title')}\nState: {issue.get('state')}" except Exception as e: return f"Error updating issue: {e}" @tool("Add a comment to an issue") -def add_issue_comment(issue_number: int, body: str) -> str: - """Adds a comment to an existing issue.""" +def add_issue_comment( + issue_number: int, + body: str, + idempotency_key: str = "", +) -> str: + """Adds a comment to an existing issue. + + Supply the stable approval/request id to make retries replay-safe. + """ try: owner, repo, token, _branch = get_repo_context() - comment = _run_async(gi.add_issue_comment(owner, repo, issue_number, body, token=token)) + comment = run_legacy_mutation( + scope=f"github.issue.comment:{owner}/{repo}:{issue_number}", + idempotency_key=idempotency_key, + arguments={"body": body}, + operation=lambda: _run_async( + gi.add_issue_comment(owner, repo, issue_number, body, token=token) + ), + ) return f"Comment added to issue #{issue_number}\nURL: {comment.get('html_url', '')}" except Exception as e: return f"Error adding comment: {e}" diff --git a/gitpilot/local_tools.py b/gitpilot/local_tools.py index ef014a0..c30ac37 100644 --- a/gitpilot/local_tools.py +++ b/gitpilot/local_tools.py @@ -12,6 +12,7 @@ from crewai.tools import tool +from .idempotency import run_legacy_mutation from .sandbox_routing import ( SandboxFallback, coerce_timeout, @@ -69,22 +70,45 @@ def read_local_file(file_path: str) -> str: @tool("Write local file") -def write_local_file(file_path: str, content: str) -> str: - """Write content to a file in the local workspace. Creates parent directories.""" +def write_local_file( + file_path: str, + content: str, + idempotency_key: str = "", +) -> str: + """Write content to a local file, creating parent directories. + + Legacy callers may omit ``idempotency_key`` for backward compatibility. + When the runtime supplies its stable approval/request id, retries of that + exact write are durably deduplicated. + """ ws = _require_workspace() try: - result = _run_async(_ws_manager.write_file(ws, file_path, content)) + result = run_legacy_mutation( + scope=f"local.file.write:{ws.path}", + idempotency_key=idempotency_key, + arguments={"file_path": file_path, "content": content}, + operation=lambda: _run_async(_ws_manager.write_file(ws, file_path, content)), + ) return f"Written {result['size']} bytes to {result['path']}" except Exception as e: return f"Error writing {file_path}: {e}" @tool("Delete local file") -def delete_local_file(file_path: str) -> str: - """Delete a file from the local workspace.""" +def delete_local_file(file_path: str, idempotency_key: str = "") -> str: + """Delete a file from the local workspace. + + With an approval/request id, retries replay the first result. Historical + direct callers may omit the key and retain their original one-shot behavior. + """ ws = _require_workspace() try: - deleted = _run_async(_ws_manager.delete_file(ws, file_path)) + deleted = run_legacy_mutation( + scope=f"local.file.delete:{ws.path}", + idempotency_key=idempotency_key, + arguments={"file_path": file_path}, + operation=lambda: _run_async(_ws_manager.delete_file(ws, file_path)), + ) return f"Deleted: {deleted}" except Exception as e: return f"Error deleting {file_path}: {e}" diff --git a/gitpilot/pr_tools.py b/gitpilot/pr_tools.py index c644070..5583d21 100644 --- a/gitpilot/pr_tools.py +++ b/gitpilot/pr_tools.py @@ -9,6 +9,7 @@ from .agent_tools import get_repo_context from . import github_pulls as gp +from .idempotency import run_legacy_mutation def _run_async(coro): @@ -74,15 +75,27 @@ def create_pull_request( base: str, body: str = "", draft: bool = False, + idempotency_key: str = "", ) -> str: - """Creates a new pull request. head=source branch, base=target branch.""" + """Creates a pull request. head=source branch, base=target branch. + + Existing callers keep the historical argument order. When supplied, + ``idempotency_key`` is the stable approval/request id and makes retries of + this exact action replay-safe. + """ try: owner, repo, token, _branch = get_repo_context() - pr = _run_async( - gp.create_pull_request( - owner, repo, title=title, head=head, base=base, - body=body or None, draft=draft, token=token, - ) + args = {"title": title, "head": head, "base": base, "body": body, "draft": draft} + pr = run_legacy_mutation( + scope=f"github.pr.create:{owner}/{repo}", + idempotency_key=idempotency_key, + arguments=args, + operation=lambda: _run_async( + gp.create_pull_request( + owner, repo, title=title, head=head, base=base, + body=body or None, draft=draft, token=token, + ) + ), ) return ( f"Created PR #{pr.get('number')}: {pr.get('title')}\n" @@ -97,17 +110,32 @@ def merge_pull_request( pull_number: int, merge_method: str = "merge", commit_title: str = "", + idempotency_key: str = "", ) -> str: - """Merges a pull request. merge_method: merge, squash, or rebase.""" + """Merges a pull request. merge_method: merge, squash, or rebase. + + Existing callers keep the historical argument order. Supply the stable + approval/request id to make retries replay-safe. + """ try: owner, repo, token, _branch = get_repo_context() - result = _run_async( - gp.merge_pull_request( - owner, repo, pull_number, - merge_method=merge_method, - commit_title=commit_title or None, - token=token, - ) + args = { + "pull_number": pull_number, + "merge_method": merge_method, + "commit_title": commit_title, + } + result = run_legacy_mutation( + scope=f"github.pr.merge:{owner}/{repo}:{pull_number}", + idempotency_key=idempotency_key, + arguments=args, + operation=lambda: _run_async( + gp.merge_pull_request( + owner, repo, pull_number, + merge_method=merge_method, + commit_title=commit_title or None, + token=token, + ) + ), ) sha = result.get("sha", "unknown") if isinstance(result, dict) else "unknown" return f"PR #{pull_number} merged successfully. Merge commit: {sha}" @@ -139,12 +167,24 @@ def create_pr_review( pull_number: int, body: str, event: str = "COMMENT", + idempotency_key: str = "", ) -> str: - """Adds a review to a PR. event: APPROVE, REQUEST_CHANGES, or COMMENT.""" + """Adds a review to a PR. event: APPROVE, REQUEST_CHANGES, or COMMENT. + + Existing callers keep the historical argument order. Supply the stable + approval/request id to make retries replay-safe. + """ try: owner, repo, token, _branch = get_repo_context() - review = _run_async( - gp.create_pr_review(owner, repo, pull_number, body=body, event=event, token=token) + review = run_legacy_mutation( + scope=f"github.pr.review:{owner}/{repo}:{pull_number}", + idempotency_key=idempotency_key, + arguments={"body": body, "event": event}, + operation=lambda: _run_async( + gp.create_pr_review( + owner, repo, pull_number, body=body, event=event, token=token + ) + ), ) return f"Review submitted on PR #{pull_number} (event={event})\nURL: {review.get('html_url', '')}" except Exception as e: @@ -152,11 +192,25 @@ def create_pr_review( @tool("Comment on a pull request") -def add_pr_comment(pull_number: int, body: str) -> str: - """Adds a general comment to a pull request.""" +def add_pr_comment( + pull_number: int, + body: str, + idempotency_key: str = "", +) -> str: + """Adds a general comment to a pull request. + + Supply the stable approval/request id to make retries replay-safe. + """ try: owner, repo, token, _branch = get_repo_context() - comment = _run_async(gp.add_pr_comment(owner, repo, pull_number, body, token=token)) + comment = run_legacy_mutation( + scope=f"github.pr.comment:{owner}/{repo}:{pull_number}", + idempotency_key=idempotency_key, + arguments={"body": body}, + operation=lambda: _run_async( + gp.add_pr_comment(owner, repo, pull_number, body, token=token) + ), + ) return f"Comment added to PR #{pull_number}\nURL: {comment.get('html_url', '')}" except Exception as e: return f"Error commenting on PR: {e}" diff --git a/gitpilot/toolkit/builtins.py b/gitpilot/toolkit/builtins.py index fbe763f..61a68c6 100644 --- a/gitpilot/toolkit/builtins.py +++ b/gitpilot/toolkit/builtins.py @@ -17,6 +17,7 @@ from . import forge, fs, git, terminal, testing, todo, web from .registry import ToolRegistry +from .runtime_registry import RuntimeToolRegistry #: Feature flag gating the *use* of this registry by execution paths. Building #: one is always safe and has no side effects. @@ -33,15 +34,19 @@ def build_default_registry( include_todo: bool = True, include_web: bool = True, ) -> ToolRegistry: - """Return a registry holding GitPilot's built-in tools. + """Return the replay-safe registry holding GitPilot's built-in tools. The include flags exist for tests and for callers that know a namespace cannot work in their context — there is no point offering ``git.*`` to a session with no checkout. They are *not* the permission mechanism: capability masks decide what a model may call (Batch V4-D1), and a tool withheld here is invisible to policy too. + + RuntimeToolRegistry behaves exactly like ToolRegistry for direct calls that + have no run/session identity; durable idempotency activates only inside real + agent runs, after the policy/approval boundary. """ - registry = ToolRegistry() + registry = RuntimeToolRegistry() if include_filesystem: fs.register(registry) if include_terminal: diff --git a/gitpilot/toolkit/runtime_registry.py b/gitpilot/toolkit/runtime_registry.py new file mode 100644 index 0000000..a3980e6 --- /dev/null +++ b/gitpilot/toolkit/runtime_registry.py @@ -0,0 +1,155 @@ +"""Production ToolRegistry with durable replay protection for mutations. + +The base :class:`ToolRegistry` remains a small, predictable library primitive. +This subclass adds the runtime-only contract GitPilot needs once a tool call +belongs to a real run/session: + +* safe/read-only tools keep the base registry's zero-ledger fast path; +* direct SDK/toolkit calls without a run/session keep byte-for-byte base + execution semantics; +* mutating runtime calls use the stable provider/tool-call id as the + idempotency key; +* the key is bound to the canonical tool id, exact arguments, run, and target; +* successful retries replay the original :class:`ToolResult` without repeating + side effects; +* ambiguous handler failures fail closed instead of automatically duplicating a + remote action whose response may simply have been lost. + +The important layering rule is that replay protection lives at the production +execution boundary, not inside every tool handler and not in the generic +registry. That keeps tests, SDK use, and read-heavy workloads inexpensive while +making the real agent path durable by construction. +""" +from __future__ import annotations + +from typing import Any, Dict, Mapping + +from ..idempotency import IdempotencyError, get_idempotency_store +from .registry import ( + Effect, + ToolCall, + ToolExecutionContext, + ToolRegistry, + ToolResult, + ToolSpec, + validate_arguments, +) + +_MUTATING_EFFECTS = frozenset({ + Effect.WRITES_FS, + Effect.GIT_LOCAL, + Effect.GIT_REMOTE, + Effect.FORGE_WRITE, + Effect.EXTERNAL_WRITE, +}) + + +def _mutating(spec: ToolSpec) -> bool: + return bool(spec.effects & _MUTATING_EFFECTS) + + +def _scope(ctx: ToolExecutionContext, spec: ToolSpec) -> str | None: + """Return a durable scope, or ``None`` for direct/ephemeral calls.""" + run = ctx.run_id or ctx.session_id + if not run: + return None + if ctx.workspace is not None: + target = f"workspace:{ctx.workspace.root}" + elif ctx.repo is not None: + target = f"repo:{ctx.repo.full_name}:{ctx.repo.branch or 'HEAD'}" + else: + target = "unbound" + return f"runtime:{run}:{target}:{spec.id}" + + +def _payload(result: ToolResult) -> Dict[str, Any]: + """JSON-friendly result persisted for exact replay after a restart.""" + return { + "tool": result.tool, + "ok": result.ok, + "content": result.content, + "data": result.data, + "error": result.error, + "denied": result.denied, + } + + +def _result_from_payload(call: ToolCall, payload: Mapping[str, Any]) -> ToolResult: + return ToolResult( + call_id=call.id, + tool=str(payload.get("tool") or call.tool), + ok=bool(payload.get("ok")), + content=str(payload.get("content") or ""), + data=dict(payload.get("data") or {}) or None, + error=payload.get("error"), + denied=bool(payload.get("denied")), + ) + + +class _MutationOutcomeUncertain(RuntimeError): + """Carry the first failure while making the ledger fail closed on retry.""" + + def __init__(self, result: ToolResult) -> None: + super().__init__(result.content) + self.result = result + + +class RuntimeToolRegistry(ToolRegistry): + """ToolRegistry whose mutating calls become replay-safe in real runs.""" + + async def execute( + self, + call: ToolCall, + ctx: ToolExecutionContext, + ) -> ToolResult: + # Unknown tools and invalid arguments are deterministic input failures, + # not attempted mutations. Let the base registry return its normal, + # corrective ToolResult without creating a ledger record. + resolved = self.resolve(call.tool) + if resolved is None: + return await super().execute(call, ctx) + + spec = self.spec(resolved) + if validate_arguments(spec.params_schema, call.arguments or {}): + return await super().execute(call, ctx) + + scope = _scope(ctx, spec) + if not _mutating(spec) or scope is None: + # This branch is deliberately the exact base execution path. In + # particular, standalone toolkit calls must not gain persistence, + # altered timeout semantics, or hidden disk I/O merely because the + # default registry happens to be runtime-capable. + return await super().execute(call, ctx) + + store = ctx.extras.get("idempotency_store") or get_idempotency_store() + identity = {"tool": resolved, "arguments": call.arguments or {}} + + async def operation() -> Dict[str, Any]: + result = await super(RuntimeToolRegistry, self).execute(call, ctx) + if not result.ok: + # A downstream failure can be ambiguous: the service may have + # committed the mutation and only lost the response. Raising + # here lets IdempotencyStore mark the key indeterminate. We then + # return the original ToolResult for this first attempt; only a + # retry is blocked pending reconciliation/new approval. + raise _MutationOutcomeUncertain(result) + return _payload(result) + + try: + stored = await store.run_once_async( + scope=scope, + idempotency_key=call.id, + arguments=identity, + operation=operation, + ) + except _MutationOutcomeUncertain as exc: + return exc.result + except IdempotencyError as exc: + return ToolResult.failure( + call, + f"{resolved} was not executed: {exc}", + error="idempotency_guard", + data={"retry_safe": False, "requires_reconciliation": True}, + ) + + return _result_from_payload(call, stored) diff --git a/tests/agent/test_approvals.py b/tests/agent/test_approvals.py index 1188e4f..a217414 100644 --- a/tests/agent/test_approvals.py +++ b/tests/agent/test_approvals.py @@ -2,6 +2,7 @@ from __future__ import annotations import asyncio +from types import SimpleNamespace import pytest @@ -12,7 +13,16 @@ reset_registry, ) from gitpilot.agent.policy import Decision -from gitpilot.toolkit import ToolCall +from gitpilot.toolkit import ( + Effect, + LocalWorkspace, + Risk, + ToolCall, + ToolExecutionContext, + ToolRegistry, + ToolResult, + ToolSpec, +) @pytest.fixture(autouse=True) @@ -26,6 +36,35 @@ def _call(tool="fs.write", **arguments): return ToolCall(id="req-1", tool=tool, arguments=arguments or {"path": "a.py"}) +def _ctx_for(tmp_path, tool: str, effects: frozenset[Effect]): + registry = ToolRegistry() + schema = { + "type": "object", + "properties": {"path": {"type": "string"}}, + "additionalProperties": True, + } + + async def noop(call, ctx): + return ToolResult.success(call, "ok") + + registry.register( + ToolSpec( + id=tool, + title=tool, + description="test", + params_schema=schema, + risk=Risk.APPROVAL, + effects=effects, + ), + noop, + ) + tool_context = ToolExecutionContext( + workspace=LocalWorkspace(root=tmp_path), + extras={"registry": registry}, + ) + return SimpleNamespace(tool_context=tool_context) + + # --------------------------------------------------------------------------- # The registry (F1) # --------------------------------------------------------------------------- @@ -96,6 +135,25 @@ async def drive(): first, second = asyncio.run(drive()) assert (first, second) == (True, False) + def test_one_shot_request_cannot_be_promoted_to_session(self): + async def drive(): + registry = ApprovalRegistry() + pending = registry.register( + request_id="r1", + session_id="s", + tool="github.pr.create", + allow_session=False, + ) + waiter = asyncio.create_task(registry.wait(pending, timeout=5)) + await asyncio.sleep(0) + registry.resolve("r1", True, "session") + return pending.to_dict(), await waiter + + payload, answer = asyncio.run(drive()) + assert payload["allow_session"] is False + assert payload["allowed_scopes"] == ["once"] + assert answer == {"approved": True, "scope": "once"} + def test_closing_a_session_denies_everything_outstanding(self): async def drive(): registry = ApprovalRegistry() @@ -145,30 +203,64 @@ async def drive(): assert asyncio.run(drive()) is True - def test_the_session_scope_is_recorded(self): - """"Allow for session" has to reach the policy engine, or the next call - asks again and the answer meant nothing.""" + def test_local_edit_session_scope_is_recorded(self, tmp_path): + """Local reversible work may still use Allow for session for good UX.""" granted = [] + ctx = _ctx_for(tmp_path, "fs.write", frozenset({Effect.WRITES_FS})) async def drive(): approver = self._approver(timeout_s=5, on_session_scope=granted.append) task = asyncio.create_task( - approver.request(_call(), Decision(verdict="ask"), None), + approver.request(_call(), Decision(verdict="ask"), ctx), ) await asyncio.sleep(0) + pending = get_registry().get("req-1") + assert pending is not None and pending.allow_session is True get_registry().resolve("req-1", True, "session") return await task assert asyncio.run(drive()) is True assert granted == ["fs.write"] - def test_a_once_scope_does_not_grant_the_session(self): + def test_external_write_is_always_one_shot(self, tmp_path): + """One click authorizes one externally visible side effect.""" granted = [] + ctx = _ctx_for( + tmp_path, + "github.pr.create", + frozenset({Effect.FORGE_WRITE, Effect.NETWORK}), + ) async def drive(): approver = self._approver(timeout_s=5, on_session_scope=granted.append) task = asyncio.create_task( - approver.request(_call(), Decision(verdict="ask"), None), + approver.request( + _call("github.pr.create", path="ignored"), + Decision(verdict="ask"), + ctx, + ), + ) + await asyncio.sleep(0) + pending = get_registry().get("req-1") + assert pending is not None + payload = pending.to_dict() + get_registry().resolve("req-1", True, "session") + result = await task + return payload, result + + payload, result = asyncio.run(drive()) + assert result is True + assert payload["allowed_scopes"] == ["once"] + assert granted == [] + + def test_a_once_scope_does_not_grant_the_session(self, tmp_path): + granted = [] + ctx = _ctx_for(tmp_path, "fs.write", frozenset({Effect.WRITES_FS})) + + async def drive(): + approver = self._approver(timeout_s=5, on_session_scope=granted.append) + task = asyncio.create_task( + approver.request(_call(), Decision(verdict="ask"), ctx), ) await asyncio.sleep(0) get_registry().resolve("req-1", True, "once") diff --git a/tests/agent/test_compaction.py b/tests/agent/test_compaction.py index 54466ad..b6730d8 100644 --- a/tests/agent/test_compaction.py +++ b/tests/agent/test_compaction.py @@ -2,6 +2,7 @@ from __future__ import annotations import asyncio +from dataclasses import replace import pytest @@ -216,9 +217,22 @@ def test_prose_without_paths_finds_nothing(self): class TestThroughTheLoop: - def _long_run(self, tmp_path, *, model="qwen2.5:1.5b", iterations=40): - """A run that would blow its window if nothing folded.""" + def _long_run( + self, + tmp_path, + *, + model="qwen2.5:1.5b", + iterations=40, + context_window=8_192, + ): + """A run that would blow its explicitly pinned window if nothing folded. + + This is a compaction stress test, not a model-catalog test. Pinning the + window keeps the scenario stable when a model's advertised context size + changes independently of compaction behavior. + """ profile = resolve_profile("ollama", model, cache_path=tmp_path / "c.json") + profile = replace(profile, context_window=context_window) registry = _registry() async def chatty(call, ctx): @@ -247,8 +261,7 @@ async def chatty(call, ctx): return result, ctx, events def test_a_forty_iteration_run_on_an_8k_window_completes(self, tmp_path): - """The case that motivates the whole batch: GitPilot's default model has - an 8k window and the loop can take forty turns.""" + """A constrained 8k run can still take forty turns via compaction.""" result, ctx, _events = self._long_run(tmp_path) assert result.status in ("completed", "degraded"), result.answer @@ -290,7 +303,10 @@ def test_the_state_change_is_visible(self, tmp_path): def test_a_large_window_never_compacts(self, tmp_path): """A run inside its budget should not pay for housekeeping.""" _result, _ctx, events = self._long_run( - tmp_path, model="llama3:8b", iterations=6, + tmp_path, + model="llama3:8b", + iterations=6, + context_window=32_768, ) assert "compaction" not in events.types() diff --git a/tests/test_idempotency.py b/tests/test_idempotency.py new file mode 100644 index 0000000..ecd4fc9 --- /dev/null +++ b/tests/test_idempotency.py @@ -0,0 +1,180 @@ +"""Production invariants for approved mutating tool retries.""" +from __future__ import annotations + +import ast +from pathlib import Path + +import pytest + +from gitpilot.idempotency import ( + IdempotencyConflict, + IdempotencyError, + IdempotencyIndeterminate, + IdempotencyInProgress, + IdempotencyStore, +) + + +def test_completed_mutation_replays_without_executing_twice(tmp_path): + store = IdempotencyStore(tmp_path / "idem.sqlite3") + calls: list[str] = [] + + def operation(): + calls.append("executed") + return {"number": 42, "url": "https://example.invalid/42"} + + first = store.run_once( + scope="github.issue.create:o/r", + idempotency_key="approval-123", + arguments={"title": "Bug"}, + operation=operation, + ) + second = store.run_once( + scope="github.issue.create:o/r", + idempotency_key="approval-123", + arguments={"title": "Bug"}, + operation=operation, + ) + + assert first == second == {"number": 42, "url": "https://example.invalid/42"} + assert calls == ["executed"] + + +def test_result_survives_a_new_store_instance(tmp_path): + path = tmp_path / "idem.sqlite3" + first_store = IdempotencyStore(path) + first_store.run_once( + scope="github.pr.create:o/r", + idempotency_key="approval-1", + arguments={"head": "fix", "base": "main"}, + operation=lambda: {"number": 9}, + ) + + restarted_store = IdempotencyStore(path) + called = False + + def must_not_run(): + nonlocal called + called = True + return {"number": 10} + + result = restarted_store.run_once( + scope="github.pr.create:o/r", + idempotency_key="approval-1", + arguments={"head": "fix", "base": "main"}, + operation=must_not_run, + ) + assert result == {"number": 9} + assert called is False + + +def test_same_key_cannot_authorize_changed_arguments(tmp_path): + store = IdempotencyStore(tmp_path / "idem.sqlite3") + store.run_once( + scope="github.issue.create:o/r", + idempotency_key="approval-1", + arguments={"title": "A"}, + operation=lambda: {"number": 1}, + ) + + with pytest.raises(IdempotencyConflict, match="different arguments"): + store.run_once( + scope="github.issue.create:o/r", + idempotency_key="approval-1", + arguments={"title": "B"}, + operation=lambda: {"number": 2}, + ) + + +def test_inflight_duplicate_is_not_executed(tmp_path): + store = IdempotencyStore(tmp_path / "idem.sqlite3") + store.reserve( + scope="github.pr.review:o/r:7", + idempotency_key="approval-7", + arguments={"event": "APPROVE"}, + ) + + with pytest.raises(IdempotencyInProgress, match="already executing"): + store.reserve( + scope="github.pr.review:o/r:7", + idempotency_key="approval-7", + arguments={"event": "APPROVE"}, + ) + + +def test_ambiguous_failure_fails_closed_on_retry(tmp_path): + store = IdempotencyStore(tmp_path / "idem.sqlite3") + calls = 0 + + def timeout_after_possible_commit(): + nonlocal calls + calls += 1 + raise TimeoutError("response lost") + + with pytest.raises(TimeoutError, match="response lost"): + store.run_once( + scope="github.issue.create:o/r", + idempotency_key="approval-timeout", + arguments={"title": "Maybe created"}, + operation=timeout_after_possible_commit, + ) + + with pytest.raises(IdempotencyIndeterminate, match="may have committed"): + store.run_once( + scope="github.issue.create:o/r", + idempotency_key="approval-timeout", + arguments={"title": "Maybe created"}, + operation=timeout_after_possible_commit, + ) + + assert calls == 1, "an ambiguous remote outcome must never be retried automatically" + + +def test_empty_key_is_rejected_before_side_effect(tmp_path): + store = IdempotencyStore(tmp_path / "idem.sqlite3") + called = False + + def operation(): + nonlocal called + called = True + + with pytest.raises(IdempotencyError, match="idempotency_key is required"): + store.run_once( + scope="github.issue.create:o/r", + idempotency_key="", + arguments={}, + operation=operation, + ) + assert called is False + + +def _function_args(path: Path, function_name: str) -> set[str]: + tree = ast.parse(path.read_text(encoding="utf-8")) + for node in ast.walk(tree): + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == function_name: + return {arg.arg for arg in [*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs]} + raise AssertionError(f"{function_name} not found in {path}") + + +@pytest.mark.parametrize( + ("relative_path", "function_name"), + [ + ("gitpilot/agent_tools.py", "edit_file"), + ("gitpilot/agent_tools.py", "apply_patch_to_file"), + ("gitpilot/agent_tools.py", "write_file"), + ("gitpilot/agent_tools.py", "delete_repo_file"), + ("gitpilot/agent_tools.py", "create_repo_branch"), + ("gitpilot/local_tools.py", "write_local_file"), + ("gitpilot/local_tools.py", "delete_local_file"), + ("gitpilot/issue_tools.py", "create_issue"), + ("gitpilot/issue_tools.py", "update_issue"), + ("gitpilot/issue_tools.py", "add_issue_comment"), + ("gitpilot/pr_tools.py", "create_pull_request"), + ("gitpilot/pr_tools.py", "merge_pull_request"), + ("gitpilot/pr_tools.py", "create_pr_review"), + ("gitpilot/pr_tools.py", "add_pr_comment"), + ], +) +def test_legacy_mutating_tools_require_runtime_idempotency_key(relative_path, function_name): + root = Path(__file__).resolve().parents[1] + assert "idempotency_key" in _function_args(root / relative_path, function_name) diff --git a/tests/test_issue_tools.py b/tests/test_issue_tools.py index 988397b..c8bc8ad 100644 --- a/tests/test_issue_tools.py +++ b/tests/test_issue_tools.py @@ -1,6 +1,7 @@ """Tests for gitpilot.issue_tools CrewAI tools.""" from __future__ import annotations +import uuid from unittest.mock import AsyncMock, patch import pytest @@ -16,6 +17,11 @@ ) +def _approval_key() -> str: + """Unique per test invocation so the durable ledger cannot cross-contaminate runs.""" + return f"test-approval-{uuid.uuid4()}" + + def test_issue_tools_exported(): """All 6 issue tools are exported.""" assert len(ISSUE_TOOLS) == 6 @@ -72,14 +78,18 @@ def test_creates_issue(self, mock_create, repo_context): "title": "New", "html_url": "https://github.com/o/r/issues/99", } - result = create_issue.run(title="New", body="Body text") + result = create_issue.run( + title="New", body="Body text", idempotency_key=_approval_key() + ) assert "#99" in result assert "Created" in result @patch("gitpilot.issue_tools.gi.create_issue", new_callable=AsyncMock) def test_parses_labels_csv(self, mock_create, repo_context): mock_create.return_value = {"number": 1, "title": "T", "html_url": ""} - create_issue.run(title="T", labels="bug, enhancement") + create_issue.run( + title="T", labels="bug, enhancement", idempotency_key=_approval_key() + ) call_kwargs = mock_create.call_args labels_arg = call_kwargs.kwargs.get("labels") or call_kwargs[1].get("labels") assert labels_arg == ["bug", "enhancement"] @@ -89,7 +99,9 @@ class TestUpdateIssueTool: @patch("gitpilot.issue_tools.gi.update_issue", new_callable=AsyncMock) def test_updates_issue(self, mock_update, repo_context): mock_update.return_value = {"number": 10, "title": "Fixed", "state": "closed"} - result = update_issue.run(issue_number=10, state="closed") + result = update_issue.run( + issue_number=10, state="closed", idempotency_key=_approval_key() + ) assert "Updated" in result @@ -97,7 +109,9 @@ class TestCommentTools: @patch("gitpilot.issue_tools.gi.add_issue_comment", new_callable=AsyncMock) def test_add_comment(self, mock_add, repo_context): mock_add.return_value = {"html_url": "https://github.com/o/r/issues/5#comment"} - result = add_issue_comment.run(issue_number=5, body="hello") + result = add_issue_comment.run( + issue_number=5, body="hello", idempotency_key=_approval_key() + ) assert "Comment added" in result @patch("gitpilot.issue_tools.gi.list_issue_comments", new_callable=AsyncMock) diff --git a/tests/test_pr_tools.py b/tests/test_pr_tools.py index 12be8a5..934040d 100644 --- a/tests/test_pr_tools.py +++ b/tests/test_pr_tools.py @@ -1,6 +1,7 @@ """Tests for gitpilot.pr_tools CrewAI tools.""" from __future__ import annotations +import uuid from unittest.mock import AsyncMock, patch import pytest @@ -17,6 +18,11 @@ ) +def _approval_key() -> str: + """Unique per test invocation so the durable ledger cannot cross-contaminate runs.""" + return f"test-approval-{uuid.uuid4()}" + + def test_pr_tools_exported(): """All 7 PR tools are exported.""" assert len(PR_TOOLS) == 7 @@ -79,7 +85,12 @@ def test_creates_pr(self, mock_create, repo_context): "title": "New Feature", "html_url": "https://github.com/o/r/pull/20", } - result = create_pull_request.run(title="New Feature", head="feat", base="main") + result = create_pull_request.run( + title="New Feature", + head="feat", + base="main", + idempotency_key=_approval_key(), + ) assert "Created PR #20" in result @@ -87,7 +98,9 @@ class TestMergePRTool: @patch("gitpilot.pr_tools.gp.merge_pull_request", new_callable=AsyncMock) def test_merges_pr(self, mock_merge, repo_context): mock_merge.return_value = {"sha": "abc123"} - result = merge_pull_request.run(pull_number=10) + result = merge_pull_request.run( + pull_number=10, idempotency_key=_approval_key() + ) assert "merged" in result.lower() assert "abc123" in result @@ -113,7 +126,12 @@ class TestReviewTool: @patch("gitpilot.pr_tools.gp.create_pr_review", new_callable=AsyncMock) def test_creates_review(self, mock_review, repo_context): mock_review.return_value = {"html_url": "https://github.com/o/r/pull/10#review"} - result = create_pr_review.run(pull_number=10, body="LGTM", event="APPROVE") + result = create_pr_review.run( + pull_number=10, + body="LGTM", + event="APPROVE", + idempotency_key=_approval_key(), + ) assert "Review submitted" in result @@ -121,5 +139,7 @@ class TestCommentTool: @patch("gitpilot.pr_tools.gp.add_pr_comment", new_callable=AsyncMock) def test_adds_comment(self, mock_comment, repo_context): mock_comment.return_value = {"html_url": "https://github.com/o/r/pull/10#comment"} - result = add_pr_comment.run(pull_number=10, body="Nice work") + result = add_pr_comment.run( + pull_number=10, body="Nice work", idempotency_key=_approval_key() + ) assert "Comment added" in result diff --git a/tests/toolkit/parity/fixture_repo.py b/tests/toolkit/parity/fixture_repo.py index 0e2b3f2..f44a207 100644 --- a/tests/toolkit/parity/fixture_repo.py +++ b/tests/toolkit/parity/fixture_repo.py @@ -2,7 +2,7 @@ Offline by construction: a real ``git init`` with pinned author, message and content, so ``git ls-files``, ``git grep``, ``git status`` and ``git log`` -behave exactly as they do against a user's checkout. Mocking git here would +behave exactly as they do against a user's checkout. Mocking git here would mean the parity gate proved the mocks matched, not the tools. """ from __future__ import annotations @@ -69,7 +69,12 @@ def build_fixture_repo(root: Path) -> WorkspaceInfo: Idempotent: called twice on the same path it returns the existing checkout rather than re-initialising, since a second ``git commit`` with nothing - staged fails. Callers that need a pristine tree pass a fresh path. + staged fails. Callers that need a pristine tree pass a fresh path. + + The repository also receives a *local* Git identity. The environment above + makes fixture creation deterministic, but later tests intentionally invoke + Git through GitPilot's normal WorkspaceManager, which must not depend on a + CI runner or developer having a global ``user.name``/``user.email``. """ if (root / ".git").is_dir(): return WorkspaceInfo( @@ -87,6 +92,8 @@ def build_fixture_repo(root: Path) -> WorkspaceInfo: target.write_text(text, encoding="utf-8") _git(["init", "--initial-branch=main"], root) + _git(["config", "user.name", "Parity Fixture"], root) + _git(["config", "user.email", "parity@example.invalid"], root) _git(["add", "-A"], root) _git(["commit", "-m", "Fixture commit"], root) diff --git a/tests/toolkit/test_fs_github_parity.py b/tests/toolkit/test_fs_github_parity.py index 43d5741..46858ff 100644 --- a/tests/toolkit/test_fs_github_parity.py +++ b/tests/toolkit/test_fs_github_parity.py @@ -1,11 +1,13 @@ """GitHub-mode parity for the filesystem tools — Batches V4-A2 / V4-A3. -The local cases in ``test_fs.py`` cover the on-disk backend. These cover the +The local cases in ``test_fs.py`` cover the on-disk backend. These cover the other one, where the tools' own logic — tree listing, glob pre-filtering, the fetch cap, read-modify-commit — is what has to keep behaving identically. """ from __future__ import annotations +import uuid + import pytest import gitpilot.agent_tools as at @@ -31,6 +33,11 @@ def repo_ctx() -> ToolExecutionContext: ) +def _approval_key() -> str: + """One request id per legacy mutation so the durable ledger cannot leak across cases.""" + return f"parity-{uuid.uuid4()}" + + PARITY_CASES = [ ParityCase( name="fs.read/github", @@ -78,7 +85,9 @@ def repo_ctx() -> ToolExecutionContext: name="fs.write/github", tool="fs.write", arguments={"path": "added.py", "content": "x = 1\n", "commit_message": "Add added.py"}, - legacy=lambda ws: at.write_file.func("added.py", "x = 1\n", "Add added.py"), + legacy=lambda ws: at.write_file.func( + "added.py", "x = 1\n", "Add added.py", idempotency_key=_approval_key() + ), ), ParityCase( name="fs.edit/github", @@ -89,7 +98,9 @@ def repo_ctx() -> ToolExecutionContext: "new_string": "return 2", "commit_message": "Bump", }, - legacy=lambda ws: at.edit_file.func("src/util.py", "return 1", "return 2", "Bump"), + legacy=lambda ws: at.edit_file.func( + "src/util.py", "return 1", "return 2", "Bump", idempotency_key=_approval_key() + ), ), ParityCase( name="fs.edit/github-conflict", @@ -100,13 +111,17 @@ def repo_ctx() -> ToolExecutionContext: "new_string": "hello", "commit_message": "Rename", }, - legacy=lambda ws: at.edit_file.func("docs/guide.md", "greet", "hello", "Rename"), + legacy=lambda ws: at.edit_file.func( + "docs/guide.md", "greet", "hello", "Rename", idempotency_key=_approval_key() + ), ), ParityCase( name="fs.delete/github", tool="fs.delete", arguments={"path": "docs/guide.md", "commit_message": "Drop guide"}, - legacy=lambda ws: at.delete_repo_file.func("docs/guide.md", "Drop guide"), + legacy=lambda ws: at.delete_repo_file.func( + "docs/guide.md", "Drop guide", idempotency_key=_approval_key() + ), ), ] diff --git a/tests/toolkit/test_registry_idempotency.py b/tests/toolkit/test_registry_idempotency.py new file mode 100644 index 0000000..6365c73 --- /dev/null +++ b/tests/toolkit/test_registry_idempotency.py @@ -0,0 +1,204 @@ +"""Idempotency invariants at the production RuntimeToolRegistry boundary.""" +from __future__ import annotations + +import asyncio + +from gitpilot.idempotency import IdempotencyStore +from gitpilot.toolkit import ( + Effect, + LocalWorkspace, + Risk, + ToolCall, + ToolExecutionContext, + ToolResult, + ToolSpec, +) +from gitpilot.toolkit.runtime_registry import RuntimeToolRegistry + + +SCHEMA = { + "type": "object", + "properties": {"path": {"type": "string"}, "content": {"type": "string"}}, + "required": ["path"], + "additionalProperties": False, +} + + +def _spec(tool_id: str, *, mutating: bool) -> ToolSpec: + return ToolSpec( + id=tool_id, + title=tool_id, + description="test tool", + params_schema=SCHEMA, + risk=Risk.APPROVAL if mutating else Risk.SAFE, + effects=( + frozenset({Effect.WRITES_FS}) + if mutating + else frozenset({Effect.READS_FS}) + ), + ) + + +def _ctx(tmp_path, store: IdempotencyStore) -> ToolExecutionContext: + return ToolExecutionContext( + workspace=LocalWorkspace(root=tmp_path), + session_id="s1", + run_id="r1", + extras={"idempotency_store": store}, + ) + + +def test_mutating_retry_replays_result_without_reexecuting(tmp_path): + registry = RuntimeToolRegistry() + calls: list[dict] = [] + + async def handler(call, ctx): + calls.append(dict(call.arguments)) + return ToolResult.success( + call, + "written", + data={"path": call.arguments["path"], "paths_written": [call.arguments["path"]]}, + ) + + registry.register(_spec("fs.write", mutating=True), handler) + store = IdempotencyStore(tmp_path / "idem.sqlite3") + ctx = _ctx(tmp_path, store) + call = ToolCall( + id="approval-123", + tool="fs.write", + arguments={"path": "a.txt", "content": "hello"}, + ) + + first = asyncio.run(registry.execute(call, ctx)) + second = asyncio.run(registry.execute(call, ctx)) + + assert first.ok and second.ok + assert first.content == second.content == "written" + assert first.data == second.data + assert calls == [{"path": "a.txt", "content": "hello"}] + + +def test_approval_id_is_bound_to_exact_arguments(tmp_path): + registry = RuntimeToolRegistry() + calls = 0 + + async def handler(call, ctx): + nonlocal calls + calls += 1 + return ToolResult.success(call, "written") + + registry.register(_spec("fs.write", mutating=True), handler) + ctx = _ctx(tmp_path, IdempotencyStore(tmp_path / "idem.sqlite3")) + + first = asyncio.run( + registry.execute( + ToolCall( + id="approval-1", + tool="fs.write", + arguments={"path": "a.txt", "content": "approved"}, + ), + ctx, + ) + ) + changed = asyncio.run( + registry.execute( + ToolCall( + id="approval-1", + tool="fs.write", + arguments={"path": "a.txt", "content": "different"}, + ), + ctx, + ) + ) + + assert first.ok + assert not changed.ok + assert changed.error == "idempotency_guard" + assert changed.data == {"retry_safe": False, "requires_reconciliation": True} + assert calls == 1 + + +def test_distinct_approval_ids_are_distinct_mutations(tmp_path): + registry = RuntimeToolRegistry() + calls = 0 + + async def handler(call, ctx): + nonlocal calls + calls += 1 + return ToolResult.success(call, f"write-{calls}") + + registry.register(_spec("fs.write", mutating=True), handler) + ctx = _ctx(tmp_path, IdempotencyStore(tmp_path / "idem.sqlite3")) + + one = asyncio.run( + registry.execute( + ToolCall(id="approval-1", tool="fs.write", arguments={"path": "a.txt"}), + ctx, + ) + ) + two = asyncio.run( + registry.execute( + ToolCall(id="approval-2", tool="fs.write", arguments={"path": "a.txt"}), + ctx, + ) + ) + + assert one.content == "write-1" + assert two.content == "write-2" + assert calls == 2 + + +def test_safe_reads_stay_on_zero_ledger_fast_path(tmp_path): + registry = RuntimeToolRegistry() + calls = 0 + + async def handler(call, ctx): + nonlocal calls + calls += 1 + return ToolResult.success(call, f"read-{calls}") + + registry.register(_spec("fs.read", mutating=False), handler) + + class ExplodingStore: + async def run_once_async(self, **kwargs): # pragma: no cover - must never run + raise AssertionError("safe reads must not touch idempotency storage") + + ctx = ToolExecutionContext( + workspace=LocalWorkspace(root=tmp_path), + session_id="s1", + run_id="r1", + extras={"idempotency_store": ExplodingStore()}, + ) + call = ToolCall(id="read-1", tool="fs.read", arguments={"path": "a.txt"}) + + first = asyncio.run(registry.execute(call, ctx)) + second = asyncio.run(registry.execute(call, ctx)) + + assert first.content == "read-1" + assert second.content == "read-2" + assert calls == 2 + + +def test_failed_mutation_becomes_indeterminate_and_is_not_retried(tmp_path): + registry = RuntimeToolRegistry() + calls = 0 + + async def handler(call, ctx): + nonlocal calls + calls += 1 + raise TimeoutError("downstream response lost") + + registry.register(_spec("fs.write", mutating=True), handler) + ctx = _ctx(tmp_path, IdempotencyStore(tmp_path / "idem.sqlite3")) + call = ToolCall(id="approval-timeout", tool="fs.write", arguments={"path": "a.txt"}) + + first = asyncio.run(registry.execute(call, ctx)) + second = asyncio.run(registry.execute(call, ctx)) + + # Preserve the base registry's first-attempt timeout contract. The durable + # runtime then refuses an automatic replay because the downstream outcome is + # unknowable after a lost/late response. + assert not first.ok and first.error == "timeout" + assert not second.ok and second.error == "idempotency_guard" + assert second.data == {"retry_safe": False, "requires_reconciliation": True} + assert calls == 1