mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/agent: checkpoint tool rounds so /retry does not redo side effects
After each completed assistant+tool batch, persist and keep that prefix on failure. Rollback only unfinished assistant/stream suffix. pending_retry and /retry also apply when history ends with tool results (continue model without re-executing tools).
This commit is contained in:
@@ -55,7 +55,7 @@ Async SQLAlchemy + aiosqlite. `MemoryStore`: schema init (+ lightweight SQLite `
|
|||||||
- **`ChatClient`** Protocol for `chat_completions`.
|
- **`ChatClient`** Protocol for `chat_completions`.
|
||||||
- **`@tool` / `ToolRegistry`**: decorator infers JSON Schema from type hints; execute tools by name.
|
- **`@tool` / `ToolRegistry`**: decorator infers JSON Schema from type hints; execute tools by name.
|
||||||
- **`run_chat_loop`**: multi-round tool loop; default **streaming** text deltas + stream tool-call merge; parallel tools; tool-result char budget; soft context compact on request (**API-calibrated** after first usage when available); cooperative cancel points; optional `on_limit`.
|
- **`run_chat_loop`**: multi-round tool loop; default **streaming** text deltas + stream tool-call merge; parallel tools; tool-result char budget; soft context compact on request (**API-calibrated** after first usage when available); cooperative cancel points; optional `on_limit`.
|
||||||
- **`ChatAgent`**: optional `MemoryStore` (user message persisted immediately; assistant/tools on success); `stream`; system prompt; `retry()` when history ends with a user message (failed/cancelled turn or orphan after resume).
|
- **`ChatAgent`**: optional `MemoryStore` (user message persisted immediately; **completed tool batches checkpointed** mid-turn; unfinished assistant suffix rolled back on failure); `stream`; system prompt; `retry()` continues incomplete turns (user-only **or** after committed tools — does not re-run those tools).
|
||||||
- **`/compact`**: soft-compact tool dumps → model summary (no tools) → **new** session seeded with summary message.
|
- **`/compact`**: soft-compact tool dumps → model summary (no tools) → **new** session seeded with summary message.
|
||||||
- Events: text_delta, **reasoning_delta**, assistant_message, tool_call/result, max_rounds, **error** (`retryable`/`source`), **cancelled** (`reason`), **usage** (`TokenUsage`).
|
- Events: text_delta, **reasoning_delta**, assistant_message, tool_call/result, max_rounds, **error** (`retryable`/`source`), **cancelled** (`reason`), **usage** (`TokenUsage`).
|
||||||
- Usage: API `usage` from completions (stream with `include_usage`); **char≈token fallback** (~4 chars/token) when omitted; **context size** = last request ``prompt_tokens`` (API preferred); `last_turn_usage` / `session_usage` are **billed sums** (tool rounds re-send history); CLI end-of-turn + `/status`.
|
- Usage: API `usage` from completions (stream with `include_usage`); **char≈token fallback** (~4 chars/token) when omitted; **context size** = last request ``prompt_tokens`` (API preferred); `last_turn_usage` / `session_usage` are **billed sums** (tool rounds re-send history); CLI end-of-turn + `/status`.
|
||||||
@@ -86,7 +86,7 @@ Click app + readline REPL. Entry: `plyngent` / `python -m plyngent`.
|
|||||||
- **`plyngent chat`**: provider/model (flags or interactive; Tab via readline in `prompting`); sessions store `provider_name`/`model` and restore on resume; SQLite via `[database]` (file DB under user data if unset/`:memory:`); workspace-bound; resume latest for cwd/`--workspace` by default (`--new` / `--session`). One-shot: `-p/--prompt` and non-TTY stdin; exit codes 0/1/2/3; `--yes`, `--stream/--no-stream`, `--quiet`. Root `--log-level`.
|
- **`plyngent chat`**: provider/model (flags or interactive; Tab via readline in `prompting`); sessions store `provider_name`/`model` and restore on resume; SQLite via `[database]` (file DB under user data if unset/`:memory:`); workspace-bound; resume latest for cwd/`--workspace` by default (`--new` / `--session`). One-shot: `-p/--prompt` and non-TTY stdin; exit codes 0/1/2/3; `--yes`, `--stream/--no-stream`, `--quiet`. Root `--log-level`.
|
||||||
- Slash: Click group in `cli/slash.py` + `awaitlet` for async work; Tab completer from registry + ParamType `shell_complete`. Multiline `"""` … `"""`; `/edit` via `$EDITOR`.
|
- Slash: Click group in `cli/slash.py` + `awaitlet` for async work; Tab completer from registry + ParamType `shell_complete`. Multiline `"""` … `"""`; `/edit` via `$EDITOR`.
|
||||||
- Explicit `/resume` or `--session` from another workspace prompts: **keep** / **update** / **abort**.
|
- Explicit `/resume` or `--session` from another workspace prompts: **keep** / **update** / **abort**.
|
||||||
- Failed/cancelled turns: user message kept; partial assistant/tool rolled back; Ctrl+C cancels turn; TTY confirms off-loop; auto-retry 10s/20s/30s then `/retry`.
|
- Failed/cancelled turns: user kept; **committed tool rounds kept** (side effects not re-run on `/retry`); only unfinished assistant rolled back; Ctrl+C cancels; TTY confirms off-loop; auto-retry 10s/20s/30s then `/retry`.
|
||||||
- **`plyngent providers`**, **`config path|edit`**. No providers + `$EDITOR` → optional edit then reload.
|
- **`plyngent providers`**, **`config path|edit`**. No providers + `$EDITOR` → optional edit then reload.
|
||||||
- Tools default on; workspace defaults to cwd; `--max-rounds` default 32. Readline history under platformdirs (`repl_history`). PTY: `close_all` on chat exit.
|
- Tools default on; workspace defaults to cwd; `--max-rounds` default 32. Readline history under platformdirs (`repl_history`). PTY: `close_all` on chat exit.
|
||||||
|
|
||||||
|
|||||||
+131
-23
@@ -2,7 +2,15 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from plyngent.lmproto.openai_compatible.model import SystemChatMessage, UserChatMessage
|
from msgspec import UNSET
|
||||||
|
|
||||||
|
from plyngent.lmproto.openai_compatible.model import (
|
||||||
|
AssistantChatMessage,
|
||||||
|
AssistantFunctionToolCall,
|
||||||
|
SystemChatMessage,
|
||||||
|
ToolChatMessage,
|
||||||
|
UserChatMessage,
|
||||||
|
)
|
||||||
|
|
||||||
from .budget import (
|
from .budget import (
|
||||||
DEFAULT_CONTEXT_MAX_TOKENS,
|
DEFAULT_CONTEXT_MAX_TOKENS,
|
||||||
@@ -26,6 +34,74 @@ if TYPE_CHECKING:
|
|||||||
type LimitContinueHook = Callable[[str], bool | Awaitable[bool]]
|
type LimitContinueHook = Callable[[str], bool | Awaitable[bool]]
|
||||||
|
|
||||||
|
|
||||||
|
def incomplete_turn_user_text(messages: Sequence[AnyChatMessage]) -> str | None:
|
||||||
|
"""User text for an incomplete turn, if history can be retried.
|
||||||
|
|
||||||
|
Incomplete when the last non-system message is a user message (failed before
|
||||||
|
any commit), or tool results after a complete tool batch (failed on a later
|
||||||
|
model round — keep tools so retry does not re-execute side effects).
|
||||||
|
"""
|
||||||
|
# Skip leading system prompt only.
|
||||||
|
start = 0
|
||||||
|
if messages and isinstance(messages[0], SystemChatMessage):
|
||||||
|
start = 1
|
||||||
|
body = list(messages[start:])
|
||||||
|
if not body:
|
||||||
|
return None
|
||||||
|
last = body[-1]
|
||||||
|
if isinstance(last, UserChatMessage):
|
||||||
|
return last.content
|
||||||
|
if isinstance(last, ToolChatMessage):
|
||||||
|
# Walk back to the user that opened this turn.
|
||||||
|
for msg in reversed(body):
|
||||||
|
if isinstance(msg, UserChatMessage):
|
||||||
|
return msg.content
|
||||||
|
return None
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def committed_prefix_end(messages: Sequence[AnyChatMessage], user_index: int) -> int:
|
||||||
|
"""Index of the first uncommitted message after a successful tool batch.
|
||||||
|
|
||||||
|
Committed = complete assistant(tool_calls) + matching tool results for that
|
||||||
|
batch, possibly repeated. An assistant without finished tools, or a partial
|
||||||
|
text-only assistant mid-stream, is not committed.
|
||||||
|
"""
|
||||||
|
end = user_index + 1
|
||||||
|
i = user_index + 1
|
||||||
|
n = len(messages)
|
||||||
|
while i < n:
|
||||||
|
msg = messages[i]
|
||||||
|
if not isinstance(msg, AssistantChatMessage):
|
||||||
|
break
|
||||||
|
tool_calls = msg.tool_calls
|
||||||
|
if tool_calls is UNSET or not tool_calls:
|
||||||
|
# Final text assistant is only committed on successful turn end.
|
||||||
|
break
|
||||||
|
expected_ids: list[str] = []
|
||||||
|
for call in tool_calls:
|
||||||
|
if isinstance(call, AssistantFunctionToolCall):
|
||||||
|
expected_ids.append(call.id)
|
||||||
|
else:
|
||||||
|
expected_ids.append(call.id)
|
||||||
|
j = i + 1
|
||||||
|
got: list[str] = []
|
||||||
|
while j < n and isinstance(messages[j], ToolChatMessage):
|
||||||
|
tool_msg = messages[j]
|
||||||
|
assert isinstance(tool_msg, ToolChatMessage)
|
||||||
|
got.append(tool_msg.tool_call_id)
|
||||||
|
j += 1
|
||||||
|
if len(got) < len(expected_ids):
|
||||||
|
# Incomplete tool batch — do not commit this assistant.
|
||||||
|
break
|
||||||
|
# Allow extra tool messages only if exact prefix match of ids (order may vary).
|
||||||
|
if sorted(got[: len(expected_ids)]) != sorted(expected_ids):
|
||||||
|
break
|
||||||
|
end = j
|
||||||
|
i = j
|
||||||
|
return end
|
||||||
|
|
||||||
|
|
||||||
class ChatAgent:
|
class ChatAgent:
|
||||||
"""Thin wrapper: chat client + optional tools + optional memory bind."""
|
"""Thin wrapper: chat client + optional tools + optional memory bind."""
|
||||||
|
|
||||||
@@ -47,6 +123,8 @@ class ChatAgent:
|
|||||||
last_turn_usage: TokenUsage
|
last_turn_usage: TokenUsage
|
||||||
last_request_usage: TokenUsage
|
last_request_usage: TokenUsage
|
||||||
last_turn_rounds: int
|
last_turn_rounds: int
|
||||||
|
# Index into messages of the first unpersisted message (checkpoint cursor).
|
||||||
|
_persist_from: int
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -84,18 +162,13 @@ class ChatAgent:
|
|||||||
self.last_turn_usage = TokenUsage()
|
self.last_turn_usage = TokenUsage()
|
||||||
self.last_request_usage = TokenUsage()
|
self.last_request_usage = TokenUsage()
|
||||||
self.last_turn_rounds = 0
|
self.last_turn_rounds = 0
|
||||||
|
self._persist_from = len(self.messages)
|
||||||
self._ensure_system_prompt()
|
self._ensure_system_prompt()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def pending_retry_text(self) -> str | None:
|
def pending_retry_text(self) -> str | None:
|
||||||
"""Text of an incomplete last user turn, if any.
|
"""User text of an incomplete turn that can be continued with :meth:`retry`."""
|
||||||
|
return incomplete_turn_user_text(self.messages)
|
||||||
Derived only: history ending with a user message means that turn never
|
|
||||||
completed (failure, cancel, or resume of an orphan user in DB).
|
|
||||||
"""
|
|
||||||
if self.messages and isinstance(self.messages[-1], UserChatMessage):
|
|
||||||
return self.messages[-1].content
|
|
||||||
return None
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def context_tokens(self) -> int:
|
def context_tokens(self) -> int:
|
||||||
@@ -123,6 +196,9 @@ class ChatAgent:
|
|||||||
if self.messages and isinstance(self.messages[0], SystemChatMessage):
|
if self.messages and isinstance(self.messages[0], SystemChatMessage):
|
||||||
return
|
return
|
||||||
self.messages.insert(0, SystemChatMessage(content=self.system_prompt))
|
self.messages.insert(0, SystemChatMessage(content=self.system_prompt))
|
||||||
|
# System inject is not a DB message unless already stored.
|
||||||
|
if self._persist_from == 0:
|
||||||
|
self._persist_from = 1
|
||||||
|
|
||||||
async def load_history(self) -> None:
|
async def load_history(self) -> None:
|
||||||
"""Replace in-memory messages from the bound memory session."""
|
"""Replace in-memory messages from the bound memory session."""
|
||||||
@@ -130,6 +206,7 @@ class ChatAgent:
|
|||||||
msg = "load_history requires memory and session_id"
|
msg = "load_history requires memory and session_id"
|
||||||
raise RuntimeError(msg)
|
raise RuntimeError(msg)
|
||||||
self.messages = await self.memory.list_messages(self.session_id)
|
self.messages = await self.memory.list_messages(self.session_id)
|
||||||
|
self._persist_from = len(self.messages)
|
||||||
self._ensure_system_prompt()
|
self._ensure_system_prompt()
|
||||||
|
|
||||||
async def bind_session(self, session_id: int, *, load: bool = True) -> None:
|
async def bind_session(self, session_id: int, *, load: bool = True) -> None:
|
||||||
@@ -145,6 +222,12 @@ class ChatAgent:
|
|||||||
if self.memory is not None and self.session_id is not None:
|
if self.memory is not None and self.session_id is not None:
|
||||||
_ = await self.memory.append_message(self.session_id, message)
|
_ = await self.memory.append_message(self.session_id, message)
|
||||||
|
|
||||||
|
async def _persist_range(self, start: int, end: int) -> None:
|
||||||
|
"""Persist messages[start:end] and advance the checkpoint cursor."""
|
||||||
|
for message in self.messages[start:end]:
|
||||||
|
await self._persist(message)
|
||||||
|
self._persist_from = max(self._persist_from, end)
|
||||||
|
|
||||||
def _user_index(self, user_msg: UserChatMessage) -> int:
|
def _user_index(self, user_msg: UserChatMessage) -> int:
|
||||||
for i in range(len(self.messages) - 1, -1, -1):
|
for i in range(len(self.messages) - 1, -1, -1):
|
||||||
if self.messages[i] is user_msg:
|
if self.messages[i] is user_msg:
|
||||||
@@ -156,16 +239,30 @@ class ChatAgent:
|
|||||||
msg = "user message not found in history"
|
msg = "user message not found in history"
|
||||||
raise RuntimeError(msg)
|
raise RuntimeError(msg)
|
||||||
|
|
||||||
def _rollback_partial(self, user_index: int) -> None:
|
def _find_turn_user_index(self) -> int:
|
||||||
"""Drop assistant/tool messages after the user; keep user for retry/DB."""
|
"""Index of the user message that opens the current incomplete turn."""
|
||||||
del self.messages[user_index + 1 :]
|
text = incomplete_turn_user_text(self.messages)
|
||||||
|
if text is None:
|
||||||
|
msg = "nothing to retry"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
for i in range(len(self.messages) - 1, -1, -1):
|
||||||
|
msg = self.messages[i]
|
||||||
|
if isinstance(msg, UserChatMessage) and msg.content == text:
|
||||||
|
return i
|
||||||
|
msg = "nothing to retry"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
|
||||||
|
def _rollback_uncommitted(self, user_index: int) -> None:
|
||||||
|
"""Drop incomplete suffix; keep committed tool rounds after the user."""
|
||||||
|
end = committed_prefix_end(self.messages, user_index)
|
||||||
|
del self.messages[end:]
|
||||||
|
|
||||||
async def _run_from_user_message(self, user_msg: UserChatMessage) -> AsyncIterator[AgentEvent]:
|
async def _run_from_user_message(self, user_msg: UserChatMessage) -> AsyncIterator[AgentEvent]:
|
||||||
"""Run the tool loop for an already-appended user message.
|
"""Run the tool loop for an already-appended user message.
|
||||||
|
|
||||||
On success, persists assistant/tool messages produced after the user.
|
After each completed tool batch, commits assistant+tool messages to the
|
||||||
On failure, keeps the trailing user message so :meth:`retry` can re-run
|
DB and keeps them on failure so :meth:`retry` continues without redoing
|
||||||
without duplicating it (see :attr:`pending_retry_text`).
|
side-effecting tools. Unfinished assistant/stream suffix is rolled back.
|
||||||
"""
|
"""
|
||||||
user_index = self._user_index(user_msg)
|
user_index = self._user_index(user_msg)
|
||||||
|
|
||||||
@@ -192,18 +289,23 @@ class ChatAgent:
|
|||||||
last_request = event.usage
|
last_request = event.usage
|
||||||
turn_usage = turn_usage.add(event.usage)
|
turn_usage = turn_usage.add(event.usage)
|
||||||
self.session_usage = self.session_usage.add(event.usage)
|
self.session_usage = self.session_usage.add(event.usage)
|
||||||
|
# After a full tool batch, commit prefix so failures do not
|
||||||
|
# discard work that already had external effects.
|
||||||
|
commit_end = committed_prefix_end(self.messages, user_index)
|
||||||
|
if commit_end > self._persist_from:
|
||||||
|
await self._persist_range(self._persist_from, commit_end)
|
||||||
yield event
|
yield event
|
||||||
completed = True
|
completed = True
|
||||||
except BaseException:
|
except BaseException:
|
||||||
if not completed:
|
if not completed:
|
||||||
self._rollback_partial(user_index)
|
self._rollback_uncommitted(user_index)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
self.last_turn_usage = turn_usage
|
self.last_turn_usage = turn_usage
|
||||||
self.last_request_usage = last_request
|
self.last_request_usage = last_request
|
||||||
self.last_turn_rounds = turn_rounds
|
self.last_turn_rounds = turn_rounds
|
||||||
for message in self.messages[user_index + 1 :]:
|
if self._persist_from < len(self.messages):
|
||||||
await self._persist(message)
|
await self._persist_range(self._persist_from, len(self.messages))
|
||||||
|
|
||||||
async def run(self, user_text: str) -> AsyncIterator[AgentEvent]:
|
async def run(self, user_text: str) -> AsyncIterator[AgentEvent]:
|
||||||
"""Append a user message (persist immediately), run the tool loop, yield events."""
|
"""Append a user message (persist immediately), run the tool loop, yield events."""
|
||||||
@@ -211,19 +313,25 @@ class ChatAgent:
|
|||||||
user_msg = UserChatMessage(content=user_text)
|
user_msg = UserChatMessage(content=user_text)
|
||||||
self.messages.append(user_msg)
|
self.messages.append(user_msg)
|
||||||
await self._persist(user_msg)
|
await self._persist(user_msg)
|
||||||
|
self._persist_from = len(self.messages)
|
||||||
async for event in self._run_from_user_message(user_msg):
|
async for event in self._run_from_user_message(user_msg):
|
||||||
yield event
|
yield event
|
||||||
|
|
||||||
async def retry(self) -> AsyncIterator[AgentEvent]:
|
async def retry(self) -> AsyncIterator[AgentEvent]:
|
||||||
"""Re-run the incomplete last user turn without appending a new user message.
|
"""Continue an incomplete turn without re-appending the user message.
|
||||||
|
|
||||||
Requires history to end with a :class:`UserChatMessage` (failed/cancelled
|
Supports:
|
||||||
turn or resumed orphan user in the session DB).
|
- History ending with the user (failed before any tool commit / orphan).
|
||||||
|
- History ending with tool results (failed after tools; continue model).
|
||||||
"""
|
"""
|
||||||
if not self.messages or not isinstance(self.messages[-1], UserChatMessage):
|
if incomplete_turn_user_text(self.messages) is None:
|
||||||
msg = "nothing to retry"
|
msg = "nothing to retry"
|
||||||
raise RuntimeError(msg)
|
raise RuntimeError(msg)
|
||||||
user_msg = self.messages[-1]
|
user_index = self._find_turn_user_index()
|
||||||
|
user_msg = self.messages[user_index]
|
||||||
|
assert isinstance(user_msg, UserChatMessage)
|
||||||
self._ensure_system_prompt()
|
self._ensure_system_prompt()
|
||||||
|
# Already-committed messages stay; only new rounds are persisted.
|
||||||
|
self._persist_from = len(self.messages)
|
||||||
async for event in self._run_from_user_message(user_msg):
|
async for event in self._run_from_user_message(user_msg):
|
||||||
yield event
|
yield event
|
||||||
|
|||||||
@@ -196,12 +196,20 @@ async def run_user_text_with_retries(
|
|||||||
|
|
||||||
|
|
||||||
async def retry_pending_with_retries(agent: ChatAgent) -> bool:
|
async def retry_pending_with_retries(agent: ChatAgent) -> bool:
|
||||||
"""Retry the last failed user turn with auto-retry."""
|
"""Retry/continue an incomplete turn with auto-retry.
|
||||||
|
|
||||||
|
Continues from committed tool results when present (does not re-run them).
|
||||||
|
"""
|
||||||
if agent.pending_retry_text is None:
|
if agent.pending_retry_text is None:
|
||||||
click.echo("nothing to retry")
|
click.echo("nothing to retry")
|
||||||
return False
|
return False
|
||||||
preview = agent.pending_retry_text
|
preview = agent.pending_retry_text
|
||||||
if len(preview) > _PREVIEW_LEN:
|
if len(preview) > _PREVIEW_LEN:
|
||||||
preview = preview[:_PREVIEW_LEN] + "…"
|
preview = preview[:_PREVIEW_LEN] + "…"
|
||||||
click.echo(f"retrying: {preview}")
|
from plyngent.lmproto.openai_compatible.model import ToolChatMessage
|
||||||
|
|
||||||
|
if agent.messages and isinstance(agent.messages[-1], ToolChatMessage):
|
||||||
|
click.echo(f"continuing after tools (user: {preview})")
|
||||||
|
else:
|
||||||
|
click.echo(f"retrying: {preview}")
|
||||||
return await run_turn_with_retries(agent, starter=agent.retry)
|
return await run_turn_with_retries(agent, starter=agent.retry)
|
||||||
|
|||||||
@@ -691,7 +691,11 @@ def history_cmd(state: ReplState, n: int | None) -> None:
|
|||||||
@slash.command("retry")
|
@slash.command("retry")
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def retry_cmd(state: ReplState) -> None:
|
def retry_cmd(state: ReplState) -> None:
|
||||||
"""Re-run incomplete last user turn (DB/orphan user; no retype)."""
|
"""Continue an incomplete turn (user-only or after committed tools).
|
||||||
|
|
||||||
|
Does not retype the user message. After tools already ran, continues the
|
||||||
|
model loop without re-executing those tool calls.
|
||||||
|
"""
|
||||||
_ = _await(retry_pending_with_retries(state.agent))
|
_ = _await(retry_pending_with_retries(state.agent))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ from plyngent.lmproto.openai_compatible.model import (
|
|||||||
DeltaMessage,
|
DeltaMessage,
|
||||||
StreamFunctionDelta,
|
StreamFunctionDelta,
|
||||||
StreamToolCallDelta,
|
StreamToolCallDelta,
|
||||||
|
ToolChatMessage,
|
||||||
UserChatMessage,
|
UserChatMessage,
|
||||||
)
|
)
|
||||||
from plyngent.memory import MemoryStore
|
from plyngent.memory import MemoryStore
|
||||||
@@ -690,6 +691,88 @@ async def test_retry_after_resume_orphan_user() -> None:
|
|||||||
await store.close()
|
await store.close()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_retry_keeps_committed_tools_after_second_round_fails() -> None:
|
||||||
|
"""Tools that already ran stay in history; retry continues without re-calling them."""
|
||||||
|
store = await MemoryStore.open(DatabaseConfig())
|
||||||
|
session = await store.create_session(name="t")
|
||||||
|
calls: list[str] = []
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def side_effect() -> str:
|
||||||
|
calls.append("ran")
|
||||||
|
return "done-once"
|
||||||
|
|
||||||
|
registry = ToolRegistry([side_effect])
|
||||||
|
|
||||||
|
class ToolThenBoom:
|
||||||
|
n: int
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.n = 0
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def chat_completions(
|
||||||
|
self, param: ChatCompletionsParam, *, stream: Literal[False] = False
|
||||||
|
) -> ChatCompletionResponse: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def chat_completions(
|
||||||
|
self, param: ChatCompletionsParam, *, stream: Literal[True]
|
||||||
|
) -> AsyncIterator[ChatCompletionChunk]: ...
|
||||||
|
|
||||||
|
async def chat_completions(
|
||||||
|
self, param: ChatCompletionsParam, *, stream: bool = False
|
||||||
|
) -> ChatCompletionResponse | AsyncIterator[ChatCompletionChunk]:
|
||||||
|
del stream
|
||||||
|
self.n += 1
|
||||||
|
if self.n == 1:
|
||||||
|
return _response(
|
||||||
|
AssistantChatMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
AssistantFunctionToolCall(
|
||||||
|
id="c1",
|
||||||
|
function=AssistantFunctionTool(
|
||||||
|
name="side_effect",
|
||||||
|
arguments="{}",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if self.n == 2:
|
||||||
|
msg = "api down after tools"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
# Retry continues: third call is next model round with tools in history.
|
||||||
|
assert any(getattr(m, "tool_call_id", None) == "c1" for m in param.messages)
|
||||||
|
return _response(AssistantChatMessage(content="finished"))
|
||||||
|
|
||||||
|
client = ToolThenBoom()
|
||||||
|
agent = ChatAgent(
|
||||||
|
client,
|
||||||
|
model="m",
|
||||||
|
tools=registry,
|
||||||
|
memory=store,
|
||||||
|
session_id=session.sid,
|
||||||
|
stream=False,
|
||||||
|
)
|
||||||
|
with pytest.raises(RuntimeError, match="api down after tools"):
|
||||||
|
_ = [e async for e in agent.run("do it")]
|
||||||
|
|
||||||
|
assert calls == ["ran"]
|
||||||
|
assert agent.pending_retry_text == "do it"
|
||||||
|
# User + assistant(tool_calls) + tool result kept; incomplete second assistant dropped.
|
||||||
|
assert isinstance(agent.messages[-1], ToolChatMessage)
|
||||||
|
loaded = await store.list_messages(session.sid)
|
||||||
|
assert len(loaded) == 3 # user, assistant, tool — committed before boom
|
||||||
|
|
||||||
|
events = [e async for e in agent.retry()]
|
||||||
|
assert any(isinstance(e, TextDeltaEvent) and e.content == "finished" for e in events)
|
||||||
|
assert calls == ["ran"] # tool not re-executed
|
||||||
|
assert agent.pending_retry_text is None
|
||||||
|
await store.close()
|
||||||
|
|
||||||
|
|
||||||
async def test_chat_agent_system_prompt_prepended() -> None:
|
async def test_chat_agent_system_prompt_prepended() -> None:
|
||||||
from plyngent.lmproto.openai_compatible.model import SystemChatMessage
|
from plyngent.lmproto.openai_compatible.model import SystemChatMessage
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user