fix: Shell injection and ssrf issues (#1155)

* fix: Shell injection and ssrf issues

* cr
This commit is contained in:
Brace Sproul 2026-04-01 12:37:35 -07:00 • committed by GitHub
parent 4856cc31d4
commit fd8e6d98ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 261 additions and 18 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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}",
]

View 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"]