From 14c5ae165059a2ca058d43f771b6bd127d079d38 Mon Sep 17 00:00:00 2001 From: Michael Ripa Date: Tue, 21 Apr 2026 14:02:45 -0400 Subject: [PATCH] dump requests to /tmp if content length exceeds threshold + throw error if request too large --- src/services/api/src/app.py | 18 ++++++++++-- src/services/api/src/config.py | 18 ++++++++++++ src/services/api/src/dependencies.py | 27 +++++++++++++++++- src/services/api/src/queue/dispatcher.py | 35 ++++++++++++++++++++++-- 4 files changed, 91 insertions(+), 7 deletions(-) diff --git a/src/services/api/src/app.py b/src/services/api/src/app.py index 2479f0e6..11e5b220 100755 --- a/src/services/api/src/app.py +++ b/src/services/api/src/app.py @@ -1,5 +1,7 @@ import asyncio +import os import pickle +import tempfile import traceback import uuid from typing import Any, Dict, Optional @@ -117,9 +119,19 @@ async def request( # Inject trace context before pickling for cross-process propagation backend_request.trace_context = TracingContext.inject() - await RedisProvider.async_client.lpush( - "queue", pickle.dumps(backend_request) - ) + # Use temp file for large requests, direct Redis for small ones + if backend_request.content_length > AppConfig.queue_temp_file_threshold: + # Dump request to tmp file, pass filename reference to redis + filename = f"request-{backend_request.id}.pkl" + tmp_file = os.path.join(tempfile.gettempdir(), filename) + with open(tmp_file, "wb") as f: + pickle.dump(backend_request, f) + await RedisProvider.async_client.lpush("queue", f"file:{filename}") + else: + # Small request - pass pickled data directly to redis + await RedisProvider.async_client.lpush( + "queue", pickle.dumps(backend_request) + ) span.add_event("request_queued") diff --git a/src/services/api/src/config.py b/src/services/api/src/config.py index f7b33520..5378880e 100644 --- a/src/services/api/src/config.py +++ b/src/services/api/src/config.py @@ -58,6 +58,8 @@ class AppConfig: min_nnsight_version_parsed: Version min_python_version_parsed: Version dev_mode: bool + queue_temp_file_threshold: int + max_request_size: int @classmethod def from_env(cls) -> None: @@ -94,6 +96,20 @@ class attributes. Called automatically on module import. cls.min_python_version_parsed = Version(cls.min_python_version) cls.dev_mode = os.environ.get("NDIF_DEV_MODE", "false").lower() == "true" + # Queue mode: requests larger than this threshold use temp files instead of direct Redis + # Default: 16 MB + threshold_str = os.environ.get("NDIF_QUEUE_TEMP_FILE_THRESHOLD", str(16 * 1024 * 1024)) + cls.queue_temp_file_threshold = cls._parse_positive_int( + threshold_str, "NDIF_QUEUE_TEMP_FILE_THRESHOLD" + ) + + # Maximum allowed request size. Requests larger than this are rejected. + # Default: 1 GB + max_size_str = os.environ.get("NDIF_MAX_REQUEST_SIZE", str(1024 * 1024 * 1024)) + cls.max_request_size = cls._parse_positive_int( + max_size_str, "NDIF_MAX_REQUEST_SIZE" + ) + @classmethod def to_env(cls) -> dict[str, object]: """Export configuration as a dictionary suitable for environment variables. @@ -109,6 +125,8 @@ def to_env(cls) -> dict[str, object]: "MIN_NNSIGHT_VERSION": cls.min_nnsight_version, "MIN_PYTHON_VERSION": cls.min_python_version, "NDIF_DEV_MODE": cls.dev_mode, + "NDIF_QUEUE_TEMP_FILE_THRESHOLD": cls.queue_temp_file_threshold, + "NDIF_MAX_REQUEST_SIZE": cls.max_request_size, } @classmethod diff --git a/src/services/api/src/dependencies.py b/src/services/api/src/dependencies.py index 7eb290ac..c4bac4ec 100755 --- a/src/services/api/src/dependencies.py +++ b/src/services/api/src/dependencies.py @@ -5,6 +5,7 @@ from starlette.status import ( HTTP_400_BAD_REQUEST, HTTP_401_UNAUTHORIZED, + HTTP_413_REQUEST_ENTITY_TOO_LARGE, HTTP_503_SERVICE_UNAVAILABLE, ) @@ -162,6 +163,27 @@ async def require_ray_connection() -> None: ) +def validate_content_length(content_length: int) -> int: + """Validate that the request content length is within allowed limits. + + Args: + content_length: The content length in bytes. + + Returns: + The validated content length in bytes. + + Raises: + HTTPException: 413 if the content length exceeds the maximum allowed size. + """ + if content_length > AppConfig.max_request_size: + raise HTTPException( + status_code=HTTP_413_REQUEST_ENTITY_TOO_LARGE, + detail=f"Request too large: {content_length} bytes exceeds maximum of {AppConfig.max_request_size} bytes", + ) + + return content_length + + async def validate_request(raw_request: Request) -> BackendRequestModel: """FastAPI dependency to validate and create a BackendRequestModel. @@ -186,11 +208,14 @@ async def validate_request(raw_request: Request) -> BackendRequestModel: api_key = raw_request.headers.get("ndif-api-key", "") nnsight_version = raw_request.headers.get("nnsight-version", "") python_version = raw_request.headers.get("python-version", "") + content_length = int(raw_request.headers.get("content-length", 0)) span.set_attribute("ndif.client.nnsight_version", nnsight_version) span.set_attribute("ndif.client.python_version", python_version) + span.set_attribute("ndif.content_length", content_length) - # # Validate using existing dependency functions (call them directly, not as dependencies) + # Validate using existing dependency functions (call them directly, not as dependencies) + validate_content_length(content_length) await authenticate_api_key(api_key) await validate_nnsight_version(nnsight_version) await validate_python_version(python_version) diff --git a/src/services/api/src/queue/dispatcher.py b/src/services/api/src/queue/dispatcher.py index b2fddff7..cdd9aba5 100755 --- a/src/services/api/src/queue/dispatcher.py +++ b/src/services/api/src/queue/dispatcher.py @@ -28,7 +28,7 @@ import time import traceback from typing import Optional - +import tempfile from enum import Enum from ..logging import set_logger @@ -164,6 +164,35 @@ def connect(self) -> None: RedisProvider.sync_client.set("ray:connected", "1") self.logger.info(f"Connected to Ray") + def _load_request(self, data: bytes) -> BackendRequestModel: + """Load a request from either a temp file reference or direct pickle data. + + Supports two modes: + - Temp file mode: data is "file:" string, load from disk + - Direct mode: data is raw pickled BackendRequestModel + + Args: + data: Raw bytes from Redis queue. + + Returns: + The deserialized BackendRequestModel. + """ + try: + decoded = data.decode("utf-8") + if decoded.startswith("file:"): + # Temp file mode - extract filename and load from disk + filename = decoded[5:] # Remove "file:" prefix + tmp_file = os.path.join(tempfile.gettempdir(), filename) + with open(tmp_file, "rb") as f: + request = pickle.load(f) + os.remove(tmp_file) + return request + except (UnicodeDecodeError, ValueError): + pass + + # Direct pickle mode (old way) + return pickle.loads(data) + async def get(self) -> list[BackendRequestModel]: """Fetch pending requests from the Redis queue. @@ -181,13 +210,13 @@ async def get(self) -> list[BackendRequestModel]: if result is None: return [] - requests = [pickle.loads(result[1])] + requests = [self._load_request(result[1])] while len(requests) < 32: item = await RedisProvider.async_client.rpop("queue") if item is None: break - requests.append(pickle.loads(item)) + requests.append(self._load_request(item)) return requests