diff --git a/apps/agent/agent/middleware/__init__.py b/apps/agent/agent/middleware/__init__.py index 5a9f1379..9a08ec69 100644 --- a/apps/agent/agent/middleware/__init__.py +++ b/apps/agent/agent/middleware/__init__.py @@ -1,9 +1,11 @@ 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 .tool_error_handler import ToolErrorMiddleware __all__ = [ "ToolErrorMiddleware", "check_message_queue_before_model", + "ensure_no_empty_msg", "open_pr_if_needed", ] diff --git a/apps/agent/agent/middleware/ensure_no_empty_msg.py b/apps/agent/agent/middleware/ensure_no_empty_msg.py new file mode 100644 index 00000000..4fd77b6b --- /dev/null +++ b/apps/agent/agent/middleware/ensure_no_empty_msg.py @@ -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 diff --git a/apps/agent/agent/server.py b/apps/agent/agent/server.py index 9ed12307..347edda0 100644 --- a/apps/agent/agent/server.py +++ b/apps/agent/agent/server.py @@ -29,6 +29,7 @@ from .integrations.langsmith import create_langsmith_sandbox from .middleware import ( ToolErrorMiddleware, check_message_queue_before_model, + ensure_no_empty_msg, open_pr_if_needed, ) from .prompt import construct_system_prompt @@ -388,6 +389,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915 middleware=[ ToolErrorMiddleware(), check_message_queue_before_model, + ensure_no_empty_msg, open_pr_if_needed, ], ).with_config(config) diff --git a/apps/agent/tests/test_ensure_no_empty_msg.py b/apps/agent/tests/test_ensure_no_empty_msg.py new file mode 100644 index 00000000..3238d7a2 --- /dev/null +++ b/apps/agent/tests/test_ensure_no_empty_msg.py @@ -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