Fix: prevent DNS-rebinding AI endpoint

This commit is contained in:
shamoon
2026-05-26 09:33:24 -07:00
parent 3149522426
commit 0665ae214e
7 changed files with 243 additions and 69 deletions
+3 -66
View File
@@ -4,70 +4,12 @@ import httpx
from celery import shared_task
from django.conf import settings
from paperless.network import format_host_for_url
from paperless.network import is_public_ip
from paperless.network import resolve_hostname_ips
from paperless.network import PinnedHostHTTPTransport
from paperless.network import validate_outbound_http_url
logger = logging.getLogger("paperless.workflows.webhooks")
class WebhookTransport(httpx.HTTPTransport):
"""
Transport that resolves/validates hostnames and rewrites to a vetted IP
while keeping Host/SNI as the original hostname.
"""
def __init__(
self,
hostname: str,
*args,
allow_internal: bool = False,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
self.hostname = hostname
self.allow_internal = allow_internal
def handle_request(self, request: httpx.Request) -> httpx.Response:
hostname = request.url.host
if not hostname:
raise httpx.ConnectError("No hostname in request URL")
try:
ips = resolve_hostname_ips(hostname)
except ValueError as e:
raise httpx.ConnectError(str(e)) from e
if not self.allow_internal:
for ip_str in ips:
if not is_public_ip(ip_str):
raise httpx.ConnectError(
f"Connection blocked: {hostname} resolves to a non-public address",
)
ip_str = ips[0]
formatted_ip = format_host_for_url(ip_str)
new_headers = httpx.Headers(request.headers)
if "host" in new_headers:
del new_headers["host"]
new_headers["Host"] = hostname
new_url = request.url.copy_with(host=formatted_ip)
request = httpx.Request(
method=request.method,
url=new_url,
headers=new_headers,
content=request.stream,
extensions=request.extensions,
)
request.extensions["sni_hostname"] = hostname
return super().handle_request(request)
@shared_task(
retry_backoff=True,
autoretry_for=(httpx.HTTPStatusError,),
@@ -83,7 +25,7 @@ def send_webhook(
as_json: bool = False,
):
try:
parsed = validate_outbound_http_url(
validate_outbound_http_url(
url,
allowed_schemes=settings.WEBHOOKS_ALLOWED_SCHEMES,
allowed_ports=settings.WEBHOOKS_ALLOWED_PORTS,
@@ -94,12 +36,7 @@ def send_webhook(
logger.warning("Webhook blocked: %s", e)
raise
hostname = parsed.hostname
if hostname is None: # pragma: no cover
raise ValueError("Invalid URL scheme or hostname.")
transport = WebhookTransport(
hostname=hostname,
transport = PinnedHostHTTPTransport(
allow_internal=settings.WEBHOOKS_ALLOW_INTERNAL_REQUESTS,
)
+120
View File
@@ -4,6 +4,8 @@ from collections.abc import Collection
from urllib.parse import ParseResult
from urllib.parse import urlparse
import httpx
def is_public_ip(ip: str | int) -> bool:
try:
@@ -74,3 +76,121 @@ def validate_outbound_http_url(
)
return parsed
def _rewrite_request_to_pinned_ip(
request: httpx.Request,
*,
allow_internal: bool,
) -> httpx.Request:
hostname = request.url.host
if not hostname:
raise httpx.ConnectError("No hostname in request URL")
try:
ips = resolve_hostname_ips(hostname)
except ValueError as e:
raise httpx.ConnectError(str(e)) from e
if not allow_internal:
for ip_str in ips:
if not is_public_ip(ip_str):
raise httpx.ConnectError(
f"Connection blocked: {hostname} resolves to a non-public address",
)
ip_str = ips[0]
formatted_ip = format_host_for_url(ip_str)
new_headers = httpx.Headers(request.headers)
if "host" in new_headers:
del new_headers["host"]
host_header = format_host_for_url(hostname)
default_port = 443 if request.url.scheme == "https" else 80
if request.url.port and request.url.port != default_port:
host_header = f"{host_header}:{request.url.port}"
new_headers["Host"] = host_header
new_url = request.url.copy_with(host=formatted_ip)
rewritten_request = httpx.Request(
method=request.method,
url=new_url,
headers=new_headers,
content=request.stream,
extensions=request.extensions,
)
rewritten_request.extensions["sni_hostname"] = hostname
return rewritten_request
class PinnedHostHTTPTransport(httpx.HTTPTransport):
"""
HTTP transport that resolves/validates hostnames per request and connects to
a vetted IP while preserving the original Host header and TLS SNI hostname.
"""
def __init__(
self,
*args,
allow_internal: bool = False,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
self.allow_internal = allow_internal
def handle_request(self, request: httpx.Request) -> httpx.Response:
request = _rewrite_request_to_pinned_ip(
request,
allow_internal=self.allow_internal,
)
return super().handle_request(request)
class PinnedHostAsyncHTTPTransport(httpx.AsyncHTTPTransport):
"""
Async variant of PinnedHostHTTPTransport.
"""
def __init__(
self,
*args,
allow_internal: bool = False,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
self.allow_internal = allow_internal
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
request = _rewrite_request_to_pinned_ip(
request,
allow_internal=self.allow_internal,
)
return await super().handle_async_request(request)
def create_pinned_httpx_client(
url: str,
*,
allow_internal: bool = False,
**kwargs,
) -> httpx.Client:
validate_outbound_http_url(url, allow_internal=allow_internal)
return httpx.Client(
transport=PinnedHostHTTPTransport(allow_internal=allow_internal),
**kwargs,
)
def create_pinned_async_httpx_client(
url: str,
*,
allow_internal: bool = False,
**kwargs,
) -> httpx.AsyncClient:
validate_outbound_http_url(url, allow_internal=allow_internal)
return httpx.AsyncClient(
transport=PinnedHostAsyncHTTPTransport(allow_internal=allow_internal),
**kwargs,
)
+50
View File
@@ -0,0 +1,50 @@
from unittest import mock
import httpx
import pytest
from paperless.network import PinnedHostHTTPTransport
def test_pinned_host_transport_blocks_internal_rebinding():
transport = PinnedHostHTTPTransport(allow_internal=False)
request = httpx.Request("GET", "http://example.com/test")
with (
mock.patch(
"paperless.network.resolve_hostname_ips",
return_value=["127.0.0.1"],
),
pytest.raises(httpx.ConnectError, match="non-public address"),
):
transport.handle_request(request)
def test_pinned_host_transport_rewrites_to_vetted_ip():
transport = PinnedHostHTTPTransport(allow_internal=False)
request = httpx.Request("GET", "https://example.com:8443/test")
def assert_rewritten_request(
self,
rewritten_request,
):
assert str(rewritten_request.url) == "https://93.184.216.34:8443/test"
assert rewritten_request.headers["Host"] == "example.com:8443"
assert rewritten_request.extensions["sni_hostname"] == "example.com"
return httpx.Response(200, request=rewritten_request)
with (
mock.patch(
"paperless.network.resolve_hostname_ips",
return_value=["93.184.216.34"],
),
mock.patch.object(
httpx.HTTPTransport,
"handle_request",
autospec=True,
side_effect=assert_rewritten_request,
),
):
response = transport.handle_request(request)
assert response.status_code == 200
+31 -1
View File
@@ -9,6 +9,10 @@ if TYPE_CHECKING:
from llama_index.llms.openai_like import OpenAILike
from paperless.config import AIConfig
from paperless.network import PinnedHostAsyncHTTPTransport
from paperless.network import PinnedHostHTTPTransport
from paperless.network import create_pinned_async_httpx_client
from paperless.network import create_pinned_httpx_client
from paperless.network import validate_outbound_http_url
from paperless_ai.base_model import DocumentClassifierSchema
@@ -27,23 +31,47 @@ class AIClient:
def get_llm(self) -> "Ollama | OpenAILike":
if self.settings.llm_backend == LLMBackend.OLLAMA:
from llama_index.llms.ollama import Ollama
from ollama import AsyncClient
from ollama import Client
endpoint = self.settings.llm_endpoint or "http://localhost:11434"
validate_outbound_http_url(
endpoint,
allow_internal=self.settings.llm_allow_internal_endpoints,
)
transport = PinnedHostHTTPTransport(
allow_internal=self.settings.llm_allow_internal_endpoints,
)
async_transport = PinnedHostAsyncHTTPTransport(
allow_internal=self.settings.llm_allow_internal_endpoints,
)
return Ollama(
model=self.settings.llm_model or "llama3.1",
base_url=endpoint,
request_timeout=120,
client=Client(
host=endpoint,
timeout=120,
transport=transport,
),
async_client=AsyncClient(
host=endpoint,
timeout=120,
transport=async_transport,
),
)
elif self.settings.llm_backend == LLMBackend.OPENAI_LIKE:
from llama_index.llms.openai_like import OpenAILike
endpoint = self.settings.llm_endpoint or None
http_client = None
async_http_client = None
if endpoint:
validate_outbound_http_url(
http_client = create_pinned_httpx_client(
endpoint,
allow_internal=self.settings.llm_allow_internal_endpoints,
)
async_http_client = create_pinned_async_httpx_client(
endpoint,
allow_internal=self.settings.llm_allow_internal_endpoints,
)
@@ -53,6 +81,8 @@ class AIClient:
api_key=self.settings.llm_api_key,
is_chat_model=True,
is_function_calling_model=True,
http_client=http_client,
async_http_client=async_http_client,
)
else:
raise ValueError(f"Unsupported LLM backend: {self.settings.llm_backend}")
+29 -2
View File
@@ -13,6 +13,10 @@ from documents.models import Document
from documents.models import Note
from paperless.config import AIConfig
from paperless.models import LLMEmbeddingBackend
from paperless.network import PinnedHostAsyncHTTPTransport
from paperless.network import PinnedHostHTTPTransport
from paperless.network import create_pinned_async_httpx_client
from paperless.network import create_pinned_httpx_client
from paperless.network import validate_outbound_http_url
OCR_LEADER_REGEX = re.compile(r"[._\-\u00b7]{4,}")
@@ -27,8 +31,14 @@ def get_embedding_model() -> "BaseEmbedding":
from llama_index.embeddings.openai_like import OpenAILikeEmbedding
endpoint = config.llm_embedding_endpoint or config.llm_endpoint or None
http_client = None
async_http_client = None
if endpoint:
validate_outbound_http_url(
http_client = create_pinned_httpx_client(
endpoint,
allow_internal=config.llm_allow_internal_endpoints,
)
async_http_client = create_pinned_async_httpx_client(
endpoint,
allow_internal=config.llm_allow_internal_endpoints,
)
@@ -36,6 +46,8 @@ def get_embedding_model() -> "BaseEmbedding":
model_name=config.llm_embedding_model or "text-embedding-3-small",
api_key=config.llm_api_key,
api_base=endpoint,
http_client=http_client,
async_http_client=async_http_client,
)
case LLMEmbeddingBackend.HUGGINGFACE:
from llama_index.embeddings.huggingface import HuggingFaceEmbedding
@@ -47,6 +59,8 @@ def get_embedding_model() -> "BaseEmbedding":
)
case LLMEmbeddingBackend.OLLAMA:
from llama_index.embeddings.ollama import OllamaEmbedding
from ollama import AsyncClient
from ollama import Client
endpoint = (
config.llm_embedding_endpoint
@@ -57,10 +71,23 @@ def get_embedding_model() -> "BaseEmbedding":
endpoint,
allow_internal=config.llm_allow_internal_endpoints,
)
return OllamaEmbedding(
embedding = OllamaEmbedding(
model_name=config.llm_embedding_model or "embeddinggemma",
base_url=endpoint,
)
embedding._client = Client(
host=endpoint,
transport=PinnedHostHTTPTransport(
allow_internal=config.llm_allow_internal_endpoints,
),
)
embedding._async_client = AsyncClient(
host=endpoint,
transport=PinnedHostAsyncHTTPTransport(
allow_internal=config.llm_allow_internal_endpoints,
),
)
return embedding
case _:
raise ValueError(
f"Unsupported embedding backend: {config.llm_embedding_backend}",
+5
View File
@@ -1,3 +1,4 @@
from unittest.mock import ANY
from unittest.mock import MagicMock
from unittest.mock import patch
@@ -40,6 +41,8 @@ def test_get_llm_ollama(mock_ai_config, mock_ollama_llm):
model="test_model",
base_url="http://test-url",
request_timeout=120,
client=ANY,
async_client=ANY,
)
assert client.llm == mock_ollama_llm.return_value
@@ -58,6 +61,8 @@ def test_get_llm_openai(mock_ai_config, mock_openai_llm):
api_key="test_api_key",
is_chat_model=True,
is_function_calling_model=True,
http_client=ANY,
async_http_client=ANY,
)
assert client.llm == mock_openai_llm.return_value
+5
View File
@@ -1,4 +1,5 @@
import json
from unittest.mock import ANY
from unittest.mock import MagicMock
from unittest.mock import patch
@@ -70,6 +71,8 @@ def test_get_embedding_model_openai(mock_ai_config):
model_name="text-embedding-3-small",
api_key="test_api_key",
api_base="http://test-url",
http_client=ANY,
async_http_client=ANY,
)
assert model == MockOpenAIEmbedding.return_value
@@ -89,6 +92,8 @@ def test_get_embedding_model_openai_prefers_embedding_endpoint(mock_ai_config):
model_name="text-embedding-3-small",
api_key="test_api_key",
api_base="http://embedding-url",
http_client=ANY,
async_http_client=ANY,
)
assert model == MockOpenAIEmbedding.return_value