mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-14 14:26:43 +00:00
Merge pull request #17769 from BerriAI/litellm_test_fix
Fix nvdia and geminin tests
This commit is contained in:
@@ -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 = {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user