mirror of
https://github.com/inventree/InvenTree.git
synced 2026-08-13 00:46:15 +00:00
[security] Additional SSRF protections (#12595)
Tighten SSRF protections when fetching from external URLs
This commit is contained in:
@@ -1,5 +1,7 @@
|
|||||||
"""Provides helper functions used throughout the InvenTree project that access the database."""
|
"""Provides helper functions used throughout the InvenTree project that access the database."""
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import contextvars
|
||||||
import io
|
import io
|
||||||
import ipaddress
|
import ipaddress
|
||||||
import socket
|
import socket
|
||||||
@@ -85,6 +87,66 @@ def construct_absolute_url(*arg, base_url=None, request=None):
|
|||||||
return urljoin(base_url, relative_url)
|
return urljoin(base_url, relative_url)
|
||||||
|
|
||||||
|
|
||||||
|
_ssrf_guard_active = contextvars.ContextVar('ssrf_guard_active', default=False)
|
||||||
|
_real_getaddrinfo = socket.getaddrinfo
|
||||||
|
|
||||||
|
|
||||||
|
def _is_unsafe_ip(host: str) -> bool:
|
||||||
|
"""Return True if the given (already-resolved) address string is private/reserved."""
|
||||||
|
ip = ipaddress.ip_address(host)
|
||||||
|
return ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved
|
||||||
|
|
||||||
|
|
||||||
|
def _guarded_getaddrinfo(host, *args, **kwargs):
|
||||||
|
"""Drop-in replacement for socket.getaddrinfo which enforces SSRF checks while active.
|
||||||
|
|
||||||
|
DNS resolution is otherwise delegated to the OS resolver on every connection
|
||||||
|
attempt an HTTP client makes - including ones made well after an earlier,
|
||||||
|
now-discarded resolution was validated by `validate_url_no_ssrf()`. An attacker
|
||||||
|
controlling authoritative DNS for their own domain can answer a first ("check")
|
||||||
|
lookup with a public IP and a later ("use") lookup with a private/internal one
|
||||||
|
(DNS rebinding), bypassing a validation step that only inspects the first answer.
|
||||||
|
|
||||||
|
Patching this at the socket layer means the *actual* resolution used to open the
|
||||||
|
real connection - whichever call that turns out to be - is the one that gets
|
||||||
|
checked, so there is no window between "validated" and "connected" for the
|
||||||
|
answer to change. The guard is only enforced while `ssrf_safe_context()` is
|
||||||
|
active in the current context, so unrelated resolutions (e.g. other threads
|
||||||
|
talking to the database/cache) are never affected.
|
||||||
|
"""
|
||||||
|
result = _real_getaddrinfo(host, *args, **kwargs)
|
||||||
|
|
||||||
|
if _ssrf_guard_active.get():
|
||||||
|
for _family, _type, _proto, _canonname, sockaddr in result:
|
||||||
|
if _is_unsafe_ip(sockaddr[0]):
|
||||||
|
raise socket.gaierror(
|
||||||
|
f'Host {host!r} resolved to a private or reserved address'
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
socket.getaddrinfo = _guarded_getaddrinfo
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def ssrf_safe_context():
|
||||||
|
"""Enforce SSRF IP validation on every DNS resolution performed within this block.
|
||||||
|
|
||||||
|
Wrap any code which resolves a user/attacker-influenced hostname and then
|
||||||
|
connects to it (e.g. `requests.get()`, or a library like WeasyPrint doing its
|
||||||
|
own resolution internally) in this context manager so that the resolution
|
||||||
|
backing the *actual* connection is validated, not just an earlier, separate
|
||||||
|
lookup whose result is discarded. See `_guarded_getaddrinfo` for why this
|
||||||
|
closes the DNS-rebinding TOCTOU gap that a standalone pre-check cannot.
|
||||||
|
"""
|
||||||
|
token = _ssrf_guard_active.set(True)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
_ssrf_guard_active.reset(token)
|
||||||
|
|
||||||
|
|
||||||
def validate_url_no_ssrf(url):
|
def validate_url_no_ssrf(url):
|
||||||
"""Validate that a URL does not point to a private/internal network address.
|
"""Validate that a URL does not point to a private/internal network address.
|
||||||
|
|
||||||
@@ -109,9 +171,7 @@ def validate_url_no_ssrf(url):
|
|||||||
raise ValueError(_('Invalid URL: hostname could not be resolved'))
|
raise ValueError(_('Invalid URL: hostname could not be resolved'))
|
||||||
|
|
||||||
for _family, _type, _proto, _canonname, sockaddr in addrinfo:
|
for _family, _type, _proto, _canonname, sockaddr in addrinfo:
|
||||||
ip = ipaddress.ip_address(sockaddr[0])
|
if _is_unsafe_ip(sockaddr[0]):
|
||||||
|
|
||||||
if ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved:
|
|
||||||
raise ValueError(_('URL points to a private or reserved IP address'))
|
raise ValueError(_('URL points to a private or reserved IP address'))
|
||||||
|
|
||||||
|
|
||||||
@@ -158,36 +218,40 @@ def download_image_from_url(
|
|||||||
headers = {'User-Agent': user_agent} if user_agent else None
|
headers = {'User-Agent': user_agent} if user_agent else None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response = requests.get(
|
# SSRF protection: guard every DNS resolution made while actually connecting
|
||||||
remote_url,
|
# (not just the pre-check above) so a hostname cannot pass validation with
|
||||||
timeout=timeout,
|
# one IP and then be connected to via a different, private one (DNS rebinding).
|
||||||
allow_redirects=False,
|
with ssrf_safe_context():
|
||||||
stream=True,
|
|
||||||
headers=headers,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Handle redirects manually to validate each destination
|
|
||||||
max_redirects = 5
|
|
||||||
redirect_count = 0
|
|
||||||
|
|
||||||
while response.is_redirect and redirect_count < max_redirects:
|
|
||||||
redirect_url = response.headers.get('Location')
|
|
||||||
if not redirect_url:
|
|
||||||
break
|
|
||||||
|
|
||||||
# Validate the redirect destination against SSRF
|
|
||||||
validator(redirect_url)
|
|
||||||
validate_url_no_ssrf(redirect_url)
|
|
||||||
|
|
||||||
redirect_count += 1
|
|
||||||
response = requests.get(
|
response = requests.get(
|
||||||
redirect_url,
|
remote_url,
|
||||||
timeout=timeout,
|
timeout=timeout,
|
||||||
allow_redirects=False,
|
allow_redirects=False,
|
||||||
stream=True,
|
stream=True,
|
||||||
headers=headers,
|
headers=headers,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Handle redirects manually to validate each destination
|
||||||
|
max_redirects = 5
|
||||||
|
redirect_count = 0
|
||||||
|
|
||||||
|
while response.is_redirect and redirect_count < max_redirects:
|
||||||
|
redirect_url = response.headers.get('Location')
|
||||||
|
if not redirect_url:
|
||||||
|
break
|
||||||
|
|
||||||
|
# Validate the redirect destination against SSRF
|
||||||
|
validator(redirect_url)
|
||||||
|
validate_url_no_ssrf(redirect_url)
|
||||||
|
|
||||||
|
redirect_count += 1
|
||||||
|
response = requests.get(
|
||||||
|
redirect_url,
|
||||||
|
timeout=timeout,
|
||||||
|
allow_redirects=False,
|
||||||
|
stream=True,
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
|
|
||||||
if redirect_count >= max_redirects:
|
if redirect_count >= max_redirects:
|
||||||
raise ValueError(_('Too many redirects'))
|
raise ValueError(_('Too many redirects'))
|
||||||
|
|
||||||
|
|||||||
@@ -776,6 +776,41 @@ class TestHelpers(TestCase):
|
|||||||
large_img, timeout=10, max_size=10 * 1024 * 1024
|
large_img, timeout=10, max_size=10 * 1024 * 1024
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_download_image_dns_rebind_blocked(self):
|
||||||
|
"""A hostname that resolves safely for validation but privately for the real fetch must be blocked.
|
||||||
|
|
||||||
|
Regression test: `validate_url_no_ssrf()` used to resolve the hostname once,
|
||||||
|
validate that result, and then discard it - the actual `requests.get()` call
|
||||||
|
resolved the same hostname again independently. An attacker controlling DNS
|
||||||
|
for their own domain could answer the first ("check") lookup with a public IP
|
||||||
|
and every subsequent ("use") lookup with a private/internal one (DNS
|
||||||
|
rebinding), so the real connection went somewhere the validator never saw.
|
||||||
|
"""
|
||||||
|
public_addrinfo = [(2, 1, 6, '', ('93.184.216.34', 0))]
|
||||||
|
private_addrinfo = [(2, 1, 6, '', ('127.0.0.1', 0))]
|
||||||
|
|
||||||
|
calls = {'n': 0}
|
||||||
|
|
||||||
|
def rebinding_getaddrinfo(host, *args, **kwargs):
|
||||||
|
calls['n'] += 1
|
||||||
|
return public_addrinfo if calls['n'] == 1 else private_addrinfo
|
||||||
|
|
||||||
|
with mock.patch.object(
|
||||||
|
InvenTree.helpers_model,
|
||||||
|
'_real_getaddrinfo',
|
||||||
|
side_effect=rebinding_getaddrinfo,
|
||||||
|
):
|
||||||
|
# The pre-check sees the safe public IP and passes; the *actual*
|
||||||
|
# connection attempt made by requests.get() must be independently
|
||||||
|
# guarded and must fail once it resolves to the private IP.
|
||||||
|
with self.assertRaisesRegex(Exception, 'Connection error'):
|
||||||
|
InvenTree.helpers_model.download_image_from_url(
|
||||||
|
'http://rebind.example.com/image.png'
|
||||||
|
)
|
||||||
|
|
||||||
|
# Both the validation-time and connection-time lookups must have happened.
|
||||||
|
self.assertEqual(calls['n'], 2)
|
||||||
|
|
||||||
def test_model_mixin(self):
|
def test_model_mixin(self):
|
||||||
"""Test the getModelsWithMixin function."""
|
"""Test the getModelsWithMixin function."""
|
||||||
from InvenTree.models import InvenTreeBarcodeMixin
|
from InvenTree.models import InvenTreeBarcodeMixin
|
||||||
|
|||||||
@@ -22,8 +22,16 @@ class InvenTreeURLFetcher(URLFetcher):
|
|||||||
scheme = parsed.scheme.lower()
|
scheme = parsed.scheme.lower()
|
||||||
|
|
||||||
if scheme in ('data', 'http', 'https'):
|
if scheme in ('data', 'http', 'https'):
|
||||||
|
from InvenTree.helpers_model import ssrf_safe_context
|
||||||
|
|
||||||
self._validate_http_url(url, parsed)
|
self._validate_http_url(url, parsed)
|
||||||
return super().fetch(url, headers)
|
|
||||||
|
# SSRF protection: guard the DNS resolution WeasyPrint performs internally
|
||||||
|
# when it actually opens the connection, not just the pre-check above - a
|
||||||
|
# hostname could otherwise pass validation with one IP and then be
|
||||||
|
# connected to via a different, private one moments later (DNS rebinding).
|
||||||
|
with ssrf_safe_context():
|
||||||
|
return super().fetch(url, headers)
|
||||||
|
|
||||||
if scheme == 'file':
|
if scheme == 'file':
|
||||||
logger.warning("InvenTreeURLFetcher: blocked file:// URL: '%s'", url)
|
logger.warning("InvenTreeURLFetcher: blocked file:// URL: '%s'", url)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Unit testing for the various report models."""
|
"""Unit testing for the various report models."""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import socket
|
||||||
import tempfile
|
import tempfile
|
||||||
from io import StringIO
|
from io import StringIO
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -1077,6 +1078,49 @@ class URLFetcherTest(TestCase):
|
|||||||
self.fetcher.fetch('data:image/png;base64,abc123')
|
self.fetcher.fetch('data:image/png;base64,abc123')
|
||||||
self.fetcher.fetch('data:text/css;base64,abc123')
|
self.fetcher.fetch('data:text/css;base64,abc123')
|
||||||
|
|
||||||
|
def test_dns_rebind_is_blocked_at_fetch_time(self):
|
||||||
|
"""A hostname that resolves safely for validation but privately for the real fetch must be blocked.
|
||||||
|
|
||||||
|
Regression test: `validate_url_no_ssrf()` used to resolve the hostname once,
|
||||||
|
validate that result, and then discard it, so `InvenTreeURLFetcher.fetch()`
|
||||||
|
would delegate straight to WeasyPrint's own fetcher, which resolves the
|
||||||
|
hostname *again* independently. An attacker controlling DNS for their own
|
||||||
|
domain could answer the first ("check") lookup with a public IP and every
|
||||||
|
subsequent ("use") lookup with a private/internal one (DNS rebinding), so the
|
||||||
|
real request WeasyPrint made went somewhere the validator never saw.
|
||||||
|
"""
|
||||||
|
set_global_setting('REPORT_FETCH_URLS', True, change_user=None)
|
||||||
|
|
||||||
|
public_addrinfo = [(2, 1, 6, '', ('93.184.216.34', 0))]
|
||||||
|
private_addrinfo = [(2, 1, 6, '', ('127.0.0.1', 0))]
|
||||||
|
|
||||||
|
calls = {'n': 0}
|
||||||
|
|
||||||
|
def rebinding_getaddrinfo(host, *args, **kwargs):
|
||||||
|
calls['n'] += 1
|
||||||
|
return public_addrinfo if calls['n'] == 1 else private_addrinfo
|
||||||
|
|
||||||
|
def fake_weasyprint_fetch(_self, url, headers=None):
|
||||||
|
# Simulate WeasyPrint's own fetch-time DNS resolution, performed
|
||||||
|
# independently of the validation InvenTreeURLFetcher already did.
|
||||||
|
socket.getaddrinfo('rebind.example.com', None)
|
||||||
|
return {'string': b'should never be reached'}
|
||||||
|
|
||||||
|
import InvenTree.helpers_model as helpers_model
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
helpers_model, '_real_getaddrinfo', side_effect=rebinding_getaddrinfo
|
||||||
|
):
|
||||||
|
with patch(
|
||||||
|
'weasyprint.urls.URLFetcher.fetch', side_effect=fake_weasyprint_fetch
|
||||||
|
):
|
||||||
|
with self.assertRaises(socket.gaierror):
|
||||||
|
self.fetcher.fetch('http://rebind.example.com/image.png')
|
||||||
|
|
||||||
|
# The validation-time lookup (safe) and the fetch-time lookup (private)
|
||||||
|
# must both have happened for this to be a meaningful regression test.
|
||||||
|
self.assertEqual(calls['n'], 2)
|
||||||
|
|
||||||
|
|
||||||
class DefaultTemplateFileTest(TestCase):
|
class DefaultTemplateFileTest(TestCase):
|
||||||
"""Unit tests for building the default report and label template files."""
|
"""Unit tests for building the default report and label template files."""
|
||||||
|
|||||||
Reference in New Issue
Block a user