diff --git a/agent/middleware/open_pr.py b/agent/middleware/open_pr.py index b5f90c13..b02a8509 100644 --- a/agent/middleware/open_pr.py +++ b/agent/middleware/open_pr.py @@ -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( diff --git a/agent/server.py b/agent/server.py index dae949d8..a9f81bb0 100644 --- a/agent/server.py +++ b/agent/server.py @@ -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: diff --git a/agent/tools/commit_and_open_pr.py b/agent/tools/commit_and_open_pr.py index 774740d6..7d5f3452 100644 --- a/agent/tools/commit_and_open_pr.py +++ b/agent/tools/commit_and_open_pr.py @@ -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, diff --git a/agent/tools/fetch_url.py b/agent/tools/fetch_url.py index 4ed35040..978a74da 100644 --- a/agent/tools/fetch_url.py +++ b/agent/tools/fetch_url.py @@ -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 diff --git a/agent/tools/http_request.py b/agent/tools/http_request.py index 908c47ad..942faba8 100644 --- a/agent/tools/http_request.py +++ b/agent/tools/http_request.py @@ -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 ''}" + 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() diff --git a/agent/utils/github.py b/agent/utils/github.py index 29326747..acd3cf59 100644 --- a/agent/utils/github.py +++ b/agent/utils/github.py @@ -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, diff --git a/tests/test_github_security.py b/tests/test_github_security.py new file mode 100644 index 00000000..c89334f7 --- /dev/null +++ b/tests/test_github_security.py @@ -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}", + ] diff --git a/tests/test_http_security.py b/tests/test_http_security.py new file mode 100644 index 00000000..02d36527 --- /dev/null +++ b/tests/test_http_security.py @@ -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"]