mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-06 17:02:14 +00:00
fix: Middleware to ensure no empty msg (#1010)
This commit is contained in:
parent
a1fcb23498
commit
c0f17d32d8
4 changed files with 240 additions and 0 deletions
|
|
@ -1,9 +1,11 @@
|
||||||
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 .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
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ToolErrorMiddleware",
|
"ToolErrorMiddleware",
|
||||||
"check_message_queue_before_model",
|
"check_message_queue_before_model",
|
||||||
|
"ensure_no_empty_msg",
|
||||||
"open_pr_if_needed",
|
"open_pr_if_needed",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
83
apps/agent/agent/middleware/ensure_no_empty_msg.py
Normal file
83
apps/agent/agent/middleware/ensure_no_empty_msg.py
Normal file
|
|
@ -0,0 +1,83 @@
|
||||||
|
from typing import Any
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from langchain.agents.middleware import AgentState, after_model
|
||||||
|
from langchain_core.messages import AnyMessage, ToolMessage
|
||||||
|
from langgraph.runtime import Runtime
|
||||||
|
|
||||||
|
|
||||||
|
def get_every_message_since_last_human(state: AgentState) -> list[AnyMessage]:
|
||||||
|
messages = state["messages"]
|
||||||
|
last_human_idx = -1
|
||||||
|
for i in range(len(messages) - 1, -1, -1):
|
||||||
|
if messages[i].type == "human":
|
||||||
|
last_human_idx = i
|
||||||
|
break
|
||||||
|
return messages[last_human_idx + 1 :]
|
||||||
|
|
||||||
|
|
||||||
|
def check_if_model_already_called_commit_and_open_pr(messages: list[AnyMessage]) -> bool:
|
||||||
|
for msg in messages:
|
||||||
|
if msg.type == "tool" and msg.name == "commit_and_open_pr":
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def check_if_model_messaged_user(messages: list[AnyMessage]) -> bool:
|
||||||
|
for msg in messages:
|
||||||
|
# see if tool name is one of slack_thread_reply or linear_comment
|
||||||
|
if msg.type == "tool" and msg.name in ["slack_thread_reply", "linear_comment"]:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def check_if_confirming_completion(messages: list[AnyMessage]) -> bool:
|
||||||
|
for msg in messages:
|
||||||
|
if msg.type == "tool" and msg.name == "confirming_completion":
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
@after_model
|
||||||
|
def ensure_no_empty_msg(state: AgentState, runtime: Runtime) -> dict[str, Any] | None:
|
||||||
|
last_msg = state["messages"][-1]
|
||||||
|
has_contents = bool(last_msg.text())
|
||||||
|
has_tool_calls = bool(last_msg.tool_calls)
|
||||||
|
if not has_tool_calls and not has_contents:
|
||||||
|
tc_id = str(uuid4())
|
||||||
|
last_msg.tool_calls = [{"name": "no_op", "args": {}, "id": tc_id}]
|
||||||
|
no_op_tool_msg = ToolMessage(
|
||||||
|
content="No operation performed."
|
||||||
|
+ "Please continue with the task, ensuring you ALWAYS call at least one tool in"
|
||||||
|
+ " every message unless you are absolutely sure the task has been fully completed.",
|
||||||
|
tool_call_id=tc_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
return {"messages": [last_msg, no_op_tool_msg]}
|
||||||
|
|
||||||
|
if has_contents and not has_tool_calls:
|
||||||
|
# See if the model already called open_pr or it sent a slack/linear message
|
||||||
|
# First, get every message since the last human message
|
||||||
|
messages_since_last_human = get_every_message_since_last_human(state)
|
||||||
|
|
||||||
|
# If it opened a PR, we don't need to do anything
|
||||||
|
if (
|
||||||
|
check_if_model_already_called_commit_and_open_pr(messages_since_last_human)
|
||||||
|
or check_if_model_messaged_user(messages_since_last_human)
|
||||||
|
or check_if_confirming_completion(messages_since_last_human)
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
|
||||||
|
tc_id = str(uuid4())
|
||||||
|
last_msg.tool_calls = [{"name": "confirming_completion", "args": {}, "id": tc_id}]
|
||||||
|
no_op_tool_msg = ToolMessage(
|
||||||
|
content="Confirming task completion. I see you did not call a tool, which would end the task, however you haven't called a tool to message the user or open a pull request."
|
||||||
|
+ "This may indicate premature termination - please ensure you fully complete the task before ending it. "
|
||||||
|
+ "If you do not call any tools it will end the task.",
|
||||||
|
name="confirming_completion",
|
||||||
|
tool_call_id=tc_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
return {"messages": [last_msg, no_op_tool_msg]}
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
@ -29,6 +29,7 @@ from .integrations.langsmith import create_langsmith_sandbox
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
ToolErrorMiddleware,
|
ToolErrorMiddleware,
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
|
ensure_no_empty_msg,
|
||||||
open_pr_if_needed,
|
open_pr_if_needed,
|
||||||
)
|
)
|
||||||
from .prompt import construct_system_prompt
|
from .prompt import construct_system_prompt
|
||||||
|
|
@ -388,6 +389,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915
|
||||||
middleware=[
|
middleware=[
|
||||||
ToolErrorMiddleware(),
|
ToolErrorMiddleware(),
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
|
ensure_no_empty_msg,
|
||||||
open_pr_if_needed,
|
open_pr_if_needed,
|
||||||
],
|
],
|
||||||
).with_config(config)
|
).with_config(config)
|
||||||
|
|
|
||||||
153
apps/agent/tests/test_ensure_no_empty_msg.py
Normal file
153
apps/agent/tests/test_ensure_no_empty_msg.py
Normal file
|
|
@ -0,0 +1,153 @@
|
||||||
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||||
|
|
||||||
|
from agent.middleware.ensure_no_empty_msg import (
|
||||||
|
check_if_confirming_completion,
|
||||||
|
check_if_model_already_called_commit_and_open_pr,
|
||||||
|
check_if_model_messaged_user,
|
||||||
|
get_every_message_since_last_human,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetEveryMessageSinceLastHuman:
|
||||||
|
def test_returns_messages_after_last_human(self) -> None:
|
||||||
|
state = {
|
||||||
|
"messages": [
|
||||||
|
HumanMessage(content="first human"),
|
||||||
|
AIMessage(content="ai response"),
|
||||||
|
HumanMessage(content="second human"),
|
||||||
|
AIMessage(content="final ai"),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
result = get_every_message_since_last_human(state)
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0].content == "final ai"
|
||||||
|
|
||||||
|
def test_returns_all_messages_when_no_human(self) -> None:
|
||||||
|
state = {
|
||||||
|
"messages": [
|
||||||
|
AIMessage(content="ai 1"),
|
||||||
|
AIMessage(content="ai 2"),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
result = get_every_message_since_last_human(state)
|
||||||
|
|
||||||
|
assert len(result) == 2
|
||||||
|
assert result[0].content == "ai 1"
|
||||||
|
assert result[1].content == "ai 2"
|
||||||
|
|
||||||
|
def test_returns_empty_when_human_is_last(self) -> None:
|
||||||
|
state = {
|
||||||
|
"messages": [
|
||||||
|
AIMessage(content="ai response"),
|
||||||
|
HumanMessage(content="human last"),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
result = get_every_message_since_last_human(state)
|
||||||
|
|
||||||
|
assert len(result) == 0
|
||||||
|
|
||||||
|
def test_returns_multiple_messages_after_human(self) -> None:
|
||||||
|
state = {
|
||||||
|
"messages": [
|
||||||
|
HumanMessage(content="human"),
|
||||||
|
AIMessage(content="ai 1"),
|
||||||
|
ToolMessage(content="tool result", tool_call_id="123"),
|
||||||
|
AIMessage(content="ai 2"),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
result = get_every_message_since_last_human(state)
|
||||||
|
|
||||||
|
assert len(result) == 3
|
||||||
|
assert result[0].content == "ai 1"
|
||||||
|
assert result[1].content == "tool result"
|
||||||
|
assert result[2].content == "ai 2"
|
||||||
|
|
||||||
|
|
||||||
|
class TestCheckIfModelAlreadyCalledCommitAndOpenPr:
|
||||||
|
def test_returns_true_when_commit_and_open_pr_called(self) -> None:
|
||||||
|
messages = [
|
||||||
|
AIMessage(content="opening pr"),
|
||||||
|
ToolMessage(content="PR opened", tool_call_id="123", name="commit_and_open_pr"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_model_already_called_commit_and_open_pr(messages) is True
|
||||||
|
|
||||||
|
def test_returns_false_when_not_called(self) -> None:
|
||||||
|
messages = [
|
||||||
|
AIMessage(content="doing something"),
|
||||||
|
ToolMessage(content="done", tool_call_id="123", name="bash"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_model_already_called_commit_and_open_pr(messages) is False
|
||||||
|
|
||||||
|
def test_returns_false_for_empty_list(self) -> None:
|
||||||
|
assert check_if_model_already_called_commit_and_open_pr([]) is False
|
||||||
|
|
||||||
|
def test_ignores_non_tool_messages(self) -> None:
|
||||||
|
messages = [
|
||||||
|
AIMessage(content="commit_and_open_pr"),
|
||||||
|
HumanMessage(content="commit_and_open_pr"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_model_already_called_commit_and_open_pr(messages) is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestCheckIfModelMessagedUser:
|
||||||
|
def test_returns_true_for_slack_thread_reply(self) -> None:
|
||||||
|
messages = [
|
||||||
|
ToolMessage(content="sent", tool_call_id="123", name="slack_thread_reply"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_model_messaged_user(messages) is True
|
||||||
|
|
||||||
|
def test_returns_true_for_linear_comment(self) -> None:
|
||||||
|
messages = [
|
||||||
|
ToolMessage(content="commented", tool_call_id="123", name="linear_comment"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_model_messaged_user(messages) is True
|
||||||
|
|
||||||
|
def test_returns_false_for_other_tools(self) -> None:
|
||||||
|
messages = [
|
||||||
|
ToolMessage(content="result", tool_call_id="123", name="bash"),
|
||||||
|
ToolMessage(content="result", tool_call_id="456", name="read_file"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_model_messaged_user(messages) is False
|
||||||
|
|
||||||
|
def test_returns_false_for_empty_list(self) -> None:
|
||||||
|
assert check_if_model_messaged_user([]) is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestCheckIfConfirmingCompletion:
|
||||||
|
def test_returns_true_when_confirming_completion_called(self) -> None:
|
||||||
|
messages = [
|
||||||
|
ToolMessage(content="confirmed", tool_call_id="123", name="confirming_completion"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_confirming_completion(messages) is True
|
||||||
|
|
||||||
|
def test_returns_false_for_other_tools(self) -> None:
|
||||||
|
messages = [
|
||||||
|
ToolMessage(content="result", tool_call_id="123", name="bash"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_confirming_completion(messages) is False
|
||||||
|
|
||||||
|
def test_returns_false_for_empty_list(self) -> None:
|
||||||
|
assert check_if_confirming_completion([]) is False
|
||||||
|
|
||||||
|
def test_finds_confirming_completion_among_other_messages(self) -> None:
|
||||||
|
messages = [
|
||||||
|
AIMessage(content="working"),
|
||||||
|
ToolMessage(content="done", tool_call_id="1", name="bash"),
|
||||||
|
ToolMessage(content="confirmed", tool_call_id="2", name="confirming_completion"),
|
||||||
|
AIMessage(content="finished"),
|
||||||
|
]
|
||||||
|
|
||||||
|
assert check_if_confirming_completion(messages) is True
|
||||||
Loading…
Add table
Reference in a new issue