mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-23 18:25:22 +00:00
Revert "Stabilise mock tests"
This commit is contained in:
@@ -761,6 +761,7 @@ async def handle_bedrock_passthrough_router_model(
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
async def handle_bedrock_count_tokens(
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
|
||||
+26
-16
@@ -2,28 +2,14 @@
|
||||
Integration tests for responses API background cost tracking
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
# Import litellm first to ensure it's in sys.modules before enterprise imports
|
||||
import litellm # noqa: E402
|
||||
|
||||
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse # noqa: E402
|
||||
|
||||
# Now import enterprise modules
|
||||
try:
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import ( # noqa: E402
|
||||
CheckResponsesCost,
|
||||
)
|
||||
except ImportError as e:
|
||||
# Skip all tests in this module if enterprise module is not available
|
||||
pytest.skip(f"Enterprise module not available: {e}", allow_module_level=True)
|
||||
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
|
||||
|
||||
|
||||
class TestResponsesBackgroundCostTracking:
|
||||
@@ -298,6 +284,10 @@ class TestCheckResponsesCost:
|
||||
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Test CheckResponsesCost initialization"""
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import (
|
||||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
checker = CheckResponsesCost(
|
||||
proxy_logging_obj=mock_proxy_logging_obj,
|
||||
prisma_client=mock_prisma_client,
|
||||
@@ -313,6 +303,10 @@ class TestCheckResponsesCost:
|
||||
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Test polling when there are no jobs"""
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import (
|
||||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
# Mock find_many to return empty list
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
@@ -340,6 +334,10 @@ class TestCheckResponsesCost:
|
||||
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Test polling with a completed job"""
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import (
|
||||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
# Create a mock job
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-123"
|
||||
@@ -393,6 +391,10 @@ class TestCheckResponsesCost:
|
||||
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Test polling with a failed job"""
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import (
|
||||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
# Create a mock job
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-456"
|
||||
@@ -433,6 +435,10 @@ class TestCheckResponsesCost:
|
||||
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Test polling with a job still in progress"""
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import (
|
||||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
# Create a mock job
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-789"
|
||||
@@ -473,6 +479,10 @@ class TestCheckResponsesCost:
|
||||
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Test that errors when querying responses are handled gracefully"""
|
||||
from litellm_enterprise.proxy.common_utils.check_responses_cost import (
|
||||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
# Create a mock job
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-error"
|
||||
+45
-59
@@ -101,69 +101,55 @@ async def test_bedrock_converse_budget_tokens_preserved():
|
||||
The bug was that the messages -> completion adapter was converting thinking to reasoning_effort
|
||||
and losing the original budget_tokens value, causing it to use the default (128) instead.
|
||||
"""
|
||||
import os
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
# Mock at httpx level for better CI compatibility
|
||||
with patch("httpx.AsyncClient.post") as mock_httpx_post:
|
||||
with patch.object(client, "post") as mock_post:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
mock_response.text = "mock response"
|
||||
mock_response.json.return_value = {
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "4"}]
|
||||
}
|
||||
},
|
||||
"stopReason": "end_turn",
|
||||
"usage": {
|
||||
"inputTokens": 10,
|
||||
"outputTokens": 5,
|
||||
"totalTokens": 15
|
||||
with patch.object(client, "post") as mock_post:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
mock_response.text = "mock response"
|
||||
mock_response.json.return_value = {
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "4"}]
|
||||
}
|
||||
},
|
||||
"stopReason": "end_turn",
|
||||
"usage": {
|
||||
"inputTokens": 10,
|
||||
"outputTokens": 5,
|
||||
"totalTokens": 15
|
||||
}
|
||||
mock_post.return_value = mock_response
|
||||
mock_httpx_post.return_value = mock_response
|
||||
|
||||
try:
|
||||
await messages.acreate(
|
||||
client=client,
|
||||
max_tokens=1024,
|
||||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||||
model="bedrock/converse/us.anthropic.claude-sonnet-4-20250514-v1:0",
|
||||
thinking={
|
||||
"budget_tokens": 1024,
|
||||
"type": "enabled"
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
pass # Expected due to mock response format
|
||||
|
||||
# Check which mock was called (client.post or httpx.AsyncClient.post)
|
||||
if mock_post.call_count == 0 and mock_httpx_post.call_count == 0:
|
||||
# Skip test if neither mock was called (CI environment issue)
|
||||
if os.getenv("CI") == "true":
|
||||
pytest.skip("Mock not intercepted in CI environment")
|
||||
else:
|
||||
pytest.fail("Expected mock to be called but it wasn't")
|
||||
|
||||
# Use whichever mock was actually called
|
||||
active_mock = mock_post if mock_post.call_count > 0 else mock_httpx_post
|
||||
|
||||
call_kwargs = active_mock.call_args.kwargs
|
||||
json_data = call_kwargs.get("json") or json.loads(call_kwargs.get("data", "{}"))
|
||||
print("Request json: ", json.dumps(json_data, indent=4, default=str))
|
||||
|
||||
additional_fields = json_data.get("additionalModelRequestFields", {})
|
||||
thinking_config = additional_fields.get("thinking", {})
|
||||
|
||||
assert "thinking" in additional_fields, "thinking parameter should be in additionalModelRequestFields"
|
||||
assert thinking_config.get("type") == "enabled", "thinking.type should be 'enabled'"
|
||||
assert thinking_config.get("budget_tokens") == 1024, f"thinking.budget_tokens should be 1024, but got {thinking_config.get('budget_tokens')}"
|
||||
}
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
try:
|
||||
await messages.acreate(
|
||||
client=client,
|
||||
max_tokens=1024,
|
||||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||||
model="bedrock/converse/us.anthropic.claude-sonnet-4-20250514-v1:0",
|
||||
thinking={
|
||||
"budget_tokens": 1024,
|
||||
"type": "enabled"
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
pass # Expected due to mock response format
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
call_kwargs = mock_post.call_args.kwargs
|
||||
json_data = call_kwargs.get("json") or json.loads(call_kwargs.get("data", "{}"))
|
||||
print("Request json: ", json.dumps(json_data, indent=4, default=str))
|
||||
|
||||
additional_fields = json_data.get("additionalModelRequestFields", {})
|
||||
thinking_config = additional_fields.get("thinking", {})
|
||||
|
||||
assert "thinking" in additional_fields, "thinking parameter should be in additionalModelRequestFields"
|
||||
assert thinking_config.get("type") == "enabled", "thinking.type should be 'enabled'"
|
||||
assert thinking_config.get("budget_tokens") == 1024, f"thinking.budget_tokens should be 1024, but got {thinking_config.get('budget_tokens')}"
|
||||
|
||||
|
||||
def test_openai_model_with_thinking_converts_to_reasoning_effort():
|
||||
|
||||
@@ -2610,6 +2610,99 @@ def test_request_metadata_not_provided():
|
||||
assert "requestMetadata" not in request_data
|
||||
|
||||
|
||||
def test_empty_assistant_message_handling():
|
||||
"""
|
||||
Test that empty assistant messages are handled correctly by replacing
|
||||
empty or whitespace-only content with a placeholder to prevent AWS Bedrock
|
||||
Converse API 400 Bad Request errors.
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
_bedrock_converse_messages_pt,
|
||||
)
|
||||
|
||||
# Test case 1: Empty string content - test with modify_params=True to prevent merging
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": ""}, # Empty content
|
||||
{"role": "user", "content": "How are you?"}
|
||||
]
|
||||
|
||||
# Enable modify_params to prevent consecutive user message merging
|
||||
original_modify_params = litellm.modify_params
|
||||
litellm.modify_params = True
|
||||
|
||||
try:
|
||||
result = _bedrock_converse_messages_pt(
|
||||
messages=messages,
|
||||
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
llm_provider="bedrock_converse"
|
||||
)
|
||||
|
||||
# Should have 3 messages: user, assistant (with placeholder), user
|
||||
assert len(result) == 3
|
||||
assert result[0]["role"] == "user"
|
||||
assert result[1]["role"] == "assistant"
|
||||
assert result[2]["role"] == "user"
|
||||
|
||||
# Assistant message should have placeholder text instead of empty content
|
||||
assert len(result[1]["content"]) == 1
|
||||
assert result[1]["content"][0]["text"] == "Please continue."
|
||||
|
||||
# Test case 2: Whitespace-only content
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": " "}, # Whitespace-only content
|
||||
{"role": "user", "content": "How are you?"}
|
||||
]
|
||||
|
||||
result = _bedrock_converse_messages_pt(
|
||||
messages=messages,
|
||||
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
llm_provider="bedrock_converse"
|
||||
)
|
||||
|
||||
# Assistant message should have placeholder text instead of whitespace
|
||||
assert len(result[1]["content"]) == 1
|
||||
assert result[1]["content"][0]["text"] == "Please continue."
|
||||
|
||||
# Test case 3: Empty list content
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": ""}]}, # Empty text in list
|
||||
{"role": "user", "content": "How are you?"}
|
||||
]
|
||||
|
||||
result = _bedrock_converse_messages_pt(
|
||||
messages=messages,
|
||||
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
llm_provider="bedrock_converse"
|
||||
)
|
||||
|
||||
# Assistant message should have placeholder text instead of empty text
|
||||
assert len(result[1]["content"]) == 1
|
||||
assert result[1]["content"][0]["text"] == "Please continue."
|
||||
|
||||
# Test case 4: Normal content should not be affected
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "I'm doing well, thank you!"}, # Normal content
|
||||
{"role": "user", "content": "How are you?"}
|
||||
]
|
||||
|
||||
result = _bedrock_converse_messages_pt(
|
||||
messages=messages,
|
||||
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
llm_provider="bedrock_converse"
|
||||
)
|
||||
|
||||
# Assistant message should keep original content
|
||||
assert len(result[1]["content"]) == 1
|
||||
assert result[1]["content"][0]["text"] == "I'm doing well, thank you!"
|
||||
|
||||
finally:
|
||||
# Restore original modify_params setting
|
||||
litellm.modify_params = original_modify_params
|
||||
|
||||
|
||||
def test_is_nova_lite_2_model():
|
||||
"""Test the _is_nova_lite_2_model() method for detecting Nova 2 models."""
|
||||
|
||||
@@ -21,51 +21,43 @@ class TestBedrockFilesIntegration:
|
||||
file_id = "s3://test-bucket/test-file.jsonl"
|
||||
expected_content = b'{"recordId": "request-1", "modelInput": {}, "modelOutput": {}}'
|
||||
|
||||
# Mock AWS credentials
|
||||
with patch.dict(
|
||||
"os.environ",
|
||||
{
|
||||
"AWS_ACCESS_KEY_ID": "test-access-key",
|
||||
"AWS_SECRET_ACCESS_KEY": "test-secret-key",
|
||||
},
|
||||
):
|
||||
# Mock the bedrock_files_instance.file_content method
|
||||
with patch(
|
||||
"litellm.files.main.bedrock_files_instance.file_content",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_file_content:
|
||||
# Create a mock HttpxBinaryResponseContent response
|
||||
import httpx
|
||||
# Mock the bedrock_files_instance.file_content method
|
||||
with patch(
|
||||
"litellm.files.main.bedrock_files_instance.file_content",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_file_content:
|
||||
# Create a mock HttpxBinaryResponseContent response
|
||||
import httpx
|
||||
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(
|
||||
method="GET", url="s3://test-bucket/test-file.jsonl"
|
||||
),
|
||||
)
|
||||
mock_file_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
)
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(
|
||||
method="GET", url="s3://test-bucket/test-file.jsonl"
|
||||
),
|
||||
)
|
||||
mock_file_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
)
|
||||
|
||||
# Call litellm.afile_content
|
||||
result = await litellm.afile_content(
|
||||
file_id=file_id,
|
||||
custom_llm_provider="bedrock",
|
||||
aws_region_name="us-west-2",
|
||||
)
|
||||
# Call litellm.afile_content
|
||||
result = await litellm.afile_content(
|
||||
file_id=file_id,
|
||||
custom_llm_provider="bedrock",
|
||||
aws_region_name="us-west-2",
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == expected_content
|
||||
assert result.response.status_code == 200
|
||||
# Verify the result
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == expected_content
|
||||
assert result.response.status_code == 200
|
||||
|
||||
# Verify the mock was called with correct parameters
|
||||
mock_file_content.assert_called_once()
|
||||
call_kwargs = mock_file_content.call_args.kwargs
|
||||
assert call_kwargs["_is_async"] is True
|
||||
assert call_kwargs["file_content_request"]["file_id"] == file_id
|
||||
# Verify the mock was called with correct parameters
|
||||
mock_file_content.assert_called_once()
|
||||
call_kwargs = mock_file_content.call_args.kwargs
|
||||
assert call_kwargs["_is_async"] is True
|
||||
assert call_kwargs["file_content_request"]["file_id"] == file_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_afile_content_bedrock_provider_with_unified_file_id(self):
|
||||
@@ -80,47 +72,39 @@ class TestBedrockFilesIntegration:
|
||||
|
||||
expected_content = b'{"recordId": "request-1", "modelInput": {}, "modelOutput": {}}'
|
||||
|
||||
# Mock AWS credentials
|
||||
with patch.dict(
|
||||
"os.environ",
|
||||
{
|
||||
"AWS_ACCESS_KEY_ID": "test-access-key",
|
||||
"AWS_SECRET_ACCESS_KEY": "test-secret-key",
|
||||
},
|
||||
):
|
||||
# Mock the bedrock_files_instance.file_content method
|
||||
with patch(
|
||||
"litellm.files.main.bedrock_files_instance.file_content",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_file_content:
|
||||
# Create a mock HttpxBinaryResponseContent response
|
||||
import httpx
|
||||
# Mock the bedrock_files_instance.file_content method
|
||||
with patch(
|
||||
"litellm.files.main.bedrock_files_instance.file_content",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_file_content:
|
||||
# Create a mock HttpxBinaryResponseContent response
|
||||
import httpx
|
||||
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(method="GET", url=s3_uri),
|
||||
)
|
||||
mock_file_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
)
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(method="GET", url=s3_uri),
|
||||
)
|
||||
mock_file_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
)
|
||||
|
||||
# Call litellm.afile_content with unified file ID
|
||||
result = await litellm.afile_content(
|
||||
file_id=encoded_file_id,
|
||||
custom_llm_provider="bedrock",
|
||||
aws_region_name="us-west-2",
|
||||
)
|
||||
# Call litellm.afile_content with unified file ID
|
||||
result = await litellm.afile_content(
|
||||
file_id=encoded_file_id,
|
||||
custom_llm_provider="bedrock",
|
||||
aws_region_name="us-west-2",
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == expected_content
|
||||
assert result.response.status_code == 200
|
||||
# Verify the result
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == expected_content
|
||||
assert result.response.status_code == 200
|
||||
|
||||
# Verify the mock was called - the handler should extract S3 URI from unified file ID
|
||||
mock_file_content.assert_called_once()
|
||||
call_kwargs = mock_file_content.call_args.kwargs
|
||||
assert call_kwargs["_is_async"] is True
|
||||
# The handler extracts S3 URI from the unified file ID
|
||||
assert call_kwargs["file_content_request"]["file_id"] == encoded_file_id
|
||||
# Verify the mock was called - the handler should extract S3 URI from unified file ID
|
||||
mock_file_content.assert_called_once()
|
||||
call_kwargs = mock_file_content.call_args.kwargs
|
||||
assert call_kwargs["_is_async"] is True
|
||||
# The handler extracts S3 URI from the unified file ID
|
||||
assert call_kwargs["file_content_request"]["file_id"] == encoded_file_id
|
||||
|
||||
@@ -41,12 +41,8 @@ def mock_embedding_async_http_handler():
|
||||
class TestHuggingFaceEmbedding:
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup(self, mock_embedding_http_handler, mock_embedding_async_http_handler):
|
||||
# Mock both sync and async versions of get_hf_task functions
|
||||
self.mock_get_task_patcher = patch("litellm.llms.huggingface.embedding.handler.get_hf_task_embedding_for_model")
|
||||
self.mock_get_task_async_patcher = patch("litellm.llms.huggingface.embedding.handler.async_get_hf_task_embedding_for_model", new_callable=AsyncMock)
|
||||
|
||||
self.mock_get_task = self.mock_get_task_patcher.start()
|
||||
self.mock_get_task_async = self.mock_get_task_async_patcher.start()
|
||||
|
||||
def mock_get_task_side_effect(model, task_type, api_base):
|
||||
if task_type is not None:
|
||||
@@ -54,7 +50,6 @@ class TestHuggingFaceEmbedding:
|
||||
return "sentence-similarity"
|
||||
|
||||
self.mock_get_task.side_effect = mock_get_task_side_effect
|
||||
self.mock_get_task_async.side_effect = mock_get_task_side_effect
|
||||
|
||||
self.model = "huggingface/BAAI/bge-m3"
|
||||
self.mock_http = mock_embedding_http_handler
|
||||
@@ -64,7 +59,6 @@ class TestHuggingFaceEmbedding:
|
||||
yield
|
||||
|
||||
self.mock_get_task_patcher.stop()
|
||||
self.mock_get_task_async_patcher.stop()
|
||||
|
||||
def test_input_type_preserved_in_optional_params(self):
|
||||
input_text = ["hello world"]
|
||||
@@ -87,3 +81,31 @@ class TestHuggingFaceEmbedding:
|
||||
# Should NOT have sentence-similarity format
|
||||
assert "source_sentence" not in str(request_data)
|
||||
assert "sentences" not in str(request_data)
|
||||
|
||||
def test_embedding_with_sentence_similarity_task(self):
|
||||
"""Test embedding when task type is sentence-similarity (requires 2+ sentences)"""
|
||||
|
||||
similarity_response = {
|
||||
"similarities": [[0, 0.9], [1, 0.8]]
|
||||
}
|
||||
|
||||
self.mock_http.return_value.json.return_value = similarity_response
|
||||
|
||||
# Test with 2+ sentences (required for sentence-similarity)
|
||||
input_text = ["This is the source sentence", "This is sentence one", "This is sentence two"]
|
||||
|
||||
response = litellm.embedding(
|
||||
model=self.model,
|
||||
input=input_text,
|
||||
# Use the model's natural task type (sentence-similarity)
|
||||
)
|
||||
|
||||
self.mock_http.assert_called_once()
|
||||
post_call_args = self.mock_http.call_args
|
||||
request_data = json.loads(post_call_args[1]["data"])
|
||||
|
||||
assert "inputs" in request_data
|
||||
assert "source_sentence" in request_data["inputs"]
|
||||
assert "sentences" in request_data["inputs"]
|
||||
assert request_data["inputs"]["source_sentence"] == input_text[0]
|
||||
assert request_data["inputs"]["sentences"] == input_text[1:]
|
||||
@@ -12,7 +12,53 @@ from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
class TestVertexAIFilesIntegration:
|
||||
"""Test integration of Vertex AI files with main litellm API"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_afile_content_vertex_ai_provider(self):
|
||||
"""Test litellm.afile_content with vertex_ai provider"""
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
# Mock the vertex_ai_files_instance.file_content method
|
||||
with patch(
|
||||
"litellm.files.main.vertex_ai_files_instance.file_content",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_file_content:
|
||||
# Create a mock HttpxBinaryResponseContent response
|
||||
import httpx
|
||||
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(
|
||||
method="GET", url="gs://test-bucket/test-file.txt"
|
||||
),
|
||||
)
|
||||
mock_file_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
)
|
||||
|
||||
# Call litellm.afile_content
|
||||
result = await litellm.afile_content(
|
||||
file_id=file_id,
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
vertex_credentials=None,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == expected_content
|
||||
assert result.response.status_code == 200
|
||||
|
||||
# Verify the mock was called with correct parameters
|
||||
mock_file_content.assert_called_once()
|
||||
call_kwargs = mock_file_content.call_args.kwargs
|
||||
assert call_kwargs["_is_async"] is True
|
||||
assert call_kwargs["file_content_request"]["file_id"] == file_id
|
||||
assert call_kwargs["vertex_project"] == "test-project"
|
||||
assert call_kwargs["vertex_location"] == "us-central1"
|
||||
|
||||
def test_litellm_file_content_vertex_ai_provider(self):
|
||||
"""Test litellm.file_content with vertex_ai provider (sync)"""
|
||||
|
||||
@@ -75,6 +75,40 @@ class TestCreateToolFunction:
|
||||
call_args[0][0]
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_leading_digit_parameter(self):
|
||||
"""Test function with parameter starting with digit (e.g., 2fa-code)."""
|
||||
operation = {
|
||||
"parameters": [
|
||||
{
|
||||
"name": "2fa-code",
|
||||
"in": "query",
|
||||
"required": False,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
func = create_tool_function(
|
||||
path="/verify",
|
||||
method="post",
|
||||
operation=operation,
|
||||
base_url="https://api.example.com",
|
||||
)
|
||||
|
||||
assert callable(func)
|
||||
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("post", "verified")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"2fa-code": "123456"})
|
||||
assert result == "verified"
|
||||
|
||||
# Verify query parameter was included
|
||||
call_args = async_client.post.call_args
|
||||
assert call_args[1]["params"]["2fa-code"] == "123456"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dot_in_parameter_name(self):
|
||||
"""Test function with dot in parameter name (e.g., user.name)."""
|
||||
|
||||
@@ -8,7 +8,7 @@ and following LiteLLM testing patterns and best practices.
|
||||
# Standard library imports
|
||||
import os
|
||||
import sys
|
||||
from typing import Dict, Any
|
||||
from typing import Dict
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
# Add parent directory to path for imports
|
||||
@@ -43,6 +43,33 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
Standard LiteLLM fixture that reloads litellm before every function
|
||||
to speed up testing by removing callbacks being chained.
|
||||
"""
|
||||
import importlib
|
||||
import asyncio
|
||||
|
||||
# Reload litellm to ensure clean state
|
||||
importlib.reload(litellm)
|
||||
|
||||
# Set up async loop
|
||||
loop = asyncio.get_event_loop_policy().new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
|
||||
# Set up litellm state
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
yield
|
||||
|
||||
# Teardown
|
||||
loop.close()
|
||||
asyncio.set_event_loop(None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def env_setup(monkeypatch):
|
||||
"""Fixture to set up environment variables for testing."""
|
||||
|
||||
@@ -1393,23 +1393,21 @@ async def test_embedding_header_forwarding_with_model_group():
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
# Verify that headers were added to the request metadata
|
||||
assert "metadata" in updated_data, "Metadata should be added to embedding request"
|
||||
assert "headers" in updated_data["metadata"], "Headers should be added to embedding request metadata"
|
||||
# Verify that headers were added to the request data
|
||||
assert "headers" in updated_data, "Headers should be added to embedding request"
|
||||
|
||||
# Verify that only x- prefixed headers (except x-stainless) were forwarded
|
||||
forwarded_headers = updated_data["metadata"]["headers"]
|
||||
forwarded_headers = updated_data["headers"]
|
||||
assert "X-Custom-Header" in forwarded_headers, "X-Custom-Header should be forwarded"
|
||||
assert forwarded_headers["X-Custom-Header"] == "custom-value"
|
||||
assert "X-Request-ID" in forwarded_headers, "X-Request-ID should be forwarded"
|
||||
assert forwarded_headers["X-Request-ID"] == "test-request-123"
|
||||
|
||||
# Verify that Authorization header is present in metadata (not filtered out at this level)
|
||||
# Note: The metadata headers contain all original headers for logging/tracking purposes
|
||||
assert "Authorization" in forwarded_headers, "Authorization header should be in metadata headers"
|
||||
# Verify that authorization header was NOT forwarded (sensitive header)
|
||||
assert "Authorization" not in forwarded_headers, "Authorization header should not be forwarded"
|
||||
|
||||
# Verify that Content-Type is present (it's included in metadata headers)
|
||||
assert "Content-Type" in forwarded_headers, "Content-Type should be in metadata headers"
|
||||
# Verify that Content-Type was NOT forwarded (doesn't start with x-)
|
||||
assert "Content-Type" not in forwarded_headers, "Content-Type should not be forwarded"
|
||||
|
||||
# Verify original data fields are preserved
|
||||
assert updated_data["model"] == "local-openai/text-embedding-3-small"
|
||||
|
||||
@@ -55,7 +55,7 @@ example_embedding_result = {
|
||||
|
||||
def mock_patch_aembedding():
|
||||
return mock.patch(
|
||||
"litellm.aembedding",
|
||||
"litellm.proxy.proxy_server.llm_router.aembedding",
|
||||
return_value=example_embedding_result,
|
||||
)
|
||||
|
||||
@@ -668,6 +668,43 @@ def test_team_info_masking():
|
||||
assert "public-test-key" not in str(exc_info.value)
|
||||
|
||||
|
||||
@mock_patch_aembedding()
|
||||
def test_embedding_input_array_of_tokens(mock_aembedding, client_no_auth):
|
||||
"""
|
||||
Test to bypass decoding input as array of tokens for selected providers
|
||||
|
||||
Ref: https://github.com/BerriAI/litellm/issues/10113
|
||||
"""
|
||||
try:
|
||||
test_data = {
|
||||
"model": "vllm_embed_model",
|
||||
"input": [[2046, 13269, 158208]],
|
||||
}
|
||||
|
||||
response = client_no_auth.post("/v1/embeddings", json=test_data)
|
||||
|
||||
# DEPRECATED - mock_aembedding.assert_called_once_with is too strict, and will fail when new kwargs are added to embeddings
|
||||
# mock_aembedding.assert_called_once_with(
|
||||
# model="vllm_embed_model",
|
||||
# input=[[2046, 13269, 158208]],
|
||||
# metadata=mock.ANY,
|
||||
# proxy_server_request=mock.ANY,
|
||||
# secret_fields=mock.ANY,
|
||||
# )
|
||||
# Assert that aembedding was called, and that input was not modified
|
||||
mock_aembedding.assert_called_once()
|
||||
call_args, call_kwargs = mock_aembedding.call_args
|
||||
assert call_kwargs["model"] == "vllm_embed_model"
|
||||
assert call_kwargs["input"] == [[2046, 13269, 158208]]
|
||||
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
print(len(result["data"][0]["embedding"]))
|
||||
assert len(result["data"][0]["embedding"]) > 10 # this usually has len==1536 so
|
||||
except Exception as e:
|
||||
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_all_team_models():
|
||||
"""
|
||||
|
||||
@@ -1231,30 +1231,18 @@ async def test_acompletion_streaming_disable_fallbacks_midstream():
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
if self.index == self.error_after_index:
|
||||
raise self.error
|
||||
if self.index >= len(self.items):
|
||||
raise StopAsyncIteration
|
||||
if self.index == self.error_after_index:
|
||||
raise self.error
|
||||
item = self.items[self.index]
|
||||
self.index += 1
|
||||
self.chunks.append(item)
|
||||
return item
|
||||
|
||||
# Create properly structured mock chunks using ModelResponse
|
||||
from litellm.types.utils import Delta, ModelResponse, StreamingChoices
|
||||
|
||||
mock_chunk = ModelResponse(
|
||||
id="chatcmpl-123",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0, delta=Delta(content="Hello", role="assistant"), finish_reason=None
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
mock_chunks = [mock_chunk]
|
||||
mock_chunks = [
|
||||
MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello"))]),
|
||||
]
|
||||
|
||||
mock_error_response = AsyncIteratorWithError(
|
||||
mock_chunks, 1, error_with_original
|
||||
|
||||
Reference in New Issue
Block a user