diff --git a/README.md b/README.md index 8bf14c419..5fe469939 100644 --- a/README.md +++ b/README.md @@ -183,7 +183,7 @@ For the exploration workflow and local-memory evaluation, see [LIBERO exploratio ### Interactive CLI mode -Add `--interactive` (`-i`) to steer the agent live from your terminal. At the `you>` prompt, the built-in task is pre-filled — press Enter to use it or replace it with your own — then type any message while it runs to steer the agent at the next turn (`/help` lists commands; `/quit` or Ctrl-D ends). Requires an interactive terminal (TTY). +With `claude_code` or `codex`, add `--interactive` (`-i`) to steer the agent live from your terminal. The built-in task is pre-filled at the `you>` prompt. Press Enter to submit it as-is. To change or replace the task, edit the input before pressing Enter. While the agent runs, type follow-up messages to steer it (`/help` lists commands; `/quit` or Ctrl-D ends). Requires a TTY. The `api` planner uses a native CLI that runs the preset task first and accepts follow-ups between runs; use `/exit` to close it. ```bash rpent --robot libero --suite libero_object_swap --task 2 --seed 0 \ diff --git a/README.zh-CN.md b/README.zh-CN.md index 02b7cf091..ca6895d72 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -184,7 +184,7 @@ rpent --robot libero --suite libero_object_swap --task 2 --seed 0 \ ### 交互模式 -加上 `--interactive`(`-i`)即可在终端里实时引导智能体。在 `you>` 提示符处,内置任务已预填——按 Enter 直接使用,或替换为你自己的任务;智能体运行时,随时输入消息即可在下一轮引导它(`/help` 查看命令,`/quit` 或 Ctrl-D 结束)。需要交互式终端(TTY)。 +使用 `claude_code` 或 `codex` 时,加上 `--interactive`(`-i`)即可在终端里实时引导智能体。内置任务会预填在 `you>` 提示符处。直接按 Enter 即可提交该任务;如需修改或替换任务,请先编辑输入内容,再按 Enter 提交。任务运行期间,可以继续输入消息引导智能体(`/help` 查看命令,`/quit` 或 Ctrl-D 结束)。需要真实终端(TTY)。`api` planner 使用原生 CLI,先运行预设任务,再在每轮运行完成后接收输入,使用 `/exit` 退出。 ```bash rpent --robot libero --suite libero_object_swap --task 2 --seed 0 \ diff --git a/docs/source-en/rst_source/development/architecture.rst b/docs/source-en/rst_source/development/architecture.rst index b38485ecc..2ec375775 100644 --- a/docs/source-en/rst_source/development/architecture.rst +++ b/docs/source-en/rst_source/development/architecture.rst @@ -34,8 +34,8 @@ to add a new primitive. **A swappable planner.** The planner is the LLM agent runtime that drives the tool-calling loop. One ``--planner`` flag switches it while the tools and -prompts stay put. Three are built in: ``api`` is RPent's own tool-calling loop -(built on pydantic-ai, the default, provider-agnostic across model APIs); +prompts stay put. Three are built in: ``api`` uses the native Pydantic AI loop +with Harness sliding-window history trimming (the default, provider-agnostic); ``claude_code`` reuses the Claude Agent SDK runtime; ``codex`` reuses the Codex SDK runtime. Because all three face the exact same tools, they can be compared head-to-head on the same physical benchmark. See diff --git a/docs/source-en/rst_source/quickstart.rst b/docs/source-en/rst_source/quickstart.rst index 0663f3dbb..664184c4d 100644 --- a/docs/source-en/rst_source/quickstart.rst +++ b/docs/source-en/rst_source/quickstart.rst @@ -68,6 +68,11 @@ streams agent reasoning, camera views, and the action timeline; submit another task after the current one finishes. Use ``--dashboard-language zh-cn`` for the Chinese UI. +.. _quickstart-interactive: + +For terminal interaction, add ``--interactive`` (``-i``) to enter follow-up +instructions. It cannot be combined with ``--dashboard``. + Key CLI options --------------- diff --git a/docs/source-en/rst_source/usage/configure_planner.rst b/docs/source-en/rst_source/usage/configure_planner.rst index 5a495a418..b8c8a0e44 100644 --- a/docs/source-en/rst_source/usage/configure_planner.rst +++ b/docs/source-en/rst_source/usage/configure_planner.rst @@ -20,11 +20,9 @@ loop is orchestrated, and which model SDK is used. - What it is - When to pick it * - ``api`` - - Provider-agnostic tool-calling loop built on - `pydantic-ai `_. It currently supports - the Anthropic Messages API, the OpenAI Responses API, and - OpenAI-compatible Chat Completions APIs. It handles prompt caching - and history-image pruning. + - A tool-calling loop built on `Pydantic AI `_ + that supports multiple model APIs. For long conversations, it sends + fewer older messages to the model. - You want the tightest control over model calls, the widest provider coverage, or the cheapest per-turn spend. * - ``claude_code`` @@ -50,8 +48,8 @@ loop is orchestrated, and which model SDK is used. The ``api`` planner (direct model API) --------------------------------------- -``--planner api`` is the default. It uses Pydantic AI to implement the -tool-calling loop and requires a provider prefix in ``--model``. The +``--planner api`` is the default. It uses the native Pydantic AI tool-calling +runtime and requires a provider prefix in ``--model``. The project currently installs the Anthropic and OpenAI integrations, so it can directly use the Anthropic Messages API, the OpenAI Responses API, and OpenAI-compatible Chat Completions APIs. @@ -79,12 +77,16 @@ needed): Relevant ``api`` planner knobs: - ``--max-tokens`` — cap each LLM reply (default ``8192``). -- ``--max-turns`` — cap the number of tool-calling turns (default - ``100``). +- ``--max-turns`` — cap model requests across the whole conversation, + including retries and follow-ups (default ``100``). Reaching the cap is + a normal stop, not a planner error or a claim of task success. Exploration + can continue with the next session and merge memory when otherwise eligible. - ``--no-images`` — never send image bytes; this is required for text-only models. The agent then reasons from textual state alone, so task performance may not be satisfactory. +For ``--interactive`` usage, see :ref:`Terminal interaction `. + .. _planner-claude-code: The ``claude_code`` planner @@ -262,17 +264,17 @@ schemas, or your context length. Add a custom planner -------------------- -If none of the three planners fit — say you want to plug in an +If none of the built-in planners fit — say you want to plug in an in-house planner, a research prototype, or a different agent SDK — -implement the ``rpent.planner.base.Planner`` protocol and add a -construction branch to ``rpent.planner.base.build_planner``: +subclass ``rpent.planner.base.Planner``, implement its abstract ``solve()`` +method, and update ``rpent.planner.base.build_planner`` to create the new backend: .. code-block:: python # rpent/planner/my_planner.py - from rpent.planner.base import PlannerResult + from rpent.planner.base import Planner, PlannerResult - class MyPlanner: + class MyPlanner(Planner): def solve( self, *, @@ -281,6 +283,7 @@ construction branch to ``rpent.planner.base.build_planner``: toolkit, max_turns, input_queue=None, + dashboard_interaction=None, ): tool_specs = toolkit.get_tools_spec() # Call the model with system_prompt, user_message, and tool_specs. diff --git a/docs/source-zh/rst_source/development/architecture.rst b/docs/source-zh/rst_source/development/architecture.rst index 3cfe7784c..1f947cd28 100644 --- a/docs/source-zh/rst_source/development/architecture.rst +++ b/docs/source-zh/rst_source/development/architecture.rst @@ -28,7 +28,7 @@ **可替换的 planner。** planner 就是驱动工具调用循环的 LLM agent 运行时。 一个 ``--planner`` 参数就能切换它,而工具和提示词保持不变。内置三种: -``api`` 是 RPent 自研的工具调用循环(基于 pydantic-ai,为默认值, +``api`` 使用 Pydantic AI 原生运行时与 Harness 滑动窗口历史裁剪(为默认值, 不绑定具体模型提供商);``claude_code`` 复用 Claude Agent SDK 运行时; ``codex`` 复用 Codex SDK 运行时。由于三者面对完全相同的工具, 可以在同一套物理基准上正面对比。配置方法见 :doc:`../usage/configure_planner`。 diff --git a/docs/source-zh/rst_source/quickstart.rst b/docs/source-zh/rst_source/quickstart.rst index a45783e85..b154c4bc5 100644 --- a/docs/source-zh/rst_source/quickstart.rst +++ b/docs/source-zh/rst_source/quickstart.rst @@ -64,6 +64,11 @@ Session 配置全部来自命令行,打开地址后直接进入实时监控; 推理过程、相机画面和动作时间线;任务结束后可以继续提交下一任务。使用 ``--dashboard-language zh-cn`` 可切换到中文界面。 +.. _quickstart-interactive: + +也可以添加 ``--interactive``(``-i``),在终端输入后续指令。 +该选项不能与 ``--dashboard`` 同时使用。 + 关键 CLI 选项 ------------- diff --git a/docs/source-zh/rst_source/usage/configure_planner.rst b/docs/source-zh/rst_source/usage/configure_planner.rst index 3ca3196f1..0df22e725 100644 --- a/docs/source-zh/rst_source/usage/configure_planner.rst +++ b/docs/source-zh/rst_source/usage/configure_planner.rst @@ -20,8 +20,7 @@ SDK。 - 什么时候选它 * - ``api`` - 基于 `Pydantic AI `_ 实现的工具调用循环, - 不绑定特定模型提供商。当前支持 Anthropic Messages API、OpenAI Responses - API 和 OpenAI 兼容的 Chat Completions API,内置 prompt 缓存和历史图片剪枝。 + 支持多种模型 API。对话过长时,会减少发送给模型的早期记录。 - 需要精细控制模型调用、支持更多模型提供商,或降低单轮调用成本。 * - ``claude_code`` - `Claude Agent SDK @@ -44,7 +43,7 @@ SDK。 ``api`` planner(直接调用模型 API) ------------------------------------- -``--planner api`` 是默认选项。它使用 Pydantic AI 实现工具调用循环,并要求 +``--planner api`` 是默认选项。它使用 Pydantic AI 原生工具调用运行时,并要求 ``--model`` 带有模型提供商前缀。当前项目安装的依赖包含 Anthropic 和 OpenAI 集成,因此可以直接使用 Anthropic Messages API、OpenAI Responses API, 以及 OpenAI 兼容的 Chat Completions API。 @@ -71,10 +70,14 @@ SDK。 ``api`` planner 的相关调节参数: - ``--max-tokens`` —— 单次 LLM 回复的 token 上限(默认 ``8192``)。 -- ``--max-turns`` —— 工具调用轮数上限(默认 ``100``)。 +- ``--max-turns`` —— 整段对话的模型请求次数上限,包含重试和后续输入 + (默认 ``100``)。达到上限时正常停止,不记为规划器错误,也不代表任务成功。 + 探索模式仍可继续下一会话,并在满足其他条件时合并记忆。 - ``--no-images`` —— 不向模型发送图片字节;纯文本模型必须加此参数。此时 智能体只依赖文本状态推理,任务表现可能不够理想。 +``--interactive`` 的用法见 :ref:`终端交互 `。 + .. _planner-claude-code: ``claude_code`` planner @@ -237,16 +240,16 @@ Dashboard 提供同一项检查:启动页的 **测试连接** 按钮会针对 接入自定义 planner ------------------ -如果三种内置 planner 都不合适,例如需要接入内部 planner、研究原型或其他 -agent SDK,可以实现 ``rpent.planner.base.Planner`` 协议,并在 -``rpent.planner.base.build_planner`` 中增加对应的构造分支: +如果内置 planner 都不合适,例如需要接入内部 planner、研究原型或其他 +agent SDK,可以继承 ``rpent.planner.base.Planner``,实现抽象方法 ``solve()``, +并修改 ``rpent.planner.base.build_planner``,使其能够创建新后端的实例: .. code-block:: python # rpent/planner/my_planner.py - from rpent.planner.base import PlannerResult + from rpent.planner.base import Planner, PlannerResult - class MyPlanner: + class MyPlanner(Planner): def solve( self, *, @@ -255,6 +258,7 @@ agent SDK,可以实现 ``rpent.planner.base.Planner`` 协议,并在 toolkit, max_turns, input_queue=None, + dashboard_interaction=None, ): tool_specs = toolkit.get_tools_spec() # 使用 system_prompt、user_message 和 tool_specs 调用模型。 diff --git a/pyproject.toml b/pyproject.toml index 70ae39906..570abdd15 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,7 +27,8 @@ classifiers = [ ] dependencies = [ - "pydantic-ai-slim[anthropic,openai]>=2.1", + "pydantic-ai-slim[anthropic,openai,cli]>=2.43", + "pydantic-ai-harness>=0.31", "pydantic>=2", "fastapi>=0.110", "uvicorn>=0.27", diff --git a/rpent/cli/main.py b/rpent/cli/main.py index 7dd8912ff..1bda6f2de 100644 --- a/rpent/cli/main.py +++ b/rpent/cli/main.py @@ -306,6 +306,7 @@ def _start_continuation_session( claude_code_max_budget_usd=args.claude_code_max_budget_usd, dashboard_events=dashboard_events, no_images=args.no_images, + interactive=args.interactive, ) system_prompt = prompt_bundle.render( "system", @@ -368,6 +369,7 @@ def main() -> int: ) if sys.stdin is None or not sys.stdin.isatty(): parser.error("This robot requires a TTY for operator confirmation.") + native_cli = args.interactive and args.planner == "api" if args.base_url and args.planner in BASE_URL_ENV_BY_PLANNER: parser.error( "--base-url applies to the 'api' planner only; " @@ -452,6 +454,7 @@ def main() -> int: claude_code_max_budget_usd=args.claude_code_max_budget_usd, dashboard_events=dashboard_events, no_images=args.no_images, + interactive=args.interactive, ) prompt_bundle = robot_spec.prompts prompt_vars = {**prompt_vars, "output_dir": output_dir} @@ -468,10 +471,13 @@ def main() -> int: if human_interactive_exploration: from rpent.tools.human_in_the_loop import HumanInTheLoopInput - operator_input = HumanInTheLoopInput(interactive=args.interactive) + # The native CLI reads between runs, leaving the TTY available to tools. + operator_input = HumanInTheLoopInput( + interactive=args.interactive and not native_cli + ) input_queue: "queue.Queue[str | None] | None" = None await_first_prompt: "Callable[[], str | None] | None" = None - if args.interactive: + if args.interactive and not native_cli: input_queue = queue.Queue() # Pre-fill the first prompt with the rendered default task (editable # preset); @@ -575,7 +581,7 @@ def main() -> int: config=run_config, ) memory_manager = toolkit.memory - if operator_input is not None and args.interactive: + if operator_input is not None and input_queue is not None: def accept_verdict(verdict: str, active_toolkit=toolkit) -> bool: if not active_toolkit.request_direct_verdict(verdict): @@ -618,7 +624,7 @@ def accept_verdict(verdict: str, active_toolkit=toolkit) -> bool: if solved and callable(write_recipe): recipe_path = write_recipe(recipe_tag) or recipe_path finally: - if operator_input is not None and args.interactive: + if operator_input is not None and input_queue is not None: operator_input.bind_verdict(None) try: if robot_spec.finalize_run is not None: diff --git a/rpent/dashboard/interaction.py b/rpent/dashboard/interaction.py index 583c9a415..e60ed512a 100644 --- a/rpent/dashboard/interaction.py +++ b/rpent/dashboard/interaction.py @@ -49,6 +49,31 @@ def as_dict(self) -> dict[str, Any]: } +class PlannerSessionDriver(Protocol): + """Backend operations used by the Dashboard planner control. + + The controller drains physical toolkit work before requesting an interrupt. + Drivers own their SDK resources, event consumers, and cleanup. + """ + + async def submit(self, message: DashboardMessage) -> int: + """Submit input and return the number of new completion events expected. + + Steering an active turn may return zero. With deferred acknowledgement, + the driver reports when the message starts or is discarded through + ``DashboardPlannerControl``; returning only confirms acceptance. + """ + ... + + async def interrupt(self) -> int: + """Interrupt execution and return completions to remove from the count. + + Count only work that will not report completion through the normal event + path, such as discarded queued input. Do not count those events twice. + """ + ... + + class DashboardInteractionPort(Protocol): """Planner-facing access to one Dashboard interaction Session.""" diff --git a/rpent/dashboard/planner_control.py b/rpent/dashboard/planner_control.py index ac9a0126c..119aa3d5b 100644 --- a/rpent/dashboard/planner_control.py +++ b/rpent/dashboard/planner_control.py @@ -18,9 +18,8 @@ import asyncio from collections.abc import Callable -from typing import Any -from rpent.dashboard.interaction import DashboardInteractionPort +from rpent.dashboard.interaction import DashboardInteractionPort, PlannerSessionDriver class DashboardPlannerControl: @@ -34,12 +33,14 @@ def __init__( emit_user: Callable[[str], None], emit_initial_user: Callable[[], None], defer_message_ack: bool = False, + submit_while_busy: bool = False, ) -> None: self._interaction = interaction self._cancel_active_and_wait = cancel_active_and_wait self._emit_user = emit_user self._emit_initial_user = emit_initial_user self._defer_message_ack = defer_message_ack + self._submit_while_busy = submit_while_busy self._lock = asyncio.Lock() self._outstanding_completions = 0 @@ -52,7 +53,7 @@ async def start(self) -> None: ) self._emit_initial_user() - async def run(self, driver: Any) -> None: + async def run(self, driver: PlannerSessionDriver) -> None: """Forward Dashboard commands until the interaction ends.""" version = self._interaction.interaction_version while self._interaction.planner_activity != "ended": @@ -62,7 +63,7 @@ async def run(self, driver: Any) -> None: version, ) - async def complete(self, driver: Any) -> None: + async def complete(self, driver: PlannerSessionDriver) -> None: """Record one completed backend request and flush queued input.""" async with self._lock: if self._interaction.planner_activity == "ended": @@ -73,7 +74,7 @@ async def complete(self, driver: Any) -> None: ) await self._flush(driver) - async def tool_completed(self, driver: Any) -> None: + async def tool_completed(self, driver: PlannerSessionDriver) -> None: """Flush input queued while the backend was running a tool.""" async with self._lock: await self._flush(driver) @@ -96,7 +97,7 @@ async def cancel_active_toolkit(self) -> None: """Cancel and drain the active toolkit operation off the event loop.""" await asyncio.to_thread(self._cancel_active_and_wait) - async def _process(self, driver: Any) -> None: + async def _process(self, driver: PlannerSessionDriver) -> None: async with self._lock: if self._interaction.planner_activity == "ended": return @@ -127,17 +128,14 @@ async def _process(self, driver: Any) -> None: ) await self._flush(driver) return - if self._interaction.planner_activity == "idle": + if self._interaction.planner_activity == "idle" or self._submit_while_busy: await self._flush(driver) - async def _flush(self, driver: Any) -> None: + async def _flush(self, driver: PlannerSessionDriver) -> None: message = self._interaction.claim_next_pending_message() while message is not None and not self._interaction.task_replacement_requested: try: - if self._defer_message_ack: - added_completions = await driver.submit_dashboard_message(message) - else: - added_completions = await driver.submit(message.text) + added_completions = await driver.submit(message) except Exception as exc: self._interaction.mark_message_failed( message.message_id, diff --git a/rpent/planner/api_loop.py b/rpent/planner/api_loop.py index fee704307..f60230ef0 100644 --- a/rpent/planner/api_loop.py +++ b/rpent/planner/api_loop.py @@ -12,92 +12,95 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Provider-independent tool-use agent loop built on pydantic-ai. +"""API planner composed from Pydantic AI and Harness capabilities. -The loop wraps the agent's :class:`~rpent.tools.toolkit.Toolkit` as -pydantic-ai function tools and drives :class:`pydantic_ai.Agent` runs, -streaming each turn so progress is logged in real time. Task completion is -signalled by the robot-provided ``finish`` tool, whose result carries ``_finish``. +The SDK owns model requests, tool scheduling, retries, queued-message delivery, +and cancellation. Harness owns compaction; clai owns terminal interaction. +RPent's runner saves the returned transcript through its existing output flow. """ -from __future__ import annotations - import asyncio import base64 -import contextlib -import dataclasses import json -import queue -from collections import deque -from collections.abc import Callable -from pathlib import Path +import threading +from collections.abc import AsyncIterator, Sequence +from contextlib import asynccontextmanager +from functools import partial from typing import Any - -from pydantic_ai import Agent, BinaryContent, ModelSettings, Tool, ToolReturn -from pydantic_ai.capabilities import ProcessHistory, Thinking -from pydantic_ai.exceptions import ModelHTTPError, UsageLimitExceeded -from pydantic_ai.messages import ( - FunctionToolCallEvent, - FunctionToolResultEvent, - ModelMessage, - ModelResponse, +from uuid import uuid4 + +from jsonschema import Draft202012Validator +from pydantic_ai import ( + Agent, + BinaryContent, + CallToolsNode, + CancellationToken, + EnqueuedMessagesEvent, + FunctionToolset, + ModelRequestNode, + ModelRetry, + ModelSettings, + PartEndEvent, + RunCancelled, + RunContext, + RunUsage, + StructuredDict, TextPart, ThinkingPart, + Tool, ToolCallPart, - UserPromptPart, -) -from pydantic_ai.models import Model -from pydantic_ai.usage import RunUsage, UsageLimits - -from rpent.cli.tui import QUIT_TOKENS -from rpent.dashboard.events import ( - DashboardEventSink, - TranscriptEvent, - UsageEvent, + ToolOutput, + ToolReturn, + UsageLimits, ) +from pydantic_ai.agent import WrapperAgent +from pydantic_ai.capabilities import AbstractCapability, Thinking, on_event +from pydantic_ai.capabilities.abstract import AgentNode, NodeResult, WrapRunHandler +from pydantic_ai.exceptions import UsageLimitExceeded, UserError +from pydantic_ai.messages import ModelMessage, UserContent +from pydantic_ai.models import Model, ModelRequestContext +from pydantic_ai.run import AgentRun, AgentRunResult +from pydantic_ai_harness.compaction import SlidingWindowCompaction + +from rpent.dashboard.events import DashboardEventSink, TranscriptEvent, UsageEvent from rpent.dashboard.interaction import DashboardInteractionPort, DashboardMessage from rpent.dashboard.planner_control import DashboardPlannerControl -from rpent.planner.base import REASONING_EFFORTS, PlannerResult -from rpent.session import EnvState -from rpent.tools.toolkit import Toolkit +from rpent.planner.base import REASONING_EFFORTS, Planner, PlannerResult +from rpent.tools.toolkit import Toolkit, ToolResult from rpent.utils.logging import get_logger -logger = get_logger("api_loop") +logger = get_logger("api") -#: Console-log truncation limits (characters). -_TEXT_LOG_LIMIT = 500 -_ARGS_LOG_LIMIT = 250 -_TOOL_LOG_LIMIT = 350 -#: Cap on cumulative decoded image bytes kept in the resent request history. -_MAX_HISTORY_IMAGE_BYTES = 4 * 1024 * 1024 -#: Always retain at least this many of the most recent images, even if a single -#: frame exceeds the byte budget, so the model never loses its current view. -_MIN_RECENT_IMAGES = 2 +class _RequestBudgetExhausted(Exception): + """The conversation has used its configured model-request budget.""" -class ApiAgentLoop: - """Planner that runs the tool-calling loop via a pydantic-ai ``Agent``.""" +class ApiAgentLoop(Planner): + """Implement ``Planner`` with a native Agent and reusable Harness capabilities.""" def __init__( self, + *, model: Model, + dashboard_events: DashboardEventSink, max_tokens: int = 8192, - no_images: bool = False, reasoning_effort: str = "none", - *, - dashboard_events: DashboardEventSink, - timeout_s: int | None = None, - ): - """Store the pydantic-ai model and the output-token cap.""" - self._model = model - self._max_tokens = max_tokens - self._dashboard_events = dashboard_events - self._no_images = no_images - self._timeout_s = timeout_s + no_images: bool = False, + timeout_s: float = 1200, + interactive: bool = False, + ) -> None: if reasoning_effort not in REASONING_EFFORTS: raise ValueError(f"unsupported reasoning effort: {reasoning_effort}") - self._reasoning_effort = reasoning_effort + if max_tokens < 1 or timeout_s <= 0: + raise ValueError("max_tokens and timeout_s must be positive") + self.model = model + self.dashboard_events = dashboard_events + self.max_tokens = max_tokens + self.reasoning_effort = reasoning_effort + self.no_images = no_images + self.timeout_s = timeout_s + self.interactive = interactive def solve( self, @@ -106,844 +109,574 @@ def solve( user_message: str, toolkit: Toolkit, max_turns: int, - input_queue: queue.Queue[str | None] | None = None, + input_queue=None, dashboard_interaction: DashboardInteractionPort | None = None, ) -> PlannerResult: - """Run the tool-calling loop until finish, normal stop, or budget.""" - if input_queue is not None and dashboard_interaction is not None: - raise ValueError( - "input_queue and dashboard_interaction cannot be used together" - ) - if dashboard_interaction is not None: - return asyncio.run( - self._solve_dashboard( - system_prompt=system_prompt, - user_message=user_message, - toolkit=toolkit, - max_turns=max_turns, - interaction=dashboard_interaction, - ) - ) - solve = self._solve( - system_prompt=system_prompt, - user_message=user_message, - toolkit=toolkit, - max_turns=max_turns, - input_queue=input_queue, - ) + """Run a conversation and return its transcript for RPent to save.""" + if max_turns < 1: + raise ValueError("max_turns must be positive") if input_queue is not None: - return asyncio.run(solve) - try: - return asyncio.run(asyncio.wait_for(solve, timeout=self._timeout_s)) - except asyncio.TimeoutError: - toolkit.cancel_active_and_wait() - return PlannerResult( - finish_result=None, - messages=[{"role": "user", "content": user_message}], - stats={}, - error=f"API planner timed out after {self._timeout_s}s", + raise ValueError("api does not use RPent's terminal input queue") + if self.interactive and dashboard_interaction is not None: + raise ValueError("clai and Dashboard cannot run together") + return asyncio.run( + self._solve( + system_prompt, + user_message, + toolkit, + max_turns, + dashboard_interaction, ) - - async def _solve( - self, - *, - system_prompt: str, - user_message: str, - toolkit: Toolkit, - max_turns: int, - input_queue: queue.Queue[str | None] | None = None, - ) -> PlannerResult: - agent = self._build_agent(system_prompt, toolkit) - - interactive = input_queue is not None - messages: list[dict[str, Any]] = [{"role": "user", "content": user_message}] - observer = _ApiRunObserver( - dashboard_events=self._dashboard_events, - messages=messages, - max_turns=max_turns, - ) - last_error: str | None = None - usage: RunUsage | None = None - quit_requested = False - - def _inject_pending(run: Any) -> bool: - """Drain queued user lines into the live run; True => end session. - - Each line is enqueued ``asap`` so it lands in the next model request - (the next turn boundary). This runs on the event-loop thread, so - mutating the run's pending-message queue here is race-free. - """ - while True: - try: - line = input_queue.get_nowait() # type: ignore[union-attr] - except queue.Empty: - return False - if line is None: - return True - line = line.strip() - if line.lower() in QUIT_TOKENS: - return True - if not line: - continue - run.enqueue(line, priority="asap") - messages.append({"role": "user", "content": line}) - logger.info("[user] %s", _clip(line, _ARGS_LOG_LIMIT)) - - async def _await_next() -> str | None: - """Block off-loop for the next user line between runs (None => end).""" - logger.info("awaiting input — type a message to continue, /quit to end") - while True: - line = await asyncio.to_thread(input_queue.get) # type: ignore[union-attr] - if line is None: - return None - line = line.strip() - if line.lower() in QUIT_TOKENS: - return None - if line: - logger.info("[user] %s", _clip(line, _ARGS_LOG_LIMIT)) - return line - - seed = user_message - history: list[ModelMessage] | None = None - try: - while True: - run_turns = 0 - # request_limit overrides pydantic-ai's default (50) so the - # manual max_turns break below is what bounds each run. - async with agent.iter( - seed, - message_history=history, - usage_limits=UsageLimits(request_limit=max_turns + 1), - ) as run: - async for node in run: - if interactive and _inject_pending(run): - quit_requested = True - break - if Agent.is_call_tools_node(node): - run_turns += 1 - observer.observe_response( - node.model_response, - run.usage, - log_turn=run_turns, - ) - - async with node.stream(run.ctx) as stream: - async for event in stream: - observer.observe_tool(event, run.usage) - - if observer.finish_result is not None: - logger.info("FINISH called: %s", observer.finish_result) - break - if observer.turns >= max_turns: - logger.info( - "reached max_turns=%d. Stopping.", max_turns - ) - break - elif Agent.is_end_node(node): - if interactive: - logger.info( - "model ended turn without a tool call " - "— awaiting your input." - ) - else: - logger.info( - "model ended turn without a tool call. Stopping." - ) - break - - usage = run.usage - if interactive: - history = run.all_messages() - - # finish, quit, non-interactive, or the cumulative turn budget is - # spent => end the whole session so max_turns is enforced across - # every run, not per run. - if ( - observer.finish_result is not None - or quit_requested - or not interactive - or observer.turns >= max_turns - ): - break - nxt = await _await_next() - if nxt is None: - break - seed = nxt - messages.append({"role": "user", "content": seed}) - except UsageLimitExceeded as e: - logger.info("usage limit reached: %s", e) - except Exception as e: # noqa: BLE001 - surfaced via PlannerResult.error - last_error = _api_error_text(e, no_images=self._no_images) - logger.error("agent run failed: %s", last_error) - - return PlannerResult( - finish_result=observer.finish_result, - messages=messages, - stats=_build_stats(usage, observer.turns, observer.tool_calls), - error=last_error, ) - async def _solve_dashboard( + async def _solve( self, - *, system_prompt: str, user_message: str, toolkit: Toolkit, max_turns: int, - interaction: DashboardInteractionPort, + interaction: DashboardInteractionPort | None, ) -> PlannerResult: - """Drive cancellable PydanticAI runs from complete history checkpoints.""" - agent = self._build_agent(system_prompt, toolkit) - messages: list[dict[str, Any]] = [{"role": "user", "content": user_message}] - - def emit_user(text: str, *, initial: bool = False) -> None: - if not initial: - messages.append({"role": "user", "content": text}) - self._dashboard_events.emit( - TranscriptEvent( - {"type": "initial_prompt"} - if initial - else {"type": "user", "text": text} - ) - ) - - control = DashboardPlannerControl( - interaction=interaction, - cancel_active_and_wait=toolkit.cancel_active_and_wait, - emit_user=emit_user, - emit_initial_user=lambda: emit_user(user_message, initial=True), - defer_message_ack=True, + conversation_id = uuid4().hex + adapter = _HarnessToolkit( + toolkit, dashboard_events=self.dashboard_events, no_images=self.no_images ) - observer = _ApiRunObserver( - dashboard_events=self._dashboard_events, - messages=messages, - max_turns=max_turns, + session = _Session( + adapter, conversation_id, max_turns, interactive=self.interactive ) - session = _ApiDashboardSession( - agent=agent, - control=control, - observer=observer, - max_turns=max_turns, - no_images=self._no_images, + agent = Agent( + self.model, + name="rpent_api", + system_prompt=system_prompt, + output_type=adapter.output_types, + toolsets=[adapter], + # Accepted finish skips sibling tool calls in the same response. + end_strategy="early", + # Finish refusals may span attempts; the request budget still applies. + retries={"output": max_turns}, + model_settings=_build_model_settings(self.model, self.max_tokens), + capabilities=[ + SlidingWindowCompaction( + max_messages=80, max_fraction=0.8, keep_messages=40 + ), + Thinking( + False if self.reasoning_effort == "none" else self.reasoning_effort + ), + session, + ], ) - error: str | None = None + if interaction is not None: + session.control = DashboardPlannerControl( + interaction=interaction, + cancel_active_and_wait=adapter.cancel_active_and_wait, + emit_user=session.emit_user, + emit_initial_user=lambda: session.emit_user(user_message), + defer_message_ack=True, + submit_while_busy=True, + ) try: await asyncio.wait_for( - session.run(user_message), - timeout=self._timeout_s, + session.run(_ConversationAgent(agent), user_message), + timeout=None if self.interactive else self.timeout_s, ) except asyncio.TimeoutError: - error = f"API planner timed out after {self._timeout_s}s" - control.end() + session.error = f"planner timed out after {self.timeout_s:g}s" + except _RequestBudgetExhausted as exc: + session.error = None + logger.info("%s", exc) except Exception as exc: - error = _api_error_text(exc, no_images=self._no_images) - control.end() + session.error = f"{type(exc).__name__}: {exc}" + logger.exception("Harness planner failed") finally: - try: - await control.cancel_active_toolkit() - except Exception as exc: - cleanup_error = ( - f"API toolkit cancellation failed: {type(exc).__name__}: {exc}" - ) - logger.warning(cleanup_error) - error = error or cleanup_error - await session.close() - + session.emit_usage() return PlannerResult( - finish_result=observer.finish_result, - messages=messages, + finish_result=adapter.finish_result, + messages=adapter.messages, stats={ - "backend": "api", - **_build_stats(session.usage, observer.turns, observer.tool_calls), + "total_input_tokens": session.usage.input_tokens, + "total_output_tokens": session.usage.output_tokens, + "total_cached_input_tokens": session.usage.cache_read_tokens, + "turns_used": session.usage.requests, + "tool_calls": adapter.tool_calls, }, - error=error or session.error, + error=session.error, ) - def _build_agent(self, system_prompt: str, toolkit: Toolkit) -> Agent: - """Build an Agent for terminal or Dashboard execution.""" - thinking_effort: str | bool = self._reasoning_effort - if thinking_effort == "none": - thinking_effort = False - return Agent( - self._model, - instructions=system_prompt or None, - tools=_build_tools(toolkit, no_images=self._no_images), - model_settings=_build_model_settings(self._model, self._max_tokens), - capabilities=[ - Thinking(effort=thinking_effort), - ProcessHistory(processor=_prune_history_images), - ], + +def _build_model_settings(model: Model, max_tokens: int) -> ModelSettings: + """Configure sequential calls and preserve Anthropic prompt caching.""" + from pydantic_ai.models.anthropic import AnthropicModel, AnthropicModelSettings + + if isinstance(model, AnthropicModel): + return AnthropicModelSettings( + max_tokens=max_tokens, + parallel_tool_calls=False, + anthropic_cache_instructions=True, + anthropic_cache_tool_definitions=True, + anthropic_cache_messages=True, ) + return ModelSettings(max_tokens=max_tokens, parallel_tool_calls=False) + +class _ConversationAgent(WrapperAgent): + """Preserve completed operations across failed and cancelled runs. + + This checkpoint takes precedence over the CLI's last successful history. + """ + + def __init__(self, agent: Agent) -> None: + super().__init__(agent) + self.history: list[ModelMessage] | None = None + + @asynccontextmanager + async def iter( + self, + user_prompt: str | Sequence[UserContent] | None = None, + *, + message_history: Sequence[ModelMessage] | None = None, + **kwargs: Any, + ) -> AsyncIterator[AgentRun]: + run: AgentRun | None = None + try: + async with self.wrapped.iter( + user_prompt, + message_history=( + self.history if self.history is not None else message_history + ), + **kwargs, + ) as run: + yield run + finally: + # Read after the SDK exits so cancellation cleanup is included. + if run is not None: + self.history = run.all_messages() -@dataclasses.dataclass -class _ApiRunObserver: - """Record model/tool events shared by terminal and Dashboard runs.""" - dashboard_events: DashboardEventSink - messages: list[dict[str, Any]] - max_turns: int - turns: int = 0 - tool_calls: int = 0 - finish_result: dict[str, Any] | None = None - pending_finish: dict[str, Any] | None = None +class _HarnessToolkit(FunctionToolset): + """Expose Toolkit tools and a read-only image artifact reader.""" - def observe_response( + def __init__( self, - response: ModelResponse, - usage: RunUsage, + toolkit: Toolkit, *, - log_turn: int | None = None, + dashboard_events: DashboardEventSink, + no_images: bool, ) -> None: - self.turns += 1 - message = _serialize_response(response) - self.messages.append(message) - _log_response( - response, - usage, - self.turns if log_turn is None else log_turn, - self.max_turns, - ) - for block in message["content"]: - if block["type"] == "text": - payload = {"type": "text", "text": block["text"]} - elif block["type"] == "thinking": - payload = {"type": "thinking", "text": block["thinking"]} + super().__init__(id="rpent") + self.toolkit = toolkit + self.dashboard_events = dashboard_events + self.no_images = no_images + self.finish_result: dict[str, Any] | None = None + self.tool_calls = 0 + self.messages: list[dict[str, Any]] = [] + self.stopping = threading.Event() + self.validators: dict[str, Draft202012Validator] = {} + self.output_types: list[Any] = [str] + self.add_tool(Tool(self.read_image, takes_ctx=True, sequential=True)) + for spec in toolkit.get_tools_spec(): + name = spec["name"] + self.validators[name] = Draft202012Validator(spec["input_schema"]) + if name == "finish": + # Eager annotations bind this toolkit's schema to the output + # function. Keep this module free of postponed annotations. + FinishArguments = StructuredDict( + spec["input_schema"], name="FinishArguments" + ) + + async def finish( + ctx: RunContext, arguments: FinishArguments + ) -> dict[str, Any]: + result = await self.execute("finish", arguments, ctx) + if not result.is_finish: + raise ModelRetry(self._text(result)) + self.finish_result = result.result + return result.result + + self.output_types.append( + ToolOutput( + finish, name="finish", description=spec.get("description") + ) + ) else: - continue - self.dashboard_events.emit(TranscriptEvent(payload)) - self.emit_usage(usage) - - def observe_tool(self, event: Any, usage: RunUsage) -> bool: - completed = False - if isinstance(event, FunctionToolCallEvent): - self.tool_calls += 1 - part = event.part - args = part.args_as_dict() - self.dashboard_events.emit( - TranscriptEvent( - {"type": "tool_call", "tool": part.tool_name, "args": args} + self.add_tool( + Tool.from_schema( + partial(self.call, name), + name=name, + description=spec.get("description"), + json_schema=spec["input_schema"], + takes_ctx=True, + sequential=True, + ) ) - ) - if part.tool_name == "finish": - self.pending_finish = {"_finish": True, **args} - elif isinstance(event, FunctionToolResultEvent): - completed = True - message = _serialize_tool_result(event) - self.messages.append(message) - _log_tool_result(message) - part = event.part - is_error = bool(getattr(part, "is_error", False)) - if self.pending_finish is not None: - if not is_error and "finish refused" not in str(message): - self.finish_result = self.pending_finish - self.pending_finish = None - self.dashboard_events.emit( - TranscriptEvent( - { - "type": "tool_result", - "tool": message.get("name") or "tool_result", - "result": { - "is_error": is_error, - "size": len(message["content"]), - }, + + async def read_image( + self, ctx: RunContext, name: str, step: int = -1 + ) -> ToolReturn: + """Read a saved image by artifact filename and step (-1 selects the latest).""" + self._record_call("read_image", {"name": name, "step": step}, ctx) + result = await asyncio.to_thread(self._read_image, name, step) + self._record_result("read_image", ctx, json.dumps(result.return_value)) + return result + + def _read_image(self, name: str, step: int) -> ToolReturn: + """Resolve an image in the step store, returning artifact errors to the model.""" + try: + state = self.toolkit.state + record = state.get(step) + path = state.artifact_path(name, step=record.step_idx) + if name not in record.artifacts or not path.is_file(): + raise FileNotFoundError( + f"image artifact {name!r} is not available at step {step}" + ) + if path.suffix.lower() not in {".png", ".jpg", ".jpeg"}: + raise ValueError(f"artifact {name!r} is not an image") + metadata = {"artifact": name, "step": record.step_idx} + if self.no_images: + return ToolReturn( + return_value={ + **metadata, + "notice": "Image omitted: --no-images is enabled.", } ) + image = BinaryContent( + data=state.load_bytes(name, step=record.step_idx), + media_type=( + "image/jpeg" + if path.suffix.lower() in {".jpg", ".jpeg"} + else "image/png" + ), ) - self.emit_usage(usage) - return completed + except Exception as exc: + return ToolReturn(return_value={"error": str(exc)}) + return ToolReturn(return_value=metadata, content=[image]) + + async def execute( + self, name: str, arguments: dict[str, Any], ctx: RunContext + ) -> ToolResult: + """Validate and run one physical operation, retaining cancellation ownership.""" + error = next(self.validators[name].iter_errors(arguments), None) + if error is not None: + raise ModelRetry(f"Invalid arguments for {name}: {error.message}") + self._record_call(name, arguments, ctx) + operation = asyncio.create_task( + asyncio.to_thread(self.toolkit.execute_tool, name, arguments) + ) + try: + result = await asyncio.shield(operation) + except asyncio.CancelledError: + self.stopping.set() + await asyncio.to_thread(self.cancel_active_and_wait) + # A cancelled asyncio task does not stop the physical worker. + # Drain it before another run may use the same toolkit. + await operation + raise + self._record_result(name, ctx, self._text(result)) + return result + + def _record_call( + self, name: str, arguments: dict[str, Any], ctx: RunContext + ) -> None: + if self.stopping.is_set(): + # Toolkit draining can finish before Dashboard sends its token. + # Mark this as application cancellation so the SDK raises RunCancelled. + ctx.cancel() + raise asyncio.CancelledError + self.dashboard_events.emit( + TranscriptEvent({"type": "tool_call", "tool": name, "args": arguments}) + ) + self.tool_calls += 1 - def emit_usage(self, usage: RunUsage) -> None: + def _record_result(self, name: str, ctx: RunContext, text: str) -> None: + self.messages.append( + { + "role": "tool", + "name": name, + "tool_call_id": ctx.tool_call_id, + "content": text, + } + ) self.dashboard_events.emit( - UsageEvent( - inp=int(usage.input_tokens or 0), - out=int(usage.output_tokens or 0), - tool_calls=self.tool_calls, - ) + TranscriptEvent({"type": "tool_result", "tool": name, "result": text}) + ) + + async def call(self, name: str, ctx: RunContext, /, **arguments: Any) -> ToolReturn: + """Return native multimodal content for an ordinary tool.""" + result = await self.execute(name, arguments, ctx) + content: list[str | BinaryContent] = [] + for block in result.content_blocks: + if block["type"] == "text": + content.append(block["text"]) + elif block["type"] == "image" and self.no_images: + content.append("[Image omitted: --no-images is enabled.]") + elif block["type"] == "image": + source = block["source"] + content.append( + BinaryContent( + data=base64.b64decode(source["data"]), + media_type=source["media_type"], + ) + ) + return ToolReturn(return_value=content) + + def cancel_active_and_wait(self) -> None: + """Stop admitting tools before requesting physical cancellation.""" + self.stopping.set() + self.toolkit.cancel_active_and_wait() + + @staticmethod + def _text(result: ToolResult) -> str: + return "\n".join( + block["text"] for block in result.content_blocks if block["type"] == "text" ) -class _ApiDashboardSession: - """Own serial, independent PydanticAI runs for one Dashboard TaskRun.""" +class _Session(AbstractCapability): + """Track native runs and bridge Dashboard input and events. + + This capability observes public lifecycle boundaries; it never replaces a + graph node or mutates message history. The conversation driver starts a new + native run only for a follow-up after completion or interruption. + """ def __init__( self, - *, - agent: Agent, - control: DashboardPlannerControl, - observer: _ApiRunObserver, + toolkit: _HarnessToolkit, + conversation_id: str, max_turns: int, - no_images: bool, + *, + interactive: bool = False, ) -> None: - self._agent = agent - self._control = control - self._max_turns = max_turns - self._no_images = no_images - self._observer = observer - self._history: list[ModelMessage] = [] + self.toolkit = toolkit + self.conversation_id = conversation_id + self.limits = UsageLimits(request_limit=max_turns) self.usage = RunUsage() - self._pending_prompts: deque[tuple[str | None, str]] = deque() - self._run_task: asyncio.Task[Any] | None = None - self._active_prompt = False - self._closing = False + self.interactive = interactive self.error: str | None = None - - async def run(self, prompt: str) -> None: - await self.submit(prompt) - await self._control.start() - await self._control.run(self) - - async def submit(self, text: str) -> int: - """Queue Dashboard input as a new independent API run.""" - return self._queue_prompt(text) - - async def submit_dashboard_message(self, message: DashboardMessage) -> int: - """Queue Dashboard input and defer acknowledgement until it starts.""" - return self._queue_prompt(message.text, message_id=message.message_id) - - def _queue_prompt(self, text: str, *, message_id: str | None = None) -> int: - if self._closing: - raise RuntimeError("API conversation is closed") - if self.error is not None: - raise RuntimeError(self.error) - self._pending_prompts.append((message_id, text)) - if self._run_task is None: - self._run_task = asyncio.create_task(self._run_pending_prompts()) - return 1 - - async def interrupt(self) -> int: - run_task = self._run_task - interrupted = int(self._active_prompt) + len(self._pending_prompts) - discarded_message_ids = tuple( - message_id - for message_id, _ in self._pending_prompts - if message_id is not None - ) - self._pending_prompts.clear() - for message_id in discarded_message_ids: - self._control.message_discarded(message_id) - if run_task is None or run_task.done(): - return interrupted - run_task.cancel() - try: - with contextlib.suppress(asyncio.CancelledError): - await run_task - finally: - if self._run_task is run_task: - self._run_task = None - return interrupted - - async def close(self) -> None: - if self._closing: - return - self._closing = True - await self.interrupt() - - async def _run_pending_prompts(self) -> None: - task = asyncio.current_task() - try: - while self._pending_prompts and not self._closing: - message_id, seed = self._pending_prompts.popleft() - self._active_prompt = True - try: - if message_id is not None: - self._control.message_started(message_id, seed) - if not await self._run_agent(seed): - self._pending_prompts.clear() - return - await self._control.complete(self) - finally: - self._active_prompt = False - finally: - if self._run_task is task: - self._run_task = None - - async def _run_agent(self, seed: str) -> bool: - run_completed = False - run: Any | None = None - node: Any | None = None + self.control: DashboardPlannerControl | None = None + self.inbox: asyncio.Queue[DashboardMessage] = asyncio.Queue() + self.enqueued: dict[str, DashboardMessage] = {} + self.run_done = asyncio.Event() + self.run_done.set() + self.cancellation: CancellationToken | None = None + self.completions = 0 + + async def wrap_run( + self, ctx: RunContext, *, handler: WrapRunHandler + ) -> AgentRunResult: + if self.toolkit.finish_result is not None: + raise UserError("The task has finished. Use /exit to close the session.") + self.toolkit.stopping.clear() + if self.interactive and ctx.prompt: + self.emit_user(ctx.prompt) try: - async with self._agent.iter( - seed, - message_history=list(self._history), - usage=self.usage, - usage_limits=UsageLimits(request_limit=self._max_turns + 1), - ) as run: - node = run.next_node - while not Agent.is_end_node(node): - if Agent.is_call_tools_node(node): - await self._process_tool_node(run, node) - if ( - self._observer.finish_result is not None - or self._observer.turns >= self._max_turns - ): - self._control.end() - return False - if self._pending_prompts: - # Dashboard input accepted at this tool boundary starts - # a fresh run from the checkpoint captured below. - node = await run.next(node) - break - node = await run.next(node) - - run_completed = True + result = await handler() + self.error = None + return result + except _RequestBudgetExhausted: + # The native CLI catches run exceptions before _solve can see them. + self.error = None + raise except Exception as exc: - self.error = _api_error_text(exc, no_images=self._no_images) - if not self._closing: - self._control.end() + self.error = f"{type(exc).__name__}: {exc}" + raise finally: - # Preserve interrupted tool results for PydanticAI to repair on the - # next run. Older supported releases can leave a bare tool-call - # response when cancellation wins before any tool returns; remove - # only that unusable frontier. - if run is not None: - history = list(run.all_messages()) - if ( - history - and isinstance(history[-1], ModelResponse) - and history[-1].tool_calls - ): - if request := getattr(node, "request", None): - history.append(request) - else: - history.pop() - self._history = history - return run_completed and not self._closing - - async def _process_tool_node(self, run: Any, node: Any) -> None: - self._observer.observe_response(node.model_response, run.usage) - - async with node.stream(run.ctx) as stream: - async for event in stream: - tool_completed = self._observer.observe_tool(event, run.usage) - if ( - tool_completed - and self._observer.finish_result is None - and self._observer.turns < self._max_turns - ): - await self._control.tool_completed(self) - - -def _build_model_settings(model: Model, max_tokens: int) -> ModelSettings: - """Build model settings, enabling prompt caching for Anthropic models.""" - from pydantic_ai.models.anthropic import AnthropicModel, AnthropicModelSettings - - if isinstance(model, AnthropicModel): - return AnthropicModelSettings( - max_tokens=max_tokens, - anthropic_cache_instructions=True, - anthropic_cache_tool_definitions=True, - anthropic_cache_messages=True, + # clai also keeps its own per-run usage for /usage. Observe it without + # replacing the SDK's counters, including failed/interrupted runs. + self.usage.incr(ctx.usage) + self.emit_usage() + + async def before_model_request( + self, ctx: RunContext, request_context: ModelRequestContext + ) -> ModelRequestContext: + try: + self.limits.check_before_request(self.usage + ctx.usage) + except UsageLimitExceeded: + # This limit only caps requests. Preserve other SDK/provider failures. + message = f"Request budget of {self.limits.request_limit} reached." + if self.interactive: + message += " Use /exit to close the session." + raise _RequestBudgetExhausted(message) from None + return request_context + + async def before_node_run(self, ctx: RunContext, *, node: AgentNode) -> AgentNode: + if self.toolkit.stopping.is_set(): + # Dashboard drains the physical operation before awaiting the SDK + # interrupt. Do not deliver queued input during that drain. + ctx.cancel() + return node + if not isinstance(node, ModelRequestNode): + return node + while not self.inbox.empty(): + submission = self.inbox.get_nowait() + enqueue_id = ctx.enqueue(submission.text) + if enqueue_id is not None: + self.enqueued[enqueue_id] = submission + self.completions += 1 + return node + + async def after_node_run( + self, ctx: RunContext, *, node: AgentNode, result: NodeResult + ) -> NodeResult: + if isinstance(node, ModelRequestNode) and isinstance(result, CallToolsNode): + # Retain completed turns even when request history is later compacted. + content: list[dict[str, Any]] = [] + for part in result.model_response.parts: + if isinstance(part, TextPart) and part.content: + content.append({"type": "text", "text": part.content}) + elif isinstance(part, ThinkingPart) and part.content: + content.append({"type": "thinking", "thinking": part.content}) + elif isinstance(part, ToolCallPart): + content.append( + { + "type": "tool_use", + "id": part.tool_call_id, + "name": part.tool_name, + "input": part.args_as_dict(), + } + ) + self.toolkit.messages.append({"role": "assistant", "content": content}) + if isinstance(node, (ModelRequestNode, CallToolsNode)): + self.emit_usage(self.usage + ctx.usage) + return result + + @on_event(EnqueuedMessagesEvent) + async def acknowledge_enqueued( + self, ctx: RunContext, event: EnqueuedMessagesEvent + ) -> None: + submission = self.enqueued.pop(event.enqueue_id, None) + if submission is not None and self.control is not None: + self.control.message_started(submission.message_id, submission.text) + + @on_event(PartEndEvent) + async def emit_part(self, ctx: RunContext, event: PartEndEvent) -> None: + part = event.part + if not isinstance(part, (TextPart, ThinkingPart)) or not part.content: + return + kind = "thinking" if isinstance(part, ThinkingPart) else "text" + self.toolkit.dashboard_events.emit( + TranscriptEvent({"type": kind, "text": part.content}) ) - return ModelSettings(max_tokens=max_tokens) - - -def _prune_history_images(messages: list[ModelMessage]) -> list[ModelMessage]: - """Drop old camera images so the resent request body stays bounded.""" - # Every image in history, oldest -> newest: (msg_idx, part_idx, item_idx, nbytes). - located: list[tuple[int, int, int, int]] = [] - for mi, message in enumerate(messages): - for pi, part in enumerate(getattr(message, "parts", ()) or ()): - if not isinstance(part, UserPromptPart) or not isinstance( - part.content, list - ): - continue - for ii, item in enumerate(part.content): - if isinstance(item, BinaryContent) and item.media_type.startswith( - "image/" - ): - located.append((mi, pi, ii, len(item.data))) - - if not located: - return messages - - # Walk newest -> oldest, keeping images while under the byte budget. - keep: set[tuple[int, int, int]] = set() - total = 0 - for rank, (mi, pi, ii, nbytes) in enumerate(reversed(located)): - if rank < _MIN_RECENT_IMAGES or total + nbytes <= _MAX_HISTORY_IMAGE_BYTES: - keep.add((mi, pi, ii)) - total += nbytes - - if len(keep) == len(located): - return messages - - drop_items_by_part: dict[tuple[int, int], set[int]] = {} - for mi, pi, ii, _ in located: - if (mi, pi, ii) not in keep: - drop_items_by_part.setdefault((mi, pi), set()).add(ii) - - new_messages = list(messages) - for (mi, pi), drop_items in drop_items_by_part.items(): - message = new_messages[mi] - part = message.parts[pi] - new_content = [ - "[earlier camera image omitted to bound request size]" - if ci in drop_items - else item - for ci, item in enumerate(part.content) - ] - new_parts = list(message.parts) - new_parts[pi] = dataclasses.replace(part, content=new_content) - new_messages[mi] = dataclasses.replace(message, parts=new_parts) - - return new_messages - - -def _is_image_rejection(e: Exception) -> bool: - """True when the provider returned a 4xx complaining about image input. - - Matches errors like ``400 {'code': 10007, 'msg': "Bad Request: [message - type 'image_url' is not supported]"}`` from OpenAI-compatible endpoints - serving text-only models. - """ - if not isinstance(e, ModelHTTPError): - return False - if not 400 <= e.status_code < 500: - return False - return "image" in str(e).lower() - - -def _api_error_text(error: Exception, *, no_images: bool) -> str: - text = f"{type(error).__name__}: {error}" - if not no_images and _is_image_rejection(error): - text += ( - "\n\nThe model rejected image input — it is likely a text-only " - "model (no vision support). Re-run with --no-images: RPent will " - "then keep every visual observation as a file-path text notice " - "instead of sending image bytes." + if not self.interactive: + logger.info("[%s] %s", kind, part.content) + + def emit_user(self, text: str) -> None: + self.toolkit.messages.append({"role": "user", "content": text}) + self.toolkit.dashboard_events.emit( + TranscriptEvent({"type": "user", "text": text}) ) - return text - - -def _build_tools(toolkit: Toolkit, *, no_images: bool = False) -> list[Tool]: - """Build the API-only image reader plus pydantic-ai toolkit wrappers.""" - image_reader = _make_image_reader(toolkit.state, no_images=no_images) - # sequential=True serializes a turn's tool calls so the toolkit's - # single-operation lock never rejects an overlapping call. - tools: list[Tool] = [Tool(image_reader, name="read_image", sequential=True)] - for spec in toolkit.get_tools_spec(): - name = spec["name"] - tools.append( - Tool.from_schema( - function=_make_tool_function(toolkit, name, no_images=no_images), - name=name, - description=spec.get("description", ""), - json_schema=spec.get("input_schema") - or {"type": "object", "properties": {}}, - takes_ctx=False, - sequential=True, + + def emit_usage(self, usage: RunUsage | None = None) -> None: + if usage is None: + usage = self.usage + self.toolkit.dashboard_events.emit( + UsageEvent( + inp=usage.input_tokens, + out=usage.output_tokens, + tool_calls=self.toolkit.tool_calls, ) ) - return tools - - -def _make_image_reader( - state: EnvState, - *, - no_images: bool, -) -> Callable[[str, int], ToolReturn | dict[str, str] | str]: - if no_images: - def read_image_tool(name: str, step: int = -1) -> str: - return read_image_text_only(name, step, state=state) - - read_image_tool.__name__ = "read_image" - read_image_tool.__doc__ = read_image_text_only.__doc__ - return read_image_tool - - def read_image_tool(name: str, step: int = -1) -> ToolReturn | dict[str, str]: - return read_image(name, step, state=state) - - read_image_tool.__name__ = "read_image" - read_image_tool.__doc__ = read_image.__doc__ - return read_image_tool - - -def read_image( - name: str, step: int = -1, *, state: EnvState -) -> ToolReturn | dict[str, str]: - """Read a step-scoped image artifact as visual input. + async def submit(self, message: DashboardMessage) -> int: + self.inbox.put_nowait(message) + return 1 - Artifact failures are returned as structured tool errors so a bad - model-supplied name or step does not abort the agent run. - """ - try: - resolved_step, path = _resolve_image_artifact(state, name, step) - content = BinaryContent( - data=state.load_bytes(name, step=resolved_step), - media_type=_image_media_type(path), - ) - except Exception as e: - return {"error": str(e)} - return ToolReturn( - return_value={"artifact": name, "step": resolved_step}, - content=[content], - ) - - -def read_image_text_only( - name: str, step: int = -1, *, state: EnvState -) -> str | dict[str, str]: - """Acknowledge an image artifact without sending bytes to the model.""" - try: - resolved_step, _ = _resolve_image_artifact(state, name, step) - except Exception as e: - return {"error": str(e)} - return ( - f"Image artifact {name!r} exists at step {resolved_step}, but image " - "input is disabled (--no-images, text-only model). Reason from textual " - "state instead: view_env_state, back_project, and numeric tool results." - ) - - -def _resolve_image_artifact( - state: EnvState, - name: str, - step: int, -) -> tuple[int, Path]: - record = state.get(step) - path = state.artifact_path(name, step=record.step_idx) - if name not in record.artifacts or not path.is_file(): - raise FileNotFoundError( - f"image artifact {name!r} is not available at step {step}" + async def interrupt(self) -> int: + """Discard unstarted submissions, then cancel and drain the native run.""" + discarded = self.inbox.qsize() + while not self.inbox.empty(): + submission = self.inbox.get_nowait() + if self.control is not None: + self.control.message_discarded(submission.message_id) + if self.cancellation is not None: + self.cancellation.cancel() + await self.run_done.wait() + # The run accounts for its own completion and SDK-enqueued messages. + # Only submissions that never entered it need to be subtracted here. + return discarded + + async def run_cli(self, agent: _ConversationAgent, prompt: str) -> None: + """Run the preset task, then hand its history to the native CLI.""" + from rich.console import Console + + console = Console() + console.print("Task context:\n" + prompt, markup=False, highlight=False) + console.print( + "Starting the task. You can enter follow-up instructions after it " + "responds, or use /exit to close.", + markup=False, ) - if path.suffix.lower() not in {".png", ".jpg", ".jpeg"}: - raise ValueError(f"artifact {name!r} is not an image") - return record.step_idx, path - - -def _image_media_type(path: Path) -> str: - return "image/jpeg" if path.suffix.lower() in {".jpg", ".jpeg"} else "image/png" - - -def _make_tool_function(toolkit: Toolkit, name: str, *, no_images: bool = False): - """Return a callable that dispatches one tool call to the toolkit.""" - - def _call(**kwargs: Any) -> Any: - result = toolkit.execute_tool(name, kwargs) - text, images = _content_blocks_to_pydantic(result.content_blocks) - if images and not no_images: - return ToolReturn(return_value=text, content=images) - return text - - _call.__name__ = name - return _call - - -def _content_blocks_to_pydantic( - blocks: list[dict[str, Any]], -) -> tuple[str, list[BinaryContent]]: - """Split Anthropic-shaped content blocks into text and image content.""" - text_parts: list[str] = [] - images: list[BinaryContent] = [] - for block in blocks: - block_type = block.get("type") - if block_type == "text": - text_parts.append(block.get("text", "")) - elif block_type == "image": - source = block.get("source") or {} - data = source.get("data") - if source.get("type") == "base64" and data: - images.append( - BinaryContent( - data=base64.b64decode(data), - media_type=source.get("media_type", "image/png"), - ) - ) - text = "\n\n".join(part for part in text_parts if part) or "{}" - return text, images - - -def _serialize_response(response: ModelResponse) -> dict[str, Any]: - """Render one assistant turn as a serialisable transcript message.""" - content: list[dict[str, Any]] = [] - for part in response.parts: - if isinstance(part, TextPart): - if part.content: - content.append({"type": "text", "text": part.content}) - elif isinstance(part, ThinkingPart): - if part.content: - content.append({"type": "thinking", "thinking": part.content}) - elif isinstance(part, ToolCallPart): - content.append( - { - "type": "tool_use", - "id": part.tool_call_id, - "name": part.tool_name, - "input": part.args_as_dict(), - } + with console.status("Running task…"): + result = await agent.run( + prompt, + conversation_id=self.conversation_id, + usage_limits=self.limits, ) - return {"role": "assistant", "content": content} - - -def _serialize_tool_result(event: FunctionToolResultEvent) -> dict[str, Any]: - """Render one tool result as a serialisable transcript message (no images).""" - part = event.part - content = getattr(part, "content", None) - if not isinstance(content, str): - content = json.dumps(content, default=str) - return { - "role": "tool", - "name": getattr(part, "tool_name", None), - "tool_call_id": getattr(part, "tool_call_id", None), - "content": content, - } - - -def _build_stats( - usage: RunUsage | None, turns: int, n_tool_calls: int -) -> dict[str, Any]: - """Assemble the run stats dict from accumulated usage and counters.""" - stats: dict[str, Any] = {"turns_used": turns, "tool_calls": n_tool_calls} - if usage is not None: - stats.update( - { - "total_input_tokens": int(usage.input_tokens or 0), - "total_output_tokens": int(usage.output_tokens or 0), - "cache_read_tokens": int(usage.cache_read_tokens or 0), - "cache_write_tokens": int(usage.cache_write_tokens or 0), - "requests": int(usage.requests or 0), - } + console.print(result.output, markup=False, highlight=False) + await agent.to_cli( + prog_name="rpent", + message_history=result.all_messages(), + usage_limits=self.limits, ) - return stats - - -def _log_response( - response: ModelResponse, usage: RunUsage, turn: int, max_turns: int -) -> None: - """Log model text, thinking, tool calls, and cumulative usage for a turn.""" - logger.info("=== turn %d/%d ===", turn, max_turns) - for part in response.parts: - if isinstance(part, TextPart): - text = (part.content or "").strip() - if text: - logger.info("[model] %s", text) - elif isinstance(part, ThinkingPart): - text = (part.content or "").strip() - if text: - logger.info("[think] %s", _clip(text, _TEXT_LOG_LIMIT)) - elif isinstance(part, ToolCallPart): - args = json.dumps(part.args_as_dict(), default=str) - logger.info("[tool>] %s(%s)", part.tool_name, _clip(args, _ARGS_LOG_LIMIT)) - logger.info( - "[usage] in=%s out=%s cache_read=%s cache_write=%s requests=%s", - usage.input_tokens, - usage.output_tokens, - usage.cache_read_tokens, - usage.cache_write_tokens, - usage.requests, - ) - - -def _log_tool_result(message: dict[str, Any]) -> None: - """Log a one-line summary of a tool result.""" - content = " ".join((message.get("content") or "").split()) - logger.info("[tool<] %s: %s", message.get("name"), _clip(content, _TOOL_LOG_LIMIT)) - - -def _clip(text: str, limit: int) -> str: - """Truncate ``text`` to ``limit`` characters with an overflow marker.""" - if len(text) <= limit: - return text - return text[:limit] + "...(+%d)" % (len(text) - limit) + + async def run(self, agent: _ConversationAgent, prompt: str) -> None: + tasks: list[asyncio.Task] = [] + try: + if self.interactive: + await self.run_cli(agent, prompt) + return + if self.control is None: + self.emit_user(prompt) + await self._run_conversation(agent, prompt) + return + await self.control.start() + tasks = [ + asyncio.create_task(self.control.run(self)), + asyncio.create_task(self._run_conversation(agent, prompt)), + ] + done, _ = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) + for task in done: + await task + finally: + if self.control is not None: + self.control.end() + # Stop new physical work before cancelling the SDK task. Toolkit + # cancellation is cooperative, and must finish before solve returns. + self.toolkit.stopping.set() + if self.cancellation is not None: + self.cancellation.cancel() + await asyncio.to_thread(self.toolkit.cancel_active_and_wait) + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + async def _run_conversation(self, agent: _ConversationAgent, prompt: str) -> None: + """Run the SDK, then wait for Dashboard follow-ups in the same conversation.""" + while True: + self.cancellation = CancellationToken() + self.run_done.clear() + self.completions = 1 + try: + await agent.run( + prompt, + usage_limits=self.limits, + conversation_id=self.conversation_id, + cancellation_token=self.cancellation, + ) + except RunCancelled: + self.error = None + finally: + self.cancellation = None + self.run_done.set() + if self.control is not None: + for pending in self.enqueued.values(): + self.control.message_discarded(pending.message_id) + self.enqueued.clear() + if self.toolkit.finish_result is not None or self.control is None: + return + if self.usage.requests >= self.limits.request_limit: + return + for _ in range(self.completions): + await self.control.complete(self) + submission = await self.inbox.get() + self.control.message_started(submission.message_id, submission.text) + prompt = submission.text diff --git a/rpent/planner/base.py b/rpent/planner/base.py index 4b6720585..972c50b2b 100644 --- a/rpent/planner/base.py +++ b/rpent/planner/base.py @@ -12,14 +12,15 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Shared protocol for high-level reasoning backends.""" +"""Shared abstract base class for high-level reasoning backends.""" from __future__ import annotations import os import queue +from abc import ABC, abstractmethod from pathlib import Path -from typing import TYPE_CHECKING, Protocol +from typing import TYPE_CHECKING from rpent.dashboard.events import DashboardEventSink from rpent.dashboard.interaction import DashboardInteractionPort @@ -79,7 +80,7 @@ def __init__( self.error = error # str | None — set when the planner raises -class Planner(Protocol): +class Planner(ABC): """A planner selects or supplies the actions used to solve a task. It is given one system prompt, one initial user message, and a set of @@ -87,6 +88,7 @@ class Planner(Protocol): finished or the turn budget is exhausted. """ + @abstractmethod def solve( self, *, @@ -114,7 +116,7 @@ def solve( ``PlannerResult`` with finish status, conversation transcript, token-usage stats, and optional error string. """ - ... + raise NotImplementedError # --------------------------------------------------------------------------- @@ -123,7 +125,7 @@ def solve( def build_api_model(model: str | None, base_url: str | None = None) -> "Model": - """Resolve the pydantic-ai model used by the ``api`` planner. + """Resolve the pydantic-ai model used by the API planner. This is the single provider-resolution path: both :func:`build_planner` and the connectivity check in :mod:`rpent.planner.check` call it, so a @@ -185,7 +187,8 @@ def build_planner( claude_code_max_budget_usd: float | None = None, dashboard_events: DashboardEventSink, no_images: bool = False, -): + interactive: bool = False, +) -> Planner: """Build a planner for the given backend, resolving credentials from env vars.""" # Imports are deferred to avoid a circular import: api_loop / claude_code / # codex all import from this module (PlannerResult). @@ -204,6 +207,7 @@ def build_planner( dashboard_events=dashboard_events, no_images=no_images, timeout_s=api_timeout_s, + interactive=interactive, ) if planner_type == "claude_code": from rpent.planner.claude_code import ClaudeCodePlanner diff --git a/rpent/planner/claude_code.py b/rpent/planner/claude_code.py index 1c2676085..7a32be6f1 100644 --- a/rpent/planner/claude_code.py +++ b/rpent/planner/claude_code.py @@ -41,10 +41,11 @@ TranscriptEvent, UsageEvent, ) -from rpent.dashboard.interaction import DashboardInteractionPort +from rpent.dashboard.interaction import DashboardInteractionPort, DashboardMessage from rpent.dashboard.planner_control import DashboardPlannerControl from rpent.planner.base import ( REASONING_EFFORTS, + Planner, PlannerResult, add_mcp_prefix, strip_mcp_prefix, @@ -62,7 +63,7 @@ # --------------------------------------------------------------------------- -class ClaudeCodePlanner: +class ClaudeCodePlanner(Planner): """Planner backed by the Claude Agent SDK.""" def __init__( @@ -433,9 +434,9 @@ async def query(self, text: str) -> None: raise RuntimeError("Claude session is not connected") await self._client.query(text) - async def submit(self, text: str) -> int: + async def submit(self, message: DashboardMessage) -> int: """Submit Dashboard input as a new Claude query.""" - await self.query(text) + await self.query(message.text) return 1 async def interrupt(self) -> int: diff --git a/rpent/planner/codex.py b/rpent/planner/codex.py index 14eb87ac7..0f5cf2b6e 100644 --- a/rpent/planner/codex.py +++ b/rpent/planner/codex.py @@ -46,9 +46,14 @@ TranscriptEvent, UsageEvent, ) -from rpent.dashboard.interaction import DashboardInteractionPort +from rpent.dashboard.interaction import DashboardInteractionPort, DashboardMessage from rpent.dashboard.planner_control import DashboardPlannerControl -from rpent.planner.base import REASONING_EFFORTS, PlannerResult, strip_mcp_prefix +from rpent.planner.base import ( + REASONING_EFFORTS, + Planner, + PlannerResult, + strip_mcp_prefix, +) from rpent.planner.utils.http_mcp_server import HttpMcpServer from rpent.tools.toolkit import Toolkit from rpent.utils.config import get_repo_root @@ -81,7 +86,7 @@ def _codex_environment() -> dict[str, str]: # --------------------------------------------------------------------------- -class CodexPlanner: +class CodexPlanner(Planner): """Planner backed by the OpenAI Codex Python SDK.""" def __init__( @@ -494,11 +499,15 @@ def __init__( async def run(self, prompt: str) -> None: self._codex = openai_codex.AsyncCodex(self._config) self._thread = await self._codex.thread_start(**self._thread_options) - await self.submit(prompt) + await self._submit_text(prompt) await self._control.start() await self._control.run(self) - async def submit(self, text: str) -> int: + async def submit(self, message: DashboardMessage) -> int: + """Steer the active turn or start a turn in the same conversation.""" + return await self._submit_text(message.text) + + async def _submit_text(self, text: str) -> int: if self._closing or self._thread is None: raise RuntimeError("Codex conversation is closed") if self._turn is not None: diff --git a/rpent/planner/flash.py b/rpent/planner/flash.py index dc03ceb8a..c15e6e52c 100644 --- a/rpent/planner/flash.py +++ b/rpent/planner/flash.py @@ -29,7 +29,7 @@ import time -from rpent.planner.base import PlannerResult +from rpent.planner.base import Planner, PlannerResult from rpent.robots.base import get_robot_spec from rpent.tools.toolkit import Toolkit from rpent.utils.logging import get_logger @@ -37,7 +37,7 @@ logger = get_logger("flash") -class FlashPlanner: +class FlashPlanner(Planner): """Replay one recorded program against the toolkit the runtime handed over.""" def __init__( @@ -64,7 +64,7 @@ def solve( A program decides the actions before the episode begins, so there is no conversation to hold and no turn to spend. The arguments are accepted to - satisfy the planner protocol. + satisfy the planner interface. """ run_flash = get_robot_spec(self._robot_name).run_flash if run_flash is None: diff --git a/tests/unit_tests/robots/dual_franka/test_exploration.py b/tests/unit_tests/robots/dual_franka/test_exploration.py index 14247334a..97e711102 100644 --- a/tests/unit_tests/robots/dual_franka/test_exploration.py +++ b/tests/unit_tests/robots/dual_franka/test_exploration.py @@ -281,8 +281,9 @@ def test_successful_memory_pair_uses_existing_merge_and_index(setup, tmp_path): assert (t.memory.root / "task-specific/dual_franka_t0_recipe.jsonl").exists() +@pytest.mark.parametrize("interactive", [False, True]) def test_cli_two_sessions_operator_feedback_and_memory_pipeline( - tmp_path, monkeypatch, dual_franka_robot_config + tmp_path, monkeypatch, interactive, dual_franka_robot_config ): import sys from dataclasses import replace @@ -297,7 +298,7 @@ def test_cli_two_sessions_operator_feedback_and_memory_pipeline( class Operator: def __init__(self, **kwargs): - pass + assert kwargs["interactive"] is False def __call__(self, prompt, cancelled): cancelled() @@ -384,6 +385,7 @@ def init_runtime(*args): "--robot-config", str(dual_franka_robot_config), "--auto-merge-memory", + *(["--interactive"] if interactive else []), ], ) assert cli.main() == 0 @@ -619,6 +621,8 @@ def solve(self, *, toolkit, input_queue, **kwargs): "rpent", "--robot", robot_name, + "--planner", + "codex", "--explore", "--interactive", "--output-dir", diff --git a/tests/unit_tests/rpent/cli/test_main_contracts.py b/tests/unit_tests/rpent/cli/test_main_contracts.py index 425073f87..a2b37885f 100644 --- a/tests/unit_tests/rpent/cli/test_main_contracts.py +++ b/tests/unit_tests/rpent/cli/test_main_contracts.py @@ -244,6 +244,7 @@ def parse_config(args) -> None: ), ) monkeypatch.setattr(sys, "argv", ["rpent", *argv]) + monkeypatch.setattr(sys.stdin, "isatty", lambda: False) with pytest.raises(SystemExit) as exc_info: cli.main() @@ -389,13 +390,18 @@ def test_handoff_message_lists_prior_attempts_deterministically(tmp_path: Path) assert "memory inbox under wip/" in message +@pytest.mark.parametrize("interactive", [False, True]) +@pytest.mark.parametrize("budget_exhausted", [False, True]) def test_full_cli_exploration_finalizes_memory_without_starting_gpu_runtime( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, + interactive: bool, + budget_exhausted: bool, ) -> None: cli = _cli_module() from rpent.planner.base import PlannerResult from rpent.robots import PromptBundle, RobotSpec, RunConfig + from rpent.tools import common from rpent.tools.toolkit import ToolResult calls: dict[str, Any] = {} @@ -419,6 +425,8 @@ def __init__(self) -> None: def execute_tool(self, name: str, args: dict[str, Any]) -> ToolResult: self.calls.append((name, args)) + if budget_exhausted: + return ToolResult(name, {"error": "finish refused"}) return ToolResult( name, { @@ -431,8 +439,14 @@ def execute_tool(self, name: str, args: dict[str, Any]) -> ToolResult: def close(self) -> None: self.closed = True + def get_tools_spec(self) -> list[dict[str, Any]]: + return [common.TOOLS_SPEC[-1]] + + def cancel_active_and_wait(self) -> None: + pass + def solved(self) -> bool: - return True + return not budget_exhausted def write_recipe(self, recipe_tag: str) -> str: calls["write_recipe"] = recipe_tag @@ -505,8 +519,27 @@ def init_runtime(*args: Any) -> tuple[list[FakeDaemon], dict[str, str]]: supports_exploration=True, ) - def build_planner(*args: Any, **kwargs: Any) -> ScriptedPlanner: + def build_planner(*args: Any, **kwargs: Any): calls["build_planner"] = (args, kwargs) + calls["planner_count"] = calls.get("planner_count", 0) + 1 + if budget_exhausted: + from pydantic_ai.models.function import DeltaToolCall, FunctionModel + + from rpent.planner.api_loop import ApiAgentLoop + + async def stream(messages, info): + yield { + 0: DeltaToolCall( + name="finish", + json_args=json.dumps({"status": "success", "summary": "done"}), + ) + } + + return ApiAgentLoop( + model=FunctionModel(stream_function=stream), + dashboard_events=kwargs["dashboard_events"], + interactive=interactive, + ) return planner def get_toolkit(*args: Any, **kwargs: Any) -> FakeToolkit: @@ -521,6 +554,12 @@ def reject_memory_sync(*args: Any, **kwargs: Any) -> None: monkeypatch.setattr(cli, "get_robot_spec", lambda name: robot_spec) monkeypatch.setattr(cli, "build_planner", build_planner) monkeypatch.setattr(cli, "get_toolkit", get_toolkit) + monkeypatch.setattr(sys.stdin, "isatty", lambda: True) + monkeypatch.setattr( + cli, + "start_interactive_reader", + lambda *args, **kwargs: pytest.fail("API mode must use the native CLI"), + ) monkeypatch.setattr("rpent.memory.MemoryManager.sync", reject_memory_sync) monkeypatch.setattr( sys, @@ -539,10 +578,30 @@ def reject_memory_sync(*args: Any, **kwargs: Any) -> None: str(tmp_path), "--max-turns", "4", + "--explore-sessions", + "2", + *(["--interactive"] if interactive else []), ], ) assert cli.main() == 0 + assert calls["build_planner"][1]["interactive"] is interactive + assert toolkit.closed is True + assert daemon.stopped is True + if budget_exhausted: + assert calls["planner_count"] == 2 + assert len(toolkit.calls) == 8 + assert "write_recipe" not in calls + assert calls["merge_memory"] == { + "cell_tag": "libero_s0", + "run_state_dir": tmp_path, + "solved": False, + } + transcript = json.loads((tmp_path / "transcript_libero_s0.json").read_text()) + assert transcript["finish"] is None + assert transcript["stats"]["turns_used"] == 4 + assert len([m for m in transcript["messages"] if m["role"] == "tool"]) == 8 + return assert calls["solve"] == { "system_prompt": "simulated system prompt\n", @@ -554,8 +613,6 @@ def reject_memory_sync(*args: Any, **kwargs: Any) -> None: assert toolkit.calls == [ ("finish", {"status": "success", "summary": "simulated task complete"}) ] - assert toolkit.closed is True - assert daemon.stopped is True assert calls["get_toolkit"][1]["runtime_kwargs"] == {"runtime": "simulated"} assert calls["get_toolkit"][1]["mode"] == "exploration" assert calls["get_toolkit"][1]["attempts_per_session"] == 2 diff --git a/tests/unit_tests/rpent/dashboard/test_planner_control_contracts.py b/tests/unit_tests/rpent/dashboard/test_planner_control_contracts.py index 80ac0056a..ac24f3655 100644 --- a/tests/unit_tests/rpent/dashboard/test_planner_control_contracts.py +++ b/tests/unit_tests/rpent/dashboard/test_planner_control_contracts.py @@ -183,21 +183,13 @@ def __init__( self.submit_error = submit_error self.interrupt_error = interrupt_error self.interrupt_completions = interrupt_completions - self.submissions: list[str] = [] - self.dashboard_submissions: list[DashboardMessage] = [] + self.submissions: list[DashboardMessage] = [] - async def submit(self, text: str) -> int: - self.events.append(f"submit:{text}") + async def submit(self, message: DashboardMessage) -> int: + self.events.append(f"submit:{message.text}") if self.submit_error is not None: raise self.submit_error - self.submissions.append(text) - return 1 - - async def submit_dashboard_message(self, message: DashboardMessage) -> int: - self.events.append(f"queue:{message.text}") - if self.submit_error is not None: - raise self.submit_error - self.dashboard_submissions.append(message) + self.submissions.append(message) return 1 async def interrupt(self) -> int: @@ -260,7 +252,10 @@ async def scenario() -> None: asyncio.run(scenario()) - assert driver.submissions == ["first", "second"] + assert [(m.message_id, m.text) for m in driver.submissions] == [ + ("one", "first"), + ("two", "second"), + ] assert [message.status for message in interaction.messages] == ["sent", "sent"] assert interaction.activity == "busy" assert events == [ @@ -302,6 +297,7 @@ async def scenario() -> None: await control.start() await control.complete(driver) assert interaction.messages[0].status == "sending" + assert driver.submissions[0].message_id == "later" control.message_discarded("later") asyncio.run(scenario()) @@ -325,7 +321,7 @@ async def scenario() -> None: asyncio.run(scenario()) assert interaction.messages[0].status == "sent" - assert events == ["initial-user", "queue:queued for API", "user:queued for API"] + assert events == ["initial-user", "submit:queued for API", "user:queued for API"] def test_interrupt_cancels_toolkit_before_backend_and_then_flushes() -> None: diff --git a/tests/unit_tests/rpent/planner/test_api_contracts.py b/tests/unit_tests/rpent/planner/test_api_contracts.py index 1d7101363..907723109 100644 --- a/tests/unit_tests/rpent/planner/test_api_contracts.py +++ b/tests/unit_tests/rpent/planner/test_api_contracts.py @@ -12,288 +12,898 @@ # See the License for the specific language governing permissions and # limitations under the License. +"""Offline integration tests against the installed SDK and Harness, not loop mocks.""" + from __future__ import annotations import asyncio -import base64 +import copy +import json import queue -from typing import Any +import threading +import time +from types import SimpleNamespace +import numpy as np import pytest -from pydantic_ai import BinaryContent, ToolReturn -from pydantic_ai.messages import ( - ModelResponse, +from pydantic_ai import ( + BinaryContent, + ModelRequest, TextPart, - ToolCallPart, ToolReturnPart, + UserPromptPart, ) -from pydantic_ai.models.function import FunctionModel -from pydantic_ai.usage import RequestUsage - -from rpent.dashboard.events import TranscriptEvent, UsageEvent -from rpent.planner.api_loop import ( - ApiAgentLoop, - _build_tools, - _content_blocks_to_pydantic, - _make_tool_function, -) -from rpent.tools.toolkit import ToolResult +from pydantic_ai.exceptions import UsageLimitExceeded +from pydantic_ai.models.function import DeltaThinkingPart, DeltaToolCall, FunctionModel +from pydantic_ai.models.test import TestModel + +from rpent.dashboard.events import RunStartedEvent, TranscriptEvent, UsageEvent +from rpent.dashboard.state import DashboardState +from rpent.planner.api_loop import ApiAgentLoop +from rpent.session import EnvState +from rpent.tools import common +from rpent.tools.toolkit import Toolkit, readonly +FINISH_ARGS = {"status": "success", "summary": "done"} -class RecordingSink: - def __init__(self) -> None: - self.events: list[Any] = [] - @property - def enabled(self) -> bool: - return True +class Events: + enabled = True - def emit(self, event: Any) -> None: + def __init__(self): + self.events = [] + + def emit(self, event): self.events.append(event) -class FakeToolkit: - state = None +class RobotToolkit(Toolkit): + def __init__(self, events, state=None): + self.calls = [] + super().__init__(dashboard_events=events, memory=SimpleNamespace(), state=state) - def __init__(self, result: dict[str, Any] | None = None) -> None: - self.result = result or {"ok": True} - self.calls: list[tuple[str, dict[str, Any]]] = [] - self.cancel_calls = 0 + def _register_common_tools(self): + self.add_tool("finish", common.TOOLS_SPEC[-1], self.finish) + self.register("observe", lambda: {"position": 1, "_image_bytes": b"image"}) - def get_tools_spec(self) -> list[dict[str, Any]]: - return [ + def register(self, name, handler, schema=None): + self.add_tool( + name, { - "name": "finish", - "description": "Finish after the environment accepts the result.", - "input_schema": { - "type": "object", - "properties": { - "status": {"type": "string"}, - "summary": {"type": "string"}, - }, - "required": ["status", "summary"], - }, - } - ] - - def execute_tool(self, name: str, args: dict[str, Any]) -> ToolResult: - self.calls.append((name, args)) - return ToolResult(name, dict(self.result)) - - def cancel_active_and_wait(self) -> None: - self.cancel_calls += 1 - - -def solve_with_model( - function: Any, - toolkit: FakeToolkit, - sink: RecordingSink, - *, - timeout_s: float = 5, -): - planner = ApiAgentLoop( - FunctionModel(function), - max_tokens=321, - dashboard_events=sink, - timeout_s=timeout_s, - ) - return planner.solve( - system_prompt="Use tools carefully.", - user_message="complete the task", + "name": name, + "description": name, + "input_schema": schema or {"type": "object", "properties": {}}, + }, + readonly(handler), + ) + + @readonly + def finish(self, status, summary): + return {"_finish": True, "status": status, "summary": summary} + + def execute_tool(self, name, input_dict): + self.calls.append((name, input_dict)) + return super().execute_tool(name, input_dict) + + +@pytest.fixture(autouse=True) +def local_tools(monkeypatch): + monkeypatch.setattr("rpent.tools.toolkit.substitute", lambda value: value) + monkeypatch.setenv("PYDANTIC_AI_NO_BANNER", "1") + + +def tool(name, args=None, index=0): + return {index: DeltaToolCall(name=name, json_args=json.dumps(args or {}))} + + +def finish(): + return tool("finish", FINISH_ARGS) + + +def solve(tmp_path, model, toolkit=None, events=None, **kwargs): + events = events or Events() + toolkit = toolkit or RobotToolkit(events) + config = { + key: kwargs.pop(key) + for key in ( + "max_tokens", + "timeout_s", + "no_images", + "reasoning_effort", + "interactive", + ) + if key in kwargs + } + planner = ApiAgentLoop(model=model, dashboard_events=events, **config) + result = planner.solve( + system_prompt="Use the robot toolkit. Finish when done.", + user_message="Do the task.", toolkit=toolkit, - max_turns=3, + max_turns=kwargs.pop("max_turns", 10), + **kwargs, ) + return result, toolkit, events + + +def dashboard(tmp_path): + state = DashboardState( + output_dir=tmp_path, + dashboard_spec={ + "task": { + "command": "/rpent-task", + "usage": "/rpent-task ", + "fields": ({"name": "seed", "kind": "integer", "minimum": 0},), + "display": "{seed}", + "output_slug": "s{seed}", + }, + "runtime_components": (), + "frame_channels": (), + "primitives": (), + }, + ) + state.shared_services_ready() + state.submit_input("/rpent-task 0") + assert state.wait_for_task(timeout=0) is not None + state.emit(RunStartedEvent()) + return state + + +def user_texts(messages): + return [ + part.content + for message in messages + if isinstance(message, ModelRequest) + for part in message.parts + if isinstance(part, UserPromptPart) + ] -def test_successful_finish_waits_for_its_tool_result() -> None: - seen_instructions: list[str | None] = [] +def test_finish_returns_transcript_for_rpent_without_step_files(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + model = TestModel(call_tools=[], custom_output_args=FINISH_ARGS) + result, toolkit, _ = solve(tmp_path, model) + assert result.error is None + assert result.finish_result == {"_finish": True, **FINISH_ARGS} + assert toolkit.calls == [("finish", FINISH_ARGS)] + assert result.stats["turns_used"] == 1 + assert [m["role"] for m in result.messages] == ["user", "assistant", "tool"] + assert result.messages[0] == {"role": "user", "content": "Do the task."} + call = result.messages[1]["content"][0] + assert call["name"] == "finish" + assert call["input"] == FINISH_ARGS + returned = result.messages[2] + assert returned["tool_call_id"] == call["id"] + assert json.loads(returned["content"]) == {"_finish": True, **FINISH_ARGS} + assert not list(tmp_path.iterdir()) + assert model.last_model_request_parameters.output_tools[0].name == "finish" + assert "finish" not in [ + t.name for t in model.last_model_request_parameters.function_tools + ] - def model(messages: list[Any], info: Any) -> ModelResponse: - seen_instructions.append(info.instructions) - assert info.model_settings["max_tokens"] == 321 - assert not any( - isinstance(part, ToolReturnPart) - for message in messages - for part in message.parts - ) - return ModelResponse( - parts=[ - ToolCallPart( - "finish", - {"status": "success", "summary": "done"}, - "finish-call", - ) - ], - usage=RequestUsage(input_tokens=7, output_tokens=3), - ) - toolkit = FakeToolkit() - sink = RecordingSink() - result = solve_with_model(model, toolkit, sink) +def test_refused_finish_retries_without_claiming_success(tmp_path): + events = Events() + toolkit = RobotToolkit(events) + attempts = [] - assert seen_instructions == ["Use tools carefully."] - assert toolkit.calls == [("finish", {"status": "success", "summary": "done"})] - assert result.finish_result == { - "_finish": True, - "status": "success", - "summary": "done", - } + def guarded_finish(**args): + attempts.append(args) + if len(attempts) <= 2: + return {"error": "finish refused; verify the environment"} + return {"_finish": True, **args} + + toolkit.register("finish", guarded_finish, common.TOOLS_SPEC[-1]["input_schema"]) + histories = [] + + async def stream(messages, info): + histories.append(copy.deepcopy(messages)) + yield finish() + + result, _, _ = solve( + tmp_path, FunctionModel(stream_function=stream), toolkit, events + ) assert result.error is None - assert result.stats == { - "turns_used": 1, - "tool_calls": 1, - "total_input_tokens": 7, - "total_output_tokens": 3, - "cache_read_tokens": 0, - "cache_write_tokens": 0, - "requests": 1, - } - assert any(isinstance(event, TranscriptEvent) for event in sink.events) - assert any(isinstance(event, UsageEvent) for event in sink.events) - - -def test_rejected_finish_does_not_end_the_run() -> None: - def model(messages: list[Any], info: Any) -> ModelResponse: - del info - if any( - isinstance(part, ToolReturnPart) - for message in messages - for part in message.parts - ): - return ModelResponse(parts=[TextPart("I could not finish.")]) - return ModelResponse( - parts=[ - ToolCallPart( - "finish", - {"status": "success", "summary": "too early"}, - "rejected-finish", - ) - ] - ) + assert len(attempts) == 3 + assert result.stats["turns_used"] == 3 + assert result.finish_result == {"_finish": True, **FINISH_ARGS} + assert "finish refused" in repr(histories[1]) + assert "finish refused" in json.dumps(result.messages) + + +def test_persistent_finish_refusal_stops_at_request_budget(tmp_path): + events = Events() + toolkit = RobotToolkit(events) + toolkit.register( + "finish", + lambda **args: {"error": "finish refused"}, + common.TOOLS_SPEC[-1]["input_schema"], + ) - toolkit = FakeToolkit({"error": "finish refused by environment"}) - result = solve_with_model(model, toolkit, RecordingSink()) + async def stream(messages, info): + yield finish() + result, _, _ = solve( + tmp_path, FunctionModel(stream_function=stream), toolkit, events, max_turns=3 + ) assert result.finish_result is None assert result.error is None - assert result.stats["tool_calls"] == 1 - assert any( - message.get("role") == "tool" - and message.get("content") == '{\n "error": "finish refused by environment"\n}' - for message in result.messages + assert result.stats["turns_used"] == 3 + assert len(toolkit.calls) == 3 + + +def test_anthropic_request_retains_all_prompt_cache_controls(tmp_path, monkeypatch): + from pydantic_ai.models.anthropic import AnthropicModel + from pydantic_ai.providers.anthropic import AnthropicProvider + + model = AnthropicModel( + "claude-sonnet-4-5", provider=AnthropicProvider(api_key="offline-test") ) + requests = [] + + async def capture(**kwargs): + requests.append(kwargs) + raise RuntimeError("offline request captured") + + monkeypatch.setattr(model.client.beta.messages, "create", capture) + result, _, _ = solve(tmp_path, model, max_tokens=321) + assert result.error == "RuntimeError: offline request captured" + assert len(requests) == 1 + request = requests[0] + assert request["max_tokens"] == 321 + assert request["system"][0]["text"] == "Use the robot toolkit. Finish when done." + assert request["system"][-1]["cache_control"]["type"] == "ephemeral" + assert request["tools"][-1]["cache_control"]["type"] == "ephemeral" + assert ( + request["messages"][-1]["content"][-1]["cache_control"]["type"] == "ephemeral" + ) + assert request["tool_choice"]["disable_parallel_tool_use"] is True -def test_backend_failure_is_returned_without_escaping() -> None: - def model(messages: list[Any], info: Any) -> ModelResponse: - del messages, info - raise RuntimeError("provider failed") +@pytest.mark.parametrize("no_images", [False, True]) +def test_multimodal_tool_results_and_dashboard_events(tmp_path, no_images): + histories = [] - result = solve_with_model(model, FakeToolkit(), RecordingSink()) + async def stream(messages, info): + histories.append(copy.deepcopy(messages)) + if len(histories) == 1: + yield {0: DeltaThinkingPart(content="Inspect the scene.")} + yield tool("observe", index=1) + else: + yield "Ready." + yield finish() - assert result.finish_result is None - assert result.error == "RuntimeError: provider failed" - assert result.messages == [{"role": "user", "content": "complete the task"}] + result, _, events = solve( + tmp_path, FunctionModel(stream_function=stream), no_images=no_images + ) + assert result.error is None + returns = [ + part + for message in histories[1] + if isinstance(message, ModelRequest) + for part in message.parts + if isinstance(part, ToolReturnPart) + ] + images = [ + item + for part in returns + for item in part.content + if isinstance(item, BinaryContent) + ] + assert bool(images) is not no_images + if images: + assert images[0].data == b"image" + else: + assert "Image omitted" in repr(returns) + transcript = [e.payload for e in events.events if isinstance(e, TranscriptEvent)] + assert {p["type"] for p in transcript} >= { + "user", + "thinking", + "text", + "tool_call", + "tool_result", + } + assert all("aW1hZ2U=" not in str(p) for p in transcript) + usage = [event for event in events.events if isinstance(event, UsageEvent)][-1] + assert usage.tool_calls == 2 + assert result.stats["total_input_tokens"] == usage.inp > 0 + assert result.stats["total_output_tokens"] == usage.out > 0 + serialized = json.dumps(result.messages) + assert "Inspect the scene." in serialized + assert "Ready." in serialized + assert "aW1hZ2U=" not in serialized + assert any(m.get("name") == "observe" for m in result.messages) + + +@pytest.fixture +def image_state(tmp_path): + state = EnvState(tmp_path / "state") + for value in (0, 255): + with state.record_step(state={}): + for name in ("camera.png", "camera.jpg", "camera.jpeg"): + image = np.full((4, 4, 3), value, dtype=np.uint8) + assert state.save(name, image) == name + assert state.save("notes.json", {"position": 1}, step=0) == "notes.json" + return state + + +@pytest.mark.parametrize("no_images", [False, True]) +@pytest.mark.parametrize( + "name,step,media_type", + [ + ("camera.png", None, "image/png"), + ("camera.jpg", 0, "image/jpeg"), + ("camera.jpeg", 1, "image/jpeg"), + ], +) +def test_read_image_returns_saved_artifact_and_records_call( + tmp_path, image_state, monkeypatch, no_images, name, step, media_type +): + events = Events() + toolkit = RobotToolkit(events, state=image_state) + arguments = {"name": name} + if step is not None: + arguments["step"] = step + resolved_step = image_state.latest_step if step is None else step + expected_bytes = image_state.load_bytes(name, step=resolved_step) + histories = [] + if no_images: -def test_timeout_cancels_active_toolkit_work() -> None: - async def model(messages: list[Any], info: Any) -> ModelResponse: - del messages, info - await asyncio.sleep(10) - return ModelResponse(parts=[TextPart("unreachable")]) + def unexpected_read(*args, **kwargs): + pytest.fail("--no-images should not read image bytes") - toolkit = FakeToolkit() - result = solve_with_model( - model, + monkeypatch.setattr(image_state, "load_bytes", unexpected_read) + + async def stream(messages, info): + histories.append(copy.deepcopy(messages)) + yield tool("read_image", arguments) if len(histories) == 1 else finish() + + result, _, _ = solve( + tmp_path, + FunctionModel(stream_function=stream), toolkit, - RecordingSink(), - timeout_s=0.01, + events, + no_images=no_images, ) + assert result.error is None + assert result.finish_result == {"_finish": True, **FINISH_ARGS} + returned = next( + part + for message in histories[1] + for part in message.parts + if isinstance(part, ToolReturnPart) and part.tool_name == "read_image" + ) + assert returned.content["artifact"] == name + assert returned.content["step"] == resolved_step + images = [ + item + for message in histories[1] + for part in message.parts + if isinstance(part, (ToolReturnPart, UserPromptPart)) + and isinstance(part.content, list) + for item in part.content + if isinstance(item, BinaryContent) + ] + if no_images: + assert not images + assert "--no-images" in returned.content["notice"] + else: + assert len(images) == 1 + assert images[0].data == expected_bytes + assert images[0].media_type == media_type + assert result.stats["tool_calls"] == 2 + recorded = next(m for m in result.messages if m.get("name") == "read_image") + assert recorded["tool_call_id"] == returned.tool_call_id + assert json.loads(recorded["content"]) == returned.content + transcript = [e.payload for e in events.events if isinstance(e, TranscriptEvent)] + assert [p["type"] for p in transcript if p.get("tool") == "read_image"] == [ + "tool_call", + "tool_result", + ] + - assert result.error == "API planner timed out after 0.01s" - assert toolkit.cancel_calls == 1 - assert result.messages == [{"role": "user", "content": "complete the task"}] +@pytest.mark.parametrize( + "name,step,error", + [ + ("missing.png", 0, "not available"), + ("notes.json", 0, "not an image"), + ("../camera.png", 0, "base filename"), + ("camera.png", 99, "not present"), + ("unregistered.png", 0, "not available"), + ("camera.png", 0, "not available"), + ], +) +def test_read_image_errors_reach_model_without_ending_run( + tmp_path, image_state, name, step, error +): + events = Events() + toolkit = RobotToolkit(events, state=image_state) + image_state.artifact_path("camera.png", step=0).unlink() + unregistered = image_state.artifact_path("unregistered.png", step=0) + unregistered.parent.mkdir() + unregistered.write_bytes(b"not registered in the step") + histories = [] + + async def stream(messages, info): + histories.append(copy.deepcopy(messages)) + if len(histories) == 1: + yield tool("read_image", {"name": name, "step": step}) + else: + yield finish() + + result, _, _ = solve( + tmp_path, FunctionModel(stream_function=stream), toolkit, events + ) + assert result.error is None + assert result.finish_result == {"_finish": True, **FINISH_ARGS} + returned = next( + part + for message in histories[1] + for part in message.parts + if isinstance(part, ToolReturnPart) and part.tool_name == "read_image" + ) + assert error in returned.content["error"] -def test_queue_and_dashboard_inputs_are_rejected_before_model_use() -> None: +def test_schema_validation_and_sequential_physical_tools(tmp_path): + events = Events() + toolkit = RobotToolkit(events) + positions = [] + toolkit.register( + "move", + lambda position: positions.append(position) or {"position": position}, + { + "type": "object", + "properties": {"position": {"type": "integer"}}, + "required": ["position"], + }, + ) calls = 0 - def model(messages: list[Any], info: Any) -> ModelResponse: + async def stream(messages, info): nonlocal calls - del messages, info calls += 1 - return ModelResponse(parts=[TextPart("unused")]) + if calls == 1: + yield tool("move", {"position": "bad"}) + elif calls == 2: + yield tool("move", {"position": 1}) | tool("move", {"position": 2}, 1) + else: + yield finish() + + result, _, _ = solve( + tmp_path, FunctionModel(stream_function=stream), toolkit, events + ) + assert result.error is None + assert positions == [1, 2] + assert len(toolkit.calls) == 3 # Only validated calls reach the Toolkit. + assert [m["name"] for m in result.messages if m["role"] == "tool"] == [ + "move", + "move", + "finish", + ] + + +def test_finish_skips_sibling_actions_using_native_end_strategy(tmp_path): + async def stream(messages, info): + yield finish() | tool("observe", index=1) - planner = ApiAgentLoop( - FunctionModel(model), - dashboard_events=RecordingSink(), + result, toolkit, _ = solve(tmp_path, FunctionModel(stream_function=stream)) + assert result.error is None + assert [name for name, _ in toolkit.calls] == ["finish"] + assert [m["name"] for m in result.messages if m["role"] == "tool"] == ["finish"] + + +@pytest.mark.parametrize("mode", ["normal", "terminal", "dashboard"]) +def test_request_budget_stops_normally_without_claiming_success(tmp_path, mode): + async def stream(messages, info): + yield tool("observe") + + state = dashboard(tmp_path) if mode == "dashboard" else None + result, toolkit, _ = solve( + tmp_path, + FunctionModel(stream_function=stream), + events=state, + max_turns=2, + interactive=mode == "terminal", + dashboard_interaction=state, ) + assert result.error is None + assert result.finish_result is None + assert result.stats["turns_used"] == 2 + assert len(toolkit.calls) == 2 + assert len([m for m in result.messages if m["role"] == "tool"]) == 2 - with pytest.raises(ValueError, match="cannot be used together"): - planner.solve( - system_prompt="", - user_message="task", - toolkit=FakeToolkit(), - max_turns=1, - input_queue=queue.Queue(), - dashboard_interaction=object(), - ) - assert calls == 0 +def test_other_usage_limit_failures_remain_errors(tmp_path): + async def stream(messages, info): + raise UsageLimitExceeded("provider-specific limit") + yield # pragma: no cover + result, _, _ = solve(tmp_path, FunctionModel(stream_function=stream), max_turns=1) + assert result.error.startswith("UsageLimitExceeded: provider-specific limit") + assert result.finish_result is None -def test_tool_schema_and_dispatch_are_mapped_to_pydantic_ai() -> None: - toolkit = FakeToolkit() - tools = _build_tools(toolkit) +def test_harness_compacts_long_history(tmp_path): + lengths = [] - assert [tool.name for tool in tools] == ["read_image", "finish"] - assert all(tool.sequential for tool in tools) - finish = tools[1] - assert finish.description == "Finish after the environment accepts the result." - assert ( - finish.function_schema.json_schema - == toolkit.get_tools_spec()[0]["input_schema"] + async def stream(messages, info): + lengths.append(len(messages)) + yield tool("observe") if len(lengths) < 45 else finish() + + result, _, _ = solve(tmp_path, FunctionModel(stream_function=stream), max_turns=50) + assert result.error is None + assert any(after < before for before, after in zip(lengths, lengths[1:])) + assert result.messages[0] == {"role": "user", "content": "Do the task."} + assert len([m for m in result.messages if m["role"] == "assistant"]) == 45 + assert len([m for m in result.messages if m.get("name") == "observe"]) == 44 + + +def test_dashboard_steering_enters_next_request_in_same_native_run(tmp_path): + state = dashboard(tmp_path) + sent = [] + histories = [] + + async def stream(messages, info): + histories.append(copy.deepcopy(messages)) + if len(histories) == 1: + sent.append(state.submit_input("Look left first.")) + await asyncio.sleep(0.05) + yield tool("observe") + else: + assert "Look left first." in user_texts(messages) + yield finish() + + result, _, _ = solve( + tmp_path, + FunctionModel(stream_function=stream), + events=state, + dashboard_interaction=state, ) + assert result.error is None + assert [m["content"] for m in result.messages if m["role"] == "user"] == [ + "Do the task.", + "Look left first.", + ] + assert len(histories) == 2 + assert state.planner_activity == "ended" + assert state.snapshot()["interaction"]["messages"][0]["status"] == "sent" -def test_tool_result_conversion_keeps_text_and_images_separate() -> None: - raw_image = b"\x89PNG\r\ncontract-image" - encoded = base64.b64encode(raw_image).decode() - blocks = [ - {"type": "text", "text": "observation"}, - { - "type": "image", - "source": { - "type": "base64", - "media_type": "image/png", - "data": encoded, - }, - }, +def test_dashboard_message_during_finish_is_unsent(tmp_path): + state = dashboard(tmp_path) + toolkit = RobotToolkit(state) + + def finish_and_submit(**args): + state.submit_input("This must not reopen the completed task.") + time.sleep(0.05) + return {"_finish": True, **args} + + toolkit.register("finish", finish_and_submit, common.TOOLS_SPEC[-1]["input_schema"]) + result, _, _ = solve( + tmp_path, + TestModel(call_tools=[], custom_output_args=FINISH_ARGS), + toolkit, + state, + dashboard_interaction=state, + ) + assert result.error is None + assert result.stats["turns_used"] == 1 + assert state.snapshot()["interaction"]["messages"][0]["status"] == "unsent" + assert "This must not reopen" not in json.dumps(result.messages) + + +def test_dashboard_interrupt_then_followup_preserves_history(tmp_path): + state = dashboard(tmp_path) + calls = 0 + + async def stream(messages, info): + nonlocal calls + calls += 1 + if calls == 1: + yield "Starting inspection." + state.request_interrupt() + state.submit_input("Continue with the new instruction.") + await asyncio.sleep(5) + else: + assert "Do the task." in user_texts(messages) + assert "Continue with the new instruction." in user_texts(messages) + yield finish() + + result, _, _ = solve( + tmp_path, + FunctionModel(stream_function=stream), + events=state, + dashboard_interaction=state, + timeout_s=3, + ) + assert result.error is None + assert result.finish_result["status"] == "success" + assert [m["content"] for m in result.messages if m["role"] == "user"] == [ + "Do the task.", + "Continue with the new instruction.", + ] + assert state.snapshot()["interaction"]["messages"][0]["status"] == "sent" + + +@pytest.fixture(params=[False, True], ids=["normal", "tool-stops-first"]) +def dashboard_interrupt_order(request, monkeypatch): + if not request.param: + return + from rpent.planner.api_loop import _Session + + interrupt = _Session.interrupt + + async def interrupt_after_tool_stop(self): + # Let the SDK reach the queued sibling tool before Dashboard sends its + # cancellation token, reproducing the physical-drain scheduling race. + await asyncio.wait_for(self.run_done.wait(), timeout=1) + return await interrupt(self) + + monkeypatch.setattr(_Session, "interrupt", interrupt_after_tool_stop) + + +def test_dashboard_interrupt_drains_tool_before_followup( + tmp_path, dashboard_interrupt_order +): + state = dashboard(tmp_path) + toolkit = RobotToolkit(state) + stopped = threading.Event() + calls = 0 + + def move(): + try: + state.submit_input("Old instruction 1.") + state.submit_input("Old instruction 2.") + while any( + m["status"] != "sending" + for m in state.snapshot()["interaction"]["messages"] + ): + toolkit.raise_if_cancelled() + state.wait_for_interaction_change( + state.interaction_version, timeout=0.05 + ) + state.request_interrupt() + state.submit_input("Continue after stopping the move.") + while True: + toolkit.raise_if_cancelled() + time.sleep(0.005) + finally: + stopped.set() + + toolkit.register("move", move) + + async def stream(messages, info): + nonlocal calls + calls += 1 + if calls == 1: + yield tool("move") | tool("observe", index=1) + else: + assert stopped.is_set() + assert user_texts(messages) == [ + "Do the task.", + "Continue after stopping the move.", + ] + yield finish() + + result, _, _ = solve( + tmp_path, + FunctionModel(stream_function=stream), + toolkit, + state, + dashboard_interaction=state, + timeout_s=3, + ) + assert result.error is None + assert result.finish_result["status"] == "success" + assert [name for name, _ in toolkit.calls] == ["move", "finish"] + assert [m["status"] for m in state.snapshot()["interaction"]["messages"]] == [ + "unsent", + "unsent", + "sent", ] - text, images = _content_blocks_to_pydantic(blocks) - assert text == "observation" - assert images == [BinaryContent(data=raw_image, media_type="image/png")] - assert blocks[1]["source"]["data"] == encoded +def test_task_replacement_cancels_and_drains_physical_work( + tmp_path, dashboard_interrupt_order +): + state = dashboard(tmp_path) + toolkit = RobotToolkit(state) + stopped = threading.Event() + + def move(): + state.submit_input("/rpent-task 1") + try: + while True: + toolkit.raise_if_cancelled() + time.sleep(0.005) + finally: + stopped.set() + + toolkit.register("move", move) + + async def stream(messages, info): + yield tool("move") | tool("observe", index=1) + + result, _, _ = solve( + tmp_path, + FunctionModel(stream_function=stream), + toolkit, + state, + dashboard_interaction=state, + timeout_s=3, + ) + assert result.error is None + assert stopped.is_set() + assert [name for name, _ in toolkit.calls] == ["move"] + assert state.planner_activity == "ended" + assert result.finish_result is None + + +def test_timeout_drains_physical_work_and_records_partial_history(tmp_path): + events = Events() + toolkit = RobotToolkit(events) + stopped = threading.Event() + + def move(): + try: + while True: + toolkit.raise_if_cancelled() + time.sleep(0.005) + finally: + stopped.set() + + toolkit.register("move", move) + async def stream(messages, info): + yield tool("move") + + result, _, _ = solve( + tmp_path, FunctionModel(stream_function=stream), toolkit, events, timeout_s=0.15 + ) + assert "timed out" in result.error + assert stopped.is_set() + assert result.messages[0] == {"role": "user", "content": "Do the task."} + assert result.messages[1]["content"][0]["name"] == "move" + json.dumps(result.messages) -def test_no_images_mode_suppresses_binary_tool_content() -> None: - toolkit = FakeToolkit({"value": "visible", "_image_bytes": b"secret pixels"}) - multimodal = _make_tool_function(toolkit, "finish")( - status="success", - summary="done", +def test_native_terminal_preserves_completed_actions_after_model_failure( + tmp_path, monkeypatch +): + import pydantic_ai._cli as cli + + requests = [] + replies = iter(["Inspect the scene.", "Continue after the error.", "/exit"]) + + async def read_prompt(*args, **kwargs): + assert requests, "The preset task must run before the first input prompt." + return next(replies) + + async def stream(messages, info): + requests.append(copy.deepcopy(messages)) + if len(requests) == 2: + yield tool("observe") + elif len(requests) == 3: + raise RuntimeError("provider failed after observation") + else: + yield "Ready." + + monkeypatch.setattr(cli, "PYDANTIC_AI_HOME", tmp_path / "cli") + monkeypatch.setattr( + cli, "PromptSession", lambda **kwargs: SimpleNamespace(prompt_async=read_prompt) ) - text_only = _make_tool_function(toolkit, "finish", no_images=True)( - status="success", - summary="done", + result, toolkit, _ = solve( + tmp_path, FunctionModel(stream_function=stream), interactive=True ) + assert result.error is None + assert any( + isinstance(part, TextPart) and part.content == "Ready." + for message in requests[-1] + for part in message.parts + ) + assert user_texts(requests[-1]) == [ + "Do the task.", + "Inspect the scene.", + "Continue after the error.", + ] + returns = [ + part + for message in requests[-1] + for part in message.parts + if isinstance(part, ToolReturnPart) and part.tool_name == "observe" + ] + assert len(returns) == 1 + assert "position" in str(returns[0].content) + assert any(isinstance(item, BinaryContent) for item in returns[0].content) + assert toolkit.calls == [("observe", {})] + assert [m["content"] for m in result.messages if m["role"] == "user"] == user_texts( + requests[-1] + ) + + +def test_native_terminal_followups_share_budget_and_clear_recovered_error( + tmp_path, monkeypatch, capsys +): + import pydantic_ai._cli as cli + + requests = [] + replies = iter(["Inspect.", "Retry.", "Try after the budget.", "/exit"]) + + async def read_prompt(*args, **kwargs): + return next(replies) + + async def stream(messages, info): + requests.append(copy.deepcopy(messages)) + if len(requests) == 1: + yield "Ready." + elif len(requests) == 2: + raise RuntimeError("temporary provider failure") + else: + yield tool("observe") + + monkeypatch.setattr(cli, "PYDANTIC_AI_HOME", tmp_path / "cli") + monkeypatch.setattr( + cli, "PromptSession", lambda **kwargs: SimpleNamespace(prompt_async=read_prompt) + ) + result, toolkit, _ = solve( + tmp_path, FunctionModel(stream_function=stream), interactive=True, max_turns=3 + ) + assert len(requests) == result.stats["turns_used"] == 3 + assert toolkit.calls == [("observe", {})] + assert result.error is None + assert result.finish_result is None + assert any(m.get("name") == "observe" for m in result.messages) + output = " ".join(capsys.readouterr().out.split()) + assert "Use /exit to close the session." in output + - assert isinstance(multimodal, ToolReturn) - assert multimodal.return_value == '{\n "value": "visible"\n}' - assert len(multimodal.content or []) == 1 - assert isinstance(multimodal.content[0], BinaryContent) - assert text_only == '{\n "value": "visible"\n}' - assert "secret" not in text_only +def test_legacy_terminal_queue_is_rejected(tmp_path): + with pytest.raises(ValueError, match="terminal input queue"): + solve(tmp_path, TestModel(), input_queue=queue.Queue()) + + +def test_native_terminal_and_dashboard_cannot_run_together(tmp_path): + with pytest.raises(ValueError, match="clai and Dashboard cannot run together"): + solve(tmp_path, TestModel(), interactive=True, dashboard_interaction=object()) + + +@pytest.mark.parametrize("after_tool", [False, True]) +def test_model_failure_preserves_rpent_transcript(tmp_path, after_tool): + calls = 0 + + async def stream(messages, info): + nonlocal calls + calls += 1 + if after_tool and calls == 1: + yield tool("observe") + return + raise RuntimeError("provider unavailable") + + result, _, _ = solve(tmp_path, FunctionModel(stream_function=stream)) + assert result.error == "RuntimeError: provider unavailable" + if after_tool: + assert [m["role"] for m in result.messages] == ["user", "assistant", "tool"] + assert result.messages[-1]["name"] == "observe" + else: + assert result.messages == [{"role": "user", "content": "Do the task."}] + json.dumps(result.messages) + + +def test_factory_passes_interactive_mode(tmp_path, monkeypatch): + from rpent.planner.base import Planner, build_planner + + model = TestModel(call_tools=[], custom_output_args=FINISH_ARGS) + monkeypatch.setattr("rpent.planner.base.build_api_model", lambda *args: model) + planner = build_planner( + "api", + output_dir=tmp_path, + recipe_tag="test", + robot_name="test", + model="test:test", + dashboard_events=Events(), + interactive=True, + ) + assert isinstance(planner, ApiAgentLoop) + assert isinstance(planner, Planner) + assert planner.interactive is True