Merge pull request #926 from langchain-ai/yogesh/extract-middleware-to-separate-files

refactor: extract inline middleware from server.py into separate files [closes #900]
This commit is contained in:
Aran Yogesh 2026-02-10 13:36:33 -08:00 • committed by GitHub
commit b09a39210a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 688 additions and 598 deletions

View file

@ -1,3 +1,11 @@
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__ = ["ToolErrorMiddleware"]
__all__ = [
"ToolErrorMiddleware",
"check_message_queue_before_model",
"open_pr_if_needed",
"post_to_linear_after_model",
]

View file

@ -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

View file

@ -0,0 +1,251 @@
"""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
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__)
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(
state: AgentState,
runtime: Runtime,
) -> 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

View file

@ -0,0 +1,115 @@
"""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 ..utils.linear import comment_on_linear_issue
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)
"""
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

View file

@ -6,15 +6,11 @@
import logging
import os
import warnings
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 +27,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
@ -93,597 +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
MIN_MESSAGES_FOR_PREV_CHECK = 2
_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
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
from .utils.sandbox_state import SANDBOX_BACKENDS
async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
@ -943,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", "")

View file

@ -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"

View file

@ -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

View file

@ -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] = {}