From 72639c4b907b7ecec6273d3c600c2f622ff2b174 Mon Sep 17 00:00:00 2001 From: h Date: Sat, 29 Aug 2026 15:27:26 +0200 Subject: [PATCH] feat(policy,claude_sdk,cli,config_loader): pretooluse policy rules per agent, tool call audit, sibling imports for config --- src/beaver_gateway/agents/claude.py | 4 + src/beaver_gateway/backends/claude_sdk.py | 53 +++++++ src/beaver_gateway/cli.py | 31 +++- src/beaver_gateway/config_loader.py | 5 + src/beaver_gateway/core/audit.py | 1 + src/beaver_gateway/core/policy.py | 174 ++++++++++++++++++++++ src/beaver_gateway/core/sessions.py | 2 + tests/test_claude_sdk_backend.py | 64 ++++++++ tests/test_config_loader_siblings.py | 16 ++ tests/test_policy.py | 71 +++++++++ 10 files changed, 420 insertions(+), 1 deletion(-) create mode 100644 src/beaver_gateway/core/policy.py create mode 100644 tests/test_config_loader_siblings.py create mode 100644 tests/test_policy.py diff --git a/src/beaver_gateway/agents/claude.py b/src/beaver_gateway/agents/claude.py index c6388dd..ab03a52 100644 --- a/src/beaver_gateway/agents/claude.py +++ b/src/beaver_gateway/agents/claude.py @@ -16,6 +16,7 @@ from pydantic import BaseModel, ConfigDict, Field, model_validator from beaver_gateway.agents.base import BaseAgent from beaver_gateway.core.kinds import KINDS, Kind +from beaver_gateway.core.policy import PolicyRule # noqa: TC001 - pydantic runtime from beaver_gateway.core.prompt import PromptSource # noqa: TC001 - pydantic runtime __all__ = ["ClaudeAgent", "ClaudeOptions", "Prompts"] @@ -86,6 +87,9 @@ class ClaudeAgent(BaseAgent): ``say``, ``schedule``, ``inject``); empty = no gateway MCP server.""" options: ClaudeOptions = Field(default_factory=ClaudeOptions) + policy: tuple[PolicyRule, ...] = () + """``PreToolUse`` rules (§3.7), run in order on every tool call; the + first ``Deny`` is what the model reads back. See ``core/policy``.""" @model_validator(mode="after") def _kinds_follow_prompts(self) -> ClaudeAgent: diff --git a/src/beaver_gateway/backends/claude_sdk.py b/src/beaver_gateway/backends/claude_sdk.py index f3651c0..3809301 100644 --- a/src/beaver_gateway/backends/claude_sdk.py +++ b/src/beaver_gateway/backends/claude_sdk.py @@ -53,6 +53,7 @@ from claude_agent_sdk import ( CanUseToolShadowedWarning, ClaudeAgentOptions, ClaudeSDKClient, + HookMatcher, MirrorErrorMessage, PermissionResultAllow, PermissionResultDeny, @@ -66,6 +67,7 @@ from claude_agent_sdk import ( project_key_for_directory, ) +from beaver_gateway.core import policy as policy_mod from beaver_gateway.core import prompt as prompt_assembly from beaver_gateway.core.events import ( StopReason, @@ -101,6 +103,7 @@ if TYPE_CHECKING: SessionStore, ToolPermissionContext, ) + from claude_agent_sdk.types import HookEvent from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.claude import ClaudeAgent @@ -116,6 +119,7 @@ warnings.filterwarnings("ignore", category=CanUseToolShadowedWarning) ASK_TOOL = "AskUserQuestion" __all__ = [ + "AuditSink", "ClaudeSdkBackend", "RunnerConfig", "SessionClient", @@ -166,6 +170,8 @@ ClientFactory = "Callable[[ClaudeAgentOptions], SessionClient]" UsageSink = "Callable[[UsageEvent], Awaitable[None]]" ToolServerFactory = "Callable[[str, str], McpSdkServerConfig | None]" """``(conversation_key, kind) -> in-process MCP server config`` or ``None``.""" +AuditSink = "Callable[[policy_mod.ToolAudit], Awaitable[None]]" +"""Receives every tool call the PreToolUse hook saw, allowed or denied.""" Asker = "Callable[[str, dict[str, Any]], Awaitable[str]]" """``(conversation_key, AskUserQuestion input) -> text the model reads as the tool result``. The only channel an answer has in bypass mode is @@ -218,8 +224,10 @@ class ClaudeSdkBackend: pool: SessionPool | None = None, tool_server: Callable[[str, str], McpSdkServerConfig | None] | None = None, asker: Callable[[str, dict[str, Any]], Awaitable[str]] | None = None, + audit_sink: Callable[[policy_mod.ToolAudit], Awaitable[None]] | None = None, ) -> None: self._agent = agent + self._audit_sink = audit_sink self._store = session_store self._runner = runner or RunnerConfig() self._usage_sink = usage_sink @@ -594,6 +602,7 @@ class ClaudeSdkBackend: cli_path=str(self._exec_wrapper(extra_keep=tuple(env))), include_partial_messages=opt.include_partial_messages, can_use_tool=self._can_use_tool(key) if self._asker else None, + hooks=self._hooks(key, spec), session_store=self._store, session_store_flush=cast("Any", opt.session_store_flush), resume=resume, @@ -629,6 +638,50 @@ class ClaudeSdkBackend: return can_use_tool + def _hooks( + self, key: str, spec: _SessionSpec + ) -> dict[HookEvent, list[HookMatcher]] | None: + """§3.7: one in-process ``PreToolUse`` hook - policy rules, then audit.""" + agent = self._agent + if not agent.policy and self._audit_sink is None: + return None + + async def pre_tool_use( + hook_input: Any, _tool_use_id: str | None, _ctx: Any + ) -> dict[str, Any]: + live = self._pool.get(key) + call = policy_mod.ToolCall( + tool=str(hook_input.get("tool_name", "")), + input=hook_input.get("tool_input") or {}, + agent=agent.name, + kind=spec.kind, + conversation=key, + cwd=agent.cwd, + state=live.state if live is not None else {}, + ) + deny = await policy_mod.evaluate(agent.policy, call) + if deny is not None: + _log.info( + "policy: %s denied %s: %s", agent.name, call.tool, deny.reason + ) + if self._audit_sink is not None: + audit = policy_mod.ToolAudit( + agent=agent.name, + conversation=key, + kind=spec.kind, + tool=call.tool, + decision="deny" if deny else "allow", + reason=deny.reason if deny else None, + brief=policy_mod.brief(call.input), + ) + try: + await self._audit_sink(audit) + except Exception: # noqa: BLE001 + _log.exception("tool audit failed for %s", call.tool) + return policy_mod.hook_output(deny) + + return {"PreToolUse": [HookMatcher(hooks=[cast("Any", pre_tool_use)])]} + def _plugins(self) -> list[dict[str, str]]: plugins: list[dict[str, str]] = [] root = self._work_dir / "plugins" / self._agent.name diff --git a/src/beaver_gateway/cli.py b/src/beaver_gateway/cli.py index 3621c25..c0e0b38 100644 --- a/src/beaver_gateway/cli.py +++ b/src/beaver_gateway/cli.py @@ -59,7 +59,13 @@ from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.root import build_root_app from beaver_gateway.mcp.internal_app import build_internal_app from beaver_gateway.settings import Settings -from beaver_gateway.storage import Database, PostgresSessionStore, Usage, append_usage +from beaver_gateway.storage import ( + Database, + PostgresSessionStore, + Usage, + append_audit, + append_usage, +) if TYPE_CHECKING: from claude_agent_sdk import McpSdkServerConfig @@ -70,6 +76,7 @@ if TYPE_CHECKING: from starlette.types import ASGIApp from beaver_gateway.backends.base import Backend + from beaver_gateway.core.policy import ToolAudit from beaver_gateway.mcp.types import McpServerT @@ -432,6 +439,27 @@ async def _build_backends( except Exception: # noqa: BLE001 _log.exception("usage write failed for %s", event.agent_name) + async def record_tool(event: ToolAudit) -> None: + detail = { + "conversation": event.conversation, + "kind": event.kind, + "tool": event.tool, + "decision": event.decision, + "reason": event.reason, + "brief": event.brief, + } + try: + async with db.session() as session: + await append_audit( + session, + actor=f"agent:{event.agent}", + kind="tool_call", + agent_name=event.agent, + detail=detail, + ) + except Exception: # noqa: BLE001 + _log.exception("tool audit write failed for %s", event.agent) + for a in agents: if isinstance(a, ClaudeAgent): adapter = ClaudeSdkBackend( @@ -444,6 +472,7 @@ async def _build_backends( pool=pool, tool_server=functools.partial(late.server, names=a.gateway_tools), asker=late.ask, + audit_sink=record_tool, ) await stack.enter_async_context(adapter) backends[a.name] = adapter diff --git a/src/beaver_gateway/config_loader.py b/src/beaver_gateway/config_loader.py index 6365e60..9b756fc 100644 --- a/src/beaver_gateway/config_loader.py +++ b/src/beaver_gateway/config_loader.py @@ -14,6 +14,7 @@ contents here before handing it back. from __future__ import annotations +import sys from typing import TYPE_CHECKING, Any from croniter import croniter @@ -59,6 +60,10 @@ def load(path: Path) -> Gateway: raise ConfigError(msg) from exc code = compile(source, str(path), "exec") + # Siblings of the config (``policy.py``, ``mcps/``) import by name. + parent = str(path.resolve().parent) + if parent not in sys.path: + sys.path.insert(0, parent) namespace: dict[str, Any] = {"__file__": str(path), **_PUBLIC_NAMES} exec(code, namespace) # noqa: S102 - exec'ing user config is the feature diff --git a/src/beaver_gateway/core/audit.py b/src/beaver_gateway/core/audit.py index 23bf26e..b36bb9c 100644 --- a/src/beaver_gateway/core/audit.py +++ b/src/beaver_gateway/core/audit.py @@ -51,6 +51,7 @@ KNOWN_KINDS: frozenset[str] = frozenset( { "messages", # POST /v1/messages accepted "mcp_call", # /mcp//... proxied + "tool_call", # a model's tool call seen by the PreToolUse hook "login_ok", "login_failed", "logout", diff --git a/src/beaver_gateway/core/policy.py b/src/beaver_gateway/core/policy.py new file mode 100644 index 0000000..b576b49 --- /dev/null +++ b/src/beaver_gateway/core/policy.py @@ -0,0 +1,174 @@ +"""PreToolUse policy - the boundary without permission prompts (§3.7). + +``bypassPermissions`` everywhere; what a model may do is decided by +mounts, ``disallowed_tools`` and the rules here. A rule is a callable +``(ToolCall) -> Deny | None`` declared per agent (``ClaudeAgent.policy``); +the SDK backend registers one in-process ``PreToolUse`` hook that runs +the rules in order and turns the first :class:`Deny` into a hook deny +whose reason the model reads as the tool result. Rules never see +secrets and never prompt - they only say no, with a reason. + +A rule that raises is a deny too: the boundary fails closed, the +traceback lands in the log. + +Every tool call - allowed or denied - is reported to the audit sink the +backend was given, so the admin audit page shows what the model touched. +""" + +from __future__ import annotations + +import inspect +import json +import logging +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field +from pathlib import Path +from typing import TYPE_CHECKING, Any, Literal + +if TYPE_CHECKING: + from collections.abc import Iterable, Mapping, MutableMapping + +__all__ = [ + "AuditSink", + "Decision", + "Deny", + "PolicyRule", + "ToolAudit", + "ToolCall", + "brief", + "evaluate", + "hook_output", +] + +_log = logging.getLogger("beaver_gateway.policy") + +PATH_KEYS: tuple[str, ...] = ("file_path", "notebook_path", "path") +"""Input keys built-in file tools use for their target.""" + +Decision = Literal["allow", "deny"] + + +@dataclass(frozen=True, slots=True) +class Deny: + """Refusal with the reason the model reads as the tool result.""" + + reason: str + + +@dataclass(frozen=True, slots=True) +class ToolCall: + """One tool invocation as the hook sees it. + + ``state`` is a mutable per-session dict rules may write to (which + skills were opened, what was already asked); it lives as long as the + claude process and is empty again after a resume. + """ + + tool: str + input: Mapping[str, Any] + agent: str + kind: str + conversation: str + cwd: Path + state: MutableMapping[str, Any] = field(default_factory=dict) + + @property + def mcp(self) -> tuple[str, str] | None: + """``(server, tool)`` for ``mcp____`` names, else ``None``.""" + if not self.tool.startswith("mcp__"): + return None + _, _, rest = self.tool.partition("__") + server, sep, name = rest.partition("__") + return (server, name) if sep else (server, "") + + def path(self, *keys: str) -> Path | None: + """Target of a file tool, absolute; relative paths resolve against ``cwd``.""" + for key in keys or PATH_KEYS: + raw = self.input.get(key) + if isinstance(raw, str) and raw: + return self.resolve(raw) + return None + + def resolve(self, raw: str) -> Path: + path = Path(raw).expanduser() + if not path.is_absolute(): + path = self.cwd / path + # ``resolve`` would follow symlinks on the gateway host, which may + # not be the model's view; normalise lexically instead. + return Path(*_normalize(path.parts)) + + +def _normalize(parts: Iterable[str]) -> list[str]: + out: list[str] = [] + for part in parts: + if part == "..": + if len(out) > 1: + out.pop() + elif part != ".": + out.append(part) + return out + + +type PolicyRule = Callable[[ToolCall], Deny | None | Awaitable[Deny | None]] +"""A rule; sync or async. First ``Deny`` wins, ``None`` passes to the next.""" + + +@dataclass(frozen=True, slots=True) +class ToolAudit: + agent: str + conversation: str + kind: str + tool: str + decision: Decision + reason: str | None + brief: str + + +AuditSink = "Callable[[ToolAudit], Awaitable[None]]" + + +async def evaluate( + rules: Iterable[Callable[[ToolCall], Any]], call: ToolCall +) -> Deny | None: + """Run ``rules`` in order; the first deny wins, an exception is a deny.""" + for rule in rules: + try: + verdict = rule(call) + if inspect.isawaitable(verdict): + verdict = await verdict + except Exception: # noqa: BLE001 - a broken rule must fail closed + name = getattr(rule, "__name__", repr(rule)) + _log.exception("policy rule %s failed on %s", name, call.tool) + return Deny(reason=f"policy rule {name} failed; the call is refused") + if isinstance(verdict, Deny): + return verdict + return None + + +def hook_output(deny: Deny | None) -> dict[str, Any]: + """``PreToolUse`` hook JSON for a verdict: empty means allow.""" + if deny is None: + return {} + return { + "hookSpecificOutput": { + "hookEventName": "PreToolUse", + "permissionDecision": "deny", + "permissionDecisionReason": deny.reason, + } + } + + +def brief(tool_input: Mapping[str, Any], *, limit: int = 240) -> str: + """Short, log-safe summary of a call: the path, the command, or the args.""" + for key in (*PATH_KEYS, "command", "skill", "pattern", "query", "url"): + value = tool_input.get(key) + if isinstance(value, str) and value: + return _cut(value, limit) + try: + return _cut(json.dumps(tool_input, ensure_ascii=False, sort_keys=True), limit) + except (TypeError, ValueError): + return _cut(repr(tool_input), limit) + + +def _cut(text: str, limit: int) -> str: + return text if len(text) <= limit else text[: limit - 1] + "…" diff --git a/src/beaver_gateway/core/sessions.py b/src/beaver_gateway/core/sessions.py index 9914a00..b27f319 100644 --- a/src/beaver_gateway/core/sessions.py +++ b/src/beaver_gateway/core/sessions.py @@ -68,6 +68,8 @@ class Session: last_used: float = field(default_factory=time.monotonic) created_at: float = field(default_factory=time.monotonic) turns: int = 0 + state: dict[str, Any] = field(default_factory=dict) + """Scratch for policy rules (``core/policy``); dies with the process.""" @property def busy(self) -> bool: diff --git a/tests/test_claude_sdk_backend.py b/tests/test_claude_sdk_backend.py index c47bf5a..3ac451c 100644 --- a/tests/test_claude_sdk_backend.py +++ b/tests/test_claude_sdk_backend.py @@ -458,3 +458,67 @@ async def test_deltas_reach_the_caller_before_the_turn_ends(cwd: Path) -> None: GatedClient.gate.set() rest = [e async for e in events] assert isinstance(rest[-1], RawMessageStopEvent) + + +async def test_policy_hook_denies_and_audits(cwd: Path) -> None: + from beaver_gateway.core.policy import Deny, ToolCall + + def no_days(c: ToolCall): + p = c.path() + if c.tool == "Write" and p is not None and "📅 дни" in p.parts: + return Deny(reason="дневник только читать") + c.state["seen"] = c.state.get("seen", 0) + 1 + return None + + audits = [] + + async def sink(a): + audits.append(a) + + backend = _backend(cwd, InMemorySessionStore(), policy=(no_days,)) + backend._audit_sink = sink + await _drain( + backend.complete( + agent=backend.agent, + messages=[{"role": "user", "content": "hi"}], + conversation_id="conv-1", + ) + ) + hook = FakeClient.instances[0].options.hooks["PreToolUse"][0].hooks[0] + denied = await hook( + {"tool_name": "Write", "tool_input": {"file_path": "/vault/📅 дни/x.md"}}, + "t1", + {}, + ) + assert denied["hookSpecificOutput"]["permissionDecision"] == "deny" + assert ( + denied["hookSpecificOutput"]["permissionDecisionReason"] + == "дневник только читать" + ) + allowed = await hook( + {"tool_name": "Write", "tool_input": {"file_path": "/vault/мета/бобер/x.md"}}, + "t2", + {}, + ) + assert allowed == {} + await hook( + {"tool_name": "Read", "tool_input": {"file_path": "/vault/a.md"}}, "t3", {} + ) + assert backend.live("conv-1").state["seen"] == 2 + assert [(a.tool, a.decision) for a in audits] == [ + ("Write", "deny"), + ("Write", "allow"), + ("Read", "allow"), + ] + assert audits[0].brief == "/vault/📅 дни/x.md" + assert audits[0].kind == "deep" and audits[0].conversation == "conv-1" + + +async def test_no_policy_no_hooks(cwd: Path) -> None: + backend = _backend(cwd, InMemorySessionStore()) + await _drain( + backend.complete( + agent=backend.agent, messages=[{"role": "user", "content": "hi"}] + ) + ) + assert FakeClient.instances[0].options.hooks is None diff --git a/tests/test_config_loader_siblings.py b/tests/test_config_loader_siblings.py new file mode 100644 index 0000000..9871885 --- /dev/null +++ b/tests/test_config_loader_siblings.py @@ -0,0 +1,16 @@ +import tempfile +from pathlib import Path + +from beaver_gateway import config_loader + + +def test_config_imports_sibling_modules(): + root = Path(tempfile.mkdtemp(prefix="beaver-cfg-")) + (root / "policy_sibling.py").write_text("RULES = ('x',)\n") + (root / "config.py").write_text( + "from policy_sibling import RULES\n" + "assert RULES == ('x',)\n" + "gateway = Gateway()\n" + ) + gw = config_loader.load(root / "config.py") + assert gw.agents == [] diff --git a/tests/test_policy.py b/tests/test_policy.py new file mode 100644 index 0000000..494182a --- /dev/null +++ b/tests/test_policy.py @@ -0,0 +1,71 @@ +from pathlib import Path + +import pytest + +from beaver_gateway.core.policy import Deny, ToolCall, brief, evaluate, hook_output + + +def call(tool: str, **tool_input) -> ToolCall: + return ToolCall( + tool=tool, + input=tool_input, + agent="a", + kind="master", + conversation="c1", + cwd=Path("/vault"), + ) + + +def test_path_resolves_relative_against_cwd(): + assert call("Write", file_path="мета/бобер/x.md").path() == Path( + "/vault/мета/бобер/x.md" + ) + assert call("Write", file_path="/vault/a/../b.md").path() == Path("/vault/b.md") + assert call("Write", file_path="../../etc/passwd").path() == Path("/etc/passwd") + assert call("Bash", command="ls").path() is None + + +def test_mcp_split(): + assert call("mcp__firefly__store_transaction").mcp == ( + "firefly", + "store_transaction", + ) + assert call("Write").mcp is None + + +async def test_first_deny_wins_and_async_rules_work(): + seen = [] + + def allow(c): + seen.append("allow") + + async def deny(c): + return Deny(reason="no") + + def never(c): + raise AssertionError + + verdict = await evaluate((allow, deny, never), call("Write")) + assert verdict == Deny(reason="no") + assert seen == ["allow"] + + +async def test_raising_rule_fails_closed(): + def boom(c): + raise RuntimeError("x") + + verdict = await evaluate((boom,), call("Write")) + assert verdict is not None and "boom" in verdict.reason + + +def test_hook_output_shapes(): + assert hook_output(None) == {} + out = hook_output(Deny(reason="r")) + assert out["hookSpecificOutput"]["permissionDecision"] == "deny" + assert out["hookSpecificOutput"]["permissionDecisionReason"] == "r" + + +def test_brief_prefers_path_then_command_and_truncates(): + assert brief({"file_path": "/v/x.md", "content": "a" * 999}) == "/v/x.md" + assert brief({"command": "ls -la"}) == "ls -la" + assert len(brief({"q": "a" * 999})) == 240