From 2b7ce6b69127e0eace18d255606aee17057cce6d Mon Sep 17 00:00:00 2001 From: aran-yogesh Date: Mon, 9 Feb 2026 15:34:18 -0800 Subject: [PATCH 1/3] refactor: extract inline middleware from server.py into separate files --- apps/agent/agent/middleware/__init__.py | 11 +- .../agent/middleware/check_message_queue.py | 106 +++++ apps/agent/agent/middleware/open_pr.py | 254 +++++++++++ apps/agent/agent/middleware/post_to_linear.py | 116 +++++ apps/agent/agent/server.py | 423 +----------------- 5 files changed, 492 insertions(+), 418 deletions(-) create mode 100644 apps/agent/agent/middleware/check_message_queue.py create mode 100644 apps/agent/agent/middleware/open_pr.py create mode 100644 apps/agent/agent/middleware/post_to_linear.py diff --git a/apps/agent/agent/middleware/__init__.py b/apps/agent/agent/middleware/__init__.py index c5cc72d3..0ce6ccde 100644 --- a/apps/agent/agent/middleware/__init__.py +++ b/apps/agent/agent/middleware/__init__.py @@ -1,3 +1,12 @@ +from .check_message_queue import LinearNotifyState, check_message_queue_before_model +from .open_pr import open_pr_if_needed +from .post_to_linear import post_to_linear_after_model from .tool_error_handler import ToolErrorMiddleware -__all__ = ["ToolErrorMiddleware"] +__all__ = [ + "LinearNotifyState", + "ToolErrorMiddleware", + "check_message_queue_before_model", + "open_pr_if_needed", + "post_to_linear_after_model", +] diff --git a/apps/agent/agent/middleware/check_message_queue.py b/apps/agent/agent/middleware/check_message_queue.py new file mode 100644 index 00000000..b85248fd --- /dev/null +++ b/apps/agent/agent/middleware/check_message_queue.py @@ -0,0 +1,106 @@ +"""Before-model middleware that injects queued messages into state. + +Checks the LangGraph store for pending messages (e.g. follow-up Linear +comments that arrived while the agent was busy) and injects them as new +human messages before the next model call. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from langchain.agents.middleware import AgentState, before_model +from langgraph.config import get_config, get_store +from langgraph.runtime import Runtime + +logger = logging.getLogger(__name__) + + +class LinearNotifyState(AgentState): + """Extended agent state for tracking Linear notifications.""" + + linear_messages_sent_count: int + + +@before_model(state_schema=LinearNotifyState) +async def check_message_queue_before_model( # noqa: PLR0911 + state: LinearNotifyState, # noqa: ARG001 + runtime: Runtime, # noqa: ARG001 +) -> dict[str, Any] | None: + """Middleware that checks for queued messages before each model call. + + If messages are found in the queue for this thread, it extracts all messages, + adds them to the conversation state as new human messages, and clears the queue. + Messages are processed in FIFO order (oldest first). + + This enables handling of follow-up comments that arrive while the agent is busy. + The agent will see the new messages and can incorporate them into its response. + """ + try: + config = get_config() + configurable = config.get("configurable", {}) + thread_id = configurable.get("thread_id") + + if not thread_id: + return None + + try: + store = get_store() + except Exception as e: # noqa: BLE001 + logger.debug("Could not get store from context: %s", e) + return None + + if store is None: + return None + + namespace = ("queue", thread_id) + + try: + queued_item = await store.aget(namespace, "pending_messages") + except Exception as e: # noqa: BLE001 + logger.warning("Failed to get queued item: %s", e) + return None + + if queued_item is None: + return None + + queued_value = queued_item.value + queued_messages = queued_value.get("messages", []) + + # Delete early to prevent duplicate processing if middleware runs again + await store.adelete(namespace, "pending_messages") + + if not queued_messages: + return None + + logger.info( + "Found %d queued message(s) for thread %s, injecting into state", + len(queued_messages), + thread_id, + ) + + content_blocks = [ + {"type": "text", "text": msg.get("content", "")} + for msg in queued_messages + if msg.get("content") + ] + + if not content_blocks: + return None + + new_message = { + "role": "user", + "content": content_blocks, + } + + logger.info( + "Injected %d queued message(s) into state for thread %s", + len(content_blocks), + thread_id, + ) + + return {"messages": [new_message]} # noqa: TRY300 + except Exception: + logger.exception("Error in check_message_queue_before_model") + return None diff --git a/apps/agent/agent/middleware/open_pr.py b/apps/agent/agent/middleware/open_pr.py new file mode 100644 index 00000000..93d40e80 --- /dev/null +++ b/apps/agent/agent/middleware/open_pr.py @@ -0,0 +1,254 @@ +"""After-agent middleware that creates a GitHub PR and comments on Linear. + +Runs once after the agent finishes. If the agent called the +``commit_and_open_pr`` tool, this middleware commits any remaining changes, +pushes to a feature branch, opens a GitHub PR, and posts a summary comment +back to the originating Linear issue. +""" + +from __future__ import annotations + +import asyncio +import json as _json +import logging +from typing import Any + +from langchain.agents.middleware import AgentState, after_agent +from langgraph.config import get_config +from langgraph.runtime import Runtime + +logger = logging.getLogger(__name__) + + +def _extract_pr_params_from_messages(messages: list) -> dict[str, str] | None: + """Extract PR title/body/commit_message from the last commit_and_open_pr tool result.""" + for msg in reversed(messages): + if isinstance(msg, dict): + content = msg.get("content", "") + name = msg.get("name", "") + else: + content = getattr(msg, "content", "") + name = getattr(msg, "name", "") + + if name == "commit_and_open_pr" and content: + try: + parsed = _json.loads(content) if isinstance(content, str) else content + if isinstance(parsed, dict) and "title" in parsed: + return parsed + except (ValueError, TypeError): + pass + return None + + +@after_agent +async def open_pr_if_needed( # noqa: PLR0912, PLR0915 + state: AgentState, + runtime: Runtime, # noqa: ARG001 +) -> dict[str, Any] | None: + """Middleware that commits/pushes changes and comments on Linear after agent runs.""" + from ..encryption import decrypt_token + from ..server import ( + _SANDBOX_BACKENDS, + comment_on_linear_issue, + create_github_pr, + get_github_default_branch, + ) + + logger.info("After-agent middleware started") + pr_url = None + pr_number = None + + try: + config = get_config() + configurable = config.get("configurable", {}) + thread_id = configurable.get("thread_id") + logger.debug("Middleware running for thread %s", thread_id) + + last_message_content = "" + messages = state.get("messages", []) + if messages: + last_message = messages[-1] + if isinstance(last_message, dict): + last_message_content = last_message.get("content", "") + elif hasattr(last_message, "content"): + last_message_content = last_message.content + + linear_issue = configurable.get("linear_issue", {}) + linear_issue_id = linear_issue.get("id") + + pr_params = _extract_pr_params_from_messages(messages) + + if not pr_params: + logger.info("No commit_and_open_pr tool call found, skipping PR creation") + if linear_issue_id and last_message_content: + comment = f""" **Agent Response** + +{last_message_content}""" + await comment_on_linear_issue(linear_issue_id, comment) + return None + + pr_title = pr_params.get("title", "feat: Open SWE PR") + pr_body = pr_params.get("body", "Automated PR created by Open SWE agent.") + commit_message = pr_params.get("commit_message", pr_title) + + if not thread_id: + if linear_issue_id and last_message_content: + comment = f"""🤖 **Agent Response** + +{last_message_content}""" + await comment_on_linear_issue(linear_issue_id, comment) + return None + + repo_config = configurable.get("repo", {}) + repo_owner = repo_config.get("owner") + repo_name = repo_config.get("name") + + sandbox_backend = _SANDBOX_BACKENDS.get(thread_id) + + repo_dir = f"/workspace/{repo_name}" + + if not sandbox_backend or not repo_dir: + if linear_issue_id and last_message_content: + comment = f"""🤖 **Agent Response** + +{last_message_content}""" + await comment_on_linear_issue(linear_issue_id, comment) + return None + + result = await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git status --porcelain" + ) + + has_uncommitted_changes = result.exit_code == 0 and result.output.strip() + + await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git fetch origin 2>/dev/null || true" + ) + git_log_cmd = ( + f"cd {repo_dir} && git log --oneline @{{upstream}}..HEAD 2>/dev/null " + "|| git log --oneline origin/HEAD..HEAD 2>/dev/null || echo ''" + ) + unpushed_result = await asyncio.to_thread(sandbox_backend.execute, git_log_cmd) + has_unpushed_commits = unpushed_result.exit_code == 0 and unpushed_result.output.strip() + + has_changes = has_uncommitted_changes or has_unpushed_commits + + if not has_changes: + logger.info("No changes detected, skipping PR creation") + if linear_issue_id and last_message_content: + comment = f"""🤖 **Agent Response** + +{last_message_content}""" + await comment_on_linear_issue(linear_issue_id, comment) + return None + + logger.info("Changes detected, preparing PR for thread %s", thread_id) + + branch_result = await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git rev-parse --abbrev-ref HEAD" + ) + current_branch = branch_result.output.strip() if branch_result.exit_code == 0 else "" + + target_branch = f"open-swe/{thread_id}" + + if current_branch != target_branch: + checkout_result = await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git checkout -b {target_branch}" + ) + if checkout_result.exit_code != 0: + await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git checkout {target_branch}" + ) + + await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git config user.name 'Open SWE[bot]'" + ) + await asyncio.to_thread( + sandbox_backend.execute, + f"cd {repo_dir} && git config user.email 'Open SWE@users.noreply.github.com'", + ) + + await asyncio.to_thread(sandbox_backend.execute, f"cd {repo_dir} && git add -A") + + safe_commit_msg = commit_message.replace("'", "'\\''") + await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git commit -m '{safe_commit_msg}'" + ) + + encrypted_token = configurable.get("github_token_encrypted") + if encrypted_token: + github_token = decrypt_token(encrypted_token) + + if github_token: + remote_result = await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git remote get-url origin" + ) + if remote_result.exit_code == 0: + remote_url = remote_result.output.strip() + if "github.com" in remote_url and "@" not in remote_url: + auth_url = remote_url.replace("https://", f"https://git:{github_token}@") + await asyncio.to_thread( + sandbox_backend.execute, + f"cd {repo_dir} && git push {auth_url} {target_branch}", + ) + else: + await asyncio.to_thread( + sandbox_backend.execute, f"cd {repo_dir} && git push origin {target_branch}" + ) + + base_branch = await get_github_default_branch(repo_owner, repo_name, github_token) + logger.info("Using base branch: %s", base_branch) + + pr_url, pr_number = await create_github_pr( + repo_owner=repo_owner, + repo_name=repo_name, + github_token=github_token, + title=pr_title, + head_branch=target_branch, + base_branch=base_branch, + body=pr_body, + ) + + linear_issue = configurable.get("linear_issue", {}) + linear_issue_id = linear_issue.get("id") + + if linear_issue_id and last_message_content: + if pr_url: + comment = f"""✅ **Pull Request Created** + +I've created a pull request to address this issue: + +**[PR #{pr_number}: {pr_title}]({pr_url})** + +--- + +🤖 **Agent Response** + +{last_message_content}""" + else: + comment = f"""🤖 **Agent Response** + +{last_message_content}""" + await comment_on_linear_issue(linear_issue_id, comment) + + logger.info("After-agent middleware completed successfully") + + except Exception as e: + logger.exception("Error in after-agent middleware") + try: + config = get_config() + configurable = config.get("configurable", {}) + linear_issue = configurable.get("linear_issue", {}) + linear_issue_id = linear_issue.get("id") + if linear_issue_id: + error_comment = f"""❌ **Agent Error** + +An error occurred while processing this issue: + +``` +{type(e).__name__}: {e} +```""" + await comment_on_linear_issue(linear_issue_id, error_comment) + except Exception: + logger.exception("Failed to post error comment to Linear") + return None diff --git a/apps/agent/agent/middleware/post_to_linear.py b/apps/agent/agent/middleware/post_to_linear.py new file mode 100644 index 00000000..52326ac9 --- /dev/null +++ b/apps/agent/agent/middleware/post_to_linear.py @@ -0,0 +1,116 @@ +"""After-model middleware that posts AI responses to Linear. + +Posts the first AI text response back to the originating Linear issue so +stakeholders can see progress without leaving Linear. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from langchain.agents.middleware import after_model +from langgraph.config import get_config +from langgraph.runtime import Runtime + +from .check_message_queue import LinearNotifyState + +logger = logging.getLogger(__name__) + +MIN_MESSAGES_FOR_PREV_CHECK = 2 + + +@after_model(state_schema=LinearNotifyState) +async def post_to_linear_after_model( # noqa: PLR0911, PLR0912 + state: LinearNotifyState, + runtime: Runtime, # noqa: ARG001 +) -> dict[str, Any] | None: + """Middleware that posts AI responses to Linear after each model call. + + Only posts if: + - This is a Linear-triggered conversation (has linear_issue in config) + - There's exactly 1 human message (initial request) + - The previous message was from human (not a tool result) + - The AI response has text content (not just tool calls) + - The message hasn't already been sent (tracked via linear_messages_sent_count) + """ + from ..server import comment_on_linear_issue + + try: + config = get_config() + configurable = config.get("configurable", {}) + + linear_issue = configurable.get("linear_issue", {}) + linear_issue_id = linear_issue.get("id") + + if not linear_issue_id: + return None + + messages = state.get("messages", []) + if not messages: + return None + + sent_count = state.get("linear_messages_sent_count", 0) + + human_message_count = 0 + for msg in messages: + if isinstance(msg, dict): + role = msg.get("role", "") + else: + role = getattr(msg, "type", "") or getattr(msg, "role", "") + if role in ("human", "user"): + human_message_count += 1 + + if human_message_count != 1: + return None + + last_message = messages[-1] + if isinstance(last_message, dict): + role = last_message.get("role", "") + content = last_message.get("content", "") + else: + role = getattr(last_message, "type", "") or getattr(last_message, "role", "") + content = getattr(last_message, "content", "") + + if role not in ("ai", "assistant"): + return None + + ai_message_count = 0 + for msg in messages: + if isinstance(msg, dict): + r = msg.get("role", "") + else: + r = getattr(msg, "type", "") or getattr(msg, "role", "") + if r in ("ai", "assistant"): + ai_message_count += 1 + + if ai_message_count <= sent_count: + return None + + if len(messages) >= MIN_MESSAGES_FOR_PREV_CHECK: + prev_message = messages[-2] + if isinstance(prev_message, dict): + prev_role = prev_message.get("role", "") + else: + prev_role = getattr(prev_message, "type", "") or getattr(prev_message, "role", "") + + if prev_role not in ("human", "user"): + return None + + if not content or not isinstance(content, str): + return None + + comment = f"""🤖 **Agent Response** + +{content}""" + logger.info("Posting AI response to Linear issue %s", linear_issue_id) + success = await comment_on_linear_issue(linear_issue_id, comment) + + if success: + logger.info("Successfully posted to Linear") + return {"linear_messages_sent_count": ai_message_count} + logger.warning("Failed to post to Linear") + + except Exception: + logger.exception("Error in post_to_linear_after_model") + return None diff --git a/apps/agent/agent/server.py b/apps/agent/agent/server.py index 9fa46e8d..1623020c 100644 --- a/apps/agent/agent/server.py +++ b/apps/agent/agent/server.py @@ -10,11 +10,8 @@ from typing import Any logger = logging.getLogger(__name__) -from langchain.agents.middleware import AgentState, after_agent, after_model, before_model -from langgraph.config import get_config, get_store from langgraph.graph.state import RunnableConfig from langgraph.pregel import Pregel -from langgraph.runtime import Runtime from langgraph_sdk import get_client warnings.filterwarnings("ignore", module="langchain_core._api.deprecation") @@ -31,7 +28,12 @@ from deepagents import create_deep_agent from langchain_anthropic import ChatAnthropic from .encryption import decrypt_token -from .middleware import ToolErrorMiddleware +from .middleware import ( + ToolErrorMiddleware, + check_message_queue_before_model, + open_pr_if_needed, + post_to_linear_after_model, +) from .prompt import construct_system_prompt from .protocol import SandboxBackendProtocol from .tools import commit_and_open_pr, fetch_url, http_request @@ -98,8 +100,6 @@ HTTP_CREATED = 201 HTTP_UNPROCESSABLE_ENTITY = 422 # Message count thresholds -MIN_MESSAGES_FOR_PREV_CHECK = 2 - _SANDBOX_BACKENDS: dict[str, Any] = {} import httpx @@ -275,417 +275,6 @@ async def comment_on_linear_issue(issue_id: str, comment_body: str) -> bool: return False -class LinearNotifyState(AgentState): - """Extended agent state for tracking Linear notifications.""" - - linear_messages_sent_count: int - - -@before_model(state_schema=LinearNotifyState) -async def check_message_queue_before_model( # noqa: PLR0911 - state: LinearNotifyState, # noqa: ARG001 - runtime: Runtime, # noqa: ARG001 -) -> dict[str, Any] | None: - """Middleware that checks for queued messages before each model call. - - If messages are found in the queue for this thread, it extracts all messages, - adds them to the conversation state as new human messages, and clears the queue. - Messages are processed in FIFO order (oldest first). - - This enables handling of follow-up comments that arrive while the agent is busy. - The agent will see the new messages and can incorporate them into its response. - """ - try: - config = get_config() - configurable = config.get("configurable", {}) - thread_id = configurable.get("thread_id") - - if not thread_id: - return None - - try: - store = get_store() - except Exception as e: # noqa: BLE001 - logger.debug("Could not get store from context: %s", e) - return None - - if store is None: - return None - - namespace = ("queue", thread_id) - - try: - queued_item = await store.aget(namespace, "pending_messages") - except Exception as e: # noqa: BLE001 - logger.warning("Failed to get queued item: %s", e) - return None - - if queued_item is None: - return None - - queued_value = queued_item.value - queued_messages = queued_value.get("messages", []) - - # Delete early to prevent duplicate processing if middleware runs again - await store.adelete(namespace, "pending_messages") - - if not queued_messages: - return None - - logger.info( - "Found %d queued message(s) for thread %s, injecting into state", - len(queued_messages), - thread_id, - ) - - content_blocks = [ - {"type": "text", "text": msg.get("content", "")} - for msg in queued_messages - if msg.get("content") - ] - - if not content_blocks: - return None - - new_message = { - "role": "user", - "content": content_blocks, - } - - logger.info( - "Injected %d queued message(s) into state for thread %s", - len(content_blocks), - thread_id, - ) - - return {"messages": [new_message]} # noqa: TRY300 - except Exception: - logger.exception("Error in check_message_queue_before_model") - return None - - -@after_model(state_schema=LinearNotifyState) -async def post_to_linear_after_model( # noqa: PLR0911, PLR0912 - state: LinearNotifyState, - runtime: Runtime, # noqa: ARG001 -) -> dict[str, Any] | None: - """Middleware that posts AI responses to Linear after each model call. - - Only posts if: - - This is a Linear-triggered conversation (has linear_issue in config) - - There's exactly 1 human message (initial request) - - The previous message was from human (not a tool result) - - The AI response has text content (not just tool calls) - - The message hasn't already been sent (tracked via linear_messages_sent_count) - """ - try: - config = get_config() - configurable = config.get("configurable", {}) - - linear_issue = configurable.get("linear_issue", {}) - linear_issue_id = linear_issue.get("id") - - if not linear_issue_id: - return None - - messages = state.get("messages", []) - if not messages: - return None - - sent_count = state.get("linear_messages_sent_count", 0) - - human_message_count = 0 - for msg in messages: - if isinstance(msg, dict): - role = msg.get("role", "") - else: - role = getattr(msg, "type", "") or getattr(msg, "role", "") - if role in ("human", "user"): - human_message_count += 1 - - if human_message_count != 1: - return None - - last_message = messages[-1] - if isinstance(last_message, dict): - role = last_message.get("role", "") - content = last_message.get("content", "") - else: - role = getattr(last_message, "type", "") or getattr(last_message, "role", "") - content = getattr(last_message, "content", "") - - if role not in ("ai", "assistant"): - return None - - ai_message_count = 0 - for msg in messages: - if isinstance(msg, dict): - r = msg.get("role", "") - else: - r = getattr(msg, "type", "") or getattr(msg, "role", "") - if r in ("ai", "assistant"): - ai_message_count += 1 - - if ai_message_count <= sent_count: - return None - - if len(messages) >= MIN_MESSAGES_FOR_PREV_CHECK: - prev_message = messages[-2] - if isinstance(prev_message, dict): - prev_role = prev_message.get("role", "") - else: - prev_role = getattr(prev_message, "type", "") or getattr(prev_message, "role", "") - - if prev_role not in ("human", "user"): - return None - - if not content or not isinstance(content, str): - return None - - comment = f"""🤖 **Agent Response** - -{content}""" - logger.info("Posting AI response to Linear issue %s", linear_issue_id) - success = await comment_on_linear_issue(linear_issue_id, comment) - - if success: - logger.info("Successfully posted to Linear") - return {"linear_messages_sent_count": ai_message_count} - logger.warning("Failed to post to Linear") - - except Exception: - logger.exception("Error in post_to_linear_after_model") - return None - - -def _extract_pr_params_from_messages(messages: list) -> dict[str, str] | None: - """Extract PR title/body/commit_message from the last commit_and_open_pr tool result.""" - for msg in reversed(messages): - if isinstance(msg, dict): - content = msg.get("content", "") - name = msg.get("name", "") - else: - content = getattr(msg, "content", "") - name = getattr(msg, "name", "") - - if name == "commit_and_open_pr" and content: - import json as _json - - try: - parsed = _json.loads(content) if isinstance(content, str) else content - if isinstance(parsed, dict) and "title" in parsed: - return parsed - except (ValueError, TypeError): - pass - return None - - -@after_agent -async def open_pr_if_needed( # noqa: PLR0912, PLR0915 - state: AgentState, - runtime: Runtime, # noqa: ARG001 -) -> dict[str, Any] | None: - """Middleware that commits/pushes changes and comments on Linear after agent runs.""" - logger.info("After-agent middleware started") - pr_url = None - pr_number = None - - try: - config = get_config() - configurable = config.get("configurable", {}) - thread_id = configurable.get("thread_id") - logger.debug("Middleware running for thread %s", thread_id) - - last_message_content = "" - messages = state.get("messages", []) - if messages: - last_message = messages[-1] - if isinstance(last_message, dict): - last_message_content = last_message.get("content", "") - elif hasattr(last_message, "content"): - last_message_content = last_message.content - - linear_issue = configurable.get("linear_issue", {}) - linear_issue_id = linear_issue.get("id") - - pr_params = _extract_pr_params_from_messages(messages) - - if not pr_params: - logger.info("No commit_and_open_pr tool call found, skipping PR creation") - if linear_issue_id and last_message_content: - comment = f""" **Agent Response** - -{last_message_content}""" - await comment_on_linear_issue(linear_issue_id, comment) - return None - - pr_title = pr_params.get("title", "feat: Open SWE PR") - pr_body = pr_params.get("body", "Automated PR created by Open SWE agent.") - commit_message = pr_params.get("commit_message", pr_title) - - if not thread_id: - if linear_issue_id and last_message_content: - comment = f"""🤖 **Agent Response** - -{last_message_content}""" - await comment_on_linear_issue(linear_issue_id, comment) - return None - - repo_config = configurable.get("repo", {}) - repo_owner = repo_config.get("owner") - repo_name = repo_config.get("name") - - sandbox_backend = _SANDBOX_BACKENDS.get(thread_id) - - repo_dir = f"/workspace/{repo_name}" - - if not sandbox_backend or not repo_dir: - if linear_issue_id and last_message_content: - comment = f"""🤖 **Agent Response** - -{last_message_content}""" - await comment_on_linear_issue(linear_issue_id, comment) - return None - - result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git status --porcelain" - ) - - has_uncommitted_changes = result.exit_code == 0 and result.output.strip() - - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git fetch origin 2>/dev/null || true" - ) - git_log_cmd = ( - f"cd {repo_dir} && git log --oneline @{{upstream}}..HEAD 2>/dev/null " - "|| git log --oneline origin/HEAD..HEAD 2>/dev/null || echo ''" - ) - unpushed_result = await asyncio.to_thread(sandbox_backend.execute, git_log_cmd) - has_unpushed_commits = unpushed_result.exit_code == 0 and unpushed_result.output.strip() - - has_changes = has_uncommitted_changes or has_unpushed_commits - - if not has_changes: - logger.info("No changes detected, skipping PR creation") - if linear_issue_id and last_message_content: - comment = f"""🤖 **Agent Response** - -{last_message_content}""" - await comment_on_linear_issue(linear_issue_id, comment) - return None - - logger.info("Changes detected, preparing PR for thread %s", thread_id) - - branch_result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git rev-parse --abbrev-ref HEAD" - ) - current_branch = branch_result.output.strip() if branch_result.exit_code == 0 else "" - - target_branch = f"open-swe/{thread_id}" - - if current_branch != target_branch: - checkout_result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git checkout -b {target_branch}" - ) - if checkout_result.exit_code != 0: - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git checkout {target_branch}" - ) - - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git config user.name 'Open SWE[bot]'" - ) - await asyncio.to_thread( - sandbox_backend.execute, - f"cd {repo_dir} && git config user.email 'Open SWE@users.noreply.github.com'", - ) - - await asyncio.to_thread(sandbox_backend.execute, f"cd {repo_dir} && git add -A") - - safe_commit_msg = commit_message.replace("'", "'\\''") - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git commit -m '{safe_commit_msg}'" - ) - - encrypted_token = configurable.get("github_token_encrypted") - if encrypted_token: - github_token = decrypt_token(encrypted_token) - - if github_token: - remote_result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git remote get-url origin" - ) - if remote_result.exit_code == 0: - remote_url = remote_result.output.strip() - if "github.com" in remote_url and "@" not in remote_url: - auth_url = remote_url.replace("https://", f"https://git:{github_token}@") - await asyncio.to_thread( - sandbox_backend.execute, - f"cd {repo_dir} && git push {auth_url} {target_branch}", - ) - else: - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git push origin {target_branch}" - ) - - base_branch = await get_github_default_branch(repo_owner, repo_name, github_token) - logger.info("Using base branch: %s", base_branch) - - pr_url, pr_number = await create_github_pr( - repo_owner=repo_owner, - repo_name=repo_name, - github_token=github_token, - title=pr_title, - head_branch=target_branch, - base_branch=base_branch, - body=pr_body, - ) - - linear_issue = configurable.get("linear_issue", {}) - linear_issue_id = linear_issue.get("id") - - if linear_issue_id and last_message_content: - if pr_url: - comment = f"""✅ **Pull Request Created** - -I've created a pull request to address this issue: - -**[PR #{pr_number}: {pr_title}]({pr_url})** - ---- - -🤖 **Agent Response** - -{last_message_content}""" - else: - comment = f"""🤖 **Agent Response** - -{last_message_content}""" - await comment_on_linear_issue(linear_issue_id, comment) - - logger.info("After-agent middleware completed successfully") - - except Exception as e: - logger.exception("Error in after-agent middleware") - try: - config = get_config() - configurable = config.get("configurable", {}) - linear_issue = configurable.get("linear_issue", {}) - linear_issue_id = linear_issue.get("id") - if linear_issue_id: - error_comment = f"""❌ **Agent Error** - -An error occurred while processing this issue: - -``` -{type(e).__name__}: {e} -```""" - await comment_on_linear_issue(linear_issue_id, error_comment) - except Exception: - logger.exception("Failed to post error comment to Linear") - return None - - async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915 sandbox_backend: SandboxBackendProtocol, owner: str, From 8c9bc67335ae9f913d03b39b9865c6316fe74fad Mon Sep 17 00:00:00 2001 From: aran-yogesh Date: Tue, 10 Feb 2026 11:22:52 -0800 Subject: [PATCH 2/3] refactor: extract shared utils from server.py into agent/utils/ package --- apps/agent/agent/middleware/__init__.py | 3 +- apps/agent/agent/middleware/open_pr.py | 25 ++- apps/agent/agent/middleware/post_to_linear.py | 3 +- apps/agent/agent/server.py | 182 +----------------- apps/agent/agent/utils/__init__.py | 0 apps/agent/agent/utils/github.py | 133 +++++++++++++ apps/agent/agent/utils/linear.py | 58 ++++++ apps/agent/agent/utils/sandbox_state.py | 8 + 8 files changed, 214 insertions(+), 198 deletions(-) create mode 100644 apps/agent/agent/utils/__init__.py create mode 100644 apps/agent/agent/utils/github.py create mode 100644 apps/agent/agent/utils/linear.py create mode 100644 apps/agent/agent/utils/sandbox_state.py diff --git a/apps/agent/agent/middleware/__init__.py b/apps/agent/agent/middleware/__init__.py index 0ce6ccde..250091d4 100644 --- a/apps/agent/agent/middleware/__init__.py +++ b/apps/agent/agent/middleware/__init__.py @@ -1,10 +1,9 @@ -from .check_message_queue import LinearNotifyState, check_message_queue_before_model +from .check_message_queue import check_message_queue_before_model from .open_pr import open_pr_if_needed from .post_to_linear import post_to_linear_after_model from .tool_error_handler import ToolErrorMiddleware __all__ = [ - "LinearNotifyState", "ToolErrorMiddleware", "check_message_queue_before_model", "open_pr_if_needed", diff --git a/apps/agent/agent/middleware/open_pr.py b/apps/agent/agent/middleware/open_pr.py index 93d40e80..0dd29489 100644 --- a/apps/agent/agent/middleware/open_pr.py +++ b/apps/agent/agent/middleware/open_pr.py @@ -17,6 +17,11 @@ from langchain.agents.middleware import AgentState, after_agent from langgraph.config import get_config from langgraph.runtime import Runtime +from ..encryption import decrypt_token +from ..utils.github import create_github_pr, get_github_default_branch +from ..utils.linear import comment_on_linear_issue +from ..utils.sandbox_state import SANDBOX_BACKENDS + logger = logging.getLogger(__name__) @@ -41,19 +46,11 @@ def _extract_pr_params_from_messages(messages: list) -> dict[str, str] | None: @after_agent -async def open_pr_if_needed( # noqa: PLR0912, PLR0915 +async def open_pr_if_needed( state: AgentState, - runtime: Runtime, # noqa: ARG001 + runtime: Runtime, ) -> dict[str, Any] | None: """Middleware that commits/pushes changes and comments on Linear after agent runs.""" - from ..encryption import decrypt_token - from ..server import ( - _SANDBOX_BACKENDS, - comment_on_linear_issue, - create_github_pr, - get_github_default_branch, - ) - logger.info("After-agent middleware started") pr_url = None pr_number = None @@ -103,7 +100,7 @@ async def open_pr_if_needed( # noqa: PLR0912, PLR0915 repo_owner = repo_config.get("owner") repo_name = repo_config.get("name") - sandbox_backend = _SANDBOX_BACKENDS.get(thread_id) + sandbox_backend = SANDBOX_BACKENDS.get(thread_id) repo_dir = f"/workspace/{repo_name}" @@ -214,7 +211,7 @@ async def open_pr_if_needed( # noqa: PLR0912, PLR0915 if linear_issue_id and last_message_content: if pr_url: - comment = f"""✅ **Pull Request Created** + comment = f"""**Pull Request Created** I've created a pull request to address this issue: @@ -226,7 +223,7 @@ I've created a pull request to address this issue: {last_message_content}""" else: - comment = f"""🤖 **Agent Response** + comment = f""" **Agent Response** {last_message_content}""" await comment_on_linear_issue(linear_issue_id, comment) @@ -241,7 +238,7 @@ I've created a pull request to address this issue: linear_issue = configurable.get("linear_issue", {}) linear_issue_id = linear_issue.get("id") if linear_issue_id: - error_comment = f"""❌ **Agent Error** + error_comment = f""" **Agent Error** An error occurred while processing this issue: diff --git a/apps/agent/agent/middleware/post_to_linear.py b/apps/agent/agent/middleware/post_to_linear.py index 52326ac9..c8a1c569 100644 --- a/apps/agent/agent/middleware/post_to_linear.py +++ b/apps/agent/agent/middleware/post_to_linear.py @@ -13,6 +13,7 @@ from langchain.agents.middleware import after_model from langgraph.config import get_config from langgraph.runtime import Runtime +from ..utils.linear import comment_on_linear_issue from .check_message_queue import LinearNotifyState logger = logging.getLogger(__name__) @@ -34,8 +35,6 @@ async def post_to_linear_after_model( # noqa: PLR0911, PLR0912 - The AI response has text content (not just tool calls) - The message hasn't already been sent (tracked via linear_messages_sent_count) """ - from ..server import comment_on_linear_issue - try: config = get_config() configurable = config.get("configurable", {}) diff --git a/apps/agent/agent/server.py b/apps/agent/agent/server.py index 1623020c..31dd5fe1 100644 --- a/apps/agent/agent/server.py +++ b/apps/agent/agent/server.py @@ -6,7 +6,6 @@ import logging import os import warnings -from typing import Any logger = logging.getLogger(__name__) @@ -95,184 +94,7 @@ SANDBOX_CREATING = "__creating__" SANDBOX_CREATION_TIMEOUT = 180 SANDBOX_POLL_INTERVAL = 1.0 -# HTTP status codes -HTTP_CREATED = 201 -HTTP_UNPROCESSABLE_ENTITY = 422 - -# Message count thresholds -_SANDBOX_BACKENDS: dict[str, Any] = {} - -import httpx - -LINEAR_API_KEY = os.environ.get("LINEAR_API_KEY", "") - - -async def create_github_pr( - repo_owner: str, - repo_name: str, - github_token: str, - title: str, - head_branch: str, - base_branch: str, - body: str, -) -> tuple[str | None, int | None]: - """Create a GitHub pull request via the API. - - Args: - repo_owner: Repository owner (e.g., "langchain-ai") - repo_name: Repository name (e.g., "deepagents") - github_token: GitHub access token - title: PR title - head_branch: Source branch name - base_branch: Target branch name - body: PR description - - Returns: - Tuple of (pr_url, pr_number) if successful, (None, None) otherwise - """ - pr_payload = { - "title": title, - "head": head_branch, - "base": base_branch, - "body": body, - } - - logger.info( - "Creating PR: head=%s, base=%s, repo=%s/%s", - head_branch, - base_branch, - repo_owner, - repo_name, - ) - - try: - async with httpx.AsyncClient() as http_client: - pr_response = await http_client.post( - f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls", - headers={ - "Authorization": f"Bearer {github_token}", - "Accept": "application/vnd.github+json", - "X-GitHub-Api-Version": "2022-11-28", - }, - json=pr_payload, - ) - - pr_data = pr_response.json() - - if pr_response.status_code == HTTP_CREATED: - pr_url = pr_data.get("html_url") - pr_number = pr_data.get("number") - logger.info("PR created successfully: %s", pr_url) - return pr_url, pr_number - - if pr_response.status_code == HTTP_UNPROCESSABLE_ENTITY: - logger.error("GitHub API validation error (422): %s", pr_data.get("message")) - else: - logger.error( - "GitHub API error (%s): %s", - pr_response.status_code, - pr_data.get("message"), - ) - - if "errors" in pr_data: - logger.error("GitHub API errors detail: %s", pr_data.get("errors")) - - return None, None - - except httpx.HTTPError: - logger.exception("Failed to create PR via GitHub API") - return None, None - - -async def get_github_default_branch( - repo_owner: str, - repo_name: str, - github_token: str, -) -> str: - """Get the default branch of a GitHub repository via the API. - - Args: - repo_owner: Repository owner (e.g., "langchain-ai") - repo_name: Repository name (e.g., "deepagents") - github_token: GitHub access token - - Returns: - The default branch name (e.g., "main" or "master") - """ - try: - async with httpx.AsyncClient() as http_client: - response = await http_client.get( - f"https://api.github.com/repos/{repo_owner}/{repo_name}", - headers={ - "Authorization": f"Bearer {github_token}", - "Accept": "application/vnd.github+json", - "X-GitHub-Api-Version": "2022-11-28", - }, - ) - - if response.status_code == 200: # noqa: PLR2004 - repo_data = response.json() - default_branch = repo_data.get("default_branch", "main") - logger.debug("Got default branch from GitHub API: %s", default_branch) - return default_branch - - logger.warning( - "Failed to get repo info from GitHub API (%s), falling back to 'main'", - response.status_code, - ) - return "main" - - except httpx.HTTPError: - logger.exception("Failed to get default branch from GitHub API, falling back to 'main'") - return "main" - - -async def comment_on_linear_issue(issue_id: str, comment_body: str) -> bool: - """Add a comment to a Linear issue. - - Args: - issue_id: The Linear issue ID - comment_body: The comment text - - Returns: - True if successful, False otherwise - """ - if not LINEAR_API_KEY: - return False - - import httpx - - url = "https://api.linear.app/graphql" - - mutation = """ - mutation CommentCreate($issueId: String!, $body: String!) { - commentCreate(input: { issueId: $issueId, body: $body }) { - success - comment { - id - } - } - } - """ - - async with httpx.AsyncClient() as http_client: - try: - response = await http_client.post( - url, - headers={ - "Authorization": LINEAR_API_KEY, - "Content-Type": "application/json", - }, - json={ - "query": mutation, - "variables": {"issueId": issue_id, "body": comment_body}, - }, - ) - response.raise_for_status() - result = response.json() - return bool(result.get("data", {}).get("commentCreate", {}).get("success")) - except Exception: # noqa: BLE001 - return False +from .utils.sandbox_state import SANDBOX_BACKENDS async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915 @@ -532,7 +354,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915 logger.exception("Failed to pull repo in existing sandbox") raise - _SANDBOX_BACKENDS[thread_id] = sandbox_backend + SANDBOX_BACKENDS[thread_id] = sandbox_backend linear_issue = config["configurable"].get("linear_issue", {}) linear_project_id = linear_issue.get("linear_project_id", "") diff --git a/apps/agent/agent/utils/__init__.py b/apps/agent/agent/utils/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/apps/agent/agent/utils/github.py b/apps/agent/agent/utils/github.py new file mode 100644 index 00000000..216cf581 --- /dev/null +++ b/apps/agent/agent/utils/github.py @@ -0,0 +1,133 @@ +"""GitHub API utilities.""" + +from __future__ import annotations + +import logging + +import httpx + +logger = logging.getLogger(__name__) + +# HTTP status codes +HTTP_CREATED = 201 +HTTP_UNPROCESSABLE_ENTITY = 422 + + +async def create_github_pr( + repo_owner: str, + repo_name: str, + github_token: str, + title: str, + head_branch: str, + base_branch: str, + body: str, +) -> tuple[str | None, int | None]: + """Create a GitHub pull request via the API. + + Args: + repo_owner: Repository owner (e.g., "langchain-ai") + repo_name: Repository name (e.g., "deepagents") + github_token: GitHub access token + title: PR title + head_branch: Source branch name + base_branch: Target branch name + body: PR description + + Returns: + Tuple of (pr_url, pr_number) if successful, (None, None) otherwise + """ + pr_payload = { + "title": title, + "head": head_branch, + "base": base_branch, + "body": body, + } + + logger.info( + "Creating PR: head=%s, base=%s, repo=%s/%s", + head_branch, + base_branch, + repo_owner, + repo_name, + ) + + try: + async with httpx.AsyncClient() as http_client: + pr_response = await http_client.post( + f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls", + headers={ + "Authorization": f"Bearer {github_token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + }, + json=pr_payload, + ) + + pr_data = pr_response.json() + + if pr_response.status_code == HTTP_CREATED: + pr_url = pr_data.get("html_url") + pr_number = pr_data.get("number") + logger.info("PR created successfully: %s", pr_url) + return pr_url, pr_number + + if pr_response.status_code == HTTP_UNPROCESSABLE_ENTITY: + logger.error("GitHub API validation error (422): %s", pr_data.get("message")) + else: + logger.error( + "GitHub API error (%s): %s", + pr_response.status_code, + pr_data.get("message"), + ) + + if "errors" in pr_data: + logger.error("GitHub API errors detail: %s", pr_data.get("errors")) + + return None, None + + except httpx.HTTPError: + logger.exception("Failed to create PR via GitHub API") + return None, None + + +async def get_github_default_branch( + repo_owner: str, + repo_name: str, + github_token: str, +) -> str: + """Get the default branch of a GitHub repository via the API. + + Args: + repo_owner: Repository owner (e.g., "langchain-ai") + repo_name: Repository name (e.g., "deepagents") + github_token: GitHub access token + + Returns: + The default branch name (e.g., "main" or "master") + """ + try: + async with httpx.AsyncClient() as http_client: + response = await http_client.get( + f"https://api.github.com/repos/{repo_owner}/{repo_name}", + headers={ + "Authorization": f"Bearer {github_token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + }, + ) + + if response.status_code == 200: # noqa: PLR2004 + repo_data = response.json() + default_branch = repo_data.get("default_branch", "main") + logger.debug("Got default branch from GitHub API: %s", default_branch) + return default_branch + + logger.warning( + "Failed to get repo info from GitHub API (%s), falling back to 'main'", + response.status_code, + ) + return "main" + + except httpx.HTTPError: + logger.exception("Failed to get default branch from GitHub API, falling back to 'main'") + return "main" diff --git a/apps/agent/agent/utils/linear.py b/apps/agent/agent/utils/linear.py new file mode 100644 index 00000000..96bb9cd6 --- /dev/null +++ b/apps/agent/agent/utils/linear.py @@ -0,0 +1,58 @@ +"""Linear API utilities.""" + +from __future__ import annotations + +import logging +import os + +import httpx + +logger = logging.getLogger(__name__) + +LINEAR_API_KEY = os.environ.get("LINEAR_API_KEY", "") + + +async def comment_on_linear_issue(issue_id: str, comment_body: str) -> bool: + """Add a comment to a Linear issue. + + Args: + issue_id: The Linear issue ID + comment_body: The comment text + + Returns: + True if successful, False otherwise + """ + if not LINEAR_API_KEY: + return False + + url = "https://api.linear.app/graphql" + + mutation = """ + mutation CommentCreate($issueId: String!, $body: String!) { + commentCreate(input: { issueId: $issueId, body: $body }) { + success + comment { + id + } + } + } + """ + + async with httpx.AsyncClient() as http_client: + try: + response = await http_client.post( + url, + headers={ + "Authorization": LINEAR_API_KEY, + "Content-Type": "application/json", + }, + json={ + "query": mutation, + "variables": {"issueId": issue_id, "body": comment_body}, + }, + ) + response.raise_for_status() + result = response.json() + return bool(result.get("data", {}).get("commentCreate", {}).get("success")) + except Exception: # noqa: BLE001 + return False diff --git a/apps/agent/agent/utils/sandbox_state.py b/apps/agent/agent/utils/sandbox_state.py new file mode 100644 index 00000000..5d8d2b9c --- /dev/null +++ b/apps/agent/agent/utils/sandbox_state.py @@ -0,0 +1,8 @@ +"""Shared sandbox state used by server and middleware.""" + +from __future__ import annotations + +from typing import Any + +# Thread ID -> SandboxBackend mapping, shared between server.py and middleware +SANDBOX_BACKENDS: dict[str, Any] = {} From c1da51f78d90ea8856dbd10495a50619922318a9 Mon Sep 17 00:00:00 2001 From: Aran Yogesh Date: Tue, 10 Feb 2026 13:30:30 -0800 Subject: [PATCH 3/3] Delete apps/agent/agent/utils/__init__.py --- apps/agent/agent/utils/__init__.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) delete mode 100644 apps/agent/agent/utils/__init__.py diff --git a/apps/agent/agent/utils/__init__.py b/apps/agent/agent/utils/__init__.py deleted file mode 100644 index e69de29b..00000000