mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-05 10:41:58 +03:00
feat(runtime): add user-controlled turn recovery
This commit is contained in:
+103
-132
@@ -79,6 +79,15 @@ from nanobot.session.model_selection import (
|
||||
SESSION_MODEL_PRESET_METADATA_KEY,
|
||||
model_preset_from_metadata,
|
||||
)
|
||||
from nanobot.session.recovery import (
|
||||
PENDING_FOLLOWUP_ID_KEY,
|
||||
RECOVERY_INBOUND_METADATA_KEY,
|
||||
RecoveryAdmission,
|
||||
acknowledge_pending_followups,
|
||||
record_pending_followup,
|
||||
restore_pending_interruption,
|
||||
restore_runtime_checkpoint,
|
||||
)
|
||||
from nanobot.session.summary import SessionSummary
|
||||
from nanobot.triggers.local_turns import LocalTriggerTurnCoordinator
|
||||
from nanobot.utils.cancellation import task_is_cancelling
|
||||
@@ -291,12 +300,14 @@ class AgentLoop:
|
||||
restart_mode: str = "auto",
|
||||
local_trigger_store: LocalTriggerStore | None = None,
|
||||
idle_compact_check_interval_seconds: int = 0,
|
||||
recovery_admission: RecoveryAdmission | None = None,
|
||||
):
|
||||
from nanobot.config.schema import ToolsConfig
|
||||
|
||||
_tc = tools_config or ToolsConfig()
|
||||
defaults = AgentDefaults()
|
||||
self.bus = bus
|
||||
self._recovery_admission = recovery_admission
|
||||
if turn_delivery_factory is not None:
|
||||
if turn_delivery_factory.bus is not bus:
|
||||
raise ValueError("turn delivery factory must use the agent message bus")
|
||||
@@ -409,6 +420,7 @@ class AgentLoop:
|
||||
# When a session has an active task, new messages for that session
|
||||
# are routed here instead of creating a new task.
|
||||
self._pending_queues: dict[str, asyncio.Queue[InboundMessage]] = {}
|
||||
self._preserve_inflight_turns_on_shutdown = False
|
||||
self._deferred_automation_turns: dict[str, list[InboundMessage]] = {}
|
||||
self._cron_turns = CronTurnCoordinator(
|
||||
publish_inbound=self.bus.publish_inbound,
|
||||
@@ -726,6 +738,9 @@ class AgentLoop:
|
||||
extra[RUNTIME_CONTEXT_HISTORY_META] = runtime_context_meta
|
||||
session.add_message("user", text, **extra)
|
||||
self._mark_pending_user_turn(session)
|
||||
followup_id = msg.metadata.get(PENDING_FOLLOWUP_ID_KEY)
|
||||
if isinstance(followup_id, str) and followup_id:
|
||||
acknowledge_pending_followups(session, [followup_id])
|
||||
self.sessions.save(session)
|
||||
return True
|
||||
return False
|
||||
@@ -1061,6 +1076,9 @@ class AgentLoop:
|
||||
row["subagent_task_id"] = task_id
|
||||
row[HIDDEN_HISTORY_META] = subagent_marker
|
||||
row["injected_event"] = "subagent_result"
|
||||
followup_id = metadata.get(PENDING_FOLLOWUP_ID_KEY)
|
||||
if isinstance(followup_id, str) and followup_id:
|
||||
row[PENDING_FOLLOWUP_ID_KEY] = followup_id
|
||||
return row
|
||||
|
||||
items: list[dict[str, Any]] = []
|
||||
@@ -1285,6 +1303,17 @@ class AgentLoop:
|
||||
break
|
||||
if deferred:
|
||||
continue
|
||||
# A newer WebUI message must supersede an explicit recovery
|
||||
# before it is injected into that recovery's pending queue.
|
||||
# Without this admission point, a recovered turn could finish
|
||||
# first and only then observe the user's newer request.
|
||||
if (
|
||||
effective_key in self._pending_queues
|
||||
and msg.channel == "websocket"
|
||||
and self._recovery_admission is not None
|
||||
and not await self._recovery_admission.admit(msg)
|
||||
):
|
||||
continue
|
||||
# If this session already has an active pending queue (i.e. a task
|
||||
# is processing this session), route the message there for mid-turn
|
||||
# injection instead of creating a competing task.
|
||||
@@ -1303,6 +1332,17 @@ class AgentLoop:
|
||||
msg,
|
||||
session_key_override=effective_key,
|
||||
)
|
||||
session = self.sessions.get_or_create(effective_key)
|
||||
followup_id = record_pending_followup(session, pending_msg)
|
||||
if followup_id is not None:
|
||||
pending_msg = dataclasses.replace(
|
||||
pending_msg,
|
||||
metadata={
|
||||
**pending_msg.metadata,
|
||||
PENDING_FOLLOWUP_ID_KEY: followup_id,
|
||||
},
|
||||
)
|
||||
self.sessions.save(session)
|
||||
try:
|
||||
self._pending_queues[effective_key].put_nowait(pending_msg)
|
||||
except asyncio.QueueFull:
|
||||
@@ -1310,6 +1350,7 @@ class AgentLoop:
|
||||
"Pending queue full for session {}, falling back to queued task",
|
||||
effective_key,
|
||||
)
|
||||
msg = pending_msg
|
||||
else:
|
||||
logger.info(
|
||||
"Routed follow-up message to pending queue for session {}",
|
||||
@@ -1319,17 +1360,45 @@ class AgentLoop:
|
||||
# Compute the effective session key before dispatching
|
||||
# This ensures /stop command can find tasks correctly when unified session is enabled
|
||||
task = asyncio.create_task(self._dispatch(msg))
|
||||
active_tasks = self._active_tasks.setdefault(effective_key, set())
|
||||
active_tasks: set[asyncio.Task[Any]] = self._active_tasks.setdefault(
|
||||
effective_key,
|
||||
set(),
|
||||
)
|
||||
active_tasks.add(task)
|
||||
task.add_done_callback(active_tasks.discard)
|
||||
finally:
|
||||
await self.aclose()
|
||||
|
||||
def preserve_inflight_turns_on_shutdown(self) -> None:
|
||||
"""Keep durable checkpoints when the owning gateway is restarting.
|
||||
|
||||
Normal cancellation intentionally materializes partial output so a
|
||||
stopped gateway leaves a readable conversation. A managed restart is
|
||||
different: RecoveryCoordinator needs the checkpoint intact to safely
|
||||
offer the unfinished turn for explicit continuation after restart.
|
||||
"""
|
||||
self._preserve_inflight_turns_on_shutdown = True
|
||||
|
||||
async def _dispatch(self, msg: InboundMessage) -> None:
|
||||
"""Process a message: per-session serial, cross-session concurrent."""
|
||||
session_key = self._effective_session_key(msg)
|
||||
if session_key != msg.session_key:
|
||||
msg = dataclasses.replace(msg, session_key_override=session_key)
|
||||
recovery_task_registered = False
|
||||
recovery_admission = self._recovery_admission
|
||||
current_task: asyncio.Task[Any] | None = None
|
||||
if recovery_admission is not None:
|
||||
recovery_id = msg.metadata.get(RECOVERY_INBOUND_METADATA_KEY)
|
||||
if isinstance(recovery_id, str) and recovery_id:
|
||||
current_task = asyncio.current_task()
|
||||
if current_task is not None:
|
||||
recovery_admission.register_recovery_task(session_key, current_task)
|
||||
recovery_task_registered = True
|
||||
if not await recovery_admission.admit(msg):
|
||||
logger.info("Skipped stale recovery for session {}", session_key)
|
||||
if recovery_task_registered and current_task is not None:
|
||||
recovery_admission.unregister_recovery_task(session_key, current_task)
|
||||
return
|
||||
lock = self._get_session_lock(session_key)
|
||||
gate = self._concurrency_gate or nullcontext()
|
||||
|
||||
@@ -1380,7 +1449,10 @@ class AgentLoop:
|
||||
# _emit_checkpoint during tool execution; materializing
|
||||
# it into session history now makes it visible in the
|
||||
# next conversation turn.
|
||||
if session_key in self._discarding_sessions:
|
||||
if (
|
||||
session_key in self._discarding_sessions
|
||||
or self._preserve_inflight_turns_on_shutdown
|
||||
):
|
||||
raise
|
||||
try:
|
||||
key = self._effective_session_key(msg)
|
||||
@@ -1437,6 +1509,12 @@ class AgentLoop:
|
||||
await delivery.idle()
|
||||
await self._publish_next_deferred_automation_turn(session_key)
|
||||
finally:
|
||||
if (
|
||||
recovery_task_registered
|
||||
and current_task is not None
|
||||
and recovery_admission is not None
|
||||
):
|
||||
recovery_admission.unregister_recovery_task(session_key, current_task)
|
||||
if pending is None:
|
||||
await delivery.idle()
|
||||
await self._publish_next_deferred_automation_turn(session_key)
|
||||
@@ -1738,7 +1816,10 @@ class AgentLoop:
|
||||
|
||||
if self._restore_runtime_checkpoint(session):
|
||||
self.sessions.save(session)
|
||||
if self._restore_pending_user_turn(session):
|
||||
if (
|
||||
RECOVERY_INBOUND_METADATA_KEY not in msg.metadata
|
||||
and restore_pending_interruption(session)
|
||||
):
|
||||
self.sessions.save(session)
|
||||
|
||||
async def _compact_session(self, ctx: TurnContext) -> None:
|
||||
@@ -2093,8 +2174,21 @@ class AgentLoop:
|
||||
if m.get("role") == "tool" and m.get("tool_call_id")
|
||||
}
|
||||
last_assistant_idx: int | None = None
|
||||
saved_followup_ids: set[str] = set()
|
||||
for m in messages[skip:]:
|
||||
entry = dict(m)
|
||||
followup_id_value = cast(object, entry.pop(PENDING_FOLLOWUP_ID_KEY, None))
|
||||
followup_ids = (
|
||||
[followup_id_value]
|
||||
if isinstance(followup_id_value, str)
|
||||
else [
|
||||
followup_id
|
||||
for followup_id in cast(list[object], followup_id_value)
|
||||
if isinstance(followup_id, str)
|
||||
]
|
||||
if isinstance(followup_id_value, list)
|
||||
else []
|
||||
)
|
||||
internal_meta = cast(object, entry.pop("_meta", None))
|
||||
runtime_context_meta = (
|
||||
cast(dict[str, Any], internal_meta).get(
|
||||
@@ -2147,6 +2241,8 @@ class AgentLoop:
|
||||
entry[RUNTIME_CONTEXT_HISTORY_META] = runtime_context_meta
|
||||
entry.setdefault("timestamp", datetime.now().isoformat())
|
||||
session.messages.append(entry)
|
||||
if role == "user":
|
||||
saved_followup_ids.update(followup_id for followup_id in followup_ids if followup_id)
|
||||
if role == "assistant":
|
||||
last_assistant_idx = len(session.messages) - 1
|
||||
declared_tool_call_ids.update(
|
||||
@@ -2161,6 +2257,8 @@ class AgentLoop:
|
||||
)
|
||||
if turn_latency_ms is not None and last_assistant_idx is not None:
|
||||
session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms)
|
||||
if saved_followup_ids:
|
||||
acknowledge_pending_followups(session, saved_followup_ids)
|
||||
session.updated_at = datetime.now()
|
||||
|
||||
def _persist_subagent_followup(self, session: Session, msg: InboundMessage) -> bool:
|
||||
@@ -2195,7 +2293,7 @@ class AgentLoop:
|
||||
def _set_runtime_checkpoint(self, session: Session, payload: dict[str, Any]) -> None:
|
||||
"""Persist the latest in-flight turn state into session metadata."""
|
||||
session.metadata[self._RUNTIME_CHECKPOINT_KEY] = payload
|
||||
self.sessions.save(session)
|
||||
self.sessions.save_runtime_checkpoint(session)
|
||||
|
||||
def _mark_pending_user_turn(self, session: Session) -> None:
|
||||
session.metadata[self._PENDING_USER_TURN_KEY] = True
|
||||
@@ -2207,136 +2305,9 @@ class AgentLoop:
|
||||
if self._RUNTIME_CHECKPOINT_KEY in session.metadata:
|
||||
session.metadata.pop(self._RUNTIME_CHECKPOINT_KEY, None)
|
||||
|
||||
@staticmethod
|
||||
def _checkpoint_message_key(message: dict[str, Any]) -> tuple[Any, ...]:
|
||||
return (
|
||||
message.get("role"),
|
||||
message.get("content"),
|
||||
message.get("tool_call_id"),
|
||||
message.get("name"),
|
||||
message.get("tool_calls"),
|
||||
message.get("reasoning_content"),
|
||||
message.get("thinking_blocks"),
|
||||
)
|
||||
|
||||
def _restore_runtime_checkpoint(self, session: Session) -> bool:
|
||||
"""Materialize an unfinished turn into session history before a new request."""
|
||||
from datetime import datetime
|
||||
|
||||
checkpoint = cast(
|
||||
object,
|
||||
session.metadata.get(self._RUNTIME_CHECKPOINT_KEY),
|
||||
)
|
||||
if not isinstance(checkpoint, dict):
|
||||
return False
|
||||
checkpoint_data = cast(dict[str, Any], checkpoint)
|
||||
|
||||
assistant_message = cast(object, checkpoint_data.get("assistant_message"))
|
||||
completed_tool_results = cast(
|
||||
Iterable[object],
|
||||
checkpoint_data.get("completed_tool_results") or [],
|
||||
)
|
||||
pending_tool_calls = cast(
|
||||
Iterable[object],
|
||||
checkpoint_data.get("pending_tool_calls") or [],
|
||||
)
|
||||
|
||||
restored_messages: list[dict[str, Any]] = []
|
||||
if isinstance(assistant_message, dict):
|
||||
restored = dict(cast(dict[str, Any], assistant_message))
|
||||
restored.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored_messages.append(restored)
|
||||
for message in completed_tool_results:
|
||||
if isinstance(message, dict):
|
||||
restored = dict(cast(dict[str, Any], message))
|
||||
restored.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored_messages.append(restored)
|
||||
for tool_call in pending_tool_calls:
|
||||
if not isinstance(tool_call, dict):
|
||||
continue
|
||||
tool_call_data = cast(dict[str, Any], tool_call)
|
||||
tool_id = tool_call_data.get("id")
|
||||
function_data = cast(
|
||||
dict[str, Any],
|
||||
tool_call_data.get("function") or {},
|
||||
)
|
||||
name = function_data.get("name") or "tool"
|
||||
restored_messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_id,
|
||||
"name": name,
|
||||
"content": "Error: Task interrupted before this tool finished.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
overlap = 0
|
||||
max_overlap = min(len(session.messages), len(restored_messages))
|
||||
for size in range(max_overlap, 0, -1):
|
||||
existing = session.messages[-size:]
|
||||
restored = restored_messages[:size]
|
||||
if all(
|
||||
self._checkpoint_message_key(left) == self._checkpoint_message_key(right)
|
||||
for left, right in zip(existing, restored)
|
||||
):
|
||||
overlap = size
|
||||
break
|
||||
appended_messages = restored_messages[overlap:]
|
||||
session.messages.extend(appended_messages)
|
||||
assistant_message_data = (
|
||||
cast(dict[str, Any], assistant_message)
|
||||
if isinstance(assistant_message, dict)
|
||||
else None
|
||||
)
|
||||
provider_state_is_synchronized = (
|
||||
checkpoint_data.get(self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY)
|
||||
== self._PROVIDER_STATE_CHECKPOINT_VERSION
|
||||
)
|
||||
phase = checkpoint_data.get("phase")
|
||||
exact_final_response = (
|
||||
phase == "final_response"
|
||||
and assistant_message_data is not None
|
||||
and assistant_message_data.get("role") == "assistant"
|
||||
and not bool(checkpoint_data.get("completed_tool_results"))
|
||||
and not bool(checkpoint_data.get("pending_tool_calls"))
|
||||
)
|
||||
exact_completed_tools = (
|
||||
phase == "tools_completed"
|
||||
and assistant_message_data is not None
|
||||
and assistant_message_data.get("role") == "assistant"
|
||||
and not bool(checkpoint_data.get("pending_tool_calls"))
|
||||
)
|
||||
if not (
|
||||
provider_state_is_synchronized
|
||||
and (exact_final_response or exact_completed_tools)
|
||||
):
|
||||
session.provider_state = None
|
||||
|
||||
self._clear_pending_user_turn(session)
|
||||
self._clear_runtime_checkpoint(session)
|
||||
return True
|
||||
|
||||
def _restore_pending_user_turn(self, session: Session) -> bool:
|
||||
"""Close a turn that only persisted the user message before crashing."""
|
||||
from datetime import datetime
|
||||
|
||||
if not session.metadata.get(self._PENDING_USER_TURN_KEY):
|
||||
return False
|
||||
|
||||
if session.messages and session.messages[-1].get("role") == "user":
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Error: Task interrupted before a response was generated.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
session.provider_state = None
|
||||
session.updated_at = datetime.now()
|
||||
|
||||
self._clear_pending_user_turn(session)
|
||||
return True
|
||||
return restore_runtime_checkpoint(session)
|
||||
|
||||
async def process_direct(
|
||||
self,
|
||||
|
||||
@@ -37,6 +37,7 @@ from nanobot.runtime_context import (
|
||||
reattach_runtime_context,
|
||||
)
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.session.recovery import PENDING_FOLLOWUP_ID_KEY
|
||||
from nanobot.utils.helpers import (
|
||||
IncrementalThinkExtractor,
|
||||
build_assistant_message,
|
||||
@@ -234,6 +235,23 @@ class AgentRunner:
|
||||
merged.get("content"),
|
||||
injection.get("content"),
|
||||
)
|
||||
followup_id = injection.get(PENDING_FOLLOWUP_ID_KEY)
|
||||
if isinstance(followup_id, str) and followup_id:
|
||||
existing = cast(object, merged.get(PENDING_FOLLOWUP_ID_KEY))
|
||||
followup_ids = (
|
||||
[existing]
|
||||
if isinstance(existing, str)
|
||||
else [
|
||||
item
|
||||
for item in cast(list[object], existing)
|
||||
if isinstance(item, str)
|
||||
]
|
||||
if isinstance(existing, list)
|
||||
else []
|
||||
)
|
||||
if followup_id not in followup_ids:
|
||||
followup_ids.append(followup_id)
|
||||
merged[PENDING_FOLLOWUP_ID_KEY] = followup_ids
|
||||
messages[-1] = merged
|
||||
continue
|
||||
messages.append(injection)
|
||||
|
||||
@@ -62,6 +62,15 @@ class TurnEndEvent(OutboundEvent):
|
||||
context_window_tokens: int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RecoveryStateEvent(OutboundEvent):
|
||||
status: str
|
||||
recovery_id: str
|
||||
reason: str | None = None
|
||||
attempts: int = 0
|
||||
can_continue: bool | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GoalStatusEvent(OutboundEvent):
|
||||
status: str
|
||||
|
||||
@@ -104,6 +104,9 @@ class ChannelManager:
|
||||
webui_mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
|
||||
webui_mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
|
||||
webui_skill_state_action: Callable[[set[str]], None] | None = None,
|
||||
webui_recovery_action: (
|
||||
Callable[[str, dict[str, Any]], Awaitable[dict[str, Any]]] | None
|
||||
) = None,
|
||||
config_path: Path | None = None,
|
||||
):
|
||||
if config_path is None:
|
||||
@@ -126,6 +129,7 @@ class ChannelManager:
|
||||
self._webui_mcp_runtime_status = webui_mcp_runtime_status
|
||||
self._webui_mcp_reload = webui_mcp_reload
|
||||
self._webui_skill_state_action = webui_skill_state_action
|
||||
self._webui_recovery_action = webui_recovery_action
|
||||
self.channels: dict[str, BaseChannel] = {}
|
||||
self._channel_owners: dict[str, str] = {}
|
||||
self._channel_runtime_specs: dict[str, tuple[str, str]] = {}
|
||||
@@ -197,6 +201,7 @@ class ChannelManager:
|
||||
mcp_runtime_status=self._webui_mcp_runtime_status,
|
||||
mcp_reload=self._webui_mcp_reload,
|
||||
skill_state_action=self._webui_skill_state_action,
|
||||
recovery_action=self._webui_recovery_action,
|
||||
logger=logger,
|
||||
)
|
||||
kwargs["gateway"] = gateway
|
||||
@@ -615,6 +620,12 @@ class ChannelManager:
|
||||
if target is None:
|
||||
logger.warning("Restart notice target channel is not enabled: {}", notice.channel)
|
||||
return
|
||||
if notice.channel == "websocket":
|
||||
# Reconnect and recovery are already represented by WebSocket
|
||||
# protocol state. A generic restart-complete notice must not
|
||||
# masquerade as a recovery transition and overwrite a real
|
||||
# awaiting-user checkpoint in connected clients.
|
||||
return
|
||||
|
||||
while not target.is_running:
|
||||
remaining = deadline - loop.time()
|
||||
|
||||
@@ -32,6 +32,7 @@ from nanobot.bus.outbound_events import (
|
||||
GoalStateSyncEvent,
|
||||
GoalStatusEvent,
|
||||
ProgressEvent,
|
||||
RecoveryStateEvent,
|
||||
RuntimeModelUpdatedEvent,
|
||||
SessionUpdatedEvent,
|
||||
TurnEndEvent,
|
||||
@@ -55,6 +56,7 @@ from nanobot.security.workspace_access import (
|
||||
)
|
||||
from nanobot.session.goal_state import goal_state_ws_blob
|
||||
from nanobot.session.model_selection import model_preset_from_metadata
|
||||
from nanobot.session.recovery import recovery_state_from_metadata
|
||||
from nanobot.session.webui_turns import (
|
||||
clear_websocket_turn_if_current,
|
||||
clear_websocket_turns,
|
||||
@@ -453,6 +455,9 @@ class WebSocketChannel(BaseChannel):
|
||||
self.logger.warning("ignoring invalid model preset metadata for chat_id={}", chat_id)
|
||||
fields["model_preset"] = None
|
||||
if isinstance(metadata, dict):
|
||||
recovery_state = recovery_state_from_metadata(metadata)
|
||||
if recovery_state is not None:
|
||||
fields["recovery_state"] = recovery_state
|
||||
usage = metadata.get("_last_usage")
|
||||
if isinstance(usage, dict):
|
||||
sanitized_usage: dict[str, int | float] = {}
|
||||
@@ -1740,6 +1745,10 @@ class WebSocketChannel(BaseChannel):
|
||||
provenance=event.provenance,
|
||||
)
|
||||
return
|
||||
if isinstance(event, RecoveryStateEvent):
|
||||
if conns:
|
||||
await self.send_recovery_state(msg.chat_id, event)
|
||||
return
|
||||
if isinstance(event, GoalStateSyncEvent):
|
||||
if conns:
|
||||
await self.send_goal_state(msg.chat_id, event.goal_state or {"active": False})
|
||||
@@ -2057,6 +2066,27 @@ class WebSocketChannel(BaseChannel):
|
||||
for connection in conns:
|
||||
await self._safe_send_to(connection, raw, label=" turn_end ")
|
||||
|
||||
async def send_recovery_state(
|
||||
self,
|
||||
chat_id: str,
|
||||
event: RecoveryStateEvent,
|
||||
) -> None:
|
||||
"""Publish one structured recovery transition without chat pollution."""
|
||||
body: dict[str, Any] = {
|
||||
"event": "recovery_state",
|
||||
"chat_id": chat_id,
|
||||
"status": event.status,
|
||||
"recovery_id": event.recovery_id,
|
||||
"attempts": event.attempts,
|
||||
}
|
||||
if event.reason:
|
||||
body["reason"] = event.reason
|
||||
if event.can_continue is not None:
|
||||
body["can_continue"] = event.can_continue
|
||||
raw = json.dumps(body, ensure_ascii=False)
|
||||
for connection in list(self._subs.get(chat_id, ())):
|
||||
await self._safe_send_to(connection, raw, label=" recovery_state ")
|
||||
|
||||
async def send_goal_state(self, chat_id: str, blob: dict[str, Any]) -> None:
|
||||
"""Push persisted goal-state snapshot for *chat_id* (multi-chat isolation)."""
|
||||
conns = list(self._subs.get(chat_id, ()))
|
||||
|
||||
@@ -27,6 +27,7 @@ from nanobot.bus.outbound_events import (
|
||||
GoalStateSyncEvent,
|
||||
GoalStatusEvent,
|
||||
ProgressEvent,
|
||||
RecoveryStateEvent,
|
||||
RuntimeModelUpdatedEvent,
|
||||
SessionUpdatedEvent,
|
||||
TurnEndEvent,
|
||||
@@ -2720,6 +2721,39 @@ async def test_send_turn_end_emits_turn_end_event() -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recovery_state_is_a_structured_event_not_assistant_text() -> None:
|
||||
bus = MagicMock()
|
||||
channel = WebSocketChannel(
|
||||
{"enabled": True, "allowFrom": ["*"]},
|
||||
bus,
|
||||
gateway=_basic_handler(bus),
|
||||
)
|
||||
mock_ws = AsyncMock()
|
||||
channel._attach(mock_ws, "chat-1")
|
||||
|
||||
await channel.send(OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
event=RecoveryStateEvent(
|
||||
status="awaiting_user",
|
||||
recovery_id="recovery-1",
|
||||
reason="tool_state_unknown",
|
||||
attempts=1,
|
||||
),
|
||||
))
|
||||
|
||||
assert _sent_ws_payloads(mock_ws) == [{
|
||||
"event": "recovery_state",
|
||||
"chat_id": "chat-1",
|
||||
"status": "awaiting_user",
|
||||
"recovery_id": "recovery-1",
|
||||
"reason": "tool_state_unknown",
|
||||
"attempts": 1,
|
||||
}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_system_command_turn_end_only_refreshes_session_metadata() -> None:
|
||||
bus = MagicMock()
|
||||
|
||||
@@ -83,6 +83,7 @@ def _make_handler(
|
||||
channel_feature_action: Any | None = None,
|
||||
channel_runtime_status: Any | None = None,
|
||||
mcp_reload: Any | None = None,
|
||||
recovery_action: Any | None = None,
|
||||
) -> GatewayServices:
|
||||
config = WebSocketConfig.model_validate(cfg) if isinstance(cfg, dict) else cfg
|
||||
workspace = workspace_path or Path.cwd()
|
||||
@@ -103,6 +104,7 @@ def _make_handler(
|
||||
channel_feature_action=channel_feature_action,
|
||||
channel_runtime_status=channel_runtime_status,
|
||||
mcp_reload=mcp_reload,
|
||||
recovery_action=recovery_action,
|
||||
)
|
||||
|
||||
|
||||
@@ -121,6 +123,7 @@ def _ch(
|
||||
channel_feature_action: Any | None = None,
|
||||
channel_runtime_status: Any | None = None,
|
||||
mcp_reload: Any | None = None,
|
||||
recovery_action: Any | None = None,
|
||||
**extra: Any,
|
||||
) -> WebSocketChannel:
|
||||
cfg: dict[str, Any] = {
|
||||
@@ -145,6 +148,7 @@ def _ch(
|
||||
channel_feature_action=channel_feature_action,
|
||||
channel_runtime_status=channel_runtime_status,
|
||||
mcp_reload=mcp_reload,
|
||||
recovery_action=recovery_action,
|
||||
)
|
||||
return InProcessHttpChannel(cfg, bus, gateway=gateway)
|
||||
|
||||
@@ -3242,6 +3246,28 @@ async def _webui_mutate(
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recovery_mutation_uses_authenticated_websocket_action(bus: MagicMock) -> None:
|
||||
recovery_action = AsyncMock(return_value={
|
||||
"status": "resuming",
|
||||
"recovery_id": "recovery-1",
|
||||
})
|
||||
channel = _ch(bus, recovery_action=recovery_action)
|
||||
|
||||
response = await _webui_mutate(
|
||||
channel,
|
||||
"recovery.continue",
|
||||
{"chat_id": "chat-1", "recovery_id": "recovery-1"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["status"] == "resuming"
|
||||
recovery_action.assert_awaited_once_with(
|
||||
"continue",
|
||||
{"chat_id": "chat-1", "recovery_id": "recovery-1"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_workspace_folder_picker_is_local_authenticated_mutation(
|
||||
bus: MagicMock,
|
||||
|
||||
@@ -322,6 +322,7 @@ def _run_gateway(
|
||||
from nanobot.providers.fallback_provider import FallbackProvider
|
||||
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.session.recovery import RecoveryCoordinator
|
||||
from nanobot.session.webui_turns import (
|
||||
WebuiTurnCoordinator,
|
||||
WebuiTurnRoutePolicy,
|
||||
@@ -422,6 +423,12 @@ def _run_gateway(
|
||||
tools = ToolRegistry()
|
||||
mcp_provider = MCPProvider.from_config(config, tools)
|
||||
|
||||
recovery = RecoveryCoordinator(
|
||||
sessions=session_manager,
|
||||
bus=bus,
|
||||
unified_session=config.agents.defaults.unified_session,
|
||||
)
|
||||
|
||||
# Create agent with cron service
|
||||
agent = AgentLoop.from_config(
|
||||
config, bus,
|
||||
@@ -440,6 +447,7 @@ def _run_gateway(
|
||||
local_trigger_store=trigger_store,
|
||||
hook_factories=[create_file_edit_activity_hook],
|
||||
tool_registry=tools,
|
||||
recovery_admission=recovery,
|
||||
)
|
||||
def _schedule_webui_background(awaitable: Awaitable[None]) -> None:
|
||||
agent.schedule_background(cast(Coroutine[Any, Any, None], awaitable))
|
||||
@@ -448,6 +456,7 @@ def _run_gateway(
|
||||
bus=bus,
|
||||
sessions=session_manager,
|
||||
schedule_background=_schedule_webui_background,
|
||||
recovery=recovery,
|
||||
)
|
||||
webui_turn_coordinator.subscribe(runtime_events)
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
@@ -683,6 +692,7 @@ def _run_gateway(
|
||||
webui_mcp_runtime_status=mcp_provider.runtime_status,
|
||||
webui_mcp_reload=mcp_provider.reload,
|
||||
webui_skill_state_action=_webui_skill_state_action,
|
||||
webui_recovery_action=recovery.handle_action,
|
||||
config_path=Path(config_path),
|
||||
)
|
||||
|
||||
@@ -849,6 +859,7 @@ def _run_gateway(
|
||||
tasks: list[asyncio.Task[Any]] = []
|
||||
shutdown_task: asyncio.Task[Any] | None = None
|
||||
runtime_tasks: asyncio.Future[list[Any]] | None = None
|
||||
startup_complete = False
|
||||
shutdown_event = asyncio.Event()
|
||||
cli_terminal._ensure_interactive_tty_mode()
|
||||
restore_shutdown_handlers = _install_gateway_shutdown_handlers(
|
||||
@@ -861,6 +872,10 @@ def _run_gateway(
|
||||
await cron.start()
|
||||
# Re-read once on first admission to close the watcher subscription window.
|
||||
agent.runtime_resolver.invalidate()
|
||||
# Recovery must finish before WebSocket and other channels begin
|
||||
# accepting new input. That makes a new user message reliably
|
||||
# supersede an old recoverable turn instead of racing its queue.
|
||||
await recovery.scan()
|
||||
async def _run_agent() -> None:
|
||||
try:
|
||||
await mcp_provider.connect()
|
||||
@@ -915,6 +930,7 @@ def _run_gateway(
|
||||
name="nanobot-webui-dev-server",
|
||||
))
|
||||
runtime_tasks = asyncio.gather(*tasks)
|
||||
startup_complete = True
|
||||
shutdown_task = asyncio.create_task(
|
||||
shutdown_event.wait(),
|
||||
name="nanobot-gateway-shutdown",
|
||||
@@ -936,6 +952,10 @@ def _run_gateway(
|
||||
|
||||
console.print("\n[red]Error: Gateway crashed unexpectedly[/red]")
|
||||
console.print(traceback.format_exc())
|
||||
if not startup_complete:
|
||||
# Do not report a successful gateway command when startup
|
||||
# failed before any runtime task or listener was created.
|
||||
raise typer.Exit(1)
|
||||
finally:
|
||||
try:
|
||||
if shutdown_task and not shutdown_task.done():
|
||||
@@ -943,6 +963,8 @@ def _run_gateway(
|
||||
with suppress(asyncio.CancelledError):
|
||||
await shutdown_task
|
||||
cron.stop()
|
||||
if gateway_runtime.preserves_inflight_turns_on_exit():
|
||||
agent.preserve_inflight_turns_on_shutdown()
|
||||
agent.stop()
|
||||
# Cancel runtime tasks first, then deterministically close
|
||||
# exec/MCP resources while the event loop is still alive.
|
||||
|
||||
@@ -201,6 +201,52 @@ class GatewayRuntime(ManagedProcessRuntime[ProcessStartOptions]):
|
||||
"""Serialize long lifecycle transitions without blocking child cleanup."""
|
||||
return FileLock(f"{self.paths.state_path}.transition.lock")
|
||||
|
||||
@property
|
||||
def _restart_intent_path(self) -> Path:
|
||||
"""Return the short-lived marker used to distinguish restart from stop.
|
||||
|
||||
A gateway receives the same operating-system termination request for a
|
||||
graceful ``restart`` and an explicit ``stop``. The marker lets the
|
||||
exiting process preserve its durable turn checkpoint only for the
|
||||
former. It is intentionally local to one gateway instance.
|
||||
"""
|
||||
return self.paths.state_path.with_name(f"{self.paths.state_path.name}.restart")
|
||||
|
||||
def preserves_inflight_turns_on_exit(self) -> bool:
|
||||
"""Whether this gateway was asked to exit as part of a managed restart."""
|
||||
try:
|
||||
raw_intent: object = json.loads(
|
||||
self._restart_intent_path.read_text(encoding="utf-8")
|
||||
)
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return False
|
||||
if not isinstance(raw_intent, dict):
|
||||
return False
|
||||
intent = cast(dict[str, object], raw_intent)
|
||||
return intent.get("pid") == os.getpid()
|
||||
|
||||
def _write_restart_intent(self, pid: int) -> None:
|
||||
self.paths.run_dir.mkdir(parents=True, exist_ok=True)
|
||||
target = self._restart_intent_path
|
||||
with tempfile.NamedTemporaryFile(
|
||||
"w",
|
||||
encoding="utf-8",
|
||||
dir=target.parent,
|
||||
prefix=f".{target.name}.",
|
||||
delete=False,
|
||||
) as handle:
|
||||
json.dump({"pid": pid}, handle)
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
temporary = Path(handle.name)
|
||||
os.replace(temporary, target)
|
||||
|
||||
def _clear_restart_intent(self) -> None:
|
||||
try:
|
||||
self._restart_intent_path.unlink()
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
def start_background(self, options: ProcessStartOptions) -> RuntimeResult:
|
||||
"""Start the gateway detached from the current terminal."""
|
||||
lease = GatewayClientLease(self, kind="gateway-background")
|
||||
@@ -353,7 +399,15 @@ class GatewayRuntime(ManagedProcessRuntime[ProcessStartOptions]):
|
||||
"gateway_foreground_restart_required",
|
||||
status,
|
||||
)
|
||||
stop_result = self._stop(timeout_s=timeout_s)
|
||||
assert status.pid is not None
|
||||
self._write_restart_intent(status.pid)
|
||||
try:
|
||||
stop_result = self._stop(timeout_s=timeout_s)
|
||||
finally:
|
||||
# The old process reads the marker while handling shutdown.
|
||||
# Never let a stale marker turn a later explicit stop into a
|
||||
# recoverable restart.
|
||||
self._clear_restart_intent()
|
||||
if not stop_result.ok:
|
||||
return self._result(stop_result)
|
||||
with self._lifecycle_lock():
|
||||
|
||||
+135
-2
@@ -48,15 +48,21 @@ _SESSION_PREVIEW_MAX_CHARS = 120
|
||||
_SESSION_LIST_PREVIEW_MAX_RECORDS = 200
|
||||
_SESSION_LIST_PREVIEW_MAX_CHARS = 1_000_000
|
||||
_SESSION_DATA_ERRORS = (ValueError, TypeError, AttributeError, KeyError)
|
||||
_RUNTIME_CHECKPOINT_DATA_ERRORS = (OSError, *_SESSION_DATA_ERRORS)
|
||||
_PROVIDER_STATE_RECORD_TYPE = "provider_state"
|
||||
_PROVIDER_STATE_RECORD_PREFIX_RE = re.compile(
|
||||
r'^\s*\{\s*"_type"\s*:\s*"provider_state"\s*(?:,|\})'
|
||||
)
|
||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||
_RUNTIME_CHECKPOINT_VERSION = 1
|
||||
_RUNTIME_CHECKPOINT_SUFFIX = ".checkpoint.json"
|
||||
_FORK_VOLATILE_METADATA_KEYS = {
|
||||
"goal_state",
|
||||
"pending_user_turn",
|
||||
"pending_user_followups",
|
||||
"runtime_checkpoint",
|
||||
"session_handle",
|
||||
"webui_recovery",
|
||||
"thread_goal",
|
||||
"title",
|
||||
"title_user_edited",
|
||||
@@ -1001,6 +1007,9 @@ class JsonlSessionStore:
|
||||
def get_session_path(self, key: str) -> Path:
|
||||
return self.sessions_dir / f"{self.storage_key(key)}.jsonl"
|
||||
|
||||
def get_runtime_checkpoint_path(self, key: str) -> Path:
|
||||
return self.sessions_dir / f"{self.storage_key(key)}{_RUNTIME_CHECKPOINT_SUFFIX}"
|
||||
|
||||
def get_legacy_lossy_path(self, key: str) -> Path:
|
||||
return self.sessions_dir / f"{safe_filename(key.replace(':', '_'))}.jsonl"
|
||||
|
||||
@@ -1066,7 +1075,7 @@ class JsonlSessionStore:
|
||||
else:
|
||||
messages.append(data)
|
||||
|
||||
return Session(
|
||||
session = Session(
|
||||
key=key,
|
||||
messages=messages,
|
||||
created_at=created_at or datetime.now(),
|
||||
@@ -1075,6 +1084,8 @@ class JsonlSessionStore:
|
||||
last_consolidated=last_consolidated,
|
||||
provider_state=provider_state,
|
||||
)
|
||||
self._overlay_runtime_checkpoint_unlocked(session, path)
|
||||
return session
|
||||
except _SESSION_DATA_ERRORS as e:
|
||||
logger.warning("Failed to load session {}: {}", key, e)
|
||||
repaired = self._repair_unlocked(key)
|
||||
@@ -1159,7 +1170,7 @@ class JsonlSessionStore:
|
||||
if not messages and not metadata and provider_state is None:
|
||||
return None
|
||||
|
||||
return Session(
|
||||
session = Session(
|
||||
key=key,
|
||||
messages=messages,
|
||||
created_at=created_at or datetime.now(),
|
||||
@@ -1168,6 +1179,8 @@ class JsonlSessionStore:
|
||||
last_consolidated=last_consolidated,
|
||||
provider_state=provider_state,
|
||||
)
|
||||
self._overlay_runtime_checkpoint_unlocked(session, path)
|
||||
return session
|
||||
except _SESSION_DATA_ERRORS as e:
|
||||
logger.warning("Repair failed for session {}: {}", key, e)
|
||||
return None
|
||||
@@ -1186,6 +1199,105 @@ class JsonlSessionStore:
|
||||
with self._session_files_lock:
|
||||
self._save_unlocked(session, fsync=fsync)
|
||||
|
||||
def save_runtime_checkpoint(self, session: Session) -> None:
|
||||
"""Atomically persist only the volatile in-flight turn state.
|
||||
|
||||
A checkpoint is written several times during a tool-heavy turn. Keeping it
|
||||
beside the append history avoids copying the full transcript at each safe
|
||||
recovery boundary.
|
||||
"""
|
||||
with self._session_files_lock:
|
||||
path = self.get_session_path(session.key)
|
||||
if not path.exists():
|
||||
# A user turn normally creates the session first. Internal callers
|
||||
# may checkpoint a fresh session, so establish the durable base once.
|
||||
self._save_unlocked(session)
|
||||
return
|
||||
|
||||
checkpoint = session.metadata.get(_RUNTIME_CHECKPOINT_KEY)
|
||||
if not isinstance(checkpoint, dict):
|
||||
self.get_runtime_checkpoint_path(session.key).unlink(missing_ok=True)
|
||||
return
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"version": _RUNTIME_CHECKPOINT_VERSION,
|
||||
"session_key": session.key,
|
||||
"base_updated_at": session.updated_at.isoformat(),
|
||||
"base_message_count": len(session.messages),
|
||||
"checkpoint": checkpoint,
|
||||
"provider_state": (
|
||||
session.provider_state.to_private_record()
|
||||
if session.provider_state is not None
|
||||
else None
|
||||
),
|
||||
}
|
||||
target = self.get_runtime_checkpoint_path(session.key)
|
||||
tmp = target.with_name(f".{target.name}.{secrets.token_hex(8)}.tmp")
|
||||
try:
|
||||
with open(tmp, "x", encoding="utf-8") as handle:
|
||||
os.chmod(tmp, 0o600)
|
||||
json.dump(
|
||||
payload,
|
||||
handle,
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
os.replace(tmp, target)
|
||||
finally:
|
||||
tmp.unlink(missing_ok=True)
|
||||
|
||||
def _overlay_runtime_checkpoint_unlocked(self, session: Session, main_path: Path) -> None:
|
||||
checkpoint_path = self.get_runtime_checkpoint_path(session.key)
|
||||
try:
|
||||
checkpoint_stat = checkpoint_path.lstat()
|
||||
if not stat.S_ISREG(checkpoint_stat.st_mode):
|
||||
logger.warning(
|
||||
"Ignoring non-regular runtime checkpoint for session {}",
|
||||
session.key,
|
||||
)
|
||||
return
|
||||
# A complete session save supersedes an older sidecar. This comparison
|
||||
# closes the small crash window between replacing the JSONL and unlinking
|
||||
# its previous checkpoint.
|
||||
if main_path.stat().st_mtime_ns > checkpoint_stat.st_mtime_ns:
|
||||
checkpoint_path.unlink(missing_ok=True)
|
||||
return
|
||||
raw = _json_object(json.loads(checkpoint_path.read_text(encoding="utf-8")))
|
||||
if (
|
||||
raw.get("version") != _RUNTIME_CHECKPOINT_VERSION
|
||||
or raw.get("session_key") != session.key
|
||||
or raw.get("base_updated_at") != session.updated_at.isoformat()
|
||||
or raw.get("base_message_count") != len(session.messages)
|
||||
or not isinstance(raw.get("checkpoint"), dict)
|
||||
):
|
||||
checkpoint_path.unlink(missing_ok=True)
|
||||
return
|
||||
provider_record = raw.get("provider_state")
|
||||
provider_state = (
|
||||
None
|
||||
if provider_record is None
|
||||
else ProviderConversationState.from_private_record(provider_record)
|
||||
)
|
||||
if provider_record is not None and provider_state is None:
|
||||
raise ValueError("invalid checkpoint provider state")
|
||||
session.metadata[_RUNTIME_CHECKPOINT_KEY] = cast(
|
||||
dict[str, Any], raw["checkpoint"]
|
||||
)
|
||||
session.provider_state = provider_state
|
||||
except FileNotFoundError:
|
||||
return
|
||||
except _RUNTIME_CHECKPOINT_DATA_ERRORS as exc:
|
||||
logger.warning(
|
||||
"Ignoring invalid runtime checkpoint for session {}: {}",
|
||||
session.key,
|
||||
exc,
|
||||
)
|
||||
# Atomic writes mean a malformed target cannot become valid later.
|
||||
# Remove it once so future loads do not repeatedly parse and log it.
|
||||
with suppress(OSError):
|
||||
if checkpoint_path.is_file() and not checkpoint_path.is_symlink():
|
||||
checkpoint_path.unlink()
|
||||
|
||||
def _save_unlocked(self, session: Session, *, fsync: bool = False) -> None:
|
||||
path = self.get_session_path(session.key)
|
||||
tmp_path = path.with_name(f".{path.name}.{secrets.token_hex(8)}.tmp")
|
||||
@@ -1215,6 +1327,10 @@ class JsonlSessionStore:
|
||||
|
||||
os.replace(tmp_path, path)
|
||||
|
||||
# The full record now contains the authoritative checkpoint state (or
|
||||
# its removal), so an older volatile overlay is no longer needed.
|
||||
self.get_runtime_checkpoint_path(session.key).unlink(missing_ok=True)
|
||||
|
||||
if fsync:
|
||||
with suppress(PermissionError):
|
||||
fd = os.open(str(path.parent), os.O_RDONLY)
|
||||
@@ -1278,6 +1394,7 @@ class JsonlSessionStore:
|
||||
def _delete_unlocked(self, key: str) -> bool:
|
||||
paths = [
|
||||
self.get_session_path(key),
|
||||
self.get_runtime_checkpoint_path(key),
|
||||
self.get_legacy_lossy_path(key),
|
||||
self.get_legacy_session_path(key),
|
||||
]
|
||||
@@ -1585,6 +1702,10 @@ class SessionManager:
|
||||
"""Get the collision-resistant workspace path for a session."""
|
||||
return self._jsonl_store.get_session_path(key)
|
||||
|
||||
def _get_runtime_checkpoint_path(self, key: str) -> Path:
|
||||
"""Get the private in-flight checkpoint path for a session."""
|
||||
return self._jsonl_store.get_runtime_checkpoint_path(key)
|
||||
|
||||
def _get_legacy_lossy_path(self, key: str) -> Path:
|
||||
"""Previous workspace session path using lossy ':' to '_' replacement."""
|
||||
return self._jsonl_store.get_legacy_lossy_path(key)
|
||||
@@ -1653,6 +1774,18 @@ class SessionManager:
|
||||
self._store.save(session, fsync=fsync)
|
||||
self._remember(session)
|
||||
|
||||
def save_runtime_checkpoint(self, session: Session) -> None:
|
||||
"""Persist volatile recovery state without rewriting long history."""
|
||||
if not session.policy.persist:
|
||||
return
|
||||
if self._store is self._jsonl_store:
|
||||
self._jsonl_store.save_runtime_checkpoint(session)
|
||||
self._remember(session)
|
||||
return
|
||||
# Third-party stores keep their existing all-or-nothing semantics until
|
||||
# they opt into a dedicated checkpoint primitive.
|
||||
self.save(session)
|
||||
|
||||
def rename_model_preset(self, old_name: str, new_name: str) -> int:
|
||||
"""Rename a session-scoped model preset across durable and live sessions."""
|
||||
if old_name == new_name:
|
||||
|
||||
@@ -0,0 +1,911 @@
|
||||
"""Durable, side-effect-safe recovery for interrupted WebUI turns.
|
||||
|
||||
The coordinator owns restart policy. AgentLoop only exposes checkpoint
|
||||
materialization and an admission hook, so transport code never has to guess
|
||||
whether an interrupted tool call is safe to replay.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import dataclasses
|
||||
import json
|
||||
from collections.abc import Iterable, Mapping
|
||||
from datetime import datetime
|
||||
from typing import Any, Protocol, cast
|
||||
from uuid import uuid4
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
RecoveryStateEvent,
|
||||
SessionUpdatedEvent,
|
||||
outbound_message_for_event,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.session import turn_continuation
|
||||
from nanobot.session.keys import UNIFIED_SESSION_KEY, last_channel_from_metadata
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
from nanobot.webui.metadata import WEBUI_TURN_METADATA_KEY
|
||||
|
||||
RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||
PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||
RECOVERY_METADATA_KEY = "webui_recovery"
|
||||
RECOVERY_INBOUND_METADATA_KEY = "_webui_recovery_id"
|
||||
PENDING_FOLLOWUPS_KEY = "pending_user_followups"
|
||||
PENDING_FOLLOWUP_ID_KEY = "_recovery_followup_id"
|
||||
PROVIDER_STATE_CHECKPOINT_VERSION_KEY = "provider_state_checkpoint_version"
|
||||
PROVIDER_STATE_CHECKPOINT_VERSION = "v1"
|
||||
|
||||
_RECOVERY_STATUSES = frozenset({"resuming", "awaiting_user", "recovered", "failed"})
|
||||
_UNCERTAIN_TOOL_PHASES = frozenset({"awaiting_tools"})
|
||||
_KNOWN_CHECKPOINT_PHASES = frozenset(
|
||||
{"final_response", "tools_completed", "awaiting_tools", "error"}
|
||||
)
|
||||
|
||||
|
||||
class RecoveryActionError(ValueError):
|
||||
"""A stale or malformed recovery action from an authenticated WebUI."""
|
||||
|
||||
def __init__(self, message: str, *, status: int = 400) -> None:
|
||||
super().__init__(message)
|
||||
self.status = status
|
||||
|
||||
|
||||
class RecoveryAdmission(Protocol):
|
||||
"""Narrow AgentLoop boundary for explicit recovery validation."""
|
||||
|
||||
async def admit(self, message: InboundMessage) -> bool: ...
|
||||
|
||||
def register_recovery_task(self, session_key: str, task: asyncio.Task[Any]) -> None: ...
|
||||
|
||||
def unregister_recovery_task(self, session_key: str, task: asyncio.Task[Any]) -> None: ...
|
||||
|
||||
|
||||
def record_pending_followup(session: Session, message: InboundMessage) -> str | None:
|
||||
"""Durably journal a WebUI follow-up before injecting it into a live turn."""
|
||||
if message.channel != "websocket":
|
||||
return None
|
||||
try:
|
||||
metadata_value: object = json.loads(json.dumps(message.metadata))
|
||||
except (TypeError, ValueError):
|
||||
logger.warning("Skipping non-serializable WebUI follow-up for recovery")
|
||||
return None
|
||||
if not isinstance(metadata_value, dict):
|
||||
return None
|
||||
metadata = cast(dict[str, Any], metadata_value)
|
||||
existing_id = metadata.pop(PENDING_FOLLOWUP_ID_KEY, None)
|
||||
followup_id = (
|
||||
existing_id
|
||||
if isinstance(existing_id, str) and existing_id
|
||||
else uuid4().hex
|
||||
)
|
||||
records = _pending_followup_records(session)
|
||||
if any(record.get("id") == followup_id for record in records):
|
||||
return followup_id
|
||||
records.append(
|
||||
{
|
||||
"id": followup_id,
|
||||
"sender_id": message.sender_id,
|
||||
"chat_id": message.chat_id,
|
||||
"content": message.content,
|
||||
"media": list(message.media or []),
|
||||
"metadata": metadata,
|
||||
}
|
||||
)
|
||||
# This journal is the recovery source of truth, not a mirror of the
|
||||
# bounded in-memory injection queue. A queued turn can receive more
|
||||
# follow-ups than the live queue accepts; dropping older journal entries
|
||||
# would make those acknowledged user messages unrecoverable after a
|
||||
# gateway restart. Entries are removed only once their user rows are
|
||||
# committed by ``acknowledge_pending_followups``.
|
||||
session.metadata[PENDING_FOLLOWUPS_KEY] = records
|
||||
session.updated_at = datetime.now()
|
||||
return followup_id
|
||||
|
||||
|
||||
def pending_followups(session: Session) -> list[InboundMessage]:
|
||||
"""Decode still-unacknowledged follow-ups from durable session metadata."""
|
||||
messages: list[InboundMessage] = []
|
||||
for record in _pending_followup_records(session):
|
||||
followup_id = cast(object, record.get("id"))
|
||||
sender_id = cast(object, record.get("sender_id"))
|
||||
chat_id = cast(object, record.get("chat_id"))
|
||||
content = cast(object, record.get("content"))
|
||||
metadata = cast(object, record.get("metadata"))
|
||||
if (
|
||||
not isinstance(followup_id, str)
|
||||
or not followup_id
|
||||
or not isinstance(sender_id, str)
|
||||
or not sender_id
|
||||
or not isinstance(chat_id, str)
|
||||
or not chat_id
|
||||
):
|
||||
continue
|
||||
if not isinstance(content, str) or not isinstance(metadata, dict):
|
||||
continue
|
||||
media_value = cast(object, record.get("media"))
|
||||
media = (
|
||||
[item for item in cast(list[object], media_value) if isinstance(item, str)]
|
||||
if isinstance(media_value, list)
|
||||
else []
|
||||
)
|
||||
messages.append(
|
||||
InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
content=content,
|
||||
media=media,
|
||||
metadata={**cast(dict[str, Any], metadata), PENDING_FOLLOWUP_ID_KEY: followup_id},
|
||||
session_key_override=session.key,
|
||||
require_existing_session=True,
|
||||
)
|
||||
)
|
||||
return messages
|
||||
|
||||
|
||||
def acknowledge_pending_followups(session: Session, followup_ids: Iterable[str]) -> None:
|
||||
"""Remove journal entries whose user rows were committed to history."""
|
||||
acknowledged = set(followup_ids)
|
||||
if not acknowledged:
|
||||
return
|
||||
records = [record for record in _pending_followup_records(session) if record.get("id") not in acknowledged]
|
||||
if records:
|
||||
session.metadata[PENDING_FOLLOWUPS_KEY] = records
|
||||
else:
|
||||
session.metadata.pop(PENDING_FOLLOWUPS_KEY, None)
|
||||
|
||||
|
||||
def _pending_followup_records(session: Session) -> list[dict[str, Any]]:
|
||||
raw = cast(object, session.metadata.get(PENDING_FOLLOWUPS_KEY))
|
||||
if not isinstance(raw, list):
|
||||
return []
|
||||
values = cast(list[object], raw)
|
||||
return [cast(dict[str, Any], value) for value in values if isinstance(value, dict)]
|
||||
|
||||
|
||||
def _checkpoint_message_key(message: Mapping[str, Any]) -> tuple[Any, ...]:
|
||||
return (
|
||||
message.get("role"),
|
||||
message.get("content"),
|
||||
message.get("tool_call_id"),
|
||||
message.get("name"),
|
||||
message.get("tool_calls"),
|
||||
message.get("reasoning_content"),
|
||||
message.get("thinking_blocks"),
|
||||
)
|
||||
|
||||
|
||||
def _checkpoint_tool_call_ids(
|
||||
value: object,
|
||||
*,
|
||||
result_rows: bool = False,
|
||||
) -> list[str] | None:
|
||||
"""Validate checkpoint tool rows and return their stable IDs."""
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
ids: list[str] = []
|
||||
for raw in cast(list[object], value):
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
row = cast(dict[str, Any], raw)
|
||||
id_key = "tool_call_id" if result_rows else "id"
|
||||
call_id = cast(object, row.get(id_key))
|
||||
if not isinstance(call_id, str) or not call_id:
|
||||
return None
|
||||
if result_rows:
|
||||
if row.get("role") != "tool":
|
||||
return None
|
||||
else:
|
||||
function_value = cast(object, row.get("function"))
|
||||
if not isinstance(function_value, dict):
|
||||
return None
|
||||
function = cast(dict[str, Any], function_value)
|
||||
name = cast(object, function.get("name"))
|
||||
if not isinstance(name, str) or not name:
|
||||
return None
|
||||
ids.append(call_id)
|
||||
return ids if len(ids) == len(set(ids)) else None
|
||||
|
||||
|
||||
def _runtime_checkpoint_is_well_formed(checkpoint: Mapping[str, Any]) -> bool:
|
||||
"""Return whether a checkpoint is safe to offer for continuation.
|
||||
|
||||
Restoration stays tolerant so Dismiss can always clear corrupt state.
|
||||
Continue is stricter: silently dropping a malformed tool result could make
|
||||
the model repeat an external side effect.
|
||||
"""
|
||||
assistant_value = cast(object, checkpoint.get("assistant_message"))
|
||||
if not isinstance(assistant_value, dict):
|
||||
return False
|
||||
assistant = cast(dict[str, Any], assistant_value)
|
||||
if assistant.get("role") != "assistant":
|
||||
return False
|
||||
|
||||
completed_ids = _checkpoint_tool_call_ids(
|
||||
cast(object, checkpoint.get("completed_tool_results")),
|
||||
result_rows=True,
|
||||
)
|
||||
pending_ids = _checkpoint_tool_call_ids(
|
||||
cast(object, checkpoint.get("pending_tool_calls")),
|
||||
)
|
||||
if completed_ids is None or pending_ids is None:
|
||||
return False
|
||||
assistant_calls_value = cast(object, assistant.get("tool_calls"))
|
||||
assistant_call_ids = (
|
||||
[]
|
||||
if assistant_calls_value is None
|
||||
else _checkpoint_tool_call_ids(assistant_calls_value)
|
||||
)
|
||||
if assistant_call_ids is None:
|
||||
return False
|
||||
|
||||
phase = checkpoint.get("phase")
|
||||
if phase == "final_response":
|
||||
content = cast(object, assistant.get("content"))
|
||||
return (
|
||||
isinstance(content, str)
|
||||
and bool(content.strip())
|
||||
and not assistant_call_ids
|
||||
and not completed_ids
|
||||
and not pending_ids
|
||||
)
|
||||
if phase == "awaiting_tools":
|
||||
return (
|
||||
bool(assistant_call_ids)
|
||||
and not completed_ids
|
||||
and len(assistant_call_ids) == len(pending_ids)
|
||||
and set(assistant_call_ids) == set(pending_ids)
|
||||
)
|
||||
if phase == "tools_completed":
|
||||
return (
|
||||
bool(assistant_call_ids)
|
||||
and not pending_ids
|
||||
and len(assistant_call_ids) == len(completed_ids)
|
||||
and set(assistant_call_ids) == set(completed_ids)
|
||||
)
|
||||
# Error checkpoints have no current producer contract. Treat legacy or
|
||||
# future instances as review-only until their exact persisted shape is
|
||||
# specified; guessing here could make a partial side effect repeat.
|
||||
return False
|
||||
|
||||
|
||||
def restore_runtime_checkpoint(session: Session) -> bool:
|
||||
"""Materialize the durable checkpoint exactly once and clear it.
|
||||
|
||||
Pending tool calls become explicit interrupted tool results. They are
|
||||
never executed here. Provider-native state is retained only for the two
|
||||
checkpoint shapes known to be synchronized with persisted history.
|
||||
"""
|
||||
checkpoint = cast(object, session.metadata.get(RUNTIME_CHECKPOINT_KEY))
|
||||
if not isinstance(checkpoint, dict):
|
||||
return False
|
||||
data = cast(dict[str, Any], checkpoint)
|
||||
assistant = cast(object, data.get("assistant_message"))
|
||||
completed_value = cast(object, data.get("completed_tool_results"))
|
||||
pending_value = cast(object, data.get("pending_tool_calls"))
|
||||
completed = cast(list[object], completed_value) if isinstance(completed_value, list) else []
|
||||
pending = cast(list[object], pending_value) if isinstance(pending_value, list) else []
|
||||
|
||||
restored: list[dict[str, Any]] = []
|
||||
if isinstance(assistant, dict):
|
||||
assistant_row = cast(dict[str, Any], assistant)
|
||||
else:
|
||||
assistant_row = {}
|
||||
if assistant_row.get("role") == "assistant":
|
||||
row = dict(assistant_row)
|
||||
row.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored.append(row)
|
||||
for value in completed:
|
||||
if not isinstance(value, dict):
|
||||
continue
|
||||
tool_result = cast(dict[str, Any], value)
|
||||
if tool_result.get("role") != "tool":
|
||||
continue
|
||||
row = dict(tool_result)
|
||||
row.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored.append(row)
|
||||
for value in pending:
|
||||
if not isinstance(value, dict):
|
||||
continue
|
||||
tool_call = cast(dict[str, Any], value)
|
||||
tool_call_id = tool_call.get("id")
|
||||
function_value = cast(object, tool_call.get("function"))
|
||||
if not isinstance(tool_call_id, str) or not tool_call_id:
|
||||
continue
|
||||
function = (
|
||||
cast(dict[str, Any], function_value)
|
||||
if isinstance(function_value, dict)
|
||||
else {}
|
||||
)
|
||||
name = function.get("name")
|
||||
restored.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call_id,
|
||||
"name": name if isinstance(name, str) and name else "tool",
|
||||
"content": "Error: Task interrupted before this tool finished.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"_recovery_interrupted": True,
|
||||
}
|
||||
)
|
||||
|
||||
overlap = 0
|
||||
for size in range(min(len(session.messages), len(restored)), 0, -1):
|
||||
if all(
|
||||
_checkpoint_message_key(left) == _checkpoint_message_key(right)
|
||||
for left, right in zip(session.messages[-size:], restored[:size])
|
||||
):
|
||||
overlap = size
|
||||
break
|
||||
session.messages.extend(restored[overlap:])
|
||||
|
||||
assistant_data = cast(dict[str, Any], assistant) if isinstance(assistant, dict) else None
|
||||
synchronized = (
|
||||
data.get(PROVIDER_STATE_CHECKPOINT_VERSION_KEY)
|
||||
== PROVIDER_STATE_CHECKPOINT_VERSION
|
||||
)
|
||||
phase = data.get("phase")
|
||||
exact_final = (
|
||||
phase == "final_response"
|
||||
and assistant_data is not None
|
||||
and assistant_data.get("role") == "assistant"
|
||||
and not data.get("completed_tool_results")
|
||||
and not data.get("pending_tool_calls")
|
||||
)
|
||||
exact_tools = (
|
||||
phase == "tools_completed"
|
||||
and assistant_data is not None
|
||||
and assistant_data.get("role") == "assistant"
|
||||
and not data.get("pending_tool_calls")
|
||||
)
|
||||
if not (synchronized and (exact_final or exact_tools)):
|
||||
session.provider_state = None
|
||||
|
||||
session.metadata.pop(PENDING_USER_TURN_KEY, None)
|
||||
session.metadata.pop(RUNTIME_CHECKPOINT_KEY, None)
|
||||
session.updated_at = datetime.now()
|
||||
return True
|
||||
|
||||
|
||||
def _discard_runtime_checkpoint(session: Session) -> bool:
|
||||
"""Drop checkpoint state that cannot be projected into valid history."""
|
||||
if RUNTIME_CHECKPOINT_KEY not in session.metadata:
|
||||
return False
|
||||
session.metadata.pop(RUNTIME_CHECKPOINT_KEY, None)
|
||||
session.provider_state = None
|
||||
session.updated_at = datetime.now()
|
||||
return True
|
||||
|
||||
|
||||
def restore_pending_interruption(session: Session, *, superseded: bool = False) -> bool:
|
||||
"""Close a persisted user-only turn without pretending it was answered."""
|
||||
if not session.metadata.get(PENDING_USER_TURN_KEY):
|
||||
return False
|
||||
if session.messages and session.messages[-1].get("role") == "user":
|
||||
content = (
|
||||
"Task recovery was superseded by a newer message."
|
||||
if superseded
|
||||
else "Error: Task interrupted before a response was generated."
|
||||
)
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"_recovery_interrupted": True,
|
||||
}
|
||||
)
|
||||
session.provider_state = None
|
||||
session.updated_at = datetime.now()
|
||||
session.metadata.pop(PENDING_USER_TURN_KEY, None)
|
||||
return True
|
||||
|
||||
|
||||
def append_recovery_interruption(session: Session, *, superseded: bool = False) -> None:
|
||||
"""Close a restored partial turn whose last durable row is not the user message."""
|
||||
if session.messages and session.messages[-1].get("_recovery_interrupted") is True:
|
||||
return
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": (
|
||||
"Task recovery was superseded by a newer message."
|
||||
if superseded
|
||||
else "Error: Task recovery was interrupted before completion."
|
||||
),
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"_recovery_interrupted": True,
|
||||
}
|
||||
)
|
||||
session.provider_state = None
|
||||
session.updated_at = datetime.now()
|
||||
|
||||
|
||||
def recovery_state_from_metadata(metadata: Mapping[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""Return a sanitized recovery state suitable for the WebSocket wire."""
|
||||
value = metadata.get(RECOVERY_METADATA_KEY) if metadata else None
|
||||
if not isinstance(value, dict):
|
||||
return None
|
||||
state = cast(dict[str, Any], value)
|
||||
status = state.get("status")
|
||||
recovery_id = state.get("recovery_id")
|
||||
if status not in _RECOVERY_STATUSES or not isinstance(recovery_id, str):
|
||||
return None
|
||||
payload: dict[str, Any] = {"status": status, "recovery_id": recovery_id}
|
||||
reason = state.get("reason")
|
||||
if isinstance(reason, str) and reason:
|
||||
payload["reason"] = reason
|
||||
attempts = state.get("attempts")
|
||||
if isinstance(attempts, int) and attempts >= 0:
|
||||
payload["attempts"] = attempts
|
||||
can_continue = state.get("can_continue")
|
||||
if isinstance(can_continue, bool):
|
||||
payload["can_continue"] = can_continue
|
||||
return payload
|
||||
|
||||
|
||||
@dataclasses.dataclass(slots=True)
|
||||
class RecoveryCoordinator:
|
||||
"""Classify, announce, and gate durable WebUI turn recovery."""
|
||||
|
||||
sessions: SessionManager
|
||||
bus: MessageBus
|
||||
unified_session: bool = False
|
||||
_active_recovery_tasks: dict[str, asyncio.Task[Any]] = dataclasses.field(
|
||||
default_factory=dict,
|
||||
init=False,
|
||||
repr=False,
|
||||
)
|
||||
|
||||
def register_recovery_task(self, session_key: str, task: asyncio.Task[Any]) -> None:
|
||||
"""Track the task that owns an explicit recovery continuation."""
|
||||
self._active_recovery_tasks[session_key] = task
|
||||
|
||||
def unregister_recovery_task(self, session_key: str, task: asyncio.Task[Any]) -> None:
|
||||
"""Drop a recovery task without removing a newer task for the same session."""
|
||||
if self._active_recovery_tasks.get(session_key) is task:
|
||||
self._active_recovery_tasks.pop(session_key, None)
|
||||
|
||||
async def _cancel_active_recovery(self, session_key: str) -> None:
|
||||
"""Stop an explicit continuation before accepting newer user input."""
|
||||
task = self._active_recovery_tasks.get(session_key)
|
||||
if task is None or task is asyncio.current_task() or task.done():
|
||||
return
|
||||
task.cancel()
|
||||
# AgentLoop's cancellation path materializes any partial checkpoint and
|
||||
# releases its pending queue. Wait for that ownership to be released
|
||||
# before the newer message is routed.
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
async def scan(self) -> None:
|
||||
"""Recover every interrupted WebUI session once at gateway startup."""
|
||||
for key in self._recovery_candidates():
|
||||
metadata_payload = self.sessions.read_session_metadata(key)
|
||||
raw_metadata = metadata_payload.get("metadata") if metadata_payload else None
|
||||
metadata = cast(dict[str, Any], raw_metadata) if isinstance(raw_metadata, dict) else {}
|
||||
route = self._websocket_route_for(key, metadata)
|
||||
if route is None:
|
||||
continue
|
||||
unfinished = self._has_unfinished_webui_transcript(key)
|
||||
if not self._needs_recovery(metadata) and not unfinished:
|
||||
continue
|
||||
session = self.sessions.get_or_create(key)
|
||||
try:
|
||||
await self._recover_session(session, route[1])
|
||||
await self._requeue_pending_followups(session)
|
||||
except Exception:
|
||||
logger.exception("failed to recover interrupted WebUI session {}", session.key)
|
||||
state = recovery_state_from_metadata(session.metadata)
|
||||
failed = self._set_state(
|
||||
session,
|
||||
status="failed",
|
||||
recovery_id=cast(str, state["recovery_id"]) if state else uuid4().hex,
|
||||
attempts=cast(int, state.get("attempts", 0)) if state else 0,
|
||||
reason="recovery_failed",
|
||||
can_continue=False,
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(route[1], failed)
|
||||
|
||||
def _recovery_candidates(self) -> list[str]:
|
||||
"""Discover canonical and transcript-only WebUI sessions cheaply."""
|
||||
candidates = dict.fromkeys(
|
||||
key
|
||||
for item in self.sessions.list_sessions()
|
||||
if isinstance((key := item.get("key")), str)
|
||||
)
|
||||
try:
|
||||
# Imported lazily because the sidebar index also projects recovery
|
||||
# metadata. The index is the owner of transcript-only discovery;
|
||||
# duplicating its filename and migration rules here would drift.
|
||||
from nanobot.webui.session_list_index import list_webui_sessions
|
||||
|
||||
for item in list_webui_sessions(self.sessions):
|
||||
key = item.get("key")
|
||||
if isinstance(key, str):
|
||||
candidates.setdefault(key, None)
|
||||
except Exception:
|
||||
# Canonical checkpoint recovery remains available even if the
|
||||
# optional display-history index is corrupt or unavailable.
|
||||
logger.exception("failed to discover transcript-only WebUI sessions")
|
||||
return list(candidates)
|
||||
|
||||
@staticmethod
|
||||
def _needs_recovery(metadata: Mapping[str, Any]) -> bool:
|
||||
if metadata.get(PENDING_USER_TURN_KEY) is True:
|
||||
return True
|
||||
if isinstance(metadata.get(RUNTIME_CHECKPOINT_KEY), dict):
|
||||
return True
|
||||
followups = metadata.get(PENDING_FOLLOWUPS_KEY)
|
||||
if isinstance(followups, list) and len(cast(list[object], followups)) > 0:
|
||||
return True
|
||||
state = recovery_state_from_metadata(metadata)
|
||||
return bool(state and state["status"] in {"resuming", "awaiting_user", "failed"})
|
||||
|
||||
async def admit(self, message: InboundMessage) -> bool:
|
||||
"""Reject stale queued recoveries and let new user input supersede them."""
|
||||
recovery_id = message.metadata.get(RECOVERY_INBOUND_METADATA_KEY)
|
||||
if isinstance(recovery_id, str):
|
||||
session = self.sessions.get_or_create(message.session_key)
|
||||
state = recovery_state_from_metadata(session.metadata)
|
||||
return bool(
|
||||
state
|
||||
and state["status"] == "resuming"
|
||||
and state["recovery_id"] == recovery_id
|
||||
)
|
||||
if message.channel != "websocket":
|
||||
return True
|
||||
session = self.sessions.get_or_create(message.session_key)
|
||||
state = recovery_state_from_metadata(session.metadata)
|
||||
if state and state["status"] in {"resuming", "awaiting_user", "failed"}:
|
||||
await self._cancel_active_recovery(message.session_key)
|
||||
restore_runtime_checkpoint(session)
|
||||
if not restore_pending_interruption(session, superseded=True):
|
||||
append_recovery_interruption(session, superseded=True)
|
||||
recovered = self._set_state(
|
||||
session,
|
||||
status="recovered",
|
||||
recovery_id=cast(str, state["recovery_id"]),
|
||||
attempts=cast(int, state.get("attempts", 0)),
|
||||
reason="superseded",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(message.chat_id, recovered)
|
||||
return True
|
||||
|
||||
async def turn_completed(self, session_key: str) -> None:
|
||||
"""Resolve a resuming state after the recovered turn commits."""
|
||||
session = self.sessions.get_or_create(session_key)
|
||||
state = recovery_state_from_metadata(session.metadata)
|
||||
if not state or state["status"] != "resuming":
|
||||
return
|
||||
route = self._websocket_route(session)
|
||||
if route is None:
|
||||
return
|
||||
recovered = self._set_state(
|
||||
session,
|
||||
status="recovered",
|
||||
recovery_id=cast(str, state["recovery_id"]),
|
||||
attempts=cast(int, state.get("attempts", 0)),
|
||||
reason="continued",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(route[1], recovered)
|
||||
|
||||
async def handle_action(self, action: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Apply an authenticated continue/dismiss operation."""
|
||||
chat_id = payload.get("chat_id")
|
||||
recovery_id = payload.get("recovery_id")
|
||||
if not isinstance(chat_id, str) or not chat_id:
|
||||
raise RecoveryActionError("missing chat_id")
|
||||
if not isinstance(recovery_id, str) or not recovery_id:
|
||||
raise RecoveryActionError("missing recovery_id")
|
||||
session = self.sessions.get_or_create(self._session_key(chat_id))
|
||||
state = recovery_state_from_metadata(session.metadata)
|
||||
if not state or state["recovery_id"] != recovery_id:
|
||||
raise RecoveryActionError("recovery state is stale", status=409)
|
||||
|
||||
if action == "dismiss":
|
||||
restore_runtime_checkpoint(session)
|
||||
restore_pending_interruption(session)
|
||||
next_state = self._set_state(
|
||||
session,
|
||||
status="recovered",
|
||||
recovery_id=recovery_id,
|
||||
attempts=cast(int, state.get("attempts", 0)),
|
||||
reason="dismissed",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, next_state)
|
||||
return next_state
|
||||
if action != "continue":
|
||||
raise RecoveryActionError("unknown recovery action")
|
||||
if state["status"] not in {"awaiting_user", "failed"}:
|
||||
raise RecoveryActionError("recovery is not waiting for confirmation", status=409)
|
||||
if state.get("can_continue") is False:
|
||||
raise RecoveryActionError("recovery context is unavailable", status=409)
|
||||
next_state = self._set_state(
|
||||
session,
|
||||
status="resuming",
|
||||
recovery_id=recovery_id,
|
||||
attempts=cast(int, state.get("attempts", 0)) + 1,
|
||||
reason="user_confirmed",
|
||||
resume_message_count=len(session.messages),
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, next_state)
|
||||
await self._queue_continuation(session, chat_id, next_state)
|
||||
return next_state
|
||||
|
||||
async def _recover_session(self, session: Session, chat_id: str) -> None:
|
||||
checkpoint_value = cast(object, session.metadata.get(RUNTIME_CHECKPOINT_KEY))
|
||||
checkpoint = (
|
||||
cast(dict[str, Any], checkpoint_value)
|
||||
if isinstance(checkpoint_value, dict)
|
||||
else None
|
||||
)
|
||||
pending = session.metadata.get(PENDING_USER_TURN_KEY) is True
|
||||
state = recovery_state_from_metadata(session.metadata)
|
||||
if not pending and checkpoint is None:
|
||||
if state and state["status"] == "resuming":
|
||||
resume_count = self._resume_message_count(session)
|
||||
if resume_count is not None and len(session.messages) > resume_count:
|
||||
next_state = self._set_state(
|
||||
session,
|
||||
status="recovered",
|
||||
recovery_id=cast(str, state["recovery_id"]),
|
||||
attempts=cast(int, state.get("attempts", 0)),
|
||||
reason="committed",
|
||||
)
|
||||
else:
|
||||
next_state = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=cast(str, state["recovery_id"]),
|
||||
attempts=cast(int, state.get("attempts", 1)),
|
||||
reason="loop_guard",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, next_state)
|
||||
elif self._has_unfinished_webui_transcript(session.key):
|
||||
# A normal last-client shutdown can materialize the checkpoint
|
||||
# before the process exits. In that path there is no pending
|
||||
# marker left to classify, but the append-only transcript still
|
||||
# contains an activity row without a turn_end. Treat it as an
|
||||
# interrupted turn instead of letting the UI resurrect it as a
|
||||
# forever-running spinner.
|
||||
waiting = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=uuid4().hex,
|
||||
attempts=0,
|
||||
reason="interrupted_without_checkpoint",
|
||||
can_continue=False,
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, waiting)
|
||||
return
|
||||
if state and state["status"] in {"awaiting_user", "failed"}:
|
||||
await self._publish(chat_id, state)
|
||||
return
|
||||
if state and state["status"] == "resuming":
|
||||
restore_runtime_checkpoint(session)
|
||||
restore_pending_interruption(session)
|
||||
waiting = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=cast(str, state["recovery_id"]),
|
||||
attempts=cast(int, state.get("attempts", 1)),
|
||||
reason="loop_guard",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, waiting)
|
||||
return
|
||||
|
||||
recovery_id = uuid4().hex
|
||||
phase = checkpoint.get("phase") if checkpoint is not None else None
|
||||
pending_calls = checkpoint.get("pending_tool_calls") if checkpoint is not None else None
|
||||
if checkpoint is not None and phase not in _KNOWN_CHECKPOINT_PHASES:
|
||||
_discard_runtime_checkpoint(session)
|
||||
restore_pending_interruption(session)
|
||||
waiting = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=recovery_id,
|
||||
attempts=0,
|
||||
reason="checkpoint_unknown",
|
||||
can_continue=False,
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, waiting)
|
||||
return
|
||||
if checkpoint is not None and not _runtime_checkpoint_is_well_formed(checkpoint):
|
||||
_discard_runtime_checkpoint(session)
|
||||
restore_pending_interruption(session)
|
||||
waiting = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=recovery_id,
|
||||
attempts=0,
|
||||
reason="checkpoint_invalid",
|
||||
can_continue=False,
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, waiting)
|
||||
return
|
||||
if phase == "final_response":
|
||||
restore_runtime_checkpoint(session)
|
||||
recovered = self._set_state(
|
||||
session,
|
||||
status="recovered",
|
||||
recovery_id=recovery_id,
|
||||
attempts=0,
|
||||
reason="answer_restored",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, recovered)
|
||||
return
|
||||
if phase in _UNCERTAIN_TOOL_PHASES or pending_calls:
|
||||
restore_runtime_checkpoint(session)
|
||||
waiting = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=recovery_id,
|
||||
attempts=0,
|
||||
reason="tool_state_unknown",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, waiting)
|
||||
return
|
||||
# A gateway restart is a lifecycle boundary. Never enqueue model work
|
||||
# implicitly: even a synchronized checkpoint may sit next to an
|
||||
# external side effect that the user should review first. The final
|
||||
# answer path above only restores persisted output; it never executes.
|
||||
restore_runtime_checkpoint(session)
|
||||
waiting = self._set_state(
|
||||
session,
|
||||
status="awaiting_user",
|
||||
recovery_id=recovery_id,
|
||||
attempts=0,
|
||||
reason="restart_requires_confirmation",
|
||||
)
|
||||
self.sessions.save(session)
|
||||
await self._publish(chat_id, waiting)
|
||||
|
||||
async def _queue_continuation(
|
||||
self,
|
||||
session: Session,
|
||||
chat_id: str,
|
||||
state: Mapping[str, Any],
|
||||
) -> None:
|
||||
recovery_id = cast(str, state["recovery_id"])
|
||||
await self.bus.publish_inbound(
|
||||
InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id="system:recovery",
|
||||
chat_id=chat_id,
|
||||
content=(
|
||||
"Continue the interrupted request from the saved conversation context. "
|
||||
"Do not repeat completed work or mention the restart unless it affects the answer."
|
||||
),
|
||||
metadata={
|
||||
"webui": True,
|
||||
"_wants_stream": True,
|
||||
WEBUI_TURN_METADATA_KEY: f"recovery:{recovery_id}",
|
||||
RECOVERY_INBOUND_METADATA_KEY: recovery_id,
|
||||
turn_continuation.INTERNAL_CONTINUATION_META: True,
|
||||
turn_continuation.SKIP_USER_PERSIST_META: True,
|
||||
},
|
||||
session_key_override=session.key,
|
||||
require_existing_session=True,
|
||||
)
|
||||
)
|
||||
|
||||
async def _requeue_pending_followups(self, session: Session) -> None:
|
||||
"""Return durable live-turn follow-ups to the bus after a restart."""
|
||||
for message in pending_followups(session):
|
||||
await self.bus.publish_inbound(message)
|
||||
|
||||
@staticmethod
|
||||
def _resume_message_count(session: Session) -> int | None:
|
||||
raw_value = cast(object, session.metadata.get(RECOVERY_METADATA_KEY))
|
||||
value = cast(dict[str, Any], raw_value) if isinstance(raw_value, dict) else None
|
||||
if value is None:
|
||||
return None
|
||||
count = value.get("resume_message_count")
|
||||
return count if isinstance(count, int) and count >= 0 else None
|
||||
|
||||
async def _publish(
|
||||
self,
|
||||
chat_id: str,
|
||||
state: Mapping[str, Any],
|
||||
) -> None:
|
||||
"""Publish the recovery state and invalidate its sidebar projection."""
|
||||
await self.bus.publish_outbound(
|
||||
outbound_message_for_event(
|
||||
channel="websocket",
|
||||
chat_id=chat_id,
|
||||
event=RecoveryStateEvent(
|
||||
status=cast(str, state["status"]),
|
||||
recovery_id=cast(str, state["recovery_id"]),
|
||||
reason=cast(str | None, state.get("reason")),
|
||||
attempts=cast(int, state.get("attempts", 0)),
|
||||
can_continue=cast(bool | None, state.get("can_continue")),
|
||||
),
|
||||
)
|
||||
)
|
||||
await self.bus.publish_outbound(
|
||||
outbound_message_for_event(
|
||||
channel="websocket",
|
||||
chat_id=chat_id,
|
||||
event=SessionUpdatedEvent(scope="thread"),
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _set_state(
|
||||
session: Session,
|
||||
*,
|
||||
status: str,
|
||||
recovery_id: str,
|
||||
attempts: int,
|
||||
reason: str,
|
||||
resume_message_count: int | None = None,
|
||||
can_continue: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
state = {
|
||||
"status": status,
|
||||
"recovery_id": recovery_id,
|
||||
"attempts": max(0, attempts),
|
||||
"reason": reason,
|
||||
"updated_at": datetime.now().isoformat(),
|
||||
}
|
||||
if not can_continue:
|
||||
state["can_continue"] = False
|
||||
if resume_message_count is not None:
|
||||
state["resume_message_count"] = max(0, resume_message_count)
|
||||
session.metadata[RECOVERY_METADATA_KEY] = state
|
||||
session.updated_at = datetime.now()
|
||||
return state
|
||||
|
||||
def _session_key(self, chat_id: str) -> str:
|
||||
return UNIFIED_SESSION_KEY if self.unified_session else f"websocket:{chat_id}"
|
||||
|
||||
@staticmethod
|
||||
def _has_unfinished_webui_transcript(session_key: str) -> bool:
|
||||
"""Detect a stale WebUI activity tail after an unclean gateway stop.
|
||||
|
||||
The transcript is intentionally consulted only as a last-resort signal:
|
||||
a durable pending turn or runtime checkpoint always takes precedence.
|
||||
This keeps browser disconnects harmless while preventing a materialized
|
||||
partial turn from being presented as active forever after a restart.
|
||||
"""
|
||||
try:
|
||||
from nanobot.webui.transcript import has_unfinished_transcript_tail
|
||||
|
||||
return has_unfinished_transcript_tail(session_key)
|
||||
except (OSError, ValueError, TypeError):
|
||||
# Recovery must fail closed if the optional display transcript is
|
||||
# corrupt or unavailable; the normal checkpoint path still applies.
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _websocket_route(session: Session) -> tuple[str, str] | None:
|
||||
return RecoveryCoordinator._websocket_route_for(session.key, session.metadata)
|
||||
|
||||
@staticmethod
|
||||
def _websocket_route_for(
|
||||
session_key: str,
|
||||
metadata: Mapping[str, Any],
|
||||
) -> tuple[str, str] | None:
|
||||
if session_key.startswith("websocket:"):
|
||||
chat_id = session_key.split(":", 1)[1]
|
||||
return ("websocket", chat_id) if chat_id else None
|
||||
if session_key == UNIFIED_SESSION_KEY:
|
||||
route = last_channel_from_metadata(metadata)
|
||||
if route and route[0] == "websocket":
|
||||
return route
|
||||
return None
|
||||
@@ -43,6 +43,7 @@ from nanobot.runtime_context import public_history_message
|
||||
from nanobot.session.goal_state import goal_state_ws_blob
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
from nanobot.session.recovery import RecoveryCoordinator
|
||||
from nanobot.session.session_handles import session_handle_for_name
|
||||
from nanobot.session.session_messages import (
|
||||
SessionMessageEnvelope,
|
||||
@@ -511,6 +512,7 @@ class WebuiTurnCoordinator:
|
||||
bus: MessageBus
|
||||
sessions: SessionManager
|
||||
schedule_background: Callable[[Awaitable[None]], None]
|
||||
recovery: RecoveryCoordinator | None = None
|
||||
|
||||
def subscribe(self, runtime_events: RuntimeEventBus) -> Callable[[], None]:
|
||||
"""Subscribe this coordinator to runtime events."""
|
||||
@@ -654,6 +656,8 @@ class WebuiTurnCoordinator:
|
||||
event.runtime.context_window_tokens if event.runtime is not None else None
|
||||
),
|
||||
)
|
||||
if self.recovery is not None:
|
||||
await self.recovery.turn_completed(event.context.session_key)
|
||||
self._schedule_title_update_from_event(event)
|
||||
|
||||
async def _handle_goal_state_changed(self, event: GoalStateChanged) -> None:
|
||||
|
||||
@@ -69,6 +69,7 @@ def build_gateway_services(
|
||||
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
|
||||
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
|
||||
skill_state_action: Callable[[set[str]], None] | None = None,
|
||||
recovery_action: Callable[[str, dict[str, Any]], Awaitable[dict[str, Any]]] | None = None,
|
||||
logger: Any = default_logger,
|
||||
) -> GatewayServices:
|
||||
settings = WebUISettingsServices.create(
|
||||
@@ -131,6 +132,7 @@ def build_gateway_services(
|
||||
mcp_runtime_status=mcp_runtime_status,
|
||||
mcp_reload=mcp_reload,
|
||||
skill_state_action=skill_state_action,
|
||||
recovery_action=recovery_action,
|
||||
log=logger,
|
||||
)
|
||||
return GatewayServices(
|
||||
|
||||
@@ -31,8 +31,9 @@ from nanobot.session.manager import (
|
||||
_metadata_title, # pyright: ignore[reportPrivateUsage]
|
||||
)
|
||||
from nanobot.session.model_selection import model_preset_from_metadata
|
||||
from nanobot.session.recovery import recovery_state_from_metadata
|
||||
|
||||
_INDEX_VERSION = 7
|
||||
_INDEX_VERSION = 8
|
||||
_INDEX_FILENAME = ".webui_session_index.json"
|
||||
_MODEL_PRESET_FIELD = "model_preset"
|
||||
_ROW_SOURCE_FIELD = "_source"
|
||||
@@ -245,6 +246,7 @@ def _public_row(sessions_dir: Path, webui_dir: Path, row: dict[str, Any]) -> dic
|
||||
"title": row.get("title", ""),
|
||||
"preview": row.get("preview", ""),
|
||||
_MODEL_PRESET_FIELD: row.get(_MODEL_PRESET_FIELD),
|
||||
"recovery_state": row.get("recovery_state"),
|
||||
_WORKSPACE_SCOPE_PRESENT_FIELD: row.get(_WORKSPACE_SCOPE_PRESENT_FIELD, False),
|
||||
_WORKSPACE_SCOPE_VALUE_FIELD: row.get(_WORKSPACE_SCOPE_VALUE_FIELD),
|
||||
"path": str(path),
|
||||
@@ -485,6 +487,7 @@ def _indexed_row_for_session(session: Session, path: Path, webui_dir: Path) -> d
|
||||
"title": _metadata_title(session.metadata),
|
||||
"preview": _preview_from_messages(session.messages),
|
||||
_MODEL_PRESET_FIELD: model_preset_from_metadata(session.metadata),
|
||||
"recovery_state": recovery_state_from_metadata(session.metadata),
|
||||
**_indexed_workspace_scope_fields(session.metadata),
|
||||
_ROW_SOURCE_FIELD: _SESSION_SOURCE,
|
||||
"file": path.name,
|
||||
@@ -601,6 +604,7 @@ def _scan_transcript_row(
|
||||
"title": "",
|
||||
"preview": preview or fallback_preview,
|
||||
_MODEL_PRESET_FIELD: None,
|
||||
"recovery_state": None,
|
||||
**_indexed_workspace_scope_fields({}),
|
||||
_ROW_SOURCE_FIELD: _TRANSCRIPT_SOURCE,
|
||||
"file": stem,
|
||||
@@ -687,6 +691,7 @@ def _scan_session_row(
|
||||
"title": _metadata_title(metadata),
|
||||
"preview": preview or fallback_preview,
|
||||
_MODEL_PRESET_FIELD: model_preset_from_metadata(metadata),
|
||||
"recovery_state": recovery_state_from_metadata(metadata),
|
||||
**_indexed_workspace_scope_fields(metadata),
|
||||
_ROW_SOURCE_FIELD: _SESSION_SOURCE,
|
||||
"file": path.name,
|
||||
|
||||
@@ -2541,6 +2541,19 @@ def has_pending_tool_calls(
|
||||
return False
|
||||
|
||||
|
||||
def has_unfinished_transcript_tail(session_key: str) -> bool:
|
||||
"""Return whether the active transcript ends in an unfinished turn.
|
||||
|
||||
Recovery runs at gateway startup and only needs the newest, still-active
|
||||
turn. Completed turns are rotated into immutable segment files, so reading
|
||||
every historical segment here would make restart cost grow with the full
|
||||
conversation history.
|
||||
"""
|
||||
return has_pending_tool_calls(
|
||||
_read_transcript_file(webui_transcript_path(session_key))
|
||||
)
|
||||
|
||||
|
||||
def completed_turn_ids(lines: list[dict[str, Any]]) -> list[str]:
|
||||
"""Return stable identities for turns with an explicitly persisted completion."""
|
||||
completed: list[str] = []
|
||||
|
||||
@@ -29,6 +29,7 @@ from nanobot.cron.session_turns import is_bound_cron_job
|
||||
from nanobot.cron.types import CronJob, CronSchedule
|
||||
from nanobot.security.workspace_access import WorkspaceScope
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.session.recovery import RecoveryActionError
|
||||
from nanobot.session.session_handles import (
|
||||
SessionHandleResolver,
|
||||
)
|
||||
@@ -145,6 +146,8 @@ _WEBUI_MUTATION_PATHS = {
|
||||
"skill.delete": "/api/webui/skills/delete",
|
||||
"sidebar.update": "/api/webui/sidebar-state/update",
|
||||
"workspace.pick_folder": "/api/workspaces/pick-folder",
|
||||
"recovery.continue": "/api/webui/recovery/continue",
|
||||
"recovery.dismiss": "/api/webui/recovery/dismiss",
|
||||
"settings.agent.update": "/api/settings/update",
|
||||
"settings.model_configuration.create": "/api/settings/model-configurations/create",
|
||||
"settings.model_configuration.update": "/api/settings/model-configurations/update",
|
||||
@@ -323,6 +326,9 @@ class GatewayHTTPHandler:
|
||||
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
|
||||
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
|
||||
skill_state_action: Callable[[set[str]], None] | None = None,
|
||||
recovery_action: (
|
||||
Callable[[str, dict[str, Any]], Awaitable[dict[str, Any]]] | None
|
||||
) = None,
|
||||
log: Any = logger,
|
||||
) -> None:
|
||||
self.config = config
|
||||
@@ -340,6 +346,7 @@ class GatewayHTTPHandler:
|
||||
disabled_skills if disabled_skills is not None else set()
|
||||
)
|
||||
self.skill_state_action = skill_state_action
|
||||
self.recovery_action = recovery_action
|
||||
self._skill_install_lock = asyncio.Lock()
|
||||
self._folder_picker_lock = asyncio.Lock()
|
||||
self.cron_service = cron_service
|
||||
@@ -454,6 +461,8 @@ class GatewayHTTPHandler:
|
||||
return True
|
||||
if re.match(r"^/api/webui/automations/(enable|disable|delete|run|update)$", path):
|
||||
return True
|
||||
if path in {"/api/webui/recovery/continue", "/api/webui/recovery/dismiss"}:
|
||||
return True
|
||||
return path in {
|
||||
"/api/webui/skills/install",
|
||||
"/api/webui/skills/update",
|
||||
@@ -507,6 +516,11 @@ class GatewayHTTPHandler:
|
||||
if response is not None:
|
||||
return response
|
||||
|
||||
# Recovery routes
|
||||
response = await self._dispatch_recovery_route(request, got)
|
||||
if response is not None:
|
||||
return response
|
||||
|
||||
# Session routes
|
||||
response = await self._dispatch_session_routes(request, got)
|
||||
if response is not None:
|
||||
@@ -700,6 +714,27 @@ class GatewayHTTPHandler:
|
||||
|
||||
return None
|
||||
|
||||
async def _dispatch_recovery_route(
|
||||
self,
|
||||
request: WsRequest,
|
||||
path: str,
|
||||
) -> Response | None:
|
||||
match = re.fullmatch(r"/api/webui/recovery/(continue|dismiss)", path)
|
||||
if match is None:
|
||||
return None
|
||||
if not getattr(request, _WEBUI_MUTATION_REQUEST_ATTR, False):
|
||||
return _http_error(405, "WebUI recovery actions require an authenticated WebSocket")
|
||||
if self.recovery_action is None:
|
||||
return _http_error(503, "WebUI recovery is unavailable")
|
||||
payload = _mutation_payload(request)
|
||||
if payload is None:
|
||||
return _http_error(400, "invalid recovery payload")
|
||||
try:
|
||||
result = await self.recovery_action(match.group(1), payload)
|
||||
except RecoveryActionError as exc:
|
||||
return _http_error(exc.status, str(exc))
|
||||
return _http_json_response(result)
|
||||
|
||||
async def _handle_session_context_get(self, request: WsRequest, key: str) -> Response:
|
||||
if not self.check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
@@ -746,6 +781,10 @@ class GatewayHTTPHandler:
|
||||
for k, v in s.items()
|
||||
if k != "path" and k not in WEBUI_SESSION_INDEX_INTERNAL_FIELDS
|
||||
}
|
||||
# Keep the additive recovery field absent for ordinary sessions so
|
||||
# older clients and compact list responses stay unchanged.
|
||||
if row.get("recovery_state") is None:
|
||||
row.pop("recovery_state", None)
|
||||
chat_id = key.split(":", 1)[1]
|
||||
started_at = websocket_turn_wall_started_at(chat_id)
|
||||
if started_at is not None:
|
||||
|
||||
Reference in New Issue
Block a user