mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-02 04:21:34 +00:00
test snowflake
This commit is contained in:
@@ -1,23 +1,27 @@
|
||||
import os
|
||||
import sys
|
||||
import asyncio
|
||||
import json
|
||||
import httpx
|
||||
from typing import Any, Dict, List
|
||||
from unittest.mock import Mock, MagicMock, patch
|
||||
from dotenv import load_dotenv
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
load_dotenv()
|
||||
import pytest
|
||||
|
||||
from litellm import completion, acompletion, responses
|
||||
from litellm.exceptions import APIConnectionError
|
||||
from litellm import completion, acompletion
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
FAKE_API_BASE = "https://fake-snowflake.example.com/api/v2/cortex/inference:chat"
|
||||
|
||||
def mock_snowflake_chat_response() -> Dict[str, Any]:
|
||||
"""
|
||||
Mock response for Snowflake chat completion.
|
||||
"""
|
||||
|
||||
def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock:
|
||||
mock = MagicMock(spec=httpx.Response)
|
||||
mock.status_code = 200
|
||||
mock.headers = {"content-type": "application/json"}
|
||||
mock.json.return_value = json_data
|
||||
mock.text = json.dumps(json_data)
|
||||
return mock
|
||||
|
||||
|
||||
def _chat_response() -> Dict[str, Any]:
|
||||
return {
|
||||
"id": "chatcmpl-snowflake-123",
|
||||
"object": "chat.completion",
|
||||
@@ -28,7 +32,7 @@ def mock_snowflake_chat_response() -> Dict[str, Any]:
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "The sky above is painted blue,\nWith clouds of white and morning dew.\nA canvas vast, serene and bright,\nThat fills my heart with pure delight.",
|
||||
"content": "The sky above is painted blue,\nWith clouds of white and morning dew.",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
@@ -41,175 +45,127 @@ def mock_snowflake_chat_response() -> Dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def mock_snowflake_streaming_response_chunks() -> List[str]:
|
||||
"""
|
||||
Mock streaming response chunks for Snowflake.
|
||||
"""
|
||||
return [
|
||||
json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-snowflake-stream-123",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": "mistral-7b",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant", "content": "The"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-snowflake-stream-123",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": "mistral-7b",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"content": " sky"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-snowflake-stream-123",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": "mistral-7b",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"content": " is blue"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
def _streaming_chunks() -> List[str]:
|
||||
base = {
|
||||
"id": "chatcmpl-snowflake-stream-123",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": "mistral-7b",
|
||||
}
|
||||
deltas = [
|
||||
{"role": "assistant", "content": "The"},
|
||||
{"content": " sky"},
|
||||
{"content": " is blue"},
|
||||
]
|
||||
chunks = []
|
||||
for i, delta in enumerate(deltas):
|
||||
finish = "stop" if i == len(deltas) - 1 else None
|
||||
chunks.append(
|
||||
json.dumps(
|
||||
{
|
||||
**base,
|
||||
"choices": [
|
||||
{"index": 0, "delta": delta, "finish_reason": finish}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
return chunks
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
def test_chat_completion_snowflake(sync_mode):
|
||||
"""
|
||||
Test Snowflake chat completion with mocked HTTP responses.
|
||||
"""
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Write me a poem about the blue sky",
|
||||
},
|
||||
]
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = mock_snowflake_chat_response()
|
||||
messages = [{"role": "user", "content": "Write me a poem about the blue sky"}]
|
||||
mock_resp = _make_mock_response(_chat_response())
|
||||
|
||||
if sync_mode:
|
||||
sync_handler = HTTPHandler()
|
||||
with patch.object(HTTPHandler, "post", return_value=mock_response):
|
||||
with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
|
||||
response = completion(
|
||||
model="snowflake/mistral-7b",
|
||||
messages=messages,
|
||||
api_base="https://exampleopenaiendpoint-production.up.railway.app/v1/chat/completions",
|
||||
client=sync_handler,
|
||||
api_base=FAKE_API_BASE,
|
||||
)
|
||||
assert response is not None
|
||||
assert response.choices[0].message.content is not None
|
||||
assert "sky" in response.choices[0].message.content.lower()
|
||||
mock_post.assert_called_once()
|
||||
else:
|
||||
async_handler = AsyncHTTPHandler()
|
||||
with patch.object(AsyncHTTPHandler, "post", return_value=mock_response):
|
||||
import asyncio
|
||||
|
||||
with patch.object(
|
||||
AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp
|
||||
) as mock_post:
|
||||
response = asyncio.run(
|
||||
acompletion(
|
||||
model="snowflake/mistral-7b",
|
||||
messages=messages,
|
||||
api_base="https://exampleopenaiendpoint-production.up.railway.app/v1/chat/completions",
|
||||
client=async_handler,
|
||||
api_base=FAKE_API_BASE,
|
||||
)
|
||||
)
|
||||
assert response is not None
|
||||
assert response.choices[0].message.content is not None
|
||||
assert "sky" in response.choices[0].message.content.lower()
|
||||
mock_post.assert_called_once()
|
||||
|
||||
assert response is not None
|
||||
assert response.choices[0].message.content is not None
|
||||
assert "sky" in response.choices[0].message.content.lower()
|
||||
assert response.usage.prompt_tokens == 10
|
||||
assert response.usage.completion_tokens == 30
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
def test_chat_completion_snowflake_stream(sync_mode):
|
||||
"""
|
||||
Test Snowflake streaming chat completion with mocked HTTP responses.
|
||||
"""
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Write me a poem about the blue sky",
|
||||
},
|
||||
]
|
||||
messages = [{"role": "user", "content": "Write me a poem about the blue sky"}]
|
||||
raw_chunks = _streaming_chunks()
|
||||
|
||||
if sync_mode:
|
||||
sync_handler = HTTPHandler()
|
||||
mock_chunks = mock_snowflake_streaming_response_chunks()
|
||||
|
||||
def mock_iter_lines():
|
||||
for chunk in mock_chunks:
|
||||
for line in [f"data: {chunk}", "data: [DONE]"]:
|
||||
yield line
|
||||
def _iter_lines():
|
||||
for chunk in raw_chunks:
|
||||
yield f"data: {chunk}"
|
||||
yield "data: [DONE]"
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.iter_lines.side_effect = mock_iter_lines
|
||||
mock_response.status_code = 200
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.iter_lines.return_value = _iter_lines()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = {"content-type": "text/event-stream"}
|
||||
|
||||
with patch.object(HTTPHandler, "post", return_value=mock_response):
|
||||
with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
|
||||
response = completion(
|
||||
model="snowflake/mistral-7b",
|
||||
messages=messages,
|
||||
max_tokens=100,
|
||||
stream=True,
|
||||
api_base="https://exampleopenaiendpoint-production.up.railway.app/v1/chat/completions",
|
||||
client=sync_handler,
|
||||
api_base=FAKE_API_BASE,
|
||||
)
|
||||
|
||||
chunks_received = []
|
||||
for chunk in response:
|
||||
chunks_received.append(chunk)
|
||||
|
||||
assert len(chunks_received) > 0
|
||||
chunks_received = list(response)
|
||||
mock_post.assert_called_once()
|
||||
else:
|
||||
async_handler = AsyncHTTPHandler()
|
||||
mock_chunks = mock_snowflake_streaming_response_chunks()
|
||||
|
||||
async def mock_iter_lines():
|
||||
for chunk in mock_chunks:
|
||||
for line in [f"data: {chunk}", "data: [DONE]"]:
|
||||
yield line
|
||||
async def _aiter_lines():
|
||||
for chunk in raw_chunks:
|
||||
yield f"data: {chunk}"
|
||||
yield "data: [DONE]"
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.iter_lines.side_effect = mock_iter_lines
|
||||
mock_response.status_code = 200
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.aiter_lines.return_value = _aiter_lines()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = {"content-type": "text/event-stream"}
|
||||
|
||||
with patch.object(AsyncHTTPHandler, "post", return_value=mock_response):
|
||||
import asyncio
|
||||
|
||||
async def test_async_stream():
|
||||
response = await acompletion(
|
||||
async def _run():
|
||||
with patch.object(
|
||||
AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp
|
||||
) as mock_post:
|
||||
resp = await acompletion(
|
||||
model="snowflake/mistral-7b",
|
||||
messages=messages,
|
||||
max_tokens=100,
|
||||
stream=True,
|
||||
api_base="https://exampleopenaiendpoint-production.up.railway.app/v1/chat/completions",
|
||||
client=async_handler,
|
||||
api_base=FAKE_API_BASE,
|
||||
)
|
||||
received = []
|
||||
async for chunk in resp:
|
||||
received.append(chunk)
|
||||
mock_post.assert_called_once()
|
||||
return received
|
||||
|
||||
chunks_received = []
|
||||
async for chunk in response:
|
||||
chunks_received.append(chunk)
|
||||
chunks_received = asyncio.run(_run())
|
||||
|
||||
assert len(chunks_received) > 0
|
||||
|
||||
asyncio.run(test_async_stream())
|
||||
assert len(chunks_received) > 0
|
||||
content = "".join(
|
||||
c.choices[0].delta.content for c in chunks_received if c.choices[0].delta.content
|
||||
)
|
||||
assert "sky" in content.lower()
|
||||
|
||||
Reference in New Issue
Block a user