Merge pull request #38752 from BerriAI/litellm_deflake_20260829

fix: bound Hugging Face config fetch and keep embedding tests off the network
This commit is contained in:
Mateo Wang 2026-08-29 13:33:06 -07:00 • committed by GitHub
commit 306daf13b5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 45 additions and 3 deletions

View file

@ -35,6 +35,7 @@ DEFAULT_COOLDOWN_TIME_SECONDS: Final = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECO
DEFAULT_REPLICATE_POLLING_RETRIES: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5))
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1))
DEFAULT_IMAGE_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250))
HF_CONFIG_FETCH_TIMEOUT_SECONDS: Final = 10.0
# Maximum wall-clock seconds a streaming response is allowed to run.
# Streams exceeding this duration are terminated with a Timeout error.

View file

@ -69,6 +69,7 @@ from litellm.constants import (
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_TRIM_RATIO,
FUNCTION_DEFINITION_TOKEN_COUNT,
HF_CONFIG_FETCH_TIMEOUT_SECONDS,
INITIAL_RETRY_DELAY,
JITTER,
MAX_RETRY_DELAY,
@ -5168,7 +5169,7 @@ def get_max_tokens(model: str) -> int | None:
config_url: Final = f"https://huggingface.co/{model_name}/raw/main/config.json"
try:
# Make the HTTP request to get the raw JSON file
response: Final = litellm.module_level_client.get(config_url)
response: Final = litellm.module_level_client.get(config_url, timeout=HF_CONFIG_FETCH_TIMEOUT_SECONDS)
response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx)
# Parse the JSON response
@ -5522,7 +5523,7 @@ def _get_max_position_embeddings(model_name: str) -> int | None:
try:
# Make the HTTP request to get the raw JSON file
response: Final = litellm.module_level_client.get(config_url)
response: Final = litellm.module_level_client.get(config_url, timeout=HF_CONFIG_FETCH_TIMEOUT_SECONDS)
response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx)
# Parse the JSON response

View file

@ -4,6 +4,7 @@ from unittest.mock import patch, MagicMock, AsyncMock
import litellm
import pytest
import respx
MOCK_EMBEDDING_RESPONSE = [[0.1, 0.2, 0.3, 0.4, 0.5]]
@ -21,6 +22,16 @@ def mock_embedding_http_handler():
yield mock_post
@pytest.fixture
def mock_hf_config_fetch():
"""Serve the Hugging Face config.json fetched during cost calculation, so no test leaves the process"""
with respx.mock(assert_all_called=False) as respx_mock:
respx_mock.get(url__regex=r"https://huggingface\.co/.*/config\.json").respond(
json={"max_position_embeddings": 512}
)
yield respx_mock
@pytest.fixture
def mock_embedding_async_http_handler():
"""Fixture to mock the async HTTP handler for embedding tests"""
@ -39,7 +50,7 @@ def mock_embedding_async_http_handler():
class TestHuggingFaceEmbedding:
@pytest.fixture(autouse=True)
def setup(self, mock_embedding_http_handler, mock_embedding_async_http_handler):
def setup(self, mock_embedding_http_handler, mock_embedding_async_http_handler, mock_hf_config_fetch):
self.mock_get_task_patcher = patch(
"litellm.llms.huggingface.embedding.handler.get_hf_task_embedding_for_model"
)

View file

@ -6,6 +6,7 @@ from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import respx
from jsonschema import validate
@ -5736,3 +5737,31 @@ class TestDefaultReasoningEffortHydration:
model_info = dict(_get_model_info_helper(model="gpt-5.6-terra", custom_llm_provider="openai"))
assert model_info.get("default_reasoning_effort") is None
class TestHuggingFaceConfigFetch:
"""The Hugging Face config.json fetch runs on background logging threads during cost
calculation, so an unbounded request can hang a whole test job; the timeout is the fix."""
@pytest.fixture
def hf_config_route(self):
with respx.mock(assert_all_called=True) as respx_mock:
yield respx_mock.get(url__regex=r"https://huggingface\.co/.*/config\.json").respond(
json={"max_position_embeddings": 512}
)
def test_get_max_tokens_reads_hf_config_with_a_bounded_timeout(self, hf_config_route):
from litellm.constants import HF_CONFIG_FETCH_TIMEOUT_SECONDS
from litellm.utils import get_max_tokens
assert get_max_tokens("huggingface/some-org/some-model") == 512
request_timeout = hf_config_route.calls.last.request.extensions["timeout"]
assert request_timeout["read"] == HF_CONFIG_FETCH_TIMEOUT_SECONDS
def test_get_max_position_embeddings_reads_hf_config_with_a_bounded_timeout(self, hf_config_route):
from litellm.constants import HF_CONFIG_FETCH_TIMEOUT_SECONDS
from litellm.utils import _get_max_position_embeddings
assert _get_max_position_embeddings("some-org/some-model") == 512
request_timeout = hf_config_route.calls.last.request.extensions["timeout"]
assert request_timeout["read"] == HF_CONFIG_FETCH_TIMEOUT_SECONDS