diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 312f46bb021..f9581aa3841 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -553,6 +553,10 @@ class HTTPResponseLimitError(ValueError): pass +class HTTPResponseEncodingError(HTTPResponseLimitError): + pass + + class MaskedHTTPStatusError(httpx.HTTPStatusError): def __init__(self, original_error, message: str | None = None, text: str | None = None): # Create a new error with the masked URL @@ -781,7 +785,7 @@ class AsyncHTTPHandler: if response.is_redirect or response.is_error: return httpx.Response(response.status_code, headers=response.headers, request=response.request) if response.headers.get("content-encoding", "identity").lower() != "identity": - raise HTTPResponseLimitError("Response size limits require an uncompressed response") + raise HTTPResponseEncodingError("Response size limits require an uncompressed response") if int(response.headers.get("content-length", "0")) > max_bytes: raise HTTPResponseLimitError("Response exceeds the configured size limit") with BytesIO() as body: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 4597872d84e..bfc894798d3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -52,6 +52,7 @@ from litellm.types.proxy.carried_budget_state import ( UserBudgetSnapshot, ) from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry +from litellm.types.proxy.rag_ingest import RagIngestSettings from litellm.types.proxy.spend_capture_rate import SpendCaptureRateCheckSettings from litellm.types.router import RouterErrors, UpdateRouterConfig from litellm.types.router_weights import validate_router_settings_dict @@ -2761,6 +2762,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="override user_api_key_auth with your own auth script - https://docs.litellm.ai/docs/proxy/virtual_keys#custom-auth", ) + rag_ingest: RagIngestSettings | None = Field( + None, + description=( + "controls run on every upload that lands in a vector store (/v1/rag/ingest and /v1/files with purpose " + "assistants or user_data); rag_ingest.malware_scanner names the scanner instance as . " + "the way custom_auth does, and unset runs the EICAR test scanner - https://docs.litellm.ai/docs/rag_ingest" + ), + ) max_parallel_requests: int | None = Field( None, description="maximum parallel requests for each api key", diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index f5e62da5962..4a6ad8adc47 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -94,6 +94,7 @@ from litellm.proxy.openai_files_endpoints.general_upload_validation import ( coerce_optional_str_list_setting, raise_upload_validation_failure, ) +from litellm.proxy.rag_endpoints.upload_security import MalwareScanner, RejectedUpload, validate_upload from litellm.proxy.utils import PrismaClient, ProxyLogging, is_known_model from litellm.repositories.table_repositories import ManagedFileRepository from litellm.router import Router @@ -527,6 +528,26 @@ async def route_create_file( return response +_VECTOR_STORE_UPLOAD_PURPOSES: Final[frozenset[str]] = frozenset({"assistants", "user_data"}) + + +async def _reject_unsafe_vector_store_upload( + purpose: str, + file_source: bytes | BinaryIO, + scanner: MalwareScanner, +) -> None: + if purpose not in _VECTOR_STORE_UPLOAD_PURPOSES or not isinstance(file_source, bytes): + return + validation: Final = await asyncio.to_thread(validate_upload, content=file_source, scanner=scanner) + if isinstance(validation, RejectedUpload): + raise ProxyException( + message=f"{validation.message} Rejection reason: {validation.reason.value}.", + type="invalid_request_error", + param="file", + code=400, + ) + + @router.post( "/{provider}/v1/files", dependencies=[Depends(user_api_key_auth)], @@ -577,6 +598,7 @@ async def create_file( llm_router, proxy_config, proxy_logging_obj, + rag_upload_malware_scanner, version, ) @@ -655,6 +677,8 @@ async def create_file( if blocked_extension_failure is not None: raise_upload_validation_failure(blocked_extension_failure) + await _reject_unsafe_vector_store_upload(purpose, file_source, rag_upload_malware_scanner) + if passthrough: _validate_passthrough_upload( purpose=purpose, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 61304f0d919..22ba313ed0f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -742,6 +742,12 @@ from litellm.proxy.prometheus_cleanup import mark_dead_workers, mark_worker_exit from litellm.proxy.public_endpoints import router as public_endpoints_router from litellm.proxy.public_endpoints.public_v1 import router as public_v1_router from litellm.proxy.rag_endpoints.endpoints import router as rag_router +from litellm.proxy.rag_endpoints.upload_security import ( + EicarTestMalwareScanner, + MalwareScanner, + MalwareScannerConfigError, + resolve_malware_scanner, +) from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router from litellm.proxy.response_api_endpoints.endpoints import router as response_router from litellm.proxy.route_llm_request import route_request @@ -1035,6 +1041,7 @@ def cleanup_router_config_variables(): user_config_file_path, \ otel_logging, \ user_custom_auth, \ + rag_upload_malware_scanner, \ user_custom_auth_path, \ user_custom_key_generate, \ user_custom_key_update, \ @@ -1053,6 +1060,7 @@ def cleanup_router_config_variables(): user_config_file_path = None otel_logging = None user_custom_auth = None + rag_upload_malware_scanner = EicarTestMalwareScanner() user_custom_auth_path = None user_custom_key_generate = None user_custom_key_update = None @@ -2564,6 +2572,7 @@ polling_via_cache_enabled: Literal["all"] | list[str] | bool = False native_background_mode: list[str] = [] # Models that should use native provider background mode instead of polling polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache user_custom_auth = None +rag_upload_malware_scanner: MalwareScanner = EicarTestMalwareScanner() user_custom_key_generate = None # Sentinel: prevents PKCE-no-Redis advisory from re-logging on config hot-reload. # Tests that need to reset it can patch 'litellm.proxy.proxy_server._pkce_no_redis_warning_emitted'. @@ -4511,6 +4520,14 @@ _DB_OVERLAY_REMOTE_MODULE_LIST_FIELDS: Final[dict[str, tuple[str, ...]]] = { "audit_log_callbacks", ), } +_DB_OVERLAY_REMOTE_MODULE_NESTED_STR_FIELDS: Final[Mapping[str, tuple[tuple[str, str], ...]]] = MappingProxyType( + { + "general_settings": ( + ("litellm_jwtauth", "custom_validate"), + ("rag_ingest", "malware_scanner"), + ), + } +) def _is_remote_module_url(value: object) -> bool: @@ -4607,16 +4624,17 @@ def _scrub_db_overlay_remote_module_loads(section: str, db_value: object) -> obj if isinstance(lp, dict): _scrub_guardrail_inner(lp) - # ``general_settings.litellm_jwtauth.custom_validate`` is a nested - # string field. - if section == "general_settings": - jwt: Final = sanitized.get("litellm_jwtauth") - if isinstance(jwt, dict) and _is_remote_module_url(jwt.get("custom_validate")): + for parent, nested_field in _DB_OVERLAY_REMOTE_MODULE_NESTED_STR_FIELDS.get(section, ()): + if isinstance(nested := sanitized.get(parent), dict) and _is_remote_module_url(nested.get(nested_field)): verbose_proxy_logger.warning( - "Refused remote-URL custom_validate from DB-overlay general_settings.litellm_jwtauth: %r", - jwt.get("custom_validate"), + "Refused remote-URL %s from DB-overlay %s.%s: %r", + nested_field, + section, + parent, + nested.get(nested_field), ) - jwt["custom_validate"] = None + nested[nested_field] = None + if section == "general_settings": # ``pass_through_endpoints`` is a list of dicts whose ``target`` # is passed through ``create_pass_through_route`` → # ``get_instance_fn``. A DB-overlay ``target: "s3://attacker/m.i"`` @@ -5931,6 +5949,7 @@ class ProxyConfig: user_config_file_path, \ otel_logging, \ user_custom_auth, \ + rag_upload_malware_scanner, \ user_custom_auth_path, \ user_custom_key_generate, \ user_custom_key_update, \ @@ -6347,6 +6366,15 @@ class ProxyConfig: TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"]) + resolved_malware_scanner: Final = resolve_malware_scanner( + general_settings.get("rag_ingest"), + config_file_path=config_file_path, + load_instance=get_instance_fn, + ) + if isinstance(resolved_malware_scanner, MalwareScannerConfigError): + raise ValueError(resolved_malware_scanner.message) + rag_upload_malware_scanner = resolved_malware_scanner + if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) if declared_proxy_ranges(general_settings) is None: @@ -6463,7 +6491,6 @@ class ProxyConfig: custom_auth_configured=custom_auth is not None, run_common_checks=bool(general_settings.get("custom_auth_run_common_checks", False)), ) - log_once_if_budget_reservation_disabled( disabled=general_settings.get("disable_budget_reservation") is True, ) diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 974bff6338a..43a1cbfc1c7 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -6,12 +6,14 @@ Provides: - /rag/query: RAG query pipeline (Search -> Rerank -> LLM Completion) """ +import asyncio import base64 import json -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Any, Final, TypeAlias +import httpx import orjson from fastapi import APIRouter, Depends, HTTPException, Request, Response, status from fastapi.responses import ORJSONResponse, StreamingResponse @@ -23,6 +25,12 @@ from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, ) +from litellm.litellm_core_utils.url_utils import async_safe_get +from litellm.llms.custom_httpx.http_handler import ( + HTTPResponseEncodingError, + HTTPResponseLimitError, + get_async_httpx_client, +) from litellm.proxy._types import * from litellm.proxy.auth.auth_utils import is_request_body_safe from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth @@ -38,9 +46,9 @@ from litellm.proxy.common_utils.http_parsing_utils import ( ) from litellm.proxy.rag_endpoints.upload_security import ( MAX_UPLOAD_SIZE_BYTES, - EicarTestMalwareScanner, MalwareScanner, RejectedUpload, + RejectionReason, validate_upload, ) from litellm.proxy.vector_store_endpoints.endpoints import ( @@ -52,6 +60,7 @@ from litellm.proxy.vector_store_endpoints.utils import ( ) from litellm.rag.main import get_ingestion_class from litellm.repositories.table_repositories import ManagedVectorStoresRepository +from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.utils import ModelResponse if TYPE_CHECKING: @@ -388,17 +397,60 @@ async def _save_vector_store_to_db_from_rag_ingest( verbose_proxy_logger.exception("Failed to save vector store %s to database: %s", vector_store_id, db_error) -def _secure_uploaded_file( - file_data: tuple[str, bytes, str], - scanner: MalwareScanner, -) -> tuple[str, bytes, str]: - validation: Final = validate_upload(content=file_data[1], scanner=scanner) +async def _secure_uploaded_file(content: bytes, scanner: MalwareScanner) -> tuple[str, bytes, str]: + validation: Final = await asyncio.to_thread(validate_upload, content=content, scanner=scanner) if isinstance(validation, RejectedUpload): raise HTTPException( status_code=400, detail={"error": validation.message, "reason": validation.reason.value}, ) - return validation.safe_filename, file_data[1], validation.content_type + return validation.safe_filename, content, validation.content_type + + +UrlFetcher: TypeAlias = Callable[[str], Awaitable[httpx.Response]] + +_FILE_URL_FETCH_TIMEOUT: Final = httpx.Timeout(connect=10.0, read=60.0, write=10.0, pool=10.0) + + +async def fetch_file_url(url: str) -> httpx.Response: + client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.RAG) + return await async_safe_get(client, url, timeout=_FILE_URL_FETCH_TIMEOUT, max_response_bytes=MAX_UPLOAD_SIZE_BYTES) + + +async def _download_file_url(file_url: str, fetch_url: UrlFetcher) -> bytes: + try: + response: Final = await fetch_url(file_url) + except HTTPResponseEncodingError as e: + raise HTTPException(status_code=400, detail={"error": f"Could not fetch file_url: {e}"}) from e + except HTTPResponseLimitError as e: + raise HTTPException( + status_code=400, + detail={"error": f"Could not fetch file_url: {e}", "reason": RejectionReason.FILE_TOO_LARGE.value}, + ) from e + except (ValueError, httpx.HTTPError) as e: + raise HTTPException(status_code=400, detail={"error": f"Could not fetch file_url: {e}"}) from e + if not response.is_success: + raise HTTPException( + status_code=400, + detail={"error": f"Could not fetch file_url: the server answered HTTP {response.status_code}"}, + ) + return response.content + + +async def fetch_and_secure_file_url( + file_url: str, + *, + scanner: MalwareScanner, + fetch_url: UrlFetcher, +) -> tuple[str, bytes, str]: + return await _secure_uploaded_file(await _download_file_url(file_url, fetch_url), scanner) + + +def _optional_str_field(body: Mapping[str, object], name: str) -> str | None: + value: Final = body.get(name) + if value is None or isinstance(value, str): + return value + raise HTTPException(status_code=400, detail={"error": f"'{name}' must be a string"}) async def parse_rag_ingest_request( @@ -415,7 +467,9 @@ async def parse_rag_ingest_request( Uploaded file bytes are validated against the vector-store upload controls (size limit, format allowlist with content inspection, archive rejection, and the injected malware scanner) and given a server-generated filename - before they are returned. + before they are returned. A ``file_url`` is returned as given; the route + fetches it through :func:`fetch_and_secure_file_url` once the rest of the + request has been validated. Returns: Tuple of (ingest_options, file_data, file_url, file_id) @@ -443,15 +497,15 @@ async def parse_rag_ingest_request( if request_json_str: request_data: Final = orjson.loads(request_json_str) ingest_options = request_data.get("ingest_options", {}) - file_url = request_data.get("file_url") - file_id = request_data.get("file_id") + file_url = _optional_str_field(request_data, "file_url") + file_id = _optional_str_field(request_data, "file_id") else: # JSON body data: Final = await _read_request_body(request) ingest_options = data.get("ingest_options", {}) - file_url = data.get("file_url") - file_id = data.get("file_id") + file_url = _optional_str_field(data, "file_url") + file_id = _optional_str_field(data, "file_id") # Handle base64-encoded file in JSON body file_obj = data.get("file") @@ -477,10 +531,6 @@ async def parse_rag_ingest_request( detail={"error": "Must provide file, file_url, or file_id"}, ) - secured_file_data: Final[tuple[str, bytes, str] | None] = ( - _secure_uploaded_file(file_data, scanner) if file_data is not None else None - ) - if "vector_store" not in ingest_options: raise HTTPException( status_code=400, @@ -522,6 +572,8 @@ async def parse_rag_ingest_request( }, ) + secured_file_data: Final = await _secure_uploaded_file(file_data[1], scanner) if file_data is not None else None + return ingest_options, secured_file_data, file_url, file_id @@ -580,13 +632,14 @@ async def rag_ingest( llm_router, prisma_client, proxy_config, + rag_upload_malware_scanner, version, ) try: # Parse request ingest_options, file_data, file_url, file_id = await parse_rag_ingest_request( - request, scanner=EicarTestMalwareScanner() + request, scanner=rag_upload_malware_scanner ) # INTERNAL_USER_VIEW_ONLY can ingest to existing vector stores only @@ -634,6 +687,12 @@ async def rag_ingest( detail={"error": provider_error}, # mutable-ok: FastAPI serializes the detail as JSON ) + uploaded_file_data: Final = ( + await fetch_and_secure_file_url(file_url, scanner=rag_upload_malware_scanner, fetch_url=fetch_file_url) + if file_data is None and file_url is not None + else file_data + ) + # Add litellm data request_data: dict[str, Any] = {} request_data = await add_litellm_data_to_request( @@ -654,8 +713,7 @@ async def rag_ingest( # Call ingest response: Final = await litellm.aingest( ingest_options=merged_ingest_options, - file_data=file_data, - file_url=file_url, + file_data=uploaded_file_data, file_id=file_id, router=llm_router, **request_data, diff --git a/litellm/proxy/rag_endpoints/upload_security.py b/litellm/proxy/rag_endpoints/upload_security.py index f0318f2f709..e7f6a5e7fef 100644 --- a/litellm/proxy/rag_endpoints/upload_security.py +++ b/litellm/proxy/rag_endpoints/upload_security.py @@ -11,16 +11,21 @@ server-generated filename so the client-controlled name never reaches storage. from __future__ import annotations import uuid -from collections.abc import Mapping +from collections.abc import Callable, Mapping from dataclasses import dataclass from enum import Enum from types import MappingProxyType from typing import Final, Protocol, TypeAlias, runtime_checkable +from pydantic import ValidationError from typing_extensions import assert_never +from litellm.types.proxy.rag_ingest import RagIngestSettings + MAX_UPLOAD_SIZE_BYTES: Final = 512 * 1024 * 1024 +MALWARE_SCANNER_OPTION: Final = "general_settings.rag_ingest.malware_scanner" + EICAR_TEST_SIGNATURE: Final = b"X5O!P%@AP[4\\PZX54(P^)7CC)7}$EICAR-STANDARD-ANTIVIRUS-TEST-FILE!$H+H*" _ARCHIVE_MAGIC_PREFIXES: Final[tuple[bytes, ...]] = ( @@ -91,6 +96,13 @@ class ScanResult: @runtime_checkable class MalwareScanner(Protocol): + """One shared instance scans every upload, from worker threads running at the same time. + + ``scan`` must therefore be safe to call concurrently and must return + :class:`ScanResult` with the ``error`` verdict instead of raising when the + engine is unavailable. + """ + def scan(self, content: bytes) -> ScanResult: ... @@ -110,6 +122,45 @@ class EicarTestMalwareScanner: return ScanResult(verdict=ScanVerdict.CLEAN) +@dataclass(frozen=True, slots=True) +class MalwareScannerConfigError: + message: str + + +InstanceLoader: TypeAlias = Callable[[str, str | None], object] + + +def resolve_malware_scanner( + rag_ingest: object, + *, + config_file_path: str | None, + load_instance: InstanceLoader, +) -> MalwareScanner | MalwareScannerConfigError: + try: + settings: Final = RagIngestSettings() if rag_ingest is None else RagIngestSettings.model_validate(rag_ingest) + except ValidationError as e: + return MalwareScannerConfigError(f"general_settings.rag_ingest is invalid: {e}") + if settings.malware_scanner is None: + return EicarTestMalwareScanner() + try: + loaded: Final = load_instance(settings.malware_scanner, config_file_path) + except Exception as e: + return MalwareScannerConfigError( + f"{MALWARE_SCANNER_OPTION}={settings.malware_scanner!r} could not be loaded: {e}" + ) + if isinstance(loaded, type): + return MalwareScannerConfigError( + f"{MALWARE_SCANNER_OPTION}={settings.malware_scanner!r} names the class {loaded.__qualname__}; " + "name an instance of it instead" + ) + if isinstance(loaded, MalwareScanner) and callable(loaded.scan): + return loaded + return MalwareScannerConfigError( + f"{MALWARE_SCANNER_OPTION}={settings.malware_scanner!r} resolved to a {type(loaded).__qualname__}, " + "which has no scan(content: bytes) -> ScanResult method" + ) + + @dataclass(frozen=True, slots=True) class AllowedContent: format: DetectedFormat diff --git a/litellm/types/proxy/rag_ingest.py b/litellm/types/proxy/rag_ingest.py new file mode 100644 index 00000000000..e091c6a0434 --- /dev/null +++ b/litellm/types/proxy/rag_ingest.py @@ -0,0 +1,17 @@ +"""``general_settings.rag_ingest``: the controls run on bytes before they land in a vector store.""" + +from pydantic import BaseModel, ConfigDict, Field + + +class RagIngestSettings(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + malware_scanner: str | None = Field( + None, + description=( + "The scanner every vector-store upload goes through, as . where the module sits " + "next to config.yaml or is importable, resolved the way custom_auth is. The instance must expose " + "scan(content: bytes) -> ScanResult and be safe to call from several threads at once. Unset runs " + "the EICAR test scanner, which flags only the EICAR test file" + ), + ) diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 84cf4ea7c32..23f4cc56371 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -4939,7 +4939,7 @@ def test_create_file_blocked_extension_unset_allows_everything(monkeypatch, llm_ try: response = client.post( "/v1/files", - files={"file": ("payload.exe", b"MZ\x90\x00", "application/octet-stream")}, + files={"file": ("payload.exe", b"plain text under an executable name\n", "application/octet-stream")}, data={"purpose": "user_data"}, headers={"Authorization": "Bearer test-key"}, ) @@ -5907,3 +5907,129 @@ def test_create_file_passthrough_fails_closed_when_guardrails_would_scan_the_bat assert error["param"] == "passthrough" assert "guardrails" in error["message"] assert forwarded_calls == [] + + +EICAR_BYTES: Final = b"X5O!P%@AP[4\\PZX54(P^)7CC)7}$EICAR-STANDARD-ANTIVIRUS-TEST-FILE!$H+H*" + + +def _post_file_as_admin(monkeypatch, llm_router: Router, *, purpose: str, name: str, content: bytes, ctype: str): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.openai_files_endpoints import files_endpoints as fe + + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) + setup_proxy_logging_object(monkeypatch, llm_router) + captured: dict = {} + + async def fake_route_create_file(*, _create_file_request, **kwargs): + captured["file"] = _create_file_request["file"] + return OpenAIFileObject( + id="file-clean", + object="file", + bytes=len(content), + created_at=1234567890, + filename=name, + purpose=purpose, + status="uploaded", + ) + + monkeypatch.setattr(fe, "route_create_file", fake_route_create_file) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user" + ) + try: + resp = client.post( + "/v1/files", + files={"file": (name, content, ctype)}, + data={"purpose": purpose}, + headers={"Authorization": "Bearer test-key"}, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + return resp, captured + + +@pytest.mark.parametrize("purpose", ["assistants", "user_data"]) +@pytest.mark.parametrize( + "name, content, ctype, reason", + [ + ("notes.txt", EICAR_BYTES, "text/plain", "malware_detected"), + ("setup.txt", b"#!/bin/sh\nrm -rf /\n", "text/plain", "executable_not_allowed"), + ("report.pdf", b"PK\x03\x04\x14\x00\x00\x00payload", "application/pdf", "archive_not_allowed"), + ("empty.txt", b"", "text/plain", "empty_file"), + ], +) +def test_create_file_for_vector_store_purposes_runs_the_upload_controls( + monkeypatch, llm_router: Router, purpose, name, content, ctype, reason +): + resp, captured = _post_file_as_admin( + monkeypatch, llm_router, purpose=purpose, name=name, content=content, ctype=ctype + ) + assert resp.status_code == 400, resp.text + error = resp.json()["error"] + assert error["type"] == "invalid_request_error" + assert error["param"] == "file" + assert f"Rejection reason: {reason}." in error["message"] + assert captured == {} + + +@pytest.mark.parametrize("purpose", ["assistants", "user_data"]) +@pytest.mark.parametrize( + "name, content, ctype", + [ + ("Quarterly Notes (final).txt", b"benign document text\n", "text/plain"), + ("handbook.pdf", b"%PDF-1.7\n1 0 obj<<>>endobj\n", "application/pdf"), + ], +) +def test_create_file_for_vector_store_purposes_keeps_the_callers_filename_on_clean_uploads( + monkeypatch, llm_router: Router, purpose, name, content, ctype +): + resp, captured = _post_file_as_admin( + monkeypatch, llm_router, purpose=purpose, name=name, content=content, ctype=ctype + ) + assert resp.status_code == 200, resp.text + assert captured["file"][0] == name + assert captured["file"][1] == content + + +def test_create_file_for_vector_store_purposes_uses_the_configured_scanner(monkeypatch, llm_router: Router): + from litellm.proxy.rag_endpoints.upload_security import ScanResult, ScanVerdict + + class _InfectedScanner: + def scan(self, content: bytes) -> ScanResult: + return ScanResult(verdict=ScanVerdict.INFECTED, signature="Custom.Sig") + + monkeypatch.setattr("litellm.proxy.proxy_server.rag_upload_malware_scanner", _InfectedScanner()) + resp, captured = _post_file_as_admin( + monkeypatch, llm_router, purpose="assistants", name="notes.txt", content=b"benign text\n", ctype="text/plain" + ) + assert resp.status_code == 400, resp.text + assert "Custom.Sig" in resp.json()["error"]["message"] + assert captured == {} + + +@pytest.mark.parametrize( + "purpose, name, content, ctype", + [ + ("batch", "batch.jsonl", VALID_BATCH_LINE, "application/jsonl"), + ("fine-tune", "train.jsonl", VALID_BATCH_LINE, "application/jsonl"), + ("vision", "photo.png", b"\x89PNG\r\n\x1a\n" + b"\x00" * 16, "image/png"), + ], +) +def test_create_file_other_purposes_skip_the_vector_store_upload_controls( + monkeypatch, llm_router: Router, purpose, name, content, ctype +): + from litellm.proxy.rag_endpoints.upload_security import ScanResult, ScanVerdict + + class _InfectedScanner: + def scan(self, content: bytes) -> ScanResult: + return ScanResult(verdict=ScanVerdict.INFECTED, signature="Custom.Sig") + + monkeypatch.setattr("litellm.proxy.proxy_server.rag_upload_malware_scanner", _InfectedScanner()) + resp, captured = _post_file_as_admin( + monkeypatch, llm_router, purpose=purpose, name=name, content=content, ctype=ctype + ) + assert resp.status_code == 200, resp.text + assert captured["file"][0] == name diff --git a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py index 2d654ea28ec..4907e5a9ba9 100644 --- a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py +++ b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py @@ -5,17 +5,28 @@ Covers: - internal_user_viewer restriction: can only ingest to existing vector stores (must provide vector_store_id) """ +import base64 import io import json +from dataclasses import dataclass from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +from fastapi import HTTPException from fastapi.testclient import TestClient import litellm +from litellm.llms.custom_httpx.http_handler import HTTPResponseEncodingError, HTTPResponseLimitError from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.proxy_server import app +from litellm.proxy.rag_endpoints.upload_security import ( + MAX_UPLOAD_SIZE_BYTES, + EicarTestMalwareScanner, + ScanResult, + ScanVerdict, +) @pytest.fixture @@ -58,19 +69,13 @@ def test_internal_user_viewer_rag_ingest_without_vector_store_id_rejected( response = client_internal_user_viewer.post( "/v1/rag/ingest", files={"file": ("sample.txt", io.BytesIO(b"test content"), "text/plain")}, - data={ - "request": '{"ingest_options":{"vector_store":{"custom_llm_provider":"openai"}}}' - }, + data={"request": '{"ingest_options":{"vector_store":{"custom_llm_provider":"openai"}}}'}, ) assert response.status_code == 403 detail = response.json() assert "detail" in detail - error_msg = ( - detail["detail"]["error"] - if isinstance(detail["detail"], dict) - else str(detail["detail"]) - ) + error_msg = detail["detail"]["error"] if isinstance(detail["detail"], dict) else str(detail["detail"]) assert "internal_user_viewer" in error_msg assert "vector_store_id" in error_msg @@ -97,8 +102,7 @@ def test_internal_user_viewer_rag_ingest_with_vector_store_id_passes_check( # Should not be 403 (role check passed) assert response.status_code != 403, ( - f"internal_user_viewer with vector_store_id should pass role check. " - f"Response: {response.json()}" + f"internal_user_viewer with vector_store_id should pass role check. Response: {response.json()}" ) @@ -114,15 +118,12 @@ def test_internal_user_rag_ingest_without_vector_store_id_allowed(client_interna response = client_internal_user.post( "/v1/rag/ingest", files={"file": ("sample.txt", io.BytesIO(b"test content"), "text/plain")}, - data={ - "request": '{"ingest_options":{"vector_store":{"custom_llm_provider":"openai"}}}' - }, + data={"request": '{"ingest_options":{"vector_store":{"custom_llm_provider":"openai"}}}'}, ) # Should not be 403 assert response.status_code != 403, ( - f"internal_user should be allowed to create new vector stores. " - f"Response: {response.json()}" + f"internal_user should be allowed to create new vector stores. Response: {response.json()}" ) @@ -170,13 +171,13 @@ def test_rag_ingest_blocks_clientside_credentials(client_internal_user, blocked_ }, }, ) - assert ( - response.status_code == 400 - ), f"Expected 400 when '{blocked_field}' is set clientside, got {response.status_code}: {response.json()}" + assert response.status_code == 400, ( + f"Expected 400 when '{blocked_field}' is set clientside, got {response.status_code}: {response.json()}" + ) body = response.json() - assert blocked_field in str( - body - ), f"Response should mention '{blocked_field}': {body}" + assert blocked_field in str(body), f"Response should mention '{blocked_field}': {body}" + + class TestRagIngestSSRFBlocked: """ aws_sts_endpoint and related credential-redirect fields must be rejected @@ -193,9 +194,7 @@ class TestRagIngestSSRFBlocked: ("aws_bedrock_runtime_endpoint", "https://attacker.example/bedrock"), ], ) - def test_ssrf_field_in_vector_store_config_rejected( - self, field, value, client_internal_user - ): + def test_ssrf_field_in_vector_store_config_rejected(self, field, value, client_internal_user): payload = { "file_url": "https://example.com/doc.pdf", "ingest_options": { @@ -215,12 +214,12 @@ class TestRagIngestSSRFBlocked: ) body = response.json() detail = body.get("detail", {}) - error_text = ( - detail.get("error", "") if isinstance(detail, dict) else str(detail) - ) + error_text = detail.get("error", "") if isinstance(detail, dict) else str(detail) assert field in error_text, f"Error should name the offending field: {error_text}" - def test_clean_bedrock_ingest_options_not_rejected(self, client_internal_user): + def test_clean_bedrock_ingest_options_not_rejected(self, client_internal_user, monkeypatch): + fetch_url, _seen = _fetcher_answering(200, b"%PDF-1.7\n1 0 obj<<>>endobj\n", content_type="application/pdf") + monkeypatch.setattr("litellm.proxy.rag_endpoints.endpoints.fetch_file_url", fetch_url) with patch( "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", new_callable=AsyncMock, @@ -230,14 +229,10 @@ class TestRagIngestSSRFBlocked: "/v1/rag/ingest", json={ "file_url": "https://example.com/doc.pdf", - "ingest_options": { - "vector_store": {"custom_llm_provider": "bedrock"} - }, + "ingest_options": {"vector_store": {"custom_llm_provider": "bedrock"}}, }, ) - assert response.status_code != 400, ( - f"Clean Bedrock ingest_options should not be rejected: {response.json()}" - ) + assert response.status_code != 400, f"Clean Bedrock ingest_options should not be rejected: {response.json()}" S3_REGISTRY_STORE = { @@ -703,12 +698,14 @@ def test_rag_query_returns_response_cost_header(client_internal_user): ) mock_response._hidden_params["response_cost"] = 3.45e-06 - with patch( - "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", - new_callable=AsyncMock, - return_value=mock_response, - ), patch("litellm.vector_store_registry", None), patch( - "litellm.proxy.proxy_server.prisma_client", None + with ( + patch( + "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", + new_callable=AsyncMock, + return_value=mock_response, + ), + patch("litellm.vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", None), ): response = client_internal_user.post( "/v1/rag/query", @@ -742,7 +739,9 @@ def test_rag_query_surfaces_upstream_status_code(client_internal_user, upstream_ new=AsyncMock(side_effect=upstream_error), ), patch("litellm.vector_store_registry", None), # test-quality-ok: proxy module global, no injection seam - patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: proxy module global, no injection seam + patch( + "litellm.proxy.proxy_server.prisma_client", None + ), # test-quality-ok: proxy module global, no injection seam ): response = client_internal_user.post( "/v1/rag/query", @@ -779,10 +778,14 @@ def test_rag_query_stream_returns_event_stream(client_internal_user): api_key="test-key", ) - with patch( - "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", - new=AsyncMock(side_effect=fake_aquery), - ), patch("litellm.vector_store_registry", None), patch("litellm.proxy.proxy_server.prisma_client", None): + with ( + patch( + "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", + new=AsyncMock(side_effect=fake_aquery), + ), + patch("litellm.vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + ): response = client_internal_user.post( "/v1/rag/query", json={ @@ -825,7 +828,9 @@ def test_rag_query_stream_pings_while_retrieval_is_still_running(client_internal new=AsyncMock(side_effect=slow_aquery), ), patch("litellm.vector_store_registry", None), # test-quality-ok: proxy module global, no injection seam - patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: proxy module global, no injection seam + patch( + "litellm.proxy.proxy_server.prisma_client", None + ), # test-quality-ok: proxy module global, no injection seam ): response = client_internal_user.post( "/v1/rag/query", @@ -849,9 +854,7 @@ def test_rag_query_stream_pings_while_retrieval_is_still_running(client_internal assert response.text.endswith("data: [DONE]\n\n") -def test_rag_query_stream_keeps_response_headers_when_retrieval_beats_the_keepalive( - client_internal_user, monkeypatch -): +def test_rag_query_stream_keeps_response_headers_when_retrieval_beats_the_keepalive(client_internal_user, monkeypatch): import litellm as litellm_module monkeypatch.setattr(litellm_module, "sse_keepalive_ping_interval_seconds", 5) @@ -873,7 +876,9 @@ def test_rag_query_stream_keeps_response_headers_when_retrieval_beats_the_keepal new=AsyncMock(side_effect=fast_aquery), ), patch("litellm.vector_store_registry", None), # test-quality-ok: proxy module global, no injection seam - patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: proxy module global, no injection seam + patch( + "litellm.proxy.proxy_server.prisma_client", None + ), # test-quality-ok: proxy module global, no injection seam ): response = client_internal_user.post( "/v1/rag/query", @@ -925,13 +930,17 @@ def test_rag_query_merges_managed_store_params(client_internal_user): model="gpt-4o-mini", ) - with patch( # test-quality-ok: aquery is the endpoint's downstream boundary; the forwarded config is what the test asserts - "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", - new_callable=AsyncMock, - return_value=mock_response, - ) as mock_aquery, patch.object(litellm, "vector_store_registry", mock_registry), patch( # test-quality-ok: seeds the managed-store registry the merge under test reads and grants access so real store resolution runs - "litellm.proxy.vector_store_endpoints.utils.can_user_access_vector_store", - new=AsyncMock(return_value=True), + with ( + patch( # test-quality-ok: aquery is the endpoint's downstream boundary; the forwarded config is what the test asserts + "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", + new_callable=AsyncMock, + return_value=mock_response, + ) as mock_aquery, + patch.object(litellm, "vector_store_registry", mock_registry), + patch( # test-quality-ok: seeds the managed-store registry the merge under test reads and grants access so real store resolution runs + "litellm.proxy.vector_store_endpoints.utils.can_user_access_vector_store", + new=AsyncMock(return_value=True), + ), ): response = client_internal_user.post( "/v1/rag/query", @@ -971,13 +980,17 @@ def test_rag_query_store_params_win_over_user_retrieval_config(client_internal_u model="gpt-4o-mini", ) - with patch( # test-quality-ok: aquery is the endpoint's downstream boundary; the forwarded config is what the test asserts - "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", - new_callable=AsyncMock, - return_value=mock_response, - ) as mock_aquery, patch.object(litellm, "vector_store_registry", mock_registry), patch( # test-quality-ok: seeds the managed-store registry the merge under test reads and grants access so real store resolution runs - "litellm.proxy.vector_store_endpoints.utils.can_user_access_vector_store", - new=AsyncMock(return_value=True), + with ( + patch( # test-quality-ok: aquery is the endpoint's downstream boundary; the forwarded config is what the test asserts + "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", + new_callable=AsyncMock, + return_value=mock_response, + ) as mock_aquery, + patch.object(litellm, "vector_store_registry", mock_registry), + patch( # test-quality-ok: seeds the managed-store registry the merge under test reads and grants access so real store resolution runs + "litellm.proxy.vector_store_endpoints.utils.can_user_access_vector_store", + new=AsyncMock(return_value=True), + ), ): response = client_internal_user.post( "/v1/rag/query", @@ -1124,6 +1137,163 @@ def _multipart_ingest_request(*, filename: str, content: bytes, content_type: st return Request(scope, receive) +async def _never_fetch(url: str): + raise AssertionError(f"no fetch expected, got {url}") + + +@dataclass(frozen=True) +class _FixedScanner: + verdict: ScanVerdict + + def scan(self, content: bytes) -> ScanResult: + return ScanResult(verdict=self.verdict, signature="Custom.Sig") + + +def _fetcher_answering(status_code: int, content: bytes, content_type: str = "text/plain"): + seen: list[str] = [] + + async def fetch_url(url: str) -> httpx.Response: + seen.append(url) + return httpx.Response(status_code, content=content, headers={"content-type": content_type}) + + return fetch_url, seen + + +FILE_URL_BODY = { + "file_url": "https://docs.example.com/reports/q3.txt", + "ingest_options": json.loads(INGEST_REQUEST)["ingest_options"], +} + + +class TestFileUrlUploadControls: + """A ``file_url`` is fetched by the proxy and goes through the same controls as an inline upload.""" + + async def _fetch(self, fetch_url, scanner=None): + from litellm.proxy.rag_endpoints.endpoints import fetch_and_secure_file_url + + return await fetch_and_secure_file_url( + FILE_URL_BODY["file_url"], + scanner=scanner or EicarTestMalwareScanner(), + fetch_url=fetch_url, + ) + + async def _expect_rejection(self, fetch_url, reason: str | None, scanner=None) -> dict: + with pytest.raises(HTTPException) as excinfo: + await self._fetch(fetch_url, scanner=scanner) + assert excinfo.value.status_code == 400 + detail = excinfo.value.detail + assert detail.get("reason") == reason, detail + return detail + + async def test_eicar_behind_a_url_is_blocked(self): + fetch_url, seen = _fetcher_answering(200, EICAR.encode()) + await self._expect_rejection(fetch_url, "malware_detected") + assert seen == [FILE_URL_BODY["file_url"]] + + async def test_shebang_behind_a_url_is_blocked(self): + fetch_url, _seen = _fetcher_answering(200, b"#!/bin/sh\nrm -rf /\n") + await self._expect_rejection(fetch_url, "executable_not_allowed") + + async def test_oversize_url_is_blocked_before_it_is_buffered(self): + async def fetch_url(url: str) -> httpx.Response: + raise HTTPResponseLimitError(f"Response exceeded {MAX_UPLOAD_SIZE_BYTES} bytes") + + detail = await self._expect_rejection(fetch_url, "file_too_large") + assert str(MAX_UPLOAD_SIZE_BYTES) in detail["error"] + + async def test_compressed_url_response_is_refused_without_calling_it_too_large(self): + async def fetch_url(url: str) -> httpx.Response: + raise HTTPResponseEncodingError("Response size limits require an uncompressed response") + + detail = await self._expect_rejection(fetch_url, None) + assert "uncompressed" in detail["error"] + + @pytest.mark.parametrize("status_code", [301, 403, 404, 500]) + async def test_non_2xx_fetch_is_a_client_error_naming_the_status(self, status_code): + fetch_url, _seen = _fetcher_answering(status_code, b"benign document text") + detail = await self._expect_rejection(fetch_url, None) + assert f"HTTP {status_code}" in detail["error"] + + async def test_ssrf_or_transport_failure_is_a_client_error(self): + async def fetch_url(url: str) -> httpx.Response: + raise ValueError("Blocked private address") + + detail = await self._expect_rejection(fetch_url, None) + assert "Blocked private address" in detail["error"] + + async def test_clean_text_behind_a_url_is_handed_on_as_bytes(self): + fetch_url, _seen = _fetcher_answering(200, b"benign document text\n", content_type="text/x-whatever") + server_filename, content_bytes, content_type = await self._fetch(fetch_url) + assert content_bytes == b"benign document text\n" + assert server_filename.endswith(".txt") and "/" not in server_filename + assert content_type == "text/plain" + + async def test_url_bytes_go_through_the_injected_scanner(self): + fetch_url, _seen = _fetcher_answering(200, b"benign document text\n") + await self._expect_rejection(fetch_url, "malware_detected", scanner=_FixedScanner(ScanVerdict.INFECTED)) + + def test_route_prefers_an_inline_file_and_never_fetches_the_url(self, client_internal_user, monkeypatch): + monkeypatch.setattr("litellm.proxy.rag_endpoints.endpoints.fetch_file_url", _never_fetch) + inline = {"filename": "inline.txt", "content": base64.b64encode(b"inline text").decode()} + with ( + patch( # test-quality-ok: aingest is the endpoint's downstream boundary; the test asserts the forwarded bytes + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new_callable=AsyncMock, + return_value={"vector_store_id": "vs_new", "file_id": "file_123"}, + ) as mock_aingest + ): + response = client_internal_user.post("/v1/rag/ingest", json={**FILE_URL_BODY, "file": inline}) + assert response.status_code == 200, response.text + assert mock_aingest.await_args.kwargs["file_data"][1] == b"inline text" + + @pytest.mark.parametrize("field", [{"file_url": ["https://docs.example.com/q3.txt"]}, {"file_id": {"id": 1}}]) + def test_route_rejects_a_non_string_file_reference(self, client_internal_user, monkeypatch, field): + monkeypatch.setattr("litellm.proxy.rag_endpoints.endpoints.fetch_file_url", _never_fetch) + body = {"ingest_options": FILE_URL_BODY["ingest_options"], **field} + response = client_internal_user.post("/v1/rag/ingest", json=body) + assert response.status_code == 400, response.text + assert f"'{next(iter(field))}' must be a string" in response.json()["detail"]["error"] + + def test_route_validates_the_request_before_fetching_the_url(self, client_internal_user, monkeypatch): + monkeypatch.setattr("litellm.proxy.rag_endpoints.endpoints.fetch_file_url", _never_fetch) + vector_store = {"custom_llm_provider": "bedrock", "aws_sts_endpoint": "https://attacker.example/sts"} + body = {**FILE_URL_BODY, "ingest_options": {"vector_store": vector_store}} + response = client_internal_user.post("/v1/rag/ingest", json=body) + assert response.status_code == 400, response.text + assert "aws_sts_endpoint" in response.json()["detail"]["error"] + + def test_route_hands_fetched_bytes_to_ingestion_and_never_the_url(self, client_internal_user, monkeypatch): + fetch_url, seen = _fetcher_answering(200, b"benign document text\n") + monkeypatch.setattr("litellm.proxy.rag_endpoints.endpoints.fetch_file_url", fetch_url) + with ( + patch( # test-quality-ok: aingest is the endpoint's downstream boundary; the test asserts the forwarded bytes + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new_callable=AsyncMock, + return_value={"vector_store_id": "vs_new", "file_id": "file_123"}, + ) as mock_aingest + ): + response = client_internal_user.post("/v1/rag/ingest", json=FILE_URL_BODY) + assert response.status_code == 200, response.text + assert seen == [FILE_URL_BODY["file_url"]] + forwarded = mock_aingest.await_args.kwargs + assert forwarded["file_data"][1] == b"benign document text\n" + assert "file_url" not in forwarded + + def test_route_rejects_eicar_behind_a_url(self, client_internal_user, monkeypatch): + fetch_url, _seen = _fetcher_answering(200, EICAR.encode()) + monkeypatch.setattr("litellm.proxy.rag_endpoints.endpoints.fetch_file_url", fetch_url) + with ( + patch( # test-quality-ok: aingest is the endpoint's downstream boundary; the test asserts it is never reached + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new_callable=AsyncMock, + ) as mock_aingest + ): + response = client_internal_user.post("/v1/rag/ingest", json=FILE_URL_BODY) + assert response.status_code == 400, response.text + assert response.json()["detail"]["reason"] == "malware_detected" + mock_aingest.assert_not_awaited() + + class TestVectorStoreUploadControls: """End-to-end enforcement of pentest M4 upload controls on /v1/rag/ingest.""" @@ -1155,6 +1325,43 @@ class TestVectorStoreUploadControls: assert response.status_code == 400, response.text assert response.json()["detail"]["reason"] == "archive_not_allowed" + def test_route_runs_the_configured_scanner(self, client_internal_user, monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.rag_upload_malware_scanner", _FixedScanner(ScanVerdict.INFECTED) + ) + with ( + patch( # test-quality-ok: aingest is the endpoint's downstream boundary; the test asserts it is never reached + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new_callable=AsyncMock, + ) as mock_aingest + ): + response = client_internal_user.post( + "/v1/rag/ingest", + files={"file": ("doc.txt", io.BytesIO(b"benign document text"), "text/plain")}, + data={"request": INGEST_REQUEST}, + ) + assert response.status_code == 400, response.text + assert response.json()["detail"]["reason"] == "malware_detected" + assert "Custom.Sig" in response.json()["detail"]["error"] + mock_aingest.assert_not_awaited() + + def test_route_lets_a_clean_verdict_from_the_configured_scanner_through(self, client_internal_user, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.rag_upload_malware_scanner", _FixedScanner(ScanVerdict.CLEAN)) + with ( + patch( # test-quality-ok: aingest is the endpoint's downstream boundary; the test asserts the forwarded bytes + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new_callable=AsyncMock, + return_value={"vector_store_id": "vs_new", "file_id": "file_123"}, + ) as mock_aingest + ): + response = client_internal_user.post( + "/v1/rag/ingest", + files={"file": ("clean_name.txt", io.BytesIO(EICAR.encode()), "text/plain")}, + data={"request": INGEST_REQUEST}, + ) + assert response.status_code == 200, response.text + assert mock_aingest.await_args.kwargs["file_data"][1] == EICAR.encode() + async def test_clean_text_upload_gets_server_generated_filename(self): from litellm.proxy.rag_endpoints.endpoints import parse_rag_ingest_request from litellm.proxy.rag_endpoints.upload_security import EicarTestMalwareScanner @@ -1164,9 +1371,7 @@ class TestVectorStoreUploadControls: content=b"benign document text\n", content_type="text/plain", ) - _options, file_data, _url, _file_id = await parse_rag_ingest_request( - request, scanner=EicarTestMalwareScanner() - ) + _options, file_data, _url, _file_id = await parse_rag_ingest_request(request, scanner=EicarTestMalwareScanner()) assert file_data is not None server_filename, content_bytes, secured_content_type = file_data assert server_filename != "../../etc/passwd" diff --git a/tests/test_litellm/proxy/rag_endpoints/test_upload_security.py b/tests/test_litellm/proxy/rag_endpoints/test_upload_security.py index 84375080904..1347fea8c1f 100644 --- a/tests/test_litellm/proxy/rag_endpoints/test_upload_security.py +++ b/tests/test_litellm/proxy/rag_endpoints/test_upload_security.py @@ -7,6 +7,7 @@ dependency-injected malware scanner validated with the EICAR test file. """ from dataclasses import dataclass +from types import SimpleNamespace import pytest @@ -14,6 +15,7 @@ from litellm.proxy.rag_endpoints.upload_security import ( EICAR_TEST_SIGNATURE, DetectedFormat, EicarTestMalwareScanner, + MalwareScannerConfigError, RejectedUpload, RejectionReason, ScanResult, @@ -21,6 +23,7 @@ from litellm.proxy.rag_endpoints.upload_security import ( SecuredUpload, generate_safe_filename, inspect_content, + resolve_malware_scanner, safe_download_headers, validate_upload, ) @@ -174,3 +177,96 @@ def test_safe_download_headers_sanitize_injection(hostile): disposition = safe_download_headers(hostile)["Content-Disposition"] assert "\r" not in disposition and "\n" not in disposition assert disposition.count('"') == 2 + + +def _loader_returning(loaded: object): + calls: list[tuple[str, str | None]] = [] + + def load_instance(value: str, config_file_path: str | None) -> object: + calls.append((value, config_file_path)) + return loaded + + return load_instance, calls + + +def _loader_raising(error: Exception): + def load_instance(value: str, config_file_path: str | None) -> object: + raise error + + return load_instance + + +def test_resolve_malware_scanner_unset_runs_the_eicar_test_scanner(): + load_instance, calls = _loader_returning(_CLEAN_SCANNER) + for rag_ingest in (None, {}, {"malware_scanner": None}): + resolved = resolve_malware_scanner( + rag_ingest, config_file_path="/etc/litellm/config.yaml", load_instance=load_instance + ) + assert isinstance(resolved, EicarTestMalwareScanner), rag_ingest + assert calls == [] + + +def test_resolve_malware_scanner_loads_the_configured_instance_next_to_the_config(): + load_instance, calls = _loader_returning(_INFECTED_SCANNER) + resolved = resolve_malware_scanner( + {"malware_scanner": "custom_scanner.scanner"}, + config_file_path="/etc/litellm/config.yaml", + load_instance=load_instance, + ) + assert resolved is _INFECTED_SCANNER + assert calls == [("custom_scanner.scanner", "/etc/litellm/config.yaml")] + assert validate_upload(content=_TEXT_BYTES, scanner=resolved).reason is RejectionReason.MALWARE_DETECTED + + +@pytest.mark.parametrize( + "loaded, expected_fragment", + [ + (object(), "no scan(content: bytes) -> ScanResult method"), + ("custom_scanner.scanner", "no scan(content: bytes) -> ScanResult method"), + (SimpleNamespace(scan="not callable"), "no scan(content: bytes) -> ScanResult method"), + (_StubScanner, "names the class _StubScanner"), + ], +) +def test_resolve_malware_scanner_rejects_an_object_that_is_not_a_scanner(loaded, expected_fragment): + load_instance, _calls = _loader_returning(loaded) + resolved = resolve_malware_scanner( + {"malware_scanner": "custom_scanner.scanner"}, + config_file_path=None, + load_instance=load_instance, + ) + assert isinstance(resolved, MalwareScannerConfigError) + assert "general_settings.rag_ingest.malware_scanner" in resolved.message + assert "custom_scanner.scanner" in resolved.message + assert expected_fragment in resolved.message + + +@pytest.mark.parametrize( + "error", + [ + ImportError("Could not import scanner from custom_scanner"), + AttributeError("module has no attribute scanner"), + ValueError("Empty module name"), + RuntimeError("clamd socket is missing"), + ], +) +def test_resolve_malware_scanner_reports_a_failed_load_with_the_option_name(error): + resolved = resolve_malware_scanner( + {"malware_scanner": "custom_scanner.scanner"}, + config_file_path=None, + load_instance=_loader_raising(error), + ) + assert isinstance(resolved, MalwareScannerConfigError) + assert "general_settings.rag_ingest.malware_scanner" in resolved.message + assert str(error) in resolved.message + + +@pytest.mark.parametrize( + "rag_ingest", + [{"malware_scaner": "custom_scanner.scanner"}, {"malware_scanner": 42}, "custom_scanner.scanner", ["x"]], +) +def test_resolve_malware_scanner_rejects_an_invalid_rag_ingest_block(rag_ingest): + load_instance, calls = _loader_returning(_CLEAN_SCANNER) + resolved = resolve_malware_scanner(rag_ingest, config_file_path=None, load_instance=load_instance) + assert isinstance(resolved, MalwareScannerConfigError) + assert "general_settings.rag_ingest" in resolved.message + assert calls == [] diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 884a9c81500..ecf19e362e4 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3518,6 +3518,41 @@ async def test_load_config_role_permissions_usable_by_jwt_auth(tmp_path): assert get_role_based_models(rbac_role="team", general_settings=settings) is None +@pytest.mark.asyncio +async def test_load_config_wires_the_configured_malware_scanner_into_uploads(tmp_path, monkeypatch): + from litellm.proxy.proxy_server import ProxyConfig + from litellm.proxy.rag_endpoints.upload_security import EicarTestMalwareScanner, ScanVerdict + + monkeypatch.setattr(proxy_server_module, "rag_upload_malware_scanner", EicarTestMalwareScanner()) + (tmp_path / "boot_scanner.py").write_text( + "from litellm.proxy.rag_endpoints.upload_security import ScanResult, ScanVerdict\n" + "class FlagEverything:\n" + " def scan(self, content):\n" + " return ScanResult(verdict=ScanVerdict.INFECTED, signature='Boot.Test')\n" + "scanner = FlagEverything()\n" + "not_a_scanner = object()\n" + ) + config_file: Final = tmp_path / "config.yaml" + + config_file.write_text( + yaml.dump({"model_list": [], "general_settings": {"rag_ingest": {"malware_scanner": "boot_scanner.scanner"}}}) + ) + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + assert proxy_server_module.rag_upload_malware_scanner.scan(b"plain text").verdict is ScanVerdict.INFECTED + + config_file.write_text(yaml.dump({"model_list": []})) + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + assert isinstance(proxy_server_module.rag_upload_malware_scanner, EicarTestMalwareScanner) + + config_file.write_text( + yaml.dump( + {"model_list": [], "general_settings": {"rag_ingest": {"malware_scanner": "boot_scanner.not_a_scanner"}}} + ) + ) + with pytest.raises(ValueError, match=re.escape("general_settings.rag_ingest.malware_scanner")): + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + + @pytest.mark.asyncio async def test_load_config_without_role_permissions_leaves_every_role_unrestricted(tmp_path): from litellm.proxy.auth.auth_checks import get_role_based_models, get_role_based_routes diff --git a/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py b/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py index 0072997a0d9..9f5ed39f232 100644 --- a/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py +++ b/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py @@ -182,6 +182,21 @@ def test_litellm_jwtauth_custom_validate_stripped(): assert cleaned["litellm_jwtauth"]["user_id_jwt_field"] == "sub" +@pytest.mark.parametrize("remote", ["s3://attacker/m.scanner", "gcs://attacker/m.scanner"]) +def test_rag_ingest_malware_scanner_stripped(remote): + overlay = {"rag_ingest": {"malware_scanner": remote}, "custom_auth": "my_auth.handler"} + cleaned = _scrub_db_overlay_remote_module_loads("general_settings", overlay) + assert cleaned["rag_ingest"]["malware_scanner"] is None + assert cleaned["custom_auth"] == "my_auth.handler" + assert overlay["rag_ingest"]["malware_scanner"] == remote + + +def test_rag_ingest_local_malware_scanner_preserved(): + overlay = {"rag_ingest": {"malware_scanner": "custom_scanner.scanner"}} + cleaned = _scrub_db_overlay_remote_module_loads("general_settings", overlay) + assert cleaned["rag_ingest"]["malware_scanner"] == "custom_scanner.scanner" + + def test_local_dotted_name_preserved(): # The scrub only targets s3:// / gcs:// scheme prefixes — legitimate # dotted module names (the documented operator flow) must pass diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1e0aed46923..f63e4657a97 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -28451,6 +28451,8 @@ export interface components { * @default 30 */ proxy_config_reload_interval_seconds: number; + /** @description controls run on every upload that lands in a vector store (/v1/rag/ingest and /v1/files with purpose assistants or user_data); rag_ingest.malware_scanner names the scanner instance as . the way custom_auth does, and unset runs the EICAR test scanner - https://docs.litellm.ai/docs/rag_ingest */ + rag_ingest?: components["schemas"]["RagIngestSettings"] | null; /** * Reject Clientside Metadata Tags * @description When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata. @@ -39051,6 +39053,14 @@ export interface components { }; } | null; }; + /** RagIngestSettings */ + RagIngestSettings: { + /** + * Malware Scanner + * @description The scanner every vector-store upload goes through, as . where the module sits next to config.yaml or is importable, resolved the way custom_auth is. The instance must expose scan(content: bytes) -> ScanResult and be safe to call from several threads at once. Unset runs the EICAR test scanner, which flags only the EICAR test file + */ + malware_scanner?: string | null; + }; /** * RankingOptions * @description Ranking options for search.