mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/cli: interactive continue prompts for loop and PTY limits
on_limit raises tool-round allowance; PTY session/budget limits can prompt to raise; MaxRoundsEvent.continued marks extended runs.
This commit is contained in:
@@ -54,7 +54,7 @@ Async SQLAlchemy + aiosqlite. `MemoryStore`: schema init, default local user, se
|
|||||||
|
|
||||||
- **`ChatClient`** Protocol for `chat_completions`.
|
- **`ChatClient`** Protocol for `chat_completions`.
|
||||||
- **`@tool` / `ToolRegistry`**: decorator infers JSON Schema from type hints; execute tools by name.
|
- **`@tool` / `ToolRegistry`**: decorator infers JSON Schema from type hints; execute tools by name.
|
||||||
- **`run_chat_loop`**: multi-round tool loop, yields `AgentEvent` stream.
|
- **`run_chat_loop`**: multi-round tool loop, yields `AgentEvent` stream; optional `on_limit` to continue past max rounds.
|
||||||
- **`ChatAgent`**: wrapper with optional `MemoryStore` bind (load/append messages).
|
- **`ChatAgent`**: wrapper with optional `MemoryStore` bind (load/append messages).
|
||||||
|
|
||||||
### Tools (`tools/`)
|
### Tools (`tools/`)
|
||||||
@@ -65,6 +65,7 @@ Module-level `@tool` handlers. Call `set_workspace_root()` before use.
|
|||||||
- **`file`**: `read_file`, `write_file`, `listdir`, `edit_replace` (first occurrence).
|
- **`file`**: `read_file`, `write_file`, `listdir`, `edit_replace` (first occurrence).
|
||||||
- **`process`**: `run_command` (argv, no shell, timeout, optional stdin/env); PTY `open_pty` / `read_pty` / `write_pty` / `close_pty` (**Unix only**: `pty`+`fork`).
|
- **`process`**: `run_command` (argv, no shell, timeout, optional stdin/env); PTY `open_pty` / `read_pty` / `write_pty` / `close_pty` (**Unix only**: `pty`+`fork`).
|
||||||
- PTY: structured status (`alive`/`exit_code`/`data`); `read_pty(..., until=)`; session limit/idle TTL/output budget; close SIGTERM→SIGKILL.
|
- PTY: structured status (`alive`/`exit_code`/`data`); `read_pty(..., until=)`; session limit/idle TTL/output budget; close SIGTERM→SIGKILL.
|
||||||
|
- CLI limit hooks: interactive confirm to raise tool-loop rounds, PTY session cap, or PTY output budget.
|
||||||
- **`DEFAULT_TOOLS`**: file + process tool list for a `ToolRegistry`.
|
- **`DEFAULT_TOOLS`**: file + process tool list for a `ToolRegistry`.
|
||||||
|
|
||||||
### CLI (`cli/`)
|
### CLI (`cli/`)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from plyngent.lmproto.openai_compatible.model import UserChatMessage
|
|||||||
from .loop import DEFAULT_MAX_ROUNDS, run_chat_loop
|
from .loop import DEFAULT_MAX_ROUNDS, run_chat_loop
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import AsyncIterator, Sequence
|
from collections.abc import AsyncIterator, Callable, Sequence
|
||||||
|
|
||||||
from plyngent.lmproto.openai_compatible.model import AnyChatMessage
|
from plyngent.lmproto.openai_compatible.model import AnyChatMessage
|
||||||
from plyngent.memory import MemoryStore
|
from plyngent.memory import MemoryStore
|
||||||
@@ -16,6 +16,8 @@ if TYPE_CHECKING:
|
|||||||
from .events import AgentEvent
|
from .events import AgentEvent
|
||||||
from .tools import ToolRegistry
|
from .tools import ToolRegistry
|
||||||
|
|
||||||
|
type LimitContinueHook = Callable[[str], bool]
|
||||||
|
|
||||||
|
|
||||||
class ChatAgent:
|
class ChatAgent:
|
||||||
"""Thin wrapper: chat client + optional tools + optional memory bind."""
|
"""Thin wrapper: chat client + optional tools + optional memory bind."""
|
||||||
@@ -27,6 +29,7 @@ class ChatAgent:
|
|||||||
session_id: int | None
|
session_id: int | None
|
||||||
max_rounds: int
|
max_rounds: int
|
||||||
temperature: float | None
|
temperature: float | None
|
||||||
|
on_limit: LimitContinueHook | None
|
||||||
messages: list[AnyChatMessage]
|
messages: list[AnyChatMessage]
|
||||||
|
|
||||||
def __init__( # noqa: PLR0913
|
def __init__( # noqa: PLR0913
|
||||||
@@ -40,6 +43,7 @@ class ChatAgent:
|
|||||||
max_rounds: int = DEFAULT_MAX_ROUNDS,
|
max_rounds: int = DEFAULT_MAX_ROUNDS,
|
||||||
temperature: float | None = None,
|
temperature: float | None = None,
|
||||||
messages: Sequence[AnyChatMessage] | None = None,
|
messages: Sequence[AnyChatMessage] | None = None,
|
||||||
|
on_limit: LimitContinueHook | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.client = client
|
self.client = client
|
||||||
self.model = model
|
self.model = model
|
||||||
@@ -48,6 +52,7 @@ class ChatAgent:
|
|||||||
self.session_id = session_id
|
self.session_id = session_id
|
||||||
self.max_rounds = max_rounds
|
self.max_rounds = max_rounds
|
||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
|
self.on_limit = on_limit
|
||||||
self.messages = list(messages) if messages is not None else []
|
self.messages = list(messages) if messages is not None else []
|
||||||
|
|
||||||
async def load_history(self) -> None:
|
async def load_history(self) -> None:
|
||||||
@@ -84,6 +89,7 @@ class ChatAgent:
|
|||||||
tools=self.tools,
|
tools=self.tools,
|
||||||
max_rounds=self.max_rounds,
|
max_rounds=self.max_rounds,
|
||||||
temperature=self.temperature,
|
temperature=self.temperature,
|
||||||
|
on_limit=self.on_limit,
|
||||||
):
|
):
|
||||||
yield event
|
yield event
|
||||||
|
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ class ToolResultEvent(Struct, tag_field="type", tag="tool_result"):
|
|||||||
|
|
||||||
class MaxRoundsEvent(Struct, tag_field="type", tag="max_rounds"):
|
class MaxRoundsEvent(Struct, tag_field="type", tag="max_rounds"):
|
||||||
rounds: int
|
rounds: int
|
||||||
|
continued: bool = False
|
||||||
|
|
||||||
|
|
||||||
type AgentEvent = TextDeltaEvent | AssistantMessageEvent | ToolCallEvent | ToolResultEvent | MaxRoundsEvent
|
type AgentEvent = TextDeltaEvent | AssistantMessageEvent | ToolCallEvent | ToolResultEvent | MaxRoundsEvent
|
||||||
|
|||||||
+60
-40
@@ -5,6 +5,7 @@ from typing import TYPE_CHECKING
|
|||||||
from msgspec import UNSET
|
from msgspec import UNSET
|
||||||
|
|
||||||
from plyngent.lmproto.openai_compatible.model import (
|
from plyngent.lmproto.openai_compatible.model import (
|
||||||
|
AnyAssistantToolCall,
|
||||||
AssistantChatMessage,
|
AssistantChatMessage,
|
||||||
AssistantFunctionToolCall,
|
AssistantFunctionToolCall,
|
||||||
ChatCompletionsParam,
|
ChatCompletionsParam,
|
||||||
@@ -21,17 +22,38 @@ from .events import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import AsyncIterator, Sequence
|
from collections.abc import AsyncIterator, Callable, Sequence
|
||||||
|
|
||||||
from plyngent.lmproto.openai_compatible.model import AnyChatMessage, AnyToolItem
|
from plyngent.lmproto.openai_compatible.model import AnyChatMessage, AnyToolItem
|
||||||
|
|
||||||
from .client import ChatClient
|
from .client import ChatClient
|
||||||
from .tools import ToolRegistry
|
from .tools import ToolRegistry
|
||||||
|
|
||||||
|
type LimitContinueHook = Callable[[str], bool]
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_MAX_ROUNDS = 32
|
DEFAULT_MAX_ROUNDS = 32
|
||||||
|
|
||||||
|
|
||||||
|
async def _execute_tool_calls(
|
||||||
|
tools: ToolRegistry,
|
||||||
|
tool_calls: Sequence[AnyAssistantToolCall],
|
||||||
|
messages: list[AnyChatMessage],
|
||||||
|
) -> AsyncIterator[AgentEvent]:
|
||||||
|
for call in tool_calls:
|
||||||
|
yield ToolCallEvent(tool_call=call)
|
||||||
|
if isinstance(call, AssistantFunctionToolCall):
|
||||||
|
result_text = await tools.execute(call.function.name, call.function.arguments)
|
||||||
|
tool_msg = ToolChatMessage(content=result_text, tool_call_id=call.id)
|
||||||
|
else:
|
||||||
|
tool_msg = ToolChatMessage(
|
||||||
|
content="error: custom tool calls are not supported",
|
||||||
|
tool_call_id=call.id,
|
||||||
|
)
|
||||||
|
messages.append(tool_msg)
|
||||||
|
yield ToolResultEvent(message=tool_msg)
|
||||||
|
|
||||||
|
|
||||||
async def run_chat_loop( # noqa: PLR0913
|
async def run_chat_loop( # noqa: PLR0913
|
||||||
client: ChatClient,
|
client: ChatClient,
|
||||||
messages: list[AnyChatMessage],
|
messages: list[AnyChatMessage],
|
||||||
@@ -40,58 +62,56 @@ async def run_chat_loop( # noqa: PLR0913
|
|||||||
tools: ToolRegistry | None = None,
|
tools: ToolRegistry | None = None,
|
||||||
max_rounds: int = DEFAULT_MAX_ROUNDS,
|
max_rounds: int = DEFAULT_MAX_ROUNDS,
|
||||||
temperature: float | None = None,
|
temperature: float | None = None,
|
||||||
|
on_limit: LimitContinueHook | None = None,
|
||||||
) -> AsyncIterator[AgentEvent]:
|
) -> AsyncIterator[AgentEvent]:
|
||||||
"""Multi-round chat/tool loop; mutates ``messages`` in place and yields events.
|
"""Multi-round chat/tool loop; mutates ``messages`` in place and yields events.
|
||||||
|
|
||||||
Continues until the model returns no tool calls, or ``max_rounds`` is hit.
|
Continues until the model returns no tool calls, or ``max_rounds`` is hit.
|
||||||
Streaming is non-stream LLM calls for reliable tool_calls; text is emitted
|
If ``on_limit`` is set and returns True when the cap is reached, another
|
||||||
as a single :class:`TextDeltaEvent` when content is present.
|
batch of ``max_rounds`` is granted and the loop continues.
|
||||||
"""
|
"""
|
||||||
tool_items: Sequence[AnyToolItem] | None = None
|
tool_items: Sequence[AnyToolItem] | None = None
|
||||||
if tools is not None and len(tools) > 0:
|
if tools is not None and len(tools) > 0:
|
||||||
tool_items = tools.tool_items()
|
tool_items = tools.tool_items()
|
||||||
|
|
||||||
for round_idx in range(max_rounds):
|
rounds_used = 0
|
||||||
param = ChatCompletionsParam(
|
allowance = max_rounds
|
||||||
messages=list(messages),
|
|
||||||
model=model,
|
|
||||||
temperature=temperature if temperature is not None else UNSET,
|
|
||||||
tools=list(tool_items) if tool_items is not None else UNSET,
|
|
||||||
)
|
|
||||||
response = await client.chat_completions(param, stream=False)
|
|
||||||
if not response.choices:
|
|
||||||
msg = "chat completion response contained no choices"
|
|
||||||
raise RuntimeError(msg)
|
|
||||||
assistant = response.choices[0].message
|
|
||||||
messages.append(assistant)
|
|
||||||
yield AssistantMessageEvent(message=assistant)
|
|
||||||
|
|
||||||
if isinstance(assistant.content, str) and assistant.content:
|
while True:
|
||||||
yield TextDeltaEvent(content=assistant.content)
|
while rounds_used < allowance:
|
||||||
|
rounds_used += 1
|
||||||
|
param = ChatCompletionsParam(
|
||||||
|
messages=list(messages),
|
||||||
|
model=model,
|
||||||
|
temperature=temperature if temperature is not None else UNSET,
|
||||||
|
tools=list(tool_items) if tool_items is not None else UNSET,
|
||||||
|
)
|
||||||
|
response = await client.chat_completions(param, stream=False)
|
||||||
|
if not response.choices:
|
||||||
|
msg = "chat completion response contained no choices"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
assistant = response.choices[0].message
|
||||||
|
messages.append(assistant)
|
||||||
|
yield AssistantMessageEvent(message=assistant)
|
||||||
|
|
||||||
tool_calls = assistant.tool_calls
|
if isinstance(assistant.content, str) and assistant.content:
|
||||||
if tool_calls is UNSET or not tool_calls:
|
yield TextDeltaEvent(content=assistant.content)
|
||||||
return
|
|
||||||
|
|
||||||
if tools is None:
|
tool_calls = assistant.tool_calls
|
||||||
return
|
if tool_calls is UNSET or not tool_calls:
|
||||||
|
return
|
||||||
|
if tools is None:
|
||||||
|
return
|
||||||
|
async for event in _execute_tool_calls(tools, tool_calls, messages):
|
||||||
|
yield event
|
||||||
|
|
||||||
for call in tool_calls:
|
reason = f"tool loop reached {allowance} rounds (used {rounds_used})"
|
||||||
yield ToolCallEvent(tool_call=call)
|
if on_limit is not None and on_limit(reason):
|
||||||
if isinstance(call, AssistantFunctionToolCall):
|
yield MaxRoundsEvent(rounds=allowance, continued=True)
|
||||||
result_text = await tools.execute(call.function.name, call.function.arguments)
|
allowance += max_rounds
|
||||||
tool_msg = ToolChatMessage(content=result_text, tool_call_id=call.id)
|
continue
|
||||||
else:
|
yield MaxRoundsEvent(rounds=allowance, continued=False)
|
||||||
tool_msg = ToolChatMessage(
|
return
|
||||||
content="error: custom tool calls are not supported",
|
|
||||||
tool_call_id=call.id,
|
|
||||||
)
|
|
||||||
messages.append(tool_msg)
|
|
||||||
yield ToolResultEvent(message=tool_msg)
|
|
||||||
|
|
||||||
_ = round_idx # used only for loop bound
|
|
||||||
|
|
||||||
yield MaxRoundsEvent(rounds=max_rounds)
|
|
||||||
|
|
||||||
|
|
||||||
def collect_assistant_messages(events: Sequence[AgentEvent]) -> list[AssistantChatMessage]:
|
def collect_assistant_messages(events: Sequence[AgentEvent]) -> list[AssistantChatMessage]:
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from plyngent.cli.editor import (
|
|||||||
open_in_editor,
|
open_in_editor,
|
||||||
resolve_config_path,
|
resolve_config_path,
|
||||||
)
|
)
|
||||||
|
from plyngent.cli.limits import install_cli_limit_hooks
|
||||||
from plyngent.cli.repl import run_repl
|
from plyngent.cli.repl import run_repl
|
||||||
from plyngent.cli.selection import select_model, select_provider
|
from plyngent.cli.selection import select_model, select_provider
|
||||||
from plyngent.cli.state import ReplState
|
from plyngent.cli.state import ReplState
|
||||||
@@ -66,6 +67,7 @@ async def _run_chat( # noqa: PLR0913
|
|||||||
raise click.ClickException(str(exc)) from exc
|
raise click.ClickException(str(exc)) from exc
|
||||||
|
|
||||||
_ = set_workspace_root(workspace)
|
_ = set_workspace_root(workspace)
|
||||||
|
install_cli_limit_hooks()
|
||||||
memory = await MemoryStore.open(_database_config(store))
|
memory = await MemoryStore.open(_database_config(store))
|
||||||
try:
|
try:
|
||||||
state = ReplState(
|
state = ReplState(
|
||||||
|
|||||||
@@ -48,7 +48,13 @@ async def render_events(events: AsyncIterator[AgentEvent]) -> None:
|
|||||||
)
|
)
|
||||||
click.secho(f"[tool result] {preview}", fg="magenta")
|
click.secho(f"[tool result] {preview}", fg="magenta")
|
||||||
elif isinstance(event, MaxRoundsEvent):
|
elif isinstance(event, MaxRoundsEvent):
|
||||||
click.secho(f"\n[max rounds reached: {event.rounds}]", fg="red")
|
if event.continued:
|
||||||
|
click.secho(
|
||||||
|
f"\n[max rounds {event.rounds} reached — continuing with a higher allowance]",
|
||||||
|
fg="yellow",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
click.secho(f"\n[max rounds reached: {event.rounds}]", fg="red")
|
||||||
else:
|
else:
|
||||||
# AssistantMessageEvent — text already shown via TextDeltaEvent.
|
# AssistantMessageEvent — text already shown via TextDeltaEvent.
|
||||||
_ = event
|
_ = event
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import click
|
||||||
|
|
||||||
|
from plyngent.tools.process.pty_session import PtyManager
|
||||||
|
|
||||||
|
|
||||||
|
def prompt_continue_limit(reason: str) -> bool:
|
||||||
|
"""Ask the user whether to raise a limit and continue (TTY)."""
|
||||||
|
click.echo()
|
||||||
|
click.secho(f"[limit] {reason}", fg="yellow")
|
||||||
|
try:
|
||||||
|
return bool(click.confirm("Raise limit and continue?", default=True))
|
||||||
|
except click.Abort:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def install_cli_limit_hooks() -> None:
|
||||||
|
"""Register interactive continue hooks for process-global tool limits."""
|
||||||
|
PtyManager.set_limit_continue_hook(prompt_continue_limit)
|
||||||
@@ -43,6 +43,8 @@ class ReplState:
|
|||||||
return ToolRegistry(list(DEFAULT_TOOLS))
|
return ToolRegistry(list(DEFAULT_TOOLS))
|
||||||
|
|
||||||
def _make_agent(self) -> ChatAgent:
|
def _make_agent(self) -> ChatAgent:
|
||||||
|
from plyngent.cli.limits import prompt_continue_limit
|
||||||
|
|
||||||
return ChatAgent(
|
return ChatAgent(
|
||||||
self.client,
|
self.client,
|
||||||
model=self.model,
|
model=self.model,
|
||||||
@@ -50,6 +52,7 @@ class ReplState:
|
|||||||
memory=self.memory,
|
memory=self.memory,
|
||||||
session_id=self.session_id,
|
session_id=self.session_id,
|
||||||
max_rounds=self.max_rounds,
|
max_rounds=self.max_rounds,
|
||||||
|
on_limit=prompt_continue_limit,
|
||||||
)
|
)
|
||||||
|
|
||||||
def rebuild_client(self) -> None:
|
def rebuild_client(self) -> None:
|
||||||
|
|||||||
@@ -10,16 +10,23 @@ import signal
|
|||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from threading import Lock
|
from threading import Lock
|
||||||
from typing import ClassVar
|
from typing import TYPE_CHECKING, ClassVar
|
||||||
|
|
||||||
from plyngent.tools.workspace import WorkspaceError, check_command_allowed, resolve_path
|
from plyngent.tools.workspace import WorkspaceError, check_command_allowed, resolve_path
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
|
type LimitContinueHook = Callable[[str], bool]
|
||||||
|
|
||||||
DEFAULT_PTY_READ_BYTES = 8192
|
DEFAULT_PTY_READ_BYTES = 8192
|
||||||
DEFAULT_PTY_POLL_TIMEOUT = 0.2
|
DEFAULT_PTY_POLL_TIMEOUT = 0.2
|
||||||
DEFAULT_MAX_SESSIONS = 8
|
DEFAULT_MAX_SESSIONS = 8
|
||||||
DEFAULT_IDLE_TTL_SECONDS = 600.0
|
DEFAULT_IDLE_TTL_SECONDS = 600.0
|
||||||
DEFAULT_SESSION_OUTPUT_BUDGET = 256_000
|
DEFAULT_SESSION_OUTPUT_BUDGET = 256_000
|
||||||
DEFAULT_CLOSE_GRACE_SECONDS = 0.5
|
DEFAULT_CLOSE_GRACE_SECONDS = 0.5
|
||||||
|
_SESSION_LIMIT_STEP = 4
|
||||||
|
_BUDGET_STEP = 256_000
|
||||||
_STDERR_FD = 2
|
_STDERR_FD = 2
|
||||||
_EXEC_FAIL_MARKER = b"plyngent-pty-exec-failed: "
|
_EXEC_FAIL_MARKER = b"plyngent-pty-exec-failed: "
|
||||||
|
|
||||||
@@ -67,6 +74,7 @@ class PtyManager:
|
|||||||
max_sessions: ClassVar[int] = DEFAULT_MAX_SESSIONS
|
max_sessions: ClassVar[int] = DEFAULT_MAX_SESSIONS
|
||||||
idle_ttl_seconds: ClassVar[float] = DEFAULT_IDLE_TTL_SECONDS
|
idle_ttl_seconds: ClassVar[float] = DEFAULT_IDLE_TTL_SECONDS
|
||||||
session_output_budget: ClassVar[int] = DEFAULT_SESSION_OUTPUT_BUDGET
|
session_output_budget: ClassVar[int] = DEFAULT_SESSION_OUTPUT_BUDGET
|
||||||
|
_limit_continue: ClassVar[LimitContinueHook | None] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def configure(
|
def configure(
|
||||||
@@ -83,6 +91,21 @@ class PtyManager:
|
|||||||
if session_output_budget is not None:
|
if session_output_budget is not None:
|
||||||
cls.session_output_budget = max(1024, session_output_budget)
|
cls.session_output_budget = max(1024, session_output_budget)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def set_limit_continue_hook(cls, hook: LimitContinueHook | None) -> None:
|
||||||
|
"""Optional interactive hook: return True to raise a limit and continue."""
|
||||||
|
cls._limit_continue = hook
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _offer_raise(cls, reason: str) -> bool:
|
||||||
|
hook = cls._limit_continue
|
||||||
|
if hook is None:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
return bool(hook(reason))
|
||||||
|
except Exception: # noqa: BLE001 — never break tools on prompt failure
|
||||||
|
return False
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def open(
|
def open(
|
||||||
cls,
|
cls,
|
||||||
@@ -100,8 +123,12 @@ class PtyManager:
|
|||||||
with cls._lock:
|
with cls._lock:
|
||||||
alive_count = sum(1 for s in cls._sessions.values() if not s.closed)
|
alive_count = sum(1 for s in cls._sessions.values() if not s.closed)
|
||||||
if alive_count >= cls.max_sessions:
|
if alive_count >= cls.max_sessions:
|
||||||
msg = f"PTY session limit reached ({cls.max_sessions}); close idle sessions"
|
reason = f"PTY session limit reached ({cls.max_sessions})"
|
||||||
raise WorkspaceError(msg)
|
if cls._offer_raise(f"{reason}; raise by {_SESSION_LIMIT_STEP}?"):
|
||||||
|
cls.max_sessions += _SESSION_LIMIT_STEP
|
||||||
|
else:
|
||||||
|
msg = f"{reason}; close idle sessions or allow a higher limit"
|
||||||
|
raise WorkspaceError(msg)
|
||||||
|
|
||||||
master_fd, slave_fd = pty.openpty()
|
master_fd, slave_fd = pty.openpty()
|
||||||
pid = os.fork()
|
pid = os.fork()
|
||||||
@@ -255,13 +282,19 @@ class PtyManager:
|
|||||||
raise WorkspaceError(msg)
|
raise WorkspaceError(msg)
|
||||||
|
|
||||||
if (cls.session_output_budget - session.bytes_read) <= 0:
|
if (cls.session_output_budget - session.bytes_read) <= 0:
|
||||||
return PtyReadResult(
|
if cls._offer_raise(
|
||||||
session_id=session_id,
|
f"PTY output budget exhausted for session {session_id} "
|
||||||
alive=session.alive,
|
f"({cls.session_output_budget} bytes); raise by {_BUDGET_STEP}?"
|
||||||
exit_code=session.exit_code,
|
):
|
||||||
data="",
|
cls.session_output_budget += _BUDGET_STEP
|
||||||
budget_exhausted=True,
|
else:
|
||||||
)
|
return PtyReadResult(
|
||||||
|
session_id=session_id,
|
||||||
|
alive=session.alive,
|
||||||
|
exit_code=session.exit_code,
|
||||||
|
data="",
|
||||||
|
budget_exhausted=True,
|
||||||
|
)
|
||||||
|
|
||||||
chunks, matched = cls._collect_chunks(
|
chunks, matched = cls._collect_chunks(
|
||||||
session, max_bytes=max_bytes, timeout=timeout, until=until
|
session, max_bytes=max_bytes, timeout=timeout, until=until
|
||||||
@@ -270,6 +303,13 @@ class PtyManager:
|
|||||||
truncated = len(data) > max_bytes
|
truncated = len(data) > max_bytes
|
||||||
if truncated:
|
if truncated:
|
||||||
data = data[:max_bytes]
|
data = data[:max_bytes]
|
||||||
|
budget_exhausted = (cls.session_output_budget - session.bytes_read) <= 0
|
||||||
|
if budget_exhausted and cls._offer_raise(
|
||||||
|
f"PTY output budget exhausted for session {session_id} "
|
||||||
|
f"({cls.session_output_budget} bytes); raise by {_BUDGET_STEP}?"
|
||||||
|
):
|
||||||
|
cls.session_output_budget += _BUDGET_STEP
|
||||||
|
budget_exhausted = False
|
||||||
return PtyReadResult(
|
return PtyReadResult(
|
||||||
session_id=session_id,
|
session_id=session_id,
|
||||||
alive=session.alive,
|
alive=session.alive,
|
||||||
@@ -277,7 +317,7 @@ class PtyManager:
|
|||||||
data=data,
|
data=data,
|
||||||
truncated=truncated,
|
truncated=truncated,
|
||||||
matched=matched,
|
matched=matched,
|
||||||
budget_exhausted=(cls.session_output_budget - session.bytes_read) <= 0,
|
budget_exhausted=budget_exhausted,
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -155,10 +155,53 @@ async def test_max_rounds() -> None:
|
|||||||
client = ScriptedClient([forever, forever, forever])
|
client = ScriptedClient([forever, forever, forever])
|
||||||
messages: list[AnyChatMessage] = [UserChatMessage(content="x")]
|
messages: list[AnyChatMessage] = [UserChatMessage(content="x")]
|
||||||
events = [e async for e in run_chat_loop(client, messages, model="m", tools=registry, max_rounds=2)]
|
events = [e async for e in run_chat_loop(client, messages, model="m", tools=registry, max_rounds=2)]
|
||||||
assert any(isinstance(e, MaxRoundsEvent) and e.rounds == 2 for e in events) # noqa: PLR2004
|
assert any(isinstance(e, MaxRoundsEvent) and e.rounds == 2 and not e.continued for e in events) # noqa: PLR2004
|
||||||
assert len(client.calls) == 2 # noqa: PLR2004
|
assert len(client.calls) == 2 # noqa: PLR2004
|
||||||
|
|
||||||
|
|
||||||
|
async def test_max_rounds_continue_hook() -> None:
|
||||||
|
@tool
|
||||||
|
def ping() -> str:
|
||||||
|
return "pong"
|
||||||
|
|
||||||
|
registry = ToolRegistry([ping])
|
||||||
|
forever = _response(
|
||||||
|
AssistantChatMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
AssistantFunctionToolCall(
|
||||||
|
id="c",
|
||||||
|
function=AssistantFunctionTool(name="ping", arguments="{}"),
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
final = _response(AssistantChatMessage(content="done"))
|
||||||
|
client = ScriptedClient([forever, forever, final])
|
||||||
|
messages: list[AnyChatMessage] = [UserChatMessage(content="x")]
|
||||||
|
asks: list[str] = []
|
||||||
|
|
||||||
|
def on_limit(reason: str) -> bool:
|
||||||
|
asks.append(reason)
|
||||||
|
return len(asks) == 1
|
||||||
|
|
||||||
|
events = [
|
||||||
|
e
|
||||||
|
async for e in run_chat_loop(
|
||||||
|
client,
|
||||||
|
messages,
|
||||||
|
model="m",
|
||||||
|
tools=registry,
|
||||||
|
max_rounds=2,
|
||||||
|
on_limit=on_limit,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
assert len(asks) == 1
|
||||||
|
assert any(isinstance(e, MaxRoundsEvent) and e.continued for e in events)
|
||||||
|
assert any(isinstance(e, TextDeltaEvent) and e.content == "done" for e in events)
|
||||||
|
assert len(client.calls) == 3 # noqa: PLR2004
|
||||||
|
|
||||||
|
|
||||||
async def test_default_max_rounds_is_generous() -> None:
|
async def test_default_max_rounds_is_generous() -> None:
|
||||||
from plyngent.agent.loop import DEFAULT_MAX_ROUNDS
|
from plyngent.agent.loop import DEFAULT_MAX_ROUNDS
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from plyngent.cli.limits import install_cli_limit_hooks, prompt_continue_limit
|
||||||
|
from plyngent.tools.process.pty_session import PtyManager
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
def test_prompt_continue_limit_yes(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
def _confirm(*_a: object, **_k: object) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
monkeypatch.setattr("click.confirm", _confirm)
|
||||||
|
assert prompt_continue_limit("hit a wall") is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_prompt_continue_limit_no(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
def _confirm(*_a: object, **_k: object) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
monkeypatch.setattr("click.confirm", _confirm)
|
||||||
|
assert prompt_continue_limit("hit a wall") is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_install_cli_limit_hooks() -> None:
|
||||||
|
install_cli_limit_hooks()
|
||||||
|
# Hook is installed process-wide for the CLI session.
|
||||||
|
assert callable(getattr(PtyManager, "_limit_continue", None))
|
||||||
|
PtyManager.set_limit_continue_hook(None)
|
||||||
@@ -147,6 +147,7 @@ def test_pty_session_limit(workspace: object) -> None:
|
|||||||
del workspace
|
del workspace
|
||||||
previous = PtyManager.max_sessions
|
previous = PtyManager.max_sessions
|
||||||
try:
|
try:
|
||||||
|
PtyManager.set_limit_continue_hook(None)
|
||||||
PtyManager.configure(max_sessions=1)
|
PtyManager.configure(max_sessions=1)
|
||||||
first = call_sync(open_pty, ["sleep", "30"])
|
first = call_sync(open_pty, ["sleep", "30"])
|
||||||
assert "session_id=" in first
|
assert "session_id=" in first
|
||||||
@@ -155,6 +156,24 @@ def test_pty_session_limit(workspace: object) -> None:
|
|||||||
finally:
|
finally:
|
||||||
PtyManager.close_all()
|
PtyManager.close_all()
|
||||||
PtyManager.configure(max_sessions=previous)
|
PtyManager.configure(max_sessions=previous)
|
||||||
|
PtyManager.set_limit_continue_hook(None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_pty_session_limit_continue(workspace: object) -> None:
|
||||||
|
del workspace
|
||||||
|
previous = PtyManager.max_sessions
|
||||||
|
try:
|
||||||
|
PtyManager.configure(max_sessions=1)
|
||||||
|
PtyManager.set_limit_continue_hook(lambda _reason: True)
|
||||||
|
first = call_sync(open_pty, ["sleep", "30"])
|
||||||
|
second = call_sync(open_pty, ["sleep", "30"])
|
||||||
|
assert "session_id=" in first
|
||||||
|
assert "session_id=" in second
|
||||||
|
assert PtyManager.max_sessions >= 2 # noqa: PLR2004
|
||||||
|
finally:
|
||||||
|
PtyManager.close_all()
|
||||||
|
PtyManager.configure(max_sessions=previous)
|
||||||
|
PtyManager.set_limit_continue_hook(None)
|
||||||
|
|
||||||
|
|
||||||
def test_pty_output_budget(workspace: object) -> None:
|
def test_pty_output_budget(workspace: object) -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user