diff --git a/sdk/python/adrian/langchain_handler.py b/sdk/python/adrian/langchain_handler.py index 4596a20..85bb859 100644 --- a/sdk/python/adrian/langchain_handler.py +++ b/sdk/python/adrian/langchain_handler.py @@ -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? diff --git a/sdk/python/tests/test_block_mode.py b/sdk/python/tests/test_block_mode.py index 949cdb2..f2fa365 100644 --- a/sdk/python/tests/test_block_mode.py +++ b/sdk/python/tests/test_block_mode.py @@ -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