mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 10:23:14 +00:00
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:
commit
b09a39210a
8 changed files with 688 additions and 598 deletions
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
106
apps/agent/agent/middleware/check_message_queue.py
Normal file
106
apps/agent/agent/middleware/check_message_queue.py
Normal 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
|
||||
251
apps/agent/agent/middleware/open_pr.py
Normal file
251
apps/agent/agent/middleware/open_pr.py
Normal 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
|
||||
115
apps/agent/agent/middleware/post_to_linear.py
Normal file
115
apps/agent/agent/middleware/post_to_linear.py
Normal 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
|
||||
|
|
@ -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", "")
|
||||
|
|
|
|||
133
apps/agent/agent/utils/github.py
Normal file
133
apps/agent/agent/utils/github.py
Normal 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"
|
||||
58
apps/agent/agent/utils/linear.py
Normal file
58
apps/agent/agent/utils/linear.py
Normal 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
|
||||
8
apps/agent/agent/utils/sandbox_state.py
Normal file
8
apps/agent/agent/utils/sandbox_state.py
Normal 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] = {}
|
||||
Loading…
Add table
Reference in a new issue