Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion sdk/python/adrian/langchain_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -614,9 +614,20 @@ def _sync_gate(tool_call_id: str) -> bool:
thread (no loop *set* there, since Python 3.10+), which would
misclassify the worker-thread case as "no loop" and skip the gate -
leaving sync tools ungated under ``create_react_agent``.

Login and policy checks belong to ``_async_gate``, which waits for the
handshake. Deciding them here races startup and lets the first sync
tool call of a run through ungated.

Args:
tool_call_id: Id of the tool call to gate.

Returns:
True if the tool should be blocked.
"""
ws = _ws_getter()
if ws is None or not ws._login_ack_received.is_set() or not ws.policy_active(): # pyright: ignore[reportPrivateUsage]

if ws is None:
return False

# Is THIS thread running an event loop?
Expand Down
59 changes: 59 additions & 0 deletions sdk/python/tests/test_block_mode.py
Original file line number Diff line number Diff line change
Expand Up @@ -559,3 +559,62 @@ def _real_tool(x: str) -> str:

assert captured == []
assert "BLOCKED" in result["messages"][0].content


class TestSyncGateBeforeLoginAck:
"""The sync gate must not skip while the handshake is still in flight.

``_async_gate`` waits for the LoginAck and blocks if it never arrives.
``_sync_gate`` used to answer that itself and skip, so the first tool
call of a run ran ungated if the tool was sync, while an identical async
tool was gated. The verdict still arrived, which made the bypass silent.
"""

async def test_sync_tool_gated_when_login_ack_is_late(self, tmp_path: Path) -> None:
"""LoginAck lands after dispatch: the tool must still be blocked."""
captured: list[str] = []

def _real_tool(x: str) -> str:
"""Sync tool stub; records execution."""
captured.append(x)

return x

adrian.init(
api_key="k",
log_file=str(tmp_path / "events.jsonl"),
auto_instrument=True,
ws_url="ws://x",
block_timeout=5.0,
)

ws = adrian._ws_client
assert ws is not None
policy = _apply_mode(ws, pb.MODE_BLOCK, policy_m4=True)
ws._connected.set()
ws._loop = asyncio.get_running_loop()
ws._tool_call_id_to_event_id["tc-1"] = "llm-evt"
fut = ws.register_pending("llm-evt")
fut.set_result(
pb.Verdict(event_id="llm-evt", mad_code="M4_a", policy=policy),
)

# Dispatch before the handshake completes.
ws._login_ack_received.clear()

async def _late_login_ack() -> None:
await asyncio.sleep(0.05)
ws._login_ack_received.set()

_ = asyncio.create_task(_late_login_ack())

ai = AIMessage(
content="",
tool_calls=[{"id": "tc-1", "name": "_real_tool", "args": {"x": "hi"}}],
)
state: dict[str, Any] = {"messages": [ai]}

result = await ToolNode([_real_tool]).ainvoke(state, config=_runtime_config()) # pyright: ignore[reportUnknownMemberType]

assert captured == []
assert "BLOCKED" in result["messages"][0].content
Loading