([^<]+)")
+
+
+@dataclass(frozen=True, slots=True)
+class _PreparedPut:
+ json_string: str
+ headers: Mapping[str, str]
+
+
+def _s3_error_code(response: httpx.Response) -> str | None:
+ text: Final = response.text
+ match: Final = _S3_ERROR_CODE.search(text) if isinstance(text, str) else None
+ return match.group(1) if match else None
+
+
+def _is_terminal(response: httpx.Response) -> bool:
+ """True only for object-specific, unrecoverable rejections (400/403 with a terminal XML code).
+ Unknown codes, empty or non-XML bodies, and every other status fail safe toward retry."""
+ return response.status_code in _BODY_CODED_STATUSES and _s3_error_code(response) in _TERMINAL_ERROR_CODES
+
+
def _s3_key_parent(s3_object_key: str) -> str:
return s3_object_key.rsplit("/", 1)[0] if "/" in s3_object_key else ""
@@ -58,11 +106,19 @@ class S3BatchUploadError(Exception):
def __init__(self, failed: int, total: int) -> None:
self.failed = failed
self.total = total
- super().__init__(f"{failed} of {total} S3 uploads failed; events kept in queue for the next flush")
+ super().__init__(f"{failed} of {total} S3 uploads failed; transient failures kept in queue for the next flush")
+
+
+_in_flush: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar("s3_v2_in_flush", default=False)
class S3Logger(CustomBatchLogger, BaseAWSLLM):
preserve_events_added_during_flush = True
+ _flush_retries: int = 0
+ _requeued_count: int = 0
+ _upload_limiter: asyncio.Semaphore | AdaptiveConcurrencyLimiter | None = None
+ s3_drop_on_terminal_error: bool = True
+ s3_max_retry_age_seconds: int | None = 3600
def __init__(
self,
@@ -92,6 +148,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
+ s3_max_queue_size: int | None = None,
+ s3_max_retry_age_seconds: int | None = 3600,
+ s3_drop_on_terminal_error: bool = True,
+ s3_adaptive_concurrency: bool = False,
+ s3_max_adaptive_concurrency: int | None = None,
s3_batch_file_upload: bool = False,
s3_callback_params_override: dict | None = None,
**kwargs,
@@ -135,9 +196,22 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_sse_kms_key_id=s3_sse_kms_key_id,
s3_log_prompts_only=s3_log_prompts_only,
s3_max_concurrent_uploads=s3_max_concurrent_uploads,
+ s3_max_queue_size=s3_max_queue_size,
+ s3_max_retry_age_seconds=s3_max_retry_age_seconds,
+ s3_drop_on_terminal_error=s3_drop_on_terminal_error,
+ s3_adaptive_concurrency=s3_adaptive_concurrency,
+ s3_max_adaptive_concurrency=s3_max_adaptive_concurrency,
s3_batch_file_upload=s3_batch_file_upload,
)
- self._upload_semaphore = asyncio.Semaphore(self.s3_max_concurrent_uploads)
+ self._upload_limiter = (
+ AdaptiveConcurrencyLimiter(
+ initial=self.s3_max_concurrent_uploads,
+ floor=self.s3_max_concurrent_uploads,
+ ceiling=max(self.s3_max_concurrent_uploads, self.s3_max_adaptive_concurrency),
+ )
+ if self.s3_adaptive_concurrency
+ else asyncio.Semaphore(self.s3_max_concurrent_uploads)
+ )
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
# IMPORTANT
@@ -158,8 +232,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
flush_lock=self.flush_lock,
flush_interval=s3_flush_interval,
batch_size=s3_batch_size,
+ max_queue_size=self.s3_max_queue_size,
)
self.log_queue: list[s3BatchLoggingElement] = []
+ self._requeued_count = 0
+ self._flush_retries = 0
+ self._flush_dropped: dict[int, s3BatchLoggingElement] = {}
# Call BaseAWSLLM's __init__
BaseAWSLLM.__init__(self)
@@ -194,6 +272,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
+ s3_max_queue_size: int | None = None,
+ s3_max_retry_age_seconds: int | None = 3600,
+ s3_drop_on_terminal_error: bool = True,
+ s3_adaptive_concurrency: bool = False,
+ s3_max_adaptive_concurrency: int | None = None,
s3_batch_file_upload: bool = False,
params_source: dict | None = None,
):
@@ -259,6 +342,37 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
)
+ configured_queue_size: Final = params.get("s3_max_queue_size")
+ constructor_queue_size: Final = resolve_s3_max_queue_size(
+ s3_max_queue_size, CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE
+ )
+ self.s3_max_queue_size = resolve_s3_max_queue_size(configured_queue_size, constructor_queue_size)
+
+ configured_retry_age: Final = params.get("s3_max_retry_age_seconds")
+ constructor_retry_age: Final = resolve_s3_max_retry_age_seconds(s3_max_retry_age_seconds, 3600)
+ self.s3_max_retry_age_seconds = (
+ constructor_retry_age
+ if configured_retry_age is None or configured_retry_age == ""
+ else resolve_s3_max_retry_age_seconds(configured_retry_age, constructor_retry_age)
+ )
+
+ configured_drop: Final = params.get("s3_drop_on_terminal_error")
+ self.s3_drop_on_terminal_error = resolve_s3_drop_on_terminal_error(
+ configured_drop if configured_drop is not None else s3_drop_on_terminal_error
+ )
+
+ self.s3_adaptive_concurrency = s3_adaptive_concurrency or resolve_s3_adaptive_concurrency(
+ params.get("s3_adaptive_concurrency")
+ )
+
+ configured_adaptive_ceiling: Final = params.get("s3_max_adaptive_concurrency")
+ self.s3_max_adaptive_concurrency = resolve_s3_max_adaptive_concurrency(
+ s3_max_adaptive_concurrency
+ if configured_adaptive_ceiling is None or configured_adaptive_ceiling == ""
+ else configured_adaptive_ceiling,
+ DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY,
+ )
+
self.s3_batch_file_upload = s3_batch_file_upload or resolve_s3_batch_file_upload(
params.get("s3_batch_file_upload")
)
@@ -310,6 +424,35 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
}
return {key: value for key, value in candidates.items() if value}
+ def _prepare_put(self, batch_logging_element: s3BatchLoggingElement) -> _PreparedPut:
+ try:
+ import base64
+ import hashlib
+ except ImportError:
+ raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
+
+ json_string: Final = (
+ batch_logging_element.body
+ if batch_logging_element.body is not None
+ else safe_dumps(batch_logging_element.payload)
+ )
+ content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
+ content_md5: Final = base64.b64encode(
+ hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
+ ).decode()
+ return _PreparedPut(
+ json_string=json_string,
+ headers={
+ "Content-Type": batch_logging_element.content_type,
+ "Content-MD5": content_md5,
+ "x-amz-content-sha256": content_hash,
+ "Content-Language": "en",
+ "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
+ "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
+ **self._sse_headers(),
+ },
+ )
+
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
await self._async_log_event_base(
kwargs=kwargs,
@@ -384,12 +527,21 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
verbose_logger.exception("s3 Layer Error - %s", e)
self.handle_callback_failure(callback_name="S3Logger")
- async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
- try:
- import base64
- import hashlib
- except ImportError:
- raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
+ @property
+ def _upload_semaphore(self) -> asyncio.Semaphore | AdaptiveConcurrencyLimiter:
+ limiter: Final = self._upload_limiter
+ if limiter is None:
+ raise AttributeError("_upload_semaphore")
+ return limiter
+
+ @_upload_semaphore.setter
+ def _upload_semaphore(self, value: asyncio.Semaphore | AdaptiveConcurrencyLimiter) -> None:
+ self._upload_limiter = value
+
+ async def async_upload_data_to_s3(
+ self,
+ batch_logging_element: s3BatchLoggingElement,
+ ) -> bool:
try:
from litellm.litellm_core_utils.asyncify import asyncify
@@ -400,31 +552,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
- # Convert JSON to string
- json_string: Final = (
- batch_logging_element.body
- if batch_logging_element.body is not None
- else safe_dumps(batch_logging_element.payload)
- )
-
- # Calculate SHA256 hash of the content
- content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
- content_md5: Final = base64.b64encode(
- hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
- ).decode()
-
- # Prepare the request
- headers: Final = {
- "Content-Type": batch_logging_element.content_type,
- "Content-MD5": content_md5,
- "x-amz-content-sha256": content_hash,
- "Content-Language": "en",
- "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
- "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
- **self._sse_headers(),
- }
-
- async def signed_put() -> httpx.Response:
+ async def signed_put(prepared: _PreparedPut) -> httpx.Response:
credentials: Final = await asyncified_get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
@@ -436,18 +564,26 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
aws_web_identity_token=self.s3_aws_web_identity_token,
aws_sts_endpoint=self.s3_aws_sts_endpoint,
)
- signed_headers: Final = await run_aws_signing(self._sign_put, credentials, url, json_string, headers)
+ signed_headers: Final = await run_aws_signing(
+ self._sign_put, credentials, url, prepared.json_string, prepared.headers
+ )
try:
- return await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
+ return await self.async_httpx_client.put(url, data=prepared.json_string, headers=signed_headers)
except httpx.HTTPStatusError as error:
return error.response
max_retries: Final = 3
+ prepared: Final = self._prepare_put(batch_logging_element)
for attempt in range(max_retries):
- response = await signed_put()
- if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
+ response = await self._recorded_put(partial(signed_put, prepared))
+ if (
+ response.status_code in _RETRYABLE_STATUSES
+ and not (self.s3_drop_on_terminal_error and _is_terminal(response))
+ and attempt < max_retries - 1
+ ):
wait_time = 2**attempt # 1s, 2s
- verbose_logger.warning(
+ verbose_logger.log(
+ logging.DEBUG if _in_flush.get() else logging.WARNING,
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
response.status_code,
wait_time,
@@ -455,6 +591,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
max_retries,
batch_logging_element.s3_object_key,
)
+ self._flush_retries += 1
await asyncio.sleep(wait_time)
continue
response.raise_for_status()
@@ -462,6 +599,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
except Exception as e:
verbose_logger.exception("Error uploading to s3: %s", e)
self.handle_callback_failure(callback_name="S3Logger")
+ if isinstance(e, httpx.HTTPStatusError) and self.s3_drop_on_terminal_error and _is_terminal(e.response):
+ verbose_logger.warning(
+ "s3 logging: dropping object %s after terminal status %s",
+ batch_logging_element.s3_object_key,
+ e.response.status_code,
+ )
+ self._flush_dropped[id(batch_logging_element)] = batch_logging_element
return False
return True
@@ -483,11 +627,64 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# see custom_batch_logger.py which triggers the flush
#########################################################
uploads: Final = self._batch_file_elements(batch) if self._batch_file_mode_active() else batch
- results: Final = await asyncio.gather(*(self._upload_bounded(element) for element in uploads))
- failed: Final = tuple(element for element, ok in zip(uploads, results, strict=True) if not ok)
- if not failed:
+ self._flush_retries = 0
+ self._flush_dropped = {} # mutable-ok: per-flush drop marks read back by _upload_bounded
+ stale: Final = min(self._requeued_count, len(uploads)) if len(uploads) == len(batch) else 0
+ order: Final = (*range(stale, len(uploads)), *range(stale))
+ ordered: Final = await asyncio.gather(*(self._upload_outcome(uploads[i]) for i in order))
+ outcomes: Final = dict(zip(order, ordered, strict=True))
+ results: Final = tuple(outcomes[i] for i in range(len(uploads)))
+ if self._flush_retries:
+ verbose_logger.warning(
+ "s3 logging: %s in-call retries across %s uploads this flush",
+ self._flush_retries,
+ len(uploads),
+ )
+ delivered: Final = sum(1 for outcome in results if outcome == "delivered")
+ bucket_wide: Final = delivered == 0
+ failed: Final = tuple(
+ (element, outcome) for element, outcome in zip(uploads, results, strict=True) if outcome != "delivered"
+ )
+ now: Final = time.monotonic()
+ requeued: Final = (
+ tuple(element for element, _ in failed)
+ if bucket_wide
+ else tuple(
+ element
+ if element.retrying_since is not None or self.s3_max_retry_age_seconds is None
+ else element.model_copy(update={"retrying_since": now})
+ for element, outcome in failed
+ if outcome != "dropped"
+ and not (
+ self.s3_max_retry_age_seconds is not None
+ and element.retrying_since is not None
+ and now - element.retrying_since > self.s3_max_retry_age_seconds
+ )
+ )
+ )
+ dropped: Final = len(failed) - len(requeued)
+ if dropped:
+ verbose_logger.warning(
+ "s3 logging: %s uploads dropped (terminal or retrying longer than s3_max_retry_age_seconds=%s)",
+ dropped,
+ self.s3_max_retry_age_seconds,
+ )
+ if not requeued:
+ self._requeued_count = 0
return
- self.log_queue = [*failed, *self.log_queue[len(batch) :]]
+ arrivals: Final = self.log_queue[len(batch) :]
+ overflow: Final = max(0, len(requeued) + len(arrivals) - self.max_queue_size)
+ if overflow:
+ verbose_logger.warning(
+ "s3 logging: queue exceeded max_queue_size=%s after a failed flush, dropped %s oldest events",
+ self.max_queue_size,
+ overflow,
+ )
+ self.log_queue = [ # mutable-ok: log_queue is the flush buffer shared with custom_batch_logger
+ *requeued,
+ *arrivals,
+ ][overflow:]
+ self._requeued_count = max(0, len(requeued) - overflow)
raise S3BatchUploadError(failed=len(failed), total=len(uploads))
def _batch_file_mode_active(self) -> bool:
@@ -502,8 +699,37 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
return True
async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool:
- async with self._upload_semaphore:
- return await self.async_upload_data_to_s3(element)
+ token: Final = _in_flush.set(True)
+ try:
+ async with self._upload_semaphore:
+ return await self.async_upload_data_to_s3(element)
+ finally:
+ _in_flush.reset(token)
+
+ async def _upload_outcome(self, element: s3BatchLoggingElement) -> UploadOutcome:
+ delivered: Final = await self._upload_bounded(element)
+ if delivered:
+ return "delivered"
+ if id(element) in self._flush_dropped:
+ return "dropped"
+ return "retry"
+
+ async def _recorded_put(self, signed_put: Callable[[], Awaitable[httpx.Response]]) -> httpx.Response:
+ limiter: Final = self._upload_limiter
+ adaptive: Final = limiter if isinstance(limiter, AdaptiveConcurrencyLimiter) else None
+ try:
+ response: Final = await signed_put()
+ except Exception:
+ if adaptive is not None:
+ adaptive.record(PutSample(throttled=True))
+ raise
+ if adaptive is not None:
+ adaptive.record(
+ PutSample(
+ throttled=response.status_code in (429, 503) or _s3_error_code(response) == "SlowDown",
+ )
+ )
+ return response
def _batch_file_elements(self, batch: tuple[s3BatchLoggingElement, ...]) -> tuple[s3BatchLoggingElement, ...]:
now: Final = datetime.now(timezone.utc)
@@ -527,6 +753,9 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
content_type="application/x-ndjson",
s3_object_key=f"{parent}/{batch_name}.jsonl" if parent else f"{batch_name}.jsonl",
s3_object_download_filename=f"{batch_name}.jsonl",
+ retrying_since=min(
+ (element.retrying_since for element in elements if element.retrying_since is not None), default=None
+ ),
)
def create_s3_batch_logging_element(
@@ -596,58 +825,35 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
)
def upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
- try:
- import base64
- import hashlib
- except ImportError:
- raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
try:
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
- # Convert JSON to string
- json_string: Final = (
- batch_logging_element.body
- if batch_logging_element.body is not None
- else safe_dumps(batch_logging_element.payload)
- )
-
- # Calculate SHA256 hash of the content
- content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
- content_md5: Final = base64.b64encode(
- hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
- ).decode()
-
- # Prepare the request
- headers: Final = {
- "Content-Type": batch_logging_element.content_type,
- "Content-MD5": content_md5,
- "x-amz-content-sha256": content_hash,
- "Content-Language": "en",
- "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
- "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
- **self._sse_headers(),
- }
+ prepared: Final = self._prepare_put(batch_logging_element)
httpx_client: Final = _get_httpx_client(
params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None)
)
- def signed_put() -> httpx.Response:
+ def signed_put(prepared_put: _PreparedPut) -> httpx.Response:
credentials: Final = self.get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
aws_session_token=self.s3_aws_session_token,
aws_region_name=self.s3_region_name,
)
- signed_headers: Final = self._sign_put(credentials, url, json_string, headers)
- return httpx_client.put(url, data=json_string, headers=signed_headers)
+ signed_headers: Final = self._sign_put(credentials, url, prepared_put.json_string, prepared_put.headers)
+ return httpx_client.put(url, data=prepared_put.json_string, headers=signed_headers)
max_retries: Final = 3
for attempt in range(max_retries):
- response = signed_put()
- if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
+ response = signed_put(prepared)
+ if (
+ response.status_code in _RETRYABLE_STATUSES
+ and not (self.s3_drop_on_terminal_error and _is_terminal(response))
+ and attempt < max_retries - 1
+ ):
wait_time = 2**attempt # 1s, 2s
verbose_logger.warning(
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
@@ -664,6 +870,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
except Exception as e:
verbose_logger.exception("Error uploading to s3: %s", e)
self.handle_callback_failure(callback_name="S3Logger")
+ if isinstance(e, httpx.HTTPStatusError) and self.s3_drop_on_terminal_error and _is_terminal(e.response):
+ verbose_logger.warning(
+ "s3 logging: dropping object %s after terminal status %s",
+ batch_logging_element.s3_object_key,
+ e.response.status_code,
+ )
async def _download_object_from_s3(self, s3_object_key: str) -> dict | None:
"""
diff --git a/litellm/types/integrations/s3_v2.py b/litellm/types/integrations/s3_v2.py
index 555b16dc141..3b0dad97e8c 100644
--- a/litellm/types/integrations/s3_v2.py
+++ b/litellm/types/integrations/s3_v2.py
@@ -11,3 +11,4 @@ class s3BatchLoggingElement(BaseModel):
s3_object_download_filename: str
body: str | None = None
content_type: str = "application/json"
+ retrying_since: float | None = None
diff --git a/tests/documentation_tests/test_env_keys.py b/tests/documentation_tests/test_env_keys.py
index 3652378503e..0713abc86d0 100644
--- a/tests/documentation_tests/test_env_keys.py
+++ b/tests/documentation_tests/test_env_keys.py
@@ -38,6 +38,7 @@ EXCLUDED_ROLLOUT_FLAGS = {
EXCLUDED_INTERNAL_TUNING_VARS = {
"ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS",
"ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE",
+ "DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY",
}
EXCLUDED_TERMINAL_VARS = {
diff --git a/tests/integration/observability/_s3_v2_support.py b/tests/integration/observability/_s3_v2_support.py
index 104c0eda863..205ac833fcc 100644
--- a/tests/integration/observability/_s3_v2_support.py
+++ b/tests/integration/observability/_s3_v2_support.py
@@ -27,11 +27,15 @@ class RecordingS3Sink:
fail_attempts: int = 0
fail_until: float = 0.0
fail_status: int = 503
+ fail_code: str = "SinkFailure"
+ fail_body: bytes | None = None
delay_seconds: float = 0.5
lock: threading.Lock = field(default_factory=threading.Lock)
in_flight: int = 0
peak: int = 0
attempts: int = 0
+ attempt_log: list[tuple[float, int]] = field(default_factory=list) # mutable-ok: appended under lock per PUT
+ attempt_counts: dict[str, int] = field(default_factory=dict) # mutable-ok: per-target PUT counts under lock
store: dict[str, bytes] = field(default_factory=dict) # mutable-ok: GET reads must see writes from earlier PUTs
def respond(self, request: Request) -> Reply:
@@ -44,20 +48,31 @@ class RecordingS3Sink:
assert request.target.startswith(f"/{BUCKET}/{PREFIX}/"), request.target
with self.lock:
self.attempts += 1
- if self.attempts <= self.fail_attempts or time.time() < self.fail_until:
- return Reply(
- status=self.fail_status,
- body=b"SinkFailure{self.fail_code}{self.reject_code} gets a single PUT and the "retry" outcome
+ for the flush-level requeue, never an in-call backoff.
"""
from unittest.mock import AsyncMock, MagicMock
@@ -501,24 +512,29 @@ async def test_async_upload_no_retry_on_4xx():
)
test_element = s3BatchLoggingElement(
- s3_object_key="2025-09-14/test-no-retry.json",
- payload={"test": "no-retry"},
- s3_object_download_filename="test-no-retry.json",
+ s3_object_key="2025-09-14/test-retry-400.json",
+ payload={"test": "retry-400"},
+ s3_object_download_filename="test-retry-400.json",
)
response_400 = MagicMock()
response_400.status_code = 400
+ response_400.text = "SomethingElse "
response_400.raise_for_status = MagicMock(side_effect=Exception("400 Bad Request"))
+ response_200 = MagicMock()
+ response_200.status_code = 200
+ response_200.text = ""
+ response_200.raise_for_status = MagicMock()
logger.async_httpx_client = AsyncMock()
- logger.async_httpx_client.put = AsyncMock(return_value=response_400)
+ logger.async_httpx_client.put = AsyncMock(side_effect=[response_400, response_200])
- with patch.object(logger, "handle_callback_failure") as mock_failure:
- await logger.async_upload_data_to_s3(test_element)
+ with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
+ outcome = await logger.async_upload_data_to_s3(test_element)
- # Only 1 attempt — no retry for 4xx
assert logger.async_httpx_client.put.call_count == 1
- mock_failure.assert_called_once_with(callback_name="S3Logger")
+ mock_sleep.assert_not_awaited()
+ assert outcome is False
_SIGV4_ACCESS_KEY = re.compile(r"Credential=(AKIA\d+)/")
@@ -657,21 +673,57 @@ async def test_async_upload_exhausts_403_retries_through_production_http_handler
@pytest.mark.asyncio
-async def test_async_upload_does_not_retry_404_through_production_http_handler(rotating_profile: str, caplog):
+async def test_async_upload_is_single_attempted_on_404_through_production_http_handler(rotating_profile: str, caplog):
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-404.json",
payload={"test": "404"},
s3_object_download_filename="test-404.json",
)
async with _s3_logger_on_production_handler(rotating_profile, [404]) as (logger, requests, mock_sleep):
- await logger.async_upload_data_to_s3(test_element)
+ outcome = await logger.async_upload_data_to_s3(test_element)
assert len(requests) == 1
+ assert outcome is False
mock_sleep.assert_not_awaited()
assert "Error uploading to s3" in caplog.text
+@pytest.mark.asyncio
+async def test_async_upload_access_denied_403_is_retried_and_then_requeued(rotating_profile: str, caplog):
+ test_element = s3BatchLoggingElement(
+ s3_object_key="2025-09-14/test-403-denied.json",
+ payload={"test": "403-denied"},
+ s3_object_download_filename="test-403-denied.json",
+ )
+ requests: list[httpx.Request] = []
+
+ def respond(request: httpx.Request) -> httpx.Response:
+ requests.append(request)
+ return httpx.Response(403, request=request, text="AccessDenied ")
+
+ handler = AsyncHTTPHandler()
+ handler.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_region_name="us-east-1",
+ s3_aws_profile_name=rotating_profile,
+ s3_flush_interval=3600,
+ )
+ logger.async_httpx_client = handler
+ with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
+ outcome = await logger.async_upload_data_to_s3(test_element)
+ await handler.client.aclose()
+
+ assert outcome is False
+ assert len(requests) == 3
+ assert mock_sleep.await_args_list == [call(1), call(2)]
+ assert "Error uploading to s3" in caplog.text
+
+
def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, monkeypatch: pytest.MonkeyPatch):
+ monkeypatch.delenv("AWS_ACCESS_KEY_ID", raising=False)
+ monkeypatch.delenv("AWS_SECRET_ACCESS_KEY", raising=False)
+ monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False)
monkeypatch.setenv("AWS_PROFILE", rotating_profile)
logger = S3Logger(s3_bucket_name="test-bucket", s3_region_name="us-east-1", s3_flush_interval=3600)
test_element = s3BatchLoggingElement(
@@ -684,7 +736,7 @@ def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, mon
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
- return httpx.Response(next(replies), request=request)
+ return httpx.Response(next(replies), request=request, text="SignatureDoesNotMatch ")
handler = HTTPHandler()
handler.client = httpx.Client(transport=httpx.MockTransport(respond))
@@ -2481,12 +2533,14 @@ def _element(payload: dict[str, object], key_suffix: str) -> s3BatchLoggingEleme
def _ok_response() -> MagicMock:
response = MagicMock()
response.status_code = 200
+ response.text = ""
response.raise_for_status = MagicMock()
return response
class _CountingPut:
- def __init__(self) -> None:
+ def __init__(self, width: int) -> None:
+ self.width = width
self.in_flight = 0
self.peak = 0
self.calls = 0
@@ -2495,7 +2549,10 @@ class _CountingPut:
self.in_flight += 1
self.peak = max(self.peak, self.in_flight)
self.calls += 1
- await asyncio.sleep(0.01)
+ for _ in range(50):
+ if self.in_flight >= self.width:
+ break
+ await _real_sleep(0)
self.in_flight -= 1
return _ok_response()
@@ -2515,36 +2572,56 @@ class _LateAppendingPut:
self.element = element
self.fail_first = fail_first
self.appended = False
+ self.failed_key: str | None = None
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
if not self.appended:
self.appended = True
self.logger.log_queue.append(self.element)
if self.fail_first:
- return _failure_response()
+ self.failed_key = url
+ if url == self.failed_key:
+ return _transient_failure_response()
return _ok_response()
+class _AppendingFailingPut:
+ def __init__(self, logger: S3Logger, elements: tuple[s3BatchLoggingElement, ...]) -> None:
+ self.logger = logger
+ self.elements = elements
+ self.appended = False
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ if not self.appended:
+ self.appended = True
+ for element in self.elements:
+ self.logger.log_queue.append(element)
+ return _transient_failure_response()
+
+
class _FailOnSuffixPut:
def __init__(self, suffixes: tuple[str, ...]) -> None:
self.failing = True
self.suffixes = suffixes
+ self.calls: tuple[str, ...] = ()
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls = (*self.calls, url)
if self.failing and url.endswith(self.suffixes):
- return _failure_response()
+ return _transient_failure_response()
return _ok_response()
class _FailUntilClearedPut:
- def __init__(self) -> None:
+ def __init__(self, status: int = 503, code: str | None = "SlowDown", raw_body: str | None = None) -> None:
self.failing = True
+ self.response: Final = _coded_failure_response(status, code, raw_body)
self.calls: tuple[tuple[str, str | None], ...] = ()
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
self.calls = (*self.calls, (url, data))
if self.failing:
- return _failure_response()
+ return self.response
return _ok_response()
@@ -2558,7 +2635,7 @@ async def test_async_send_batch_bounds_concurrent_uploads() -> None:
s3_max_concurrent_uploads=4,
)
- put = _CountingPut()
+ put = _CountingPut(logger.s3_max_concurrent_uploads)
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
@@ -2652,14 +2729,14 @@ def test_invalid_concurrency_falls_back_to_default(bad: object) -> None:
logger = _override_logger(s3_max_concurrent_uploads=bad)
assert logger.s3_max_concurrent_uploads == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
- assert logger._upload_semaphore._value == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
+ assert logger._upload_limiter._value == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
def test_env_backed_concurrency_string_is_parsed() -> None:
logger = _override_logger(s3_max_concurrent_uploads="4")
assert logger.s3_max_concurrent_uploads == 4
- assert logger._upload_semaphore._value == 4
+ assert logger._upload_limiter._value == 4
@pytest.mark.parametrize("empty", [None, ""])
@@ -2674,16 +2751,30 @@ def test_empty_config_concurrency_falls_back_to_constructor_value(empty: object)
)
assert logger.s3_max_concurrent_uploads == 4
- assert logger._upload_semaphore._value == 4
+ assert logger._upload_limiter._value == 4
-def _failure_response() -> MagicMock:
+def _coded_failure_response(status: int, code: str | None, raw_body: str | None = None) -> MagicMock:
+ body: Final = (
+ raw_body if raw_body is not None else (f"{code} " if code is not None else "")
+ )
response = MagicMock()
- response.status_code = 400
- response.raise_for_status = MagicMock(side_effect=Exception("s3 rejected the object"))
+ response.status_code = status
+ response.text = body
+ response.raise_for_status = MagicMock(
+ side_effect=httpx.HTTPStatusError(str(status), request=MagicMock(), response=response)
+ )
return response
+def _transient_failure_response(status: int = 503) -> MagicMock:
+ return _coded_failure_response(status, "SlowDown")
+
+
+def _terminal_failure_response() -> MagicMock:
+ return _coded_failure_response(400, "EntityTooLarge")
+
+
@pytest.mark.asyncio
async def test_failed_uploads_stay_queued_for_next_flush() -> None:
logger = S3Logger(
@@ -2700,12 +2791,16 @@ async def test_failed_uploads_stay_queued_for_next_flush() -> None:
logger.async_httpx_client.put = put
logger.log_queue = list(elements)
- await logger.flush_queue()
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
- assert logger.log_queue == [elements[2], elements[4]]
+ assert [element.s3_object_key for element in logger.log_queue] == [
+ elements[2].s3_object_key,
+ elements[4].s3_object_key,
+ ]
- put.failing = False
- await logger.flush_queue()
+ put.failing = False
+ await logger.flush_queue()
assert logger.log_queue == []
@@ -2728,9 +2823,10 @@ async def test_batch_file_upload_failure_keeps_whole_batch() -> None:
elements = [_element({"i": i}, f"{i}") for i in range(3)]
logger.log_queue = list(elements)
- await logger.flush_queue()
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
- assert len(put.calls) == 1
+ assert len(put.calls) == 3
assert len(logger.log_queue) == 1
assert logger.log_queue[0].body == "\n".join(json.dumps(element.payload) for element in elements)
@@ -2752,9 +2848,16 @@ async def test_events_appended_during_failed_flush_survive() -> None:
first = _element({"id": "first"}, "first")
logger.log_queue = [first]
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert [element.s3_object_key for element in logger.log_queue] == [first.s3_object_key, late.s3_object_key]
+ assert logger.log_queue[0].retrying_since is None
+
+ logger.async_httpx_client.put.failed_key = None
await logger.flush_queue()
- assert logger.log_queue == [first, late]
+ assert logger.log_queue == []
@pytest.mark.asyncio
@@ -2841,7 +2944,8 @@ async def test_failed_batch_file_is_requeued_and_resent_unchanged() -> None:
logger.async_httpx_client.put = put
logger.log_queue = [_element({"i": i}, f"{i}") for i in range(3)]
- await logger.flush_queue()
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
assert len(logger.log_queue) == 1
assert logger.log_queue[0].body is not None
@@ -2851,7 +2955,7 @@ async def test_failed_batch_file_is_requeued_and_resent_unchanged() -> None:
await logger.flush_queue()
assert logger.log_queue == []
- assert len(put.calls) == 2
+ assert len(put.calls) == 4
assert put.calls[0] == put.calls[1]
@@ -2871,7 +2975,8 @@ async def test_elements_appended_after_failed_batch_file_get_their_own_file() ->
logger.async_httpx_client.put = put
logger.log_queue = [_element({"id": "first"}, "first")]
- await logger.flush_queue()
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
late = _element({"id": "late"}, "late")
logger.log_queue.append(late)
@@ -2880,10 +2985,13 @@ async def test_elements_appended_after_failed_batch_file_get_their_own_file() ->
await logger.flush_queue()
assert logger.log_queue == []
- assert len(put.calls) == 3
+ assert len(put.calls) == 5
assert put.calls[0] == put.calls[1]
- assert put.calls[2][0] != put.calls[0][0]
- assert put.calls[2][1] == json.dumps({"id": "late"})
+ second_flush: Final = put.calls[3:]
+ assert put.calls[0] in second_flush
+ late_call: Final = next(call for call in second_flush if call != put.calls[0])
+ assert late_call[0] != put.calls[0][0]
+ assert late_call[1] == json.dumps({"id": "late"})
@pytest.mark.asyncio
@@ -2918,3 +3026,1769 @@ async def test_batch_file_mode_disabled_when_s3_v2_is_cold_storage_logger(monkey
assert len(put.calls) == 2
assert put.calls[1][0].endswith(".jsonl")
+
+
+class _FailOnSuffixCodedPut:
+ def __init__(
+ self, suffixes: tuple[str, ...], status: int, code: str | None = None, raw_body: str | None = None
+ ) -> None:
+ self.suffixes = suffixes
+ self.response: Final = _coded_failure_response(status, code, raw_body)
+ self.calls: tuple[str, ...] = ()
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls = (*self.calls, url)
+ if url.endswith(self.suffixes):
+ return self.response
+ return _ok_response()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ("status", "code", "raw_body", "puts_per_element"),
+ [
+ pytest.param(403, "AccessDenied", None, 3, id="access-denied-403"),
+ pytest.param(403, None, None, 3, id="empty-403"),
+ pytest.param(403, None, "Forbidden", 3, id="html-403"),
+ pytest.param(400, "KMS.DisabledException", None, 1, id="kms-disabled-400"),
+ pytest.param(404, "NoSuchBucket", None, 1, id="no-such-bucket-404"),
+ ],
+)
+async def test_non_terminal_failure_is_requeued_and_delivered_on_recovery(
+ status: int, code: str | None, raw_body: str | None, puts_per_element: int
+) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ elements = [_element({"i": i}, f"{i}") for i in range(5)]
+ put = _FailUntilClearedPut(status=status, code=code, raw_body=raw_body)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = list(elements)
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 5
+ assert len(put.calls) == 5 * puts_per_element
+
+ put.failing = False
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ assert len(put.calls) == 5 * puts_per_element + 5
+ landed: Final = frozenset(
+ element.s3_object_key
+ for element in elements
+ if any(call[0].endswith(element.s3_object_key) for call in put.calls[-5:])
+ )
+ assert landed == frozenset(element.s3_object_key for element in elements)
+
+
+@pytest.mark.asyncio
+async def test_persistent_500_stays_queued_through_a_dozen_failed_flushes() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailUntilClearedPut(status=500, code="InternalError")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": i}, f"{i}") for i in range(5)]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ for _ in range(12):
+ await logger.flush_queue()
+ assert len(logger.log_queue) == 5
+
+ assert len(put.calls) == 12 * 15
+
+ put.failing = False
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ assert len(put.calls) == 12 * 15 + 5
+
+
+@pytest.mark.asyncio
+async def test_terminal_object_is_dropped_once_next_to_delivered_siblings_when_opted_in() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=True,
+ )
+
+ elements = [_element({"i": i}, f"{i}") for i in range(5)]
+ put = _FailOnSuffixCodedPut(("test-1.json",), 400, "EntityTooLarge")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = list(elements)
+
+ await logger.flush_queue()
+
+ assert len(put.calls) == 5
+ assert logger.log_queue == []
+
+
+@pytest.mark.asyncio
+async def test_terminal_object_is_requeued_when_opted_out() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=False,
+ )
+
+ put = _FailOnSuffixCodedPut(("test-1.json",), 400, "EntityTooLarge")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": i}, f"{i}") for i in range(5)]
+
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 1
+ assert logger.log_queue[0].s3_object_key.endswith("test-1.json")
+ assert sum(call.endswith("test-1.json") for call in put.calls) == 1
+ assert len(put.calls) == 5
+
+
+@pytest.mark.asyncio
+async def test_terminal_objects_are_requeued_when_every_upload_in_the_flush_fails() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=True,
+ )
+
+ put = _FailUntilClearedPut(status=400, code="EntityTooLarge")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": i}, f"{i}") for i in range(5)]
+
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 5
+
+
+@pytest.mark.asyncio
+async def test_retrying_past_the_opted_in_budget_is_dropped_only_next_to_delivered_siblings() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=60,
+ )
+
+ aged = _element({"id": "aged"}, "aged").model_copy(update={"retrying_since": _NOW - 120})
+ fresh = _element({"id": "fresh"}, "fresh")
+ put = _FailOnSuffixCodedPut(("test-aged.json",), 503, "SlowDown")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [aged, fresh]
+
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=_NOW),
+ ):
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ assert len(put.calls) == 4
+
+
+@pytest.mark.asyncio
+async def test_retrying_past_the_budget_stays_queued_when_the_whole_flush_fails() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=60,
+ )
+
+ aged = _element({"id": "aged"}, "aged").model_copy(update={"retrying_since": _NOW - 120})
+ put = _FailUntilClearedPut(status=503, code="SlowDown")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [aged]
+
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=_NOW),
+ ):
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 1
+
+
+@pytest.mark.asyncio
+async def test_overflow_after_a_failed_flush_trims_failed_first_and_counts_upload_failures_only(
+ caplog,
+) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size=4,
+ )
+
+ late = tuple(_element({"id": f"late-{index}"}, f"late-{index}") for index in range(3))
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _AppendingFailingPut(logger, late)
+ logger.log_queue = [_element({"id": "first"}, "first"), _element({"id": "second"}, "second")]
+
+ with (
+ patch.object(logger, "handle_callback_failure") as mock_failure,
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ pytest.raises(S3BatchUploadError),
+ ):
+ await logger.async_send_batch()
+
+ assert [element.payload["id"] for element in logger.log_queue] == ["second", "late-0", "late-1", "late-2"]
+ failed_uploads: Final = 2
+ assert mock_failure.call_count == failed_uploads
+ mock_failure.assert_called_with(callback_name="S3Logger")
+ assert "dropped 1 oldest events" in caplog.text
+
+
+@pytest.mark.asyncio
+async def test_default_logger_ages_out_elements_retrying_longer_than_an_hour(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ elements = [
+ _element({"i": index}, f"{index}").model_copy(update={"retrying_since": _NOW - 7200}) for index in range(3)
+ ]
+ put = _FailOnSuffixPut(("test-1.json", "test-2.json"))
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = list(elements)
+
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=_NOW),
+ ):
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ assert "uploads dropped" in caplog.text
+
+
+@pytest.mark.asyncio
+async def test_opted_out_logger_never_ages_out_long_retrying_elements(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=0,
+ )
+
+ elements = [
+ _element({"i": index}, f"{index}").model_copy(update={"retrying_since": _NOW - 7200}) for index in range(3)
+ ]
+ put = _FailOnSuffixPut(("test-1.json", "test-2.json"))
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = list(elements)
+
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=_NOW),
+ ):
+ await logger.flush_queue()
+
+ assert [element.s3_object_key for element in logger.log_queue] == [
+ elements[1].s3_object_key,
+ elements[2].s3_object_key,
+ ]
+ assert "uploads dropped" not in caplog.text
+
+ put.failing = False
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ landed: Final = frozenset(call_url.rsplit("/", 1)[-1] for call_url in put.calls)
+ assert landed == frozenset(f"test-{index}.json" for index in range(3))
+
+
+@pytest.mark.asyncio
+async def test_queue_grows_past_the_cap_while_the_sink_fails_and_everything_lands() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size=5,
+ )
+
+ elements = [_element({"i": index}, f"{index}") for index in range(8)]
+ put = _FailUntilClearedPut()
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = list(elements[:5])
+
+ with (
+ patch.object(logger, "handle_callback_failure") as mock_failure,
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ ):
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 5
+ upload_failures: Final = 5
+ assert mock_failure.call_count == upload_failures
+
+ for element in elements[5:]:
+ logger.log_queue.append(element)
+
+ put.failing = False
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ landed: Final = frozenset(call[0].rsplit("/", 1)[-1] for call in put.calls[-8:])
+ assert landed == frozenset(f"test-{index}.json" for index in range(8)) # calls are (url, data) pairs
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ("status", "code"),
+ [
+ pytest.param(404, "NoSuchKey", id="404"),
+ pytest.param(401, None, id="401"),
+ pytest.param(400, None, id="uncoded-400"),
+ ],
+)
+async def test_unlisted_status_gets_one_put_and_stays_queued(status: int, code: str | None) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailUntilClearedPut(status=status, code=code)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": index}, f"{index}") for index in range(4)]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
+ await logger.flush_queue()
+
+ assert len(put.calls) == 4
+ mock_sleep.assert_not_awaited()
+ assert len(logger.log_queue) == 4
+
+
+class _SyncRecordingClient:
+ def __init__(self, response: httpx.Response) -> None:
+ self.response: Final = response
+ self.put_calls: list = [] # mutable-ok: call log appended once per PUT
+
+ def put(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.put_calls.append(url)
+ return self.response
+
+
+def test_sync_upload_404_is_single_attempt_without_sleep() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ sync_client: Final = _SyncRecordingClient(_coded_failure_response(404, "NoSuchKey"))
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=sync_client),
+ patch("time.sleep") as mock_sleep,
+ ):
+ logger.upload_data_to_s3(_element({"id": "sync-404"}, "sync-404"))
+
+ assert len(sync_client.put_calls) == 1
+ mock_sleep.assert_not_called()
+
+
+@pytest.mark.parametrize(
+ ("status", "expected_puts", "expected_sleeps"),
+ [
+ pytest.param(429, 1, [], id="429-single"),
+ pytest.param(408, 1, [], id="408-single"),
+ pytest.param(502, 1, [], id="502-single"),
+ pytest.param(504, 1, [], id="504-single"),
+ pytest.param(503, 3, [call(1), call(2)], id="503-backoff"),
+ ],
+)
+def test_sync_upload_retry_set_matches_base(status: int, expected_puts: int, expected_sleeps: list) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ sync_client: Final = _SyncRecordingClient(_coded_failure_response(status, "SlowDown"))
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=sync_client),
+ patch("time.sleep") as mock_sleep,
+ ):
+ logger.upload_data_to_s3(_element({"id": "sync"}, "sync"))
+
+ assert len(sync_client.put_calls) == expected_puts
+ assert mock_sleep.call_args_list == expected_sleeps
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ("status", "code"),
+ [
+ pytest.param(503, "SlowDown", id="503"),
+ pytest.param(500, "InternalError", id="500"),
+ pytest.param(403, "AccessDenied", id="access-denied-403"),
+ ],
+)
+async def test_retryable_statuses_back_off_three_attempts(status: int, code: str | None) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailUntilClearedPut(status=status, code=code)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"id": "req"}, "req")]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
+ await logger.flush_queue()
+
+ assert len(put.calls) == 3
+ assert mock_sleep.await_args_list == [call(1), call(2)]
+ assert len(logger.log_queue) == 1
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ("status", "code"),
+ [
+ pytest.param(429, "TooManyRequests", id="429"),
+ pytest.param(408, None, id="408"),
+ pytest.param(502, None, id="502"),
+ pytest.param(504, None, id="504"),
+ ],
+)
+async def test_non_base_statuses_are_not_retried_in_call(status: int, code: str | None) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailUntilClearedPut(status=status, code=code)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"id": "req"}, "req")]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
+ await logger.flush_queue()
+
+ assert len(put.calls) == 1
+ assert mock_sleep.await_args_list == []
+ assert len(logger.log_queue) == 1
+
+
+class _FirstFailThenOkPut:
+ def __init__(self, fail_suffix: str) -> None:
+ self.fail_suffix = fail_suffix
+ self.failed_once = False
+ self.calls: tuple[str, ...] = ()
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls = (*self.calls, url)
+ if url.endswith(self.fail_suffix) and not self.failed_once:
+ self.failed_once = True
+ return _transient_failure_response()
+ return _ok_response()
+
+
+@pytest.mark.asyncio
+async def test_retry_finishes_before_the_next_first_attempt() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_concurrent_uploads=1,
+ )
+
+ put = _FirstFailThenOkPut("test-a.json")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"id": "a"}, "a"), _element({"id": "b"}, "b")]
+
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ await logger.flush_queue()
+
+ assert [call_url.rsplit("/", 1)[-1] for call_url in put.calls] == ["test-a.json", "test-a.json", "test-b.json"]
+ assert logger.log_queue == []
+
+
+@pytest.mark.asyncio
+async def test_objects_in_backoff_are_bounded_by_the_slot_width() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_concurrent_uploads=2,
+ )
+
+ put = _FailUntilClearedPut(status=503)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": i}, f"{i}") for i in range(20)]
+
+ sleeping: Final[list[int]] = [0]
+ peak: Final[list[int]] = [0]
+
+ async def counting_sleep(delay: float) -> None:
+ sleeping[0] += 1
+ peak[0] = max(peak[0], sleeping[0])
+ for _ in range(10):
+ await _real_sleep(0)
+ sleeping[0] -= 1
+
+ with patch("asyncio.sleep", new=counting_sleep):
+ await logger.flush_queue()
+
+ assert peak[0] <= 2, f"{peak[0]} objects slept at once, slot width is 2"
+ assert len(put.calls) == 60
+
+
+@pytest.mark.asyncio
+async def test_subclass_returning_true_drains_the_queue() -> None:
+ class _TrueUploadLogger(S3Logger):
+ async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
+ return True
+
+ logger = _TrueUploadLogger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ logger.async_httpx_client = AsyncMock()
+ logger.log_queue = [_element({"id": "a"}, "a")]
+
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ logger.async_httpx_client.put.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_failed_direct_upload_returns_false() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailUntilClearedPut(status=500)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ test_element = _element({"id": "x"}, "x")
+
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ outcome = await logger.async_upload_data_to_s3(test_element)
+
+ assert outcome is False
+ assert len(put.calls) == 3
+
+
+@pytest.mark.asyncio
+async def test_terminal_drop_of_one_element_does_not_drop_a_sibling_with_the_same_key() -> None:
+ class _TerminalForMarkerPut:
+ def __init__(self) -> None:
+ self.calls: tuple[str | None, ...] = ()
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls = (*self.calls, data)
+ if data is not None and "terminal-marker" in data:
+ return _terminal_failure_response()
+ return _transient_failure_response()
+
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=True,
+ )
+
+ put = _TerminalForMarkerPut()
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+
+ shared_key = "2025-09-14/shared.json"
+ dropped = s3BatchLoggingElement(
+ s3_object_key=shared_key, payload={"m": "terminal-marker"}, s3_object_download_filename="shared.json"
+ )
+ sibling = s3BatchLoggingElement(
+ s3_object_key=shared_key, payload={"m": "healthy"}, s3_object_download_filename="shared.json"
+ )
+
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ assert await logger._upload_outcome(dropped) == "dropped"
+ assert await logger._upload_outcome(sibling) == "retry"
+
+
+def test_upload_semaphore_alias_is_the_limiter() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ assert logger._upload_semaphore is logger._upload_limiter
+
+
+@pytest.mark.asyncio
+async def test_overridden_upload_stays_bounded_by_the_configured_width() -> None:
+ class _InFlightUploadLogger(S3Logger):
+ def __init__(self, **kwargs: object) -> None:
+ super().__init__(**kwargs)
+ self.in_flight = 0
+ self.peak = 0
+
+ async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
+ self.in_flight += 1
+ self.peak = max(self.peak, self.in_flight)
+ for _ in range(10):
+ await _real_sleep(0)
+ self.in_flight -= 1
+ return True
+
+ logger = _InFlightUploadLogger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_concurrent_uploads=4,
+ )
+ logger.log_queue = [_element({"i": index}, f"{index}") for index in range(40)]
+
+ await logger.flush_queue()
+
+ assert logger.peak <= 4
+ assert logger.log_queue == []
+
+
+@pytest.mark.asyncio
+async def test_holding_the_semaphore_during_a_direct_upload_does_not_deadlock() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_concurrent_uploads=1,
+ )
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _RecordingPut()
+ element = _element({"id": "x"}, "x")
+
+ async def held_upload() -> bool:
+ async with logger._upload_semaphore:
+ return await logger.async_upload_data_to_s3(element)
+
+ assert await asyncio.wait_for(held_upload(), timeout=5) is True
+
+
+@pytest.mark.asyncio
+async def test_assigning_a_semaphore_changes_the_upload_width() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ logger._upload_semaphore = asyncio.Semaphore(3)
+
+ put = _CountingPut(width=3)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": index}, f"{index}") for index in range(30)]
+
+ await logger.flush_queue()
+
+ assert put.peak == 3
+ assert logger.log_queue == []
+
+
+@pytest.mark.asyncio
+async def test_terminal_code_is_retried_like_base_when_the_drop_flag_is_off() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=False,
+ )
+ put = _StatusPut([_coded_failure_response(403, "InvalidRequest"), _ok_response()])
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ failures = AsyncMock()
+ logger.handle_callback_failure = failures
+
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ assert await logger.async_upload_data_to_s3(_element({"id": "x"}, "x")) is True
+
+ assert put.calls == 2
+ failures.assert_not_called()
+
+ dropping = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ put.calls = 0
+ dropping.async_httpx_client = AsyncMock()
+ dropping.async_httpx_client.put = put
+
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ assert await dropping.async_upload_data_to_s3(_element({"id": "x"}, "x")) is False
+
+ assert put.calls == 1
+
+
+def test_sync_terminal_code_is_retried_like_base_when_the_drop_flag_is_off() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=False,
+ )
+ mock_sync_client = MagicMock()
+ mock_sync_client.put = MagicMock(side_effect=[_coded_failure_response(403, "InvalidRequest"), _ok_response()])
+ failures = MagicMock()
+ logger.handle_callback_failure = failures
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep"),
+ ):
+ logger.upload_data_to_s3(_element({"id": "x"}, "x"))
+
+ assert mock_sync_client.put.call_count == 2
+ failures.assert_not_called()
+
+ dropping = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ mock_sync_client.put = MagicMock(side_effect=[_coded_failure_response(403, "InvalidRequest"), _ok_response()])
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep"),
+ ):
+ dropping.upload_data_to_s3(_element({"id": "x"}, "x"))
+
+ assert mock_sync_client.put.call_count == 1
+
+
+def test_sync_retry_lines_stay_at_warning_level(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ mock_sync_client = MagicMock()
+ mock_sync_client.put = MagicMock(
+ side_effect=[_transient_failure_response(503), _transient_failure_response(503), _ok_response()]
+ )
+
+ with (
+ caplog.at_level("WARNING"),
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep"),
+ ):
+ logger.upload_data_to_s3(_element({"id": "x"}, "x"))
+
+ assert mock_sync_client.put.call_count == 3
+ assert sum(1 for record in caplog.records if "retrying in" in record.getMessage()) == 2
+
+
+@pytest.mark.asyncio
+async def test_direct_async_upload_logs_retry_lines_at_warning_level(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ put = _StatusPut([_transient_failure_response(503), _ok_response()])
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+
+ with caplog.at_level("WARNING"), patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ assert await logger.async_upload_data_to_s3(_element({"id": "x"}, "x")) is True
+
+ assert put.calls == 2
+ assert sum(1 for record in caplog.records if "retrying in" in record.getMessage()) == 1
+
+
+def _init_bypassed_logger() -> S3Logger:
+ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
+
+ logger = S3Logger.__new__(S3Logger)
+ logger.iam_cache = BaseAWSLLM._shared_iam_cache
+ logger.s3_endpoint_url = None
+ logger.s3_bucket_name = "test-bucket"
+ logger.s3_region_name = "us-east-1"
+ logger.s3_use_virtual_hosted_style = False
+ logger.s3_verify = None
+ logger.s3_aws_access_key_id = "test-key"
+ logger.s3_aws_secret_access_key = "test-secret"
+ logger.s3_aws_session_token = None
+ logger.s3_aws_session_name = None
+ logger.s3_aws_profile_name = None
+ logger.s3_aws_role_name = None
+ logger.s3_aws_web_identity_token = None
+ logger.s3_aws_sts_endpoint = None
+ logger.s3_server_side_encryption = None
+ logger.s3_sse_kms_key_id = None
+ logger.s3_log_prompts_only = None
+ return logger
+
+
+@pytest.mark.asyncio
+async def test_init_bypassed_logger_retries_a_503_and_reports_a_404() -> None:
+ logger = _init_bypassed_logger()
+ put = _StatusPut([_transient_failure_response(503), _ok_response()])
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ assert await logger.async_upload_data_to_s3(_element({"id": "x"}, "x")) is True
+
+ assert put.calls == 2
+
+ put.calls = 0
+ put.responses = [_coded_failure_response(404, None)]
+ with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))):
+ assert await logger.async_upload_data_to_s3(_element({"id": "y"}, "y")) is False
+
+ assert put.calls == 1
+
+
+def test_init_bypassed_sync_logger_retries_a_503_and_reports_a_404() -> None:
+ logger = _init_bypassed_logger()
+ mock_sync_client = MagicMock()
+ mock_sync_client.put = MagicMock(side_effect=[_transient_failure_response(503), _ok_response()])
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep"),
+ ):
+ logger.upload_data_to_s3(_element({"id": "x"}, "x"))
+
+ assert mock_sync_client.put.call_count == 2
+ retried_headers: Final = dict(mock_sync_client.put.call_args.kwargs["headers"])
+ assert "X-Amz-Date" in retried_headers
+
+ mock_sync_client.put = MagicMock(return_value=_coded_failure_response(404, None))
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep"),
+ ):
+ logger.upload_data_to_s3(_element({"id": "y"}, "y"))
+
+ assert mock_sync_client.put.call_count == 1
+ failed_url: Final = str(mock_sync_client.put.call_args[0][0])
+ assert "test-y.json" in failed_url
+
+
+@pytest.mark.asyncio
+async def test_subclass_with_base_style_upload_bounded_drains_the_queue() -> None:
+ class _BaseStyleLogger(S3Logger):
+ async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool:
+ return True
+
+ logger = _BaseStyleLogger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+ logger.async_httpx_client = AsyncMock()
+ logger.log_queue = [_element({"id": "a"}, "a")]
+
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ logger.async_httpx_client.put.assert_not_called()
+
+
+def test_bool_config_values_fall_back_to_the_default() -> None:
+ from litellm.integrations.s3 import (
+ resolve_s3_max_concurrent_uploads,
+ resolve_s3_max_queue_size,
+ resolve_s3_max_retry_age_seconds,
+ )
+
+ assert resolve_s3_max_concurrent_uploads(True, 16) == 1
+ assert resolve_s3_max_queue_size(True, 50000) == 50000
+ assert resolve_s3_max_retry_age_seconds(True, 3600) == 3600
+
+
+def test_int_env_helper_falls_back_on_non_numeric(monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.litellm_core_utils.env_utils import get_env_int
+
+ monkeypatch.setenv("TEST_S3_INT_ENV", "abc")
+ assert get_env_int("TEST_S3_INT_ENV", 3) == 3
+ monkeypatch.setenv("TEST_S3_INT_ENV", "7")
+ assert get_env_int("TEST_S3_INT_ENV", 3) == 7
+
+
+class _FailOncePerKeyPut:
+ def __init__(self) -> None:
+ self.failed: set[str] = set()
+ self.calls: tuple[str, ...] = ()
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls = (*self.calls, url)
+ if url not in self.failed:
+ self.failed.add(url)
+ return _transient_failure_response()
+ return _ok_response()
+
+
+class _SlowFailOncePerKeyPut:
+ def __init__(self, dumps_count) -> None:
+ self.failed: set[str] = set()
+ self.dumps_count = dumps_count
+ self.first_completed: int | None = None
+ self.calls: tuple[str, ...] = ()
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls = (*self.calls, url)
+ await _real_sleep(0)
+ if self.first_completed is None:
+ self.first_completed = self.dumps_count()
+ if url not in self.failed:
+ self.failed.add(url)
+ return _transient_failure_response()
+ return _ok_response()
+
+
+@pytest.mark.asyncio
+async def test_peak_serialized_bodies_bounded_by_upload_width() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps as real_safe_dumps
+
+ dumps_calls: list[object] = []
+
+ def counting_dumps(*args, **kwargs):
+ dumps_calls.append(args)
+ return real_safe_dumps(*args, **kwargs)
+
+ put = _SlowFailOncePerKeyPut(lambda: len(dumps_calls))
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": index}, f"{index}") for index in range(64)]
+
+ with (
+ patch("litellm.integrations.s3_v2.safe_dumps", side_effect=counting_dumps),
+ patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))),
+ ):
+ await logger.flush_queue()
+
+ assert put.first_completed is not None
+ assert put.first_completed <= logger.s3_max_concurrent_uploads
+ assert len(dumps_calls) == 64
+ assert len(put.calls) == 128
+
+
+@pytest.mark.asyncio
+async def test_send_batch_calls_upload_with_one_positional_arg() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ uploaded: list[str] = [] # mutable-ok: appended once per upload by the double
+
+ async def mock_upload(batch_logging_element) -> str:
+ uploaded.append(batch_logging_element.s3_object_key)
+ return "delivered"
+
+ logger.async_upload_data_to_s3 = mock_upload
+ logger.log_queue = [_element({"id": "a"}, "a"), _element({"id": "b"}, "b")]
+
+ await logger.flush_queue()
+
+ assert sorted(key.rsplit("/", 1)[-1] for key in uploaded) == ["test-a.json", "test-b.json"]
+ assert logger.log_queue == []
+
+
+@pytest.mark.asyncio
+async def test_retries_serialize_the_body_once_per_element_per_flush() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps as real_safe_dumps
+
+ dumps_calls: list[object] = []
+
+ def counting_dumps(*args, **kwargs):
+ dumps_calls.append(args)
+ return real_safe_dumps(*args, **kwargs)
+
+ put = _FailUntilClearedPut(status=503)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": index}, f"{index}") for index in range(8)]
+
+ with (
+ patch("litellm.integrations.s3_v2.safe_dumps", side_effect=counting_dumps),
+ patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))),
+ ):
+ await logger.flush_queue()
+
+ assert len(dumps_calls) == 8
+ assert len(put.calls) == 24
+ assert len(logger.log_queue) == 8
+
+
+@pytest.mark.asyncio
+async def test_async_flush_logs_one_retry_warning(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailOncePerKeyPut()
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"i": index}, f"{index}") for index in range(8)]
+
+ with caplog.at_level("WARNING"), patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ assert sum(1 for record in caplog.records if "in-call retries" in record.getMessage()) == 1
+ assert all("retrying in" not in record.getMessage() for record in caplog.records)
+
+
+class _AppendingSuffixFailingPut:
+ def __init__(self, logger: S3Logger, element: s3BatchLoggingElement, fail_suffixes: tuple[str, ...]) -> None:
+ self.logger = logger
+ self.element = element
+ self.fail_suffixes = fail_suffixes
+ self.appended = False
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ if not self.appended:
+ self.appended = True
+ self.logger.log_queue.append(self.element)
+ if url.endswith(self.fail_suffixes):
+ return _transient_failure_response()
+ return _ok_response()
+
+
+@pytest.mark.asyncio
+async def test_failed_elements_stay_oldest_first_when_requeued() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=3600,
+ )
+
+ late = _element({"id": "late"}, "late")
+ failed = _element({"id": "f2"}, "f2")
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _AppendingSuffixFailingPut(logger, late, ("test-f2.json",))
+ logger.log_queue = [_element({"id": "f1"}, "f1"), failed]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert [element.s3_object_key for element in logger.log_queue] == [failed.s3_object_key, late.s3_object_key]
+ assert logger.log_queue[0].retrying_since is not None
+
+
+@pytest.mark.asyncio
+async def test_overflow_prefers_arrivals_over_failed_elements_without_counting_the_trim(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size=1,
+ )
+
+ late = _element({"id": "late"}, "late")
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _AppendingSuffixFailingPut(logger, late, ("test-f2.json",))
+ logger.log_queue = [_element({"id": "f1"}, "f1"), _element({"id": "f2"}, "f2")]
+
+ with (
+ patch.object(logger, "handle_callback_failure") as mock_failure,
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ ):
+ await logger.flush_queue()
+
+ assert [element.s3_object_key for element in logger.log_queue] == [late.s3_object_key]
+ failed_uploads: Final = 1
+ assert mock_failure.call_count == failed_uploads
+ assert "dropped 1 oldest events" in caplog.text
+
+
+@pytest.mark.asyncio
+async def test_fresh_elements_upload_before_stale_retries_after_a_failed_flush() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_concurrent_uploads=1,
+ )
+
+ late = _element({"id": "late"}, "late")
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _AppendingSuffixFailingPut(logger, late, ("test-f2.json",))
+ logger.log_queue = [_element({"id": "f1"}, "f1"), _element({"id": "f2"}, "f2")]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ recovered = _FailOnSuffixPut(("never-matches",))
+ logger.async_httpx_client.put = recovered
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert [call_url.rsplit("/", 1)[-1] for call_url in recovered.calls] == ["test-late.json", "test-f2.json"]
+ assert logger.log_queue == []
+
+
+@pytest.mark.asyncio
+async def test_repeated_overflow_trims_oldest_across_failed_flushes(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size=3,
+ )
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _AppendingFailingPut(logger, (_element({"id": "d"}, "d"),))
+ logger.log_queue = [_element({"id": name}, name) for name in ("a", "b", "c")]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert [element.payload["id"] for element in logger.log_queue] == ["b", "c", "d"]
+
+ logger.async_httpx_client.put = _AppendingFailingPut(logger, (_element({"id": "e"}, "e"),))
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert [element.payload["id"] for element in logger.log_queue] == ["c", "d", "e"]
+ assert caplog.text.count("dropped 1 oldest events") == 2
+
+
+@pytest.mark.asyncio
+async def test_retry_age_budget_drops_after_the_clock_set_by_a_partial_failure(caplog) -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=1,
+ )
+
+ put = _FailOnSuffixCodedPut(("test-poison.json",), 503, "SlowDown")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"id": "poison"}, "poison"), _element({"id": "good"}, "good")]
+
+ t0: Final = _NOW
+ with patch.object(logger, "handle_callback_failure") as mock_failure:
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=t0),
+ ):
+ await logger.flush_queue()
+
+ assert [element.s3_object_key for element in logger.log_queue] == ["2025-09-14/test-poison.json"]
+ assert logger.log_queue[0].retrying_since == t0
+
+ logger.log_queue.append(_element({"id": "good-2"}, "good-2"))
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=t0 + 2),
+ ):
+ await logger.flush_queue()
+
+ assert logger.log_queue == []
+ assert "retrying longer than s3_max_retry_age_seconds=1" in caplog.text
+ poison_puts: Final = sum(1 for call_url in put.calls if call_url.endswith("test-poison.json"))
+ assert poison_puts == 6
+ upload_failures: Final = 2
+ assert mock_failure.call_count == upload_failures
+
+
+@pytest.mark.asyncio
+async def test_the_retry_clock_starts_at_the_first_partial_failure_not_first_seen() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=1,
+ )
+
+ put = _FailUntilClearedPut(status=503, code="SlowDown")
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"id": "poison"}, "poison")]
+
+ t0: Final = _NOW
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=t0),
+ ):
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 1
+ assert logger.log_queue[0].retrying_since is None
+
+ logger.log_queue.append(_element({"id": "good"}, "good"))
+ put.failing = False
+ failing_poison: Final = _FailOnSuffixCodedPut(("test-poison.json",), 503, "SlowDown")
+ logger.async_httpx_client.put = failing_poison
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=t0 + 500),
+ ):
+ await logger.flush_queue()
+
+ assert [element.s3_object_key for element in logger.log_queue] == ["2025-09-14/test-poison.json"]
+ assert logger.log_queue[0].retrying_since == t0 + 500
+
+
+def test_sync_upload_retries_access_denied_403(caplog):
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ test_element = s3BatchLoggingElement(
+ s3_object_key="2025-09-14/test-sync-403.json",
+ payload={"test": "sync-403"},
+ s3_object_download_filename="test-sync-403.json",
+ )
+
+ mock_sync_client = MagicMock()
+ mock_sync_client.put = MagicMock(return_value=_coded_failure_response(403, "AccessDenied"))
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep") as mock_sleep,
+ ):
+ logger.upload_data_to_s3(test_element)
+
+ assert mock_sync_client.put.call_count == 3
+ assert mock_sleep.call_args_list == [call(1), call(2)]
+ assert "dropping object" not in caplog.text
+
+
+def test_sync_upload_drops_terminal_object_once_and_logs_it_only_when_opted_in(caplog):
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=True,
+ )
+
+ test_element = s3BatchLoggingElement(
+ s3_object_key="2025-09-14/test-sync-terminal.json",
+ payload={"test": "sync-terminal"},
+ s3_object_download_filename="test-sync-terminal.json",
+ )
+
+ mock_sync_client = MagicMock()
+ mock_sync_client.put = MagicMock(return_value=_coded_failure_response(400, "EntityTooLarge"))
+
+ with (
+ patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client),
+ patch("time.sleep") as mock_sleep,
+ ):
+ logger.upload_data_to_s3(test_element)
+
+ assert mock_sync_client.put.call_count == 1
+ mock_sleep.assert_not_called()
+ assert "dropping object" in caplog.text
+
+
+@pytest.mark.asyncio
+async def test_requeued_batch_file_keeps_the_earliest_member_retrying_since() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_batch_file_upload=True,
+ )
+
+ put = _FailUntilClearedPut()
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+
+ stale: Final = _NOW - 30
+ retried = _element({"id": "retried"}, "retried").model_copy(update={"retrying_since": stale})
+ fresh = _element({"id": "fresh"}, "fresh")
+ logger.log_queue = [retried, fresh]
+
+ with (
+ patch("asyncio.sleep", new_callable=AsyncMock),
+ patch("time.monotonic", return_value=_NOW),
+ ):
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 1
+ assert logger.log_queue[0].s3_object_key.endswith(".jsonl")
+ assert logger.log_queue[0].retrying_since == stale
+
+
+@pytest.mark.asyncio
+async def test_an_unlisted_5xx_is_requeued_without_an_extra_attempt() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ put = _FailUntilClearedPut(status=507)
+
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = put
+ logger.log_queue = [_element({"id": "req-507"}, "507")]
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert len(logger.log_queue) == 1
+ assert len(put.calls) == 1
+
+
+@pytest.mark.parametrize("configured", [0, "0", None, ""])
+def test_retry_age_resolution_disables_the_budget(configured: object) -> None:
+ from litellm.integrations.s3 import resolve_s3_max_retry_age_seconds
+
+ assert resolve_s3_max_retry_age_seconds(configured, 3600) is None
+
+
+@pytest.mark.parametrize("configured", ["abc", -5, True])
+def test_invalid_retry_age_resolution_falls_back_with_a_warning(configured: object, caplog) -> None:
+ from litellm.integrations.s3 import resolve_s3_max_retry_age_seconds
+
+ assert resolve_s3_max_retry_age_seconds(configured, 3600) == 3600
+ assert "s3_max_retry_age_seconds" in caplog.text
+
+
+def test_retry_age_resolution_accepts_a_positive_int() -> None:
+ from litellm.integrations.s3 import resolve_s3_max_retry_age_seconds
+
+ assert resolve_s3_max_retry_age_seconds(30, 3600) == 30
+
+
+def test_default_logger_sets_a_one_hour_retry_age_budget() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ )
+
+ assert logger.s3_max_retry_age_seconds == 3600
+
+
+def test_constructor_zero_disables_the_retry_age_budget() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=0,
+ )
+
+ assert logger.s3_max_retry_age_seconds is None
+
+
+def test_invalid_callback_params_retry_age_falls_back_to_the_default() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_callback_params_override={"s3_max_retry_age_seconds": "abc"},
+ )
+
+ assert logger.s3_max_retry_age_seconds == 3600
+
+
+def test_callback_params_retry_age_wins_over_constructor() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_retry_age_seconds=30,
+ s3_callback_params_override={"s3_max_retry_age_seconds": 60},
+ )
+
+ assert logger.s3_max_retry_age_seconds == 60
+
+
+def test_callback_params_drop_terminal_error_wins_over_constructor() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=False,
+ s3_callback_params_override={"s3_drop_on_terminal_error": True},
+ )
+
+ assert logger.s3_drop_on_terminal_error is True
+
+
+def test_invalid_callback_params_drop_terminal_error_falls_back_to_constructor_value() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_drop_on_terminal_error=True,
+ s3_callback_params_override={"s3_drop_on_terminal_error": "banana"},
+ )
+
+ assert logger.s3_drop_on_terminal_error is True
+
+
+def test_callback_params_adaptive_concurrency_wins_over_constructor() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_adaptive_concurrency=False,
+ s3_callback_params_override={"s3_adaptive_concurrency": "true"},
+ )
+
+ assert logger.s3_adaptive_concurrency is True
+ assert logger._upload_limiter._ceiling > logger._upload_limiter.limit
+
+
+def test_invalid_callback_params_max_adaptive_concurrency_falls_back_to_default() -> None:
+ from litellm.constants import DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY
+
+ logger = _override_logger(s3_adaptive_concurrency=True, s3_max_adaptive_concurrency="abc")
+
+ assert logger.s3_max_adaptive_concurrency == DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY
+ assert logger._upload_limiter._ceiling == DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY
+
+
+def test_callback_params_queue_size_wins_over_constructor() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size=7,
+ s3_callback_params_override={"s3_max_queue_size": 4},
+ )
+
+ assert logger.s3_max_queue_size == 4
+ assert logger.max_queue_size == 4
+
+
+def test_invalid_callback_params_queue_size_falls_back_to_constructor_value() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size=7,
+ s3_callback_params_override={"s3_max_queue_size": "abc"},
+ )
+
+ assert logger.s3_max_queue_size == 7
+
+
+def test_invalid_constructor_queue_size_falls_back_to_default() -> None:
+ from litellm.integrations.custom_batch_logger import CustomBatchLogger
+
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_queue_size="abc",
+ )
+
+ assert logger.s3_max_queue_size == CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE
+ assert logger.max_queue_size == CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE
+
+
+class _StatusPut:
+ def __init__(self, responses: "list[MagicMock | Exception]") -> None:
+ self.responses = responses
+ self.calls = 0
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls += 1
+ outcome = self.responses[min(self.calls - 1, len(self.responses) - 1)]
+ if isinstance(outcome, Exception):
+ raise outcome
+ return outcome
+
+
+def _slow_down_response(status: int = 200) -> MagicMock:
+ response = _ok_response() if status == 200 else _transient_failure_response(status)
+ response.text = "SlowDown "
+ return response
+
+
+@pytest.mark.asyncio
+async def test_503_response_lowers_the_adaptive_limit() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _StatusPut([_transient_failure_response(503), _ok_response()])
+
+ logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ logger.log_queue = [_element({"i": 0}, "0")]
+ await logger.flush_queue()
+
+ assert logger._upload_limiter.limit == 32
+
+
+@pytest.mark.asyncio
+async def test_429_response_lowers_the_adaptive_limit() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _StatusPut([_transient_failure_response(429), _ok_response()])
+
+ logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ logger.log_queue = [_element({"i": 0}, "0")]
+ await logger.flush_queue()
+
+ assert logger._upload_limiter.limit == 32
+
+
+@pytest.mark.asyncio
+async def test_slow_down_body_code_lowers_the_adaptive_limit() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _StatusPut([_slow_down_response()])
+
+ logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth
+ logger.log_queue = [_element({"i": 0}, "0")]
+ await logger.async_send_batch()
+
+ assert logger._upload_limiter.limit == 32
+
+
+@pytest.mark.asyncio
+async def test_transport_error_lowers_the_adaptive_limit() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _StatusPut(
+ [httpx.ConnectError("connect refused", request=MagicMock()), _ok_response()]
+ )
+
+ logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ logger.log_queue = [_element({"i": 0}, "0")]
+ await logger.flush_queue()
+
+ assert logger._upload_limiter.limit == 32
+
+
+@pytest.mark.asyncio
+async def test_fast_uploads_raise_the_adaptive_limit() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _RecordingPut()
+
+ before: Final = logger._upload_limiter.limit
+ logger.log_queue = [_element({"i": i}, f"{i}") for i in range(before)]
+ await logger.async_send_batch()
+
+ assert logger._upload_limiter.limit > before
+
+
+@pytest.mark.asyncio
+async def test_configured_concurrency_is_the_fixed_limit_when_adaptive_is_off() -> None:
+ logger = S3Logger(
+ s3_bucket_name="test-bucket",
+ s3_aws_access_key_id="test-key",
+ s3_aws_secret_access_key="test-secret",
+ s3_region_name="us-east-1",
+ s3_max_concurrent_uploads=64,
+ )
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _StatusPut([_transient_failure_response(503), _ok_response()])
+
+ assert logger._upload_limiter._value == 64
+
+ logger.log_queue = [_element({"i": 0}, "0")]
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ await logger.flush_queue()
+
+ assert logger._upload_limiter._value == 64
+
+
+@pytest.mark.asyncio
+async def test_the_limit_never_falls_below_the_configured_width() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True, s3_max_concurrent_uploads=8)
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _StatusPut(
+ [_transient_failure_response(503), _transient_failure_response(503), _transient_failure_response(503)]
+ )
+
+ with patch("asyncio.sleep", new_callable=AsyncMock):
+ logger.log_queue = [_element({"i": 0}, "0")]
+ await logger.flush_queue()
+
+ assert logger._upload_limiter.limit == 8
+
+
+def test_default_upload_width_is_16() -> None:
+ from litellm.constants import DEFAULT_S3_MAX_CONCURRENT_UPLOADS
+
+ logger = _override_logger()
+
+ assert logger._upload_semaphore._value == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
+ assert DEFAULT_S3_MAX_CONCURRENT_UPLOADS == 16
+
+
+@pytest.mark.asyncio
+async def test_a_slow_put_does_not_lower_the_adaptive_limit() -> None:
+ logger = _override_logger(s3_adaptive_concurrency=True)
+ logger.async_httpx_client = AsyncMock()
+
+ async def slow_put(url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ await _real_sleep(0)
+ return _ok_response()
+
+ logger.async_httpx_client.put = slow_put
+
+ before: Final = logger._upload_limiter.limit
+ logger.log_queue = [_element({"i": 0}, "0")]
+ await logger.async_send_batch()
+
+ assert logger._upload_limiter.limit >= before
+
+
+class _FastOkPut:
+ def __init__(self) -> None:
+ self.calls = 0
+
+ async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
+ self.calls += 1
+ await _real_sleep(0)
+ return _ok_response()
+
+
+async def _timed_send_batch(size: int) -> float:
+ logger = _override_logger()
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.put = _FastOkPut()
+ logger.log_queue = [_element({"i": i}, f"{i}") for i in range(size)]
+ started = time.perf_counter()
+ await logger.async_send_batch()
+ return time.perf_counter() - started
+
+
+@pytest.mark.asyncio
+async def test_send_batch_time_grows_linearly_with_the_batch() -> None:
+ baseline: Final = await _timed_send_batch(2_000)
+ quadrupled: Final = await _timed_send_batch(8_000)
+
+ assert quadrupled / baseline < 8, f"2k took {baseline:.3f}s, 8k took {quadrupled:.3f}s"
From 1474ea53e6e81bdfa6b991aea387ddbbd4a0f16a Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 15:07:56 -0700
Subject: [PATCH 13/39] feat(proxy): add
maximum_daily_tag_spend_retention_period cleanup setting (#39221)
* feat(proxy): add maximum_daily_tag_spend_retention_period cleanup setting
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* chore(ui): regenerate schema.d.ts for new retention setting
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* feat(proxy): rebase daily tag spend retention onto the run-budgeted cleanup job
Reworks the cleanup on top of the refactored SpendLogCleanup: the daily tag spend table is pruned through the shared batched delete with a text cutoff on the indexed ISO date column, the setting is picked up by /config/update and the scheduler registration, and an integration test proves rows older than the period are pruned while the cutoff day and unset retention are left alone
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): schedule the cleanup job when a retention db row lands before the side effects run
A config reload applies the db row to the SettingsStore before _update_general_settings snapshots the previous retention values, so the before/after compare saw no change and a retention period first set through /config/update never scheduled the cleanup job. Also reschedule when the job is missing but a retention period is set
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): accept list-valued top-level keys in the base integration proxy config
The shared tests/integration/proxy_config.yaml now carries list-valued top-level keys, so the retention config helper validates only the mapping it merges into. Also drops a SQL-shape assertion from the unit test in favor of the behavioral cutoff-day check
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): cover runtime update, invalid value, independent horizons and worker loss for daily tag spend retention
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): restore the shared retention setting, capture seeded days once and kill a listening worker
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): retry a failed cleanup schedule only when its settings change
_apply_retention_settings rescheduled whenever retention was set and no job existed, so an unparseable cleanup cron was retried on every config reload. Remember the last attempted retention, cron and interval tuple and retry only when it differs
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): give daily tag spend retention tests a 240s timeout
Each node boots a proxy and waits for a whole-minute cleanup cron tick, so the global 90s pytest-timeout can expire during teardown on a slow runner, as integration-accounting did on pipeline 90302
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): reschedule cleanup when only the cron or interval changes at runtime
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): reschedule cleanup when the first db sync changes only the cron or interval
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* feat(ui): add text input for String general settings so retention periods can be set from the Admin UI
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): record a cleanup schedule attempt only after it did not raise
Records _last_cleanup_schedule_attempt after _reschedule_spend_log_cleanup_job returns, so a transient add_job error is retried on the next config sync while an invalid cron, which is caught and logged inside the reschedule, is still attempted once per settings value
Also adds --num_workers 2 to the dev proxy command in AGENTS.md as requested on the PR
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* docs: revert unrelated AGENTS.md dev command change
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* Revert "docs: revert unrelated AGENTS.md dev command change"
This reverts commit 047706a623c8725935377fde76b40bc674a6dc12.
* Revert "fix(proxy): record a cleanup schedule attempt only after it did not raise"
This reverts commit 678f7c72b4b9d683d8aaad9c7f7473ca8f1eabe0.
* Revert "feat(ui): add text input for String general settings so retention periods can be set from the Admin UI"
This reverts commit 24e49d71d7673748ade9e4fe8e51b9d71da57427.
* fix(proxy): record a cleanup schedule attempt only after it did not raise
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): validate the cleanup schedule before swapping the job and leave startup registration to the startup block
_reschedule_spend_log_cleanup_job builds the new trigger first and only touches the live job once it parsed, so an invalid cron or interval (including a non string value) keeps the previous schedule running instead of removing it. An error raised while rescheduling is logged and retried on the next sync, so it no longer stops the rest of the general settings sync. _apply_retention_settings skips the job-missing path while the scheduler is still stopped, so the startup block is the only registration before start and the cross-replica stagger it applies to pending jobs survives
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): skip cleanup rescheduling while the scheduler is stopped and retry a failed replacement
The stopped-scheduler guard only covered the missing-job path, so the first DB sync (which runs before the startup block) still registered the cleanup job whenever the DB schedule differed from yaml, and startup then replaced it. Every runtime path now defers to the startup block while the scheduler is stopped.
A raised add_job that was replacing a live job was never retried because the live job kept wants_job == has_job; the sync now remembers the failure and retries on the next sync until the schedule is applied.
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* refactor(proxy): drop redundant docstring on _spend_log_cleanup_trigger
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): schedule DB-only retention at boot and log overflowing cleanup intervals once
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* refactor(proxy): drop explanatory comment from startup cleanup block
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): reject a non-string cleanup cron at startup and drop legacy covers markers
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): reschedule spend log cleanup when the reload path already applied a DB cron or interval edit
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(proxy): assert cleanup scheduling on a real paused scheduler instead of mock call counts
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: yucheng
---
litellm/proxy/_types.py | 8 +
.../db_transaction_queue/spend_log_cleanup.py | 72 +++-
litellm/proxy/proxy_server.py | 210 ++++++-----
.../spend/test_daily_tag_spend_retention.py | 247 +++++++++++++
.../config_resolvers/test_settings_rules.py | 1 +
.../proxy/proxy_server/test_proxy_config.py | 328 +++++++++++++++++-
tests/test_litellm/proxy/test_proxy_server.py | 83 +++++
.../proxy/test_spend_log_cleanup.py | 23 ++
ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 +
9 files changed, 884 insertions(+), 93 deletions(-)
create mode 100644 tests/integration/spend/test_daily_tag_spend_retention.py
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 0cf6d34bd6e..14aa42afefd 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -2983,6 +2983,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"Set this well above health_check_interval because /health and the UI read the latest row per model."
),
)
+ maximum_daily_tag_spend_retention_period: str | None = Field(
+ None,
+ description=(
+ "Maximum retention period for per-day tag spend aggregate rows (e.g., '90d'). Rows whose day is older "
+ "than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rows are never "
+ "deleted. Only historical tag usage analytics are affected; tag budgets read the lifetime counter."
+ ),
+ )
use_spend_logs_partitioning: bool | None = Field(
None,
description="If True and LiteLLM_SpendLogs has been converted to a range-partitioned table (db_scripts/partition_spend_logs.sql), retention cleanup drops expired partitions instead of deleting rows, and pre-creates upcoming partitions. Default is False.",
diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py
index c6f52bf074b..06e4d06fca4 100644
--- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py
+++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py
@@ -32,6 +32,17 @@ from litellm.proxy.utils import PrismaClient
StopReason: TypeAlias = Literal["exhausted", "budget_exhausted", "batch_cap_reached", "aborted"]
+Cutoff: TypeAlias = datetime | str
+"""Rows strictly older than this are expired: a timestamp, or an ISO calendar day for tables keyed by day"""
+
+
+def _cutoff_cast(cutoff: Cutoff) -> str:
+ return "timestamptz" if isinstance(cutoff, datetime) else "text"
+
+
+def _cutoff_text(cutoff: Cutoff) -> str:
+ return cutoff.isoformat() if isinstance(cutoff, datetime) else cutoff
+
@dataclass(frozen=True, slots=True)
class TableCleanupResult:
@@ -278,7 +289,7 @@ class SpendLogCleanup:
return remaining
async def _execute_delete_batch(
- self, prisma_client: PrismaClient, delete_sql: str, cutoff_date: datetime, deadline: float
+ self, prisma_client: PrismaClient, delete_sql: str, cutoff_date: Cutoff, deadline: float
) -> int | None:
"""
Run one delete batch under a Postgres statement and lock timeout.
@@ -301,7 +312,7 @@ class SpendLogCleanup:
return deleted_result if isinstance(deleted_result, int) else None
async def _count_remaining(
- self, prisma_client: PrismaClient, cutoff_date: datetime, table_name: str, time_column: str, deadline: float
+ self, prisma_client: PrismaClient, cutoff_date: Cutoff, table_name: str, time_column: str, deadline: float
) -> int | None:
"""
Count expired rows still outstanding, stopping at a cap.
@@ -314,7 +325,7 @@ class SpendLogCleanup:
count_sql: Final = f"""
SELECT count(*)::int AS remaining FROM (
SELECT 1 FROM "{table_name}"
- WHERE "{time_column}" < $1::timestamptz
+ WHERE "{time_column}" < $1::{_cutoff_cast(cutoff_date)}
LIMIT $2
) capped
"""
@@ -332,7 +343,7 @@ class SpendLogCleanup:
async def _delete_old_rows_batched(
self,
prisma_client: PrismaClient,
- cutoff_date: datetime,
+ cutoff_date: Cutoff,
table_name: str,
key_columns: tuple[str, ...],
time_column: str,
@@ -350,7 +361,7 @@ class SpendLogCleanup:
DELETE FROM "{table_name}"
WHERE ({key_list}) IN (
SELECT {key_list} FROM "{table_name}"
- WHERE "{time_column}" < $1::timestamptz
+ WHERE "{time_column}" < $1::{_cutoff_cast(cutoff_date)}
LIMIT $2
)
"""
@@ -406,7 +417,7 @@ class SpendLogCleanup:
run_count,
consecutive_failures,
self.batch_size,
- cutoff_date.isoformat(),
+ _cutoff_text(cutoff_date),
total_deleted,
type(batch_exc).__name__,
batch_exc,
@@ -454,7 +465,7 @@ class SpendLogCleanup:
async def _finish_table(
self,
prisma_client: PrismaClient,
- cutoff_date: datetime,
+ cutoff_date: Cutoff,
table_name: str,
time_column: str,
rows_deleted: int,
@@ -541,6 +552,18 @@ class SpendLogCleanup:
deadline=deadline,
)
+ async def _delete_old_daily_tag_spend_rows(
+ self, prisma_client: PrismaClient, cutoff_day: str, deadline: float
+ ) -> TableCleanupResult:
+ return await self._delete_old_rows_batched(
+ prisma_client,
+ cutoff_day,
+ table_name="LiteLLM_DailyTagSpend",
+ key_columns=("id",),
+ time_column="date",
+ deadline=deadline,
+ )
+
async def _clean_spend_log_tables(
self, prisma_client: PrismaClient, deadline: float
) -> tuple[TableCleanupResult, ...]:
@@ -624,6 +647,18 @@ class SpendLogCleanup:
)
return (health_checks_result,)
+ async def _clean_daily_tag_spend(
+ self, prisma_client: PrismaClient, retention_seconds: int, deadline: float
+ ) -> tuple[TableCleanupResult, ...]:
+ """
+ Prune per-day tag spend rows whose ISO day sorts before the horizon day; the horizon day itself is kept.
+ """
+ horizon: Final = datetime.now(timezone.utc) - timedelta(seconds=float(retention_seconds))
+ cutoff_day: Final = horizon.date().isoformat()
+ result: Final = await self._delete_old_daily_tag_spend_rows(prisma_client, cutoff_day, deadline)
+ verbose_proxy_logger.info("Deleted %s expired daily tag spend rows", result.rows_deleted)
+ return (result,)
+
@staticmethod
def _run_outcome(results: tuple[TableCleanupResult, ...]) -> RunOutcome:
"""
@@ -671,10 +706,14 @@ class SpendLogCleanup:
"maximum_autorouter_session_retention_period"
)
health_check_retention_seconds: Final = self._retention_seconds_for("maximum_health_check_retention_period")
+ daily_tag_spend_retention_seconds: Final = self._retention_seconds_for(
+ "maximum_daily_tag_spend_retention_period"
+ )
if (
not delete_spend_logs
and autorouter_retention_seconds is None
and health_check_retention_seconds is None
+ and daily_tag_spend_retention_seconds is None
):
SpendLogCleanupMetrics.record_run("skipped_disabled")
return
@@ -706,6 +745,7 @@ class SpendLogCleanup:
int(delete_spend_logs and self.retention_seconds is not None)
+ int(autorouter_retention_seconds is not None)
+ int(health_check_retention_seconds is not None)
+ + int(daily_tag_spend_retention_seconds is not None)
)
spend_log_results: Final = (
@@ -716,8 +756,13 @@ class SpendLogCleanup:
if delete_spend_logs and self.retention_seconds is not None
else ()
)
- remaining_groups_after_spend_logs: Final = int(autorouter_retention_seconds is not None) + int(
- health_check_retention_seconds is not None
+ remaining_groups_after_spend_logs: Final = (
+ int(autorouter_retention_seconds is not None)
+ + int(health_check_retention_seconds is not None)
+ + int(daily_tag_spend_retention_seconds is not None)
+ )
+ remaining_groups_after_sessions: Final = int(health_check_retention_seconds is not None) + int(
+ daily_tag_spend_retention_seconds is not None
)
session_results: Final = (
await self._clean_session_rollup(
@@ -732,13 +777,18 @@ class SpendLogCleanup:
await self._clean_health_checks(
prisma_client,
health_check_retention_seconds,
- deadline,
+ self._group_deadline(deadline, remaining_groups_after_sessions),
)
if health_check_retention_seconds is not None
else ()
)
+ daily_tag_spend_results: Final = (
+ await self._clean_daily_tag_spend(prisma_client, daily_tag_spend_retention_seconds, deadline)
+ if daily_tag_spend_retention_seconds is not None
+ else ()
+ )
- results: Final = spend_log_results + session_results + health_check_results
+ results: Final = spend_log_results + session_results + health_check_results + daily_tag_spend_results
outcome: Final = self._run_outcome(results)
SpendLogCleanupMetrics.record_run(outcome)
self._log_run_summary(outcome, results, time.monotonic() - run_started_at)
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index d2f9a4d7d93..f842f2e1e4a 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -213,6 +213,8 @@ try:
import orjson
import yaml
from apscheduler.schedulers.asyncio import AsyncIOScheduler
+ from apscheduler.schedulers.base import STATE_STOPPED
+ from apscheduler.triggers.base import BaseTrigger
from apscheduler.triggers.interval import IntervalTrigger
except ImportError as e:
raise ImportError(f"Missing dependency {e}. Run `pip install 'litellm[proxy]'`")
@@ -5078,6 +5080,20 @@ def _current_general_settings() -> Mapping[str, object]:
return general_settings
+_CLEANUP_SCHEDULE_KEYS: Final = (
+ "maximum_spend_logs_retention_period",
+ "maximum_autorouter_session_retention_period",
+ "maximum_health_check_retention_period",
+ "maximum_daily_tag_spend_retention_period",
+ "maximum_spend_logs_cleanup_cron",
+ "maximum_spend_logs_retention_interval",
+)
+
+
+def _cleanup_schedule_of(settings: Mapping[str, object]) -> tuple[object, ...]:
+ return tuple(settings.get(key) for key in _CLEANUP_SCHEDULE_KEYS)
+
+
@lru_cache(maxsize=4096)
def _log_ignored_cost_map_copy(model_id: str, fields: tuple[str, ...]) -> None:
verbose_proxy_logger.warning(
@@ -5100,6 +5116,8 @@ class ProxyConfig:
self._last_websearch_interception_config: dict[str, object] | None = None
self._last_hashicorp_vault_config: dict[str, object] | None = None
self._last_cyberark_config: dict[str, object] | None = None # mutable-ok: change-detection cache
+ self._last_cleanup_schedule_attempt: tuple[object, ...] | None = None
+ self._cleanup_reschedule_failed: bool = False
self._cyberark_boot_env: dict[str, str | None] | None = None # mutable-ok: deployment env snapshot, set once
self.worker_registry: list[WorkerRegistryEntry] = []
self.config_sync_subscriber: ConfigSyncSubscriber | None = None
@@ -7455,69 +7473,67 @@ class ProxyConfig:
if scheduler is None:
return
- # Remove existing job if it exists
- try:
- scheduler.remove_job("spend_log_cleanup_job")
- verbose_proxy_logger.info("Removed existing spend log cleanup job")
- except Exception:
- pass # Job might not exist, which is fine
-
- # Schedule new job if retention period is set (not None)
- retention_period: Final = general_settings.get("maximum_spend_logs_retention_period")
- autorouter_retention: Final = general_settings.get("maximum_autorouter_session_retention_period")
- health_check_retention: Final = general_settings.get("maximum_health_check_retention_period")
- if retention_period is not None or autorouter_retention is not None or health_check_retention is not None:
- from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import (
- SpendLogCleanup,
+ wants_job: Final = any(
+ general_settings.get(key) is not None
+ for key in (
+ "maximum_spend_logs_retention_period",
+ "maximum_autorouter_session_retention_period",
+ "maximum_health_check_retention_period",
+ "maximum_daily_tag_spend_retention_period",
)
+ )
+ if not wants_job:
+ if scheduler.get_job("spend_log_cleanup_job") is not None:
+ scheduler.remove_job("spend_log_cleanup_job")
+ verbose_proxy_logger.info("Removed existing spend log cleanup job")
+ return
- spend_log_cleanup: Final = SpendLogCleanup()
- cleanup_cron: Final = general_settings.get("maximum_spend_logs_cleanup_cron")
+ trigger: Final = self._spend_log_cleanup_trigger()
+ if trigger is None:
+ return
+ from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import (
+ SpendLogCleanup,
+ )
- if cleanup_cron:
- from apscheduler.triggers.cron import CronTrigger
+ scheduler.add_job(
+ SpendLogCleanup().cleanup_old_spend_logs,
+ trigger,
+ args=[prisma_client],
+ id="spend_log_cleanup_job",
+ replace_existing=True,
+ misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
+ )
+ verbose_proxy_logger.info("Spend log cleanup rescheduled with trigger: %s", trigger)
- try:
- cron_trigger: Final = CronTrigger.from_crontab(cleanup_cron)
- scheduler.add_job(
- spend_log_cleanup.cleanup_old_spend_logs,
- cron_trigger,
- args=[prisma_client],
- id="spend_log_cleanup_job",
- replace_existing=True,
- misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
- )
- verbose_proxy_logger.info("Spend log cleanup rescheduled with cron: %s", cleanup_cron)
- except ValueError:
- verbose_proxy_logger.error("Invalid maximum_spend_logs_cleanup_cron value: %s", cleanup_cron)
- else:
- # Interval-based scheduling (existing behavior)
- from litellm.litellm_core_utils.duration_parser import (
- duration_in_seconds,
- )
+ def _spend_log_cleanup_trigger(self) -> BaseTrigger | None:
+ cleanup_cron: Final[object] = general_settings.get("maximum_spend_logs_cleanup_cron")
+ if cleanup_cron:
+ from apscheduler.triggers.cron import CronTrigger
- retention_interval: Final = general_settings.get("maximum_spend_logs_retention_interval", "1d")
- try:
- interval_seconds: Final = duration_in_seconds(retention_interval)
- # this runs against a started scheduler, which the startup stagger sweep
- # cannot reach, so the offset is applied here or the job reconverges across
- # replicas the first time an admin edits the retention settings
- scheduler.add_job(
- spend_log_cleanup.cleanup_old_spend_logs,
- stagger_trigger(
- job_id="spend_log_cleanup_job",
- trigger=IntervalTrigger(seconds=interval_seconds),
- period_seconds=interval_seconds,
- settings=parse_stagger_settings(general_settings),
- ),
- args=[prisma_client],
- id="spend_log_cleanup_job",
- replace_existing=True,
- misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
- )
- verbose_proxy_logger.info("Spend log cleanup rescheduled with interval: %s", retention_interval)
- except ValueError:
- verbose_proxy_logger.error("Invalid maximum_spend_logs_retention_interval value")
+ try:
+ cron_trigger: Final[BaseTrigger] = CronTrigger.from_crontab(cleanup_cron)
+ except (ValueError, TypeError, AttributeError):
+ verbose_proxy_logger.error("Invalid maximum_spend_logs_cleanup_cron value: %s", cleanup_cron)
+ return None
+ return cron_trigger
+ retention_interval: Final[object] = general_settings.get("maximum_spend_logs_retention_interval", "1d")
+ if not isinstance(retention_interval, str):
+ verbose_proxy_logger.error("Invalid maximum_spend_logs_retention_interval value: %r", retention_interval)
+ return None
+ # this runs against a started scheduler, which the startup stagger sweep
+ # cannot reach, so the offset is applied here or the job reconverges across
+ # replicas the first time an admin edits the retention settings
+ try:
+ interval_seconds: Final = duration_in_seconds(retention_interval)
+ return stagger_trigger(
+ job_id="spend_log_cleanup_job",
+ trigger=IntervalTrigger(seconds=interval_seconds),
+ period_seconds=interval_seconds,
+ settings=parse_stagger_settings(general_settings),
+ )
+ except (ValueError, OverflowError):
+ verbose_proxy_logger.error("Invalid maximum_spend_logs_retention_interval value: %r", retention_interval)
+ return None
async def _update_general_settings(self, db_general_settings: Mapping[str, SettingsJsonValue] | None) -> None:
global general_settings
@@ -7526,32 +7542,28 @@ class ProxyConfig:
if not isinstance(general_settings, SettingsStore):
self.settings.load_yaml(_as_settings_mapping(general_settings))
cache_size_was_db: Final = self.settings.source("user_api_key_cache_max_size") == "db"
- previous_retention_values: Final = self._resolved_retention_values()
+ previous_cleanup_schedule: Final = self._resolved_cleanup_schedule()
previous_pass_through_endpoints: Final = self.settings.get("pass_through_endpoints")
self.settings.apply_db_row("general_settings", db_general_settings)
_bind_general_settings_store(self.settings)
await self._apply_general_settings_side_effects(
db_general_settings,
cache_size_was_db,
- previous_retention_values,
+ previous_cleanup_schedule,
previous_pass_through_endpoints,
)
- def _resolved_retention_values(self) -> tuple[SettingsJsonValue | None, ...]:
- return tuple(
- self.settings.get(key)
- for key in (
- "maximum_spend_logs_retention_period",
- "maximum_autorouter_session_retention_period",
- "maximum_health_check_retention_period",
- )
- )
+ def _resolved_cleanup_schedule(self) -> tuple[object, ...]:
+ return _cleanup_schedule_of(self.settings)
+
+ def record_cleanup_schedule_attempt(self, settings: Mapping[str, object]) -> None:
+ self._last_cleanup_schedule_attempt = _cleanup_schedule_of(settings)
async def _apply_general_settings_side_effects(
self,
db_values: Mapping[str, SettingsJsonValue],
cache_size_was_db: bool,
- previous_retention_values: tuple[SettingsJsonValue | None, ...],
+ previous_cleanup_schedule: tuple[object, ...],
previous_pass_through_endpoints: SettingsJsonValue | None,
) -> None:
effects: Final = (
@@ -7560,7 +7572,7 @@ class ProxyConfig:
self._apply_boolean_settings,
partial(self._apply_cache_size_setting, cache_size_was_db=cache_size_was_db),
self._apply_store_model_in_db_setting,
- partial(self._apply_retention_settings, previous_retention_values=previous_retention_values),
+ partial(self._apply_retention_settings, previous_cleanup_schedule=previous_cleanup_schedule),
self._apply_ssrf_settings,
)
for effect in effects:
@@ -7655,10 +7667,36 @@ class ProxyConfig:
async def _apply_retention_settings(
self,
db_values: Mapping[str, SettingsJsonValue],
- previous_retention_values: tuple[SettingsJsonValue | None, ...],
+ previous_cleanup_schedule: tuple[object, ...],
) -> None:
- if previous_retention_values != self._resolved_retention_values():
+ # while the scheduler is still stopped the startup block owns the first registration
+ if scheduler is not None and scheduler.state == STATE_STOPPED:
+ return
+ schedule: Final = self._resolved_cleanup_schedule()
+ wants_job: Final = any(value is not None for value in schedule[:4])
+ has_job: Final = scheduler is not None and scheduler.get_job("spend_log_cleanup_job") is not None
+ baseline: Final = (
+ self._last_cleanup_schedule_attempt
+ if has_job and self._last_cleanup_schedule_attempt is not None
+ else previous_cleanup_schedule
+ )
+ retry_due: Final = (
+ wants_job
+ and (not has_job or self._cleanup_reschedule_failed)
+ and schedule != self._last_cleanup_schedule_attempt
+ )
+ if not (baseline != schedule or retry_due or (has_job and not wants_job)):
+ return
+ try:
await self._reschedule_spend_log_cleanup_job()
+ except Exception as exc:
+ self._cleanup_reschedule_failed = True
+ verbose_proxy_logger.exception(
+ "Spend log cleanup could not be rescheduled, will retry on next sync: %s", exc
+ )
+ return
+ self._cleanup_reschedule_failed = False
+ self._last_cleanup_schedule_attempt = schedule
async def _apply_ssrf_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None:
_apply_ssrf_general_settings(db_values)
@@ -10535,15 +10573,19 @@ class ProxyStartupEvent:
)
### SPEND LOG CLEANUP ###
+ cleanup_settings: Final = _current_general_settings()
if (
- general_settings.get("maximum_spend_logs_retention_period") is not None
- or general_settings.get("maximum_autorouter_session_retention_period") is not None
- or general_settings.get("maximum_health_check_retention_period") is not None
+ cleanup_settings.get("maximum_spend_logs_retention_period") is not None
+ or cleanup_settings.get("maximum_autorouter_session_retention_period") is not None
+ or cleanup_settings.get("maximum_health_check_retention_period") is not None
+ or cleanup_settings.get("maximum_daily_tag_spend_retention_period") is not None
):
spend_log_cleanup: Final = SpendLogCleanup()
- cleanup_cron: Final = general_settings.get("maximum_spend_logs_cleanup_cron")
+ cleanup_cron: Final = cleanup_settings.get("maximum_spend_logs_cleanup_cron")
- if cleanup_cron:
+ if cleanup_cron and not isinstance(cleanup_cron, str):
+ verbose_proxy_logger.error("Invalid maximum_spend_logs_cleanup_cron value: %r", cleanup_cron)
+ elif isinstance(cleanup_cron, str) and cleanup_cron:
from apscheduler.triggers.cron import CronTrigger
try:
@@ -10561,8 +10603,10 @@ class ProxyStartupEvent:
verbose_proxy_logger.error("Invalid maximum_spend_logs_cleanup_cron value: %s", cleanup_cron)
else:
# Interval-based scheduling (existing behavior)
- retention_interval: Final = general_settings.get("maximum_spend_logs_retention_interval", "1d")
+ retention_interval: Final = cleanup_settings.get("maximum_spend_logs_retention_interval", "1d")
try:
+ if not isinstance(retention_interval, str):
+ raise ValueError(retention_interval)
interval_seconds: Final = duration_in_seconds(retention_interval)
scheduler.add_job(
spend_log_cleanup.cleanup_old_spend_logs,
@@ -10573,8 +10617,11 @@ class ProxyStartupEvent:
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
- except ValueError:
- verbose_proxy_logger.error("Invalid maximum_spend_logs_retention_interval value")
+ except (ValueError, OverflowError):
+ verbose_proxy_logger.error(
+ "Invalid maximum_spend_logs_retention_interval value: %r", retention_interval
+ )
+ proxy_config.record_cleanup_schedule_attempt(cleanup_settings)
### CHECK BATCH COST ###
if llm_router is not None and PROXY_BATCH_POLLING_ENABLED:
try:
@@ -17909,6 +17956,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"store_prompts_in_spend_logs": "Boolean",
"maximum_spend_logs_retention_period": "String",
"maximum_health_check_retention_period": "String",
+ "maximum_daily_tag_spend_retention_period": "String",
"maximum_spend_logs_cleanup_batch_size": "Integer",
"maximum_spend_logs_cleanup_max_batches": "Integer",
"maximum_spend_logs_cleanup_run_budget": "String",
diff --git a/tests/integration/spend/test_daily_tag_spend_retention.py b/tests/integration/spend/test_daily_tag_spend_retention.py
new file mode 100644
index 00000000000..fcb624cb188
--- /dev/null
+++ b/tests/integration/spend/test_daily_tag_spend_retention.py
@@ -0,0 +1,247 @@
+import json
+import os
+import signal
+import uuid
+from datetime import datetime, timedelta, timezone
+from pathlib import Path
+from typing import Final
+
+import psutil
+import psycopg
+import pytest
+import yaml
+from pydantic import JsonValue, TypeAdapter
+
+from tests.integration._support.client import Gateway, eventually, string_value
+from tests.integration._support.database import read_rows
+from tests.integration._support.process import OwnedProxy, owned_proxy, owned_proxy_process
+
+CLEANUP_EVERY_MINUTE: Final = "* * * * *"
+RETENTION_SETTING: Final = "maximum_daily_tag_spend_retention_period"
+_MAPPING: Final = TypeAdapter(dict[str, JsonValue])
+_SETTINGS: Final = TypeAdapter(list[dict[str, JsonValue]])
+
+
+def _day(days_ago: int) -> str:
+ return (datetime.now(timezone.utc) - timedelta(days=days_ago)).strftime("%Y-%m-%d")
+
+
+def _seed_daily_tag_spend(tag: str, days: tuple[str, ...]) -> None:
+ with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
+ for day in days:
+ connection.execute(
+ 'INSERT INTO "LiteLLM_DailyTagSpend" (id, tag, date, api_key, model, spend, updated_at) '
+ "VALUES (%s, %s, %s, %s, %s, 1.0, now())",
+ (uuid.uuid4().hex, tag, day, f"integration-{tag}", "gpt-4o-mini"),
+ )
+
+
+def _seed_old_spend_log(request_id: str, days_ago: int) -> None:
+ with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
+ connection.execute(
+ 'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, spend, "startTime", "endTime") '
+ "VALUES (%s, 'acompletion', %s, 0, now() - make_interval(days => %s), now() - make_interval(days => %s))",
+ (request_id, f"integration-{request_id}", str(days_ago), str(days_ago)),
+ )
+
+
+def _delete_daily_tag_spend(tag: str) -> None:
+ with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
+ connection.execute('DELETE FROM "LiteLLM_DailyTagSpend" WHERE tag = %s', (tag,))
+
+
+def _remaining_days(tag: str) -> tuple[str, ...]:
+ rows: Final = read_rows('SELECT date FROM "LiteLLM_DailyTagSpend" WHERE tag = %s ORDER BY date', (tag,))
+ return tuple(str(row["date"]) for row in rows)
+
+
+def _spend_log_present(request_id: str) -> bool:
+ return bool(read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (request_id,)))
+
+
+def _stored_retention_setting() -> JsonValue:
+ rows: Final = read_rows(
+ 'SELECT param_value -> %s AS value FROM "LiteLLM_Config" WHERE param_name = %s',
+ (RETENTION_SETTING, "general_settings"),
+ )
+ return rows[0]["value"] if rows else None
+
+
+def _store_retention_setting(value: JsonValue) -> None:
+ with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
+ if value is None:
+ connection.execute(
+ 'UPDATE "LiteLLM_Config" SET param_value = param_value - %s WHERE param_name = %s',
+ (RETENTION_SETTING, "general_settings"),
+ )
+ return
+ connection.execute(
+ 'UPDATE "LiteLLM_Config" SET param_value = jsonb_set(param_value, ARRAY[%s], %s::jsonb) '
+ "WHERE param_name = %s",
+ (RETENTION_SETTING, json.dumps(value), "general_settings"),
+ )
+
+
+def _listening_workers(owned: OwnedProxy) -> tuple[psutil.Process, ...]:
+ port: Final = owned.gateway.client.base_url.port
+ return tuple(
+ child
+ for child in psutil.Process(owned.process.pid).children(recursive=True)
+ if any(conn.status == psutil.CONN_LISTEN and conn.laddr.port == port for conn in child.net_connections("inet"))
+ )
+
+
+def _listed_retention_value(gateway: Gateway) -> JsonValue:
+ listed: Final = _SETTINGS.validate_json(
+ gateway.request("GET", "/config/list", params={"config_type": "general_settings"}).content
+ )
+ matching: Final = tuple(entry for entry in listed if entry["field_name"] == RETENTION_SETTING)
+ return matching[0]["field_value"] if matching else "not listed"
+
+
+def _completion_id(gateway: Gateway, model: str) -> str:
+ return string_value(gateway.chat(model, text=f"retention audit {uuid.uuid4().hex}")["id"])
+
+
+def _cleanup_config(tmp_path: Path, retention: dict[str, JsonValue]) -> Path:
+ base: Final = _MAPPING.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
+ config: Final = {
+ **base,
+ "general_settings": {
+ **_MAPPING.validate_python(base["general_settings"]),
+ **retention,
+ "maximum_spend_logs_cleanup_cron": CLEANUP_EVERY_MINUTE,
+ "scheduled_job_stagger": {"enabled": False},
+ },
+ }
+ path: Final = tmp_path / "retention.yaml"
+ path.write_text(yaml.safe_dump(config))
+ return path
+
+
+@pytest.mark.timeout(240)
+def test_daily_tag_spend_retention_prunes_only_rows_older_than_the_period(gateway: Gateway, tmp_path: Path) -> None:
+ tag: Final = f"integration-retention-{uuid.uuid4().hex}"
+ expired, on_the_cutoff, today = _day(200), _day(30), _day(0)
+ _seed_daily_tag_spend(tag, (expired, on_the_cutoff, today))
+ try:
+ config: Final = _cleanup_config(tmp_path, {"maximum_daily_tag_spend_retention_period": "30d"})
+ with owned_proxy(gateway, tmp_path, {}, config=config):
+ remaining: Final = eventually(
+ lambda: _remaining_days(tag),
+ lambda days: expired not in days,
+ seconds=150,
+ )
+ assert remaining == (on_the_cutoff, today), remaining
+ finally:
+ _delete_daily_tag_spend(tag)
+
+
+@pytest.mark.timeout(240)
+def test_config_update_turns_on_daily_tag_spend_cleanup_without_a_restart(gateway: Gateway, tmp_path: Path) -> None:
+ tag: Final = f"integration-retention-{uuid.uuid4().hex}"
+ expired, yesterday_of_cutoff, on_the_cutoff, today = _day(200), _day(31), _day(30), _day(0)
+ _seed_daily_tag_spend(tag, (expired, yesterday_of_cutoff, on_the_cutoff, today))
+ previously_stored: Final = _stored_retention_setting()
+ _store_retention_setting(None)
+ try:
+ config: Final = _cleanup_config(tmp_path, {})
+ with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as owned, owned.scenario() as scenario:
+ model: Final = scenario.model()
+ assert _listed_retention_value(owned) is None
+ owned.post("/config/update", {"general_settings": {RETENTION_SETTING: "30d"}})
+ assert _listed_retention_value(owned) == "30d"
+ remaining: Final = eventually(
+ lambda: _remaining_days(tag),
+ lambda days: yesterday_of_cutoff not in days,
+ seconds=150,
+ )
+ assert remaining == (on_the_cutoff, today), remaining
+ assert _completion_id(owned, model).startswith("chatcmpl-")
+ finally:
+ _store_retention_setting(previously_stored)
+ _delete_daily_tag_spend(tag)
+
+
+@pytest.mark.timeout(240)
+def test_unparseable_daily_tag_spend_retention_deletes_nothing_and_keeps_serving(
+ gateway: Gateway, tmp_path: Path
+) -> None:
+ tag: Final = f"integration-retention-{uuid.uuid4().hex}"
+ request_id: Final = f"integration-retention-{uuid.uuid4().hex}"
+ expired: Final = _day(200)
+ _seed_daily_tag_spend(tag, (expired,))
+ _seed_old_spend_log(request_id, days_ago=200)
+ try:
+ config: Final = _cleanup_config(
+ tmp_path, {RETENTION_SETTING: "soon", "maximum_spend_logs_retention_period": "30d"}
+ )
+ with owned_proxy(gateway, tmp_path, {}, config=config) as owned, owned.scenario() as scenario:
+ model: Final = scenario.model()
+ eventually(lambda: _spend_log_present(request_id), lambda present: not present, seconds=150)
+ assert _remaining_days(tag) == (expired,)
+ assert _completion_id(owned, model).startswith("chatcmpl-")
+ finally:
+ _delete_daily_tag_spend(tag)
+
+
+@pytest.mark.timeout(240)
+def test_daily_tag_spend_keeps_days_the_shorter_spend_log_horizon_already_pruned(
+ gateway: Gateway, tmp_path: Path
+) -> None:
+ tag: Final = f"integration-retention-{uuid.uuid4().hex}"
+ request_id: Final = f"integration-retention-{uuid.uuid4().hex}"
+ expired, inside_tag_horizon = _day(200), _day(60)
+ _seed_daily_tag_spend(tag, (expired, inside_tag_horizon))
+ _seed_old_spend_log(request_id, days_ago=60)
+ try:
+ config: Final = _cleanup_config(
+ tmp_path, {RETENTION_SETTING: "90d", "maximum_spend_logs_retention_period": "30d"}
+ )
+ with owned_proxy(gateway, tmp_path, {}, config=config):
+ eventually(lambda: _spend_log_present(request_id), lambda present: not present, seconds=150)
+ remaining: Final = eventually(lambda: _remaining_days(tag), lambda days: expired not in days, seconds=150)
+ assert remaining == (inside_tag_horizon,), remaining
+ finally:
+ _delete_daily_tag_spend(tag)
+
+
+@pytest.mark.timeout(240)
+def test_daily_tag_spend_cleanup_completes_after_one_of_two_workers_is_killed(gateway: Gateway, tmp_path: Path) -> None:
+ tag: Final = f"integration-retention-{uuid.uuid4().hex}"
+ expired, today = _day(200), _day(0)
+ _seed_daily_tag_spend(tag, (expired, today))
+ try:
+ config: Final = _cleanup_config(tmp_path, {RETENTION_SETTING: "30d"})
+ with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
+ with owned.gateway.scenario() as scenario:
+ model: Final = scenario.model()
+ workers: Final = eventually(
+ lambda: _listening_workers(owned), lambda found: len(found) == 2, seconds=30
+ )
+ workers[0].send_signal(signal.SIGKILL)
+ eventually(lambda: workers[0].is_running(), lambda alive: not alive, seconds=10)
+ ids: Final = tuple(_completion_id(owned.gateway, model) for _ in range(6))
+ assert len(set(ids)) == 6 and all(identity.startswith("chatcmpl-") for identity in ids), ids
+ remaining: Final = eventually(
+ lambda: _remaining_days(tag), lambda days: expired not in days, seconds=150
+ )
+ assert remaining == (today,), remaining
+ finally:
+ _delete_daily_tag_spend(tag)
+
+
+@pytest.mark.timeout(240)
+def test_daily_tag_spend_is_kept_forever_when_its_retention_is_unset(gateway: Gateway, tmp_path: Path) -> None:
+ tag: Final = f"integration-retention-{uuid.uuid4().hex}"
+ request_id: Final = f"integration-retention-{uuid.uuid4().hex}"
+ expired: Final = _day(200)
+ _seed_daily_tag_spend(tag, (expired,))
+ _seed_old_spend_log(request_id, days_ago=200)
+ try:
+ config: Final = _cleanup_config(tmp_path, {"maximum_spend_logs_retention_period": "30d"})
+ with owned_proxy(gateway, tmp_path, {}, config=config):
+ eventually(lambda: _spend_log_present(request_id), lambda present: not present, seconds=150)
+ assert _remaining_days(tag) == (expired,)
+ finally:
+ _delete_daily_tag_spend(tag)
diff --git a/tests/test_litellm/proxy/config_resolvers/test_settings_rules.py b/tests/test_litellm/proxy/config_resolvers/test_settings_rules.py
index ea5ebe6cf12..40e5870c804 100644
--- a/tests/test_litellm/proxy/config_resolvers/test_settings_rules.py
+++ b/tests/test_litellm/proxy/config_resolvers/test_settings_rules.py
@@ -79,6 +79,7 @@ _PREVIOUSLY_DB_WINS: Final[tuple[str, ...]] = (
"maximum_spend_logs_retention_period",
"maximum_autorouter_session_retention_period",
"maximum_health_check_retention_period",
+ "maximum_daily_tag_spend_retention_period",
"maximum_spend_logs_cleanup_batch_size",
"maximum_spend_logs_cleanup_max_batches",
"maximum_spend_logs_cleanup_run_budget",
diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
index b2ef327f50e..7378564f7a8 100644
--- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
+++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
@@ -3862,6 +3862,7 @@ async def test_ProxyConfig__reschedule_spend_log_cleanup_job_health_check_retent
async def test_ProxyConfig__update_general_settings_updates_health_check_retention(monkeypatch):
settings = {}
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", settings)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", MagicMock(**{"get_job.return_value": None}))
pc = ProxyConfig()
reschedule = AsyncMock()
monkeypatch.setattr(pc, "_reschedule_spend_log_cleanup_job", reschedule)
@@ -3872,6 +3873,329 @@ async def test_ProxyConfig__update_general_settings_updates_health_check_retenti
reschedule.assert_awaited_once()
+def _paused_scheduler(monkeypatch):
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ real_scheduler = AsyncIOScheduler()
+ real_scheduler.start(paused=True)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ return real_scheduler
+
+
+def _scheduler_whose_first_add_job_raises(monkeypatch):
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ class FirstAddJobRaises(AsyncIOScheduler):
+ raised = False
+
+ def add_job(self, *args, **kwargs):
+ if not self.raised:
+ self.raised = True
+ raise RuntimeError("scheduler busy")
+ return super().add_job(*args, **kwargs)
+
+ real_scheduler = FirstAddJobRaises()
+ real_scheduler.start(paused=True)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ return real_scheduler
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__reschedule_spend_log_cleanup_job_daily_tag_spend_retention(monkeypatch):
+ real_scheduler = _paused_scheduler(monkeypatch)
+ monkeypatch.setattr(
+ "litellm.proxy.proxy_server.general_settings",
+ {"maximum_daily_tag_spend_retention_period": "90d"},
+ )
+ pc = ProxyConfig()
+ try:
+ await pc._reschedule_spend_log_cleanup_job()
+ job = real_scheduler.get_job("spend_log_cleanup_job")
+ assert job is not None, "daily tag spend retention alone did not schedule the cleanup job"
+ assert job.func.__name__ == "cleanup_old_spend_logs"
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_updates_daily_tag_spend_retention(monkeypatch):
+ real_scheduler = _paused_scheduler(monkeypatch)
+ pc = ProxyConfig()
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d"})
+ from litellm.proxy import proxy_server
+
+ assert proxy_server.general_settings["maximum_daily_tag_spend_retention_period"] == "90d"
+ assert real_scheduler.get_job("spend_log_cleanup_job") is not None, "runtime retention did not schedule cleanup"
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_schedules_cleanup_when_db_row_was_already_applied(monkeypatch):
+ """A config reload applies the db row to the store before the side effects run, so the
+ before/after snapshot is equal; the job must still be scheduled when none is running."""
+ real_scheduler = _paused_scheduler(monkeypatch)
+ pc = ProxyConfig()
+ pc.settings.apply_db_row("general_settings", {"maximum_daily_tag_spend_retention_period": "90d"})
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d"})
+ assert real_scheduler.get_job("spend_log_cleanup_job") is not None, "DB-only retention never scheduled cleanup"
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_retries_a_failed_schedule_once_per_settings_value(
+ monkeypatch, caplog
+):
+ """An unparseable cron leaves no job behind; reloads must not retry it every tick, only when the
+ cron or a retention value changes."""
+ real_scheduler = _paused_scheduler(monkeypatch)
+ pc = ProxyConfig()
+ bad_cron = {"maximum_daily_tag_spend_retention_period": "90d", "maximum_spend_logs_cleanup_cron": "not a cron"}
+ pc.settings.apply_db_row("general_settings", bad_cron)
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
+ for _ in range(3):
+ await pc._update_general_settings(bad_cron)
+ assert real_scheduler.get_job("spend_log_cleanup_job") is None
+ cron_errors = [r for r in caplog.records if "maximum_spend_logs_cleanup_cron" in r.getMessage()]
+ assert len(cron_errors) == 1, f"invalid cron was retried on every reload: {len(cron_errors)} error lines"
+
+ await pc._update_general_settings({**bad_cron, "maximum_spend_logs_cleanup_cron": "* * * * *"})
+ job = real_scheduler.get_job("spend_log_cleanup_job")
+ assert job is not None, "a corrected cron did not schedule cleanup"
+ assert "minute='*'" in str(job.trigger)
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_retries_a_schedule_that_raised(monkeypatch):
+ """A transient add_job failure must not be remembered as a completed attempt; the next
+ reload with the same settings tries again."""
+ real_scheduler = _scheduler_whose_first_add_job_raises(monkeypatch)
+ pc = ProxyConfig()
+ retention = {"maximum_daily_tag_spend_retention_period": "90d"}
+ pc.settings.apply_db_row("general_settings", retention)
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ await pc._update_general_settings(retention)
+ assert real_scheduler.get_job("spend_log_cleanup_job") is None
+ await pc._update_general_settings(retention)
+ assert real_scheduler.get_job("spend_log_cleanup_job") is not None, "raised add_job was not retried"
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_retries_a_failed_replacement_of_the_live_job(monkeypatch):
+ """A cron change whose add_job raised keeps the old job running, so the next reload with the
+ same settings must try the replacement again instead of leaving the new cron unapplied."""
+ real_scheduler = _scheduler_whose_first_add_job_raises(monkeypatch)
+ pc = ProxyConfig()
+ pc.settings.load_yaml({"maximum_daily_tag_spend_retention_period": "90d"})
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ real_scheduler.raised = True
+ await pc._reschedule_spend_log_cleanup_job()
+ real_scheduler.raised = False
+ try:
+ new_cron = {"maximum_spend_logs_cleanup_cron": "0 3 * * *"}
+ await pc._update_general_settings(new_cron)
+ assert "hour='3'" not in str(real_scheduler.get_job("spend_log_cleanup_job").trigger), "old job was lost"
+ await pc._update_general_settings(new_cron)
+ assert "hour='3'" in str(real_scheduler.get_job("spend_log_cleanup_job").trigger), (
+ "failed replacement was not retried on the next sync"
+ )
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_leaves_a_changed_db_schedule_to_startup_while_scheduler_is_stopped(
+ monkeypatch,
+):
+ """The first DB sync runs before the scheduler starts and usually differs from the yaml; it
+ must still leave registration to the startup block instead of adding a job it will replace."""
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ real_scheduler = AsyncIOScheduler()
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ pc = ProxyConfig()
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d"})
+ assert real_scheduler.get_jobs() == [], "DB sync registered the cleanup job before the scheduler started"
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_leaves_first_registration_to_startup_while_scheduler_is_stopped(
+ monkeypatch,
+):
+ """The DB sync that runs before the scheduler starts must not register the cleanup job; the
+ startup block does, once, so the cross-replica stagger it applies to pending jobs survives."""
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ real_scheduler = AsyncIOScheduler()
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ pc = ProxyConfig()
+ pc.settings.load_yaml({"maximum_daily_tag_spend_retention_period": "90d"})
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ await pc._update_general_settings({"unrelated_key": "value"})
+ assert real_scheduler.get_jobs() == [], "DB sync registered the cleanup job before the scheduler started"
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_runtime_interval_job_carries_the_stagger_offset(monkeypatch):
+ """Once the scheduler is running the sync owns registration and the job it adds is staggered."""
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ from litellm.proxy.common_utils.scheduled_job_stagger import _OffsetTrigger
+
+ real_scheduler = AsyncIOScheduler()
+ real_scheduler.start(paused=True)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ pc = ProxyConfig()
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d"})
+ jobs = real_scheduler.get_jobs()
+ assert [job.id for job in jobs] == ["spend_log_cleanup_job"]
+ assert isinstance(jobs[0].trigger, _OffsetTrigger), repr(jobs[0].trigger)
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "bad_schedule",
+ [
+ {"maximum_spend_logs_cleanup_cron": "not a cron"},
+ {"maximum_spend_logs_cleanup_cron": "0 0 * * * *"},
+ {"maximum_spend_logs_retention_interval": "soon"},
+ {"maximum_spend_logs_retention_interval": 86400},
+ ],
+)
+async def test_ProxyConfig__update_general_settings_keeps_the_live_cleanup_job_when_the_new_schedule_is_invalid(
+ monkeypatch, bad_schedule
+):
+ """A schedule edit that does not parse must leave the old cleanup job running and must not
+ stop the rest of the general settings sync."""
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ real_scheduler = AsyncIOScheduler()
+ real_scheduler.start(paused=True)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ ssrf_sync = MagicMock()
+ monkeypatch.setattr("litellm.proxy.proxy_server._apply_ssrf_general_settings", ssrf_sync)
+ pc = ProxyConfig()
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d"})
+ old_trigger = real_scheduler.get_job("spend_log_cleanup_job").trigger
+ ssrf_sync.reset_mock()
+ for _ in range(2):
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d", **bad_schedule})
+ live_job = real_scheduler.get_job("spend_log_cleanup_job")
+ assert live_job is not None, "invalid schedule removed the cleanup job"
+ assert live_job.trigger is old_trigger
+ assert ssrf_sync.call_count == 2, "schedule error blocked the rest of the settings sync"
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_logs_an_overflowing_interval_once(monkeypatch, caplog):
+ """An interval that parses but overflows the trigger must keep the live job and log one
+ error, not a traceback on every sync."""
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ real_scheduler = AsyncIOScheduler()
+ real_scheduler.start(paused=True)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ pc = ProxyConfig()
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ await pc._update_general_settings({"maximum_daily_tag_spend_retention_period": "90d"})
+ old_trigger = real_scheduler.get_job("spend_log_cleanup_job").trigger
+ overflowing = {
+ "maximum_daily_tag_spend_retention_period": "90d",
+ "maximum_spend_logs_retention_interval": "99999999999d",
+ }
+ with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
+ for _ in range(5):
+ await pc._update_general_settings(overflowing)
+ errors = [record for record in caplog.records if record.levelno >= logging.ERROR]
+ assert len(errors) == 1, [record.getMessage() for record in errors]
+ assert real_scheduler.get_job("spend_log_cleanup_job").trigger is old_trigger
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_reschedules_when_only_the_cron_changes(monkeypatch):
+ real_scheduler = _paused_scheduler(monkeypatch)
+ pc = ProxyConfig()
+ pc.settings.load_yaml({"maximum_daily_tag_spend_retention_period": "90d"})
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ await pc._reschedule_spend_log_cleanup_job()
+ try:
+ interval_job = real_scheduler.get_job("spend_log_cleanup_job")
+ assert interval_job is not None and "hour='3'" not in str(interval_job.trigger)
+
+ await pc._update_general_settings({"maximum_spend_logs_cleanup_cron": "0 3 * * *"})
+ cron_job = real_scheduler.get_job("spend_log_cleanup_job")
+ assert "hour='3'" in str(cron_job.trigger), "cron-only change did not reschedule"
+
+ await pc._update_general_settings({"maximum_spend_logs_cleanup_cron": "0 3 * * *"})
+ assert real_scheduler.get_job("spend_log_cleanup_job") is cron_job, "unchanged cron replaced the job"
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_ProxyConfig__update_general_settings_reschedules_a_cron_edit_the_reload_path_already_applied(
+ monkeypatch,
+):
+ """The periodic reload applies the DB row through _update_config_from_db before
+ _update_general_settings snapshots the previous schedule, so a cron edited in the DB must
+ still replace the live job's trigger."""
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ real_scheduler = AsyncIOScheduler()
+ real_scheduler.start(paused=True)
+ monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", real_scheduler)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
+ pc = ProxyConfig()
+ monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", pc.settings)
+ try:
+ first_row = {"maximum_daily_tag_spend_retention_period": "90d", "maximum_spend_logs_cleanup_cron": "0 3 * * *"}
+ pc.settings.apply_db_row("general_settings", first_row)
+ await pc._update_general_settings(first_row)
+ assert "hour='3'" in str(real_scheduler.get_job("spend_log_cleanup_job").trigger)
+
+ edited_row = {**first_row, "maximum_spend_logs_cleanup_cron": "0 5 * * *"}
+ pc.settings.apply_db_row("general_settings", edited_row)
+ await pc._update_general_settings(edited_row)
+ assert "hour='5'" in str(real_scheduler.get_job("spend_log_cleanup_job").trigger), "DB cron edit was ignored"
+
+ pc.settings.apply_db_row("general_settings", edited_row)
+ await pc._update_general_settings(edited_row)
+ assert "hour='5'" in str(real_scheduler.get_job("spend_log_cleanup_job").trigger)
+ finally:
+ real_scheduler.shutdown(wait=False)
+
+
# ---------------------------------------------------------------------------
# ProxyConfig._update_general_settings
# ---------------------------------------------------------------------------
@@ -4003,6 +4327,7 @@ async def test_ProxyConfig__update_general_settings_skips_redundant_retention_re
pc = ProxyConfig()
reschedule: Final = AsyncMock()
monkeypatch.setattr(proxy_server, "general_settings", {})
+ monkeypatch.setattr(proxy_server, "scheduler", MagicMock())
monkeypatch.setattr(pc, "_reschedule_spend_log_cleanup_job", reschedule)
await pc._update_general_settings({"maximum_health_check_retention_period": "30d"})
@@ -4021,6 +4346,7 @@ async def test_ProxyConfig__update_general_settings_reschedules_after_retention_
pc = ProxyConfig()
reschedule: Final = AsyncMock()
monkeypatch.setattr(proxy_server, "general_settings", {})
+ monkeypatch.setattr(proxy_server, "scheduler", MagicMock(**{"get_job.return_value": None}))
monkeypatch.setattr(pc, "_reschedule_spend_log_cleanup_job", reschedule)
await pc._update_general_settings({"maximum_health_check_retention_period": "30d"})
@@ -4052,7 +4378,7 @@ async def test_ProxyConfig__update_general_settings_dispatches_every_side_effect
if name == "_apply_cache_size_setting":
handler.assert_awaited_once_with({}, cache_size_was_db=False)
elif name == "_apply_retention_settings":
- handler.assert_awaited_once_with({}, previous_retention_values=())
+ handler.assert_awaited_once_with({}, previous_cleanup_schedule=())
elif name == "_apply_pass_through_settings":
handler.assert_awaited_once_with({}, previous_endpoints=None)
else:
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index 89156cd19a0..df8feb74305 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -935,6 +935,89 @@ async def test_periodic_reload_job_scheduled_without_store_model_in_db(monkeypat
scheduler.shutdown(wait=False)
+@pytest.mark.asyncio
+async def test_initialize_scheduled_jobs_registers_cleanup_when_retention_lives_only_in_the_db(monkeypatch):
+ """With no config file, the startup DB sync rebinds general_settings to a store holding the
+ retention period; the cleanup job must be registered from that live value, not the stale
+ empty dict the caller passed in."""
+ monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
+ monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ from litellm.proxy.proxy_server import ProxyStartupEvent
+ from litellm.proxy.utils import ProxyLogging
+
+ mock_prisma_client = MagicMock()
+ mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
+ mock_proxy_logging = MagicMock(spec=ProxyLogging)
+ mock_proxy_logging.slack_alerting_instance = MagicMock()
+ mock_proxy_logging.db_spend_update_writer = MagicMock()
+ mock_proxy_config = _mock_scheduled_proxy_config()
+ db_settings = proxy_server_module.ProxyConfig().settings
+ db_settings.apply_db_row("general_settings", {"maximum_daily_tag_spend_retention_period": "30d"})
+
+ async def sync_from_db(*args: object, **kwargs: object) -> None:
+ proxy_server_module._bind_general_settings_store(db_settings)
+
+ mock_proxy_config.add_deployment.side_effect = sync_from_db
+ scheduler = AsyncIOScheduler()
+ try:
+ with (
+ patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
+ patch("litellm.proxy.proxy_server.store_model_in_db", True),
+ patch("litellm.proxy.proxy_server.general_settings", {}),
+ patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=scheduler),
+ ):
+ await ProxyStartupEvent.initialize_scheduled_background_jobs(
+ general_settings={},
+ prisma_client=mock_prisma_client,
+ proxy_budget_rescheduler_min_time=1,
+ proxy_budget_rescheduler_max_time=2,
+ proxy_batch_write_at=5,
+ proxy_logging_obj=mock_proxy_logging,
+ )
+ assert scheduler.get_job("spend_log_cleanup_job") is not None, "DB-only retention was not scheduled at boot"
+ finally:
+ scheduler.shutdown(wait=False)
+
+
+@pytest.mark.asyncio
+async def test_initialize_scheduled_jobs_does_not_fall_back_to_the_interval_for_a_non_string_cron(monkeypatch):
+ """A truthy non-string cron is invalid, so startup must log it and register no cleanup job
+ rather than silently pruning on the default interval the admin never configured."""
+ monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
+ monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
+ from apscheduler.schedulers.asyncio import AsyncIOScheduler
+
+ from litellm.proxy.proxy_server import ProxyStartupEvent
+ from litellm.proxy.utils import ProxyLogging
+
+ mock_prisma_client = MagicMock()
+ mock_proxy_logging = MagicMock(spec=ProxyLogging)
+ mock_proxy_logging.slack_alerting_instance = MagicMock()
+ mock_proxy_logging.db_spend_update_writer = MagicMock()
+ settings = {"maximum_daily_tag_spend_retention_period": "30d", "maximum_spend_logs_cleanup_cron": 5}
+ scheduler = AsyncIOScheduler()
+ try:
+ with (
+ patch("litellm.proxy.proxy_server.proxy_config", _mock_scheduled_proxy_config()),
+ patch("litellm.proxy.proxy_server.store_model_in_db", False),
+ patch("litellm.proxy.proxy_server.general_settings", settings),
+ patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=scheduler),
+ ):
+ await ProxyStartupEvent.initialize_scheduled_background_jobs(
+ general_settings=settings,
+ prisma_client=mock_prisma_client,
+ proxy_budget_rescheduler_min_time=1,
+ proxy_budget_rescheduler_max_time=2,
+ proxy_batch_write_at=5,
+ proxy_logging_obj=mock_proxy_logging,
+ )
+ assert scheduler.get_job("spend_log_cleanup_job") is None, "invalid cron fell back to the interval"
+ finally:
+ scheduler.shutdown(wait=False)
+
+
@pytest.mark.asyncio
async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(monkeypatch):
"""
diff --git a/tests/test_litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py
index 72463e17c6b..46ac1234615 100644
--- a/tests/test_litellm/proxy/test_spend_log_cleanup.py
+++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py
@@ -827,6 +827,29 @@ async def test_health_check_retention_alone_cleans_only_the_health_check_table()
assert abs((cutoff_date - expected_cutoff).total_seconds()) < 1
+@pytest.mark.asyncio
+async def test_daily_tag_spend_retention_alone_prunes_only_that_table_by_calendar_day():
+ client = _mock_prisma_for_retention([0])
+ cleaner = SpendLogCleanup(general_settings={"maximum_daily_tag_spend_retention_period": "90d"})
+ cleaner.pod_lock_manager = None
+ await cleaner.cleanup_old_spend_logs(client)
+ tables = [call[0][0] for call in client.db.execute_raw.call_args_list]
+ assert len(tables) == 1
+ assert '"LiteLLM_DailyTagSpend"' in tables[0]
+ cutoff_day = client.db.execute_raw.call_args[0][1]
+ assert cutoff_day == (datetime.now(timezone.utc) - timedelta(days=90)).date().isoformat()
+
+
+@pytest.mark.asyncio
+async def test_spend_logs_retention_alone_keeps_daily_tag_spend_forever():
+ client = _mock_prisma_for_retention([0, 0])
+ cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"})
+ cleaner.pod_lock_manager = None
+ await cleaner.cleanup_old_spend_logs(client)
+ tables = [call[0][0] for call in client.db.execute_raw.call_args_list]
+ assert not any('"LiteLLM_DailyTagSpend"' in sql for sql in tables)
+
+
@pytest.mark.asyncio
async def test_each_retention_key_cuts_off_at_its_own_horizon():
client = _mock_prisma_for_retention([0, 0, 0, 0, 0])
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index c77f7d84bd8..0500aeb95c8 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -28325,6 +28325,11 @@ export interface components {
* @description Maximum retention period for auto-router benchmark session rollup rows (e.g., '365d'). Rows whose last turn is older than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rollup rows are never deleted.
*/
maximum_autorouter_session_retention_period?: string | null;
+ /**
+ * Maximum Daily Tag Spend Retention Period
+ * @description Maximum retention period for per-day tag spend aggregate rows (e.g., '90d'). Rows whose day is older than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rows are never deleted. Only historical tag usage analytics are affected; tag budgets read the lifetime counter.
+ */
+ maximum_daily_tag_spend_retention_period?: string | null;
/**
* Maximum Health Check Retention Period
* @description Maximum retention period for health-check rows (e.g., '30d'). Rows whose checked_at is older than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rows are never deleted. Set this well above health_check_interval because /health and the UI read the latest row per model.
From 89061aa1f248f572968140a3e3a22268c20c7fa2 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 15:10:10 -0700
Subject: [PATCH 14/39] fix(cost-map): sync OpenRouter, Together, Cohere and
Azure AI registry values with official sources (#43337)
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
...odel_prices_and_context_window_backup.json | 114 ++++++++++++------
model_prices_and_context_window.json | 114 ++++++++++++------
2 files changed, 148 insertions(+), 80 deletions(-)
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index 82bb1c84dbe..8f61dc91adf 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -12216,7 +12216,8 @@
"text",
"image"
],
- "supports_embedding_image_input": true
+ "supports_embedding_image_input": true,
+ "input_cost_per_image_token": 4.7e-07
},
"azure_ai/grok-4": {
"input_cost_per_token": 3e-06,
@@ -15804,7 +15805,9 @@
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1536,
- "supports_embedding_image_input": true
+ "supports_embedding_image_input": true,
+ "input_cost_per_image_token": 4.7e-07,
+ "source": "https://cohere.com/pricing"
},
"cohere/parse-v5.0": {
"litellm_provider": "cohere",
@@ -41913,14 +41916,14 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
- "cache_read_input_token_cost": 3.828e-08,
- "input_cost_per_token": 4.5936e-07,
+ "cache_read_input_token_cost": 2.9e-08,
+ "input_cost_per_token": 3.48e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
- "output_cost_per_token": 9.1872e-07,
+ "output_cost_per_token": 6.96e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@@ -41933,14 +41936,14 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4.1-flash": {
- "cache_read_input_token_cost": 4.2e-09,
- "input_cost_per_token": 1.4e-07,
+ "cache_read_input_token_cost": 1e-09,
+ "input_cost_per_token": 3.5e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
- "max_output_tokens": 393216,
- "max_tokens": 393216,
+ "max_output_tokens": 384000,
+ "max_tokens": 384000,
"mode": "chat",
- "output_cost_per_token": 4.2e-07,
+ "output_cost_per_token": 2.9e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@@ -43052,7 +43055,6 @@
"supports_web_search": false
},
"openrouter/openai/gpt-oss-20b": {
- "cache_read_input_token_cost": 3e-08,
"input_cost_per_token": 1.8e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 131072,
@@ -43322,8 +43324,8 @@
"input_cost_per_token": 2.6e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 65536,
- "max_tokens": 65536,
+ "max_output_tokens": 235929,
+ "max_tokens": 235929,
"mode": "chat",
"output_cost_per_token": 2.08e-06,
"source": "https://openrouter.ai/api/v1/models",
@@ -43603,8 +43605,8 @@
"input_cost_per_token": 9.646e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 204800,
- "max_output_tokens": 128000,
- "max_tokens": 128000,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 3.0316e-06,
"source": "https://openrouter.ai/api/v1/models",
@@ -66825,14 +66827,14 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3-flash": {
- "cache_read_input_token_cost": 1e-08,
- "input_cost_per_token": 4.5e-08,
+ "cache_read_input_token_cost": 1.5e-08,
+ "input_cost_per_token": 4e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
- "max_output_tokens": 943718,
- "max_tokens": 943718,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 1.4e-07,
+ "output_cost_per_token": 5e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@@ -66865,13 +66867,13 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3": {
- "input_cost_per_token": 1.4e-06,
- "output_cost_per_token": 4.4e-06,
- "cache_read_input_token_cost": 2.6e-07,
+ "input_cost_per_token": 3.794e-07,
+ "output_cost_per_token": 1.1924e-06,
+ "cache_read_input_token_cost": 7.046e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
- "max_output_tokens": 943717,
- "max_tokens": 943717,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
@@ -67520,8 +67522,8 @@
"input_cost_per_token": 3.2e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 262140,
- "max_tokens": 262140,
+ "max_output_tokens": 81920,
+ "max_tokens": 81920,
"mode": "chat",
"output_cost_per_token": 3.2e-06,
"source": "https://openrouter.ai/api/v1/models",
@@ -68297,8 +68299,8 @@
"input_cost_per_token": 1.5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 32768,
- "max_tokens": 32768,
+ "max_output_tokens": 16384,
+ "max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 6e-07,
"source": "https://openrouter.ai/api/v1/models",
@@ -68476,8 +68478,8 @@
"cache_read_input_token_cost": 7e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
+ "max_output_tokens": 235929,
+ "max_tokens": 235929,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
@@ -68637,8 +68639,8 @@
"output_cost_per_token": 3e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 32000,
- "max_tokens": 32000,
+ "max_output_tokens": 235929,
+ "max_tokens": 235929,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
@@ -69502,7 +69504,9 @@
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 1.5e-07,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 128000,
+ "max_tokens": 128000
},
"vertex_ai/gemini-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
@@ -69613,42 +69617,72 @@
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 3.5e-06,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 4096,
+ "max_tokens": 4096
},
"together_ai/meta-llama/Llama-3.2-1B-Instruct": {
"input_cost_per_token": 6e-08,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 6e-08,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 131072,
+ "max_tokens": 131072
},
"together_ai/meta-llama/Llama-3.2-3B-Instruct": {
"input_cost_per_token": 6e-08,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 6e-08,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 131072,
+ "max_tokens": 131072
},
"together_ai/Qwen/Qwen2-1.5B-Instruct": {
"input_cost_per_token": 2e-08,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 2e-08,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 32768,
+ "max_tokens": 32768
},
"together_ai/Qwen/Qwen2.5-14B-Instruct": {
"input_cost_per_token": 8e-07,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 8e-07,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 32768,
+ "max_tokens": 32768
},
"together_ai/Qwen/Qwen2.5-72B-Instruct": {
"input_cost_per_token": 1.2e-06,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 1.2e-06,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 32768,
+ "max_tokens": 32768
+ },
+ "together_ai/Salesforce/Llama-Rank-V1": {
+ "input_cost_per_token": 1e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "rerank",
+ "output_cost_per_token": 0.0,
+ "source": "https://api.together.xyz/v1/models"
+ },
+ "together_ai/meta-llama/Meta-Llama-3.1-8B": {
+ "input_cost_per_token": 2e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 16384,
+ "max_tokens": 16384,
+ "mode": "completion",
+ "output_cost_per_token": 2e-07,
+ "source": "https://api.together.xyz/v1/models"
},
"together_ai/together/Tev1-4B-experimental": {
"cache_read_input_token_cost": 4.2e-08,
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index 82bb1c84dbe..8f61dc91adf 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -12216,7 +12216,8 @@
"text",
"image"
],
- "supports_embedding_image_input": true
+ "supports_embedding_image_input": true,
+ "input_cost_per_image_token": 4.7e-07
},
"azure_ai/grok-4": {
"input_cost_per_token": 3e-06,
@@ -15804,7 +15805,9 @@
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1536,
- "supports_embedding_image_input": true
+ "supports_embedding_image_input": true,
+ "input_cost_per_image_token": 4.7e-07,
+ "source": "https://cohere.com/pricing"
},
"cohere/parse-v5.0": {
"litellm_provider": "cohere",
@@ -41913,14 +41916,14 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
- "cache_read_input_token_cost": 3.828e-08,
- "input_cost_per_token": 4.5936e-07,
+ "cache_read_input_token_cost": 2.9e-08,
+ "input_cost_per_token": 3.48e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
- "output_cost_per_token": 9.1872e-07,
+ "output_cost_per_token": 6.96e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@@ -41933,14 +41936,14 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4.1-flash": {
- "cache_read_input_token_cost": 4.2e-09,
- "input_cost_per_token": 1.4e-07,
+ "cache_read_input_token_cost": 1e-09,
+ "input_cost_per_token": 3.5e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
- "max_output_tokens": 393216,
- "max_tokens": 393216,
+ "max_output_tokens": 384000,
+ "max_tokens": 384000,
"mode": "chat",
- "output_cost_per_token": 4.2e-07,
+ "output_cost_per_token": 2.9e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@@ -43052,7 +43055,6 @@
"supports_web_search": false
},
"openrouter/openai/gpt-oss-20b": {
- "cache_read_input_token_cost": 3e-08,
"input_cost_per_token": 1.8e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 131072,
@@ -43322,8 +43324,8 @@
"input_cost_per_token": 2.6e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 65536,
- "max_tokens": 65536,
+ "max_output_tokens": 235929,
+ "max_tokens": 235929,
"mode": "chat",
"output_cost_per_token": 2.08e-06,
"source": "https://openrouter.ai/api/v1/models",
@@ -43603,8 +43605,8 @@
"input_cost_per_token": 9.646e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 204800,
- "max_output_tokens": 128000,
- "max_tokens": 128000,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 3.0316e-06,
"source": "https://openrouter.ai/api/v1/models",
@@ -66825,14 +66827,14 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3-flash": {
- "cache_read_input_token_cost": 1e-08,
- "input_cost_per_token": 4.5e-08,
+ "cache_read_input_token_cost": 1.5e-08,
+ "input_cost_per_token": 4e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
- "max_output_tokens": 943718,
- "max_tokens": 943718,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 1.4e-07,
+ "output_cost_per_token": 5e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@@ -66865,13 +66867,13 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3": {
- "input_cost_per_token": 1.4e-06,
- "output_cost_per_token": 4.4e-06,
- "cache_read_input_token_cost": 2.6e-07,
+ "input_cost_per_token": 3.794e-07,
+ "output_cost_per_token": 1.1924e-06,
+ "cache_read_input_token_cost": 7.046e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
- "max_output_tokens": 943717,
- "max_tokens": 943717,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
@@ -67520,8 +67522,8 @@
"input_cost_per_token": 3.2e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 262140,
- "max_tokens": 262140,
+ "max_output_tokens": 81920,
+ "max_tokens": 81920,
"mode": "chat",
"output_cost_per_token": 3.2e-06,
"source": "https://openrouter.ai/api/v1/models",
@@ -68297,8 +68299,8 @@
"input_cost_per_token": 1.5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 32768,
- "max_tokens": 32768,
+ "max_output_tokens": 16384,
+ "max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 6e-07,
"source": "https://openrouter.ai/api/v1/models",
@@ -68476,8 +68478,8 @@
"cache_read_input_token_cost": 7e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
+ "max_output_tokens": 235929,
+ "max_tokens": 235929,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
@@ -68637,8 +68639,8 @@
"output_cost_per_token": 3e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
- "max_output_tokens": 32000,
- "max_tokens": 32000,
+ "max_output_tokens": 235929,
+ "max_tokens": 235929,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
@@ -69502,7 +69504,9 @@
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 1.5e-07,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 128000,
+ "max_tokens": 128000
},
"vertex_ai/gemini-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
@@ -69613,42 +69617,72 @@
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 3.5e-06,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 4096,
+ "max_tokens": 4096
},
"together_ai/meta-llama/Llama-3.2-1B-Instruct": {
"input_cost_per_token": 6e-08,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 6e-08,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 131072,
+ "max_tokens": 131072
},
"together_ai/meta-llama/Llama-3.2-3B-Instruct": {
"input_cost_per_token": 6e-08,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 6e-08,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 131072,
+ "max_tokens": 131072
},
"together_ai/Qwen/Qwen2-1.5B-Instruct": {
"input_cost_per_token": 2e-08,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 2e-08,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 32768,
+ "max_tokens": 32768
},
"together_ai/Qwen/Qwen2.5-14B-Instruct": {
"input_cost_per_token": 8e-07,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 8e-07,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 32768,
+ "max_tokens": 32768
},
"together_ai/Qwen/Qwen2.5-72B-Instruct": {
"input_cost_per_token": 1.2e-06,
"litellm_provider": "together_ai",
"mode": "chat",
"output_cost_per_token": 1.2e-06,
- "source": "https://api.together.ai/v1/models"
+ "source": "https://api.together.ai/v1/models",
+ "max_input_tokens": 32768,
+ "max_tokens": 32768
+ },
+ "together_ai/Salesforce/Llama-Rank-V1": {
+ "input_cost_per_token": 1e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "rerank",
+ "output_cost_per_token": 0.0,
+ "source": "https://api.together.xyz/v1/models"
+ },
+ "together_ai/meta-llama/Meta-Llama-3.1-8B": {
+ "input_cost_per_token": 2e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 16384,
+ "max_tokens": 16384,
+ "mode": "completion",
+ "output_cost_per_token": 2e-07,
+ "source": "https://api.together.xyz/v1/models"
},
"together_ai/together/Tev1-4B-experimental": {
"cache_read_input_token_cost": 4.2e-08,
From f8870b64e9a654085c586a28b1c7b4050ceb775e Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 15:14:14 -0700
Subject: [PATCH 15/39] docs(pr-template): add the backport-stable label only
for a P0 regression (#43351)
* docs(pr-template): add the backport-stable label only for a P0 regression
* docs(pr-template): keep a narrow security regression eligible for backport-stable
---------
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
---
.github/pull_request_template.md | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md
index db46114715d..6beb6e99e0e 100644
--- a/.github/pull_request_template.md
+++ b/.github/pull_request_template.md
@@ -54,7 +54,7 @@ After: the same request comes back with real token counts, so the dashboard show
## Affected release
-
+
## Linear ticket
From 9413b82477be61af07a41cdd70694c738111df25 Mon Sep 17 00:00:00 2001
From: shrey-berri
Date: Sat, 26 Sep 2026 15:20:09 -0700
Subject: [PATCH 16/39] fix(params): keep _litellm_* kwargs out of provider
request bodies by construction (#43221)
Kwargs LiteLLM code introduces for its own use were only kept out of provider
bodies if someone also listed them in all_litellm_params. Undeclared ones went
into extra_body or optional_params, reached the provider, and the provider
rejected the request. is_litellm_owned_kwarg in types/utils.py now defines
LiteLLM-owned once: a registered name, or any name starting with
INTERNAL_KWARG_PREFIX from litellm/constants.py. Every filter that builds
provider params from kwargs uses it: chat completion, transcription,
embedding, image generation and edit, search and video, ElevenLabs text to
speech, and the Bedrock batch mapper. The two untyped shared filters now take
Mapping[str, object]
The stream_chunk_size wire test becomes test_internal_params_wire.py. It also
sends an undeclared _litellm_ kwarg and asserts that no _litellm_ key reaches
any of the six provider bodies, while extra_body passthrough keeps working
Refs LIT-8318, LIT-8319
---
litellm/constants.py | 1 +
litellm/images/main.py | 14 +++-----
litellm/llms/bedrock/files/transformation.py | 4 +--
.../text_to_speech/transformation.py | 6 ++--
litellm/main.py | 13 +++----
litellm/types/utils.py | 5 +++
litellm/utils.py | 36 +++++--------------
...e_wire.py => test_internal_params_wire.py} | 4 ++-
.../images/test_image_edit_extra_params.py | 20 +++++++++++
tests/unit/images/test_main.py | 29 +++++++++++++++
.../test_bedrock_files_transformation.py | 23 ++++++++++++
...levenlabs_text_to_speech_transformation.py | 34 +++++++++++++++---
tests/unit/test_main.py | 28 +++++++++++++++
tests/unit/types/test_litellm_params.py | 23 ++++++++----
14 files changed, 178 insertions(+), 62 deletions(-)
rename tests/integration/providers/{test_stream_chunk_size_wire.py => test_internal_params_wire.py} (98%)
create mode 100644 tests/unit/images/test_main.py
diff --git a/litellm/constants.py b/litellm/constants.py
index a5be2f6568d..dac15c01fbf 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -1610,6 +1610,7 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = {
# e.g., 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value'
# Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.)
PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-"
+INTERNAL_KWARG_PREFIX: Final = "_litellm_"
AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech"
AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech"
diff --git a/litellm/images/main.py b/litellm/images/main.py
index 1f722eb752a..5ca8a726a69 100644
--- a/litellm/images/main.py
+++ b/litellm/images/main.py
@@ -52,7 +52,7 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
LITELLM_IMAGE_VARIATION_PROVIDERS,
LlmProviders,
- all_litellm_params,
+ is_litellm_owned_kwarg,
)
from litellm.utils import (
ImageResponse,
@@ -249,11 +249,9 @@ def image_generation(
"size",
"style",
]
- litellm_params: Final = all_litellm_params
- default_params: Final = openai_params + litellm_params
non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in default_params
- } # model-specific params - pass them straight to the model/provider
+ k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k)
+ }
image_generation_config: BaseImageGenerationConfig | None = None
if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values():
@@ -757,11 +755,9 @@ def image_edit(
"style",
"async_call",
]
- litellm_params_list: Final = all_litellm_params
- default_params: Final = openai_params + litellm_params_list
non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in default_params
- } # model-specific params - pass them straight to the model/provider
+ k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k)
+ }
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
model_info: Final = kwargs.get("model_info", None)
diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py
index ce4c955a884..fdc8e34ed3d 100644
--- a/litellm/llms/bedrock/files/transformation.py
+++ b/litellm/llms/bedrock/files/transformation.py
@@ -58,7 +58,7 @@ from litellm.types.llms.openai import (
OpenAIFileObject,
PathLike,
)
-from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums, all_litellm_params
+from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums, is_litellm_owned_kwarg
from litellm.utils import get_llm_provider, get_optional_params
from ..base_aws_llm import BaseAWSLLM
@@ -907,7 +907,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
{
k: v
for k, v in optional_params.items()
- if k not in all_litellm_params or k in _LITELLM_PARAMS_THE_MAPPER_TAKES
+ if not is_litellm_owned_kwarg(k) or k in _LITELLM_PARAMS_THE_MAPPER_TAKES
}
),
)
diff --git a/litellm/llms/elevenlabs/text_to_speech/transformation.py b/litellm/llms/elevenlabs/text_to_speech/transformation.py
index 3cf9a983efe..eb93543df46 100644
--- a/litellm/llms/elevenlabs/text_to_speech/transformation.py
+++ b/litellm/llms/elevenlabs/text_to_speech/transformation.py
@@ -18,7 +18,7 @@ from litellm.llms.base_llm.text_to_speech.transformation import (
TextToSpeechRequestData,
)
from litellm.secret_managers.main import get_secret_str
-from litellm.types.utils import all_litellm_params
+from litellm.types.utils import is_litellm_owned_kwarg
from ..common_utils import ElevenLabsException
@@ -241,7 +241,7 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig):
continue
mapped_params[key] = value
- reserved_kwarg_keys: Final = set(all_litellm_params) | {
+ reserved_kwarg_keys: Final = {
self.ELEVENLABS_QUERY_PARAMS_KEY,
self.ELEVENLABS_VOICE_ID_KEY,
"voice",
@@ -260,7 +260,7 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig):
mapped_params[key] = value
for key in list(kwargs.keys()):
- if key in reserved_kwarg_keys:
+ if key in reserved_kwarg_keys or is_litellm_owned_kwarg(key):
continue
value = kwargs[key]
if value is None:
diff --git a/litellm/main.py b/litellm/main.py
index 12854db15d0..8c2afe4429a 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -284,7 +284,7 @@ from .types.utils import (
LlmProviders,
PromptTokensDetails,
ProviderSpecificHeader,
- all_litellm_params,
+ is_litellm_owned_kwarg,
)
####### ENVIRONMENT VARIABLES ###################
@@ -6351,15 +6351,10 @@ def embedding(
"max_retries",
"encoding_format",
]
- litellm_params: Final = [
- "aembedding",
- "extra_headers",
- ] + all_litellm_params
-
- default_params: Final = openai_params + litellm_params
+ default_params: Final = [*openai_params, "aembedding", "extra_headers"]
non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in default_params
- } # model-specific params - pass them straight to the model/provider
+ k: v for k, v in kwargs.items() if k not in default_params and not is_litellm_owned_kwarg(k)
+ }
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
model=model,
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 749ef229fbe..f8b57139b37 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -48,6 +48,7 @@ from typing_extensions import NotRequired, ReadOnly, Required, TypedDict
from litellm._logging import verbose_logger
from litellm._uuid import uuid
+from litellm.constants import INTERNAL_KWARG_PREFIX
from litellm.types.llms.base import (
BaseLiteLLMOpenAIResponseObject,
CachedTokensDetails,
@@ -3937,6 +3938,10 @@ all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-
]
+def is_litellm_owned_kwarg(name: str) -> bool:
+ return name in all_litellm_params or name.startswith(INTERNAL_KWARG_PREFIX)
+
+
class KeyGenerationConfig(TypedDict, total=False):
required_params: list[str] # specify params that must be present in the key generation request
diff --git a/litellm/utils.py b/litellm/utils.py
index e5eea562c11..092fe936cf9 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -257,7 +257,7 @@ from litellm.types.utils import (
TextCompletionResponse,
TranscriptionResponse,
Usage,
- all_litellm_params,
+ is_litellm_owned_kwarg,
)
_CALL_TYPE_ENUM_MAP: Final[dict] = {ct.value: ct for ct in CallTypes}
@@ -4161,26 +4161,8 @@ def _remove_unsupported_params(non_default_params: dict, supported_openai_params
return non_default_params
-def filter_out_litellm_params(kwargs: dict) -> dict:
- """
- Filter out LiteLLM internal parameters from kwargs dict.
-
- Returns a new dict containing only non-LiteLLM parameters that should be
- passed to external provider APIs.
-
- Args:
- kwargs: Dictionary that may contain LiteLLM internal parameters
-
- Returns:
- Dictionary with LiteLLM internal parameters filtered out
-
- Example:
- >>> kwargs = {"query": "test", "shared_session": session_obj, "metadata": {}}
- >>> filtered = filter_out_litellm_params(kwargs)
- >>> # filtered = {"query": "test"}
- """
-
- return {key: value for key, value in kwargs.items() if key not in all_litellm_params}
+def filter_out_litellm_params(kwargs: Mapping[str, object]) -> dict:
+ return {key: value for key, value in kwargs.items() if not is_litellm_owned_kwarg(key)}
def _provider_supports_vertex_params(custom_llm_provider: str) -> bool:
@@ -10152,10 +10134,9 @@ def get_standard_openai_params(params: Mapping[str, object]) -> dict:
def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict:
openai_params: Final = litellm.OPENAI_CHAT_COMPLETION_PARAMS
- default_params: Final = openai_params + all_litellm_params
non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in default_params
- } # model-specific params - pass them straight to the model/provider
+ k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k)
+ }
return non_default_params
@@ -10203,11 +10184,12 @@ def strip_reasoning_summary_aliases_from_optional_params(
return op, rs_val
-def get_non_default_transcription_params(kwargs: dict) -> dict:
+def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict:
from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS
- default_params: Final = OPENAI_TRANSCRIPTION_PARAMS + all_litellm_params
- non_default_params: Final = {k: v for k, v in kwargs.items() if k not in default_params}
+ non_default_params: Final = {
+ k: v for k, v in kwargs.items() if k not in OPENAI_TRANSCRIPTION_PARAMS and not is_litellm_owned_kwarg(k)
+ }
return non_default_params
diff --git a/tests/integration/providers/test_stream_chunk_size_wire.py b/tests/integration/providers/test_internal_params_wire.py
similarity index 98%
rename from tests/integration/providers/test_stream_chunk_size_wire.py
rename to tests/integration/providers/test_internal_params_wire.py
index 3681da0e3d4..17b0fc9d815 100644
--- a/tests/integration/providers/test_stream_chunk_size_wire.py
+++ b/tests/integration/providers/test_internal_params_wire.py
@@ -276,7 +276,7 @@ def provider_wire_environment(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -
@pytest.mark.parametrize("provider", PROVIDERS)
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("stream", [False, True])
-async def test_stream_chunk_size_never_reaches_provider_body(
+async def test_internal_params_never_reach_provider_body(
monkeypatch: pytest.MonkeyPatch,
provider_wire_environment: None,
provider: str,
@@ -289,6 +289,7 @@ async def test_stream_chunk_size_never_reaches_provider_body(
**_request_parameters(provider, wire.url),
"stream": stream,
"stream_chunk_size": 64,
+ "_litellm_undeclared_sentinel": "internal",
"extra_body": {"custom_provider_key": 1},
"max_tokens": 16,
"timeout": 5,
@@ -313,4 +314,5 @@ async def test_stream_chunk_size_never_reaches_provider_body(
keys: Final = keys_at_every_depth(body)
assert "stream_chunk_size" not in keys
assert not INTERNAL_FIELDS.intersection(keys)
+ assert not frozenset(key for key in keys if key.startswith("_litellm_")), keys
assert _custom_key(body, provider) == 1
diff --git a/tests/unit/images/test_image_edit_extra_params.py b/tests/unit/images/test_image_edit_extra_params.py
index 088faafa9f3..c3b0a5d2828 100644
--- a/tests/unit/images/test_image_edit_extra_params.py
+++ b/tests/unit/images/test_image_edit_extra_params.py
@@ -58,6 +58,26 @@ def test_image_edit_forwards_provider_params_and_extra_body():
assert response.data
+def test_image_edit_keeps_an_internal_prefixed_kwarg_out_of_the_provider_request():
+ captured = {}
+ client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_image_edit_request(captured))))
+
+ litellm.image_edit(
+ model="openai/gpt-image-1",
+ image=PNG_BYTES,
+ prompt="add a hat",
+ api_key="sk-test",
+ api_base="https://edit.example/v1",
+ client=client,
+ seed=42,
+ _litellm_undeclared_sentinel="internal",
+ )
+
+ fields = _multipart_text_fields(captured["content_type"], captured["body"])
+ assert "_litellm_undeclared_sentinel" not in fields
+ assert fields["seed"] == "42"
+
+
def test_image_edit_extra_body_takes_precedence_over_kwargs():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_image_edit_request(captured))))
diff --git a/tests/unit/images/test_main.py b/tests/unit/images/test_main.py
new file mode 100644
index 00000000000..d65e5d929b5
--- /dev/null
+++ b/tests/unit/images/test_main.py
@@ -0,0 +1,29 @@
+import json
+from typing import Final
+
+import httpx
+import respx
+
+import litellm
+
+
+def test_image_generation_keeps_an_internal_prefixed_kwarg_out_of_the_provider_request(
+ respx_mock: respx.MockRouter,
+) -> None:
+ api_base: Final = "http://localhost:12346/v1"
+ mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/images/generations.*").mock(
+ return_value=httpx.Response(status_code=200, json={"created": 1712697600, "data": [{"b64_json": "aW1n"}]})
+ )
+
+ litellm.image_generation(
+ model="openai/gpt-image-1",
+ prompt="a red circle",
+ api_base=api_base,
+ api_key="fake_openai_api_key",
+ _litellm_undeclared_sentinel="internal",
+ )
+
+ assert mock_route.called
+ sent: Final = json.loads(respx_mock.calls[0].request.content)
+ assert "_litellm_undeclared_sentinel" not in sent, sent
+ assert sent["prompt"] == "a red circle"
diff --git a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py
index b2ce4ab2dde..12275df404f 100644
--- a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py
+++ b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py
@@ -84,6 +84,29 @@ class TestBedrockFilesTransformation:
"max_tokens" in model_input
), f"Record {i+1} should have max_tokens"
+ def test_batch_keeps_an_internal_prefixed_key_out_of_the_bedrock_model_input(self):
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ result: Final = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content(
+ [
+ {
+ "custom_id": "internal-key-1",
+ "method": "POST",
+ "url": "/v1/chat/completions",
+ "body": {
+ "model": "anthropic.claude-3-5-sonnet-20240620-v1:0",
+ "messages": [{"role": "user", "content": "hi"}],
+ "max_tokens": 10,
+ "_litellm_undeclared_sentinel": "internal",
+ },
+ }
+ ]
+ )
+
+ model_input: Final = json.dumps(result[0]["modelInput"])
+ assert "_litellm_undeclared_sentinel" not in model_input, model_input
+ assert result[0]["modelInput"]["max_tokens"] == 10
+
def test_nova_text_only_uses_converse_format(self):
"""
Test that Nova models produce Converse API format in batch modelInput.
diff --git a/tests/unit/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py b/tests/unit/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py
index 54e689dea6b..d05371d7df9 100644
--- a/tests/unit/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py
+++ b/tests/unit/llms/elevenlabs/test_elevenlabs_text_to_speech_transformation.py
@@ -1,5 +1,11 @@
-import pytest
+import json
+from typing import Final
+import httpx
+import pytest
+import respx
+
+import litellm
from litellm.llms.elevenlabs.text_to_speech.transformation import (
ElevenLabsTextToSpeechConfig,
)
@@ -16,10 +22,7 @@ def test_should_encode_elevenlabs_voice_id_path_segment():
},
)
- assert (
- url
- == "https://api.elevenlabs.io/v1/text-to-speech/voice%2F..%2F..%2Fmodels%3Fx%3D1%23frag"
- )
+ assert url == "https://api.elevenlabs.io/v1/text-to-speech/voice%2F..%2F..%2Fmodels%3Fx%3D1%23frag"
def test_should_reject_dot_segment_elevenlabs_voice_id():
@@ -31,3 +34,24 @@ def test_should_reject_dot_segment_elevenlabs_voice_id():
api_base="https://api.elevenlabs.io",
litellm_params={config.ELEVENLABS_VOICE_ID_KEY: ".."},
)
+
+
+def test_speech_keeps_an_internal_prefixed_kwarg_out_of_the_elevenlabs_request(respx_mock: respx.MockRouter) -> None:
+ api_base: Final = "http://localhost:12346"
+ mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/v1/text-to-speech/.*").mock(
+ return_value=httpx.Response(status_code=200, content=b"audio", headers={"content-type": "audio/mpeg"})
+ )
+
+ litellm.speech(
+ model="elevenlabs/eleven_multilingual_v2",
+ input="hi",
+ voice="21m00Tcm4TlvDq8ikWAM",
+ api_base=api_base,
+ api_key="fake_elevenlabs_api_key",
+ _litellm_undeclared_sentinel="internal",
+ )
+
+ assert mock_route.called
+ sent: Final = json.loads(respx_mock.calls[0].request.content)
+ assert "_litellm_undeclared_sentinel" not in sent, sent
+ assert sent["text"] == "hi"
diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py
index 57200a79a8c..7bef35d8559 100644
--- a/tests/unit/test_main.py
+++ b/tests/unit/test_main.py
@@ -395,6 +395,34 @@ def test_completion_strips_eager_input_streaming_before_openai(respx_mock: respx
assert sent_tool["function"]["name"] == "write_file"
+def test_embedding_keeps_an_internal_prefixed_kwarg_out_of_the_provider_request(respx_mock: respx.MockRouter) -> None:
+ api_base: Final = "http://localhost:12346/v1"
+ mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/embeddings.*").mock(
+ return_value=httpx.Response(
+ status_code=200,
+ json={
+ "object": "list",
+ "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
+ "model": "text-embedding-3-small",
+ "usage": {"prompt_tokens": 1, "total_tokens": 1},
+ },
+ )
+ )
+
+ litellm.embedding(
+ model="openai/text-embedding-3-small",
+ input="hi",
+ api_base=api_base,
+ api_key="fake_openai_api_key",
+ _litellm_undeclared_sentinel="internal",
+ )
+
+ assert mock_route.called
+ sent: Final = json.loads(respx_mock.calls[0].request.content)
+ assert "_litellm_undeclared_sentinel" not in sent, sent
+ assert sent["model"] == "text-embedding-3-small"
+
+
def test_custom_provider_with_extra_headers():
with patch.object(
diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py
index e421321aaaa..33467a78aa8 100644
--- a/tests/unit/types/test_litellm_params.py
+++ b/tests/unit/types/test_litellm_params.py
@@ -262,7 +262,7 @@ OWNED_NAMES: Final = (
*PRICING_NAMES,
)
-Classifier: TypeAlias = Callable[[dict[str, object]], dict[str, object]] # mutable-ok: classifiers use dict
+Classifier: TypeAlias = Callable[[Mapping[str, object]], Mapping[str, object]]
CLASSIFIERS: Final[Mapping[str, Classifier]] = MappingProxyType(
{ # pyright: ignore[reportUnknownArgumentType] # untyped legacy classifiers
@@ -279,18 +279,31 @@ def test_owned_name_is_kept_out_of_provider_params(name: str, classifier_name: s
provider_value: Final = object()
classify: Final = CLASSIFIERS[classifier_name]
- result: Final = classify({name: object(), PROVIDER_KNOB: provider_value}) # mutable-ok: classifiers take a dict
+ result: Final = classify(MappingProxyType({name: object(), PROVIDER_KNOB: provider_value}))
assert result == MappingProxyType({PROVIDER_KNOB: provider_value})
assert result[PROVIDER_KNOB] is provider_value
def test_a_name_no_object_declares_reaches_the_provider() -> None:
- result: Final = CLASSIFIERS["completion"]({PROVIDER_KNOB: 1}) # mutable-ok: classifier input type
+ result: Final = CLASSIFIERS["completion"](MappingProxyType({PROVIDER_KNOB: 1}))
assert result == MappingProxyType({PROVIDER_KNOB: 1})
+@pytest.mark.parametrize("classifier_name", CLASSIFIERS)
+def test_an_undeclared_internal_prefixed_name_is_kept_out_of_provider_params(classifier_name: str) -> None:
+ undeclared: Final = "_litellm_never_declared_anywhere"
+ lookalike: Final = "provider_litellm_knob"
+ assert undeclared not in all_litellm_params
+
+ result: Final = CLASSIFIERS[classifier_name](
+ MappingProxyType({undeclared: object(), PROVIDER_KNOB: 1, lookalike: 2})
+ )
+
+ assert result == MappingProxyType({PROVIDER_KNOB: 1, lookalike: 2})
+
+
def _cache_key_for_model_group(cache: Cache, model_group: str, options: CachingOptions) -> str:
return cache.get_cache_key( # pyright: ignore[reportUnknownMemberType] # untyped legacy key builder
model=model_group,
@@ -421,9 +434,7 @@ CARRIED_PARAMS: Final = tuple(
def test_every_param_get_litellm_params_carries_is_kept_out_of_provider_params(name: str) -> None:
provider_value: Final = object()
- result: Final = CLASSIFIERS["completion"](
- {name: object(), PROVIDER_KNOB: provider_value} # mutable-ok: classifier input type
- )
+ result: Final = CLASSIFIERS["completion"](MappingProxyType({name: object(), PROVIDER_KNOB: provider_value}))
assert result == MappingProxyType({PROVIDER_KNOB: provider_value})
From 3afcd176b372d7262eb619ea65ffa227c2efbed8 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 22:28:28 +0000
Subject: [PATCH 17/39] test: remove substring guard test_default_api_base
(#43355)
It asserted no provider name is a substring of any other provider's default api_base, so any new provider whose name sits inside an existing hostname (sail vs parasail) broke main without a bug in our code. The litellm_proxy default api_base fix it originally guarded is covered by the explicit api_base tests in the same file
Co-authored-by: kerry
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
tests/local_testing/test_get_llm_provider.py | 40 --------------------
1 file changed, 40 deletions(-)
diff --git a/tests/local_testing/test_get_llm_provider.py b/tests/local_testing/test_get_llm_provider.py
index 4ac7cecb97a..982e14660b7 100644
--- a/tests/local_testing/test_get_llm_provider.py
+++ b/tests/local_testing/test_get_llm_provider.py
@@ -133,46 +133,6 @@ def test_get_llm_provider_azure_o1():
assert model == "o1-mini"
-def test_default_api_base():
- from litellm.litellm_core_utils.get_llm_provider_logic import (
- _get_openai_compatible_provider_info,
- )
- from litellm.types.utils import LlmProviders
-
- # Patch environment variable to remove API base if it's set
- with patch.dict(os.environ, {}, clear=True):
- for provider in litellm.openai_compatible_providers:
- # Get the API base for the given provider
- if provider == "github_copilot":
- continue
- # Skip chatgpt as it requires OAuth authentication
- if provider == "chatgpt":
- continue
- # Skip ragflow as it requires specific model format: ragflow/chat/{id}/{model} or ragflow/agent/{id}/{model}
- if provider == "ragflow":
- continue
- _, _, _, api_base = _get_openai_compatible_provider_info(
- model=f"{provider}/*", api_base=None, api_key=None, dynamic_api_key=None
- )
- if api_base is None:
- continue
-
- for other_provider in LlmProviders:
- if other_provider.value != provider and provider != "{}_chat".format(
- other_provider.value
- ):
- if provider == "codestral" and other_provider.value == "mistral":
- continue
- elif provider == "github" and other_provider.value == "azure":
- continue
- elif (
- provider in ("qwencloud", "qwen_ai_platform")
- and other_provider.value == "dashscope"
- ):
- continue
- assert other_provider.value not in api_base.replace("/openai", "")
-
-
def test_hosted_vllm_default_api_key():
from litellm.litellm_core_utils.get_llm_provider_logic import (
_get_openai_compatible_provider_info,
From 635a718ba1e0a868458338374ad4b75c4cb6eb38 Mon Sep 17 00:00:00 2001
From: yuneng-jiang
Date: Sat, 26 Sep 2026 15:34:53 -0700
Subject: [PATCH 18/39] ci: cut CircleCI wall time without loosening test
isolation (#43347)
* ci: cut CircleCI wall time without loosening test isolation
* fix(ci): parse integration split files that follow --results
The CircleCI machine image ships Python 3.12.2, whose argparse leaves the
files positional empty when it follows an option and another positional, so
every extensions node exited with 'unrecognized arguments'. Reproduced on
3.12.2; parse_intermixed_args selects the files on 3.12.2, 3.12.13 and 3.13
* test(ci): resolve command references in the Rust toolchain guard
The Windows rustup install moved into the install_windows_toolchain command,
which the guard only recognized for install_rust. It now accepts any command
that installs a pinned rustup and reads the Windows toolchain pin from it
* ci: cache the Windows release cargo build from main
windows_release_wheel rebuilt every dependency with fat LTO on each run. It now
restores the release target and cargo registry saved by main's scheduled run,
drops the workspace crates' fingerprints so they always rebuild from the
checked-out source, and still runs the full LTO link
* ci: run the Windows release wheel build on windows.xlarge
The fat-LTO release build is the slowest job in the pipeline; more cores
speed up the dependency compile ahead of the final link
* ci: skip the Windows fingerprint cleanup when the cargo cache missed
On a cold cache the release fingerprint directory does not exist, and the
CircleCI PowerShell wrapper failed the step on the suppressed not-found error
---
.circleci/config.yml | 198 +++++++++++++-----
.circleci/scripts/classify_changes.sh | 10 +-
.circleci/scripts/run_integration.sh | 11 +-
tests/integration/README.md | 4 +-
tests/integration/conftest.py | 10 +-
tests/integration/run.py | 15 +-
.../test_router_tag_routing.py | 11 +
tests/unit/test_circleci_path_filter.py | 11 +
tests/unit/test_circleci_rust_toolchain.py | 32 ++-
tests/unit/test_pre_commit_lint.py | 1 +
.../check_windows_wheel_install.py | 7 +-
.../test_check_windows_wheel_install.py | 22 ++
12 files changed, 263 insertions(+), 69 deletions(-)
diff --git a/.circleci/config.yml b/.circleci/config.yml
index d9c85cfa042..7d4e2e40769 100644
--- a/.circleci/config.yml
+++ b/.circleci/config.yml
@@ -141,7 +141,7 @@ commands:
node --version
npm --version
install_rust:
- description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself."
+ description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself. Also restores the dev-profile cargo cache that save_cargo_target writes on main, minus the workspace crates' fingerprints so those always rebuild from the checked-out source."
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
@@ -167,9 +167,29 @@ commands:
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
rm -f /tmp/rustup-init
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
+ echo 'export CARGO_INCREMENTAL=0' >> "$BASH_ENV"
export PATH="$HOME/.cargo/bin:$PATH"
rustc --version
cargo --version
+ { rustc -vV; cc --version; cat /etc/os-release; } > /tmp/cargo-build-env
+ - restore_cache:
+ keys:
+ - v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
+ - v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-
+ - run:
+ name: Force a rebuild of the workspace crates restored from the cargo cache
+ command: rm -rf litellm-rust/target/debug/.fingerprint/litellm-*
+ save_cargo_target:
+ steps:
+ - when:
+ condition:
+ equal: [main, << pipeline.git.branch >>]
+ steps:
+ - save_cache:
+ key: v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
+ paths:
+ - ~/.cargo/registry
+ - ~/project/litellm-rust/target/debug
start_postgres:
description: "Start a postgres-db container on port 5432 and wait until it accepts connections."
parameters:
@@ -281,51 +301,11 @@ commands:
# `uv sync --package litellm-enterprise` here — that overwrites the
# shared .venv and strips out dev/test deps (pytest, prisma, etc.).
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
- setup_litellm_test_deps:
+ install_windows_toolchain:
steps:
- - checkout
- - setup_google_dns
- - install_uv
- - install_rust
- - restore_cache:
- keys:
- - v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
- name: Install Dependencies
- command: |
- uv sync --frozen --all-groups --all-extras --python 3.12
- - setup_litellm_enterprise_pip
- - save_cache:
- paths:
- - ~/.cache/uv
- key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
-
-jobs:
- # Add Windows testing job
- using_litellm_on_windows:
- executor:
- name: win/default
- shell: powershell.exe
- working_directory: ~/project
- environment:
- UV_PYTHON: "3.11"
- CARGO_HTTP_MULTIPLEXING: "false"
- CARGO_NET_RETRY: "5"
- steps:
- - checkout
- - run:
- name: Install Python
- command: |
- choco install python --version=3.11.0 -y --no-progress --force
- refreshenv
- python --version
- environment:
- CHOCOLATEY_CONFIRM_ALL: "true"
- - run:
- name: Install Dependencies
+ name: Install Rust and uv
no_output_timeout: 30m
- environment:
- UV_HTTP_TIMEOUT: "300"
command: |
$rustupInit = Join-Path $env:TEMP "rustup-init.exe"
$rustupVersion = "1.28.2"
@@ -365,6 +345,55 @@ jobs:
if (-not (Select-String -Path $PROFILE -SimpleMatch $cargoBin -Quiet)) {
Add-Content -Path $PROFILE -Value "`$env:Path = `"$cargoBin;`$env:Path`""
}
+ setup_litellm_test_deps:
+ steps:
+ - checkout
+ - setup_google_dns
+ - install_uv
+ - install_rust
+ - restore_cache:
+ keys:
+ - v3-integration-uv-cache-{{ checksum "uv.lock" }}
+ - run:
+ name: Install Dependencies
+ command: |
+ uv sync --frozen --all-groups --all-extras --python 3.12
+ - setup_litellm_enterprise_pip
+ - save_cache:
+ paths:
+ - ~/.cache/uv
+ key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
+ - save_cargo_target
+
+jobs:
+ # Add Windows testing job
+ using_litellm_on_windows:
+ executor:
+ name: win/default
+ shell: powershell.exe
+ working_directory: ~/project
+ environment:
+ UV_PYTHON: "3.11"
+ CARGO_HTTP_MULTIPLEXING: "false"
+ CARGO_NET_RETRY: "5"
+ steps:
+ - checkout
+ - run:
+ name: Install Python
+ command: |
+ choco install python --version=3.11.0 -y --no-progress --force
+ refreshenv
+ python --version
+ environment:
+ CHOCOLATEY_CONFIRM_ALL: "true"
+ - install_windows_toolchain
+ - run:
+ name: Install Dependencies
+ no_output_timeout: 30m
+ environment:
+ UV_HTTP_TIMEOUT: "300"
+ command: |
+ $env:Path = "$HOME\.cargo\bin;$HOME\.local\bin;$env:Path"
for ($attempt = 1; $attempt -le 5; $attempt++) {
Write-Host "uv sync attempt $attempt/5"
uv sync --frozen --group dev --python 3.11
@@ -380,17 +409,68 @@ jobs:
name: Run Windows-specific test
command: |
uv run --no-sync python -m pytest tests/windows_tests/ -v
+
+ windows_release_wheel:
+ executor:
+ name: win/default
+ shell: powershell.exe
+ size: xlarge
+ working_directory: ~/project
+ environment:
+ UV_PYTHON: "3.11"
+ CARGO_HTTP_MULTIPLEXING: "false"
+ CARGO_NET_RETRY: "5"
+ steps:
+ - checkout
- run:
- name: Guard against MAX_PATH-busting packaged wheel paths
+ name: Skip job when no windows-release-relevant files changed
+ shell: bash.exe
+ command: bash .circleci/scripts/path_filter.sh windows-release
+ - run:
+ name: Install Python
+ command: |
+ choco install python --version=3.11.0 -y --no-progress --force
+ refreshenv
+ python --version
+ environment:
+ CHOCOLATEY_CONFIRM_ALL: "true"
+ - install_windows_toolchain
+ - run:
+ name: Record the Rust build environment for the release cargo cache key
+ command: |
+ & "$HOME\.cargo\bin\rustc.exe" -vV | Out-File -Encoding ascii .cargo-build-env
+ - restore_cache:
+ keys:
+ - v1-cargo-release-windows-{{ checksum ".cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
+ - v1-cargo-release-windows-{{ checksum ".cargo-build-env" }}-
+ - run:
+ name: Force a rebuild of the workspace crates restored from the cargo cache
+ command: |
+ $fingerprints = "litellm-rust/target/release/.fingerprint"
+ if (Test-Path $fingerprints) {
+ Get-ChildItem -Path $fingerprints -Filter "litellm-*" | Remove-Item -Recurse -Force
+ }
+ - run:
+ name: Build the release wheel and install it under a worst-case MAX_PATH prefix
no_output_timeout: 30m
environment:
UV_HTTP_TIMEOUT: "300"
command: |
$env:Path = "$HOME\.cargo\bin;$HOME\.local\bin;$env:Path"
- cargo --version
- Get-ChildItem -Path "litellm\rust_bridge" -Filter "_native*" -File -ErrorAction SilentlyContinue | Remove-Item -Force
uv build --wheel --out-dir dist
- uv run --no-sync python tests/windows_tests/check_windows_wheel_install.py
+ if ($LASTEXITCODE -ne 0) {
+ exit $LASTEXITCODE
+ }
+ python tests/windows_tests/check_windows_wheel_install.py
+ - when:
+ condition:
+ equal: [main, << pipeline.git.branch >>]
+ steps:
+ - save_cache:
+ key: v1-cargo-release-windows-{{ checksum ".cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
+ paths:
+ - ~/.cargo/registry
+ - ~/project/litellm-rust/target/release
base_sdk_install:
docker:
@@ -418,6 +498,10 @@ jobs:
uv venv /tmp/base-sdk --python 3.12
VIRTUAL_ENV=/tmp/base-sdk uv pip install dist/*.whl
/tmp/base-sdk/bin/python tests/base_sdk_tests/check_base_sdk_install.py
+ - run:
+ name: Guard against MAX_PATH-busting packaged wheel paths
+ command: |
+ python3 tests/windows_tests/check_windows_wheel_install.py --lengths-only
local_testing_part1:
docker:
@@ -446,6 +530,7 @@ jobs:
paths:
- ~/.cache/uv
key: v1-uv-cache-{{ checksum "uv.lock" }}
+ - save_cargo_target
- run:
name: Run prisma ./docker/entrypoint.sh
command: |
@@ -3120,10 +3205,14 @@ jobs:
type: enum
enum: [standard, replica]
default: standard
+ parallelism:
+ type: integer
+ default: 1
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
+ parallelism: << parameters.parallelism >>
steps:
- setup_litellm_test_deps
- when:
@@ -3249,6 +3338,7 @@ jobs:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
+ parallelism: 4
steps:
- setup_litellm_test_deps
- run:
@@ -3258,10 +3348,11 @@ jobs:
name: Run unit tests
command: |
mkdir -p test-results/unit
- mapfile -t files < <(find tests/unit -name 'test_*.py' | sort)
- if [ "${#files[@]}" -eq 0 ]; then echo "tests/unit holds no test_*.py files; nothing to run"; exit 0; fi
+ shard="$(find tests/unit -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)"
+ if [ -z "${shard}" ]; then echo "shard ${CIRCLE_NODE_INDEX} received no tests/unit files; nothing to run"; exit 0; fi
+ mapfile -t files < <(printf '%s\n' "${shard}")
set +e
- LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --junitxml=test-results/unit/junit.xml
+ LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short -o junit_family=xunit1 --junitxml=test-results/unit/junit.xml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from tests/unit; passing"; exit 0; fi
@@ -3328,7 +3419,11 @@ workflows:
name: integration-<< matrix.suite >>
matrix:
parameters:
- suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
+ suite: [management, accounting, database, providers, mcp, sdk, cost, browser]
+ - integration_contracts:
+ name: integration-extensions
+ suite: extensions
+ parallelism: 4
- integration_contracts:
name: integration-<< matrix.suite >>-replica
matrix:
@@ -3343,6 +3438,7 @@ workflows:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- using_litellm_on_windows
+ - windows_release_wheel
- unit
- provider_replay_harness
- base_sdk_install
diff --git a/.circleci/scripts/classify_changes.sh b/.circleci/scripts/classify_changes.sh
index 8c2ac019b99..387197b65d7 100755
--- a/.circleci/scripts/classify_changes.sh
+++ b/.circleci/scripts/classify_changes.sh
@@ -1,7 +1,7 @@
#!/usr/bin/env bash
set -uo pipefail
-category="${1:?usage: classify_changes.sh }"
+category="${1:?usage: classify_changes.sh }"
has_client=false
has_backend=false
@@ -9,6 +9,7 @@ has_ci=false
has_provider_harness=false
has_cost_map=false
has_mcp_dependencies=false
+has_windows_release=false
outside_cost_map_set=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
@@ -22,6 +23,10 @@ while IFS= read -r file || [ -n "$file" ]; do
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
has_provider_harness=true ;;
esac
+ case "$file" in
+ litellm-rust/* | litellm/rust_bridge/* | rust-toolchain.toml | pyproject.toml | uv.lock | tests/windows_tests/* | .circleci/*)
+ has_windows_release=true ;;
+ esac
case "$file" in
ui/* | tests/e2e/ui/*) has_client=true ;;
docs/* | *.md | *.mdx) : ;;
@@ -46,6 +51,9 @@ case "$category" in
provider-harness)
[ "$has_provider_harness" = true ] && echo run || echo skip
;;
+ windows-release)
+ [ "$has_windows_release" = true ] && echo run || echo skip
+ ;;
backend)
[ "$has_backend" = true ] && echo run || echo skip
;;
diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh
index 984419717a3..47ad2274e2f 100644
--- a/.circleci/scripts/run_integration.sh
+++ b/.circleci/scripts/run_integration.sh
@@ -212,6 +212,15 @@ if [ "$suite" = browser ]; then
exit 0
fi
+node_files=()
+if [ "${CIRCLE_NODE_TOTAL:-1}" -gt 1 ]; then
+ split="$(.venv/bin/python tests/integration/run.py "$suite" --list \
+ | circleci tests split --split-by=timings --timings-type=filename)"
+ read -r -a node_files <<< "$(printf '%s' "$split" | tr '\n' ' ')"
+ test "${#node_files[@]}" -gt 0
+ printf '%s\n' "${node_files[@]}" > "$results/node-files.txt"
+fi
+
env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
@@ -225,7 +234,7 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_PROXY_DATABASE_URL="$INTEGRATION_PROXY_DATABASE_URL" \
INTEGRATION_PROXY_READ_REPLICA_URL="$INTEGRATION_PROXY_READ_REPLICA_URL" \
INTEGRATION_ROUTING="$INTEGRATION_ROUTING" \
- .venv/bin/python tests/integration/run.py "$suite" --results "$results"
+ .venv/bin/python tests/integration/run.py "$suite" --results "$results" "${node_files[@]}"
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
for covered_pid in "$proxy_pid" "$peer_pid"; do
diff --git a/tests/integration/README.md b/tests/integration/README.md
index ac9b01786b9..c09904597ab 100644
--- a/tests/integration/README.md
+++ b/tests/integration/README.md
@@ -8,7 +8,7 @@ Use `tests/integration/run.py management`, `accounting`, `database`, `providers`
Management also requires `INTEGRATION_PEER_URL`, `REDIS_HOST` and `REDIS_PORT`. CircleCI starts two directly addressed proxy processes sharing only that job's stores. The test-only CLI wrapper supplies enterprise route entitlement, following the existing behavior suite's convention. It does not qualify license validation; run it with one worker and no reload
-The generated lifecycle models use 20 examples, eight steps, generation and shrinking, with isolated resources per example. HTTP operation caps include generation and shrinking and exempt cleanup. Local qualification defaults to seed 4106601 and canonical order; CircleCI derives exploration and ordering seeds from the checked-out revision and workflow ID. Use `--seed` and `--order-seed` to reproduce a run. Actual installed Hypothesis version, settings, seeds and collected order are written beside the execution manifest
+The generated lifecycle models use 20 examples, eight steps, generation and shrinking, with isolated resources per example. HTTP operation caps include generation and shrinking and exempt cleanup. Local qualification defaults to seed 4106601 and canonical order; CircleCI derives exploration and ordering seeds from the checked-out revision and workflow ID. The ordering seed shuffles the file order and the test order inside each file but keeps each file's tests together, so module fixtures are built once per file. Use `--seed` and `--order-seed` to reproduce a run. Actual installed Hypothesis version, settings, seeds and collected order are written beside the execution manifest
Reuse the existing canned provider handlers through `_support/upstream.py`. It rejects internal request fields and exposes actual received requests for independent assertions. Register every created resource for cleanup immediately, keep expected values independent of production calculations, and assert readback plus the runtime effect of a change
@@ -30,7 +30,7 @@ Streaming checks send real HTTP transfer chunks, including one-byte partitions,
The sdk shard exercises the SDK's own HTTP clients against local protocol peers with no gateway in the path, so a case here fails only when the client library or its wire behavior changes. The HTTP/2 case runs a hypercorn TLS peer offering h2 and http/1.1 over ALPN, drives the sync and async httpx handlers at it with `LITELLM_HTTP2` off and on, and asserts the version both the client and the peer observed on the wire. Put a test here only when it needs no proxy, database or Redis; a case that reaches the gateway belongs in one of the other shards
-The extensions shard uses the built-in generic callback and guardrail transports. It checks callback correlation and credential exclusion, guardrail rewriting and denial, retained OpenAI consumers and A2A wire versions
+The extensions shard uses the built-in generic callback and guardrail transports. It checks callback correlation and credential exclusion, guardrail rewriting and denial, retained OpenAI consumers and A2A wire versions. CircleCI runs it on parallel nodes, and each node starts its own database, Redis, upstream and proxy and runs its share of the group's files serially, split by recorded timings with `circleci tests split`. Tests keep the isolation of a serial run; they still must not assume a particular set of sibling files. `run.py --list` prints a group's files and `run.py ...` runs a subset of them
The mcp shard runs the MCP gateway against SDK peers owned by each test (`_support/mcp.py`): streamable HTTP, SSE and stdio peers, an OpenAPI-spec app, and an OAuth 2.1 authorization-server double. Every peer records the requests it receives so a test can assert what reached the peer, not only what the proxy answered. The shard runs with `INTEGRATION_WORKERS` set and with `INTEGRATION_COVERAGE=1`, which starts the proxy under `coverage run --parallel-mode` limited to the MCP modules and stores `coverage.txt` plus an HTML report with the job artifacts. A test that fails because the product is wrong is skipped with `pytest.skip("BUG: ")` so the skip list in `execution.json` is the open MCP bug list
diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py
index 368b3ebee75..1b39eb81b01 100644
--- a/tests/integration/conftest.py
+++ b/tests/integration/conftest.py
@@ -51,10 +51,18 @@ def _owned(nodeid: str) -> bool:
return parts[:2] == ("tests", "integration") and len(parts) > 3 and parts[2] in OWNED_DIRECTORIES
+def _digest(seed: int, identity: str) -> bytes:
+ return hashlib.sha256(f"{seed}:{identity}".encode()).digest()
+
+
+def _order_key(seed: int, nodeid: str) -> tuple[bytes, bytes]:
+ return _digest(seed, nodeid.split("::", 1)[0]), _digest(seed, nodeid)
+
+
def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None:
order_seed: Final = config.getoption("integration_order_seed")
if order_seed:
- items.sort(key=lambda item: hashlib.sha256(f"{order_seed}:{item.nodeid}".encode()).digest())
+ items.sort(key=lambda item: _order_key(order_seed, item.nodeid))
root: Final = Path(__file__).parent
owned: Final = tuple(
item
diff --git a/tests/integration/run.py b/tests/integration/run.py
index 9ef585def3d..87f93873267 100644
--- a/tests/integration/run.py
+++ b/tests/integration/run.py
@@ -30,13 +30,22 @@ def main() -> int:
parser.add_argument("--seed", type=int, default=int(os.environ.get("INTEGRATION_SEED", "4106601")))
parser.add_argument("--order-seed", type=int, default=int(os.environ.get("INTEGRATION_ORDER_SEED", "0")))
parser.add_argument("--workers", type=int, default=int(os.environ.get("INTEGRATION_WORKERS", "1")))
- options: Final = parser.parse_args()
+ parser.add_argument("--list", action="store_true", help="print the group's test files and exit")
+ parser.add_argument("files", nargs="*", help="run only these files of the group")
+ options: Final = parser.parse_intermixed_args()
root: Final = Path(__file__).resolve().parents[2]
- selected: Final = tuple(
+ group_files: Final = tuple(
str(path.relative_to(root))
for folder in GROUPS[options.group]
for path in sorted((root / "tests/integration" / folder).glob("test_*.py"))
)
+ if options.list:
+ print("\n".join(group_files))
+ return 0
+ foreign: Final = sorted(set(options.files) - set(group_files))
+ if foreign:
+ parser.error(f"Not in the {options.group} group: {', '.join(foreign)}")
+ selected: Final = tuple(options.files) or group_files
if not selected:
parser.error(f"No integration test files selected for {options.group}")
output: Final = options.results.resolve()
@@ -65,6 +74,8 @@ def main() -> int:
f"--hypothesis-seed={options.seed}",
f"--integration-order-seed={options.order_seed}",
f"--junitxml={output / 'junit.xml'}",
+ "-o",
+ "junit_family=xunit1",
*(("-n", str(options.workers)) if options.workers > 1 else ()),
],
cwd=root,
diff --git a/tests/unit/router_strategy/test_router_tag_routing.py b/tests/unit/router_strategy/test_router_tag_routing.py
index e4b8860a7a6..d46b12a338f 100644
--- a/tests/unit/router_strategy/test_router_tag_routing.py
+++ b/tests/unit/router_strategy/test_router_tag_routing.py
@@ -647,6 +647,7 @@ async def test_negation_with_positive_tag():
@pytest.mark.asyncio()
async def test_negation_all_excluded_raises():
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "gpt-4",
@@ -907,6 +908,7 @@ async def test_positive_tags_unchanged_by_negation():
@pytest.mark.asyncio()
async def test_negation_skips_banned_group_and_uses_fallback():
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "primary",
@@ -943,6 +945,7 @@ async def test_negation_skips_banned_group_and_uses_fallback():
@pytest.mark.asyncio()
async def test_negation_exhausts_entire_fallback_chain():
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "primary",
@@ -1696,6 +1699,7 @@ async def test_required_and_single_tag_matches_trivially():
async def test_required_and_unmatched_raises_by_default():
# allow_fail_open unset -> unmatched required-AND raises, same as today's "!" behavior.
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "gpt-4",
@@ -1728,6 +1732,7 @@ async def test_required_and_combined_with_positive_unmatched_raises_by_default()
# &A eliminates every candidate before the positive-tag preference even runs;
# this must be gated by allow_fail_open too, not just the required-AND-only path.
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "gpt-4",
@@ -1858,6 +1863,7 @@ async def test_allow_fail_open_per_hop_across_fallback_chain():
# required-AND fail-open must be re-evaluated fresh on every hop, the same
# per-hop guarantee the negation feature already established.
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "primary",
@@ -1950,6 +1956,7 @@ async def test_allow_fail_open_resolves_locally_without_triggering_external_fall
@pytest.mark.asyncio()
async def test_negation_combined_with_positive_unmatched_raises_by_default():
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "gpt-4",
@@ -2287,6 +2294,7 @@ async def test_required_and_exhausts_primary_group_falls_through_to_fallback_gro
# where the tag is satisfiable. No allow_fail_open involved; this is the plain
# fallback-chain mechanics already established for "!" extended to "&".
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "primary",
@@ -2332,6 +2340,7 @@ async def test_required_and_negation_and_allow_fail_open_combine_across_three_mo
# carrier is legitimately excluded, not hidden behind an invented tag, so the
# opted-in allow_fail_open falls back to the group's own default deployment.
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "primary",
@@ -2393,6 +2402,7 @@ async def test_unknown_tag_denial_is_scoped_per_hop_not_leaked_across_fallback_g
# discover what its own group knows; a deny decision from a prior hop's group
# must not leak forward and block a later hop that has no relevant knowledge.
router = litellm.Router(
+ num_retries=0,
model_list=[
{
"model_name": "primary",
@@ -2868,6 +2878,7 @@ def _tagged_marker_router(tier_tags=None):
},
],
enable_tag_filtering=True,
+ num_retries=0,
)
router.auto_routers = {
"gpt4o": [TaggedPreRoutingStrategy(tags=("route",), strategy=_RewriteToTierStrategy("gemini-flash"))]
diff --git a/tests/unit/test_circleci_path_filter.py b/tests/unit/test_circleci_path_filter.py
index dcce7f57113..3776aea28e4 100644
--- a/tests/unit/test_circleci_path_filter.py
+++ b/tests/unit/test_circleci_path_filter.py
@@ -73,6 +73,17 @@ CI = [".github/workflows/test-litellm-ui-unit.yml"]
("provider-harness", ["tests/e2e/quota_management/test_quota.py"], "skip"),
("provider-harness", ["litellm/main.py"], "skip"),
("provider-harness", ["ui/litellm-dashboard/src/App.tsx"], "skip"),
+ ("windows-release", ["litellm-rust/crates/core/src/lib.rs"], "run"),
+ ("windows-release", ["litellm/rust_bridge/dispatch.py"], "run"),
+ ("windows-release", ["rust-toolchain.toml"], "run"),
+ ("windows-release", ["pyproject.toml"], "run"),
+ ("windows-release", ["uv.lock"], "run"),
+ ("windows-release", ["tests/windows_tests/check_windows_wheel_install.py"], "run"),
+ ("windows-release", [".circleci/config.yml"], "run"),
+ ("windows-release", ["litellm/main.py"], "skip"),
+ ("windows-release", ["tests/unit/test_utils.py"], "skip"),
+ ("windows-release", ["ui/litellm-dashboard/src/App.tsx"], "skip"),
+ ("windows-release", ["docs/my-website/docs/index.md"], "skip"),
# docs-only: skip everything
("backend", DOCS, "skip"),
("client", DOCS, "skip"),
diff --git a/tests/unit/test_circleci_rust_toolchain.py b/tests/unit/test_circleci_rust_toolchain.py
index 800ca21b95d..039854a75ac 100644
--- a/tests/unit/test_circleci_rust_toolchain.py
+++ b/tests/unit/test_circleci_rust_toolchain.py
@@ -13,8 +13,8 @@ Two invariants are pinned here:
1. No step list (job or reusable command) reaches a `uv sync` / `uv build`
without a Rust toolchain already provisioned ahead of it. That is the
- `install_rust` command on Linux and an inline pinned rustup install in the
- Windows job, so the check accepts either. A new job that syncs without one
+ `install_rust` command on Linux and `install_windows_toolchain` on Windows,
+ so the check accepts any command or step that installs a pinned rustup. A new job that syncs without one
falls back to the unpinned path, which is exactly the regression a static
check catches at PR time and a green CI run does not.
2. Both installers pin what they download: an explicit rustup version, a
@@ -66,13 +66,23 @@ def _without_comments(text: str) -> str:
return "\n".join(line for line in text.splitlines() if not line.lstrip().startswith("#"))
-def _provisions_rust(step: object) -> bool:
- if step == "install_rust":
- return True
+def _installs_pinned_rustup(step: object) -> bool:
text = _step_text(step)
return "rustup-init" in text and ("sha256sum" in text or "SHA256" in text)
+def _provisioning_commands() -> frozenset[str]:
+ return frozenset(
+ name.removeprefix("command ")
+ for name, steps in _step_lists().items()
+ if name.startswith("command ") and any(_installs_pinned_rustup(step) for step in steps)
+ )
+
+
+def _provisions_rust(step: object, provisioning_commands: frozenset[str]) -> bool:
+ return (isinstance(step, str) and step in provisioning_commands) or _installs_pinned_rustup(step)
+
+
def _step_lists() -> dict[str, list[object]]:
config = _config()
lists: dict[str, list[object]] = {}
@@ -87,11 +97,11 @@ def _step_lists() -> dict[str, list[object]]:
return lists
-def _first_unprovisioned_build(steps: list[object]) -> str | None:
+def _first_unprovisioned_build(steps: list[object], provisioning_commands: frozenset[str]) -> str | None:
"""Return the shell text of the first workspace build reached without Rust, if any."""
rust_ready = False
for step in steps:
- if _provisions_rust(step):
+ if _provisions_rust(step, provisioning_commands):
rust_ready = True
text = _step_text(step)
if BUILDS_WORKSPACE.search(_without_comments(text)) and not rust_ready:
@@ -111,8 +121,12 @@ def test_step_lists_exist() -> None:
def test_no_workspace_build_without_a_provisioned_rust_toolchain() -> None:
+ provisioning_commands: Final = _provisioning_commands()
+ assert {"install_rust", "install_windows_toolchain"} <= provisioning_commands
offenders = {
- name: build for name, steps in _step_lists().items() if (build := _first_unprovisioned_build(steps)) is not None
+ name: build
+ for name, steps in _step_lists().items()
+ if (build := _first_unprovisioned_build(steps, provisioning_commands)) is not None
}
assert not offenders, (
"these CircleCI step lists run `uv sync`/`uv build` with no Rust toolchain provisioned first, "
@@ -156,7 +170,7 @@ def test_install_rust_pins_an_exact_toolchain_version(install_rust_command: str)
def test_windows_installer_matches_the_repo_toolchain() -> None:
- windows_steps: Final = _step_lists()["job using_litellm_on_windows"]
+ windows_steps: Final = _step_lists()["command install_windows_toolchain"]
windows_command: Final = "\n".join(_step_text(step) for step in windows_steps)
match: Final = EXACT_TOOLCHAIN.search(windows_command)
assert match is not None
diff --git a/tests/unit/test_pre_commit_lint.py b/tests/unit/test_pre_commit_lint.py
index 56f98d0e05e..471c8b41b5c 100644
--- a/tests/unit/test_pre_commit_lint.py
+++ b/tests/unit/test_pre_commit_lint.py
@@ -388,6 +388,7 @@ def test_interrupt_spares_the_invoking_process(tmp_path: Path) -> None:
)
try:
assert _wait_until((hang_dir / "make.started").exists, 10)
+ assert _wait_until((hang_dir / "eslint_report.started").exists, 10)
os.killpg(proc.pid, signal.SIGINT)
assert proc.wait(timeout=10) == 0
assert _wait_until(marker.exists, 5)
diff --git a/tests/windows_tests/check_windows_wheel_install.py b/tests/windows_tests/check_windows_wheel_install.py
index 6dbb9da6288..d0b448f35f6 100644
--- a/tests/windows_tests/check_windows_wheel_install.py
+++ b/tests/windows_tests/check_windows_wheel_install.py
@@ -35,7 +35,7 @@ def _run(cmd):
return subprocess.call(cmd)
-def main():
+def main(argv):
wheels = glob.glob(os.path.join("dist", "*.whl"))
if not wheels:
print("::error::no wheel in dist/; run `uv build --wheel --out-dir dist` first")
@@ -51,6 +51,9 @@ def main():
for n in offenders[:15]:
print(f" on-disk {WORST_CASE_PREFIX + len(n):4} {n}")
return 1
+ if "--lengths-only" in argv:
+ print(f"ok: every path in {os.path.basename(wheel)} fits MAX_PATH at a {WORST_CASE_PREFIX}-char prefix")
+ return 0
venv = _deep_venv_dir()
os.makedirs(os.path.dirname(venv), exist_ok=True)
@@ -73,4 +76,4 @@ def main():
if __name__ == "__main__":
- sys.exit(main())
+ sys.exit(main(sys.argv[1:]))
diff --git a/tests/windows_tests/test_check_windows_wheel_install.py b/tests/windows_tests/test_check_windows_wheel_install.py
index 22a197604ed..204bcb2f5e2 100644
--- a/tests/windows_tests/test_check_windows_wheel_install.py
+++ b/tests/windows_tests/test_check_windows_wheel_install.py
@@ -3,6 +3,7 @@ import zipfile
from check_windows_wheel_install import (
MAX_PATH,
WORST_CASE_PREFIX,
+ main,
overlong_install_paths,
)
@@ -34,3 +35,24 @@ def test_orders_offenders_longest_first(tmp_path):
longer,
shorter,
]
+
+
+def _dist_with(tmp_path, *entry_names):
+ dist = tmp_path / "dist"
+ dist.mkdir()
+ with zipfile.ZipFile(dist / "litellm-0-py3-none-any.whl", "w") as zf:
+ for name in entry_names:
+ zf.writestr(name, "{}")
+
+
+def test_lengths_only_passes_without_installing(tmp_path, monkeypatch):
+ _dist_with(tmp_path, "litellm/__init__.py")
+ monkeypatch.chdir(tmp_path)
+ monkeypatch.setenv("PATH", "")
+ assert main(["--lengths-only"]) == 0
+
+
+def test_lengths_only_fails_on_an_overlong_path(tmp_path, monkeypatch):
+ _dist_with(tmp_path, "a" * (MAX_PATH - WORST_CASE_PREFIX + 1))
+ monkeypatch.chdir(tmp_path)
+ assert main(["--lengths-only"]) == 1
From 69d2a3c24f915898876c5d1c52f8f36b361f4630 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 22:40:16 +0000
Subject: [PATCH 19/39] fix(cost-map): correct azure gpt-4o-mini tts,
transcribe, alias and MAI-Image-2.5 prices (#43357)
Co-authored-by: kerry
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/model_prices_and_context_window_backup.json | 10 +++++-----
model_prices_and_context_window.json | 10 +++++-----
2 files changed, 10 insertions(+), 10 deletions(-)
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index 8f61dc91adf..45f5967d372 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -5853,13 +5853,13 @@
"azure/gpt-4o-mini": {
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 7.5e-08,
- "input_cost_per_token": 1.65e-07,
+ "input_cost_per_token": 1.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 6.6e-07,
+ "output_cost_per_token": 6e-07,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
@@ -6215,7 +6215,7 @@
},
"azure/gpt-4o-mini-transcribe": {
"deprecation_date": "2027-06-15",
- "input_cost_per_audio_token": 1.25e-06,
+ "input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure",
"max_input_tokens": 16000,
@@ -6228,7 +6228,7 @@
},
"azure/gpt-4o-mini-tts": {
"deprecation_date": "2027-06-15",
- "input_cost_per_token": 2.5e-06,
+ "input_cost_per_token": 6e-07,
"litellm_provider": "azure",
"mode": "audio_speech",
"output_cost_per_audio_token": 1.2e-05,
@@ -11699,7 +11699,7 @@
"input_cost_per_token": 5e-06,
"litellm_provider": "azure_ai",
"mode": "image_generation",
- "output_cost_per_image": 0.05,
+ "output_cost_per_image": 0.048,
"output_cost_per_image_token": 4.7e-05,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
"supported_endpoints": [
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index 8f61dc91adf..45f5967d372 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -5853,13 +5853,13 @@
"azure/gpt-4o-mini": {
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 7.5e-08,
- "input_cost_per_token": 1.65e-07,
+ "input_cost_per_token": 1.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 6.6e-07,
+ "output_cost_per_token": 6e-07,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
@@ -6215,7 +6215,7 @@
},
"azure/gpt-4o-mini-transcribe": {
"deprecation_date": "2027-06-15",
- "input_cost_per_audio_token": 1.25e-06,
+ "input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure",
"max_input_tokens": 16000,
@@ -6228,7 +6228,7 @@
},
"azure/gpt-4o-mini-tts": {
"deprecation_date": "2027-06-15",
- "input_cost_per_token": 2.5e-06,
+ "input_cost_per_token": 6e-07,
"litellm_provider": "azure",
"mode": "audio_speech",
"output_cost_per_audio_token": 1.2e-05,
@@ -11699,7 +11699,7 @@
"input_cost_per_token": 5e-06,
"litellm_provider": "azure_ai",
"mode": "image_generation",
- "output_cost_per_image": 0.05,
+ "output_cost_per_image": 0.048,
"output_cost_per_image_token": 4.7e-05,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
"supported_endpoints": [
From 8d166258a65ba272546c5e62c3aac79cc7831ae3 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 15:49:39 -0700
Subject: [PATCH 20/39] fix(tests): stop VCR recording and replaying a test's
own localhost upstream (#43346)
* fix(tests): stop VCR recording and replaying a test's own localhost upstream
* test(vcr): prove a localhost response an earlier run stored is never replayed
* test(vcr): drive the localhost cassette checks in-process instead of through a loopback server
---------
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
---
tests/_vcr_conftest_common.py | 1 +
tests/llm_translation/Readme.md | 5 ++
tests/unit/test_vcr_safe_body_matcher.py | 69 ++++++++++++++++++++++++
3 files changed, 75 insertions(+)
diff --git a/tests/_vcr_conftest_common.py b/tests/_vcr_conftest_common.py
index 3adc671021b..36ae70497e7 100644
--- a/tests/_vcr_conftest_common.py
+++ b/tests/_vcr_conftest_common.py
@@ -1090,6 +1090,7 @@ def vcr_config_dict() -> dict:
"decode_compressed_response": True,
"record_mode": "new_episodes",
"allow_playback_repeats": True,
+ "ignore_localhost": True,
"match_on": (
"method",
"scheme",
diff --git a/tests/llm_translation/Readme.md b/tests/llm_translation/Readme.md
index 813c188ee7b..f0a32f6c989 100644
--- a/tests/llm_translation/Readme.md
+++ b/tests/llm_translation/Readme.md
@@ -16,6 +16,11 @@ The persister, header scrubbing, and 2xx-only filtering are defined in
patches the same httpx transport vcrpy does) are excluded from the
auto-marker — see `_RESPX_CONFLICTING_FILES` in `conftest.py`.
+Requests to `localhost`, `127.0.0.1`, or `0.0.0.0` are never recorded or
+replayed (`ignore_localhost` in `vcr_config_dict()`): a server the test
+process starts itself on an ephemeral port is not a provider, and a cassette
+entry for it would replay against whichever later test lands on that port
+
The same VCR cache is used by other test directories that exercise live
provider APIs. The reusable conftest plumbing lives in
`tests/_vcr_conftest_common.py` and is wired into:
diff --git a/tests/unit/test_vcr_safe_body_matcher.py b/tests/unit/test_vcr_safe_body_matcher.py
index cf4e4a1c276..71ae97e69d9 100644
--- a/tests/unit/test_vcr_safe_body_matcher.py
+++ b/tests/unit/test_vcr_safe_body_matcher.py
@@ -1,10 +1,15 @@
from __future__ import annotations
+import json
import os
import sys
+from pathlib import Path
from types import SimpleNamespace
+from typing import Final
import pytest
+import vcr
+from vcr.request import Request
_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))
if _REPO_ROOT not in sys.path:
@@ -384,3 +389,67 @@ def test_before_record_request_is_idempotent_on_the_same_request_object():
_before_record_request(req)
assert req.headers[KEY_FINGERPRINT_HEADER] == fp_after_first
assert fp_after_first != "no-key"
+
+
+LOCAL_UPSTREAM: Final = "http://127.0.0.1:54321/v1/moderations"
+REMOTE_UPSTREAM: Final = "https://api.openai.com/v1/moderations"
+
+
+def _recorder_with_repo_matchers(cassette_dir: Path) -> vcr.VCR:
+ recorder: Final = vcr.VCR(cassette_library_dir=str(cassette_dir))
+ recorder.register_matcher(SAFE_BODY_MATCHER_NAME, _safe_body_matcher)
+ recorder.register_matcher(KEY_FINGERPRINT_MATCHER_NAME, _key_fingerprint_matcher)
+ recorder.register_matcher(TOLERANT_QUERY_MATCHER_NAME, _tolerant_query_matcher)
+ recorder.register_matcher(TOLERANT_PATH_MATCHER_NAME, _tolerant_path_matcher)
+ return recorder
+
+
+def _request_to(uri: str) -> Request:
+ return Request(
+ method="POST",
+ uri=uri,
+ body=b'{"model":"omni-moderation-latest","input":"hi"}',
+ headers={"content-type": "application/json"},
+ )
+
+
+def _response_served_by(server: str) -> dict[str, object]:
+ payload: Final = json.dumps({"served_by": server}).encode()
+ return {
+ "status": {"code": 200, "message": "OK"},
+ "headers": {"content-type": ["application/json"]},
+ "body": {"string": payload},
+ }
+
+
+def _stored_uris(session: vcr.cassette.Cassette) -> list[str]:
+ return [request.uri for request in session.requests]
+
+
+def test_config_never_records_a_test_owned_local_upstream(tmp_path: Path):
+ recorder: Final = _recorder_with_repo_matchers(tmp_path)
+
+ with recorder.use_cassette("local_upstream.yaml", **vcr_config_dict()) as session:
+ session.append(_request_to(LOCAL_UPSTREAM), _response_served_by("the test's own server"))
+ session.append(_request_to(REMOTE_UPSTREAM), _response_served_by("a real provider"))
+
+ assert _stored_uris(session) == [REMOTE_UPSTREAM]
+ assert (tmp_path / "local_upstream.yaml").exists()
+
+
+def test_config_never_replays_a_localhost_response_an_earlier_run_stored(tmp_path: Path):
+ recorder: Final = _recorder_with_repo_matchers(tmp_path)
+ config_that_recorded_localhost: Final = vcr_config_dict() | {"ignore_localhost": False}
+
+ with recorder.use_cassette("stored_by_an_earlier_run.yaml", **config_that_recorded_localhost) as earlier_run:
+ earlier_run.append(_request_to(LOCAL_UPSTREAM), _response_served_by("an earlier run's server"))
+ earlier_run.append(_request_to(REMOTE_UPSTREAM), _response_served_by("a real provider"))
+ assert _stored_uris(earlier_run) == [LOCAL_UPSTREAM, REMOTE_UPSTREAM]
+
+ with recorder.use_cassette("stored_by_an_earlier_run.yaml", **vcr_config_dict()) as session:
+ replayable: Final = tuple(
+ bool(session.can_play_response_for(_request_to(uri))) for uri in (LOCAL_UPSTREAM, REMOTE_UPSTREAM)
+ )
+
+ assert replayable == (False, True)
+ assert _stored_uris(session) == [REMOTE_UPSTREAM]
From f12f7b5a037ab5357643ed9e56a95cc36ba0b0b5 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 16:00:32 -0700
Subject: [PATCH 21/39] test(integration): group /v1/messages contracts under
tests/integration/messages_endpoint (#43352)
* test(integration): group /v1/messages contracts under tests/integration/messages
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): make ci coverage census collect nested test dirs
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): nest /v1/messages contracts under messages_endpoint/providers
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: kerry
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
.github/scripts/assert_ci_coverage.py | 2 +-
tests/integration/README.md | 2 ++
tests/integration/_support/manifest.py | 1 +
.../test_anthropic_messages_fireworks_stop_wire.py | 0
.../providers/anthropic}/test_anthropic_advisor_wire.py | 0
.../anthropic}/test_anthropic_legacy_thinking_budget_wire.py | 0
.../anthropic}/test_anthropic_messages_timeout_wire.py | 0
.../test_anthropic_thinking_signature_retry_wire.py | 0
.../providers/anthropic}/test_anthropic_wire.py | 0
.../providers/anthropic}/test_websearch_interception_wire.py | 0
.../bedrock}/test_bedrock_invoke_tool_search_wire.py | 0
.../bedrock}/test_bedrock_messages_web_search_replay_wire.py | 0
.../gemini}/test_gemini_messages_cache_control_wire.py | 0
.../test_anthropic_messages_claude_code_cache_key_wire.py | 0
.../test_anthropic_messages_openai_bridge_wire.py | 0
.../test_anthropic_messages_openai_tools_wire.py | 0
.../responses_bridge}/test_responses_bridge_stream_options.py | 0
tests/integration/run.py | 4 ++--
18 files changed, 6 insertions(+), 3 deletions(-)
rename tests/integration/{providers => messages_endpoint/chat_bridge}/test_anthropic_messages_fireworks_stop_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_advisor_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_legacy_thinking_budget_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_messages_timeout_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_thinking_signature_retry_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_websearch_interception_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/bedrock}/test_bedrock_invoke_tool_search_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/bedrock}/test_bedrock_messages_web_search_replay_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/providers/gemini}/test_gemini_messages_cache_control_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_anthropic_messages_claude_code_cache_key_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_anthropic_messages_openai_bridge_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_anthropic_messages_openai_tools_wire.py (100%)
rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_responses_bridge_stream_options.py (100%)
diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py
index 01a01b1034b..a483dcec9d7 100644
--- a/.github/scripts/assert_ci_coverage.py
+++ b/.github/scripts/assert_ci_coverage.py
@@ -516,7 +516,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
str(path.relative_to(repo_root))
for folders in groups.values()
for folder in folders
- for path in (integration_root / folder).glob("test_*.py")
+ for path in (integration_root / folder).rglob("test_*.py")
)
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
diff --git a/tests/integration/README.md b/tests/integration/README.md
index c09904597ab..c559e7545e0 100644
--- a/tests/integration/README.md
+++ b/tests/integration/README.md
@@ -28,6 +28,8 @@ Provider contracts exercise actual TCP requests with synthetic credentials and l
Streaming checks send real HTTP transfer chunks, including one-byte partitions, fragmented tools, incomplete transfers and a cancellation barrier. They assert meaningful text, tool arguments, final usage and persisted cost. The Redis recovery case owns a separate database and Redis process, uses the supported one-second circuit-breaker recovery setting, waits for the real subscriber and verifies response data in Redis after restart. CircleCI reuses its existing Redis image for that extra process; it never pulls an image during tests
+The `messages_endpoint/` directory holds `/v1/messages` endpoint contracts: native-provider backends under `providers/` (`anthropic`, `bedrock`, `gemini`) and the translation bridges (`responses_bridge`, `chat_bridge`) at the top level. It runs in the providers shard; `run.py` selects test files recursively under each scheduled directory
+
The sdk shard exercises the SDK's own HTTP clients against local protocol peers with no gateway in the path, so a case here fails only when the client library or its wire behavior changes. The HTTP/2 case runs a hypercorn TLS peer offering h2 and http/1.1 over ALPN, drives the sync and async httpx handlers at it with `LITELLM_HTTP2` off and on, and asserts the version both the client and the peer observed on the wire. Put a test here only when it needs no proxy, database or Redis; a case that reaches the gateway belongs in one of the other shards
The extensions shard uses the built-in generic callback and guardrail transports. It checks callback correlation and credential exclusion, guardrail rewriting and denial, retained OpenAI consumers and A2A wire versions. CircleCI runs it on parallel nodes, and each node starts its own database, Redis, upstream and proxy and runs its share of the group's files serially, split by recorded timings with `circleci tests split`. Tests keep the isolation of a serial run; they still must not assume a particular set of sibling files. `run.py --list` prints a group's files and `run.py ...` runs a subset of them
diff --git a/tests/integration/_support/manifest.py b/tests/integration/_support/manifest.py
index aa0b27eceda..376c2a515b7 100644
--- a/tests/integration/_support/manifest.py
+++ b/tests/integration/_support/manifest.py
@@ -10,6 +10,7 @@ OWNED_DIRECTORIES: Final = frozenset(
"routing",
"providers",
"streaming",
+ "messages_endpoint",
"configuration",
"mcp",
"observability",
diff --git a/tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py b/tests/integration/messages_endpoint/chat_bridge/test_anthropic_messages_fireworks_stop_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py
rename to tests/integration/messages_endpoint/chat_bridge/test_anthropic_messages_fireworks_stop_wire.py
diff --git a/tests/integration/providers/test_anthropic_advisor_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_advisor_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_advisor_wire.py
rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_advisor_wire.py
diff --git a/tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_legacy_thinking_budget_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py
rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_legacy_thinking_budget_wire.py
diff --git a/tests/integration/providers/test_anthropic_messages_timeout_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_timeout_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_messages_timeout_wire.py
rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_timeout_wire.py
diff --git a/tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_thinking_signature_retry_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py
rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_thinking_signature_retry_wire.py
diff --git a/tests/integration/providers/test_anthropic_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_wire.py
rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_wire.py
diff --git a/tests/integration/providers/test_websearch_interception_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_websearch_interception_wire.py
similarity index 100%
rename from tests/integration/providers/test_websearch_interception_wire.py
rename to tests/integration/messages_endpoint/providers/anthropic/test_websearch_interception_wire.py
diff --git a/tests/integration/providers/test_bedrock_invoke_tool_search_wire.py b/tests/integration/messages_endpoint/providers/bedrock/test_bedrock_invoke_tool_search_wire.py
similarity index 100%
rename from tests/integration/providers/test_bedrock_invoke_tool_search_wire.py
rename to tests/integration/messages_endpoint/providers/bedrock/test_bedrock_invoke_tool_search_wire.py
diff --git a/tests/integration/providers/test_bedrock_messages_web_search_replay_wire.py b/tests/integration/messages_endpoint/providers/bedrock/test_bedrock_messages_web_search_replay_wire.py
similarity index 100%
rename from tests/integration/providers/test_bedrock_messages_web_search_replay_wire.py
rename to tests/integration/messages_endpoint/providers/bedrock/test_bedrock_messages_web_search_replay_wire.py
diff --git a/tests/integration/providers/test_gemini_messages_cache_control_wire.py b/tests/integration/messages_endpoint/providers/gemini/test_gemini_messages_cache_control_wire.py
similarity index 100%
rename from tests/integration/providers/test_gemini_messages_cache_control_wire.py
rename to tests/integration/messages_endpoint/providers/gemini/test_gemini_messages_cache_control_wire.py
diff --git a/tests/integration/providers/test_anthropic_messages_claude_code_cache_key_wire.py b/tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_claude_code_cache_key_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_messages_claude_code_cache_key_wire.py
rename to tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_claude_code_cache_key_wire.py
diff --git a/tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py b/tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_bridge_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py
rename to tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_bridge_wire.py
diff --git a/tests/integration/providers/test_anthropic_messages_openai_tools_wire.py b/tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_tools_wire.py
similarity index 100%
rename from tests/integration/providers/test_anthropic_messages_openai_tools_wire.py
rename to tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_tools_wire.py
diff --git a/tests/integration/providers/test_responses_bridge_stream_options.py b/tests/integration/messages_endpoint/responses_bridge/test_responses_bridge_stream_options.py
similarity index 100%
rename from tests/integration/providers/test_responses_bridge_stream_options.py
rename to tests/integration/messages_endpoint/responses_bridge/test_responses_bridge_stream_options.py
diff --git a/tests/integration/run.py b/tests/integration/run.py
index 87f93873267..8f1ff1f4a92 100644
--- a/tests/integration/run.py
+++ b/tests/integration/run.py
@@ -14,7 +14,7 @@ GROUPS: Final = MappingProxyType(
"management": ("management", "authorization", "configuration"),
"accounting": ("pricing", "spend"),
"database": ("database",),
- "providers": ("providers", "routing", "streaming"),
+ "providers": ("providers", "routing", "streaming", "messages_endpoint"),
"extensions": ("observability", "compatibility"),
"mcp": ("mcp",),
"sdk": ("sdk",),
@@ -37,7 +37,7 @@ def main() -> int:
group_files: Final = tuple(
str(path.relative_to(root))
for folder in GROUPS[options.group]
- for path in sorted((root / "tests/integration" / folder).glob("test_*.py"))
+ for path in sorted((root / "tests/integration" / folder).rglob("test_*.py"))
)
if options.list:
print("\n".join(group_files))
From e4947231058394bcfaac14f6e3f0700d4be5644c Mon Sep 17 00:00:00 2001
From: shrey-berri
Date: Sat, 26 Sep 2026 16:01:20 -0700
Subject: [PATCH 22/39] fix(params): validate stream_chunk_size once, before
any provider call (#43222)
* fix(params): validate stream_chunk_size once and carry it as typed control options
Checks stream_chunk_size at the top of completion() and acompletion(), accepts
digit strings, returns a 400 naming the param unless drop_params is set, and
stores the checked value under _litellm_control. Bedrock Converse and Invoke
read it from litellm_params; the Bedrock-only checker and the dead Invoke pops
are gone. Owned-kwarg filtering now runs through one helper everywhere.
Refs LIT-8317
* test(bedrock): drop tests for the removed stream_chunk_size_from helper
Refs LIT-8317
* fix(params): check stream_chunk_size before the MCP gateway branch
Refs LIT-8317
* fix(params): return assert_never in the exhaustive control-options match
Refs LIT-8317
* fix(params): address council review of the control options change
Read all_litellm_params live so names registered after import stay
LiteLLM-owned, make litellm_params a required keyword on the stream
wrapper hooks, give digit strings and ints the same 18-digit range,
share the default-chunking test table, test the Responses bridge through
litellm.responses, and revert formatting-only churn in existing tests.
Refs LIT-8317
* fix(params): address the second council review of control options
Keep the Responses bridge on its original all_litellm_params forwarding,
narrow _int_from_decimal_string inline so it type-checks, bound nested
huge ints in the error message, store _litellm_control only when a value
is set, simplify the parser to its single field, drop the one-caller
wrapper, and tighten the tests.
Refs LIT-8317
* fix(params): keep the 18-digit length check on stream_chunk_size strings
A 19-character string with leading zeros such as 0000000000000000001 would
otherwise pass as 1, although the rule and the error message say at most
18 digits.
Refs LIT-8317
* test(params): tidy control options tests after council sign-off
Move the Responses bridge test into the existing bridge test file, drop the
rebind test that pinned an implementation detail, assert through
stored_control_options instead of the storage key, and cover
drop_params="true" through Bedrock streaming.
Refs LIT-8317
* test(params): wrap a chunking test row that went past 120 characters
Refs LIT-8317
---
litellm/caching/caching.py | 5 +-
litellm/constants.py | 1 +
litellm/images/main.py | 11 +-
.../litellm_core_utils/get_litellm_params.py | 53 +++-
litellm/llms/base_llm/chat/transformation.py | 4 +
.../bedrock/chat/agentcore/transformation.py | 6 +-
litellm/llms/bedrock/chat/converse_handler.py | 5 +-
.../anthropic_claude3_transformation.py | 1 -
.../base_invoke_transformation.py | 13 +-
litellm/llms/bedrock/common_utils.py | 11 +-
litellm/llms/bytez/chat/transformation.py | 5 +
litellm/llms/custom_httpx/llm_http_handler.py | 2 +
litellm/llms/langgraph/chat/transformation.py | 5 +
litellm/llms/oci/chat/transformation.py | 6 +-
litellm/llms/sagemaker/chat/transformation.py | 5 +
.../vertex_ai/agent_engine/transformation.py | 5 +
litellm/main.py | 37 ++-
litellm/types/litellm_params.py | 28 ++-
litellm/utils.py | 24 +-
tests/_support/stream_chunk_size.py | 34 +--
.../providers/test_internal_params_wire.py | 9 +-
tests/unit/caching/test_caching.py | 14 ++
.../test_get_litellm_params.py | 84 ++++++-
.../test_base_invoke_transformation.py | 169 +++++++------
tests/unit/llms/bedrock/test_common_utils.py | 20 --
tests/unit/llms/chat/test_converse_handler.py | 135 ++++------
.../oci/chat/test_oci_chat_transformation.py | 2 +
.../unit/llms/oci/test_oci_coverage_boost.py | 2 +
.../test_sagemaker_chat_transformation.py | 3 +
.../test_sagemaker_nova_transformation.py | 2 +
.../test_responses_api_bridge_flag.py | 18 ++
tests/unit/test_filter_out_litellm_params.py | 20 ++
tests/unit/test_main.py | 233 +++++++++++++++++-
tests/unit/types/test_litellm_params.py | 10 +-
34 files changed, 695 insertions(+), 287 deletions(-)
delete mode 100644 tests/unit/llms/bedrock/test_common_utils.py
diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py
index a31cad4af29..d766d1a58bc 100644
--- a/litellm/caching/caching.py
+++ b/litellm/caching/caching.py
@@ -25,7 +25,7 @@ from litellm._logging import verbose_logger
from litellm.constants import CACHED_STREAMING_CHUNK_DELAY
from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
from litellm.types.caching import *
-from litellm.types.utils import EmbeddingResponse, all_litellm_params
+from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg
from .azure_blob_cache import AzureBlobCache
from .base_cache import BaseCache
@@ -377,7 +377,6 @@ class Cache:
return preset_cache_key
combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params()
- litellm_param_kwargs: Final = all_litellm_params
is_semantic_cache: Final = self._is_semantic_cache()
scope_excluded_params: Final = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset()
for param in kwargs:
@@ -387,7 +386,7 @@ class Cache:
param_value: str | None = self._get_param_value(param, kwargs)
if param_value is not None:
cache_key += f"{param}: {param_value}"
- elif param not in litellm_param_kwargs: # check if user passed in optional param - e.g. top_k
+ elif not is_litellm_owned_kwarg(param):
if litellm.enable_caching_on_provider_specific_optional_params is True: # feature flagged for now
if kwargs[param] is None:
continue # ignore None params
diff --git a/litellm/constants.py b/litellm/constants.py
index dac15c01fbf..a292b654778 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -1611,6 +1611,7 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = {
# Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.)
PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-"
INTERNAL_KWARG_PREFIX: Final = "_litellm_"
+CONTROL_OPTIONS_KEY: Final = f"{INTERNAL_KWARG_PREFIX}control"
AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech"
AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech"
diff --git a/litellm/images/main.py b/litellm/images/main.py
index 5ca8a726a69..7dc68dafecc 100644
--- a/litellm/images/main.py
+++ b/litellm/images/main.py
@@ -25,7 +25,7 @@ from litellm.llms.base_llm import BaseImageEditConfig, BaseImageGenerationConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.custom_llm import CustomLLM
-from litellm.utils import exception_type, get_litellm_params
+from litellm.utils import exception_type, filter_out_litellm_params, get_litellm_params
#################### Initialize provider clients ####################
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
@@ -52,7 +52,6 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
LITELLM_IMAGE_VARIATION_PROVIDERS,
LlmProviders,
- is_litellm_owned_kwarg,
)
from litellm.utils import (
ImageResponse,
@@ -249,9 +248,7 @@ def image_generation(
"size",
"style",
]
- non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k)
- }
+ non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params)
image_generation_config: BaseImageGenerationConfig | None = None
if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values():
@@ -755,9 +752,7 @@ def image_edit(
"style",
"async_call",
]
- non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k)
- }
+ non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params)
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
model_info: Final = kwargs.get("model_info", None)
diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py
index f7aaef3a51f..f28259a1b7f 100644
--- a/litellm/litellm_core_utils/get_litellm_params.py
+++ b/litellm/litellm_core_utils/get_litellm_params.py
@@ -1,9 +1,15 @@
+import reprlib
from collections.abc import Mapping, MutableMapping
+from dataclasses import dataclass, fields
from types import MappingProxyType
from typing import Final
+from pydantic import TypeAdapter, ValidationError
+
+from litellm.constants import CONTROL_OPTIONS_KEY
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.llms.openai.data_residency import infer_openai_data_residency
+from litellm.types.litellm_params import MAX_CONTROL_INT_DIGITS, ControlOptions
from litellm.types.router import CustomPricingLiteLLMParams
AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
@@ -70,6 +76,51 @@ OPTIONAL_KWARGS_KEYS: Final = (
# Backward-compatible alias for existing imports/tests.
_OPTIONAL_KWARGS_KEYS: Final = OPTIONAL_KWARGS_KEYS
+_CONTROL_OPTIONS: Final = TypeAdapter(ControlOptions)
+_CONTROL_OPTION_NAMES: Final = tuple(field.name for field in fields(ControlOptions))
+_MAX_SHOWN_INT_BITS: Final = 64
+_EXPECTED: Final = f"expected a positive integer of at most {MAX_CONTROL_INT_DIGITS} digits"
+
+
+class _BoundedRepr(reprlib.Repr):
+ def repr_int(self, x: int, level: int) -> str:
+ if x.bit_length() > _MAX_SHOWN_INT_BITS:
+ return f""
+ return super().repr_int(x, level)
+
+
+_BOUNDED_REPR: Final = _BoundedRepr()
+
+
+@dataclass(frozen=True, slots=True)
+class InvalidControlOption:
+ param: str
+ message: str
+
+
+def parse_control_options(kwargs: Mapping[str, object]) -> ControlOptions | InvalidControlOption:
+ given: Final = { # mutable-ok: TypeAdapter.validate_python takes a dict
+ name: kwargs[name] for name in _CONTROL_OPTION_NAMES if name in kwargs
+ }
+ try:
+ return _CONTROL_OPTIONS.validate_python(given)
+ except ValidationError as e:
+ param: Final = str(e.errors(include_url=False)[0]["loc"][0])
+ return InvalidControlOption(
+ param=param, message=f"Invalid {param}={_BOUNDED_REPR.repr(given[param])}: {_EXPECTED}"
+ )
+
+
+def stored_control_options(litellm_params: Mapping[str, object]) -> ControlOptions:
+ control: Final = litellm_params.get(CONTROL_OPTIONS_KEY)
+ return control if isinstance(control, ControlOptions) else ControlOptions()
+
+
+def with_control_options(litellm_params: Mapping[str, object], control: ControlOptions) -> dict[str, object]:
+ if control == ControlOptions():
+ return dict(litellm_params) # mutable-ok: completion() hands litellm_params to provider code typed as dict
+ return {**litellm_params, CONTROL_OPTIONS_KEY: control} # mutable-ok: same dict contract as above
+
def _get_base_model_from_litellm_call_metadata(
metadata: dict | None,
@@ -130,7 +181,6 @@ def get_litellm_params(
api_version: str | None = None,
max_retries: int | None = None,
litellm_request_debug: bool | None = None,
- stream_chunk_size: int | None = None,
**kwargs,
) -> dict:
_litellm_metadata_dict: Final = litellm_metadata if isinstance(litellm_metadata, dict) else None
@@ -193,7 +243,6 @@ def get_litellm_params(
"max_retries": max_retries,
"use_litellm_proxy": use_litellm_proxy,
"litellm_request_debug": litellm_request_debug,
- "stream_chunk_size": stream_chunk_size,
}
# Sparse extraction: only add kwargs keys that are actually present
diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py
index 948a9bc6852..f1b41a2302d 100644
--- a/litellm/llms/base_llm/chat/transformation.py
+++ b/litellm/llms/base_llm/chat/transformation.py
@@ -393,6 +393,8 @@ class BaseConfig(ABC):
client: AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "CustomStreamWrapper":
raise NotImplementedError
@@ -408,6 +410,8 @@ class BaseConfig(ABC):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "CustomStreamWrapper":
raise NotImplementedError
diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py
index 29133bcfaf9..2ad23e84e9f 100644
--- a/litellm/llms/bedrock/chat/agentcore/transformation.py
+++ b/litellm/llms/bedrock/chat/agentcore/transformation.py
@@ -5,7 +5,7 @@ https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agentcore_InvokeAgen
"""
import json
-from collections.abc import AsyncGenerator
+from collections.abc import AsyncGenerator, Mapping
from typing import TYPE_CHECKING, Any, Final, Optional, Union
from urllib.parse import quote
@@ -643,6 +643,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "CustomStreamWrapper":
"""
Simplified sync streaming - returns a generator that yields ModelResponse chunks.
@@ -862,6 +864,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
client: Optional["AsyncHTTPHandler"] = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "CustomStreamWrapper":
"""
Simplified async streaming - returns an async generator that yields ModelResponse chunks.
diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py
index 65a34f72167..28df1af4bc0 100644
--- a/litellm/llms/bedrock/chat/converse_handler.py
+++ b/litellm/llms/bedrock/chat/converse_handler.py
@@ -7,6 +7,7 @@ import litellm
from litellm.anthropic_beta_headers_manager import (
update_headers_with_filtered_beta,
)
+from litellm.litellm_core_utils.get_litellm_params import stored_control_options
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@@ -18,7 +19,7 @@ from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, pop_aws_auth_params, run_aws_signing
-from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text, stream_chunk_size_from
+from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
@@ -280,7 +281,7 @@ class BedrockConverseLLM(BaseAWSLLM):
):
## SETUP ##
stream: Final = optional_params.pop("stream", None)
- stream_chunk_size: Final = stream_chunk_size_from(litellm_params) if stream is True else None
+ stream_chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size if stream is True else None
unencoded_model_id: Final = optional_params.pop("model_id", None)
fake_stream = optional_params.pop("fake_stream", False)
json_mode: Final = optional_params.get("json_mode", False)
diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py
index a8b94fb5703..c5abb5e9a1c 100644
--- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py
+++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py
@@ -225,7 +225,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
anthropic_request.pop("model", None)
anthropic_request.pop("stream", None)
- anthropic_request.pop("stream_chunk_size", None)
apply_bedrock_invoke_structured_output(
model=model,
request_body=anthropic_request,
diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py
index 629806b58e2..9baf8110b4e 100644
--- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py
+++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py
@@ -1,6 +1,7 @@
import copy
import json
import time
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, cast, get_args
import httpx
@@ -9,6 +10,7 @@ from pydantic import TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import map_finish_reason
+from litellm.litellm_core_utils.get_litellm_params import stored_control_options
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.litellm_core_utils.prompt_templates.factory import (
cohere_message_pt,
@@ -18,7 +20,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.chat.invoke_handler import make_call, make_sync_call
-from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from
+from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.request_metadata import (
bedrock_request_metadata_headers,
merge_bedrock_invoke_headers,
@@ -180,7 +182,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
) -> dict:
## SETUP ##
stream: Final = optional_params.pop("stream", None)
- optional_params.pop("stream_chunk_size", None)
custom_prompt_dict: Final[dict] = litellm_params.pop("custom_prompt_dict", None) or {}
hf_model_name: Final = litellm_params.get("hf_model_name", None)
@@ -452,8 +453,10 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
client: AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> CustomStreamWrapper:
- chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params)
+ chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size
completion_stream, response_headers = await make_call(
client=client,
api_base=api_base,
@@ -489,11 +492,13 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> CustomStreamWrapper:
sync_client: Final = (
_get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client
)
- chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params)
+ chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size
completion_stream, response_headers = make_sync_call(
client=sync_client,
api_base=api_base,
diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py
index 5f044897b2c..ccc4309fc5d 100644
--- a/litellm/llms/bedrock/common_utils.py
+++ b/litellm/llms/bedrock/common_utils.py
@@ -18,7 +18,7 @@ if TYPE_CHECKING:
from litellm.types.llms.bedrock import BedrockCreateBatchRequest
import httpx
-from pydantic import ConfigDict, TypeAdapter, ValidationError
+from pydantic import TypeAdapter, ValidationError
import litellm
from litellm import verbose_logger
@@ -86,15 +86,6 @@ class BedrockError(BaseLLMException):
_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (*AWS_AUTH_PARAM_KEYS, "aws_region_name")
-_STREAM_CHUNK_SIZE_VALIDATOR: Final[TypeAdapter[int | None]] = TypeAdapter(int | None, config=ConfigDict(strict=True))
-
-
-def stream_chunk_size_from(litellm_params: Mapping[str, object]) -> int | None:
- raw: Final = litellm_params.get("stream_chunk_size")
- try:
- return _STREAM_CHUNK_SIZE_VALIDATOR.validate_python(raw)
- except ValidationError as e:
- raise BedrockError(status_code=400, message=f"Invalid stream_chunk_size={raw!r}. Expected int. Error: {e}")
def merge_bedrock_aws_request_params(
diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py
index e622761dd7f..5846ba560a8 100644
--- a/litellm/llms/bytez/chat/transformation.py
+++ b/litellm/llms/bytez/chat/transformation.py
@@ -1,6 +1,7 @@
import json
import time
import traceback
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
@@ -258,6 +259,8 @@ class BytezChatConfig(BaseConfig):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "BytezCustomStreamWrapper":
if client is None or isinstance(client, AsyncHTTPHandler):
client = _get_httpx_client(params={})
@@ -300,6 +303,8 @@ class BytezChatConfig(BaseConfig):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "BytezCustomStreamWrapper":
if client is None or isinstance(client, HTTPHandler):
client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={})
diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py
index 0cb1416db3f..8aa38ff3341 100644
--- a/litellm/llms/custom_httpx/llm_http_handler.py
+++ b/litellm/llms/custom_httpx/llm_http_handler.py
@@ -790,6 +790,7 @@ class BaseLLMHTTPHandler:
messages=messages,
client=client,
json_mode=json_mode,
+ litellm_params=litellm_params,
)
completion_stream, headers = self.make_sync_call(
provider_config=provider_config,
@@ -953,6 +954,7 @@ class BaseLLMHTTPHandler:
client=client,
json_mode=json_mode,
signed_json_body=signed_json_body,
+ litellm_params=litellm_params,
)
completion_stream, _response_headers = await self.make_async_call_stream_helper(
diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py
index c9388ee472f..293672f1ca9 100644
--- a/litellm/llms/langgraph/chat/transformation.py
+++ b/litellm/llms/langgraph/chat/transformation.py
@@ -9,6 +9,7 @@ Non-streaming endpoint: POST /runs/wait
"""
import json
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast
import httpx
@@ -285,6 +286,8 @@ class LangGraphConfig(BaseConfig):
client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> CustomStreamWrapper:
"""
Get a CustomStreamWrapper for synchronous streaming.
@@ -344,6 +347,8 @@ class LangGraphConfig(BaseConfig):
client: Optional["AsyncHTTPHandler"] = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> CustomStreamWrapper:
"""
Get a CustomStreamWrapper for asynchronous streaming.
diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py
index ecff823a18d..24f3ddd5162 100644
--- a/litellm/llms/oci/chat/transformation.py
+++ b/litellm/llms/oci/chat/transformation.py
@@ -10,7 +10,7 @@ implement the LiteLLM BaseConfig interface. Heavy-lifting lives in:
"""
import json
-from collections.abc import AsyncIterator, Callable, Iterator
+from collections.abc import AsyncIterator, Callable, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
@@ -642,6 +642,8 @@ class OCIChatConfig(BaseConfig):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "OCIStreamWrapper":
if client is None or isinstance(client, AsyncHTTPHandler):
client = _get_httpx_client(params={})
@@ -681,6 +683,8 @@ class OCIChatConfig(BaseConfig):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "OCIStreamWrapper":
if client is None or isinstance(client, HTTPHandler):
client = get_async_httpx_client(llm_provider=LlmProviders.OCI, params={})
diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py
index 04995f32d97..f99a3f9e1bc 100644
--- a/litellm/llms/sagemaker/chat/transformation.py
+++ b/litellm/llms/sagemaker/chat/transformation.py
@@ -7,6 +7,7 @@ LiteLLM Docs: https://docs.litellm.ai/docs/providers/aws_sagemaker#sagemaker-mes
Huggingface Docs: https://huggingface.co/docs/text-generation-inference/en/messages_api
"""
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, cast
import httpx
@@ -149,6 +150,8 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> CustomStreamWrapper:
if client is None or isinstance(client, AsyncHTTPHandler):
client = _get_httpx_client(params={})
@@ -191,6 +194,8 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
client: HTTPHandler | AsyncHTTPHandler | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> CustomStreamWrapper:
if client is None or isinstance(client, HTTPHandler):
try:
diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py
index b37bf473731..cf889c481a3 100644
--- a/litellm/llms/vertex_ai/agent_engine/transformation.py
+++ b/litellm/llms/vertex_ai/agent_engine/transformation.py
@@ -10,6 +10,7 @@ API Reference:
"""
import json
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast
import httpx
@@ -365,6 +366,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase):
client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "CustomStreamWrapper":
"""Get a CustomStreamWrapper for synchronous streaming."""
from litellm.llms.custom_httpx.http_handler import (
@@ -423,6 +426,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase):
client: Optional["AsyncHTTPHandler"] = None,
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
+ *,
+ litellm_params: Mapping[str, object],
) -> "CustomStreamWrapper":
"""Get a CustomStreamWrapper for asynchronous streaming."""
from litellm.llms.custom_httpx.http_handler import (
diff --git a/litellm/main.py b/litellm/main.py
index 8c2afe4429a..6c85adf3ae8 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -38,7 +38,7 @@ import dotenv
import httpx
import openai
from pydantic import BaseModel
-from typing_extensions import overload
+from typing_extensions import assert_never, overload
import litellm
@@ -48,6 +48,7 @@ from litellm import client
# Other utils are imported directly to avoid circular imports
from litellm.utils import (
exception_type,
+ filter_out_litellm_params,
get_litellm_params,
get_optional_params,
peek_reasoning_summary_aliases,
@@ -83,6 +84,9 @@ from litellm.litellm_core_utils.get_litellm_params import (
AWS_CREDENTIAL_KWARGS_KEYS,
OPTIONAL_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
+ InvalidControlOption,
+ parse_control_options,
+ with_control_options,
)
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
@@ -127,7 +131,7 @@ from litellm.types.completion import (
_CompletionDispatchContext,
_CompletionDispatchResult,
)
-from litellm.types.litellm_params import RetryStrategy
+from litellm.types.litellm_params import ControlOptions, RetryStrategy
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
CustomPricingLiteLLMParams,
@@ -178,7 +182,7 @@ from litellm.utils import (
from ._logging import verbose_logger
from .caching.caching import disable_cache, enable_cache, update_cache
-from .litellm_core_utils.core_helpers import safe_deep_copy
+from .litellm_core_utils.core_helpers import normalize_drop_params, safe_deep_copy
from .litellm_core_utils.fallback_utils import (
async_completion_with_fallbacks,
completion_with_fallbacks,
@@ -284,7 +288,6 @@ from .types.utils import (
LlmProviders,
PromptTokensDetails,
ProviderSpecificHeader,
- is_litellm_owned_kwarg,
)
####### ENVIRONMENT VARIABLES ###################
@@ -335,6 +338,21 @@ ovhcloud_transformation: Final = OVHCloudChatConfig()
lemonade_transformation: Final = LemonadeChatConfig()
MOCK_RESPONSE_TYPE = str | Exception | dict | ModelResponse | ModelResponseStream
+
+
+def _resolve_control_options(kwargs: Mapping[str, object], model: str) -> ControlOptions:
+ control: Final = parse_control_options(kwargs)
+ match control:
+ case ControlOptions():
+ return control
+ case InvalidControlOption(param=param, message=message):
+ if litellm.drop_params is True or normalize_drop_params(kwargs.get("drop_params")) is True:
+ return ControlOptions()
+ raise litellm.BadRequestError(message=message, model=model, llm_provider=None, body={"param": param})
+ case _:
+ return assert_never(control)
+
+
####### COMPLETION ENDPOINTS ################
@@ -501,6 +519,7 @@ async def acompletion(
loop: Final = asyncio.get_event_loop()
custom_llm_provider = kwargs.get("custom_llm_provider", None)
+ _ = _resolve_control_options(kwargs, model)
## PROMPT MANAGEMENT HOOKS ##
#########################################################
@@ -5230,6 +5249,7 @@ def completion(
# Responses API config (get_provider_responses_api_config -> None).
skip_responses_api_bridge: Final = kwargs.pop("_skip_responses_api_bridge", False)
+ control_options: Final = _resolve_control_options(kwargs, model)
skip_mcp_handler: Final = kwargs.pop("_skip_mcp_handler", False)
if not skip_mcp_handler and tools:
from litellm.responses.mcp.chat_completions_handler import acompletion_with_mcp
@@ -5370,7 +5390,6 @@ def completion(
)
######## end of unpacking kwargs ###########
non_default_params: Final = get_non_default_completion_params(kwargs=kwargs)
- litellm_params: dict[str, object] = {} # used to prevent unbound var errors
## PROMPT MANAGEMENT HOOKS ##
from litellm.integrations.anthropic_cache_control_hook import (
@@ -5622,7 +5641,7 @@ def completion(
messages = function_call_prompt(messages=messages, functions=functions_unsupported_model)
# For logging - save the values of the litellm-specific params passed in
- litellm_params = get_litellm_params(
+ requested_litellm_params: Final = get_litellm_params(
acompletion=acompletion,
api_key=api_key,
force_timeout=force_timeout,
@@ -5670,7 +5689,6 @@ def completion(
max_retries=max_retries,
timeout=timeout,
litellm_request_debug=kwargs.get("litellm_request_debug", False),
- stream_chunk_size=kwargs.get("stream_chunk_size"),
tpm=kwargs.get("tpm"),
rpm=kwargs.get("rpm"),
use_xai_oauth=kwargs.get("use_xai_oauth", False),
@@ -5683,6 +5701,7 @@ def completion(
if key in kwargs
},
)
+ litellm_params: Final = with_control_options(requested_litellm_params, control_options)
if litellm_params.get("provider_affinity_header") is not None:
try:
headers = add_provider_affinity_header(
@@ -6352,9 +6371,7 @@ def embedding(
"encoding_format",
]
default_params: Final = [*openai_params, "aembedding", "extra_headers"]
- non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in default_params and not is_litellm_owned_kwarg(k)
- }
+ non_default_params: Final = filter_out_litellm_params(kwargs, excluding=default_params)
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
model=model,
diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py
index f5ba9ebd3da..439858ea2b5 100644
--- a/litellm/types/litellm_params.py
+++ b/litellm/types/litellm_params.py
@@ -5,7 +5,10 @@ from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequenc
from dataclasses import dataclass, field, fields, is_dataclass
from itertools import chain
from types import MappingProxyType
-from typing import TYPE_CHECKING, Final, Literal, TypeAlias
+from typing import TYPE_CHECKING, Annotated, Final, Literal, TypeAlias
+
+from pydantic import BeforeValidator, Field
+from pydantic.dataclasses import dataclass as pydantic_dataclass
if TYPE_CHECKING:
import httpx
@@ -234,11 +237,31 @@ class ResponseOptions:
merge_reasoning_content_in_choices: bool | None = None
enable_json_schema_validation: bool | None = None
complete_response: bool | None = None
- stream_chunk_size: int | None = None
keepalive_seconds: float | None = None
allow_client_keepalive_override: bool | None = None
+MAX_CONTROL_INT_DIGITS: Final = 18
+
+
+def _int_from_decimal_string(value: object) -> object:
+ if isinstance(value, str) and value.isascii() and value.isdecimal() and len(value) <= MAX_CONTROL_INT_DIGITS:
+ return int(value)
+ return value
+
+
+@pydantic_dataclass(frozen=True, slots=True, kw_only=True)
+class ControlOptions:
+ stream_chunk_size: (
+ Annotated[
+ int,
+ BeforeValidator(_int_from_decimal_string),
+ Field(strict=True, gt=0, lt=10**MAX_CONTROL_INT_DIGITS),
+ ]
+ | None
+ ) = None
+
+
@dataclass(frozen=True, slots=True, kw_only=True)
class MockOptions:
mock_response: "MockResponse | None" = None
@@ -258,6 +281,7 @@ class LiteLLMOptions:
guardrails: GuardrailOptions
prompt: PromptOptions
response: ResponseOptions
+ control: ControlOptions
mock: MockOptions
diff --git a/litellm/utils.py b/litellm/utils.py
index 092fe936cf9..7ce412e818c 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -288,7 +288,7 @@ except (ImportError, AttributeError, TypeError):
# Convert to str (if necessary)
claude_json_str = json.dumps(json_data)
import importlib.metadata
-from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
+from collections.abc import AsyncIterator, Callable, Collection, Iterable, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, runtime_checkable
from typing_extensions import assert_never
@@ -4161,8 +4161,10 @@ def _remove_unsupported_params(non_default_params: dict, supported_openai_params
return non_default_params
-def filter_out_litellm_params(kwargs: Mapping[str, object]) -> dict:
- return {key: value for key, value in kwargs.items() if not is_litellm_owned_kwarg(key)}
+def filter_out_litellm_params(
+ kwargs: Mapping[str, object], excluding: Collection[str] = frozenset()
+) -> dict[str, object]:
+ return {key: value for key, value in kwargs.items() if key not in excluding and not is_litellm_owned_kwarg(key)}
def _provider_supports_vertex_params(custom_llm_provider: str) -> bool:
@@ -10132,13 +10134,8 @@ def get_standard_openai_params(params: Mapping[str, object]) -> dict:
return {k: v for k, v in params.items() if k in litellm.OPENAI_CHAT_COMPLETION_PARAMS and v is not None}
-def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict:
- openai_params: Final = litellm.OPENAI_CHAT_COMPLETION_PARAMS
- non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k)
- }
-
- return non_default_params
+def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict[str, object]:
+ return filter_out_litellm_params(kwargs, excluding=litellm.OPENAI_CHAT_COMPLETION_PARAMS)
def peek_reasoning_summary_aliases(optional_params: dict) -> object | None:
@@ -10184,13 +10181,10 @@ def strip_reasoning_summary_aliases_from_optional_params(
return op, rs_val
-def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict:
+def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict[str, object]:
from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS
- non_default_params: Final = {
- k: v for k, v in kwargs.items() if k not in OPENAI_TRANSCRIPTION_PARAMS and not is_litellm_owned_kwarg(k)
- }
- return non_default_params
+ return filter_out_litellm_params(kwargs, excluding=OPENAI_TRANSCRIPTION_PARAMS)
def add_openai_metadata(
diff --git a/tests/_support/stream_chunk_size.py b/tests/_support/stream_chunk_size.py
index 051f552e282..6e6256637f0 100644
--- a/tests/_support/stream_chunk_size.py
+++ b/tests/_support/stream_chunk_size.py
@@ -1,26 +1,28 @@
from collections.abc import Mapping
+from types import MappingProxyType
from typing import Final
-import litellm
import pytest
-from litellm.integrations.custom_logger import CustomLogger
+from litellm.constants import CONTROL_OPTIONS_KEY
+from litellm.types.litellm_params import ControlOptions
-class LitellmParamsRecorder(CustomLogger):
- def __init__(self) -> None:
- super().__init__()
- self.seen: tuple[Mapping[str, object], ...] = ()
+DEFAULT_CHUNKING_REQUESTS: Final = (
+ pytest.param(MappingProxyType({}), id="unset"),
+ pytest.param(MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": True}), id="dropped"),
+ pytest.param(
+ MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": "true"}), id="dropped_by_string_flag"
+ ),
+ pytest.param(MappingProxyType({CONTROL_OPTIONS_KEY: ControlOptions(stream_chunk_size=1)}), id="forged_options"),
+ pytest.param(MappingProxyType({CONTROL_OPTIONS_KEY: {"stream_chunk_size": 1}}), id="forged_mapping"),
+)
- def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
- params: Final = kwargs["litellm_params"]
- assert isinstance(params, Mapping)
- self.seen = (*self.seen, params)
-
-
-def record_litellm_params(monkeypatch: pytest.MonkeyPatch) -> LitellmParamsRecorder:
- recorder: Final = LitellmParamsRecorder()
- monkeypatch.setattr(litellm, "input_callback", [recorder])
- return recorder
+ROUTER_CHUNK_SIZE_CASES: Final = (
+ pytest.param(MappingProxyType({"stream_chunk_size": 64}), 64, id="int"),
+ pytest.param(MappingProxyType({"stream_chunk_size": "64"}), 64, id="digit_string"),
+ pytest.param(MappingProxyType({}), None, id="unset"),
+ pytest.param(MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": True}), None, id="dropped"),
+)
def keys_at_every_depth(value: object) -> frozenset[str]:
diff --git a/tests/integration/providers/test_internal_params_wire.py b/tests/integration/providers/test_internal_params_wire.py
index 17b0fc9d815..f9f5d1e5478 100644
--- a/tests/integration/providers/test_internal_params_wire.py
+++ b/tests/integration/providers/test_internal_params_wire.py
@@ -8,11 +8,12 @@ from collections.abc import Callable, Mapping
from pathlib import Path
from typing import Final
-import litellm
import pytest
from integration._support.upstream import INTERNAL_FIELDS
from integration._support.wire import Reply, Request, wire_server
-from tests._support.stream_chunk_size import keys_at_every_depth, record_litellm_params
+
+import litellm
+from tests._support.stream_chunk_size import keys_at_every_depth
TEXT: Final = "wire control"
OPENAI_RESPONSE: Final = {
@@ -277,13 +278,11 @@ def provider_wire_environment(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("stream", [False, True])
async def test_internal_params_never_reach_provider_body(
- monkeypatch: pytest.MonkeyPatch,
provider_wire_environment: None,
provider: str,
asynchronous: bool,
stream: bool,
) -> None:
- recorder: Final = record_litellm_params(monkeypatch)
with wire_server(_peer(provider)) as wire:
parameters: Final = {
**_request_parameters(provider, wire.url),
@@ -308,8 +307,6 @@ async def test_internal_params_never_reach_provider_body(
assert result.choices[0].message.content == TEXT
requests: Final = wire.drain()
assert len(requests) == 1
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == 64
body: Final = json.loads(requests[0].body)
keys: Final = keys_at_every_depth(body)
assert "stream_chunk_size" not in keys
diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py
index 2e4122530d8..0e0f2b7eac6 100644
--- a/tests/unit/caching/test_caching.py
+++ b/tests/unit/caching/test_caching.py
@@ -1,10 +1,12 @@
import asyncio
import logging
import re
+from typing import Final
from unittest.mock import MagicMock
import pytest
+import litellm
import litellm.caching.redis_cache as redis_cache_module
from litellm.caching.caching import Cache
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
@@ -389,3 +391,15 @@ async def test_embedding_cache_serves_base64_string_embeddings_on_repeat(monkeyp
assert embedder.provider_calls == 1, "a string embedding written to the cache must be served on repeat"
assert [item["embedding"] for item in second.data] == [item["embedding"] for item in first.data] == ["AACAPwAAAEA="]
+
+
+def test_provider_specific_cache_key_ignores_litellm_owned_kwargs(monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setattr(litellm, "enable_caching_on_provider_specific_optional_params", True)
+ cache: Final = Cache(type=LiteLLMCacheType.LOCAL)
+ request: Final = {"model": "gpt-4.1-mini", "messages": [{"role": "user", "content": "hi"}], "top_k": 5}
+
+ base_key: Final = cache.get_cache_key(**request)
+
+ assert cache.get_cache_key(**request, _litellm_control={"stream_chunk_size": 64}) == base_key
+ assert cache.get_cache_key(**request, litellm_trace_id="trace-1") == base_key
+ assert cache.get_cache_key(**{**request, "top_k": 6}) != base_key
diff --git a/tests/unit/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py
index 39bc2688ae0..9b5771092ac 100644
--- a/tests/unit/litellm_core_utils/test_get_litellm_params.py
+++ b/tests/unit/litellm_core_utils/test_get_litellm_params.py
@@ -7,12 +7,18 @@ Ensures backward compatibility after sparse kwargs extraction optimization.
from typing import Final
import pytest
+from pydantic import ValidationError
+from litellm.constants import CONTROL_OPTIONS_KEY
from litellm.litellm_core_utils.get_litellm_params import (
_OPTIONAL_KWARGS_KEYS,
+ InvalidControlOption,
_get_base_model_from_litellm_call_metadata,
get_litellm_params,
+ parse_control_options,
+ stored_control_options,
)
+from litellm.types.litellm_params import ControlOptions
NAMED_PRICE_PARAMS: Final = frozenset(
{"input_cost_per_token", "output_cost_per_token", "input_cost_per_second", "output_cost_per_second"}
@@ -90,9 +96,8 @@ class TestGetLitellmParamsKwargsExtraction:
assert "s3_endpoint_url" not in result_without_s3_kwargs
assert "s3_region_name" not in result_without_s3_kwargs
- def test_stream_chunk_size_is_carried_as_a_litellm_param(self) -> None:
- assert get_litellm_params(stream_chunk_size=64)["stream_chunk_size"] == 64
- assert get_litellm_params()["stream_chunk_size"] is None
+ def test_a_caller_supplied_control_options_key_is_not_carried(self) -> None:
+ assert CONTROL_OPTIONS_KEY not in get_litellm_params(**{CONTROL_OPTIONS_KEY: {"stream_chunk_size": 64}})
def test_s3_credential_kwargs_are_forwarded_for_s3_signing(self):
result = get_litellm_params(s3_access_key_id="s3-key", s3_secret_access_key="s3-secret")
@@ -122,6 +127,79 @@ class TestGetLitellmParamsKwargsExtraction:
assert result[key] == f"val_{key}"
+@pytest.mark.parametrize(
+ "kwargs,expected",
+ [
+ ({"stream_chunk_size": 64, "temperature": 0.2}, ControlOptions(stream_chunk_size=64)),
+ ({"stream_chunk_size": "64"}, ControlOptions(stream_chunk_size=64)),
+ ({"stream_chunk_size": None}, ControlOptions()),
+ ({"temperature": 0.2}, ControlOptions()),
+ ],
+)
+def test_control_options_are_read_from_the_request_kwargs(kwargs: dict[str, object], expected: ControlOptions) -> None:
+ assert parse_control_options(kwargs) == expected
+
+
+@pytest.mark.parametrize(
+ "raw,shown",
+ [
+ ("sixty-four", "'sixty-four'"),
+ (" 64", "' 64'"),
+ ("-1", "'-1'"),
+ ("\uff16\uff14", "'\uff16\uff14'"),
+ ("x" * 500, "'xxxxxxxxxxxx...xxxxxxxxxxxxx'"),
+ pytest.param(-(10**5000), "", id="huge_negative_int"),
+ pytest.param(-(2**64 - 1), "-18446744073709551615", id="64_bit_negative_int"),
+ pytest.param(-(2**64), "", id="65_bit_negative_int"),
+ pytest.param([-(10**5000)], "[]", id="nested_huge_int"),
+ pytest.param(10**18, "1000000000000000000", id="19_digit_int"),
+ pytest.param("1" + "0" * 18, "'1000000000000000000'", id="19_digit_string"),
+ pytest.param("9" * 5000, "'999999999999...9999999999999'", id="5000_digit_string"),
+ pytest.param("0" * 18 + "1", "'0000000000000000001'", id="19_digit_string_with_leading_zeros"),
+ (64.0, "64.0"),
+ (True, "True"),
+ (0, "0"),
+ ("0", "'0'"),
+ (-1, "-1"),
+ ],
+)
+def test_control_options_reject_a_stream_chunk_size_that_is_not_a_positive_int(raw: object, shown: str) -> None:
+ assert parse_control_options({"stream_chunk_size": raw}) == InvalidControlOption(
+ param="stream_chunk_size",
+ message=f"Invalid stream_chunk_size={shown}: expected a positive integer of at most 18 digits",
+ )
+
+
+@pytest.mark.parametrize("raw", [10**18 - 1, "9" * 18], ids=["int", "digit_string"])
+def test_control_options_accept_the_largest_18_digit_value(raw: object) -> None:
+ assert parse_control_options({"stream_chunk_size": raw}) == ControlOptions(stream_chunk_size=10**18 - 1)
+
+
+def test_control_options_accept_an_18_digit_string_with_leading_zeros() -> None:
+ assert parse_control_options({"stream_chunk_size": "0" * 17 + "1"}) == ControlOptions(stream_chunk_size=1)
+
+
+@pytest.mark.parametrize("raw", [0, -1, "sixty-four", 64.0, True])
+def test_control_options_enforce_their_rule_at_construction(raw: object) -> None:
+ with pytest.raises(ValidationError):
+ ControlOptions(stream_chunk_size=raw) # pyright: ignore[reportArgumentType] # the invalid type is the input
+
+
+@pytest.mark.parametrize(
+ "litellm_params,expected",
+ [
+ ({CONTROL_OPTIONS_KEY: ControlOptions(stream_chunk_size=64)}, ControlOptions(stream_chunk_size=64)),
+ ({}, ControlOptions()),
+ ({CONTROL_OPTIONS_KEY: {"stream_chunk_size": 64}}, ControlOptions()),
+ ({"stream_chunk_size": 64}, ControlOptions()),
+ ],
+)
+def test_stored_control_options_reads_only_the_validated_options(
+ litellm_params: dict[str, object], expected: ControlOptions
+) -> None:
+ assert stored_control_options(litellm_params) == expected
+
+
class TestGetLitellmParamsBaseModel:
"""Verify base_model resolution precedence."""
diff --git a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py
index ed172fdfbff..d1748e1b38d 100644
--- a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py
+++ b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py
@@ -1,4 +1,6 @@
import json
+from collections.abc import Mapping
+from types import MappingProxyType
from typing import Final
from unittest.mock import AsyncMock, MagicMock
@@ -6,44 +8,40 @@ import httpx
import pytest
import litellm
-from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
- AmazonAnthropicClaudeConfig,
-)
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
-from tests._support.stream_chunk_size import (
- LitellmParamsRecorder,
- keys_at_every_depth,
- record_litellm_params,
-)
+from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth
@pytest.mark.parametrize(
- "config,model",
+ "model",
[
- (AmazonInvokeConfig, "anthropic.claude-3-sonnet-20240229-v1:0"),
- (AmazonInvokeConfig, "amazon.titan-text-express-v1"),
- (AmazonInvokeConfig, "mistral.mistral-7b-instruct-v0:2"),
- (AmazonAnthropicClaudeConfig, "anthropic.claude-sonnet-4-6"),
+ "anthropic.claude-sonnet-4-6",
+ "amazon.titan-text-express-v1",
+ "mistral.mistral-7b-instruct-v0:2",
],
)
-def test_transform_request_drops_stream_chunk_size(config, model):
- """stream_chunk_size is a LiteLLM-internal knob for re-chunking the HTTP
- response stream. Leaking it into the provider request body makes Bedrock
- reject the whole request: ValidationException 'stream_chunk_size: Extra
- inputs are not permitted'."""
- request_body = config().transform_request(
- model=model,
+def test_completion_keeps_stream_chunk_size_out_of_invoke_bodies(model: str) -> None:
+ send: Final = MagicMock(return_value=httpx.Response(200))
+ client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send)))
+
+ litellm.completion(
+ model=f"bedrock/invoke/{model}",
messages=[{"role": "user", "content": "hi"}],
- optional_params={"stream": True, "stream_chunk_size": 2048, "max_tokens": 10},
- litellm_params={},
- headers={},
+ stream=True,
+ max_tokens=10,
+ client=client,
+ aws_access_key_id="fake",
+ aws_secret_access_key="fake",
+ aws_region_name="us-east-1",
+ stream_chunk_size=2048,
)
- assert "stream_chunk_size" not in json.dumps(request_body)
+ request: Final = send.call_args.args[0]
+ assert "stream_chunk_size" not in keys_at_every_depth(json.loads(request.content)), request.content
def test_validate_environment_maps_guardrail_config_to_invoke_headers():
@@ -243,10 +241,7 @@ def test_transform_response_hands_json_mode_to_nova():
assert json.loads(result.choices[0].message.content) == {"city": "Paris", "temperature": 21}
-def _stream_invoke_completion_with_spied_client(
- monkeypatch: pytest.MonkeyPatch, **kwargs
-) -> tuple[MagicMock, MagicMock, LitellmParamsRecorder]:
- recorder: Final = record_litellm_params(monkeypatch)
+def _stream_invoke_completion_with_spied_client(**kwargs: object) -> tuple[MagicMock, MagicMock]:
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.iter_bytes = MagicMock(return_value=iter([]))
@@ -263,39 +258,33 @@ def _stream_invoke_completion_with_spied_client(
aws_region_name="us-east-1",
**kwargs,
)
- return mock_response.iter_bytes, client.post, recorder
+ return mock_response.iter_bytes, client.post
-def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_invoke_body(
- monkeypatch: pytest.MonkeyPatch,
-):
- iter_bytes_spy, post_spy, recorder = _stream_invoke_completion_with_spied_client(monkeypatch, stream_chunk_size=64)
+def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_invoke_body() -> None:
+ iter_bytes_spy, post_spy = _stream_invoke_completion_with_spied_client(stream_chunk_size=64)
iter_bytes_spy.assert_called_once_with(chunk_size=64)
data: Final = post_spy.call_args.kwargs["data"]
assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == 64
-def test_completion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch):
- iter_bytes_spy, _, recorder = _stream_invoke_completion_with_spied_client(monkeypatch)
+@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS)
+def test_completion_uses_default_chunking_unless_a_valid_size_is_requested(
+ request_kwargs: Mapping[str, object],
+) -> None:
+ iter_bytes_spy, _ = _stream_invoke_completion_with_spied_client(**request_kwargs)
iter_bytes_spy.assert_called_once_with(chunk_size=None)
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] is None
-async def _astream_invoke_completion_with_spied_client(
- monkeypatch: pytest.MonkeyPatch, **kwargs
-) -> tuple[MagicMock, AsyncMock, LitellmParamsRecorder]:
+async def _astream_invoke_completion_with_spied_client(**kwargs: object) -> tuple[MagicMock, AsyncMock]:
async def _no_bytes():
return
yield b""
mock_response = MagicMock()
mock_response.status_code = 200
- recorder: Final = record_litellm_params(monkeypatch)
mock_response.aiter_bytes = MagicMock(return_value=_no_bytes())
aiter_bytes_spy = mock_response.aiter_bytes
client = AsyncHTTPHandler()
@@ -311,57 +300,49 @@ async def _astream_invoke_completion_with_spied_client(
aws_region_name="us-east-1",
**kwargs,
)
- return aiter_bytes_spy, client.post, recorder
+ return aiter_bytes_spy, client.post
@pytest.mark.asyncio
-async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_invoke_body(
- monkeypatch: pytest.MonkeyPatch,
-):
- aiter_bytes_spy, post_spy, recorder = await _astream_invoke_completion_with_spied_client(
- monkeypatch, stream_chunk_size=64
- )
+async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_invoke_body() -> None:
+ aiter_bytes_spy, post_spy = await _astream_invoke_completion_with_spied_client(stream_chunk_size=64)
aiter_bytes_spy.assert_called_once_with(chunk_size=64)
data: Final = post_spy.call_args.kwargs["data"]
assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == 64
@pytest.mark.asyncio
-async def test_acompletion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch):
- aiter_bytes_spy, _, recorder = await _astream_invoke_completion_with_spied_client(monkeypatch)
+@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS)
+async def test_acompletion_uses_default_chunking_unless_a_valid_size_is_requested(
+ request_kwargs: Mapping[str, object],
+) -> None:
+ aiter_bytes_spy, _ = await _astream_invoke_completion_with_spied_client(**request_kwargs)
aiter_bytes_spy.assert_called_once_with(chunk_size=None)
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] is None
-@pytest.mark.parametrize("stream_chunk_size,expected_chunk_size", [(64, 64), (None, None)])
-def test_router_deployment_stream_chunk_size_reaches_iter_bytes(
- monkeypatch: pytest.MonkeyPatch, stream_chunk_size, expected_chunk_size
-):
- recorder: Final = record_litellm_params(monkeypatch)
- mock_response = MagicMock()
- mock_response.status_code = 200
- mock_response.iter_bytes = MagicMock(return_value=iter([]))
- client = HTTPHandler()
- client.post = MagicMock(return_value=mock_response)
- deployment_params = {
+INVOKE_DEPLOYMENT: Final = MappingProxyType(
+ {
"model": "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0",
"aws_access_key_id": "fake",
"aws_secret_access_key": "fake",
"aws_region_name": "us-east-1",
}
- router = litellm.Router(
- model_list=[
- {
- "model_name": "invoke-chunked",
- "litellm_params": deployment_params
- | ({} if stream_chunk_size is None else {"stream_chunk_size": stream_chunk_size}),
- }
- ]
+)
+
+
+@pytest.mark.parametrize("deployment_extras,expected_chunk_size", ROUTER_CHUNK_SIZE_CASES)
+def test_router_deployment_stream_chunk_size_reaches_iter_bytes(
+ deployment_extras: Mapping[str, object], expected_chunk_size: int | None
+) -> None:
+ mock_response: Final = MagicMock()
+ mock_response.status_code = 200
+ mock_response.iter_bytes = MagicMock(return_value=iter([]))
+ client: Final = HTTPHandler()
+ client.post = MagicMock(return_value=mock_response)
+ router: Final = litellm.Router(
+ model_list=[{"model_name": "invoke-chunked", "litellm_params": {**INVOKE_DEPLOYMENT, **deployment_extras}}]
)
router.completion(
@@ -374,17 +355,11 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes(
mock_response.iter_bytes.assert_called_once_with(chunk_size=expected_chunk_size)
data: Final = client.post.call_args.kwargs["data"]
assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == stream_chunk_size
-def test_stream_wrapper_rejects_non_int_stream_chunk_size(monkeypatch: pytest.MonkeyPatch):
- record_litellm_params(monkeypatch)
- mock_response = MagicMock()
- mock_response.status_code = 200
- mock_response.iter_bytes = MagicMock(return_value=iter([]))
- client = HTTPHandler()
- client.post = MagicMock(return_value=mock_response)
+def test_invoke_stream_rejects_non_int_stream_chunk_size_before_calling_bedrock() -> None:
+ send: Final = MagicMock(return_value=httpx.Response(200))
+ client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send)))
with pytest.raises(litellm.BadRequestError):
litellm.completion(
@@ -398,4 +373,28 @@ def test_stream_wrapper_rejects_non_int_stream_chunk_size(monkeypatch: pytest.Mo
stream_chunk_size="sixty-four",
)
- client.post.assert_not_called()
+ send.assert_not_called()
+
+
+def test_router_deployment_with_a_non_numeric_stream_chunk_size_gets_a_400_before_calling_bedrock() -> None:
+ send: Final = MagicMock(return_value=httpx.Response(200))
+ client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send)))
+ router: Final = litellm.Router(
+ model_list=[
+ {
+ "model_name": "invoke-chunked",
+ "litellm_params": {**INVOKE_DEPLOYMENT, "stream_chunk_size": "sixty-four"},
+ }
+ ]
+ )
+
+ with pytest.raises(litellm.BadRequestError) as exc_info:
+ router.completion(
+ model="invoke-chunked",
+ messages=[{"role": "user", "content": "hi"}],
+ stream=True,
+ client=client,
+ )
+
+ assert exc_info.value.status_code == 400
+ send.assert_not_called()
diff --git a/tests/unit/llms/bedrock/test_common_utils.py b/tests/unit/llms/bedrock/test_common_utils.py
deleted file mode 100644
index cfcc15f186b..00000000000
--- a/tests/unit/llms/bedrock/test_common_utils.py
+++ /dev/null
@@ -1,20 +0,0 @@
-import pytest
-
-from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from
-
-
-def test_stream_chunk_size_from_absent_is_none():
- assert stream_chunk_size_from({}) is None
-
-
-def test_stream_chunk_size_from_int_is_returned():
- assert stream_chunk_size_from({"stream_chunk_size": 64}) == 64
-
-
-@pytest.mark.parametrize("bad_value", ["64", 6.4, True])
-def test_stream_chunk_size_from_rejects_non_int_with_400(bad_value):
- with pytest.raises(BedrockError) as excinfo:
- stream_chunk_size_from({"stream_chunk_size": bad_value})
-
- assert excinfo.value.status_code == 400
- assert repr(bad_value) in excinfo.value.message
diff --git a/tests/unit/llms/chat/test_converse_handler.py b/tests/unit/llms/chat/test_converse_handler.py
index cbb8e3acf78..57bb9ab771f 100644
--- a/tests/unit/llms/chat/test_converse_handler.py
+++ b/tests/unit/llms/chat/test_converse_handler.py
@@ -1,5 +1,6 @@
import json
-from collections.abc import AsyncIterator
+from collections.abc import AsyncIterator, Mapping
+from types import MappingProxyType
from typing import Final
from unittest.mock import AsyncMock, MagicMock
@@ -11,11 +12,7 @@ from litellm.llms.bedrock.chat import BedrockConverseLLM
from litellm.llms.bedrock.chat.converse_handler import make_sync_call
from litellm.llms.bedrock.common_utils import _get_all_bedrock_regions
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
-from tests._support.stream_chunk_size import (
- LitellmParamsRecorder,
- keys_at_every_depth,
- record_litellm_params,
-)
+from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth
def test_encode_model_id_with_inference_profile():
@@ -319,10 +316,7 @@ def test_completion_plumbs_stream_chunk_size_through_converse() -> None:
iter_bytes_spy.assert_called_once_with(chunk_size=2048)
-def _stream_converse_completion_with_spied_client(
- monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None = None
-) -> tuple[MagicMock, MagicMock, LitellmParamsRecorder]:
- recorder: Final = record_litellm_params(monkeypatch)
+def _stream_converse_completion_with_spied_client(**request: object) -> tuple[MagicMock, MagicMock]:
mock_response: Final = MagicMock()
mock_response.status_code = 200
mock_response.iter_bytes = MagicMock(return_value=iter([]))
@@ -337,43 +331,35 @@ def _stream_converse_completion_with_spied_client(
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
- stream_chunk_size=stream_chunk_size,
+ **request,
)
- return mock_response.iter_bytes, client.post, recorder
+ return mock_response.iter_bytes, client.post
-def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_converse_body(
- monkeypatch: pytest.MonkeyPatch,
-) -> None:
- iter_bytes_spy, post_spy, recorder = _stream_converse_completion_with_spied_client(
- monkeypatch, stream_chunk_size=64
- )
+def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_converse_body() -> None:
+ iter_bytes_spy, post_spy = _stream_converse_completion_with_spied_client(stream_chunk_size=64)
iter_bytes_spy.assert_called_once_with(chunk_size=64)
data: Final = post_spy.call_args.kwargs["data"]
assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == 64
-def test_completion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch) -> None:
- iter_bytes_spy, _, recorder = _stream_converse_completion_with_spied_client(monkeypatch)
+@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS)
+def test_completion_uses_default_chunking_unless_a_valid_size_is_requested(
+ request_kwargs: Mapping[str, object],
+) -> None:
+ iter_bytes_spy, _ = _stream_converse_completion_with_spied_client(**request_kwargs)
iter_bytes_spy.assert_called_once_with(chunk_size=None)
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] is None
-async def _astream_converse_completion_with_spied_client(
- monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None = None
-) -> tuple[MagicMock, AsyncMock, LitellmParamsRecorder]:
+async def _astream_converse_completion_with_spied_client(**request: object) -> tuple[MagicMock, AsyncMock]:
async def _no_bytes(chunk_size: int | None = None) -> AsyncIterator[bytes]:
return
yield b""
mock_response: Final = MagicMock()
mock_response.status_code = 200
- recorder: Final = record_litellm_params(monkeypatch)
mock_response.aiter_bytes = MagicMock(return_value=_no_bytes())
aiter_bytes_spy: Final = mock_response.aiter_bytes
client: Final = AsyncHTTPHandler()
@@ -387,61 +373,51 @@ async def _astream_converse_completion_with_spied_client(
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
- stream_chunk_size=stream_chunk_size,
+ **request,
)
- return aiter_bytes_spy, client.post, recorder
+ return aiter_bytes_spy, client.post
@pytest.mark.asyncio
-async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_converse_body(
- monkeypatch: pytest.MonkeyPatch,
-) -> None:
- aiter_bytes_spy, post_spy, recorder = await _astream_converse_completion_with_spied_client(
- monkeypatch, stream_chunk_size=64
- )
+async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_converse_body() -> None:
+ aiter_bytes_spy, post_spy = await _astream_converse_completion_with_spied_client(stream_chunk_size=64)
aiter_bytes_spy.assert_called_once_with(chunk_size=64)
data: Final = post_spy.call_args.kwargs["data"]
assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == 64
@pytest.mark.asyncio
-async def test_acompletion_without_stream_chunk_size_uses_default_chunking(
- monkeypatch: pytest.MonkeyPatch,
+@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS)
+async def test_acompletion_uses_default_chunking_unless_a_valid_size_is_requested(
+ request_kwargs: Mapping[str, object],
) -> None:
- aiter_bytes_spy, _, recorder = await _astream_converse_completion_with_spied_client(monkeypatch)
+ aiter_bytes_spy, _ = await _astream_converse_completion_with_spied_client(**request_kwargs)
aiter_bytes_spy.assert_called_once_with(chunk_size=None)
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] is None
-@pytest.mark.parametrize("stream_chunk_size,expected_chunk_size", [(64, 64), (None, None)])
-def test_router_deployment_stream_chunk_size_reaches_iter_bytes(
- monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None, expected_chunk_size: int | None
-) -> None:
- recorder: Final = record_litellm_params(monkeypatch)
- mock_response: Final = MagicMock()
- mock_response.status_code = 200
- mock_response.iter_bytes = MagicMock(return_value=iter([]))
- client: Final = HTTPHandler()
- client.post = MagicMock(return_value=mock_response)
- deployment_params: Final = {
+CONVERSE_DEPLOYMENT: Final = MappingProxyType(
+ {
"model": "bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
"aws_access_key_id": "fake",
"aws_secret_access_key": "fake",
"aws_region_name": "us-east-1",
}
+)
+
+
+@pytest.mark.parametrize("deployment_extras,expected_chunk_size", ROUTER_CHUNK_SIZE_CASES)
+def test_router_deployment_stream_chunk_size_reaches_iter_bytes(
+ deployment_extras: Mapping[str, object], expected_chunk_size: int | None
+) -> None:
+ mock_response: Final = MagicMock()
+ mock_response.status_code = 200
+ mock_response.iter_bytes = MagicMock(return_value=iter([]))
+ client: Final = HTTPHandler()
+ client.post = MagicMock(return_value=mock_response)
router: Final = litellm.Router(
- model_list=[
- {
- "model_name": "converse-chunked",
- "litellm_params": deployment_params
- | ({} if stream_chunk_size is None else {"stream_chunk_size": stream_chunk_size}),
- }
- ]
+ model_list=[{"model_name": "converse-chunked", "litellm_params": {**CONVERSE_DEPLOYMENT, **deployment_extras}}]
)
router.completion(
@@ -454,20 +430,18 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes(
mock_response.iter_bytes.assert_called_once_with(chunk_size=expected_chunk_size)
data: Final = client.post.call_args.kwargs["data"]
assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data
- assert len(recorder.seen) == 1
- assert recorder.seen[0]["stream_chunk_size"] == stream_chunk_size
-def test_converse_stream_rejects_non_int_stream_chunk_size_before_calling_bedrock(monkeypatch: pytest.MonkeyPatch):
- record_litellm_params(monkeypatch)
- client = HTTPHandler()
- client.post = MagicMock()
+@pytest.mark.parametrize("stream", [True, False], ids=["stream", "non_stream"])
+def test_converse_rejects_non_int_stream_chunk_size_before_calling_bedrock(stream: bool) -> None:
+ send: Final = MagicMock(return_value=httpx.Response(200))
+ client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send)))
with pytest.raises(litellm.BadRequestError):
litellm.completion(
model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
- stream=True,
+ stream=stream,
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
@@ -475,30 +449,7 @@ def test_converse_stream_rejects_non_int_stream_chunk_size_before_calling_bedroc
stream_chunk_size="sixty-four",
)
- client.post.assert_not_called()
-
-
-def test_converse_non_stream_ignores_invalid_stream_chunk_size():
- mock_response = MagicMock()
- mock_response.status_code = 200
- mock_response.json = MagicMock(return_value=_converse_response_body())
- mock_response.text = json.dumps(_converse_response_body())
- mock_response.headers = httpx.Headers()
- client = HTTPHandler()
- client.post = MagicMock(return_value=mock_response)
-
- response = litellm.completion(
- model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
- messages=[{"role": "user", "content": "hi"}],
- client=client,
- aws_access_key_id="fake",
- aws_secret_access_key="fake",
- aws_region_name="us-east-1",
- stream_chunk_size="64",
- )
-
- assert response.choices[0].message.content == "hi"
- client.post.assert_called_once()
+ send.assert_not_called()
def _bedrock_error_response(status_code: int, request_id: str) -> httpx.Response:
diff --git a/tests/unit/llms/oci/chat/test_oci_chat_transformation.py b/tests/unit/llms/oci/chat/test_oci_chat_transformation.py
index 708187b8ae1..462b1d6ea72 100644
--- a/tests/unit/llms/oci/chat/test_oci_chat_transformation.py
+++ b/tests/unit/llms/oci/chat/test_oci_chat_transformation.py
@@ -1247,6 +1247,7 @@ class TestOCIStreamingSignedBody:
mock_logging = MagicMock()
config.get_sync_custom_stream_wrapper(
+ litellm_params={},
api_base="https://example.com",
headers={},
data={"key": "value"},
@@ -1286,6 +1287,7 @@ class TestOCIStreamingSignedBody:
payload = {"key": "value"}
config.get_sync_custom_stream_wrapper(
+ litellm_params={},
api_base="https://example.com",
headers={},
data=payload,
diff --git a/tests/unit/llms/oci/test_oci_coverage_boost.py b/tests/unit/llms/oci/test_oci_coverage_boost.py
index 7c91ece70b5..8f7588c5de7 100644
--- a/tests/unit/llms/oci/test_oci_coverage_boost.py
+++ b/tests/unit/llms/oci/test_oci_coverage_boost.py
@@ -1111,6 +1111,7 @@ def test_get_sync_custom_stream_wrapper_returns_wrapper():
mock_client.post.return_value = mock_response
wrapper = config.get_sync_custom_stream_wrapper(
+ litellm_params={},
model=_GENERIC_MODEL,
custom_llm_provider="oci",
logging_obj=MagicMock(),
@@ -1143,6 +1144,7 @@ async def test_get_async_custom_stream_wrapper_returns_wrapper():
mock_client.post = AsyncMock(return_value=mock_response)
wrapper = await config.get_async_custom_stream_wrapper(
+ litellm_params={},
model=_GENERIC_MODEL,
custom_llm_provider="oci",
logging_obj=MagicMock(),
diff --git a/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py
index 697f5a7ff59..cb20b3390bb 100644
--- a/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py
+++ b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py
@@ -125,6 +125,7 @@ def test_sync_first_event_emitted_after_a_single_frame():
response = httpx.Response(200, stream=stream)
wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper(
+ litellm_params={},
model="phi-4",
custom_llm_provider="sagemaker_chat",
logging_obj=MagicMock(),
@@ -147,6 +148,7 @@ def test_sync_events_emitted_incrementally_without_bursting():
response = httpx.Response(200, stream=stream)
wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper(
+ litellm_params={},
model="phi-4",
custom_llm_provider="sagemaker_chat",
logging_obj=MagicMock(),
@@ -171,6 +173,7 @@ async def test_async_first_event_emitted_after_a_single_frame():
response = httpx.Response(200, stream=stream)
wrapper = await SagemakerChatConfig().get_async_custom_stream_wrapper(
+ litellm_params={},
model="phi-4",
custom_llm_provider="sagemaker_chat",
logging_obj=MagicMock(),
diff --git a/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py
index 5cc414819e3..d878bc70a09 100644
--- a/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py
+++ b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py
@@ -309,6 +309,7 @@ class TestSagemakerChatBackwardsCompatibility:
) as mock_csw:
mock_csw.return_value = MagicMock()
self.config.get_sync_custom_stream_wrapper(
+ litellm_params={},
model="my-hf-endpoint",
custom_llm_provider="sagemaker_chat",
logging_obj=MagicMock(),
@@ -348,6 +349,7 @@ class TestSagemakerChatBackwardsCompatibility:
mock_csw.return_value = MagicMock()
asyncio.run(
self.config.get_async_custom_stream_wrapper(
+ litellm_params={},
model="my-hf-endpoint",
custom_llm_provider="sagemaker_chat",
logging_obj=MagicMock(),
diff --git a/tests/unit/responses/test_responses_api_bridge_flag.py b/tests/unit/responses/test_responses_api_bridge_flag.py
index 642495fab86..fb1361c1f49 100644
--- a/tests/unit/responses/test_responses_api_bridge_flag.py
+++ b/tests/unit/responses/test_responses_api_bridge_flag.py
@@ -12,6 +12,7 @@ from typing import Final
from unittest.mock import MagicMock, patch
import httpx
+import openai
import pytest
import respx
@@ -592,3 +593,20 @@ class TestUseResponsesApiBridgeFlag:
mock_native_handler.assert_called_once()
assert result is not None
+
+ def test_bridge_still_rejects_an_invalid_stream_chunk_size(self) -> None:
+ send: Final = MagicMock(return_value=httpx.Response(200))
+ client: Final = openai.OpenAI(api_key="fake-key", http_client=httpx.Client(transport=httpx.MockTransport(send)))
+
+ with pytest.raises(litellm.BadRequestError) as exc_info:
+ litellm.responses(
+ model="openai/gpt-4.1-mini",
+ input="hi",
+ use_chat_completions_api=True,
+ stream_chunk_size="sixty-four",
+ client=client,
+ num_retries=0,
+ )
+
+ assert exc_info.value.param == "stream_chunk_size"
+ send.assert_not_called()
diff --git a/tests/unit/test_filter_out_litellm_params.py b/tests/unit/test_filter_out_litellm_params.py
index 72f8f5f1478..342e251d611 100644
--- a/tests/unit/test_filter_out_litellm_params.py
+++ b/tests/unit/test_filter_out_litellm_params.py
@@ -2,6 +2,10 @@
Test filter_out_litellm_params helper function.
"""
+from typing import Final
+
+
+import litellm
from litellm.utils import filter_out_litellm_params
@@ -34,3 +38,19 @@ def test_filter_out_litellm_params():
assert "litellm_trace_id" not in filtered
assert "proxy_server_request" not in filtered
assert "secret_fields" not in filtered
+
+
+def test_filter_out_litellm_params_also_drops_the_excluded_names():
+ kwargs = {"temperature": 0.2, "top_k": 5, "litellm_trace_id": "trace-1", "_litellm_control": object()}
+
+ assert filter_out_litellm_params(kwargs, excluding=("temperature",)) == {"top_k": 5}
+
+
+def test_filter_out_litellm_params_sees_a_name_appended_to_the_public_list_after_import():
+ litellm.all_litellm_params.append("registered_later")
+ try:
+ filtered: Final = filter_out_litellm_params({"registered_later": 1, "top_k": 2})
+ finally:
+ litellm.all_litellm_params.remove("registered_later")
+
+ assert filtered == {"top_k": 2}
diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py
index 7bef35d8559..e0e1fcfe105 100644
--- a/tests/unit/test_main.py
+++ b/tests/unit/test_main.py
@@ -22,10 +22,16 @@ from unittest.mock import MagicMock, patch
import litellm
from litellm import main as litellm_main
+from litellm.constants import CONTROL_OPTIONS_KEY
from litellm.integrations.custom_logger import CustomLogger
+from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
+from litellm.litellm_core_utils.get_litellm_params import stored_control_options
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
-from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage
+from litellm.types.litellm_params import ControlOptions
+from litellm.types.llms.openai import AllMessageValues
+from litellm.types.prompts.init_prompts import PromptSpec
+from litellm.types.utils import Delta, ModelResponseStream, StandardCallbackDynamicParams, StreamingChoices, Usage
@pytest.fixture(autouse=True)
@@ -4273,3 +4279,228 @@ def test_completion_rejects_untranslatable_tool_choice_with_a_400(tool_choice):
)
assert exc_info.value.status_code == 400
assert f"tool_choice={tool_choice}" in str(exc_info.value)
+
+
+@pytest.mark.parametrize("raw", ["sixty-four", 0, -1])
+def test_completion_rejects_an_invalid_stream_chunk_size_with_a_400_naming_the_param(raw: object) -> None:
+ with pytest.raises(litellm.BadRequestError) as exc_info:
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ stream_chunk_size=raw,
+ mock_response="unused",
+ )
+ assert exc_info.value.status_code == 400
+ assert exc_info.value.param == "stream_chunk_size"
+ assert f"Invalid stream_chunk_size={raw!r}: expected a positive integer of at most 18 digits" in str(exc_info.value)
+
+
+class _PromptHookRecorder(CustomPromptManagement):
+ def __init__(self, on_prompt: MagicMock) -> None:
+ super().__init__()
+ self.on_prompt: Final = on_prompt
+
+ def get_chat_completion_prompt(
+ self,
+ model: str,
+ messages: list[AllMessageValues],
+ non_default_params: dict,
+ prompt_id: str | None,
+ prompt_variables: dict | None,
+ dynamic_callback_params: StandardCallbackDynamicParams,
+ prompt_spec: PromptSpec | None = None,
+ prompt_label: str | None = None,
+ prompt_version: int | None = None,
+ ignore_prompt_manager_model: bool | None = False,
+ ignore_prompt_manager_optional_params: bool | None = False,
+ ) -> tuple[str, list[AllMessageValues], dict]:
+ self.on_prompt("sync")
+ return model, messages, non_default_params
+
+ async def async_get_chat_completion_prompt(
+ self,
+ model: str,
+ messages: list[AllMessageValues],
+ non_default_params: dict,
+ prompt_id: str | None,
+ prompt_variables: dict | None,
+ dynamic_callback_params: StandardCallbackDynamicParams,
+ litellm_logging_obj: LiteLLMLogging,
+ prompt_spec: PromptSpec | None = None,
+ tools: list[dict] | None = None,
+ prompt_label: str | None = None,
+ prompt_version: int | None = None,
+ ignore_prompt_manager_model: bool | None = False,
+ ignore_prompt_manager_optional_params: bool | None = False,
+ ) -> tuple[str, list[AllMessageValues], dict]:
+ self.on_prompt("async")
+ return model, messages, non_default_params
+
+
+async def _call_completion(is_async: bool, **kwargs: object) -> None:
+ if is_async:
+ await litellm.acompletion(**kwargs)
+ else:
+ litellm.completion(**kwargs)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("is_async,hook", [(False, "sync"), (True, "async")], ids=["completion", "acompletion"])
+async def test_the_prompt_hook_runs_when_stream_chunk_size_is_valid(
+ monkeypatch: pytest.MonkeyPatch, is_async: bool, hook: str
+) -> None:
+ on_prompt: Final = MagicMock()
+ monkeypatch.setattr(litellm, "callbacks", [_PromptHookRecorder(on_prompt)])
+
+ await _call_completion(
+ is_async,
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ prompt_id="greeting",
+ stream_chunk_size=64,
+ mock_response="hi",
+ )
+
+ on_prompt.assert_any_call(hook)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("is_async", [False, True], ids=["completion", "acompletion"])
+async def test_an_invalid_stream_chunk_size_is_rejected_before_any_prompt_hook_runs(
+ monkeypatch: pytest.MonkeyPatch, is_async: bool
+) -> None:
+ on_prompt: Final = MagicMock()
+ monkeypatch.setattr(litellm, "callbacks", [_PromptHookRecorder(on_prompt)])
+
+ with pytest.raises(litellm.BadRequestError):
+ await _call_completion(
+ is_async,
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ prompt_id="greeting",
+ stream_chunk_size="sixty-four",
+ mock_response="hi",
+ )
+
+ on_prompt.assert_not_called()
+
+
+def _completion_logging_obj(call_id: str) -> LiteLLMLogging:
+ return LiteLLMLogging(
+ model="gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ stream=False,
+ call_type="completion",
+ start_time=datetime(2026, 1, 1),
+ litellm_call_id=call_id,
+ function_id=f"{call_id}-function",
+ )
+
+
+def test_completion_carries_the_control_options_into_the_logged_litellm_params() -> None:
+ logging_obj: Final = _completion_logging_obj("control-params")
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ stream_chunk_size=64,
+ mock_response="hi",
+ litellm_logging_obj=logging_obj,
+ )
+ assert stored_control_options(logging_obj.litellm_params) == ControlOptions(stream_chunk_size=64)
+
+
+def test_completion_ignores_a_caller_supplied_control_options_key() -> None:
+ logging_obj: Final = _completion_logging_obj("control-params-injection")
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ mock_response="hi",
+ litellm_logging_obj=logging_obj,
+ **{CONTROL_OPTIONS_KEY: {"stream_chunk_size": 1}},
+ )
+ assert stored_control_options(logging_obj.litellm_params) == ControlOptions()
+
+
+@pytest.mark.parametrize("drop_params", [True, "true"])
+def test_drop_params_drops_an_invalid_stream_chunk_size_instead_of_rejecting_it(drop_params: object) -> None:
+ logging_obj: Final = _completion_logging_obj(f"drop-params-{drop_params}")
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ stream_chunk_size="sixty-four",
+ drop_params=drop_params,
+ mock_response="hi",
+ litellm_logging_obj=logging_obj,
+ )
+ assert stored_control_options(logging_obj.litellm_params) == ControlOptions()
+
+
+def test_drop_params_keeps_a_dropped_stream_chunk_size_out_of_the_provider_request(
+ respx_mock: respx.MockRouter,
+) -> None:
+ api_base: Final = "http://localhost:12346/v1"
+ mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/chat/completions.*").mock(
+ return_value=httpx.Response(
+ status_code=200,
+ json={
+ "id": "chatcmpl-drop",
+ "object": "chat.completion",
+ "created": 1712697600,
+ "model": "gpt-4.1-mini",
+ "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
+ "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
+ },
+ )
+ )
+
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ api_base=api_base,
+ api_key="fake_openai_api_key",
+ stream_chunk_size="sixty-four",
+ drop_params=True,
+ )
+
+ assert mock_route.called
+ sent: Final = json.loads(respx_mock.calls[0].request.content)
+ assert "stream_chunk_size" not in sent, sent
+ assert sent["model"] == "gpt-4.1-mini"
+
+
+@pytest.mark.asyncio
+async def test_global_drop_params_drops_an_invalid_stream_chunk_size_on_acompletion(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ monkeypatch.setattr(litellm, "drop_params", True)
+ logging_obj: Final = _completion_logging_obj("global-drop-params")
+ await litellm.acompletion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ stream_chunk_size=0,
+ mock_response="hi",
+ litellm_logging_obj=logging_obj,
+ )
+ assert stored_control_options(logging_obj.litellm_params) == ControlOptions()
+
+
+def test_completion_rejects_an_invalid_stream_chunk_size_before_the_mcp_gateway() -> None:
+ with pytest.raises(litellm.BadRequestError) as exc_info:
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ tools=[{"type": "mcp", "server_label": "gateway", "server_url": "litellm_proxy"}],
+ stream_chunk_size="sixty-four",
+ )
+ assert exc_info.value.param == "stream_chunk_size"
+
+
+def test_drop_params_false_still_rejects_an_invalid_stream_chunk_size() -> None:
+ with pytest.raises(litellm.BadRequestError):
+ litellm.completion(
+ model="openai/gpt-4.1-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ stream_chunk_size="sixty-four",
+ drop_params=False,
+ mock_response="hi",
+ )
diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py
index 33467a78aa8..a2d944fcf39 100644
--- a/tests/unit/types/test_litellm_params.py
+++ b/tests/unit/types/test_litellm_params.py
@@ -504,7 +504,8 @@ LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = {
litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2},
litellm_params.GuardrailOptions: {"guardrails": ("default",)},
litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}},
- litellm_params.ResponseOptions: {"stream_chunk_size": 64},
+ litellm_params.ResponseOptions: {"keepalive_seconds": 1.5},
+ litellm_params.ControlOptions: {"stream_chunk_size": 64},
litellm_params.MockOptions: {"mock_timeout": True},
litellm_params.CallState: {
"completion_call_id": "call",
@@ -533,7 +534,8 @@ LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = {
litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"},
litellm_params.GuardrailOptions: {"guardrails": (1,)},
litellm_params.PromptOptions: {"prompt_id": 1},
- litellm_params.ResponseOptions: {"stream_chunk_size": "64"},
+ litellm_params.ResponseOptions: {"keepalive_seconds": "1.5"},
+ litellm_params.ControlOptions: {"stream_chunk_size": "sixty-four"},
litellm_params.MockOptions: {"mock_timeout": "true"},
litellm_params.CallState: {"completion_call_id": 1},
litellm_params.AgenticLoopState: {"depth": "1"},
@@ -583,10 +585,8 @@ def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, samp
@pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id)
def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None:
- instance: Final = _leaf_instance(leaf, sample)
-
with pytest.raises(ValidationError):
- _strict_leaf_validation(leaf, instance)
+ _strict_leaf_validation(leaf, _leaf_instance(leaf, sample))
@pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id)
From b248b1c7dc12c194c4e1e176b9432391ef5ae5ab Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 16:09:09 -0700
Subject: [PATCH 23/39] fix(openai): exclude fine-tuned and custom gpt-5-chat
aliases from gpt-5 reasoning path (#43185)
* fix(openai): exclude fine-tuned and custom gpt-5-chat aliases from gpt-5 reasoning path
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(openai): keep gpt-5-chat alias regression test diff minimal
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(openai): cover temperature pass-through for gpt-5-chat aliases
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(openai): annotate locals and wrap long lines in gpt-5-chat alias test
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: mateo
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
.../llms/openai/chat/gpt_5_transformation.py | 2 +-
.../llms/openai/test_is_model_gpt_5_model.py | 34 +++++++++++++++++--
2 files changed, 32 insertions(+), 4 deletions(-)
diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py
index d0e5ff01e71..bf6b52225f2 100644
--- a/litellm/llms/openai/chat/gpt_5_transformation.py
+++ b/litellm/llms/openai/chat/gpt_5_transformation.py
@@ -69,7 +69,7 @@ GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6")
def is_gpt_reasoning_series_name(model: str) -> bool:
normalized: Final = model.split("/")[-1]
- return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and not normalized.startswith("gpt-5-chat")
+ return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and "gpt-5-chat" not in normalized
class OpenAIGPT5Config(OpenAIGPTConfig):
diff --git a/tests/unit/llms/openai/test_is_model_gpt_5_model.py b/tests/unit/llms/openai/test_is_model_gpt_5_model.py
index 0bb8425d95e..f6fef92fc6a 100644
--- a/tests/unit/llms/openai/test_is_model_gpt_5_model.py
+++ b/tests/unit/llms/openai/test_is_model_gpt_5_model.py
@@ -26,14 +26,18 @@ There are two distinct families:
``gpt-5.3-chat``, …) — ARE GPT-5 reasoning models and must stay on the GPT-5
path.
-The fix uses a prefix check (``startswith("gpt-5-chat")``) on the normalised model
-name instead of a substring check, which correctly distinguishes the two families.
+The fix uses a substring check for ``gpt-5-chat`` on the normalised model
+name (not a prefix check), which correctly distinguishes the two families.
"""
+from typing import Final
+
import pytest
-from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
+import litellm
from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
+from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
+from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
# ---------------------------------------------------------------------------
# Parametrized fixtures
@@ -73,6 +77,9 @@ NON_GPT5_MODELS = [
"gpt-5-chat", # gpt-5-chat family — regular chat path
"gpt-5-chat-latest", # gpt-5-chat family with alias suffix
"gpt-5-chat-2025-08-07", # gpt-5-chat family with date suffix
+ "ft:gpt-5-chat-latest:org:abc",
+ "my-custom-gpt-5-chat",
+ "openai/ft:gpt-5-chat-latest:org:abc",
"gpt-4",
"gpt-4o",
"gpt-4-turbo",
@@ -117,6 +124,27 @@ class TestOpenAIGPT5ConfigIsModelGpt5Model:
model
), f"Expected '{model}' (gpt-5-chat family) NOT to be on the GPT-5 path"
+ def test_responses_api_gpt5_chat_aliases_are_not_gpt5(self):
+ for model in ["ft:gpt-5-chat-latest:org:abc", "openai/my-custom-gpt-5-chat"]:
+ assert not OpenAIResponsesAPIConfig._is_gpt_5_model(
+ model
+ ), f"Expected Responses API '{model}' NOT to be on the GPT-5 path"
+
+ @pytest.mark.parametrize("model", ["ft:gpt-5-chat-latest:org:abc", "my-custom-gpt-5-chat"])
+ def test_gpt5_chat_aliases_keep_non_default_temperature(self, model: str):
+ chat_params: Final = litellm.get_optional_params(
+ model=model, custom_llm_provider="openai", temperature=0.7
+ )
+ responses_params: Final = OpenAIResponsesAPIConfig().map_openai_params(
+ response_api_optional_params={"temperature": 0.7}, model=model, drop_params=False
+ )
+ assert chat_params["temperature"] == 0.7, (
+ f"chat completions dropped or rejected temperature for '{model}'"
+ )
+ assert responses_params["temperature"] == 0.7, (
+ f"responses dropped or rejected temperature for '{model}'"
+ )
+
# Models that are gpt-5.4 or newer. main.py gates the automatic switch to the
# /v1/responses bridge (when reasoning_effort is set and tools are passed) on
From 849f3037b4f43c4e4f60233ae6a5f62905d8e192 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 26 Sep 2026 16:31:52 -0700
Subject: [PATCH 24/39] fix(langtrace): deliver spans to
app.langtrace.ai/api/trace with x-api-key (#43322)
* test(langtrace): integration test for the built-in callback wire (path, x-api-key)
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(langtrace): deliver spans to app.langtrace.ai/api/trace with x-api-key
The built-in langtrace callback posted to the dead host langtrace.ai, sent
the key as api_key instead of x-api-key, and let the OTLP endpoint
normalizer append /v1/traces to the complete /api/trace path, so every
export returned 404. Default the host to https://app.langtrace.ai, honor
LANGTRACE_API_HOST for self-hosted servers, pass the key as an exporter
header instead of a process-wide env var, and keep the /api/trace path
unchanged for traces on the langtrace callback only
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(langtrace): keep LANGTRACE_API_HOST that already ends in /api/trace
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(langtrace): deterministic audit inventory for the built-in callback (surfaces, failures, endpoints, chaos)
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(langtrace): build the repeated-request body once
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(langtrace): outage cell asserts at-most-once delivery, not span loss
The OTLP HTTP exporter reposts once on ConnectionError and the batch
processor may still be flushing the previous burst when the sink closes,
so whether the outage burst is lost or delivered after revival depends
on timing. The invariant is no duplicate and recovery on the same port
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(langtrace): delivered spans must carry the prompt in the gen_ai.content.prompt event
The upstream echoes the marker into the completion, so a whole-span
match alone would still pass if the prompt event disappeared
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(langtrace): disable model info refresh so the scripted upstream only sees completion requests
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: yucheng
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/integrations/langtrace.py | 8 +
litellm/integrations/opentelemetry.py | 4 +
litellm/litellm_core_utils/litellm_logging.py | 5 +-
.../observability/test_langtrace_delivery.py | 684 ++++++++++++++++++
tests/unit/integrations/test_opentelemetry.py | 38 +
.../test_litellm_logging.py | 42 ++
6 files changed, 779 insertions(+), 2 deletions(-)
create mode 100644 tests/integration/observability/test_langtrace_delivery.py
diff --git a/litellm/integrations/langtrace.py b/litellm/integrations/langtrace.py
index 0b4e1393ee6..53f5d2a0318 100644
--- a/litellm/integrations/langtrace.py
+++ b/litellm/integrations/langtrace.py
@@ -10,6 +10,14 @@ if TYPE_CHECKING:
else:
Span = Any
+LANGTRACE_DEFAULT_HOST: Final = "https://app.langtrace.ai"
+LANGTRACE_TRACE_PATH: Final = "/api/trace"
+
+
+def langtrace_trace_endpoint(api_host: str | None) -> str:
+ host: Final = (api_host or LANGTRACE_DEFAULT_HOST).rstrip("/")
+ return host if host.endswith(LANGTRACE_TRACE_PATH) else host + LANGTRACE_TRACE_PATH
+
class LangtraceAttributes:
"""
diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py
index 948e3113337..8d588896b2f 100644
--- a/litellm/integrations/opentelemetry.py
+++ b/litellm/integrations/opentelemetry.py
@@ -15,6 +15,7 @@ from litellm.integrations._types.open_inference import (
SpanAttributes,
)
from litellm.integrations.custom_logger import CustomLogger
+from litellm.integrations.langtrace import LANGTRACE_TRACE_PATH
from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
OTEL_SEMCONV_STABILITY_OPT_IN_ENV,
OTELGenAISemconvMixin,
@@ -3334,6 +3335,9 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
if signal_type == "traces" and "/v2/trace/otlp" in endpoint:
return endpoint
+ if signal_type == "traces" and self.callback_name == "langtrace" and endpoint.endswith(LANGTRACE_TRACE_PATH):
+ return endpoint
+
# Check if endpoint already ends with the correct signal path
target_path: Final = f"/v1/{signal_type}"
if endpoint.endswith(target_path):
diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py
index f6211869913..9ee7a7b0a7a 100644
--- a/litellm/litellm_core_utils/litellm_logging.py
+++ b/litellm/litellm_core_utils/litellm_logging.py
@@ -60,6 +60,7 @@ from litellm.integrations.arize.arize import ArizeLogger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.deepeval.deepeval import DeepEvalLogger
+from litellm.integrations.langtrace import langtrace_trace_endpoint
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.sqs import SQSLogger
from litellm.litellm_core_utils.classifier_logging import (
@@ -4920,9 +4921,9 @@ def _init_custom_logger_compatible_class(
otel_config = OpenTelemetryConfig(
exporter="otlp_http",
- endpoint="https://langtrace.ai/api/trace",
+ endpoint=langtrace_trace_endpoint(os.getenv("LANGTRACE_API_HOST")),
+ headers=f"x-api-key={os.environ['LANGTRACE_API_KEY']}",
)
- os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
for callback in _in_memory_loggers:
if isinstance(callback, OpenTelemetry) and callback.callback_name == "langtrace":
return callback
diff --git a/tests/integration/observability/test_langtrace_delivery.py b/tests/integration/observability/test_langtrace_delivery.py
new file mode 100644
index 00000000000..84c9ed5bec0
--- /dev/null
+++ b/tests/integration/observability/test_langtrace_delivery.py
@@ -0,0 +1,684 @@
+import asyncio
+import json
+import re
+import signal
+import time
+import uuid
+from collections.abc import Callable, Iterator, Sequence
+from concurrent.futures import ThreadPoolExecutor
+from contextlib import contextmanager
+from dataclasses import dataclass, field
+from itertools import repeat
+from pathlib import Path
+from queue import SimpleQueue
+from typing import Final
+
+import httpx
+import psutil
+import pytest
+import yaml
+from anthropic import Anthropic
+from integration._support.client import Gateway, eventually
+from integration._support.process import OwnedProxy, owned_proxy, owned_proxy_process
+from integration._support.wire import Reply, Request, Wire, wire_server
+from openai import AsyncOpenAI, OpenAI
+from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest
+from opentelemetry.proto.trace.v1.trace_pb2 import Span, Status
+from pydantic import TypeAdapter
+
+TRACE_PATH: Final = "/api/trace"
+STOCK_CONFIG: Final = Path("tests/integration/proxy_config.yaml")
+_PROXY_CONFIG: Final = TypeAdapter(dict[str, object])
+_SETTINGS: Final = TypeAdapter(dict[str, object])
+_MARKER: Final = re.compile(rb"lt[0-9a-f]{32}")
+_USAGE: Final = {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}
+
+
+def _marker() -> str:
+ return "lt" + uuid.uuid4().hex
+
+
+def _sse(events: Sequence[object]) -> tuple[bytes, ...]:
+ return tuple(b"data: " + json.dumps(event).encode() + b"\n\n" for event in events) + (b"data: [DONE]\n\n",)
+
+
+def _chat_reply(marker: str, stream: bool) -> Reply:
+ identity: Final = "chatcmpl-" + marker
+ if not stream:
+ return Reply(
+ body=json.dumps(
+ {
+ "id": identity,
+ "object": "chat.completion",
+ "created": 1,
+ "model": "gpt-4o-mini",
+ "choices": [
+ {
+ "index": 0,
+ "message": {"role": "assistant", "content": "echo " + marker},
+ "finish_reason": "stop",
+ }
+ ],
+ "usage": _USAGE,
+ }
+ ).encode()
+ )
+ head: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"}
+ return Reply(
+ content_type="text/event-stream",
+ chunks=_sse(
+ (
+ {**head, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "echo "}}]},
+ {**head, "choices": [{"index": 0, "delta": {"content": marker}}]},
+ {**head, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]},
+ {**head, "choices": [], "usage": _USAGE},
+ )
+ ),
+ )
+
+
+def _responses_reply(marker: str, stream: bool) -> Reply:
+ completed: Final = {
+ "id": "resp_" + marker,
+ "object": "response",
+ "created_at": 1,
+ "status": "completed",
+ "model": "gpt-4o-mini",
+ "output": [
+ {
+ "type": "message",
+ "id": "msg_" + marker,
+ "status": "completed",
+ "role": "assistant",
+ "content": [{"type": "output_text", "text": "echo " + marker, "annotations": []}],
+ }
+ ],
+ "parallel_tool_calls": False,
+ "tool_choice": "auto",
+ "tools": [],
+ "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
+ }
+ if not stream:
+ return Reply(body=json.dumps(completed).encode())
+ events: Final = (
+ {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}},
+ {
+ "type": "response.output_text.delta",
+ "item_id": "msg_" + marker,
+ "output_index": 0,
+ "content_index": 0,
+ "delta": "echo " + marker,
+ },
+ {"type": "response.completed", "response": completed},
+ )
+ return Reply(
+ content_type="text/event-stream",
+ chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
+ )
+
+
+def _upstream(request: Request) -> Reply:
+ match: Final = _MARKER.search(request.body)
+ assert match is not None, request.body[:300]
+ marker: Final = match.group().decode()
+ if b'"fail"' in request.body:
+ return Reply(status=401, body=json.dumps({"error": {"message": "bad provider key " + marker}}).encode())
+ stream: Final = json.loads(request.body).get("stream") is True
+ if request.target.endswith("/responses"):
+ return _responses_reply(marker, stream)
+ return _chat_reply(marker, stream)
+
+
+def _config(tmp_path: Path, **litellm_settings: object) -> Path:
+ config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text()))
+ settings: Final = {**_SETTINGS.validate_python(config["litellm_settings"]), **litellm_settings}
+ general: Final = {**_SETTINGS.validate_python(config["general_settings"]), "disable_model_info_refresh": True}
+ path: Final = tmp_path / "langtrace.yaml"
+ path.write_text(yaml.safe_dump({**config, "litellm_settings": settings, "general_settings": general}))
+ return path
+
+
+def _spans(batches: Sequence[Request]) -> tuple[Span, ...]:
+ return tuple(
+ span
+ for batch in batches
+ for resource_spans in ExportTraceServiceRequest.FromString(batch.body).resource_spans
+ for scope_spans in resource_spans.scope_spans
+ for span in scope_spans.spans
+ )
+
+
+def _prompt_events(span: Span) -> tuple[str, ...]:
+ return tuple(
+ attribute.value.string_value
+ for event in span.events
+ if event.name == "gen_ai.content.prompt"
+ for attribute in event.attributes
+ if attribute.key == "gen_ai.prompt"
+ )
+
+
+def _assert_prompted_with(span: Span, marker: str) -> Span:
+ prompts: Final = _prompt_events(span)
+ assert any(marker in prompt for prompt in prompts), (span.name, prompts)
+ return span
+
+
+def _spans_carrying(batches: Sequence[Request], marker: str, name: str | None = "litellm_request") -> tuple[Span, ...]:
+ return tuple(
+ span for span in _spans(batches) if name in (None, span.name) and marker.encode() in span.SerializeToString()
+ )
+
+
+def _streamed_text(sse: str, key: str) -> str:
+ def strings(node: object) -> Iterator[str]:
+ if isinstance(node, dict):
+ for field_name, value in node.items():
+ if field_name == key and isinstance(value, str):
+ yield value
+ else:
+ yield from strings(value)
+ if isinstance(node, list):
+ for item in node:
+ yield from strings(item)
+
+ events: Final = tuple(
+ json.loads(line.removeprefix("data: "))
+ for line in sse.splitlines()
+ if line.startswith("data: ") and line != "data: [DONE]"
+ )
+ return "".join(text for event in events for text in strings(event))
+
+
+def _accepted(request: Request) -> Reply:
+ return Reply(body=b'{"message":"Traces added successfully"}')
+
+
+@dataclass(frozen=True, slots=True)
+class _Sink:
+ wire: Wire
+ api_key: str
+ # mutable-ok: drain() consumes, so batches accumulate across polls
+ received: list[Request] = field(default_factory=list)
+
+ def collect(self) -> tuple[Request, ...]:
+ self.received.extend(self.wire.drain())
+ return tuple(self.received)
+
+ def spans_for(self, marker: str) -> tuple[Span, ...]:
+ return _spans_carrying(self.collect(), marker)
+
+ def assert_wire_contract(self, batches: Sequence[Request], target: str = TRACE_PATH) -> None:
+ for request in batches:
+ assert (request.method, request.target) == ("POST", target), (request.method, request.target)
+ assert request.headers.get("x-api-key") == self.api_key, request.headers
+ assert "api_key" not in request.headers, request.headers
+ assert request.headers.get("content-type") == "application/x-protobuf", request.headers
+ assert self.api_key.encode() not in b"".join(batch.body for batch in batches)
+
+ def delivered_once(self, marker: str, seconds: float = 20) -> Span:
+ batches: Final = eventually(
+ self.collect, lambda value: len(_spans_carrying(value, marker)) >= 1, seconds=seconds
+ )
+ self.assert_wire_contract(batches)
+ settled: Final = eventually(
+ self.collect, lambda value: len(_spans_carrying(value, marker)) >= 2, seconds=1, return_last_on_timeout=True
+ )
+ spans: Final = _spans_carrying(settled, marker)
+ assert len(spans) == 1, [span.span_id for span in spans]
+ return _assert_prompted_with(spans[0], marker)
+
+
+@dataclass(frozen=True, slots=True)
+class _Rig:
+ proxy: Gateway
+ model: str
+ provider: Wire
+ sink: _Sink
+
+ def provider_hits(self, marker: str) -> int:
+ return sum(marker.encode() in request.body for request in self.provider.drain())
+
+
+@contextmanager
+def _langtrace_rig(
+ gateway: Gateway,
+ tmp_path: Path,
+ *,
+ mode: str = "callbacks",
+ host: Callable[[str], str] = lambda url: url,
+ api_key: str | None = None,
+ workers: int = 1,
+ respond: Callable[[Request], Reply] = _accepted,
+ sink_port: int = 0,
+) -> Iterator[_Rig]:
+ key: Final = "synthetic-langtrace-key-" + uuid.uuid4().hex if api_key is None else api_key
+ with wire_server(_upstream) as provider, wire_server(respond, port=sink_port) as sink:
+ overrides: Final = {
+ "LANGTRACE_API_KEY": key,
+ "LANGTRACE_API_HOST": host(sink.url),
+ "OTEL_BSP_SCHEDULE_DELAY": "300",
+ }
+ config: Final = _config(tmp_path, **{mode: ["langtrace"]})
+ with (
+ owned_proxy(gateway, tmp_path, overrides, config=config, workers=workers) as proxy,
+ proxy.scenario() as scenario,
+ ):
+ yield _Rig(proxy, scenario.model(api_base=provider.url + "/v1"), provider, _Sink(sink, key))
+
+
+def _chat_httpx(rig: _Rig, marker: str, stream: bool) -> str:
+ response: Final = rig.proxy.request(
+ "POST",
+ "/v1/chat/completions",
+ {"model": rig.model, "messages": [{"role": "user", "content": marker}], "stream": stream},
+ )
+ assert response.status_code == 200, response.text
+ return _streamed_text(response.text, "content") if stream else response.text
+
+
+def _chat_openai_sync_stream(rig: _Rig, marker: str, stream: bool) -> str:
+ with OpenAI(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client:
+ chunks: Final = client.chat.completions.create(
+ model=rig.model, messages=[{"role": "user", "content": marker}], stream=True
+ )
+ return "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
+
+
+def _chat_openai_async(rig: _Rig, marker: str, stream: bool) -> str:
+ async def call() -> str:
+ async with AsyncOpenAI(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client:
+ completion: Final = await client.chat.completions.create(
+ model=rig.model, messages=[{"role": "user", "content": marker}]
+ )
+ return completion.model_dump_json()
+
+ return asyncio.run(call())
+
+
+def _messages_anthropic(rig: _Rig, marker: str, stream: bool) -> str:
+ with Anthropic(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client:
+ message: Final = client.messages.create(
+ model=rig.model, max_tokens=64, messages=[{"role": "user", "content": marker}]
+ )
+ return message.model_dump_json()
+
+
+def _messages_httpx(rig: _Rig, marker: str, stream: bool) -> str:
+ response: Final = rig.proxy.request(
+ "POST",
+ "/v1/messages",
+ {"model": rig.model, "max_tokens": 64, "messages": [{"role": "user", "content": marker}], "stream": stream},
+ )
+ assert response.status_code == 200, response.text
+ return _streamed_text(response.text, "text") if stream else response.text
+
+
+def _responses_openai(rig: _Rig, marker: str, stream: bool) -> str:
+ with OpenAI(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client:
+ return client.responses.create(model=rig.model, input=marker).model_dump_json()
+
+
+def _responses_httpx(rig: _Rig, marker: str, stream: bool) -> str:
+ response: Final = rig.proxy.request(
+ "POST", "/v1/responses", {"model": rig.model, "input": marker, "stream": stream}
+ )
+ assert response.status_code == 200, response.text
+ return _streamed_text(response.text, "delta") if stream else response.text
+
+
+@dataclass(frozen=True, slots=True)
+class _Surface:
+ call: Callable[[_Rig, str, bool], str]
+ stream: bool
+
+
+_SURFACES: Final = (
+ pytest.param(_Surface(_chat_httpx, False), id="chat-httpx"),
+ pytest.param(_Surface(_chat_openai_sync_stream, True), id="chat-openai-sync-stream"),
+ pytest.param(_Surface(_chat_openai_async, False), id="chat-openai-async"),
+ pytest.param(_Surface(_messages_anthropic, False), id="messages-anthropic"),
+ pytest.param(_Surface(_messages_httpx, True), id="messages-httpx-stream"),
+ pytest.param(_Surface(_responses_openai, False), id="responses-openai"),
+ pytest.param(_Surface(_responses_httpx, True), id="responses-httpx-stream"),
+)
+
+
+def _assert_delivered(rig: _Rig, surface: _Surface, marker: str) -> Span:
+ text: Final = surface.call(rig, marker, surface.stream)
+ assert "echo " + marker in text, text
+ assert rig.provider_hits(marker) == 1
+ return rig.sink.delivered_once(marker)
+
+
+@pytest.mark.parametrize("surface", _SURFACES)
+def test_langtrace_span_reaches_api_trace_with_x_api_key(gateway: Gateway, tmp_path: Path, surface: _Surface) -> None:
+ with _langtrace_rig(gateway, tmp_path) as rig:
+ _assert_delivered(rig, surface, _marker())
+
+
+def test_langtrace_exports_cache_hit_twin_as_its_own_span(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with _langtrace_rig(gateway, tmp_path) as rig:
+ first: Final = rig.proxy.request(
+ "POST", "/v1/chat/completions", {"model": rig.model, "messages": [{"role": "user", "content": marker}]}
+ )
+ assert first.status_code == 200, first.text
+ rig.sink.delivered_once(marker)
+ second: Final = rig.proxy.request(
+ "POST", "/v1/chat/completions", {"model": rig.model, "messages": [{"role": "user", "content": marker}]}
+ )
+ assert second.status_code == 200 and second.headers.get("x-litellm-cache-key"), second.headers
+ assert second.json()["id"] == first.json()["id"], second.text
+ assert rig.provider_hits(marker) == 1
+ batches: Final = eventually(
+ rig.sink.collect, lambda value: len(_spans_carrying(value, marker)) >= 2, seconds=20
+ )
+ rig.sink.assert_wire_contract(batches)
+ assert len(_spans_carrying(batches, marker)) == 2
+
+
+def test_langtrace_success_callback_mode_delivers(gateway: Gateway, tmp_path: Path) -> None:
+ with _langtrace_rig(gateway, tmp_path, mode="success_callback") as rig:
+ _assert_delivered(rig, _Surface(_chat_httpx, False), _marker())
+
+
+def test_langtrace_failure_callback_mode_exports_provider_error_span(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with _langtrace_rig(gateway, tmp_path, mode="failure_callback") as rig:
+ response: Final = rig.proxy.request(
+ "POST",
+ "/v1/chat/completions",
+ {"model": rig.model, "messages": [{"role": "user", "content": marker + " fail"}], "user": "fail"},
+ )
+ assert response.status_code == 401, response.text
+ assert "bad provider key " + marker in response.text, response.text
+ assert rig.provider_hits(marker) == 1
+ span: Final = rig.sink.delivered_once(marker)
+ assert span.status.code == Status.STATUS_CODE_ERROR, span.status
+
+
+@pytest.mark.parametrize("status", (403, 404), ids=("forbidden", "not-found"))
+def test_langtrace_rejecting_sink_leaves_callers_and_later_exports_intact(
+ gateway: Gateway, tmp_path: Path, status: int
+) -> None:
+ scripted: Final[SimpleQueue[int]] = SimpleQueue()
+
+ def respond(request: Request) -> Reply:
+ return Reply(status=scripted.get_nowait()) if not scripted.empty() else _accepted(request)
+
+ rejected: Final = _marker()
+ accepted: Final = _marker()
+ with _langtrace_rig(gateway, tmp_path, respond=respond) as rig:
+ scripted.put(status)
+ assert "echo " + rejected in _chat_httpx(rig, rejected, False)
+ batches: Final = eventually(
+ rig.sink.collect, lambda value: len(_spans_carrying(value, rejected)) >= 1, seconds=20
+ )
+ rig.sink.assert_wire_contract(batches)
+ assert scripted.empty()
+ assert "echo " + accepted in _chat_httpx(rig, accepted, False)
+ rig.sink.delivered_once(accepted)
+ assert rig.proxy.request("GET", "/health/liveliness").status_code == 200
+
+
+def test_langtrace_missing_api_key_logs_startup_error_and_exports_nothing(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with wire_server(_upstream) as provider, wire_server(_accepted) as sink:
+ overrides: Final = {"LANGTRACE_API_HOST": sink.url, "OTEL_BSP_SCHEDULE_DELAY": "300"}
+ with (
+ owned_proxy_process(
+ gateway,
+ tmp_path,
+ overrides,
+ config=_config(tmp_path, callbacks=["langtrace"]),
+ remove_environment=("LANGTRACE_API_KEY",),
+ ) as owned,
+ owned.gateway.scenario() as scenario,
+ ):
+ assert "LANGTRACE_API_KEY not found in environment variables" in owned.log.read_text()
+ rig: Final = _Rig(owned.gateway, scenario.model(api_base=provider.url + "/v1"), provider, _Sink(sink, ""))
+ assert "echo " + marker in _chat_httpx(rig, marker, False)
+ assert rig.provider_hits(marker) == 1
+ batches: Final = eventually(
+ rig.sink.collect, lambda value: len(value) >= 1, seconds=2, return_last_on_timeout=True
+ )
+ assert batches == (), batches
+
+
+def test_langtrace_empty_api_key_still_posts_to_api_trace(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with _langtrace_rig(gateway, tmp_path, api_key="") as rig:
+ assert "echo " + marker in _chat_httpx(rig, marker, False)
+ batches: Final = eventually(
+ rig.sink.collect, lambda value: len(_spans_carrying(value, marker)) >= 1, seconds=20
+ )
+ for request in batches:
+ assert (request.method, request.target) == ("POST", TRACE_PATH), (request.method, request.target)
+ assert "api_key" not in request.headers, request.headers
+ assert request.headers.get("x-api-key", "") == "", request.headers
+
+
+@pytest.mark.parametrize(
+ "host",
+ (lambda url: url + "/", lambda url: url + TRACE_PATH, lambda url: url + TRACE_PATH + "/"),
+ ids=("trailing-slash", "already-suffixed", "suffixed-trailing-slash"),
+)
+def test_langtrace_api_host_variants_append_api_trace_exactly_once(
+ gateway: Gateway, tmp_path: Path, host: Callable[[str], str]
+) -> None:
+ with _langtrace_rig(gateway, tmp_path, host=host) as rig:
+ _assert_delivered(rig, _Surface(_chat_httpx, False), _marker())
+
+
+def test_langtrace_logs_repeated_identical_requests_once_each(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with _langtrace_rig(gateway, tmp_path) as rig:
+ body: Final = {
+ "model": rig.model,
+ "messages": [{"role": "user", "content": marker}],
+ "cache": {"no-cache": True},
+ }
+ responses: Final = tuple(rig.proxy.request("POST", "/v1/chat/completions", body) for _ in range(2))
+ assert [response.status_code for response in responses] == [200, 200], [r.text for r in responses]
+ assert rig.provider_hits(marker) == 2
+ batches: Final = eventually(
+ rig.sink.collect, lambda value: len(_spans_carrying(value, marker)) >= 2, seconds=20
+ )
+ rig.sink.assert_wire_contract(batches)
+ settled: Final = eventually(
+ rig.sink.collect,
+ lambda value: len(_spans_carrying(value, marker)) >= 3,
+ seconds=1,
+ return_last_on_timeout=True,
+ )
+ assert len(_spans_carrying(settled, marker)) == 2
+
+
+def test_generic_otel_callback_keeps_v1_traces_suffix_on_api_trace_endpoint(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with wire_server(_upstream) as provider, wire_server(_accepted) as sink:
+ overrides: Final = {
+ "OTEL_EXPORTER": "otlp_http",
+ "OTEL_ENDPOINT": sink.url + TRACE_PATH,
+ "OTEL_HEADERS": "x-api-key=generic-otel-key",
+ "OTEL_BSP_SCHEDULE_DELAY": "300",
+ }
+ with (
+ owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["otel"])) as proxy,
+ proxy.scenario() as scenario,
+ ):
+ model: Final = scenario.model(api_base=provider.url + "/v1")
+ response: Final = proxy.request(
+ "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": marker}]}
+ )
+ assert response.status_code == 200, response.text
+ collector: Final = _Sink(sink, "generic-otel-key")
+ batches: Final = eventually(
+ collector.collect, lambda value: len(_spans_carrying(value, marker)) >= 1, seconds=20
+ )
+ collector.assert_wire_contract(batches, target=TRACE_PATH + "/v1/traces")
+
+
+def test_langtrace_otel_v2_route_still_targets_collector_v1_traces(gateway: Gateway, tmp_path: Path) -> None:
+ marker: Final = _marker()
+ with wire_server(_upstream) as provider, wire_server(_accepted) as collector:
+ overrides: Final = {
+ "LITELLM_OTEL_V2": "true",
+ "LANGTRACE_API_KEY": "unused-by-the-collector-route",
+ "OTEL_EXPORTER_OTLP_ENDPOINT": collector.url,
+ "OTEL_EXPORTER_OTLP_PROTOCOL": "http/protobuf",
+ "OTEL_EXPORTER_OTLP_HEADERS": "x-api-key=collector-key",
+ "OTEL_BSP_SCHEDULE_DELAY": "300",
+ }
+ with (
+ owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["langtrace"])) as proxy,
+ proxy.scenario() as scenario,
+ ):
+ model: Final = scenario.model(api_base=provider.url + "/v1")
+ response: Final = proxy.request(
+ "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": marker}]}
+ )
+ assert response.status_code == 200, response.text
+ sink: Final = _Sink(collector, "collector-key")
+ batches: Final = eventually(
+ sink.collect, lambda value: len(_spans_carrying(value, marker, name=None)) >= 1, seconds=20
+ )
+ sink.assert_wire_contract(batches, target="/v1/traces")
+
+
+_BURST: Final = (
+ (_chat_httpx, False),
+ (_chat_httpx, True),
+ (_messages_httpx, False),
+ (_messages_httpx, True),
+ (_responses_httpx, False),
+ (_responses_httpx, True),
+)
+
+
+def _burst_call(rig: _Rig, index: int, marker: str) -> str:
+ call, stream = _BURST[index % len(_BURST)]
+ return call(rig, marker, stream)
+
+
+def _burst(rig: _Rig, size: int) -> tuple[str, ...]:
+ markers: Final = tuple(_marker() for _ in range(size))
+ with ThreadPoolExecutor(max_workers=size) as pool:
+ texts: Final = tuple(pool.map(_burst_call, repeat(rig), range(size), markers))
+ for marker, text in zip(markers, texts, strict=True):
+ assert "echo " + marker in text, text
+ return markers
+
+
+def _assert_each_once(sink: _Sink, markers: Sequence[str], seconds: float = 30) -> None:
+ batches: Final = eventually(
+ sink.collect, lambda value: all(_spans_carrying(value, marker) for marker in markers), seconds=seconds
+ )
+ sink.assert_wire_contract(batches)
+ settled: Final = eventually(
+ sink.collect,
+ lambda value: any(len(_spans_carrying(value, marker)) > 1 for marker in markers),
+ seconds=1,
+ return_last_on_timeout=True,
+ )
+ counts: Final = {marker: len(_spans_carrying(settled, marker)) for marker in markers}
+ assert all(count == 1 for count in counts.values()), counts
+ for marker in markers:
+ _assert_prompted_with(_spans_carrying(settled, marker)[0], marker)
+
+
+def test_langtrace_two_workers_deliver_every_burst_span_exactly_once(gateway: Gateway, tmp_path: Path) -> None:
+ with _langtrace_rig(gateway, tmp_path, workers=2) as rig:
+ markers: Final = _burst(rig, 24)
+ _assert_each_once(rig.sink, markers)
+
+
+def test_langtrace_sink_outage_mid_burst_recovers_on_the_same_port(gateway: Gateway, tmp_path: Path) -> None:
+ key: Final = "synthetic-langtrace-key-" + uuid.uuid4().hex
+ with wire_server(_accepted) as probe:
+ port: Final = int(probe.url.rsplit(":", 1)[1])
+ host: Final = f"http://127.0.0.1:{port}"
+ with wire_server(_upstream) as provider:
+ overrides: Final = {"LANGTRACE_API_KEY": key, "LANGTRACE_API_HOST": host, "OTEL_BSP_SCHEDULE_DELAY": "300"}
+ with (
+ owned_proxy_process(
+ gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["langtrace"]), workers=2
+ ) as owned,
+ owned.gateway.scenario() as scenario,
+ ):
+ model: Final = scenario.model(api_base=provider.url + "/v1")
+ with wire_server(_accepted, port=port) as sink:
+ rig: Final = _Rig(owned.gateway, model, provider, _Sink(sink, key))
+ _assert_each_once(rig.sink, _burst(rig, 6))
+ log: Final = owned.log
+ failures_before: Final = log.read_text().count("Exception while exporting Span batch")
+ outage: Final = _burst(rig, 12)
+ eventually(
+ lambda: log.read_text().count("Exception while exporting Span batch"),
+ lambda value: value > failures_before,
+ seconds=20,
+ )
+ assert owned.gateway.request("GET", "/health/liveliness").status_code == 200
+ with wire_server(_accepted, port=port) as revived:
+ recovered: Final = _Rig(owned.gateway, model, provider, _Sink(revived, key))
+ _assert_each_once(recovered.sink, _burst(recovered, 6))
+ counts: Final = {marker: len(recovered.sink.spans_for(marker)) for marker in outage}
+ assert all(count <= 1 for count in counts.values()), counts
+
+
+def test_langtrace_slow_sink_does_not_delay_callers_or_duplicate_spans(gateway: Gateway, tmp_path: Path) -> None:
+ def slow(request: Request) -> Reply:
+ time.sleep(1)
+ return _accepted(request)
+
+ with _langtrace_rig(gateway, tmp_path, respond=slow) as rig:
+ started: Final = time.monotonic()
+ markers: Final = _burst(rig, 6)
+ assert time.monotonic() - started < 5
+ _assert_each_once(rig.sink, markers, seconds=40)
+
+
+def test_langtrace_survives_a_killed_worker(gateway: Gateway, tmp_path: Path) -> None:
+ key: Final = "synthetic-langtrace-key-" + uuid.uuid4().hex
+ with wire_server(_upstream) as provider, wire_server(_accepted) as sink:
+ overrides: Final = {"LANGTRACE_API_KEY": key, "LANGTRACE_API_HOST": sink.url, "OTEL_BSP_SCHEDULE_DELAY": "300"}
+ with (
+ owned_proxy_process(
+ gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["langtrace"]), workers=2
+ ) as owned,
+ httpx.Client(
+ base_url=owned.gateway.client.base_url,
+ timeout=15,
+ trust_env=False,
+ limits=httpx.Limits(max_keepalive_connections=0),
+ ) as fresh_connections,
+ ):
+ proxy: Final = Gateway(fresh_connections, owned.gateway.key, owned.gateway.upstream_url)
+ with proxy.scenario() as scenario:
+ rig: Final = _Rig(proxy, scenario.model(api_base=provider.url + "/v1"), provider, _Sink(sink, key))
+ _assert_kill_and_recovery(owned, rig)
+
+
+def _cmdline(process: psutil.Process) -> str:
+ try:
+ return " ".join(process.cmdline())
+ except psutil.Error:
+ return ""
+
+
+def _assert_kill_and_recovery(owned: OwnedProxy, rig: _Rig) -> None:
+ _assert_delivered(rig, _Surface(_chat_httpx, False), _marker())
+
+ def uvicorn_workers() -> tuple[psutil.Process, ...]:
+ return tuple(child for child in psutil.Process(owned.process.pid).children() if "spawn_main" in _cmdline(child))
+
+ workers: Final = uvicorn_workers()
+ assert len(workers) == 2, workers
+ workers[0].send_signal(signal.SIGKILL)
+ eventually(
+ uvicorn_workers,
+ lambda value: len(value) == 2 and workers[0].pid not in {child.pid for child in value},
+ seconds=30,
+ )
+ _assert_each_once(rig.sink, _burst(rig, 6))
diff --git a/tests/unit/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py
index 52eeec31e71..175bd95c263 100644
--- a/tests/unit/integrations/test_opentelemetry.py
+++ b/tests/unit/integrations/test_opentelemetry.py
@@ -1949,6 +1949,44 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
expected,
)
+ @parameterized.expand(
+ [
+ ("https://app.langtrace.ai/api/trace", "https://app.langtrace.ai/api/trace"),
+ ("https://app.langtrace.ai/api/trace/", "https://app.langtrace.ai/api/trace"),
+ ("http://localhost:3000/api/trace", "http://localhost:3000/api/trace"),
+ ]
+ )
+ def test_langtrace_callback_keeps_api_trace_endpoint_unchanged(self, input_url: str, expected: str) -> None:
+ """Langtrace ingests OTLP at the complete /api/trace path, so no /v1/traces is appended."""
+ otel = OpenTelemetry(callback_name="langtrace")
+ self.assertEqual(otel._normalize_otel_endpoint(input_url, "traces"), expected)
+
+ @parameterized.expand(
+ [
+ (None, "https://app.langtrace.ai/api/trace", "https://app.langtrace.ai/api/trace/v1/traces"),
+ ("otel", "https://app.langtrace.ai/api/trace", "https://app.langtrace.ai/api/trace/v1/traces"),
+ ("otel", "https://collector.example.com/api/trace", "https://collector.example.com/api/trace/v1/traces"),
+ ("langtrace", "https://app.langtrace.ai", "https://app.langtrace.ai/v1/traces"),
+ ]
+ )
+ def test_api_trace_exemption_is_scoped_to_langtrace_callback(
+ self, callback_name: str | None, input_url: str, expected: str
+ ) -> None:
+ """Any other callback, or a Langtrace host without the /api/trace path, keeps OTLP normalization."""
+ otel = OpenTelemetry(callback_name=callback_name)
+ self.assertEqual(otel._normalize_otel_endpoint(input_url, "traces"), expected)
+
+ def test_langtrace_callback_still_normalizes_logs_and_metrics(self) -> None:
+ otel = OpenTelemetry(callback_name="langtrace")
+ self.assertEqual(
+ otel._normalize_otel_endpoint("https://app.langtrace.ai/api/trace", "logs"),
+ "https://app.langtrace.ai/api/trace/v1/logs",
+ )
+ self.assertEqual(
+ otel._normalize_otel_endpoint("https://app.langtrace.ai/api/trace", "metrics"),
+ "https://app.langtrace.ai/api/trace/v1/metrics",
+ )
+
def test_normalize_endpoint_none(self):
"""Test that None endpoint returns None"""
otel = OpenTelemetry()
diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py
index f211505d06d..c8b02ebc790 100644
--- a/tests/unit/litellm_core_utils/test_litellm_logging.py
+++ b/tests/unit/litellm_core_utils/test_litellm_logging.py
@@ -1616,6 +1616,48 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch):
logging_module._in_memory_loggers.clear()
+@pytest.mark.parametrize(
+ ("api_host", "expected_endpoint"),
+ [
+ (None, "https://app.langtrace.ai/api/trace"),
+ ("http://langtrace.internal:3000/", "http://langtrace.internal:3000/api/trace"),
+ ("http://langtrace.internal:3000/api/trace", "http://langtrace.internal:3000/api/trace"),
+ ],
+)
+def test_langtrace_callback_exports_to_api_trace_with_x_api_key(
+ monkeypatch: pytest.MonkeyPatch, api_host: str | None, expected_endpoint: str
+) -> None:
+ """The exporter must post to Langtrace's complete /api/trace path with the key in x-api-key,
+ without leaking it into the process-wide OTEL_EXPORTER_OTLP_TRACES_HEADERS."""
+ from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
+
+ from litellm.integrations.opentelemetry import OpenTelemetry
+ from litellm.litellm_core_utils import litellm_logging as logging_module
+
+ api_key: Final = "synthetic-langtrace-key"
+ monkeypatch.setenv("LANGTRACE_API_KEY", api_key)
+ monkeypatch.delenv("LANGTRACE_API_HOST", raising=False)
+ monkeypatch.delenv("OTEL_EXPORTER_OTLP_TRACES_HEADERS", raising=False)
+ if api_host is not None:
+ monkeypatch.setenv("LANGTRACE_API_HOST", api_host)
+ logging_module._in_memory_loggers.clear()
+ try:
+ logger: Final = logging_module._init_custom_logger_compatible_class(
+ logging_integration="langtrace",
+ internal_usage_cache=None,
+ llm_router=None,
+ custom_logger_init_args={},
+ )
+ assert type(logger) is OpenTelemetry and logger.callback_name == "langtrace"
+ exporter: Final = logger._get_span_processor().span_exporter
+ assert isinstance(exporter, OTLPSpanExporter)
+ assert exporter._endpoint == expected_endpoint
+ assert exporter._headers == {"x-api-key": api_key}
+ assert "OTEL_EXPORTER_OTLP_TRACES_HEADERS" not in os.environ
+ finally:
+ logging_module._in_memory_loggers.clear()
+
+
@pytest.mark.asyncio
async def test_logging_result_for_bridge_calls(logging_obj):
"""
From b2e82cf3beeb1cf87645ef2b91d214664895e0ea Mon Sep 17 00:00:00 2001
From: yuneng-jiang
Date: Sat, 26 Sep 2026 16:44:30 -0700
Subject: [PATCH 25/39] fix(caching): stand default cache points down when
extra_body hides a direct client mark (#43341)
* fix(caching): stand default cache points down when extra_body hides a direct client mark
On native /v1/messages the extra_body envelope is dropped, so a client tool mark or root cache_control reaches Anthropic even when extra_body overrides it. The stand-down check only counted the envelope-merged view and injected two default marks on top of the client's.
* fix(caching): keep chat completions on the envelope-merged mark count for the default stand-down
Chat completions merge extra_body over the request, so a direct tool mark that extra_body replaces never reaches the provider there. Only /v1/messages, where the native transforms drop the envelope, needs to count marks on both sides.
---
.../anthropic_cache_control_hook.py | 19 +++++++---
.../test_anthropic_cache_control_hook.py | 36 +++++++++++++++++++
2 files changed, 50 insertions(+), 5 deletions(-)
diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py
index 0d6cbc2232e..1db144b5fdc 100644
--- a/litellm/integrations/anthropic_cache_control_hook.py
+++ b/litellm/integrations/anthropic_cache_control_hook.py
@@ -695,6 +695,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
tools: list | None = None,
cache_control: object = None,
request_kwargs: object = None,
+ on_messages_route: bool = False,
) -> bool:
"""Return True if the request already carries any client-supplied cache_control.
@@ -704,10 +705,14 @@ class AnthropicCacheControlHook(CustomPromptManagement):
envelope. Configured injection points are an explicit instruction and are
applied alongside the client's marks, bounded by the provider cap.
"""
- return (
- AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system)
- + AnthropicCacheControlHook.count_external_cache_breakpoints(tools, cache_control, request_kwargs)
- ) > 0
+ external_breakpoints: Final = (
+ AnthropicCacheControlHook.count_external_cache_breakpoints_on_messages_route(
+ tools, cache_control, request_kwargs
+ )
+ if on_messages_route
+ else AnthropicCacheControlHook.count_external_cache_breakpoints(tools, cache_control, request_kwargs)
+ )
+ return AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) + external_breakpoints > 0
@staticmethod
def get_default_injection_points(
@@ -719,6 +724,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
enable_prompt_caching: bool | None = None,
cache_control: object = None,
request_kwargs: object = None,
+ on_messages_route: bool = False,
) -> list[CacheControlInjectionPoint]:
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
@@ -739,7 +745,9 @@ class AnthropicCacheControlHook(CustomPromptManagement):
if not supports_anthropic_cache_control(model, custom_llm_provider):
return []
- if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools, cache_control, request_kwargs):
+ if AnthropicCacheControlHook._request_has_cache_control(
+ messages, system, tools, cache_control, request_kwargs, on_messages_route
+ ):
return []
if is_claude_code_one_shot_subagent_request(
@@ -968,6 +976,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
enable_prompt_caching=enable_prompt_caching,
cache_control=cache_control,
request_kwargs=kwargs,
+ on_messages_route=True,
)
if model is not None
else ()
diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py
index f787d370f04..1d70a21af7b 100644
--- a/tests/unit/integrations/test_anthropic_cache_control_hook.py
+++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py
@@ -2633,6 +2633,42 @@ class TestConfiguredInjectionPointsSurviveClientMarks:
assert kwargs["cache_control"] is root_cache_control
assert "litellm_gateway_injected_cache" not in kwargs["litellm_metadata"]
+ @pytest.mark.parametrize(
+ "tools,kwargs,injected",
+ [
+ ([MARKED_V1_TOOL], {"extra_body": {"tools": [UNMARKED_V1_TOOL]}}, False),
+ (None, {"cache_control": EPHEMERAL, "extra_body": {"cache_control": None}}, False),
+ ([UNMARKED_V1_TOOL], {"extra_body": {"tools": [UNMARKED_V1_TOOL]}}, True),
+ ],
+ ids=["extra_body_unmarks_direct_tool", "extra_body_nulls_root_cache_control", "no_client_mark_anywhere"],
+ )
+ def test_v1_messages_automatic_defaults_stand_down_for_a_direct_mark_extra_body_hides(
+ self, monkeypatch, tools, kwargs, injected
+ ):
+ monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
+ request_kwargs = {**copy.deepcopy(kwargs), "litellm_metadata": {}}
+
+ result_messages, result_system = self._inject(
+ copy.deepcopy(self.V1_MESSAGES), request_kwargs, tools=copy.deepcopy(tools)
+ )
+
+ assert AnthropicCacheControlHook.count_request_cache_breakpoints(result_messages, result_system) == (
+ 2 if injected else 0
+ )
+ assert ("litellm_gateway_injected_cache" in request_kwargs["litellm_metadata"]) is injected
+
+ def test_chat_automatic_defaults_apply_when_extra_body_drops_the_only_client_mark(self, monkeypatch):
+ monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
+ params = {"extra_body": {"tools": [self.UNMARKED_TOOL]}}
+
+ self._seed(params, copy.deepcopy(self.CLEAN_MESSAGES), tools=[self.MARKED_TOOL_TOP_LEVEL])
+ affinity = AnthropicCacheControlHook.messages_with_default_injections(
+ copy.deepcopy(self.CLEAN_MESSAGES), ["claude-sonnet-4-5"], tools=[self.MARKED_TOOL_TOP_LEVEL], request_kwargs=params
+ )
+
+ assert [p["index"] for p in params["cache_control_injection_points"]] == [None, -1]
+ assert AnthropicCacheControlHook.count_request_cache_breakpoints(affinity) == 2
+
@pytest.mark.parametrize(
"marked_turns,expected_system",
[(2, [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]), (3, "sys")],
From 7244040908658ce94fbb072bb77243a6ed123efd Mon Sep 17 00:00:00 2001
From: tin-berri
Date: Sat, 26 Sep 2026 16:44:49 -0700
Subject: [PATCH 26/39] fix(mcp): report reachability without stored
credentials (#43240)
---
litellm/models/mcp_server.py | 4 +-
.../mcp_server/mcp_server_manager.py | 62 ++-
litellm/proxy/_lazy_openapi_snapshot.json | 30 +-
.../mcp_management_endpoints.py | 32 +-
.../mcp_server/test_mcp_env_vars.py | 40 +-
.../mcp_server/test_mcp_server_manager.py | 361 +++++++++++++-----
.../test_mcp_management_endpoints.py | 169 +++++++-
.../_components/MCPServerCard.test.tsx | 14 +
.../mcp-servers/_components/MCPServerCard.tsx | 7 +-
.../_components/mcp_servers.test.tsx | 2 +
.../mcp-servers/_components/mcp_servers.tsx | 5 +-
.../AIHub/MCPHubTableColumns.test.tsx | 13 +-
.../components/AIHub/MCPHubTableColumns.tsx | 8 +-
.../src/components/mcp_tools/types.tsx | 4 +-
.../src/components/networking.test.ts | 24 ++
.../src/components/networking.tsx | 1 +
ui/litellm-dashboard/src/lib/http/schema.d.ts | 11 +-
17 files changed, 644 insertions(+), 143 deletions(-)
diff --git a/litellm/models/mcp_server.py b/litellm/models/mcp_server.py
index 9125d708e79..efc8574932f 100644
--- a/litellm/models/mcp_server.py
+++ b/litellm/models/mcp_server.py
@@ -73,9 +73,9 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
mcp_info: MCPInfo | None = None
static_headers: dict[str, str] | None = None
env_vars: list[MCPEnvVar] | None = None
- status: Literal["healthy", "unhealthy", "unknown"] | None = Field(
+ status: Literal["healthy", "reachable", "unhealthy", "unknown"] | None = Field(
default="unknown",
- description="Health status: 'healthy', 'unhealthy', 'unknown'",
+ description="Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)",
)
last_health_check: datetime | None = None
health_check_error: str | None = None
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index 8befc99cad4..d0d9100971d 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -897,6 +897,41 @@ def _sanitized_error_text(exc: Exception) -> str:
return re.sub(r"https?://\S+", "", str(exc))[:200]
+async def _mcp_server_reachability(
+ server: MCPServer, *, timeout: float
+) -> tuple[Literal["reachable", "unhealthy", "unknown"], str | None]:
+ if server.transport not in (MCPTransport.http, MCPTransport.sse) or not server.url:
+ return "unknown", "Server reachability requires an HTTP or SSE URL"
+ try:
+ url: Final = httpx.URL(server.url)
+ except (httpx.InvalidURL, ValueError):
+ return "unknown", "Server reachability requires an HTTP URL without embedded credentials"
+ if url.scheme not in ("http", "https") or not url.host or url.userinfo:
+ return "unknown", "Server reachability requires an HTTP URL without embedded credentials"
+
+ async def probe() -> None:
+ handler: Final = get_async_httpx_client(llm_provider="mcp_reachability")
+ async with handler.client.stream(
+ "GET",
+ url,
+ headers={"Accept": "text/event-stream, application/json"},
+ auth=None,
+ follow_redirects=False,
+ timeout=timeout,
+ ):
+ pass
+
+ try:
+ await asyncio.wait_for(probe(), timeout=timeout)
+ except (asyncio.TimeoutError, httpx.TimeoutException):
+ return "unhealthy", f"Reachability check timed out after {timeout} seconds"
+ except asyncio.CancelledError:
+ return "unknown", "Reachability check was cancelled"
+ except Exception as exc:
+ return "unhealthy", f"Reachability check failed ({type(exc).__name__})"
+ return "reachable", None
+
+
async def _openapi_spec_health(
spec_path: str, *, timeout: float
) -> tuple[Literal["healthy", "unhealthy", "unknown"], str | None]:
@@ -6986,13 +7021,9 @@ class MCPServerManager:
)
)
- status: Literal["healthy", "unhealthy", "unknown"] = "unknown"
+ status: Literal["healthy", "reachable", "unhealthy", "unknown"] = "unknown"
health_check_error = None
- # Check if we should skip health check based on auth configuration
- should_skip_health_check = False
-
- # Skip if server requires per-user authentication (OAuth2 or passthrough auth)
if (
server.requires_per_user_auth
or (
@@ -7003,9 +7034,8 @@ class MCPServerManager:
)
or self._references_per_user_env_var(server)
):
- should_skip_health_check = True
-
- if not should_skip_health_check:
+ status, health_check_error = await _mcp_server_reachability(server, timeout=MCP_HEALTH_CHECK_TIMEOUT)
+ else:
try:
resolved_static_headers: Final = await self._resolve_static_headers_with_env_vars(
server=server,
@@ -7081,6 +7111,8 @@ class MCPServerManager:
self,
user_api_key_auth: UserAPIKeyAuth | None = None,
server_ids: list[str] | None = None,
+ *,
+ checked_server_ids: frozenset[str] = frozenset(),
) -> list[LiteLLM_MCPServerTable]:
"""
Get all MCP servers that the user has access to, with health status and team information.
@@ -7105,7 +7137,7 @@ class MCPServerManager:
# Check all accessible servers
target_server_ids = allowed_server_ids
- return await self._run_health_checks(target_server_ids)
+ return await self._run_health_checks([sid for sid in target_server_ids if sid not in checked_server_ids])
async def get_all_allowed_mcp_servers(
self,
@@ -7236,9 +7268,15 @@ class MCPServerManager:
if not target_server_ids:
return []
- tasks: Final = [self.health_check_server(server_id) for server_id in target_server_ids]
- results: Final = await asyncio.gather(*tasks)
- return [server for server in results if server is not None]
+ unique_server_ids: Final = tuple(dict.fromkeys(target_server_ids))
+ batch_size: Final = 10
+ batches: Final = [
+ await asyncio.gather(
+ *(self.health_check_server(server_id) for server_id in unique_server_ids[offset : offset + batch_size])
+ )
+ for offset in range(0, len(unique_server_ids), batch_size)
+ ]
+ return [server for batch in batches for server in batch if server is not None]
global_mcp_server_manager: Final[MCPServerManager] = MCPServerManager()
diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index 3d44315341b..ded4db6d2aa 100644
--- a/litellm/proxy/_lazy_openapi_snapshot.json
+++ b/litellm/proxy/_lazy_openapi_snapshot.json
@@ -32168,6 +32168,7 @@
{
"enum": [
"healthy",
+ "reachable",
"unhealthy",
"unknown"
],
@@ -32178,7 +32179,7 @@
}
],
"default": "unknown",
- "description": "Health status: 'healthy', 'unhealthy', 'unknown'",
+ "description": "Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)",
"title": "Status"
},
"subject_token_type": {
@@ -35224,6 +35225,7 @@
{
"enum": [
"healthy",
+ "reachable",
"unhealthy",
"unknown"
],
@@ -35234,7 +35236,7 @@
}
],
"default": "unknown",
- "description": "Health status: 'healthy', 'unhealthy', 'unknown'",
+ "description": "Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)",
"title": "Status"
},
"subject_token_type": {
@@ -38095,6 +38097,18 @@
"description": "Server IDs to check. If not provided, checks all accessible servers.",
"title": "Server Ids"
}
+ },
+ {
+ "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.",
+ "in": "query",
+ "name": "include_reachability",
+ "required": false,
+ "schema": {
+ "default": false,
+ "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.",
+ "title": "Include Reachability",
+ "type": "boolean"
+ }
}
],
"responses": {
@@ -38389,6 +38403,18 @@
"title": "Server Id",
"type": "string"
}
+ },
+ {
+ "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.",
+ "in": "query",
+ "name": "include_reachability",
+ "required": false,
+ "schema": {
+ "default": false,
+ "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.",
+ "title": "Include Reachability",
+ "type": "boolean"
+ }
}
],
"responses": {
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index d92443b104b..346b1a75ac6 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -1296,6 +1296,12 @@ if MCP_AVAILABLE:
return redacted_mcp_servers
+ def _mcp_health_status_for_response(
+ health_status: Literal["healthy", "reachable", "unhealthy", "unknown"] | None,
+ include_reachability: bool,
+ ) -> Literal["healthy", "reachable", "unhealthy", "unknown"] | None:
+ return "unknown" if health_status == "reachable" and not include_reachability else health_status
+
@router.get(
"/server/health",
description="Health check for MCP servers",
@@ -1307,6 +1313,10 @@ if MCP_AVAILABLE:
description="Server IDs to check. If not provided, checks all accessible servers.",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+ include_reachability: Annotated[
+ bool,
+ Query(description="Allow the 'reachable' status for responding servers whose authentication is unchecked."),
+ ] = False,
):
"""
Perform health checks on one or more MCP servers.
@@ -1331,21 +1341,31 @@ if MCP_AVAILABLE:
if user_mcp_management_mode == "view_all" and not _is_restricted_virtual_key_request(user_api_key_dict):
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(server_ids=server_ids)
- return [{"server_id": server.server_id, "status": server.status} for server in servers]
+ return [
+ {
+ "server_id": server.server_id,
+ "status": _mcp_health_status_for_response(server.status, include_reachability),
+ }
+ for server in servers
+ ]
auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict)
- server_status_map: Final[dict[str, Literal["healthy", "unhealthy", "unknown"] | None]] = {}
+ server_status_map: Final[dict[str, Literal["healthy", "reachable", "unhealthy", "unknown"] | None]] = {}
for auth_context in auth_contexts:
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams(
user_api_key_auth=auth_context,
server_ids=server_ids,
+ checked_server_ids=frozenset(server_status_map),
)
for server in servers:
if server.server_id not in server_status_map:
server_status_map[server.server_id] = server.status
- return [{"server_id": server_id, "status": status} for server_id, status in server_status_map.items()]
+ return [
+ {"server_id": server_id, "status": _mcp_health_status_for_response(status, include_reachability)}
+ for server_id, status in server_status_map.items()
+ ]
@router.post(
"/server/register",
@@ -1615,6 +1635,10 @@ if MCP_AVAILABLE:
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+ include_reachability: Annotated[
+ bool,
+ Query(description="Allow the 'reachable' status for responding servers whose authentication is unchecked."),
+ ] = False,
):
"""
Get the info on the mcp server specified by the `server_id`
@@ -1672,7 +1696,7 @@ if MCP_AVAILABLE:
try:
health_result: Final = await global_mcp_server_manager.health_check_server(server_id)
# Update the server object with health check results
- mcp_server.status = health_result.status if health_result.status else "unknown"
+ mcp_server.status = _mcp_health_status_for_response(health_result.status, include_reachability) or "unknown"
mcp_server.last_health_check = health_result.last_health_check
mcp_server.health_check_error = health_result.health_check_error
except Exception as e:
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py
index 93b894f7645..fff4221f243 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py
@@ -6,7 +6,13 @@ connection. The DB-backed per-user flow is exercised in higher-level
tests in tests/mcp_tests.
"""
+from typing import Final
+from unittest.mock import AsyncMock
+
import pytest
+from respx import MockRouter
+
+from litellm.types.mcp_server.mcp_server_manager import MCPServer
# Look up these names lazily on every access. Tests in this directory call
# ``importlib.reload`` on the utils module to exercise registration logic,
@@ -568,7 +574,7 @@ async def test_resolve_static_headers_user_value_wins_over_empty_global(
assert headers == {"Authorization": "Bearer user-secret"}
-# ── health-check skip for per-user-env-var-backed headers ──────────────────
+# ── health-check reachability for per-user-env-var-backed headers ───────────
@pytest.mark.parametrize(
@@ -615,32 +621,26 @@ def test_references_per_user_env_var(static_headers, env_vars, expected):
@pytest.mark.asyncio
-async def test_health_check_skips_servers_referencing_per_user_env_var(
- mock_server, monkeypatch
-):
- """A userless health probe cannot fill per-user ${NAME} placeholders, so a
- server whose static_headers reference one must report 'unknown' without
- connecting. Otherwise it forwards the literal placeholder upstream, gets a
- 401, and flips to 'unhealthy' even though real user calls succeed."""
+async def test_health_check_reaches_servers_without_forwarding_per_user_env_vars(
+ mock_server: MCPServer, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter
+) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
- manager = MCPServerManager()
+ manager: Final = MCPServerManager()
manager.registry[mock_server.server_id] = mock_server
+ create_client: Final = AsyncMock()
+ monkeypatch.setattr(manager, "_create_mcp_client", create_client)
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ route: Final = respx_mock.get(mock_server.url).respond(401)
- created = []
+ result: Final = await manager.health_check_server(mock_server.server_id)
- async def fake_create_client(*args, **kwargs):
- created.append((args, kwargs))
- raise RuntimeError("upstream rejected literal ${NAME}")
-
- monkeypatch.setattr(manager, "_create_mcp_client", fake_create_client)
-
- result = await manager.health_check_server(mock_server.server_id)
-
- assert created == []
- assert result.status == "unknown"
+ create_client.assert_not_called()
+ assert route.call_count == 1
+ assert not {"x-db-url", "x-other"}.intersection(route.calls[0].request.headers)
+ assert result.status == "reachable"
assert result.health_check_error is None
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
index 16bffa1a356..db476e86043 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
@@ -5,6 +5,7 @@ import json
import logging
import os
import sys
+from collections.abc import AsyncIterator
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Final, Literal, Optional
@@ -4894,69 +4895,258 @@ class TestMCPServerManager:
assert result.last_health_check is not None
@pytest.mark.asyncio
- async def test_health_check_server_oauth2_skips_check(self):
- """Test that health check is skipped for OAuth2 servers and returns unknown status"""
- manager = MCPServerManager()
-
- # Mock OAuth2 server
- server = MCPServer(
+ @pytest.mark.parametrize("oauth2_flow", [None, "authorization_code", "client_credentials"])
+ async def test_health_check_server_oauth2_reports_reachability(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, oauth2_flow: Literal["authorization_code", "client_credentials"] | None
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
server_id="oauth2-server",
name="oauth2-server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
url="http://oauth2-server.com",
+ oauth2_flow=oauth2_flow,
+ client_id="client-id",
+ client_secret="stored-client-secret",
+ static_headers={"Authorization": "Bearer static-secret", "X-API-Key": "key-secret", "Cookie": "secret"},
)
-
- manager.get_mcp_server_by_id = MagicMock(return_value=server)
-
- # _create_mcp_client should not be called for OAuth2 servers
+ manager.registry[server.server_id] = server
manager._create_mcp_client = AsyncMock()
+ route: Final = respx_mock.get(server.url).respond(401)
- # Perform health check
- result = await manager.health_check_server("oauth2-server")
+ result: Final = await manager.health_check_server(server.server_id, mcp_auth_header="caller-secret")
- # Verify that client was not created (health check was skipped)
manager._create_mcp_client.assert_not_called()
+ assert result.status == "reachable"
+ assert result.health_check_error is None
+ assert result.last_health_check is not None
+ assert route.call_count == 1
+ assert not {"authorization", "x-api-key", "cookie"}.intersection(route.calls[0].request.headers)
- # Verify results
- assert isinstance(result, LiteLLM_MCPServerTable)
- assert result.server_id == "oauth2-server"
- assert result.status == "unknown"
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("auth_type", [
+ MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token,
+ MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, MCPAuth.true_passthrough, MCPAuth.oauth_delegate,
+ ])
+ @pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse])
+ @pytest.mark.parametrize("response_code", [200, 204, 302, 401, 403, 405, 503])
+ async def test_health_check_without_credentials_accepts_any_http_response(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, auth_type: MCPAuthType, transport: Literal[MCPTransport.http, MCPTransport.sse],
+ response_code: int,
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="no-token-server",
+ name="no-token-server",
+ transport=transport,
+ auth_type=auth_type,
+ authentication_token=None,
+ url="http://no-token-server.com",
+ )
+ manager.registry[server.server_id] = server
+ manager._create_mcp_client = AsyncMock()
+ route: Final = respx_mock.get(server.url).respond(response_code)
+
+ result: Final = await manager.health_check_server(server.server_id)
+
+ manager._create_mcp_client.assert_not_called()
+ assert route.call_count == 1
+ assert result.status == "reachable"
assert result.health_check_error is None
assert result.last_health_check is not None
@pytest.mark.asyncio
- async def test_health_check_server_no_token_skips_check(self):
- """Test that health check is skipped when auth_type is set but authentication_token is missing"""
- manager = MCPServerManager()
+ @pytest.mark.parametrize("response_code", [200, 302])
+ async def test_health_reachability_closes_sse_without_body_redirect_or_cookie_reuse(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, response_code: int
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ class UnreadBody(httpx.AsyncByteStream):
+ def __init__(self) -> None:
+ self.read = False
+ self.closed = False
- # Mock server with auth_type but no authentication_token
- server = MCPServer(
- server_id="no-token-server",
- name="no-token-server",
- transport=MCPTransport.http,
- auth_type=MCPAuth.bearer_token,
- authentication_token=None, # No token
- url="http://no-token-server.com",
+ async def __aiter__(self) -> AsyncIterator[bytes]:
+ self.read = True
+ yield b"secret SSE body"
+
+ async def aclose(self) -> None:
+ self.closed = True
+
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="streaming-health", name="streaming-health", transport=MCPTransport.sse,
+ auth_type=MCPAuth.oauth2, url="https://mcp.example.test/events",
+ )
+ manager.registry[server.server_id] = server
+ bodies: Final = (UnreadBody(), UnreadBody())
+ route: Final = respx_mock.get(server.url).mock(side_effect=[
+ httpx.Response(response_code, stream=body, headers={
+ "Content-Type": "text/event-stream", "Set-Cookie": "health=secret; Path=/",
+ "Location": "http://127.0.0.1/private",
+ }) for body in bodies
+ ])
+
+ first: Final = await manager.health_check_server(server.server_id)
+ second: Final = await manager.health_check_server(server.server_id)
+
+ assert (first.status, second.status) == ("reachable", "reachable")
+ assert route.call_count == len(respx_mock.calls) == 2
+ assert all(body.closed and not body.read for body in bodies)
+ assert all("cookie" not in call.request.headers for call in route.calls)
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(("transport", "url"), [
+ (MCPTransport.stdio, "https://mcp.example.test"),
+ (MCPTransport.http, None), (MCPTransport.http, ""), (MCPTransport.http, "not-a-url"),
+ (MCPTransport.http, "ftp://mcp.example.test"),
+ (MCPTransport.http, "https://user:secret@mcp.example.test"),
+ (MCPTransport.http, "https://mcp.example.test:bad/mcp"),
+ ])
+ async def test_health_reachability_rejects_unprobeable_urls_without_requests(
+ self, respx_mock: MockRouter, transport: Literal[MCPTransport.http, MCPTransport.stdio], url: str | None
+ ) -> None:
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="unprobeable", name="unprobeable", transport=transport, auth_type=MCPAuth.oauth2, url=url,
+ )
+ manager.registry[server.server_id] = server
+
+ result: Final = await manager.health_check_server(server.server_id)
+
+ assert result.status == "unknown"
+ assert result.health_check_error and "secret" not in result.health_check_error
+ assert not respx_mock.calls
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("failure", [
+ httpx.ConnectError("TLS/connection failure with secret details"),
+ httpx.ReadTimeout("secret timeout details"),
+ httpx.RemoteProtocolError("secret malformed response"),
+ ])
+ async def test_health_reachability_reports_no_response_without_secret_details(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, failure: httpx.RequestError
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="failed-health", name="failed-health", transport=MCPTransport.http,
+ auth_type=MCPAuth.bearer_token, is_byok=True, url="https://mcp.example.test/secret?token=secret",
+ )
+ manager.registry[server.server_id] = server
+ route: Final = respx_mock.get(server.url).mock(side_effect=failure)
+
+ result: Final = await manager.health_check_server(server.server_id)
+
+ assert result.status == "unhealthy"
+ assert result.health_check_error and "secret" not in result.health_check_error
+ assert route.call_count == 1
+
+ @pytest.mark.asyncio
+ async def test_health_reachability_contains_ssl_setup_errors(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setenv("SSL_SECURITY_LEVEL", "invalid-secret-cipher")
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="bad-tls", name="bad-tls", transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2, url="https://mcp.example.test",
+ )
+ manager.registry[server.server_id] = server
+
+ result: Final = await manager.health_check_server(server.server_id)
+
+ assert result.status == "unhealthy"
+ assert result.health_check_error == "Reachability check failed (SSLError)"
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("cancel", [False, True])
+ async def test_health_reachability_timeout_and_cancellation_clean_up(
+ self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, cancel: bool
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.1)
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="slow-health", name="slow-health", transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2, url="https://mcp.example.test/slow",
+ )
+ manager.registry[server.server_id] = server
+ started: Final = asyncio.Event()
+ stopped: Final = asyncio.Event()
+
+ async def slow_response(request: httpx.Request) -> httpx.Response:
+ started.set()
+ try:
+ await asyncio.Event().wait()
+ return httpx.Response(200)
+ finally:
+ stopped.set()
+
+ respx_mock.get(server.url).mock(side_effect=slow_response)
+ task: Final = asyncio.create_task(manager.health_check_server(server.server_id))
+ await asyncio.wait_for(started.wait(), timeout=1)
+ if cancel:
+ task.cancel()
+ result: Final = await task
+
+ assert result.status == ("unknown" if cancel else "unhealthy")
+ assert result.health_check_error == (
+ "Reachability check was cancelled" if cancel else "Reachability check timed out after 0.1 seconds"
+ )
+ assert stopped.is_set()
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("server_count", [0, 1, 10, 11, 25])
+ @pytest.mark.parametrize("filtered", [False, True])
+ async def test_bulk_health_checks_deduplicate_and_bound_upstream_requests(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, server_count: int, filtered: bool
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+
+ class Probe:
+ def __init__(self) -> None:
+ self.active = 0
+ self.peak = 0
+
+ async def respond(self, request: httpx.Request) -> httpx.Response:
+ self.active += 1
+ self.peak = max(self.peak, self.active)
+ try:
+ await asyncio.sleep(0)
+ return httpx.Response(401)
+ finally:
+ self.active -= 1
+
+ manager: Final = MCPServerManager()
+ server_ids: Final = [f"health-{index}" for index in range(server_count)]
+ manager.registry = {
+ server_id: MCPServer(
+ server_id=server_id, name=server_id, transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2, url=f"https://health.example.test/{server_id}",
+ )
+ for server_id in server_ids
+ }
+ probe: Final = Probe()
+ route: Final = respx_mock.get(host="health.example.test").mock(side_effect=probe.respond)
+ requested_ids: Final = [*server_ids, *reversed(server_ids), *server_ids, "not-registered"]
+
+ results: Final = (
+ await manager.get_all_mcp_servers_with_health_and_teams(
+ user_api_key_auth=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
+ server_ids=requested_ids,
+ )
+ if filtered
+ else await manager.get_all_mcp_servers_with_health_unfiltered(server_ids=requested_ids)
)
- manager.get_mcp_server_by_id = MagicMock(return_value=server)
-
- # _create_mcp_client should not be called
- manager._create_mcp_client = AsyncMock()
-
- # Perform health check
- result = await manager.health_check_server("no-token-server")
-
- # Verify that client was not created (health check was skipped)
- manager._create_mcp_client.assert_not_called()
-
- # Verify results
- assert isinstance(result, LiteLLM_MCPServerTable)
- assert result.server_id == "no-token-server"
- assert result.status == "unknown"
- assert result.health_check_error is None
- assert result.last_health_check is not None
+ assert [(server.server_id, server.status) for server in results] == [
+ (server_id, "reachable") for server_id in server_ids
+ ]
+ assert route.call_count == server_count
+ assert probe.peak == min(server_count, 10)
+ assert probe.active == 0
@pytest.mark.asyncio
async def test_health_check_server_with_static_headers(self):
@@ -5003,70 +5193,58 @@ class TestMCPServerManager:
assert result.health_check_error is None
@pytest.mark.asyncio
- async def test_health_check_skips_passthrough_auth_with_authorization_header(self):
- """Test that health check is skipped for servers with passthrough Authorization header"""
- manager = MCPServerManager()
-
- # Mock server with auth_type=none and Authorization in extra_headers (passthrough auth)
- server = MCPServer(
+ async def test_health_check_reaches_passthrough_auth_with_authorization_header(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
server_id="github-server",
name="github-server",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
authentication_token=None,
url="http://github-server.com",
- extra_headers=["Authorization"], # Passthrough auth configured
+ extra_headers=["Authorization"],
)
-
- manager.get_mcp_server_by_id = MagicMock(return_value=server)
-
- # _create_mcp_client should not be called (health check should be skipped)
+ manager.registry[server.server_id] = server
manager._create_mcp_client = AsyncMock()
+ route: Final = respx_mock.get(server.url).respond(401)
- # Perform health check
- result = await manager.health_check_server("github-server")
+ result: Final = await manager.health_check_server(server.server_id)
- # Verify that client was not created (health check was skipped)
manager._create_mcp_client.assert_not_called()
-
- # Verify results
- assert isinstance(result, LiteLLM_MCPServerTable)
- assert result.server_id == "github-server"
- assert result.status == "unknown"
+ assert route.call_count == 1
+ assert "authorization" not in route.calls[0].request.headers
+ assert result.status == "reachable"
assert result.health_check_error is None
assert result.last_health_check is not None
@pytest.mark.asyncio
- async def test_health_check_skips_passthrough_auth_with_api_key_header(self):
- """Test that health check is skipped for servers with passthrough x-api-key header"""
- manager = MCPServerManager()
-
- # Mock server with auth_type=none and x-api-key in extra_headers
- server = MCPServer(
+ async def test_health_check_reaches_passthrough_auth_with_api_key_header(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
server_id="sourcegraph-server",
name="sourcegraph-server",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
authentication_token=None,
url="http://sourcegraph-server.com",
- extra_headers=["x-api-key"], # Passthrough auth configured
+ extra_headers=["x-api-key"],
)
-
- manager.get_mcp_server_by_id = MagicMock(return_value=server)
-
- # _create_mcp_client should not be called
+ manager.registry[server.server_id] = server
manager._create_mcp_client = AsyncMock()
+ route: Final = respx_mock.get(server.url).respond(403)
- # Perform health check
- result = await manager.health_check_server("sourcegraph-server")
+ result: Final = await manager.health_check_server(server.server_id)
- # Verify that client was not created (health check was skipped)
manager._create_mcp_client.assert_not_called()
-
- # Verify results
- assert isinstance(result, LiteLLM_MCPServerTable)
- assert result.server_id == "sourcegraph-server"
- assert result.status == "unknown"
+ assert route.call_count == 1
+ assert "x-api-key" not in route.calls[0].request.headers
+ assert result.status == "reachable"
assert result.health_check_error is None
assert result.last_health_check is not None
@@ -9239,16 +9417,19 @@ class TestRegistryTableConversionPreservesEnvVars:
self._assert_env_vars_round_tripped(table)
@pytest.mark.asyncio
- async def test_health_check_server_preserves_env_vars(self):
- # OAuth2 without client credentials needs a per-user token, so the
- # health check is skipped (no network) and we exercise the table
- # construction path directly.
- manager = MCPServerManager()
- server = self._server_with_env_vars()
+ async def test_health_check_server_preserves_env_vars(
+ self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter
+ ) -> None:
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = MCPServerManager()
+ server: Final = self._server_with_env_vars()
assert server.requires_per_user_auth is True
manager.registry[server.server_id] = server
- table = await manager.health_check_server(server.server_id)
+ route: Final = respx_mock.get(server.url).respond(401)
+ table: Final = await manager.health_check_server(server.server_id)
self._assert_env_vars_round_tripped(table)
+ assert route.call_count == 1
+ assert "x-db-url" not in route.calls[0].request.headers
class TestHealthCheckInterpolatesGlobalEnvVars:
diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
index a9ec575e99b..8aa16b817fc 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
@@ -9,11 +9,12 @@ from contextlib import ExitStack, contextmanager
from dataclasses import dataclass, field
from datetime import datetime, timedelta
from types import SimpleNamespace
-from typing import Final, List, Optional, cast
+from typing import Final, List, Literal, Optional, cast
from unittest.mock import AsyncMock, MagicMock, patch
+import httpx
import pytest
-from pydantic import BaseModel
+from pydantic import BaseModel, TypeAdapter, ValidationError
from respx import MockRouter
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
@@ -4311,6 +4312,170 @@ async def test_health_discovery_respects_route_restricted_key_grants(
assert all(row["status"] == expected_status for row in result)
+@pytest.mark.asyncio
+@pytest.mark.respx(assert_all_called=False)
+@pytest.mark.parametrize("include_reachability", [False, True])
+@pytest.mark.parametrize(
+ ("requested", "expected"),
+ [
+ (None, ("shared", "first", "second")),
+ ((), ("shared", "first", "second")),
+ (("shared", "shared", "denied"), ("shared",)),
+ (("second", "first"), ("first", "second")),
+ (("denied",), ()),
+ ],
+)
+async def test_health_checks_probe_shared_servers_once_across_auth_contexts(
+ respx_mock: MockRouter,
+ monkeypatch: pytest.MonkeyPatch,
+ requested: tuple[str, ...] | None,
+ expected: tuple[str, ...],
+ include_reachability: bool,
+) -> None:
+ from litellm.proxy._experimental.mcp_server import mcp_server_manager
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = mcp_server_manager.MCPServerManager()
+ manager.registry = {
+ server_id: MCPServer(
+ server_id=server_id,
+ name=server_id,
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ url=f"https://mcp.example.test/{server_id}",
+ )
+ for server_id in ("shared", "first", "second", "denied")
+ }
+ routes: Final = {
+ server_id: respx_mock.get(server.url).respond(401)
+ for server_id, server in manager.registry.items()
+ }
+ contexts: Final = [
+ UserAPIKeyAuth(
+ user_role=LitellmUserRoles.INTERNAL_USER,
+ api_key=f"test-health-{index}",
+ object_permission=LiteLLM_ObjectPermissionTable(
+ object_permission_id=f"health-{index}", mcp_servers=list(grants)
+ ),
+ )
+ for index, grants in enumerate((("shared", "first"), ("shared", "second")))
+ ]
+ with (
+ patch.object(
+ mgmt_endpoints, "global_mcp_server_manager", manager
+ ),
+ patch.object(
+ mcp_server_manager, "global_mcp_server_manager", manager
+ ),
+ patch.object(
+ mgmt_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=contexts)
+ ),
+ patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "restricted"}),
+ ):
+ result: Final = await mgmt_endpoints.health_check_servers(
+ server_ids=list(requested) if requested is not None else None,
+ user_api_key_dict=contexts[0],
+ include_reachability=include_reachability,
+ )
+
+ expected_status: Final = "reachable" if include_reachability else "unknown"
+ assert sorted(result, key=lambda row: row["server_id"]) == [
+ {"server_id": server_id, "status": expected_status} for server_id in sorted(expected)
+ ]
+ if requested:
+ assert [row["server_id"] for row in result] == list(expected)
+ assert {server_id: route.call_count for server_id, route in routes.items()} == {
+ server_id: int(server_id in expected) for server_id in routes
+ }
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("mode", ["restricted", "view_all"])
+@pytest.mark.parametrize("detail", [False, True])
+@pytest.mark.parametrize("flag", [None, "false", "true"])
+async def test_health_reachability_requires_explicit_api_opt_in(
+ respx_mock: MockRouter,
+ monkeypatch: pytest.MonkeyPatch,
+ mode: str,
+ detail: bool,
+ flag: str | None,
+) -> None:
+ from litellm.proxy._experimental.mcp_server import mcp_server_manager
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+
+ class HealthResponse(BaseModel):
+ server_id: str
+ status: str | None
+
+ class LegacyHealthResponse(BaseModel):
+ server_id: str
+ status: Literal["healthy", "unhealthy", "unknown"] | None
+
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ manager: Final = mcp_server_manager.MCPServerManager()
+ server: Final = MCPServer(
+ server_id="health-compatibility",
+ name="health-compatibility",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ url="https://mcp.example.test/mcp",
+ )
+ manager.registry[server.server_id] = server
+ route: Final = respx_mock.get(server.url).respond(401)
+ caller: Final = UserAPIKeyAuth(
+ user_role=LitellmUserRoles.PROXY_ADMIN,
+ api_key="test-health-compatibility",
+ object_permission=LiteLLM_ObjectPermissionTable(
+ object_permission_id="health-compatibility", mcp_servers=[server.server_id]
+ ),
+ )
+
+ def authenticated_caller() -> UserAPIKeyAuth:
+ return caller
+
+ app: Final = FastAPI()
+ app.include_router(mgmt_endpoints.router)
+ app.dependency_overrides[mgmt_endpoints.user_api_key_auth] = authenticated_caller
+ suffix: Final = server.server_id if detail else "health"
+ query: Final = {} if flag is None else {"include_reachability": flag}
+ with (
+ patch.object( # test-quality-ok: TQ008 inject the real registry into the legacy route binding
+ mgmt_endpoints, "global_mcp_server_manager", manager
+ ),
+ patch.object( # test-quality-ok: TQ008 permission resolution uses the shared registry
+ mcp_server_manager, "global_mcp_server_manager", manager
+ ),
+ patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": mode}),
+ patch.object( # test-quality-ok: TQ008 select the config-backed detail path without a database
+ mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
+ ),
+ patch.object( # test-quality-ok: TQ008 a missing database row falls back to the real registry
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
+ ):
+ async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://gateway") as client:
+ response: Final = await client.get(f"/v1/mcp/server/{suffix}", params=query)
+
+ assert response.status_code == 200, response.text
+ rows: Final = (
+ [HealthResponse.model_validate_json(response.content)]
+ if detail else TypeAdapter(list[HealthResponse]).validate_json(response.content)
+ )
+ expected_status: Final = "reachable" if flag == "true" else "unknown"
+ assert [row.model_dump() for row in rows] == [{"server_id": server.server_id, "status": expected_status}]
+ assert route.call_count == 1
+ legacy_parser: Final = (
+ LegacyHealthResponse.model_validate_json
+ if detail else TypeAdapter(list[LegacyHealthResponse]).validate_json
+ )
+ if flag == "true":
+ with pytest.raises(ValidationError, match="literal_error"):
+ legacy_parser(response.content)
+ else:
+ legacy_parser(response.content)
+
+
class TestMCPRegistryEndpoint:
def test_registry_returns_404_when_flag_missing(self):
client = create_mcp_router_test_client()
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx
index 71c2e107774..100b0ea93d3 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx
@@ -1,5 +1,6 @@
import React from "react";
import { fireEvent, render, screen } from "@testing-library/react";
+import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi, afterEach } from "vitest";
import MCPServerCard from "./MCPServerCard";
import type { MCPServer } from "@/components/mcp_tools/types";
@@ -18,6 +19,19 @@ function renderCard(overrides: Partial) {
render( );
}
+describe("MCPServerCard health", () => {
+ it("explains that reachable does not verify authentication or tools", async () => {
+ const user = userEvent.setup();
+ renderCard({ status: "reachable", oauth2_flow: "authorization_code" });
+
+ await user.hover(screen.getByText("Reachable"));
+
+ expect(await screen.findByText("Server responded. Authentication and tools were not checked")).toBeInTheDocument();
+ expect(screen.queryByText("No health data")).not.toBeInTheDocument();
+ expect(screen.queryByText("Healthy")).not.toBeInTheDocument();
+ });
+});
+
describe("MCPServerCard OAuth flow indicator", () => {
it("shows the 'OAuth flow not set' badge for an oauth2 server with no oauth2_flow", () => {
renderCard({ auth_type: "oauth2", oauth2_flow: null });
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx
index 42fb95d5951..775809e3670 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx
@@ -11,7 +11,7 @@ import {
} from "@/components/ui/dropdown-menu";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { cn } from "@/lib/cva.config";
-import { AUTH_TYPE, type MCPServer } from "@/components/mcp_tools/types";
+import { AUTH_TYPE, MCP_REACHABLE_DESCRIPTION, type MCPServer } from "@/components/mcp_tools/types";
import { Logo } from "@/components/molecules/logo/Logo";
import { getMaskedAndFullUrl } from "./utils";
@@ -33,6 +33,7 @@ interface MCPServerCardProps {
const HEALTH_TONE: Record = {
healthy: { dot: "bg-success" },
+ reachable: { dot: "bg-info" },
unhealthy: { dot: "bg-destructive" },
unknown: { dot: "bg-border" },
};
@@ -332,6 +333,7 @@ const HealthChip: FC = ({
);
}
+ const hasHealthData = Boolean(lastCheck || error || status === "reachable");
return (
= ({
/>
Health: {status}
+ {status === "reachable" && {MCP_REACHABLE_DESCRIPTION}}
{lastCheck && Last check: {new Date(lastCheck).toLocaleString()}}
{error && (
@@ -362,7 +365,7 @@ const HealthChip: FC = ({
{error}
)}
- {!lastCheck && !error && No health data}
+ {!hasHealthData && No health data}
{onRecheck && Click to recheck}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx
index 1217d878489..9a3ff0cc6cb 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx
@@ -113,12 +113,14 @@ describe("compareServers", () => {
it("sorts health before recency and display name", () => {
const servers: MCPServer[] = [
{ ...server("healthy", "aaa", "2026-03-01T00:00:00Z"), status: "healthy" },
+ { ...server("reachable", "aaa", "2026-04-01T00:00:00Z"), status: "reachable" },
{ ...server("unknown", "bbb", "2026-02-01T00:00:00Z"), status: "unknown" },
{ ...server("unhealthy", "zzz", "2026-01-01T00:00:00Z"), status: "unhealthy" },
];
expect(servers.sort((a, b) => compareServers(a, b, "health")).map((s) => s.server_id)).toEqual([
"unhealthy",
"unknown",
+ "reachable",
"healthy",
]);
});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx
index b4b7ab6b3c8..56a18e0ca4c 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx
@@ -61,7 +61,8 @@ const SORT_OPTIONS: { value: SortKey; label: string }[] = [
const HEALTH_RANK: Record = {
unhealthy: 0,
unknown: 1,
- healthy: 2,
+ reachable: 2,
+ healthy: 3,
};
const compareByName = (a: MCPServer, b: MCPServer): number => {
@@ -191,7 +192,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i
const healthStatus = healthMap.get(server.server_id);
return {
...server,
- status: healthStatus ? (healthStatus as "healthy" | "unhealthy" | "unknown") : server.status,
+ status: healthStatus ? (healthStatus as MCPServer["status"]) : server.status,
};
});
}, [mcpServers, healthStatuses]);
diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx
index e32c861f13a..1032b03a3ce 100644
--- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx
@@ -28,10 +28,10 @@ const mockServer: MCPServerData = {
env: {},
};
-function renderTable(onServerClick = vi.fn()) {
+function renderTable(onServerClick = vi.fn(), servers = [mockServer]) {
render(
server.server_id}
sortingMode="client"
@@ -42,6 +42,15 @@ function renderTable(onServerClick = vi.fn()) {
}
describe("getMCPHubTableColumns", () => {
+ it("explains the limited check for a reachable server", async () => {
+ const user = userEvent.setup();
+ renderTable(vi.fn(), [{ ...mockServer, status: "reachable" }]);
+
+ await user.hover(screen.getByText("reachable"));
+
+ expect(await screen.findByText("Server responded. Authentication and tools were not checked")).toBeInTheDocument();
+ });
+
it("renders the server row", () => {
renderTable();
expect(screen.getByText("exa_test")).toBeInTheDocument();
diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx
index 6a1ede11201..20a14bcb476 100644
--- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx
@@ -4,6 +4,7 @@ import { ColumnDef } from "@tanstack/react-table";
import { Copy, Info, MoreHorizontal } from "lucide-react";
import { DataTableSortHeader } from "@/components/shared/DataTable";
+import { MCP_REACHABLE_DESCRIPTION } from "@/components/mcp_tools/types";
import { IdentityCell, StatusBadge, type StatusTone } from "@/components/shared/table_cells";
import { Badge } from "@/components/ui/badge";
import { buttonVariants } from "@/components/ui/button";
@@ -49,6 +50,7 @@ const STATUS_TONES: Record = {
inactive: "error",
unknown: "neutral",
healthy: "success",
+ reachable: "info",
unhealthy: "error",
};
@@ -150,7 +152,11 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps)
enableSorting: true,
sortingFn: "alphanumeric",
cell: ({ row }) => (
-
+
),
},
{
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx
index 47df369fb8a..be7d39616ca 100644
--- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx
+++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx
@@ -405,6 +405,8 @@ export interface MCPToolsViewerProps {
extraHeaders?: string[] | null;
}
+export const MCP_REACHABLE_DESCRIPTION = "Server responded. Authentication and tools were not checked";
+
export interface MCPServer {
server_id: string;
is_config?: boolean;
@@ -435,7 +437,7 @@ export interface MCPServer {
updated_by: string;
extra_headers?: string[] | null;
static_headers?: Record | null;
- status?: "healthy" | "unhealthy" | "unknown";
+ status?: "healthy" | "reachable" | "unhealthy" | "unknown";
last_health_check?: string | null;
health_check_error?: string | null;
teams?: Team[];
diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts
index e14f1939ee1..b964231804e 100644
--- a/ui/litellm-dashboard/src/components/networking.test.ts
+++ b/ui/litellm-dashboard/src/components/networking.test.ts
@@ -706,6 +706,30 @@ describe("testMCPToolsListRequest auth headers", () => {
});
});
+describe("fetchMCPServerHealth", () => {
+ const originalFetch = global.fetch;
+
+ afterEach(() => {
+ global.fetch = originalFetch;
+ });
+
+ it.each([{ serverIds: undefined }, { serverIds: [] }, { serverIds: ["server one", "server&two"] }])(
+ "opts into reachability while preserving requested servers: $serverIds",
+ async ({ serverIds }) => {
+ const mockFetch = vi.fn().mockResolvedValue(new Response("[]", { status: 200 }));
+ global.fetch = mockFetch;
+
+ await Networking.fetchMCPServerHealth("test-token", serverIds);
+
+ expect(mockFetch).toHaveBeenCalledOnce();
+ const url = new URL(String(mockFetch.mock.calls[0][0]), "http://localhost");
+ expect(url.pathname).toMatch(/\/v1\/mcp\/server\/health$/);
+ expect(url.searchParams.get("include_reachability")).toBe("true");
+ expect(url.searchParams.getAll("server_ids")).toEqual(serverIds ?? []);
+ },
+ );
+});
+
describe("getAutoRouterClassifierDefaultPromptCall", () => {
const originalFetch = global.fetch;
diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx
index c4008d512c3..e1271b9151f 100644
--- a/ui/litellm-dashboard/src/components/networking.tsx
+++ b/ui/litellm-dashboard/src/components/networking.tsx
@@ -4968,6 +4968,7 @@ export const fetchMCPServerHealth = async (accessToken: string, serverIds?: stri
return await apiClient.get(`/v1/mcp/server/health`, {
accessToken,
query: {
+ include_reachability: true,
server_ids: serverIds && serverIds.length > 0 ? serverIds : undefined,
},
});
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index 0500aeb95c8..b5d515bd1fb 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -32504,10 +32504,10 @@ export interface components {
} | null;
/**
* Status
- * @description Health status: 'healthy', 'unhealthy', 'unknown'
+ * @description Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)
* @default unknown
*/
- status: ("healthy" | "unhealthy" | "unknown") | null;
+ status: ("healthy" | "reachable" | "unhealthy" | "unknown") | null;
/** Subject Token Type */
subject_token_type?: string | null;
/** Submitted At */
@@ -72212,6 +72212,8 @@ export interface operations {
query?: {
/** @description Server IDs to check. If not provided, checks all accessible servers. */
server_ids?: string[] | null;
+ /** @description Allow the 'reachable' status for responding servers whose authentication is unchecked. */
+ include_reachability?: boolean;
};
header?: never;
path?: never;
@@ -72363,7 +72365,10 @@ export interface operations {
};
fetch_mcp_server_v1_mcp_server__server_id__get: {
parameters: {
- query?: never;
+ query?: {
+ /** @description Allow the 'reachable' status for responding servers whose authentication is unchecked. */
+ include_reachability?: boolean;
+ };
header?: never;
path: {
server_id: string;
From 5d777c16d9e59c690886d978f7546b130bd2432d Mon Sep 17 00:00:00 2001
From: tin-berri
Date: Sat, 26 Sep 2026 16:53:02 -0700
Subject: [PATCH 27/39] fix(mcp): align hub publication status and controls
(#43241)
* fix(mcp): align hub publication status and controls
* refactor(mcp): keep hub visibility guard outside table rendering
---------
Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
---
cookbook/litellm_proxy_server/mcp/README.md | 37 ++++
.../mcp_server/mcp_server_manager.py | 24 +--
.../mcp_management_endpoints.py | 35 ++--
.../public_endpoints/public_endpoints.py | 14 +-
.../mcp_server/test_mcp_server_manager.py | 54 ++++++
.../test_mcp_management_endpoints.py | 156 ++++++++++++++++
.../public_endpoints/test_public_endpoints.py | 66 +++++--
.../_components/MCPPermissionManagement.tsx | 4 +-
.../_components/MCPServerCard.test.tsx | 12 ++
.../mcp-servers/_components/MCPServerCard.tsx | 19 +-
.../_components/mcp_server_view.test.tsx | 11 +-
.../_components/mcp_server_view.tsx | 21 +--
.../mcp-servers/_components/utils.test.tsx | 29 +++
.../mcp-servers/_components/utils.tsx | 35 +++-
.../AIHub/MCPHubTableColumns.test.tsx | 19 +-
.../components/AIHub/MCPHubTableColumns.tsx | 6 +-
.../components/AIHub/ModelHubTable.test.tsx | 30 ++-
.../src/components/AIHub/ModelHubTable.tsx | 15 +-
.../AIHub/forms/MakeMCPPublicForm.test.tsx | 174 +++++++++++++-----
.../AIHub/forms/MakeMCPPublicForm.tsx | 121 ++++++++----
.../src/components/mcp_tools/types.tsx | 2 +
21 files changed, 720 insertions(+), 164 deletions(-)
create mode 100644 cookbook/litellm_proxy_server/mcp/README.md
diff --git a/cookbook/litellm_proxy_server/mcp/README.md b/cookbook/litellm_proxy_server/mcp/README.md
new file mode 100644
index 00000000000..aeee0719019
--- /dev/null
+++ b/cookbook/litellm_proxy_server/mcp/README.md
@@ -0,0 +1,37 @@
+# Publish MCP servers in the AI Hub
+
+Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments
+
+```yaml
+mcp_servers:
+ documentation:
+ server_id: documentation-mcp
+ url: https://mcp.example.com/mcp
+ transport: http
+ available_on_public_internet: true
+
+litellm_settings:
+ public_mcp_hub_strict_whitelist: true
+ public_mcp_servers:
+ - documentation-mcp
+```
+
+Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server`
+
+The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file
+
+To remove all explicit entries, save an empty selection in the dialog or configure:
+
+```yaml
+litellm_settings:
+ public_mcp_hub_strict_whitelist: true
+ public_mcp_servers: []
+```
+
+## Hub listing and network access
+
+The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list
+
+Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply
+
+The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index d0d9100971d..31896d9ddc5 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -6799,6 +6799,16 @@ class MCPServerManager:
return server
return None
+ @staticmethod
+ def _is_public_mcp_server(server: MCPServer, public_ids: Container[str]) -> bool:
+ return server.server_id in public_ids or (
+ not litellm.public_mcp_hub_strict_whitelist and server.available_on_public_internet
+ )
+
+ def is_mcp_server_public(self, server_id: str) -> bool:
+ server: Final = self.registry.get(server_id) or self.config_mcp_servers.get(server_id)
+ return server is not None and self._is_public_mcp_server(server, litellm.public_mcp_servers or ())
+
def get_public_mcp_servers(self) -> list[MCPServer]:
"""
Return the MCP servers published to the AI Hub via /v1/mcp/make_public.
@@ -6816,18 +6826,8 @@ class MCPServerManager:
deployments that relied on the OR-with-default semantics; will be
removed in a future release.
"""
- if litellm.public_mcp_hub_strict_whitelist:
- if litellm.public_mcp_servers is None:
- return []
- public_ids = set(litellm.public_mcp_servers)
- return [server for server in self.get_registry().values() if server.server_id in public_ids]
-
- public_ids = set(litellm.public_mcp_servers or [])
- return [
- server
- for server in self.get_registry().values()
- if server.available_on_public_internet or server.server_id in public_ids
- ]
+ public_ids: Final = frozenset(litellm.public_mcp_servers or ())
+ return [server for server in self.get_registry().values() if self._is_public_mcp_server(server, public_ids)]
def expand_permission_list(self, identifiers: list[str]) -> list[str]:
"""
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index 346b1a75ac6..deb0e00ff9b 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -650,7 +650,16 @@ if MCP_AVAILABLE:
if hasattr(redacted_server, "credentials"):
setattr(redacted_server, "credentials", _preserved_admin_config_credentials(redacted_server.credentials))
- return redacted_server
+ is_public: Final = global_mcp_server_manager.is_mcp_server_public(redacted_server.server_id)
+ return redacted_server.model_copy(
+ update={
+ "mcp_info": {
+ **(redacted_server.mcp_info or {}),
+ "is_public": is_public,
+ "is_public_explicit": is_public and redacted_server.server_id in (litellm.public_mcp_servers or ()),
+ }
+ }
+ )
def _preserved_admin_config_credentials(
credentials: "MCPCredentials | str | None",
@@ -832,10 +841,10 @@ if MCP_AVAILABLE:
sanitized.updated_at = None
# `mcp_info` is arbitrary metadata; keep only an explicit safe subset.
- is_public = False
- if isinstance(sanitized.mcp_info, dict):
- is_public = bool(sanitized.mcp_info.get("is_public"))
- sanitized.mcp_info = {"is_public": True} if is_public else None
+ sanitized.mcp_info = {
+ "is_public": (sanitized.mcp_info or {}).get("is_public") is True,
+ "is_public_explicit": (sanitized.mcp_info or {}).get("is_public_explicit") is True,
+ }
return sanitized
@@ -1260,14 +1269,6 @@ if MCP_AVAILABLE:
for server in redacted_mcp_servers:
server.connected_app_reachable = server.server_id in reachable_ids
- # augment the mcp servers with public status
- if litellm.public_mcp_servers is not None:
- for server in redacted_mcp_servers:
- if server.server_id in litellm.public_mcp_servers:
- if server.mcp_info is None:
- server.mcp_info = {}
- server.mcp_info["is_public"] = True
-
# Annotate has_user_credential for BYOK servers (single batched query)
from litellm.proxy.proxy_server import prisma_client as _byok_prisma_client
@@ -3041,9 +3042,6 @@ if MCP_AVAILABLE:
},
)
- if litellm.public_mcp_servers is None:
- litellm.public_mcp_servers = []
-
for server_id in request.mcp_server_ids:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id=server_id)
if server is None:
@@ -3052,16 +3050,15 @@ if MCP_AVAILABLE:
detail=f"MCP Server with ID {server_id} not found",
)
- litellm.public_mcp_servers = request.mcp_server_ids
-
# Update config with new settings
if "litellm_settings" not in config or config["litellm_settings"] is None:
config["litellm_settings"] = {}
- config["litellm_settings"]["public_mcp_servers"] = litellm.public_mcp_servers
+ config["litellm_settings"]["public_mcp_servers"] = request.mcp_server_ids
# Save the updated config
await proxy_config.save_config(new_config=config)
+ litellm.public_mcp_servers = request.mcp_server_ids
verbose_proxy_logger.debug(
"Updated public mcp servers to: %s by user: %s", litellm.public_mcp_servers, user_api_key_dict.user_id
diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py
index 26a5c44fce1..bba5ef681d0 100644
--- a/litellm/proxy/public_endpoints/public_endpoints.py
+++ b/litellm/proxy/public_endpoints/public_endpoints.py
@@ -300,7 +300,19 @@ async def get_mcp_servers():
)
public_mcp_servers: Final = global_mcp_server_manager.get_public_mcp_servers()
- return [MCPPublicServer.model_validate(server.model_dump()) for server in public_mcp_servers]
+ return [
+ MCPPublicServer.model_validate(
+ {
+ **server.model_dump(),
+ "mcp_info": {
+ **(server.mcp_info or {}),
+ "is_public": True,
+ "is_public_explicit": server.server_id in (litellm.public_mcp_servers or ()),
+ },
+ }
+ )
+ for server in public_mcp_servers
+ ]
@router.get(
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
index db476e86043..70ef4312f4c 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
@@ -9566,6 +9566,60 @@ class TestGetPublicMCPServers:
manager.config_mcp_servers[s.server_id] = s
return manager
+ @pytest.mark.parametrize("registered_in", ("config", "database", "both", "neither"))
+ @pytest.mark.parametrize("public_ids", (None, [], ["server-id"], ["server-alias"], ["Server Name"]))
+ @pytest.mark.parametrize(
+ "strict,network_access,implicitly_public",
+ ((True, True, False), (True, False, False), (False, True, True), (False, False, False)),
+ )
+ def test_public_status_agrees_with_hub_membership(
+ self,
+ registered_in: Literal["config", "database", "both", "neither"],
+ public_ids: list[str] | None,
+ strict: bool,
+ network_access: bool,
+ implicitly_public: bool,
+ ) -> None:
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="server-id",
+ name="server-alias",
+ alias="server-alias",
+ server_name="Server Name",
+ transport=MCPTransport.http,
+ available_on_public_internet=network_access,
+ mcp_info={"is_public": True, "description": "Preserve custom metadata"},
+ )
+ config_server: Final = (
+ server.model_copy(update={"available_on_public_internet": not network_access})
+ if registered_in == "both"
+ else server
+ )
+ manager.config_mcp_servers = (
+ {server.server_id: config_server} if registered_in in ("config", "both") else {}
+ )
+ manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {}
+ original_server: Final = server.model_dump()
+ original_config_server: Final = config_server.model_dump()
+ expected_public: Final = registered_in != "neither" and (
+ public_ids == [server.server_id] or implicitly_public
+ )
+
+ with (
+ patch("litellm.public_mcp_servers", public_ids),
+ patch("litellm.public_mcp_hub_strict_whitelist", strict),
+ ):
+ public_servers: Final = manager.get_public_mcp_servers()
+ assert manager.is_mcp_server_public(server.server_id) is expected_public
+ assert [item.server_id for item in public_servers] == (
+ [server.server_id] if expected_public else []
+ )
+ assert manager.is_mcp_server_public("server-alias") is False
+ assert manager.is_mcp_server_public("missing-server") is False
+
+ assert server.model_dump() == original_server
+ assert config_server.model_dump() == original_config_server
+
@patch("litellm.public_mcp_servers", None)
def test_returns_empty_when_whitelist_is_none(self):
"""No /make_public call yet → hub returns nothing, regardless of
diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
index 8aa16b817fc..11b3dcf54bc 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
@@ -33,6 +33,7 @@ from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
LiteLLM_MCPServerTable,
LitellmUserRoles,
+ MakeMCPServersPublicRequest,
MCPTransport,
MCPUserCredentialResponse,
NewMCPServerRequest,
@@ -154,6 +155,161 @@ def patch_proxy_general_settings(settings: dict):
)
+@pytest.mark.asyncio
+@pytest.mark.parametrize("from_db", (False, True))
+@pytest.mark.parametrize(
+ "strict,explicit,expected_public",
+ ((True, True, True), (True, False, False), (False, False, True)),
+)
+async def test_mcp_publication_list_and_detail_derive_current_status(
+ from_db: bool, strict: bool, explicit: bool, expected_public: bool
+) -> None:
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
+
+ manager: Final = MCPServerManager()
+ server: Final = MCPServer(
+ server_id="publication-server",
+ name="publication-server",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.api_key,
+ available_on_public_internet=True,
+ mcp_info={
+ "is_public": not expected_public,
+ "is_public_explicit": not explicit,
+ "description": "Keep this description",
+ },
+ )
+ manager.registry = {server.server_id: server} if from_db else {}
+ manager.config_mcp_servers = {} if from_db else {server.server_id: server}
+ record: Final = manager._build_mcp_server_table(server)
+ original_metadata: Final = dict(server.mcp_info or {})
+ admin: Final = generate_mock_user_api_key_auth()
+
+ with (
+ patch("litellm.public_mcp_servers", [server.server_id] if explicit else []),
+ patch("litellm.public_mcp_hub_strict_whitelist", strict),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
+ patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()),
+ patch.object(mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=record if from_db else None)),
+ patch("litellm.proxy.proxy_server.prisma_client", None),
+ patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "view_all"}),
+ ):
+ listing: Final = await mgmt_endpoints.fetch_all_mcp_servers(
+ user_api_key_dict=admin, team_id=None, connected_app_view=False
+ )
+ detail: Final = await mgmt_endpoints.fetch_mcp_server(
+ request=_make_mock_request(), server_id=server.server_id, user_api_key_dict=admin
+ )
+ assert len(listing) == 1
+ for projected in (listing[0], detail):
+ assert projected.mcp_info == {
+ "is_public": expected_public,
+ "is_public_explicit": explicit,
+ "description": "Keep this description",
+ }
+ assert bool(manager.get_public_mcp_servers()) is expected_public
+
+ assert server.mcp_info == original_metadata
+ assert record.mcp_info == original_metadata
+
+
+@pytest.mark.parametrize("approval_status", ("pending_review", "rejected", "draft", "active"))
+@pytest.mark.parametrize("strict", (False, True))
+def test_mcp_publication_projection_excludes_unregistered_lifecycle_records(
+ approval_status: str, strict: bool
+) -> None:
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
+
+ record: Final = LiteLLM_MCPServerTable(
+ server_id="unregistered-server",
+ transport=MCPTransport.http,
+ approval_status=approval_status,
+ credentials={"auth_value": "test-secret"},
+ available_on_public_internet=True,
+ mcp_info={"is_public": True, "is_public_explicit": True},
+ )
+ original: Final = record.model_dump()
+ with (
+ patch("litellm.public_mcp_servers", [record.server_id]),
+ patch("litellm.public_mcp_hub_strict_whitelist", strict),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", MCPServerManager()),
+ ):
+ for project in (
+ mgmt_endpoints._redact_mcp_credentials,
+ mgmt_endpoints._sanitize_mcp_server_for_non_admin,
+ mgmt_endpoints._sanitize_mcp_server_for_virtual_key,
+ ):
+ projected: Final = project(record)
+ assert projected.mcp_info == {"is_public": False, "is_public_explicit": False}
+ assert projected.credentials is None
+ assert record.model_dump() == original
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("previous_ids", (None, ["old-server"]))
+@pytest.mark.parametrize(
+ "selected_ids,save_error,role,error_status",
+ (
+ (["new-server"], None, LitellmUserRoles.PROXY_ADMIN, None),
+ ([], None, LitellmUserRoles.PROXY_ADMIN, None),
+ (["new-server"], HTTPException(400, "Owned by config file"), LitellmUserRoles.PROXY_ADMIN, 400),
+ (["new-server"], RuntimeError("Database write failed"), LitellmUserRoles.PROXY_ADMIN, 500),
+ (["missing-server"], None, LitellmUserRoles.PROXY_ADMIN, 404),
+ (["new-server"], None, LitellmUserRoles.INTERNAL_USER, 403),
+ ),
+)
+async def test_mcp_publication_updates_runtime_only_after_successful_save(
+ previous_ids: list[str] | None,
+ selected_ids: list[str],
+ save_error: HTTPException | RuntimeError | None,
+ role: LitellmUserRoles,
+ error_status: int | None,
+) -> None:
+ import litellm
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
+
+ manager: Final = MCPServerManager()
+ server: Final = generate_mock_mcp_server_config_record(server_id="new-server")
+ manager.config_mcp_servers = {server.server_id: server}
+ expected_config: Final = {"litellm_settings": {"drop_params": True, "public_mcp_servers": selected_ids}}
+
+ async def save_config(new_config: Mapping[str, object]) -> None:
+ assert litellm.public_mcp_servers is previous_ids
+ assert new_config == expected_config
+ if save_error is not None:
+ raise save_error
+
+ save: Final = AsyncMock(side_effect=save_config)
+ proxy_config: Final = SimpleNamespace(
+ get_config=AsyncMock(return_value={"litellm_settings": {"drop_params": True}}),
+ save_config=save,
+ )
+ request: Final = MakeMCPServersPublicRequest(mcp_server_ids=selected_ids)
+ caller: Final = generate_mock_user_api_key_auth(user_role=role)
+ with (
+ patch("litellm.public_mcp_servers", previous_ids),
+ patch("litellm.proxy.proxy_server.proxy_config", proxy_config),
+ patch(
+ "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
+ manager,
+ ),
+ ):
+ if error_status is None:
+ response: Final = await mgmt_endpoints.make_mcp_servers_public(request, caller)
+ assert response["public_mcp_servers"] == selected_ids
+ assert litellm.public_mcp_servers == selected_ids
+ else:
+ with pytest.raises(HTTPException) as error:
+ await mgmt_endpoints.make_mcp_servers_public(request, caller)
+ assert error.value.status_code == error_status
+ assert litellm.public_mcp_servers is previous_ids
+
+ if error_status in (403, 404):
+ save.assert_not_awaited()
+ else:
+ save.assert_awaited_once_with(new_config=expected_config)
+
+
class TestMCPCredentialsTokenExchangeProfile:
"""token_exchange_profile must be a declared MCPCredentials field so the management API can
persist the entra_obo profile. An undeclared key is silently stripped by pydantic when the
diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py
index 0dec44af402..18839a65d62 100644
--- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py
+++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py
@@ -1086,43 +1086,73 @@ def test_clean_display_name_passthrough_when_no_suffix():
assert _clean_display_name("") == ""
-def test_public_mcp_hub_returns_only_whitelisted_servers():
- """Regression: /public/mcp_hub must gate strictly on
- litellm.public_mcp_servers, mirroring /public/model_hub and
- /public/agent_hub. Servers with available_on_public_internet=True that
- are not on the whitelist must not leak."""
+@pytest.mark.parametrize(
+ "strict,explicit,expected_listed",
+ ((True, True, True), (True, False, False), (False, True, True), (False, False, True)),
+)
+@pytest.mark.parametrize("stored_public", (None, False, True))
+def test_public_mcp_hub_derives_publication_metadata_without_mutating_registry(
+ strict: bool,
+ explicit: bool,
+ expected_listed: bool,
+ stored_public: bool | None,
+) -> None:
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._types import MCPTransport
- app = FastAPI()
+ app: Final = FastAPI()
app.include_router(router)
- app.dependency_overrides[user_api_key_auth] = lambda: MagicMock()
- client = TestClient(app)
+ client: Final = TestClient(app)
- listed = MCPServer(
+ server: Final = MCPServer(
server_id="listed",
name="listed",
server_name="listed",
transport=MCPTransport.http,
available_on_public_internet=True,
+ mcp_info=(
+ {
+ "is_public": stored_public,
+ "is_public_explicit": not explicit,
+ "description": "Preserve custom metadata",
+ }
+ if stored_public is not None
+ else None
+ ),
)
-
- mock_manager = MagicMock()
- mock_manager.get_public_mcp_servers.return_value = [listed]
+ unlisted: Final = MCPServer(
+ server_id="unlisted",
+ name="unlisted",
+ transport=MCPTransport.http,
+ available_on_public_internet=False,
+ mcp_info={"is_public": True, "is_public_explicit": True},
+ )
+ manager: Final = MCPServerManager()
+ manager.config_mcp_servers = {server.server_id: server}
+ manager.registry = {unlisted.server_id: unlisted}
+ original_registry: Final = {key: value.model_dump() for key, value in manager.get_registry().items()}
with (
- patch("litellm.public_mcp_servers", ["listed"]),
+ patch("litellm.public_mcp_servers", [server.server_id] if explicit else []),
+ patch("litellm.public_mcp_hub_strict_whitelist", strict),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
- mock_manager,
+ manager,
),
):
- response = client.get("/public/mcp_hub")
+ response: Final = client.get("/public/mcp_hub")
assert response.status_code == 200
- data = response.json()
- assert [item["server_id"] for item in data] == ["listed"]
- app.dependency_overrides.clear()
+ data: Final = response.json()
+ assert [item["server_id"] for item in data] == ([server.server_id] if expected_listed else [])
+ if expected_listed:
+ assert data[0]["mcp_info"] == {
+ **({"description": "Preserve custom metadata"} if stored_public is not None else {}),
+ "is_public": True,
+ "is_public_explicit": explicit,
+ }
+ assert {key: value.model_dump() for key, value in manager.get_registry().items()} == original_registry
def test_public_mcp_hub_returns_empty_when_whitelist_unset():
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx
index cb423b435ae..a48c991bc7a 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx
@@ -217,12 +217,12 @@ const MCPPermissionManagement: React.FC = ({
Internal network only
-
+
- Turn on to restrict access to callers within your internal network only.
+ Turn on to restrict public IPs. Explicitly published server IDs remain accessible from public IPs.
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx
index 100b0ea93d3..d298d9d8145 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx
@@ -126,3 +126,15 @@ describe("MCPServerCard per-user credentials", () => {
expect(screen.queryByRole("button", { name: "Set" })).not.toBeInTheDocument();
});
});
+
+describe("MCPServerCard network access", () => {
+ it("shows effective network access without a hub listing badge", () => {
+ renderCard({
+ available_on_public_internet: false,
+ mcp_info: { server_name: "demo_server", is_public: true, is_public_explicit: true },
+ });
+
+ expect(screen.getByText("All Networks")).toBeInTheDocument();
+ expect(screen.queryByText(/^Hub:/)).not.toBeInTheDocument();
+ });
+});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx
index 775809e3670..bb153f94665 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx
@@ -13,7 +13,7 @@ import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/comp
import { cn } from "@/lib/cva.config";
import { AUTH_TYPE, MCP_REACHABLE_DESCRIPTION, type MCPServer } from "@/components/mcp_tools/types";
import { Logo } from "@/components/molecules/logo/Logo";
-import { getMaskedAndFullUrl } from "./utils";
+import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils";
interface MCPServerCardProps {
server: MCPServer;
@@ -70,7 +70,7 @@ const MCPServerCard: FC = ({
server.auth_type === AUTH_TYPE.OAUTH2 && !server.oauth2_flow && !server.delegate_auth_to_upstream;
const status = server.status || "unknown";
const healthTone = HEALTH_TONE[status] ?? HEALTH_TONE.unknown;
- const isPublic = server.available_on_public_internet;
+ const networkAccess = getMCPNetworkAccess(server);
const accessGroups = (server.mcp_access_groups ?? []).filter((g): g is string => typeof g === "string");
const missing = missingUserFields ?? [];
@@ -236,10 +236,17 @@ const MCPServerCard: FC = ({
)}
-
-
- {isPublic ? "Public" : "Internal"}
-
+
+
+
+ {networkAccess.label}
+
+ }
+ />
+ {networkAccess.description}
+
{accessGroups.slice(0, 2).map((g) => (
{
});
it("shows the read-only settings summary before editing", async () => {
- renderView({ allow_all_keys: true, available_on_public_internet: false });
+ renderView({
+ allow_all_keys: true,
+ available_on_public_internet: false,
+ mcp_info: { server_name: "demo server", is_public: true, is_public_explicit: true },
+ });
await userEvent.click(screen.getByRole("tab", { name: "Settings" }));
expect(await screen.findByText("MCP Server Settings")).toBeInTheDocument();
expect(screen.getByText("Allow All Keys")).toBeInTheDocument();
expect(screen.getByText("Enabled")).toBeInTheDocument();
- expect(screen.getByText("Internal only")).toBeInTheDocument();
+ expect(screen.getByText("Network access")).toBeInTheDocument();
+ expect(screen.getByText("All Networks")).toBeInTheDocument();
+ expect(screen.queryByText("MCP Hub")).not.toBeInTheDocument();
+ expect(screen.queryByText("Listed")).not.toBeInTheDocument();
expect(screen.queryByText("edit form")).not.toBeInTheDocument();
});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx
index a7ff34301a0..c97596ce0f6 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx
@@ -13,7 +13,7 @@ import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel";
import { getSecureItem } from "@/utils/secureStorage";
import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles";
import MCPServerCostDisplay from "./mcp_server_cost_display";
-import { getMaskedAndFullUrl } from "./utils";
+import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils";
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils";
import { CheckIcon, CopyIcon } from "lucide-react";
@@ -68,6 +68,7 @@ export const MCPServerView: React.FC = ({
const returningFromEditOAuth = isReturningFromEditOAuth(canEdit, mcpServer.server_id);
const [editing, setEditing] = useState(isEditing || returningFromEditOAuth);
const [showFullUrl, setShowFullUrl] = useState(false);
+ const networkAccess = getMCPNetworkAccess(mcpServer);
const [copiedStates, setCopiedStates] = useState>({});
const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex);
const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole);
@@ -318,19 +319,13 @@ export const MCPServerView: React.FC = ({
- Network Access
+ Network access
- {mcpServer.available_on_public_internet ? (
-
-
- Public
-
- ) : (
-
-
- Internal only
-
- )}
+
+
+ {networkAccess.label}
+
+ {networkAccess.description}
{handleAuth(mcpServer.auth_type) === "oauth2" && (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx
index 3b4fda400c2..bf30821d73a 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx
@@ -3,11 +3,40 @@ import {
extractMCPToken,
maskUrl,
getMaskedAndFullUrl,
+ getMCPNetworkAccess,
validateMCPServerUrl,
validateMCPServerName,
normalizeToolOverrideMap,
} from "./utils";
+describe("getMCPNetworkAccess", () => {
+ it.each([
+ { publicIp: true, explicit: false, label: "All Networks" },
+ { publicIp: false, explicit: true, label: "All Networks" },
+ { publicIp: true, explicit: true, label: "All Networks" },
+ { publicIp: false, explicit: false, label: "Internal Only" },
+ { publicIp: true, explicit: undefined, label: "All Networks" },
+ { publicIp: false, explicit: undefined, label: "Unknown" },
+ { publicIp: undefined, explicit: false, label: "Unknown" },
+ ])("reports $label for network=$publicIp and publication=$explicit", ({ publicIp, explicit, label }) => {
+ expect(
+ getMCPNetworkAccess({
+ available_on_public_internet: publicIp,
+ mcp_info: { server_name: "demo", is_public: true, is_public_explicit: explicit },
+ }).label,
+ ).toBe(label);
+ });
+
+ it("explains when hub publication permits public IPs", () => {
+ expect(
+ getMCPNetworkAccess({
+ available_on_public_internet: false,
+ mcp_info: { server_name: "demo", is_public_explicit: true },
+ }).description,
+ ).toContain("because this server is published in MCP Hub");
+ });
+});
+
describe("extractMCPToken", () => {
it("should extract token after /mcp/", () => {
const result = extractMCPToken("https://example.com/mcp/abc123");
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx
index 4738e1e8fba..bb72831d92e 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx
@@ -1,4 +1,37 @@
-import { MCPEnvVar, MCPEnvVarScope } from "@/components/mcp_tools/types";
+import { MCPEnvVar, MCPEnvVarScope, type MCPServer } from "@/components/mcp_tools/types";
+
+export const getMCPNetworkAccess = (
+ server: Pick,
+): {
+ readonly label: "All Networks" | "Internal Only" | "Unknown";
+ readonly dotClassName: string;
+ readonly description: string;
+} => {
+ const explicitlyPublished = server.mcp_info?.is_public_explicit;
+ if (server.available_on_public_internet === true || explicitlyPublished === true) {
+ return {
+ label: "All Networks",
+ dotClassName: "bg-success",
+ description:
+ server.available_on_public_internet === true
+ ? "Allows requests from public and internal IPs. Authentication and access permissions still apply"
+ : "Allows requests from public and internal IPs because this server is published in MCP Hub. Authentication and access permissions still apply",
+ };
+ }
+ if (server.available_on_public_internet === false && explicitlyPublished === false) {
+ return {
+ label: "Internal Only",
+ dotClassName: "bg-warning",
+ description:
+ "Allows requests only from internal/private IP ranges. Authentication and access permissions still apply",
+ };
+ }
+ return {
+ label: "Unknown",
+ dotClassName: "bg-border",
+ description: "The proxy did not report enough network and publication settings to determine allowed client IPs",
+ };
+};
export const extractMCPToken = (url: string): { token: string | null; baseUrl: string } => {
try {
diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx
index 1032b03a3ce..fee58e10fda 100644
--- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx
@@ -1,4 +1,4 @@
-import { render, screen } from "@testing-library/react";
+import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import { DataTable } from "@/components/shared/DataTable";
@@ -63,6 +63,23 @@ describe("getMCPHubTableColumns", () => {
expect(screen.getByText("Auth Type")).toBeInTheDocument();
});
+ it("shows hub membership separately from the network setting", () => {
+ renderTable(vi.fn(), [
+ { ...mockServer, available_on_public_internet: false, mcp_info: { is_public: true } },
+ {
+ ...mockServer,
+ server_id: "network-only",
+ server_name: "Network-only server",
+ available_on_public_internet: true,
+ mcp_info: { is_public: false },
+ },
+ ]);
+
+ expect(screen.getByText("Hub listing")).toBeInTheDocument();
+ expect(within(screen.getByRole("row", { name: /exa_test/ })).getByText("Listed")).toBeInTheDocument();
+ expect(within(screen.getByRole("row", { name: /Network-only server/ })).getByText("Unlisted")).toBeInTheDocument();
+ });
+
it("does not expose a URL column", () => {
renderTable();
expect(screen.queryByText("URL")).not.toBeInTheDocument();
diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx
index 20a14bcb476..db53e97569b 100644
--- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx
@@ -203,8 +203,8 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps)
{
id: "is_public",
accessorFn: (row) => row.mcp_info?.is_public === true,
- meta: { title: "Public", skeleton: "badge", className: "hidden md:table-cell" },
- header: ({ column }) => ,
+ meta: { title: "Hub listing", skeleton: "badge", className: "hidden md:table-cell" },
+ header: ({ column }) => ,
size: 100,
enableSorting: true,
sortingFn: (rowA, rowB) => {
@@ -214,7 +214,7 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps)
},
cell: ({ row }) => {
const isPublic = row.original.mcp_info?.is_public === true;
- return ;
+ return ;
},
},
{
diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx
index 27fe2330acd..1f052b8932a 100644
--- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx
@@ -1,5 +1,7 @@
import * as networking from "@/components/networking";
import userEvent from "@testing-library/user-event";
+import { act } from "@testing-library/react";
+import type { MCPServerData } from "./MCPHubTableColumns";
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils";
import ModelHubTable from "./ModelHubTable";
@@ -18,6 +20,7 @@ vi.mock("@/components/networking", () => ({
getProxyBaseUrl: vi.fn(() => "http://localhost:4000"),
getAgentsList: vi.fn(),
fetchMCPServers: vi.fn(),
+ makeMCPPublicCall: vi.fn(),
getUiSettings: vi.fn(),
getClaudeCodePluginsList: vi.fn(() => Promise.resolve({ plugins: [] })),
}));
@@ -202,13 +205,13 @@ describe("ModelHubTable", () => {
});
describe("hub tabs", () => {
- const renderHub = async (agents: object[] = []) => {
+ const renderHub = async (agents: object[] = [], mcpServers: Promise = Promise.resolve([])) => {
vi.mocked(networking.modelHubCall).mockResolvedValue({
data: [{ model_group: "claude-opus-4-8", providers: ["anthropic"], mode: "chat" }],
});
vi.mocked(networking.getConfigFieldSetting).mockResolvedValue({ field_value: false });
vi.mocked(networking.getAgentsList).mockResolvedValue({ agents });
- vi.mocked(networking.fetchMCPServers).mockResolvedValue([]);
+ vi.mocked(networking.fetchMCPServers).mockReturnValue(mcpServers);
vi.mocked(networking.getUiSettings).mockResolvedValue({ values: {} });
mockUseUISettings.mockReturnValue({ data: { values: {} }, isLoading: false });
@@ -219,6 +222,29 @@ describe("ModelHubTable", () => {
return { user, search: await screen.findByPlaceholderText("Search model names...") };
};
+ it("requires a fresh MCP publication list before and after saving", async () => {
+ const servers = Promise.withResolvers();
+ const { user } = await renderHub([], servers.promise);
+ await user.click(screen.getByRole("tab", { name: "MCP Hub" }));
+
+ const manageVisibility = screen.getByRole("button", { name: "Manage MCP Hub Visibility" });
+ expect(manageVisibility).toBeDisabled();
+ await act(async () => servers.resolve([]));
+ expect(manageVisibility).toBeEnabled();
+
+ const refresh = Promise.withResolvers();
+ vi.mocked(networking.makeMCPPublicCall).mockResolvedValueOnce({});
+ vi.mocked(networking.fetchMCPServers).mockReturnValueOnce(refresh.promise);
+ await user.click(manageVisibility);
+ await user.click(screen.getByRole("button", { name: "Next" }));
+ await user.click(screen.getByRole("button", { name: "Save Publication List" }));
+
+ expect(networking.makeMCPPublicCall).toHaveBeenCalledWith("test-token", []);
+ expect(manageVisibility).toBeDisabled();
+ await act(async () => refresh.reject(new Error("Unable to reload the publication list")));
+ expect(manageVisibility).toBeDisabled();
+ });
+
it("keeps the model filter typed on the Model Hub tab after visiting another hub", async () => {
const { user, search } = await renderHub();
diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx
index 063850b3e72..c4d776ea2fb 100644
--- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx
@@ -49,6 +49,10 @@ interface ModelHubTableProps {
userRole: string | null;
}
+function isMCPHubVisibilityDisabled(isLoading: boolean, servers: readonly MCPServerData[] | null): boolean {
+ return isLoading || servers === null;
+}
+
function HubEmptyState({ title, body }: { title: string; body: string }) {
return (
@@ -359,10 +363,14 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage,
if (accessToken) {
const fetchMcpData = async () => {
try {
+ setMcpLoading(true);
const response = await fetchMCPServers(accessToken);
setMcpHubData(response);
} catch (error) {
+ setMcpHubData(null);
console.error("Error refreshing MCP server data:", error);
+ } finally {
+ setMcpLoading(false);
}
};
fetchMcpData();
@@ -567,7 +575,12 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage,
{/* Header with Make Public Button */}
{publicPage == false && canModify && (
-
+
)}
diff --git a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx
index 5b96e9ad194..881711c1668 100644
--- a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx
@@ -1,6 +1,8 @@
import { render, screen, fireEvent, act, waitFor } from "@testing-library/react";
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
import MakeMCPPublicForm from "./MakeMCPPublicForm";
+import userEvent from "@testing-library/user-event";
+import { toast } from "@/lib/toast";
import { MCPServerData } from "@/components/AIHub/MCPHubTableColumns";
// Mock the networking function
@@ -8,6 +10,10 @@ vi.mock("../../networking", () => ({
makeMCPPublicCall: vi.fn(),
}));
+vi.mock("@/lib/toast", () => ({
+ toast: { success: vi.fn(), fromError: vi.fn() },
+}));
+
// Import the mocked function
import { makeMCPPublicCall } from "../../networking";
const mockMakeMCPPublicCall = vi.mocked(makeMCPPublicCall);
@@ -28,7 +34,7 @@ describe("MakeMCPPublicForm", () => {
url: "http://example.com/server1",
transport: "http",
status: "active",
- mcp_info: { is_public: false },
+ mcp_info: { is_public: false, is_public_explicit: false },
allowed_tools: ["tool-1", "tool-2"],
auth_type: "bearer",
credentials: {},
@@ -50,7 +56,7 @@ describe("MakeMCPPublicForm", () => {
url: "http://example.com/server2",
transport: "websocket",
status: "inactive",
- mcp_info: { is_public: true },
+ mcp_info: { is_public: true, is_public_explicit: true },
allowed_tools: [],
auth_type: "none",
credentials: {},
@@ -80,16 +86,16 @@ describe("MakeMCPPublicForm", () => {
it("should render the component", () => {
render( );
- expect(screen.getByText("Make MCP Servers Public")).toBeInTheDocument();
- expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument();
+ expect(screen.getByText("Manage MCP Hub Visibility")).toBeInTheDocument();
+ expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument();
});
it("should initialize with correct state", () => {
render( );
// Check that the component renders with the correct title and content
- expect(screen.getByText("Make MCP Servers Public")).toBeInTheDocument();
- expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument();
+ expect(screen.getByText("Manage MCP Hub Visibility")).toBeInTheDocument();
+ expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument();
// Check that all server checkboxes are present
const checkboxes = screen.getAllByRole("checkbox");
@@ -104,7 +110,7 @@ describe("MakeMCPPublicForm", () => {
render( );
// Initially on step 1
- expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument();
+ expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument();
// Select all servers using the select all checkbox
const selectAllCheckbox = screen.getByRole("checkbox", { name: "Select All (2)" });
@@ -123,7 +129,7 @@ describe("MakeMCPPublicForm", () => {
// Should move to step 2
await waitFor(() => {
- expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument();
+ expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument();
});
});
@@ -145,10 +151,10 @@ describe("MakeMCPPublicForm", () => {
// Wait for navigation to complete
await waitFor(() => {
- expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument();
+ expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument();
});
- const submitButton = screen.getByRole("button", { name: "Make Public" });
+ const submitButton = screen.getByRole("button", { name: "Save Publication List" });
await act(async () => {
fireEvent.click(submitButton);
});
@@ -187,29 +193,105 @@ describe("MakeMCPPublicForm", () => {
expect(checkboxes[2]).not.toBeChecked();
});
- it("should show error when no servers selected", async () => {
+ it("submits an empty publication list after the last server is deselected", async () => {
+ mockMakeMCPPublicCall.mockResolvedValueOnce({});
render( );
- // Deselect all servers first
- const checkboxes = screen.getAllByRole("checkbox");
- await act(async () => {
- fireEvent.click(checkboxes[0]); // Click select all to select all
- });
- await act(async () => {
- fireEvent.click(checkboxes[0]); // Click select all again to deselect all
- });
+ fireEvent.click(screen.getAllByRole("checkbox")[2]);
+ expect(screen.getByRole("button", { name: "Next" })).toBeEnabled();
+ fireEvent.click(screen.getByRole("button", { name: "Next" }));
+ fireEvent.click(screen.getByRole("button", { name: "Save Publication List" }));
- // Try to go to next step
- const nextButton = screen.getByRole("button", { name: "Next" });
- await act(async () => {
- fireEvent.click(nextButton);
- });
-
- // Should stay on same step
- expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument();
+ await waitFor(() => expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", []));
+ expect(mockProps.onSuccess).toHaveBeenCalled();
});
- it("should display empty state when no servers are available", () => {
+ it("keeps legacy listings separate from explicitly published selections", () => {
+ render(
+ ,
+ );
+
+ expect(screen.getAllByRole("checkbox")[1]).not.toBeChecked();
+ expect(screen.getAllByRole("checkbox")[2]).toBeChecked();
+ expect(screen.getByText("Listed by legacy mode")).toBeInTheDocument();
+ });
+
+ it.each([
+ { mode: "all missing, stale true", info: { is_public: true }, mixed: false },
+ { mode: "all missing, stale false", info: { is_public: false }, mixed: false },
+ { mode: "mixed, stale true", info: { is_public: true }, mixed: true },
+ { mode: "mixed, stale false", info: { is_public: false }, mixed: true },
+ { mode: "null explicit status", info: { is_public: true, is_public_explicit: null }, mixed: true },
+ { mode: "nonboolean explicit status", info: { is_public: true, is_public_explicit: "true" }, mixed: true },
+ ])("blocks unknown explicit publication metadata: $mode", ({ info, mixed }) => {
+ const unknownServer = { ...mockProps.mcpHubData[0], mcp_info: info };
+ const catalog = mixed ? [unknownServer, mockProps.mcpHubData[1]] : [unknownServer];
+ render( );
+
+ expect(screen.getByRole("alert")).toHaveTextContent("explicit publication status");
+ expect(screen.queryByRole("checkbox")).not.toBeInTheDocument();
+ expect(screen.queryByText("Configure in YAML")).not.toBeInTheDocument();
+ expect(screen.queryByRole("button", { name: "Copy code" })).not.toBeInTheDocument();
+ const nextButton = screen.getByRole("button", { name: "Next" });
+ expect(nextButton).toBeDisabled();
+ fireEvent.click(nextButton);
+ expect(screen.queryByText("Confirm MCP Hub Publication")).not.toBeInTheDocument();
+ expect(mockMakeMCPPublicCall).not.toHaveBeenCalled();
+ });
+
+ it.each([true, false])("blocks confirmation when explicit metadata disappears with stale listing %s", (listed) => {
+ const { rerender } = render( );
+ fireEvent.click(screen.getByRole("button", { name: "Next" }));
+ expect(screen.getByRole("button", { name: "Save Publication List" })).toBeEnabled();
+
+ const catalog = [{ ...mockProps.mcpHubData[0], mcp_info: { is_public: listed } }, mockProps.mcpHubData[1]];
+ rerender( );
+
+ expect(screen.getByRole("alert")).toHaveTextContent("explicit publication status");
+ expect(screen.queryByText("Confirm MCP Hub Publication")).not.toBeInTheDocument();
+ const saveButton = screen.getByRole("button", { name: "Save Publication List" });
+ expect(saveButton).toBeDisabled();
+ fireEvent.click(saveButton);
+ expect(mockMakeMCPPublicCall).not.toHaveBeenCalled();
+ expect(screen.queryByRole("button", { name: "Copy code" })).not.toBeInTheDocument();
+
+ const refreshedCatalog = [
+ { ...mockProps.mcpHubData[0], mcp_info: { is_public: true, is_public_explicit: true } },
+ { ...mockProps.mcpHubData[1], mcp_info: { is_public: false, is_public_explicit: false } },
+ ];
+ rerender( );
+ expect(screen.queryByRole("alert")).not.toBeInTheDocument();
+ expect(screen.getByRole("button", { name: "Next" })).toBeEnabled();
+ expect(screen.getByRole("checkbox", { name: "Publish Test Server 1" })).toBeChecked();
+ expect(screen.getByRole("checkbox", { name: "Publish Test Server 2" })).not.toBeChecked();
+ });
+
+ it("copies publication YAML using the selected server IDs", async () => {
+ const user = userEvent.setup();
+ render( );
+
+ await user.click(screen.getByText("Configure in YAML"));
+ await user.click(screen.getByRole("button", { name: "Copy code" }));
+
+ expect(await navigator.clipboard.readText()).toBe(
+ 'litellm_settings:\n public_mcp_hub_strict_whitelist: true\n public_mcp_servers:\n - "server-2"',
+ );
+
+ await user.click(screen.getByRole("checkbox", { name: "Publish Test Server 2" }));
+ await user.click(screen.getByRole("button", { name: "Copy code" }));
+ expect(await navigator.clipboard.readText()).toBe(
+ "litellm_settings:\n public_mcp_hub_strict_whitelist: true\n public_mcp_servers: []",
+ );
+ });
+
+ it("allows clearing publication IDs when the loaded server catalog is empty", async () => {
+ mockMakeMCPPublicCall.mockResolvedValueOnce({});
const emptyProps = {
...mockProps,
mcpHubData: [] as MCPServerData[],
@@ -223,9 +305,13 @@ describe("MakeMCPPublicForm", () => {
const selectAllCheckbox = screen.getByRole("checkbox", { name: "Select All" });
expectDisabledControl(selectAllCheckbox);
- // Next button should be disabled
const nextButton = screen.getByRole("button", { name: "Next" });
- expect(nextButton).toBeDisabled();
+ expect(nextButton).toBeEnabled();
+ fireEvent.click(nextButton);
+ fireEvent.click(screen.getByRole("button", { name: "Save Publication List" }));
+
+ await waitFor(() => expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", []));
+ expect(mockProps.onSuccess).toHaveBeenCalled();
});
it("should handle Cancel button functionality", async () => {
@@ -252,7 +338,7 @@ describe("MakeMCPPublicForm", () => {
// Verify we're on step 1
await waitFor(() => {
- expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument();
+ expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument();
});
// Click Previous button
@@ -262,7 +348,7 @@ describe("MakeMCPPublicForm", () => {
});
// Should go back to step 0
- expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument();
+ expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument();
});
it("should handle individual server selection", async () => {
@@ -322,8 +408,8 @@ describe("MakeMCPPublicForm", () => {
});
it("should handle submit error properly", async () => {
- const errorMessage = "Network error";
- mockMakeMCPPublicCall.mockRejectedValueOnce(new Error(errorMessage));
+ const error = new Error("Update litellm_settings.public_mcp_servers in your YAML configuration");
+ mockMakeMCPPublicCall.mockRejectedValueOnce(error);
render( );
@@ -333,10 +419,10 @@ describe("MakeMCPPublicForm", () => {
});
await waitFor(() => {
- expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument();
+ expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument();
});
- const submitButton = screen.getByRole("button", { name: "Make Public" });
+ const submitButton = screen.getByRole("button", { name: "Save Publication List" });
await act(async () => {
fireEvent.click(submitButton);
});
@@ -346,6 +432,8 @@ describe("MakeMCPPublicForm", () => {
expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", ["server-2"]);
});
+ expect(toast.fromError).toHaveBeenCalledWith(error);
+
// Should not call onSuccess or onClose on error
expect(mockProps.onSuccess).not.toHaveBeenCalled();
expect(mockProps.onClose).not.toHaveBeenCalled();
@@ -366,10 +454,10 @@ describe("MakeMCPPublicForm", () => {
});
await waitFor(() => {
- expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument();
+ expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument();
});
- const submitButton = screen.getByRole("button", { name: "Make Public" });
+ const submitButton = screen.getByRole("button", { name: "Save Publication List" });
await act(async () => {
fireEvent.click(submitButton);
});
@@ -381,7 +469,7 @@ describe("MakeMCPPublicForm", () => {
expect(mockMakeMCPPublicCall).toHaveBeenCalledTimes(1);
expect(mockProps.onSuccess).not.toHaveBeenCalled();
expect(mockProps.onClose).not.toHaveBeenCalled();
- expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument();
+ expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument();
resolvePromise({});
await waitFor(() => {
@@ -400,7 +488,7 @@ describe("MakeMCPPublicForm", () => {
// Modal should not be rendered
expect(screen.queryByRole("dialog")).not.toBeInTheDocument();
- expect(screen.queryByText("Make MCP Servers Public")).not.toBeInTheDocument();
+ expect(screen.queryByText("Manage MCP Hub Visibility")).not.toBeInTheDocument();
});
it("should preselect already public servers when modal opens", () => {
@@ -415,7 +503,7 @@ describe("MakeMCPPublicForm", () => {
url: "http://example.com/server1",
transport: "http",
status: "active",
- mcp_info: { is_public: false }, // Not public
+ mcp_info: { is_public: false, is_public_explicit: false }, // Not public
allowed_tools: [],
auth_type: "bearer",
credentials: {},
@@ -437,7 +525,7 @@ describe("MakeMCPPublicForm", () => {
url: "http://example.com/server2",
transport: "websocket",
status: "inactive",
- mcp_info: { is_public: true }, // Already public
+ mcp_info: { is_public: true, is_public_explicit: true }, // Already public
allowed_tools: [],
auth_type: "none",
credentials: {},
@@ -459,7 +547,7 @@ describe("MakeMCPPublicForm", () => {
url: "http://example.com/server3",
transport: "sse",
status: "healthy",
- mcp_info: { is_public: true }, // Already public
+ mcp_info: { is_public: true, is_public_explicit: true }, // Already public
allowed_tools: [],
auth_type: "oauth",
credentials: {},
diff --git a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx
index 8287cf47f1a..2448732a236 100644
--- a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx
+++ b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx
@@ -1,5 +1,6 @@
import React, { useState, useEffect } from "react";
import { Loader2 } from "lucide-react";
+import CodeBlock from "@/components/CodeBlock";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Checkbox } from "@/components/ui/checkbox";
@@ -29,6 +30,11 @@ interface MakeMCPPublicFormProps {
onSuccess: () => void;
}
+interface PublicationSelection {
+ readonly catalog: MCPServerData[];
+ readonly serverIds: Set;
+}
+
const MakeMCPPublicForm: React.FC = ({
visible,
onClose,
@@ -37,21 +43,28 @@ const MakeMCPPublicForm: React.FC = ({
onSuccess,
}) => {
const [currentStep, setCurrentStep] = useState(0);
- const [selectedServers, setSelectedServers] = useState>(new Set());
+ const [selection, setSelection] = useState(null);
const [loading, setLoading] = useState(false);
+ const selectedServers = selection?.serverIds ?? new Set();
+ const hasPublicationMetadata = mcpHubData.every((server) => typeof server.mcp_info?.is_public_explicit === "boolean");
+ const canManagePublication = hasPublicationMetadata && selection?.catalog === mcpHubData;
+ const publicationYaml = [
+ "litellm_settings:",
+ " public_mcp_hub_strict_whitelist: true",
+ selectedServers.size === 0
+ ? " public_mcp_servers: []"
+ : ` public_mcp_servers:\n${Array.from(selectedServers, (id) => ` - ${JSON.stringify(id)}`).join("\n")}`,
+ ].join("\n");
const handleClose = () => {
setCurrentStep(0);
- setSelectedServers(new Set());
+ setSelection(null);
onClose();
};
const handleNext = () => {
+ if (!canManagePublication) return;
if (currentStep === 0) {
- if (selectedServers.size === 0) {
- toast.fromError("Please select at least one MCP server to make public");
- return;
- }
setCurrentStep(1);
}
};
@@ -69,37 +82,32 @@ const MakeMCPPublicForm: React.FC = ({
} else {
newSelection.delete(serverId);
}
- setSelectedServers(newSelection);
+ setSelection({ catalog: mcpHubData, serverIds: newSelection });
};
const handleSelectAll = (checked: boolean) => {
if (checked) {
const allServerIds = mcpHubData.map((server) => server.server_id);
- setSelectedServers(new Set(allServerIds));
+ setSelection({ catalog: mcpHubData, serverIds: new Set(allServerIds) });
} else {
- setSelectedServers(new Set());
+ setSelection({ catalog: mcpHubData, serverIds: new Set() });
}
};
- // Initialize and preselect already public servers when modal opens
useEffect(() => {
- if (visible && mcpHubData.length > 0) {
- // Extract server IDs from servers that are already public
- const publicServerIds = mcpHubData
- .filter((server) => server.mcp_info?.is_public === true)
- .map((server) => server.server_id);
-
- // Preselect servers that are already public
- setSelectedServers(new Set(publicServerIds));
- }
- }, [visible]); // Only re-run when modal visibility changes, not when mcpHubData updates
-
- const handleSubmit = async () => {
- if (selectedServers.size === 0) {
- toast.fromError("Please select at least one MCP server to make public");
+ if (!visible || !hasPublicationMetadata) {
+ setSelection(null);
return;
}
+ const publicServerIds = mcpHubData
+ .filter((server) => server.mcp_info.is_public_explicit === true)
+ .map((server) => server.server_id);
+ setSelection({ catalog: mcpHubData, serverIds: new Set(publicServerIds) });
+ setCurrentStep(0);
+ }, [visible, mcpHubData, hasPublicationMetadata]);
+ const handleSubmit = async () => {
+ if (!canManagePublication) return;
setLoading(true);
try {
const serverIdsToMakePublic = Array.from(selectedServers);
@@ -107,12 +115,12 @@ const MakeMCPPublicForm: React.FC = ({
// Make batch API call for all servers
await makeMCPPublicCall(accessToken, serverIdsToMakePublic);
- toast.success(`Successfully made ${serverIdsToMakePublic.length} MCP server(s) public!`);
+ toast.success("MCP Hub publication list updated");
handleClose();
onSuccess();
} catch (error) {
console.error("Error making MCP servers public:", error);
- toast.fromError("Failed to make MCP servers public. Please try again.");
+ toast.fromError(error);
} finally {
setLoading(false);
}
@@ -126,7 +134,7 @@ const MakeMCPPublicForm: React.FC = ({
return (
- Select MCP Servers to Make Public
+ Select MCP Servers for the Hub
- Select the MCP servers you want to be visible on the public model hub. Users will still require a valid
- Virtual Key to use these servers.
+ Select the complete list of MCP servers to publish on the public hub. Uncheck a server to remove it from this
+ list, or uncheck all to clear it. Authentication and access permissions still apply
+
+
+
+ Legacy mode also lists servers with public IP access enabled. Set public_mcp_hub_strict_whitelist to true in
+ your configuration to use only the publication list
@@ -160,16 +173,22 @@ const MakeMCPPublicForm: React.FC = ({
className="flex items-center space-x-3 p-3 border rounded-lg hover:bg-accent"
>
handleServerSelection(server.server_id, checked === true)}
/>
{server.server_name}
- {isPublic && Public }
+ {isPublic && (
+
+ {server.mcp_info?.is_public_explicit === false ? "Listed by legacy mode" : "Listed"}
+
+ )}
{server.transport}
{server.status || "unknown"}
+ {server.server_id}
{server.description || server.url}
@@ -193,6 +212,18 @@ const MakeMCPPublicForm: React.FC = ({
+
+ Configure in YAML
+
+
+ Merge these settings into your proxy configuration and reload it. Entries use the server IDs shown above,
+ not names or aliases. For servers defined in YAML, pin server_id in each existing mcp_servers entry so the
+ publication list stays stable
+
+
+
+
+
{selectedServers.size > 0 && (
@@ -207,19 +238,20 @@ const MakeMCPPublicForm: React.FC = ({
const renderStep2Content = () => {
return (
- Confirm Making MCP Servers Public
+ Confirm MCP Hub Publication
- Warning: Once you make these MCP servers public, anyone who can go to the{" "}
- /ui/model_hub_table will be able to know they exist on the proxy.
+ Anyone who can open /ui/model_hub_table can discover published servers. Explicitly published
+ server IDs also allow requests from public IPs. Authentication and access permissions still apply
- MCP Servers to be made public:
+ MCP servers in the publication list:
+ {selectedServers.size === 0 && No explicitly published servers
}
{Array.from(selectedServers).map((serverId) => {
const server = mcpHubData.find((s) => s.server_id === serverId);
return (
@@ -248,8 +280,8 @@ const MakeMCPPublicForm: React.FC = ({
- Total: {selectedServers.size} MCP server{selectedServers.size !== 1 ? "s" : ""} will be
- made public
+ Saving replaces the publication list with {selectedServers.size} MCP server
+ {selectedServers.size !== 1 ? "s" : ""}. Legacy mode may still list servers with public IP access enabled
@@ -257,6 +289,15 @@ const MakeMCPPublicForm: React.FC = ({
};
const renderStepContent = () => {
+ if (!hasPublicationMetadata) {
+ return (
+
+ This proxy does not provide explicit publication status for every MCP server. Update the proxy to manage
+ visibility here, or edit litellm_settings.public_mcp_servers in its existing configuration
+
+ );
+ }
+ if (!canManagePublication) return Loading publication settings
;
switch (currentStep) {
case 0:
return renderStep1Content();
@@ -276,15 +317,15 @@ const MakeMCPPublicForm: React.FC = ({
{currentStep === 0 && (
-
@@ -296,7 +337,7 @@ const MakeMCPPublicForm: React.FC = ({