mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-09-11 02:32:20 +02:00
fix(skills): harden skill import against DNS rebinding and SSRF TOCTOU (#5986)
* fix(skill-importer): validate URL scheme and improve skills.sh handling * fix(skill-importer): enhance DNS resolution and SSRF protection in fetch URL handling * fix(url-safety): add allowed_dist parameter to check_outbound_url for flexible private blocking * test(skill-importer): add comprehensive tests for URL parsing and outbound checks * ensure newline at end of file in test_check_outbound_url_allows_public_ip * fix(skill-importer): improve TLS certificate handling in _get_checked function * fix(skill-importer): enhance _check_fetch_url to handle both hostnames and full URLs * fix(skill-importer): enhance parse_skill_source to support skills.sh URLs in path and netloc * fix(skill-importer): simplify skills.sh hostname check in parse_skill_source * fix(skill-importer): enhance parse_skill_source to identify skills.sh URLs in path and handle localhost/IP addresses * fix(skill-importer): enhance _resolve_and_check_url to validate all resolved IP addresses and prevent TOCTOU vulnerabilities * fix(skill-importer): enhance parse_skill_source to support schemeless GitHub and skills.sh URLs * fix(memory): resolve CodeQL URL sanitization warning and restore _check_fetch_url test alias * fix(memory): pin skill fetch sockets without rewriting URLs * fix(memory): reject unsupported skill wrapper hosts * refactor(url-safety): remove unused importer exception * test(memory): keep redirect regression hermetic * test(dns-rebinding): add test for _PinnedTransport to ensure connection to pinned IP * fix(skill-importer): enhance skills.sh support to extract GitHub links from page content * fix(skill-importer): improve URL scheme validation for GitHub and skills.sh links * fix(skills): reject unusable skill URLs instead of guessing Resolving a skills.sh link by scraping the first github.com URL out of the page body cannot work. Skill pages only ever link the repository root, never the skill's subdirectory, so every skill in a repo resolved to the same bundle: importing skills.sh/anthropics/skills/pdf walked the whole monorepo, saturated the 64-file cap, and installed algorithmic-art behind an ok:true response. Restore the redirect-target unwrap and fail with a message that says what to do instead. Also report the real reason a URL is rejected. The scheme check keyed off "://" appearing anywhere in the string, so a supplied-but-unusable URL came back as "URL is required", and a schemeless URL carrying "://" in its query was reported as an unsupported scheme. Key off the parsed scheme and let opaque schemes (mailto:, javascript:) and a schemeless host:port fall through to the host check. * test(skills): tighten the real-socket pinning regression The handler swallowed its own exceptions, so a failure inside it surfaced as a confusing assertion on the captured client address. Record the exception and assert on it, run the thread as a daemon, and close the listening socket from the test so a hang cannot outlive the run. Also drop the duplicate ipaddress import and the missing newline. * fix(skills): require exact GitHub skill URLs * test(skills): read complete pinned request headers --------- Co-authored-by: RaresKeY <158580472+RaresKeY@users.noreply.github.com> Co-authored-by: Léo <leograndcontact@gmail.com>
This commit is contained in:
@@ -1,16 +1,18 @@
|
||||
"""Import SKILL.md bundles from public GitHub (or skills.sh → GitHub) URLs."""
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
from typing import Dict, Iterable, List, Optional, Tuple, cast
|
||||
from urllib.parse import quote, urljoin, urlparse
|
||||
|
||||
import httpcore
|
||||
import httpx
|
||||
|
||||
from src.url_safety import check_outbound_url
|
||||
from src.url_safety import _default_resolver, check_outbound_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -25,6 +27,7 @@ TEXT_NAMES = {"skill.md", "license", "license.md", "readme.md"}
|
||||
_GITHUB_HOSTS = frozenset({
|
||||
"github.com", "www.github.com", "api.github.com", "raw.githubusercontent.com",
|
||||
})
|
||||
_SKILLS_SH_HOSTS = frozenset({"skills.sh", "www.skills.sh"})
|
||||
|
||||
|
||||
def _github_host(url: str) -> str:
|
||||
@@ -72,18 +75,158 @@ def _is_text_file(name: str) -> bool:
|
||||
_MAX_FETCH_REDIRECTS = 5
|
||||
|
||||
|
||||
def _check_fetch_url(url: str) -> None:
|
||||
"""SSRF guard for skill-import fetches (defense-in-depth).
|
||||
def _validated_ips(raw_ips: List[str]) -> List[ipaddress._BaseAddress]:
|
||||
"""Parse and de-duplicate one resolver snapshot in resolver order."""
|
||||
ips: List[ipaddress._BaseAddress] = []
|
||||
seen = set()
|
||||
for raw in raw_ips:
|
||||
if not isinstance(raw, str):
|
||||
continue
|
||||
try:
|
||||
ip = ipaddress.ip_address(raw.split("%", 1)[0])
|
||||
except ValueError:
|
||||
continue
|
||||
if ip in seen:
|
||||
continue
|
||||
seen.add(ip)
|
||||
ips.append(ip)
|
||||
return ips
|
||||
|
||||
Skill bundles only ever come from public GitHub, never an internal
|
||||
address, so block private/loopback/link-local targets on every hop —
|
||||
matching the hardened web-fetch path in
|
||||
``services/search/content.py:_get_public_url`` rather than the lenient
|
||||
default used for admin-configured model endpoints.
|
||||
"""
|
||||
ok, reason = check_outbound_url(url, block_private=True)
|
||||
|
||||
def _resolve_and_check_url(url: str) -> List[ipaddress._BaseAddress]:
|
||||
"""Return the exact address snapshot approved for one fetch hop."""
|
||||
resolved_ips: List[str] = []
|
||||
|
||||
def _recording_resolver(host: str) -> List[str]:
|
||||
answers = list(_default_resolver(host))
|
||||
resolved_ips[:] = answers
|
||||
return answers
|
||||
|
||||
ok, reason = check_outbound_url(
|
||||
url,
|
||||
block_private=True,
|
||||
resolver=_recording_resolver,
|
||||
)
|
||||
if not ok:
|
||||
raise SkillImportError(reason)
|
||||
raise SkillImportError(f"outbound URL blocked: {reason}")
|
||||
|
||||
pinned_ips = _validated_ips(resolved_ips)
|
||||
if not pinned_ips:
|
||||
raise SkillImportError("outbound URL blocked: host did not resolve to a usable address")
|
||||
return pinned_ips
|
||||
|
||||
|
||||
# Backward compatibility alias for tests importing _check_fetch_url directly
|
||||
_check_fetch_url = _resolve_and_check_url
|
||||
|
||||
|
||||
class _PinnedBackend(httpcore.NetworkBackend):
|
||||
"""Connect only to addresses from one validated DNS snapshot."""
|
||||
|
||||
def __init__(self, ips: List[ipaddress._BaseAddress]):
|
||||
self._ips = [str(ip) for ip in ips]
|
||||
self._real = httpcore.SyncBackend()
|
||||
|
||||
def connect_tcp(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
timeout: float | None = None,
|
||||
local_address: str | None = None,
|
||||
socket_options=None,
|
||||
):
|
||||
deadline = None if timeout is None else time.monotonic() + timeout
|
||||
last_exc: Optional[Exception] = None
|
||||
for ip in self._ips:
|
||||
remaining = None if deadline is None else max(0.0, deadline - time.monotonic())
|
||||
try:
|
||||
return self._real.connect_tcp(
|
||||
ip,
|
||||
port,
|
||||
remaining,
|
||||
local_address,
|
||||
socket_options,
|
||||
)
|
||||
except (httpcore.ConnectError, httpcore.ConnectTimeout) as exc:
|
||||
last_exc = exc
|
||||
if deadline is not None and time.monotonic() >= deadline:
|
||||
break
|
||||
if last_exc is not None:
|
||||
raise last_exc
|
||||
raise httpcore.ConnectError("no validated address available")
|
||||
|
||||
def connect_unix_socket(self, path, timeout=None, socket_options=None):
|
||||
return self._real.connect_unix_socket(path, timeout, socket_options)
|
||||
|
||||
def sleep(self, seconds: float) -> None:
|
||||
return self._real.sleep(seconds)
|
||||
|
||||
|
||||
_HTTPCORE_TO_HTTPX_EXC = {
|
||||
httpcore.ConnectError: httpx.ConnectError,
|
||||
httpcore.ConnectTimeout: httpx.ConnectTimeout,
|
||||
httpcore.LocalProtocolError: httpx.LocalProtocolError,
|
||||
httpcore.NetworkError: httpx.NetworkError,
|
||||
httpcore.PoolTimeout: httpx.PoolTimeout,
|
||||
httpcore.ProtocolError: httpx.ProtocolError,
|
||||
httpcore.ProxyError: httpx.ProxyError,
|
||||
httpcore.ReadError: httpx.ReadError,
|
||||
httpcore.ReadTimeout: httpx.ReadTimeout,
|
||||
httpcore.RemoteProtocolError: httpx.RemoteProtocolError,
|
||||
httpcore.TimeoutException: httpx.TimeoutException,
|
||||
httpcore.UnsupportedProtocol: httpx.UnsupportedProtocol,
|
||||
httpcore.WriteError: httpx.WriteError,
|
||||
httpcore.WriteTimeout: httpx.WriteTimeout,
|
||||
}
|
||||
|
||||
|
||||
class _PinnedTransport(httpx.BaseTransport):
|
||||
"""Pin socket connects while preserving URL authority, Host, and TLS SNI."""
|
||||
|
||||
def __init__(self, ips: List[ipaddress._BaseAddress]):
|
||||
self._pinned_ips = list(ips)
|
||||
self._pool = httpcore.ConnectionPool(
|
||||
ssl_context=httpx.create_ssl_context(),
|
||||
http1=True,
|
||||
http2=False,
|
||||
network_backend=_PinnedBackend(ips),
|
||||
)
|
||||
|
||||
def handle_request(self, request: httpx.Request) -> httpx.Response:
|
||||
core_request = httpcore.Request(
|
||||
method=request.method,
|
||||
url=httpcore.URL(
|
||||
scheme=request.url.raw_scheme,
|
||||
host=request.url.raw_host,
|
||||
port=request.url.port,
|
||||
target=request.url.raw_path,
|
||||
),
|
||||
headers=request.headers.raw,
|
||||
content=request.stream,
|
||||
extensions=request.extensions,
|
||||
)
|
||||
core_response = None
|
||||
try:
|
||||
core_response = self._pool.handle_request(core_request)
|
||||
content = b"".join(cast(Iterable[bytes], core_response.stream))
|
||||
except Exception as exc:
|
||||
mapped = _HTTPCORE_TO_HTTPX_EXC.get(type(exc))
|
||||
if mapped is not None:
|
||||
raise mapped(str(exc)) from exc
|
||||
raise
|
||||
finally:
|
||||
if core_response is not None:
|
||||
core_response.close()
|
||||
|
||||
return httpx.Response(
|
||||
status_code=core_response.status,
|
||||
headers=core_response.headers,
|
||||
content=content,
|
||||
extensions=core_response.extensions,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self._pool.close()
|
||||
|
||||
|
||||
def _get_checked(
|
||||
@@ -100,49 +243,76 @@ def _get_checked(
|
||||
hand lets us re-validate every hop, closing that blind-SSRF gap.
|
||||
"""
|
||||
current = url
|
||||
with httpx.Client(follow_redirects=False, timeout=timeout) as client:
|
||||
for _ in range(_MAX_FETCH_REDIRECTS + 1):
|
||||
_check_fetch_url(current)
|
||||
for _ in range(_MAX_FETCH_REDIRECTS + 1):
|
||||
pinned_ips = _resolve_and_check_url(current)
|
||||
with httpx.Client(
|
||||
transport=_PinnedTransport(pinned_ips),
|
||||
follow_redirects=False,
|
||||
timeout=timeout,
|
||||
) as client:
|
||||
r = client.get(current, headers=headers)
|
||||
if r.status_code in (301, 302, 303, 307, 308):
|
||||
location = r.headers.get("location")
|
||||
if not location:
|
||||
return r
|
||||
current = urljoin(str(r.url), location)
|
||||
continue
|
||||
return r
|
||||
|
||||
if r.status_code in (301, 302, 303, 307, 308):
|
||||
location = r.headers.get("location")
|
||||
if not location:
|
||||
return r
|
||||
current = urljoin(str(r.url), location)
|
||||
continue
|
||||
return r
|
||||
raise SkillImportError("too many redirects while fetching skill bundle")
|
||||
|
||||
|
||||
def parse_skill_source(url: str) -> ResolvedSource:
|
||||
"""Normalize skills.sh / GitHub web URLs into owner/repo/ref/path."""
|
||||
raw = (url or "").strip()
|
||||
if not raw:
|
||||
url = (url or "").strip()
|
||||
if not url:
|
||||
raise SkillImportError("URL is required")
|
||||
|
||||
# skills.sh often links to GitHub; try to unwrap ?url= or redirect target later.
|
||||
if "skills.sh" in raw and "github.com" not in raw:
|
||||
r = _get_checked(raw, timeout=20.0)
|
||||
# ``urlparse`` only reports an unambiguous scheme when the URL carries the
|
||||
# ``scheme://`` form. Opaque schemes (``mailto:``, ``javascript:``) and a
|
||||
# schemeless ``host:port`` both parse a "scheme" that is not one, so they
|
||||
# fall through to the host check below and are rejected on the host instead.
|
||||
scheme = urlparse(url).scheme.lower()
|
||||
if scheme not in ("http", "https"):
|
||||
if scheme and url.lower().startswith(f"{scheme}://"):
|
||||
raise SkillImportError(f"unsupported URL scheme: {scheme}")
|
||||
# Schemeless "github.com/owner/repo" — accept only a supported host.
|
||||
rough_host = (urlparse("//" + url).hostname or "").lower()
|
||||
if rough_host not in _GITHUB_HOSTS and rough_host not in _SKILLS_SH_HOSTS:
|
||||
raise SkillImportError("Only GitHub or skills.sh URLs are supported")
|
||||
url = "https://" + url
|
||||
|
||||
parsed = urlparse(url)
|
||||
hostname = (parsed.hostname or "").lower()
|
||||
if hostname not in _GITHUB_HOSTS and hostname not in _SKILLS_SH_HOSTS:
|
||||
raise SkillImportError("Only GitHub or skills.sh URLs are supported")
|
||||
|
||||
# A skills.sh link is only usable if it redirects to an exact supported
|
||||
# GitHub host. Scraping the page body for a github.com link cannot work:
|
||||
# skill pages only ever link the repository root, never the skill's
|
||||
# subdirectory, so the scrape resolves every skill in a repo to the same
|
||||
# (wrong) bundle. Fail with an actionable message instead.
|
||||
if hostname in _SKILLS_SH_HOSTS:
|
||||
r = _get_checked(url, timeout=20.0)
|
||||
if r.status_code >= 400:
|
||||
raise _github_response_error(r)
|
||||
final = str(r.url)
|
||||
_assert_github_url(final, context="redirect target")
|
||||
# Page may embed a github link; prefer final URL if redirected.
|
||||
if "github.com" in final:
|
||||
raw = final
|
||||
else:
|
||||
m = re.search(r"https?://github\.com/[^\s\"')]+", r.text or "")
|
||||
if m:
|
||||
raw = m.group(0).rstrip(".,)")
|
||||
if _github_host(final) not in _GITHUB_HOSTS:
|
||||
raise SkillImportError(
|
||||
"skills.sh did not redirect to GitHub — open the skill's "
|
||||
"repository on GitHub, navigate to the exact skill folder or "
|
||||
"SKILL.md file, and paste that URL; the repository-root link "
|
||||
"alone is not sufficient"
|
||||
)
|
||||
url = final
|
||||
|
||||
parsed = urlparse(raw)
|
||||
host = _github_host(raw)
|
||||
if host not in _GITHUB_HOSTS:
|
||||
raise SkillImportError(
|
||||
"Only GitHub URLs are supported (https://github.com/... or raw.githubusercontent.com/...)"
|
||||
)
|
||||
# Update parsed and hostname to reflect the new GitHub URL
|
||||
parsed = urlparse(url)
|
||||
hostname = (parsed.hostname or "").lower()
|
||||
|
||||
if host == "raw.githubusercontent.com":
|
||||
_assert_github_url(url)
|
||||
|
||||
if hostname == "raw.githubusercontent.com":
|
||||
# /owner/repo/ref/path/to/file
|
||||
bits = [p for p in parsed.path.split("/") if p]
|
||||
if len(bits) < 4:
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
"""Skill URL importer — GitHub path parsing."""
|
||||
import ipaddress
|
||||
|
||||
import pytest
|
||||
|
||||
from services.memory.skill_importer import (
|
||||
@@ -11,6 +13,13 @@ from services.memory.skill_importer import (
|
||||
)
|
||||
|
||||
|
||||
def _allow_fetch(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.memory.skill_importer._resolve_and_check_url",
|
||||
lambda url: [ipaddress.ip_address("93.184.216.34")],
|
||||
)
|
||||
|
||||
|
||||
def test_parse_github_blob_skill_md():
|
||||
src = parse_skill_source(
|
||||
"https://github.com/anthropics/skills/blob/main/skills/pdf/SKILL.md"
|
||||
@@ -69,10 +78,7 @@ def test_fetch_bytes_rejects_cross_host_redirect(monkeypatch):
|
||||
return _Resp()
|
||||
|
||||
monkeypatch.setattr("services.memory.skill_importer.httpx.Client", _Client)
|
||||
monkeypatch.setattr(
|
||||
"services.memory.skill_importer.check_outbound_url",
|
||||
lambda url, **kwargs: (True, ""),
|
||||
)
|
||||
_allow_fetch(monkeypatch)
|
||||
with pytest.raises(SkillImportError, match="redirect target"):
|
||||
_fetch_bytes("https://raw.githubusercontent.com/o/r/main/SKILL.md")
|
||||
|
||||
@@ -89,10 +95,7 @@ def test_list_github_dir_accepts_api_github_response(monkeypatch):
|
||||
"services.memory.skill_importer._fetch_text",
|
||||
lambda url: "# skill\n",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"services.memory.skill_importer.check_outbound_url",
|
||||
lambda url, **kwargs: (True, ""),
|
||||
)
|
||||
_allow_fetch(monkeypatch)
|
||||
|
||||
class _Resp:
|
||||
url = "https://api.github.com/repos/o/r/contents?ref=main"
|
||||
@@ -144,10 +147,7 @@ def _mock_httpx_client(monkeypatch, response):
|
||||
return response
|
||||
|
||||
monkeypatch.setattr("services.memory.skill_importer.httpx.Client", _Client)
|
||||
monkeypatch.setattr(
|
||||
"services.memory.skill_importer.check_outbound_url",
|
||||
lambda url, **kwargs: (True, ""),
|
||||
)
|
||||
_allow_fetch(monkeypatch)
|
||||
|
||||
|
||||
def test_list_github_dir_surfaces_rate_limit(monkeypatch):
|
||||
|
||||
@@ -0,0 +1,217 @@
|
||||
"""Deterministic regressions for skill-import DNS validation and pinning."""
|
||||
|
||||
import gzip
|
||||
import ipaddress
|
||||
import socket
|
||||
import threading
|
||||
|
||||
import httpcore
|
||||
import httpx
|
||||
|
||||
from services.memory import skill_importer
|
||||
|
||||
|
||||
PUBLIC_A = ipaddress.ip_address("93.184.216.34")
|
||||
PUBLIC_B = ipaddress.ip_address("1.1.1.1")
|
||||
|
||||
|
||||
def test_validation_snapshot_is_the_only_connect_destination(monkeypatch):
|
||||
answers = iter([
|
||||
[str(PUBLIC_A)],
|
||||
["127.0.0.1"],
|
||||
])
|
||||
resolver_calls = []
|
||||
|
||||
def _flipping_resolver(host):
|
||||
resolver_calls.append(host)
|
||||
return next(answers)
|
||||
|
||||
monkeypatch.setattr(skill_importer, "_default_resolver", _flipping_resolver)
|
||||
pinned_ips = skill_importer._check_fetch_url("https://rebind.example/skill")
|
||||
|
||||
connected = []
|
||||
|
||||
class _RecordingBackend:
|
||||
def connect_tcp(self, host, port, timeout, local_address, socket_options):
|
||||
connected.append((host, port))
|
||||
return object()
|
||||
|
||||
backend = skill_importer._PinnedBackend(pinned_ips)
|
||||
backend._real = _RecordingBackend()
|
||||
backend.connect_tcp("rebind.example", 443, timeout=1.0)
|
||||
|
||||
assert resolver_calls == ["rebind.example"]
|
||||
assert pinned_ips == [PUBLIC_A]
|
||||
assert connected == [(str(PUBLIC_A), 443)]
|
||||
|
||||
|
||||
def test_pinned_backend_falls_back_only_within_validated_snapshot():
|
||||
attempts = []
|
||||
|
||||
class _FallbackBackend:
|
||||
def connect_tcp(self, host, port, timeout, local_address, socket_options):
|
||||
attempts.append((host, timeout))
|
||||
if host == str(PUBLIC_A):
|
||||
raise httpcore.ConnectError("first address unavailable")
|
||||
return "connected"
|
||||
|
||||
backend = skill_importer._PinnedBackend([PUBLIC_A, PUBLIC_B])
|
||||
backend._real = _FallbackBackend()
|
||||
|
||||
assert backend.connect_tcp("rebind.example", 443, timeout=1.0) == "connected"
|
||||
assert [host for host, _ in attempts] == [str(PUBLIC_A), str(PUBLIC_B)]
|
||||
assert all(timeout is not None and 0 <= timeout <= 1.0 for _, timeout in attempts)
|
||||
|
||||
|
||||
def test_transport_preserves_request_authority_and_response_url():
|
||||
recorded = []
|
||||
|
||||
class _CoreResponse:
|
||||
status = 200
|
||||
headers = [(b"content-type", b"text/plain")]
|
||||
stream = [b"ok"]
|
||||
extensions = {}
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
class _RecordingPool:
|
||||
def handle_request(self, request):
|
||||
recorded.append(request)
|
||||
return _CoreResponse()
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
transport = skill_importer._PinnedTransport([PUBLIC_A])
|
||||
transport._pool.close()
|
||||
transport._pool = _RecordingPool()
|
||||
|
||||
url = "https://github.com:444/octocat/repo?q=1"
|
||||
with httpx.Client(transport=transport) as client:
|
||||
response = client.get(url)
|
||||
|
||||
core_request = recorded[0]
|
||||
assert core_request.url.host == b"github.com"
|
||||
assert core_request.url.port == 444
|
||||
assert core_request.url.target == b"/octocat/repo?q=1"
|
||||
assert (b"host", b"github.com:444") in [
|
||||
(name.lower(), value) for name, value in core_request.headers
|
||||
]
|
||||
assert str(response.url) == url
|
||||
|
||||
|
||||
def test_get_checked_uses_fresh_transport_per_redirect_hop(monkeypatch):
|
||||
first = "https://github.com/owner/repo"
|
||||
second = "https://raw.githubusercontent.com/owner/repo/main/SKILL.md"
|
||||
snapshots = {
|
||||
first: [PUBLIC_A],
|
||||
second: [PUBLIC_B],
|
||||
}
|
||||
clients = []
|
||||
requested = []
|
||||
|
||||
monkeypatch.setattr(
|
||||
skill_importer,
|
||||
"_resolve_and_check_url",
|
||||
lambda url: snapshots[url],
|
||||
)
|
||||
|
||||
class _Client:
|
||||
def __init__(self, *, transport, follow_redirects, timeout):
|
||||
assert follow_redirects is False
|
||||
clients.append((transport._pinned_ips, timeout))
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *args):
|
||||
return False
|
||||
|
||||
def get(self, url, headers=None):
|
||||
requested.append((url, headers))
|
||||
request = httpx.Request("GET", url)
|
||||
if url == first:
|
||||
return httpx.Response(
|
||||
302,
|
||||
headers={"location": second},
|
||||
request=request,
|
||||
)
|
||||
return httpx.Response(200, content=b"ok", request=request)
|
||||
|
||||
monkeypatch.setattr(skill_importer.httpx, "Client", _Client)
|
||||
|
||||
response = skill_importer._get_checked(first, headers={"Accept": "text/plain"})
|
||||
|
||||
assert [ips for ips, _ in clients] == [[PUBLIC_A], [PUBLIC_B]]
|
||||
assert requested == [
|
||||
(first, {"Accept": "text/plain"}),
|
||||
(second, {"Accept": "text/plain"}),
|
||||
]
|
||||
assert str(response.url) == second
|
||||
|
||||
|
||||
def test_dns_rebinding_pinned_transport_dials_pinned_ip():
|
||||
"""The real pool must dial the pinned IP and keep the logical request intact.
|
||||
|
||||
Everything above this test replaces the pool or the backend, so this is the
|
||||
only case that exercises ``httpcore.ConnectionPool`` end to end: the socket
|
||||
goes to the pinned address while URL, ``Host``, and the decoded body stay on
|
||||
the original hostname.
|
||||
"""
|
||||
server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
server_socket.bind(("127.0.0.1", 0))
|
||||
server_socket.listen(1)
|
||||
port = server_socket.getsockname()[1]
|
||||
|
||||
captured_request = b""
|
||||
client_address = None
|
||||
server_error = None
|
||||
|
||||
def handle_client():
|
||||
nonlocal captured_request, client_address, server_error
|
||||
try:
|
||||
conn, client_address = server_socket.accept()
|
||||
with conn:
|
||||
conn.settimeout(2.0)
|
||||
while b"\r\n\r\n" not in captured_request:
|
||||
if len(captured_request) >= 16_384:
|
||||
raise AssertionError("request headers exceeded 16 KiB")
|
||||
chunk = conn.recv(min(4096, 16_384 - len(captured_request)))
|
||||
if not chunk:
|
||||
break
|
||||
captured_request += chunk
|
||||
if b"\r\n\r\n" not in captured_request:
|
||||
raise AssertionError(
|
||||
"connection closed before request headers completed"
|
||||
)
|
||||
body = gzip.compress(b"successfully decoded gzip body")
|
||||
conn.sendall(
|
||||
b"HTTP/1.1 200 OK\r\n"
|
||||
b"Content-Encoding: gzip\r\n"
|
||||
b"Content-Length: " + str(len(body)).encode() + b"\r\n"
|
||||
b"\r\n" + body
|
||||
)
|
||||
except Exception as exc: # surfaced by the assertion below
|
||||
server_error = exc
|
||||
|
||||
server_thread = threading.Thread(target=handle_client, daemon=True)
|
||||
server_thread.start()
|
||||
|
||||
# Any hostname works: only the pinned snapshot decides where the socket goes.
|
||||
url = f"http://rebind.example:{port}/secret-metadata"
|
||||
try:
|
||||
with httpx.Client(
|
||||
transport=skill_importer._PinnedTransport([ipaddress.ip_address("127.0.0.1")])
|
||||
) as client:
|
||||
response = client.get(url)
|
||||
server_thread.join(timeout=5.0)
|
||||
finally:
|
||||
server_socket.close()
|
||||
|
||||
assert server_error is None, server_error
|
||||
assert not server_thread.is_alive()
|
||||
assert client_address is not None and client_address[0] == "127.0.0.1"
|
||||
assert f"host: rebind.example:{port}".encode() in captured_request.lower()
|
||||
assert str(response.url) == url
|
||||
assert response.text == "successfully decoded gzip body"
|
||||
@@ -0,0 +1,155 @@
|
||||
import re
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from services.memory.skill_importer import (
|
||||
ResolvedSource,
|
||||
SkillImportError,
|
||||
check_outbound_url,
|
||||
parse_skill_source,
|
||||
)
|
||||
|
||||
## 1. Tests for Hostname Dispatch & Substring Spoofing
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url",
|
||||
[
|
||||
"https://skills.sh.attacker.com/owner/repo",
|
||||
"https://evilskills.sh/owner/repo",
|
||||
"https://notskills.sh/owner/repo",
|
||||
"https://api.skills.sh/owner/repo",
|
||||
"https://1.1.1.1/skills.sh/owner/repo",
|
||||
"http://localhost/skills.sh/owner/repo",
|
||||
],
|
||||
)
|
||||
def test_parse_skill_source_rejects_unsupported_host_before_fetch(url):
|
||||
"""Unsupported authorities must never reach the network unwrap path."""
|
||||
with patch("services.memory.skill_importer._get_checked") as mock_get:
|
||||
with pytest.raises(SkillImportError):
|
||||
parse_skill_source(url)
|
||||
mock_get.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ["https://skills.sh/my-skill", "https://www.skills.sh/my-skill"])
|
||||
def test_parse_skill_source_unwraps_skills_sh_redirect_to_github(entry):
|
||||
"""Both skills.sh spellings unwrap when the fetch lands on a GitHub host."""
|
||||
with patch("services.memory.skill_importer._get_checked") as mock_get:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.url = "https://github.com/test-owner/test-repo"
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
source = parse_skill_source(entry)
|
||||
assert source.owner == "test-owner"
|
||||
assert source.repo == "test-repo"
|
||||
|
||||
|
||||
def test_parse_skill_source_rejects_skills_sh_page_that_never_reaches_github():
|
||||
"""A skills.sh page that does not redirect to GitHub must fail loudly.
|
||||
|
||||
The live site serves the skill page from ``www.skills.sh`` and only ever
|
||||
links the repository root, never the skill's subdirectory. Scraping a
|
||||
``github.com`` link out of the body therefore resolves every skill in a
|
||||
repo to the same bundle, so the importer must refuse rather than guess.
|
||||
"""
|
||||
body = '<html><body><a href="https://github.com/test-owner/test-repo">Repository</a></body></html>'
|
||||
with patch("services.memory.skill_importer._get_checked") as mock_get:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.url = "https://www.skills.sh/anthropics/skills/pdf"
|
||||
mock_response.text = body
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
with pytest.raises(
|
||||
SkillImportError, match="did not redirect to GitHub"
|
||||
) as exc_info:
|
||||
parse_skill_source("https://skills.sh/anthropics/skills/pdf")
|
||||
|
||||
message = str(exc_info.value)
|
||||
assert "exact skill folder or SKILL.md file" in message
|
||||
assert "repository-root link alone is not sufficient" in message
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("url", "expected"),
|
||||
[
|
||||
("ftp://github.com/o/r", "unsupported URL scheme: ftp"),
|
||||
("file:///etc/passwd", "unsupported URL scheme: file"),
|
||||
("gopher://github.com/o/r", "unsupported URL scheme: gopher"),
|
||||
("javascript:alert(1)", "Only GitHub or skills.sh URLs are supported"),
|
||||
("mailto:x@y.z", "Only GitHub or skills.sh URLs are supported"),
|
||||
("data:text/html,<b>x</b>", "Only GitHub or skills.sh URLs are supported"),
|
||||
("https://evil.example/o/r", "Only GitHub or skills.sh URLs are supported"),
|
||||
],
|
||||
)
|
||||
def test_parse_skill_source_reports_the_real_reason_for_rejection(url, expected):
|
||||
"""A supplied-but-unusable URL must not be reported as a missing URL."""
|
||||
with patch("services.memory.skill_importer._get_checked") as mock_get:
|
||||
with pytest.raises(SkillImportError, match=re.escape(expected)):
|
||||
parse_skill_source(url)
|
||||
mock_get.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("url", "owner", "repo"),
|
||||
[
|
||||
("github.com/octocat/Hello-World", "octocat", "Hello-World"),
|
||||
("HTTPS://github.com/octocat/Hello-World", "octocat", "Hello-World"),
|
||||
("github.com:443/octocat/Hello-World", "octocat", "Hello-World"),
|
||||
("github.com/octocat/Hello-World?q=a://b", "octocat", "Hello-World"),
|
||||
],
|
||||
)
|
||||
def test_parse_skill_source_accepts_schemeless_and_uppercase_github(url, owner, repo):
|
||||
"""A schemeless host, an uppercase scheme, and a ``://`` in the query all parse."""
|
||||
source = parse_skill_source(url)
|
||||
assert (source.owner, source.repo) == (owner, repo)
|
||||
|
||||
|
||||
def test_parse_skill_source_valid_github():
|
||||
"""Ensure standard GitHub URLs parse into the correct ResolvedSource fields."""
|
||||
source = parse_skill_source("https://github.com/octocat/Hello-World/tree/main/docs")
|
||||
assert isinstance(source, ResolvedSource)
|
||||
assert source.owner == "octocat"
|
||||
assert source.repo == "Hello-World"
|
||||
assert source.ref == "main"
|
||||
assert source.path == "docs"
|
||||
|
||||
|
||||
## 2. Tests for SSRF Guard & CGNAT
|
||||
|
||||
def test_check_outbound_url_blocks_cgnat():
|
||||
"""Ensure Carrier-Grade NAT (RFC 6598) block 100.64.0.0/10 is blocked."""
|
||||
def mock_resolver(host):
|
||||
return ["100.64.5.10"]
|
||||
|
||||
ok, reason = check_outbound_url("http://example.com", block_private=True, resolver=mock_resolver)
|
||||
assert not ok
|
||||
assert "private/shared/loopback" in reason # Updated to match your codebase's error string
|
||||
|
||||
def test_check_outbound_url_blocks_loopback():
|
||||
"""Ensure loopback IPs (127.0.0.1) are blocked by default."""
|
||||
def mock_resolver(host):
|
||||
return ["127.0.0.1"]
|
||||
|
||||
ok, reason = check_outbound_url("http://localhost", block_private=True, resolver=mock_resolver)
|
||||
assert not ok
|
||||
|
||||
|
||||
def test_check_outbound_url_blocks_metadata():
|
||||
"""Ensure cloud metadata endpoints (169.254.169.254) are blocked."""
|
||||
def mock_resolver(host):
|
||||
return ["169.254.169.254"]
|
||||
|
||||
ok, reason = check_outbound_url("http://metadata.google.internal", block_private=True, resolver=mock_resolver)
|
||||
assert not ok
|
||||
|
||||
|
||||
def test_check_outbound_url_allows_public_ip():
|
||||
"""Ensure public routable IPs pass successfully."""
|
||||
def mock_resolver(host):
|
||||
return ["93.184.216.34"]
|
||||
|
||||
ok, reason = check_outbound_url("http://example.com", block_private=True, resolver=mock_resolver)
|
||||
assert ok
|
||||
assert reason == "ok"
|
||||
@@ -6,10 +6,12 @@ path in ``services/search/content.py:_get_public_url``. Previously it used
|
||||
``httpx``'s ``follow_redirects=True`` with the lenient guard on the *initial*
|
||||
URL only, so a ``3xx`` to an internal/metadata address was still connected to.
|
||||
|
||||
These tests are hermetic: every host is an IP literal, so ``check_outbound_url``
|
||||
resolves them locally (``getaddrinfo`` on a numeric address does no DNS) and no
|
||||
network access is required. The HTTP layer is faked so no real request is made.
|
||||
These tests are hermetic: public and internal guard cases use IP literals, while
|
||||
the exact ``skills.sh`` case injects its validated address snapshot. The HTTP
|
||||
layer is faked, so no real DNS lookup or request is made.
|
||||
"""
|
||||
import ipaddress
|
||||
|
||||
import pytest
|
||||
|
||||
from services.memory import skill_importer
|
||||
@@ -108,10 +110,20 @@ def test_fetch_bytes_blocks_redirect_to_internal(monkeypatch, internal):
|
||||
|
||||
def test_skills_sh_entry_blocks_redirect_to_metadata(monkeypatch):
|
||||
# The skills.sh unwrap path (user-supplied host) must also revalidate hops.
|
||||
raw = "http://1.1.1.1/skills.sh" # contains "skills.sh", not "github.com"
|
||||
raw = "https://skills.sh/example/skill"
|
||||
checked = []
|
||||
|
||||
def _check_hop(url):
|
||||
checked.append(url)
|
||||
if url == raw:
|
||||
return [ipaddress.ip_address("1.1.1.1")]
|
||||
raise SkillImportError("outbound URL blocked: private target")
|
||||
|
||||
monkeypatch.setattr(skill_importer, "_resolve_and_check_url", _check_hop)
|
||||
_install_fake_client(monkeypatch, redirect_from=raw, redirect_to=METADATA)
|
||||
with pytest.raises(SkillImportError, match="blocked"):
|
||||
parse_skill_source(raw)
|
||||
assert checked == [raw, METADATA]
|
||||
|
||||
|
||||
# --- Positive: a legitimate public->public redirect is still followed --------
|
||||
|
||||
Reference in New Issue
Block a user