jev-firewall 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- jev_firewall/__init__.py +48 -0
- jev_firewall/_runtime.py +75 -0
- jev_firewall/adapters/__init__.py +1 -0
- jev_firewall/adapters/_common.py +40 -0
- jev_firewall/adapters/google_adk.py +126 -0
- jev_firewall/adapters/langchain.py +137 -0
- jev_firewall/adapters/langgraph.py +161 -0
- jev_firewall/adapters/openai_agents.py +131 -0
- jev_firewall/adapters/semantic_kernel.py +83 -0
- jev_firewall/approval/__init__.py +33 -0
- jev_firewall/approval/base.py +46 -0
- jev_firewall/approval/cli.py +73 -0
- jev_firewall/approval/timeout_deny.py +28 -0
- jev_firewall/approval/webhook.py +74 -0
- jev_firewall/audit.py +88 -0
- jev_firewall/cli.py +122 -0
- jev_firewall/engine.py +215 -0
- jev_firewall/errors.py +44 -0
- jev_firewall/guard.py +136 -0
- jev_firewall/jev/__init__.py +86 -0
- jev_firewall/jev/cloudflare.py +86 -0
- jev_firewall/jev/fake.py +124 -0
- jev_firewall/jev/sdk_backend.py +88 -0
- jev_firewall/jev/types.py +79 -0
- jev_firewall/policy.example.yaml +94 -0
- jev_firewall/policy.py +232 -0
- jev_firewall/py.typed +0 -0
- jev_firewall/redact.py +190 -0
- jev_firewall/rubric.py +138 -0
- jev_firewall/verdict.py +96 -0
- jev_firewall-0.1.0.dist-info/METADATA +258 -0
- jev_firewall-0.1.0.dist-info/RECORD +35 -0
- jev_firewall-0.1.0.dist-info/WHEEL +4 -0
- jev_firewall-0.1.0.dist-info/entry_points.txt +2 -0
- jev_firewall-0.1.0.dist-info/licenses/LICENSE +202 -0
jev_firewall/__init__.py
ADDED
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
"""jev-firewall: a runtime action firewall for AI agents, powered by TypeSafe's Jev."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from jev_firewall.approval import (
|
|
6
|
+
ApprovalChannel,
|
|
7
|
+
ApprovalRequest,
|
|
8
|
+
ApprovalResult,
|
|
9
|
+
CLIApproval,
|
|
10
|
+
TimeoutDenyApproval,
|
|
11
|
+
WebhookApproval,
|
|
12
|
+
)
|
|
13
|
+
from jev_firewall.audit import AuditLog, read_audit
|
|
14
|
+
from jev_firewall.engine import PolicyEngine
|
|
15
|
+
from jev_firewall.errors import ActionBlocked, FirewallError, JevUnavailable, PolicyConfigError
|
|
16
|
+
from jev_firewall.guard import Firewall
|
|
17
|
+
from jev_firewall.jev import FakeJevClient, JevClient, build_jev_client
|
|
18
|
+
from jev_firewall.policy import Policy, RiskThresholds
|
|
19
|
+
from jev_firewall.verdict import Decision, JevAssessment, SeverityTier, ToolCall, Verdict
|
|
20
|
+
|
|
21
|
+
__version__ = "0.1.0"
|
|
22
|
+
|
|
23
|
+
__all__ = [
|
|
24
|
+
"ActionBlocked",
|
|
25
|
+
"ApprovalChannel",
|
|
26
|
+
"ApprovalRequest",
|
|
27
|
+
"ApprovalResult",
|
|
28
|
+
"AuditLog",
|
|
29
|
+
"CLIApproval",
|
|
30
|
+
"Decision",
|
|
31
|
+
"FakeJevClient",
|
|
32
|
+
"Firewall",
|
|
33
|
+
"FirewallError",
|
|
34
|
+
"JevAssessment",
|
|
35
|
+
"JevClient",
|
|
36
|
+
"JevUnavailable",
|
|
37
|
+
"Policy",
|
|
38
|
+
"PolicyConfigError",
|
|
39
|
+
"PolicyEngine",
|
|
40
|
+
"RiskThresholds",
|
|
41
|
+
"SeverityTier",
|
|
42
|
+
"TimeoutDenyApproval",
|
|
43
|
+
"ToolCall",
|
|
44
|
+
"Verdict",
|
|
45
|
+
"WebhookApproval",
|
|
46
|
+
"build_jev_client",
|
|
47
|
+
"read_audit",
|
|
48
|
+
]
|
jev_firewall/_runtime.py
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""Event-loop plumbing so one engine serves sync and async frameworks alike.
|
|
2
|
+
|
|
3
|
+
Jev calls always run on a single background loop owned by the engine. HTTP connection pools
|
|
4
|
+
are bound to the loop that created them, and agent frameworks call tools from a mix of
|
|
5
|
+
threads, executor pools, and their own loops; pinning the client to one loop avoids
|
|
6
|
+
"attached to a different loop" failures and keeps connections warm across calls.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import contextvars
|
|
13
|
+
import threading
|
|
14
|
+
from collections.abc import Coroutine
|
|
15
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
16
|
+
from typing import Any, TypeVar
|
|
17
|
+
|
|
18
|
+
T = TypeVar("T")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class LoopThread:
|
|
22
|
+
def __init__(self, name: str = "jev-firewall-loop") -> None:
|
|
23
|
+
self._name = name
|
|
24
|
+
self._loop: asyncio.AbstractEventLoop | None = None
|
|
25
|
+
self._thread: threading.Thread | None = None
|
|
26
|
+
self._lock = threading.Lock()
|
|
27
|
+
|
|
28
|
+
def _ensure(self) -> asyncio.AbstractEventLoop:
|
|
29
|
+
with self._lock:
|
|
30
|
+
if self._loop is None or self._loop.is_closed():
|
|
31
|
+
loop = asyncio.new_event_loop()
|
|
32
|
+
ready = threading.Event()
|
|
33
|
+
|
|
34
|
+
def run() -> None:
|
|
35
|
+
asyncio.set_event_loop(loop)
|
|
36
|
+
loop.call_soon(ready.set)
|
|
37
|
+
loop.run_forever()
|
|
38
|
+
|
|
39
|
+
self._thread = threading.Thread(target=run, name=self._name, daemon=True)
|
|
40
|
+
self._thread.start()
|
|
41
|
+
ready.wait()
|
|
42
|
+
self._loop = loop
|
|
43
|
+
return self._loop
|
|
44
|
+
|
|
45
|
+
def run(self, coro: Coroutine[Any, Any, T]) -> T:
|
|
46
|
+
"""Run `coro` on the background loop and block the calling thread for the result."""
|
|
47
|
+
return asyncio.run_coroutine_threadsafe(coro, self._ensure()).result()
|
|
48
|
+
|
|
49
|
+
async def arun(self, coro: Coroutine[Any, Any, T]) -> T:
|
|
50
|
+
"""Run `coro` on the background loop and await it from the caller's loop."""
|
|
51
|
+
return await asyncio.wrap_future(asyncio.run_coroutine_threadsafe(coro, self._ensure()))
|
|
52
|
+
|
|
53
|
+
def close(self) -> None:
|
|
54
|
+
with self._lock:
|
|
55
|
+
loop, self._loop = self._loop, None
|
|
56
|
+
if loop is not None and not loop.is_closed():
|
|
57
|
+
loop.call_soon_threadsafe(loop.stop)
|
|
58
|
+
if self._thread is not None:
|
|
59
|
+
self._thread.join(timeout=2)
|
|
60
|
+
loop.close()
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def run_sync(coro: Coroutine[Any, Any, T]) -> T:
|
|
64
|
+
"""Run a coroutine to completion from sync code, in the caller's context.
|
|
65
|
+
|
|
66
|
+
Unlike `LoopThread.run`, this preserves the caller's contextvars (LangGraph's `interrupt()`
|
|
67
|
+
depends on them). Works even if the calling thread already has a running loop.
|
|
68
|
+
"""
|
|
69
|
+
try:
|
|
70
|
+
asyncio.get_running_loop()
|
|
71
|
+
except RuntimeError:
|
|
72
|
+
return asyncio.run(coro)
|
|
73
|
+
ctx = contextvars.copy_context()
|
|
74
|
+
with ThreadPoolExecutor(max_workers=1) as pool:
|
|
75
|
+
return pool.submit(ctx.run, asyncio.run, coro).result()
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Thin framework adapters. Each imports its framework lazily; the core never imports them."""
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
"""Framework-free helpers shared by adapters (never imported by the core package)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Iterable, Mapping, Sequence
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def first_user_text(messages: Iterable[Any] | None) -> str | None:
|
|
10
|
+
"""Text of the first user/human message in a list of message objects or dicts.
|
|
11
|
+
|
|
12
|
+
Used as the default `agent_goal`. Only the user's own words count; tool output and
|
|
13
|
+
retrieved content come later in the conversation and cannot rewrite the goal.
|
|
14
|
+
"""
|
|
15
|
+
for m in messages or ():
|
|
16
|
+
if isinstance(m, Mapping):
|
|
17
|
+
role = m.get("role") or m.get("type")
|
|
18
|
+
content = m.get("content")
|
|
19
|
+
else:
|
|
20
|
+
role = getattr(m, "role", None) or getattr(m, "type", None)
|
|
21
|
+
content = getattr(m, "content", None)
|
|
22
|
+
if role in ("human", "user"):
|
|
23
|
+
return content_text(content)
|
|
24
|
+
return None
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def content_text(content: Any) -> str | None:
|
|
28
|
+
if isinstance(content, str):
|
|
29
|
+
return content
|
|
30
|
+
if isinstance(content, Sequence):
|
|
31
|
+
parts = []
|
|
32
|
+
for p in content:
|
|
33
|
+
if isinstance(p, Mapping):
|
|
34
|
+
parts.append(str(p.get("text") or ""))
|
|
35
|
+
elif isinstance(p, str):
|
|
36
|
+
parts.append(p)
|
|
37
|
+
else:
|
|
38
|
+
parts.append(str(getattr(p, "text", "") or ""))
|
|
39
|
+
return " ".join(x for x in parts if x) or None
|
|
40
|
+
return None
|
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
"""Google ADK adapter: `before_tool_callback` for one agent, or a plugin for a whole Runner.
|
|
2
|
+
|
|
3
|
+
Per agent:
|
|
4
|
+
|
|
5
|
+
fw_cb = JevFirewallCallbacks(firewall)
|
|
6
|
+
agent = LlmAgent(..., before_tool_callback=fw_cb.before_tool_callback,
|
|
7
|
+
after_tool_callback=fw_cb.after_tool_callback)
|
|
8
|
+
|
|
9
|
+
Every agent in an app:
|
|
10
|
+
|
|
11
|
+
app = App(name="ops", root_agent=root_agent, plugins=[JevFirewallPlugin(firewall)])
|
|
12
|
+
runner = InMemoryRunner(app=app)
|
|
13
|
+
|
|
14
|
+
On a block, `on_block="result"` skips the tool and returns `{"error": ..., "jev_firewall":
|
|
15
|
+
{...}}` to the model as the tool result (ADK's idiomatic veto). `on_block="raise"` raises
|
|
16
|
+
`ActionBlocked`: from agent-level callbacks it propagates as is; ADK's plugin manager wraps
|
|
17
|
+
plugin exceptions in `RuntimeError`, so from the plugin it arrives as that error's
|
|
18
|
+
`__cause__`. Defaults: callbacks raise, the plugin returns a result.
|
|
19
|
+
|
|
20
|
+
Requires `pip install jev-firewall[google_adk]`.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
from __future__ import annotations
|
|
24
|
+
|
|
25
|
+
from collections.abc import Callable
|
|
26
|
+
from typing import Any, Literal
|
|
27
|
+
|
|
28
|
+
try:
|
|
29
|
+
from google.adk.plugins.base_plugin import BasePlugin
|
|
30
|
+
from google.adk.tools.base_tool import BaseTool
|
|
31
|
+
from google.adk.tools.tool_context import ToolContext
|
|
32
|
+
except ImportError as exc: # pragma: no cover
|
|
33
|
+
raise ImportError(
|
|
34
|
+
"jev_firewall.adapters.google_adk needs its framework: pip install 'jev-firewall[google_adk]'"
|
|
35
|
+
) from exc
|
|
36
|
+
|
|
37
|
+
from jev_firewall.adapters._common import content_text
|
|
38
|
+
from jev_firewall.errors import ActionBlocked
|
|
39
|
+
from jev_firewall.guard import Firewall
|
|
40
|
+
from jev_firewall.verdict import ToolCall
|
|
41
|
+
|
|
42
|
+
FRAMEWORK = "google_adk"
|
|
43
|
+
|
|
44
|
+
AdkGoal = str | Callable[[ToolContext], str | None] | None
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def user_content_goal(tool_context: ToolContext) -> str | None:
|
|
48
|
+
"""Default `agent_goal`: the user message that started this invocation."""
|
|
49
|
+
content = tool_context.user_content
|
|
50
|
+
if content is None or not content.parts:
|
|
51
|
+
return None
|
|
52
|
+
return content_text([p.text for p in content.parts if p.text])
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class JevFirewallCallbacks:
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
firewall: Firewall,
|
|
59
|
+
*,
|
|
60
|
+
goal: AdkGoal = user_content_goal,
|
|
61
|
+
on_block: Literal["raise", "result"] = "raise",
|
|
62
|
+
) -> None:
|
|
63
|
+
self.firewall = firewall
|
|
64
|
+
self.goal = goal
|
|
65
|
+
self.on_block = on_block
|
|
66
|
+
|
|
67
|
+
def _call(self, tool: BaseTool, args: dict[str, Any], ctx: ToolContext) -> ToolCall:
|
|
68
|
+
extra: dict[str, Any] = {}
|
|
69
|
+
if ctx.function_call_id:
|
|
70
|
+
extra["call_id"] = ctx.function_call_id
|
|
71
|
+
goal = self.goal(ctx) if callable(self.goal) else self.goal
|
|
72
|
+
return ToolCall(tool.name, dict(args), goal, framework=FRAMEWORK, **extra)
|
|
73
|
+
|
|
74
|
+
async def before_tool_callback(
|
|
75
|
+
self, tool: BaseTool, args: dict[str, Any], tool_context: ToolContext
|
|
76
|
+
) -> dict[str, Any] | None:
|
|
77
|
+
try:
|
|
78
|
+
await self.firewall.acheck(self._call(tool, args, tool_context))
|
|
79
|
+
except ActionBlocked as exc:
|
|
80
|
+
if self.on_block == "raise":
|
|
81
|
+
raise
|
|
82
|
+
return {
|
|
83
|
+
"error": f"Action blocked by jev-firewall: {exc}",
|
|
84
|
+
"jev_firewall": {
|
|
85
|
+
"decision": exc.verdict.decision.value,
|
|
86
|
+
"reasons": list(exc.verdict.reasons),
|
|
87
|
+
},
|
|
88
|
+
}
|
|
89
|
+
return None # only None lets ADK run the tool
|
|
90
|
+
|
|
91
|
+
async def after_tool_callback(
|
|
92
|
+
self, tool: BaseTool, args: dict[str, Any], tool_context: ToolContext, tool_response: Any
|
|
93
|
+
) -> dict[str, Any] | None:
|
|
94
|
+
if tool_context.function_call_id:
|
|
95
|
+
self.firewall.record_outcome(tool_context.function_call_id, "executed")
|
|
96
|
+
return None
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
class JevFirewallPlugin(BasePlugin):
|
|
100
|
+
"""Runner-wide plugin. Plugin callbacks run before agent-level tool callbacks."""
|
|
101
|
+
|
|
102
|
+
def __init__(
|
|
103
|
+
self,
|
|
104
|
+
firewall: Firewall,
|
|
105
|
+
*,
|
|
106
|
+
goal: AdkGoal = user_content_goal,
|
|
107
|
+
on_block: Literal["raise", "result"] = "result",
|
|
108
|
+
name: str = "jev_firewall",
|
|
109
|
+
) -> None:
|
|
110
|
+
super().__init__(name=name)
|
|
111
|
+
self._cb = JevFirewallCallbacks(firewall, goal=goal, on_block=on_block)
|
|
112
|
+
|
|
113
|
+
async def before_tool_callback(
|
|
114
|
+
self, *, tool: BaseTool, tool_args: dict[str, Any], tool_context: ToolContext
|
|
115
|
+
) -> dict[str, Any] | None:
|
|
116
|
+
return await self._cb.before_tool_callback(tool, tool_args, tool_context)
|
|
117
|
+
|
|
118
|
+
async def after_tool_callback(
|
|
119
|
+
self,
|
|
120
|
+
*,
|
|
121
|
+
tool: BaseTool,
|
|
122
|
+
tool_args: dict[str, Any],
|
|
123
|
+
tool_context: ToolContext,
|
|
124
|
+
result: dict[str, Any],
|
|
125
|
+
) -> dict[str, Any] | None:
|
|
126
|
+
return await self._cb.after_tool_callback(tool, tool_args, tool_context, result)
|
|
@@ -0,0 +1,137 @@
|
|
|
1
|
+
"""LangChain adapter.
|
|
2
|
+
|
|
3
|
+
Two hooks, same core:
|
|
4
|
+
|
|
5
|
+
`JevFirewallCallbackHandler` works with any LangChain tool invocation (LCEL chains, legacy
|
|
6
|
+
agents, `tool.invoke`). It runs in `on_tool_start` and raises `ActionBlocked` before the
|
|
7
|
+
tool body executes:
|
|
8
|
+
|
|
9
|
+
handler = JevFirewallCallbackHandler(firewall, goal="Summarize the Q3 report")
|
|
10
|
+
tool.invoke(args, config={"callbacks": [handler]})
|
|
11
|
+
# or pass the goal per run: config={"callbacks": [handler], "metadata": {"agent_goal": "..."}}
|
|
12
|
+
|
|
13
|
+
`JevFirewallMiddleware` plugs into `langchain.agents.create_agent` (LangChain 1.x):
|
|
14
|
+
|
|
15
|
+
agent = create_agent(model, tools, middleware=[JevFirewallMiddleware(firewall)])
|
|
16
|
+
|
|
17
|
+
Requires `pip install jev-firewall[langchain]`.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from __future__ import annotations
|
|
21
|
+
|
|
22
|
+
import threading
|
|
23
|
+
from collections.abc import Awaitable, Callable, Mapping
|
|
24
|
+
from typing import Any, Literal
|
|
25
|
+
from uuid import UUID
|
|
26
|
+
|
|
27
|
+
try:
|
|
28
|
+
from langchain_core.callbacks import BaseCallbackHandler
|
|
29
|
+
except ImportError as exc: # pragma: no cover
|
|
30
|
+
raise ImportError(
|
|
31
|
+
"jev_firewall.adapters.langchain needs its framework: pip install 'jev-firewall[langchain]'"
|
|
32
|
+
) from exc
|
|
33
|
+
|
|
34
|
+
from jev_firewall.guard import Firewall
|
|
35
|
+
from jev_firewall.verdict import ToolCall
|
|
36
|
+
|
|
37
|
+
FRAMEWORK = "langchain"
|
|
38
|
+
|
|
39
|
+
CallbackGoal = str | Callable[[Mapping[str, Any]], str | None] | None
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class JevFirewallCallbackHandler(BaseCallbackHandler):
|
|
43
|
+
"""Vetoes tool calls from `on_tool_start`.
|
|
44
|
+
|
|
45
|
+
This is a *sync* handler on purpose: LangChain swallows exceptions from async handlers on
|
|
46
|
+
the sync tool path, so an async handler could not block anything there. Sync handlers run
|
|
47
|
+
on both paths and `raise_error = True` makes their exceptions abort the tool.
|
|
48
|
+
"""
|
|
49
|
+
|
|
50
|
+
raise_error: bool = True
|
|
51
|
+
|
|
52
|
+
def __init__(self, firewall: Firewall, *, goal: CallbackGoal = None) -> None:
|
|
53
|
+
super().__init__()
|
|
54
|
+
self.firewall = firewall
|
|
55
|
+
self.goal = goal
|
|
56
|
+
self._runs: dict[UUID, str] = {}
|
|
57
|
+
self._lock = threading.Lock()
|
|
58
|
+
|
|
59
|
+
def _goal(self, metadata: Mapping[str, Any]) -> str | None:
|
|
60
|
+
if callable(self.goal):
|
|
61
|
+
return self.goal(metadata)
|
|
62
|
+
if self.goal is not None:
|
|
63
|
+
return self.goal
|
|
64
|
+
g = metadata.get("agent_goal")
|
|
65
|
+
return None if g is None else str(g)
|
|
66
|
+
|
|
67
|
+
def on_tool_start(
|
|
68
|
+
self,
|
|
69
|
+
serialized: dict[str, Any],
|
|
70
|
+
input_str: str,
|
|
71
|
+
*,
|
|
72
|
+
run_id: UUID,
|
|
73
|
+
parent_run_id: UUID | None = None,
|
|
74
|
+
tags: list[str] | None = None,
|
|
75
|
+
metadata: dict[str, Any] | None = None,
|
|
76
|
+
inputs: dict[str, Any] | None = None,
|
|
77
|
+
**kwargs: Any,
|
|
78
|
+
) -> None:
|
|
79
|
+
name = str(serialized.get("name") or kwargs.get("name") or "unknown_tool")
|
|
80
|
+
args: dict[str, Any] = dict(inputs) if inputs is not None else {"input": input_str}
|
|
81
|
+
extra: dict[str, Any] = {}
|
|
82
|
+
if kwargs.get("tool_call_id"):
|
|
83
|
+
extra["call_id"] = str(kwargs["tool_call_id"])
|
|
84
|
+
call = ToolCall(name, args, self._goal(metadata or {}), framework=FRAMEWORK, **extra)
|
|
85
|
+
self.firewall.check(call) # raises ActionBlocked on DENY or a rejected HOLD
|
|
86
|
+
with self._lock:
|
|
87
|
+
self._runs[run_id] = call.call_id
|
|
88
|
+
|
|
89
|
+
def on_tool_end(self, output: Any, *, run_id: UUID, **kwargs: Any) -> None:
|
|
90
|
+
with self._lock:
|
|
91
|
+
call_id = self._runs.pop(run_id, None)
|
|
92
|
+
if call_id is not None:
|
|
93
|
+
self.firewall.record_outcome(call_id, "executed")
|
|
94
|
+
|
|
95
|
+
def on_tool_error(self, error: BaseException, *, run_id: UUID, **kwargs: Any) -> None:
|
|
96
|
+
with self._lock:
|
|
97
|
+
call_id = self._runs.pop(run_id, None)
|
|
98
|
+
if call_id is not None:
|
|
99
|
+
self.firewall.record_outcome(call_id, "error", f"{type(error).__name__}: {error}")
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
try:
|
|
103
|
+
from langchain.agents.middleware import AgentMiddleware
|
|
104
|
+
from langchain.tools.tool_node import ToolCallRequest
|
|
105
|
+
from langchain_core.messages import ToolMessage
|
|
106
|
+
from langgraph.types import Command
|
|
107
|
+
except ImportError: # pragma: no cover - only `langchain-core` installed
|
|
108
|
+
pass
|
|
109
|
+
else:
|
|
110
|
+
from jev_firewall.adapters.langgraph import GoalSource, JevToolCallWrapper, first_human_message
|
|
111
|
+
|
|
112
|
+
class JevFirewallMiddleware(AgentMiddleware):
|
|
113
|
+
"""`create_agent` middleware; the same interception as the LangGraph adapter."""
|
|
114
|
+
|
|
115
|
+
def __init__(
|
|
116
|
+
self,
|
|
117
|
+
firewall: Firewall,
|
|
118
|
+
*,
|
|
119
|
+
goal: GoalSource = first_human_message,
|
|
120
|
+
on_block: Literal["raise", "message"] = "message",
|
|
121
|
+
) -> None:
|
|
122
|
+
super().__init__()
|
|
123
|
+
self._wrapper = JevToolCallWrapper(firewall, goal=goal, on_block=on_block)
|
|
124
|
+
|
|
125
|
+
def wrap_tool_call(
|
|
126
|
+
self,
|
|
127
|
+
request: ToolCallRequest,
|
|
128
|
+
handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]],
|
|
129
|
+
) -> ToolMessage | Command[Any]:
|
|
130
|
+
return self._wrapper(request, handler)
|
|
131
|
+
|
|
132
|
+
async def awrap_tool_call(
|
|
133
|
+
self,
|
|
134
|
+
request: ToolCallRequest,
|
|
135
|
+
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]],
|
|
136
|
+
) -> ToolMessage | Command[Any]:
|
|
137
|
+
return await self._wrapper.acall(request, handler)
|
|
@@ -0,0 +1,161 @@
|
|
|
1
|
+
"""LangGraph adapter: intercept tool calls in `ToolNode` via `wrap_tool_call`.
|
|
2
|
+
|
|
3
|
+
from jev_firewall import Firewall
|
|
4
|
+
from jev_firewall.adapters.langgraph import firewall_tool_node
|
|
5
|
+
|
|
6
|
+
firewall = Firewall.from_yaml("policy.yaml")
|
|
7
|
+
tools_node = firewall_tool_node([terminal, send_email], firewall)
|
|
8
|
+
builder.add_node("tools", tools_node)
|
|
9
|
+
|
|
10
|
+
HOLD is resolved by the firewall's approval channel. To pause the graph instead (needs a
|
|
11
|
+
checkpointer), pass `LangGraphInterruptApproval()` as the Firewall's approval channel and
|
|
12
|
+
resume with `Command(resume={"approved": True, "resolver": "alice"})`.
|
|
13
|
+
|
|
14
|
+
Requires `pip install jev-firewall[langgraph]`.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
from collections import OrderedDict
|
|
20
|
+
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
21
|
+
from typing import Any, Literal
|
|
22
|
+
|
|
23
|
+
try:
|
|
24
|
+
from langchain_core.messages import ToolMessage
|
|
25
|
+
from langchain_core.tools import BaseTool
|
|
26
|
+
from langgraph.prebuilt import ToolNode
|
|
27
|
+
from langgraph.prebuilt.tool_node import ToolCallRequest
|
|
28
|
+
from langgraph.types import Command, interrupt
|
|
29
|
+
except ImportError as exc: # pragma: no cover
|
|
30
|
+
raise ImportError(
|
|
31
|
+
"jev_firewall.adapters.langgraph needs its framework: pip install 'jev-firewall[langgraph]'"
|
|
32
|
+
) from exc
|
|
33
|
+
|
|
34
|
+
from jev_firewall.adapters._common import first_user_text
|
|
35
|
+
from jev_firewall.approval import ApprovalRequest, ApprovalResult
|
|
36
|
+
from jev_firewall.errors import ActionBlocked
|
|
37
|
+
from jev_firewall.guard import Firewall
|
|
38
|
+
from jev_firewall.verdict import Decision, ToolCall, Verdict
|
|
39
|
+
|
|
40
|
+
GoalSource = str | Callable[[Any], str | None] | None
|
|
41
|
+
ToolResult = ToolMessage | Command[Any]
|
|
42
|
+
|
|
43
|
+
FRAMEWORK = "langgraph"
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def first_human_message(state: Any) -> str | None:
|
|
47
|
+
"""Default `agent_goal`: the first human message in `state["messages"]`."""
|
|
48
|
+
messages = state.get("messages") if isinstance(state, Mapping) else getattr(state, "messages", None)
|
|
49
|
+
return first_user_text(messages)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
class JevToolCallWrapper:
|
|
53
|
+
"""A `wrap_tool_call` / `awrap_tool_call` pair for `ToolNode` (or any compatible hook).
|
|
54
|
+
|
|
55
|
+
`on_block="raise"` lets `ActionBlocked` propagate and stops the run.
|
|
56
|
+
`on_block="message"` returns an error `ToolMessage` instead so the model can recover.
|
|
57
|
+
"""
|
|
58
|
+
|
|
59
|
+
def __init__(
|
|
60
|
+
self,
|
|
61
|
+
firewall: Firewall,
|
|
62
|
+
*,
|
|
63
|
+
goal: GoalSource = first_human_message,
|
|
64
|
+
on_block: Literal["raise", "message"] = "raise",
|
|
65
|
+
cache_size: int = 1024,
|
|
66
|
+
) -> None:
|
|
67
|
+
self.firewall = firewall
|
|
68
|
+
self.goal = goal
|
|
69
|
+
self.on_block = on_block
|
|
70
|
+
# Verdicts for held calls, keyed by tool_call id. A LangGraph interrupt re-runs the node
|
|
71
|
+
# on resume; reusing the verdict avoids a second Jev call and a second audit record.
|
|
72
|
+
self._held: OrderedDict[str, Verdict] = OrderedDict()
|
|
73
|
+
self._cache_size = cache_size
|
|
74
|
+
|
|
75
|
+
def _call(self, request: ToolCallRequest) -> ToolCall:
|
|
76
|
+
tc = request.tool_call
|
|
77
|
+
goal = self.goal(request.state) if callable(self.goal) else self.goal
|
|
78
|
+
kwargs: dict[str, Any] = {}
|
|
79
|
+
if tc.get("id"):
|
|
80
|
+
kwargs["call_id"] = str(tc["id"])
|
|
81
|
+
return ToolCall(tc["name"], dict(tc.get("args") or {}), goal, framework=FRAMEWORK, **kwargs)
|
|
82
|
+
|
|
83
|
+
def _remember(self, v: Verdict) -> None:
|
|
84
|
+
if v.decision is Decision.HOLD:
|
|
85
|
+
self._held[v.call_id] = v
|
|
86
|
+
while len(self._held) > self._cache_size:
|
|
87
|
+
self._held.popitem(last=False)
|
|
88
|
+
|
|
89
|
+
def _blocked(self, request: ToolCallRequest, exc: ActionBlocked) -> ToolMessage:
|
|
90
|
+
if self.on_block == "raise":
|
|
91
|
+
raise exc
|
|
92
|
+
return ToolMessage(
|
|
93
|
+
content=f"Action blocked by jev-firewall: {exc}",
|
|
94
|
+
name=request.tool_call["name"],
|
|
95
|
+
tool_call_id=request.tool_call["id"] or "",
|
|
96
|
+
status="error",
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
def __call__(
|
|
100
|
+
self, request: ToolCallRequest, execute: Callable[[ToolCallRequest], ToolResult]
|
|
101
|
+
) -> ToolResult:
|
|
102
|
+
call = self._call(request)
|
|
103
|
+
verdict = self._held.get(call.call_id) or self.firewall.engine.evaluate(call)
|
|
104
|
+
self._remember(verdict)
|
|
105
|
+
try:
|
|
106
|
+
self.firewall.resolve(verdict)
|
|
107
|
+
except ActionBlocked as exc:
|
|
108
|
+
self._held.pop(call.call_id, None)
|
|
109
|
+
return self._blocked(request, exc)
|
|
110
|
+
self._held.pop(call.call_id, None)
|
|
111
|
+
return self.firewall.guard_executed(call, lambda: execute(request))
|
|
112
|
+
|
|
113
|
+
async def acall(
|
|
114
|
+
self, request: ToolCallRequest, execute: Callable[[ToolCallRequest], Awaitable[ToolResult]]
|
|
115
|
+
) -> ToolResult:
|
|
116
|
+
call = self._call(request)
|
|
117
|
+
verdict = self._held.get(call.call_id) or await self.firewall.engine.aevaluate(call)
|
|
118
|
+
self._remember(verdict)
|
|
119
|
+
try:
|
|
120
|
+
await self.firewall.aresolve(verdict)
|
|
121
|
+
except ActionBlocked as exc:
|
|
122
|
+
self._held.pop(call.call_id, None)
|
|
123
|
+
return self._blocked(request, exc)
|
|
124
|
+
self._held.pop(call.call_id, None)
|
|
125
|
+
return await self.firewall.aguard_executed(call, lambda: execute(request))
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def firewall_tool_node(
|
|
129
|
+
tools: Sequence[BaseTool | Callable[..., Any]],
|
|
130
|
+
firewall: Firewall,
|
|
131
|
+
*,
|
|
132
|
+
goal: GoalSource = first_human_message,
|
|
133
|
+
on_block: Literal["raise", "message"] = "raise",
|
|
134
|
+
**tool_node_kwargs: Any,
|
|
135
|
+
) -> ToolNode:
|
|
136
|
+
"""A `ToolNode` whose every tool call goes through the firewall first."""
|
|
137
|
+
wrapper = JevToolCallWrapper(firewall, goal=goal, on_block=on_block)
|
|
138
|
+
return ToolNode(tools, wrap_tool_call=wrapper, awrap_tool_call=wrapper.acall, **tool_node_kwargs)
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
class LangGraphInterruptApproval:
|
|
142
|
+
"""Approval channel that pauses the graph with `interrupt()` (requires a checkpointer).
|
|
143
|
+
|
|
144
|
+
The interrupt value is `{"type": "jev_firewall_hold", ...ApprovalRequest}`. Resume with
|
|
145
|
+
`Command(resume=True)`, `Command(resume="yes")` or
|
|
146
|
+
`Command(resume={"approved": True, "resolver": "alice", "note": "..."})`.
|
|
147
|
+
"""
|
|
148
|
+
|
|
149
|
+
async def request(self, req: ApprovalRequest) -> ApprovalResult:
|
|
150
|
+
answer = interrupt({"type": "jev_firewall_hold", **req.to_dict()})
|
|
151
|
+
if isinstance(answer, Mapping):
|
|
152
|
+
return ApprovalResult(
|
|
153
|
+
approved=answer.get("approved") is True,
|
|
154
|
+
resolver=str(answer.get("resolver") or "langgraph:resume"),
|
|
155
|
+
note=None if answer.get("note") is None else str(answer["note"]),
|
|
156
|
+
)
|
|
157
|
+
if isinstance(answer, str):
|
|
158
|
+
return ApprovalResult(
|
|
159
|
+
answer.strip().lower() in {"y", "yes", "approve", "approved"}, "langgraph:resume"
|
|
160
|
+
)
|
|
161
|
+
return ApprovalResult(answer is True, "langgraph:resume")
|