diff --git a/strix/config/settings.py b/strix/config/settings.py index 9d16e083..84339218 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -56,6 +56,8 @@ class RuntimeSettings(BaseSettings): # on large repos). Above this, the user must bind-mount via ``--mount``. # Set to 0 (or less) to disable the pre-flight check entirely. max_local_copy_mb: int = Field(default=1024, alias="STRIX_MAX_LOCAL_COPY_MB") + # Max screenshot/image tool outputs kept live per agent context (0 = none). + max_context_images: int = Field(default=3, ge=0, alias="STRIX_MAX_CONTEXT_IMAGES") class TelemetrySettings(BaseSettings): diff --git a/strix/core/agents.py b/strix/core/agents.py index 24a7185f..5cb48d2f 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -10,6 +10,8 @@ from dataclasses import dataclass, field from pathlib import Path from typing import TYPE_CHECKING, Any, Literal, cast +from strix.core.sessions import session_write_lock + if TYPE_CHECKING: from agents.items import TResponseInputItem @@ -137,7 +139,8 @@ class AgentCoordinator: ) return False try: - await session.add_items([self._message_to_session_item(message)]) + async with session_write_lock(session): + await session.add_items([self._message_to_session_item(message)]) except Exception: logger.exception( "agent.send failed to append to SDK session target=%s", diff --git a/strix/core/execution.py b/strix/core/execution.py index 06dc3ddf..d26486de 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -17,7 +17,11 @@ from openai import APIError from strix.core.hooks import BudgetExceededError from strix.core.inputs import child_initial_input -from strix.core.sessions import open_agent_session, strip_all_images_from_session +from strix.core.sessions import ( + enforce_image_budget, + open_agent_session, + strip_all_images_from_session, +) if TYPE_CHECKING: @@ -349,6 +353,13 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 while True: try: await coordinator.mark_running(agent_id) + if session is not None: + max_images = context.get("max_context_images") + if isinstance(max_images, int): + try: + await enforce_image_budget(session, max_images) + except Exception: + logger.exception("image-budget enforcement failed for %s", agent_id) stream = Runner.run_streamed( agent, input=input_data, diff --git a/strix/core/inputs.py b/strix/core/inputs.py index e53ae321..1712fbd3 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -13,6 +13,7 @@ from strix.config.models import ( is_known_openai_bare_model, model_supports_reasoning, ) +from strix.core.sessions import scrub_images_from_items if TYPE_CHECKING: @@ -161,7 +162,11 @@ def child_initial_input( """ parts: list[str] = [] if parent_history: - rendered = json.dumps(parent_history, ensure_ascii=False, default=str) + rendered = json.dumps( + scrub_images_from_items(parent_history), + ensure_ascii=False, + default=str, + ) parts.append( "== Inherited context from parent (background only) ==\n" f"{rendered}\n" diff --git a/strix/core/runner.py b/strix/core/runner.py index 0ebcc48b..e6a15dc3 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -287,6 +287,7 @@ async def run_strix_scan( "parent_id": None, "interactive": interactive, "spawn_child_agent": spawn_child_agent, + "max_context_images": settings.runtime.max_context_images, } root_session = open_agent_session(root_id, agents_db) diff --git a/strix/core/sessions.py b/strix/core/sessions.py index 8e4f4142..67fbba95 100644 --- a/strix/core/sessions.py +++ b/strix/core/sessions.py @@ -2,64 +2,149 @@ from __future__ import annotations -import contextlib +import asyncio +import logging from typing import TYPE_CHECKING, Any, cast +from weakref import WeakKeyDictionary from agents.memory import SQLiteSession if TYPE_CHECKING: + from collections.abc import Callable from pathlib import Path from agents.items import TResponseInputItem from agents.memory import Session +logger = logging.getLogger(__name__) + + def open_agent_session(agent_id: str, path: Path) -> SQLiteSession: path.parent.mkdir(parents=True, exist_ok=True) return SQLiteSession(session_id=agent_id, db_path=path) _IMAGE_REJECTED_TEXT = "[image rejected by the model]" +_IMAGE_ELIDED_TEXT = "[older screenshot elided to bound context memory]" +_INHERITED_IMAGE_TEXT = "[screenshot omitted from inherited context]" + + +def _output_has_image(item_dict: dict[str, Any]) -> bool: + return ( + item_dict.get("type") == "function_call_output" + and isinstance(item_dict.get("output"), list) + and any(isinstance(b, dict) and b.get("type") == "input_image" for b in item_dict["output"]) + ) + + +def _elided_output(item_dict: dict[str, Any], text: str) -> dict[str, Any]: + # Replace only image blocks; sibling text blocks are preserved. + output = item_dict.get("output") + blocks = output if isinstance(output, list) else [] + return { + "type": "function_call_output", + "call_id": item_dict.get("call_id"), + "output": [ + {"type": "input_text", "text": text} + if isinstance(block, dict) and block.get("type") == "input_image" + else block + for block in blocks + ], + } + + +_session_write_locks: WeakKeyDictionary[Session, asyncio.Lock] = WeakKeyDictionary() + + +def session_write_lock(session: Session) -> asyncio.Lock: + """Lock serialising all out-of-band writes to ``session``.""" + lock = _session_write_locks.get(session) + if lock is None: + lock = asyncio.Lock() + _session_write_locks[session] = lock + return lock + + +async def _rewrite_session( + session: Session, + transform: Callable[[list[Any]], tuple[list[Any], bool]], +) -> bool: + """Read-modify-write a session under its write lock, restoring on failure.""" + async with session_write_lock(session): + items = await session.get_items() + if not items: + return False + rebuilt, changed = transform(list(items)) + if not changed: + return False + rebuilt_items = cast("list[TResponseInputItem]", rebuilt) + original_items = cast("list[TResponseInputItem]", list(items)) + await session.clear_session() + try: + await session.add_items(rebuilt_items) + except Exception: + logger.exception("session rewrite failed; restoring original items") + await session.clear_session() + await session.add_items(original_items) + raise + return True async def strip_all_images_from_session(session: Session) -> bool: - items = await session.get_items() - if not items: + """Replace every image tool output with a text placeholder (rejection recovery).""" + + def _transform(items: list[Any]) -> tuple[list[Any], bool]: + rebuilt: list[Any] = [] + changed = False + for item in items: + item_dict = cast("dict[str, Any]", item) if isinstance(item, dict) else None + if item_dict is not None and _output_has_image(item_dict): + rebuilt.append(_elided_output(item_dict, _IMAGE_REJECTED_TEXT)) + changed = True + else: + rebuilt.append(item) + return rebuilt, changed + + return await _rewrite_session(session, _transform) + + +async def enforce_image_budget(session: Session, max_images: int) -> bool: + """Keep only the most recent ``max_images`` image outputs; elide older ones.""" + if max_images < 0: return False - rebuilt: list[Any] = [] - changed = False - for item in items: - item_dict = cast("dict[str, Any]", item) if isinstance(item, dict) else None - if ( - item_dict is not None - and item_dict.get("type") == "function_call_output" - and isinstance(item_dict.get("output"), list) - and any( - isinstance(b, dict) and b.get("type") == "input_image" for b in item_dict["output"] - ) - ): - rebuilt.append( - { - "type": "function_call_output", - "call_id": item_dict.get("call_id"), - "output": [{"type": "input_text", "text": _IMAGE_REJECTED_TEXT}], - }, - ) - changed = True - else: - rebuilt.append(item) + def _transform(items: list[Any]) -> tuple[list[Any], bool]: + image_indices = [ + i + for i, item in enumerate(items) + if isinstance(item, dict) and _output_has_image(cast("dict[str, Any]", item)) + ] + if len(image_indices) <= max_images: + return items, False + to_elide = set(image_indices[: len(image_indices) - max_images]) + rebuilt = [ + _elided_output(cast("dict[str, Any]", item), _IMAGE_ELIDED_TEXT) + if i in to_elide + else item + for i, item in enumerate(items) + ] + return rebuilt, True - if not changed: - return False + return await _rewrite_session(session, _transform) - rebuilt_items = cast("list[TResponseInputItem]", rebuilt) - await session.clear_session() - try: - await session.add_items(rebuilt_items) - except Exception: - with contextlib.suppress(Exception): - await session.add_items(rebuilt_items) - raise - return True + +def scrub_images_from_items(items: list[Any]) -> list[Any]: + """Return a copy of ``items`` with every image block replaced by text.""" + + def _scrub(obj: Any) -> Any: + if isinstance(obj, dict): + if obj.get("type") == "input_image": + return {"type": "input_text", "text": _INHERITED_IMAGE_TEXT} + return {k: _scrub(v) for k, v in obj.items()} + if isinstance(obj, list): + return [_scrub(v) for v in obj] + return obj + + return [_scrub(item) for item in items] diff --git a/tests/test_runner_rate_limit.py b/tests/test_runner_rate_limit.py index 7e708eb5..401caa15 100644 --- a/tests/test_runner_rate_limit.py +++ b/tests/test_runner_rate_limit.py @@ -37,7 +37,8 @@ async def test_persistent_rate_limit_stops_gracefully( model="openai/gpt-4o", reasoning_effort="high", force_required_tool_choice=False, - ) + ), + runtime=types.SimpleNamespace(max_context_images=3), ) monkeypatch.setattr(runner, "load_settings", lambda: settings) monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _settings: None) @@ -54,8 +55,8 @@ async def test_persistent_rate_limit_stops_gracefully( async def _cleanup(*_args: Any, **_kwargs: Any) -> None: return None - monkeypatch.setattr(runner.session_manager, "create_or_reuse", _create_or_reuse) - monkeypatch.setattr(runner.session_manager, "cleanup", _cleanup) + monkeypatch.setattr(runner.session_manager, "create_or_reuse", _create_or_reuse) # type: ignore[attr-defined] + monkeypatch.setattr(runner.session_manager, "cleanup", _cleanup) # type: ignore[attr-defined] monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task") monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: "")