diff --git a/agent/middleware/__init__.py b/agent/middleware/__init__.py index 9a08ec69..18f11fa3 100644 --- a/agent/middleware/__init__.py +++ b/agent/middleware/__init__.py @@ -1,5 +1,6 @@ from .check_message_queue import check_message_queue_before_model from .ensure_no_empty_msg import ensure_no_empty_msg +from .notify_step_limit import notify_step_limit_reached from .open_pr import open_pr_if_needed from .tool_error_handler import ToolErrorMiddleware @@ -7,5 +8,6 @@ __all__ = [ "ToolErrorMiddleware", "check_message_queue_before_model", "ensure_no_empty_msg", + "notify_step_limit_reached", "open_pr_if_needed", ] diff --git a/agent/middleware/notify_step_limit.py b/agent/middleware/notify_step_limit.py new file mode 100644 index 00000000..9a88e9a0 --- /dev/null +++ b/agent/middleware/notify_step_limit.py @@ -0,0 +1,88 @@ +"""After-agent middleware that notifies users when the step limit is reached.""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import Any + +from langchain.agents.middleware import AgentState, after_agent +from langgraph.config import get_config +from langgraph.runtime import Runtime + +from ..utils.slack import post_slack_thread_reply + +logger = logging.getLogger(__name__) + +_LIMIT_MARKER = "Model call limits exceeded" + + +def _content_to_text(content: object) -> str: + if isinstance(content, str): + return content + if not isinstance(content, list): + return str(content) + + parts: list[str] = [] + for block in content: + if isinstance(block, Mapping): + text = block.get("text", "") + parts.append(text if isinstance(text, str) else str(text)) + else: + parts.append(str(block)) + return " ".join(parts) + + +@after_agent +async def notify_step_limit_reached( + state: AgentState, + runtime: Runtime, +) -> dict[str, Any] | None: + """Notify the user via Slack when the agent hits its step limit. + + Runs after the agent exits. Checks whether the last AI message contains + the ``ModelCallLimitMiddleware`` marker text; if so, posts a Slack thread + reply so the user is not left wondering what happened. + """ + messages = state.get("messages", []) + if not messages: + return None + + last_msg = messages[-1] + content = _content_to_text(getattr(last_msg, "content", "") or "") + + if _LIMIT_MARKER not in content: + return None + + config = get_config() + configurable = config.get("configurable", {}) + slack_thread = configurable.get("slack_thread") if isinstance(configurable, dict) else None + if not isinstance(slack_thread, dict): + logger.info("No Slack thread config — cannot send step-limit notification") + return None + + channel_id = slack_thread.get("channel_id") + thread_ts = slack_thread.get("thread_ts") + + if ( + not isinstance(channel_id, str) + or not isinstance(thread_ts, str) + or not channel_id + or not thread_ts + ): + logger.info("No Slack thread config — cannot send step-limit notification") + return None + + message = ( + "I've reached my maximum step limit and had to stop. " + "The task may be incomplete. You can retry with a more focused request, " + "or ask me to continue from where I left off." + ) + + try: + await post_slack_thread_reply(channel_id, thread_ts, message) + logger.info("Sent step-limit notification to Slack thread %s", thread_ts) + except Exception: + logger.exception("Failed to send step-limit notification") + + return None diff --git a/agent/server.py b/agent/server.py index bdd3becf..9c6ca69d 100644 --- a/agent/server.py +++ b/agent/server.py @@ -23,6 +23,7 @@ warnings.filterwarnings("ignore", message=".*Pydantic V1.*", category=UserWarnin # Now safe to import agent (which imports LangChain modules) from deepagents import create_deep_agent from deepagents.backends.protocol import SandboxBackendProtocol +from langchain.agents.middleware import ModelCallLimitMiddleware from langsmith.sandbox import SandboxClientError from .integrations.langsmith import _configure_github_proxy @@ -30,6 +31,7 @@ from .middleware import ( ToolErrorMiddleware, check_message_queue_before_model, ensure_no_empty_msg, + notify_step_limit_reached, open_pr_if_needed, ) from .prompt import construct_system_prompt @@ -322,9 +324,12 @@ async def get_agent(config: RunnableConfig) -> Pregel: ], backend=sandbox_backend, middleware=[ + ModelCallLimitMiddleware(run_limit=60, exit_behavior="end"), ToolErrorMiddleware(), check_message_queue_before_model, ensure_no_empty_msg, + # after_agent hooks run in reverse list order; notify after the PR safety net. + notify_step_limit_reached, open_pr_if_needed, ], ).with_config(config) diff --git a/tests/test_notify_step_limit_middleware.py b/tests/test_notify_step_limit_middleware.py new file mode 100644 index 00000000..207798bf --- /dev/null +++ b/tests/test_notify_step_limit_middleware.py @@ -0,0 +1,96 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from langchain_core.messages import AIMessage, HumanMessage + +from agent.middleware.notify_step_limit import notify_step_limit_reached + + +class TestNotifyStepLimitReached: + def _make_runtime(self) -> MagicMock: + return MagicMock() + + @pytest.mark.asyncio + async def test_posts_slack_reply_when_limit_marker_present(self) -> None: + state = {"messages": [AIMessage(content="Model call limits exceeded: run limit reached")]} + + with ( + patch( + "agent.middleware.notify_step_limit.get_config", + return_value={ + "configurable": {"slack_thread": {"channel_id": "C123", "thread_ts": "171.123"}} + }, + ), + patch( + "agent.middleware.notify_step_limit.post_slack_thread_reply", + new_callable=AsyncMock, + ) as mock_post, + ): + result = await notify_step_limit_reached.aafter_agent(state, self._make_runtime()) + + assert result is None + mock_post.assert_awaited_once() + assert mock_post.await_args.args[0:2] == ("C123", "171.123") + assert "maximum step limit" in mock_post.await_args.args[2] + + @pytest.mark.asyncio + async def test_posts_slack_reply_for_list_content_with_limit_marker(self) -> None: + state = { + "messages": [ + AIMessage( + content=[ + {"type": "text", "text": "Model call limits exceeded:"}, + {"type": "text", "text": "run limit reached"}, + ] + ) + ] + } + + with ( + patch( + "agent.middleware.notify_step_limit.get_config", + return_value={ + "configurable": {"slack_thread": {"channel_id": "C123", "thread_ts": "171.123"}} + }, + ), + patch( + "agent.middleware.notify_step_limit.post_slack_thread_reply", + new_callable=AsyncMock, + ) as mock_post, + ): + result = await notify_step_limit_reached.aafter_agent(state, self._make_runtime()) + + assert result is None + mock_post.assert_awaited_once() + + @pytest.mark.asyncio + async def test_skips_when_limit_marker_absent(self) -> None: + state = {"messages": [HumanMessage(content="keep going")]} + + with patch( + "agent.middleware.notify_step_limit.post_slack_thread_reply", + new_callable=AsyncMock, + ) as mock_post: + result = await notify_step_limit_reached.aafter_agent(state, self._make_runtime()) + + assert result is None + mock_post.assert_not_called() + + @pytest.mark.asyncio + async def test_skips_when_slack_thread_config_missing(self) -> None: + state = {"messages": [AIMessage(content="Model call limits exceeded: run limit reached")]} + + with ( + patch( + "agent.middleware.notify_step_limit.get_config", + return_value={"configurable": {}}, + ), + patch( + "agent.middleware.notify_step_limit.post_slack_thread_reply", + new_callable=AsyncMock, + ) as mock_post, + ): + result = await notify_step_limit_reached.aafter_agent(state, self._make_runtime()) + + assert result is None + mock_post.assert_not_called()