mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 08:03:15 +00:00
* feat: auto-load scoped AGENTS on reads (#1684) Adapted from upstream langchain-ai/open-swe #1684 to the fork's direct-import middleware registry and middleware stack ordering. Adds SubdirAgentsReadMiddleware, which appends applicable ancestor AGENTS.md instructions to read_file results once per run, so scoped rules are visible before edits. Wired into get_agent immediately after ToolErrorMiddleware, matching upstream's relative position. Note: this changes file-read behavior for every main-agent run. The reviewer graph uses its own leaner middleware stack and is unaffected. (cherry picked from commit 7f7af71547be2199cea699284676f8ceefba7691) * feat: add platform issue reporting tool (#1685) Adapted from upstream langchain-ai/open-swe #1685 to the fork's direct-import tool registry. Adds the report_platform_issue tool (stdlib-only: returns a locally generated UUIDv7 report id, no external network call) and wires it into get_agent's curated tool list. Dropped upstream's test_task_retry_wraps_inside_tool_error_middleware assertion, which references ToolRetryMiddleware that this fork does not wire into the middleware stack. (cherry picked from commit 88b62322b44103335773002d0745704cd96e9160) --------- Co-authored-by: Ramon Nogueira <ramon.nogueira@langchain.dev>
This commit is contained in:
parent
a53a96d37b
commit
b65a99f3e9
9 changed files with 430 additions and 7 deletions
15
AGENTS.md
15
AGENTS.md
|
|
@ -66,13 +66,14 @@ Configured in `agent/server.py:get_agent`, runs around every model call (in this
|
|||
1. `SanitizeToolInputsMiddleware` — strips/normalizes tool inputs before they reach tools.
|
||||
2. `ModelCallLimitMiddleware` (from `langchain.agents.middleware`) — caps model calls at `MODEL_CALL_RECURSION_LIMIT` (~half of `DEFAULT_RECURSION_LIMIT`); `exit_behavior="end"`.
|
||||
3. `ToolErrorMiddleware` — catches tool exceptions and surfaces them as tool messages.
|
||||
4. `check_message_queue_before_model` — pulls Linear comments / Slack messages that arrived mid-run from the thread queue and injects them as user messages before the next LLM call. This is what makes "message the agent while it's working" work.
|
||||
5. `SlackAssistantStatusMiddleware` — keeps the Slack "assistant is typing"-style status up to date around model calls.
|
||||
6. `ensure_no_empty_msg` — after-model hook; when the model emits a message with no tool call (and hasn't already messaged the user or confirmed completion) it re-injects a synthetic `no_op` / `confirming_completion` tool call so the run continues instead of ending prematurely.
|
||||
7. `notify_step_limit_reached` — after-agent hook that posts a Slack reply when the agent hits the step limit, so the user gets a clear signal instead of silence.
|
||||
8. `SandboxCircuitBreakerMiddleware` — trips the agent out of repeated sandbox failures instead of looping.
|
||||
9. `ModelFallbackMiddleware` (optional) — added only when `LLM_FALLBACK_MODEL_ID` or the per-model default fallback differs from the primary model.
|
||||
10. `SanitizeThinkingBlocksMiddleware` — strips malformed empty Anthropic thinking blocks immediately before provider calls.
|
||||
4. `SubdirAgentsReadMiddleware` — appends applicable ancestor `AGENTS.md` instructions to `read_file` results once per run, so scoped rules are visible before edits.
|
||||
5. `check_message_queue_before_model` — pulls Linear comments / Slack messages that arrived mid-run from the thread queue and injects them as user messages before the next LLM call. This is what makes "message the agent while it's working" work.
|
||||
6. `SlackAssistantStatusMiddleware` — keeps the Slack "assistant is typing"-style status up to date around model calls.
|
||||
7. `ensure_no_empty_msg` — after-model hook; when the model emits a message with no tool call (and hasn't already messaged the user or confirmed completion) it re-injects a synthetic `no_op` / `confirming_completion` tool call so the run continues instead of ending prematurely.
|
||||
8. `notify_step_limit_reached` — after-agent hook that posts a Slack reply when the agent hits the step limit, so the user gets a clear signal instead of silence.
|
||||
9. `SandboxCircuitBreakerMiddleware` — trips the agent out of repeated sandbox failures instead of looping.
|
||||
10. `ModelFallbackMiddleware` (optional) — added only when `LLM_FALLBACK_MODEL_ID` or the per-model default fallback differs from the primary model.
|
||||
11. `SanitizeThinkingBlocksMiddleware` — strips malformed empty Anthropic thinking blocks immediately before provider calls.
|
||||
|
||||
The system prompt instructs the agent to call a tool every turn, and `ensure_no_empty_msg` re-injects a tool call when it doesn't — together these keep runs from stopping partway through a task.
|
||||
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from .sandbox_circuit_breaker import SandboxCircuitBreakerMiddleware
|
|||
from .sanitize_thinking_blocks import SanitizeThinkingBlocksMiddleware
|
||||
from .sanitize_tool_inputs import SanitizeToolInputsMiddleware
|
||||
from .settle_review_check import settle_review_check_on_exit
|
||||
from .subdir_agents import SubdirAgentsReadMiddleware
|
||||
from .tool_artifact import ToolArtifactMiddleware
|
||||
from .tool_error_handler import ToolErrorMiddleware
|
||||
from .workflow_push_guard import WorkflowPushGuardMiddleware
|
||||
|
|
@ -22,6 +23,7 @@ __all__ = [
|
|||
"RepairOrphanedToolCallsMiddleware",
|
||||
"SanitizeThinkingBlocksMiddleware",
|
||||
"SanitizeToolInputsMiddleware",
|
||||
"SubdirAgentsReadMiddleware",
|
||||
"ToolArtifactMiddleware",
|
||||
"ToolErrorMiddleware",
|
||||
"WorkflowPushGuardMiddleware",
|
||||
|
|
|
|||
197
agent/middleware/subdir_agents.py
Normal file
197
agent/middleware/subdir_agents.py
Normal file
|
|
@ -0,0 +1,197 @@
|
|||
"""Auto-load applicable ``AGENTS.md`` files after file reads."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import posixpath
|
||||
from collections import defaultdict
|
||||
from collections.abc import Awaitable, Callable, Iterable, Mapping
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import AgentMiddleware, AgentState
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langgraph.config import get_config
|
||||
from langgraph.prebuilt.tool_node import ToolCallRequest
|
||||
from langgraph.types import Command
|
||||
|
||||
from ..utils.sandbox_state import SANDBOX_BACKENDS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_READ_FILE = "read_file"
|
||||
_AGENTS_MD = "AGENTS.md"
|
||||
_MAX_AGENTS_LINES = 1_000
|
||||
_MAX_AGENTS_BYTES = 64 * 1024
|
||||
|
||||
|
||||
def _tool_name(request: ToolCallRequest) -> str | None:
|
||||
tool_call = getattr(request, "tool_call", None)
|
||||
if isinstance(tool_call, Mapping):
|
||||
name = tool_call.get("name")
|
||||
return name if isinstance(name, str) and name else None
|
||||
return None
|
||||
|
||||
|
||||
def _tool_args(request: ToolCallRequest) -> dict[str, Any]:
|
||||
tool_call = getattr(request, "tool_call", None)
|
||||
if isinstance(tool_call, Mapping):
|
||||
args = tool_call.get("args")
|
||||
if isinstance(args, Mapping):
|
||||
return dict(args)
|
||||
return {}
|
||||
|
||||
|
||||
def _thread_id(request: ToolCallRequest) -> str | None:
|
||||
runtime_config = getattr(getattr(request, "runtime", None), "config", None)
|
||||
config: Mapping[str, Any] | None = (
|
||||
runtime_config if isinstance(runtime_config, Mapping) else None
|
||||
)
|
||||
if config is None:
|
||||
try:
|
||||
maybe_config = get_config()
|
||||
except Exception:
|
||||
return None
|
||||
config = maybe_config if isinstance(maybe_config, Mapping) else None
|
||||
if config is None:
|
||||
return None
|
||||
configurable = config.get("configurable", {})
|
||||
if not isinstance(configurable, Mapping):
|
||||
return None
|
||||
thread_id = configurable.get("thread_id")
|
||||
return thread_id if isinstance(thread_id, str) and thread_id else None
|
||||
|
||||
|
||||
def _file_path(args: Mapping[str, Any]) -> str | None:
|
||||
raw = args.get("file_path")
|
||||
if not isinstance(raw, str):
|
||||
return None
|
||||
path = raw.strip()
|
||||
if not path.startswith("/"):
|
||||
return None
|
||||
return posixpath.normpath(path)
|
||||
|
||||
|
||||
def _candidate_agents_paths(file_path: str) -> list[str]:
|
||||
current = posixpath.dirname(file_path)
|
||||
candidates: list[str] = []
|
||||
while current and current != "/":
|
||||
candidates.append(posixpath.join(current, _AGENTS_MD))
|
||||
parent = posixpath.dirname(current)
|
||||
if parent == current:
|
||||
break
|
||||
current = parent
|
||||
candidates.reverse()
|
||||
return [candidate for candidate in candidates if candidate != file_path]
|
||||
|
||||
|
||||
def _extract_text(result: Any) -> str | None:
|
||||
error = getattr(result, "error", None)
|
||||
if error:
|
||||
return None
|
||||
|
||||
file_data = getattr(result, "file_data", None)
|
||||
if file_data is None:
|
||||
return None
|
||||
|
||||
if isinstance(file_data, Mapping):
|
||||
encoding = file_data.get("encoding")
|
||||
content = file_data.get("content")
|
||||
else:
|
||||
encoding = getattr(file_data, "encoding", None)
|
||||
content = getattr(file_data, "content", None)
|
||||
|
||||
if encoding is not None and encoding != "utf-8":
|
||||
return None
|
||||
if not isinstance(content, str) or not content.strip():
|
||||
return None
|
||||
|
||||
encoded = content.encode("utf-8")
|
||||
if len(encoded) <= _MAX_AGENTS_BYTES:
|
||||
return content
|
||||
return encoded[:_MAX_AGENTS_BYTES].decode("utf-8", errors="ignore") + "\n\n[truncated]"
|
||||
|
||||
|
||||
def _system_reminder(file_path: str, loaded: Iterable[tuple[str, str]]) -> str:
|
||||
sections = [f"Instructions from: {path}\n{content.rstrip()}" for path, content in loaded]
|
||||
body = "\n\n".join(sections)
|
||||
return (
|
||||
"<system-reminder>\n"
|
||||
f"Loaded applicable AGENTS.md instructions for `{file_path}`. "
|
||||
"Follow these before editing files under their scopes; more deeply nested "
|
||||
"instructions take precedence.\n\n"
|
||||
f"{body}\n"
|
||||
"</system-reminder>"
|
||||
)
|
||||
|
||||
|
||||
def _can_append_reminder(result: ToolMessage | Command) -> bool:
|
||||
return (
|
||||
isinstance(result, ToolMessage)
|
||||
and getattr(result, "status", None) != "error"
|
||||
and isinstance(result.content, str)
|
||||
)
|
||||
|
||||
|
||||
def _append_reminder(result: ToolMessage | Command, reminder: str | None) -> ToolMessage | Command:
|
||||
if reminder is not None and _can_append_reminder(result):
|
||||
result.content = f"{result.content}\n\n{reminder}"
|
||||
return result
|
||||
|
||||
|
||||
class SubdirAgentsReadMiddleware(AgentMiddleware):
|
||||
"""Append applicable ancestor ``AGENTS.md`` files to ``read_file`` results."""
|
||||
|
||||
state_schema = AgentState
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._loaded: defaultdict[str, set[str]] = defaultdict(set)
|
||||
|
||||
def _thread_key(self, request: ToolCallRequest) -> str:
|
||||
return _thread_id(request) or "__unknown_thread__"
|
||||
|
||||
def _backend(self, request: ToolCallRequest) -> Any | None:
|
||||
thread_id = _thread_id(request)
|
||||
if not thread_id:
|
||||
return None
|
||||
return SANDBOX_BACKENDS.get(thread_id)
|
||||
|
||||
def _mark_direct_agents_read(self, request: ToolCallRequest, file_path: str) -> bool:
|
||||
if posixpath.basename(file_path) != _AGENTS_MD:
|
||||
return False
|
||||
self._loaded[self._thread_key(request)].add(file_path)
|
||||
return True
|
||||
|
||||
async def _load_async(self, request: ToolCallRequest, file_path: str) -> str | None:
|
||||
if self._mark_direct_agents_read(request, file_path):
|
||||
return None
|
||||
backend = self._backend(request)
|
||||
if backend is None:
|
||||
return None
|
||||
loaded_paths = self._loaded[self._thread_key(request)]
|
||||
loaded: list[tuple[str, str]] = []
|
||||
for path in _candidate_agents_paths(file_path):
|
||||
if path in loaded_paths:
|
||||
continue
|
||||
loaded_paths.add(path)
|
||||
try:
|
||||
text = _extract_text(await backend.aread(path, offset=0, limit=_MAX_AGENTS_LINES))
|
||||
except Exception:
|
||||
logger.debug("subdir_agents: aread failed for %s", path, exc_info=True)
|
||||
continue
|
||||
if text is None:
|
||||
continue
|
||||
loaded.append((path, text))
|
||||
return _system_reminder(file_path, loaded) if loaded else None
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
|
||||
) -> ToolMessage | Command:
|
||||
if _tool_name(request) != _READ_FILE:
|
||||
return await handler(request)
|
||||
result = await handler(request)
|
||||
file_path = _file_path(_tool_args(request))
|
||||
if file_path is None or not _can_append_reminder(result):
|
||||
return result
|
||||
return _append_reminder(result, await self._load_async(request, file_path))
|
||||
|
|
@ -61,6 +61,7 @@ from .middleware import (
|
|||
SanitizeThinkingBlocksMiddleware,
|
||||
SanitizeToolInputsMiddleware,
|
||||
SlackAssistantStatusMiddleware,
|
||||
SubdirAgentsReadMiddleware,
|
||||
ToolArtifactMiddleware,
|
||||
ToolErrorMiddleware,
|
||||
WorkflowPushGuardMiddleware,
|
||||
|
|
@ -82,6 +83,7 @@ from .tools import (
|
|||
linear_list_teams,
|
||||
linear_update_issue,
|
||||
open_pull_request,
|
||||
report_platform_issue,
|
||||
request_pr_review,
|
||||
save_plan,
|
||||
schedule_thread_wakeup,
|
||||
|
|
@ -957,6 +959,7 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
|||
linear_update_issue,
|
||||
open_pull_request,
|
||||
request_pr_review,
|
||||
report_platform_issue,
|
||||
schedule_thread_wakeup,
|
||||
slack_add_reaction,
|
||||
slack_read_thread_messages,
|
||||
|
|
@ -973,6 +976,7 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
|||
SanitizeToolInputsMiddleware(),
|
||||
ModelCallLimitMiddleware(run_limit=MODEL_CALL_RECURSION_LIMIT, exit_behavior="end"),
|
||||
ToolErrorMiddleware(),
|
||||
SubdirAgentsReadMiddleware(),
|
||||
ToolArtifactMiddleware(),
|
||||
WorkflowPushGuardMiddleware(),
|
||||
refresh_github_proxy_before_model,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from .open_pull_request import open_pull_request
|
|||
from .publish_review import publish_review
|
||||
from .read_repo_file import read_repo_file
|
||||
from .reply_to_finding_thread import reply_to_finding_thread
|
||||
from .report_platform_issue import report_platform_issue
|
||||
from .request_pr_review import request_pr_review
|
||||
from .resolve_finding_thread import resolve_finding_thread
|
||||
from .save_plan import save_plan
|
||||
|
|
@ -44,6 +45,7 @@ __all__ = [
|
|||
"open_pull_request",
|
||||
"publish_review",
|
||||
"read_repo_file",
|
||||
"report_platform_issue",
|
||||
"request_pr_review",
|
||||
"reply_to_finding_thread",
|
||||
"resolve_finding_thread",
|
||||
|
|
|
|||
16
agent/tools/report_platform_issue.py
Normal file
16
agent/tools/report_platform_issue.py
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
import secrets
|
||||
import time
|
||||
import uuid
|
||||
|
||||
|
||||
def _uuid7() -> str:
|
||||
timestamp_ms = (time.time_ns() // 1_000_000) & ((1 << 48) - 1)
|
||||
rand_a = secrets.randbits(12)
|
||||
rand_b = secrets.randbits(62)
|
||||
uuid_int = (timestamp_ms << 80) | (0x7 << 76) | (rand_a << 64) | (0b10 << 62) | rand_b
|
||||
return str(uuid.UUID(int=uuid_int))
|
||||
|
||||
|
||||
async def report_platform_issue() -> dict[str, str]:
|
||||
"""Report an issue with the sandbox or execution environment."""
|
||||
return {"report_id": _uuid7()}
|
||||
|
|
@ -92,6 +92,16 @@ async def test_agent_does_not_add_custom_repair_middleware() -> None:
|
|||
assert "RepairOrphanedToolCallsMiddleware" not in names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_wires_subdir_agents_middleware_after_tool_error() -> None:
|
||||
captured = await _capture_create_deep_agent_kwargs()
|
||||
middleware = captured["middleware"]
|
||||
assert isinstance(middleware, list)
|
||||
names = [type(m).__name__ for m in middleware]
|
||||
assert "SubdirAgentsReadMiddleware" in names
|
||||
assert names.index("ToolErrorMiddleware") < names.index("SubdirAgentsReadMiddleware")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_keeps_message_queue_and_step_limit_middleware() -> None:
|
||||
captured = await _capture_create_deep_agent_kwargs()
|
||||
|
|
@ -101,3 +111,13 @@ async def test_agent_keeps_message_queue_and_step_limit_middleware() -> None:
|
|||
present = {type(m).__name__ for m in middleware}
|
||||
assert "check_message_queue_before_model" in present
|
||||
assert "notify_step_limit_reached" in present
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_includes_report_platform_issue_tool() -> None:
|
||||
from agent.tools import report_platform_issue
|
||||
|
||||
captured = await _capture_create_deep_agent_kwargs()
|
||||
tools = captured["tools"]
|
||||
assert isinstance(tools, list)
|
||||
assert report_platform_issue in tools
|
||||
|
|
|
|||
20
tests/test_report_platform_issue_tool.py
Normal file
20
tests/test_report_platform_issue_tool.py
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
|
||||
async def test_report_platform_issue_returns_uuid7_report_id() -> None:
|
||||
from agent.tools.report_platform_issue import report_platform_issue
|
||||
|
||||
result = await report_platform_issue()
|
||||
|
||||
assert set(result) == {"report_id"}
|
||||
report_id = uuid.UUID(result["report_id"])
|
||||
assert report_id.version == 7
|
||||
assert report_id.variant == uuid.RFC_4122
|
||||
|
||||
|
||||
def test_report_platform_issue_exported() -> None:
|
||||
from agent.tools import report_platform_issue
|
||||
|
||||
assert callable(report_platform_issue)
|
||||
161
tests/test_subdir_agents_middleware.py
Normal file
161
tests/test_subdir_agents_middleware.py
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
"""Tests for subdirectory AGENTS.md auto-loading on read_file."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import ToolMessage
|
||||
|
||||
from agent.middleware.subdir_agents import SubdirAgentsReadMiddleware
|
||||
from agent.utils import sandbox_state
|
||||
|
||||
|
||||
class FakeReadResult:
|
||||
def __init__(
|
||||
self, *, content: str | None = None, encoding: str = "utf-8", error: str | None = None
|
||||
) -> None:
|
||||
self.error = error
|
||||
self.file_data = None if content is None else {"content": content, "encoding": encoding}
|
||||
|
||||
|
||||
class FakeBackend:
|
||||
def __init__(self, results: dict[str, Any]) -> None:
|
||||
self.results = results
|
||||
self.reads: list[str] = []
|
||||
|
||||
async def aread(self, file_path: str, offset: int = 0, limit: int = 2000) -> Any:
|
||||
self.reads.append(file_path)
|
||||
result = self.results.get(file_path)
|
||||
if isinstance(result, Exception):
|
||||
raise result
|
||||
if result is None:
|
||||
return FakeReadResult(error="File not found")
|
||||
return result
|
||||
|
||||
|
||||
def _request(name: str, args: dict[str, Any], thread_id: str = "t1") -> Any:
|
||||
return SimpleNamespace(
|
||||
tool_call={"name": name, "args": args, "id": "call-1"},
|
||||
runtime=SimpleNamespace(config={"configurable": {"thread_id": thread_id}}),
|
||||
)
|
||||
|
||||
|
||||
def _ok(name: str, content: str = "file content") -> ToolMessage:
|
||||
return ToolMessage(content=content, tool_call_id="call-1", status="success", name=name)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def register_backend():
|
||||
registered: list[str] = []
|
||||
|
||||
def _register(thread_id: str, backend: Any) -> Any:
|
||||
sandbox_state.SANDBOX_BACKENDS[thread_id] = backend
|
||||
registered.append(thread_id)
|
||||
return backend
|
||||
|
||||
yield _register
|
||||
for thread_id in registered:
|
||||
sandbox_state.SANDBOX_BACKENDS.pop(thread_id, None)
|
||||
|
||||
|
||||
async def test_read_file_appends_applicable_agents_in_root_to_leaf_order(register_backend) -> None:
|
||||
backend = FakeBackend(
|
||||
{
|
||||
"/repo/AGENTS.md": FakeReadResult(content="root rules"),
|
||||
"/repo/pkg/AGENTS.md": FakeReadResult(content="pkg rules"),
|
||||
}
|
||||
)
|
||||
register_backend("t1", backend)
|
||||
request = _request("read_file", {"file_path": "/repo/pkg/src/app.py"})
|
||||
|
||||
async def handler(_req: Any) -> ToolMessage:
|
||||
return _ok("read_file")
|
||||
|
||||
result = await SubdirAgentsReadMiddleware().awrap_tool_call(request, handler)
|
||||
|
||||
assert isinstance(result.content, str)
|
||||
assert "<system-reminder>" in result.content
|
||||
assert "Loaded applicable AGENTS.md instructions for `/repo/pkg/src/app.py`" in result.content
|
||||
assert result.content.index("Instructions from: /repo/AGENTS.md") < result.content.index(
|
||||
"Instructions from: /repo/pkg/AGENTS.md"
|
||||
)
|
||||
assert "root rules" in result.content
|
||||
assert "pkg rules" in result.content
|
||||
assert backend.reads == [
|
||||
"/repo/AGENTS.md",
|
||||
"/repo/pkg/AGENTS.md",
|
||||
"/repo/pkg/src/AGENTS.md",
|
||||
]
|
||||
|
||||
|
||||
async def test_read_file_does_not_reload_same_agents_file(register_backend) -> None:
|
||||
backend = FakeBackend({"/repo/pkg/AGENTS.md": FakeReadResult(content="pkg rules")})
|
||||
register_backend("t1", backend)
|
||||
middleware = SubdirAgentsReadMiddleware()
|
||||
|
||||
async def handler(_req: Any) -> ToolMessage:
|
||||
return _ok("read_file")
|
||||
|
||||
first = await middleware.awrap_tool_call(
|
||||
_request("read_file", {"file_path": "/repo/pkg/a.py"}), handler
|
||||
)
|
||||
second = await middleware.awrap_tool_call(
|
||||
_request("read_file", {"file_path": "/repo/pkg/b.py"}), handler
|
||||
)
|
||||
|
||||
assert isinstance(first.content, str)
|
||||
assert "pkg rules" in first.content
|
||||
assert isinstance(second.content, str)
|
||||
assert "system-reminder" not in second.content
|
||||
assert backend.reads == ["/repo/AGENTS.md", "/repo/pkg/AGENTS.md"]
|
||||
|
||||
|
||||
async def test_reading_agents_md_directly_does_not_append_reminder(register_backend) -> None:
|
||||
backend = FakeBackend({"/repo/pkg/AGENTS.md": FakeReadResult(content="pkg rules")})
|
||||
register_backend("t1", backend)
|
||||
request = _request("read_file", {"file_path": "/repo/pkg/AGENTS.md"})
|
||||
|
||||
async def handler(_req: Any) -> ToolMessage:
|
||||
return _ok("read_file", "pkg rules")
|
||||
|
||||
result = await SubdirAgentsReadMiddleware().awrap_tool_call(request, handler)
|
||||
|
||||
assert result.content == "pkg rules"
|
||||
assert backend.reads == []
|
||||
|
||||
|
||||
async def test_non_read_file_tool_is_untouched(register_backend) -> None:
|
||||
backend = FakeBackend({"/repo/AGENTS.md": FakeReadResult(content="root rules")})
|
||||
register_backend("t1", backend)
|
||||
request = _request("edit_file", {"file_path": "/repo/a.py"})
|
||||
|
||||
async def handler(_req: Any) -> ToolMessage:
|
||||
return _ok("edit_file")
|
||||
|
||||
result = await SubdirAgentsReadMiddleware().awrap_tool_call(request, handler)
|
||||
|
||||
assert result.content == "file content"
|
||||
assert backend.reads == []
|
||||
|
||||
|
||||
async def test_missing_backend_is_graceful() -> None:
|
||||
request = _request("read_file", {"file_path": "/repo/a.py"}, thread_id="absent-thread")
|
||||
|
||||
async def handler(_req: Any) -> ToolMessage:
|
||||
return _ok("read_file")
|
||||
|
||||
result = await SubdirAgentsReadMiddleware().awrap_tool_call(request, handler)
|
||||
|
||||
assert result.content == "file content"
|
||||
|
||||
|
||||
def test_sync_tool_call_is_not_supported() -> None:
|
||||
request = _request("read_file", {"file_path": "/repo/a.py"})
|
||||
|
||||
def handler(_req: Any) -> ToolMessage:
|
||||
return _ok("read_file")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
SubdirAgentsReadMiddleware().wrap_tool_call(request, handler)
|
||||
Loading…
Add table
Reference in a new issue