mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-03 15:03:27 +00:00
fix: notify users via Slack when agent hits model call step limit (#1204)
* fix: notify users via Slack when agent hits model call step limit - Root cause: GraphRecursionError at 1000 steps bypassed all @after_agent middleware including open_pr_if_needed, leaving users with no notification - Change: Added ModelCallLimitMiddleware(run_limit=60) to intercept gracefully before the hard recursion limit, and added notify_step_limit_reached @after_agent middleware to post a Slack thread reply when the limit fires - Verified: 107 existing tests pass, no regressions * fix: harden step-limit Slack notification Ensure the step-limit notification runs after the PR safety net and cover the new middleware behavior with focused unit tests. --------- Co-authored-by: LangSmith Forge <forge-agent@langsmith.ai> Co-authored-by: Johannes du Plessis <johannes@langchain.dev>
This commit is contained in:
parent
0464126c83
commit
e5bc27a0ad
4 changed files with 191 additions and 0 deletions
|
|
@ -1,5 +1,6 @@
|
||||||
from .check_message_queue import check_message_queue_before_model
|
from .check_message_queue import check_message_queue_before_model
|
||||||
from .ensure_no_empty_msg import ensure_no_empty_msg
|
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 .open_pr import open_pr_if_needed
|
||||||
from .tool_error_handler import ToolErrorMiddleware
|
from .tool_error_handler import ToolErrorMiddleware
|
||||||
|
|
||||||
|
|
@ -7,5 +8,6 @@ __all__ = [
|
||||||
"ToolErrorMiddleware",
|
"ToolErrorMiddleware",
|
||||||
"check_message_queue_before_model",
|
"check_message_queue_before_model",
|
||||||
"ensure_no_empty_msg",
|
"ensure_no_empty_msg",
|
||||||
|
"notify_step_limit_reached",
|
||||||
"open_pr_if_needed",
|
"open_pr_if_needed",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
88
agent/middleware/notify_step_limit.py
Normal file
88
agent/middleware/notify_step_limit.py
Normal file
|
|
@ -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
|
||||||
|
|
@ -23,6 +23,7 @@ warnings.filterwarnings("ignore", message=".*Pydantic V1.*", category=UserWarnin
|
||||||
# Now safe to import agent (which imports LangChain modules)
|
# Now safe to import agent (which imports LangChain modules)
|
||||||
from deepagents import create_deep_agent
|
from deepagents import create_deep_agent
|
||||||
from deepagents.backends.protocol import SandboxBackendProtocol
|
from deepagents.backends.protocol import SandboxBackendProtocol
|
||||||
|
from langchain.agents.middleware import ModelCallLimitMiddleware
|
||||||
from langsmith.sandbox import SandboxClientError
|
from langsmith.sandbox import SandboxClientError
|
||||||
|
|
||||||
from .integrations.langsmith import _configure_github_proxy
|
from .integrations.langsmith import _configure_github_proxy
|
||||||
|
|
@ -30,6 +31,7 @@ from .middleware import (
|
||||||
ToolErrorMiddleware,
|
ToolErrorMiddleware,
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
ensure_no_empty_msg,
|
ensure_no_empty_msg,
|
||||||
|
notify_step_limit_reached,
|
||||||
open_pr_if_needed,
|
open_pr_if_needed,
|
||||||
)
|
)
|
||||||
from .prompt import construct_system_prompt
|
from .prompt import construct_system_prompt
|
||||||
|
|
@ -322,9 +324,12 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
||||||
],
|
],
|
||||||
backend=sandbox_backend,
|
backend=sandbox_backend,
|
||||||
middleware=[
|
middleware=[
|
||||||
|
ModelCallLimitMiddleware(run_limit=60, exit_behavior="end"),
|
||||||
ToolErrorMiddleware(),
|
ToolErrorMiddleware(),
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
ensure_no_empty_msg,
|
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,
|
open_pr_if_needed,
|
||||||
],
|
],
|
||||||
).with_config(config)
|
).with_config(config)
|
||||||
|
|
|
||||||
96
tests/test_notify_step_limit_middleware.py
Normal file
96
tests/test_notify_step_limit_middleware.py
Normal file
|
|
@ -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()
|
||||||
Loading…
Add table
Reference in a new issue