core/tools/process: timeout partial output, per-session PTY budget

run_command keeps stdout/stderr on timeout with timed_out=true;
open_pty uses shlex.join; PTY output budget is per-session.
This commit is contained in:
2026-07-14 20:07:07 +08:00
parent a634021a80
commit 4f1e99fa4b
4 changed files with 102 additions and 38 deletions
+3 -1
View File
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
import shlex
from plyngent.agent import tool from plyngent.agent import tool
from plyngent.tools.workspace import WorkspaceError from plyngent.tools.workspace import WorkspaceError
@@ -24,6 +26,6 @@ def open_pty(command: list[str], *, cwd: str = ".") -> str:
f"session_id={session.session_id}", f"session_id={session.session_id}",
"alive=true", "alive=true",
"exit_code=", "exit_code=",
f"cmd={' '.join(session.command)}", f"cmd={shlex.join(session.command)}",
] ]
) )
+20 -17
View File
@@ -42,6 +42,7 @@ class PtySession:
created_at: float = field(default_factory=time.time) created_at: float = field(default_factory=time.time)
last_activity: float = field(default_factory=time.time) last_activity: float = field(default_factory=time.time)
bytes_read: int = 0 bytes_read: int = 0
output_budget: int = DEFAULT_SESSION_OUTPUT_BUDGET
command: tuple[str, ...] = () command: tuple[str, ...] = ()
@@ -157,6 +158,7 @@ class PtyManager:
master_fd=master_fd, master_fd=master_fd,
pid=pid, pid=pid,
command=tuple(command), command=tuple(command),
output_budget=cls.session_output_budget,
) )
cls._sessions[session_id] = session cls._sessions[session_id] = session
return session return session
@@ -200,9 +202,9 @@ class PtyManager:
@classmethod @classmethod
def _read_once(cls, session: PtySession, *, max_bytes: int, timeout: float) -> bytes: def _read_once(cls, session: PtySession, *, max_bytes: int, timeout: float) -> bytes:
if session.closed or (cls.session_output_budget - session.bytes_read) <= 0: if session.closed or (session.output_budget - session.bytes_read) <= 0:
return b"" return b""
to_read = min(max_bytes, cls.session_output_budget - session.bytes_read) to_read = min(max_bytes, session.output_budget - session.bytes_read)
data = b"" data = b""
try: try:
ready, _, _ = select.select([session.master_fd], [], [], timeout) ready, _, _ = select.select([session.master_fd], [], [], timeout)
@@ -250,10 +252,23 @@ class PtyManager:
matched = True matched = True
break break
cls._poll_exit(session) cls._poll_exit(session)
if not session.alive or (cls.session_output_budget - session.bytes_read) <= 0: if not session.alive or (session.output_budget - session.bytes_read) <= 0:
break break
return chunks, matched return chunks, matched
@classmethod
def _maybe_raise_budget(cls, session: PtySession) -> bool:
"""If budget is exhausted, offer to raise this session's ceiling. Returns whether still exhausted."""
if (session.output_budget - session.bytes_read) > 0:
return False
if cls._offer_raise(
f"PTY output budget exhausted for session {session.session_id} "
f"({session.output_budget} bytes); raise by {_BUDGET_STEP}?"
):
session.output_budget += _BUDGET_STEP
return False
return True
@classmethod @classmethod
def read( def read(
cls, cls,
@@ -281,13 +296,7 @@ class PtyManager:
msg = f"closed PTY session: {session_id}" msg = f"closed PTY session: {session_id}"
raise WorkspaceError(msg) raise WorkspaceError(msg)
if (cls.session_output_budget - session.bytes_read) <= 0: if cls._maybe_raise_budget(session):
if 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
else:
return PtyReadResult( return PtyReadResult(
session_id=session_id, session_id=session_id,
alive=session.alive, alive=session.alive,
@@ -303,13 +312,7 @@ 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 budget_exhausted = cls._maybe_raise_budget(session)
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,
+29 -12
View File
@@ -26,16 +26,18 @@ def _truncate(text: str, label: str) -> str:
return text[:DEFAULT_MAX_OUTPUT_CHARS] + f"\n...[{label} truncated]" return text[:DEFAULT_MAX_OUTPUT_CHARS] + f"\n...[{label} truncated]"
def _format_result( def _format_result( # noqa: PLR0913
*, *,
returncode: int | None, returncode: int | None,
workdir_display: str, workdir_display: str,
command: list[str], command: list[str],
stdout: str, stdout: str,
stderr: str, stderr: str,
timed_out: bool = False,
) -> str: ) -> str:
parts = [ parts = [
f"exit_code={returncode}", f"exit_code={'' if returncode is None else returncode}",
f"timed_out={'true' if timed_out else 'false'}",
f"cwd={workdir_display}", f"cwd={workdir_display}",
f"cmd={shlex.join(command)}", f"cmd={shlex.join(command)}",
] ]
@@ -46,13 +48,20 @@ def _format_result(
return "\n".join(parts) return "\n".join(parts)
def _validate_env(env: object) -> str | None:
"""Runtime guard for tool-JSON args (may not match static typing at the boundary)."""
if env is None:
return None
if type(env) is not dict:
return "error: env must be an object of string keys and values"
# Tool schema should already enforce dict[str, str]; keep a light check.
return None
def _merge_env(overrides: dict[str, str] | None) -> dict[str, str] | None: def _merge_env(overrides: dict[str, str] | None) -> dict[str, str] | None:
if overrides is None: if overrides is None:
return None return None
merged = dict(os.environ) return {**os.environ, **overrides}
for key, value in overrides.items():
merged[str(key)] = str(value)
return merged
async def _run_exec( async def _run_exec(
@@ -62,7 +71,8 @@ async def _run_exec(
timeout_seconds: float, timeout_seconds: float,
stdin_data: bytes | None, stdin_data: bytes | None,
env: dict[str, str] | None, env: dict[str, str] | None,
) -> tuple[int | None, str, str] | str: ) -> tuple[int | None, str, str, bool] | str:
"""Return ``(returncode, stdout, stderr, timed_out)`` or an error string."""
stdin = asyncio.subprocess.PIPE if stdin_data is not None else None stdin = asyncio.subprocess.PIPE if stdin_data is not None else None
try: try:
proc = await asyncio.create_subprocess_exec( proc = await asyncio.create_subprocess_exec(
@@ -78,19 +88,20 @@ async def _run_exec(
except OSError as exc: except OSError as exc:
return f"error: failed to start command: {exc}" return f"error: failed to start command: {exc}"
timed_out = False
try: try:
stdout_b, stderr_b = await asyncio.wait_for( stdout_b, stderr_b = await asyncio.wait_for(
proc.communicate(input=stdin_data), proc.communicate(input=stdin_data),
timeout=timeout_seconds, timeout=timeout_seconds,
) )
except TimeoutError: except TimeoutError:
timed_out = True
proc.kill() proc.kill()
_ = await proc.communicate() stdout_b, stderr_b = await proc.communicate()
return f"error: command timed out after {timeout_seconds}s: {shlex.join(command)}"
stdout = _truncate(stdout_b.decode(errors="replace"), "stdout") stdout = _truncate(stdout_b.decode(errors="replace"), "stdout")
stderr = _truncate(stderr_b.decode(errors="replace"), "stderr") stderr = _truncate(stderr_b.decode(errors="replace"), "stderr")
return proc.returncode, stdout, stderr return proc.returncode, stdout, stderr, timed_out
def _validate_run_args( def _validate_run_args(
@@ -112,7 +123,9 @@ def _validate_run_args(
return f"error: cwd is not a directory: {cwd}" return f"error: cwd is not a directory: {cwd}"
if timeout_seconds <= 0: if timeout_seconds <= 0:
return "error: timeout_seconds must be > 0" return "error: timeout_seconds must be > 0"
_ = env # typed as dict[str, str] | None; runtime JSON already constrained by tool schema env_error = _validate_env(env)
if env_error is not None:
return env_error
return workdir return workdir
@@ -130,6 +143,9 @@ async def run_command(
``cwd`` is relative to or under the workspace root. Optional ``stdin`` is ``cwd`` is relative to or under the workspace root. Optional ``stdin`` is
written to the process stdin. Optional ``env`` overlays process environment written to the process stdin. Optional ``env`` overlays process environment
variables (merged with the current environment). Output is truncated. variables (merged with the current environment). Output is truncated.
On timeout the process is killed and any partial stdout/stderr is still
returned with ``timed_out=true``.
""" """
workdir = _validate_run_args(command, cwd=cwd, timeout_seconds=timeout_seconds, env=env) workdir = _validate_run_args(command, cwd=cwd, timeout_seconds=timeout_seconds, env=env)
if isinstance(workdir, str): if isinstance(workdir, str):
@@ -145,11 +161,12 @@ async def run_command(
) )
if isinstance(result, str): if isinstance(result, str):
return result return result
returncode, stdout, stderr = result returncode, stdout, stderr, timed_out = result
return _format_result( return _format_result(
returncode=returncode, returncode=returncode,
workdir_display=str(workdir.relative_to(get_workspace_root())), workdir_display=str(workdir.relative_to(get_workspace_root())),
command=command, command=command,
stdout=stdout, stdout=stdout,
stderr=stderr, stderr=stderr,
timed_out=timed_out,
) )
+43 -1
View File
@@ -51,7 +51,19 @@ async def test_run_command_cwd(workspace: object) -> None:
async def test_run_command_timeout(workspace: object) -> None: async def test_run_command_timeout(workspace: object) -> None:
del workspace del workspace
out = await call_async(run_command, ["sleep", "5"], timeout_seconds=0.2) out = await call_async(run_command, ["sleep", "5"], timeout_seconds=0.2)
assert "timed out" in out assert "timed_out=true" in out
async def test_run_command_timeout_keeps_partial_output(workspace: object) -> None:
del workspace
# Print then sleep past the timeout so communicate has partial stdout after kill.
out = await call_async(
run_command,
["sh", "-c", "printf partial-out; sleep 5"],
timeout_seconds=0.3,
)
assert "timed_out=true" in out
assert "partial-out" in out
async def test_run_command_stdin(workspace: object) -> None: async def test_run_command_stdin(workspace: object) -> None:
@@ -180,6 +192,7 @@ def test_pty_output_budget(workspace: object) -> None:
del workspace del workspace
previous = PtyManager.session_output_budget previous = PtyManager.session_output_budget
try: try:
PtyManager.set_limit_continue_hook(None)
PtyManager.configure(session_output_budget=64) PtyManager.configure(session_output_budget=64)
opened = call_sync(open_pty, ["sh", "-c", "yes x | head -c 1000"]) opened = call_sync(open_pty, ["sh", "-c", "yes x | head -c 1000"])
session_id = _session_id(opened) session_id = _session_id(opened)
@@ -199,3 +212,32 @@ def test_pty_output_budget(workspace: object) -> None:
finally: finally:
PtyManager.close_all() PtyManager.close_all()
PtyManager.configure(session_output_budget=previous) PtyManager.configure(session_output_budget=previous)
PtyManager.set_limit_continue_hook(None)
def test_pty_output_budget_is_per_session(workspace: object) -> None:
del workspace
previous = PtyManager.session_output_budget
try:
# configure clamps budget to >= 1024
PtyManager.configure(session_output_budget=1024)
class_budget = PtyManager.session_output_budget
PtyManager.set_limit_continue_hook(lambda _reason: True)
opened = call_sync(open_pty, ["sh", "-c", "yes x | head -c 200"])
session_id = _session_id(opened)
session = PtyManager.get(session_id)
assert session is not None
before = session.output_budget
# Force budget exhaustion path by setting bytes_read high.
session.bytes_read = session.output_budget
_ = call_sync(read_pty, session_id, timeout=0.1)
session2 = PtyManager.get(session_id)
assert session2 is not None
assert session2.output_budget > before
# Raising is per-session; class default for new sessions stays put.
assert PtyManager.session_output_budget == class_budget
_ = call_sync(close_pty, session_id)
finally:
PtyManager.close_all()
PtyManager.configure(session_output_budget=previous)
PtyManager.set_limit_continue_hook(None)