Merge pull request #17769 from BerriAI/litellm_test_fix

Fix nvdia and geminin tests
This commit is contained in:
Sameer Kankute
2025-12-10 22:29:04 +05:30
committed by GitHub
3 changed files with 25 additions and 4 deletions
@@ -28,9 +28,13 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig):
"""
def _get_clean_model_name(self, model: str) -> str:
"""Strip 'ranking/' prefix from model name."""
"""Strip 'nvidia_nim/' and 'ranking/' prefixes from model name."""
# First strip nvidia_nim/ prefix if present
if model.startswith("nvidia_nim/"):
model = model[len("nvidia_nim/"):]
# Then strip ranking/ prefix if present
if model.startswith("ranking/"):
return model[len("ranking/"):]
model = model[len("ranking/"):]
return model
def get_complete_url(
@@ -55,6 +55,12 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
def __init__(self) -> None:
pass
def _get_clean_model_name(self, model: str) -> str:
"""Strip 'nvidia_nim/' prefix from model name if present."""
if model.startswith("nvidia_nim/"):
return model[len("nvidia_nim/"):]
return model
def get_complete_url(
self,
api_base: Optional[str],
@@ -82,7 +88,10 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
if api_base.endswith("/v1"):
api_base = api_base[:-3]
return f"{api_base}/v1/retrieval/{model}/reranking"
# Strip nvidia_nim/ prefix from model name if present
clean_model = self._get_clean_model_name(model)
return f"{api_base}/v1/retrieval/{clean_model}/reranking"
def get_supported_cohere_rerank_params(self, model: str) -> list:
"""
@@ -210,9 +219,12 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
else:
passages.append({"text": str(doc)})
# Strip nvidia_nim/ prefix from model name if present
clean_model = self._get_clean_model_name(model)
# Note: URL path uses underscores (llama-3_2) but JSON body uses periods (llama-3.2)
# Convert underscores back to periods for the model field in request body
model_for_body = model.replace("_", ".")
model_for_body = clean_model.replace("_", ".")
# Build request using TypedDict
request_data: NvidiaNimRerankRequest = {
+5
View File
@@ -738,6 +738,11 @@ async def test_gemini_image_generation_async():
CONTENT = response.choices[0].message.content
# Check if images list exists and has items before accessing
assert hasattr(response.choices[0].message, "images"), "Response message should have images attribute"
assert response.choices[0].message.images is not None, "Images should not be None"
assert len(response.choices[0].message.images) > 0, "Images list should not be empty"
IMAGE_URL = response.choices[0].message.images[0]["image_url"]
print("IMAGE_URL: ", IMAGE_URL)