mirror of
https://github.com/rookiestar28/ComfyUI-OpenClaw.git
synced 2026-08-14 00:48:07 +00:00
refactor(egress): consolidate safe outbound executor
This commit is contained in:
+272
-241
@@ -6,15 +6,19 @@ S51: Outbound endpoint policy v2 (scheme+port constraints).
|
||||
Any module that touches filesystem or outbound HTTP MUST use this layer.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import http.client
|
||||
import ipaddress
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import tempfile
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, FrozenSet, Optional, Set, Tuple
|
||||
from typing import Any, Callable, Dict, FrozenSet, Optional, Set, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
logger = logging.getLogger("ComfyUI-OpenClaw.services.safe_io")
|
||||
@@ -247,6 +251,203 @@ def _http_error_body_preview(error: Exception, max_bytes: int = 4096) -> Optiona
|
||||
return str(raw)
|
||||
|
||||
|
||||
def _get_pack_version() -> str:
|
||||
try:
|
||||
from ..config import PACK_VERSION
|
||||
except ImportError: # pragma: no cover
|
||||
try:
|
||||
from config import PACK_VERSION # type: ignore
|
||||
except ImportError:
|
||||
PACK_VERSION = "0.0.0"
|
||||
return str(PACK_VERSION)
|
||||
|
||||
|
||||
def _apply_outbound_headers(
|
||||
request: urllib.request.Request,
|
||||
*,
|
||||
headers: Optional[dict],
|
||||
content_type: Optional[str],
|
||||
) -> None:
|
||||
request.add_header("User-Agent", f"ComfyUI-OpenClaw/{_get_pack_version()}")
|
||||
if content_type:
|
||||
request.add_header("Content-Type", content_type)
|
||||
|
||||
if not headers:
|
||||
return
|
||||
|
||||
for key, value in headers.items():
|
||||
key_lower = key.lower()
|
||||
if any(key_lower.startswith(p) for p in ALLOWED_OUTBOUND_HEADER_PREFIXES):
|
||||
request.add_header(key, value)
|
||||
else:
|
||||
logger.debug("Skipping disallowed outbound header.")
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RedirectState:
|
||||
url: str
|
||||
method: str
|
||||
body: Optional[bytes]
|
||||
redirects_followed: int
|
||||
|
||||
|
||||
def _next_redirect_state(
|
||||
*,
|
||||
response: Any,
|
||||
code: int,
|
||||
current_url: str,
|
||||
current_method: str,
|
||||
current_body: Optional[bytes],
|
||||
redirects_followed: int,
|
||||
max_redirects: int,
|
||||
redirect_error_factory: Callable[[int], Exception],
|
||||
) -> Optional[_RedirectState]:
|
||||
if code not in (301, 302, 303, 307, 308):
|
||||
return None
|
||||
if not (max_redirects > 0 and redirects_followed < max_redirects):
|
||||
raise redirect_error_factory(max_redirects)
|
||||
|
||||
new_loc = getattr(response, "headers", {}).get("Location")
|
||||
if not new_loc:
|
||||
raise redirect_error_factory(code)
|
||||
|
||||
next_url = urllib.parse.urljoin(current_url, new_loc)
|
||||
next_method = current_method
|
||||
next_body = current_body
|
||||
if code in (301, 302, 303):
|
||||
next_method = "GET"
|
||||
next_body = None
|
||||
return _RedirectState(
|
||||
url=next_url,
|
||||
method=next_method,
|
||||
body=next_body,
|
||||
redirects_followed=redirects_followed + 1,
|
||||
)
|
||||
|
||||
|
||||
def _raise_safe_io_http_error(error: urllib.error.HTTPError, method: str, url: str):
|
||||
raise SafeIOHTTPError(
|
||||
status_code=error.code,
|
||||
reason=str(getattr(error, "reason", "HTTPError")),
|
||||
method=method,
|
||||
url=url,
|
||||
headers=_headers_to_dict(getattr(error, "headers", None)),
|
||||
body=_http_error_body_preview(error),
|
||||
)
|
||||
|
||||
|
||||
def _raise_safe_fetch_http_error(error: urllib.error.HTTPError, method: str, url: str):
|
||||
raise SSRFError(f"Fetch failed: {error}")
|
||||
|
||||
|
||||
def _raise_safe_request_url_error(error: urllib.error.URLError) -> None:
|
||||
if isinstance(error.reason, SSRFError):
|
||||
raise error.reason
|
||||
raise RuntimeError(f"Request failed: {error}")
|
||||
|
||||
|
||||
def _raise_safe_fetch_url_error(error: urllib.error.URLError) -> None:
|
||||
if isinstance(error.reason, SSRFError):
|
||||
raise error.reason
|
||||
raise SSRFError(f"Fetch failed: {error}")
|
||||
|
||||
|
||||
def _open_outbound_response(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
body: Optional[bytes] = None,
|
||||
allow_hosts: Optional[Set[str]] = None,
|
||||
allow_any_public_host: bool = False,
|
||||
allow_loopback_hosts: Optional[Set[str]] = None,
|
||||
allow_insecure_base_url: bool = False,
|
||||
headers: Optional[dict] = None,
|
||||
content_type: Optional[str] = None,
|
||||
timeout_sec: int = 10,
|
||||
max_redirects: int = 0,
|
||||
policy: Optional[OutboundPolicy] = None,
|
||||
redirect_error_factory: Callable[[int], Exception] = lambda limit: RuntimeError(
|
||||
f"Too many redirects: {limit}"
|
||||
),
|
||||
http_error_mapper: Callable[
|
||||
[urllib.error.HTTPError, str, str], None
|
||||
] = _raise_safe_io_http_error,
|
||||
url_error_mapper: Callable[
|
||||
[urllib.error.URLError], None
|
||||
] = _raise_safe_request_url_error,
|
||||
):
|
||||
current_url = url
|
||||
current_method = method
|
||||
current_body = body
|
||||
redirects_followed = 0
|
||||
|
||||
# CRITICAL: all outbound wrappers must use this seam so every redirect hop is
|
||||
# re-validated and re-pinned before connect; bypassing it reintroduces SSRF drift.
|
||||
while True:
|
||||
_scheme, _host, _port, pinned_ips = validate_outbound_url(
|
||||
current_url,
|
||||
allow_hosts=allow_hosts,
|
||||
allow_any_public_host=allow_any_public_host,
|
||||
allow_loopback_hosts=allow_loopback_hosts,
|
||||
allow_insecure_base_url=allow_insecure_base_url,
|
||||
policy=policy,
|
||||
)
|
||||
|
||||
request = urllib.request.Request(
|
||||
current_url, data=current_body, method=current_method
|
||||
)
|
||||
_apply_outbound_headers(
|
||||
request,
|
||||
headers=headers,
|
||||
content_type=content_type,
|
||||
)
|
||||
opener = _build_pinned_opener(pinned_ips)
|
||||
|
||||
try:
|
||||
response = opener.open(request, timeout=timeout_sec)
|
||||
enter = getattr(response, "__enter__", None)
|
||||
if callable(enter):
|
||||
entered = enter()
|
||||
if entered is not None:
|
||||
response = entered
|
||||
except urllib.error.HTTPError as error:
|
||||
http_error_mapper(error, current_method, current_url)
|
||||
raise AssertionError("http_error_mapper must raise") # pragma: no cover
|
||||
except urllib.error.URLError as error:
|
||||
url_error_mapper(error)
|
||||
raise AssertionError("url_error_mapper must raise") # pragma: no cover
|
||||
|
||||
code = response.getcode()
|
||||
try:
|
||||
redirect_state = _next_redirect_state(
|
||||
response=response,
|
||||
code=code,
|
||||
current_url=current_url,
|
||||
current_method=current_method,
|
||||
current_body=current_body,
|
||||
redirects_followed=redirects_followed,
|
||||
max_redirects=max_redirects,
|
||||
redirect_error_factory=redirect_error_factory,
|
||||
)
|
||||
if redirect_state is None:
|
||||
return response, current_url, current_method
|
||||
except Exception:
|
||||
try:
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
raise
|
||||
|
||||
try:
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
current_url = redirect_state.url
|
||||
current_method = redirect_state.method
|
||||
current_body = redirect_state.body
|
||||
redirects_followed = redirect_state.redirects_followed
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# S51: Outbound Endpoint Policy v2
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -548,61 +749,25 @@ def safe_fetch(
|
||||
"""
|
||||
Safely fetch a URL with SSRF protections and IP pinning.
|
||||
"""
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
|
||||
current_url = url
|
||||
redirects_followed = 0
|
||||
|
||||
while True:
|
||||
# Validate initial URL and resolve IPs
|
||||
scheme, host, port, pinned_ips = validate_outbound_url(
|
||||
current_url, allow_hosts=allow_hosts
|
||||
)
|
||||
|
||||
# Build request
|
||||
request = urllib.request.Request(current_url)
|
||||
response, _current_url, _current_method = _open_outbound_response(
|
||||
"GET",
|
||||
url,
|
||||
allow_hosts=allow_hosts,
|
||||
timeout_sec=timeout_sec,
|
||||
max_redirects=max_redirects,
|
||||
redirect_error_factory=lambda limit: SSRFError(
|
||||
f"Steps limit exceeded or redirects disabled: {limit}"
|
||||
),
|
||||
http_error_mapper=_raise_safe_fetch_http_error,
|
||||
url_error_mapper=_raise_safe_fetch_url_error,
|
||||
)
|
||||
try:
|
||||
return response.read(max_bytes)
|
||||
finally:
|
||||
try:
|
||||
from ..config import PACK_VERSION
|
||||
except ImportError: # pragma: no cover
|
||||
try:
|
||||
from config import PACK_VERSION # type: ignore
|
||||
except ImportError:
|
||||
PACK_VERSION = "0.0.0"
|
||||
|
||||
request.add_header("User-Agent", f"ComfyUI-OpenClaw/{PACK_VERSION}")
|
||||
|
||||
# Build S37-hardened pinned opener
|
||||
opener = _build_pinned_opener(pinned_ips)
|
||||
|
||||
try:
|
||||
with opener.open(request, timeout=timeout_sec) as response:
|
||||
code = response.getcode()
|
||||
|
||||
# Handle redirects manually (NoRedirectHandler returns a 3xx response object).
|
||||
if code in (301, 302, 303, 307, 308):
|
||||
if max_redirects > 0 and redirects_followed < max_redirects:
|
||||
redirects_followed += 1
|
||||
new_loc = response.headers.get("Location")
|
||||
if not new_loc:
|
||||
raise SSRFError(f"Redirect without Location header: {code}")
|
||||
|
||||
# Resolve relative URL
|
||||
current_url = urllib.parse.urljoin(current_url, new_loc)
|
||||
continue
|
||||
raise SSRFError(
|
||||
f"Steps limit exceeded or redirects disabled: {max_redirects}"
|
||||
)
|
||||
|
||||
return response.read(max_bytes)
|
||||
|
||||
except urllib.error.HTTPError as e:
|
||||
# Should mostly catch 4xx/5xx only
|
||||
raise SSRFError(f"Fetch failed: {e}")
|
||||
except urllib.error.URLError as e:
|
||||
if isinstance(e.reason, SSRFError):
|
||||
raise e.reason
|
||||
raise SSRFError(f"Fetch failed: {e}")
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def safe_request_json(
|
||||
@@ -626,108 +791,37 @@ def safe_request_json(
|
||||
Perform a safe HTTP request with JSON body (e.g., POST callback).
|
||||
"""
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
|
||||
current_url = url
|
||||
current_method = method
|
||||
if json_body is not None and raw_body is not None:
|
||||
raise ValueError("safe_request_json accepts either json_body or raw_body")
|
||||
current_body = raw_body
|
||||
if json_body is not None:
|
||||
current_body = json.dumps(json_body).encode("utf-8")
|
||||
redirects_followed = 0
|
||||
|
||||
while True:
|
||||
# Validate URL + Pin IPs
|
||||
# IMPORTANT:
|
||||
# Keep these controls aligned with any caller pre-validation. Divergence
|
||||
# between pre-check and request-time check caused S65 regressions.
|
||||
scheme, host, port, pinned_ips = validate_outbound_url(
|
||||
current_url,
|
||||
allow_hosts=allow_hosts,
|
||||
allow_any_public_host=allow_any_public_host,
|
||||
allow_loopback_hosts=allow_loopback_hosts,
|
||||
allow_insecure_base_url=allow_insecure_base_url,
|
||||
policy=policy,
|
||||
)
|
||||
|
||||
# Build request
|
||||
request = urllib.request.Request(
|
||||
current_url, data=current_body, method=current_method
|
||||
)
|
||||
response, _current_url, _current_method = _open_outbound_response(
|
||||
method,
|
||||
url,
|
||||
body=current_body,
|
||||
allow_hosts=allow_hosts,
|
||||
allow_any_public_host=allow_any_public_host,
|
||||
allow_loopback_hosts=allow_loopback_hosts,
|
||||
allow_insecure_base_url=allow_insecure_base_url,
|
||||
headers=headers,
|
||||
content_type=content_type,
|
||||
timeout_sec=timeout_sec,
|
||||
max_redirects=max_redirects,
|
||||
policy=policy,
|
||||
)
|
||||
try:
|
||||
data = response.read(max_response_bytes)
|
||||
finally:
|
||||
try:
|
||||
from ..config import PACK_VERSION
|
||||
except ImportError: # pragma: no cover
|
||||
try:
|
||||
from config import PACK_VERSION # type: ignore
|
||||
except ImportError:
|
||||
PACK_VERSION = "0.0.0"
|
||||
|
||||
request.add_header("User-Agent", f"ComfyUI-OpenClaw/{PACK_VERSION}")
|
||||
if content_type:
|
||||
request.add_header("Content-Type", content_type)
|
||||
|
||||
# Add safe headers
|
||||
# R106: external control-plane adapter requires Authorization header support.
|
||||
# Keep this allowlist narrow to avoid leaking arbitrary caller headers.
|
||||
if headers:
|
||||
for key, value in headers.items():
|
||||
key_lower = key.lower()
|
||||
if any(
|
||||
key_lower.startswith(p) for p in ALLOWED_OUTBOUND_HEADER_PREFIXES
|
||||
):
|
||||
request.add_header(key, value)
|
||||
else:
|
||||
# IMPORTANT: do not log caller-supplied header names verbatim here.
|
||||
logger.debug("Skipping disallowed outbound header.")
|
||||
|
||||
# Build Pinned Opener
|
||||
opener = _build_pinned_opener(pinned_ips)
|
||||
|
||||
try:
|
||||
with opener.open(request, timeout=timeout_sec) as response:
|
||||
code = response.getcode()
|
||||
|
||||
if code in (301, 302, 303, 307, 308):
|
||||
if max_redirects > 0 and redirects_followed < max_redirects:
|
||||
redirects_followed += 1
|
||||
new_loc = response.headers.get("Location")
|
||||
if not new_loc:
|
||||
raise RuntimeError(f"Redirect without Location: {code}")
|
||||
|
||||
current_url = urllib.parse.urljoin(current_url, new_loc)
|
||||
|
||||
# Handle Method/Body transformation rules
|
||||
if code in (301, 302, 303):
|
||||
current_method = "GET"
|
||||
current_body = None
|
||||
|
||||
continue
|
||||
raise RuntimeError(f"Too many redirects: {max_redirects}")
|
||||
|
||||
data = response.read(max_response_bytes)
|
||||
try:
|
||||
return json.loads(data.decode("utf-8"))
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
return {
|
||||
"raw_response": data.decode("utf-8", errors="replace")[:1000]
|
||||
}
|
||||
|
||||
except urllib.error.HTTPError as e:
|
||||
raise SafeIOHTTPError(
|
||||
status_code=e.code,
|
||||
reason=str(getattr(e, "reason", "HTTPError")),
|
||||
method=current_method,
|
||||
url=current_url,
|
||||
headers=_headers_to_dict(getattr(e, "headers", None)),
|
||||
body=_http_error_body_preview(e),
|
||||
)
|
||||
|
||||
except urllib.error.URLError as e:
|
||||
if isinstance(e.reason, SSRFError):
|
||||
raise e.reason
|
||||
raise RuntimeError(f"Request failed: {e}")
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
return json.loads(data.decode("utf-8"))
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
return {"raw_response": data.decode("utf-8", errors="replace")[:1000]}
|
||||
|
||||
|
||||
def safe_request_text_stream(
|
||||
@@ -751,98 +845,35 @@ def safe_request_text_stream(
|
||||
Intended for SSE/event-stream style provider responses.
|
||||
"""
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
|
||||
current_url = url
|
||||
current_method = method
|
||||
current_body = json.dumps(json_body).encode("utf-8") if json_body else None
|
||||
redirects_followed = 0
|
||||
response, _current_url, _current_method = _open_outbound_response(
|
||||
method,
|
||||
url,
|
||||
body=current_body,
|
||||
allow_hosts=allow_hosts,
|
||||
allow_any_public_host=allow_any_public_host,
|
||||
allow_loopback_hosts=allow_loopback_hosts,
|
||||
allow_insecure_base_url=allow_insecure_base_url,
|
||||
headers=headers,
|
||||
content_type="application/json",
|
||||
timeout_sec=timeout_sec,
|
||||
max_redirects=max_redirects,
|
||||
policy=policy,
|
||||
)
|
||||
|
||||
while True:
|
||||
_scheme, _host, _port, pinned_ips = validate_outbound_url(
|
||||
current_url,
|
||||
allow_hosts=allow_hosts,
|
||||
allow_any_public_host=allow_any_public_host,
|
||||
allow_loopback_hosts=allow_loopback_hosts,
|
||||
allow_insecure_base_url=allow_insecure_base_url,
|
||||
policy=policy,
|
||||
)
|
||||
|
||||
request = urllib.request.Request(
|
||||
current_url, data=current_body, method=current_method
|
||||
)
|
||||
try:
|
||||
while True:
|
||||
line = response.readline(max_line_bytes + 1)
|
||||
if not line:
|
||||
break
|
||||
if len(line) > max_line_bytes:
|
||||
raise RuntimeError(
|
||||
f"Stream line exceeds max_line_bytes ({max_line_bytes})"
|
||||
)
|
||||
yield line.decode("utf-8", errors="replace")
|
||||
finally:
|
||||
try:
|
||||
from ..config import PACK_VERSION
|
||||
except ImportError: # pragma: no cover
|
||||
try:
|
||||
from config import PACK_VERSION # type: ignore
|
||||
except ImportError:
|
||||
PACK_VERSION = "0.0.0"
|
||||
|
||||
request.add_header("User-Agent", f"ComfyUI-OpenClaw/{PACK_VERSION}")
|
||||
request.add_header("Content-Type", "application/json")
|
||||
|
||||
if headers:
|
||||
for key, value in headers.items():
|
||||
key_lower = key.lower()
|
||||
if any(
|
||||
key_lower.startswith(p) for p in ALLOWED_OUTBOUND_HEADER_PREFIXES
|
||||
):
|
||||
request.add_header(key, value)
|
||||
else:
|
||||
logger.debug("Skipping disallowed outbound header.")
|
||||
|
||||
opener = _build_pinned_opener(pinned_ips)
|
||||
|
||||
try:
|
||||
response = opener.open(request, timeout=timeout_sec)
|
||||
code = response.getcode()
|
||||
|
||||
if code in (301, 302, 303, 307, 308):
|
||||
try:
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
if max_redirects > 0 and redirects_followed < max_redirects:
|
||||
redirects_followed += 1
|
||||
new_loc = getattr(response, "headers", {}).get("Location")
|
||||
if not new_loc:
|
||||
raise RuntimeError(f"Redirect without Location: {code}")
|
||||
current_url = urllib.parse.urljoin(current_url, new_loc)
|
||||
if code in (301, 302, 303):
|
||||
current_method = "GET"
|
||||
current_body = None
|
||||
continue
|
||||
raise RuntimeError(f"Too many redirects: {max_redirects}")
|
||||
|
||||
try:
|
||||
while True:
|
||||
line = response.readline(max_line_bytes + 1)
|
||||
if not line:
|
||||
break
|
||||
if len(line) > max_line_bytes:
|
||||
raise RuntimeError(
|
||||
f"Stream line exceeds max_line_bytes ({max_line_bytes})"
|
||||
)
|
||||
yield line.decode("utf-8", errors="replace")
|
||||
finally:
|
||||
try:
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
return
|
||||
|
||||
except urllib.error.HTTPError as e:
|
||||
raise SafeIOHTTPError(
|
||||
status_code=e.code,
|
||||
reason=str(getattr(e, "reason", "HTTPError")),
|
||||
method=current_method,
|
||||
url=current_url,
|
||||
headers=_headers_to_dict(getattr(e, "headers", None)),
|
||||
body=_http_error_body_preview(e),
|
||||
)
|
||||
except urllib.error.URLError as e:
|
||||
if isinstance(e.reason, SSRFError):
|
||||
raise e.reason
|
||||
raise RuntimeError(f"Request failed: {e}")
|
||||
response.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -118,6 +118,24 @@ class TestF58SlackOAuthInstallations(unittest.IsolatedAsyncioTestCase):
|
||||
headers = mock_session.post.call_args.kwargs["headers"]
|
||||
self.assertEqual(headers["Authorization"], "Bearer xoxb-workspace")
|
||||
|
||||
async def test_exchange_code_uses_safe_io_form_encoded_contract(self):
|
||||
with patch(
|
||||
"connector.platforms.slack_installation_manager.safe_request_json",
|
||||
return_value={"ok": True, "access_token": "xoxb-oauth"},
|
||||
) as mock_safe_request:
|
||||
out = await self.manager.exchange_code("oauth-code")
|
||||
|
||||
self.assertEqual(out["ok"], True)
|
||||
kwargs = mock_safe_request.call_args.kwargs
|
||||
self.assertEqual(kwargs["method"], "POST")
|
||||
self.assertEqual(kwargs["url"], "https://slack.com/api/oauth.v2.access")
|
||||
self.assertEqual(kwargs["content_type"], "application/x-www-form-urlencoded")
|
||||
self.assertEqual(kwargs["headers"], {"Accept": "application/json"})
|
||||
self.assertEqual(kwargs["allow_hosts"], {"slack.com"})
|
||||
self.assertIn(b"code=oauth-code", kwargs["raw_body"])
|
||||
self.assertIn(b"client_id=client-id", kwargs["raw_body"])
|
||||
self.assertIn(b"client_secret=client-secret", kwargs["raw_body"])
|
||||
|
||||
async def test_tokens_revoked_event_marks_installation_invalid(self):
|
||||
self.manager.upsert_from_oauth_payload(self._oauth_payload("xoxb-workspace"))
|
||||
router = MagicMock()
|
||||
|
||||
@@ -3,7 +3,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
class TestR167AssetApiContract(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_callback_delivery_preserves_asset_api_only_refs_without_view_fetch(self):
|
||||
async def test_callback_delivery_preserves_asset_api_only_refs_without_view_fetch(
|
||||
self,
|
||||
):
|
||||
import services.callback_delivery as callback_delivery
|
||||
|
||||
sent_payloads = []
|
||||
|
||||
@@ -16,6 +16,7 @@ from services.safe_io import (
|
||||
_normalize_host,
|
||||
is_private_ip,
|
||||
resolve_under_root,
|
||||
safe_fetch,
|
||||
safe_read_bytes,
|
||||
safe_read_text,
|
||||
safe_request_json,
|
||||
@@ -127,6 +128,27 @@ class TestPathSafety(unittest.TestCase):
|
||||
|
||||
class TestURLSafety(unittest.TestCase):
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, code, *, headers=None, body=b"", lines=None):
|
||||
self._code = code
|
||||
self.headers = headers or {}
|
||||
self._body = body
|
||||
self._lines = list(lines or [])
|
||||
|
||||
def getcode(self):
|
||||
return self._code
|
||||
|
||||
def read(self, _max_bytes=None):
|
||||
return self._body
|
||||
|
||||
def readline(self, _max_bytes):
|
||||
if self._lines:
|
||||
return self._lines.pop(0)
|
||||
return b""
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
def test_reject_no_allowlist(self):
|
||||
"""Test that URLs are rejected when no allowlist is provided."""
|
||||
with self.assertRaises(SSRFError) as ctx:
|
||||
@@ -330,6 +352,100 @@ class TestURLSafety(unittest.TestCase):
|
||||
)
|
||||
self.assertEqual(header_map.get("accept"), "application/json")
|
||||
|
||||
@patch("services.safe_io._build_pinned_opener")
|
||||
@patch("services.safe_io.validate_outbound_url")
|
||||
def test_safe_fetch_revalidates_each_redirect_hop(self, mock_validate, mock_build):
|
||||
mock_validate.side_effect = [
|
||||
("https", "example.com", 443, ["93.184.216.34"]),
|
||||
("https", "cdn.example.com", 443, ["93.184.216.35"]),
|
||||
]
|
||||
opener_one = MagicMock()
|
||||
opener_one.open.return_value = self._FakeResponse(
|
||||
302, headers={"Location": "https://cdn.example.com/file.bin"}
|
||||
)
|
||||
opener_two = MagicMock()
|
||||
opener_two.open.return_value = self._FakeResponse(200, body=b"payload")
|
||||
mock_build.side_effect = [opener_one, opener_two]
|
||||
|
||||
out = safe_fetch(
|
||||
"https://example.com/file.bin",
|
||||
allow_hosts={"example.com", "cdn.example.com"},
|
||||
max_redirects=1,
|
||||
)
|
||||
|
||||
self.assertEqual(out, b"payload")
|
||||
self.assertEqual(mock_validate.call_count, 2)
|
||||
self.assertEqual(
|
||||
mock_validate.call_args_list[1].args[0], "https://cdn.example.com/file.bin"
|
||||
)
|
||||
|
||||
@patch("services.safe_io._build_pinned_opener")
|
||||
@patch("services.safe_io.validate_outbound_url")
|
||||
def test_safe_request_json_revalidates_redirect_and_rewrites_post_to_get(
|
||||
self, mock_validate, mock_build
|
||||
):
|
||||
mock_validate.side_effect = [
|
||||
("https", "example.com", 443, ["93.184.216.34"]),
|
||||
("https", "api.example.com", 443, ["93.184.216.35"]),
|
||||
]
|
||||
opener_one = MagicMock()
|
||||
opener_one.open.return_value = self._FakeResponse(
|
||||
302, headers={"Location": "https://api.example.com/final"}
|
||||
)
|
||||
opener_two = MagicMock()
|
||||
opener_two.open.return_value = self._FakeResponse(200, body=b'{"ok": true}')
|
||||
mock_build.side_effect = [opener_one, opener_two]
|
||||
|
||||
out = safe_request_json(
|
||||
method="POST",
|
||||
url="https://example.com/start",
|
||||
json_body={"x": 1},
|
||||
allow_hosts={"example.com", "api.example.com"},
|
||||
max_redirects=1,
|
||||
)
|
||||
|
||||
self.assertEqual(out["ok"], True)
|
||||
self.assertEqual(mock_validate.call_count, 2)
|
||||
second_request = opener_two.open.call_args.args[0]
|
||||
self.assertEqual(second_request.get_method(), "GET")
|
||||
self.assertIsNone(second_request.data)
|
||||
|
||||
@patch("services.safe_io._build_pinned_opener")
|
||||
@patch("services.safe_io.validate_outbound_url")
|
||||
def test_safe_request_text_stream_revalidates_redirect_and_rewrites_post_to_get(
|
||||
self, mock_validate, mock_build
|
||||
):
|
||||
mock_validate.side_effect = [
|
||||
("https", "example.com", 443, ["93.184.216.34"]),
|
||||
("https", "stream.example.com", 443, ["93.184.216.35"]),
|
||||
]
|
||||
opener_one = MagicMock()
|
||||
opener_one.open.return_value = self._FakeResponse(
|
||||
303, headers={"Location": "https://stream.example.com/events"}
|
||||
)
|
||||
opener_two = MagicMock()
|
||||
opener_two.open.return_value = self._FakeResponse(
|
||||
200,
|
||||
lines=[b"data: one\n", b""],
|
||||
)
|
||||
mock_build.side_effect = [opener_one, opener_two]
|
||||
|
||||
lines = list(
|
||||
safe_request_text_stream(
|
||||
method="POST",
|
||||
url="https://example.com/start",
|
||||
json_body={"x": 1},
|
||||
allow_hosts={"example.com", "stream.example.com"},
|
||||
max_redirects=1,
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(lines, ["data: one\n"])
|
||||
self.assertEqual(mock_validate.call_count, 2)
|
||||
second_request = opener_two.open.call_args.args[0]
|
||||
self.assertEqual(second_request.get_method(), "GET")
|
||||
self.assertIsNone(second_request.data)
|
||||
|
||||
@patch("services.safe_io._build_pinned_opener")
|
||||
@patch("services.safe_io.validate_outbound_url")
|
||||
def test_safe_request_text_stream_accept_header_is_allowed(
|
||||
|
||||
Reference in New Issue
Block a user