feat(rag): configurable malware scanner and upload controls on every vector-store upload route

This commit is contained in:
mateo-berri 2026-09-25 16:50:45 -07:00
parent e065a2575b
commit cc26342a3a
13 changed files with 777 additions and 100 deletions

View file

@ -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:

View file

@ -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 <module>.<instance> "
"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",

View file

@ -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,

View file

@ -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,
)

View file

@ -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,

View file

@ -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

View file

@ -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 <module>.<instance> 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"
),
)

View file

@ -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

View file

@ -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"

View file

@ -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 == []

View file

@ -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

View file

@ -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

View file

@ -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 <module>.<instance> 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 <module>.<instance> 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.