fix(sessions): harden cross-process turn lease wait and refresh
Honor interrupts while waiting for admission, stop the turn when refresh loses the lease, poll once per second under contention, and test dead-PID reclaim.
This commit is contained in:
+15
-1
@@ -5818,9 +5818,10 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
*,
|
||||
ttl_seconds: float = 300.0,
|
||||
wait_seconds: float = 1800.0,
|
||||
poll_interval_seconds: float = 0.1,
|
||||
poll_interval_seconds: float = 1.0,
|
||||
on_wait=None,
|
||||
wait_notice_interval_seconds: float = 15.0,
|
||||
should_abort=None,
|
||||
) -> bool:
|
||||
"""Wait for a cross-process turn lease without holding a SQLite lock.
|
||||
|
||||
@@ -5828,12 +5829,25 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
attempt fails (elapsed ~0) and again about every
|
||||
``wait_notice_interval_seconds`` while still waiting, so UIs can show
|
||||
that another process holds the conversation.
|
||||
|
||||
When ``should_abort()`` returns True (for example the agent received
|
||||
``/stop`` while waiting), acquisition stops immediately and returns
|
||||
False without consuming the full ``wait_seconds`` budget.
|
||||
"""
|
||||
deadline = time.monotonic() + max(0.0, float(wait_seconds))
|
||||
wait_started = None
|
||||
last_notice_at = None
|
||||
notice_every = max(0.0, float(wait_notice_interval_seconds))
|
||||
while True:
|
||||
if should_abort is not None:
|
||||
try:
|
||||
if should_abort():
|
||||
return False
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"session turn lease should_abort callback failed",
|
||||
exc_info=True,
|
||||
)
|
||||
if self.try_acquire_session_turn_lease(
|
||||
session_id, holder, ttl_seconds=ttl_seconds
|
||||
):
|
||||
|
||||
+29
-4
@@ -8138,7 +8138,21 @@ class AIAgent:
|
||||
ttl_seconds=_lease_ttl,
|
||||
wait_seconds=1800.0,
|
||||
on_wait=_on_session_turn_lease_wait,
|
||||
should_abort=lambda: getattr(self, "_interrupt_requested", False),
|
||||
):
|
||||
if getattr(self, "_interrupt_requested", False):
|
||||
logger.info(
|
||||
"session turn lease wait aborted by interrupt: %s",
|
||||
session_id,
|
||||
)
|
||||
relay_outcome = "cancelled"
|
||||
return {
|
||||
"final_response": "",
|
||||
"messages": list(conversation_history or []),
|
||||
"api_calls": 0,
|
||||
"completed": False,
|
||||
"interrupted": True,
|
||||
}
|
||||
# Fail closed like gateway TurnLeaseTimeoutError: do not
|
||||
# enter load/run/flush, and surface a resend notice instead
|
||||
# of a bare TimeoutError that looks like a hang.
|
||||
@@ -8192,24 +8206,35 @@ class AIAgent:
|
||||
# in a daemon thread; holder-qualified UPDATE and DELETE fence a
|
||||
# late refresher/release from a successor lease.
|
||||
durable_turn_lease_stop = threading.Event()
|
||||
_lease_refresh_interval = float(
|
||||
getattr(self, "_session_turn_lease_refresh_interval", 60.0)
|
||||
)
|
||||
|
||||
def _refresh_durable_turn_lease() -> None:
|
||||
while not durable_turn_lease_stop.wait(60.0):
|
||||
while not durable_turn_lease_stop.wait(_lease_refresh_interval):
|
||||
try:
|
||||
if not _turn_db.refresh_session_turn_lease(
|
||||
session_id,
|
||||
getattr(self, "session_id", None) or session_id,
|
||||
durable_turn_lease,
|
||||
ttl_seconds=_lease_ttl,
|
||||
):
|
||||
logger.error(
|
||||
"Lost session turn lease while turn is active: %s",
|
||||
session_id,
|
||||
getattr(self, "session_id", None) or session_id,
|
||||
)
|
||||
try:
|
||||
self.interrupt(
|
||||
"Session turn lease lost; stopping to "
|
||||
"protect the transcript.",
|
||||
hard_cancel=True,
|
||||
)
|
||||
except Exception:
|
||||
self._interrupt_requested = True
|
||||
return
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to refresh session turn lease: %s",
|
||||
session_id,
|
||||
getattr(self, "session_id", None) or session_id,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
@@ -166,3 +168,83 @@ def test_run_conversation_lease_timeout_returns_resend_notice(monkeypatch):
|
||||
kind == "warn" and text and "send it again" in text
|
||||
for kind, text in status_events
|
||||
)
|
||||
|
||||
|
||||
def test_run_conversation_lease_wait_honors_interrupt(monkeypatch):
|
||||
db = _DB()
|
||||
agent = _agent_with_db(db)
|
||||
agent._interrupt_requested = False
|
||||
|
||||
def acquire_with_abort(session_id, holder, **kwargs):
|
||||
db.events.append(("acquire", session_id, holder))
|
||||
should_abort = kwargs.get("should_abort")
|
||||
assert callable(should_abort)
|
||||
agent._interrupt_requested = True
|
||||
assert should_abort()
|
||||
return False
|
||||
|
||||
db.acquire_session_turn_lease = acquire_with_abort
|
||||
|
||||
def boom(*_args, **_kwargs):
|
||||
raise AssertionError("turn must not start when lease wait is aborted")
|
||||
|
||||
monkeypatch.setattr("agent.conversation_loop.run_conversation", boom)
|
||||
result = AIAgent.run_conversation(
|
||||
agent,
|
||||
"new message",
|
||||
conversation_history=[{"role": "user", "content": "stale"}],
|
||||
)
|
||||
|
||||
assert result.get("interrupted") is True
|
||||
assert result.get("failed") is not True
|
||||
assert "session_turn_lease_timeout" not in str(result.get("error", ""))
|
||||
assert [event[0] for event in db.events] == ["acquire"]
|
||||
|
||||
|
||||
def test_run_conversation_interrupts_when_lease_refresh_lost(monkeypatch):
|
||||
db = _DB()
|
||||
agent = _agent_with_db(db)
|
||||
agent._session_turn_lease_refresh_interval = 0.01
|
||||
interrupt_calls = []
|
||||
|
||||
def track_interrupt(message=None, hard_cancel=False):
|
||||
interrupt_calls.append((message, hard_cancel))
|
||||
agent._interrupt_requested = True
|
||||
|
||||
agent.interrupt = track_interrupt
|
||||
|
||||
def refresh_lost(session_id, holder, **kwargs):
|
||||
return False
|
||||
|
||||
db.refresh_session_turn_lease = refresh_lost
|
||||
|
||||
observed = {"started": False}
|
||||
|
||||
def fake_run(_agent, _message, _system, history, *_args, **_kwargs):
|
||||
observed["started"] = True
|
||||
deadline = time.monotonic() + 2.0
|
||||
while time.monotonic() < deadline:
|
||||
if getattr(_agent, "_interrupt_requested", False):
|
||||
return {
|
||||
"final_response": "",
|
||||
"messages": history,
|
||||
"api_calls": 0,
|
||||
"completed": False,
|
||||
"interrupted": True,
|
||||
}
|
||||
time.sleep(0.01)
|
||||
raise AssertionError("refresh loss did not interrupt the turn")
|
||||
|
||||
monkeypatch.setattr("agent.conversation_loop.run_conversation", fake_run)
|
||||
|
||||
result = AIAgent.run_conversation(
|
||||
agent,
|
||||
"new message",
|
||||
conversation_history=[{"role": "user", "content": "seed"}],
|
||||
)
|
||||
|
||||
assert observed["started"] is True
|
||||
assert result.get("interrupted") is True
|
||||
assert interrupt_calls
|
||||
assert interrupt_calls[0][1] is True
|
||||
assert "lease lost" in str(interrupt_calls[0][0]).lower()
|
||||
|
||||
@@ -5,7 +5,11 @@ from __future__ import annotations
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
import hermes_state
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
@@ -164,4 +168,65 @@ def test_acquire_turn_lease_notifies_wait_callback(tmp_path):
|
||||
|
||||
assert notices
|
||||
assert notices[0] < 0.05
|
||||
second.release_session_turn_lease("shared", second_holder)
|
||||
second.release_session_turn_lease("shared", second_holder)
|
||||
|
||||
|
||||
def test_acquire_turn_lease_honors_should_abort(tmp_path):
|
||||
"""Waiters stop immediately when should_abort() returns True."""
|
||||
path = tmp_path / "state.db"
|
||||
first = SessionDB(path)
|
||||
second = SessionDB(path)
|
||||
first.create_session("shared", source="test")
|
||||
|
||||
first_holder = f"pid={os.getpid()}:turn=first"
|
||||
second_holder = f"pid={os.getpid()}:turn=second"
|
||||
assert first.try_acquire_session_turn_lease(
|
||||
"shared", first_holder, ttl_seconds=60
|
||||
)
|
||||
|
||||
abort_checks = {"count": 0}
|
||||
|
||||
def should_abort():
|
||||
abort_checks["count"] += 1
|
||||
return True
|
||||
|
||||
started = time.monotonic()
|
||||
assert not second.acquire_session_turn_lease(
|
||||
"shared",
|
||||
second_holder,
|
||||
wait_seconds=30,
|
||||
poll_interval_seconds=0.05,
|
||||
should_abort=should_abort,
|
||||
)
|
||||
assert time.monotonic() - started < 1.0
|
||||
assert abort_checks["count"] >= 1
|
||||
first.release_session_turn_lease("shared", first_holder)
|
||||
|
||||
|
||||
def test_non_expired_turn_lease_from_dead_pid_is_reclaimed(
|
||||
tmp_path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""A holder whose structured pid= no longer exists can be reclaimed early."""
|
||||
db = SessionDB(tmp_path / "state.db")
|
||||
db.create_session("shared", source="test")
|
||||
|
||||
dead_holder = "pid=424242:turn=dead:platform=test"
|
||||
assert db.try_acquire_session_turn_lease(
|
||||
"shared", dead_holder, ttl_seconds=300
|
||||
) is True
|
||||
|
||||
probed: list[int] = []
|
||||
|
||||
def pid_exists(pid: int) -> bool:
|
||||
probed.append(pid)
|
||||
return False
|
||||
|
||||
monkeypatch.setattr(
|
||||
hermes_state, "psutil", SimpleNamespace(pid_exists=pid_exists)
|
||||
)
|
||||
|
||||
fresh_holder = "pid=525252:turn=fresh:platform=test"
|
||||
assert db.try_acquire_session_turn_lease(
|
||||
"shared", fresh_holder, ttl_seconds=300
|
||||
) is True
|
||||
assert probed == [424242]
|
||||
Reference in New Issue
Block a user