mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 14:14:57 +08:00
core/agent+cli: defer DB write until turn success; auto/manual retry
Failed turns roll back memory and keep pending_retry_text; persist only after success. CLI auto-retries 10s/20s/30s (cancellable) and /retry.
This commit is contained in:
+43
-16
@@ -31,6 +31,7 @@ class ChatAgent:
|
||||
temperature: float | None
|
||||
on_limit: LimitContinueHook | None
|
||||
messages: list[AnyChatMessage]
|
||||
pending_retry_text: str | None
|
||||
|
||||
def __init__( # noqa: PLR0913
|
||||
self,
|
||||
@@ -54,6 +55,7 @@ class ChatAgent:
|
||||
self.temperature = temperature
|
||||
self.on_limit = on_limit
|
||||
self.messages = list(messages) if messages is not None else []
|
||||
self.pending_retry_text = None
|
||||
|
||||
async def load_history(self) -> None:
|
||||
"""Replace in-memory messages from the bound memory session."""
|
||||
@@ -61,6 +63,7 @@ class ChatAgent:
|
||||
msg = "load_history requires memory and session_id"
|
||||
raise RuntimeError(msg)
|
||||
self.messages = await self.memory.list_messages(self.session_id)
|
||||
self.pending_retry_text = None
|
||||
|
||||
async def bind_session(self, session_id: int, *, load: bool = True) -> None:
|
||||
"""Attach a memory session id; optionally load existing messages."""
|
||||
@@ -75,24 +78,48 @@ class ChatAgent:
|
||||
if self.memory is not None and self.session_id is not None:
|
||||
_ = await self.memory.append_message(self.session_id, message)
|
||||
|
||||
async def _run_from_user_message(self, user_msg: UserChatMessage) -> AsyncIterator[AgentEvent]:
|
||||
"""Run the tool loop for an already-appended user message; persist only on success."""
|
||||
pre_len = len(self.messages) - 1
|
||||
if pre_len < 0 or self.messages[pre_len] is not user_msg:
|
||||
pre_len = len(self.messages)
|
||||
self.messages.append(user_msg)
|
||||
|
||||
try:
|
||||
async for event in run_chat_loop(
|
||||
self.client,
|
||||
self.messages,
|
||||
model=self.model,
|
||||
tools=self.tools,
|
||||
max_rounds=self.max_rounds,
|
||||
temperature=self.temperature,
|
||||
on_limit=self.on_limit,
|
||||
):
|
||||
yield event
|
||||
except Exception:
|
||||
# Roll back the whole turn so a failed user message is not left half-applied.
|
||||
del self.messages[pre_len:]
|
||||
self.pending_retry_text = user_msg.content
|
||||
raise
|
||||
|
||||
for message in self.messages[pre_len:]:
|
||||
await self._persist(message)
|
||||
self.pending_retry_text = None
|
||||
|
||||
async def run(self, user_text: str) -> AsyncIterator[AgentEvent]:
|
||||
"""Append a user message, run the tool loop, yield events, persist new messages."""
|
||||
"""Append a user message, run the tool loop, yield events, persist only on success."""
|
||||
user_msg = UserChatMessage(content=user_text)
|
||||
self.messages.append(user_msg)
|
||||
await self._persist(user_msg)
|
||||
|
||||
start_len = len(self.messages)
|
||||
async for event in run_chat_loop(
|
||||
self.client,
|
||||
self.messages,
|
||||
model=self.model,
|
||||
tools=self.tools,
|
||||
max_rounds=self.max_rounds,
|
||||
temperature=self.temperature,
|
||||
on_limit=self.on_limit,
|
||||
):
|
||||
async for event in self._run_from_user_message(user_msg):
|
||||
yield event
|
||||
|
||||
# Persist messages appended by the loop after the user message.
|
||||
for message in self.messages[start_len:]:
|
||||
await self._persist(message)
|
||||
async def retry(self) -> AsyncIterator[AgentEvent]:
|
||||
"""Re-run the last failed user turn (no duplicate user message in history/DB)."""
|
||||
text = self.pending_retry_text
|
||||
if text is None:
|
||||
msg = "nothing to retry"
|
||||
raise RuntimeError(msg)
|
||||
user_msg = UserChatMessage(content=text)
|
||||
self.messages.append(user_msg)
|
||||
async for event in self._run_from_user_message(user_msg):
|
||||
yield event
|
||||
|
||||
@@ -28,6 +28,7 @@ SLASH_COMMANDS: tuple[str, ...] = (
|
||||
"/model",
|
||||
"/tools",
|
||||
"/rounds",
|
||||
"/retry",
|
||||
)
|
||||
|
||||
_TOOLS_ARGS: tuple[str, ...] = ("on", "off")
|
||||
|
||||
@@ -7,8 +7,8 @@ from typing import TYPE_CHECKING
|
||||
import click
|
||||
from msgspec import UNSET
|
||||
|
||||
from plyngent.cli.display import render_events
|
||||
from plyngent.cli.readline_setup import setup_readline
|
||||
from plyngent.cli.retry import retry_pending_with_retries, run_user_text_with_retries
|
||||
from plyngent.cli.selection import select_model, select_provider
|
||||
from plyngent.lmproto.openai_compatible.model import (
|
||||
AssistantChatMessage,
|
||||
@@ -35,6 +35,10 @@ Commands:
|
||||
/model [id] Show or switch model
|
||||
/tools [on|off] Show or toggle tools
|
||||
/rounds [n] Show or set max tool-loop rounds
|
||||
/retry Retry the last failed user turn (after errors)
|
||||
|
||||
On network/API errors, the user turn is not saved to the DB. Auto-retry
|
||||
waits 10s, 20s, then 30s (Ctrl+C cancels waits; use /retry later).
|
||||
|
||||
Tab completes slash commands and some arguments (provider, model, tools).
|
||||
Use --session ID or /resume to continue a prior chat after restart.
|
||||
@@ -141,6 +145,7 @@ def _cmd_rounds(state: ReplState, arg: str) -> None:
|
||||
|
||||
def _cmd_clear(state: ReplState) -> None:
|
||||
state.agent.messages.clear()
|
||||
state.agent.pending_retry_text = None
|
||||
click.echo("conversation cleared (in-memory only; DB history kept)")
|
||||
|
||||
|
||||
@@ -197,6 +202,15 @@ def _cmd_history(state: ReplState, arg: str) -> None:
|
||||
click.echo(f"session={state.session_id} messages={len(messages)} showing={len(messages) - start}")
|
||||
for offset, message in enumerate(messages[start:]):
|
||||
click.echo(_format_history_message(start + offset, message))
|
||||
if state.agent.pending_retry_text is not None:
|
||||
click.secho(
|
||||
f"(pending retry) user: {_preview_content(state.agent.pending_retry_text)}",
|
||||
fg="yellow",
|
||||
)
|
||||
|
||||
|
||||
async def _cmd_retry(state: ReplState) -> None:
|
||||
_ = await retry_pending_with_retries(state.agent)
|
||||
|
||||
|
||||
async def _dispatch_slash(state: ReplState, command: str, arg: str) -> bool:
|
||||
@@ -214,6 +228,7 @@ async def _dispatch_slash(state: ReplState, command: str, arg: str) -> bool:
|
||||
"model": lambda: _cmd_model(state, arg),
|
||||
"tools": lambda: _cmd_tools(state, arg),
|
||||
"rounds": lambda: _cmd_rounds(state, arg),
|
||||
"retry": lambda: _cmd_retry(state),
|
||||
}
|
||||
handler = handlers.get(command)
|
||||
if handler is None:
|
||||
@@ -269,8 +284,4 @@ async def run_repl(state: ReplState) -> None:
|
||||
|
||||
click.secho("user: ", fg="green", nl=False)
|
||||
click.echo(line)
|
||||
try:
|
||||
await render_events(state.agent.run(line))
|
||||
except Exception as exc: # noqa: BLE001 — show API/runtime errors in REPL
|
||||
click.secho(f"error: {exc}", fg="red")
|
||||
click.echo() # match render_events spacing before next prompt
|
||||
_ = await run_user_text_with_retries(state.agent, line)
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import click
|
||||
|
||||
from plyngent.cli.display import render_events
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
|
||||
from plyngent.agent import AgentEvent
|
||||
from plyngent.agent.chat import ChatAgent
|
||||
|
||||
# Wait before retry attempt 1, 2, and 3 (after the first failure).
|
||||
DEFAULT_RETRY_DELAYS_SECONDS: tuple[float, ...] = (10.0, 20.0, 30.0)
|
||||
_PREVIEW_LEN = 80
|
||||
|
||||
|
||||
async def sleep_cancellable(seconds: float) -> bool:
|
||||
"""Sleep in short steps so Ctrl+C can cancel. Returns False if interrupted."""
|
||||
deadline = time.monotonic() + seconds
|
||||
try:
|
||||
while True:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
return True
|
||||
await asyncio.sleep(min(0.5, remaining))
|
||||
except asyncio.CancelledError:
|
||||
return False
|
||||
except KeyboardInterrupt:
|
||||
return False
|
||||
|
||||
|
||||
async def run_turn_with_retries(
|
||||
agent: ChatAgent,
|
||||
*,
|
||||
starter: Callable[[], AsyncIterator[AgentEvent]],
|
||||
delays: tuple[float, ...] = DEFAULT_RETRY_DELAYS_SECONDS,
|
||||
) -> bool:
|
||||
"""Run a chat turn with automatic retries on failure.
|
||||
|
||||
``starter`` produces the event stream for the current attempt (``agent.run``
|
||||
or ``agent.retry``). Returns True if the turn completed successfully.
|
||||
"""
|
||||
max_retries = len(delays)
|
||||
attempt = 0
|
||||
while True:
|
||||
try:
|
||||
await render_events(starter())
|
||||
except KeyboardInterrupt:
|
||||
click.echo()
|
||||
click.secho("interrupted", fg="yellow")
|
||||
return False
|
||||
except Exception as exc: # noqa: BLE001 — surface and optionally retry
|
||||
click.secho(f"error: {exc}", fg="red")
|
||||
if attempt >= max_retries:
|
||||
if agent.pending_retry_text is not None:
|
||||
click.secho(
|
||||
"auto-retry exhausted; use /retry to try again, or send a new message",
|
||||
fg="yellow",
|
||||
)
|
||||
click.echo()
|
||||
return False
|
||||
|
||||
wait = delays[attempt]
|
||||
attempt += 1
|
||||
click.secho(
|
||||
f"auto-retry {attempt}/{max_retries} in {wait:g}s "
|
||||
f"(Ctrl+C to cancel; then /retry later)",
|
||||
fg="yellow",
|
||||
)
|
||||
try:
|
||||
ok = await sleep_cancellable(wait)
|
||||
except KeyboardInterrupt:
|
||||
ok = False
|
||||
if not ok:
|
||||
click.secho("auto-retry cancelled; use /retry to try again", fg="yellow")
|
||||
click.echo()
|
||||
return False
|
||||
click.secho(f"retrying ({attempt}/{max_retries})…", fg="yellow")
|
||||
else:
|
||||
return True
|
||||
|
||||
|
||||
async def run_user_text_with_retries(agent: ChatAgent, text: str) -> bool:
|
||||
"""Send a new user message with auto-retry."""
|
||||
return await run_turn_with_retries(agent, starter=lambda: agent.run(text))
|
||||
|
||||
|
||||
async def retry_pending_with_retries(agent: ChatAgent) -> bool:
|
||||
"""Retry the last failed user turn with auto-retry."""
|
||||
if agent.pending_retry_text is None:
|
||||
click.echo("nothing to retry")
|
||||
return False
|
||||
preview = agent.pending_retry_text
|
||||
if len(preview) > _PREVIEW_LEN:
|
||||
preview = preview[:_PREVIEW_LEN] + "…"
|
||||
click.echo(f"retrying: {preview}")
|
||||
return await run_turn_with_retries(agent, starter=agent.retry)
|
||||
@@ -66,6 +66,7 @@ class ReplState:
|
||||
session = await self.memory.create_session(name=name)
|
||||
self.session_id = session.sid
|
||||
self.agent = self._make_agent()
|
||||
self.agent.pending_retry_text = None
|
||||
|
||||
async def resume_session(self, session_id: int) -> None:
|
||||
row = await self.memory.get_session(session_id)
|
||||
|
||||
Reference in New Issue
Block a user