fix(agent): use async aupdate_state for checkpoint repair
The repair paths (send-path _repair_pending_tool_calls and resume fallback)
called graph.update_state — the SYNC method — which internally invokes
AsyncPostgresSaver.get_tuple() and raises InvalidStateError ('Synchronous
calls to AsyncPostgresSaver are only allowed from a different thread').
The repair therefore always failed (surfacing as 'Event loop is closed' +
a dangling aget_tuple coroutine) and broken threads stayed broken.
Switch both repair sites to agent.aupdate_state (async checkpointer
interface) and align the mocked agent in tests. Verified live: aupdate_state
repaired the broken checkpoint of conversation 69651ca1 (pending calls → 0).
This commit is contained in:
@@ -711,10 +711,15 @@ async def handle_resume( # noqa: C901
|
||||
# ToolMessages so the thread is consistent and the NEXT user
|
||||
# message does not fail LangGraph INVALID_CHAT_HISTORY, and a
|
||||
# repeated confirm cannot re-execute the same tool calls.
|
||||
# NB: must use the async aupdate_state — with AsyncPostgresSaver
|
||||
# the sync update_state raises InvalidStateError ("Synchronous
|
||||
# calls to AsyncPostgresSaver are only allowed from a different
|
||||
# thread"), so the repair silently failed and the thread stayed
|
||||
# broken (observed as "Event loop is closed" + dangling aget_tuple).
|
||||
if repaired_msgs:
|
||||
try:
|
||||
await _await_if_async(
|
||||
agent.update_state(config, {"messages": [*state_messages, *repaired_msgs]})
|
||||
await agent.aupdate_state(
|
||||
config, {"messages": [*state_messages, *repaired_msgs]}
|
||||
)
|
||||
logger.reason(
|
||||
"Checkpoint repaired after resume fallback",
|
||||
|
||||
@@ -38,7 +38,6 @@ from openai import APIConnectionError, APITimeoutError, AuthenticationError, Rat
|
||||
|
||||
from ss_tools.agent._config import GRADIO_ROOT_PATH, GRADIO_SERVER_NAME, GRADIO_SERVER_PORT, STORAGE_ROOT as _STORAGE_ROOT
|
||||
from ss_tools.agent._confirmation import (
|
||||
_await_if_async,
|
||||
_pending_confirmations,
|
||||
confirmation_payload,
|
||||
handle_resume,
|
||||
@@ -198,7 +197,11 @@ async def _repair_pending_tool_calls(agent: Any, config: dict[str, Any]) -> int:
|
||||
)
|
||||
for tool_name, _args, tcid in pending
|
||||
]
|
||||
await _await_if_async(agent.update_state(config, {"messages": [*current_msgs, *repaired]}))
|
||||
# NB: use the ASYNC aupdate_state — with AsyncPostgresSaver the sync
|
||||
# update_state calls checkpointer.get_tuple() which raises
|
||||
# InvalidStateError ("Synchronous calls to AsyncPostgresSaver are only
|
||||
# allowed from a different thread"), leaving the thread permanently broken.
|
||||
await agent.aupdate_state(config, {"messages": [*current_msgs, *repaired]})
|
||||
return len(repaired)
|
||||
# #endregion AgentChat.GradioApp.RepairBrokenThread
|
||||
|
||||
|
||||
@@ -283,7 +283,7 @@ class TestHandleResumeIntegration:
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.astream_events = MagicMock(side_effect=ValueError("INVALID_CHAT_HISTORY"))
|
||||
mock_agent.aget_state = AsyncMock(return_value=pending_state)
|
||||
mock_agent.update_state = AsyncMock(return_value=None)
|
||||
mock_agent.aupdate_state = AsyncMock(return_value=None)
|
||||
|
||||
_pending_confirmations["conv-unknown"] = {
|
||||
"tool_name": "nonexistent_tool_xyz", "tool_args": {}, "_fast_path": False,
|
||||
@@ -303,8 +303,8 @@ class TestHandleResumeIntegration:
|
||||
assert len(error_chunk) == 1
|
||||
assert "Unknown tool" in error_chunk[0]["metadata"]["error"]
|
||||
# Checkpoint repair ran with a ToolMessage for the pending call.
|
||||
assert mock_agent.update_state.called
|
||||
repaired = mock_agent.update_state.call_args.args[1]["messages"]
|
||||
assert mock_agent.aupdate_state.called
|
||||
repaired = mock_agent.aupdate_state.call_args.args[1]["messages"]
|
||||
assert repaired[-1].tool_call_id == "call_xyz"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -324,7 +324,7 @@ class TestHandleResumeIntegration:
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.astream_events = MagicMock(side_effect=ValueError("INVALID_CHAT_HISTORY"))
|
||||
mock_agent.aget_state = AsyncMock(return_value=pending_state)
|
||||
mock_agent.update_state = AsyncMock(return_value=None)
|
||||
mock_agent.aupdate_state = AsyncMock(return_value=None)
|
||||
|
||||
tool = MagicMock()
|
||||
tool.ainvoke = AsyncMock(side_effect=RuntimeError("API timeout"))
|
||||
@@ -343,8 +343,8 @@ class TestHandleResumeIntegration:
|
||||
assert len(error_chunk) == 1
|
||||
assert "API timeout" in error_chunk[0]["metadata"]["error"]
|
||||
# Checkpoint repaired with the error ToolMessage.
|
||||
assert mock_agent.update_state.called
|
||||
repaired = mock_agent.update_state.call_args.args[1]["messages"]
|
||||
assert mock_agent.aupdate_state.called
|
||||
repaired = mock_agent.aupdate_state.call_args.args[1]["messages"]
|
||||
assert repaired[-1].tool_call_id == "call_env"
|
||||
assert "Error: API timeout" in repaired[-1].content
|
||||
|
||||
@@ -583,7 +583,7 @@ class TestHandleResumeIntegration:
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.astream_events = MagicMock(side_effect=ValueError("INVALID_CHAT_HISTORY"))
|
||||
mock_agent.aget_state = AsyncMock(return_value=pending_state)
|
||||
mock_agent.update_state = AsyncMock(return_value=None)
|
||||
mock_agent.aupdate_state = AsyncMock(return_value=None)
|
||||
|
||||
tool = self._scenario_tool_mock()
|
||||
|
||||
@@ -610,8 +610,8 @@ class TestHandleResumeIntegration:
|
||||
"agent_run_id must be injected into tool args from the resolved run"
|
||||
)
|
||||
# Repair is awaited with the ToolMessage appended.
|
||||
assert mock_agent.update_state.called
|
||||
repaired = mock_agent.update_state.call_args.args[1]["messages"]
|
||||
assert mock_agent.aupdate_state.called
|
||||
repaired = mock_agent.aupdate_state.call_args.args[1]["messages"]
|
||||
assert repaired[-1].tool_call_id == "call_scn"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -634,7 +634,7 @@ class TestHandleResumeIntegration:
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.astream_events = MagicMock(side_effect=ValueError("INVALID_CHAT_HISTORY"))
|
||||
mock_agent.aget_state = AsyncMock(return_value=pending_state)
|
||||
mock_agent.update_state = AsyncMock(return_value=None)
|
||||
mock_agent.aupdate_state = AsyncMock(return_value=None)
|
||||
|
||||
tool = self._scenario_tool_mock()
|
||||
|
||||
@@ -651,8 +651,8 @@ class TestHandleResumeIntegration:
|
||||
assert error_chunk[0]["metadata"]["code"] == "SCENARIO_RUN_REQUIRED"
|
||||
assert not tool.ainvoke.called, "Scenario tool must not execute without a run"
|
||||
# The thread is still repaired so a later message cannot wedge the thread.
|
||||
assert mock_agent.update_state.called
|
||||
repaired = mock_agent.update_state.call_args.args[1]["messages"]
|
||||
assert mock_agent.aupdate_state.called
|
||||
repaired = mock_agent.aupdate_state.call_args.args[1]["messages"]
|
||||
assert repaired[-1].tool_call_id == "call_scn2"
|
||||
# #endregion Test.AgentChat.TestHandleResumeIntegration
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ class _FakeAgent:
|
||||
async def aget_state(self, config):
|
||||
return self._state
|
||||
|
||||
async def update_state(self, config, values):
|
||||
async def aupdate_state(self, config, values):
|
||||
self.updated.append(values)
|
||||
return {"configurable": config}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user