mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 06:53:14 +00:00
fix: Shell injection and ssrf issues (#1155)
* fix: Shell injection and ssrf issues * cr
This commit is contained in:
parent
4856cc31d4
commit
fd8e6d98ee
8 changed files with 261 additions and 18 deletions
|
|
@ -28,6 +28,7 @@ from ..utils.github import (
|
|||
get_github_default_branch,
|
||||
git_add_all,
|
||||
git_checkout_branch,
|
||||
git_checkout_existing_branch,
|
||||
git_commit,
|
||||
git_config_user,
|
||||
git_current_branch,
|
||||
|
|
@ -142,8 +143,7 @@ async def open_pr_if_needed(
|
|||
if branch_name:
|
||||
# Existing branch — plain checkout, do not create or reset
|
||||
await asyncio.to_thread(
|
||||
sandbox_backend.execute,
|
||||
f"cd {repo_dir} && git checkout {target_branch}",
|
||||
git_checkout_existing_branch, sandbox_backend, repo_dir, target_branch
|
||||
)
|
||||
else:
|
||||
await asyncio.to_thread(
|
||||
|
|
|
|||
|
|
@ -70,7 +70,9 @@ from .utils.agents_md import read_agents_md_in_sandbox
|
|||
from .utils.github import (
|
||||
_CRED_FILE_PATH,
|
||||
cleanup_git_credentials,
|
||||
git_current_branch,
|
||||
git_has_uncommitted_changes,
|
||||
git_pull_branch,
|
||||
is_valid_git_repo,
|
||||
remove_directory,
|
||||
setup_git_credentials,
|
||||
|
|
@ -108,7 +110,7 @@ async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
|
|||
work_dir = await aresolve_sandbox_work_dir(sandbox_backend)
|
||||
repo_dir = await aresolve_repo_dir(sandbox_backend, repo)
|
||||
clean_url = f"https://github.com/{owner}/{repo}.git"
|
||||
cred_helper_arg = f"-c credential.helper='store --file={_CRED_FILE_PATH}'"
|
||||
cred_helper = shlex.quote(f"store --file={_CRED_FILE_PATH}")
|
||||
safe_repo_dir = shlex.quote(repo_dir)
|
||||
safe_clean_url = shlex.quote(clean_url)
|
||||
|
||||
|
|
@ -140,12 +142,22 @@ async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
|
|||
|
||||
logger.info("Repo is clean, pulling latest changes from %s/%s", owner, repo)
|
||||
|
||||
await loop.run_in_executor(None, setup_git_credentials, sandbox_backend, token)
|
||||
try:
|
||||
current_branch = await loop.run_in_executor(
|
||||
None, git_current_branch, sandbox_backend, repo_dir
|
||||
)
|
||||
if not current_branch:
|
||||
msg = f"Failed to determine current branch for repo at {repo_dir}"
|
||||
logger.error(msg)
|
||||
raise RuntimeError(msg)
|
||||
|
||||
pull_result = await loop.run_in_executor(
|
||||
None,
|
||||
sandbox_backend.execute,
|
||||
f"cd {repo_dir} && git {cred_helper_arg} pull origin $(git rev-parse --abbrev-ref HEAD)",
|
||||
git_pull_branch,
|
||||
sandbox_backend,
|
||||
repo_dir,
|
||||
current_branch,
|
||||
token,
|
||||
)
|
||||
logger.debug("Git pull result: exit_code=%s", pull_result.exit_code)
|
||||
if pull_result.exit_code != 0:
|
||||
|
|
@ -157,8 +169,6 @@ async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
|
|||
except Exception:
|
||||
logger.exception("Failed to execute git pull")
|
||||
raise
|
||||
finally:
|
||||
await loop.run_in_executor(None, cleanup_git_credentials, sandbox_backend)
|
||||
|
||||
logger.info("Repo updated at %s", repo_dir)
|
||||
return repo_dir
|
||||
|
|
@ -169,7 +179,7 @@ async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
|
|||
result = await loop.run_in_executor(
|
||||
None,
|
||||
sandbox_backend.execute,
|
||||
f"git {cred_helper_arg} clone {safe_clean_url} {safe_repo_dir}",
|
||||
f"git -c credential.helper={cred_helper} clone {safe_clean_url} {safe_repo_dir}",
|
||||
)
|
||||
logger.debug("Git clone result: exit_code=%s", result.exit_code)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from ..utils.github import (
|
|||
get_github_default_branch,
|
||||
git_add_all,
|
||||
git_checkout_branch,
|
||||
git_checkout_existing_branch,
|
||||
git_commit,
|
||||
git_config_user,
|
||||
git_current_branch,
|
||||
|
|
@ -165,7 +166,7 @@ def commit_and_open_pr(
|
|||
if current_branch != target_branch:
|
||||
if branch_name:
|
||||
# Existing branch — plain checkout, do not create or reset
|
||||
result = sandbox_backend.execute(f"cd {repo_dir} && git checkout {target_branch}")
|
||||
result = git_checkout_existing_branch(sandbox_backend, repo_dir, target_branch)
|
||||
if result.exit_code != 0:
|
||||
return {
|
||||
"success": False,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ from typing import Any
|
|||
import requests
|
||||
from markdownify import markdownify
|
||||
|
||||
from .http_request import _request_with_safe_redirects
|
||||
|
||||
|
||||
def fetch_url(url: str, timeout: int = 30) -> dict[str, Any]:
|
||||
"""Fetch content from a URL and convert HTML to markdown format.
|
||||
|
|
@ -30,11 +32,19 @@ def fetch_url(url: str, timeout: int = 30) -> dict[str, Any]:
|
|||
4. NEVER show the raw markdown to the user unless specifically requested
|
||||
"""
|
||||
try:
|
||||
response = requests.get(
|
||||
response, blocked = _request_with_safe_redirects(
|
||||
"GET",
|
||||
url,
|
||||
timeout=timeout,
|
||||
headers={"User-Agent": "Mozilla/5.0 (compatible; DeepAgents/1.0)"},
|
||||
)
|
||||
if blocked:
|
||||
return {
|
||||
"error": blocked["content"],
|
||||
"status_code": blocked["status_code"],
|
||||
"url": blocked["url"],
|
||||
}
|
||||
|
||||
response.raise_for_status()
|
||||
|
||||
# Convert HTML content to markdown
|
||||
|
|
|
|||
|
|
@ -1,15 +1,20 @@
|
|||
import ipaddress
|
||||
import socket
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
from urllib.parse import urljoin, urlparse
|
||||
|
||||
import requests
|
||||
|
||||
_MAX_REDIRECTS = 5
|
||||
|
||||
|
||||
def _is_url_safe(url: str) -> tuple[bool, str]:
|
||||
"""Check if a URL is safe to request (not targeting private/internal networks)."""
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in {"http", "https"}:
|
||||
return False, f"Unsupported URL scheme: {parsed.scheme or '<missing>'}"
|
||||
|
||||
hostname = parsed.hostname
|
||||
if not hostname:
|
||||
return False, "Could not parse hostname from URL"
|
||||
|
|
@ -44,6 +49,54 @@ def _blocked_response(url: str, reason: str) -> dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def _request_with_safe_redirects(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
timeout: int,
|
||||
**kwargs: Any,
|
||||
) -> tuple[requests.Response | None, dict[str, Any] | None]:
|
||||
"""Issue a request while validating every redirect target before following it."""
|
||||
current_method = method.upper()
|
||||
current_url = url
|
||||
request_kwargs = dict(kwargs)
|
||||
|
||||
for redirect_count in range(_MAX_REDIRECTS + 1):
|
||||
is_safe, reason = _is_url_safe(current_url)
|
||||
if not is_safe:
|
||||
return None, _blocked_response(current_url, reason)
|
||||
|
||||
response = requests.request(
|
||||
current_method,
|
||||
current_url,
|
||||
timeout=timeout,
|
||||
allow_redirects=False,
|
||||
**request_kwargs,
|
||||
)
|
||||
|
||||
if not response.is_redirect and not response.is_permanent_redirect:
|
||||
return response, None
|
||||
|
||||
location = response.headers.get("Location")
|
||||
if not location:
|
||||
return response, None
|
||||
|
||||
if redirect_count == _MAX_REDIRECTS:
|
||||
return None, _blocked_response(current_url, "Too many redirects")
|
||||
|
||||
current_url = urljoin(str(response.url), location)
|
||||
|
||||
if response.status_code == requests.codes.see_other or (
|
||||
response.status_code in {requests.codes.moved, requests.codes.found}
|
||||
and current_method not in {"GET", "HEAD"}
|
||||
):
|
||||
current_method = "GET"
|
||||
request_kwargs.pop("data", None)
|
||||
request_kwargs.pop("json", None)
|
||||
|
||||
return None, _blocked_response(current_url, "Too many redirects")
|
||||
|
||||
|
||||
def http_request(
|
||||
url: str,
|
||||
method: str = "GET",
|
||||
|
|
@ -65,10 +118,6 @@ def http_request(
|
|||
Returns:
|
||||
Dictionary with response data including status, headers, and content
|
||||
"""
|
||||
is_safe, reason = _is_url_safe(url)
|
||||
if not is_safe:
|
||||
return _blocked_response(url, reason)
|
||||
|
||||
try:
|
||||
kwargs: dict[str, Any] = {}
|
||||
|
||||
|
|
@ -82,7 +131,14 @@ def http_request(
|
|||
else:
|
||||
kwargs["data"] = data
|
||||
|
||||
response = requests.request(method.upper(), url, timeout=timeout, **kwargs)
|
||||
response, blocked = _request_with_safe_redirects(
|
||||
method,
|
||||
url,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
if blocked:
|
||||
return blocked
|
||||
|
||||
try:
|
||||
content = response.json()
|
||||
|
|
|
|||
|
|
@ -19,7 +19,8 @@ def _run_git(
|
|||
sandbox_backend: SandboxBackendProtocol, repo_dir: str, command: str
|
||||
) -> ExecuteResponse:
|
||||
"""Run a git command in the sandbox repo directory."""
|
||||
return sandbox_backend.execute(f"cd {repo_dir} && {command}")
|
||||
safe_repo_dir = shlex.quote(repo_dir)
|
||||
return sandbox_backend.execute(f"cd {safe_repo_dir} && {command}")
|
||||
|
||||
|
||||
def is_valid_git_repo(sandbox_backend: SandboxBackendProtocol, repo_dir: str) -> bool:
|
||||
|
|
@ -79,6 +80,14 @@ def git_checkout_branch(
|
|||
return fallback.exit_code == 0
|
||||
|
||||
|
||||
def git_checkout_existing_branch(
|
||||
sandbox_backend: SandboxBackendProtocol, repo_dir: str, branch: str
|
||||
) -> ExecuteResponse:
|
||||
"""Checkout an existing branch without creating or resetting it."""
|
||||
safe_branch = shlex.quote(branch)
|
||||
return _run_git(sandbox_backend, repo_dir, f"git checkout {safe_branch}")
|
||||
|
||||
|
||||
def git_config_user(
|
||||
sandbox_backend: SandboxBackendProtocol,
|
||||
repo_dir: str,
|
||||
|
|
@ -158,6 +167,23 @@ def git_push(
|
|||
cleanup_git_credentials(sandbox_backend)
|
||||
|
||||
|
||||
def git_pull_branch(
|
||||
sandbox_backend: SandboxBackendProtocol,
|
||||
repo_dir: str,
|
||||
branch: str,
|
||||
github_token: str | None = None,
|
||||
) -> ExecuteResponse:
|
||||
"""Pull a specific branch from origin, using a token if needed."""
|
||||
safe_branch = shlex.quote(branch)
|
||||
if not github_token:
|
||||
return _run_git(sandbox_backend, repo_dir, f"git pull origin {safe_branch}")
|
||||
setup_git_credentials(sandbox_backend, github_token)
|
||||
try:
|
||||
return _git_with_credentials(sandbox_backend, repo_dir, f"pull origin {safe_branch}")
|
||||
finally:
|
||||
cleanup_git_credentials(sandbox_backend)
|
||||
|
||||
|
||||
async def create_github_pr(
|
||||
repo_owner: str,
|
||||
repo_name: str,
|
||||
|
|
|
|||
48
tests/test_github_security.py
Normal file
48
tests/test_github_security.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import shlex
|
||||
from types import SimpleNamespace
|
||||
|
||||
from agent.utils import github
|
||||
|
||||
|
||||
class FakeSandboxBackend:
|
||||
def __init__(self) -> None:
|
||||
self.commands: list[str] = []
|
||||
self.writes: list[tuple[str, str]] = []
|
||||
|
||||
def execute(self, command: str) -> SimpleNamespace:
|
||||
self.commands.append(command)
|
||||
return SimpleNamespace(exit_code=0, output="")
|
||||
|
||||
def write(self, path: str, content: str) -> None:
|
||||
self.writes.append((path, content))
|
||||
|
||||
|
||||
def test_git_checkout_existing_branch_quotes_repo_dir_and_branch() -> None:
|
||||
sandbox = FakeSandboxBackend()
|
||||
repo_dir = "/tmp/repo; curl attacker"
|
||||
branch = "main; curl attacker"
|
||||
|
||||
github.git_checkout_existing_branch(sandbox, repo_dir, branch)
|
||||
|
||||
assert sandbox.commands == [f"cd {shlex.quote(repo_dir)} && git checkout {shlex.quote(branch)}"]
|
||||
|
||||
|
||||
def test_git_pull_branch_quotes_repo_dir_and_branch_when_using_credentials() -> None:
|
||||
sandbox = FakeSandboxBackend()
|
||||
repo_dir = "/tmp/repo; curl attacker"
|
||||
branch = "main; curl attacker"
|
||||
|
||||
github.git_pull_branch(sandbox, repo_dir, branch, github_token="secret-token")
|
||||
|
||||
assert sandbox.writes == [(github._CRED_FILE_PATH, "https://git:secret-token@github.com\n")]
|
||||
assert sandbox.commands == [
|
||||
f"chmod 600 {github._CRED_FILE_PATH}",
|
||||
(
|
||||
f"cd {shlex.quote(repo_dir)} && git -c "
|
||||
f"credential.helper={shlex.quote(f'store --file={github._CRED_FILE_PATH}')}"
|
||||
f" pull origin {shlex.quote(branch)}"
|
||||
),
|
||||
f"rm -f {github._CRED_FILE_PATH}",
|
||||
]
|
||||
92
tests/test_http_security.py
Normal file
92
tests/test_http_security.py
Normal file
|
|
@ -0,0 +1,92 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
import types
|
||||
|
||||
import requests
|
||||
|
||||
exa_py_stub = types.ModuleType("exa_py")
|
||||
exa_py_stub.Exa = object
|
||||
sys.modules.setdefault("exa_py", exa_py_stub)
|
||||
|
||||
importlib.import_module("agent.tools.fetch_url")
|
||||
importlib.import_module("agent.tools.http_request")
|
||||
fetch_url_tool = sys.modules["agent.tools.fetch_url"]
|
||||
http_request_tool = sys.modules["agent.tools.http_request"]
|
||||
|
||||
_REDIRECT_CODES = {301, 302, 303, 307, 308}
|
||||
_PERMANENT_REDIRECT_CODES = {301, 308}
|
||||
_NO_JSON = object()
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
status_code: int,
|
||||
url: str,
|
||||
headers: dict[str, str] | None = None,
|
||||
text: str = "",
|
||||
json_data: object = _NO_JSON,
|
||||
) -> None:
|
||||
self.status_code = status_code
|
||||
self.url = url
|
||||
self.headers = headers or {}
|
||||
self.text = text
|
||||
self._json_data = json_data
|
||||
|
||||
@property
|
||||
def is_redirect(self) -> bool:
|
||||
return self.status_code in _REDIRECT_CODES and "Location" in self.headers
|
||||
|
||||
@property
|
||||
def is_permanent_redirect(self) -> bool:
|
||||
return self.status_code in _PERMANENT_REDIRECT_CODES and "Location" in self.headers
|
||||
|
||||
def json(self) -> object:
|
||||
if self._json_data is _NO_JSON:
|
||||
raise ValueError("response is not json")
|
||||
return self._json_data
|
||||
|
||||
def raise_for_status(self) -> None:
|
||||
if self.status_code >= 400:
|
||||
raise requests.exceptions.HTTPError(f"{self.status_code} error")
|
||||
|
||||
|
||||
def test_fetch_url_blocks_private_ip_without_issuing_a_request(monkeypatch) -> None:
|
||||
def fail_request(*args, **kwargs): # type: ignore[no-untyped-def]
|
||||
raise AssertionError("request should not be issued for blocked URLs")
|
||||
|
||||
monkeypatch.setattr(http_request_tool.requests, "request", fail_request)
|
||||
|
||||
result = fetch_url_tool.fetch_url(
|
||||
"http://169.254.169.254/latest/meta-data/iam/security-credentials/"
|
||||
)
|
||||
|
||||
assert result["status_code"] == 0
|
||||
assert "Request blocked" in result["error"]
|
||||
assert result["url"].startswith("http://169.254.169.254/")
|
||||
|
||||
|
||||
def test_fetch_url_blocks_redirects_to_private_ips(monkeypatch) -> None:
|
||||
calls: list[tuple[str, str, bool]] = []
|
||||
|
||||
def fake_request(
|
||||
method: str, url: str, *, timeout: int, allow_redirects: bool, **kwargs
|
||||
) -> FakeResponse: # type: ignore[no-untyped-def]
|
||||
calls.append((method, url, allow_redirects))
|
||||
return FakeResponse(
|
||||
status_code=302,
|
||||
url=url,
|
||||
headers={"Location": "http://169.254.169.254/latest/meta-data"},
|
||||
)
|
||||
|
||||
monkeypatch.setattr(http_request_tool.requests, "request", fake_request)
|
||||
|
||||
result = fetch_url_tool.fetch_url("https://example.com/start")
|
||||
|
||||
assert calls == [("GET", "https://example.com/start", False)]
|
||||
assert result["status_code"] == 0
|
||||
assert result["url"] == "http://169.254.169.254/latest/meta-data"
|
||||
assert "Request blocked" in result["error"]
|
||||
Loading…
Add table
Reference in a new issue