mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-06 19:22:13 +00:00
feat: add size caps for PR diff, fetch_url, Slack threads, pagination, message queue [closes OPE-51] (#1567)
* feat: add size caps for PR diff, fetch_url, Slack threads, pagination, message queue Per-source byte/token caps with explicit truncation markers to prevent unbounded payloads from blowing up LLM context/memory. - reviewer_diff.py: cap PR diff at 200K chars with head+tail truncation - fetch_url.py: cap markdownify output at 100K chars - slack.py: cap thread message fetch at 500 messages - github_comments.py: cap _fetch_paginated at 50 pages - thread_ops.py: cap queued messages at 100 (drop oldest) Closes OPE-51 Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com> * fix: compute diff line set from full diff, keep most recent Slack messages Address PR review comments: 1. Truncated diffs rejected valid findings: fetch_pr_diff now returns the full diff; truncate_diff is called separately in reviewer.py so the line set used for add_finding/publish_review validation is computed from the complete diff, not the truncated prompt text. 2. Slack cap dropped recent thread context: fetch_slack_thread_messages now keeps the most recent SLACK_THREAD_MAX_MESSAGES messages (was keeping the oldest). The tool surfaces a truncation marker in the formatted output so the LLM knows the thread was truncated. Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com> --------- Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
parent
7dd758f845
commit
98b824bd54
7 changed files with 81 additions and 9 deletions
|
|
@ -47,7 +47,7 @@ from .middleware import (
|
||||||
refresh_github_proxy_before_model,
|
refresh_github_proxy_before_model,
|
||||||
settle_review_check_on_exit,
|
settle_review_check_on_exit,
|
||||||
)
|
)
|
||||||
from .reviewer_diff import compute_diff_line_set, fetch_pr_diff, fetch_pr_metadata
|
from .reviewer_diff import compute_diff_line_set, fetch_pr_diff, fetch_pr_metadata, truncate_diff
|
||||||
from .reviewer_findings import (
|
from .reviewer_findings import (
|
||||||
list_findings as list_findings_async,
|
list_findings as list_findings_async,
|
||||||
)
|
)
|
||||||
|
|
@ -847,9 +847,14 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel:
|
||||||
and bool(github_api_token)
|
and bool(github_api_token)
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _fetch_diff_context() -> tuple[str, dict[str, dict[str, set[int]]] | None]:
|
async def _fetch_diff_context() -> tuple[str, str, dict[str, dict[str, set[int]]] | None]:
|
||||||
|
"""Return (full_diff, truncated_diff, line_set).
|
||||||
|
|
||||||
|
The line set is computed from the full diff so findings on any changed
|
||||||
|
line are accepted. Only the prompt-facing text is truncated.
|
||||||
|
"""
|
||||||
if not can_fetch_pr or github_api_token is None or not isinstance(pr_number, int):
|
if not can_fetch_pr or github_api_token is None or not isinstance(pr_number, int):
|
||||||
return "", None
|
return "", "", None
|
||||||
fetched_diff = await fetch_pr_diff(
|
fetched_diff = await fetch_pr_diff(
|
||||||
owner=repo_owner,
|
owner=repo_owner,
|
||||||
repo=repo_name,
|
repo=repo_name,
|
||||||
|
|
@ -857,8 +862,8 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel:
|
||||||
token=github_api_token,
|
token=github_api_token,
|
||||||
)
|
)
|
||||||
if fetched_diff is None:
|
if fetched_diff is None:
|
||||||
return "", None
|
return "", "", None
|
||||||
return fetched_diff, compute_diff_line_set(fetched_diff)
|
return fetched_diff, truncate_diff(fetched_diff), compute_diff_line_set(fetched_diff)
|
||||||
|
|
||||||
async def _fetch_pr_overview() -> tuple[str, str]:
|
async def _fetch_pr_overview() -> tuple[str, str]:
|
||||||
if not can_fetch_pr or github_api_token is None or not isinstance(pr_number, int):
|
if not can_fetch_pr or github_api_token is None or not isinstance(pr_number, int):
|
||||||
|
|
@ -952,7 +957,7 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel:
|
||||||
_fetch_org_guidelines(),
|
_fetch_org_guidelines(),
|
||||||
fetch_api_standards_skill(),
|
fetch_api_standards_skill(),
|
||||||
)
|
)
|
||||||
pr_diff_text, pr_diff_line_set = diff_context
|
_, pr_diff_text, pr_diff_line_set = diff_context
|
||||||
pr_title, pr_body = pr_overview
|
pr_title, pr_body = pr_overview
|
||||||
config["configurable"]["diff_text"] = pr_diff_text
|
config["configurable"]["diff_text"] = pr_diff_text
|
||||||
config["configurable"]["diff_line_set"] = pr_diff_line_set
|
config["configurable"]["diff_line_set"] = pr_diff_line_set
|
||||||
|
|
|
||||||
|
|
@ -197,6 +197,10 @@ def is_range_in_diff(
|
||||||
return all(line in side_lines for line in range(start_line, end_line + 1))
|
return all(line in side_lines for line in range(start_line, end_line + 1))
|
||||||
|
|
||||||
|
|
||||||
|
PR_DIFF_MAX_CHARS = 200_000
|
||||||
|
PR_DIFF_TRUNCATION_MARKER = "\n... [PR diff truncated: {kept}/{total} chars]\n"
|
||||||
|
|
||||||
|
|
||||||
async def fetch_pr_diff(
|
async def fetch_pr_diff(
|
||||||
*,
|
*,
|
||||||
owner: str,
|
owner: str,
|
||||||
|
|
@ -211,6 +215,9 @@ async def fetch_pr_diff(
|
||||||
validates against when posting inline review comments, so it's the right
|
validates against when posting inline review comments, so it's the right
|
||||||
source for ``add_finding``'s in-diff anchor validation and for
|
source for ``add_finding``'s in-diff anchor validation and for
|
||||||
``publish_review``'s 422 retry filter.
|
``publish_review``'s 422 retry filter.
|
||||||
|
|
||||||
|
The full diff is returned — callers must use ``truncate_diff`` if they
|
||||||
|
need to cap the text before feeding it to the LLM.
|
||||||
"""
|
"""
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
|
|
@ -233,6 +240,19 @@ async def fetch_pr_diff(
|
||||||
return response.text
|
return response.text
|
||||||
|
|
||||||
|
|
||||||
|
def truncate_diff(diff_text: str) -> str:
|
||||||
|
"""Truncate diff text to ``PR_DIFF_MAX_CHARS`` with a visible marker.
|
||||||
|
|
||||||
|
Keeps the first half from the beginning and the second half from the end
|
||||||
|
so file headers and recent changes are both visible.
|
||||||
|
"""
|
||||||
|
if len(diff_text) <= PR_DIFF_MAX_CHARS:
|
||||||
|
return diff_text
|
||||||
|
half = PR_DIFF_MAX_CHARS // 2
|
||||||
|
marker = PR_DIFF_TRUNCATION_MARKER.format(kept=PR_DIFF_MAX_CHARS, total=len(diff_text))
|
||||||
|
return diff_text[:half] + marker + diff_text[-half:]
|
||||||
|
|
||||||
|
|
||||||
async def fetch_pr_metadata(
|
async def fetch_pr_metadata(
|
||||||
*,
|
*,
|
||||||
owner: str,
|
owner: str,
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,8 @@ from markdownify import markdownify
|
||||||
|
|
||||||
from .http_request import _request_with_safe_redirects
|
from .http_request import _request_with_safe_redirects
|
||||||
|
|
||||||
|
FETCH_URL_MAX_CHARS = 100_000
|
||||||
|
|
||||||
|
|
||||||
def fetch_url(url: str, timeout: int = 30) -> dict[str, Any]:
|
def fetch_url(url: str, timeout: int = 30) -> dict[str, Any]:
|
||||||
"""Fetch content from a URL and convert HTML to markdown format.
|
"""Fetch content from a URL and convert HTML to markdown format.
|
||||||
|
|
@ -50,6 +52,12 @@ def fetch_url(url: str, timeout: int = 30) -> dict[str, Any]:
|
||||||
# Convert HTML content to markdown
|
# Convert HTML content to markdown
|
||||||
markdown_content = markdownify(response.text)
|
markdown_content = markdownify(response.text)
|
||||||
|
|
||||||
|
if len(markdown_content) > FETCH_URL_MAX_CHARS:
|
||||||
|
markdown_content = (
|
||||||
|
markdown_content[:FETCH_URL_MAX_CHARS] + "\n... [content truncated: "
|
||||||
|
f"{FETCH_URL_MAX_CHARS}/{len(markdown_content)} chars]\n"
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"url": str(response.url),
|
"url": str(response.url),
|
||||||
"markdown_content": markdown_content,
|
"markdown_content": markdown_content,
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ import asyncio
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ..utils.slack import (
|
from ..utils.slack import (
|
||||||
|
SLACK_THREAD_MAX_MESSAGES,
|
||||||
fetch_slack_thread_messages,
|
fetch_slack_thread_messages,
|
||||||
format_slack_messages_for_prompt,
|
format_slack_messages_for_prompt,
|
||||||
get_slack_user_names,
|
get_slack_user_names,
|
||||||
|
|
@ -19,8 +20,18 @@ async def _fetch_and_format(channel_id: str, message_ts: str) -> dict[str, Any]:
|
||||||
]
|
]
|
||||||
user_names = await get_slack_user_names(user_ids) if user_ids else {}
|
user_names = await get_slack_user_names(user_ids) if user_ids else {}
|
||||||
|
|
||||||
|
truncated = len(messages) >= SLACK_THREAD_MAX_MESSAGES
|
||||||
formatted = format_slack_messages_for_prompt(messages, user_names)
|
formatted = format_slack_messages_for_prompt(messages, user_names)
|
||||||
return {"success": True, "formatted": formatted, "count": len(messages)}
|
if truncated:
|
||||||
|
formatted = (
|
||||||
|
f"[thread truncated — showing most recent {len(messages)} messages]\n{formatted}"
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"formatted": formatted,
|
||||||
|
"count": len(messages),
|
||||||
|
"truncated": truncated,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def slack_read_thread_messages(channel_id: str, message_ts: str) -> dict[str, Any]:
|
def slack_read_thread_messages(channel_id: str, message_ts: str) -> dict[str, Any]:
|
||||||
|
|
|
||||||
|
|
@ -44,6 +44,8 @@ _REACTION_ENDPOINTS: dict[str, str] = {
|
||||||
"pull_request_review": "https://api.github.com/repos/{owner}/{repo}/pulls/{pull_number}/reviews/{comment_id}/reactions",
|
"pull_request_review": "https://api.github.com/repos/{owner}/{repo}/pulls/{pull_number}/reviews/{comment_id}/reactions",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
PAGINATED_MAX_PAGES = 50
|
||||||
|
|
||||||
|
|
||||||
def verify_github_signature(body: bytes, signature: str, *, secret: str) -> bool:
|
def verify_github_signature(body: bytes, signature: str, *, secret: str) -> bool:
|
||||||
"""Verify the GitHub webhook signature (X-Hub-Signature-256).
|
"""Verify the GitHub webhook signature (X-Hub-Signature-256).
|
||||||
|
|
@ -464,6 +466,9 @@ async def _fetch_paginated(
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Fetch all pages from a GitHub paginated endpoint.
|
"""Fetch all pages from a GitHub paginated endpoint.
|
||||||
|
|
||||||
|
Caps at ``PAGINATED_MAX_PAGES`` pages to avoid unbounded fetching on
|
||||||
|
pathological PRs with thousands of comments.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
client: An active httpx async client.
|
client: An active httpx async client.
|
||||||
url: The GitHub API endpoint URL.
|
url: The GitHub API endpoint URL.
|
||||||
|
|
@ -475,7 +480,7 @@ async def _fetch_paginated(
|
||||||
results: list[dict[str, Any]] = []
|
results: list[dict[str, Any]] = []
|
||||||
params: dict[str, Any] = {"per_page": 100, "page": 1}
|
params: dict[str, Any] = {"per_page": 100, "page": 1}
|
||||||
|
|
||||||
while True:
|
while params["page"] <= PAGINATED_MAX_PAGES:
|
||||||
try:
|
try:
|
||||||
response = await client.get(url, headers=headers, params=params)
|
response = await client.get(url, headers=headers, params=params)
|
||||||
if response.status_code == 401:
|
if response.status_code == 401:
|
||||||
|
|
|
||||||
|
|
@ -24,6 +24,7 @@ logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
SLACK_API_BASE_URL = "https://slack.com/api"
|
SLACK_API_BASE_URL = "https://slack.com/api"
|
||||||
SLACK_BOT_TOKEN = os.environ.get("SLACK_BOT_TOKEN", "")
|
SLACK_BOT_TOKEN = os.environ.get("SLACK_BOT_TOKEN", "")
|
||||||
|
SLACK_THREAD_MAX_MESSAGES = 500
|
||||||
DEFAULT_ASSISTANT_STATUS = "is thinking…"
|
DEFAULT_ASSISTANT_STATUS = "is thinking…"
|
||||||
|
|
||||||
# Curated rotating loading strings shown by Slack while the indicator is active.
|
# Curated rotating loading strings shown by Slack while the indicator is active.
|
||||||
|
|
@ -504,12 +505,13 @@ async def get_slack_user_names(user_ids: list[str]) -> dict[str, str]:
|
||||||
|
|
||||||
|
|
||||||
async def fetch_slack_thread_messages(channel_id: str, thread_ts: str) -> list[dict[str, Any]]:
|
async def fetch_slack_thread_messages(channel_id: str, thread_ts: str) -> list[dict[str, Any]]:
|
||||||
"""Fetch all messages for a Slack thread."""
|
"""Fetch messages for a Slack thread, keeping the most recent window."""
|
||||||
if not SLACK_BOT_TOKEN:
|
if not SLACK_BOT_TOKEN:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
messages: list[dict[str, Any]] = []
|
messages: list[dict[str, Any]] = []
|
||||||
cursor: str | None = None
|
cursor: str | None = None
|
||||||
|
truncated = False
|
||||||
|
|
||||||
async with httpx.AsyncClient() as http_client:
|
async with httpx.AsyncClient() as http_client:
|
||||||
while True:
|
while True:
|
||||||
|
|
@ -537,6 +539,16 @@ async def fetch_slack_thread_messages(channel_id: str, thread_ts: str) -> list[d
|
||||||
if isinstance(batch, list):
|
if isinstance(batch, list):
|
||||||
messages.extend(item for item in batch if isinstance(item, dict))
|
messages.extend(item for item in batch if isinstance(item, dict))
|
||||||
|
|
||||||
|
if len(messages) >= SLACK_THREAD_MAX_MESSAGES:
|
||||||
|
truncated = True
|
||||||
|
logger.warning(
|
||||||
|
"Slack thread %s/%s capped at %d messages",
|
||||||
|
channel_id,
|
||||||
|
thread_ts,
|
||||||
|
SLACK_THREAD_MAX_MESSAGES,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
|
||||||
response_metadata = payload.get("response_metadata", {})
|
response_metadata = payload.get("response_metadata", {})
|
||||||
cursor = (
|
cursor = (
|
||||||
response_metadata.get("next_cursor") if isinstance(response_metadata, dict) else ""
|
response_metadata.get("next_cursor") if isinstance(response_metadata, dict) else ""
|
||||||
|
|
@ -544,6 +556,8 @@ async def fetch_slack_thread_messages(channel_id: str, thread_ts: str) -> list[d
|
||||||
if not cursor:
|
if not cursor:
|
||||||
break
|
break
|
||||||
|
|
||||||
|
if truncated:
|
||||||
|
messages = messages[-SLACK_THREAD_MAX_MESSAGES:]
|
||||||
messages.sort(key=lambda item: _parse_ts(item.get("ts")))
|
messages.sort(key=lambda item: _parse_ts(item.get("ts")))
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,8 @@ from langgraph_sdk import get_client
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
MAX_QUEUED_MESSAGES = 100
|
||||||
|
|
||||||
|
|
||||||
def langgraph_url() -> str:
|
def langgraph_url() -> str:
|
||||||
return os.environ.get("LANGGRAPH_URL") or os.environ.get(
|
return os.environ.get("LANGGRAPH_URL") or os.environ.get(
|
||||||
|
|
@ -57,6 +59,13 @@ async def queue_message_for_thread(
|
||||||
logger.debug("No existing queued messages for thread %s", thread_id)
|
logger.debug("No existing queued messages for thread %s", thread_id)
|
||||||
|
|
||||||
existing_messages.append(new_message)
|
existing_messages.append(new_message)
|
||||||
|
if len(existing_messages) > MAX_QUEUED_MESSAGES:
|
||||||
|
existing_messages = existing_messages[-MAX_QUEUED_MESSAGES:]
|
||||||
|
logger.warning(
|
||||||
|
"Thread %s queue capped at %d messages (dropped oldest)",
|
||||||
|
thread_id,
|
||||||
|
MAX_QUEUED_MESSAGES,
|
||||||
|
)
|
||||||
await client.store.put_item(namespace, key, {"messages": existing_messages})
|
await client.store.put_item(namespace, key, {"messages": existing_messages})
|
||||||
logger.info(
|
logger.info(
|
||||||
"Queued message for thread %s (total queued: %d)",
|
"Queued message for thread %s (total queued: %d)",
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue