refactor(test_router_caching.py): add tests for router caching

This commit is contained in:
Krrish Dholakia 2023-12-15 20:38:51 -08:00
parent 4d8376a8e9
commit 5fe5149070
3 changed files with 213 additions and 204 deletions

View file

@ -1,174 +1,174 @@
##### THESE TESTS CAN ONLY RUN LOCALLY WITH THE OLLAMA SERVER RUNNING ######
# https://ollama.ai/
# ##### THESE TESTS CAN ONLY RUN LOCALLY WITH THE OLLAMA SERVER RUNNING ######
# # https://ollama.ai/
import sys, os
import traceback
from dotenv import load_dotenv
load_dotenv()
import os
sys.path.insert(0, os.path.abspath('../..')) # Adds the parent directory to the system path
import pytest
import litellm
from litellm import embedding, completion
import asyncio
# import sys, os
# import traceback
# from dotenv import load_dotenv
# load_dotenv()
# import os
# sys.path.insert(0, os.path.abspath('../..')) # Adds the parent directory to the system path
# import pytest
# import litellm
# from litellm import embedding, completion
# import asyncio
user_message = "respond in 20 words. who are you?"
messages = [{ "content": user_message,"role": "user"}]
# user_message = "respond in 20 words. who are you?"
# messages = [{ "content": user_message,"role": "user"}]
async def test_async_ollama_streaming():
try:
litellm.set_verbose = True
response = await litellm.acompletion(model="ollama/mistral-openorca",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
stream=True)
async for chunk in response:
print(chunk)
except Exception as e:
print(e)
# async def test_async_ollama_streaming():
# try:
# litellm.set_verbose = True
# response = await litellm.acompletion(model="ollama/mistral-openorca",
# messages=[{"role": "user", "content": "Hey, how's it going?"}],
# stream=True)
# async for chunk in response:
# print(chunk)
# except Exception as e:
# print(e)
asyncio.run(test_async_ollama_streaming())
# asyncio.run(test_async_ollama_streaming())
def test_completion_ollama():
try:
response = completion(
model="ollama/llama2",
messages=messages,
max_tokens=200,
request_timeout = 10,
# def test_completion_ollama():
# try:
# response = completion(
# model="ollama/llama2",
# messages=messages,
# max_tokens=200,
# request_timeout = 10,
)
print(response)
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# )
# print(response)
# except Exception as e:
# pytest.fail(f"Error occurred: {e}")
# test_completion_ollama()
# # test_completion_ollama()
def test_completion_ollama_with_api_base():
try:
response = completion(
model="ollama/llama2",
messages=messages,
api_base="http://localhost:11434"
)
print(response)
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# def test_completion_ollama_with_api_base():
# try:
# response = completion(
# model="ollama/llama2",
# messages=messages,
# api_base="http://localhost:11434"
# )
# print(response)
# except Exception as e:
# pytest.fail(f"Error occurred: {e}")
# test_completion_ollama_with_api_base()
# # test_completion_ollama_with_api_base()
def test_completion_ollama_custom_prompt_template():
user_message = "what is litellm?"
litellm.register_prompt_template(
model="ollama/llama2",
roles={
"system": {"pre_message": "System: "},
"user": {"pre_message": "User: "},
"assistant": {"pre_message": "Assistant: "}
}
)
messages = [{ "content": user_message,"role": "user"}]
litellm.set_verbose = True
try:
response = completion(
model="ollama/llama2",
messages=messages,
stream=True
)
print(response)
for chunk in response:
print(chunk)
# print(chunk['choices'][0]['delta'])
# def test_completion_ollama_custom_prompt_template():
# user_message = "what is litellm?"
# litellm.register_prompt_template(
# model="ollama/llama2",
# roles={
# "system": {"pre_message": "System: "},
# "user": {"pre_message": "User: "},
# "assistant": {"pre_message": "Assistant: "}
# }
# )
# messages = [{ "content": user_message,"role": "user"}]
# litellm.set_verbose = True
# try:
# response = completion(
# model="ollama/llama2",
# messages=messages,
# stream=True
# )
# print(response)
# for chunk in response:
# print(chunk)
# # print(chunk['choices'][0]['delta'])
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
# except Exception as e:
# traceback.print_exc()
# pytest.fail(f"Error occurred: {e}")
# test_completion_ollama_custom_prompt_template()
# # test_completion_ollama_custom_prompt_template()
async def test_completion_ollama_async_stream():
user_message = "what is the weather"
messages = [{ "content": user_message,"role": "user"}]
try:
response = await litellm.acompletion(
model="ollama/llama2",
messages=messages,
api_base="http://localhost:11434",
stream=True
)
async for chunk in response:
print(chunk['choices'][0]['delta'])
# async def test_completion_ollama_async_stream():
# user_message = "what is the weather"
# messages = [{ "content": user_message,"role": "user"}]
# try:
# response = await litellm.acompletion(
# model="ollama/llama2",
# messages=messages,
# api_base="http://localhost:11434",
# stream=True
# )
# async for chunk in response:
# print(chunk['choices'][0]['delta'])
print("TEST ASYNC NON Stream")
response = await litellm.acompletion(
model="ollama/llama2",
messages=messages,
api_base="http://localhost:11434",
)
print(response)
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# print("TEST ASYNC NON Stream")
# response = await litellm.acompletion(
# model="ollama/llama2",
# messages=messages,
# api_base="http://localhost:11434",
# )
# print(response)
# except Exception as e:
# pytest.fail(f"Error occurred: {e}")
# import asyncio
# asyncio.run(test_completion_ollama_async_stream())
# # import asyncio
# # asyncio.run(test_completion_ollama_async_stream())
def prepare_messages_for_chat(text: str) -> list:
messages = [
{"role": "user", "content": text},
]
return messages
# def prepare_messages_for_chat(text: str) -> list:
# messages = [
# {"role": "user", "content": text},
# ]
# return messages
async def ask_question():
params = {
"messages": prepare_messages_for_chat("What is litellm? tell me 10 things about it who is sihaan.write an essay"),
"api_base": "http://localhost:11434",
"model": "ollama/llama2",
"stream": True,
}
response = await litellm.acompletion(**params)
return response
# async def ask_question():
# params = {
# "messages": prepare_messages_for_chat("What is litellm? tell me 10 things about it who is sihaan.write an essay"),
# "api_base": "http://localhost:11434",
# "model": "ollama/llama2",
# "stream": True,
# }
# response = await litellm.acompletion(**params)
# return response
async def main():
response = await ask_question()
async for chunk in response:
print(chunk)
# async def main():
# response = await ask_question()
# async for chunk in response:
# print(chunk)
print("test async completion without streaming")
response = await litellm.acompletion(
model="ollama/llama2",
messages=prepare_messages_for_chat("What is litellm? respond in 2 words"),
)
print("response", response)
# print("test async completion without streaming")
# response = await litellm.acompletion(
# model="ollama/llama2",
# messages=prepare_messages_for_chat("What is litellm? respond in 2 words"),
# )
# print("response", response)
def test_completion_expect_error():
# this tests if we can exception map correctly for ollama
print("making ollama request")
# litellm.set_verbose=True
user_message = "what is litellm?"
messages = [{ "content": user_message,"role": "user"}]
try:
response = completion(
model="ollama/invalid",
messages=messages,
stream=True
)
print(response)
for chunk in response:
print(chunk)
# print(chunk['choices'][0]['delta'])
# def test_completion_expect_error():
# # this tests if we can exception map correctly for ollama
# print("making ollama request")
# # litellm.set_verbose=True
# user_message = "what is litellm?"
# messages = [{ "content": user_message,"role": "user"}]
# try:
# response = completion(
# model="ollama/invalid",
# messages=messages,
# stream=True
# )
# print(response)
# for chunk in response:
# print(chunk)
# # print(chunk['choices'][0]['delta'])
except Exception as e:
pass
pytest.fail(f"Error occurred: {e}")
# except Exception as e:
# pass
# pytest.fail(f"Error occurred: {e}")
# test_completion_expect_error()
# # test_completion_expect_error()
# if __name__ == "__main__":
# import asyncio
# asyncio.run(main())
# # if __name__ == "__main__":
# # import asyncio
# # asyncio.run(main())

View file

@ -371,66 +371,6 @@ def test_function_calling():
router.reset()
print(response)
def test_acompletion_on_router():
# tests acompletion + caching on router
try:
litellm.set_verbose = True
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo-0613",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 100000,
"rpm": 10000,
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/chatgpt-v-2",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION")
},
"tpm": 100000,
"rpm": 10000,
}
]
messages = [
{"role": "user", "content": f"write a one sentence poem {time.time()}?"}
]
start_time = time.time()
router = Router(model_list=model_list,
redis_host=os.environ["REDIS_HOST"],
redis_password=os.environ["REDIS_PASSWORD"],
redis_port=os.environ["REDIS_PORT"],
cache_responses=True,
timeout=30,
routing_strategy="simple-shuffle")
async def get_response():
print("Testing acompletion + caching on router")
response1 = await router.acompletion(model="gpt-3.5-turbo", messages=messages, temperature=1)
print(f"response1: {response1}")
await asyncio.sleep(1) # add cache is async, async sleep for cache to get set
response2 = await router.acompletion(model="gpt-3.5-turbo", messages=messages, temperature=1)
print(f"response2: {response2}")
assert response1.id == response2.id
assert len(response1.choices[0].message.content) > 0
assert response1.choices[0].message.content == response2.choices[0].message.content
asyncio.run(get_response())
router.reset()
except litellm.Timeout as e:
end_time = time.time()
print(f"timeout error occurred: {end_time - start_time}")
pass
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
# test_acompletion_on_router()
def test_function_calling_on_router():

View file

@ -0,0 +1,69 @@
#### What this tests ####
# This tests caching on the router
import sys, os, time
import traceback, asyncio
import pytest
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from litellm import Router
## Scenarios
## 1. 2 models - openai + azure - 1 model group "gpt-3.5-turbo", assert cache key is the model group
@pytest.mark.asyncio
async def test_acompletion_caching_on_router():
# tests acompletion + caching on router
try:
litellm.set_verbose = True
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo-0613",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 100000,
"rpm": 10000,
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/chatgpt-v-2",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION")
},
"tpm": 100000,
"rpm": 10000,
}
]
messages = [
{"role": "user", "content": f"write a one sentence poem {time.time()}?"}
]
start_time = time.time()
router = Router(model_list=model_list,
redis_host=os.environ["REDIS_HOST"],
redis_password=os.environ["REDIS_PASSWORD"],
redis_port=os.environ["REDIS_PORT"],
cache_responses=True,
timeout=30,
routing_strategy="simple-shuffle")
response1 = await router.acompletion(model="gpt-3.5-turbo", messages=messages, temperature=1)
print(f"response1: {response1}")
await asyncio.sleep(1) # add cache is async, async sleep for cache to get set
response2 = await router.acompletion(model="gpt-3.5-turbo", messages=messages, temperature=1)
print(f"response2: {response2}")
assert response1.id == response2.id
assert len(response1.choices[0].message.content) > 0
assert response1.choices[0].message.content == response2.choices[0].message.content
router.reset()
except litellm.Timeout as e:
end_time = time.time()
print(f"timeout error occurred: {end_time - start_time}")
pass
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")