diff --git a/src/documents/workflows/webhooks.py b/src/documents/workflows/webhooks.py index 0c510a35d..e2eff558d 100644 --- a/src/documents/workflows/webhooks.py +++ b/src/documents/workflows/webhooks.py @@ -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, ) diff --git a/src/paperless/network.py b/src/paperless/network.py index 4af00052f..8a5b28a91 100644 --- a/src/paperless/network.py +++ b/src/paperless/network.py @@ -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, + ) diff --git a/src/paperless/tests/test_network.py b/src/paperless/tests/test_network.py new file mode 100644 index 000000000..a306c52b4 --- /dev/null +++ b/src/paperless/tests/test_network.py @@ -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 diff --git a/src/paperless_ai/client.py b/src/paperless_ai/client.py index ce1874461..d4bcef0c8 100644 --- a/src/paperless_ai/client.py +++ b/src/paperless_ai/client.py @@ -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}") diff --git a/src/paperless_ai/embedding.py b/src/paperless_ai/embedding.py index 7ef841e4b..407dd4c0e 100644 --- a/src/paperless_ai/embedding.py +++ b/src/paperless_ai/embedding.py @@ -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}", diff --git a/src/paperless_ai/tests/test_client.py b/src/paperless_ai/tests/test_client.py index 35a881400..ae903b8a0 100644 --- a/src/paperless_ai/tests/test_client.py +++ b/src/paperless_ai/tests/test_client.py @@ -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 diff --git a/src/paperless_ai/tests/test_embedding.py b/src/paperless_ai/tests/test_embedding.py index f8bc98f7d..d3eff080e 100644 --- a/src/paperless_ai/tests/test_embedding.py +++ b/src/paperless_ai/tests/test_embedding.py @@ -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