mirror of
https://github.com/usestrix/strix.git
synced 2026-08-18 09:49:17 +02:00
feat(context): spill oversized tool output into the sandbox workspace (#882)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
This commit is contained in:
co-authored by
Ahmed Allam
parent
66a283b71b
commit
3b79e97f00
@@ -34,7 +34,7 @@ from strix.tools.notes.tools import (
|
||||
list_notes,
|
||||
update_note,
|
||||
)
|
||||
from strix.tools.output_store import bound_text
|
||||
from strix.tools.output_store import bound_and_store, bound_text
|
||||
from strix.tools.proxy.tools import (
|
||||
list_requests,
|
||||
list_sitemap,
|
||||
@@ -115,11 +115,11 @@ def _tool_output_limits() -> tuple[int, int]:
|
||||
return context.tool_output_max_lines, context.tool_output_max_bytes
|
||||
|
||||
|
||||
def _bound_result(result: Any) -> Any:
|
||||
async def _bound_result(result: Any) -> Any:
|
||||
if not isinstance(result, str):
|
||||
return result
|
||||
max_lines, max_bytes = _tool_output_limits()
|
||||
return bound_text(result, max_lines=max_lines, max_bytes=max_bytes)
|
||||
return await bound_and_store(result, max_lines=max_lines, max_bytes=max_bytes)
|
||||
|
||||
|
||||
def _format_tool_error(exc: Exception) -> str:
|
||||
@@ -135,7 +135,7 @@ def _with_bounded_result(tool: FunctionTool) -> FunctionTool:
|
||||
invoke_tool = tool.on_invoke_tool
|
||||
|
||||
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||
return _bound_result(await invoke_tool(ctx, raw_input))
|
||||
return await _bound_result(await invoke_tool(ctx, raw_input))
|
||||
|
||||
tool.on_invoke_tool = invoke
|
||||
tool._strix_bounded = True # type: ignore[attr-defined]
|
||||
@@ -147,7 +147,7 @@ def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
||||
|
||||
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||
try:
|
||||
return _bound_result(await invoke_tool(ctx, raw_input))
|
||||
return await _bound_result(await invoke_tool(ctx, raw_input))
|
||||
except Exception as exc: # noqa: BLE001 - tool errors should be model-visible results.
|
||||
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
||||
return _format_tool_error(exc)
|
||||
@@ -162,7 +162,7 @@ def _custom_tool_as_function_tool(tool: CustomTool) -> FunctionTool:
|
||||
if not custom_input:
|
||||
return f"`{_custom_tool_input_field(tool)}` must be a non-empty string."
|
||||
try:
|
||||
return _bound_result(await tool.on_invoke_tool(ctx, custom_input))
|
||||
return await _bound_result(await tool.on_invoke_tool(ctx, custom_input))
|
||||
except Exception as exc: # noqa: BLE001 - matches SDK CustomTool error-as-result behavior.
|
||||
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
||||
return _format_tool_error(exc)
|
||||
@@ -199,7 +199,7 @@ def _bound_custom_tool(tool: CustomTool) -> CustomTool:
|
||||
invoke_tool = tool.on_invoke_tool
|
||||
|
||||
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||
return _bound_result(await invoke_tool(ctx, raw_input))
|
||||
return await _bound_result(await invoke_tool(ctx, raw_input))
|
||||
|
||||
tool.on_invoke_tool = invoke
|
||||
return tool
|
||||
|
||||
@@ -3,10 +3,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agents import RunConfig
|
||||
@@ -40,6 +42,10 @@ from strix.core.paths import run_dir_for, runtime_state_dir
|
||||
from strix.core.sessions import open_agent_session
|
||||
from strix.runtime import session_manager
|
||||
from strix.telemetry.logging import set_scan_id, setup_scan_logging
|
||||
from strix.tools.output_store import (
|
||||
WORKSPACE_SPILL_DIR,
|
||||
configure_spill_writer,
|
||||
)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -203,6 +209,20 @@ async def run_strix_scan(
|
||||
)
|
||||
logger.info("Sandbox ready for scan %s", scan_id)
|
||||
|
||||
sandbox_session = bundle["session"]
|
||||
|
||||
async def _spill_to_workspace(output_id: str, text: str) -> str | None:
|
||||
"""Write an oversized tool result into the sandbox; return its path or None."""
|
||||
path = f"{WORKSPACE_SPILL_DIR}/{output_id}.txt"
|
||||
try:
|
||||
await sandbox_session.write(Path(path), io.BytesIO(text.encode("utf-8")))
|
||||
except Exception:
|
||||
logger.exception("failed to spill tool output to sandbox workspace")
|
||||
return None
|
||||
return path
|
||||
|
||||
configure_spill_writer(_spill_to_workspace)
|
||||
|
||||
sessions_to_close: list[SQLiteSession] = []
|
||||
|
||||
try:
|
||||
@@ -399,6 +419,7 @@ async def run_strix_scan(
|
||||
await coordinator.set_status(root_id, "failed")
|
||||
raise
|
||||
finally:
|
||||
configure_spill_writer(None)
|
||||
for s in sessions_to_close:
|
||||
with contextlib.suppress(Exception):
|
||||
s.close()
|
||||
|
||||
+103
-11
@@ -1,13 +1,47 @@
|
||||
"""Bound oversized tool results before they enter agent history.
|
||||
|
||||
Keeps a head + tail slice and drops the middle, replacing it with a notice of
|
||||
how much was removed.
|
||||
Oversized results are spilled into the sandbox at
|
||||
``/workspace/.strix/tool-output/<id>.txt``; the agent sees a head + tail slice
|
||||
plus the path and reads the rest back with its own file tools. The spill writer
|
||||
is injected by the runner via :func:`configure_spill_writer`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TRUNCATION_NOTICE = "[... {lines} lines ({bytes} bytes) truncated ...]"
|
||||
_WORKSPACE_SPILL_NOTICE = (
|
||||
"[... {lines} lines ({bytes} bytes) truncated — full output saved to {path} "
|
||||
"in the sandbox; read it with exec_command (e.g. `sed -n`, `grep`, `cat`) ...]"
|
||||
)
|
||||
|
||||
WORKSPACE_SPILL_DIR = "/workspace/.strix/tool-output"
|
||||
|
||||
# Longest possible workspace path, used only to reserve notice bytes.
|
||||
_SAMPLE_WORKSPACE_PATH = f"{WORKSPACE_SPILL_DIR}/{'0' * 32}.txt"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
SpillWriter = Callable[[str, str], Awaitable[str | None]]
|
||||
|
||||
_spill: dict[str, SpillWriter] = {}
|
||||
|
||||
|
||||
def configure_spill_writer(writer: SpillWriter | None) -> None:
|
||||
"""Install (or clear) the sandbox-workspace spill writer."""
|
||||
if writer is None:
|
||||
_spill.pop("writer", None)
|
||||
else:
|
||||
_spill["writer"] = writer
|
||||
|
||||
|
||||
def _byte_len(text: str) -> int:
|
||||
@@ -39,20 +73,37 @@ def _take_suffix(text: str, max_bytes: int) -> str:
|
||||
return "".join(out)
|
||||
|
||||
|
||||
def bound_text(text: str, *, max_lines: int, max_bytes: int) -> str:
|
||||
"""Return ``text`` unchanged when small, else a head+tail preview.
|
||||
def _head_tail(
|
||||
text: str,
|
||||
max_lines: int,
|
||||
max_bytes: int,
|
||||
*,
|
||||
notice_templates: tuple[str, ...] = (_TRUNCATION_NOTICE,),
|
||||
) -> tuple[str, str, int, int] | None:
|
||||
"""Head/tail slices plus dropped line/byte counts, or ``None`` if small.
|
||||
|
||||
Truncates on whichever limit is hit first (line count or UTF-8 byte size).
|
||||
``max_bytes`` bounds the entire joined result, notice and separators
|
||||
included.
|
||||
``max_bytes`` bounds the entire joined result; the largest of
|
||||
``notice_templates`` (plus separators) is reserved before slicing.
|
||||
"""
|
||||
lines = text.split("\n")
|
||||
total_bytes = _byte_len(text)
|
||||
if len(lines) <= max_lines and total_bytes <= max_bytes:
|
||||
return text
|
||||
return None
|
||||
|
||||
# Reserve notice + separator bytes up front; ``+ 4`` covers the two "\n\n".
|
||||
notice_overhead = _byte_len(_TRUNCATION_NOTICE.format(lines=len(lines), bytes=total_bytes)) + 4
|
||||
# Reserve using the largest counts/path; ``+ 4`` covers the two "\n\n".
|
||||
notice_overhead = (
|
||||
max(
|
||||
_byte_len(
|
||||
template.format(
|
||||
lines=len(lines),
|
||||
bytes=total_bytes,
|
||||
path=_SAMPLE_WORKSPACE_PATH,
|
||||
)
|
||||
)
|
||||
for template in notice_templates
|
||||
)
|
||||
+ 4
|
||||
)
|
||||
byte_budget = max(2, max_bytes - notice_overhead)
|
||||
|
||||
head_lines = max(1, max_lines // 2)
|
||||
@@ -70,5 +121,46 @@ def bound_text(text: str, *, max_lines: int, max_bytes: int) -> str:
|
||||
kept_lines = len(head.split("\n")) + (len(tail.split("\n")) if tail else 0)
|
||||
dropped_lines = max(0, len(lines) - kept_lines)
|
||||
dropped_bytes = max(0, total_bytes - _byte_len(head) - _byte_len(tail))
|
||||
notice = _TRUNCATION_NOTICE.format(lines=dropped_lines, bytes=dropped_bytes)
|
||||
return head, tail, dropped_lines, dropped_bytes
|
||||
|
||||
|
||||
def _join(head: str, tail: str, notice: str) -> str:
|
||||
return f"{head}\n\n{notice}\n\n{tail}" if tail else f"{head}\n\n{notice}"
|
||||
|
||||
|
||||
def bound_text(text: str, *, max_lines: int, max_bytes: int) -> str:
|
||||
"""Return ``text`` unchanged when small, else a head+tail preview.
|
||||
|
||||
Nothing is persisted; use :func:`bound_and_store` to keep the full output.
|
||||
"""
|
||||
parts = _head_tail(text, max_lines, max_bytes)
|
||||
if parts is None:
|
||||
return text
|
||||
head, tail, dropped_lines, dropped_bytes = parts
|
||||
return _join(head, tail, _TRUNCATION_NOTICE.format(lines=dropped_lines, bytes=dropped_bytes))
|
||||
|
||||
|
||||
async def bound_and_store(text: str, *, max_lines: int, max_bytes: int) -> str:
|
||||
"""Like :func:`bound_text`, but spill the full output into the sandbox and
|
||||
point the agent at its path. Degrades to a plain preview if the spill fails.
|
||||
"""
|
||||
parts = _head_tail(
|
||||
text,
|
||||
max_lines,
|
||||
max_bytes,
|
||||
notice_templates=(_WORKSPACE_SPILL_NOTICE, _TRUNCATION_NOTICE),
|
||||
)
|
||||
if parts is None:
|
||||
return text
|
||||
head, tail, dropped_lines, dropped_bytes = parts
|
||||
|
||||
writer = _spill.get("writer")
|
||||
if writer is not None:
|
||||
path = await writer(uuid.uuid4().hex, text)
|
||||
if path is not None:
|
||||
notice = _WORKSPACE_SPILL_NOTICE.format(
|
||||
lines=dropped_lines, bytes=dropped_bytes, path=path
|
||||
)
|
||||
return _join(head, tail, notice)
|
||||
|
||||
return _join(head, tail, _TRUNCATION_NOTICE.format(lines=dropped_lines, bytes=dropped_bytes))
|
||||
|
||||
@@ -108,3 +108,10 @@ async def test_chat_completions_filesystem_custom_tool_becomes_function_tool() -
|
||||
factory._configure_filesystem_tools(toolset, chat_completions=True)
|
||||
|
||||
assert isinstance(toolset.read_file, FunctionTool)
|
||||
|
||||
|
||||
def test_function_tools_are_result_bounded() -> None:
|
||||
agent = factory.build_strix_agent(is_root=True)
|
||||
by_name = {t.name: t for t in agent.tools}
|
||||
|
||||
assert getattr(by_name["think"], "_strix_bounded", False) is True
|
||||
|
||||
@@ -1,10 +1,22 @@
|
||||
"""Tests for per-tool-output bounding."""
|
||||
"""Tests for per-tool-output bounding and sandbox-workspace spill."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
from strix.tools.output_store import bound_text
|
||||
import pytest
|
||||
|
||||
from strix.tools.output_store import (
|
||||
WORKSPACE_SPILL_DIR,
|
||||
bound_and_store,
|
||||
bound_text,
|
||||
configure_spill_writer,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_spill_writer() -> None:
|
||||
configure_spill_writer(None)
|
||||
|
||||
|
||||
def test_small_output_passes_through_unchanged() -> None:
|
||||
@@ -58,3 +70,87 @@ def test_dropped_line_count_accounts_for_byte_trimming() -> None:
|
||||
kept = [ln for ln in bounded.splitlines() if ln and "truncated" not in ln]
|
||||
assert dropped == 200 - len(kept)
|
||||
assert dropped > 200 - 20
|
||||
|
||||
|
||||
async def test_bound_and_store_small_output_not_spilled() -> None:
|
||||
written: dict[str, str] = {}
|
||||
|
||||
async def writer(output_id: str, text: str) -> str | None:
|
||||
written[output_id] = text
|
||||
return f"{WORKSPACE_SPILL_DIR}/{output_id}.txt"
|
||||
|
||||
configure_spill_writer(writer)
|
||||
text = "just a few lines\nsecond line"
|
||||
assert await bound_and_store(text, max_lines=100, max_bytes=10_000) == text
|
||||
assert written == {}
|
||||
|
||||
|
||||
async def test_bound_and_store_spills_full_output_to_workspace() -> None:
|
||||
written: dict[str, str] = {}
|
||||
|
||||
async def writer(output_id: str, text: str) -> str | None:
|
||||
written[output_id] = text
|
||||
return f"{WORKSPACE_SPILL_DIR}/{output_id}.txt"
|
||||
|
||||
configure_spill_writer(writer)
|
||||
text = "\n".join(f"secret-line-{i}" for i in range(1000))
|
||||
bounded = await bound_and_store(text, max_lines=10, max_bytes=1_000_000)
|
||||
|
||||
assert WORKSPACE_SPILL_DIR in bounded
|
||||
assert "exec_command" in bounded
|
||||
assert "read_tool_output" not in bounded
|
||||
assert len(bounded.encode("utf-8")) <= 1_000_000
|
||||
assert list(written.values()) == [text]
|
||||
stored = next(iter(written.values()))
|
||||
assert stored.splitlines() == text.splitlines()
|
||||
# A buried line elided from the preview is still present in the spilled file.
|
||||
assert "secret-line-500" not in bounded
|
||||
assert "secret-line-500" in stored
|
||||
|
||||
|
||||
async def test_workspace_notice_carries_the_returned_path() -> None:
|
||||
async def writer(output_id: str, _text: str) -> str | None:
|
||||
return f"{WORKSPACE_SPILL_DIR}/{output_id}.txt"
|
||||
|
||||
configure_spill_writer(writer)
|
||||
text = "\n".join(f"line-{i}" for i in range(1000))
|
||||
bounded = await bound_and_store(text, max_lines=10, max_bytes=1_000_000)
|
||||
|
||||
match = re.search(rf"{re.escape(WORKSPACE_SPILL_DIR)}/([0-9a-f]{{32}})\.txt", bounded)
|
||||
assert match is not None, bounded
|
||||
|
||||
|
||||
async def test_no_writer_degrades_to_plain_preview() -> None:
|
||||
text = "\n".join(f"line-{i}" for i in range(1000))
|
||||
bounded = await bound_and_store(text, max_lines=10, max_bytes=1_000_000)
|
||||
|
||||
assert "truncated" in bounded
|
||||
assert WORKSPACE_SPILL_DIR not in bounded
|
||||
assert "read_tool_output" not in bounded
|
||||
assert len(bounded.encode("utf-8")) <= 1_000_000
|
||||
|
||||
|
||||
async def test_writer_failure_degrades_to_plain_preview() -> None:
|
||||
async def failing_writer(_output_id: str, _text: str) -> str | None:
|
||||
return None
|
||||
|
||||
configure_spill_writer(failing_writer)
|
||||
text = "\n".join(f"line-{i}" for i in range(1000))
|
||||
bounded = await bound_and_store(text, max_lines=10, max_bytes=1_000_000)
|
||||
|
||||
assert "truncated" in bounded
|
||||
assert WORKSPACE_SPILL_DIR not in bounded
|
||||
assert len(bounded.encode("utf-8")) <= 1_000_000
|
||||
|
||||
|
||||
async def test_workspace_preview_honours_byte_budget() -> None:
|
||||
# The workspace notice is longer than a plain notice; the preview reserves for it.
|
||||
async def writer(output_id: str, _text: str) -> str | None:
|
||||
return f"{WORKSPACE_SPILL_DIR}/{output_id}.txt"
|
||||
|
||||
configure_spill_writer(writer)
|
||||
text = "\n".join("x" * 500 for _ in range(200))
|
||||
bounded = await bound_and_store(text, max_lines=2_000, max_bytes=2_000)
|
||||
|
||||
assert WORKSPACE_SPILL_DIR in bounded
|
||||
assert len(bounded.encode("utf-8")) <= 2_000
|
||||
|
||||
Reference in New Issue
Block a user