fix(gateway): offload delivery ledger I/O
This commit is contained in:
@@ -6066,13 +6066,14 @@ class BasePlatformAdapter(ABC):
|
||||
record_obligation,
|
||||
)
|
||||
|
||||
if ledger_enabled():
|
||||
if await asyncio.to_thread(ledger_enabled):
|
||||
_obligation_id = compute_obligation_id(
|
||||
session_key,
|
||||
str(getattr(event, "message_id", "") or ""),
|
||||
text_content,
|
||||
)
|
||||
record_obligation(
|
||||
await asyncio.to_thread(
|
||||
record_obligation,
|
||||
obligation_id=_obligation_id,
|
||||
session_key=session_key,
|
||||
platform=str(
|
||||
@@ -6083,7 +6084,7 @@ class BasePlatformAdapter(ABC):
|
||||
thread_id=getattr(event.source, "thread_id", None),
|
||||
content=text_content,
|
||||
)
|
||||
mark_attempting(_obligation_id)
|
||||
await asyncio.to_thread(mark_attempting, _obligation_id)
|
||||
except Exception:
|
||||
logger.debug("delivery ledger record failed", exc_info=True)
|
||||
_obligation_id = None
|
||||
@@ -6102,9 +6103,10 @@ class BasePlatformAdapter(ABC):
|
||||
)
|
||||
|
||||
if getattr(result, "success", False):
|
||||
mark_delivered(_obligation_id)
|
||||
await asyncio.to_thread(mark_delivered, _obligation_id)
|
||||
else:
|
||||
mark_failed(
|
||||
await asyncio.to_thread(
|
||||
mark_failed,
|
||||
_obligation_id,
|
||||
str(getattr(result, "error", "") or ""),
|
||||
)
|
||||
|
||||
+4
-3
@@ -10231,7 +10231,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
sweep_recoverable,
|
||||
)
|
||||
|
||||
if not ledger_enabled():
|
||||
if not await asyncio.to_thread(ledger_enabled):
|
||||
return 0
|
||||
# Only claim rows we can actually send this boot: self.adapters
|
||||
# holds a platform only after its connect() succeeded, and each
|
||||
@@ -10283,7 +10283,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
result = None
|
||||
try:
|
||||
if result is not None and getattr(result, "success", False):
|
||||
mark_delivered(row["obligation_id"])
|
||||
await asyncio.to_thread(mark_delivered, row["obligation_id"])
|
||||
redelivered += 1
|
||||
logger.info(
|
||||
"Redelivered recovered final response to %s:%s "
|
||||
@@ -10292,7 +10292,8 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
row["obligation_id"], row["attempts"],
|
||||
)
|
||||
else:
|
||||
mark_failed(
|
||||
await asyncio.to_thread(
|
||||
mark_failed,
|
||||
row["obligation_id"],
|
||||
str(getattr(result, "error", "") or "send failed"),
|
||||
)
|
||||
|
||||
@@ -10,6 +10,7 @@ id stability, and the startup redelivery sweep's contract:
|
||||
"""
|
||||
|
||||
import time
|
||||
import threading
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -50,6 +51,30 @@ def _row(oid):
|
||||
}
|
||||
|
||||
|
||||
def _blocking_probe():
|
||||
"""Return a blocking ledger call and an event-loop progress witness."""
|
||||
ledger_started = threading.Event()
|
||||
event_loop_progressed = threading.Event()
|
||||
blocked_event_loop = []
|
||||
|
||||
def _slow_ledger_call(*args, **kwargs):
|
||||
ledger_started.set()
|
||||
if not event_loop_progressed.wait(timeout=0.5):
|
||||
blocked_event_loop.append(True)
|
||||
|
||||
async def _event_loop_witness():
|
||||
import asyncio
|
||||
|
||||
deadline = asyncio.get_running_loop().time() + 1
|
||||
while not ledger_started.is_set():
|
||||
if asyncio.get_running_loop().time() >= deadline:
|
||||
raise AssertionError("ledger call never started")
|
||||
await asyncio.sleep(0)
|
||||
event_loop_progressed.set()
|
||||
|
||||
return _slow_ledger_call, _event_loop_witness, blocked_event_loop
|
||||
|
||||
|
||||
def _orphan(oid):
|
||||
"""Make the row look like it belongs to a dead process."""
|
||||
with dl._connect() as conn:
|
||||
@@ -232,6 +257,28 @@ class TestUnconnectedPlatformKeepsItsBudget:
|
||||
runner._async_session_store = _store
|
||||
return runner
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("send_success", "ledger_method"),
|
||||
[(True, "mark_delivered"), (False, "mark_failed")],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_slow_state_update_does_not_block_event_loop(
|
||||
self, send_success, ledger_method
|
||||
):
|
||||
import asyncio
|
||||
|
||||
_record()
|
||||
_orphan("ob-1")
|
||||
runner = self._runner(self._adapter(success=send_success))
|
||||
slow_update, event_loop_witness, blocked_event_loop = _blocking_probe()
|
||||
|
||||
with patch.object(dl, ledger_method, side_effect=slow_update):
|
||||
await asyncio.gather(
|
||||
runner._redeliver_pending_obligations(), event_loop_witness()
|
||||
)
|
||||
|
||||
assert blocked_event_loop == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_row_survives_boots_where_its_platform_is_down(self):
|
||||
_record(platform="slack")
|
||||
|
||||
@@ -8,6 +8,7 @@ block the send.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -65,6 +66,28 @@ def _rows():
|
||||
).fetchall()
|
||||
|
||||
|
||||
def _blocking_probe():
|
||||
"""Return a blocking ledger call and an event-loop progress witness."""
|
||||
ledger_started = threading.Event()
|
||||
event_loop_progressed = threading.Event()
|
||||
blocked_event_loop = []
|
||||
|
||||
def _slow_ledger_call(*args, **kwargs):
|
||||
ledger_started.set()
|
||||
if not event_loop_progressed.wait(timeout=0.5):
|
||||
blocked_event_loop.append(True)
|
||||
|
||||
async def _event_loop_witness():
|
||||
deadline = asyncio.get_running_loop().time() + 1
|
||||
while not ledger_started.is_set():
|
||||
if asyncio.get_running_loop().time() >= deadline:
|
||||
raise AssertionError("ledger call never started")
|
||||
await asyncio.sleep(0)
|
||||
event_loop_progressed.set()
|
||||
|
||||
return _slow_ledger_call, _event_loop_witness, blocked_event_loop
|
||||
|
||||
|
||||
async def _run(adapter, event, response="final answer"):
|
||||
adapter._message_handler = AsyncMock(return_value=response)
|
||||
session_key = "agent:main:slack:channel:C1"
|
||||
@@ -97,6 +120,36 @@ class TestProducerHook:
|
||||
assert rows[0][1] == "failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slow_ledger_record_does_not_block_event_loop(self):
|
||||
adapter = _Adapter()
|
||||
slow_record, event_loop_witness, blocked_event_loop = _blocking_probe()
|
||||
|
||||
with patch(
|
||||
"gateway.delivery_ledger.record_obligation",
|
||||
side_effect=slow_record,
|
||||
), patch("gateway.delivery_ledger.mark_attempting"):
|
||||
await asyncio.gather(_run(adapter, _event()), event_loop_witness())
|
||||
|
||||
assert blocked_event_loop == []
|
||||
assert adapter.sent == ["final answer"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slow_ledger_update_does_not_block_event_loop(self):
|
||||
adapter = _Adapter()
|
||||
slow_delivered, event_loop_witness, blocked_event_loop = _blocking_probe()
|
||||
|
||||
with patch("gateway.delivery_ledger.record_obligation"), patch(
|
||||
"gateway.delivery_ledger.mark_attempting"
|
||||
), patch(
|
||||
"gateway.delivery_ledger.mark_delivered",
|
||||
side_effect=slow_delivered,
|
||||
):
|
||||
await asyncio.gather(_run(adapter, _event()), event_loop_witness())
|
||||
|
||||
assert blocked_event_loop == []
|
||||
assert adapter.sent == ["final answer"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_crash_between_attempting_and_ack_is_recoverable(self):
|
||||
"""The core scenario (#58818): process dies mid-send. The row must
|
||||
|
||||
Reference in New Issue
Block a user