mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 09:26:39 +02:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
10376b412b | ||
|
|
95046a6cea | ||
|
|
aac59de1e5 | ||
|
|
ce358aa879 | ||
|
|
dab93bcc12 | ||
|
|
3b8980c47b |
+83
-13
@@ -16,6 +16,7 @@ from agents.tool import CustomTool, FunctionTool, Tool
|
|||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
|
|
||||||
from strix.agents.prompt import render_system_prompt
|
from strix.agents.prompt import render_system_prompt
|
||||||
|
from strix.config import load_settings
|
||||||
from strix.tools.agents_graph.tools import (
|
from strix.tools.agents_graph.tools import (
|
||||||
agent_finish,
|
agent_finish,
|
||||||
create_agent,
|
create_agent,
|
||||||
@@ -33,6 +34,7 @@ from strix.tools.notes.tools import (
|
|||||||
list_notes,
|
list_notes,
|
||||||
update_note,
|
update_note,
|
||||||
)
|
)
|
||||||
|
from strix.tools.output_store import bound_text
|
||||||
from strix.tools.proxy.tools import (
|
from strix.tools.proxy.tools import (
|
||||||
list_requests,
|
list_requests,
|
||||||
list_sitemap,
|
list_sitemap,
|
||||||
@@ -103,8 +105,36 @@ def _extract_custom_input(tool: CustomTool, raw_input: str | dict[str, Any]) ->
|
|||||||
return value if isinstance(value, str) else ""
|
return value if isinstance(value, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_output_limits() -> tuple[int, int]:
|
||||||
|
context = load_settings().context
|
||||||
|
return context.tool_output_max_lines, context.tool_output_max_bytes
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
def _format_tool_error(exc: Exception) -> str:
|
def _format_tool_error(exc: Exception) -> str:
|
||||||
return str(exc) or exc.__class__.__name__
|
message = str(exc) or exc.__class__.__name__
|
||||||
|
max_lines, max_bytes = _tool_output_limits()
|
||||||
|
return bound_text(message, max_lines=max_lines, max_bytes=max_bytes)
|
||||||
|
|
||||||
|
|
||||||
|
def _with_bounded_result(tool: FunctionTool) -> FunctionTool:
|
||||||
|
"""Cap a tool's result size before it enters history (idempotent)."""
|
||||||
|
if getattr(tool, "_strix_bounded", False):
|
||||||
|
return tool
|
||||||
|
invoke_tool = tool.on_invoke_tool
|
||||||
|
|
||||||
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
|
return _bound_result(await invoke_tool(ctx, raw_input))
|
||||||
|
|
||||||
|
tool.on_invoke_tool = invoke
|
||||||
|
tool._strix_bounded = True # type: ignore[attr-defined]
|
||||||
|
return tool
|
||||||
|
|
||||||
|
|
||||||
def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
||||||
@@ -112,7 +142,7 @@ def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
|||||||
|
|
||||||
async def invoke(ctx: Any, raw_input: str) -> Any:
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
try:
|
try:
|
||||||
return await invoke_tool(ctx, raw_input)
|
return _bound_result(await invoke_tool(ctx, raw_input))
|
||||||
except Exception as exc: # noqa: BLE001 - tool errors should be model-visible results.
|
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)
|
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
||||||
return _format_tool_error(exc)
|
return _format_tool_error(exc)
|
||||||
@@ -127,7 +157,7 @@ def _custom_tool_as_function_tool(tool: CustomTool) -> FunctionTool:
|
|||||||
if not custom_input:
|
if not custom_input:
|
||||||
return f"`{_custom_tool_input_field(tool)}` must be a non-empty string."
|
return f"`{_custom_tool_input_field(tool)}` must be a non-empty string."
|
||||||
try:
|
try:
|
||||||
return await tool.on_invoke_tool(ctx, custom_input)
|
return _bound_result(await tool.on_invoke_tool(ctx, custom_input))
|
||||||
except Exception as exc: # noqa: BLE001 - matches SDK CustomTool error-as-result behavior.
|
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)
|
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
||||||
return _format_tool_error(exc)
|
return _format_tool_error(exc)
|
||||||
@@ -159,12 +189,35 @@ def _custom_tool_as_function_tool(tool: CustomTool) -> FunctionTool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _configure_chat_completions_filesystem_tools(toolset: Any) -> None:
|
def _bound_custom_tool(tool: CustomTool) -> CustomTool:
|
||||||
|
"""Bound a native ``CustomTool`` result in place (Responses path)."""
|
||||||
|
invoke_tool = tool.on_invoke_tool
|
||||||
|
|
||||||
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
|
return _bound_result(await invoke_tool(ctx, raw_input))
|
||||||
|
|
||||||
|
tool.on_invoke_tool = invoke
|
||||||
|
return tool
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_filesystem_tools(toolset: Any, *, chat_completions: bool) -> None:
|
||||||
for name, tool in vars(toolset).items():
|
for name, tool in vars(toolset).items():
|
||||||
if isinstance(tool, CustomTool):
|
if chat_completions:
|
||||||
setattr(toolset, name, _custom_tool_as_function_tool(tool))
|
if isinstance(tool, CustomTool):
|
||||||
|
setattr(toolset, name, _custom_tool_as_function_tool(tool))
|
||||||
|
elif isinstance(tool, FunctionTool):
|
||||||
|
setattr(toolset, name, _function_tool_with_error_result(tool))
|
||||||
|
elif isinstance(tool, CustomTool):
|
||||||
|
setattr(toolset, name, _bound_custom_tool(tool))
|
||||||
elif isinstance(tool, FunctionTool):
|
elif isinstance(tool, FunctionTool):
|
||||||
setattr(toolset, name, _function_tool_with_error_result(tool))
|
setattr(toolset, name, _with_bounded_result(tool))
|
||||||
|
|
||||||
|
|
||||||
|
def _make_filesystem_configurator(*, chat_completions: bool) -> Any:
|
||||||
|
def configure(toolset: Any) -> None:
|
||||||
|
_configure_filesystem_tools(toolset, chat_completions=chat_completions)
|
||||||
|
|
||||||
|
return configure
|
||||||
|
|
||||||
|
|
||||||
_CHARS_ESCAPE_RE = re.compile(r"\\(?:u[0-9a-fA-F]{4}|x[0-9a-fA-F]{2}|[0abtnvfr\\])")
|
_CHARS_ESCAPE_RE = re.compile(r"\\(?:u[0-9a-fA-F]{4}|x[0-9a-fA-F]{2}|[0abtnvfr\\])")
|
||||||
@@ -205,6 +258,16 @@ def _format_validation_error(tool_name: str, exc: ValidationError) -> str:
|
|||||||
return f"{tool_name}: invalid arguments — " + "; ".join(parts)
|
return f"{tool_name}: invalid arguments — " + "; ".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_shell_output_cap(parsed: dict[str, Any]) -> None:
|
||||||
|
"""Clamp the SDK shell tools' ``max_output_tokens`` to the configured
|
||||||
|
ceiling; a smaller explicit value is respected."""
|
||||||
|
ceiling = load_settings().context.tool_output_max_tokens
|
||||||
|
requested = parsed.get("max_output_tokens")
|
||||||
|
parsed["max_output_tokens"] = (
|
||||||
|
ceiling if not isinstance(requested, int) or requested > ceiling else requested
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _wrap_exec_command(tool: FunctionTool) -> FunctionTool:
|
def _wrap_exec_command(tool: FunctionTool) -> FunctionTool:
|
||||||
invoke_tool = tool.on_invoke_tool
|
invoke_tool = tool.on_invoke_tool
|
||||||
|
|
||||||
@@ -213,8 +276,10 @@ def _wrap_exec_command(tool: FunctionTool) -> FunctionTool:
|
|||||||
parsed = json.loads(raw_input)
|
parsed = json.loads(raw_input)
|
||||||
except (json.JSONDecodeError, TypeError):
|
except (json.JSONDecodeError, TypeError):
|
||||||
parsed = None
|
parsed = None
|
||||||
if isinstance(parsed, dict) and "shell" not in parsed:
|
if isinstance(parsed, dict):
|
||||||
parsed["shell"] = "bash"
|
if "shell" not in parsed:
|
||||||
|
parsed["shell"] = "bash"
|
||||||
|
_apply_shell_output_cap(parsed)
|
||||||
raw_input = json.dumps(parsed)
|
raw_input = json.dumps(parsed)
|
||||||
try:
|
try:
|
||||||
return await invoke_tool(ctx, raw_input)
|
return await invoke_tool(ctx, raw_input)
|
||||||
@@ -240,8 +305,10 @@ def _wrap_write_stdin(tool: FunctionTool) -> FunctionTool:
|
|||||||
parsed = json.loads(raw_input)
|
parsed = json.loads(raw_input)
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
parsed = None
|
parsed = None
|
||||||
if isinstance(parsed, dict) and isinstance(parsed.get("chars"), str):
|
if isinstance(parsed, dict):
|
||||||
parsed["chars"] = _decode_chars_escape(parsed["chars"])
|
if isinstance(parsed.get("chars"), str):
|
||||||
|
parsed["chars"] = _decode_chars_escape(parsed["chars"])
|
||||||
|
_apply_shell_output_cap(parsed)
|
||||||
raw_input = json.dumps(parsed)
|
raw_input = json.dumps(parsed)
|
||||||
try:
|
try:
|
||||||
return await invoke_tool(ctx, raw_input)
|
return await invoke_tool(ctx, raw_input)
|
||||||
@@ -440,6 +507,9 @@ def build_strix_agent(
|
|||||||
else:
|
else:
|
||||||
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
|
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
|
||||||
_ensure_unique_tool_names(tools)
|
_ensure_unique_tool_names(tools)
|
||||||
|
tools = [
|
||||||
|
_with_bounded_result(tool) if isinstance(tool, FunctionTool) else tool for tool in tools
|
||||||
|
]
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Built %s agent '%s' (skills=%d, tools=%d, scan_mode=%s, whitebox=%s)",
|
"Built %s agent '%s' (skills=%d, tools=%d, scan_mode=%s, whitebox=%s)",
|
||||||
@@ -459,8 +529,8 @@ def build_strix_agent(
|
|||||||
model=None,
|
model=None,
|
||||||
capabilities=[
|
capabilities=[
|
||||||
Filesystem(
|
Filesystem(
|
||||||
configure_tools=(
|
configure_tools=_make_filesystem_configurator(
|
||||||
_configure_chat_completions_filesystem_tools if chat_completions_tools else None
|
chat_completions=chat_completions_tools,
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
Shell(
|
Shell(
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from strix.config.loader import (
|
|||||||
persist_current,
|
persist_current,
|
||||||
)
|
)
|
||||||
from strix.config.settings import (
|
from strix.config.settings import (
|
||||||
|
ContextSettings,
|
||||||
DedupeSettings,
|
DedupeSettings,
|
||||||
IntegrationSettings,
|
IntegrationSettings,
|
||||||
LlmSettings,
|
LlmSettings,
|
||||||
@@ -27,6 +28,7 @@ from strix.config.settings import (
|
|||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"ContextSettings",
|
||||||
"DedupeSettings",
|
"DedupeSettings",
|
||||||
"IntegrationSettings",
|
"IntegrationSettings",
|
||||||
"LlmSettings",
|
"LlmSettings",
|
||||||
|
|||||||
@@ -55,6 +55,26 @@ class DedupeSettings(BaseSettings):
|
|||||||
api_base: str | None = Field(default=None, alias="DEDUPE_LLM_API_BASE")
|
api_base: str | None = Field(default=None, alias="DEDUPE_LLM_API_BASE")
|
||||||
|
|
||||||
|
|
||||||
|
class ContextSettings(BaseSettings):
|
||||||
|
"""Context-window management: per-tool-output caps and history compaction."""
|
||||||
|
|
||||||
|
model_config = _BASE_CONFIG
|
||||||
|
|
||||||
|
auto_compact: bool = Field(default=True, alias="STRIX_CONTEXT_AUTO_COMPACT")
|
||||||
|
compact_buffer_tokens: int = Field(default=20_000, gt=0, alias="STRIX_CONTEXT_BUFFER_TOKENS")
|
||||||
|
keep_tokens: int = Field(default=8_000, gt=0, alias="STRIX_CONTEXT_KEEP_TOKENS")
|
||||||
|
fallback_context_tokens: int = Field(
|
||||||
|
default=200_000, gt=0, alias="STRIX_CONTEXT_FALLBACK_TOKENS"
|
||||||
|
)
|
||||||
|
summary_max_tokens: int = Field(default=4_096, gt=0, alias="STRIX_CONTEXT_SUMMARY_TOKENS")
|
||||||
|
tool_output_max_tokens: int = Field(default=8_000, gt=0, alias="STRIX_TOOL_OUTPUT_MAX_TOKENS")
|
||||||
|
tool_output_max_lines: int = Field(default=2_000, gt=0, alias="STRIX_TOOL_OUTPUT_MAX_LINES")
|
||||||
|
# Floor above the truncation-notice size so a preview always fits.
|
||||||
|
tool_output_max_bytes: int = Field(
|
||||||
|
default=50 * 1024, ge=1024, alias="STRIX_TOOL_OUTPUT_MAX_BYTES"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class RuntimeSettings(BaseSettings):
|
class RuntimeSettings(BaseSettings):
|
||||||
model_config = _BASE_CONFIG
|
model_config = _BASE_CONFIG
|
||||||
|
|
||||||
@@ -99,6 +119,7 @@ class Settings(BaseSettings):
|
|||||||
llm: LlmSettings = Field(default_factory=LlmSettings)
|
llm: LlmSettings = Field(default_factory=LlmSettings)
|
||||||
dedupe: DedupeSettings = Field(default_factory=DedupeSettings)
|
dedupe: DedupeSettings = Field(default_factory=DedupeSettings)
|
||||||
runtime: RuntimeSettings = Field(default_factory=RuntimeSettings)
|
runtime: RuntimeSettings = Field(default_factory=RuntimeSettings)
|
||||||
|
context: ContextSettings = Field(default_factory=ContextSettings)
|
||||||
telemetry: TelemetrySettings = Field(default_factory=TelemetrySettings)
|
telemetry: TelemetrySettings = Field(default_factory=TelemetrySettings)
|
||||||
integrations: IntegrationSettings = Field(default_factory=IntegrationSettings)
|
integrations: IntegrationSettings = Field(default_factory=IntegrationSettings)
|
||||||
viewer: ViewerSettings = Field(default_factory=ViewerSettings)
|
viewer: ViewerSettings = Field(default_factory=ViewerSettings)
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ import re
|
|||||||
import tempfile
|
import tempfile
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any, cast
|
||||||
|
|
||||||
from pygments.lexers import PythonLexer, get_lexer_by_name, guess_lexer
|
from pygments.lexers import PythonLexer, get_lexer_by_name, guess_lexer
|
||||||
from pygments.lexers.special import TextLexer
|
from pygments.lexers.special import TextLexer
|
||||||
@@ -74,10 +74,10 @@ def resolve_lexer(language: str | None, code: str) -> Lexer:
|
|||||||
try:
|
try:
|
||||||
lexer = guess_lexer(code)
|
lexer = guess_lexer(code)
|
||||||
except ClassNotFound:
|
except ClassNotFound:
|
||||||
return PythonLexer()
|
return cast("Lexer", PythonLexer())
|
||||||
# ``guess_lexer`` returns the plain-text lexer when it can't detect anything.
|
# ``guess_lexer`` returns the plain-text lexer when it can't detect anything.
|
||||||
if isinstance(lexer, TextLexer):
|
if isinstance(lexer, TextLexer):
|
||||||
return PythonLexer()
|
return cast("Lexer", PythonLexer())
|
||||||
return lexer
|
return lexer
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
|
_TRUNCATION_NOTICE = "[... {lines} lines ({bytes} bytes) truncated ...]"
|
||||||
|
|
||||||
|
|
||||||
|
def _byte_len(text: str) -> int:
|
||||||
|
return len(text.encode("utf-8"))
|
||||||
|
|
||||||
|
|
||||||
|
def _take_prefix(text: str, max_bytes: int) -> str:
|
||||||
|
budget = 0
|
||||||
|
out: list[str] = []
|
||||||
|
for char in text:
|
||||||
|
size = len(char.encode("utf-8"))
|
||||||
|
if budget + size > max_bytes:
|
||||||
|
break
|
||||||
|
out.append(char)
|
||||||
|
budget += size
|
||||||
|
return "".join(out)
|
||||||
|
|
||||||
|
|
||||||
|
def _take_suffix(text: str, max_bytes: int) -> str:
|
||||||
|
budget = 0
|
||||||
|
out: list[str] = []
|
||||||
|
for char in reversed(text):
|
||||||
|
size = len(char.encode("utf-8"))
|
||||||
|
if budget + size > max_bytes:
|
||||||
|
break
|
||||||
|
out.append(char)
|
||||||
|
budget += size
|
||||||
|
out.reverse()
|
||||||
|
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.
|
||||||
|
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
lines = text.split("\n")
|
||||||
|
total_bytes = _byte_len(text)
|
||||||
|
if len(lines) <= max_lines and total_bytes <= max_bytes:
|
||||||
|
return text
|
||||||
|
|
||||||
|
# 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
|
||||||
|
byte_budget = max(2, max_bytes - notice_overhead)
|
||||||
|
|
||||||
|
head_lines = max(1, max_lines // 2)
|
||||||
|
tail_lines = max_lines - head_lines
|
||||||
|
head = "\n".join(lines[:head_lines])
|
||||||
|
tail = "\n".join(lines[len(lines) - tail_lines :]) if tail_lines > 0 else ""
|
||||||
|
|
||||||
|
half_bytes = max(1, byte_budget // 2)
|
||||||
|
if _byte_len(head) > half_bytes:
|
||||||
|
head = _take_prefix(head, half_bytes)
|
||||||
|
if tail and _byte_len(tail) > half_bytes:
|
||||||
|
tail = _take_suffix(tail, half_bytes)
|
||||||
|
|
||||||
|
# Count from the final slices; the byte pass may have dropped whole lines.
|
||||||
|
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 f"{head}\n\n{notice}\n\n{tail}" if tail else f"{head}\n\n{notice}"
|
||||||
@@ -3,12 +3,14 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
from types import SimpleNamespace
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from agents.tool import FunctionTool
|
from agents.tool import CustomTool, FunctionTool
|
||||||
|
|
||||||
from strix.agents import factory
|
from strix.agents import factory
|
||||||
|
from strix.config import load_settings
|
||||||
|
|
||||||
|
|
||||||
def _capturing_exec_tool(captured: dict[str, str]) -> FunctionTool:
|
def _capturing_exec_tool(captured: dict[str, str]) -> FunctionTool:
|
||||||
@@ -32,10 +34,37 @@ async def test_wrap_exec_command_defaults_shell_to_bash() -> None:
|
|||||||
result = await wrapped.on_invoke_tool(cast("Any", None), json.dumps({"cmd": "source /tmp/env"}))
|
result = await wrapped.on_invoke_tool(cast("Any", None), json.dumps({"cmd": "source /tmp/env"}))
|
||||||
|
|
||||||
assert result == "ok"
|
assert result == "ok"
|
||||||
assert json.loads(captured["raw_input"]) == {
|
parsed = json.loads(captured["raw_input"])
|
||||||
"cmd": "source /tmp/env",
|
assert parsed["cmd"] == "source /tmp/env"
|
||||||
"shell": "bash",
|
assert parsed["shell"] == "bash"
|
||||||
}
|
expected_cap = load_settings().context.tool_output_max_tokens
|
||||||
|
assert parsed["max_output_tokens"] == expected_cap
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrap_exec_command_preserves_smaller_explicit_output_cap() -> None:
|
||||||
|
captured: dict[str, str] = {}
|
||||||
|
wrapped = factory._wrap_exec_command(_capturing_exec_tool(captured))
|
||||||
|
|
||||||
|
await wrapped.on_invoke_tool(
|
||||||
|
cast("Any", None), json.dumps({"cmd": "echo hi", "max_output_tokens": 42})
|
||||||
|
)
|
||||||
|
|
||||||
|
assert json.loads(captured["raw_input"])["max_output_tokens"] == 42
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrap_exec_command_clamps_oversized_explicit_output_cap() -> None:
|
||||||
|
captured: dict[str, str] = {}
|
||||||
|
wrapped = factory._wrap_exec_command(_capturing_exec_tool(captured))
|
||||||
|
ceiling = load_settings().context.tool_output_max_tokens
|
||||||
|
|
||||||
|
await wrapped.on_invoke_tool(
|
||||||
|
cast("Any", None),
|
||||||
|
json.dumps({"cmd": "echo hi", "max_output_tokens": ceiling * 100}),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert json.loads(captured["raw_input"])["max_output_tokens"] == ceiling
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -49,3 +78,33 @@ async def test_wrap_exec_command_preserves_explicit_shell(shell: str) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert json.loads(captured["raw_input"])["shell"] == shell
|
assert json.loads(captured["raw_input"])["shell"] == shell
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_responses_filesystem_custom_tool_output_is_bounded() -> None:
|
||||||
|
async def invoke(_ctx: Any, _inp: str) -> str:
|
||||||
|
return "line\n" * 50_000
|
||||||
|
|
||||||
|
toolset = SimpleNamespace(
|
||||||
|
read_file=CustomTool(name="read_file", description="read", on_invoke_tool=invoke)
|
||||||
|
)
|
||||||
|
factory._configure_filesystem_tools(toolset, chat_completions=False)
|
||||||
|
|
||||||
|
assert isinstance(toolset.read_file, CustomTool)
|
||||||
|
result = await toolset.read_file.on_invoke_tool(cast("Any", None), "{}")
|
||||||
|
|
||||||
|
assert "truncated" in result
|
||||||
|
assert len(result) < len("line\n" * 50_000)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_completions_filesystem_custom_tool_becomes_function_tool() -> None:
|
||||||
|
async def invoke(_ctx: Any, _inp: str) -> str:
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
toolset = SimpleNamespace(
|
||||||
|
read_file=CustomTool(name="read_file", description="read", on_invoke_tool=invoke)
|
||||||
|
)
|
||||||
|
factory._configure_filesystem_tools(toolset, chat_completions=True)
|
||||||
|
|
||||||
|
assert isinstance(toolset.read_file, FunctionTool)
|
||||||
|
|||||||
@@ -6,10 +6,11 @@ import json
|
|||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from pydantic import AliasChoices, Field
|
from pydantic import AliasChoices, Field, ValidationError
|
||||||
from pydantic.fields import FieldInfo
|
from pydantic.fields import FieldInfo
|
||||||
|
|
||||||
from strix.config import loader
|
from strix.config import loader
|
||||||
|
from strix.config.settings import ContextSettings
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -120,6 +121,15 @@ def test_read_json_overrides_uses_json_when_no_alias_in_environ(tmp_path: Path)
|
|||||||
assert loader._read_json_overrides(path) == {"llm": {"api_key": "sk-file"}}
|
assert loader._read_json_overrides(path) == {"llm": {"api_key": "sk-file"}}
|
||||||
|
|
||||||
|
|
||||||
|
def test_tool_output_max_bytes_rejects_sub_notice_values() -> None:
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
ContextSettings(STRIX_TOOL_OUTPUT_MAX_BYTES=64)
|
||||||
|
|
||||||
|
|
||||||
|
def test_tool_output_max_bytes_accepts_floor() -> None:
|
||||||
|
assert ContextSettings(STRIX_TOOL_OUTPUT_MAX_BYTES=1024).tool_output_max_bytes == 1024
|
||||||
|
|
||||||
|
|
||||||
# --------------------------------------------------------------------------- #
|
# --------------------------------------------------------------------------- #
|
||||||
# _aliases_for
|
# _aliases_for
|
||||||
# --------------------------------------------------------------------------- #
|
# --------------------------------------------------------------------------- #
|
||||||
|
|||||||
@@ -0,0 +1,60 @@
|
|||||||
|
"""Tests for per-tool-output bounding."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
|
||||||
|
from strix.tools.output_store import bound_text
|
||||||
|
|
||||||
|
|
||||||
|
def test_small_output_passes_through_unchanged() -> None:
|
||||||
|
text = "line 1\nline 2\nline 3"
|
||||||
|
assert bound_text(text, max_lines=100, max_bytes=10_000) == text
|
||||||
|
|
||||||
|
|
||||||
|
def test_line_limit_keeps_head_and_tail() -> None:
|
||||||
|
text = "\n".join(str(i) for i in range(1000))
|
||||||
|
bounded = bound_text(text, max_lines=10, max_bytes=1_000_000)
|
||||||
|
|
||||||
|
assert bounded.startswith("0\n1\n2\n3\n4")
|
||||||
|
assert bounded.rstrip().endswith("999")
|
||||||
|
assert "truncated" in bounded
|
||||||
|
assert len(bounded.splitlines()) < 30
|
||||||
|
|
||||||
|
|
||||||
|
def test_byte_limit_enforced_on_single_long_line() -> None:
|
||||||
|
text = "x" * 100_000
|
||||||
|
bounded = bound_text(text, max_lines=2_000, max_bytes=1_000)
|
||||||
|
|
||||||
|
assert "truncated" in bounded
|
||||||
|
assert len(bounded.encode("utf-8")) <= 1_000
|
||||||
|
|
||||||
|
|
||||||
|
def test_multibyte_characters_not_split() -> None:
|
||||||
|
text = "😀" * 50_000
|
||||||
|
bounded = bound_text(text, max_lines=2_000, max_bytes=1_000)
|
||||||
|
|
||||||
|
# Must remain valid UTF-8 (no mid-character cut).
|
||||||
|
assert bounded == bounded.encode("utf-8").decode("utf-8")
|
||||||
|
assert "truncated" in bounded
|
||||||
|
|
||||||
|
|
||||||
|
def test_notice_reports_dropped_counts() -> None:
|
||||||
|
text = "\n".join("y" * 10 for _ in range(500))
|
||||||
|
bounded = bound_text(text, max_lines=10, max_bytes=1_000_000)
|
||||||
|
|
||||||
|
assert "lines" in bounded
|
||||||
|
assert "bytes" in bounded
|
||||||
|
|
||||||
|
|
||||||
|
def test_dropped_line_count_accounts_for_byte_trimming() -> None:
|
||||||
|
# Tight byte budget drops whole lines from head/tail; the notice must count them.
|
||||||
|
text = "\n".join(f"line-{i}" for i in range(200))
|
||||||
|
bounded = bound_text(text, max_lines=20, max_bytes=40)
|
||||||
|
|
||||||
|
match = re.search(r"\[\.\.\. (\d+) lines", bounded)
|
||||||
|
assert match is not None, bounded
|
||||||
|
dropped = int(match.group(1))
|
||||||
|
kept = [ln for ln in bounded.splitlines() if ln and "truncated" not in ln]
|
||||||
|
assert dropped == 200 - len(kept)
|
||||||
|
assert dropped > 200 - 20
|
||||||
Reference in New Issue
Block a user