fix(plugins): gate hook callbacks by call identity, not by tool name alone

Concurrent invocations of the same tool in one session collapsed into a single
busy key (hook_name, id(cb)): the second invocation was reported as 'still
running' and dropped. For pre_tool_call a drop is a fail-closed block, so the
gate silenced itself on an ordinary, healthy callback.

Measured on a busy profile: 3574 skip lines and 0 timeout lines in one hour —
every skip was the 'while still running' branch, i.e. pure key collision, not
slowness.

The gate now keys on the call identity that is already in the payload
(tool_call_id, else turn_id, else none — the last case behaves exactly as
before). Suppression stays keyed coarsely on (hook_name, id(cb)): a hung
callback is a fact about the callback, so its back-off must not be diluted
per call.

Refs #98382. Independent of #107894 (that one releases the slot on timeout;
this one stops healthy concurrency from colliding).

(cherry picked from commit 53b3dacd008418fcdf5fa6dfcadde575a35a776e)
This commit is contained in:
deadczarvc
2026-09-14 04:49:37 +03:00
committed by kshitij
parent b18140576d
commit 4121aa295a
2 changed files with 90 additions and 8 deletions

View File

@@ -140,6 +140,22 @@ _MAX_HOOK_CALLBACK_TIMEOUT_SECS = 600.0
_HOOK_SKIPPED = object() # returned by _run_hook_callback_bounded on skip/timeout
def _hook_call_identity(kwargs: Dict[str, Any]) -> Any:
"""Identity of the call this callback fires for, or ``None`` when the event has none.
Concurrent invocations of the same tool in one session must not collapse into one
gate key: they are different work, and treating the second as a duplicate drops the
hook as if a callback had timed out (upstream #98382). The identity is already in the
payload; nothing new is plumbed. Deliberately not ``api_request_id`` — one API request
carries many tool calls, which would re-collapse the keys.
"""
for field in ("tool_call_id", "turn_id"):
value = kwargs.get(field)
if isinstance(value, str) and value:
return value
return None
def _hook_uses_callback_timeout(hook_name: str, timeout: float) -> bool:
"""Whether *hook_name* should run under the non-blocking timeout path."""
if timeout <= 0 or hook_name in _HOOK_CALLER_THREAD_HOOKS:
@@ -211,19 +227,22 @@ class PluginDispatchMixin:
suppressed, still running, timed out (worker abandoned, never joined), or the worker
could not be started. Exceptions propagate."""
callback_name = getattr(cb, "__name__", repr(cb))
callback_key = (hook_name, id(cb))
# Suppression is a fact about the CALLBACK — a hung one must keep its back-off —
# so that key stays coarse. The gate must instead tell CONCURRENT CALLS apart.
suppression_key = (hook_name, id(cb))
gate_key = (*suppression_key, _hook_call_identity(kwargs))
token = object()
with self._hook_timeout_lock:
suppressed_until = self._hook_timeout_suppressed_until.get(callback_key)
running = callback_key in self._hook_running_callbacks
suppressed_until = self._hook_timeout_suppressed_until.get(suppression_key)
running = gate_key in self._hook_running_callbacks
if (suppressed_until is not None and suppressed_until > time.monotonic()) or running:
logger.warning(
"Hook '%s' callback %s skipped after previous "
"timeout or while still running", hook_name, callback_name)
return _HOOK_SKIPPED
if suppressed_until is not None:
self._hook_timeout_suppressed_until.pop(callback_key, None)
self._hook_running_callbacks[callback_key] = token
self._hook_timeout_suppressed_until.pop(suppression_key, None)
self._hook_running_callbacks[gate_key] = token
context = contextvars.copy_context()
done = threading.Event()
@@ -232,8 +251,8 @@ class PluginDispatchMixin:
def _release_token() -> None:
with self._hook_timeout_lock:
if self._hook_running_callbacks.get(callback_key) is token:
self._hook_running_callbacks.pop(callback_key, None)
if self._hook_running_callbacks.get(gate_key) is token:
self._hook_running_callbacks.pop(gate_key, None)
def _runner() -> None:
try:
@@ -256,7 +275,7 @@ class PluginDispatchMixin:
if not done.wait(timeout=timeout): # do not join — that would reintroduce the hang
with self._hook_timeout_lock:
# See #6622.
self._hook_timeout_suppressed_until[callback_key] = (
self._hook_timeout_suppressed_until[suppression_key] = (
time.monotonic() + self._hook_timeout_suppression_seconds)
logger.warning(
"Hook '%s' callback %s timed out after %gs — skipping", hook_name, callback_name, timeout)

View File

@@ -1182,6 +1182,69 @@ class TestForceReloadSymmetry:
assert elapsed < 5.0
hold.set()
def test_concurrent_same_tool_calls_with_distinct_ids_both_run(self, monkeypatch):
"""Two concurrent calls of one tool are different work, not a duplicate (#98382)."""
import time
monkeypatch.setattr(
"hermes_cli.plugins._resolve_hook_callback_timeout", lambda: 5.0
)
hold = threading.Event()
starts = []
def recorder(**_kwargs):
starts.append(1)
hold.wait(timeout=10.0)
return "ok"
mgr = PluginManager()
mgr._hooks["pre_tool_call"] = [recorder]
def fire(call_id):
mgr.invoke_hook(
"pre_tool_call",
tool_name="read_file",
tool_input={},
session_id="s1",
tool_call_id=call_id,
)
first = threading.Thread(target=fire, args=("call-a",), daemon=True)
first.start()
time.sleep(0.1) # let the first invocation occupy the gate
second = threading.Thread(target=fire, args=("call-b",), daemon=True)
second.start()
time.sleep(0.4)
hold.set()
first.join(5.0)
second.join(5.0)
assert len(starts) == 2
def test_repeated_same_call_identity_still_deduplicated(self, monkeypatch):
"""Negative control: the same call identity stays a duplicate, so a hung
worker is never restarted by a repeat of the very same call."""
monkeypatch.setattr(
"hermes_cli.plugins._resolve_hook_callback_timeout", lambda: 0.1
)
hold = threading.Event()
starts = []
def blocker(**_kwargs):
starts.append(1)
hold.wait(timeout=10.0)
return "late"
mgr = PluginManager()
mgr._hooks["post_tool_call"] = [blocker]
assert mgr.invoke_hook("post_tool_call", tool_name="read_file", tool_call_id="same-call") == []
assert mgr.invoke_hook("post_tool_call", tool_name="read_file", tool_call_id="same-call") == []
assert len(starts) == 1
hold.set()
def test_pre_tool_call_timeout_fail_closed(self, monkeypatch):
"""Timed-out pre_tool_call must return a block directive, not allow."""
import time